Author SHA1 Message Date
jiang 682c26fddd feat(timeseries): unify element history queries
Generic Container CI/CD / test-build-publish (push) Successful in 2m10s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 2m10s
2026-09-14 12:30:41 +08:00
jiang 0685f6dd17 fix(sensor): use published coordinates for custom SRIDs
Generic Container CI/CD / test-build-publish (push) Successful in 2m11s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 2m11s
Result validation bypassed the project-specific GIS transform and called ST_Transform on custom engineering SRIDs. Read project and map coordinates from gis.junctions so every project uses its configured publication transform.
2026-09-11 11:29:51 +08:00
jiang 90b02057bc feat(projects): automate project infrastructure provisioning
Generic Container CI/CD / test-build-publish (push) Successful in 1m13s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 1m13s
2026-09-11 10:57:51 +08:00
jiang 10a7a66a41 fix(metadata): hide inactive projects from user list
Project selection previously relied only on membership, so inactive projects remained visible. Filter at the repository boundary and add regression coverage.
2026-09-10 14:49:53 +08:00
jiang bd857ea5e1 fix(ci): test backend in container environment
Generic Container CI/CD / test-build-publish (push) Successful in 1m11s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 1m11s
2026-09-08 18:26:20 +08:00
jiang 8942541759 feat(scada): serve project devices through pooled API 2026-09-08 18:18:30 +08:00
jiang 5966d039de refactor(backend)!: separate algorithm and data layers
Reorganize algorithm packages by business responsibility, move orchestration into services, and keep database access behind pooled repositories.

Harden analysis API validation, remove unsafe legacy simulation endpoints, and add regression and architecture boundary coverage.

BREAKING CHANGE: legacy algorithm module paths and obsolete simulation endpoints are removed.
2026-09-04 17:30:55 +08:00
jiang 9b095c7439 refactor(db)!: clean up business SQL access
- make realtime replacement and analysis result writes transactional\n- consolidate SCADA repositories and remove process-global project state\n- validate SCADA batches and use indexed GIS-backed business queries\n\nBREAKING CHANGE: remove the public analysis result writer and the pipeline-health network_name query parameter.
2026-08-28 11:37:36 +08:00
jiang b74799a39d refactor(db)!: finalize pooled WNDB v2 migration 2026-08-27 17:26:22 +08:00
jiang fa188af0b1 refactor(db)!: adopt project-routed pooled databases
Reorganize WNDB by responsibility and remove legacy scheme endpoints.\n\nRoute analysis and time-series access through project pools, preserve transactional realtime replacement, and refresh GIS materialized views after writes.\n\nAdd database architecture documentation, live pooling coverage, API contract updates, and executable container verification.\n\nBREAKING CHANGE: legacy scheme APIs and flat app.native.wndb module imports are removed.
2026-08-25 18:35:05 +08:00
jiang fdbcc5c033 docs: clarify production configuration handling 2026-08-20 16:19:44 +08:00
jiang d63d7ef1b6 merge: route project DSNs and remove legacy storage backends
Merge PR #2 after isolated OpenAPI, test, container build, and runtime smoke verification.
2026-08-18 18:34:47 +08:00
jiang 6b09662de6 refactor(storage): route project DSNs and remove legacy backends 2026-08-18 18:29:09 +08:00
jiang b21eaffe40 merge: integrate agent-mvp into master
Merge PR #1 after backend security and contract gates passed.
2026-08-18 17:56:43 +08:00
jiang 8853877fcd fix(security): close backend merge blockers 2026-08-18 17:51:29 +08:00
jiang 2581631b51 feat(simulation): 支持冲洗阀门状态与设置值
Generic Container CI/CD / test-build-publish (push) Successful in 2m39s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 2m40s
2026-08-17 18:28:57 +08:00
jiang c250e97b87 ci(backend): block releases missing frontend API contract 2026-08-11 11:11:39 +08:00
jiang 69a7d53aff ci: replace webhook deployment with v2 workflow
Generic Container CI/CD / test-build-publish (push) Successful in 2m32s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 2m32s
2026-08-11 10:17:01 +08:00
jiang e4975b7be3 fix(container): exclude runtime configuration from image 2026-08-11 10:10:21 +08:00
jiang a7e1ce6ef4 fix(db): enforce metadata membership foreign keys
Server CI/CD / docker-image (push) Failing after 34s
Server CI/CD / deploy-fallback-log (push) Successful in 1s
2026-08-06 20:32:57 +08:00
jiang 4350612807 refactor(cli): remove deprecated Python implementation 2026-08-06 15:58:47 +08:00
jiang b0a23a8012 fix(cli): map timeseries element types for backend 2026-08-05 18:39:00 +08:00
jiang 70c5b0e445 fix(keycloak): design logout completion page 2026-08-05 17:57:49 +08:00
jiang 29f691731c fix(auth): enforce Keycloak access token age 2026-08-05 17:10:59 +08:00
jiang 1432934f12 feat(sensor): 返回候选节点最大管径 2026-08-03 19:00:00 +08:00
jiang b0e9d480ef refactor(sensor): 优化灵敏度监测点布置算法 2026-08-03 18:34:13 +08:00
jiang 87d922ea61 refactor(db): retain scheme start time as text 2026-08-03 11:18:38 +08:00
jiang c126e99b60 docs(api): align project lock contract 2026-08-03 11:12:31 +08:00
jiang b6a6527bab refactor(auth): remove project-local user API 2026-08-03 10:39:41 +08:00
jiang f010f071eb refactor(db): align project template schema 2026-08-03 10:39:31 +08:00
jiang 0be31869b5 fix(api): encode untyped datetime responses 2026-07-31 18:39:36 +08:00
jiang eac6b78598 fix(ci): align backend image with deployment
Server CI/CD / docker-image (push) Failing after 29s
Server CI/CD / deploy-fallback-log (push) Successful in 1s
2026-07-31 00:01:21 +08:00
jiang 1d88f8efbe fix(api): wrap pre-paginated list responses
Server CI/CD / docker-image (push) Failing after 28s
Server CI/CD / deploy-fallback-log (push) Successful in 1s
2026-07-30 21:50:21 +08:00
jiang ba947b616b feat(api): standardize REST contracts and auth 2026-07-30 20:38:51 +08:00
jiang ae1a657554 feat(server): add project RBAC and guarded workflows 2026-07-30 16:45:09 +08:00
jiang 3fbb17bb30 fix(sensor-placement): enforce project write boundaries
Bind every scheme request to ProjectContext, keep viewer access read-only, reject concurrent optimization jobs without blocking worker threads, and cap export/update payload sizes. Run optimization and workbook work off the event loop.
2026-07-30 16:21:38 +08:00
jiang ddbb50173c feat(sensor-placement): add editable scheme APIs 2026-07-30 16:16:51 +08:00
jiang 437eb5a19a fix(auth): require preferred username claim 2026-07-30 14:21:21 +08:00
jiang 31e2728db1 refactor(api): unify scheme query endpoints 2026-07-30 11:01:45 +08:00
jiang 03bb2d75c2 docs: 编写中文 README 2026-07-22 11:26:06 +08:00
jiang b977bf6725 fix(db): validate cached project connections
Server CI/CD / docker-image (push) Successful in 23s
Server CI/CD / deploy-fallback-log (push) Has been skipped
2026-07-21 11:26:21 +08:00
jiang 045d6c5b49 fix(simulation): use current user for stored schemes
Server CI/CD / docker-image (push) Successful in 23s
Server CI/CD / deploy-fallback-log (push) Has been skipped
2026-07-17 16:49:20 +08:00
jiang db6032bd84 feat(burst-detection): update scada analysis flow
Server CI/CD / docker-image (push) Successful in 24s
Server CI/CD / deploy-fallback-log (push) Has been cancelled
2026-07-17 16:31:46 +08:00
jiang b4ecfbb87a fix(scada): use project-scoped metadata 2026-07-17 16:28:40 +08:00
jiang a204980944 fix(leakage): accept display flow units 2026-07-17 11:45:58 +08:00
jiang 2b5f9b8514 fix(simulation): use report step for schemes 2026-07-16 15:47:26 +08:00
jiang ca1579dcc2 fix(api): include simulation burst ids 2026-07-16 14:50:22 +08:00
jiang 775ecb8a58 fix(simulation): use hydraulic timestep 2026-07-16 14:16:06 +08:00
jiang f72b56845f fix(wndb): refresh closed project connections 2026-07-16 12:07:44 +08:00
jiang baeaa8a2e1 fix(burst-location): correct normal data window
Server CI/CD / docker-image (push) Successful in 22s
Server CI/CD / deploy-fallback-log (push) Has been skipped
2026-07-09 11:51:40 +08:00
jiang ca97de2e51 fix(burst-location): tolerate partial SCADA gaps 2026-07-09 10:49:53 +08:00
jiang 71fa2ae18c ci: unzip health model in image build
Server CI/CD / docker-image (push) Successful in 22s
Server CI/CD / deploy-fallback-log (push) Has been skipped
2026-07-08 19:02:54 +08:00
jiang 76cf6c32bc fix(agent): expose network context
Server CI/CD / docker-image (push) Successful in 16s
Server CI/CD / deploy-fallback-log (push) Has been skipped
2026-07-08 18:41:56 +08:00
jiang 5a91da0904 fix(burst-location): use correct data sources
Simulation mode reads scheme data for both burst and normal observations. Monitoring mode reuses the burst window when no normal window is provided.
2026-07-08 17:51:08 +08:00
jiang d62bcae85e fix(burst-location): normalize simulation ids 2026-07-08 17:17:41 +08:00
jiang 4c0a4b29e9 refactor(metadata): drop geoserver config refs 2026-06-13 14:57:33 +08:00
jiang 80ca985c28 fix(cli): use renamed backend APIs 2026-06-13 13:56:44 +08:00
jiang 5a55d65002 refactor(api): add kebab-case legacy aliases 2026-06-13 13:07:16 +08:00
jiang d99f4cec6a refactor(admin): remove geoserver config 2026-06-12 15:28:14 +08:00
jiang a6e7a2e75c feat(admin): add project metadata config 2026-06-12 15:08:37 +08:00
jiang 23c008f602 feat(auth): migrate to Keycloak metadata auth 2026-06-12 10:18:41 +08:00
jiang 2a762e63a7 添加 Gitea 服务器 URL 和用户解析功能
Server CI/CD / docker-image (push) Successful in 49s
Server CI/CD / deploy-fallback-log (push) Has been skipped
2026-06-10 16:27:42 +08:00
jiang f6939f5516 移除 db_inp 目录的复制,添加临时文件夹创建
Server CI/CD / docker-image (push) Failing after 1m42s
Server CI/CD / deploy-fallback-log (push) Successful in 0s
2026-06-10 16:22:38 +08:00
jiang bbf6a0f7ba 更新代码检出步骤和工作区验证提示信息
Server CI/CD / docker-image (push) Failing after 13s
Server CI/CD / deploy-fallback-log (push) Successful in 0s
2026-06-10 16:16:26 +08:00
jiang 5fd82b8e7c 添加代码检出步骤以确保工作区有效
Server CI/CD / docker-image (push) Failing after 1m31s
Server CI/CD / deploy-fallback-log (push) Successful in 1s
2026-06-10 15:44:41 +08:00
jiang 2a823b2616 移除代码检出步骤,添加工作区验证
Server CI/CD / docker-image (push) Failing after 1s
Server CI/CD / deploy-fallback-log (push) Successful in 0s
2026-06-10 15:41:59 +08:00
jiang 2af89eea1c 优化环境变量检查逻辑,移除不必要的密钥
Server CI/CD / docker-image (push) Failing after 27s
Server CI/CD / deploy-fallback-log (push) Successful in 1s
2026-06-10 15:18:10 +08:00
jiang 26643d68c7 feat(ci): 添加 Gitea 仓库密钥 TJWATER_SERVER_ENV 检查
Server CI/CD / docker-image (push) Failing after 11s
Server CI/CD / deploy-fallback-log (push) Successful in 0s
2026-06-10 15:08:47 +08:00
jiang f35287d3cf 更新 dockerfile
Server CI/CD / docker-image (push) Failing after 11s
Server CI/CD / deploy-fallback-log (push) Successful in 0s
2026-06-10 11:45:22 +08:00
jiang 4fa8e55748 删除 copilot 自述文件 2026-06-09 18:24:14 +08:00
jiang 7a9fcaae81 ci: add deployment trigger script
Server CI/CD / docker-image (push) Has been cancelled
Server CI/CD / deploy-fallback-log (push) Has been cancelled
2026-06-09 18:22:16 +08:00
jiang a1e9673d9a ci: add Gitea package workflow 2026-06-09 18:18:22 +08:00
jiang e588d1cf33 feat(api): add Tianditu geocoding 2026-06-09 17:09:42 +08:00
jiang 1712ecd4c7 feat(api): add web search endpoint 2026-06-09 16:13:24 +08:00
jiang 441979f581 修改默认超时时间 2026-06-05 19:11:53 +08:00
jiang e336ffcd46 移除存在无效数据的 cli 命令 2026-06-05 16:42:03 +08:00
jiang 52b8f07abd 更新 cli 命令,新增 network 其他元素的属性查询 2026-06-05 15:48:53 +08:00
jiang 7efaeb41e8 新增pyclipper依赖 2026-06-05 13:43:53 +08:00
jiang 9a7aad2d36 fix(cli): constrain timeseries option values 2026-06-05 13:43:32 +08:00
jiang b7872f29a9 优化 CLI 命令,增加获取所有节点和管道属性的功能 2026-06-03 17:31:49 +08:00
jiang 233960d8db 明确时间模拟需要 scheme_name 参数 2026-06-03 17:31:44 +08:00
jiang b9410b0ff3 统一前后端时间时区请求 2026-06-03 11:17:37 +08:00
jiang 4982efba5e 更新tjwater-cli network参数;更新metadb health方法 2026-06-03 10:48:01 +08:00
jiang f87dd91b2b 修复--auth-stdin读取失败的bug 2026-06-02 18:41:39 +08:00
jiang c16e6e3d0c 移除 --auth-context,改为 --auth-stdin,结构化传递解析认证信息 2026-06-02 17:17:00 +08:00
jiang 40e699e173 拆分代码;约束cli命令 2026-06-02 14:54:08 +08:00
jiang 9b8a517092 更新文件夹命名 2026-06-02 11:13:07 +08:00
jiang f274cf5122 整理 tjwater-cli 代码和文档 2026-06-02 11:11:56 +08:00
jiang 60db2a7193 优化 cli 命令设计 2026-06-01 17:05:26 +08:00
jiang b72e42521c 优化时间范围查询,添加 UTC 时间标准化处理 2026-06-01 16:46:51 +08:00
jiang c2ccb7bc4e 移除实时数据和仿真结果接口,优化代码结构 2026-05-26 18:49:25 +08:00
jiang 88be97ddeb 修正单元测试失败代码 2026-05-25 17:51:45 +08:00
jiang 2317f4d527 新增 API 测试用例,修复失效接口问题 2026-05-21 15:32:12 +08:00
jiang 751950e5b5 调整函数说明 2026-05-20 11:45:01 +08:00
jiang a1dcbd4230 更新 dockerfile,提高打包效率 2026-04-30 13:06:09 +08:00
jiang 3b712ea467 优化传感器布置算法,修复数据库更新逻辑 2026-04-17 17:21:50 +08:00
jiang bf2aaa5ff7 后端统一时区为 UTC 2026-04-14 14:46:51 +08:00
jiang 51b481d174 优化临时文件管理,增强错误日志记录 2026-04-08 11:47:46 +08:00
jiang 644babf77e 将环境设置为生产模式;更新网络名称配置 2026-04-08 10:49:01 +08:00
jiang 6b09c6b20d 删除 Dockerfile 中的临时文件复制指令 2026-04-03 14:53:55 +08:00
jiang 93cbd7e7b3 独立 copilot 服务 2026-03-27 13:52:12 +08:00
jiang 0196206ed3 创建层级化目录的 skills 2026-03-27 13:05:22 +08:00
jiang 88eec2787b 整理 api tags 2026-03-27 12:31:52 +08:00
jiang 621cd9d2f9 删除 router 中多余的tags 2026-03-26 16:09:17 +08:00
jiang 600ddd329c 添加流式 Copilot 请求处理及审计中间件优化 2026-03-24 16:01:22 +08:00
jiang c184610035 添加 Copilot 聊天流式响应接口及测试 2026-03-24 11:22:00 +08:00
jiang 21dd393aee 添加 Copilot 聊天流式响应功能及相关配置 2026-03-23 18:03:00 +08:00
398 changed files with 68624 additions and 42291 deletions
+22
View File
@@ -0,0 +1,22 @@
.git
.github
.gitea
__pycache__/
.pytest_cache/
.mypy_cache/
.venv/
venv/
build/
dist/
package/
temp/
data/
db_inp/
inp/
.env
.env.*
logs/
coverage/
*.pyc
*.dump
app/algorithms/pipe_health_prediction/model/my_survival_forest_model_quxi.joblib
+48 -9
View File
@@ -1,19 +1,16 @@
# TJWater Server 环境变量配置模板
# 复制此文件为 .env 并填写实际值
ENVIRONMENT="local"
# CI/CD: 生产环境变量由 Dev 主机的受控 backend.env 注入,不要将完整 .env 保存为 Gitea 仓库密钥。
ENVIRONMENT="production"
NETWORK_NAME="tjwater"
# ============================================
# 安全配置 (必填)
# 敏感配置加密 (必填)
# ============================================
# JWT 密钥 - 用于生成和验证 Token
# 生成方式: openssl rand -hex 32
SECRET_KEY=your-secret-key-here-change-in-production-use-openssl-rand-hex-32
# 数据加密密钥 - 用于敏感数据加密
# Fernet 格式,生产环境必须替换为独立密钥
# 生成方式: python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())"
ENCRYPTION_KEY=
DATABASE_ENCRYPTION_KEY="rJC2VqLg4KrlSq+DGJcYm869q4v5KB2dFAeuQTe0I50="
# 用于项目数据库 DSN、GeoServer 管理密码等敏感配置
DATABASE_ENCRYPTION_KEY="replace-with-generated-fernet-key"
# ============================================
# 数据库配置 (PostgreSQL)
@@ -42,10 +39,52 @@ METADATA_DB_PORT="5432"
METADATA_DB_USER="tjwater"
METADATA_DB_PASSWORD="password"
# Per-project synchronous connection pools
PROJECT_PG_CACHE_SIZE="16"
PROJECT_TS_CACHE_SIZE="16"
PROJECT_PG_POOL_MIN_SIZE="0"
PROJECT_PG_POOL_SIZE="4"
PROJECT_PG_MAX_OVERFLOW="2"
PROJECT_TS_POOL_MIN_SIZE="0"
PROJECT_TS_POOL_MAX_SIZE="4"
WNDB_SCHEMA_TEMPLATE_DB_NAME="tjwater_v2_schema_template"
TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME="tjwater_v2_timescale_template"
WNDB_TEMP_DB_MAX_COUNT="8"
# ============================================
# GeoServer 项目供应
# ============================================
GEOSERVER_URL="http://localhost:8080/geoserver"
GEOSERVER_USERNAME="admin"
GEOSERVER_PASSWORD="password"
# 留空时分别复用 DB_HOST / DB_PORT / DB_USER / DB_PASSWORD。
# GeoServer 在容器内运行时,GEOSERVER_DB_HOST 通常应填写数据库服务名。
GEOSERVER_DB_HOST=""
GEOSERVER_DB_PORT=""
GEOSERVER_DB_USER=""
GEOSERVER_DB_PASSWORD=""
# 30 天客户端缓存,单位为秒。
GEOSERVER_CLIENT_CACHE_SECONDS="2592000"
# ============================================
# Keycloak JWT (可选)
# ============================================
KEYCLOAK_PUBLIC_KEY="-----BEGIN PUBLIC KEY-----\n...\n-----END PUBLIC KEY-----"
KEYCLOAK_ALGORITHM=RS256
KEYCLOAK_AUDIENCE="account"
KEYCLOAK_ACCESS_TOKEN_MAX_AGE_SECONDS=900
# ============================================
# Bocha Web Search API
# ============================================
BOCHA_API_KEY="sk-your-bocha-api-key"
BOCHA_WEB_SEARCH_URL="https://api.bochaai.com/v1/web-search"
BOCHA_WEB_SEARCH_TIMEOUT_SECONDS=30
# ============================================
# Tianditu Geocoding API
# ============================================
TIANDITU_GEOCODER_TOKEN="your-tianditu-geocoder-token"
TIANDITU_GEOCODER_URL="https://api.tianditu.gov.cn/geocoder"
TIANDITU_GEOCODER_TIMEOUT_SECONDS=30
+22
View File
@@ -0,0 +1,22 @@
name: Server CI/CD v2
on:
push:
tags:
- "v*"
workflow_dispatch: {}
jobs:
build-test-publish-and-deploy:
uses: OrgTJWater/ci-templates/.gitea/workflows/container-cd.yml@68c8a9855391baa31d31523674f5812cd24ec604
with:
image_name: gitea.waternetwork.cn/orgtjwater/tjwater-backend
dockerfile: Dockerfile
build_context: .
test_target: test
deploy_service: backend
deploy_host: 192.168.1.114
secrets:
REGISTRY_USERNAME: ${{ secrets.REGISTRY_USERNAME }}
REGISTRY_PASSWORD: ${{ secrets.REGISTRY_PASSWORD }}
DEV_DEPLOY_SSH_KEY: ${{ secrets.DEV_DEPLOY_SSH_KEY }}
-82
View File
@@ -1,82 +0,0 @@
# Copilot Instructions for TJWater Server
This repository contains the backend code for the TJWater Server, a water distribution network management system built with FastAPI.
## High-Level Architecture
The application follows a layered architecture:
- **Entry Point**: `app/main.py` initializes the FastAPI application, database connections (PostgreSQL & TimescaleDB), and middleware.
- **API Layer**: `app/api/v1` contains the route handlers.
- **Service Layer**: `app/services` contains business logic and orchestration.
- **Infrastructure Layer**: `app/infra` handles database connections (`db`), audit logging (`audit`), and external integrations.
- **Domain Layer**: `app/domain` likely contains core domain models.
- **Native/Algorithms**: `app/native` and `app/algorithms` handle specialized water network calculations (possibly using EPANET/WNTR).
## Build, Test, and Run Commands
### Environment Setup
- Dependencies are listed in `requirements.txt`.
- Configuration is managed via environment variables (see `.env.example` if available, or `app/core/config.py`).
- **Important**: Ensure `.env` is configured with correct database credentials for both PostgreSQL and TimescaleDB.
If first time setting up, you may want to create a Conda environment:
```bash
conda create -n server python=3.12
conda activate server
pip install uv
uv pip install -r requirements.txt
conda install -c conda-forge pymetis
```
### Running the Server
The preferred way to run the server locally is using the helper script which sets up the Python path correctly:
```bash
conda activate server
python scripts/run_server.py
```
Alternatively, you can run directly with uvicorn (ensure PYTHONPATH includes the root):
```bash
conda activate server
uvicorn app.main:app --host 0.0.0.0 --port 8000 --reload
```
### Running Tests
Use `pytest` to run tests. The `tests/conftest.py` handles path setup.
```bash
# Run all tests
pytest
# Run a specific test file
pytest tests/unit/test_specific_file.py
# Run a specific test case
pytest tests/unit/test_specific_file.py::test_function_name
```
### Building (Optional)
The project includes scripts to compile Python modules to `.pyd` files using Cython (see `scripts/build_pyd.py`). This is likely for distribution/performance but not required for standard development.
## Key Conventions
- **Async/Await**: The codebase heavily uses `async` and `await` for I/O operations, especially database interactions.
- **Database Management**:
- Connections are managed globally in `app.infra.db` and initialized in `lifespan` (app/main.py).
- Use `app.infra.db.dynamic_manager` for project-specific database connections (multi-tenancy/dynamic projects).
- **Pydantic**: extensively used for data validation and settings management.
- **Scripts**: The `scripts/` directory contains many utility scripts for maintenance, data processing, and server management. Check there before writing new operational scripts.
- **Water Network Modeling**: Interactions with water network models often involve `epanet` or `wntr` libraries. Be aware of domain-specific terminology (nodes, links, junctions, tanks).
## Code Style
- Follow standard PEP 8 guidelines.
- No specific linter configuration was found, so default to standard Python formatting.
-128
View File
@@ -1,128 +0,0 @@
name: Build And Package
on:
push:
tags:
- "v*"
jobs:
build-package:
runs-on: ${{ matrix.os }}
env:
FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true
strategy:
fail-fast: false
matrix:
os: [ubuntu-latest, windows-latest]
steps:
- name: Checkout source
uses: actions/checkout@v5
- name: Setup Python
uses: actions/setup-python@v6
with:
python-version: "3.12"
- name: Install system build tools
if: runner.os == 'Linux'
run: |
sudo apt-get update
sudo apt-get install -y build-essential
- name: Install compile dependencies
run: |
python -m pip install --upgrade pip
pip install cython setuptools wheel
- name: Run Cython compile
run: |
python scripts/compile.py
- name: Prepare package and archive
run: |
python - <<'PY'
import os
import shutil
import tarfile
import zipfile
import sys
from pathlib import Path
root = Path.cwd()
package_dir = root / "package"
dist_dir = root / "dist"
for d in [package_dir, dist_dir]:
if d.exists():
shutil.rmtree(d)
d.mkdir(parents=True, exist_ok=True)
# Define directories with compiled artifacts
compile_dirs = ["app/services", "app/native/wndb", "app/algorithms"]
# Global ignore list
ignore_names = {
".git",
".github",
"__pycache__",
".pytest_cache",
".mypy_cache",
".venv",
"venv",
"temp",
"tests",
"package",
"dist",
}
def ignore_func(directory, names):
rel_dir = os.path.relpath(directory, root).replace("\\", "/")
is_in_compile_path = any(rel_dir.startswith(d) for d in compile_dirs)
ignored = []
for name in names:
if name in ignore_names or name.endswith(".pyc"):
ignored.append(name)
# Exclude source .py files only in compiled directories
elif is_in_compile_path and name.endswith(".py"):
ignored.append(name)
return ignored
for item in root.iterdir():
if item.name in ignore_names:
continue
target = package_dir / item.name
if item.is_dir():
shutil.copytree(item, target, ignore=ignore_func)
else:
shutil.copy2(item, target)
# Safety guard: ensure no .github directory remains
github_paths = [p for p in package_dir.rglob(".github") if p.is_dir()]
for p in github_paths:
shutil.rmtree(p, ignore_errors=True)
sha = os.environ["GITHUB_SHA"]
run_os = os.environ["RUNNER_OS"].lower()
if run_os == "windows":
archive_path = dist_dir / f"tjwater-server-{run_os}-{sha}.zip"
with zipfile.ZipFile(archive_path, "w", compression=zipfile.ZIP_DEFLATED) as zf:
for f in package_dir.rglob("*"):
if f.is_file():
zf.write(f, f.relative_to(package_dir))
else:
archive_path = dist_dir / f"tjwater-server-{run_os}-{sha}.tar.gz"
with tarfile.open(archive_path, "w:gz") as tf:
tf.add(package_dir, arcname=".")
print(f"Archive created: {archive_path}")
PY
shell: bash
- name: Upload package artifact
uses: actions/upload-artifact@v5
with:
name: tjwater-server-package-${{ runner.os }}
path: dist/*
retention-days: 14
+2 -1
View File
@@ -6,4 +6,5 @@ build/
.env
*.dump
.vscode/
app/algorithms/health/model/my_survival_forest_model_quxi.joblib
app/algorithms/pipe_health_prediction/model/my_survival_forest_model_quxi.joblib
/inp/
+38
View File
@@ -0,0 +1,38 @@
# Repository Guidelines
## Project Structure & Module Organization
This repository contains the TJWater Python backend. Main application code lives in `app/`: API routes under `app/api`, authentication in `app/auth`, configuration in `app/core`, database and repository code in `app/infra`, domain models/schemas in `app/domain`, and business logic in `app/services` and `app/algorithms`.
Tests are under `tests/`, split into `tests/unit`, `tests/api`, and `tests/auth`. SQL and sample assets are stored in `resources/`; deployment files are in `Dockerfile`, `.gitea/workflows/package.yml`, and `infra/docker/docker-compose.yml`. Local data directories such as `db_inp/`, `temp/`, `data/`, and `.env` are ignored and should not be committed.
## Build, Test, and Development Commands
Use the existing conda environment when available:
```bash
conda run -n server python -m pytest tests/unit tests/auth -q
conda run -n server uvicorn app.main:app --host 0.0.0.0 --port 8000 --reload
docker build -t tjwater-server:local .
docker compose -f infra/docker/docker-compose.yml config
```
`pytest` runs backend tests. `uvicorn` starts the FastAPI app locally. `docker build` verifies the container image. `docker compose config` validates compose syntax and variable expansion.
## Coding Style & Naming Conventions
Use Python 3.12, four-space indentation, type hints for new public functions, and explicit imports. Keep API endpoint modules grouped by domain under `app/api/v1/endpoints`. Use `snake_case` for files, functions, and variables; `PascalCase` for classes and Pydantic models. Prefer existing repository/service patterns in `app/infra/db` and `app/services` over introducing new abstractions.
## Testing Guidelines
The project uses `pytest`. Name test files `test_*.py` and test functions `test_*`. Keep unit tests isolated with fakes or monkeypatching from `tests/conftest.py`. Some existing tests depend on local data outside the repository; avoid adding new tests that require untracked files. For API changes, add or update tests in `tests/api`.
## Commit & Pull Request Guidelines
History uses a mix of Conventional Commit prefixes and concise Chinese messages, for example `feat(api): add Tianditu geocoding` or `fix(auth): validate project context`. Prefer `feat(scope): ...`, `fix(scope): ...`, or a clear Chinese summary.
Pull requests should describe the behavior change, list verification commands, mention configuration or migration impacts, and link related issues. Include API examples or screenshots only when they clarify user-facing behavior.
## Security & Configuration Tips
Do not commit `.env`, database dumps, generated caches, or local project data. Use `.env.example` as the configuration template. CI/CD only uses Gitea repository secrets `REGISTRY_USERNAME`, `REGISTRY_PASSWORD`, and `DEV_DEPLOY_SSH_KEY`; production application settings are injected on the Dev host.
+117
View File
@@ -0,0 +1,117 @@
# TJWater Authentication and Metadata Management
## Ownership
Keycloak owns login identity, credentials, token issuance, and token expiry.
TJWater metadata stores only business snapshots and authorization data:
- `users.keycloak_id` is the stable identity binding.
- `users.username`, `users.email`, and `users.last_login_at` are Keycloak claim caches.
- `users.role`, `users.is_active`, and `users.is_superuser` control TJWater system access.
- `user_project_membership.project_role` controls project access.
The backend does not accept passwords, does not issue local JWTs, and does not
trust frontend-supplied user IDs.
## Fixed Project RBAC
Project roles are stored directly in
`user_project_membership.project_role`; there is no separate role table or
user-defined permission editor in this delivery.
| Role | Main access |
| --- | --- |
| `modeler` | Model upload/import, simulation, burst, risk, and optimization analysis |
| `dispatcher` | SCADA cleaning, simulation and burst analysis |
| `auditor` | Project read access and project-scoped audit logs |
| `viewer` | WebGIS and read-only risk results |
Legacy `owner`, `admin`, and `member` values remain supported for existing
records. The backend is the authorization boundary; the frontend uses
`GET /api/v1/access/context` only to hide unavailable menus and guard routes.
System admins receive environment, membership, and global-audit permissions,
but still need a project membership for project business APIs.
## Login Snapshot Refresh
Every authenticated metadata-user resolution validates the Keycloak access token
and reads `sub`, `preferred_username`, and `email` claims. The
backend finds `users` by `keycloak_id = sub`, rejects inactive or missing users,
then refreshes `username`, `email`, and `last_login_at`.
This keeps local display data current without changing the identity binding.
There is no Keycloak webhook requirement; second-level user or permission sync is
out of scope unless explicitly requested later.
## Admin APIs
All admin APIs require metadata admin access: `users.is_superuser = true` or
`users.role = 'admin'`.
User and membership management:
- `GET /api/v1/admin/me`
- `POST /api/v1/admin/users/sync`
- `POST /api/v1/admin/users/sync/batch`
- `GET /api/v1/admin/users`
- `GET /api/v1/admin/users/{user_id}`
- `PATCH /api/v1/admin/users/{user_id}`
- `GET /api/v1/admin/projects/{project_id}/members`
- `POST /api/v1/admin/projects/{project_id}/members`
- `PATCH /api/v1/admin/projects/{project_id}/members/{user_id}`
- `DELETE /api/v1/admin/projects/{project_id}/members/{user_id}`
Project configuration:
- `GET /api/v1/admin/projects`
- `POST /api/v1/admin/projects`
- `PATCH /api/v1/admin/projects/{project_id}`
- `GET /api/v1/admin/projects/{project_id}/databases`
- `PUT /api/v1/admin/projects/{project_id}/databases`
- `DELETE /api/v1/admin/projects/{project_id}/databases/{db_role}`
- `POST /api/v1/admin/projects/{project_id}/databases/{db_role}/health`
## Secret Handling
Admins submit plaintext DSNs only through HTTPS admin APIs. Operators should not
write encrypted columns manually.
- `project_databases.dsn_encrypted` is encrypted with `DATABASE_ENCRYPTION_KEY`.
- Admin responses return only `has_dsn`.
- Audit logs record whether a secret was updated, but never store plaintext DSNs
or other secrets.
Generate the database encryption key with:
```bash
python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())"
```
Keep keys stable for the lifetime of encrypted metadata. Rotating a key requires
decrypting with the old key and re-encrypting with the new key.
## Metadata Schema Patches
Apply metadata patches in order:
1. `resources/sql/004_metadata_auth_management.sql`
2. `resources/sql/005_metadata_project_configuration.sql`
3. `resources/sql/006_metadata_rbac_roles.sql`
`004` creates Keycloak-backed metadata users and project memberships. `005`
creates project and project database routing tables with uniqueness, role/type,
and pool-size constraints. `006` extends existing membership constraints with
the fixed delivery roles.
## Frontend System Management
`/system-admin` is shown only when `GET /api/v1/access/context` returns
`environment.manage`. The page lets admins maintain metadata users, project
members, projects, project database routing for `biz_data` and `iot_data`, and
connection health checks. This replaces direct SQL editing for normal project
onboarding.
Hydraulic model authoring is outside the Web application. Models are prepared
in the desktop modeling client and uploaded/imported by an authorized modeler;
the system administrator configures the project environment and database
routing.
+77
View File
@@ -0,0 +1,77 @@
# Backend Naming Audit
DOC-003 audit for the internal `TJWaterServerBinary` backend.
## Scope
Reviewed FastAPI route decorators under `app/api/v1/endpoints`, router prefixes in `app/api/v1/router.py`, and public request/response schema fields in `app/api` and `app/domain`.
The backend is mounted only under `/api/v1` from `app/main.py`; the old no-prefix router include remains commented out.
## Current Good Surface
These newer routes already follow the naming rule for public HTTP paths:
- Metadata/admin: `/api/v1/admin/projects`, `/api/v1/admin/users/sync`, `/api/v1/admin/projects/{project_id}/members`
- Audit: `/api/v1/audit/logs`, `/api/v1/audit/logs/count`
- Agent auth: `/api/v1/agent/auth/context`
- Business APIs: `/api/v1/burst-detection/detect`, `/api/v1/burst-location/locate`, `/api/v1/leakage/identify`
- Time-series APIs: `/api/v1/scada/by-ids-time-range`, `/api/v1/scada/by-ids-field-time-range`, `/api/v1/composite/clean-scada`
- Project data APIs: `/api/v1/scada-info`, `/api/v1/scheme-list`, `/api/v1/burst-locate-result`
- Web integrations: `/api/v1/web-search`, `/api/v1/geocode`
Path template parameters such as `{project_id}`, `{user_id}`, `{device_id}`, `{scheme_name}`, and `{link_id}` intentionally remain `snake_case`.
## Legacy URL Categories
### Keep With Compatibility
These now have `kebab-case` aliases. The frontend has been migrated to the replacement paths; keep the old paths as deprecated compatibility aliases for Agent planning, tests, customer scripts, or external callers:
| Current URL | Suggested replacement |
| --- | --- |
| `/api/v1/openproject/` | `/api/v1/projects/open` |
| `/api/v1/project_info/` | `/api/v1/project-info` |
| `/api/v1/getallschemes/` | `/api/v1/schemes` |
| `/api/v1/getallsensorplacements/` | `/api/v1/sensor-placement-schemes` |
| `/api/v1/sensorplacementscheme/create` | `/api/v1/sensor-placement-schemes` |
| `/api/v1/burst_analysis/` | `/api/v1/burst-analysis` |
| `/api/v1/valve_isolation_analysis/` | `/api/v1/valve-isolation-analysis` |
| `/api/v1/flushing_analysis/` | `/api/v1/flushing-analysis` |
| `/api/v1/contaminant_simulation/` | `/api/v1/contaminant-simulation` |
| `/api/v1/runsimulationmanuallybydate/` | `/api/v1/simulations/run-by-date` |
### Broad Legacy Surface
These route groups expose many command-style concatenated paths. They should not be copied into new work; replace only when a caller migration is planned:
- Project lifecycle: `listprojects`, `createproject`, `deleteproject`, `isprojectopen`, `closeproject`, `copyproject`, `importinp`, `exportinp`, `readinp`, `dumpinp`, `lockproject`, `unlockproject`
- Network object CRUD: `addjunction`, `getjunctionelevation`, `setpipediameter`, `getvalvesetting`, and similar junction/pipe/pump/tank/reservoir/valve routes
- Region/DMA/VD commands: `calculatedistrictmeteringareaforregion`, `getdistrictmeteringarea`, `generatevirtualdistrict`, and related routes
- SCADA native CRUD: `getscadadevice`, `setscadadevicedata`, `cleanscadaelement`, and related routes
- Snapshot/synchronization utilities: `takesnapshotforoperation`, `syncwithserver`
- Advanced simulation endpoints with underscore paths: `pressure_regulation`, `daily_scheduling_analysis`, `network_update`, `pressure_sensor_placement_kmeans`
### Direct Cleanup Candidates
These are likely safe only after confirming no caller uses them:
- `/api/v1/test_dict/`: development/test utility in `misc.py`.
- `/api/v1/takenapshotforcurrentoperation`: typo compatibility path; keep deprecated if any client may still call it.
- `/api/v1/getpumpenergyproperties//` and `/api/v1/setpumpenergyproperties//`: double-slash paths in options endpoints.
## Field Naming
Most public JSON, query, and SSE fields are already `snake_case`, including `project_id`, `user_id`, `scheme_name`, `scheme_type`, `start_time`, `end_time`, `device_ids`, `session_id`, and `request_id`.
Known legacy exception:
- `BurstAnalysis.burst_ID` in `app/api/v1/endpoints/simulation.py` should become `burst_id` on a new API contract. Preserve `burst_ID` only for the legacy body shape.
Headers keep standard HTTP casing:
- `X-Project-Id`
## Recommendation
Do not rename existing legacy routes in place. For each active legacy route, keep the new `kebab-case` alias as the documented path, keep the old route marked deprecated, migrate remaining Agent/customer/script callers, then remove only after a documented compatibility window.
+26 -9
View File
@@ -1,25 +1,42 @@
FROM continuumio/miniconda3:latest
FROM condaforge/miniforge3:latest AS base
WORKDIR /app
ENV PIP_INDEX_URL=https://pypi.tuna.tsinghua.edu.cn/simple \
PIP_TRUSTED_HOST=pypi.tuna.tsinghua.edu.cn \
UV_INDEX_URL=https://pypi.tuna.tsinghua.edu.cn/simple
# 安装 Python 3.12 和 pymetis (通过 conda-forge 避免编译问题)
RUN conda install -y -c conda-forge python=3.12 pymetis && \
conda clean -afy
RUN mamba install -y python=3.12 pymetis && \
mamba clean -afy
COPY requirements.txt .
RUN pip install uv
RUN pip install --no-cache-dir uv
RUN uv pip install --system --no-cache-dir -r requirements.txt
# 将代码放入子目录 'app',将数据放入子目录 'db_inp'
# 这样临时文件默认会生成在 /app 下,而代码在 /app/app 下,实现了分离
# 本地数据目录和环境变量在运行时通过 Compose 挂载或注入,
# 不应进入镜像构建上下文。
COPY app ./app
COPY db_inp ./db_inp
COPY temp ./temp
COPY .env .
COPY contracts ./contracts
COPY infra ./infra
COPY resources ./resources
RUN python -c "from pathlib import Path; from zipfile import ZipFile; model_dir = Path('app/algorithms/pipe_health_prediction/model'); zip_path = model_dir / 'my_survival_forest_model_quxi.zip'; joblib_name = 'my_survival_forest_model_quxi.joblib'; joblib_path = model_dir / joblib_name; assert zip_path.exists(), f'Model archive not found: {zip_path}'; archive = ZipFile(zip_path); archive.extract(joblib_name, model_dir); archive.close(); assert joblib_path.exists(), f'Model file not extracted: {joblib_path}'" && \
rm -f app/algorithms/pipe_health_prediction/model/my_survival_forest_model_quxi.zip
RUN mkdir -p db_inp temp data inp
# 设置 PYTHONPATH 以便 uvicorn 找到 app 模块
ENV PYTHONPATH=/app
FROM base AS test
COPY scripts ./scripts
COPY tests ./tests
RUN python -m compileall -q app && \
python -c "import app.main" && \
python scripts/check_openapi.py && \
pytest -q tests
FROM base AS runner
EXPOSE 8000
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", "--workers", "4"]
+108
View File
@@ -0,0 +1,108 @@
# TJWaterServerBinary 内部后端
`TJWaterServerBinary` 是 TJWater 内部版 Python 后端,基于 FastAPI 提供认证、项目、管网、模拟、爆管、漏损、SCADA 和地图服务集成能力。该仓库用于内部开发和完整功能维护。
## 技术栈
- Python 3.12
- FastAPI / Uvicorn
- Pydantic / SQLAlchemy / psycopg
- PostgreSQL、PostGIS、TimescaleDB
- WNTR、EPANET、Cython、科学计算与空间分析依赖
- pytest
## 目录结构
```text
app/main.py FastAPI 入口
app/api/ HTTP API 路由
app/auth/ 认证和权限上下文
app/core/ 配置、日志和基础设施初始化
app/domain/ 领域模型和 Pydantic schema
app/infra/ 数据库、EPANET 和外部集成
app/services/ 业务服务编排
app/algorithms/ 管网算法、模拟、爆管、漏损、清洗和健康分析
app/native/ 本地管网数据读写与转换
tests/ 后端测试
resources/ SQL、模板和示例资源
infra/docker/ Docker Compose 编排
```
## 本地开发
推荐使用已有 conda 环境:
```bash
conda run -n server python -m pytest tests/unit tests/auth -q
conda run -n server uvicorn app.main:app --host 0.0.0.0 --port 8000 --reload
```
如需要进入环境:
```bash
conda activate server
```
## 常用命令
```bash
conda run -n server python -m pytest tests -q
conda run -n server python scripts/run_server.py
docker build -t tjwater-server:local .
docker compose -f infra/docker/docker-compose.yml config
```
- `pytest`:运行自动化测试。
- `scripts/run_server.py`:使用项目脚本启动服务。
- `docker build`:构建后端镜像。
- `docker compose config`:检查 compose 配置和变量展开。
## 开发规范
- Python 文件、函数、变量、Pydantic 字段、JSON body 字段和 query 参数使用 `snake_case`
- Python 类和 Pydantic 模型使用 `PascalCase`
- 新 HTTP 路径使用 `kebab-case`,例如 `/api/v1/pressure-status/analyze`
- 优先复用现有 FastAPI/service/repository 边界。
- 不要把临时数据、数据库 dump、日志或本地运行产物纳入提交。
## 项目数据库路由
项目级 REST 请求通过 `X-Project-Id` 解析元数据中的数据库配置:
- `biz_data` DSN 用于管网业务数据。每个物理业务库使用同名 `_template` 数据库,通过逻辑订阅只同步 `network` schema;模拟临时库从该项目模板克隆。`WNDB_SCHEMA_TEMPLATE_DB_NAME`(当前为 `tjwater_v2_schema_template`)仅用于创建空业务库和 INP 导入暂存库。
- `iot_data` DSN 用于 TimescaleDB,始终使用元数据配置的完整 DSN,不再从项目代码推导数据库名。
- 元数据、业务库和 TimescaleDB 可以部署在同一主机,也可以分别部署。
完整新建供水项目使用 `POST /api/v1/admin/project-provisions`,以
`multipart/form-data` 同时提交 `name`、小写 `code`、可选的
`description``gs_workspace``map_zoom` 和 INP `file`。工作流会按顺序完成:
1. EPANET 校验 INP,并从空结构模板创建业务库、导入模型;
2. 创建同名 `_template`,复制 32 张 `network` 表并建立逻辑订阅;
3.`TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME` 创建空时序库;
4. 创建 GeoServer 工作空间、PostGIS 数据存储和 7 个 GIS 图层,将客户端缓存设为 `GEOSERVER_CLIENT_CACHE_SECONDS`
5. 最后在一个元数据事务中写入项目、两条加密数据库路由和创建者成员关系,并将项目设为 `active`
基础设施任一步失败时按 GeoServer、时序库、管网模板、业务库的逆序清理;元数据提交失败也执行同样清理。旧 `POST /admin/projects` 仅保留给已经由外部流程创建好的资源登记使用,并已标记为 deprecated。
使用模板复制或临时方案库的模拟功能时,`biz_data` 账号必须具备数据库创建和删除权限;只有显式删除项目时才会终止该项目的现有数据库会话,普通复制不会主动中断复制源会话。
## 测试与发布
提交前根据改动范围运行最小有效测试:
```bash
conda run -n server python -m pytest tests/unit tests/auth -q
```
发布镜像前建议运行:
```bash
docker build -t tjwater-server:local .
```
Gitea 包工作流位于 `.gitea/workflows/package.yml`,通常由 tag 触发构建、推送镜像并通知部署 webhook。
## 安全规则
不要提交 `.env`、客户数据、数据库 dump、日志、生成缓存、`db_inp/``temp/``data/` 或本地密钥。CI/CD 凭据应放在 Gitea secrets 和仓库变量中。
+4 -35
View File
@@ -1,36 +1,5 @@
from app.algorithms.cleaning import flow_data_clean, pressure_data_clean
from app.algorithms.sensor import (
pressure_sensor_placement_sensitivity,
pressure_sensor_placement_kmeans,
)
from app.algorithms.isolation.valve import valve_isolation_analysis
from app.algorithms.leakage import LeakageIdentifier
from app.algorithms.health import PipelineHealthAnalyzer
from app.algorithms.burst_location import run_burst_location
from app.algorithms.simulation.scenarios import (
convert_to_local_unit,
burst_analysis,
valve_close_analysis,
flushing_analysis,
contaminant_simulation,
age_analysis,
pressure_regulation,
)
"""Pure water-network calculation packages.
__all__ = [
"flow_data_clean",
"pressure_data_clean",
"pressure_sensor_placement_sensitivity",
"pressure_sensor_placement_kmeans",
"convert_to_local_unit",
"burst_analysis",
"valve_close_analysis",
"flushing_analysis",
"contaminant_simulation",
"age_analysis",
"pressure_regulation",
"valve_isolation_analysis",
"LeakageIdentifier",
"PipelineHealthAnalyzer",
"run_burst_location",
]
Application workflows belong in :mod:`app.services`; database and external
system access belongs in :mod:`app.infra` or :mod:`app.native`.
"""
+2 -2
View File
@@ -1,3 +1,3 @@
from app.algorithms.burst_detection.burst_detector import BurstDetector
from app.algorithms.burst_detection.pressure_anomaly import PressureAnomalyDetector
__all__ = ["BurstDetector"]
__all__ = ["PressureAnomalyDetector"]
@@ -17,7 +17,7 @@ PressureDataInput = (
IGNORED_OBSERVATION_COLUMNS = {"time", "timestamp", "datetime", "date"}
class BurstDetector:
class PressureAnomalyDetector:
"""FFT + IsolationForest based burst detection for daily aligned pressure data."""
def __init__(
@@ -0,0 +1,3 @@
from .pipeline import run_burst_location
__all__ = ["run_burst_location"]
@@ -11,13 +11,13 @@ import networkx as nx
import numpy as np
import pandas as pd
from .leak_simulator import cal_signature_pipe_multi_pf
from .network_partitioner import (
from .leak_signature import cal_signature_pipe_multi_pf
from .topology_partitioning import (
cal_group_num,
metis_grouping_pipe_weight,
visualize_metis_partition,
)
from .similarity_calculator import (
from .similarity_metrics import (
adjust_ratio,
cal_similarity_all_multi_new_sq_improve_double_lzr,
decode_mode,
@@ -769,4 +769,3 @@ def DN_search_multi_simple_add_flow_count_new(
final_candidates_csv,
)
@@ -1,5 +1,3 @@
import argparse
import json
import logging
from multiprocessing import cpu_count
from pathlib import Path
@@ -7,12 +5,12 @@ from typing import Any, Iterable
import pandas as pd
from app.algorithms.burst_location import leak_simulator
from app.algorithms.burst_localization import leak_signature
from .burst_locator import (
from .candidate_ranking import (
DN_search_multi_simple_add_flow_count_new,
)
from .network_model import (
from .topology_model import (
_build_node_pipe_maps,
cal_node_coordinate,
construct_graph,
@@ -26,35 +24,6 @@ DEFAULT_N_WORKERS = max(1, min(cpu_count() - 1, 4))
logger = logging.getLogger(__name__)
def _read_id_list_json(path):
if path is None:
return None
data = json.loads(Path(path).read_text(encoding="utf-8"))
if isinstance(data, list):
return [str(item) for item in data]
if isinstance(data, dict):
if "ids" in data and isinstance(data["ids"], list):
return [str(item) for item in data["ids"]]
raise ValueError(f"ID JSON must be list or dict with key 'ids': {path}")
raise ValueError(f"Unsupported ID JSON format: {path}")
def _read_series_csv(path):
if path is None:
return None
df = pd.read_csv(path)
if df.shape[1] < 2:
raise ValueError(f"CSV must contain at least two columns (id,value): {path}")
if {"id", "value"}.issubset(df.columns):
id_col, value_col = "id", "value"
else:
id_col, value_col = df.columns[0], df.columns[1]
series = pd.Series(
df[value_col].values, index=df[id_col].astype(str).values, dtype=float
)
return series
def _align_scada_series(
series: pd.Series, ids: Iterable[str], series_name: str
) -> pd.Series:
@@ -121,7 +90,7 @@ def run_burst_location(
basic_pressure: float = 10.0,
n_workers: int = DEFAULT_N_WORKERS,
partition_on_full_graph: bool = True,
visualize_partition: bool = True,
visualize_partition: bool = False,
visualize_pause_seconds: float = 0.3,
final_candidates_csv_path: (
str | None
@@ -159,7 +128,7 @@ def run_burst_location(
pipe_diameter,
) = read_inf_inp(wn)
candidate_pipe, _ = leak_simulator.cal_possible_pipe(
candidate_pipe, _ = leak_signature.cal_possible_pipe(
burst_leakage, all_pipe, pipe_diameter
)
@@ -267,76 +236,3 @@ def run_burst_location(
"final_candidates_csv": final_candidates_csv,
"stage_timing_seconds": stage_timing,
}
def _parse_args():
parser = argparse.ArgumentParser(description="爆管定位主函数入口")
parser.add_argument("--wn-inp", required=True, help="EPANET inp 文件路径")
parser.add_argument(
"--pressure-ids-json", required=True, help="压力SCADA ID列表 JSON 文件"
)
parser.add_argument(
"--flow-ids-json", default=None, help="(可选)流量SCADA ID列表 JSON 文件"
)
parser.add_argument(
"--burst-pressure-csv", required=True, help="爆管时压力 CSVid,value"
)
parser.add_argument(
"--normal-pressure-csv", required=True, help="正常时压力 CSVid,value"
)
parser.add_argument(
"--burst-flow-csv", default=None, help="(可选)爆管时流量 CSV(id,value"
)
parser.add_argument(
"--normal-flow-csv", default=None, help="(可选)正常时流量 CSV(id,value"
)
parser.add_argument(
"--burst-leakage", type=float, required=True, help="爆管漏损流量"
)
parser.add_argument(
"--min-dpressure",
type=float,
default=2.0,
help="(可选)最小压降阈值,默认 2.0",
)
parser.add_argument(
"--basic-pressure",
type=float,
default=10.0,
help="(可选)基础服务压力,默认 10.0",
)
parser.add_argument(
"--n-workers",
type=int,
default=DEFAULT_N_WORKERS,
help="(可选)特征中心模拟进程数,默认 max(1, min(cpu_count()-1, 4))",
)
parser.add_argument(
"--final-candidates-csv-path",
default="temp/burst_location/final_round_candidates.csv",
help="(可选)最后一轮候选管道明细 CSV 输出路径",
)
return parser.parse_args()
def main():
args = _parse_args()
result = run_burst_location(
wn_inp_path=args.wn_inp,
pressure_scada_ids=_read_id_list_json(args.pressure_ids_json),
burst_pressure=_read_series_csv(args.burst_pressure_csv),
normal_pressure=_read_series_csv(args.normal_pressure_csv),
burst_leakage=args.burst_leakage,
flow_scada_ids=_read_id_list_json(args.flow_ids_json),
burst_flow=_read_series_csv(args.burst_flow_csv),
normal_flow=_read_series_csv(args.normal_flow_csv),
min_dpressure=args.min_dpressure,
basic_pressure=args.basic_pressure,
n_workers=args.n_workers,
final_candidates_csv_path=args.final_candidates_csv_path,
)
print(json.dumps(result, ensure_ascii=False))
if __name__ == "__main__":
main()
@@ -6,7 +6,7 @@ import random
import numpy as np
import pandas as pd
from .leak_simulator import simple_add_leak, simple_recover_wn, simple_simulation_pf
from .leak_signature import simple_add_leak, simple_recover_wn, simple_simulation_pf
def add_noise_pd(data, noise_type, noise_para):
@@ -195,4 +195,3 @@ def change_para_of_wn(wn, pipe_roughness_change):
pipe.roughness = pipe_roughness_change[pipe_name]
return wn
@@ -1,3 +0,0 @@
from .burst_location import run_burst_location
__all__ = ["run_burst_location"]
-59
View File
@@ -1,59 +0,0 @@
import os
from app.algorithms.cleaning import flow as _flow_module
from app.algorithms.cleaning import pressure as _pressure_module
############################################################
# 流量监测数据清洗 ***卡尔曼滤波法***
############################################################
# 2025/08/21 hxyan
def flow_data_clean(input_csv_file: str) -> str:
"""
读取 input_csv_path 中的每列时间序列,使用一维 Kalman 滤波平滑并用预测值替换基于 3σ 检测出的异常点。
保存输出为:<input_filename>_cleaned.xlsx(与输入同目录),并返回输出文件的绝对路径。如有同名文件存在,则覆盖。
:param: input_csv_file: 输入的 CSV 文件明或路径
:return: 输出文件的绝对路径
"""
# 提供的 input_csv_path 绝对路径,以下为 默认脚本目录下同名 CSV 文件,构建绝对路径,可根据情况修改
# 使用 algorithms 根目录保持与原 data_cleaning.py 一致的行为
script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
input_csv_path = os.path.join(script_dir, input_csv_file)
# 检查文件是否存在
if not os.path.exists(input_csv_path):
raise FileNotFoundError(f"指定的文件不存在: {input_csv_path}")
# 调用 clean_flow_data_kf 函数进行数据清洗
out_xlsx_path = _flow_module.clean_flow_data_kf(input_csv_path)
print("清洗后的数据已保存到:", out_xlsx_path)
############################################################
# 压力监测数据清洗 ***kmean++法***
############################################################
# 2025/08/21 hxyan
def pressure_data_clean(input_csv_file: str) -> str:
"""
读取 input_csv_path 中的每列时间序列,使用Kmean++清洗数据。
保存输出为:<input_filename>_cleaned.xlsx(与输入同目录),并返回输出文件的绝对路径。如有同名文件存在,则覆盖。
原始数据在 sheet 'raw_pressure_data',处理后数据在 sheet 'cleaned_pressusre_data'
:param input_csv_path: 输入的 CSV 文件路径
:return: 输出文件的绝对路径
"""
# 提供的 input_csv_path 绝对路径,以下为 默认脚本目录下同名 CSV 文件,构建绝对路径,可根据情况修改
# 使用 algorithms 根目录保持与原 data_cleaning.py 一致的行为
script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
input_csv_path = os.path.join(script_dir, input_csv_file)
# 检查文件是否存在
if not os.path.exists(input_csv_path):
raise FileNotFoundError(f"指定的文件不存在: {input_csv_path}")
# 调用 clean_pressure_data_km 函数进行数据清洗
out_xlsx_path = _pressure_module.clean_pressure_data_km(input_csv_path)
print("清洗后的数据已保存到:", out_xlsx_path)
-283
View File
@@ -1,283 +0,0 @@
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
from sklearn.cluster import KMeans
from sklearn.impute import SimpleImputer
import os
from app.algorithms._utils import fill_time_gaps
def clean_pressure_data_km(
input_csv_path: str, show_plot: bool = False, fill_gaps: bool = True
) -> str:
"""
读取输入 CSV,基于 KMeans 检测异常并用滚动平均修复。输出为 <input_basename>_cleaned.xlsx(同目录)。
原始数据在 sheet 'raw_pressure_data',处理后数据在 sheet 'cleaned_pressusre_data'
返回输出文件的绝对路径。
Args:
input_csv_path: CSV 文件路径
show_plot: 是否显示可视化
fill_gaps: 是否先补齐时间缺口(默认 True)
"""
# 读取 CSV
input_csv_path = os.path.abspath(input_csv_path)
data = pd.read_csv(input_csv_path, header=0, index_col=None, encoding="utf-8")
# 补齐时间缺口(如果数据包含 time 列)
if fill_gaps and "time" in data.columns:
data = fill_time_gaps(
data, time_col="time", freq="1min", short_gap_threshold=10
)
# 分离时间列和数值列
time_col_data = None
if "time" in data.columns:
time_col_data = data["time"]
data = data.drop(columns=["time"])
# 标准化
data_norm = (data - data.mean()) / data.std()
# 聚类与异常检测
k = 3
kmeans = KMeans(n_clusters=k, init="k-means++", n_init=50, random_state=42)
clusters = kmeans.fit_predict(data_norm)
centers = kmeans.cluster_centers_
distances = np.linalg.norm(data_norm.values - centers[clusters], axis=1)
threshold = distances.mean() + 3 * distances.std()
anomaly_pos = np.where(distances > threshold)[0]
anomaly_indices = data.index[anomaly_pos]
anomaly_details = {}
for pos in anomaly_pos:
row_norm = data_norm.iloc[pos]
cluster_idx = clusters[pos]
center = centers[cluster_idx]
diff = abs(row_norm - center)
main_sensor = diff.idxmax()
anomaly_details[data.index[pos]] = main_sensor
# 修复:滚动平均(窗口可调)
data_rolled = data.rolling(window=13, center=True, min_periods=1).mean()
data_repaired = data.copy()
for pos in anomaly_pos:
label = data.index[pos]
sensor = anomaly_details[label]
data_repaired.loc[label, sensor] = data_rolled.loc[label, sensor]
# 可选可视化(使用位置作为 x 轴)
plt.rcParams["font.sans-serif"] = ["SimHei"]
plt.rcParams["axes.unicode_minus"] = False
if show_plot and len(data.columns) > 0:
n = len(data)
time = np.arange(n)
plt.figure(figsize=(12, 8))
for col in data.columns:
plt.plot(time, data[col].values, marker="o", markersize=3, label=col)
for pos in anomaly_pos:
sensor = anomaly_details[data.index[pos]]
plt.plot(pos, data.iloc[pos][sensor], "ro", markersize=8)
plt.xlabel("时间点(序号)")
plt.ylabel("压力监测值")
plt.title("各传感器折线图(红色标记主要异常点)")
plt.legend()
plt.show()
plt.figure(figsize=(12, 8))
for col in data_repaired.columns:
plt.plot(
time, data_repaired[col].values, marker="o", markersize=3, label=col
)
for pos in anomaly_pos:
sensor = anomaly_details[data.index[pos]]
plt.plot(pos, data_repaired.iloc[pos][sensor], "go", markersize=8)
plt.xlabel("时间点(序号)")
plt.ylabel("修复后压力监测值")
plt.title("修复后各传感器折线图(绿色标记修复值)")
plt.legend()
plt.show()
# 保存到 Excel:两个 sheet
input_dir = os.path.dirname(os.path.abspath(input_csv_path))
input_base = os.path.splitext(os.path.basename(input_csv_path))[0]
output_filename = f"{input_base}_cleaned.xlsx"
output_path = os.path.join(input_dir, output_filename)
# 如果原始数据包含时间列,将其添加回结果
data_for_save = data.copy()
data_repaired_for_save = data_repaired.copy()
if time_col_data is not None:
data_for_save.insert(0, "time", time_col_data)
data_repaired_for_save.insert(0, "time", time_col_data)
if os.path.exists(output_path):
os.remove(output_path) # 覆盖同名文件
with pd.ExcelWriter(output_path, engine="openpyxl") as writer:
data_for_save.to_excel(writer, sheet_name="raw_pressure_data", index=False)
data_repaired_for_save.to_excel(
writer, sheet_name="cleaned_pressusre_data", index=False
)
# 返回输出文件的绝对路径
return os.path.abspath(output_path)
def clean_pressure_data_df_km(data: pd.DataFrame, show_plot: bool = False) -> dict:
"""
接收一个 DataFrame 数据结构,使用KMeans聚类检测异常并用滚动平均修复。
返回清洗后的字典数据结构。
Args:
data: 输入 DataFrame(可包含 time 列)
show_plot: 是否显示可视化
"""
# 使用传入的 DataFrame
data = data.copy()
# 补齐时间缺口(如果启用且数据包含 time 列)
data_filled = fill_time_gaps(
data, time_col="time", freq="1min", short_gap_threshold=10
)
# 保存 time 列用于最后合并
time_col_series = None
if "time" in data_filled.columns:
time_col_series = data_filled["time"]
# 移除 time 列用于后续清洗
data_filled = data_filled.drop(columns=["time"])
# 标准化(使用填充后的数据)
data_norm = (data_filled - data_filled.mean()) / data_filled.std()
# 添加:处理标准化后的 NaN(例如,标准差为0的列),防止异常数据,时间段内所有数据都相同导致计算结果为 NaN
imputer = SimpleImputer(
strategy="constant", fill_value=0, keep_empty_features=True
) # 用 0 填充 NaN,包括全 NaN,并保留空特征
data_norm = pd.DataFrame(
imputer.fit_transform(data_norm),
columns=data_norm.columns,
index=data_norm.index,
)
# 聚类与异常检测
k = 3
kmeans = KMeans(n_clusters=k, init="k-means++", n_init=50, random_state=42)
clusters = kmeans.fit_predict(data_norm)
centers = kmeans.cluster_centers_
distances = np.linalg.norm(data_norm.values - centers[clusters], axis=1)
threshold = distances.mean() + 3 * distances.std()
anomaly_pos = np.where(distances > threshold)[0]
anomaly_indices = data_filled.index[anomaly_pos]
anomaly_details = {}
for pos in anomaly_pos:
row_norm = data_norm.iloc[pos]
cluster_idx = clusters[pos]
center = centers[cluster_idx]
diff = abs(row_norm - center)
main_sensor = diff.idxmax()
anomaly_details[data_filled.index[pos]] = main_sensor
# 修复:滚动平均(窗口可调)
data_rolled = data_filled.rolling(window=13, center=True, min_periods=1).mean()
data_repaired = data_filled.copy()
for pos in anomaly_pos:
label = data_filled.index[pos]
sensor = anomaly_details[label]
data_repaired.loc[label, sensor] = data_rolled.loc[label, sensor]
# 可选可视化(使用位置作为 x 轴)
plt.rcParams["font.sans-serif"] = ["SimHei"]
plt.rcParams["axes.unicode_minus"] = False
if show_plot and len(data.columns) > 0:
n = len(data)
time = np.arange(n)
n_filled = len(data_filled)
time_filled = np.arange(n_filled)
plt.figure(figsize=(12, 8))
for col in data.columns:
plt.plot(
time, data[col].values, marker="o", markersize=3, label=col, alpha=0.5
)
for col in data_filled.columns:
plt.plot(
time_filled,
data_filled[col].values,
marker="x",
markersize=3,
label=f"{col}_filled",
linestyle="--",
)
for pos in anomaly_pos:
sensor = anomaly_details[data_filled.index[pos]]
plt.plot(pos, data_filled.iloc[pos][sensor], "ro", markersize=8)
plt.xlabel("时间点(序号)")
plt.ylabel("压力监测值")
plt.title("各传感器折线图(红色标记主要异常点,虚线为0值填充后)")
plt.legend()
plt.show()
plt.figure(figsize=(12, 8))
for col in data_repaired.columns:
plt.plot(
time_filled, data_repaired[col].values, marker="o", markersize=3, label=col
)
for pos in anomaly_pos:
sensor = anomaly_details[data_filled.index[pos]]
plt.plot(pos, data_repaired.iloc[pos][sensor], "go", markersize=8)
plt.xlabel("时间点(序号)")
plt.ylabel("修复后压力监测值")
plt.title("修复后各传感器折线图(绿色标记修复值)")
plt.legend()
plt.show()
# 将 time 列添加回结果
if time_col_series is not None:
data_repaired.insert(0, "time", time_col_series)
# 返回清洗后的字典
return data_repaired
# 测试
# if __name__ == "__main__":
# # 默认使用脚本目录下的 pressure_raw_data.csv
# script_dir = os.path.dirname(os.path.abspath(__file__))
# default_csv = os.path.join(script_dir, "pressure_raw_data.csv")
# out_path = clean_pressure_data_km(default_csv, show_plot=False)
# print("保存路径:", out_path)
# 测试 clean_pressure_data_dict_km 函数
if __name__ == "__main__":
import random
# 读取 szh_pressure_scada.csv 文件
script_dir = os.path.dirname(os.path.abspath(__file__))
csv_path = os.path.join(script_dir, "szh_pressure_scada.csv")
data = pd.read_csv(csv_path, header=0, index_col=None, encoding="utf-8")
# 排除 Time 列,随机选择 5 列
columns_to_exclude = ["Time"]
available_columns = [col for col in data.columns if col not in columns_to_exclude]
selected_columns = random.sample(available_columns, 5)
# 将选中的列转换为字典
data_dict = {col: data[col].tolist() for col in selected_columns}
print("选中的列:", selected_columns)
print("原始数据长度:", len(data_dict[selected_columns[0]]))
# 调用函数进行清洗
cleaned_dict = clean_pressure_data_df_km(data_dict, show_plot=True)
print("清洗后的字典键:", list(cleaned_dict.keys()))
print("清洗后的数据长度:", len(cleaned_dict[selected_columns[0]]))
print("测试完成:函数运行正常")
@@ -0,0 +1,5 @@
"""Demand allocation calculations."""
from .pipe_length_weighted import allocate_demand_by_pipe_length
__all__ = ["allocate_demand_by_pipe_length"]
@@ -0,0 +1,36 @@
"""Pipe-length-weighted demand allocation.
This module deliberately accepts plain topology data and performs no database
or file access. Application services are responsible for loading topology.
"""
from typing import Any, Mapping
def allocate_demand_by_pipe_length(
demand: float,
topology_nodes: Mapping[str, Mapping[str, Any]],
topology_links: Mapping[str, Mapping[str, Any]],
) -> dict[str, float]:
"""Allocate total demand to junctions by half of each incident link length."""
if not topology_nodes or not topology_links or demand == 0.0:
return {}
total_link_length = sum(
abs(float(link["length"])) for link in topology_links.values()
)
if total_link_length <= 0.0:
return {}
demand_per_length = demand / total_link_length
result: dict[str, float] = {}
for node_id, node in topology_nodes.items():
if node["type"] != "junction":
continue
incident_length = sum(
abs(float(topology_links[link_id]["length"]))
for link_id in node["links"]
)
result[node_id] = incident_length * demand_per_length * 0.5
return result
@@ -0,0 +1,3 @@
from app.algorithms.dma_leakage_estimation.genetic_optimizer import DmaLeakageOptimizer
__all__ = ["DmaLeakageOptimizer"]
@@ -3,7 +3,6 @@ import numpy as np
import pandas as pd
import os
import time
import argparse
from multiprocessing import Pool, cpu_count
from typing import Any, List, Dict, Union
@@ -70,7 +69,7 @@ def _worker_init(
def _worker_evaluate(raw_ratios: np.ndarray) -> float:
d = _worker_data
effective_ratio_map = LeakageIdentifier._effective_area_ratios(
effective_ratio_map = DmaLeakageOptimizer._effective_area_ratios(
raw_ratios,
d["area_ids"],
d["nodes_by_area"],
@@ -121,13 +120,15 @@ def _worker_evaluate(raw_ratios: np.ndarray) -> float:
_cleanup_temp_files(prefix)
class LeakageIdentifier:
FLOW_UNIT_TO_M3S = {
"m3/s": 1.0,
"m3/h": 1.0 / 3600.0,
"L/s": 1.0 / 1000.0,
"L/min": 1.0 / 60000.0,
}
class DmaLeakageOptimizer:
FLOW_UNIT_TO_M3S = {
"m3/s": 1.0,
"m³/s": 1.0,
"m3/h": 1.0 / 3600.0,
"m³/h": 1.0 / 3600.0,
"L/s": 1.0 / 1000.0,
"L/min": 1.0 / 60000.0,
}
@classmethod
def _flow_to_m3s(cls, value: float, unit: str) -> float:
@@ -541,7 +542,7 @@ class LeakageProblem(Problem):
leak_ratios = x
# 将漏损分布归一化
effective_ratio_map = LeakageIdentifier._effective_area_ratios(
effective_ratio_map = DmaLeakageOptimizer._effective_area_ratios(
leak_ratios,
self.area_ids,
self.nodes_by_area,
@@ -603,51 +604,6 @@ class LeakageProblem(Problem):
def close(self) -> None:
if self._pool is not None:
self._pool.close()
self._pool.join()
self._pool = None
def main() -> int:
parser = argparse.ArgumentParser(description="漏损区域识别")
parser.add_argument("--inp", required=True, help=".inp 文件路径")
parser.add_argument("--map", help="节点-区域映射 CSV 路径")
parser.add_argument("--scada", help="SCADA 压力 CSV 路径 (观测数据)")
parser.add_argument("--sensors", help="传感器节点 ID 列表 (逗号分隔)")
parser.add_argument("--output", default="Results", help="输出目录")
parser.add_argument("--pop_size", type=int, default=50, help="种群大小")
parser.add_argument("--max_gen", type=int, default=100, help="最大代数")
parser.add_argument("--duration", type=float, default=24, help="模拟时长(小时)")
parser.add_argument("--q_sum", type=float, default=0.241, help="总漏损流量")
parser.add_argument(
"--q_sum_unit",
default="m3/s",
choices=list(LeakageIdentifier.FLOW_UNIT_TO_M3S.keys()),
help="q_sum 输入单位(建议与现场习惯一致,内部统一换算为 m3/s)",
)
args = parser.parse_args()
if not args.map or not args.scada or not args.sensors:
parser.error("--map、--scada、--sensors 为必填")
q_sum_m3s = LeakageIdentifier._flow_to_m3s(args.q_sum, args.q_sum_unit)
sensors = [sensor.strip() for sensor in args.sensors.split(",") if sensor.strip()]
identifier = LeakageIdentifier(
args.inp, sensors, args.map, duration=args.duration, q_sum=q_sum_m3s
)
identifier.run_identification(
args.scada,
args.output,
pop_size=args.pop_size,
max_gen=args.max_gen,
output_flow_unit=args.q_sum_unit,
)
return 0
if __name__ == "__main__":
raise SystemExit(main())
self._pool.close()
self._pool.join()
self._pool = None
@@ -0,0 +1,206 @@
"""Pure topology partitioning used by DMA leakage estimation."""
import math
from collections import deque
from typing import Any, Iterable, Mapping
import numpy as np
def build_dma_partitions(
sensor_nodes: list[str],
node_coords: Mapping[str, Mapping[str, Any]],
link_entries: Iterable[str],
dma_count: int | None,
) -> tuple[dict[str, str], list[dict[str, Any]]]:
"""Assign every topology node to a sensor-seeded virtual DMA."""
all_nodes = list(node_coords)
if not all_nodes:
raise ValueError("管网中未获取到可分区节点。")
available_sensors = [node for node in sensor_nodes if node in node_coords]
if not available_sensors:
raise ValueError("无可用压力传感器,无法生成虚拟分区。")
area_count = _resolve_dma_count(dma_count, available_sensors, all_nodes)
sensor_area_map = _cluster_sensors_to_areas(
available_sensors, node_coords, area_count
)
adjacency = _build_adjacency(link_entries, all_nodes)
distance_by_sensor = {
sensor: _bfs_distances(adjacency, sensor) for sensor in available_sensors
}
assignment_count = {sensor: 0 for sensor in available_sensors}
area_map: dict[str, str] = {}
for node_id in sorted(all_nodes):
sensor = _choose_sensor_for_node(
node_id,
available_sensors,
node_coords,
distance_by_sensor,
assignment_count,
)
assignment_count[sensor] += 1
area_map[node_id] = sensor_area_map[sensor]
return area_map, _build_area_meta(area_map, sensor_area_map)
def _resolve_dma_count(
dma_count: int | None, sensor_nodes: list[str], all_nodes: list[str]
) -> int:
if dma_count is None:
return min(len(sensor_nodes), len(all_nodes))
if dma_count <= 0:
raise ValueError("dma_count 必须大于 0。")
if dma_count > len(all_nodes):
raise ValueError("dma_count 不能大于可分区节点数量。")
if dma_count > len(sensor_nodes):
raise ValueError("dma_count 不能大于可用传感器数量。")
return dma_count
def _cluster_sensors_to_areas(
sensor_nodes: list[str],
node_coords: Mapping[str, Mapping[str, Any]],
area_count: int,
) -> dict[str, str]:
if area_count >= len(sensor_nodes):
return {sensor: str(index + 1) for index, sensor in enumerate(sensor_nodes)}
points = np.array(
[
[float(node_coords[sensor]["x"]), float(node_coords[sensor]["y"])]
for sensor in sensor_nodes
],
dtype=float,
)
centers = points[:area_count].copy()
labels = np.full(points.shape[0], -1, dtype=int)
for _ in range(20):
distances_squared = (
(points[:, None, :] - centers[None, :, :]) ** 2
).sum(axis=2)
next_labels = distances_squared.argmin(axis=1)
if np.array_equal(labels, next_labels):
break
labels = next_labels
for index in range(area_count):
cluster_points = points[labels == index]
if cluster_points.size > 0:
centers[index] = cluster_points.mean(axis=0)
labels = _restore_empty_area_labels(labels, points, centers, area_count)
return {
sensor: str(int(labels[index]) + 1)
for index, sensor in enumerate(sensor_nodes)
}
def _restore_empty_area_labels(
labels: np.ndarray,
points: np.ndarray,
centers: np.ndarray,
area_count: int,
) -> np.ndarray:
"""Keep every requested area represented when coordinates are degenerate."""
labels = labels.copy()
for missing_area in sorted(set(range(area_count)) - set(labels.tolist())):
area_sizes = {
area: int(np.count_nonzero(labels == area)) for area in range(area_count)
}
donor_area = max(
(area for area, size in area_sizes.items() if size > 1),
key=lambda area: (area_sizes[area], -area),
)
donor_indices = np.flatnonzero(labels == donor_area)
replacement_index = max(
(int(index) for index in donor_indices),
key=lambda index: (
float(np.sum((points[index] - centers[donor_area]) ** 2)),
index,
),
)
labels[replacement_index] = missing_area
centers[missing_area] = points[replacement_index]
return labels
def _build_adjacency(
link_entries: Iterable[str], all_nodes: list[str]
) -> dict[str, set[str]]:
adjacency: dict[str, set[str]] = {node: set() for node in all_nodes}
for link in link_entries:
parts = str(link).split(":")
if len(parts) < 4:
continue
node1, node2 = parts[-2], parts[-1]
if node1 in adjacency and node2 in adjacency:
adjacency[node1].add(node2)
adjacency[node2].add(node1)
return adjacency
def _bfs_distances(adjacency: Mapping[str, set[str]], start: str) -> dict[str, int]:
distances = {start: 0}
queue: deque[str] = deque([start])
while queue:
node = queue.popleft()
for neighbor in adjacency.get(node, set()):
if neighbor in distances:
continue
distances[neighbor] = distances[node] + 1
queue.append(neighbor)
return distances
def _choose_sensor_for_node(
node_id: str,
sensors: list[str],
node_coords: Mapping[str, Mapping[str, Any]],
distance_by_sensor: Mapping[str, Mapping[str, int]],
assignment_count: Mapping[str, int],
) -> str:
min_distance: int | None = None
candidates: list[str] = []
for sensor in sensors:
distance = distance_by_sensor.get(sensor, {}).get(node_id)
if distance is None:
continue
if min_distance is None or distance < min_distance:
min_distance = distance
candidates = [sensor]
elif distance == min_distance:
candidates.append(sensor)
if not candidates:
node_coord = node_coords[node_id]
return min(
sensors,
key=lambda sensor: math.hypot(
float(node_coord["x"]) - float(node_coords[sensor]["x"]),
float(node_coord["y"]) - float(node_coords[sensor]["y"]),
),
)
return min(candidates, key=lambda sensor: (assignment_count[sensor], sensor))
def _build_area_meta(
area_map: Mapping[str, str], sensor_area_map: Mapping[str, str]
) -> list[dict[str, Any]]:
nodes_by_area: dict[str, list[str]] = {}
for node_id, area_id in area_map.items():
nodes_by_area.setdefault(area_id, []).append(node_id)
sensors_by_area: dict[str, list[str]] = {}
for sensor, area_id in sensor_area_map.items():
sensors_by_area.setdefault(area_id, []).append(sensor)
return [
{
"area_id": area_id,
"sensor_nodes": sorted(sensors_by_area.get(area_id, [])),
"node_ids": sorted(nodes_by_area[area_id]),
"node_count": len(nodes_by_area[area_id]),
}
for area_id in sorted(nodes_by_area, key=int)
]
-3
View File
@@ -1,3 +0,0 @@
from app.algorithms.health.analyzer import PipelineHealthAnalyzer
__all__ = ["PipelineHealthAnalyzer"]
-3
View File
@@ -1,3 +0,0 @@
from app.algorithms.isolation.valve import valve_isolation_analysis
__all__ = ["valve_isolation_analysis"]
-165
View File
@@ -1,165 +0,0 @@
from collections import defaultdict, deque
from functools import lru_cache
from typing import Any
from app.services.tjnetwork import (
get_network_link_nodes,
is_node,
get_link_properties,
)
VALVE_LINK_TYPE = "valve"
def _parse_link_entry(link_entry: str) -> tuple[str, str, str, str]:
parts = link_entry.split(":", 3)
if len(parts) != 4:
raise ValueError(f"Invalid link entry format: {link_entry}")
return parts[0], parts[1], parts[2], parts[3]
@lru_cache(maxsize=16)
def _get_network_topology(network: str):
"""
解析并缓存网络拓扑,大幅减少重复的 API 调用和字符串解析开销。
返回:
- pipe_adj: 永久连通的管道/泵邻接表 (dict[str, set])
- all_valves: 所有阀门字典 {id: (n1, n2)}
- link_lookup: 链路快速查表 {id: (n1, n2, type)} 用于快速定位事故点
- node_set: 所有已知节点集合
"""
pipe_adj = defaultdict(set)
all_valves = {}
link_lookup = {}
node_set = set()
# 此处假设 get_network_link_nodes 获取全网数据
for link_entry in get_network_link_nodes(network):
link_id, link_type, node1, node2 = _parse_link_entry(link_entry)
link_type_name = str(link_type).lower()
link_lookup[link_id] = (node1, node2, link_type_name)
node_set.add(node1)
node_set.add(node2)
if link_type_name == VALVE_LINK_TYPE:
all_valves[link_id] = (node1, node2)
else:
# 只有非阀门(管道/泵)才进入永久连通图
pipe_adj[node1].add(node2)
pipe_adj[node2].add(node1)
return pipe_adj, all_valves, link_lookup, node_set
def valve_isolation_analysis(
network: str, accident_elements: str | list[str], disabled_valves: list[str] = None
) -> dict[str, Any]:
"""
关阀搜索/分析:基于拓扑结构确定事故隔离所需关阀。
:param network: 模型名称
:param accident_elements: 事故点(节点或管道/泵/阀门ID),可以是单个ID字符串或ID列表
:param disabled_valves: 故障/无法关闭的阀门ID列表
:return: dict,包含受影响节点、必须关闭阀门、可选阀门等信息
"""
if disabled_valves is None:
disabled_valves_set = set()
else:
disabled_valves_set = set(disabled_valves)
if isinstance(accident_elements, str):
target_elements = [accident_elements]
else:
target_elements = accident_elements
# 1. 获取缓存拓扑 (极快,无 IO)
pipe_adj, all_valves, link_lookup, node_set = _get_network_topology(network)
# 2. 确定起点,优先查表避免 API 调用
start_nodes = set()
for element in target_elements:
if element in node_set:
start_nodes.add(element)
elif element in link_lookup:
n1, n2, _ = link_lookup[element]
start_nodes.add(n1)
start_nodes.add(n2)
else:
# 仅当缓存中没找到时(极少见),才回退到慢速 API
if is_node(network, element):
start_nodes.add(element)
else:
props = get_link_properties(network, element)
n1, n2 = props.get("node1"), props.get("node2")
if n1 and n2:
start_nodes.add(n1)
start_nodes.add(n2)
else:
raise ValueError(
f"Accident element {element} invalid or missing endpoints"
)
# 3. 处理故障阀门 (构建临时增量图)
# 我们不修改 cached pipe_adj,而是建立一个 extra_adj
extra_adj = defaultdict(list)
boundary_valves = {} # 当前有效的边界阀门
for vid, (n1, n2) in all_valves.items():
if vid in disabled_valves_set:
# 故障阀门:视为连通管道
extra_adj[n1].append(n2)
extra_adj[n2].append(n1)
else:
# 正常阀门:视为潜在边界
boundary_valves[vid] = (n1, n2)
# 4. BFS 搜索 (叠加 pipe_adj 和 extra_adj)
affected_nodes: set[str] = set()
queue = deque(start_nodes)
while queue:
node = queue.popleft()
if node in affected_nodes:
continue
affected_nodes.add(node)
# 遍历永久管道邻居
if node in pipe_adj:
for neighbor in pipe_adj[node]:
if neighbor not in affected_nodes:
queue.append(neighbor)
# 遍历故障阀门带来的额外邻居
if node in extra_adj:
for neighbor in extra_adj[node]:
if neighbor not in affected_nodes:
queue.append(neighbor)
# 5. 结果聚合
must_close_valves: list[str] = []
optional_valves: list[str] = []
for valve_id, (n1, n2) in boundary_valves.items():
in_n1 = n1 in affected_nodes
in_n2 = n2 in affected_nodes
if in_n1 and in_n2:
optional_valves.append(valve_id)
elif in_n1 or in_n2:
must_close_valves.append(valve_id)
must_close_valves.sort()
optional_valves.sort()
result = {
"accident_elements": target_elements,
"disabled_valves": disabled_valves,
"affected_nodes": sorted(affected_nodes),
"must_close_valves": must_close_valves,
"optional_valves": optional_valves,
"isolatable": len(must_close_valves) > 0,
}
if len(target_elements) == 1:
result["accident_element"] = target_elements[0]
return result
-3
View File
@@ -1,3 +0,0 @@
from app.algorithms.leakage.identifier import LeakageIdentifier
__all__ = ["LeakageIdentifier"]
@@ -0,0 +1,5 @@
from app.algorithms.pipe_health_prediction.survival_predictor import (
PipeHealthSurvivalPredictor,
)
__all__ = ["PipeHealthSurvivalPredictor"]
@@ -4,7 +4,7 @@ import pandas as pd
import matplotlib.pyplot as plt
class PipelineHealthAnalyzer:
class PipeHealthSurvivalPredictor:
"""
管道健康分析器类使用随机生存森林模型预测管道的生存概率
@@ -28,11 +28,6 @@ class PipelineHealthAnalyzer:
"model",
"my_survival_forest_model_quxi.joblib",
)
# 确保 model 目录存在
model_dir = os.path.dirname(model_path)
if model_dir and not os.path.exists(model_dir):
os.makedirs(model_dir, exist_ok=True)
if not os.path.exists(model_path):
raise FileNotFoundError(f"模型文件未找到: {model_path}")
@@ -102,7 +97,7 @@ class PipelineHealthAnalyzer:
# 调用说明示例
"""
在其他项目中使用PipelineHealthAnalyzer类的步骤
在其他项目中使用 PipeHealthSurvivalPredictor 类的步骤
1. 安装依赖在requirements.txt中添加
joblib==1.5.0
@@ -112,34 +107,29 @@ class PipelineHealthAnalyzer:
matplotlib==3.9.4
2. 导入类
from pipeline_health_analyzer import PipelineHealthAnalyzer
from survival_predictor import PipeHealthSurvivalPredictor
3. 初始化分析器替换为实际模型路径
analyzer = PipelineHealthAnalyzer(model_path='path/to/my_survival_forest_model3-10.joblib')
predictor = PipeHealthSurvivalPredictor(model_path='path/to/model.joblib')
4. 准备数据pandas DataFrame包含9个特征列
4. 准备数据pandas DataFrame包含4个特征列
import pandas as pd
data = pd.DataFrame({
'Material': [1, 2], # 示例数据
'Diameter': [100, 150],
'Flow Velocity': [1.5, 2.0],
'Pressure': [50, 60],
'Temperature': [20, 25],
'Precipitation': [0.1, 0.2],
'Location': [1, 2],
'Structural Defects': [0, 1],
'Functional Defects': [0, 0]
'Pressure': [50, 60]
})
5. 进行预测
survival_funcs = analyzer.predict_survival(data)
survival_funcs = predictor.predict_survival(data)
6. 查看结果每个样本的生存概率随时间变化
for i, sf in enumerate(survival_funcs):
print(f"样本 {i+1}: 时间点: {sf.x[:5]}..., 生存概率: {sf.y[:5]}...")
7. 可视化可选
analyzer.plot_survival(survival_funcs, save_path='survival_plot.png')
predictor.plot_survival(survival_funcs, save_path='survival_plot.png')
注意
- 数据格式必须匹配特征列表特征值为数值型
@@ -0,0 +1 @@
"""Pressure sensor placement calculation implementations."""
@@ -0,0 +1,96 @@
import matplotlib.pyplot as plt
import numpy as np
import sklearn.cluster
import wntr
class KMeansPlacement:
def __init__(self, wn, num_monitors: int, min_diameter_mm: float):
self.cluster_num = num_monitors
self.wn = wn
self.monitor_nodes: list[str] = []
self.coords: list[tuple[float, float]] = []
self.candidate_nodes: list[str] = []
self.min_diameter_mm = min_diameter_mm
def get_junctions_coordinates(self) -> None:
eligible_nodes: set[str] = set()
junction_names = set(self.wn.junction_name_list)
for pipe_name in self.wn.pipe_name_list:
pipe = self.wn.get_link(pipe_name)
if float(pipe.diameter) * 1000 < self.min_diameter_mm:
continue
eligible_nodes.update(
node_id
for node_id in (pipe.start_node_name, pipe.end_node_name)
if node_id in junction_names
)
for junction_name in self.wn.junction_name_list:
if junction_name not in eligible_nodes:
continue
junction = self.wn.get_node(junction_name)
self.candidate_nodes.append(junction_name)
self.coords.append(junction.coordinates)
def select_monitoring_points(self) -> list[str]:
if not self.coords:
self.get_junctions_coordinates()
if self.cluster_num <= 0:
raise ValueError("sensor_count must be greater than zero")
if self.cluster_num > len(self.candidate_nodes):
raise ValueError("符合最小管径条件的候选节点数量少于请求的监测点数量")
coords = np.array(self.coords)
coordinate_span = coords.max(axis=0) - coords.min(axis=0)
coordinate_span[coordinate_span == 0] = 1.0
coords_normalized = (coords - coords.min(axis=0)) / coordinate_span
kmeans = sklearn.cluster.KMeans(n_clusters=self.cluster_num, random_state=42)
kmeans.fit(coords_normalized)
selected_indices: set[int] = set()
for cluster_index, center in enumerate(kmeans.cluster_centers_):
cluster_indices = np.flatnonzero(kmeans.labels_ == cluster_index)
available_indices = [
int(index)
for index in cluster_indices
if int(index) not in selected_indices
]
if not available_indices:
available_indices = [
index
for index in range(len(self.candidate_nodes))
if index not in selected_indices
]
nearest_index = min(
available_indices,
key=lambda index: (
float(np.sum((coords_normalized[index] - center) ** 2)),
index,
),
)
selected_indices.add(nearest_index)
nearest_node = self.candidate_nodes[nearest_index]
self.monitor_nodes.append(nearest_node)
return self.monitor_nodes
def visualize_network(self) -> None:
"""Visualize network with monitoring points."""
wntr.graphics.plot_network(
self.wn,
node_attribute=self.monitor_nodes,
node_size=30,
title="Optimal sensor",
)
plt.show()
def optimize_sensor_placement(
network_model: wntr.network.WaterNetworkModel,
sensor_count: int,
min_diameter_mm: float,
) -> list[str]:
"""Select sensor nodes from an already loaded network model."""
placement = KMeansPlacement(network_model, sensor_count, min_diameter_mm)
return placement.select_monitoring_points()
@@ -0,0 +1,905 @@
"""Pressure sensor placement based on scalable sensitivity analysis.
The original implementation expanded a sparse water network into several dense
``node x node``, ``node x pipe``, and ``pipe x pipe`` matrices. That made the
memory requirement quadratic and the explicit matrix inverse cubic in time.
This module keeps one algorithm for every network size:
* run EPANET once and reuse the first hydraulic state;
* keep incidence and hydraulic graphs sparse;
* estimate the row-wise L1 pressure sensitivity with deterministic Cauchy
projections and one sparse factorization;
* estimate total directed hydraulic distance from a deterministic spatial
coreset without materialising an all-pairs distance matrix;
* balance sensitivity score with geographic and pipe-network coverage without
materialising candidate-to-candidate distances.
The random seed and sample counts are fixed, so the same model and request
produce the same placement on every run.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from pathlib import Path
from tempfile import TemporaryDirectory
from time import perf_counter
import numpy as np
import wntr
from scipy.sparse import csr_matrix, eye
from scipy.sparse.csgraph import connected_components, dijkstra
from scipy.sparse.linalg import splu
from sklearn.cluster import MiniBatchKMeans
logger = logging.getLogger(__name__)
_RANDOM_SEED = 42
_SENSITIVITY_PROJECTIONS = 256
_HYDRAULIC_LANDMARKS = 256
_PROJECTION_BLOCK_SIZE = 16
_DIJKSTRA_BLOCK_SIZE = 16
_HEADLOSS_EPSILON = 1e-10
_DIAMETER_TOLERANCE_MM = 1e-9
_COVERAGE_ELIGIBILITY_RATIO = 0.70
_COVERAGE_EDGE_EPSILON = 1e-9
@dataclass(frozen=True)
class _PreparedNetwork:
"""Sparse data required by the placement pipeline."""
node_names: tuple[str, ...]
full_node_indices: np.ndarray
candidate_indices: np.ndarray
coordinates: np.ndarray
incidence: csr_matrix
conductance: np.ndarray
roughness_response: np.ndarray
distance_graph: csr_matrix
coverage_graph: csr_matrix
@dataclass(frozen=True)
class _CandidatePool:
"""Aligned candidate arrays consumed by the placement stage."""
full_indices: np.ndarray
coordinates: np.ndarray
names: np.ndarray
scores: np.ndarray
def _run_hydraulic_simulation(
wn: wntr.network.WaterNetworkModel,
):
"""Run only the initial EPANET state without shared ``temp.*`` files."""
original_duration = wn.options.time.duration
try:
# Every downstream calculation reads ``iloc[0]``. Running an extended
# simulation only allocates unused time-series results, which is
# especially expensive for daily models with tens of thousands of
# nodes. Restore the caller's model even when EPANET fails.
wn.options.time.duration = 0
with TemporaryDirectory(prefix="tjwater-sensitivity-") as temp_dir:
file_prefix = str(Path(temp_dir) / "simulation")
return wntr.sim.EpanetSimulator(wn).run_sim(file_prefix=file_prefix)
finally:
wn.options.time.duration = original_duration
def _excluded_elements(
wn: wntr.network.WaterNetworkModel,
) -> tuple[set[str], set[str]]:
"""Return nodes that cannot host sensors and source-connected pipes.
Reservoirs, tanks, pump/valve endpoints, and the junction immediately next
to a reservoir or tank are treated as hydraulic boundary nodes. Pipes
connected directly to a source are removed from the perturbation set, as
in the legacy algorithm.
"""
source_nodes = set(wn.reservoir_name_list) | set(wn.tank_name_list)
excluded_nodes = set(source_nodes)
source_pipes: set[str] = set()
for pipe_name, pipe in wn.pipes():
endpoints = {pipe.start_node_name, pipe.end_node_name}
if endpoints & source_nodes:
source_pipes.add(pipe_name)
excluded_nodes.update(endpoints)
for _link_name, link in list(wn.pumps()) + list(wn.valves()):
excluded_nodes.add(link.start_node_name)
excluded_nodes.add(link.end_node_name)
return excluded_nodes, source_pipes
def _minimum_weight_csr(
rows: list[int],
columns: list[int],
weights: list[float],
*,
shape: tuple[int, int],
) -> csr_matrix:
"""Build a CSR graph while retaining the lightest parallel edge."""
if not rows:
return csr_matrix(shape, dtype=np.float64)
row_array = np.asarray(rows, dtype=np.int64)
column_array = np.asarray(columns, dtype=np.int64)
weight_array = np.asarray(weights, dtype=np.float64)
order = np.lexsort((column_array, row_array))
row_array = row_array[order]
column_array = column_array[order]
weight_array = weight_array[order]
group_start = np.empty(len(row_array), dtype=bool)
group_start[0] = True
group_start[1:] = (row_array[1:] != row_array[:-1]) | (
column_array[1:] != column_array[:-1]
)
starts = np.flatnonzero(group_start)
minimum_weights = np.minimum.reduceat(weight_array, starts)
return csr_matrix(
(minimum_weights, (row_array[starts], column_array[starts])),
shape=shape,
)
def _node_coordinates(
wn: wntr.network.WaterNetworkModel,
node_names: tuple[str, ...],
) -> np.ndarray:
coordinate_series = wn.query_node_attribute("coordinates")
coordinates = np.asarray(
[coordinate_series.loc[node_name] for node_name in node_names],
dtype=np.float64,
)
if coordinates.ndim != 2 or coordinates.shape[1] < 2:
raise ValueError("管网节点缺少二维坐标,无法进行监测点空间布置")
coordinates = coordinates[:, :2]
if not np.isfinite(coordinates).all():
raise ValueError("管网节点坐标包含非有限值,无法进行监测点空间布置")
return coordinates
def _build_coverage_graph(
wn: wntr.network.WaterNetworkModel,
results,
full_node_index: dict[str, int],
) -> csr_matrix:
"""Build the active undirected physical graph used to spread sensors."""
status_series = results.link["status"].iloc[0]
rows: list[int] = []
columns: list[int] = []
weights: list[float] = []
for link_name, link in wn.links():
if float(status_series.loc[link_name]) <= 0:
continue
start = full_node_index[link.start_node_name]
end = full_node_index[link.end_node_name]
# Pipes carry their physical length. Pumps and valves are point
# devices, so a tiny positive length preserves connectivity without
# dominating shortest-path distance.
weight = max(
float(getattr(link, "length", 0.0)),
_COVERAGE_EDGE_EPSILON,
)
rows.extend((start, end))
columns.extend((end, start))
weights.extend((weight, weight))
return _minimum_weight_csr(
rows,
columns,
weights,
shape=(len(full_node_index), len(full_node_index)),
)
def _prepare_network(
wn: wntr.network.WaterNetworkModel,
results,
*,
min_diameter: int,
) -> _PreparedNetwork:
excluded_nodes, source_pipes = _excluded_elements(wn)
full_node_names = tuple(wn.node_name_list)
full_node_index = {
node_name: index for index, node_name in enumerate(full_node_names)
}
node_names = tuple(
node_name for node_name in full_node_names if node_name not in excluded_nodes
)
if not node_names:
raise ValueError("管网中没有可参与灵敏度分析的节点")
node_index = {node_name: index for index, node_name in enumerate(node_names)}
full_node_indices = np.asarray(
[full_node_index[node_name] for node_name in node_names],
dtype=np.int64,
)
coordinates = _node_coordinates(wn, node_names)
flow_series = results.link["flowrate"].iloc[0]
headloss_series = results.link["headloss"].iloc[0]
head_series = results.node["head"].iloc[0]
candidate_nodes: set[str] = set()
for _pipe_name, pipe in wn.pipes():
diameter_mm = float(pipe.diameter) * 1000.0
if diameter_mm + _DIAMETER_TOLERANCE_MM < min_diameter:
continue
if pipe.start_node_name in node_index:
candidate_nodes.add(pipe.start_node_name)
if pipe.end_node_name in node_index:
candidate_nodes.add(pipe.end_node_name)
incidence_rows: list[int] = []
incidence_columns: list[int] = []
incidence_values: list[float] = []
conductance: list[float] = []
roughness_response: list[float] = []
distance_rows: list[int] = []
distance_columns: list[int] = []
distance_weights: list[float] = []
kept_pipe_count = 0
for pipe_name, pipe in wn.pipes():
if pipe_name in source_pipes:
continue
start_name = pipe.start_node_name
end_name = pipe.end_node_name
if start_name not in node_index and end_name not in node_index:
continue
flow = float(flow_series.loc[pipe_name])
absolute_flow = abs(flow)
headloss = abs(float(headloss_series.loc[pipe_name]))
roughness = float(pipe.roughness)
if roughness <= 0:
raise ValueError(f"管道 {pipe_name} 的粗糙度必须大于 0")
orientation = -1.0 if flow < 0 else 1.0
if start_name in node_index:
incidence_rows.append(node_index[start_name])
incidence_columns.append(kept_pipe_count)
incidence_values.append(-orientation)
if end_name in node_index:
incidence_rows.append(node_index[end_name])
incidence_columns.append(kept_pipe_count)
incidence_values.append(orientation)
conductance.append(
absolute_flow / (1.852 * headloss + _HEADLOSS_EPSILON)
)
roughness_response.append(absolute_flow / roughness)
if flow > 0:
upstream_name, downstream_name = start_name, end_name
else:
upstream_name, downstream_name = end_name, start_name
hydraulic_weight = (
abs(float(head_series.loc[start_name]) - float(head_series.loc[end_name]))
* float(pipe.length)
)
distance_rows.append(full_node_index[upstream_name])
distance_columns.append(full_node_index[downstream_name])
distance_weights.append(hydraulic_weight)
kept_pipe_count += 1
if kept_pipe_count == 0:
raise ValueError("管网中没有可用于灵敏度分析的管道")
incidence = csr_matrix(
(
np.asarray(incidence_values, dtype=np.float64),
(
np.asarray(incidence_rows, dtype=np.int64),
np.asarray(incidence_columns, dtype=np.int64),
),
),
shape=(len(node_names), kept_pipe_count),
)
conductance_array = np.asarray(conductance, dtype=np.float64)
response_array = np.asarray(roughness_response, dtype=np.float64)
if not np.isfinite(conductance_array).all() or not np.isfinite(
response_array
).all():
raise ValueError("水力结果产生了非有限灵敏度系数")
distance_graph = _minimum_weight_csr(
distance_rows,
distance_columns,
distance_weights,
shape=(len(full_node_names), len(full_node_names)),
)
coverage_graph = _build_coverage_graph(wn, results, full_node_index)
candidate_indices = np.asarray(
[
index
for index, node_name in enumerate(node_names)
if node_name in candidate_nodes
],
dtype=np.int64,
)
return _PreparedNetwork(
node_names=node_names,
full_node_indices=full_node_indices,
candidate_indices=candidate_indices,
coordinates=coordinates,
incidence=incidence,
conductance=conductance_array,
roughness_response=response_array,
distance_graph=distance_graph,
coverage_graph=coverage_graph,
)
def _axis_normalized_coordinates(coordinates: np.ndarray) -> np.ndarray:
"""Scale each axis independently for MiniBatchKMeans."""
minimum = coordinates.min(axis=0)
span = np.ptp(coordinates, axis=0)
span[span == 0] = 1.0
return (coordinates - minimum) / span
def _isotropic_coordinates(coordinates: np.ndarray) -> np.ndarray:
"""Normalize coordinates without distorting the network aspect ratio."""
minimum = coordinates.min(axis=0)
scale = float(np.max(np.ptp(coordinates, axis=0), initial=0.0))
if scale == 0:
scale = 1.0
return (coordinates - minimum) / scale
def _cluster_labels(
coordinates: np.ndarray,
cluster_count: int,
*,
random_seed: int,
) -> tuple[np.ndarray, np.ndarray]:
"""Cluster coordinates deterministically with one implementation at all sizes."""
normalized = _axis_normalized_coordinates(coordinates)
if cluster_count == 1:
return np.zeros(len(coordinates), dtype=np.int64), normalized[[0]]
if cluster_count >= len(coordinates):
return np.arange(len(coordinates), dtype=np.int64), normalized.copy()
model = MiniBatchKMeans(
n_clusters=cluster_count,
random_state=random_seed,
n_init=3,
batch_size=min(len(coordinates), max(1024, cluster_count * 3)),
max_iter=100,
max_no_improvement=20,
reassignment_ratio=0.0,
)
labels = model.fit_predict(normalized).astype(np.int64, copy=False)
return labels, np.asarray(model.cluster_centers_, dtype=np.float64)
def _estimate_log_pressure_sensitivity(prepared: _PreparedNetwork) -> np.ndarray:
"""Estimate each row's L1 sensitivity using streaming Cauchy projections."""
weighted_incidence = prepared.incidence.multiply(prepared.conductance)
laplacian = (weighted_incidence @ prepared.incidence.T).tocsc()
diagonal = np.asarray(laplacian.diagonal(), dtype=np.float64)
diagonal_scale = float(np.max(np.abs(diagonal), initial=0.0))
if diagonal_scale == 0:
raise ValueError("水力雅可比矩阵为空,无法计算压力灵敏度")
regularization = diagonal_scale * np.sqrt(np.finfo(np.float64).eps)
laplacian = laplacian + eye(
laplacian.shape[0], format="csc", dtype=np.float64
) * regularization
factor = splu(
laplacian,
permc_spec="MMD_AT_PLUS_A",
diag_pivot_thresh=0.0,
options={"SymmetricMode": True},
)
random = np.random.default_rng(_RANDOM_SEED)
log_absolute_sum = np.zeros(len(prepared.node_names), dtype=np.float64)
projection_count = 0
float_epsilon = np.finfo(np.float64).eps
float_tiny = np.finfo(np.float64).tiny
while projection_count < _SENSITIVITY_PROJECTIONS:
block_size = min(
_PROJECTION_BLOCK_SIZE,
_SENSITIVITY_PROJECTIONS - projection_count,
)
uniform = random.random((prepared.incidence.shape[1], block_size))
np.clip(uniform, float_epsilon, 1.0 - float_epsilon, out=uniform)
cauchy_projection = np.tan(np.pi * (uniform - 0.5))
projected_response = prepared.incidence @ (
prepared.roughness_response[:, None] * cauchy_projection
)
solution = factor.solve(np.asarray(projected_response, dtype=np.float64))
log_absolute_sum += np.log(
np.maximum(np.abs(solution), float_tiny)
).sum(axis=1)
projection_count += block_size
# For a standard Cauchy variable E[log(abs(X))] is zero. Therefore this
# streaming geometric mean estimates log(||row||_1) without retaining the
# node-by-projection matrix. A finite-sample bias is common to all rows and
# does not affect ranking.
return log_absolute_sum / _SENSITIVITY_PROJECTIONS
def _landmark_coreset(prepared: _PreparedNetwork) -> tuple[np.ndarray, np.ndarray]:
landmark_count = min(_HYDRAULIC_LANDMARKS, len(prepared.node_names))
labels, centers = _cluster_labels(
prepared.coordinates,
landmark_count,
random_seed=_RANDOM_SEED + 1,
)
normalized = _axis_normalized_coordinates(prepared.coordinates)
landmarks: list[int] = []
weights: list[float] = []
for label in np.unique(labels):
members = np.flatnonzero(labels == label)
center = centers[int(label)]
squared_distance = np.square(normalized[members] - center).sum(axis=1)
landmarks.append(int(members[int(np.argmin(squared_distance))]))
weights.append(float(len(members)))
return (
np.asarray(landmarks, dtype=np.int64),
np.asarray(weights, dtype=np.float64),
)
def _estimate_hydraulic_distance_sums(prepared: _PreparedNetwork) -> np.ndarray:
"""Estimate outbound distance sums without an all-pairs distance matrix."""
landmark_indices, landmark_weights = _landmark_coreset(prepared)
full_landmark_indices = prepared.full_node_indices[landmark_indices]
reversed_graph = prepared.distance_graph.transpose().tocsr()
distance_sums = np.zeros(len(prepared.node_names), dtype=np.float64)
for start in range(0, len(landmark_indices), _DIJKSTRA_BLOCK_SIZE):
stop = min(start + _DIJKSTRA_BLOCK_SIZE, len(landmark_indices))
distances = dijkstra(
reversed_graph,
directed=True,
indices=full_landmark_indices[start:stop],
return_predecessors=False,
)
distances = np.atleast_2d(distances)[:, prepared.full_node_indices]
# The legacy matrix represented unreachable pairs as zero. Retaining
# that convention prevents disconnected branches from receiving an
# artificial infinite score.
distances[~np.isfinite(distances)] = 0.0
distance_sums += landmark_weights[start:stop] @ distances
return distance_sums
def _build_candidate_pool(
prepared: _PreparedNetwork,
log_sensitivity: np.ndarray,
hydraulic_distance_sums: np.ndarray,
) -> _CandidatePool:
candidate_indices = prepared.candidate_indices
candidate_distance = hydraulic_distance_sums[candidate_indices]
with np.errstate(divide="ignore", invalid="ignore"):
scores = log_sensitivity[candidate_indices] + np.log(candidate_distance)
scores = np.nan_to_num(
scores,
nan=-np.inf,
neginf=-np.inf,
posinf=np.finfo(np.float64).max,
)
return _CandidatePool(
full_indices=prepared.full_node_indices[candidate_indices],
coordinates=_isotropic_coordinates(prepared.coordinates[candidate_indices]),
names=np.asarray(
[prepared.node_names[index] for index in candidate_indices],
dtype=str,
),
scores=scores,
)
def _highest_scoring_position(
names: np.ndarray,
scores: np.ndarray,
positions: np.ndarray,
) -> int:
"""Return the best position, breaking score ties by node name."""
order = np.lexsort((names[positions], -scores[positions]))
return int(positions[order[0]])
def _relative_gap(
distances: np.ndarray,
available: np.ndarray,
) -> np.ndarray:
"""Normalize available distances to their current finite maximum."""
maximum = float(np.max(distances[available], initial=0.0))
if not np.isfinite(maximum) or maximum <= np.finfo(np.float64).eps:
return np.zeros(len(distances), dtype=np.float64)
return distances / maximum
def _eligible_gap_positions(
relative_gap: np.ndarray,
available: np.ndarray,
) -> np.ndarray:
"""Return positions within the configured fraction of the largest gap."""
maximum = float(np.max(relative_gap[available], initial=0.0))
if maximum <= np.finfo(np.float64).eps:
return np.flatnonzero(available)
threshold = _COVERAGE_ELIGIBILITY_RATIO * maximum
return np.flatnonzero(
available & (relative_gap >= threshold - np.finfo(np.float64).eps)
)
def _allocate_component_quotas(
coverage_graph: csr_matrix,
candidate_full_indices: np.ndarray,
candidate_scores: np.ndarray,
*,
sensor_num: int,
) -> tuple[np.ndarray, np.ndarray]:
"""Allocate sensor counts by active pipe length with candidate caps."""
component_count, node_components = connected_components(
coverage_graph,
directed=False,
return_labels=True,
)
candidate_components = node_components[candidate_full_indices]
capacities = np.bincount(
candidate_components,
minlength=component_count,
).astype(np.int64, copy=False)
# The graph is symmetric. Summed row weights count every physical edge
# twice, hence the division by two after aggregation by component.
node_lengths = np.asarray(coverage_graph.sum(axis=1)).ravel()
component_lengths = np.bincount(
node_components,
weights=node_lengths,
minlength=component_count,
) / 2.0
component_best_scores = np.full(component_count, -np.inf, dtype=np.float64)
np.maximum.at(
component_best_scores,
candidate_components,
candidate_scores,
)
active_components = np.flatnonzero(capacities)
quotas = np.zeros(component_count, dtype=np.int64)
if len(active_components) > sensor_num:
order = np.lexsort(
(
active_components,
-component_best_scores[active_components],
-component_lengths[active_components],
)
)
quotas[active_components[order[:sensor_num]]] = 1
return candidate_components, quotas
quotas[active_components] = 1
remaining = sensor_num - len(active_components)
while remaining > 0:
available = active_components[
quotas[active_components] < capacities[active_components]
]
if len(available) == 0:
raise ValueError("连通区域中的候选节点不足,无法分配监测点名额")
weights = component_lengths[available]
if float(weights.sum()) <= 0:
weights = (capacities[available] - quotas[available]).astype(
np.float64,
copy=False,
)
ideal = remaining * weights / float(weights.sum())
whole = np.minimum(
np.floor(ideal).astype(np.int64),
capacities[available] - quotas[available],
)
whole_count = int(whole.sum())
if whole_count:
quotas[available] += whole
remaining -= whole_count
continue
fractional = ideal - np.floor(ideal)
order = np.lexsort(
(
available,
-component_best_scores[available],
-weights,
-fractional,
)
)
for component in available[order]:
quotas[component] += 1
remaining -= 1
if remaining == 0:
break
return candidate_components, quotas
def _select_component_positions(
coverage_graph: csr_matrix,
candidates: _CandidatePool,
component_positions: np.ndarray,
existing_positions: list[int],
*,
quota: int,
) -> list[int]:
"""Select one component's sensors with score-aware farthest-first search."""
local_coordinates = candidates.coordinates[component_positions]
local_names = candidates.names[component_positions]
local_scores = candidates.scores[component_positions]
nearest_geographic = np.full(len(component_positions), np.inf)
nearest_topological = np.full(len(component_positions), np.inf)
for position in existing_positions:
nearest_geographic = np.minimum(
nearest_geographic,
np.linalg.norm(
local_coordinates - candidates.coordinates[position],
axis=1,
),
)
if existing_positions:
all_local = np.ones(len(component_positions), dtype=bool)
seed_eligible = _eligible_gap_positions(
_relative_gap(nearest_geographic, all_local),
all_local,
)
else:
seed_eligible = np.arange(len(component_positions), dtype=np.int64)
seed = _highest_scoring_position(
local_names,
local_scores,
seed_eligible,
)
selected_local = [seed]
remaining = np.ones(len(component_positions), dtype=bool)
remaining[seed] = False
while len(selected_local) < quota:
newest = selected_local[-1]
geographic_distance = np.linalg.norm(
local_coordinates - local_coordinates[newest],
axis=1,
)
nearest_geographic = np.minimum(
nearest_geographic,
geographic_distance,
)
source = int(candidates.full_indices[component_positions[newest]])
topological_distance = dijkstra(
coverage_graph,
directed=False,
indices=source,
return_predecessors=False,
)[candidates.full_indices[component_positions]]
nearest_topological = np.minimum(
nearest_topological,
topological_distance,
)
coverage_gap = np.maximum(
_relative_gap(nearest_geographic, remaining),
_relative_gap(nearest_topological, remaining),
)
eligible_local = _eligible_gap_positions(coverage_gap, remaining)
next_local = _highest_scoring_position(
local_names,
local_scores,
eligible_local,
)
selected_local.append(next_local)
remaining[next_local] = False
return [int(component_positions[position]) for position in selected_local]
def _geographic_coverage_metrics(
candidate_coordinates: np.ndarray,
selected_positions: list[int],
) -> tuple[float, float, float]:
normalized = _isotropic_coordinates(candidate_coordinates)
nearest = np.full(len(normalized), np.inf)
for position in selected_positions:
nearest = np.minimum(
nearest,
np.linalg.norm(normalized - normalized[position], axis=1),
)
selected_coordinates = normalized[selected_positions]
if len(selected_positions) < 2:
minimum_gap = 0.0
else:
pairwise = np.linalg.norm(
selected_coordinates[:, None, :] - selected_coordinates[None, :, :],
axis=2,
)
np.fill_diagonal(pairwise, np.inf)
minimum_gap = float(pairwise.min())
return (
float(nearest.max()),
float(np.quantile(nearest, 0.95)),
minimum_gap,
)
def _select_sensor_nodes(
prepared: _PreparedNetwork,
log_sensitivity: np.ndarray,
hydraulic_distance_sums: np.ndarray,
*,
sensor_num: int,
) -> list[str]:
candidate_indices = prepared.candidate_indices
if len(candidate_indices) < sensor_num:
raise ValueError(
"满足最小管径要求的候选节点少于请求的监测点数量:"
f"候选 {len(candidate_indices)} 个,请求 {sensor_num}"
)
candidates = _build_candidate_pool(
prepared,
log_sensitivity,
hydraulic_distance_sums,
)
candidate_components, component_quotas = _allocate_component_quotas(
prepared.coverage_graph,
candidates.full_indices,
candidates.scores,
sensor_num=sensor_num,
)
selected_positions: list[int] = []
quota_components = np.flatnonzero(component_quotas)
component_order = np.lexsort(
(quota_components, -component_quotas[quota_components])
)
for component in quota_components[component_order]:
component_positions = np.flatnonzero(candidate_components == component)
selected_positions.extend(
_select_component_positions(
prepared.coverage_graph,
candidates,
component_positions,
selected_positions,
quota=int(component_quotas[component]),
)
)
selected_array = np.asarray(selected_positions, dtype=np.int64)
selected_order = np.lexsort(
(
candidates.names[selected_array],
-candidates.scores[selected_array],
)
)
selected_positions = selected_array[selected_order].tolist()
maximum_radius, p95_radius, minimum_gap = _geographic_coverage_metrics(
prepared.coordinates[candidate_indices],
selected_positions,
)
logger.info(
"Sensitivity placement coverage: components=%d max_radius=%.6f "
"p95_radius=%.6f min_sensor_gap=%.6f",
int(np.count_nonzero(component_quotas)),
maximum_radius,
p95_radius,
minimum_gap,
)
return [str(candidates.names[position]) for position in selected_positions]
def optimize_sensor_placement(
wn: wntr.network.WaterNetworkModel,
sensor_num: int,
min_diameter: int,
) -> list[str]:
"""Return deterministic pressure monitoring nodes for a loaded network.
``min_diameter`` is expressed in millimetres, matching the HTTP contract.
A node is a valid installation candidate when at least one incident pipe
meets the threshold. All valid hydraulic nodes still participate in the
sensitivity calculation so small pipes continue to influence the result.
"""
if sensor_num <= 0:
raise ValueError("监测点数量必须大于 0")
if min_diameter < 0:
raise ValueError("最小管径不能小于 0")
total_started = perf_counter()
simulation_started = total_started
results = _run_hydraulic_simulation(wn)
simulation_seconds = perf_counter() - simulation_started
preparation_started = perf_counter()
prepared = _prepare_network(wn, results, min_diameter=min_diameter)
preparation_seconds = perf_counter() - preparation_started
sensitivity_started = perf_counter()
log_sensitivity = _estimate_log_pressure_sensitivity(prepared)
sensitivity_seconds = perf_counter() - sensitivity_started
distance_started = perf_counter()
hydraulic_distance_sums = _estimate_hydraulic_distance_sums(prepared)
distance_seconds = perf_counter() - distance_started
selection_started = perf_counter()
selected = _select_sensor_nodes(
prepared,
log_sensitivity,
hydraulic_distance_sums,
sensor_num=sensor_num,
)
selection_seconds = perf_counter() - selection_started
logger.info(
"Sensitivity placement completed: nodes=%d pipes=%d candidates=%d "
"sensors=%d seconds=%.3f "
"(simulation=%.3f preparation=%.3f sensitivity=%.3f "
"distance=%.3f selection=%.3f)",
len(prepared.node_names),
prepared.incidence.shape[1],
len(prepared.candidate_indices),
len(selected),
perf_counter() - total_started,
simulation_seconds,
preparation_seconds,
sensitivity_seconds,
distance_seconds,
selection_seconds,
)
return selected
def optimize_sensor_placement_from_inp(
inp_path: str | Path,
sensor_num: int,
min_diameter: int,
) -> list[str]:
"""Load an EPANET INP model and run the unified placement algorithm."""
wn = wntr.network.WaterNetworkModel(str(inp_path))
return optimize_sensor_placement(
wn,
sensor_num=sensor_num,
min_diameter=min_diameter,
)
@@ -0,0 +1,6 @@
"""SCADA time-series cleaning algorithms."""
from .flow_series import clean_flow_data_df_kf
from .pressure_series import clean_pressure_data_df_km
__all__ = ["clean_flow_data_df_kf", "clean_pressure_data_df_km"]
@@ -142,11 +142,13 @@ def clean_flow_data_kf(
return os.path.abspath(output_path)
def clean_flow_data_df_kf(data: pd.DataFrame, show_plot: bool = False) -> dict:
def clean_flow_data_df_kf(
data: pd.DataFrame, show_plot: bool = False
) -> pd.DataFrame:
"""
接收一个 DataFrame 数据结构使用一维 Kalman 滤波平滑并用预测值替换基于 IQR 检测出的异常点
区分合理的0值流量转换和异常的0值连续多个0或孤立0
返回完整的清洗后的字典数据结构
返回完整的清洗后 DataFrame
Args:
data: 输入 DataFrame可包含 time
@@ -305,42 +307,3 @@ def clean_flow_data_df_kf(data: pd.DataFrame, show_plot: bool = False) -> dict:
# 返回完整的修复后字典
return cleaned_data
# # 测试
# if __name__ == "__main__":
# # 默认:脚本目录下同名 CSV 文件
# script_dir = os.path.dirname(os.path.abspath(__file__))
# default_csv = os.path.join(script_dir, "pipe_flow_data_to_clean2.0.csv")
# out = clean_flow_data_kf(default_csv)
# print("清洗后的数据已保存到:", out)
# 测试 clean_flow_data_dict 函数
if __name__ == "__main__":
import random
# 读取 szh_flow_scada.csv 文件
script_dir = os.path.dirname(os.path.abspath(__file__))
csv_path = os.path.join(script_dir, "szh_flow_scada.csv")
data = pd.read_csv(csv_path, header=0, index_col=None, encoding="utf-8")
# 排除 Time 列,随机选择 5 列
columns_to_exclude = ["Time"]
available_columns = [col for col in data.columns if col not in columns_to_exclude]
selected_columns = random.sample(available_columns, 1)
# 将选中的列转换为字典
data_dict = {col: data[col].tolist() for col in selected_columns}
print("选中的列:", selected_columns)
print("原始数据长度:", len(data_dict[selected_columns[0]]))
# 调用函数进行清洗
cleaned_dict = clean_flow_data_df_kf(data_dict, show_plot=True)
# 将清洗后的字典写回 CSV
out_csv = os.path.join(script_dir, f"{selected_columns[0]}_clean.csv")
pd.DataFrame(cleaned_dict).to_csv(out_csv, index=False, encoding="utf-8-sig")
print("已保存清洗结果到:", out_csv)
print("清洗后的字典键:", list(cleaned_dict.keys()))
print("清洗后的数据长度:", len(cleaned_dict[selected_columns[0]]))
print("测试完成:函数运行正常")
@@ -0,0 +1,543 @@
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import os
ID_LIKE_COLUMNS = {
"id",
"device_id",
"node_id",
"sensor_id",
"monitor_id",
"junction_id",
}
def _normalize_time_frame(data: pd.DataFrame) -> pd.DataFrame:
"""返回按时间排序的副本,并尽量将 time 列解析为时间类型。"""
data = data.copy()
if "time" in data.columns:
data["time"] = pd.to_datetime(data["time"], errors="coerce")
data = data.sort_values(["time"]).reset_index(drop=True)
return data
def _select_pressure_columns(data: pd.DataFrame) -> tuple[list[str], list[str]]:
"""区分需要清洗的数值列与需要原样保留的列。"""
value_cols: list[str] = []
keep_cols: list[str] = []
for col in data.columns:
if col == "time":
continue
col_key = col.lower()
if col_key in ID_LIKE_COLUMNS or col_key.endswith("_id"):
keep_cols.append(col)
continue
numeric = pd.to_numeric(data[col], errors="coerce")
if numeric.notna().sum() == 0 or numeric.nunique(dropna=True) <= 1:
keep_cols.append(col)
else:
value_cols.append(col)
return value_cols, keep_cols
def _robust_scale(values: pd.Series) -> float:
"""基于 MAD 计算稳健尺度。"""
series = pd.to_numeric(values, errors="coerce").dropna()
if series.empty:
return 1.0
median = series.median()
mad = (series - median).abs().median()
if pd.notna(mad) and mad > 0:
return float(1.4826 * mad)
iqr = series.quantile(0.75) - series.quantile(0.25)
if pd.notna(iqr) and iqr > 0:
return float(iqr / 1.349)
std = series.std()
if pd.notna(std) and std > 0:
return float(std)
return 1.0
def _shrink_toward_baseline(observed: float, baseline: float, scale: float) -> float:
"""把观测值向基线值收缩,scale 越小,修复越强。"""
if pd.isna(observed):
return baseline
if pd.isna(baseline):
return observed
diff = observed - baseline
weight = scale / (abs(diff) + scale)
return float(baseline + diff * weight)
def _infer_time_frequency(time_values: pd.Series | pd.Index) -> pd.Timedelta:
"""从时间序列中推断采样频率,失败时默认 15 分钟。"""
parsed = pd.to_datetime(pd.Series(time_values), errors="coerce").dropna().sort_values()
if len(parsed) < 2:
return pd.Timedelta(minutes=15)
diffs = parsed.diff().dropna()
diffs = diffs[diffs > pd.Timedelta(0)]
if diffs.empty:
return pd.Timedelta(minutes=15)
mode = diffs.mode()
return mode.iloc[0] if not mode.empty else diffs.median()
def _build_local_pressure_baseline(series: pd.Series) -> pd.Series:
"""基于局部插值与中值滤波构造平滑基线。"""
baseline = _safe_time_interpolate(series)
baseline = baseline.rolling(window=5, center=True, min_periods=1).median()
baseline = _safe_time_interpolate(baseline)
return baseline.ffill().bfill()
def _build_seasonal_pressure_baseline(series: pd.Series) -> pd.Series:
"""按一天内的同一时刻构造季节性基线,适合日周期压力数据。"""
if not isinstance(series.index, pd.DatetimeIndex):
return pd.Series(np.nan, index=series.index, dtype=float)
slot_labels = pd.Series(series.index.strftime("%H:%M:%S"), index=series.index)
return series.groupby(slot_labels).transform("median")
def _detect_pressure_spikes(series: pd.Series, local_baseline: pd.Series) -> pd.Series:
"""识别单点异常上升/下降尖峰,避免过度修正正常波动。"""
residual = series - local_baseline
neighbor_center = (series.shift(1) + series.shift(-1)) / 2
curvature = series - neighbor_center
residual_scale = max(_robust_scale(residual), 1e-6)
curvature_scale = max(_robust_scale(curvature), 1e-6)
direction_flip = ((series - series.shift(1)) * (series.shift(-1) - series) < 0).fillna(False)
return (
residual.abs() > 3.5 * residual_scale
) & (
curvature.abs() > 3.0 * curvature_scale
) & direction_flip
def _fill_pressure_gaps(
original: pd.Series,
repaired: pd.Series,
local_baseline: pd.Series,
seasonal_baseline: pd.Series,
) -> pd.Series:
"""短缺口用局部插值,长缺口优先使用同一时刻的季节性轨迹。"""
missing_mask = original.isna()
if not missing_mask.any():
return repaired
gap_groups = (missing_mask != missing_mask.shift(fill_value=False)).cumsum()
gap_lengths = missing_mask.groupby(gap_groups).transform("sum").where(missing_mask, 0)
filled = repaired.copy()
short_gap_mask = missing_mask & (gap_lengths < 4)
long_gap_mask = missing_mask & ~short_gap_mask
filled[short_gap_mask] = local_baseline[short_gap_mask]
long_gap_fill = seasonal_baseline.where(seasonal_baseline.notna(), local_baseline)
filled[long_gap_mask] = long_gap_fill[long_gap_mask]
return filled
def _clean_pressure_series(series: pd.Series) -> pd.Series:
"""清洗单个压力时间序列。"""
series = pd.to_numeric(series, errors="coerce").astype(float)
local_baseline = _build_local_pressure_baseline(series)
spike_mask = _detect_pressure_spikes(series, local_baseline)
repaired = series.copy()
repaired[spike_mask] = local_baseline[spike_mask]
seasonal_baseline = _build_seasonal_pressure_baseline(repaired)
repaired = _fill_pressure_gaps(series, repaired, local_baseline, seasonal_baseline)
if repaired.isna().any():
repaired = repaired.where(repaired.notna(), local_baseline)
return repaired.ffill().bfill()
def _format_time_column(data: pd.DataFrame) -> pd.DataFrame:
"""统一输出时间格式,方便下游直接按 ISO 字符串解析。"""
if "time" not in data.columns:
return data
formatted = data.copy()
time_values = pd.to_datetime(formatted["time"], errors="coerce")
if time_values.isna().all():
return formatted
if time_values.dt.tz is not None:
time_strings = time_values.dt.strftime("%Y-%m-%dT%H:%M:%S%z")
time_strings = time_strings.str.replace(
r"([+-]\d{2})(\d{2})$",
r"\1:\2",
regex=True,
)
else:
time_strings = time_values.dt.strftime("%Y-%m-%dT%H:%M:%S")
formatted["time"] = time_strings.where(time_values.notna(), formatted["time"])
return formatted
def _expand_snapshot_time_grid(data: pd.DataFrame, freq: pd.Timedelta) -> pd.DataFrame:
"""仅补齐时间轴,不提前填充值,避免长缺口丢失原始形状特征。"""
expanded = data.copy()
expanded["time"] = pd.to_datetime(expanded["time"], errors="coerce")
expanded = expanded.dropna(subset=["time"]).sort_values("time")
if expanded.empty:
return data
indexed = expanded.set_index("time")
full_index = pd.date_range(indexed.index.min(), indexed.index.max(), freq=freq)
indexed = indexed.reindex(full_index)
indexed.index.name = "time"
return indexed.reset_index()
def _safe_datetime_index(values: pd.Series | pd.Index | list[object]) -> pd.DatetimeIndex | None:
"""尽量把时间值标准化为 DatetimeIndex;失败则返回 None。"""
parsed = pd.to_datetime(values, errors="coerce")
try:
datetime_index = pd.DatetimeIndex(parsed)
except (TypeError, ValueError):
return None
if datetime_index.isna().all():
return None
return datetime_index
def _safe_time_interpolate(series: pd.Series) -> pd.Series:
"""仅在索引确实是 DatetimeIndex 时使用 time interpolation。"""
if isinstance(series.index, pd.DatetimeIndex):
return series.interpolate(method="time", limit_direction="both")
return series.interpolate(limit_direction="both")
def _detect_long_form_identifier(data: pd.DataFrame, value_cols: list[str], keep_cols: list[str]) -> str | None:
"""识别 time/id/value 长表结构。"""
if "time" not in data.columns or len(value_cols) != 1:
return None
identifier_candidates = [
col
for col in keep_cols
if col.lower() in ID_LIKE_COLUMNS or col.lower().endswith("_id")
]
if len(identifier_candidates) != 1:
return None
if not data["time"].duplicated().any():
return None
return identifier_candidates[0]
def _clean_long_form_pressure(
data: pd.DataFrame,
value_col: str,
identifier_col: str,
keep_cols: list[str],
fill_gaps: bool,
) -> pd.DataFrame:
"""按测点拆分 long-form 压力数据,再逐列清洗后恢复原结构。"""
data = _normalize_time_frame(data)
wide_df = (
data[[identifier_col, "time", value_col]]
.pivot(index="time", columns=identifier_col, values=value_col)
.reset_index()
)
sensor_cols = [col for col in wide_df.columns if col != "time"]
cleaned_wide = _clean_snapshot_pressure(wide_df, sensor_cols, keep_cols=[], fill_gaps=fill_gaps)
cleaned_long = cleaned_wide.melt(
id_vars="time",
var_name=identifier_col,
value_name=value_col,
)
passthrough_cols = [col for col in keep_cols if col != identifier_col]
if passthrough_cols:
metadata = data[[identifier_col] + passthrough_cols].drop_duplicates(subset=[identifier_col])
cleaned_long = cleaned_long.merge(metadata, on=identifier_col, how="left")
try:
cleaned_long[identifier_col] = cleaned_long[identifier_col].astype(data[identifier_col].dtype)
except (TypeError, ValueError):
pass
cleaned_long = cleaned_long.sort_values(["time", identifier_col]).reset_index(drop=True)
ordered_cols = ["time", identifier_col] + passthrough_cols + [value_col]
cleaned_long = cleaned_long[[col for col in ordered_cols if col in cleaned_long.columns]]
return cleaned_long
def _build_time_slot_frame(
data: pd.DataFrame, value_col: str, expected_slots: int
) -> pd.DataFrame:
"""把重复时间点整理成 time x slot 的矩阵。"""
grouped = data.groupby("time", sort=True)
times = list(grouped.groups.keys())
slot_frame = pd.DataFrame(index=pd.Index(times, name="time"), columns=range(expected_slots), dtype=float)
for time_value, group in grouped:
values = pd.to_numeric(group[value_col], errors="coerce").tolist()
for slot_idx, value in enumerate(values[:expected_slots]):
slot_frame.loc[time_value, slot_idx] = value
return slot_frame
def _slot_baseline(slot_frame: pd.DataFrame) -> pd.DataFrame:
"""对每个槽位做时间插值和平滑,得到基线轨迹。"""
baseline = pd.DataFrame(index=slot_frame.index, columns=slot_frame.columns, dtype=float)
for col in slot_frame.columns:
series = slot_frame[col].astype(float)
series = _safe_time_interpolate(series)
series = series.rolling(window=5, center=True, min_periods=1).median()
series = _safe_time_interpolate(series).ffill().bfill()
baseline[col] = series
return baseline
def _choose_insertion_position(
observed: list[float], baseline_row: pd.Series, expected_slots: int
) -> int:
"""为少一个观测值的时间组选择最合理的插入位置。"""
missing_count = expected_slots - len(observed)
if missing_count <= 0:
return 0
best_pos = 0
best_cost = float("inf")
for insert_pos in range(expected_slots):
cost = 0.0
obs_idx = 0
for slot_idx in range(expected_slots):
if slot_idx == insert_pos:
continue
obs_value = observed[obs_idx]
base_value = float(baseline_row.iloc[slot_idx])
if pd.notna(obs_value) and pd.notna(base_value):
cost += abs(obs_value - base_value)
obs_idx += 1
if cost < best_cost:
best_cost = cost
best_pos = insert_pos
return best_pos
def _clean_repeated_timestamp_pressure(
data: pd.DataFrame, value_col: str, keep_cols: list[str]
) -> pd.DataFrame:
"""针对同一时间点重复采样的压力数据进行修复。"""
data = _normalize_time_frame(data)
grouped_sizes = data.groupby("time").size()
if grouped_sizes.empty:
return data
expected_slots = int(grouped_sizes.mode().iloc[0]) if not grouped_sizes.mode().empty else int(grouped_sizes.max())
expected_slots = max(expected_slots, int(grouped_sizes.max()))
slot_frame = _build_time_slot_frame(data, value_col, expected_slots)
baseline_frame = _slot_baseline(slot_frame)
residuals = slot_frame - baseline_frame
slot_scales = {
col: max(_robust_scale(residuals[col]), 1e-6) for col in residuals.columns
}
cleaned_rows: list[dict[str, object]] = []
grouped = data.groupby("time", sort=True)
for time_value, group in grouped:
observed_values = pd.to_numeric(group[value_col], errors="coerce").tolist()
baseline_row = baseline_frame.loc[time_value]
insert_pos = _choose_insertion_position(observed_values, baseline_row, expected_slots)
cleaned_values: list[float] = []
obs_idx = 0
for slot_idx in range(expected_slots):
if slot_idx == insert_pos and len(observed_values) < expected_slots:
cleaned_values.append(float(baseline_row.iloc[slot_idx]))
continue
if obs_idx >= len(observed_values):
cleaned_values.append(float(baseline_row.iloc[slot_idx]))
continue
observed = observed_values[obs_idx]
baseline = float(baseline_row.iloc[slot_idx])
cleaned_values.append(
_shrink_toward_baseline(observed, baseline, slot_scales.get(slot_idx, 1.0))
)
obs_idx += 1
# 其余字段原样保留;常量列(如 id)直接复制第一条记录即可
template_row = group.iloc[0].to_dict()
for slot_idx, cleaned_value in enumerate(cleaned_values):
row = dict(template_row)
row["time"] = time_value
row[value_col] = cleaned_value
cleaned_rows.append(row)
cleaned_df = pd.DataFrame(cleaned_rows)
cleaned_df = cleaned_df.sort_values(["time"]).reset_index(drop=True)
ordered_cols = ["time"] + keep_cols + [value_col]
ordered_cols = [col for col in ordered_cols if col in cleaned_df.columns]
remaining_cols = [col for col in cleaned_df.columns if col not in ordered_cols]
cleaned_df = cleaned_df[ordered_cols + remaining_cols]
return _format_time_column(cleaned_df)
def _clean_snapshot_pressure(
data: pd.DataFrame, value_cols: list[str], keep_cols: list[str], fill_gaps: bool
) -> pd.DataFrame:
"""针对单条时间序列或多列快照数据进行稳健修复。"""
data = _normalize_time_frame(data)
if fill_gaps and "time" in data.columns:
freq = _infer_time_frequency(data["time"])
data = _expand_snapshot_time_grid(data, freq)
data["time"] = pd.to_datetime(data["time"], errors="coerce")
data = data.sort_values(["time"]).reset_index(drop=True)
cleaned_df = data.copy()
time_index = (
_safe_datetime_index(cleaned_df["time"])
if "time" in cleaned_df.columns
else None
)
if time_index is None:
time_index = pd.RangeIndex(start=0, stop=len(cleaned_df))
for col in value_cols:
series = pd.Series(
pd.to_numeric(cleaned_df[col], errors="coerce").to_numpy(),
index=time_index,
dtype=float,
)
cleaned_df[col] = _clean_pressure_series(series).to_numpy()
ordered_cols = ["time"] + keep_cols + value_cols
ordered_cols = [col for col in ordered_cols if col in cleaned_df.columns]
remaining_cols = [col for col in cleaned_df.columns if col not in ordered_cols]
cleaned_df = cleaned_df[ordered_cols + remaining_cols]
return _format_time_column(cleaned_df)
def clean_pressure_data_km(
input_csv_path: str, show_plot: bool = False, fill_gaps: bool = True
) -> str:
"""
读取输入 CSV,基于时间结构进行稳健修复。输出为 <input_basename>_cleaned.xlsx(同目录)。
原始数据在 sheet 'raw_pressure_data',处理后数据在 sheet 'cleaned_pressusre_data'
返回输出文件的绝对路径。
Args:
input_csv_path: CSV 文件路径
show_plot: 是否显示可视化
fill_gaps: 是否先补齐时间缺口(默认 True)
"""
# 读取 CSV
input_csv_path = os.path.abspath(input_csv_path)
data = pd.read_csv(input_csv_path, header=0, index_col=None, encoding="utf-8")
data = _normalize_time_frame(data)
value_cols, keep_cols = _select_pressure_columns(data)
has_repeated_time = "time" in data.columns and data["time"].duplicated().any()
identifier_col = _detect_long_form_identifier(data, value_cols, keep_cols)
if identifier_col is not None:
data_repaired = _clean_long_form_pressure(
data,
value_cols[0],
identifier_col,
keep_cols,
fill_gaps,
)
elif has_repeated_time and len(value_cols) == 1:
data_repaired = _clean_repeated_timestamp_pressure(data, value_cols[0], keep_cols)
else:
data_repaired = _clean_snapshot_pressure(data, value_cols, keep_cols, fill_gaps)
# 可选可视化(只展示首个数值列)
plt.rcParams["font.sans-serif"] = ["SimHei"]
plt.rcParams["axes.unicode_minus"] = False
if show_plot and value_cols:
plot_col = value_cols[0]
if "time" in data_repaired.columns:
x = pd.to_datetime(data_repaired["time"], errors="coerce")
else:
x = np.arange(len(data_repaired))
plt.figure(figsize=(12, 6))
plt.plot(x, pd.to_numeric(data_repaired[plot_col], errors="coerce"), label="cleaned")
plt.xlabel("时间" if "time" in data_repaired.columns else "序号")
plt.ylabel("压力监测值")
plt.title(f"{plot_col} 清洗结果")
plt.legend()
plt.show()
# 保存到 Excel:两个 sheet
input_dir = os.path.dirname(os.path.abspath(input_csv_path))
input_base = os.path.splitext(os.path.basename(input_csv_path))[0]
output_filename = f"{input_base}_cleaned.xlsx"
output_path = os.path.join(input_dir, output_filename)
# 如果原始数据包含时间列,将其添加回结果
data_for_save = data.copy()
data_repaired_for_save = data_repaired.copy()
if os.path.exists(output_path):
os.remove(output_path) # 覆盖同名文件
with pd.ExcelWriter(output_path, engine="openpyxl") as writer:
data_for_save.to_excel(writer, sheet_name="raw_pressure_data", index=False)
data_repaired_for_save.to_excel(
writer, sheet_name="cleaned_pressusre_data", index=False
)
# 返回输出文件的绝对路径
return os.path.abspath(output_path)
def clean_pressure_data_df_km(data: pd.DataFrame, show_plot: bool = False) -> pd.DataFrame:
"""
接收一个 DataFrame 数据结构,使用时间感知的稳健修复方法清洗压力数据。
返回清洗后的 DataFrame。
Args:
data: 输入 DataFrame(可包含 time 列)
show_plot: 是否显示可视化
"""
# 使用传入的 DataFrame
data = data.copy()
data = _normalize_time_frame(data)
value_cols, keep_cols = _select_pressure_columns(data)
has_repeated_time = "time" in data.columns and data["time"].duplicated().any()
identifier_col = _detect_long_form_identifier(data, value_cols, keep_cols)
if identifier_col is not None:
data_repaired = _clean_long_form_pressure(
data,
value_cols[0],
identifier_col,
keep_cols,
fill_gaps=True,
)
elif has_repeated_time and len(value_cols) == 1:
data_repaired = _clean_repeated_timestamp_pressure(data, value_cols[0], keep_cols)
else:
data_repaired = _clean_snapshot_pressure(data, value_cols, keep_cols, fill_gaps=True)
if show_plot and value_cols:
plt.rcParams["font.sans-serif"] = ["SimHei"]
plt.rcParams["axes.unicode_minus"] = False
plot_col = value_cols[0]
x = pd.to_datetime(data_repaired["time"], errors="coerce") if "time" in data_repaired.columns else np.arange(len(data_repaired))
plt.figure(figsize=(12, 6))
plt.plot(x, pd.to_numeric(data_repaired[plot_col], errors="coerce"), label="cleaned")
plt.xlabel("时间" if "time" in data_repaired.columns else "序号")
plt.ylabel("压力监测值")
plt.title(f"{plot_col} 清洗结果")
plt.legend()
plt.show()
return data_repaired
-91
View File
@@ -1,91 +0,0 @@
import psycopg
from app.algorithms.sensor import kmeans as kmeans_sensor
from app.algorithms.sensor import sensitivity
from app.core.config import get_pgconn_string
from app.services.tjnetwork import dump_inp
def pressure_sensor_placement_sensitivity(
name: str, scheme_name: str, sensor_number: int, min_diameter: int, username: str
) -> None:
"""
基于改进灵敏度法进行压力监测点优化布置
:param name: 数据库名称
:param scheme_name: 监测优化布置方案名称
:param sensor_number: 传感器数目
:param min_diameter: 最小管径
:param username: 用户名
:return:
"""
sensor_location = sensitivity.get_ID(
name=name, sensor_num=sensor_number, min_diameter=min_diameter
)
try:
conn_string = get_pgconn_string(db_name=name)
with psycopg.connect(conn_string) as conn:
with conn.cursor() as cur:
sql = """
INSERT INTO sensor_placement (scheme_name, sensor_number, min_diameter, username, sensor_location)
VALUES (%s, %s, %s, %s, %s)
"""
cur.execute(
sql,
(
scheme_name,
sensor_number,
min_diameter,
username,
sensor_location,
),
)
conn.commit()
print("方案信息存储成功!")
except Exception as e:
print(f"存储方案信息时出错:{e}")
# 2025/08/21
# 基于kmeans聚类法进行压力监测点优化布置
def pressure_sensor_placement_kmeans(
name: str, scheme_name: str, sensor_number: int, min_diameter: int, username: str
) -> None:
"""
基于聚类法进行压力监测点优化布置
:param name: 数据库名称(注意,此处数据库名称也是inp文件名称,inp文件与pg库名要一样)
:param scheme_name: 监测优化布置方案名称
:param sensor_number: 传感器数目
:param min_diameter: 最小管径
:param username: 用户名
:return:
"""
# dump_inp
inp_name = f"./db_inp/{name}.db.inp"
dump_inp(name, inp_name, "2")
sensor_location = kmeans_sensor.kmeans_sensor_placement(
name=name, sensor_num=sensor_number, min_diameter=min_diameter
)
try:
conn_string = get_pgconn_string(db_name=name)
with psycopg.connect(conn_string) as conn:
with conn.cursor() as cur:
sql = """
INSERT INTO sensor_placement (scheme_name, sensor_number, min_diameter, username, sensor_location)
VALUES (%s, %s, %s, %s, %s)
"""
cur.execute(
sql,
(
scheme_name,
sensor_number,
min_diameter,
username,
sensor_location,
),
)
conn.commit()
print("方案信息存储成功!")
except Exception as e:
print(f"存储方案信息时出错:{e}")
-109
View File
@@ -1,109 +0,0 @@
import wntr
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import sklearn.cluster
import os
class QD_KMeans(object):
def __init__(self, wn, num_monitors):
# self.inp = inp
self.cluster_num = num_monitors # 聚类中心个数,也即测压点个数
self.wn=wn
self.monitor_nodes = []
self.coords = []
self.junction_nodes = {} # Added missing initialization
def get_junctions_coordinates(self):
for junction_name in self.wn.junction_name_list:
junction = self.wn.get_node(junction_name)
self.junction_nodes[junction_name] = junction.coordinates
self.coords.append(junction.coordinates )
# print(f"Total junctions: {self.junction_coordinates}")
def select_monitoring_points(self):
if not self.coords: # Add check if coordinates are collected
self.get_junctions_coordinates()
coords = np.array(self.coords)
coords_normalized = (coords - coords.min(axis=0)) / (coords.max(axis=0) - coords.min(axis=0))
kmeans = sklearn.cluster.KMeans(n_clusters= self.cluster_num, random_state=42)
kmeans.fit(coords_normalized)
for center in kmeans.cluster_centers_:
distances = np.sum((coords_normalized - center) ** 2, axis=1)
nearest_node = self.wn.junction_name_list[np.argmin(distances)]
self.monitor_nodes.append(nearest_node)
return self.monitor_nodes
def visualize_network(self):
"""Visualize network with monitoring points"""
ax=wntr.graphics.plot_network(self.wn,
node_attribute=self.monitor_nodes,
node_size=30,
title='Optimal sensor')
plt.show()
def kmeans_sensor_placement(name: str, sensor_num: int, min_diameter: int) -> list:
inp_name = f'./db_inp/{name}.db.inp'
wn= wntr.network.WaterNetworkModel(inp_name)
wn_cluster=QD_KMeans(wn, sensor_num)
# Select monitoring pointse
sensor_ids= wn_cluster.select_monitoring_points()
# wn_cluster.visualize_network()
return sensor_ids
if __name__ == "__main__":
#sensorindex = get_ID(name='suzhouhe_2024_cloud_0817', sensor_num=30, min_diameter=500)
sensorindex = kmeans_sensor_placement(name='szh', sensor_num=50, min_diameter=300)
print(sensorindex)
-653
View File
@@ -1,653 +0,0 @@
# 改进灵敏度法
import networkx
import numpy as np
import pandas
import wntr
import pandas as pd
import copy
import matplotlib.pyplot as plt
import networkx as nx
from sklearn.cluster import KMeans
from wntr.epanet.toolkit import EpanetException
from numpy.linalg import slogdet
import random
from matplotlib.lines import Line2D
from sklearn.cluster import SpectralClustering
import libpysal as ps
from spopt.region import Skater
from shapely.geometry import Point
import geopandas as gpd
from sklearn.metrics import pairwise_distances
import app.services.project_info as project_info
# 2025/03/12
# Step1: 获取节点坐标
def getCoor(wn: wntr.network.WaterNetworkModel) -> pandas.DataFrame:
"""
获取管网模型的节点坐标
:param wn: 由wntr生成的模型
:return: 节点坐标
"""
# site: pandas.Series
# index:节点名称(wn.node_name_list
# values:每个节点的坐标,格式为 tuple(如 (x, y) 或 (x, y, z)
site = wn.query_node_attribute('coordinates')
# Coor: pandas.Series
# index:与site相同(节点名称)。
# values:坐标转换为numpy.ndarray(如array([10.5, 20.3])
Coor = site.apply(lambda x: np.array(x)) # 将节点坐标转换为numpy数组
# x, y: list[float]
x = [] # 存储所有节点的 x 坐标
y = [] # 存储所有节点的 y 坐标
for i in range(0, len(Coor)):
x.append(Coor.values[i][0]) # 将 x 坐标存入 x 列表。
y.append(Coor.values[i][1]) # 将 y 坐标存入 y 列表
# xy: dict[str, list], x、y 坐标的字典
xy = {'x': x, 'y': y}
# Coor_node: pandas.DataFrame, 存储节点 x, y 坐标的 DataFrame
Coor_node = pd.DataFrame(xy, index=wn.node_name_list, columns=['x', 'y'])
return Coor_node
# 2025/03/12
# Step2: KMeans 聚类
# 将节点用kmeans根据坐标分为k组,存入字典g
def kgroup(coor: pandas.DataFrame, knum: int) -> dict[int, list[str]]:
"""
使用KMeans聚类,将节点坐标分组
:param coor: 存储所有节点的坐标数据
:param knum: 需要分成的聚类数
:return: 聚类结果字典
"""
g = {}
# estimator: sklearn.cluster.KMeans,KMeans 聚类模型
estimator = KMeans(n_clusters=knum)
estimator.fit(coor)
# label_pred: numpy.ndarrayint,每个点的类别标签
label_pred = estimator.labels_
for i in range(0, knum):
g[i] = coor[label_pred == i].index.tolist()
return g
def skater_partition(G, n_clusters):
"""
使用 SKATER 算法对输入的无向图 G 进行区域划分,
保证每个划分区域在图论意义上是连通的,
同时依据节点坐标的空间信息进行划分。
参数:
G: networkx.Graph
带有节点坐标属性(键为 'pos')的无向图。
n_clusters: int
希望划分的区域数量。
返回:
groups: dict
字典形式的聚类结果,键为区域编号,值为该区域内的节点列表。
"""
# 1. 获取所有节点坐标,假设每个节点都有 'pos' 属性
pos = nx.get_node_attributes(G, 'pos')
nodes = list(G.nodes())
# 构造坐标数组:每行为 [x, y]
coords = np.array([pos[node] for node in nodes])
# 2. 构造 GeoDataFrame:创建 DataFrame 并生成 geometry 列
df = pd.DataFrame(coords, columns=['x', 'y'], index=nodes)
# 利用 shapely 的 Point 构造空间位置
df['geometry'] = df.apply(lambda row: Point(row['x'], row['y']), axis=1)
gdf = gpd.GeoDataFrame(df, geometry='geometry')
# 3. 构造空间权重矩阵,使用 4 近邻方法(k=4,可根据实际情况调整)
w = ps.weights.KNN.from_array(coords, k=4)
w.transform = 'R'
# 4. 调用 SKATER:新版本 API 要求传入 gdf, w 以及 attrs_name(这里使用 'x' 和 'y' 作为属性)
skater = Skater(gdf, w, attrs_name=['x', 'y'], n_clusters=n_clusters)
skater.solve()
# 5. 获取聚类标签,构造成字典格式
labels = skater.labels_
groups = {}
for label, node in zip(labels, nodes):
groups.setdefault(label, []).append(node)
return groups
def spectral_partition(G, n_clusters):
"""
利用谱聚类算法对图 G 进行分区:
1. 根据所有节点的空间坐标计算欧氏距离矩阵;
2. 利用高斯核函数构造相似度矩阵;
3. 使用 SpectralClustering 进行归一化割,返回分区结果。
参数:
G: networkx.Graph
每个节点需要有 'pos' 属性,其值为 (x, y) 坐标。
n_clusters: int
希望划分的聚类数目。
返回:
groups: dict
键为聚类标签,值为该聚类对应的节点列表。
"""
# 1. 获取节点空间坐标,注意保证每个节点都有 'pos' 属性
pos_dict = nx.get_node_attributes(G, 'pos')
nodes = list(G.nodes())
coords = np.array([pos_dict[node] for node in nodes])
# 2. 计算节点之间的欧氏距离矩阵
D = pairwise_distances(coords, metric='euclidean')
# 3. 计算 sigma 值:这里取所有距离的均值,当然也可以根据实际情况调整
sigma = np.mean(D)
# 4. 构造相似度矩阵:使用高斯核函数
# A(i, j) = exp( -d(i,j)^2 / (2*sigma^2) )
A = np.exp(- (D ** 2) / (2 * sigma ** 2))
# 5. 使用谱聚类进行图分区
clustering = SpectralClustering(n_clusters=n_clusters,
affinity='precomputed',
random_state=0)
labels = clustering.fit_predict(A)
# 6. 构造字典形式的分区结果
groups = {}
for label, node in zip(labels, nodes):
groups.setdefault(label, []).append(node)
return groups
# 2025/03/12
# Step3: wn_func类,水力计算
# wn_func 主要用于计算:
# 水力距离(hydraulic length):即节点之间的水力阻力。
# 灵敏度分析(sensitivity analysis):用于优化测压点的布置。
# 一些与水力相关的函数,包括 CtoS:求水力距离,stafun:求状态函数F
# # diff:求F对P的导数,返回灵敏度矩阵A
# # sensitivity:返回灵敏度和总灵敏度
class wn_func(object):
# Step3.1: 初始化
def __init__(self, wn: wntr.network.WaterNetworkModel, min_diameter: int):
"""
获取管网模型信息
:param wn: 由wntr生成的模型
:param min_diameter: 安装的最小管径
"""
# self.results: wntr.sim.results.SimulationResults,仿真结果,包含压力、流量、水头等数据
self.results = wntr.sim.EpanetSimulator(wn).run_sim() # 存储运行结果
self.wn = wn
# self.qpandas.DataFrame,管道流量,索引为时间步长,列为管道名称
self.q = self.results.link['flowrate']
# ReservoirIndex / Tankindex: list[str],水库 / 水箱节点名称列表
ReservoirIndex = wn.reservoir_name_list
Tankindex = wn.tank_name_list
# 删除水库节点,删除与直接水库相连的虚拟管道
# self.pipes: list[str],所有管道的名称
self.pipes = wn.pipe_name_list
# self.nodes: list[str],所有节点的名称
self.nodes = wn.node_name_list
# self.coordinatespandas.Series,节点坐标,索引为节点名,值为 (x, y) 坐标的 tuple
self.coordinates = wn.query_node_attribute('coordinates')
# allpumps / allvalves: list[str],所有泵/阀门名称列表
allpumps = wn.pump_name_list
allvalves = wn.valve_name_list
# pumpstnode / pumpednode / valvestnode / valveednode: list[str],存储泵和阀门 起终点节点的名称
pumpstnode = []
pumpednode = []
valvestnode = []
valveednode = []
# Reservoirpipe / Reservoirednode: list[str],记录与水库相关的管道和节点
Reservoirpipe = []
Reservoirednode = []
for pump in allpumps:
pumpstnode.append(wn.links[pump].start_node.name)
pumpednode.append(wn.links[pump].end_node.name)
for valve in allvalves:
valvestnode.append(wn.links[valve].start_node.name)
valveednode.append(wn.links[valve].end_node.name)
for pipe in self.pipes:
if wn.links[pipe].start_node.name in ReservoirIndex:
Reservoirpipe.append(pipe)
Reservoirednode.append(wn.links[pipe].end_node.name)
if wn.links[pipe].start_node.name in Tankindex:
Reservoirpipe.append(pipe)
Reservoirednode.append(wn.links[pipe].end_node.name)
if wn.links[pipe].end_node.name in Tankindex:
Reservoirpipe.append(pipe)
Reservoirednode.append(wn.links[pipe].start_node.name)
# 泵的起终点、tank、reservoir
# self.delnodes: list[str],需要删除的节点(包括水库、泵、阀门连接的节点)
self.delnodes = list(
set(ReservoirIndex).union(Tankindex, pumpstnode, pumpednode, valvestnode, valveednode, Reservoirednode))
# 泵、起终点为tank、reservoir的管道
# self.delpipes: list[str],需要删除的管道(包括水库、泵、阀门连接的管道)
self.delpipes = list(set(wn.pump_name_list).union(wn.valve_name_list).union(Reservoirpipe))
self.pipes = [pipe for pipe in wn.pipe_name_list if pipe not in self.delpipes]
# self.L: list[float],所有管道的长度(以米为单位)
self.L = wn.query_link_attribute('length')[self.pipes].tolist()
self.n = len(self.nodes)
self.m = len(self.pipes)
# self.unit_headloss: list[float],单位水头损失(headloss 数据的第一行,单位:米/km)
self.unit_headloss = self.results.link['headloss'].iloc[0, :].tolist()
##
self.delnodes1 = list(set(ReservoirIndex).union(Tankindex))
# === 改动新增部分:筛选管径小于 min_diameter 的管道节点 ===
self.less_than_min_diameter_junction_list = []
for pipe in self.pipes:
diameter = wn.links[pipe].diameter
if diameter < min_diameter:
start_node = wn.links[pipe].start_node.name
end_node = wn.links[pipe].end_node.name
self.less_than_min_diameter_junction_list.extend([start_node, end_node])
# 去重
self.less_than_min_diameter_junction_list = list(set(self.less_than_min_diameter_junction_list))
# Step3.2: 计算水力距离
def CtoS(self):
"""
计算水力距离矩阵
:return:
"""
# 水力距离:当行索引对应的节点为控制点时,列索引对应的节点距离控制点的(路径*水头损失)的最小值
# nodeslist[str](节点名称)
nodes = copy.deepcopy(self.nodes)
# pipeslist[str](管道名称)
pipes = self.pipes
wn = self.wn
# n / m:int(节点数 / 管道数)
n = self.n
m = self.m
s1 = [0] * m
q = self.q
L = self.L
# H1pandas.DataFrame,水头数据,索引为时间步长,列为节点名
H1 = self.results.node['head'].T
# hhlist[float],计算管道两端水头之差
hh = []
# 水头损失
for p in pipes:
h1 = self.wn.links[p].start_node.name
h1 = H1.loc[str(h1)]
h2 = self.wn.links[p].end_node.name
h2 = H1.loc[str(h2)]
hh.append(abs(h1 - h2))
hh = np.array(hh)
# headlosspandas.DataFrame,管道水头损失矩阵
headloss = pd.DataFrame(hh, index=pipes).T
# s1:管道阻力系数,s2:将管道阻力系数与管道的起始节点和终止节点对应
hf = pd.DataFrame(np.array([0] * (n ** 2)).reshape(n, n), index=nodes, columns=nodes, dtype=float)
weightL = pd.DataFrame(np.array([0] * (n ** 2)).reshape(n, n), index=nodes, columns=nodes, dtype=float)
# s2为对应管道起始节点与终止节点的粗糙度系数矩阵,index代表起始节点,columns代表终止节点
G = nx.DiGraph()
for i in range(0, m):
pipe = pipes[i]
a = wn.links[pipe].start_node.name
b = wn.links[pipe].end_node.name
if q.loc[0, pipe] > 0:
hf.loc[a, b] = headloss.loc[0, pipe]
weightL.loc[a, b] = headloss.loc[0, pipe] * L[i]
G.add_weighted_edges_from([(a, b, weightL.loc[a, b])])
else:
hf.loc[b, a] = headloss.loc[0, pipe]
weightL.loc[b, a] = headloss.loc[0, pipe] * L[i]
G.add_weighted_edges_from([(b, a, weightL.loc[b, a])])
hydraulicL = pd.DataFrame(np.array([0] * (n ** 2)).reshape(n, n), index=nodes, columns=nodes, dtype=float)
for a in nodes:
if a in G.nodes:
d = nx.shortest_path_length(G, source=a, weight='weight')
for b in list(d.keys()):
hydraulicL.loc[a, b] = d[b]
hydraulicL = hydraulicL.drop(self.delnodes)
hydraulicL = hydraulicL.drop(self.delnodes, axis=1)
# 求加权水力距离
return hydraulicL, G
# Step3.3: 计算灵敏度矩阵
# 获取关系矩阵
def get_Conn(self):
"""
计算管网连接关系矩阵
:return:
"""
m = self.wn.num_links
n = self.wn.num_nodes
p = self.wn.num_pumps
v = self.wn.num_valves
self.nonjunc_index = []
self.non_link_index = []
for r in self.wn.reservoirs():
self.nonjunc_index.append(r[0])
for t in self.wn.tanks():
self.nonjunc_index.append(t[0])
# Connnumpy.matrix,节点-管道连接矩阵,起点 -1,终点 1
Conn = np.mat(np.zeros([n, m - p - v])) # 节点和管道的关系矩阵,行为节点,列为管道,起点为-1,终点为1
# NConnnumpy.matrix,节点-节点连接矩阵,有管道相连的地方设为 1
NConn = np.mat(np.zeros([n, n])) # 节点之间的关系,之间有管道为1,反之为0
# pipeslist[str],去除泵和阀门的管道列表
pipes = [pipe for pipe in self.wn.pipes() if pipe not in self.wn.pumps() and pipe not in self.wn.valves()]
for pipe_name, pipe in pipes:
start = self.wn.node_name_list.index(pipe.start_node_name)
end = self.wn.node_name_list.index(pipe.end_node_name)
p_index = self.wn.link_name_list.index(pipe_name)
Conn[start, p_index] = -1
Conn[end, p_index] = 1
NConn[start, end] = 1
NConn[end, start] = 1
self.A = Conn
link_name_list = [link for link in self.wn.link_name_list if
link not in self.wn.pump_name_list and link not in self.wn.valve_name_list]
self.A2 = pd.DataFrame(self.A, index=self.wn.node_name_list, columns=link_name_list)
self.A2 = self.A2.drop(self.delnodes)
for pipe in self.delpipes:
if pipe not in self.wn.pump_name_list and pipe not in self.wn.valve_name_list:
self.A2 = self.A2.drop(columns=pipe)
self.junc_list = self.A2.index
self.A2 = np.mat(self.A2) # 节点管道关系
self.A3 = NConn
def Jaco(self, hL: pandas.DataFrame):
"""
计算灵敏度矩阵(节点压力对粗糙度变化的响应)
:param hL: 水力距离矩阵
:return:
"""
# global result
# Anumpy.matrix, 节点-管道关系矩阵
A = self.A2
wn = self.wn
try:
result = wntr.sim.EpanetSimulator(wn).run_sim()
except EpanetException:
pass
finally:
h = result.link['headloss'][self.pipes].values[0]
q = result.link['flowrate'][self.pipes].values[0]
l = self.wn.query_link_attribute('length')[self.pipes]
C = self.wn.query_link_attribute('roughness')[self.pipes]
# headlossnumpy.ndarray,水头损失数组
headloss = np.array(h)
# 调整流量方向
for i in range(0, len(q)):
if q[i] < 0:
A[:, i] = -A[:, i]
# qnumpy.ndarray,流量数组
q = np.abs(q)
# 两个灵敏度矩阵
# B / Snumpy.matrix,灵敏度计算的中间矩阵
B = np.mat(np.diag(q / ((1.852 * headloss) + 1e-10)))
S = np.mat(np.diag(q / C))
# Xnumpy.matrix, 灵敏度矩阵
X = A * B * A.T
try:
det = np.linalg.det(X)
except RuntimeError as e:
sign, logdet = slogdet(X) # 防止溢出
det = sign * np.exp(logdet)
if det != 0:
J_H_Cw = X.I * A * S
# J_H_Q = -X.I
J_q_Cw = S - B * A.T * X.I * A * S # 去掉了delnodes和delpipes
# J_q_Q = B * A.T * X.I
else: # 当X不可逆
J_H_Cw = np.linalg.pinv(X) @ A @ S
# J_H_Q = -np.linalg.pinv(X)
J_q_Cw = S - B * A.T * np.linalg.pinv(X) * A * S
# J_q_Q = B * A.T * np.linalg.pinv(X)
Sen_pressure = []
S_pressure = np.abs(J_H_Cw).sum(axis=1).tolist() # 修改为绝对值
for ss in S_pressure:
Sen_pressure.append(ss[0])
# 求总灵敏度
SS_pressure = copy.deepcopy(hL)
for i in range(0, len(Sen_pressure)):
SS_pressure.iloc[i, :] = SS_pressure.iloc[i, :] * Sen_pressure[i]
SS = copy.deepcopy(hL)
for i in range(0, len(Sen_pressure)):
SS.iloc[i, :] = SS.iloc[i, :] * Sen_pressure[i]
# SS[i,j]:节点nodes[i]的灵敏度*该节点到nodes[j]的水力距离
return SS
# 2025/03/12
# Step4: 传感器布置优化
# Sensorplacement
# weight:分配权重
# sensor:传感器布置的位置
class Sensorplacement(wn_func):
"""
Sensorplacement 类继承了 wn_func 类,并且用于计算和优化传感器布置的位置。
"""
def __init__(self, wn: wntr.network.WaterNetworkModel, sensornum: int, min_diameter: int):
"""
:param wn: 由wntr生成的模型
:param sensornum: 传感器的数量
:param min_diameter: 安装的最小管径
"""
wn_func.__init__(self, wn, min_diameter=min_diameter)
self.sensornum = sensornum
# 1.某个节点到所有节点的加权距离之和
# 2.某个节点到该组内所有节点的加权距离之和
def sensor(self, SS: pandas.DataFrame, G: networkx.Graph, group: dict[int, list[str]]):
"""
sensor 方法是用来根据灵敏度矩阵 SS 和加权图 G 来确定传感器布置位置的
:param SS: 灵敏度矩阵,每个节点的行和列代表不同节点,矩阵元素表示节点间的灵敏度。SS.iloc[i, :] 表示第 i 行对应节点 i 到所有其他节点的灵敏度
:param G: 加权图,表示管网的拓扑结构,每个节点通过管道连接。图的边的权重通常是根据水力距离或者流量等计算的
:param group: 节点分组,字典的键是分组编号,值是该组的节点名称列表
:return:
"""
# 传感器布置个数以及位置
# W = self.weight()
n = self.n - len(self.delnodes)
nodes = copy.deepcopy(self.nodes)
for node in self.delnodes:
nodes.remove(node)
# sumSSlist[float],每个节点到其他节点的灵敏度之和。SS.iloc[i, :] 返回第 i 个节点与所有其他节点的灵敏度值,sum(SS.iloc[i, :]) 计算这些灵敏度值的总和。
sumSS = []
for i in range(0, n):
sumSS.append(sum(SS.iloc[i, :]))
# 一个整数范围,表示每个节点的索引,用作sumSS_ DataFrame的索引
indices = range(0, n)
# sumSS_pandas.DataFrame,将 sumSS 转换成 DataFrame 格式,并且将节点的总灵敏度保存到 CSV 文件 sumSS_data.csv 中
sumSS_ = pd.DataFrame(np.array(sumSS), index=indices)
# sumSS_.to_csv('sumSS_data.csv') # 存储节点总灵敏度
# sumSSpandas.DataFrame,sumSS 被转换为 DataFrame 类型,并且按总灵敏度(即灵敏度之和)降序排列。此时,sumSS 是按节点的灵敏度之和排序的 DataFrame
sumSS = pd.DataFrame(np.array(sumSS), index=nodes)
sumSS = sumSS.sort_values(by=[0], ascending=[False])
# sensorindexlist[str],用于存储根据灵敏度排序选出的传感器位置的节点名称,存储根据总灵敏度排序的节点列表,用于传感器布置
sensorindex = []
# sensorindex_2list[str],用于存储每组内根据灵敏度排序选出的传感器位置的节点名称,存储每个组内根据灵敏度排序选择的传感器节点
sensorindex_2 = []
# group_Sdict[int, pandas.DataFrame],存储每个组内的灵敏度矩阵
group_S = {}
# group_sumSSdict[int, list[float]],存储每个组内节点的总灵敏度,值为每个组内节点灵敏度之和的列表
group_sumSS = {}
# 改动
for i in range(0, len(group)):
for node in self.delnodes:
# 这里的group[i]是每个组的节点列表,代码首先去除已经被标记为删除的节点self.delnodes
if node in group[i]:
group[i].remove(node)
group_S[i] = SS.loc[group[i], group[i]]
# 对每个组内的节点,计算组内节点的总灵敏度(group_sumSS[i])。它将每个组内节点的灵敏度值相加,并且按灵敏度降序排序
group_sumSS[i] = []
for j in range(0, len(group[i])):
group_sumSS[i].append(sum(group_S[i].iloc[j, :]))
group_sumSS[i] = pd.DataFrame(np.array(group_sumSS[i]), index=group[i])
group_sumSS[i] = group_sumSS[i].sort_values(by=[0], ascending=[False])
for node in self.less_than_min_diameter_junction_list:
# 这里的group_sumSS[i]是每个分组的灵敏度节点排序列表,去除已经被标记为删除的节点self.less_than_min_diameter_junction_list
if node in group_sumSS[i]:
group_sumSS[i].remove(node)
pass
# 1.选sumSS最大的节点,然后把这个节点所在的那个组删掉,就可以不再从这个组选点。再重新排序选sumSS最大的;
# 2.在每组内选group_sumSS最大的节点
# 在这个循环中,首先选择灵敏度最高的节点Smaxnode并添加到sensorindex。然后根据灵敏度排序,删除已选的节点并继续选择下一个灵敏度最大的节点。这个过程用于选择传感器的位置
sensornum = self.sensornum
for i in range(0, sensornum):
# Smaxnodestr,最大灵敏度节点,sumSS.index[0] 表示灵敏度最高的节点
Smaxnode = sumSS.index[0]
sensorindex.append(Smaxnode)
sensorindex_2.append(group_sumSS[i].index[0])
for key, value in group.items():
if Smaxnode in value:
sumSS = sumSS.drop(index=group[key])
continue
sumSS = sumSS.sort_values(by=[0], ascending=[False])
return sensorindex, sensorindex_2
# 2025/03/13
def get_ID(name: str, sensor_num: int, min_diameter: int) -> list[str]:
"""
获取布置测压点的坐标,初始测压点布置根据灵敏度来布置,计算初始情况下的校准过程的error
:param name: 数据库名称
:param sensor_num: 测压点数目
:param min_diameter: 安装的最小管径
:return: 测压点节点ID
"""
# inp_file_realstr,输入文件名,表示原始水力模型文件的路径,该文件格式为 EPANET 输入文件(.inp),包含管网的结构信息、节点、管道、泵等数据
inp_file_real = f'./db_inp/{name}.db.inp'
# sensornumint,需要布置的传感器数量
# sensornum = sensor_num
# wn_realwntr.network.WaterNetworkModel,加载 EPANET 水力模型
wn_real = wntr.network.WaterNetworkModel(inp_file_real) # 真实粗糙度的原始管网
# sim_realwntr.sim.EpanetSimulator,创建一个水力仿真器对象
sim_real = wntr.sim.EpanetSimulator(wn_real)
# results_realwntr.sim.results.SimulationResults,运行仿真并返回结果
results_real = sim_real.run_sim()
# real_Clist[float],包含所有管道粗糙度的列表
real_C = wn_real.query_link_attribute('roughness').tolist()
# wn_fun1wn_func(继承自 object),创建 wn_func 类的实例,传入 wn_real 水力模型对象。wn_func 用于计算管网相关的水力属性,比如水力距离、灵敏度等
wn_fun1 = wn_func(wn_real, min_diameter=min_diameter)
# nodeslist[str],管网的节点名称列表
nodes = wn_fun1.nodes
# delnodeslist[str],被删除的节点(如水库、泵、阀门连接的节点等)
delnodes = wn_fun1.delnodes
# Coor_nodepandas.DataFrame
Coor_node = getCoor(wn_real)
Coor_node = Coor_node.drop(wn_fun1.delnodes)
nodes = [node for node in wn_fun1.nodes if node not in delnodes]
# coordinatespandas.Series,存储所有节点的坐标,类型为 Series,索引为节点名称,值为 (x, y) 坐标对
coordinates = wn_fun1.coordinates
# 随机产生监测点
# junctionnumintnodes 的长度,表示节点的数量
junctionnum = len(nodes)
# random_numberslist[int],使用 random.sample 随机选择 sensornum(20)个节点的编号。它返回一个不重复的随机编号列表
# random_numbers = random.sample(range(junctionnum), sensor_num)
# for i in range(sensor_num):
# # print(random_numbers[i])
wn_fun1.get_Conn()
# hLpandas.DataFrame,水力距离矩阵,表示每个节点到其他节点的水力阻力
# Gnetworkx.DiGraph,加权有向图,表示管网的拓扑结构,节点之间的边带有权重
hL, G = wn_fun1.CtoS()
# SSpandas.DataFrame,灵敏度矩阵,表示每个节点对管网变化(如粗糙度、流量等)的响应
SS = wn_fun1.Jaco(hL)
# groupdict[int, list[str]],使用 kgroup 函数将节点按坐标分成若干组,每组包含的节点数不一定相同。group 是一个字典,键为分组编号,值为节点名列表
G1 = wn_real.to_graph()
G1 = G1.to_undirected() # 变为无向图
group = kgroup(Coor_node, sensor_num)
# group = skater_partition(G1, sensor_num)
# group = spectral_partition(G1, sensor_num)
# print(group)
# --------------------- 保存 group 数据 ---------------------
# 将 group 数据转换为一个“长格式”的 DataFrame,
# 每一行记录一个节点及其所属的分组
# group_data = []
# for group_id, node_list in group.items():
# for node in node_list:
# group_data.append({"Group": group_id, "Node": node})
#
# df_group = pd.DataFrame(group_data)
#
# # 保存为 Excel 文件,文件名为 "group.xlsx"index=False 表示不保存行索引
# df_group.to_excel("group.xlsx", index=False)
# wn_funSensorplacement(继承自wn_func
# 创建Sensorplacement类的实例,传入水力网络模型wn_real和传感器数量sensornum。Sensorplacement用于计算和布置传感器
wn_fun = Sensorplacement(wn_real, sensor_num, min_diameter=min_diameter)
wn_fun.__dict__.update(wn_fun1.__dict__)
# sensorindexlist[str],初始传感器布置位置的节点名称
# sensorindex_2list[str],根据分组选择的传感器位置
sensorindex, sensorindex_2 = wn_fun.sensor(SS, G, group) # 初始的sensorindex
# print(str(sensor_num), "个测压点,测压点位置:", sensorindex)
# 重新打开数据库
# if is_project_open(name=name):
# close_project(name=name)
# open_project(name=name)
# for node_id in sensorindex :
# sensor_coord[node_id] = get_node_coord(name=name, node_id=node_id)
# close_project(name=name)
# print(sensor_coord)
# # 分区画图
# colorlist = ['lightpink', 'coral', 'rosybrown', 'olive', 'powderblue', 'lightskyblue', 'steelblue', 'peachpuff','brown','silver','indigo','lime','gold','violet','maroon','navy','teal','magenta','cyan',
# 'burlywood', 'tan', 'slategrey', 'thistle', 'lightseagreen', 'lightgreen', 'red','blue','yellow','orange','purple','grey','green','pink','lightblue','beige','chartreuse','turquoise','lavender','fuchsia','coral']
# G = wn_real.to_graph()
# G = G.to_undirected() # 变为无向图
# pos = nx.get_node_attributes(G, 'pos')
# pass
#
# for i in range(0, sensor_num):
# ax = plt.gca()
# ax.set_title(inp_file_real + str(sensor_num))
# nodes = nx.draw_networkx_nodes(G, pos, nodelist=group[i], node_color=colorlist[i], node_size=10)
# nodes = nx.draw_networkx_nodes(G, pos,
# nodelist=sensorindex_2, node_color='red', node_size=70, node_shape='*'
# )
# edges = nx.draw_networkx_edges(G, pos)
# ax.spines['top'].set_visible(False)
# ax.spines['right'].set_visible(False)
# ax.spines['bottom'].set_visible(False)
# ax.spines['left'].set_visible(False)
# plt.savefig(inp_file_real + str(sensor_num) + ".png", dpi=300)
# plt.show()
#
# wntr.graphics.plot_network(wn_real, node_attribute=sensorindex_2, node_size=50, node_labels=False,
# title=inp_file_real + '_Projetion' + str(sensor_num))
# plt.savefig(inp_file_real + '_S' + str(sensor_num) + ".png", dpi=300)
# plt.show()
return sensorindex
if __name__ == '__main__':
sensorindex = get_ID(name=project_info.name, sensor_num=20, min_diameter=300)
print(sensorindex)
# 将 sensor_coord 字典转换为 DataFrame
# 使用 orient='index' 表示字典的键作为 DataFrame 的行索引,
# 数据中每个键对应的 value 是一个子字典,其键 'x' 和 'y' 成为 DataFrame 的列名
# df_sensor_coord = pd.DataFrame.from_dict(sensor_coord, orient='index')
#
# # 将索引名称设为 'Node'
# df_sensor_coord.index.name = 'Node'
#
# # 保存到 Excel 文件
# df_sensor_coord.to_excel("sensor_coord.xlsx", index=True)
-19
View File
@@ -1,19 +0,0 @@
from app.algorithms.simulation.scenarios import (
convert_to_local_unit,
burst_analysis,
valve_close_analysis,
flushing_analysis,
contaminant_simulation,
age_analysis,
pressure_regulation,
)
__all__ = [
"convert_to_local_unit",
"burst_analysis",
"valve_close_analysis",
"flushing_analysis",
"contaminant_simulation",
"age_analysis",
"pressure_regulation",
]
-876
View File
@@ -1,876 +0,0 @@
import numpy as np
from app.services.tjnetwork import (
ChangeSet,
close_project,
copy_project,
delete_project,
get_pattern,
get_patterns,
get_pump,
get_reservoir,
get_status,
get_tank,
get_time,
have_project,
is_project_open,
open_project,
read_all,
run_project,
set_pattern,
set_status,
set_tank,
set_time,
)
# from get_real_status import *
from datetime import datetime,timedelta
from math import modf
import json
import pytz
import requests
import time
import app.services.project_info as project_info
url_path = 'http://10.101.15.16:9000/loong' # 内网
# url_path = 'http://183.64.62.100:9057/loong' # 外网
url_real = url_path + '/api/mpoints/realValue'
url_hist = url_path + '/api/curves/data'
PATTERN_TIME_STEP=15.0
DN_900_ID='2498'
DN_500_ID='3854'
DN_1000_ID='3853'
H_RESSURE='2510'
L_PRESURE='2514'
H_TANK='4780'
L_TANK='4854'
H_REGION_1='SA_ZBBDJSCP000002'
H_REGION_2='' #to do
L_REGION_1='SA_ZBBDTJSC000001'
L_REGION_2='SA_R00003'
# reservoir basic height
RESERVOIR_BASIC_HEIGHT = float(250.35)
# regions
regions = ['hp', 'lp']
regions_demand_patterns = {'hp': ['DN900', 'DN500'], 'lp': ['DN1000']} # 出厂水量近似表示用水量
regions_patterns = {'hp': ['ChuanYiJiXiao', 'BeiQuanHuaYuan', 'ZhuangYuanFuDi', 'JingNingJiaYuan',
'308', 'JiaYinYuan', 'XinChengGuoJi', 'YiJingBeiChen', 'ZhongYangXinDu',
'XinHaiJiaYuan', 'DongFengJie', 'DingYaXinYu', 'ZiYunTai', 'XieMaGuangChang',
'YongJinFu', 'BianDianZhan', 'BeiNanDaDao', 'TianShengLiJie', 'XueYuanXiaoQu',
'YunHuaLu', 'GaoJiaQiao', 'LuZuoFuLuXiaDuan', 'TianRunCheng', 'CaoJiaBa',
'PuLingChang', 'QiLongXiaoQu', 'TuanXiao',
'TuanShanBaoZhongShiHua', 'XieMa', 'BeiWenQuanJiuHaoErQi', 'LaiYinHuSiQi',
'DN500', 'DN900'],
'lp': ['PanXiMingDu', 'WanKeJinYuHuaFuGaoCeng', 'KeJiXiao',
'LuGouQiao', 'LongJiangHuaYuan', 'LaoQiZhongDui', 'ShiYanCun', 'TianQiDaSha',
'TianShengPaiChuSuo', 'TianShengShangPin', 'JiaoTang', 'RenMinHuaYuan',
'TaiJiBinJiangYiQi', 'TianQiHuaYuan', 'TaiJiBinJiangErQi', '122Zhong',
'WanKeJinYuHuaFuYangFang', 'ChengBeiCaiShiKou', 'WenXingShe', 'YueLiangTianBBGJCZ',
'YueLiangTian', 'YueLiangTian200', 'ChengTaoChang', 'HuoCheZhan', 'LiangKu', 'QunXingLu',
'JiuYuanErTongYiYuan', 'TangDouHua', 'TaiJiBinJiangErQi(SanJi)',
'ZhangDouHua', 'JinYunXiaoQuDN400',
'DN1000']}
# nodes
monitor_single_patterns = ['ChuanYiJiXiao', 'BeiQuanHuaYuan', 'ZhuangYuanFuDi', 'JingNingJiaYuan',
'308', 'JiaYinYuan', 'XinChengGuoJi', 'YiJingBeiChen', 'ZhongYangXinDu',
'XinHaiJiaYuan', 'DongFengJie', 'DingYaXinYu', 'ZiYunTai', 'XieMaGuangChang',
'YongJinFu', 'PanXiMingDu', 'WanKeJinYuHuaFuGaoCeng', 'KeJiXiao',
'LuGouQiao', 'LongJiangHuaYuan', 'LaoQiZhongDui', 'ShiYanCun', 'TianQiDaSha',
'TianShengPaiChuSuo', 'TianShengShangPin', 'JiaoTang', 'RenMinHuaYuan',
'TaiJiBinJiangYiQi', 'TianQiHuaYuan', 'TaiJiBinJiangErQi', '122Zhong',
'WanKeJinYuHuaFuYangFang']
monitor_single_patterns_id = {'ChuanYiJiXiao': '7338', 'BeiQuanHuaYuan': '7315', 'ZhuangYuanFuDi': '7316',
'JingNingJiaYuan': '7528', '308': '8272', 'JiaYinYuan': '7304',
'XinChengGuoJi': '7325', 'YiJingBeiChen': '7328', 'ZhongYangXinDu': '7329',
'XinHaiJiaYuan': '9138', 'DongFengJie': '7302', 'DingYaXinYu': '7331',
'ZiYunTai': '7420,9059', 'XieMaGuangChang': '7326', 'YongJinFu': '9059',
'PanXiMingDu': '7320', 'WanKeJinYuHuaFuGaoCeng': '7419',
'KeJiXiao': '7305', 'LuGouQiao': '7306', 'LongJiangHuaYuan': '7318',
'LaoQiZhongDui': '9075', 'ShiYanCun': '7309', 'TianQiDaSha': '7323',
'TianShengPaiChuSuo': '7335', 'TianShengShangPin': '7324', 'JiaoTang': '7332',
'RenMinHuaYuan': '7322', 'TaiJiBinJiangYiQi': '7333', 'TianQiHuaYuan': '8235',
'TaiJiBinJiangErQi': '7334', '122Zhong': '7314', 'WanKeJinYuHuaFuYangFang': '7418'}
monitor_unity_patterns = ['BianDianZhan', 'BeiNanDaDao', 'TianShengLiJie', 'XueYuanXiaoQu',
'YunHuaLu', 'GaoJiaQiao', 'LuZuoFuLuXiaDuan', 'TianRunCheng',
'CaoJiaBa', 'PuLingChang', 'QiLongXiaoQu', 'TuanXiao',
'ChengBeiCaiShiKou', 'WenXingShe', 'YueLiangTianBBGJCZ',
'YueLiangTian', 'YueLiangTian200',
'ChengTaoChang', 'HuoCheZhan', 'LiangKu', 'QunXingLu',
'TuanShanBaoZhongShiHua', 'XieMa', 'BeiWenQuanJiuHaoErQi', 'LaiYinHuSiQi',
'JiuYuanErTongYiYuan', 'TangDouHua', 'TaiJiBinJiangErQi(SanJi)',
'ZhangDouHua', 'JinYunXiaoQuDN400',
'DN500', 'DN900', 'DN1000']
monitor_unity_patterns_id = {'BianDianZhan': '7339', 'BeiNanDaDao': '7319', 'TianShengLiJie': '8242',
'XueYuanXiaoQu': '7327', 'YunHuaLu': '7312', 'GaoJiaQiao': '7340',
'LuZuoFuLuXiaDuan': '7343', 'TianRunCheng': '7310', 'CaoJiaBa': '7300',
'PuLingChang': '7307', 'QiLongXiaoQu': '7321', 'TuanXiao': '8963',
'ChengBeiCaiShiKou': '7330', 'WenXingShe': '7311',
'YueLiangTianBBGJCZ': '7313', 'YueLiangTian': '7313', 'YueLiangTian200': '7313',
'ChengTaoChang': '7301', 'HuoCheZhan': '7303',
'LiangKu': '7296', 'QunXingLu': '7308',
'DN500': '3854', 'DN900': '2498', 'DN1000': '3853'}
monitor_patterns = monitor_single_patterns + monitor_unity_patterns
monitor_patterns_id = {**monitor_single_patterns_id, **monitor_unity_patterns_id}
# pumps
pumps_name = ['1#', '2#', '3#', '4#', '5#', '6#', '7#']
pumps = ['PU00000', 'PU00001', 'PU00002', 'PU00003', 'PU00004', 'PU00005', 'PU00006']
variable_frequency_pumps = ['PU00004', 'PU00005', 'PU00006']
pumps_id = {'PU00000': '2747', 'PU00001': '2776', 'PU00002': '2730', 'PU00003': '2787',
'PU00004': '2500', 'PU00005': '2502', 'PU00006': '2504'}
# reservoirs
reservoirs = ['ZBBDJSCP000002', 'R00003']
reservoirs_id = {'ZBBDJSCP000002': '2497', 'R00003': '2571'}
# tanks
tanks = ['ZBBDTJSC000002', 'ZBBDTJSC000001']
tanks_id = {'ZBBDTJSC000002': '4780', 'ZBBDTJSC000001': '9774'}
class DataLoader:
"""数据加载器"""
def __init__(self, project_name, start_time: datetime, end_time: datetime,
pumps_control: dict = None, tank_initial_level_control: dict = None,
region_demand_control: dict = None, downloading_prohibition: bool = False):
self.project_name = project_name # 数据库名
self.current_time = self.round_time(datetime.now(pytz.timezone('Asia/Shanghai')), 1) # 圆整至整分钟
self.current_round_time = self.round_time(self.current_time, int(PATTERN_TIME_STEP))
self.updating_data_flag = True \
if self.current_round_time == self.round_time(start_time, int(PATTERN_TIME_STEP)) \
else False # 判断是否从当前时刻开始模拟(是否更新最新监测数据)
self.downloading_prohibition = downloading_prohibition # 是否禁止下载数据(默认False: 允许下载)
self.updating_data_flag = False if self.downloading_prohibition else self.updating_data_flag
self.pattern_start_index = get_pattern_index(
self.round_time(start_time, int(PATTERN_TIME_STEP)).strftime("%Y-%m-%d %H:%M:%S")) # pattern起始索引
self.pattern_end_index = get_pattern_index(
self.round_time(end_time, int(PATTERN_TIME_STEP)).strftime("%Y-%m-%d %H:%M:%S")) # pattern结束索引
self.pattern_index_list = list(range(self.pattern_start_index, self.pattern_end_index + 1)) # pattern索引列表
self.download_id = self.get_download_id() # 数据下载接口id '7338,7315,7316,...'
self.current_time_download_data = dict(
zip(self.download_id.split(','),
[np.nan]*len(list(self.download_id.split(','))))
) # {id(str): value(float)}
self.current_time_download_data_flag = dict(
zip(self.download_id.split(','),
[False]*len(list(self.download_id.split(','))))
) # 下载数据是否具备实时性, {id(str): flag(bool)}
self.old_flow_data = self.init_dict_of_list(dict(
zip(monitor_patterns,
[[np.nan]] * (len(monitor_patterns)))
)) # {pattern_name(str): flow(float)}
self.old_pattern_factor = self.init_dict_of_list(dict(
zip(monitor_patterns,
[[np.nan]] * (len(monitor_patterns)))
)) # {pattern_name(str): [pattern_factor(float)]}
self.new_flow_data = self.init_dict_of_list(dict(
zip(monitor_patterns,
[[np.nan]] * (len(monitor_patterns)))
)) # {pattern_name(str): flow(float)}
self.new_pattern_factor = self.init_dict_of_list(dict(
zip(monitor_patterns,
[[np.nan]] * (len(monitor_patterns)))
)) # {pattern_name(str): [pattern_factor(float)]}
self.reservoir_data = dict(zip(reservoirs, [np.nan]*len(reservoirs))) # {reservoir_name(str): level(float)}
self.tank_data = dict(zip(tanks, [np.nan] * len(tanks))) # {tank_name(str): level(float)}
self.pump_data = self.init_dict_of_list(
dict(zip(pumps, [[np.nan]]*len(pumps)))) # {pump_name(str): [frequency(float)]}
self.pump_control = pumps_control # {pump_name(str): [frequency(float)]}
self.tank_initial_level_control = tank_initial_level_control # {tank_name(str): level(float)}
self.region_demand_current = dict(zip(regions, [0]*len(regions))) # {region_name(str): total_demand(float)}
self.region_demand_control = region_demand_control # {region_name(str): total_demand(float)}
self.region_demand_control_factor = dict(
zip(regions, [1]*len(regions))) # 区域流量控制系数(用于调整用水量), {region_name(str): factor(float)}
def load_data(self):
"""生成数据集"""
self.download_data() # 下载实时数据
self.get_old_pattern_and_flow() # 读取历史记录pattern信息
self.cal_demand_convert_factor() # 计算用水量转换系数(设定用水量时)
self.set_new_flow() # 设置'更新'流量
self.set_new_pattern_factor() # 设置'更新'pattern factors
self.set_reservoirs() # 设置清水池
self.set_tanks() # 设置调节池
self.set_pumps() # 设置水泵
return self.pattern_start_index
def download_data(self):
"""下载数据"""
if self.updating_data_flag is True:
print('{} -- Start downloading data.'.format(
datetime.now().strftime('%Y-%m-%d %H:%M:%S')))
data_wait_flag = True
while data_wait_flag:
try:
newest_data_time = self.download_real_data(self.download_id) # 获取实时数据
except Exception as e:
print('{}\nWaiting for real data.'.format(e))
time.sleep(1)
else:
print('{} -- Downloading data ok. Newest timestamp: {}.'.format(
datetime.now().strftime('%Y-%m-%d %H:%M:%S'),
newest_data_time.strftime('%Y-%m-%d %H:%M:%S')))
data_wait_flag = False
def cal_current_region_demand(self):
"""计算区域当前用水量"""
if self.updating_data_flag is True:
for region in self.region_demand_current.keys():
total_demand = 0
for pipe in regions_demand_patterns[region]:
total_demand += self.current_time_download_data[monitor_patterns_id[pipe]] # 出厂流量
self.region_demand_current[region] = total_demand
def cal_history_region_demand(self, pattern_index_list):
"""计算区域历史用水量(对应记录的pattern)"""
old_demand = {}
for region in regions:
total_demand = 0
for pipe_pattern_name in regions_demand_patterns[region]:
old_flows, old_patterns = self.get_history_pattern_info(self.project_name, pipe_pattern_name)
for idx in pattern_index_list:
total_demand += old_flows[idx] / 4 # 15分钟水量
old_demand[region] = total_demand
return old_demand
def cal_demand_convert_factor(self):
"""计算用水量转换系数(设定用水量时)"""
self.cal_current_region_demand() # 计算区域当前时刻用水量
old_demand_moment = self.cal_history_region_demand([self.pattern_start_index]) # 计算区域目标时刻总用水量
old_demand_period = self.cal_history_region_demand(self.pattern_index_list) # 计算区域目标时段总用水量
for region in regions:
self.region_demand_control_factor[region] \
= (self.region_demand_current[region] / 4) / old_demand_moment[region] \
if self.updating_data_flag is True else 1
self.region_demand_control_factor[region] = self.region_demand_control[region] / old_demand_period[region] \
if (self.region_demand_control is not None) and (region in self.region_demand_control.keys()) \
else self.region_demand_control_factor[region]
def get_old_pattern_and_flow(self):
"""获取所有pattern的选定时段的历史记录的pattern和flow"""
for idx in monitor_patterns: # 遍历patterns
old_flows, old_patterns = self.get_history_pattern_info(self.project_name, idx)
for pattern_idx in self.pattern_index_list:
old_flow_data = old_flows[pattern_idx]
old_pattern_factor = old_patterns[pattern_idx]
if pattern_idx == self.pattern_start_index: # 起始时刻
self.old_flow_data[idx][0] = old_flow_data
self.old_pattern_factor[idx][0] = old_pattern_factor
else:
self.old_flow_data[idx].append(old_flow_data)
self.old_pattern_factor[idx].append(old_pattern_factor)
def set_new_flow(self):
"""计算模拟时段新流量(相较于历史记录)"""
for idx in self.new_flow_data.keys(): # 遍历patterns
region_name = None
for region in regions_patterns.keys():
if idx in regions_patterns[region]:
region_name = region # pattern所属分区
break
# 实时流量
if self.updating_data_flag is True:
if idx in monitor_unity_patterns[-3:]: # 出水管流量
self.new_flow_data[idx][0] = self.current_time_download_data[monitor_patterns_id[idx]]
else: # 其余流量
self.new_flow_data[idx][0] \
= self.region_demand_control_factor[region_name] * self.old_flow_data[idx][0]
# if idx == 'ZiYunTai':
# idx_a, idx_b = monitor_patterns_id[idx].split(',')
# self.new_flow_data[idx][0] \
# = self.current_time_download_data[idx_a] - self.current_time_download_data[idx_b]
# else:
# self.new_flow_data[idx][0] = self.current_time_download_data[monitor_patterns_id[idx]]
# for data_id in monitor_patterns_id[idx].split(','):
# if (self.current_time_download_data_flag[data_id] is False) \
# and (idx not in [pipe for pipe_list in regions_demand_patterns.values()
# for pipe in pipe_list]): # 无法获取实时数据
# self.new_flow_data[idx][0] \
# = self.region_demand_control_factor[region_name] * self.old_flow_data[idx][0]
# break
# 根据设定用水量修改新流量
if (self.region_demand_control is not None) \
and (region_name in self.region_demand_control.keys()):
for pattern_idx in self.pattern_index_list:
if pattern_idx == self.pattern_start_index: # 起始时刻
self.new_flow_data[idx][0] \
= self.region_demand_control_factor[region_name] * self.old_flow_data[idx][0]
else:
self.new_flow_data[idx].append(
self.region_demand_control_factor[region_name]
* self.old_flow_data[idx][self.pattern_index_list.index(pattern_idx)]
)
def set_new_pattern_factor(self):
"""更新计算选定时段(设定用水量)/时刻的pattern factor"""
pattern_index_list = self.pattern_index_list \
if self.region_demand_control is not None \
else [self.pattern_start_index]
for idx in monitor_patterns: # 遍历patterns
for pattern_idx in pattern_index_list: # 遍历需要修改的pattern(index)
pattern_idx_cls = pattern_index_list.index(pattern_idx) # 转换index(类表存储结构)
old_flow_data = self.old_flow_data[idx][pattern_idx_cls]
old_pattern_factor = self.old_pattern_factor[idx][pattern_idx_cls]
if pattern_idx_cls == 0: # 起始时刻
if idx in monitor_single_patterns:
if not np.isnan(self.new_flow_data[idx][0]):
self.new_pattern_factor[idx][0] = (self.new_flow_data[idx][0] * 1000 / 3600) # m3/h to L/s
if idx in monitor_unity_patterns:
if not np.isnan(self.new_flow_data[idx][0]):
self.new_pattern_factor[idx][0] \
= old_pattern_factor * self.new_flow_data[idx][0] / old_flow_data
else:
if idx in monitor_single_patterns:
if len(self.new_flow_data[idx]) > pattern_idx_cls:
self.new_pattern_factor[idx].append(
(self.new_flow_data[idx][pattern_idx_cls] * 1000 / 3600)) # m3/h to L/s
if idx in monitor_unity_patterns:
if len(self.new_flow_data[idx]) > pattern_idx_cls:
self.new_pattern_factor[idx].append(
old_pattern_factor
* self.new_flow_data[idx][pattern_idx_cls]
/ old_flow_data)
def set_reservoirs(self):
"""设置清水池"""
if self.updating_data_flag is True:
for idx in self.reservoir_data.keys():
if self.current_time_download_data_flag[reservoirs_id[idx]] is False: # 无法获取实时数据
print('There is no current data of reservoir: {}.'.format(idx))
else:
self.reservoir_data[idx] \
= self.current_time_download_data[reservoirs_id[idx]] + RESERVOIR_BASIC_HEIGHT
def set_tanks(self):
"""设置调节池"""
for idx in self.tank_data.keys():
if self.updating_data_flag is True:
if self.current_time_download_data_flag[tanks_id[idx]] is False: # 无法获取实时数据
print('There is no current data of tank: {}.'.format(idx))
else:
self.tank_data[idx] = self.current_time_download_data[tanks_id[idx]]
self.tank_data[idx] = self.tank_initial_level_control[idx] \
if (self.tank_initial_level_control is not None) and (idx in self.tank_initial_level_control) \
else self.tank_data[idx]
def set_pumps(self):
"""设置水泵"""
for idx in self.pump_data.keys():
if self.updating_data_flag is True:
if self.current_time_download_data_flag[pumps_id[idx]] is False: # 无法获取实时数据
print('There is no current data of pump: {}.'.format(idx))
if (self.pump_control is not None) and (idx in self.pump_control.keys()):
self.pump_data[idx] = self.pump_control[idx]
else:
self.pump_data[idx] = [self.current_time_download_data[pumps_id[idx]]]
if (self.pump_control is not None) and (idx in self.pump_control.keys()):
self.pump_data[idx] = self.pump_data[idx] + self.pump_control[idx] \
if len(self.pump_control[idx]) < len(self.pattern_index_list) \
else self.pump_control[idx] # 水泵设定
else:
if (self.pump_control is not None) and (idx in self.pump_control.keys()):
self.pump_data[idx] = self.pump_control[idx]
self.pump_data[idx] \
= list(np.array(self.pump_data[idx]) / 50) \
if idx in variable_frequency_pumps else self.pump_data[idx]
def set_valves(self):
"""设置阀门"""
pass
def download_real_data(self, ids: str):
"""加载实时数据"""
# 数据接口的地址
global url_real
# 设置GET请求的参数
params = {'ids': ids}
# 发送GET请求获取数据
response = requests.get(url_real, params=params)
# 检查响应状态码,200表示请求成功
if response.status_code == 200:
newest_data_time = None # 下载记录数据的最新时间
# 解析响应的JSON数据
data = response.json()
for realValue in data: # 取出逐个id的数据
data_time = convert_utc_to_bj(realValue['datadt']) # datetime
self.current_time_download_data[str(realValue['id'])] \
= float(realValue['realValue']) # {id(str): value(float)}
if data_time > self.current_round_time.replace(tzinfo=None) - timedelta(minutes=5): # 下载数据为实时数据
self.current_time_download_data_flag[str(realValue['id'])] = True
if newest_data_time is None:
newest_data_time = data_time
else:
newest_data_time = data_time if data_time > newest_data_time else newest_data_time # 更新最新时间
if newest_data_time <= self.current_round_time.replace(tzinfo=None) - timedelta(minutes=5): # 最新记录时间早于当前时间
warning_text = 'There is no current data with newest timestamp: {}.'.format(
newest_data_time.strftime('%Y-%m-%d %H:%M:%S'))
delta_time = self.current_round_time.replace(tzinfo=None) - newest_data_time
if delta_time < timedelta(minutes=PATTERN_TIME_STEP): # 时间接近(可等待再次下载)
raise Exception(warning_text)
else:
print(warning_text)
self.updating_data_flag = False
else:
for idx in monitor_unity_patterns[-3:]: # 出水管流量
if self.current_time_download_data_flag[monitor_patterns_id[idx]] is False: # 无法获取出水管流量的实时数据
print('There is no current data of outflow: {}.'.format(idx))
self.updating_data_flag = False
if self.updating_data_flag is False:
print('Abandon updating data with downloaded data.')
return newest_data_time
else:
# 如果请求不成功,打印错误信息
print("请求失败,状态码:", response.status_code)
raise ConnectionError('Cannot download data.')
@ staticmethod
def init_dict_of_list(dict_of_list):
"""初始化值为列表的字典(重新生成列表地址, 防止指向同一列表)"""
for idx in dict_of_list.keys():
dict_of_list[idx] = dict_of_list[idx].copy()
return dict_of_list
@ staticmethod
def get_download_id():
"""生成下载数据项的id"""
# id_list = (list(monitor_single_patterns_id.values())
# + list(monitor_unity_patterns_id.values())
# + list(tanks_id.values())
# + list(reservoirs_id.values())
# + list(pumps_id.values()))
id_list = (list(monitor_unity_patterns_id.values())[-3:]
+ list(tanks_id.values())
+ list(reservoirs_id.values())
+ list(pumps_id.values()))
id_list = sorted(set(id_list), key=id_list.index)
if None in id_list:
id_list.remove(None)
return ','.join(id_list)
@ staticmethod
def get_history_pattern_info(project_name, pattern_name):
"""读取选定pattern的保存的历史pattern信息(flow, factor)"""
factors_list = []
flow_list = []
patterns_info = read_all(project_name,
f"select * from history_patterns_flows where id = '{pattern_name}' order by _order")
for item in patterns_info:
flow_list.append(float(item['flow']))
factors_list.append(float(item['factor']))
return flow_list, factors_list
@ staticmethod
def judge_time(current_time, time_index_list):
"""时间判断"""
current_index \
= time_index_list.index(current_time) if (current_time in time_index_list) else None
return current_index
@staticmethod
def get_time_index_list(start_time: datetime, end_time: datetime, step: int):
"""生成时间索引"""
time_index_list = [] # 时间索引[str]
time_index = start_time
while time_index <= end_time:
time_index_list.append(time_index)
time_index += timedelta(minutes=step)
return time_index_list
@ staticmethod
def round_time(time_: datetime, interval=5):
"""时间向下取整到整n分钟(北京时间): 四舍六入五留双/向下取整"""
# return datetime.fromtimestamp(round(time_.timestamp() / (60 * interval)) * (60 * interval))
return datetime.fromtimestamp(int((time_.timestamp()) // (60 * interval)) * (60 * interval))
def convert_utc_to_bj(utc_time_str):
"""将utc时间(str)转换成北京时间(datetime)"""
# 解析UTC时间字符串为datetime对象
utc_time = datetime.strptime(utc_time_str, '%Y-%m-%dT%H:%M:%SZ')
# 设定UTC时区
utc_timezone = pytz.timezone('UTC')
# 转换为北京时间
beijing_timezone = pytz.timezone('Asia/Shanghai')
beijing_time = utc_time.replace(tzinfo=utc_timezone).astimezone(beijing_timezone).replace(tzinfo=None)
return beijing_time
def get_datetime(cur_datetime:str):
str_format = "%Y-%m-%d %H:%M:%S"
return datetime.strptime(cur_datetime, str_format)
def get_strftime(cur_datetime: datetime):
str_format = "%Y-%m-%d %H:%M:%S"
return cur_datetime.strftime(str_format)
def step_time(cur_datetime:str, step=5):
str_format="%Y-%m-%d %H:%M:%S"
dt=datetime.strptime(cur_datetime,str_format)
dt=dt+timedelta(minutes=step)
return datetime.strftime(dt,str_format)
def get_pattern_index(cur_datetime:str)->int:
str_format="%Y-%m-%d %H:%M:%S"
dt=datetime.strptime(cur_datetime,str_format)
hr=dt.hour
mnt=dt.minute
i=int((hr*60+mnt)/PATTERN_TIME_STEP)
return i
def get_pattern_index_str(cur_datetime:str)->str:
i=get_pattern_index(cur_datetime)
[minN,hrN]=modf(i*PATTERN_TIME_STEP/60)
minN_str=str(int(minN*60))
minN_str=minN_str.zfill(2)
hrN_str=str(int(hrN))
hrN_str=hrN_str.zfill(2)
str_i='{}:{}:00'.format(hrN_str,minN_str)
return str_i
def from_seconds_to_clock (secs: int)->str:
hrs=int(secs/3600)
minutes=int((secs-hrs*3600)/60)
seconds=(secs-hrs*3600-minutes*60)
hrs_str=str(hrs).zfill(2)
minutes_str=str(minutes).zfill(2)
seconds_str=str(seconds).zfill(2)
str_clock='{}:{}:{}'.format(hrs_str,minutes_str,seconds_str)
return str_clock
def from_clock_to_seconds (clock: str)->int:
str_format="%Y-%m-%d %H:%M:%S"
dt=datetime.strptime(clock,str_format)
hr=dt.hour
mnt=dt.minute
seconds=dt.second
return hr*3600+mnt*60+seconds
def from_clock_to_seconds_2 (clock: str)->int:
str_format="%H:%M:%S"
dt=datetime.strptime(clock,str_format)
hr=dt.hour
mnt=dt.minute
seconds=dt.second
return hr*3600+mnt*60+seconds
def from_clock_to_seconds_3 (clock: str)->int:
str_format = "%H:%M" # 更新时间格式以适应 "小时:分钟" 格式
dt = datetime.strptime(clock,str_format)
hr = dt.hour
mnt = dt.minute
seconds = dt.second
return hr * 3600 + mnt * 60
###convert datetimestring
##"XXXX-XX-XXT00:00:00Z" ->"XXXX-XX-XX 00:00:00"
def trim_time_flag(url_date_time:str)->str:
str_datetime=str.replace(url_date_time,'T',' ')
str_datetime=str.replace(str_datetime,'Z','')
return str_datetime
# 单时间步长模拟
def run_simulation(name:str,start_datetime:str,end_datetime:str=None, duration:int=900)->str:
if(is_project_open(name)):
close_project(name)
open_project(name)
#get_current_data(cur_datetime)
#extract the patternindex from datetime
#e.g. 0: the first time step for 00:00-00:14; 1: the second step for 00:15-00:30
start_datetime=trim_time_flag(start_datetime)
if(end_datetime!=None):
end_datetime=trim_time_flag(end_datetime)
## redistribute the basedemand according to the currentTotalQ and the base_totalQ
# step 1. get_real _data
if end_datetime==None or start_datetime==end_datetime:
end_datetime=step_time(start_datetime)
# # ids=['2498','3854','3853','2510','2514','4780','4854']
# # real_data=get_real_data(ids,start_datetime,end_datetime)
# print(datetime.now(pytz.timezone('Asia/Shanghai')).strftime("%Y-%m-%d %H:%M:%S")+"--获取实时数据完毕\n")
# #step 2. re-distribute the real q to base demand of the node region_sa by region_sa
# regions=get_all_service_area_ids(name)
# total_demands={}
# for region in regions:
# total_demands[region]=get_total_base_demand(name,region)
# region_demand_factor={}
# #Region_ID:SA_ZBBDJSCP000002 高区;SA_R00003+SA_ZBBDTJSC000001 低区
# H_region_real_demands=real_data[DN_900_ID][start_datetime]+real_data[DN_500_ID][start_datetime]
# L_region_real_demands=real_data[DN_1000_ID][start_datetime]
# factor_H_zone=H_region_real_demands/total_demands[H_REGION_1]/3.6 #3.6: m3/h->L/s
# factor_L_zone=L_region_real_demands/(total_demands[L_REGION_1]+total_demands[L_REGION_2])/3.6
# print(datetime.now(pytz.timezone('Asia/Shanghai')).strftime("%Y-%m-%d %H:%M:%S")+"--流量因子计算完毕完毕\n")
# for region in regions:
# region_nodes=get_nodes_in_region(name,region)
# factor=1
# if region==H_REGION_1 or H_REGION_2:
# factor=factor_H_zone
# else:
# factor=factor_L_zone
#
# for node in region_nodes:
# d=get_demand(name,node)
# for r in d['demands']:
# r['demand']=factor*r['demand']
# cs=ChangeSet()
# cs.append(d)
# set_demand(name,cs)
#
# #
# #
# print(datetime.now(pytz.timezone('Asia/Shanghai')).strftime("%Y-%m-%d %H:%M:%S")+"--节点流量重分配完毕\n")
#step 3. set pattern index to the current time,and set duration to 300 secs
#
str_pattern_start=get_pattern_index_str(start_datetime)
dic_time=get_time(name)
dic_time['PATTERN START']=str_pattern_start
if duration !=None:
dic_time['DURATION']=from_seconds_to_clock(duration)
else:
dic_time['DURATION']=dic_time['HYDRAULIC TIMESTEP']
cs=ChangeSet()
cs.operations.append(dic_time)
set_time(name,cs)
# step4. run simulation and save the result to name-time.out for download
#inp_file = 'inp\\'+name+'.inp'
#db_name=name
#dump_inp(db_name,inp_file,'2')
# result=run_inp(db_name)
result=run_project(name)
#json string format
# simulation_result, output, report
result_data=json.loads(result)
#print(result_data['simulation_result'])
print(datetime.now(pytz.timezone('Asia/Shanghai')).strftime("%Y-%m-%d %H:%M:%S")+'run finished successfully\n')
#print(result_data['report'])
return result
# 在线模拟
def run_simulation_ex(name: str, simulation_type: str, start_datetime: str,
end_datetime: str = None, duration: int = 0,
pump_control: dict[str, list] = None, tank_initial_level_control: dict[str, float] = None,
region_demand_control: dict[str, float] = None, valve_control: dict[str, dict] = None,
downloading_prohibition: bool = False) -> str:
time_cost_start = time.perf_counter()
print('{} -- Hydraulic simulation started.'.format(
datetime.now(pytz.timezone('Asia/Shanghai')).strftime('%Y-%m-%d %H:%M:%S')))
if is_project_open(name):
close_project(name)
if simulation_type.upper() == 'REALTIME': # 实时模拟(修改原数据库)
name_c = name
elif simulation_type.upper() == 'EXTENDED': # 扩展模拟(复制数据库)
name_c = '_'.join([name, 'c'])
if have_project(name_c):
if is_project_open(name_c):
close_project(name_c)
delete_project(name_c)
copy_project(name, name_c) # 备份项目
else:
raise Exception('Incorrect simulation type, choose in (realtime, extended)')
open_project(name_c)
# 时间处理
# extract the pattern index from datetime
# e.g. 0: the first time step for 00:00-00:14; 1: the second step for 00:15-00:30
# start_datetime = get_strftime(convert_utc_to_bj(start_datetime))
start_datetime = trim_time_flag(start_datetime)
if end_datetime is not None:
# end_datetime = get_strftime(convert_utc_to_bj(end_datetime))
end_datetime = trim_time_flag(end_datetime)
# pump name转化/输入值规范化
if pump_control is not None:
for key in list(pump_control.keys()):
pump_control[key] = [pump_control[key]] if type(pump_control[key]) is not list else pump_control[key]
pump_control[pumps[pumps_name.index(key)]] = pump_control.pop(key)
# 重新分配节点(nodes)水量
# 1) (single)base_demand_new=1, pattern_new=real_data
# 2) (unity)base_demand_new=base_demand_old, pattern_new=factor*pattern_old(factor=flow_new/flow_old)
# 获取需水量数据
# a) 历史pattern对应水量(读取保存数据库)
# b) 实时水量(数据接口下载)
# 修改node demand = 1, pattern factor *= demand(monitor single patterns对应node)
# nodes = get_nodes(name_c) # nodes
# for node_name in nodes: # 遍历nodes
# demands_dict = get_demand(name_c, node_name) # {'demands':[{'demand':, 'pattern':}]}
# for demands in demands_dict['demands']:
# if (demands['pattern'] in monitor_single_patterns) and (demands['demand'] != 1): # 1)
# pattern = get_pattern(name_c, demands['pattern'])
# pattern['factors'] = list(demands['demand'] * np.array(pattern['factors'])) # 修改pattern
# cs = ChangeSet()
# cs.append(pattern)
# set_pattern(name_c, cs)
# demands_dict['demands'][
# demands_dict['demands'].index(demands)
# ]['demand'] = 1 # 修改demand
# cs = ChangeSet()
# cs.append(demands_dict)
# set_demand(name_c, cs)
start_time = get_datetime(start_datetime) # datetime
end_time = get_datetime(end_datetime) \
if end_datetime is not None \
else get_datetime(start_datetime) + timedelta(seconds=duration) # datetime
# modify_pattern_start_index = get_pattern_index(start_datetime) # 待修改pattern的起始索引(int)
dataset_loader = DataLoader(project_name=name_c,
start_time=start_time, end_time=end_time,
pumps_control=pump_control, tank_initial_level_control=tank_initial_level_control,
region_demand_control=region_demand_control,
downloading_prohibition=downloading_prohibition) # 实例化数据加载器
modify_index \
= dataset_loader.load_data() # 加载数据(index: 需要修改pattern的factor index, None: 无需修改除水泵和调节池外pattern)
new_patterns \
= dataset_loader.new_pattern_factor # {name: float,} pattern factor(实时: 更新, 其他: 保持/更新(设定用水量时))
tank_init_level = dataset_loader.tank_data # {name: float,} 调节池初始液位(实时: 更新, 其他: 保持/更新(设定液位时))
reservoir_level = dataset_loader.reservoir_data # {name: float,} 水库液位(实时: 更新, 其他: 保持)
pump_freq = dataset_loader.pump_data # {name: [float,]} 水泵频率(实时: 更新, 其他: 保持/更新(设定状态时))
print(datetime.now(pytz.timezone('Asia/Shanghai')).strftime("%Y-%m-%d %H:%M:%S") + " -- Loading data ok.\n")
pattern_name_list = get_patterns(name_c) # 所有pattern
# 修改node pattern/demand
# nodes = get_nodes(name_c) # nodes
# for node_name in nodes: # 遍历nodes
# demands_dict = get_demand(name_c, node_name) # {'demands':[{'demand':, 'pattern':}]}
# for demands in demands_dict['demands']:
# if demands['pattern'] in monitor_single_patterns: # 1)
# demands_dict['demands'][
# demands_dict['demands'].index(demands)
# ]['demand'] = 1 # 修改demand
# pattern = get_pattern(name_c, demands['pattern'])
# pattern['factors'][modify_index] = flow_new[demands['pattern']] # 修改pattern
# cs = ChangeSet()
# cs.append(pattern)
# set_pattern(name_c, cs)
# if demands['pattern'] in pattern_name_list:
# pattern_name_list.remove(demands['pattern']) # 移出待修改pattern列表
# else: # 2)
# continue
# cs = ChangeSet()
# cs.append(demands_dict)
# set_demand(name_c, cs)
for pattern_name in monitor_patterns: # 遍历patterns
if not np.isnan(new_patterns[pattern_name][0]):
pattern = get_pattern(name_c, pattern_name)
pattern['factors'][modify_index:
modify_index + len(new_patterns[pattern_name])] \
= new_patterns[pattern_name]
cs = ChangeSet()
cs.append(pattern)
set_pattern(name_c, cs)
if pattern_name in pattern_name_list:
pattern_name_list.remove(pattern_name) # 移出待修改pattern列表
# 修改清水池(reservoir)液位pattern
for reservoir_name in reservoirs: # 遍历reservoirs
if (not np.isnan(reservoir_level[reservoir_name])) and (reservoir_level[reservoir_name] != 0):
reservoir_pattern = get_pattern(name_c, get_reservoir(name_c, reservoir_name)['pattern'])
reservoir_pattern['factors'][modify_index] = reservoir_level[reservoir_name]
cs = ChangeSet()
cs.append(reservoir_pattern)
set_pattern(name_c, cs)
if reservoir_pattern['id'] in pattern_name_list:
pattern_name_list.remove(reservoir_pattern['id']) # 移出待修改pattern列表
# 修改调节池(tank)初始液位
for tank_name in tanks: # 遍历tanks
if (not np.isnan(tank_init_level[tank_name])) and (tank_init_level[tank_name] != 0):
tank = get_tank(name_c, tank_name)
tank['init_level'] = tank_init_level[tank_name]
cs = ChangeSet()
cs.append(tank)
set_tank(name_c, cs)
# 修改水泵(pump)pattern
for pump_name in pumps: # 遍历pumps
if not np.isnan(pump_freq[pump_name][0]):
pump_pattern = get_pattern(name_c, get_pump(name_c, pump_name)['pattern'])
pump_pattern['factors'][modify_index
:modify_index + len(pump_freq[pump_name])] \
= pump_freq[pump_name]
cs = ChangeSet()
cs.append(pump_pattern)
set_pattern(name_c, cs)
if pump_pattern['id'] in pattern_name_list:
pattern_name_list.remove(pump_pattern['id']) # 移出待修改pattern列表
# 修改阀门(valve)status和setting
if valve_control is not None:
for valve in valve_control.keys():
status = get_status(name_c, valve)
if 'status' in valve_control[valve].keys():
status['status'] = valve_control[valve]['status']
if 'setting' in valve_control[valve].keys():
status['setting'] = valve_control[valve]['setting']
if 'k' in valve_control[valve].keys():
valve_k = valve_control[valve]['k']
if valve_k == 0:
status['status'] = 'CLOSED'
else:
status['setting'] = 0.1036 * pow(valve_k, -3.105)
cs = ChangeSet()
cs.append(status)
set_status(name_c, cs)
print('Finish demands amending, unmodified patterns: {}.'.format(pattern_name_list))
# 修改时间信息
str_pattern_start = get_pattern_index_str(
DataLoader.round_time(start_time, int(PATTERN_TIME_STEP)).strftime("%Y-%m-%d %H:%M:%S"))
dic_time = get_time(name_c)
dic_time['PATTERN START'] = str_pattern_start
if duration is not None:
dic_time['DURATION'] = from_seconds_to_clock(duration)
else:
dic_time['DURATION'] = dic_time['HYDRAULIC TIMESTEP']
cs = ChangeSet()
cs.operations.append(dic_time)
set_time(name_c, cs)
# 运行并返回结果
result = run_project(name_c)
time_cost_end = time.perf_counter()
print('{} -- Hydraulic simulation finished, cost time: {:.2f} s.'.format(
datetime.now(pytz.timezone('Asia/Shanghai')).strftime('%Y-%m-%d %H:%M:%S'),
time_cost_end - time_cost_start))
close_project(name_c)
return result
if __name__ == '__main__':
# if get_current_data()==True:
# tQ=get_current_total_Q()
# print(f"the current tQ is {tQ}\n")
# data=get_hist_data(ids,conver_beingtime_to_ucttime('2024-04-10 15:05:00'),conver_beingtime_to_ucttime('2024-04-10 15:10:00'))
# open_project("beibeizone")
# read_inp("beibeizone","beibeizone-export_nochinese.inp")
# run_simulation("beibeizone","2024-04-01T08:00:00Z")
# read_inp('bb_server', 'model20_en.inp')
run_simulation_ex(
name=project_info.name, simulation_type='extended', start_datetime='2024-11-09T02:30:00Z',
# end_datetime='2024-05-30T16:00:00Z',
# duration=0,
# pump_control={'PU00006': [45, 40]}
# region_demand_control={'hp': 6000, 'lp': 2000}
)
@@ -0,0 +1,3 @@
from app.algorithms.valve_isolation.topology_search import valve_isolation_analysis
__all__ = ["valve_isolation_analysis"]
@@ -0,0 +1,103 @@
"""Topology-only valve isolation search."""
from collections import defaultdict, deque
from typing import Any, Iterable
VALVE_LINK_TYPE = "valve"
def _parse_link_entry(link_entry: str) -> tuple[str, str, str, str]:
parts = link_entry.split(":", 3)
if len(parts) != 4:
raise ValueError(f"Invalid link entry format: {link_entry}")
return parts[0], parts[1], parts[2], parts[3]
def valve_isolation_analysis(
link_entries: Iterable[str],
accident_elements: str | list[str],
disabled_valves: list[str] | None = None,
) -> dict[str, Any]:
"""Determine boundary valves and affected nodes from a topology snapshot."""
disabled_valves_set = set(disabled_valves or [])
target_elements = (
[accident_elements]
if isinstance(accident_elements, str)
else accident_elements
)
pipe_adj: dict[str, set[str]] = defaultdict(set)
all_valves: dict[str, tuple[str, str]] = {}
link_lookup: dict[str, tuple[str, str, str]] = {}
node_set: set[str] = set()
for link_entry in link_entries:
link_id, link_type, node1, node2 = _parse_link_entry(link_entry)
link_type_name = str(link_type).lower()
link_lookup[link_id] = (node1, node2, link_type_name)
node_set.update((node1, node2))
if link_type_name == VALVE_LINK_TYPE:
all_valves[link_id] = (node1, node2)
else:
pipe_adj[node1].add(node2)
pipe_adj[node2].add(node1)
start_nodes: set[str] = set()
for element in target_elements:
if element in node_set:
start_nodes.add(element)
elif element in link_lookup:
node1, node2, _ = link_lookup[element]
start_nodes.update((node1, node2))
else:
raise ValueError(f"Accident element {element} was not found in topology")
extra_adj: dict[str, list[str]] = defaultdict(list)
boundary_valves: dict[str, tuple[str, str]] = {}
for valve_id, (node1, node2) in all_valves.items():
if valve_id in disabled_valves_set:
extra_adj[node1].append(node2)
extra_adj[node2].append(node1)
else:
boundary_valves[valve_id] = (node1, node2)
affected_nodes: set[str] = set()
queue = deque(start_nodes)
while queue:
node = queue.popleft()
if node in affected_nodes:
continue
affected_nodes.add(node)
queue.extend(pipe_adj.get(node, set()) - affected_nodes)
queue.extend(
neighbor
for neighbor in extra_adj.get(node, ())
if neighbor not in affected_nodes
)
must_close_valves: list[str] = []
optional_valves: list[str] = []
for valve_id, (node1, node2) in boundary_valves.items():
node1_affected = node1 in affected_nodes
node2_affected = node2 in affected_nodes
if node1_affected and node2_affected:
optional_valves.append(valve_id)
elif node1_affected or node2_affected:
must_close_valves.append(valve_id)
must_close_valves.sort()
optional_valves.sort()
isolatable = bool(must_close_valves)
result: dict[str, Any] = {
"accident_elements": target_elements,
"disabled_valves": disabled_valves,
"affected_nodes": sorted(affected_nodes) if isolatable else [],
"affected_node_count": len(affected_nodes),
"must_close_valves": must_close_valves,
"optional_valves": optional_valves,
"isolatable": isolatable,
}
if len(target_elements) == 1:
result["accident_element"] = target_elements[0]
return result
+14
View File
@@ -0,0 +1,14 @@
from __future__ import annotations
from collections.abc import Iterable
from typing import Generic, TypeVar
T = TypeVar("T")
class PaginatedList(list[T], Generic[T]):
"""A page of items carrying the total count from its data source."""
def __init__(self, items: Iterable[T], *, total: int) -> None:
super().__init__(items)
self.total = total
+109
View File
@@ -0,0 +1,109 @@
from __future__ import annotations
from typing import Any
from uuid import uuid4
from fastapi import FastAPI, HTTPException, Request
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field
from app.native.wndb.core.database import MaterializedViewRefreshAfterCommitError
class ProblemDetails(BaseModel):
"""RFC 9457 compatible error response used by the REST contract."""
type: str
title: str
status: int
detail: str
instance: str
code: str
trace_id: str
errors: list[dict[str, Any]] = Field(default_factory=list)
def _trace_id(request: Request) -> str:
return request.headers.get("X-Request-Id") or str(uuid4())
def _problem_response(
request: Request,
*,
status_code: int,
title: str,
detail: str,
code: str,
errors: list[dict[str, Any]] | None = None,
) -> JSONResponse:
problem = ProblemDetails(
type=f"https://tjwater.example/problems/{code.replace('_', '-')}",
title=title,
status=status_code,
detail=detail,
instance=request.url.path,
code=code,
trace_id=_trace_id(request),
errors=errors or [],
)
return JSONResponse(
status_code=status_code,
content=problem.model_dump(mode="json"),
media_type="application/problem+json",
)
def install_problem_details_handlers(app: FastAPI) -> None:
@app.exception_handler(MaterializedViewRefreshAfterCommitError)
async def materialized_view_refresh_error_handler(
request: Request,
exc: MaterializedViewRefreshAfterCommitError,
) -> JSONResponse:
response = _problem_response(
request,
status_code=503,
title="Materialized view refresh failed",
detail=(
f"Project {exc.project!r} changes were committed, but GIS query "
"views could not be refreshed. Do not repeat the write blindly."
),
code="materialized_view_refresh_failed_after_commit",
)
response.headers["X-TJWater-Changes-Committed"] = "true"
return response
@app.exception_handler(RequestValidationError)
async def validation_error_handler(
request: Request,
exc: RequestValidationError,
) -> JSONResponse:
return _problem_response(
request,
status_code=422,
title="Validation error",
detail="Request validation failed",
code="validation_error",
errors=exc.errors(),
)
@app.exception_handler(HTTPException)
async def http_error_handler(request: Request, exc: HTTPException) -> JSONResponse:
detail = exc.detail if isinstance(exc.detail, str) else str(exc.detail)
code_by_status = {
401: "unauthenticated",
403: "forbidden",
404: "not_found",
409: "conflict",
422: "validation_error",
503: "dependency_unavailable",
}
return _problem_response(
request,
status_code=exc.status_code,
title=code_by_status.get(exc.status_code, "request_error")
.replace("_", " ")
.title(),
detail=detail,
code=code_by_status.get(exc.status_code, "request_error"),
)
+39
View File
@@ -0,0 +1,39 @@
from fastapi import APIRouter, Depends, Header
from app.auth.metadata_dependencies import (
get_current_metadata_user,
get_metadata_repository,
)
from app.auth.permissions import resolve_permissions
from app.auth.project_dependencies import resolve_project_context
from app.domain.schemas.access import AccessContextResponse
from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository
router = APIRouter()
@router.get("/access-context", response_model=AccessContextResponse)
async def get_access_context(
x_project_id: str | None = Header(default=None, alias="X-Project-Id"),
current_user=Depends(get_current_metadata_user),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> AccessContextResponse:
project_context = (
await resolve_project_context(x_project_id, current_user, metadata_repo)
if x_project_id
else None
)
permissions = resolve_permissions(
project_role=project_context.project_role if project_context else None,
system_role=current_user.role,
is_superuser=current_user.is_superuser,
)
return AccessContextResponse(
user_id=current_user.id,
username=current_user.username,
system_role=current_user.role,
is_system_admin=current_user.is_superuser or current_user.role == "admin",
project_id=project_context.project_id if project_context else None,
project_role=project_context.project_role if project_context else None,
permissions=sorted(permissions),
)
+702
View File
@@ -0,0 +1,702 @@
from typing import List
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Path, Query, Response, status
from sqlalchemy import text
from sqlalchemy.engine.url import make_url
from sqlalchemy.exc import IntegrityError, SQLAlchemyError
from sqlalchemy.ext.asyncio import create_async_engine
from app.auth.metadata_dependencies import (
get_current_metadata_admin,
get_metadata_repository,
)
from app.core.audit import AuditAction, log_audit_event
from app.domain.schemas.admin_metadata import (
AdminProjectCreateRequest,
AdminProjectResponse,
AdminProjectUpdateRequest,
MetadataUsersBatchSyncRequest,
MetadataUserResponse,
MetadataUserSyncRequest,
MetadataUserSyncResult,
MetadataUserUpdateRequest,
ProjectDatabaseHealthResponse,
ProjectDatabaseHealthRequest,
ProjectDatabaseResponse,
ProjectDatabaseUpsertRequest,
ProjectDbRole,
ProjectMemberCreateRequest,
ProjectMemberResponse,
ProjectMemberUpdateRequest,
)
from app.infra.db.metadb import models
from app.infra.db.metadb.repositories.metadata_repository import (
MetadataRepository,
ProjectDbRouting,
)
router = APIRouter()
def _project_response(project: models.Project) -> AdminProjectResponse:
return 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,
)
def _project_database_response(
record: models.ProjectDatabase,
) -> ProjectDatabaseResponse:
return ProjectDatabaseResponse(
id=record.id,
project_id=record.project_id,
db_role=record.db_role,
db_type=record.db_type,
pool_min_size=record.pool_min_size,
pool_max_size=record.pool_max_size,
has_dsn=bool(record.dsn_encrypted),
)
def _database_audit_payload(payload: ProjectDatabaseUpsertRequest) -> dict:
return {
"db_role": payload.db_role,
"db_type": _db_type_for_role(payload.db_role),
"pool_min_size": payload.pool_min_size,
"pool_max_size": payload.pool_max_size,
"dsn_updated": payload.dsn is not None,
}
def _to_async_sqlalchemy_url(dsn: str) -> str:
parsed = make_url(dsn)
if parsed.drivername in {"postgresql", "postgres"}:
parsed = parsed.set(drivername="postgresql+psycopg")
return parsed.render_as_string(hide_password=False)
def _db_type_for_role(db_role: str) -> str:
if db_role == "iot_data":
return "timescaledb"
return "postgresql"
def _status_for_config_value_error(exc: ValueError) -> int:
if "DATABASE_ENCRYPTION_KEY" in str(exc):
return status.HTTP_503_SERVICE_UNAVAILABLE
return status.HTTP_400_BAD_REQUEST
async def _check_database_connection(routing: ProjectDbRouting) -> None:
engine = create_async_engine(
_to_async_sqlalchemy_url(routing.dsn),
pool_size=1,
max_overflow=0,
pool_pre_ping=True,
)
try:
async with engine.connect() as conn:
await conn.execute(text("SELECT 1"))
finally:
await engine.dispose()
def _database_health_error_detail(exc: Exception) -> str:
message = str(exc)
lower_message = message.lower()
if "password authentication failed" in lower_message:
return "连通性测试失败:用户名或密码错误,请检查 DSN 中的账号密码。"
if "connection refused" in lower_message:
return "连通性测试失败:目标主机或端口拒绝连接,请检查地址、端口和服务状态。"
if "timeout" in lower_message or "timed out" in lower_message:
return "连通性测试失败:连接超时,请检查网络、防火墙和数据库服务状态。"
if "could not translate host name" in lower_message or "name or service not known" in lower_message:
return "连通性测试失败:数据库主机名无法解析,请检查 DSN 中的主机地址。"
first_line = message.splitlines()[0] if message else exc.__class__.__name__
return f"连通性测试失败:{first_line}"
async def _upsert_and_audit_metadata_user(
payload: MetadataUserSyncRequest,
*,
current_user,
metadata_repo: MetadataRepository,
response_status: int,
) -> MetadataUserResponse:
user = await metadata_repo.upsert_user_from_keycloak(
keycloak_id=payload.keycloak_id,
username=payload.username,
email=str(payload.email),
role=payload.role,
is_active=payload.is_active,
)
await log_audit_event(
action=AuditAction.UPDATE,
user_id=current_user.id,
resource_type="metadata_user",
resource_id=str(user.id),
request_data=payload.model_dump(mode="json"),
response_status=response_status,
session=metadata_repo.session,
)
return MetadataUserResponse.model_validate(user)
@router.get("/admin/users/me", response_model=MetadataUserResponse)
async def get_metadata_admin_me(
current_user=Depends(get_current_metadata_admin),
) -> MetadataUserResponse:
return MetadataUserResponse.model_validate(current_user)
@router.post("/admin/user-syncs", response_model=MetadataUserResponse)
async def sync_metadata_user(
payload: MetadataUserSyncRequest,
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> MetadataUserResponse:
try:
return await _upsert_and_audit_metadata_user(
payload,
current_user=current_user,
metadata_repo=metadata_repo,
response_status=status.HTTP_200_OK,
)
except IntegrityError as exc:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="User keycloak_id, username, or email conflicts with an existing user",
) from exc
except SQLAlchemyError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"Metadata database error: {exc}",
) from exc
@router.post("/admin/user-syncs/batches", response_model=List[MetadataUserSyncResult])
async def sync_metadata_users_batch(
payload: MetadataUsersBatchSyncRequest,
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> List[MetadataUserSyncResult]:
results: list[MetadataUserSyncResult] = []
for item in payload.users:
try:
user = await _upsert_and_audit_metadata_user(
item,
current_user=current_user,
metadata_repo=metadata_repo,
response_status=status.HTTP_200_OK,
)
except IntegrityError as exc:
results.append(
MetadataUserSyncResult(
keycloak_id=item.keycloak_id,
success=False,
error="User keycloak_id, username, or email conflicts with an existing user",
)
)
await metadata_repo.session.rollback()
except SQLAlchemyError as exc:
results.append(
MetadataUserSyncResult(
keycloak_id=item.keycloak_id,
success=False,
error=f"Metadata database error: {exc}",
)
)
await metadata_repo.session.rollback()
else:
results.append(
MetadataUserSyncResult(
keycloak_id=item.keycloak_id,
success=True,
user=user,
)
)
return results
@router.get("/admin/users", response_model=List[MetadataUserResponse])
async def list_metadata_users(
skip: int = Query(0, ge=0),
limit: int = Query(100, ge=1, le=1000),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> List[MetadataUserResponse]:
users = await metadata_repo.list_users(skip=skip, limit=limit)
return [MetadataUserResponse.model_validate(user) for user in users]
@router.get("/admin/projects", response_model=List[AdminProjectResponse])
async def list_admin_projects(
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> List[AdminProjectResponse]:
projects = await metadata_repo.list_project_records()
return [_project_response(project) for project in projects]
@router.post(
"/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,
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> AdminProjectResponse:
try:
project = await metadata_repo.create_project(
name=payload.name,
code=payload.code,
description=payload.description,
gs_workspace=payload.gs_workspace,
map_extent=payload.map_extent,
status=payload.status,
creator_user_id=current_user.id,
)
except IntegrityError as exc:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="Project code or workspace conflicts with an existing project",
) from exc
except SQLAlchemyError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"Metadata database error: {exc}",
) from exc
await log_audit_event(
action=AuditAction.CREATE,
user_id=current_user.id,
project_id=project.id,
resource_type="project",
resource_id=str(project.id),
request_data=payload.model_dump(mode="json"),
response_status=status.HTTP_201_CREATED,
session=metadata_repo.session,
)
return _project_response(project)
@router.patch(
"/admin/projects/{project_id}",
response_model=AdminProjectResponse,
)
async def update_admin_project(
payload: AdminProjectUpdateRequest,
project_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> AdminProjectResponse:
updates = payload.model_dump(mode="json", exclude_unset=True)
try:
project = await metadata_repo.update_project(project_id, updates=updates)
except IntegrityError as exc:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="Project code or workspace conflicts with an existing project",
) from exc
except SQLAlchemyError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"Metadata database error: {exc}",
) from exc
if project is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
await log_audit_event(
action=AuditAction.UPDATE,
user_id=current_user.id,
project_id=project.id,
resource_type="project",
resource_id=str(project.id),
request_data=updates,
response_status=status.HTTP_200_OK,
session=metadata_repo.session,
)
return _project_response(project)
@router.get(
"/admin/projects/{project_id}/databases",
response_model=List[ProjectDatabaseResponse],
)
async def list_project_databases(
project_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> List[ProjectDatabaseResponse]:
project = await metadata_repo.get_project_by_id(project_id)
if project is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
records = await metadata_repo.list_project_databases(project_id)
return [_project_database_response(record) for record in records]
@router.put(
"/admin/projects/{project_id}/databases",
response_model=ProjectDatabaseResponse,
)
async def upsert_project_database(
payload: ProjectDatabaseUpsertRequest,
project_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> ProjectDatabaseResponse:
project = await metadata_repo.get_project_by_id(project_id)
if project is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
try:
routing = (
ProjectDbRouting(
project_id=project_id,
db_role=payload.db_role,
db_type=_db_type_for_role(payload.db_role),
dsn=payload.dsn,
pool_min_size=payload.pool_min_size,
pool_max_size=payload.pool_max_size,
)
if payload.dsn
else await metadata_repo.get_project_db_routing(project_id, payload.db_role)
)
if routing is None:
raise ValueError("dsn is required when creating project database config")
await _check_database_connection(routing)
except ValueError as exc:
raise HTTPException(
status_code=_status_for_config_value_error(exc),
detail=str(exc),
) from exc
except Exception as exc:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=_database_health_error_detail(exc),
) from exc
try:
record = await metadata_repo.upsert_project_database_config(
project_id,
db_role=payload.db_role,
db_type=_db_type_for_role(payload.db_role),
dsn=payload.dsn,
pool_min_size=payload.pool_min_size,
pool_max_size=payload.pool_max_size,
)
except IntegrityError as exc:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="Project database role conflicts with an existing config",
) from exc
except SQLAlchemyError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"Metadata database error: {exc}",
) from exc
await log_audit_event(
action=AuditAction.CONFIG_CHANGE,
user_id=current_user.id,
project_id=project_id,
resource_type="project_database",
resource_id=payload.db_role,
request_data=_database_audit_payload(payload),
response_status=status.HTTP_200_OK,
session=metadata_repo.session,
)
return _project_database_response(record)
@router.delete(
"/admin/projects/{project_id}/databases/{db_role}",
status_code=status.HTTP_204_NO_CONTENT,
)
async def delete_project_database(
project_id: UUID = Path(...),
db_role: ProjectDbRole = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> None:
removed = await metadata_repo.delete_project_database_config(project_id, db_role)
if not removed:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Project database config not found",
)
await log_audit_event(
action=AuditAction.CONFIG_CHANGE,
user_id=current_user.id,
project_id=project_id,
resource_type="project_database",
resource_id=db_role,
request_data={"deleted": True},
response_status=status.HTTP_204_NO_CONTENT,
session=metadata_repo.session,
)
@router.post(
"/admin/projects/{project_id}/databases/{db_role}/health-checks",
response_model=ProjectDatabaseHealthResponse,
)
async def check_project_database_health(
response: Response,
project_id: UUID = Path(...),
db_role: ProjectDbRole = Path(...),
payload: ProjectDatabaseHealthRequest | None = None,
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> ProjectDatabaseHealthResponse:
dsn_to_test = payload.dsn if payload and payload.dsn else None
if dsn_to_test:
routing = ProjectDbRouting(
project_id=project_id,
db_role=db_role,
db_type=_db_type_for_role(db_role),
dsn=dsn_to_test,
pool_min_size=1,
pool_max_size=1,
)
else:
try:
routing = await metadata_repo.get_project_db_routing(project_id, db_role)
except ValueError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"Project database routing DSN is invalid: {exc}",
) from exc
if routing is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Project database config not found",
)
try:
await _check_database_connection(routing)
except Exception as exc: # health endpoint should return diagnostic status
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return ProjectDatabaseHealthResponse(
project_id=project_id,
db_role=db_role,
db_type=routing.db_type,
ok=False,
detail=_database_health_error_detail(exc),
)
return ProjectDatabaseHealthResponse(
project_id=project_id,
db_role=db_role,
db_type=routing.db_type,
ok=True,
detail="连通性测试通过",
)
@router.get("/admin/users/{user_id}", response_model=MetadataUserResponse)
async def get_metadata_user(
user_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> MetadataUserResponse:
user = await metadata_repo.get_user_by_id(user_id)
if user is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
return MetadataUserResponse.model_validate(user)
@router.patch("/admin/users/{user_id}", response_model=MetadataUserResponse)
async def update_metadata_user(
payload: MetadataUserUpdateRequest,
user_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> MetadataUserResponse:
updates = payload.model_dump(mode="json", exclude_unset=True)
if user_id == current_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Users cannot modify themselves",
)
user = await metadata_repo.update_user_admin(
user_id,
updates=updates,
)
if user is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
await log_audit_event(
action=AuditAction.UPDATE,
user_id=current_user.id,
resource_type="metadata_user",
resource_id=str(user.id),
request_data=updates,
response_status=status.HTTP_200_OK,
session=metadata_repo.session,
)
return MetadataUserResponse.model_validate(user)
@router.get(
"/admin/projects/{project_id}/members",
response_model=List[ProjectMemberResponse],
)
async def list_project_members(
project_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> List[ProjectMemberResponse]:
project = await metadata_repo.get_project_by_id(project_id)
if project is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="Project not found"
)
members = await metadata_repo.list_project_members(project_id)
return [ProjectMemberResponse(**member.__dict__) for member in members]
@router.post(
"/admin/projects/{project_id}/members",
response_model=ProjectMemberResponse,
status_code=status.HTTP_201_CREATED,
)
async def add_project_member(
payload: ProjectMemberCreateRequest,
project_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> ProjectMemberResponse:
if payload.user_id == current_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Users cannot modify their own project membership",
)
project = await metadata_repo.get_project_by_id(project_id)
if project is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="Project not found"
)
user = await metadata_repo.get_user_by_id(payload.user_id)
if user is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
existing = await metadata_repo.get_project_membership(project_id, payload.user_id)
if existing is not None:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="User is already a project member",
)
membership = await metadata_repo.add_project_member(
project_id, payload.user_id, payload.project_role
)
await log_audit_event(
action=AuditAction.PERMISSION_CHANGE,
user_id=current_user.id,
project_id=project_id,
resource_type="project_member",
resource_id=str(payload.user_id),
request_data=payload.model_dump(mode="json"),
response_status=status.HTTP_201_CREATED,
session=metadata_repo.session,
)
return ProjectMemberResponse(
id=membership.id,
user_id=membership.user_id,
project_id=membership.project_id,
project_role=membership.project_role,
username=user.username,
email=user.email,
is_active=user.is_active,
)
@router.patch(
"/admin/projects/{project_id}/members/{user_id}",
response_model=ProjectMemberResponse,
)
async def update_project_member(
payload: ProjectMemberUpdateRequest,
project_id: UUID = Path(...),
user_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> ProjectMemberResponse:
if user_id == current_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Users cannot modify their own project membership",
)
user = await metadata_repo.get_user_by_id(user_id)
if user is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
membership = await metadata_repo.update_project_member_role(
project_id, user_id, payload.project_role
)
if membership is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="Project member not found"
)
await log_audit_event(
action=AuditAction.PERMISSION_CHANGE,
user_id=current_user.id,
project_id=project_id,
resource_type="project_member",
resource_id=str(user_id),
request_data=payload.model_dump(mode="json"),
response_status=status.HTTP_200_OK,
session=metadata_repo.session,
)
return ProjectMemberResponse(
id=membership.id,
user_id=membership.user_id,
project_id=membership.project_id,
project_role=membership.project_role,
username=user.username,
email=user.email,
is_active=user.is_active,
)
@router.delete("/admin/projects/{project_id}/members/{user_id}", status_code=status.HTTP_204_NO_CONTENT)
async def remove_project_member(
project_id: UUID = Path(...),
user_id: UUID = Path(...),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> None:
if user_id == current_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Users cannot modify their own project membership",
)
removed = await metadata_repo.remove_project_member(project_id, user_id)
if not removed:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="Project member not found"
)
await log_audit_event(
action=AuditAction.PERMISSION_CHANGE,
user_id=current_user.id,
project_id=project_id,
resource_type="project_member",
resource_id=str(user_id),
response_status=status.HTTP_204_NO_CONTENT,
session=metadata_repo.session,
)
+53
View File
@@ -0,0 +1,53 @@
from datetime import datetime, timezone
from fastapi import APIRouter, Depends
from pydantic import BaseModel
from app.auth.keycloak_dependencies import get_current_keycloak_payload
from app.auth.metadata_dependencies import get_current_metadata_user
from app.auth.project_dependencies import (
ProjectContext,
get_project_context,
)
from app.auth.permissions import permissions_for_context
router = APIRouter()
class AgentAuthContextResponse(BaseModel):
user_id: str
keycloak_sub: str
username: str
role: str
is_superuser: bool
project_id: str
network: str
project_role: str
permissions: list[str]
token_expires_at: str | None = None
@router.get("/agent-auth-context", response_model=AgentAuthContextResponse)
async def get_agent_auth_context(
ctx: ProjectContext = Depends(get_project_context),
current_user=Depends(get_current_metadata_user),
keycloak_payload: dict = Depends(get_current_keycloak_payload),
) -> AgentAuthContextResponse:
exp = keycloak_payload.get("exp")
token_expires_at = (
datetime.fromtimestamp(exp, tz=timezone.utc).isoformat()
if isinstance(exp, int)
else None
)
return AgentAuthContextResponse(
user_id=str(current_user.id),
keycloak_sub=str(current_user.keycloak_id),
username=current_user.username,
role=current_user.role,
is_superuser=current_user.is_superuser,
project_id=str(ctx.project_id),
network=ctx.project_code,
project_role=ctx.project_role,
permissions=sorted(permissions_for_context(ctx)),
token_expires_at=token_expires_at,
)
+77 -56
View File
@@ -1,56 +1,53 @@
"""
审计日志 API 接口
仅管理员可访问
"""
from typing import List, Optional
from uuid import UUID
from datetime import datetime
from fastapi import APIRouter, Depends, Query, Path
from app.domain.schemas.audit import AuditLogResponse
from app.infra.db.metadb.repositories.audit_repository import AuditRepository
from typing import Literal
from uuid import UUID
from fastapi import APIRouter, Depends, Query, Request, status
from pydantic import BaseModel
from sqlalchemy.ext.asyncio import AsyncSession
from app.auth.metadata_dependencies import (
get_current_metadata_admin,
get_current_metadata_user,
)
from app.api.pagination import PaginatedList
from app.core.audit import AuditAction, log_audit_event
from app.domain.schemas.audit import AuditLogResponse
from app.infra.db.metadb.database import get_metadata_session
from sqlalchemy.ext.asyncio import AsyncSession
from app.infra.db.metadb.repositories.audit_repository import AuditRepository
router = APIRouter()
class SessionAuditEventRequest(BaseModel):
event: Literal["login", "logout"]
async def get_audit_repository(
session: AsyncSession = Depends(get_metadata_session),
) -> AuditRepository:
"""获取审计日志仓储"""
return AuditRepository(session)
@router.get(
"/logs",
"/audit-logs",
summary="查询审计日志",
description="查询审计日志(仅管理员)",
response_model=List[AuditLogResponse],
response_model=list[AuditLogResponse],
)
async def get_audit_logs(
user_id: Optional[UUID] = Query(None, description="按用户ID过滤"),
project_id: Optional[UUID] = Query(None, description="按项目ID过滤"),
action: Optional[str] = Query(None, description="按操作类型过滤"),
resource_type: Optional[str] = Query(None, description="按资源类型过滤"),
start_time: Optional[datetime] = Query(None, description="开始时间"),
end_time: Optional[datetime] = Query(None, description="结束时间"),
user_id: UUID | None = Query(None, description="按用户ID过滤"),
project_id: UUID | None = Query(None, description="按项目ID过滤"),
action: str | None = Query(None, description="按操作类型过滤"),
resource_type: str | None = Query(None, description="按资源类型过滤"),
start_time: datetime | None = Query(None, description="开始时间"),
end_time: datetime | None = Query(None, description="结束时间"),
skip: int = Query(0, ge=0, description="跳过记录数"),
limit: int = Query(100, ge=1, le=1000, description="限制记录数"),
current_user=Depends(get_current_metadata_admin),
_current_user=Depends(get_current_metadata_admin),
audit_repo: AuditRepository = Depends(get_audit_repository),
) -> List[AuditLogResponse]:
"""
查询审计日志
支持按用户、时间、操作类型等条件过滤,仅管理员可访问
"""
logs = await audit_repo.get_logs(
) -> list[AuditLogResponse]:
items = await audit_repo.get_logs(
user_id=user_id,
project_id=project_id,
action=action,
@@ -60,29 +57,32 @@ async def get_audit_logs(
skip=skip,
limit=limit,
)
return logs
total = await audit_repo.get_log_count(
user_id=user_id,
project_id=project_id,
action=action,
resource_type=resource_type,
start_time=start_time,
end_time=end_time,
)
return PaginatedList(items, total=total)
@router.get(
"/logs/count",
"/audit-logs/count",
summary="获取审计日志总数",
description="获取审计日志总数(仅管理员)",
)
async def get_audit_logs_count(
user_id: Optional[UUID] = Query(None, description="按用户ID过滤"),
project_id: Optional[UUID] = Query(None, description="按项目ID过滤"),
action: Optional[str] = Query(None, description="按操作类型过滤"),
resource_type: Optional[str] = Query(None, description="按资源类型过滤"),
start_time: Optional[datetime] = Query(None, description="开始时间"),
end_time: Optional[datetime] = Query(None, description="结束时间"),
current_user=Depends(get_current_metadata_admin),
user_id: UUID | None = Query(None, description="按用户ID过滤"),
project_id: UUID | None = Query(None, description="按项目ID过滤"),
action: str | None = Query(None, description="按操作类型过滤"),
resource_type: str | None = Query(None, description="按资源类型过滤"),
start_time: datetime | None = Query(None, description="开始时间"),
end_time: datetime | None = Query(None, description="结束时间"),
_current_user=Depends(get_current_metadata_admin),
audit_repo: AuditRepository = Depends(get_audit_repository),
) -> dict:
"""
获取审计日志总数
获取符合条件的审计日志的总数,仅管理员可访问
"""
count = await audit_repo.get_log_count(
user_id=user_id,
project_id=project_id,
@@ -94,27 +94,42 @@ async def get_audit_logs_count(
return {"count": count}
@router.post("/audit-events", status_code=status.HTTP_204_NO_CONTENT)
async def record_session_event(
payload: SessionAuditEventRequest,
request: Request,
current_user=Depends(get_current_metadata_user),
session: AsyncSession = Depends(get_metadata_session),
) -> None:
await log_audit_event(
action=AuditAction.LOGIN if payload.event == "login" else AuditAction.LOGOUT,
user_id=current_user.id,
resource_type="session",
resource_id=str(current_user.keycloak_id),
ip_address=request.client.host if request.client else None,
request_method=request.method,
request_path=request.url.path,
response_status=status.HTTP_204_NO_CONTENT,
session=session,
)
@router.get(
"/logs/my",
"/audit-logs/mine",
summary="查询我的审计日志",
description="查询当前用户的审计日志",
response_model=List[AuditLogResponse],
response_model=list[AuditLogResponse],
)
async def get_my_audit_logs(
action: Optional[str] = Query(None, description="按操作类型过滤"),
start_time: Optional[datetime] = Query(None, description="开始时间"),
end_time: Optional[datetime] = Query(None, description="结束时间"),
action: str | None = Query(None, description="按操作类型过滤"),
start_time: datetime | None = Query(None, description="开始时间"),
end_time: datetime | None = Query(None, description="结束时间"),
skip: int = Query(0, ge=0, description="跳过记录数"),
limit: int = Query(100, ge=1, le=1000, description="限制记录数"),
current_user=Depends(get_current_metadata_user),
audit_repo: AuditRepository = Depends(get_audit_repository),
) -> List[AuditLogResponse]:
"""
查询当前用户的审计日志
普通用户只能查看自己的操作记录
"""
logs = await audit_repo.get_logs(
) -> list[AuditLogResponse]:
items = await audit_repo.get_logs(
user_id=current_user.id,
action=action,
start_time=start_time,
@@ -122,4 +137,10 @@ async def get_my_audit_logs(
skip=skip,
limit=limit,
)
return logs
total = await audit_repo.get_log_count(
user_id=current_user.id,
action=action,
start_time=start_time,
end_time=end_time,
)
return PaginatedList(items, total=total)
-190
View File
@@ -1,190 +0,0 @@
from typing import Annotated
from datetime import timedelta
from fastapi import APIRouter, Depends, HTTPException, status
from fastapi.security import OAuth2PasswordRequestForm
from app.core.config import settings
from app.core.security import create_access_token, create_refresh_token, verify_password
from app.domain.schemas.user import UserCreate, UserResponse, UserLogin, Token
from app.infra.db.metadb.repositories.user_repository import UserRepository
from app.auth.dependencies import get_user_repository, get_current_active_user
from app.domain.schemas.user import UserInDB
import logging
logger = logging.getLogger(__name__)
router = APIRouter()
@router.post(
"/register", response_model=UserResponse, status_code=status.HTTP_201_CREATED
)
async def register(
user_data: UserCreate, user_repo: UserRepository = Depends(get_user_repository)
) -> UserResponse:
"""
用户注册
创建新用户账号
"""
# 检查用户名和邮箱是否已存在
if await user_repo.user_exists(username=user_data.username):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Username already registered",
)
if await user_repo.user_exists(email=user_data.email):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail="Email already registered"
)
# 创建用户
try:
user = await user_repo.create_user(user_data)
if not user:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to create user",
)
return UserResponse.model_validate(user)
except Exception as e:
logger.error(f"Error during user registration: {e}")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Registration failed",
)
@router.post("/login", response_model=Token)
async def login(
form_data: Annotated[OAuth2PasswordRequestForm, Depends()],
user_repo: UserRepository = Depends(get_user_repository),
) -> Token:
"""
用户登录(OAuth2 标准格式)
返回 JWT Access Token 和 Refresh Token
"""
# 验证用户(支持用户名或邮箱登录)
user = await user_repo.get_user_by_username(form_data.username)
if not user:
# 尝试用邮箱登录
user = await user_repo.get_user_by_email(form_data.username)
if not user or not verify_password(form_data.password, user.hashed_password):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Incorrect username or password",
headers={"WWW-Authenticate": "Bearer"},
)
if not user.is_active:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail="Inactive user account"
)
# 生成 Token
access_token = create_access_token(subject=user.username)
refresh_token = create_refresh_token(subject=user.username)
return Token(
access_token=access_token,
refresh_token=refresh_token,
token_type="bearer",
expires_in=settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
)
@router.post("/login/simple", response_model=Token)
async def login_simple(
username: str,
password: str,
user_repo: UserRepository = Depends(get_user_repository),
) -> Token:
"""
简化版登录接口(保持向后兼容)
直接使用 username 和 password 参数
"""
# 验证用户
user = await user_repo.get_user_by_username(username)
if not user:
user = await user_repo.get_user_by_email(username)
if not user or not verify_password(password, user.hashed_password):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Incorrect username or password",
)
if not user.is_active:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail="Inactive user account"
)
# 生成 Token
access_token = create_access_token(subject=user.username)
refresh_token = create_refresh_token(subject=user.username)
return Token(
access_token=access_token,
refresh_token=refresh_token,
token_type="bearer",
expires_in=settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
)
@router.get("/me", response_model=UserResponse)
async def get_current_user_info(
current_user: UserInDB = Depends(get_current_active_user),
) -> UserResponse:
"""
获取当前登录用户信息
"""
return UserResponse.model_validate(current_user)
@router.post("/refresh", response_model=Token)
async def refresh_token(
refresh_token: str, user_repo: UserRepository = Depends(get_user_repository)
) -> Token:
"""
刷新 Access Token
使用 Refresh Token 获取新的 Access Token
"""
from jose import jwt, JWTError
credentials_exception = HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Could not validate refresh token",
headers={"WWW-Authenticate": "Bearer"},
)
try:
payload = jwt.decode(
refresh_token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]
)
username: str = payload.get("sub")
token_type: str = payload.get("type")
if username is None or token_type != "refresh":
raise credentials_exception
except JWTError:
raise credentials_exception
# 验证用户仍然存在且激活
user = await user_repo.get_user_by_username(username)
if not user or not user.is_active:
raise credentials_exception
# 生成新的 Access Token
new_access_token = create_access_token(subject=user.username)
return Token(
access_token=new_access_token,
refresh_token=refresh_token, # 保持原 refresh token
token_type="bearer",
expires_in=settings.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
)
+20 -68
View File
@@ -1,13 +1,13 @@
from datetime import datetime
from typing import Any
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body
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_detection import (
get_burst_detection_scheme_detail,
list_burst_detection_schemes,
run_burst_detection,
)
@@ -30,17 +30,26 @@ class BurstDetectionRequest(BaseModel):
points_per_day: int = Field(1440, description="每天的数据点数")
mu: int = Field(100, description="异常值检测的参数")
iforest_params: dict[str, Any] | None = Field(None, description="隔离森林算法参数")
target_time: datetime | None = Field(
None,
description="目标侦测时刻;为空时自动使用最近一个完整的监测时刻",
)
sampling_interval_minutes: int | None = Field(
None,
ge=1,
le=1440,
description="采样间隔(分钟);为空时根据压力 SCADA 传输频率自动推断",
)
scada_start: datetime | None = Field(None, description="SCADA数据起始时间")
scada_end: datetime | None = Field(None, description="SCADA数据结束时间")
sensor_nodes: list[str] | None = Field(None, description="传感器节点列表")
scheme_name: str | None = Field(None, description="方案名称")
data_source: str = Field("monitoring", description="数据来源:monitoring(监测)或simulation(模拟)")
simulation_scheme_name: str | None = Field(None, description="模拟方案名称")
simulation_scheme_type: str | None = Field(None, description="模拟方案类型")
simulation_run_id: UUID | None = Field(None, description="分析模拟运行 ID")
@router.post(
"/detect/",
"/burst-detections",
summary="执行爆管检测",
description="基于压力观测数据和其他参数执行爆管检测分析"
)
@@ -65,67 +74,10 @@ async def detect_burst(
HTTPException: 当处理过程中发生错误时
"""
try:
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)
return await run_in_threadpool(
run_burst_detection,
**data.model_dump(),
username=username,
)
except Exception as exc:
raise HTTPException(status_code=400, detail=str(exc))
+15 -71
View File
@@ -2,14 +2,14 @@ from typing import Any
from datetime import datetime
from typing import Literal
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body
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 (
get_burst_location_scheme_detail,
list_burst_location_schemes,
run_burst_location_by_network,
)
@@ -29,16 +29,17 @@ class BurstLocationRequest(BaseModel):
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爆管开始时间")
scada_burst_end: datetime | None = Field(None, description="SCADA爆管结束时间")
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_scheme_name: str | None = Field(None, description="模拟方案名称")
simulation_scheme_type: str | None = Field(None, description="模拟方案类型")
scheme_name: str | None = Field(None, description="爆管定位运行名称")
simulation_run_id: UUID | None = Field(None, description="分析模拟运行 ID")
@router.post(
"/locate/",
"/burst-locations",
summary="执行爆管定位",
description="基于压力和流量数据定位管网中的爆管位置"
)
@@ -63,67 +64,10 @@ async def locate_burst(
HTTPException: 当数据类型或值不正确时
"""
try:
return run_burst_location_by_network(**data.model_dump(), username=username)
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))
@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))
-57
View File
@@ -1,57 +0,0 @@
from fastapi import APIRouter, Query
from app.infra.cache.redis_client import redis_client
router = APIRouter()
@router.post("/clearrediskey/", summary="清除单个缓存键", description="根据键名清除单个Redis缓存")
async def fastapi_clear_redis_key(key: str = Query(..., description="缓存键名")):
"""
清除单个缓存键
根据指定的键名删除Redis中对应的缓存
"""
redis_client.delete(key)
return True
@router.post("/clearrediskeys/", summary="清除匹配的缓存键", description="根据模式清除匹配的Redis缓存键")
async def fastapi_clear_redis_keys(keys: str = Query(..., description="缓存键模式(支持通配符)")):
"""
清除匹配的缓存键
根据指定的模式删除Redis中所有匹配的缓存键
"""
# delete keys contains the key
matched_keys = redis_client.keys(f"*{keys}*")
if matched_keys:
redis_client.delete(*matched_keys)
return True
@router.post("/clearallredis/", summary="清除所有缓存", description="清空整个Redis数据库的所有缓存")
async def fastapi_clear_all_redis():
"""
清除所有缓存
清空Redis数据库中的所有缓存键值对
"""
redis_client.flushdb()
return True
@router.get("/queryredis/", summary="查询缓存键列表", description="获取Redis中所有的缓存键")
async def fastapi_query_redis():
"""
查询缓存键列表
获取Redis数据库中所有的缓存键列表
"""
# Helper to decode bytes to str for JSON response if needed,
# but original just returned keys (which might be bytes in redis-py unless decode_responses=True)
# create_redis_client usually sets decode_responses=False by default.
# We will assume user handles bytes or we should decode.
# Original just returned redis_client.keys("*")
keys = redis_client.keys("*")
# Clean output for API
return [k.decode('utf-8') if isinstance(k, bytes) else k for k in keys]
+19 -19
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
get_control,
get_control_schema,
@@ -13,58 +13,58 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/getcontrolschema/", summary="获取控制架构", description="获取网络中控制对象的架构定义")
async def fastapi_get_control_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/control", summary="获取控制架构", description="获取网络中控制对象的架构定义")
def fastapi_get_control_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取控制架构。
返回指定网络中控制对象的属性架构定义。
"""
return get_control_schema(network)
@router.get("/getcontrolproperties/", summary="获取控制属性", description="获取指定网络中的控制属性信息")
async def fastapi_get_control_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
@router.get("/controls/properties", summary="获取控制属性", description="获取指定网络中的控制属性信息")
def fastapi_get_control_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取控制属性。
返回指定网络中的控制对象属性信息。
"""
return get_control(network)
@router.post("/setcontrolproperties/", response_model=None, summary="设置控制属性", description="更新指定网络中的控制属性")
async def fastapi_set_control_properties(
@router.patch("/controls/properties", response_model=None, summary="设置控制属性", description="更新指定网络中的控制属性")
def fastapi_set_control_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置控制属性。
更新指定网络中的控制属性值。
"""
props = await req.json()
props = payload
return set_control(network, ChangeSet(props))
@router.get("/getruleschema/", summary="获取规则架构", description="获取网络中规则对象的架构定义")
async def fastapi_get_rule_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/rule-schemas", summary="获取规则架构", description="获取网络中规则对象的架构定义")
def fastapi_get_rule_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取规则架构。
返回指定网络中规则对象的属性架构定义。
"""
return get_rule_schema(network)
@router.get("/getruleproperties/", summary="获取规则属性", description="获取指定网络中的规则属性信息")
async def fastapi_get_rule_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
@router.get("/rule-properties", summary="获取规则属性", description="获取指定网络中的规则属性信息")
def fastapi_get_rule_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取规则属性。
返回指定网络中的规则对象属性信息。
"""
return get_rule(network)
@router.post("/setruleproperties/", response_model=None, summary="设置规则属性", description="更新指定网络中的规则属性")
async def fastapi_set_rule_properties(
@router.patch("/rule-properties", response_model=None, summary="设置规则属性", description="更新指定网络中的规则属性")
def fastapi_set_rule_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置规则属性。
更新指定网络中的规则属性值。
"""
props = await req.json()
props = payload
return set_rule(network, ChangeSet(props))
+21 -21
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_curve,
delete_curve,
@@ -14,32 +14,32 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/getcurveschema", summary="获取曲线架构", description="获取网络中曲线对象的架构定义")
async def fastapi_get_curve_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/curve", summary="获取曲线架构", description="获取网络中曲线对象的架构定义")
def fastapi_get_curve_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取曲线架构。
返回指定网络中曲线对象的属性架构定义。
"""
return get_curve_schema(network)
@router.post("/addcurve/", response_model=None, summary="添加曲线", description="在网络中添加一条新的曲线")
async def fastapi_add_curve(
@router.post("/curves", response_model=None, summary="添加曲线", description="在网络中添加一条新的曲线")
def fastapi_add_curve(
network: str = Query(..., description="管网名称(或数据库名称)"),
curve: str = Query(..., description="曲线ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""添加曲线。
在指定网络中创建一条新的曲线,并设置其初始属性。
"""
props = await req.json()
props = payload
ps = {
"id": curve,
} | props
return add_curve(network, ChangeSet(ps))
@router.post("/deletecurve/", response_model=None, summary="删除曲线", description="从网络中删除指定的曲线")
async def fastapi_delete_curve(
@router.delete("/curves", response_model=None, summary="删除曲线", description="从网络中删除指定的曲线")
def fastapi_delete_curve(
network: str = Query(..., description="管网名称(或数据库名称)"),
curve: str = Query(..., description="曲线ID")
) -> ChangeSet:
@@ -50,8 +50,8 @@ async def fastapi_delete_curve(
ps = {"id": curve}
return delete_curve(network, ChangeSet(ps))
@router.get("/getcurveproperties/", summary="获取曲线属性", description="获取指定曲线的属性信息")
async def fastapi_get_curve_properties(
@router.get("/curves/properties", summary="获取曲线属性", description="获取指定曲线的属性信息")
def fastapi_get_curve_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
curve: str = Query(..., description="曲线ID")
) -> dict[str, Any]:
@@ -61,30 +61,30 @@ async def fastapi_get_curve_properties(
"""
return get_curve(network, curve)
@router.post("/setcurveproperties/", response_model=None, summary="设置曲线属性", description="更新指定曲线的属性")
async def fastapi_set_curve_properties(
@router.patch("/curves/properties", response_model=None, summary="设置曲线属性", description="更新指定曲线的属性")
def fastapi_set_curve_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
curve: str = Query(..., description="曲线ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置曲线属性。
更新指定曲线的属性值。
"""
props = await req.json()
props = payload
ps = {"id": curve} | props
return set_curve(network, ChangeSet(ps))
@router.get("/getcurves/", summary="获取所有曲线", description="获取网络中的所有曲线列表")
async def fastapi_get_curves(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
@router.get("/curves", summary="获取所有曲线", description="获取网络中的所有曲线列表")
def fastapi_get_curves(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
"""获取所有曲线。
返回指定网络中的所有曲线ID列表。
"""
return get_curves(network)
@router.get("/iscurve/", summary="检查曲线存在性", description="检查指定的曲线是否存在")
async def fastapi_is_curve(
@router.get("/curves/existence", summary="检查曲线存在性", description="检查指定的曲线是否存在")
def fastapi_is_curve(
network: str = Query(..., description="管网名称(或数据库名称)"),
curve: str = Query(..., description="曲线ID")
) -> bool:
+35 -35
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
get_energy,
get_energy_schema,
@@ -19,72 +19,72 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/gettimeschema", summary="获取时间选项架构", description="获取网络中时间选项的架构定义")
async def fastapi_get_time_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/time", summary="获取时间选项架构", description="获取网络中时间选项的架构定义")
def fastapi_get_time_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取时间选项架构。
返回指定网络中时间相关选项的属性架构定义。
"""
return get_time_schema(network)
@router.get("/gettimeproperties/", summary="获取时间选项属性", description="获取指定网络中的时间选项属性信息")
async def fastapi_get_time_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
@router.get("/network-options/time", summary="获取时间选项属性", description="获取指定网络中的时间选项属性信息")
def fastapi_get_time_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取时间选项属性。
返回指定网络中的时间相关选项属性。
"""
return get_time(network)
@router.post("/settimeproperties/", response_model=None, summary="设置时间选项属性", description="更新指定网络中的时间选项属性")
async def fastapi_set_time_properties(
@router.patch("/time-properties", response_model=None, summary="设置时间选项属性", description="更新指定网络中的时间选项属性")
def fastapi_set_time_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置时间选项属性。
更新指定网络中的时间相关选项属性值。
"""
props = await req.json()
props = payload
return set_time(network, ChangeSet(props))
@router.get("/getenergyschema/", summary="获取能耗选项架构", description="获取网络中能耗选项的架构定义")
async def fastapi_get_energy_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/energy", summary="获取能耗选项架构", description="获取网络中能耗选项的架构定义")
def fastapi_get_energy_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取能耗选项架构。
返回指定网络中能耗相关选项的属性架构定义。
"""
return get_energy_schema(network)
@router.get("/getenergyproperties/", summary="获取能耗选项属性", description="获取指定网络中的能耗选项属性信息")
async def fastapi_get_energy_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
@router.get("/network-options/energy", summary="获取能耗选项属性", description="获取指定网络中的能耗选项属性信息")
def fastapi_get_energy_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取能耗选项属性。
返回指定网络中的能耗相关选项属性。
"""
return get_energy(network)
@router.post("/setenergyproperties/", response_model=None, summary="设置能耗选项属性", description="更新指定网络中的能耗选项属性")
async def fastapi_set_energy_properties(
@router.patch("/energy-properties", response_model=None, summary="设置能耗选项属性", description="更新指定网络中的能耗选项属性")
def fastapi_set_energy_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置能耗选项属性。
更新指定网络中的能耗相关选项属性值。
"""
props = await req.json()
props = payload
return set_energy(network, ChangeSet(props))
@router.get("/getpumpenergyschema/", summary="获取泵能耗选项架构", description="获取网络中泵能耗选项的架构定义")
async def fastapi_get_pump_energy_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/pump-energy", summary="获取泵能耗选项架构", description="获取网络中泵能耗选项的架构定义")
def fastapi_get_pump_energy_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取泵能耗选项架构。
返回指定网络中泵能耗相关选项的属性架构定义。
"""
return get_pump_energy_schema(network)
@router.get("/getpumpenergyproperties//", summary="获取泵能耗属性", description="获取指定泵的能耗属性信息")
async def fastapi_get_pump_energy_proeprties(
@router.get("/network-options/pump-energy", summary="获取泵能耗属性", description="获取指定泵的能耗属性信息")
def fastapi_get_pump_energy_proeprties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="泵ID")
) -> dict[str, Any]:
@@ -94,44 +94,44 @@ async def fastapi_get_pump_energy_proeprties(
"""
return get_pump_energy(network, pump)
@router.get("/setpumpenergyproperties//", response_model=None, summary="设置泵能耗属性", description="更新指定泵的能耗属性")
async def fastapi_set_pump_energy_properties(
@router.patch("/network-options/pump-energy", response_model=None, summary="设置泵能耗属性", description="更新指定泵的能耗属性")
def fastapi_set_pump_energy_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="泵ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置泵能耗属性。
更新指定泵的能耗相关属性值。
"""
props = await req.json()
props = payload
ps = {"id": pump} | props
return set_pump_energy(network, ChangeSet(ps))
@router.get("/getoptionschema/", summary="获取选项架构", description="获取网络中选项对象的架构定义")
async def fastapi_get_option_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/option", summary="获取选项架构", description="获取网络中选项对象的架构定义")
def fastapi_get_option_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取选项架构。
返回指定网络中选项对象的属性架构定义。
"""
return get_option_v3_schema(network)
@router.get("/getoptionproperties/", summary="获取选项属性", description="获取指定网络中的选项属性信息")
async def fastapi_get_option_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
@router.get("/network-options", summary="获取选项属性", description="获取指定网络中的选项属性信息")
def fastapi_get_option_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取选项属性。
返回指定网络中的选项对象属性信息。
"""
return get_option_v3(network)
@router.post("/setoptionproperties/", response_model=None, summary="设置选项属性", description="更新指定网络中的选项属性")
async def fastapi_set_option_properties(
@router.patch("/network-options", response_model=None, summary="设置选项属性", description="更新指定网络中的选项属性")
def fastapi_set_option_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置选项属性。
更新指定网络中的选项属性值。
"""
props = await req.json()
props = payload
return set_option_v3(network, ChangeSet(props))
+21 -21
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_pattern,
delete_pattern,
@@ -14,32 +14,32 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/getpatternschema", summary="获取模式架构", description="获取网络中模式对象的架构定义")
async def fastapi_get_pattern_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/pattern", summary="获取模式架构", description="获取网络中模式对象的架构定义")
def fastapi_get_pattern_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取模式架构。
返回指定网络中模式对象的属性架构定义。
"""
return get_pattern_schema(network)
@router.post("/addpattern/", response_model=None, summary="添加模式", description="在网络中添加一个新的模式")
async def fastapi_add_pattern(
@router.post("/patterns", response_model=None, summary="添加模式", description="在网络中添加一个新的模式")
def fastapi_add_pattern(
network: str = Query(..., description="管网名称(或数据库名称)"),
pattern: str = Query(..., description="模式ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""添加模式。
在指定网络中创建一个新的模式,并设置其初始属性。
"""
props = await req.json()
props = payload
ps = {
"id": pattern,
} | props
return add_pattern(network, ChangeSet(ps))
@router.post("/deletepattern/", response_model=None, summary="删除模式", description="从网络中删除指定的模式")
async def fastapi_delete_pattern(
@router.delete("/patterns", response_model=None, summary="删除模式", description="从网络中删除指定的模式")
def fastapi_delete_pattern(
network: str = Query(..., description="管网名称(或数据库名称)"),
pattern: str = Query(..., description="模式ID")
) -> ChangeSet:
@@ -50,8 +50,8 @@ async def fastapi_delete_pattern(
ps = {"id": pattern}
return delete_pattern(network, ChangeSet(ps))
@router.get("/getpatternproperties/", summary="获取模式属性", description="获取指定模式的属性信息")
async def fastapi_get_pattern_properties(
@router.get("/patterns/properties", summary="获取模式属性", description="获取指定模式的属性信息")
def fastapi_get_pattern_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pattern: str = Query(..., description="模式ID")
) -> dict[str, Any]:
@@ -61,22 +61,22 @@ async def fastapi_get_pattern_properties(
"""
return get_pattern(network, pattern)
@router.post("/setpatternproperties/", response_model=None, summary="设置模式属性", description="更新指定模式的属性")
async def fastapi_set_pattern_properties(
@router.patch("/patterns/properties", response_model=None, summary="设置模式属性", description="更新指定模式的属性")
def fastapi_set_pattern_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pattern: str = Query(..., description="模式ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置模式属性。
更新指定模式的属性值。
"""
props = await req.json()
props = payload
ps = {"id": pattern} | props
return set_pattern(network, ChangeSet(ps))
@router.get("/ispattern/", summary="检查模式存在性", description="检查指定的模式是否存在")
async def fastapi_is_pattern(
@router.get("/patterns/existence", summary="检查模式存在性", description="检查指定的模式是否存在")
def fastapi_is_pattern(
network: str = Query(..., description="管网名称(或数据库名称)"),
pattern: str = Query(..., description="模式ID")
) -> bool:
@@ -86,8 +86,8 @@ async def fastapi_is_pattern(
"""
return is_pattern(network, pattern)
@router.get("/getpatterns/", summary="获取所有模式", description="获取网络中的所有模式列表")
async def fastapi_get_patterns(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
@router.get("/patterns", summary="获取所有模式", description="获取网络中的所有模式列表")
def fastapi_get_patterns(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
"""获取所有模式。
返回指定网络中的所有模式ID列表。
+75 -75
View File
@@ -1,11 +1,10 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_mixing,
add_source,
api,
delete_mixing,
delete_source,
get_emitter,
@@ -23,6 +22,7 @@ from app.services.tjnetwork import (
get_tank_reaction,
get_tank_reaction_schema,
set_emitter,
set_mixing,
set_pipe_reaction,
set_quality,
set_reaction,
@@ -32,16 +32,16 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/getqualityschema/", summary="获取水质架构", description="获取网络中水质对象的架构定义")
async def fastapi_get_quality_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/quality", summary="获取水质架构", description="获取网络中水质对象的架构定义")
def fastapi_get_quality_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取水质架构。
返回指定网络中水质对象的属性架构定义。
"""
return get_quality_schema(network)
@router.get("/getqualityproperties/", summary="获取水质属性", description="获取指定节点的水质属性信息")
async def fastapi_get_quality_properties(
@router.get("/quality-configurations/properties", summary="获取水质属性", description="获取指定节点的水质属性信息")
def fastapi_get_quality_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> dict[str, Any]:
@@ -51,28 +51,28 @@ async def fastapi_get_quality_properties(
"""
return get_quality(network, node)
@router.post("/setqualityproperties/", response_model=None, summary="设置水质属性", description="更新指定节点的水质属性")
async def fastapi_set_quality_properties(
@router.patch("/quality-configurations/properties", response_model=None, summary="设置水质属性", description="更新指定节点的水质属性")
def fastapi_set_quality_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置水质属性。
更新指定节点的水质属性值。
"""
props = await req.json()
props = payload
return set_quality(network, ChangeSet(props))
@router.get("/getemitterschema", summary="获取发射器架构", description="获取网络中发射器对象的架构定义")
async def fastapi_get_emitter_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/emitter", summary="获取发射器架构", description="获取网络中发射器对象的架构定义")
def fastapi_get_emitter_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取发射器架构。
返回指定网络中发射器对象的属性架构定义。
"""
return get_emitter_schema(network)
@router.get("/getemitterproperties/", summary="获取发射器属性", description="获取指定连接点的发射器属性信息")
async def fastapi_get_emitter_properties(
@router.get("/emitters/properties", summary="获取发射器属性", description="获取指定连接点的发射器属性信息")
def fastapi_get_emitter_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="连接点ID")
) -> dict[str, Any]:
@@ -82,30 +82,30 @@ async def fastapi_get_emitter_properties(
"""
return get_emitter(network, junction)
@router.post("/setemitterproperties/", response_model=None, summary="设置发射器属性", description="更新指定连接点的发射器属性")
async def fastapi_set_emitter_properties(
@router.patch("/emitters/properties", response_model=None, summary="设置发射器属性", description="更新指定连接点的发射器属性")
def fastapi_set_emitter_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="连接点ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置发射器属性。
更新指定连接点的发射器属性值。
"""
props = await req.json()
props = payload
ps = {"junction": junction} | props
return set_emitter(network, ChangeSet(ps))
@router.get("/getsourcechema/", summary="获取水源架构", description="获取网络中水源对象的架构定义")
async def fastapi_get_source_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/source", summary="获取水源架构", description="获取网络中水源对象的架构定义")
def fastapi_get_source_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取水源架构。
返回指定网络中水源对象的属性架构定义。
"""
return get_source_schema(network)
@router.get("/getsource/", summary="获取水源属性", description="获取指定节点的水源属性信息")
async def fastapi_get_source(
@router.get("/sources/detail", summary="获取水源属性", description="获取指定节点的水源属性信息")
def fastapi_get_source(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> dict[str, Any]:
@@ -115,32 +115,32 @@ async def fastapi_get_source(
"""
return get_source(network, node)
@router.post("/setsource/", response_model=None, summary="设置水源属性", description="更新指定节点的水源属性")
async def fastapi_set_source(
@router.patch("/sources", response_model=None, summary="设置水源属性", description="更新指定节点的水源属性")
def fastapi_set_source(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置水源属性。
更新指定节点的水源属性值。
"""
props = await req.json()
props = payload
return set_source(network, ChangeSet(props))
@router.post("/addsource/", response_model=None, summary="添加水源", description="在网络中添加一个新的水源")
async def fastapi_add_source(
@router.post("/sources", response_model=None, summary="添加水源", description="在网络中添加一个新的水源")
def fastapi_add_source(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""添加水源。
在指定网络中创建一个新的水源,并设置其初始属性。
"""
props = await req.json()
props = payload
return add_source(network, ChangeSet(props))
@router.post("/deletesource/", response_model=None, summary="删除水源", description="从网络中删除指定节点的水源")
async def fastapi_delete_source(
@router.delete("/sources", response_model=None, summary="删除水源", description="从网络中删除指定节点的水源")
def fastapi_delete_source(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> ChangeSet:
@@ -151,44 +151,44 @@ async def fastapi_delete_source(
props = {"node": node}
return delete_source(network, ChangeSet(props))
@router.get("/getreactionschema/", summary="获取反应架构", description="获取网络中反应对象的架构定义")
async def fastapi_get_reaction_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/reaction", summary="获取反应架构", description="获取网络中反应对象的架构定义")
def fastapi_get_reaction_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取反应架构。
返回指定网络中反应对象的属性架构定义。
"""
return get_reaction_schema(network)
@router.get("/getreaction/", summary="获取反应属性", description="获取指定网络中的反应属性信息")
async def fastapi_get_reaction(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
@router.get("/reactions/detail", summary="获取反应属性", description="获取指定网络中的反应属性信息")
def fastapi_get_reaction(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取反应属性。
返回指定网络中的反应属性信息。
"""
return get_reaction(network)
@router.post("/setreaction/", response_model=None, summary="设置反应属性", description="更新指定网络中的反应属性")
async def fastapi_set_reaction(
@router.patch("/reactions", response_model=None, summary="设置反应属性", description="更新指定网络中的反应属性")
def fastapi_set_reaction(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置反应属性。
更新指定网络中的反应属性值。
"""
props = await req.json()
props = payload
return set_reaction(network, ChangeSet(props))
@router.get("/getpipereactionschema/", summary="获取管道反应架构", description="获取网络中管道反应对象的架构定义")
async def fastapi_get_pipe_reaction_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/pipe-reaction", summary="获取管道反应架构", description="获取网络中管道反应对象的架构定义")
def fastapi_get_pipe_reaction_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取管道反应架构。
返回指定网络中管道反应对象的属性架构定义。
"""
return get_pipe_reaction_schema(network)
@router.get("/getpipereaction/", summary="获取管道反应属性", description="获取指定管道的反应属性信息")
async def fastapi_get_pipe_reaction(
@router.get("/pipe-reactions/detail", summary="获取管道反应属性", description="获取指定管道的反应属性信息")
def fastapi_get_pipe_reaction(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> dict[str, Any]:
@@ -198,28 +198,28 @@ async def fastapi_get_pipe_reaction(
"""
return get_pipe_reaction(network, pipe)
@router.post("/setpipereaction/", response_model=None, summary="设置管道反应属性", description="更新指定管道的反应属性")
async def fastapi_set_pipe_reaction(
@router.patch("/pipe-reactions", response_model=None, summary="设置管道反应属性", description="更新指定管道的反应属性")
def fastapi_set_pipe_reaction(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置管道反应属性。
更新指定管道的反应属性值。
"""
props = await req.json()
props = payload
return set_pipe_reaction(network, ChangeSet(props))
@router.get("/gettankreactionschema/", summary="获取水池反应架构", description="获取网络中水池反应对象的架构定义")
async def fastapi_get_tank_reaction_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/tank-reaction", summary="获取水池反应架构", description="获取网络中水池反应对象的架构定义")
def fastapi_get_tank_reaction_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取水池反应架构。
返回指定网络中水池反应对象的属性架构定义。
"""
return get_tank_reaction_schema(network)
@router.get("/gettankreaction/", summary="获取水池反应属性", description="获取指定水池的反应属性信息")
async def fastapi_get_tank_reaction(
@router.get("/tank-reactions/detail", summary="获取水池反应属性", description="获取指定水池的反应属性信息")
def fastapi_get_tank_reaction(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水池ID")
) -> dict[str, Any]:
@@ -229,28 +229,28 @@ async def fastapi_get_tank_reaction(
"""
return get_tank_reaction(network, tank)
@router.post("/settankreaction/", response_model=None, summary="设置水池反应属性", description="更新指定水池的反应属性")
async def fastapi_set_tank_reaction(
@router.patch("/tank-reactions", response_model=None, summary="设置水池反应属性", description="更新指定水池的反应属性")
def fastapi_set_tank_reaction(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置水池反应属性。
更新指定水池的反应属性值。
"""
props = await req.json()
props = payload
return set_tank_reaction(network, ChangeSet(props))
@router.get("/getmixingschema/", summary="获取混合架构", description="获取网络中混合对象的架构定义")
async def fastapi_get_mixing_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/mixing", summary="获取混合架构", description="获取网络中混合对象的架构定义")
def fastapi_get_mixing_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取混合架构。
返回指定网络中混合对象的属性架构定义。
"""
return get_mixing_schema(network)
@router.get("/getmixing/", summary="获取混合属性", description="获取指定水池的混合属性信息")
async def fastapi_get_mixing(
@router.get("/mixing-configurations/detail", summary="获取混合属性", description="获取指定水池的混合属性信息")
def fastapi_get_mixing(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水池ID")
) -> dict[str, Any]:
@@ -260,38 +260,38 @@ async def fastapi_get_mixing(
"""
return get_mixing(network, tank)
@router.post("/setmixing/", response_model=None, summary="设置混合属性", description="更新指定水池的混合属性")
async def fastapi_set_mixing(
@router.patch("/mixing-configurations", response_model=None, summary="设置混合属性", description="更新指定水池的混合属性")
def fastapi_set_mixing(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置混合属性。
更新指定水池的混合属性值。
"""
props = await req.json()
return api.set_mixing(network, ChangeSet(props))
props = payload
return set_mixing(network, ChangeSet(props))
@router.post("/addmixing/", response_model=None, summary="添加混合", description="在网络中添加一个新的混合")
async def fastapi_add_mixing(
@router.post("/mixing-configurations", response_model=None, summary="添加混合", description="在网络中添加一个新的混合")
def fastapi_add_mixing(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""添加混合。
在指定网络中创建一个新的混合,并设置其初始属性。
"""
props = await req.json()
props = payload
return add_mixing(network, ChangeSet(props))
@router.post("/deletemixing/", response_model=None, summary="删除混合", description="从网络中删除指定的混合")
async def fastapi_delete_mixing(
@router.delete("/mixing-configurations", response_model=None, summary="删除混合", description="从网络中删除指定的混合")
def fastapi_delete_mixing(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""删除混合。
从指定网络中删除指定的混合及其相关数据。
"""
props = await req.json()
props = payload
return delete_mixing(network, ChangeSet(props))
+47 -47
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body, Response
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_label,
add_vertex,
@@ -24,16 +24,16 @@ import json
router = APIRouter()
@router.get("/getvertexschema/", summary="获取图形元素架构", description="获取网络中图形元素对象的架构定义")
async def fastapi_get_vertex_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/vertex", summary="获取图形元素架构", description="获取网络中图形元素对象的架构定义")
def fastapi_get_vertex_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取图形元素架构。
返回指定网络中图形元素对象的属性架构定义。
"""
return get_vertex_schema(network)
@router.get("/getvertexproperties/", summary="获取图形元素属性", description="获取指定图形元素的属性信息")
async def fastapi_get_vertex_properties(
@router.get("/visual-elements/properties", summary="获取图形元素属性", description="获取指定图形元素的属性信息")
def fastapi_get_vertex_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="图形元素链接")
) -> dict[str, Any]:
@@ -43,68 +43,68 @@ async def fastapi_get_vertex_properties(
"""
return get_vertex(network, link)
@router.post("/setvertexproperties/", response_model=None, summary="设置图形元素属性", description="更新指定图形元素的属性")
async def fastapi_set_vertex_properties(
@router.patch("/visual-elements/properties", response_model=None, summary="设置图形元素属性", description="更新指定图形元素的属性")
def fastapi_set_vertex_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置图形元素属性。
更新指定图形元素的属性值。
"""
props = await req.json()
props = payload
return set_vertex(network, ChangeSet(props))
@router.post("/addvertex/", response_model=None, summary="添加图形元素", description="在网络中添加一个新的图形元素")
async def fastapi_add_vertex(
@router.post("/visual-elements", response_model=None, summary="添加图形元素", description="在网络中添加一个新的图形元素")
def fastapi_add_vertex(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""添加图形元素。
在指定网络中创建一个新的图形元素,并设置其初始属性。
"""
props = await req.json()
props = payload
return add_vertex(network, ChangeSet(props))
@router.post("/deletevertex/", response_model=None, summary="删除图形元素", description="从网络中删除指定的图形元素")
async def fastapi_delete_vertex(
@router.delete("/visual-elements", response_model=None, summary="删除图形元素", description="从网络中删除指定的图形元素")
def fastapi_delete_vertex(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""删除图形元素。
从指定网络中删除指定的图形元素及其相关数据。
"""
props = await req.json()
props = payload
return delete_vertex(network, ChangeSet(props))
@router.get("/getallvertexlinks/", response_class=PlainTextResponse, summary="获取所有图形元素链接", description="获取网络中的所有图形元素链接列表")
async def fastapi_get_all_vertex_links(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
@router.get("/visual-elements/links", response_class=PlainTextResponse, summary="获取所有图形元素链接", description="获取网络中的所有图形元素链接列表")
def fastapi_get_all_vertex_links(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
"""获取所有图形元素链接。
返回指定网络中的所有图形元素链接列表。
"""
return json.dumps(get_all_vertex_links(network))
@router.get("/getallvertices/", response_class=PlainTextResponse, summary="获取所有图形元素", description="获取网络中的所有图形元素详细信息")
async def fastapi_get_all_vertices(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[dict[str, Any]]:
@router.get("/all-vertices", response_class=PlainTextResponse, summary="获取所有图形元素", description="获取网络中的所有图形元素详细信息")
def fastapi_get_all_vertices(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[dict[str, Any]]:
"""获取所有图形元素。
返回指定网络中的所有图形元素详细信息。
"""
return json.dumps(get_all_vertices(network))
@router.get("/getlabelschema/", summary="获取标签架构", description="获取网络中标签对象的架构定义")
async def fastapi_get_label_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/label", summary="获取标签架构", description="获取网络中标签对象的架构定义")
def fastapi_get_label_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取标签架构。
返回指定网络中标签对象的属性架构定义。
"""
return get_label_schema(network)
@router.get("/getlabelproperties/", summary="获取标签属性", description="获取指定坐标处的标签属性信息")
async def fastapi_get_label_properties(
@router.get("/labels/properties", summary="获取标签属性", description="获取指定坐标处的标签属性信息")
def fastapi_get_label_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
x: float = Query(..., description="X坐标"),
y: float = Query(..., description="Y坐标")
@@ -115,66 +115,66 @@ async def fastapi_get_label_properties(
"""
return get_label(network, x, y)
@router.post("/setlabelproperties/", response_model=None, summary="设置标签属性", description="更新指定标签的属性")
async def fastapi_set_label_properties(
@router.patch("/labels/properties", response_model=None, summary="设置标签属性", description="更新指定标签的属性")
def fastapi_set_label_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置标签属性。
更新指定标签的属性值。
"""
props = await req.json()
props = payload
return set_label(network, ChangeSet(props))
@router.post("/addlabel/", response_model=None, summary="添加标签", description="在网络中添加一个新的标签")
async def fastapi_add_label(
@router.post("/labels", response_model=None, summary="添加标签", description="在网络中添加一个新的标签")
def fastapi_add_label(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""添加标签。
在指定网络中创建一个新的标签,并设置其初始属性。
"""
props = await req.json()
props = payload
return add_label(network, ChangeSet(props))
@router.post("/deletelabel/", response_model=None, summary="删除标签", description="从网络中删除指定的标签")
async def fastapi_delete_label(
@router.delete("/labels", response_model=None, summary="删除标签", description="从网络中删除指定的标签")
def fastapi_delete_label(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""删除标签。
从指定网络中删除指定的标签及其相关数据。
"""
props = await req.json()
props = payload
return delete_label(network, ChangeSet(props))
@router.get("/getbackdropschema/", summary="获取背景架构", description="获取网络中背景对象的架构定义")
async def fastapi_get_backdrop_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/backdrop", summary="获取背景架构", description="获取网络中背景对象的架构定义")
def fastapi_get_backdrop_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""获取背景架构。
返回指定网络中背景对象的属性架构定义。
"""
return get_backdrop_schema(network)
@router.get("/getbackdropproperties/", summary="获取背景属性", description="获取指定网络的背景属性信息")
async def fastapi_get_backdrop_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
@router.get("/backdrops/properties", summary="获取背景属性", description="获取指定网络的背景属性信息")
def fastapi_get_backdrop_properties(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取背景属性。
返回指定网络的背景属性信息。
"""
return get_backdrop(network)
@router.post("/setbackdropproperties/", response_model=None, summary="设置背景属性", description="更新指定网络的背景属性")
async def fastapi_set_backdrop_properties(
@router.patch("/backdrops/properties", response_model=None, summary="设置背景属性", description="更新指定网络的背景属性")
def fastapi_set_backdrop_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置背景属性。
更新指定网络的背景属性值。
"""
props = await req.json()
props = payload
return set_backdrop(network, ChangeSet(props))
-388
View File
@@ -1,388 +0,0 @@
from typing import Any, List, Dict, Optional
import logging
from datetime import datetime, timedelta, timezone, time as dt_time
import msgpack
from fastapi import APIRouter
from pydantic import BaseModel
from py_linq import Enumerable
import app.infra.db.influxdb.api as influxdb_api
import app.services.time_api as time_api
from app.infra.cache.redis_client import redis_client, encode_datetime, decode_datetime
router = APIRouter()
logger = logging.getLogger(__name__)
# Basic Node/Link Latest Record Queries
@router.get("/querynodelatestrecordbyid/")
async def fastapi_query_node_latest_record_by_id(id: str) -> Any:
return influxdb_api.query_latest_record_by_ID(id, type="node")
@router.get("/querylinklatestrecordbyid/")
async def fastapi_query_link_latest_record_by_id(id: str) -> Any:
return influxdb_api.query_latest_record_by_ID(id, type="link")
@router.get("/queryscadalatestrecordbyid/")
async def fastapi_query_scada_latest_record_by_id(id: str) -> Any:
return influxdb_api.query_latest_record_by_ID(id, type="scada")
# Time-based Queries
@router.get("/queryallrecordsbytime/")
async def fastapi_query_all_records_by_time(querytime: str) -> dict[str, list]:
results: tuple = influxdb_api.query_all_records_by_time(query_time=querytime)
return {"nodes": results[0], "links": results[1]}
@router.get("/queryallrecordsbytimeproperty/")
async def fastapi_query_all_record_by_time_property(
querytime: str, type: str, property: str, bucket: str = "realtime_simulation_result"
) -> dict[str, list]:
results: tuple = influxdb_api.query_all_record_by_time_property(
query_time=querytime, type=type, property=property, bucket=bucket
)
return {"results": results}
@router.get("/queryallschemerecordsbytimeproperty/")
async def fastapi_query_all_scheme_record_by_time_property(
querytime: str,
type: str,
property: str,
schemename: str,
bucket: str = "scheme_simulation_result",
) -> dict[str, list]:
"""
查询指定方案某一时刻的所有记录,查询 'node''link' 的某一属性值
"""
results: list = influxdb_api.query_all_scheme_record_by_time_property(
query_time=querytime,
type=type,
property=property,
scheme_name=schemename,
bucket=bucket,
)
return {"results": results}
@router.get("/querysimulationrecordsbyidtime/")
async def fastapi_query_simulation_record_by_ids_time(
id: str, querytime: str, type: str, bucket: str = "realtime_simulation_result"
) -> dict[str, list]:
results: tuple = influxdb_api.query_simulation_result_by_ID_time(
ID=id, type=type, query_time=querytime, bucket=bucket
)
return {"results": results}
@router.get("/queryschemesimulationrecordsbyidtime/")
async def fastapi_query_scheme_simulation_record_by_ids_time(
scheme_name: str,
id: str,
querytime: str,
type: str,
bucket: str = "scheme_simulation_result",
) -> dict[str, list]:
results: tuple = influxdb_api.query_scheme_simulation_result_by_ID_time(
scheme_name=scheme_name, ID=id, type=type, query_time=querytime, bucket=bucket
)
return {"results": results}
# Date-based Queries with Caching
@router.get("/queryallrecordsbydate/")
async def fastapi_query_all_records_by_date(querydate: str) -> dict:
is_today_or_future = time_api.is_today_or_future(querydate)
logger.info(f"isToday or future: {is_today_or_future}")
cache_key = f"queryallrecordsbydate_{querydate}"
if not is_today_or_future:
data = redis_client.get(cache_key)
if data:
results = msgpack.unpackb(data, object_hook=decode_datetime)
logger.info("return from cache redis")
return results
logger.info("query from influxdb")
nodes_links: tuple = influxdb_api.query_all_records_by_date(query_date=querydate)
results = {"nodes": nodes_links[0], "links": nodes_links[1]}
if not is_today_or_future:
logger.info("save to cache redis")
redis_client.set(cache_key, msgpack.packb(results, default=encode_datetime))
logger.info("return results")
return results
@router.get("/queryallrecordsbytimerange/")
async def fastapi_query_all_records_by_time_range(
starttime: str, endtime: str
) -> dict[str, list]:
cache_key = f"queryallrecordsbytimerange_{starttime}_{endtime}"
if not time_api.is_today_or_future(starttime):
data = redis_client.get(cache_key)
if data:
loaded_dict = msgpack.unpackb(data, object_hook=decode_datetime)
return loaded_dict
nodes_links: tuple = influxdb_api.query_all_records_by_time_range(
starttime=starttime, endtime=endtime
)
results = {"nodes": nodes_links[0], "links": nodes_links[1]}
if not time_api.is_today_or_future(starttime):
redis_client.set(cache_key, msgpack.packb(results, default=encode_datetime))
return results
@router.get("/queryallrecordsbydatewithtype/")
async def fastapi_query_all_records_by_date_with_type(
querydate: str, querytype: str
) -> list:
cache_key = f"queryallrecordsbydatewithtype_{querydate}_{querytype}"
data = redis_client.get(cache_key)
if data:
loaded_dict = msgpack.unpackb(data, object_hook=decode_datetime)
return loaded_dict
results = influxdb_api.query_all_records_by_date_with_type(
query_date=querydate, query_type=querytype
)
packed = msgpack.packb(results, default=encode_datetime)
redis_client.set(cache_key, packed)
return results
@router.get("/queryallrecordsbyidsdatetype/")
async def fastapi_query_all_records_by_ids_date_type(
ids: str, querydate: str, querytype: str
) -> list:
cache_key = f"queryallrecordsbydatewithtype_{querydate}_{querytype}"
data = redis_client.get(cache_key)
results = []
if data:
results = msgpack.unpackb(data, object_hook=decode_datetime)
else:
results = influxdb_api.query_all_records_by_date_with_type(
query_date=querydate, query_type=querytype
)
packed = msgpack.packb(results, default=encode_datetime)
redis_client.set(cache_key, packed)
query_ids = ids.split(",")
# Using Enumerable from py_linq as in original code
e_results = Enumerable(results)
lst_results = e_results.where(lambda x: x["ID"] in query_ids).to_list()
return lst_results
@router.get("/queryallrecordsbydateproperty/")
async def fastapi_query_all_records_by_date_property(
querydate: str, querytype: str, property: str
) -> list[dict]:
cache_key = f"queryallrecordsbydateproperty_{querydate}_{querytype}_{property}"
data = redis_client.get(cache_key)
if data:
loaded_dict = msgpack.unpackb(data, object_hook=decode_datetime)
return loaded_dict
result_dict = influxdb_api.query_all_record_by_date_property(
query_date=querydate, type=querytype, property=property
)
packed = msgpack.packb(result_dict, default=encode_datetime)
redis_client.set(cache_key, packed)
return result_dict
# Curve Queries
@router.get("/querynodecurvebyidpropertydaterange/")
async def fastapi_query_node_curve_by_id_property_daterange(
id: str, prop: str, startdate: str, enddate: str
):
return influxdb_api.query_curve_by_ID_property_daterange(
id, type="node", property=prop, start_date=startdate, end_date=enddate
)
@router.get("/querylinkcurvebyidpropertydaterange/")
async def fastapi_query_link_curve_by_id_property_daterange(
id: str, prop: str, startdate: str, enddate: str
):
return influxdb_api.query_curve_by_ID_property_daterange(
id, type="link", property=prop, start_date=startdate, end_date=enddate
)
# SCADA Data Queries
@router.get("/queryscadadatabydeviceidandtime/")
async def fastapi_query_scada_data_by_device_id_and_time(ids: str, querytime: str):
query_ids = ids.split(",")
logger.info(querytime)
return influxdb_api.query_SCADA_data_by_device_ID_and_time(
query_ids_list=query_ids, query_time=querytime
)
@router.get("/queryscadadatabydeviceidandtimerange/")
async def fastapi_query_scada_data_by_device_id_and_time_range(
ids: str, starttime: str, endtime: str
):
print(f"query_ids: {ids}, starttime: {starttime}, endtime: {endtime}")
query_ids = ids.split(",")
return influxdb_api.query_SCADA_data_by_device_ID_and_timerange(
query_ids_list=query_ids, start_time=starttime, end_time=endtime
)
@router.get("/queryfillingscadadatabydeviceidandtimerange/")
async def fastapi_query_filling_scada_data_by_device_id_and_time_range(
ids: str, starttime: str, endtime: str
):
print(f"query_ids: {ids}, starttime: {starttime}, endtime: {endtime}")
query_ids = ids.split(",")
return influxdb_api.query_filling_SCADA_data_by_device_ID_and_timerange(
query_ids_list=query_ids, start_time=starttime, end_time=endtime
)
@router.get("/querycleaningscadadatabydeviceidandtimerange/")
async def fastapi_query_cleaning_scada_data_by_device_id_and_time_range(
ids: str, starttime: str, endtime: str
):
print(f"query_ids: {ids}, starttime: {starttime}, endtime: {endtime}")
query_ids = ids.split(",")
return influxdb_api.query_cleaning_SCADA_data_by_device_ID_and_timerange(
query_ids_list=query_ids, start_time=starttime, end_time=endtime
)
@router.get("/querysimulationscadadatabydeviceidandtimerange/")
async def fastapi_query_simulation_scada_data_by_device_id_and_time_range(
ids: str, starttime: str, endtime: str
):
print(f"query_ids: {ids}, starttime: {starttime}, endtime: {endtime}")
query_ids = ids.split(",")
return influxdb_api.query_simulation_SCADA_data_by_device_ID_and_timerange(
query_ids_list=query_ids, start_time=starttime, end_time=endtime
)
@router.get("/querycleanedscadadatabydeviceidandtimerange/")
async def fastapi_query_cleaned_scada_data_by_device_id_and_time_range(
ids: str, starttime: str, endtime: str
):
print(f"query_ids: {ids}, starttime: {starttime}, endtime: {endtime}")
query_ids = ids.split(",")
return influxdb_api.query_cleaned_SCADA_data_by_device_ID_and_timerange(
query_ids_list=query_ids, start_time=starttime, end_time=endtime
)
@router.get("/queryscadadatabydeviceidanddate/")
async def fastapi_query_scada_data_by_device_id_and_date(ids: str, querydate: str):
query_ids = ids.split(",")
return influxdb_api.query_SCADA_data_by_device_ID_and_date(
query_ids_list=query_ids, query_date=querydate
)
@router.get("/queryallscadarecordsbydate/")
async def fastapi_query_all_scada_records_by_date(querydate: str):
is_today_or_future = time_api.is_today_or_future(querydate)
logger.info(f"isToday or future: {is_today_or_future}")
cache_key = f"queryallscadarecordsbydate_{querydate}"
if not is_today_or_future:
data = redis_client.get(cache_key)
if data:
loaded_dict = msgpack.unpackb(data, object_hook=decode_datetime)
logger.info("return from cache redis")
return loaded_dict
logger.info("query from influxdb")
result_dict = influxdb_api.query_all_SCADA_records_by_date(query_date=querydate)
if not is_today_or_future:
logger.info("save to cache redis")
packed = msgpack.packb(result_dict, default=encode_datetime)
redis_client.set(cache_key, packed)
logger.info("return results")
return result_dict
@router.get("/queryallschemeallrecords/")
async def fastapi_query_all_scheme_all_records(
schemetype: str, schemename: str, querydate: str
) -> tuple:
cache_key = f"queryallschemeallrecords_{schemetype}_{schemename}_{querydate}"
data = redis_client.get(cache_key)
if data:
loaded_dict = msgpack.unpackb(data, object_hook=decode_datetime)
return loaded_dict
results = influxdb_api.query_scheme_all_record(
scheme_type=schemetype, scheme_name=schemename, query_date=querydate
)
packed = msgpack.packb(results, default=encode_datetime)
redis_client.set(cache_key, packed)
return results
@router.get("/queryschemeallrecordsproperty/")
async def fastapi_query_all_scheme_all_records_property(
schemetype: str, schemename: str, querydate: str, querytype: str, queryproperty: str
) -> Optional[List]:
cache_key = f"queryallschemeallrecords_{schemetype}_{schemename}_{querydate}"
data = redis_client.get(cache_key)
all_results = None
if data:
all_results = msgpack.unpackb(data, object_hook=decode_datetime)
else:
all_results = influxdb_api.query_scheme_all_record(
scheme_type=schemetype, scheme_name=schemename, query_date=querydate
)
packed = msgpack.packb(all_results, default=encode_datetime)
redis_client.set(cache_key, packed)
results = None
if querytype == "node":
results = all_results[0]
elif querytype == "link":
results = all_results[1]
return results
@router.get("/queryinfluxdbbuckets/")
async def fastapi_query_influxdb_buckets():
return influxdb_api.query_buckets()
@router.get("/queryinfluxdbbucketmeasurements/")
async def fastapi_query_influxdb_bucket_measurements(bucket: str):
return influxdb_api.query_measurements(bucket=bucket)
############################################################
# download history data
############################################################
class Download_History_Data_Manually(BaseModel):
"""
download_date:样式如 datetime(2025, 5, 4)
"""
download_date: datetime
@router.post("/download_history_data_manually/")
async def fastapi_download_history_data_manually(
data: Download_History_Data_Manually,
) -> None:
item = data.dict()
tz = timezone(timedelta(hours=8))
begin_dt = datetime.combine(item.get("download_date").date(), dt_time.min).replace(
tzinfo=tz
)
end_dt = datetime.combine(item.get("download_date").date(), dt_time(23, 59, 59)).replace(
tzinfo=tz
)
begin_time = begin_dt.isoformat()
end_time = end_dt.isoformat()
influxdb_api.download_history_data_manually(
begin_time=begin_time, end_time=end_time
)
-104
View File
@@ -1,104 +0,0 @@
from typing import List, Any
from fastapi import APIRouter, Request, HTTPException, Query, Body
from app.services.tjnetwork import (
ChangeSet,
get_all_extension_data_keys,
get_all_extension_data,
get_extension_data,
set_extension_data
)
router = APIRouter()
@router.get(
"/getallextensiondatakeys/",
summary="获取所有扩展数据键",
description="获取指定网络的所有扩展数据的键列表"
)
async def get_all_extension_data_keys_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[str]:
"""
获取所有扩展数据键。
返回指定网络中所有可用的扩展数据键。
Args:
network: 管网名称(或数据库名称)
Returns:
扩展数据键列表
"""
return get_all_extension_data_keys(network)
@router.get(
"/getallextensiondata/",
summary="获取所有扩展数据",
description="获取指定网络的所有扩展数据"
)
async def get_all_extension_data_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, Any]:
"""
获取所有扩展数据。
返回指定网络的所有扩展数据及其值。
Args:
network: 管网名称(或数据库名称)
Returns:
扩展数据字典
"""
return get_all_extension_data(network)
@router.get(
"/getextensiondata/",
summary="获取指定扩展数据",
description="获取指定网络中指定键的扩展数据值"
)
async def get_extension_data_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
key: str = Query(..., description="扩展数据键")
) -> str | None:
"""
获取指定扩展数据。
返回指定网络中指定键对应的扩展数据值。
Args:
network: 管网名称(或数据库名称)
key: 扩展数据键
Returns:
扩展数据值,如果不存在返回None
"""
return get_extension_data(network, key)
@router.post(
"/setextensiondata/",
response_model=None,
summary="设置扩展数据",
description="设置指定网络中的扩展数据"
)
async def set_extension_data_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
设置扩展数据。
在指定网络中设置扩展数据,并返回变更集信息。
Args:
network: 管网名称(或数据库名称)
req: 包含扩展数据的请求体
Returns:
变更集信息
"""
props = await req.json()
print(props)
cs = set_extension_data(network, ChangeSet(props))
print(cs.operations[0])
return cs
+29
View File
@@ -0,0 +1,29 @@
from typing import Any
from fastapi import APIRouter, HTTPException, status
from app.services.geocoding import (
TiandituGeocodeRequest,
TiandituGeocodingAPIError,
TiandituGeocodingConfigError,
geocode_tianditu,
)
router = APIRouter()
@router.post(
"/geocoding-requests",
summary="Tianditu Geocoding",
description="调用天地图地理编码服务,将结构化地址转换为经纬度",
)
async def tianditu_geocode(request: TiandituGeocodeRequest) -> dict[str, Any]:
try:
return await geocode_tianditu(request)
except TiandituGeocodingConfigError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=str(exc),
) from exc
except TiandituGeocodingAPIError as exc:
raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc
+34 -80
View File
@@ -2,37 +2,50 @@ import os
from typing import Any
from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body
from pydantic import BaseModel, Field
from fastapi import APIRouter, Depends, HTTPException, Body
from pydantic import BaseModel, ConfigDict, Field
from starlette.concurrency import run_in_threadpool
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,
from app.services.dma_leakage_estimation import (
run_leakage_identification,
)
router = APIRouter()
DEFAULT_N_WORKERS = max(1, min((os.cpu_count() or 1) - 1, 4))
MAX_POPULATION_SIZE = 1_000
MAX_GENERATIONS = 1_000
MAX_DURATION_HOURS = 168
class LeakageIdentifyRequest(BaseModel):
"""漏损识别请求模型"""
model_config = ConfigDict(extra="forbid")
network: str = Field(..., description="管网名称(或数据库名称)")
observed_pressure_data: str | dict[str, list[Any]] | list[dict[str, Any]] | None = Field(
None, description="观测的压力数据"
observed_pressure_data: dict[str, list[Any]] | list[dict[str, Any]] | None = (
Field(None, description="观测的压力数据;文件路径不属于公共 API 输入")
)
start_time: float = Field(0, description="起始时间(小时)")
duration: float = Field(24, description="持续时间(小时)")
timestep: float = Field(5, description="时间步长(分钟")
q_sum: float = Field(0.2, description="总流量(m3/s")
start_time: float = Field(0, ge=0, description="起始时间(小时)")
duration: float = Field(
24, gt=0, le=MAX_DURATION_HOURS, description="持续时间(小时"
)
timestep: float = Field(5, gt=0, le=1440, description="时间步长(分钟)")
q_sum: float = Field(0.2, ge=0, description="总流量(m3/s")
q_sum_unit: str = Field("m3/s", description="流量单位")
output_dir: str = Field("db_inp", description="输出目录")
pop_size: int = Field(50, description="种群大小")
max_gen: int = Field(100, description="最大代数")
n_workers: int = Field(DEFAULT_N_WORKERS, description="工作线程")
pop_size: int = Field(
50, ge=2, le=MAX_POPULATION_SIZE, description="种群大小"
)
max_gen: int = Field(100, ge=1, le=MAX_GENERATIONS, description="最大代")
n_workers: int = Field(
DEFAULT_N_WORKERS,
ge=1,
le=DEFAULT_N_WORKERS,
description="工作进程数",
)
output_flow_unit: str = Field("m3/s", description="输出流量单位")
dma_count: int | None = Field(None, description="DMA区域数量")
dma_count: int | None = Field(None, ge=1, description="DMA区域数量")
scada_start: datetime | None = Field(None, description="SCADA数据起始时间")
scada_end: datetime | None = Field(None, description="SCADA数据结束时间")
sensor_nodes: list[str] | None = Field(None, description="传感器节点列表")
@@ -40,7 +53,7 @@ class LeakageIdentifyRequest(BaseModel):
@router.post(
"/identify/",
"/leakage-identifications",
summary="执行漏损识别",
description="基于压力观测数据和遗传算法识别管网中的漏损位置和大小"
)
@@ -65,69 +78,10 @@ async def identify_leakage(
HTTPException: 当处理过程中发生错误时
"""
try:
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
return await run_in_threadpool(
run_leakage_identification,
**data.model_dump(),
username=username,
)
except Exception as exc:
raise HTTPException(status_code=400, detail=str(exc))
+29 -26
View File
@@ -1,21 +1,18 @@
import logging
from fastapi import APIRouter, Depends, HTTPException, status, Query, Path
import psycopg
from psycopg import AsyncConnection
from sqlalchemy import text
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from app.auth.project_dependencies import (
ProjectContext,
get_project_context,
get_project_pg_session,
get_project_pg_connection,
get_project_timescale_connection,
get_metadata_repository,
)
from app.auth.metadata_dependencies import get_current_metadata_user
from app.core.config import settings
from app.domain.schemas.metadata import (
GeoServerConfigResponse,
ProjectMetaResponse,
ProjectSummaryResponse,
)
@@ -25,7 +22,7 @@ router = APIRouter()
logger = logging.getLogger(__name__)
@router.get("/meta/project", summary="获取项目元数据", description="获取当前项目的元数据和配置信息", response_model=ProjectMetaResponse)
@router.get("/projects/current/metadata", summary="获取项目元数据", description="获取当前项目的元数据和配置信息", response_model=ProjectMetaResponse)
async def get_project_metadata(
ctx: ProjectContext = Depends(get_project_context),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
@@ -33,38 +30,26 @@ async def get_project_metadata(
"""
获取项目元数据
返回当前项目的完整元数据,包括项目基本信息和GeoServer配置
返回当前项目的完整元数据,包括项目基本信息和项目权限
"""
project = await metadata_repo.get_project_by_id(ctx.project_id)
if not project:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="Project not found"
)
geoserver = await metadata_repo.get_geoserver_config(ctx.project_id)
geoserver_payload = (
GeoServerConfigResponse(
gs_base_url=geoserver.gs_base_url,
gs_admin_user=geoserver.gs_admin_user,
gs_datastore_name=geoserver.gs_datastore_name,
default_extent=geoserver.default_extent,
srid=geoserver.srid,
)
if geoserver
else None
)
return ProjectMetaResponse(
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,
project_role=ctx.project_role,
geoserver=geoserver_payload,
)
@router.get("/meta/projects", summary="列出用户项目", description="获取当前用户有权限的所有项目列表", response_model=list[ProjectSummaryResponse])
@router.get("/projects", summary="列出用户项目", description="获取当前用户有权限的所有项目列表", response_model=list[ProjectSummaryResponse])
async def list_user_projects(
current_user=Depends(get_current_metadata_user),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
@@ -100,9 +85,9 @@ async def list_user_projects(
]
@router.get("/meta/db/health", summary="检查数据库健康状态", description="检查项目数据库连接的健康状况")
@router.get("/projects/current/database-health", summary="检查数据库健康状态", description="检查项目数据库连接的健康状况")
async def project_db_health(
pg_session: AsyncSession = Depends(get_project_pg_session),
pg_conn: AsyncConnection = Depends(get_project_pg_connection),
ts_conn: AsyncConnection = Depends(get_project_timescale_connection),
):
"""
@@ -110,7 +95,25 @@ async def project_db_health(
检查PostgreSQL和TimescaleDB数据库的连接状态
"""
await pg_session.execute(text("SELECT 1"))
async with ts_conn.cursor() as cur:
await cur.execute("SELECT 1")
try:
async with pg_conn.cursor() as cur:
await cur.execute("SELECT 1")
await cur.fetchone()
except psycopg.Error as exc:
logger.error("Project PostgreSQL health check failed", exc_info=True)
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"Project PostgreSQL health check failed: {exc}",
) from exc
try:
async with ts_conn.cursor() as cur:
await cur.execute("SELECT 1")
except psycopg.Error as exc:
logger.error("Project TimescaleDB health check failed", exc_info=True)
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=f"Project TimescaleDB health check failed: {exc}",
) from exc
return {"postgres": "ok", "timescale": "ok"}
-86
View File
@@ -1,86 +0,0 @@
from typing import Any
import random
from fastapi import APIRouter, Query
from fastapi.responses import JSONResponse
from fastapi import status
from pydantic import BaseModel
from app.services.tjnetwork import (
get_all_sensor_placements,
get_all_burst_locate_results,
)
router = APIRouter()
@router.get("/getjson/", summary="获取JSON示例", description="获取JSON格式响应示例")
async def fastapi_get_json():
"""
获取JSON示例
返回示例JSON格式的响应
"""
return JSONResponse(
status_code=status.HTTP_400_BAD_REQUEST,
content={
"code": 400,
"message": "this is message",
"data": 123,
},
)
@router.get("/getallsensorplacements/", summary="获取所有传感器位置", description="获取网络中所有传感器的放置位置信息")
async def fastapi_get_all_sensor_placements(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[dict[Any, Any]]:
"""
获取所有传感器位置
返回网络中所有传感器的放置位置及其配置信息
"""
return get_all_sensor_placements(network)
@router.get("/getallburstlocateresults/", summary="获取所有爆管定位结果", description="获取网络中所有爆管定位的分析结果")
async def fastapi_get_all_burst_locate_results(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[dict[Any, Any]]:
"""
获取所有爆管定位结果
返回网络中所有的爆管定位分析结果
"""
return get_all_burst_locate_results(network)
class Item(BaseModel):
"""测试数据模型"""
str_info: str
@router.post("/test_dict/", summary="测试字典处理", description="测试处理字典类型数据")
async def fastapi_test_dict(data: Item) -> dict[str, str]:
"""
测试字典处理
接收Item模型,返回其字典格式
"""
item = data.dict()
return item
@router.get("/getrealtimedata/", summary="获取实时数据", description="获取实时监测数据")
async def fastapi_get_realtimedata():
"""
获取实时数据
返回随机生成的实时监测数据示例
"""
data = [random.randint(0, 100) for _ in range(100)]
return data
@router.get("/getsimulationresult/", summary="获取模拟结果", description="获取仿真计算结果")
async def fastapi_get_simulationresult():
"""
获取仿真结果
返回随机生成的仿真计算结果示例
"""
data = [random.randint(0, 100) for _ in range(100)]
return data
+391
View File
@@ -0,0 +1,391 @@
import json
from pathlib import Path
from tempfile import NamedTemporaryFile
from uuid import UUID, uuid4
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 (
get_current_metadata_admin,
get_metadata_repository,
)
from app.auth.project_dependencies import (
ProjectContext,
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()
MAX_INP_FILE_BYTES = 50 * 1024 * 1024
INP_SECTIONS = ("[TITLE]", "[JUNCTIONS]", "[RESERVOIRS]", "[TANKS]", "[PIPES]")
async def _get_active_project(project_id: UUID, metadata_repo: MetadataRepository):
project = await metadata_repo.get_project_by_id(project_id)
if project is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Project not found",
)
if project.status != "active":
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Project is not active",
)
return project
def _validate_inp_bytes(content: bytes, filename: str) -> str:
if Path(filename).suffix.lower() != ".inp":
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Only .inp model files are accepted",
)
if not content:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="INP file is empty",
)
if len(content) > MAX_INP_FILE_BYTES:
raise HTTPException(
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
detail="INP file exceeds the 50 MiB limit",
)
for encoding in ("utf-8-sig", "gb18030"):
try:
text = content.decode(encoding)
break
except UnicodeDecodeError:
continue
else:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="INP file encoding is not supported",
)
upper_text = text.upper()
if not any(section in upper_text for section in INP_SECTIONS):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Invalid INP file structure",
)
return text
async def _read_upload(file: UploadFile) -> tuple[bytes, str]:
filename = Path(file.filename or "").name
content = await file.read(MAX_INP_FILE_BYTES + 1)
normalized = _validate_inp_bytes(content, filename).encode("utf-8")
return normalized, filename
async def _audit_model_change(
*,
request: Request,
current_user,
metadata_repo: MetadataRepository,
project_id: UUID,
action: str,
) -> None:
await log_audit_event(
action=AuditAction.UPDATE,
user_id=current_user.id,
project_id=project_id,
resource_type="hydraulic_model",
resource_id=action,
request_data={"operation": action},
ip_address=request.client.host if request.client else None,
request_method=request.method,
request_path=request.url.path,
response_status=status.HTTP_200_OK,
session=metadata_repo.session,
)
def _run_uploaded_inp_sync(content: bytes) -> str:
target_dir = Path("inp")
target_dir.mkdir(parents=True, exist_ok=True)
model_name = f"admin_model_{uuid4().hex}"
target_path = target_dir / f"{model_name}.inp"
target_path.write_bytes(content)
return run_inp(model_name)
async def _run_uploaded_inp(content: bytes) -> str:
return await run_in_threadpool(_run_uploaded_inp_sync, content)
def _update_from_inp_sync(content: bytes, project_code: str) -> None:
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)
network_update(str(temp_path), project_code)
finally:
if temp_path is not None:
temp_path.unlink(missing_ok=True)
async def _update_from_inp(content: bytes, project_code: str) -> None:
await run_in_threadpool(_update_from_inp_sync, content, project_code)
async def _apply_model_update(content: bytes, project_code: str) -> None:
try:
await _update_from_inp(content, project_code)
except MaterializedViewRefreshAfterCommitError:
raise
except Exception as exc:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"数据库操作失败: {exc}",
) 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="导入桌面端水力模型",
)
async def import_project_model(
request: Request,
project_id: UUID = ApiPath(...),
file: UploadFile = File(..., description="桌面端导出的 INP 模型文件"),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> dict:
project = await _get_active_project(project_id, metadata_repo)
content, filename = await _read_upload(file)
result = await _run_uploaded_inp(content)
await _audit_model_change(
request=request,
current_user=current_user,
metadata_repo=metadata_repo,
project_id=project.id,
action="import",
)
return {"project_id": str(project.id), "filename": filename, "result": result}
@router.patch(
"/admin/projects/{project_id}/model-imports",
summary="更新桌面端水力模型",
)
async def update_project_model(
request: Request,
project_id: UUID = ApiPath(...),
file: UploadFile = File(..., description="桌面端导出的 INP 模型文件"),
current_user=Depends(get_current_metadata_admin),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
) -> dict:
project = await _get_active_project(project_id, metadata_repo)
content, filename = await _read_upload(file)
routing = await resolve_project_business_routing(
ProjectContext(
project_id=project.id,
project_code=project.code,
user_id=current_user.id,
project_role="owner",
system_role=current_user.role,
is_superuser=current_user.is_superuser,
),
metadata_repo,
)
with activate_project_routing(routing):
await _apply_model_update(content, project.code)
await _audit_model_change(
request=request,
current_user=current_user,
metadata_repo=metadata_repo,
project_id=project.id,
action="update",
)
return {"project_id": str(project.id), "filename": filename, "updated": True}
+25 -25
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
calculate_demand_to_network,
calculate_demand_to_nodes,
@@ -18,11 +18,11 @@ router = APIRouter()
############################################################
@router.get(
"/getdemandschema",
"/network-schemas/demand",
summary="获取需水量属性架构",
description="获取指定水网中需水量(Demand)的属性架构定义"
)
async def fastapi_get_demand_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
def fastapi_get_demand_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""
获取需水量属性架构。
@@ -32,11 +32,11 @@ async def fastapi_get_demand_schema(network: str = Query(..., description="管
@router.get(
"/getdemandproperties/",
"/demands/properties",
summary="获取需水量属性",
description="获取指定水网中节点的需水量属性信息"
)
async def fastapi_get_demand_properties(
def fastapi_get_demand_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点ID")
) -> dict[str, Any]:
@@ -49,37 +49,37 @@ async def fastapi_get_demand_properties(
# example: set_demand(p, ChangeSet({'junction': 'j1', 'demands': [{'demand': 10.0, 'pattern': None, 'category': 'x'}, {'demand': 20.0, 'pattern': None, 'category': None}]}))
@router.post(
"/setdemandproperties/",
@router.patch(
"/demands/properties",
response_model=None,
summary="设置需水量属性",
description="设置指定水网中节点的需水量属性信息"
)
async def fastapi_set_demand_properties(
def fastapi_set_demand_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""
设置节点的需水量属性。
修改指定节点的需水量信息。请求体应包含需水量值、水压等级等属性。
"""
props = await req.json()
props = payload
ps = {"junction": junction} | props
return set_demand(network, ChangeSet(ps))
############################################################
# water distribution 36.[Water Distribution]
############################################################
@router.get(
"/calculatedemandtonodes/",
@router.post(
"/demands/to-nodes",
summary="计算需水量到节点分配",
description="将总需水量按指定方式分配到多个节点"
)
async def fastapi_calculate_demand_to_nodes(
def fastapi_calculate_demand_to_nodes(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> dict[str, float]:
"""
计算需水量到节点分配。
@@ -92,19 +92,19 @@ async def fastapi_calculate_demand_to_nodes(
"nodes": 节点ID列表(list[str])
}
"""
props = await req.json()
props = payload
demand = props["demand"]
nodes = props["nodes"]
return calculate_demand_to_nodes(network, demand, nodes)
@router.get(
"/calculatedemandtoregion/",
@router.post(
"/demands/to-region",
summary="计算需水量到区域分配",
description="将总需水量按区域特征分配到该区域内的节点"
)
async def fastapi_calculate_demand_to_region(
def fastapi_calculate_demand_to_region(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> dict[str, float]:
"""
计算需水量到区域分配。
@@ -117,17 +117,17 @@ async def fastapi_calculate_demand_to_region(
"region": 区域ID(str)
}
"""
props = await req.json()
props = payload
demand = props["demand"]
region = props["region"]
return calculate_demand_to_region(network, demand, region)
@router.get(
"/calculatedemandtonetwork/",
@router.post(
"/demands/to-network",
summary="计算需水量到整网分配",
description="将需水量均匀分配到整个水网的所有需水节点"
)
async def fastapi_calculate_demand_to_network(
def fastapi_calculate_demand_to_network(
network: str = Query(..., description="管网名称(或数据库名称)"),
demand: float = Query(..., description="总需水量(m³/h)", gt=0)
) -> dict[str, float]:
+68 -68
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
delete_junction,
delete_pipe,
@@ -45,11 +45,11 @@ router = APIRouter()
############################################################
@router.get(
"/isnode/",
"/nodes/existence",
summary="检查节点有效性",
description="检查指定ID是否为水网中的有效节点"
)
async def fastapi_is_node(
def fastapi_is_node(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> bool:
@@ -57,11 +57,11 @@ async def fastapi_is_node(
return is_node(network, node)
@router.get(
"/isjunction/",
"/junctions/existence",
summary="检查是否为接点",
description="检查指定ID是否为水网中的接点(需求点)"
)
async def fastapi_is_junction(
def fastapi_is_junction(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> bool:
@@ -69,11 +69,11 @@ async def fastapi_is_junction(
return is_junction(network, node)
@router.get(
"/isreservoir/",
"/reservoirs/existence",
summary="检查是否为水源",
description="检查指定ID是否为水网中的水源(水库/河流)"
)
async def fastapi_is_reservoir(
def fastapi_is_reservoir(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> bool:
@@ -81,11 +81,11 @@ async def fastapi_is_reservoir(
return is_reservoir(network, node)
@router.get(
"/istank/",
"/tanks/existence",
summary="检查是否为蓄水池",
description="检查指定ID是否为水网中的蓄水池"
)
async def fastapi_is_tank(
def fastapi_is_tank(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> bool:
@@ -93,11 +93,11 @@ async def fastapi_is_tank(
return is_tank(network, node)
@router.get(
"/islink/",
"/links/existence",
summary="检查管线有效性",
description="检查指定ID是否为水网中的有效管线"
)
async def fastapi_is_link(
def fastapi_is_link(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> bool:
@@ -105,11 +105,11 @@ async def fastapi_is_link(
return is_link(network, link)
@router.get(
"/ispipe/",
"/pipes/existence",
summary="检查是否为管道",
description="检查指定ID是否为水网中的管道"
)
async def fastapi_is_pipe(
def fastapi_is_pipe(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> bool:
@@ -117,11 +117,11 @@ async def fastapi_is_pipe(
return is_pipe(network, link)
@router.get(
"/ispump/",
"/pumps/existence",
summary="检查是否为泵",
description="检查指定ID是否为水网中的泵"
)
async def fastapi_is_pump(
def fastapi_is_pump(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> bool:
@@ -129,11 +129,11 @@ async def fastapi_is_pump(
return is_pump(network, link)
@router.get(
"/isvalve/",
"/valves/existence",
summary="检查是否为阀门",
description="检查指定ID是否为水网中的阀门"
)
async def fastapi_is_valve(
def fastapi_is_valve(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> bool:
@@ -141,11 +141,11 @@ async def fastapi_is_valve(
return is_valve(network, link)
@router.get(
"/getnodetype/",
"/node-types",
summary="获取节点类型",
description="获取指定节点的类型(接点/水源/蓄水池)"
)
async def fastapi_get_node_type(
def fastapi_get_node_type(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> str:
@@ -153,11 +153,11 @@ async def fastapi_get_node_type(
return get_node_type(network, node)
@router.get(
"/getlinktype/",
"/link-types",
summary="获取管线类型",
description="获取指定管线的类型(管道/泵/阀门)"
)
async def fastapi_get_link_type(
def fastapi_get_link_type(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> str:
@@ -165,11 +165,11 @@ async def fastapi_get_link_type(
return get_link_type(network, link)
@router.get(
"/getelementtype/",
"/element-types",
summary="获取元素类型",
description="获取指定元素的类型(节点或管线)"
)
async def fastapi_get_element_type(
def fastapi_get_element_type(
network: str = Query(..., description="管网名称(或数据库名称)"),
element: str = Query(..., description="元素ID")
) -> str:
@@ -177,11 +177,11 @@ async def fastapi_get_element_type(
return get_element_type(network, element)
@router.get(
"/getelementtypevalue/",
"/element-type-values",
summary="获取元素类型值",
description="获取指定元素的类型数值标识"
)
async def fastapi_get_element_type_value(
def fastapi_get_element_type_value(
network: str = Query(..., description="管网名称(或数据库名称)"),
element: str = Query(..., description="元素ID")
) -> int:
@@ -189,25 +189,25 @@ async def fastapi_get_element_type_value(
return get_element_type_value(network, element)
@router.get(
"/getnodes/",
"/nodes",
summary="获取所有节点",
description="获取指定水网中的所有节点ID列表"
)
async def fastapi_get_nodes(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
def fastapi_get_nodes(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
"""获取水网中所有节点的ID列表。"""
return get_nodes(network)
@router.get(
"/getlinks/",
"/links",
summary="获取所有管线",
description="获取指定水网中的所有管线ID列表"
)
async def fastapi_get_links(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
def fastapi_get_links(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[str]:
"""获取水网中所有管线的ID列表。"""
return get_links(network)
@router.get(
"/getnodelinks/",
"/node-links",
summary="获取节点的关联管线",
description="获取指定节点连接的所有管线ID列表"
)
@@ -223,11 +223,11 @@ def get_node_links_endpoint(
############################################################
@router.get(
"/getnodeproperties/",
"/node-properties",
summary="获取节点属性",
description="获取指定节点的所有属性信息"
)
async def fast_get_node_properties(
def fast_get_node_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> dict[str, Any]:
@@ -235,11 +235,11 @@ async def fast_get_node_properties(
return get_node_properties(network, node)
@router.get(
"/getlinkproperties/",
"/link-properties",
summary="获取管线属性",
description="获取指定管线的所有属性信息"
)
async def fast_get_link_properties(
def fast_get_link_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> dict[str, Any]:
@@ -247,11 +247,11 @@ async def fast_get_link_properties(
return get_link_properties(network, link)
@router.get(
"/getscadaproperties/",
"/scada-properties",
summary="获取SCADA点属性",
description="获取指定SCADA点的属性信息"
)
async def fast_get_scada_properties(
def fast_get_scada_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
scada: str = Query(..., description="SCADA点ID")
) -> dict[str, Any]:
@@ -259,22 +259,22 @@ async def fast_get_scada_properties(
return get_scada_info(network, scada)
@router.get(
"/getallscadaproperties/",
"/all-scada-properties",
summary="获取所有SCADA点属性",
description="获取指定水网中所有SCADA点的属性信息"
)
async def fast_get_all_scada_properties(
def fast_get_all_scada_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""获取水网中所有SCADA点的属性列表。"""
return get_all_scada_info(network)
@router.get(
"/getelementpropertieswithtype/",
"/element-properties-with-types",
summary="获取指定类型元素属性",
description="获取指定类型的元素属性信息"
)
async def fast_get_element_properties_with_type(
def fast_get_element_properties_with_type(
network: str = Query(..., description="管网名称(或数据库名称)"),
elementtype: str = Query(..., description="元素类型"),
element: str = Query(..., description="元素ID")
@@ -283,11 +283,11 @@ async def fast_get_element_properties_with_type(
return get_element_properties_with_type(network, elementtype, element)
@router.get(
"/getelementproperties/",
"/element-properties",
summary="获取元素属性",
description="获取指定元素的属性信息"
)
async def fast_get_element_properties(
def fast_get_element_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
element: str = Query(..., description="元素ID")
) -> dict[str, Any]:
@@ -299,37 +299,37 @@ async def fast_get_element_properties(
############################################################
@router.get(
"/gettitleschema/",
"/title-schemas",
summary="获取标题属性架构",
description="获取指定水网的标题(标题)属性架构定义"
)
async def fast_get_title_schema(
def fast_get_title_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""获取水网标题的属性架构。"""
return get_title_schema(network)
@router.get(
"/gettitle/",
"/titles",
summary="获取水网标题属性",
description="获取指定水网的标题(Title)信息"
)
async def fast_get_title(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
def fast_get_title(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, Any]:
"""获取水网的标题属性。"""
return get_title(network)
@router.get(
"/settitle/",
@router.patch(
"/titles",
response_model=None,
summary="设置水网标题属性",
description="设置指定水网的标题(Title)信息"
)
async def fastapi_set_title(
def fastapi_set_title(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置水网的标题属性。"""
props = await req.json()
props = payload
return set_title(network, ChangeSet(props))
############################################################
@@ -337,41 +337,41 @@ async def fastapi_set_title(
############################################################
@router.get(
"/getstatusschema",
"/status-schemas",
summary="获取状态属性架构",
description="获取指定水网的状态(Status)属性架构定义"
)
async def fastapi_get_status_schema(
def fastapi_get_status_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""获取水网状态的属性架构。"""
return get_status_schema(network)
@router.get(
"/getstatus/",
"/status",
summary="获取管线状态",
description="获取指定管线的状态信息"
)
async def fastapi_get_status(
def fastapi_get_status(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> dict[str, Any]:
"""获取管线的状态属性。"""
return get_status(network, link)
@router.post(
"/setstatus/",
@router.patch(
"/status-properties",
response_model=None,
summary="设置管线状态",
description="设置指定管线的状态信息"
)
async def fastapi_set_status_properties(
def fastapi_set_status_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置管线的状态属性。"""
props = await req.json()
props = payload
ps = {"link": link} | props
return set_status(network, ChangeSet(ps))
@@ -379,13 +379,13 @@ async def fastapi_set_status_properties(
# General Deletion
############################################################
@router.post(
"/deletenode/",
@router.delete(
"/nodes",
response_model=None,
summary="删除节点",
description="删除指定的节点(接点/水源/蓄水池)"
)
async def fastapi_delete_node(
def fastapi_delete_node(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> ChangeSet:
@@ -399,13 +399,13 @@ async def fastapi_delete_node(
return delete_tank(network, ChangeSet(ps))
return ChangeSet() # Should probably raise error or return empty
@router.post(
"/deletelink/",
@router.delete(
"/links",
response_model=None,
summary="删除管线",
description="删除指定的管线(管道/泵/阀门)"
)
async def fastapi_delete_link(
def fastapi_delete_link(
network: str = Query(..., description="管网名称(或数据库名称)"),
link: str = Query(..., description="管线ID")
) -> ChangeSet:
+15 -46
View File
@@ -1,18 +1,14 @@
from fastapi import APIRouter, Request, Depends, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Query
from app.services.tjnetwork import (
Any,
get_all_scada_info,
get_major_node_coords,
get_major_pipe_nodes,
get_network_in_extent,
get_network_link_nodes,
get_network_node_coords,
get_node_coord,
)
from app.auth.dependencies import get_current_user as verify_token
from app.infra.cache.redis_client import redis_client, encode_datetime, decode_datetime
import msgpack
router = APIRouter()
@@ -31,15 +27,15 @@ router = APIRouter()
# # example: set_coord(p, ChangeSet({'node': 'j1', 'x': 1.0, 'y': 2.0}))
# @router.post("/setcoord/", response_model=None)
# async def fastapi_set_coord(network: str, req: Request) -> ChangeSet:
# props = await req.json()
# props = payload
# return set_coord(network, ChangeSet(props))
@router.get(
"/getnodecoord/",
"/node-coords",
summary="获取节点坐标",
description="获取指定节点的地理坐标(X, Y)"
)
async def fastapi_get_node_coord(
def fastapi_get_node_coord(
network: str = Query(..., description="管网名称(或数据库名称)"),
node: str = Query(..., description="节点ID")
) -> dict[str, float] | None:
@@ -48,11 +44,11 @@ async def fastapi_get_node_coord(
# Additional geometry queries found in main.py logic (implicit or explicit)
@router.get(
"/getnetworkinextent/",
"/network-in-extents",
summary="获取范围内的网络元素",
description="获取指定地理范围内的网络节点和管线"
)
async def fastapi_get_network_in_extent(
def fastapi_get_network_in_extent(
network: str = Query(..., description="管网名称(或数据库名称)"),
x1: float = Query(..., description="范围左下角X坐标", alias="x1"),
y1: float = Query(..., description="范围左下角Y坐标", alias="y1"),
@@ -63,38 +59,11 @@ async def fastapi_get_network_in_extent(
return get_network_in_extent(network, x1, y1, x2, y2)
@router.get(
"/getnetworkgeometries/",
dependencies=[Depends(verify_token)],
summary="获取完整网络几何信息",
description="获取整个水网的所有节点、管线和SCADA点的几何信息(需要身份验证)"
)
async def fastapi_get_network_geometries(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, Any] | None:
"""获取完整的网络几何信息,包括所有节点、管线和SCADA点。结果从缓存返回。"""
cache_key = f"getnetworkgeometries_{network}"
data = redis_client.get(cache_key)
if data:
loaded_dict = msgpack.unpackb(data, object_hook=decode_datetime)
return loaded_dict
coords = get_network_node_coords(network)
nodes = []
for node_id, coord in coords.items():
nodes.append(f"{node_id}:{coord['type']}:{coord['x']}:{coord['y']}")
links = get_network_link_nodes(network)
scadas = get_all_scada_info(network)
results = {"nodes": nodes, "links": links, "scadas": scadas}
redis_client.set(cache_key, msgpack.packb(results, default=encode_datetime))
return results
@router.get(
"/getmajornodecoords/",
"/majornode-coords",
summary="获取主要节点坐标",
description="获取直径大于等于指定值的节点坐标"
)
async def fastapi_get_majornode_coords(
def fastapi_get_majornode_coords(
network: str = Query(..., description="管网名称(或数据库名称)"),
diameter: int = Query(..., description="最小直径(mm)", gt=0)
) -> dict[str, dict[str, float]]:
@@ -102,11 +71,11 @@ async def fastapi_get_majornode_coords(
return get_major_node_coords(network, diameter)
@router.get(
"/getmajorpipenodes/",
"/major-pipe-nodes",
summary="获取主要管道节点",
description="获取直径大于等于指定值的管道的节点ID"
)
async def fastapi_get_major_pipe_nodes(
def fastapi_get_major_pipe_nodes(
network: str = Query(..., description="管网名称(或数据库名称)"),
diameter: int = Query(..., description="最小直径(mm)", gt=0)
) -> list[str] | None:
@@ -114,11 +83,11 @@ async def fastapi_get_major_pipe_nodes(
return get_major_pipe_nodes(network, diameter)
@router.get(
"/getnetworklinknodes/",
"/network-link-nodes",
summary="获取网络管线节点",
description="获取指定水网所有管线的起点和终点节点"
)
async def fastapi_get_network_link_nodes(
def fastapi_get_network_link_nodes(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[str] | None:
"""获取网络中所有管线的连接节点。"""
+41 -42
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_junction,
delete_junction,
@@ -13,8 +13,8 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/getjunctionschema", summary="获取节点架构", description="获取指定项目的节点属性架构和数据类型定义。")
async def fast_get_junction_schema(
@router.get("/network-schemas/junction", summary="获取节点架构", description="获取指定项目的节点属性架构和数据类型定义。")
def fast_get_junction_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""
@@ -27,8 +27,8 @@ async def fast_get_junction_schema(
"""
return get_junction_schema(network)
@router.post("/addjunction/", response_model=None, summary="添加节点", description="在供水网络中添加新的节点,指定节点ID和空间坐标。")
async def fastapi_add_junction(
@router.post("/junctions", response_model=None, summary="添加节点", description="在供水网络中添加新的节点,指定节点ID和空间坐标。")
def fastapi_add_junction(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
x: float = Query(..., description="X 坐标"),
@@ -51,8 +51,8 @@ async def fastapi_add_junction(
ps = {"id": junction, "x": x, "y": y, "elevation": z}
return add_junction(network, ChangeSet(ps))
@router.post("/deletejunction/", response_model=None, summary="删除节点", description="从供水网络中删除指定的节点。")
async def fastapi_delete_junction(
@router.delete("/junctions", response_model=None, summary="删除节点", description="从供水网络中删除指定的节点。")
def fastapi_delete_junction(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> ChangeSet:
@@ -69,8 +69,8 @@ async def fastapi_delete_junction(
ps = {"id": junction}
return delete_junction(network, ChangeSet(ps))
@router.get("/getjunctionelevation/", summary="获取节点标高", description="获取指定节点的标高(海拔高度)。")
async def fastapi_get_junction_elevation(
@router.get("/junctions/elevation", summary="获取节点标高", description="获取指定节点的标高(海拔高度)。")
def fastapi_get_junction_elevation(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> float:
@@ -87,8 +87,8 @@ async def fastapi_get_junction_elevation(
ps = get_junction(network, junction)
return ps["elevation"]
@router.get("/getjunctionx/", summary="获取节点 X 坐标", description="获取指定节点的 X 坐标值。")
async def fastapi_get_junction_x(
@router.get("/junctions/x", summary="获取节点 X 坐标", description="获取指定节点的 X 坐标值。")
def fastapi_get_junction_x(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> float:
@@ -105,8 +105,8 @@ async def fastapi_get_junction_x(
ps = get_junction(network, junction)
return ps["x"]
@router.get("/getjunctiony/", summary="获取节点 Y 坐标", description="获取指定节点的 Y 坐标值。")
async def fastapi_get_junction_y(
@router.get("/junctions/y", summary="获取节点 Y 坐标", description="获取指定节点的 Y 坐标值。")
def fastapi_get_junction_y(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> float:
@@ -123,8 +123,8 @@ async def fastapi_get_junction_y(
ps = get_junction(network, junction)
return ps["y"]
@router.get("/getjunctioncoord/", summary="获取节点坐标", description="获取指定节点的 X 和 Y 坐标。")
async def fastapi_get_junction_coord(
@router.get("/junctions/coord", summary="获取节点坐标", description="获取指定节点的 X 和 Y 坐标。")
def fastapi_get_junction_coord(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> dict[str, float]:
@@ -142,8 +142,8 @@ async def fastapi_get_junction_coord(
coord = {"x": ps["x"], "y": ps["y"]}
return coord
@router.get("/getjunctiondemand/", summary="获取节点需水量", description="获取指定节点的需水量。")
async def fastapi_get_junction_demand(
@router.get("/junctions/demand", summary="获取节点需水量", description="获取指定节点的需水量。")
def fastapi_get_junction_demand(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> float:
@@ -160,8 +160,8 @@ async def fastapi_get_junction_demand(
ps = get_junction(network, junction)
return ps["demand"]
@router.get("/getjunctionpattern/", summary="获取节点需水模式", description="获取指定节点的需水模式标识。")
async def fastapi_get_junction_pattern(
@router.get("/junctions/pattern", summary="获取节点需水模式", description="获取指定节点的需水模式标识。")
def fastapi_get_junction_pattern(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> str:
@@ -178,8 +178,8 @@ async def fastapi_get_junction_pattern(
ps = get_junction(network, junction)
return ps["pattern"]
@router.post("/setjunctionelevation/", response_model=None, summary="设置节点标高", description="设置指定节点的标高值。")
async def fastapi_set_junction_elevation(
@router.patch("/junctions/elevation", response_model=None, summary="设置节点标高", description="设置指定节点的标高值。")
def fastapi_set_junction_elevation(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
elevation: float = Query(..., description="标高(海拔高度)")
@@ -198,8 +198,8 @@ async def fastapi_set_junction_elevation(
ps = {"id": junction, "elevation": elevation}
return set_junction(network, ChangeSet(ps))
@router.post("/setjunctionx/", response_model=None, summary="设置节点 X 坐标", description="设置指定节点的 X 坐标值。")
async def fastapi_set_junction_x(
@router.patch("/junctions/x", response_model=None, summary="设置节点 X 坐标", description="设置指定节点的 X 坐标值。")
def fastapi_set_junction_x(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
x: float = Query(..., description="X 坐标值")
@@ -218,8 +218,8 @@ async def fastapi_set_junction_x(
ps = {"id": junction, "x": x}
return set_junction(network, ChangeSet(ps))
@router.post("/setjunctiony/", response_model=None, summary="设置节点 Y 坐标", description="设置指定节点的 Y 坐标值。")
async def fastapi_set_junction_y(
@router.patch("/junctions/y", response_model=None, summary="设置节点 Y 坐标", description="设置指定节点的 Y 坐标值。")
def fastapi_set_junction_y(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
y: float = Query(..., description="Y 坐标值")
@@ -238,8 +238,8 @@ async def fastapi_set_junction_y(
ps = {"id": junction, "y": y}
return set_junction(network, ChangeSet(ps))
@router.post("/setjunctioncoord/", response_model=None, summary="设置节点坐标", description="设置指定节点的 X 和 Y 坐标。")
async def fastapi_set_junction_coord(
@router.patch("/junctions/coord", response_model=None, summary="设置节点坐标", description="设置指定节点的 X 和 Y 坐标。")
def fastapi_set_junction_coord(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
x: float = Query(..., description="X 坐标值"),
@@ -260,8 +260,8 @@ async def fastapi_set_junction_coord(
ps = {"id": junction, "x": x, "y": y}
return set_junction(network, ChangeSet(ps))
@router.post("/setjunctiondemand/", response_model=None, summary="设置节点需水量", description="设置指定节点的需水量。")
async def fastapi_set_junction_demand(
@router.patch("/junctions/demand", response_model=None, summary="设置节点需水量", description="设置指定节点的需水量。")
def fastapi_set_junction_demand(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
demand: float = Query(..., description="需水量值")
@@ -280,8 +280,8 @@ async def fastapi_set_junction_demand(
ps = {"id": junction, "demand": demand}
return set_junction(network, ChangeSet(ps))
@router.post("/setjunctionpattern/", response_model=None, summary="设置节点需水模式", description="设置指定节点的需水模式标识。")
async def fastapi_set_junction_pattern(
@router.patch("/junctions/pattern", response_model=None, summary="设置节点需水模式", description="设置指定节点的需水模式标识。")
def fastapi_set_junction_pattern(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
pattern: str = Query(..., description="需水模式标识")
@@ -300,8 +300,8 @@ async def fastapi_set_junction_pattern(
ps = {"id": junction, "pattern": pattern}
return set_junction(network, ChangeSet(ps))
@router.get("/getjunctionproperties/", summary="获取节点属性", description="获取指定节点的所有属性信息。")
async def fastapi_get_junction_properties(
@router.get("/junctions/properties", summary="获取节点属性", description="获取指定节点的所有属性信息。")
def fastapi_get_junction_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID")
) -> dict[str, Any]:
@@ -317,8 +317,8 @@ async def fastapi_get_junction_properties(
"""
return get_junction(network, junction)
@router.get("/getalljunctionproperties/", summary="获取所有节点属性", description="获取指定项目中所有节点的属性信息。")
async def fastapi_get_all_junction_properties(
@router.get("/junctions", summary="获取所有节点属性", description="获取指定项目中所有节点的属性信息。")
def fastapi_get_all_junction_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
@@ -333,15 +333,14 @@ async def fastapi_get_all_junction_properties(
list: 包含所有节点属性的列表
"""
# 缓存查询结果提高性能
# global redis_client # Redis logic removed for clean split, can be re-added if needed or imported
results = get_all_junctions(network)
return results
@router.post("/setjunctionproperties/", response_model=None, summary="批量设置节点属性", description="批量设置指定节点的多个属性。")
async def fastapi_set_junction_properties(
@router.patch("/junctions/properties", response_model=None, summary="批量设置节点属性", description="批量设置指定节点的多个属性。")
def fastapi_set_junction_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
junction: str = Query(..., description="节点 ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""
批量设置节点属性。
@@ -356,6 +355,6 @@ async def fastapi_set_junction_properties(
Returns:
ChangeSet: 包含变更信息的结果
"""
props = await req.json()
props = payload
ps = {"id": junction} | props
return set_junction(network, ChangeSet(ps))
+45 -46
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
PIPE_STATUS_OPEN,
add_pipe,
@@ -14,8 +14,8 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/getpipeschema", summary="获取管道模式", description="获取管道对象的模式定义,包含所有可用字段及其类型")
async def fastapi_get_pipe_schema(
@router.get("/network-schemas/pipe", summary="获取管道模式", description="获取管道对象的模式定义,包含所有可用字段及其类型")
def fastapi_get_pipe_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""
@@ -29,8 +29,8 @@ async def fastapi_get_pipe_schema(
"""
return get_pipe_schema(network)
@router.post("/addpipe/", response_model=None, summary="添加管道", description="向网络中添加新的管道,需要提供管道的基本参数如长度、管径、粗糙度等")
async def fastapi_add_pipe(
@router.post("/pipes", response_model=None, summary="添加管道", description="向网络中添加新的管道,需要提供管道的基本参数如长度、管径、粗糙度等")
def fastapi_add_pipe(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道标识符"),
node1: str = Query(..., description="管道起始节点ID"),
@@ -70,8 +70,8 @@ async def fastapi_add_pipe(
}
return add_pipe(network, ChangeSet(ps))
@router.post("/deletepipe/", response_model=None, summary="删除管道", description="从网络中删除指定的管道")
async def fastapi_delete_pipe(
@router.delete("/pipes", response_model=None, summary="删除管道", description="从网络中删除指定的管道")
def fastapi_delete_pipe(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="要删除的管道ID")
) -> ChangeSet:
@@ -88,8 +88,8 @@ async def fastapi_delete_pipe(
ps = {"id": pipe}
return delete_pipe(network, ChangeSet(ps))
@router.get("/getpipenode1/", summary="获取管道起始节点", description="获取指定管道的起始节点ID")
async def fastapi_get_pipe_node1(
@router.get("/pipes/node1", summary="获取管道起始节点", description="获取指定管道的起始节点ID")
def fastapi_get_pipe_node1(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> str | None:
@@ -106,8 +106,8 @@ async def fastapi_get_pipe_node1(
ps = get_pipe(network, pipe)
return ps["node1"]
@router.get("/getpipenode2/", summary="获取管道终止节点", description="获取指定管道的终止节点ID")
async def fastapi_get_pipe_node2(
@router.get("/pipes/node2", summary="获取管道终止节点", description="获取指定管道的终止节点ID")
def fastapi_get_pipe_node2(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> str | None:
@@ -124,8 +124,8 @@ async def fastapi_get_pipe_node2(
ps = get_pipe(network, pipe)
return ps["node2"]
@router.get("/getpipelength/", summary="获取管道长度", description="获取指定管道的长度")
async def fastapi_get_pipe_length(
@router.get("/pipes/length", summary="获取管道长度", description="获取指定管道的长度")
def fastapi_get_pipe_length(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> float | None:
@@ -142,8 +142,8 @@ async def fastapi_get_pipe_length(
ps = get_pipe(network, pipe)
return ps["length"]
@router.get("/getpipediameter/", summary="获取管道管径", description="获取指定管道的管径")
async def fastapi_get_pipe_diameter(
@router.get("/pipes/diameter", summary="获取管道管径", description="获取指定管道的管径")
def fastapi_get_pipe_diameter(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> float | None:
@@ -160,8 +160,8 @@ async def fastapi_get_pipe_diameter(
ps = get_pipe(network, pipe)
return ps["diameter"]
@router.get("/getpiperoughness/", summary="获取管道粗糙度", description="获取指定管道的粗糙度")
async def fastapi_get_pipe_roughness(
@router.get("/pipes/roughness", summary="获取管道粗糙度", description="获取指定管道的粗糙度")
def fastapi_get_pipe_roughness(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> float | None:
@@ -178,8 +178,8 @@ async def fastapi_get_pipe_roughness(
ps = get_pipe(network, pipe)
return ps["roughness"]
@router.get("/getpipeminorloss/", summary="获取管道局部阻力系数", description="获取指定管道的局部阻力系数")
async def fastapi_get_pipe_minor_loss(
@router.get("/pipes/minor-loss", summary="获取管道局部阻力系数", description="获取指定管道的局部阻力系数")
def fastapi_get_pipe_minor_loss(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> float | None:
@@ -196,8 +196,8 @@ async def fastapi_get_pipe_minor_loss(
ps = get_pipe(network, pipe)
return ps["minor_loss"]
@router.get("/getpipestatus/", summary="获取管道状态", description="获取指定管道的状态(开启或关闭)")
async def fastapi_get_pipe_status(
@router.get("/pipes/status", summary="获取管道状态", description="获取指定管道的状态(开启或关闭)")
def fastapi_get_pipe_status(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> str | None:
@@ -214,8 +214,8 @@ async def fastapi_get_pipe_status(
ps = get_pipe(network, pipe)
return ps["status"]
@router.post("/setpipenode1/", response_model=None, summary="设置管道起始节点", description="设置指定管道的起始节点")
async def fastapi_set_pipe_node1(
@router.patch("/pipes/node1", response_model=None, summary="设置管道起始节点", description="设置指定管道的起始节点")
def fastapi_set_pipe_node1(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
node1: str = Query(..., description="新的起始节点ID")
@@ -234,8 +234,8 @@ async def fastapi_set_pipe_node1(
ps = {"id": pipe, "node1": node1}
return set_pipe(network, ChangeSet(ps))
@router.post("/setpipenode2/", response_model=None, summary="设置管道终止节点", description="设置指定管道的终止节点")
async def fastapi_set_pipe_node2(
@router.patch("/pipes/node2", response_model=None, summary="设置管道终止节点", description="设置指定管道的终止节点")
def fastapi_set_pipe_node2(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
node2: str = Query(..., description="新的终止节点ID")
@@ -254,8 +254,8 @@ async def fastapi_set_pipe_node2(
ps = {"id": pipe, "node2": node2}
return set_pipe(network, ChangeSet(ps))
@router.post("/setpipelength/", response_model=None, summary="设置管道长度", description="设置指定管道的长度")
async def fastapi_set_pipe_length(
@router.patch("/pipes/length", response_model=None, summary="设置管道长度", description="设置指定管道的长度")
def fastapi_set_pipe_length(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
length: float = Query(..., description="新的管道长度(单位:米)")
@@ -274,8 +274,8 @@ async def fastapi_set_pipe_length(
ps = {"id": pipe, "length": length}
return set_pipe(network, ChangeSet(ps))
@router.post("/setpipediameter/", response_model=None, summary="设置管道管径", description="设置指定管道的管径")
async def fastapi_set_pipe_diameter(
@router.patch("/pipes/diameter", response_model=None, summary="设置管道管径", description="设置指定管道的管径")
def fastapi_set_pipe_diameter(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
diameter: float = Query(..., description="新的管道管径(单位:毫米)")
@@ -294,8 +294,8 @@ async def fastapi_set_pipe_diameter(
ps = {"id": pipe, "diameter": diameter}
return set_pipe(network, ChangeSet(ps))
@router.post("/setpiperoughness/", response_model=None, summary="设置管道粗糙度", description="设置指定管道的粗糙度")
async def fastapi_set_pipe_roughness(
@router.patch("/pipes/roughness", response_model=None, summary="设置管道粗糙度", description="设置指定管道的粗糙度")
def fastapi_set_pipe_roughness(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
roughness: float = Query(..., description="新的管道粗糙度值")
@@ -314,8 +314,8 @@ async def fastapi_set_pipe_roughness(
ps = {"id": pipe, "roughness": roughness}
return set_pipe(network, ChangeSet(ps))
@router.post("/setpipeminorloss/", response_model=None, summary="设置管道局部阻力系数", description="设置指定管道的局部阻力系数")
async def fastapi_set_pipe_minor_loss(
@router.patch("/pipes/minor-loss", response_model=None, summary="设置管道局部阻力系数", description="设置指定管道的局部阻力系数")
def fastapi_set_pipe_minor_loss(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
minor_loss: float = Query(..., description="新的局部阻力系数值")
@@ -334,8 +334,8 @@ async def fastapi_set_pipe_minor_loss(
ps = {"id": pipe, "minor_loss": minor_loss}
return set_pipe(network, ChangeSet(ps))
@router.post("/setpipestatus/", response_model=None, summary="设置管道状态", description="设置指定管道的状态(开启或关闭)")
async def fastapi_set_pipe_status(
@router.patch("/pipes/status", response_model=None, summary="设置管道状态", description="设置指定管道的状态(开启或关闭)")
def fastapi_set_pipe_status(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
status: str = Query(..., description="新的管道状态(开启/关闭)")
@@ -354,8 +354,8 @@ async def fastapi_set_pipe_status(
ps = {"id": pipe, "status": status}
return set_pipe(network, ChangeSet(ps))
@router.get("/getpipeproperties/", summary="获取管道属性", description="获取指定管道的所有属性信息")
async def fastapi_get_pipe_properties(
@router.get("/pipes/properties", summary="获取管道属性", description="获取指定管道的所有属性信息")
def fastapi_get_pipe_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID")
) -> dict[str, Any]:
@@ -371,8 +371,8 @@ async def fastapi_get_pipe_properties(
"""
return get_pipe(network, pipe)
@router.get("/getallpipeproperties/", summary="获取所有管道属性", description="获取网络中所有管道的属性信息列表")
async def fastapi_get_all_pipe_properties(
@router.get("/pipes", summary="获取所有管道属性", description="获取网络中所有管道的属性信息列表")
def fastapi_get_all_pipe_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
@@ -385,15 +385,14 @@ async def fastapi_get_all_pipe_properties(
包含所有管道属性的字典列表
"""
# 缓存查询结果提高性能
# global redis_client
results = get_all_pipes(network)
return results
@router.post("/setpipeproperties/", response_model=None, summary="设置管道属性", description="批量设置指定管道的多个属性")
async def fastapi_set_pipe_properties(
@router.patch("/pipes/properties", response_model=None, summary="设置管道属性", description="批量设置指定管道的多个属性")
def fastapi_set_pipe_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe: str = Query(..., description="管道ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""
批量设置管道属性。
@@ -406,6 +405,6 @@ async def fastapi_set_pipe_properties(
Returns:
ChangeSet对象,包含本次修改的变更信息
"""
props = await req.json()
props = payload
ps = {"id": pipe} | props
return set_pipe(network, ChangeSet(ps))
+25 -26
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_pump,
delete_pump,
@@ -13,8 +13,8 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/getpumpschema", summary="获取水泵模式", description="获取水泵对象的模式定义,包含所有可用字段及其类型")
async def fastapi_get_pump_schema(
@router.get("/network-schemas/pump", summary="获取水泵模式", description="获取水泵对象的模式定义,包含所有可用字段及其类型")
def fastapi_get_pump_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""
@@ -28,8 +28,8 @@ async def fastapi_get_pump_schema(
"""
return get_pump_schema(network)
@router.post("/addpump/", response_model=None, summary="添加水泵", description="向网络中添加新的水泵,需要提供水泵的基本参数如功率等")
async def fastapi_add_pump(
@router.post("/pumps", response_model=None, summary="添加水泵", description="向网络中添加新的水泵,需要提供水泵的基本参数如功率等")
def fastapi_add_pump(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="水泵标识符"),
node1: str = Query(..., description="水泵起始节点ID"),
@@ -52,8 +52,8 @@ async def fastapi_add_pump(
ps = {"id": pump, "node1": node1, "node2": node2, "power": power}
return add_pump(network, ChangeSet(ps))
@router.post("/deletepump/", response_model=None, summary="删除水泵", description="从网络中删除指定的水泵")
async def fastapi_delete_pump(
@router.delete("/pumps", response_model=None, summary="删除水泵", description="从网络中删除指定的水泵")
def fastapi_delete_pump(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="要删除的水泵ID")
) -> ChangeSet:
@@ -70,8 +70,8 @@ async def fastapi_delete_pump(
ps = {"id": pump}
return delete_pump(network, ChangeSet(ps))
@router.get("/getpumpnode1/", summary="获取水泵起始节点", description="获取指定水泵的起始节点ID")
async def fastapi_get_pump_node1(
@router.get("/pumps/node1", summary="获取水泵起始节点", description="获取指定水泵的起始节点ID")
def fastapi_get_pump_node1(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="水泵ID")
) -> str | None:
@@ -88,8 +88,8 @@ async def fastapi_get_pump_node1(
ps = get_pump(network, pump)
return ps["node1"]
@router.get("/getpumpnode2/", summary="获取水泵终止节点", description="获取指定水泵的终止节点ID")
async def fastapi_get_pump_node2(
@router.get("/pumps/node2", summary="获取水泵终止节点", description="获取指定水泵的终止节点ID")
def fastapi_get_pump_node2(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="水泵ID")
) -> str | None:
@@ -106,8 +106,8 @@ async def fastapi_get_pump_node2(
ps = get_pump(network, pump)
return ps["node2"]
@router.post("/setpumpnode1/", response_model=None, summary="设置水泵起始节点", description="设置指定水泵的起始节点")
async def fastapi_set_pump_node1(
@router.patch("/pumps/node1", response_model=None, summary="设置水泵起始节点", description="设置指定水泵的起始节点")
def fastapi_set_pump_node1(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="水泵ID"),
node1: str = Query(..., description="新的起始节点ID")
@@ -126,8 +126,8 @@ async def fastapi_set_pump_node1(
ps = {"id": pump, "node1": node1}
return set_pump(network, ChangeSet(ps))
@router.post("/setpumpnode2/", response_model=None, summary="设置水泵终止节点", description="设置指定水泵的终止节点")
async def fastapi_set_pump_node2(
@router.patch("/pumps/node2", response_model=None, summary="设置水泵终止节点", description="设置指定水泵的终止节点")
def fastapi_set_pump_node2(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="水泵ID"),
node2: str = Query(..., description="新的终止节点ID")
@@ -146,8 +146,8 @@ async def fastapi_set_pump_node2(
ps = {"id": pump, "node2": node2}
return set_pump(network, ChangeSet(ps))
@router.get("/getpumpproperties/", summary="获取水泵属性", description="获取指定水泵的所有属性信息")
async def fastapi_get_pump_properties(
@router.get("/pumps/properties", summary="获取水泵属性", description="获取指定水泵的所有属性信息")
def fastapi_get_pump_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="水泵ID")
) -> dict[str, Any]:
@@ -163,8 +163,8 @@ async def fastapi_get_pump_properties(
"""
return get_pump(network, pump)
@router.get("/getallpumpproperties/", summary="获取所有水泵属性", description="获取网络中所有水泵的属性信息列表")
async def fastapi_get_all_pump_properties(
@router.get("/pumps", summary="获取所有水泵属性", description="获取网络中所有水泵的属性信息列表")
def fastapi_get_all_pump_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
@@ -177,15 +177,14 @@ async def fastapi_get_all_pump_properties(
包含所有水泵属性的字典列表
"""
# 缓存查询结果提高性能
# global redis_client
results = get_all_pumps(network)
return results
@router.post("/setpumpproperties/", response_model=None, summary="设置水泵属性", description="批量设置指定水泵的多个属性")
async def fastapi_set_pump_properties(
@router.patch("/pumps/properties", response_model=None, summary="设置水泵属性", description="批量设置指定水泵的多个属性")
def fastapi_set_pump_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
pump: str = Query(..., description="水泵ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""
批量设置水泵属性。
@@ -198,6 +197,6 @@ async def fastapi_set_pump_properties(
Returns:
ChangeSet对象,包含本次修改的变更信息
"""
props = await req.json()
props = payload
ps = {"id": pump} | props
return set_pump(network, ChangeSet(ps))
+39 -569
View File
@@ -1,602 +1,72 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_district_metering_area,
add_region,
add_service_area,
add_virtual_district,
# calculate_district_metering_area,
calculate_district_metering_area_for_network,
calculate_district_metering_area_for_nodes,
calculate_district_metering_area_for_region,
# calculate_region,
calculate_service_area,
calculate_virtual_district,
delete_district_metering_area,
delete_region,
delete_service_area,
delete_virtual_district,
generate_district_metering_area,
# generate_region,
generate_service_area,
generate_sub_district_metering_area,
generate_virtual_district,
get_all_district_metering_area_ids,
get_all_district_metering_areas,
# get_all_regions,
get_all_service_areas,
get_all_virtual_districts,
get_district_metering_area,
get_district_metering_area_schema,
get_nodes_in_region,
get_region,
get_region_schema,
get_service_area,
get_service_area_schema,
get_virtual_district,
get_virtual_district_schema,
set_district_metering_area,
get_regions,
set_region,
set_service_area,
set_virtual_district,
)
router = APIRouter()
############################################################
# region 32
############################################################
@router.get(
"/calculateregion/",
summary="计算区域",
description="计算指定水网在指定时间步长的区域分区"
)
async def fastapi_calculate_region(
@router.get("/network-schemas/region", summary="获取区域属性架构")
def get_region_schema_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
time_index: int = Query(..., description="时间步长索引", ge=0)
) -> dict[str, Any]:
"""计算区域分区。"""
return calculate_region(network, time_index)
@router.get(
"/getregionschema/",
summary="获取区域属性架构",
description="获取指定水网的区域属性架构定义"
)
async def fastapi_get_region_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""获取区域的属性架构。"""
return get_region_schema(network)
@router.get(
"/getregion/",
summary="获取区域信息",
description="获取指定ID的区域详细信息"
)
async def fastapi_get_region(
@router.get("/regions", summary="获取区域列表")
def get_regions_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="区域ID")
) -> list[dict[str, Any]]:
return [get_region(network, region_id) for region_id in get_regions(network)]
@router.get("/regions/detail", summary="获取区域信息")
def get_region_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="区域 ID"),
) -> dict[str, Any]:
"""获取区域的详细信息。"""
return get_region(network, id)
@router.post(
"/setregion/",
response_model=None,
summary="设置区域属性",
description="修改指定区域的属性信息"
)
async def fastapi_set_region(
@router.get("/regions/nodes", summary="获取区域节点")
def get_region_nodes_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""设置区域属性。"""
props = await req.json()
return set_region(network, ChangeSet(props))
@router.post(
"/addregion/",
response_model=None,
summary="添加新区域",
description="向水网添加一个新的区域"
)
async def fastapi_add_region(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""添加新的区域。"""
props = await req.json()
return add_region(network, ChangeSet(props))
@router.post(
"/deleteregion/",
response_model=None,
summary="删除区域",
description="删除指定的区域"
)
async def fastapi_delete_region(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""删除区域。"""
props = await req.json()
return delete_region(network, ChangeSet(props))
@router.get(
"/getallregions/",
summary="获取所有区域",
description="获取指定水网中的所有区域信息"
)
async def fastapi_get_all_regions(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""获取所有区域的信息列表。"""
return get_all_regions(network)
@router.post(
"/generateregion/",
response_model=None,
summary="生成区域分区",
description="根据参数自动生成水网的区域分区"
)
async def fastapi_generate_region(
network: str = Query(..., description="管网名称(或数据库名称)"),
inflate_delta: float = Query(..., description="膨胀参数")
) -> ChangeSet:
"""生成区域分区。"""
return generate_region(network, inflate_delta)
############################################################
# district_metering_area 33
############################################################
@router.get(
"/calculatedistrictmeteringarea/",
summary="计算DMA分区",
description="计算指定节点集的区域计量(DMA)分区方案"
)
async def fastapi_calculate_district_metering_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> list[list[str]]:
"""
计算DMA分区。
请求体格式:
{
"nodes": 节点ID列表(list[str]),
"part_count": 分区数量(int),
"part_type": 分区类型(int)
}
"""
props = await req.json()
nodes = props["nodes"]
part_count = props["part_count"]
part_type = props["part_type"]
return calculate_district_metering_area(
network, nodes, part_count, part_type
)
@router.get(
"/calculatedistrictmeteringareaforregion/",
summary="计算区域内DMA分区",
description="为指定区域计算区域计量(DMA)分区方案"
)
async def fastapi_calculate_district_metering_area_for_region(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> list[list[str]]:
"""
计算区域内DMA分区。
请求体格式:
{
"region": 区域ID(str),
"part_count": 分区数量(int),
"part_type": 分区类型(int)
}
"""
props = await req.json()
region = props["region"]
part_count = props["part_count"]
part_type = props["part_type"]
return calculate_district_metering_area_for_region(
network, region, part_count, part_type
)
@router.get(
"/calculatedistrictmeteringareafornetwork/",
summary="计算整网DMA分区",
description="为整个水网计算区域计量(DMA)分区方案"
)
async def fastapi_calculate_district_metering_area_for_network(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> list[list[str]]:
"""
计算整网DMA分区。
请求体格式:
{
"part_count": 分区数量(int),
"part_type": 分区类型(int)
}
"""
props = await req.json()
part_count = props["part_count"]
part_type = props["part_type"]
return calculate_district_metering_area_for_network(network, part_count, part_type)
@router.get(
"/getdistrictmeteringareaschema/",
summary="获取DMA属性架构",
description="获取指定水网的区域计量(DMA)属性架构定义"
)
async def fastapi_get_district_metering_area_schema(
network: str = Query(..., description="管网名称(或数据库名称)"),
) -> dict[str, dict[str, Any]]:
"""获取DMA的属性架构。"""
return get_district_metering_area_schema(network)
@router.get(
"/getdistrictmeteringarea/",
summary="获取DMA信息",
description="获取指定ID的区域计量(DMA)详细信息"
)
async def fastapi_get_district_metering_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="DMA ID")
) -> dict[str, Any]:
"""获取DMA的详细信息。"""
return get_district_metering_area(network, id)
@router.post(
"/setdistrictmeteringarea/",
response_model=None,
summary="设置DMA属性",
description="修改指定DMA的属性信息"
)
async def fastapi_set_district_metering_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""设置DMA属性。"""
props = await req.json()
return set_district_metering_area(network, ChangeSet(props))
@router.post(
"/adddistrictmeteringarea/",
response_model=None,
summary="添加新DMA",
description="向水网添加一个新的区域计量(DMA)"
)
async def fastapi_add_district_metering_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""添加新的DMA。"""
props = await req.json()
# boundary should be [(x,y), (x,y)]
boundary = props.get("boundary", [])
newBoundary = []
for pt in boundary:
if len(pt) >= 2:
newBoundary.append((pt[0], pt[1]))
props["boundary"] = newBoundary
return add_district_metering_area(network, ChangeSet(props))
@router.post(
"/deletedistrictmeteringarea/",
response_model=None,
summary="删除DMA",
description="删除指定的区域计量(DMA)"
)
async def fastapi_delete_district_metering_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""删除DMA。"""
props = await req.json()
return delete_district_metering_area(network, ChangeSet(props))
@router.get(
"/getalldistrictmeteringareaids/",
summary="获取所有DMA ID",
description="获取指定水网中所有DMA的ID列表"
)
async def fastapi_get_all_district_metering_area_ids(
network: str = Query(..., description="管网名称(或数据库名称)")
id: str = Query(..., description="区域 ID"),
) -> list[str]:
"""获取所有DMA的ID列表。"""
return get_all_district_metering_area_ids(network)
return get_nodes_in_region(network, id)
@router.get(
"/getalldistrictmeteringareas/",
summary="获取所有DMA",
description="获取指定水网中所有DMA的详细信息"
)
async def getalldistrictmeteringareas(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""获取所有DMA的详细信息列表。"""
return get_all_district_metering_areas(network)
@router.post(
"/generatedistrictmeteringarea/",
response_model=None,
summary="生成DMA分区",
description="根据参数自动生成水网的DMA分区方案"
)
async def fastapi_generate_district_metering_area(
@router.patch("/regions", summary="修改区域", response_model=None)
def set_region_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
part_count: int = Query(..., description="分区数量", gt=0),
part_type: int = Query(..., description="分区类型"),
inflate_delta: float = Query(..., description="膨胀参数")
payload: dict[str, Any] = Body(...),
) -> ChangeSet:
"""生成DMA分区。"""
return generate_district_metering_area(
network, part_count, part_type, inflate_delta
)
return set_region(network, ChangeSet(payload))
@router.post(
"/generatesubdistrictmeteringarea/",
response_model=None,
summary="生成DMA子分区",
description="为指定DMA生成子DMA分区"
)
async def fastapi_generate_sub_district_metering_area(
@router.post("/regions", summary="添加区域", response_model=None)
def add_region_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
dma: str = Query(..., description="DMA ID"),
part_count: int = Query(..., description="分区数量", gt=0),
part_type: int = Query(..., description="分区类型"),
inflate_delta: float = Query(..., description="膨胀参数")
payload: dict[str, Any] = Body(...),
) -> ChangeSet:
"""生成DMA子分区。"""
return generate_sub_district_metering_area(
network, dma, part_count, part_type, inflate_delta
)
payload["boundary"] = [tuple(point[:2]) for point in payload.get("boundary", [])]
return add_region(network, ChangeSet(payload))
############################################################
# service_area 34
############################################################
@router.get(
"/calculateservicearea/",
summary="计算服务区",
description="计算指定水网在指定时间步长的服务区分区"
)
async def fastapi_calculate_service_area(
@router.delete("/regions", summary="删除区域", response_model=None)
def delete_region_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
time_index: int = Query(..., description="时间步长索引", ge=0)
) -> dict[str, Any]:
"""计算服务区分区。"""
return calculate_service_area(network, time_index)
@router.get(
"/getserviceareaschema/",
summary="获取服务区属性架构",
description="获取指定水网的服务区属性架构定义"
)
async def fastapi_get_service_area_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""获取服务区的属性架构。"""
return get_service_area_schema(network)
@router.get(
"/getservicearea/",
summary="获取服务区信息",
description="获取指定ID的服务区详细信息"
)
async def fastapi_get_service_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="服务区ID")
) -> dict[str, Any]:
"""获取服务区的详细信息。"""
return get_service_area(network, id)
@router.post(
"/setservicearea/",
response_model=None,
summary="设置服务区属性",
description="修改指定服务区的属性信息"
)
async def fastapi_set_service_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...),
) -> ChangeSet:
"""设置服务区属性。"""
props = await req.json()
return set_service_area(network, ChangeSet(props))
@router.post(
"/addservicearea/",
response_model=None,
summary="添加新服务区",
description="向水网添加一个新的服务区"
)
async def fastapi_add_service_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""添加新的服务区。"""
props = await req.json()
return add_service_area(network, ChangeSet(props))
@router.post(
"/deleteservicearea/",
response_model=None,
summary="删除服务区",
description="删除指定的服务区"
)
async def fastapi_delete_service_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""删除服务区。"""
props = await req.json()
return delete_service_area(network, ChangeSet(props))
@router.get(
"/getallserviceareas/",
summary="获取所有服务区",
description="获取指定水网中的所有服务区信息"
)
async def fastapi_get_all_service_areas(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""获取所有服务区的信息列表。"""
return get_all_service_areas(network)
@router.post(
"/generateservicearea/",
response_model=None,
summary="生成服务区分区",
description="根据参数自动生成水网的服务区分区"
)
async def fastapi_generate_service_area(
network: str = Query(..., description="管网名称(或数据库名称)"),
inflate_delta: float = Query(..., description="膨胀参数")
) -> ChangeSet:
"""生成服务区分区。"""
return generate_service_area(network, inflate_delta)
############################################################
# virtual_district 35
############################################################
@router.get(
"/calculatevirtualdistrict/",
summary="计算虚拟分区",
description="根据指定的中心节点计算虚拟分区方案"
)
async def fastapi_calculate_virtual_district(
network: str = Query(..., description="管网名称(或数据库名称)"),
centers: list[str] = Query(..., description="中心节点ID列表")
) -> dict[str, list[Any]]:
"""计算虚拟分区。"""
return calculate_virtual_district(network, centers)
@router.get(
"/getvirtualdistrictschema/",
summary="获取虚拟分区属性架构",
description="获取指定水网的虚拟分区属性架构定义"
)
async def fastapi_get_virtual_district_schema(
network: str = Query(..., description="管网名称(或数据库名称)"),
) -> dict[str, dict[str, Any]]:
"""获取虚拟分区的属性架构。"""
return get_virtual_district_schema(network)
@router.get(
"/getvirtualdistrict/",
summary="获取虚拟分区信息",
description="获取指定ID的虚拟分区详细信息"
)
async def fastapi_get_virtual_district(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="虚拟分区ID")
) -> dict[str, Any]:
"""获取虚拟分区的详细信息。"""
return get_virtual_district(network, id)
@router.post(
"/setvirtualdistrict/",
response_model=None,
summary="设置虚拟分区属性",
description="修改指定虚拟分区的属性信息"
)
async def fastapi_set_virtual_district(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""设置虚拟分区属性。"""
props = await req.json()
return set_virtual_district(network, ChangeSet(props))
@router.post(
"/addvirtualdistrict/",
response_model=None,
summary="添加新虚拟分区",
description="向水网添加一个新的虚拟分区"
)
async def fastapi_add_virtual_district(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""添加新的虚拟分区。"""
props = await req.json()
return add_virtual_district(network, ChangeSet(props))
@router.post(
"/deletevirtualdistrict/",
response_model=None,
summary="删除虚拟分区",
description="删除指定的虚拟分区"
)
async def fastapi_delete_virtual_district(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""删除虚拟分区。"""
props = await req.json()
return delete_virtual_district(network, ChangeSet(props))
@router.get(
"/getallvirtualdistrict/",
summary="获取所有虚拟分区",
description="获取指定水网中的所有虚拟分区信息"
)
async def fastapi_get_all_virtual_district(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""获取所有虚拟分区的信息列表。"""
return get_all_virtual_districts(network)
@router.post(
"/generatevirtualdistrict/",
response_model=None,
summary="生成虚拟分区",
description="根据参数自动生成虚拟分区方案"
)
async def fastapi_generate_virtual_district(
network: str = Query(..., description="管网名称(或数据库名称)"),
inflate_delta: float = Query(..., description="膨胀参数"),
req: Request = None
) -> ChangeSet:
"""生成虚拟分区。"""
props = await req.json()
return generate_virtual_district(network, props["centers"], inflate_delta)
@router.get(
"/calculatedistrictmeteringareafornodes/",
summary="计算节点DMA分区",
description="为指定节点集计算区域计量(DMA)分区方案"
)
async def fastapi_calculate_district_metering_area_for_nodes(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> list[list[str]]:
"""
计算节点DMA分区。
请求体格式:
{
"nodes": 节点ID列表(list[str]),
"part_count": 分区数量(int),
"part_type": 分区类型(int)
}
"""
props = await req.json()
nodes = props["nodes"]
part_count = props["part_count"]
part_type = props["part_type"]
return calculate_district_metering_area_for_nodes(
network, nodes, part_count, part_type
)
return delete_region(network, ChangeSet(payload))
+44 -44
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_reservoir,
delete_reservoir,
@@ -14,11 +14,11 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get(
"/getreservoirschema",
"/network-schemas/reservoir",
summary="获取水库模式",
description="获取指定供水网络中所有水库的模式/属性字段定义"
)
async def fast_get_reservoir_schema(
def fast_get_reservoir_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""
@@ -35,12 +35,12 @@ async def fast_get_reservoir_schema(
return get_reservoir_schema(network)
@router.post(
"/addreservoir/",
"/reservoirs",
response_model=None,
summary="添加水库",
description="在指定供水网络中添加新的水库/水源节点"
)
async def fastapi_add_reservoir(
def fastapi_add_reservoir(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符"),
x: float = Query(..., description="水库的X坐标"),
@@ -65,13 +65,13 @@ async def fastapi_add_reservoir(
ps = {"id": reservoir, "x": x, "y": y, "head": head}
return add_reservoir(network, ChangeSet(ps))
@router.post(
"/deletereservoir/",
@router.delete(
"/reservoirs",
response_model=None,
summary="删除水库",
description="从指定供水网络中删除指定的水库/水源节点"
)
async def fastapi_delete_reservoir(
def fastapi_delete_reservoir(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="要删除的水库的唯一标识符")
) -> ChangeSet:
@@ -91,11 +91,11 @@ async def fastapi_delete_reservoir(
return delete_reservoir(network, ChangeSet(ps))
@router.get(
"/getreservoirhead/",
"/reservoirs/head",
summary="获取水库水头",
description="获取指定水库的供水水头/总水头值"
)
async def fastapi_get_reservoir_head(
def fastapi_get_reservoir_head(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符")
) -> float | None:
@@ -115,11 +115,11 @@ async def fastapi_get_reservoir_head(
return ps["head"]
@router.get(
"/getreservoirpattern/",
"/reservoirs/pattern",
summary="获取水库模式",
description="获取指定水库的运行模式/供水模式"
)
async def fastapi_get_reservoir_pattern(
def fastapi_get_reservoir_pattern(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符")
) -> str | None:
@@ -139,11 +139,11 @@ async def fastapi_get_reservoir_pattern(
return ps["pattern"]
@router.get(
"/getreservoirx/",
"/reservoirs/x",
summary="获取水库X坐标",
description="获取指定水库的X坐标位置"
)
async def fastapi_get_reservoir_x(
def fastapi_get_reservoir_x(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符")
) -> dict[str, float] | None:
@@ -163,11 +163,11 @@ async def fastapi_get_reservoir_x(
return ps["x"]
@router.get(
"/getreservoiry/",
"/reservoirs/y",
summary="获取水库Y坐标",
description="获取指定水库的Y坐标位置"
)
async def fastapi_get_reservoir_y(
def fastapi_get_reservoir_y(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符")
) -> dict[str, float] | None:
@@ -187,11 +187,11 @@ async def fastapi_get_reservoir_y(
return ps["y"]
@router.get(
"/getreservoircoord/",
"/reservoirs/coord",
summary="获取水库坐标",
description="获取指定水库的平面坐标(X和Y坐标)"
)
async def fastapi_get_reservoir_coord(
def fastapi_get_reservoir_coord(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符")
) -> dict[str, float] | None:
@@ -211,13 +211,13 @@ async def fastapi_get_reservoir_coord(
coord = {"id": reservoir, "x": ps["x"], "y": ps["y"]}
return coord
@router.post(
"/setreservoirhead/",
@router.patch(
"/reservoirs/head",
response_model=None,
summary="设置水库水头",
description="更新指定水库的供水水头/总水头值"
)
async def fastapi_set_reservoir_head(
def fastapi_set_reservoir_head(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符"),
head: float = Query(..., description="新的水头值(米)")
@@ -238,13 +238,13 @@ async def fastapi_set_reservoir_head(
ps = {"id": reservoir, "head": head}
return set_reservoir(network, ChangeSet(ps))
@router.post(
"/setreservoirpattern/",
@router.patch(
"/reservoirs/pattern",
response_model=None,
summary="设置水库模式",
description="更新指定水库的运行模式/供水模式"
)
async def fastapi_set_reservoir_pattern(
def fastapi_set_reservoir_pattern(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符"),
pattern: str = Query(..., description="新的运行模式")
@@ -265,13 +265,13 @@ async def fastapi_set_reservoir_pattern(
ps = {"id": reservoir, "pattern": pattern}
return set_reservoir(network, ChangeSet(ps))
@router.post(
"/setreservoirx/",
@router.patch(
"/reservoirs/x",
response_model=None,
summary="设置水库X坐标",
description="更新指定水库的X坐标位置"
)
async def fastapi_set_reservoir_x(
def fastapi_set_reservoir_x(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符"),
x: float = Query(..., description="新的X坐标值")
@@ -292,13 +292,13 @@ async def fastapi_set_reservoir_x(
ps = {"id": reservoir, "x": x}
return set_reservoir(network, ChangeSet(ps))
@router.post(
"/setreservoiry/",
@router.patch(
"/reservoirs/y",
response_model=None,
summary="设置水库Y坐标",
description="更新指定水库的Y坐标位置"
)
async def fastapi_set_reservoir_y(
def fastapi_set_reservoir_y(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符"),
y: float = Query(..., description="新的Y坐标值")
@@ -319,13 +319,13 @@ async def fastapi_set_reservoir_y(
ps = {"id": reservoir, "y": y}
return set_reservoir(network, ChangeSet(ps))
@router.post(
"/setreservoircoord/",
@router.patch(
"/reservoirs/coord",
response_model=None,
summary="设置水库坐标",
description="更新指定水库的平面坐标(X和Y坐标)"
)
async def fastapi_set_reservoir_coord(
def fastapi_set_reservoir_coord(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符"),
x: float = Query(..., description="新的X坐标值"),
@@ -349,11 +349,11 @@ async def fastapi_set_reservoir_coord(
return set_reservoir(network, ChangeSet(ps))
@router.get(
"/getreservoirproperties/",
"/reservoirs/properties",
summary="获取水库属性",
description="获取指定水库的所有属性"
)
async def fastapi_get_reservoir_properties(
def fastapi_get_reservoir_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符")
) -> dict[str, Any]:
@@ -372,11 +372,11 @@ async def fastapi_get_reservoir_properties(
return get_reservoir(network, reservoir)
@router.get(
"/getallreservoirproperties/",
"/reservoirs",
summary="获取所有水库属性",
description="获取指定供水网络中所有水库的属性"
)
async def fastapi_get_all_reservoir_properties(
def fastapi_get_all_reservoir_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
@@ -393,16 +393,16 @@ async def fastapi_get_all_reservoir_properties(
results = get_all_reservoirs(network)
return results
@router.post(
"/setreservoirproperties/",
@router.patch(
"/reservoirs/properties",
response_model=None,
summary="设置水库属性",
description="批量更新指定水库的多个属性"
)
async def fastapi_set_reservoir_properties(
def fastapi_set_reservoir_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
reservoir: str = Query(..., description="水库的唯一标识符"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""
设置水库的多个属性。
@@ -417,6 +417,6 @@ async def fastapi_set_reservoir_properties(
Returns:
包含操作变更集的ChangeSet对象
"""
props = await req.json()
props = payload
ps = {"id": reservoir} | props
return set_reservoir(network, ChangeSet(ps))
+14 -14
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
get_tag,
get_tag_schema,
@@ -16,22 +16,22 @@ router = APIRouter()
############################################################
@router.get(
"/gettagschema/",
"/network-schemas/tag",
summary="获取标签属性架构",
description="获取指定水网的标签(Tag)属性架构定义"
)
async def fastapi_get_tag_schema(
def fastapi_get_tag_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""获取标签的属性架构。"""
return get_tag_schema(network)
@router.get(
"/gettag/",
"/tags/detail",
summary="获取标签信息",
description="获取指定类型和ID的标签信息"
)
async def fastapi_get_tag(
def fastapi_get_tag(
network: str = Query(..., description="管网名称(或数据库名称)"),
t_type: str = Query(..., description="标签类型"),
id: str = Query(..., description="元素ID")
@@ -40,27 +40,27 @@ async def fastapi_get_tag(
return get_tag(network, t_type, id)
@router.get(
"/gettags/",
"/tags",
summary="获取所有标签",
description="获取指定水网中的所有标签信息"
)
async def fastapi_get_tags(
def fastapi_get_tags(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""获取水网中所有标签的列表。"""
tags = get_tags(network)
return tags
@router.post(
"/settag/",
@router.patch(
"/tags",
response_model=None,
summary="设置标签",
description="为指定元素设置或修改标签信息"
)
async def fastapi_set_tag(
def fastapi_set_tag(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""设置标签信息。"""
props = await req.json()
props = payload
return set_tag(network, ChangeSet(props))
+61 -62
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
add_tank,
delete_tank,
@@ -13,8 +13,8 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get("/gettankschema", summary="获取水箱模式", description="获取指定网络的水箱数据结构模式定义")
async def fast_get_tank_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
@router.get("/network-schemas/tank", summary="获取水箱模式", description="获取指定网络的水箱数据结构模式定义")
def fast_get_tank_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[str, Any]]:
"""
获取水箱的数据结构模式。
@@ -26,8 +26,8 @@ async def fast_get_tank_schema(network: str = Query(..., description="管网名
"""
return get_tank_schema(network)
@router.post("/addtank/", summary="新增水箱", description="向指定网络中新增一个水箱", response_model=None)
async def fastapi_add_tank(
@router.post("/tanks", summary="新增水箱", description="向指定网络中新增一个水箱", response_model=None)
def fastapi_add_tank(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
x: float = Query(..., description="X坐标"),
@@ -70,8 +70,8 @@ async def fastapi_add_tank(
}
return add_tank(network, ChangeSet(ps))
@router.post("/deletetank/", summary="删除水箱", description="删除指定网络中的水箱", response_model=None)
async def fastapi_delete_tank(
@router.delete("/tanks", summary="删除水箱", description="删除指定网络中的水箱", response_model=None)
def fastapi_delete_tank(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> ChangeSet:
@@ -88,8 +88,8 @@ async def fastapi_delete_tank(
ps = {"id": tank}
return delete_tank(network, ChangeSet(ps))
@router.get("/gettankelevation/", summary="获取水箱标高", description="获取指定水箱的标高值")
async def fastapi_get_tank_elevation(
@router.get("/tanks/elevation", summary="获取水箱标高", description="获取指定水箱的标高值")
def fastapi_get_tank_elevation(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float | None:
@@ -106,8 +106,8 @@ async def fastapi_get_tank_elevation(
ps = get_tank(network, tank)
return ps["elevation"]
@router.get("/gettankinitlevel/", summary="获取水箱初始水位", description="获取指定水箱的初始水位值")
async def fastapi_get_tank_init_level(
@router.get("/tanks/init-level", summary="获取水箱初始水位", description="获取指定水箱的初始水位值")
def fastapi_get_tank_init_level(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float | None:
@@ -124,8 +124,8 @@ async def fastapi_get_tank_init_level(
ps = get_tank(network, tank)
return ps["init_level"]
@router.get("/gettankminlevel/", summary="获取水箱最小水位", description="获取指定水箱的最小水位值")
async def fastapi_get_tank_min_level(
@router.get("/tanks/min-level", summary="获取水箱最小水位", description="获取指定水箱的最小水位值")
def fastapi_get_tank_min_level(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float | None:
@@ -142,8 +142,8 @@ async def fastapi_get_tank_min_level(
ps = get_tank(network, tank)
return ps["min_level"]
@router.get("/gettankmaxlevel/", summary="获取水箱最大水位", description="获取指定水箱的最大水位值")
async def fastapi_get_tank_max_level(
@router.get("/tanks/max-level", summary="获取水箱最大水位", description="获取指定水箱的最大水位值")
def fastapi_get_tank_max_level(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float | None:
@@ -160,8 +160,8 @@ async def fastapi_get_tank_max_level(
ps = get_tank(network, tank)
return ps["max_level"]
@router.get("/gettankdiameter/", summary="获取水箱直径", description="获取指定水箱的直径值")
async def fastapi_get_tank_diameter(
@router.get("/tanks/diameter", summary="获取水箱直径", description="获取指定水箱的直径值")
def fastapi_get_tank_diameter(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float | None:
@@ -178,8 +178,8 @@ async def fastapi_get_tank_diameter(
ps = get_tank(network, tank)
return ps["diameter"]
@router.get("/gettankminvol/", summary="获取水箱最小体积", description="获取指定水箱的最小体积值")
async def fastapi_get_tank_min_vol(
@router.get("/tanks/min-vol", summary="获取水箱最小体积", description="获取指定水箱的最小体积值")
def fastapi_get_tank_min_vol(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float | None:
@@ -196,8 +196,8 @@ async def fastapi_get_tank_min_vol(
ps = get_tank(network, tank)
return ps["min_vol"]
@router.get("/gettankvolcurve/", summary="获取水箱容积曲线", description="获取指定水箱的容积曲线标识")
async def fastapi_get_tank_vol_curve(
@router.get("/tanks/vol-curve", summary="获取水箱容积曲线", description="获取指定水箱的容积曲线标识")
def fastapi_get_tank_vol_curve(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> str | None:
@@ -214,8 +214,8 @@ async def fastapi_get_tank_vol_curve(
ps = get_tank(network, tank)
return ps["vol_curve"]
@router.get("/gettankoverflow/", summary="获取水箱溢流口", description="获取指定水箱的溢流口配置")
async def fastapi_get_tank_overflow(
@router.get("/tanks/overflow", summary="获取水箱溢流口", description="获取指定水箱的溢流口配置")
def fastapi_get_tank_overflow(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> str | None:
@@ -232,8 +232,8 @@ async def fastapi_get_tank_overflow(
ps = get_tank(network, tank)
return ps["overflow"]
@router.get("/gettankx/", summary="获取水箱X坐标", description="获取指定水箱的X坐标值")
async def fastapi_get_tank_x(
@router.get("/tanks/x", summary="获取水箱X坐标", description="获取指定水箱的X坐标值")
def fastapi_get_tank_x(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float:
@@ -250,8 +250,8 @@ async def fastapi_get_tank_x(
ps = get_tank(network, tank)
return ps["x"]
@router.get("/gettanky/", summary="获取水箱Y坐标", description="获取指定水箱的Y坐标值")
async def fastapi_get_tank_y(
@router.get("/tanks/y", summary="获取水箱Y坐标", description="获取指定水箱的Y坐标值")
def fastapi_get_tank_y(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> float:
@@ -268,8 +268,8 @@ async def fastapi_get_tank_y(
ps = get_tank(network, tank)
return ps["y"]
@router.get("/gettankcoord/", summary="获取水箱坐标", description="获取指定水箱的X和Y坐标")
async def fastapi_get_tank_coord(
@router.get("/tanks/coord", summary="获取水箱坐标", description="获取指定水箱的X和Y坐标")
def fastapi_get_tank_coord(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> dict[str, float]:
@@ -287,8 +287,8 @@ async def fastapi_get_tank_coord(
coord = {"x": ps["x"], "y": ps["y"]}
return coord
@router.post("/settankelevation/", summary="设置水箱标高", description="设置指定水箱的标高值", response_model=None)
async def fastapi_set_tank_elevation(
@router.patch("/tanks/elevation", summary="设置水箱标高", description="设置指定水箱的标高值", response_model=None)
def fastapi_set_tank_elevation(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
elevation: float = Query(..., description="新的标高值")
@@ -307,8 +307,8 @@ async def fastapi_set_tank_elevation(
ps = {"id": tank, "elevation": elevation}
return set_tank(network, ChangeSet(ps))
@router.post("/settankinitlevel/", summary="设置水箱初始水位", description="设置指定水箱的初始水位值", response_model=None)
async def fastapi_set_tank_init_level(
@router.patch("/tanks/init-level", summary="设置水箱初始水位", description="设置指定水箱的初始水位值", response_model=None)
def fastapi_set_tank_init_level(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
init_level: float = Query(..., description="新的初始水位值")
@@ -327,8 +327,8 @@ async def fastapi_set_tank_init_level(
ps = {"id": tank, "init_level": init_level}
return set_tank(network, ChangeSet(ps))
@router.post("/settankminlevel/", summary="设置水箱最小水位", description="设置指定水箱的最小水位值", response_model=None)
async def fastapi_set_tank_min_level(
@router.patch("/tanks/min-level", summary="设置水箱最小水位", description="设置指定水箱的最小水位值", response_model=None)
def fastapi_set_tank_min_level(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
min_level: float = Query(..., description="新的最小水位值")
@@ -347,8 +347,8 @@ async def fastapi_set_tank_min_level(
ps = {"id": tank, "min_level": min_level}
return set_tank(network, ChangeSet(ps))
@router.post("/settankmaxlevel/", summary="设置水箱最大水位", description="设置指定水箱的最大水位值", response_model=None)
async def fastapi_set_tank_max_level(
@router.patch("/tanks/max-level", summary="设置水箱最大水位", description="设置指定水箱的最大水位值", response_model=None)
def fastapi_set_tank_max_level(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
max_level: float = Query(..., description="新的最大水位值")
@@ -367,8 +367,8 @@ async def fastapi_set_tank_max_level(
ps = {"id": tank, "max_level": max_level}
return set_tank(network, ChangeSet(ps))
@router.post("/settankdiameter/", summary="设置水箱直径", description="设置指定水箱的直径值", response_model=None)
async def fastapi_set_tank_diameter(
@router.patch("/tanks/diameter", summary="设置水箱直径", description="设置指定水箱的直径值", response_model=None)
def fastapi_set_tank_diameter(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
diameter: float = Query(..., description="新的直径值")
@@ -387,8 +387,8 @@ async def fastapi_set_tank_diameter(
ps = {"id": tank, "diameter": diameter}
return set_tank(network, ChangeSet(ps))
@router.post("/settankminvol/", summary="设置水箱最小体积", description="设置指定水箱的最小体积值", response_model=None)
async def fastapi_set_tank_min_vol(
@router.patch("/tanks/min-vol", summary="设置水箱最小体积", description="设置指定水箱的最小体积值", response_model=None)
def fastapi_set_tank_min_vol(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
min_vol: float = Query(..., description="新的最小体积值")
@@ -407,8 +407,8 @@ async def fastapi_set_tank_min_vol(
ps = {"id": tank, "min_vol": min_vol}
return set_tank(network, ChangeSet(ps))
@router.post("/settankvolcurve/", summary="设置水箱容积曲线", description="设置指定水箱的容积曲线标识", response_model=None)
async def fastapi_set_tank_vol_curve(
@router.patch("/tanks/vol-curve", summary="设置水箱容积曲线", description="设置指定水箱的容积曲线标识", response_model=None)
def fastapi_set_tank_vol_curve(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
vol_curve: str = Query(..., description="新的容积曲线标识")
@@ -427,8 +427,8 @@ async def fastapi_set_tank_vol_curve(
ps = {"id": tank, "vol_curve": vol_curve}
return set_tank(network, ChangeSet(ps))
@router.post("/settankoverflow/", summary="设置水箱溢流口", description="设置指定水箱的溢流口配置", response_model=None)
async def fastapi_set_tank_overflow(
@router.patch("/tanks/overflow", summary="设置水箱溢流口", description="设置指定水箱的溢流口配置", response_model=None)
def fastapi_set_tank_overflow(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
overflow: str = Query(..., description="新的溢流口配置")
@@ -447,8 +447,8 @@ async def fastapi_set_tank_overflow(
ps = {"id": tank, "overflow": overflow}
return set_tank(network, ChangeSet(ps))
@router.post("/settankx/", summary="设置水箱X坐标", description="设置指定水箱的X坐标值", response_model=None)
async def fastapi_set_tank_x(
@router.patch("/tanks/x", summary="设置水箱X坐标", description="设置指定水箱的X坐标值", response_model=None)
def fastapi_set_tank_x(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
x: float = Query(..., description="新的X坐标值")
@@ -467,8 +467,8 @@ async def fastapi_set_tank_x(
ps = {"id": tank, "x": x}
return set_tank(network, ChangeSet(ps))
@router.post("/settanky/", summary="设置水箱Y坐标", description="设置指定水箱的Y坐标值", response_model=None)
async def fastapi_set_tank_y(
@router.patch("/tanks/y", summary="设置水箱Y坐标", description="设置指定水箱的Y坐标值", response_model=None)
def fastapi_set_tank_y(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
y: float = Query(..., description="新的Y坐标值")
@@ -487,8 +487,8 @@ async def fastapi_set_tank_y(
ps = {"id": tank, "y": y}
return set_tank(network, ChangeSet(ps))
@router.post("/settankcoord/", summary="设置水箱坐标", description="设置指定水箱的X和Y坐标", response_model=None)
async def fastapi_set_tank_coord(
@router.patch("/tanks/coord", summary="设置水箱坐标", description="设置指定水箱的X和Y坐标", response_model=None)
def fastapi_set_tank_coord(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
x: float = Query(..., description="新的X坐标值"),
@@ -509,8 +509,8 @@ async def fastapi_set_tank_coord(
ps = {"id": tank, "x": x, "y": y}
return set_tank(network, ChangeSet(ps))
@router.get("/gettankproperties/", summary="获取水箱属性", description="获取指定水箱的所有属性")
async def fastapi_get_tank_properties(
@router.get("/tanks/properties", summary="获取水箱属性", description="获取指定水箱的所有属性")
def fastapi_get_tank_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID")
) -> dict[str, Any]:
@@ -526,8 +526,8 @@ async def fastapi_get_tank_properties(
"""
return get_tank(network, tank)
@router.get("/getalltankproperties/", summary="获取所有水箱属性", description="获取指定网络中所有水箱的属性")
async def fastapi_get_all_tank_properties(
@router.get("/tanks", summary="获取所有水箱属性", description="获取指定网络中所有水箱的属性")
def fastapi_get_all_tank_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
@@ -540,15 +540,14 @@ async def fastapi_get_all_tank_properties(
包含所有水箱属性的字典列表
"""
# 缓存查询结果提高性能
# global redis_client
results = get_all_tanks(network)
return results
@router.post("/settankproperties/", summary="设置水箱属性", description="批量设置指定水箱的多个属性", response_model=None)
async def fastapi_set_tank_properties(
@router.patch("/tanks/properties", summary="设置水箱属性", description="批量设置指定水箱的多个属性", response_model=None)
def fastapi_set_tank_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
tank: str = Query(..., description="水箱ID"),
req: Request = None
payload: dict[str, Any] = Body(...)
) -> ChangeSet:
"""
批量设置水箱的属性。
@@ -561,6 +560,6 @@ async def fastapi_set_tank_properties(
Returns:
包含变更信息的ChangeSet对象
"""
props = await req.json()
props = payload
ps = {"id": tank} | props
return set_tank(network, ChangeSet(ps))
+46 -47
View File
@@ -1,7 +1,7 @@
from fastapi import APIRouter, Request, Query, Path, Body
from typing import Any, List, Dict, Union
from typing import Any
from fastapi import APIRouter, Body, Query
from app.services.tjnetwork import (
Any,
ChangeSet,
VALVES_TYPE_PRV,
add_valve,
@@ -15,11 +15,11 @@ from app.services.tjnetwork import (
router = APIRouter()
@router.get(
"/getvalveschema",
"/network-schemas/valve",
summary="获取阀门架构",
description="获取指定水网中所有阀门的架构和字段定义",
)
async def fastapi_get_valve_schema(
def fastapi_get_valve_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""
@@ -30,12 +30,12 @@ async def fastapi_get_valve_schema(
return get_valve_schema(network)
@router.post(
"/addvalve/",
"/valves",
response_model=None,
summary="添加阀门",
description="在指定的水网中添加新的阀门",
)
async def fastapi_add_valve(
def fastapi_add_valve(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
node1: str = Query(..., description="起点节点ID"),
@@ -62,13 +62,13 @@ async def fastapi_add_valve(
return add_valve(network, ChangeSet(ps))
@router.post(
"/deletevalve/",
@router.delete(
"/valves",
response_model=None,
summary="删除阀门",
description="从指定的水网中删除指定的阀门",
)
async def fastapi_delete_valve(
def fastapi_delete_valve(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> ChangeSet:
@@ -81,11 +81,11 @@ async def fastapi_delete_valve(
return delete_valve(network, ChangeSet(ps))
@router.get(
"/getvalvenode1/",
"/valves/node1",
summary="获取阀门起点节点",
description="获取指定阀门连接的起点节点ID",
)
async def fastapi_get_valve_node1(
def fastapi_get_valve_node1(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> str | None:
@@ -98,11 +98,11 @@ async def fastapi_get_valve_node1(
return ps["node1"]
@router.get(
"/getvalvenode2/",
"/valves/node2",
summary="获取阀门终点节点",
description="获取指定阀门连接的终点节点ID",
)
async def fastapi_get_valve_node2(
def fastapi_get_valve_node2(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> str | None:
@@ -115,11 +115,11 @@ async def fastapi_get_valve_node2(
return ps["node2"]
@router.get(
"/getvalvediameter/",
"/valves/diameter",
summary="获取阀门直径",
description="获取指定阀门的直径",
)
async def fastapi_get_valve_diameter(
def fastapi_get_valve_diameter(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> float | None:
@@ -132,11 +132,11 @@ async def fastapi_get_valve_diameter(
return ps["diameter"]
@router.get(
"/getvalvetype/",
"/valves/type",
summary="获取阀门类型",
description="获取指定阀门的类型",
)
async def fastapi_get_valve_type(
def fastapi_get_valve_type(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> str | None:
@@ -149,11 +149,11 @@ async def fastapi_get_valve_type(
return ps["type"]
@router.get(
"/getvalvesetting/",
"/valves/setting",
summary="获取阀门开度",
description="获取指定阀门的开度/设置值",
)
async def fastapi_get_valve_setting(
def fastapi_get_valve_setting(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> float | None:
@@ -166,11 +166,11 @@ async def fastapi_get_valve_setting(
return ps["setting"]
@router.get(
"/getvalveminorloss/",
"/valves/minor-loss",
summary="获取阀门损失系数",
description="获取指定阀门的损失系数",
)
async def fastapi_get_valve_minor_loss(
def fastapi_get_valve_minor_loss(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> float | None:
@@ -182,13 +182,13 @@ async def fastapi_get_valve_minor_loss(
ps = get_valve(network, valve)
return ps["minor_loss"]
@router.post(
"/setvalvenode1/",
@router.patch(
"/valves/node1",
response_model=None,
summary="设置阀门起点节点",
description="设置指定阀门的起点节点",
)
async def fastapi_set_valve_node1(
def fastapi_set_valve_node1(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
node1: str = Query(..., description="新的起点节点ID"),
@@ -201,13 +201,13 @@ async def fastapi_set_valve_node1(
ps = {"id": valve, "node1": node1}
return set_valve(network, ChangeSet(ps))
@router.post(
"/setvalvenode2/",
@router.patch(
"/valves/node2",
response_model=None,
summary="设置阀门终点节点",
description="设置指定阀门的终点节点",
)
async def fastapi_set_valve_node2(
def fastapi_set_valve_node2(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
node2: str = Query(..., description="新的终点节点ID"),
@@ -220,13 +220,13 @@ async def fastapi_set_valve_node2(
ps = {"id": valve, "node2": node2}
return set_valve(network, ChangeSet(ps))
@router.post(
"/setvalvenodediameter/",
@router.patch(
"/valves/diameter",
response_model=None,
summary="设置阀门直径",
description="设置指定阀门的直径",
)
async def fastapi_set_valve_diameter(
def fastapi_set_valve_diameter(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
diameter: float = Query(..., description="新的直径值(mm"),
@@ -239,13 +239,13 @@ async def fastapi_set_valve_diameter(
ps = {"id": valve, "diameter": diameter}
return set_valve(network, ChangeSet(ps))
@router.post(
"/setvalvetype/",
@router.patch(
"/valves/type",
response_model=None,
summary="设置阀门类型",
description="设置指定阀门的类型",
)
async def fastapi_set_valve_type(
def fastapi_set_valve_type(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
type: str = Query(..., description="新的阀门类型"),
@@ -258,13 +258,13 @@ async def fastapi_set_valve_type(
ps = {"id": valve, "type": type}
return set_valve(network, ChangeSet(ps))
@router.post(
"/setvalvesetting/",
@router.patch(
"/valves/setting",
response_model=None,
summary="设置阀门开度",
description="设置指定阀门的开度/设置值",
)
async def fastapi_set_valve_setting(
def fastapi_set_valve_setting(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
setting: float = Query(..., description="新的开度值"),
@@ -278,11 +278,11 @@ async def fastapi_set_valve_setting(
return set_valve(network, ChangeSet(ps))
@router.get(
"/getvalveproperties/",
"/valves/properties",
summary="获取阀门所有属性",
description="获取指定阀门的所有属性",
)
async def fastapi_get_valve_properties(
def fastapi_get_valve_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
) -> dict[str, Any]:
@@ -294,11 +294,11 @@ async def fastapi_get_valve_properties(
return get_valve(network, valve)
@router.get(
"/getallvalveproperties/",
"/valves",
summary="获取所有阀门属性",
description="获取指定水网中所有阀门的属性",
)
async def fastapi_get_all_valve_properties(
def fastapi_get_all_valve_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
@@ -307,26 +307,25 @@ async def fastapi_get_all_valve_properties(
返回指定水网中所有阀门的完整属性列表。
"""
# 缓存查询结果提高性能
# global redis_client
results = get_all_valves(network)
return results
@router.post(
"/setvalveproperties/",
@router.patch(
"/valves/properties",
response_model=None,
summary="批量设置阀门属性",
description="批量设置指定阀门的多个属性",
)
async def fastapi_set_valve_properties(
def fastapi_set_valve_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
valve: str = Query(..., description="阀门ID"),
req: Request = None,
payload: dict[str, Any] = Body(...),
) -> ChangeSet:
"""
批量设置阀门的属性。
更新指定阀门的一个或多个属性,通过JSON请求体传递要更新的属性。
"""
props = await req.json()
props = payload
ps = {"id": valve} | props
return set_valve(network, ChangeSet(ps))
+10 -540
View File
@@ -1,48 +1,20 @@
import json
from fastapi import APIRouter, Request, HTTPException, Query, Path, Body, Depends
from fastapi.responses import PlainTextResponse
from typing import Any, Dict, List
from fastapi import APIRouter, HTTPException, Query, Depends
from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository
from app.auth.project_dependencies import get_metadata_repository
from app.domain.schemas.metadata import ProjectMetaResponse, GeoServerConfigResponse
import app.services.project_info as project_info
from app.infra.db.postgresql.database import get_database_instance as get_pg_db
from app.infra.db.timescaledb.database import get_database_instance as get_ts_db
from app.auth.project_dependencies import (
get_metadata_repository,
)
from app.domain.schemas.metadata import ProjectMetaResponse
from app.services.tjnetwork import (
ChangeSet,
list_project,
have_project,
create_project,
delete_project,
is_project_open,
open_project,
close_project,
copy_project,
import_inp,
export_inp,
read_inp,
dump_inp,
get_all_vertices,
get_all_scada_elements,
get_all_district_metering_areas,
get_all_service_areas,
get_all_virtual_districts,
get_extension_data,
convert_inp_v3_to_v2,
get_all_scada_info,
)
# For inp file upload/download
import os
from fastapi import Response, status
from fastapi.responses import FileResponse
inpDir = "data/" # Assuming data directory exists or is defined somewhere.
# In main.py it was likely global. For safety, let's use a relative path or get from config.
# But let's stick to what main.py probably used or a default.
router = APIRouter()
lockedPrjs: Dict[str, str] = {}
@router.get("/project_info/", summary="获取项目信息", description="从数据库获取项目的详细信息,包括地图范围等。", response_model=ProjectMetaResponse)
@router.get("/projects/current", summary="获取项目信息", description="从数据库获取项目的详细信息,包括地图范围等。", response_model=ProjectMetaResponse)
async def get_project_info_endpoint(
network: str = Query(..., description="管网名称(或项目代码)"),
metadata_repo: MetadataRepository = Depends(get_metadata_repository),
@@ -55,17 +27,6 @@ async def get_project_info_endpoint(
project_detail = await metadata_repo.get_project_detail_by_code(network)
if not project_detail:
raise HTTPException(status_code=404, detail=f"Project {network} not found")
geoserver_payload = None
if project_detail.geoserver:
geoserver_payload = GeoServerConfigResponse(
gs_base_url=project_detail.geoserver.gs_base_url,
gs_admin_user=project_detail.geoserver.gs_admin_user,
gs_datastore_name=project_detail.geoserver.gs_datastore_name,
default_extent=project_detail.geoserver.default_extent,
srid=project_detail.geoserver.srid,
)
return ProjectMetaResponse(
project_id=project_detail.project_id,
name=project_detail.name,
@@ -75,144 +36,10 @@ async def get_project_info_endpoint(
map_extent=project_detail.map_extent,
status=project_detail.status,
project_role="viewer", # Default role for public access
geoserver=geoserver_payload
)
@router.get("/listprojects/", summary="获取项目列表", description="获取服务器上所有可用的供水管网项目名称列表")
async def list_projects_endpoint() -> list[str]:
"""
获取项目列表
返回所有已创建项目的名称列表。
"""
return list_project()
@router.get("/haveproject/", summary="检查项目是否存在", description="检查指定名称的项目是否存在。")
async def have_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
):
"""
检查项目是否存在
- **network**: 管网名称(或数据库名称)
"""
return have_project(network)
@router.post("/createproject/", summary="创建新项目", description="创建一个新的供水管网项目。如果项目已存在,可能会覆盖或报错(取决于底层实现)。")
async def create_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
):
"""
创建新项目
- **network**: 管网名称(或数据库名称)
"""
create_project(network)
return network
@router.post("/deleteproject/", summary="删除项目", description="永久删除指定的供水管网项目。此操作不可恢复。")
async def delete_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
):
"""
删除项目
- **network**: 管网名称(或数据库名称)
"""
delete_project(network)
return True
@router.get("/isprojectopen/", summary="检查项目是否已打开", description="检查指定项目是否已被加载到内存中。")
async def is_project_open_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
):
"""
检查项目是否已打开
- **network**: 管网名称(或数据库名称)
"""
return is_project_open(network)
@router.post("/openproject/", summary="打开项目", description="将指定项目加载到内存中,并初始化数据库连接池。")
async def open_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
):
"""
打开项目
- **network**: 管网名称(或数据库名称)
"""
open_project(network)
# 尝试连接指定数据库
try:
# 初始化 PostgreSQL 连接池
pg_instance = await get_pg_db(network)
async with pg_instance.get_connection() as conn:
async with conn.cursor() as cur:
await cur.execute("SELECT 1")
# 初始化 TimescaleDB 连接池
ts_instance = await get_ts_db(network)
async with ts_instance.get_connection() as conn:
async with conn.cursor() as cur:
await cur.execute("SELECT 1")
except Exception as e:
# 记录错误但不阻断项目打开,或者根据需求决定是否阻断
# 这里选择打印错误,因为 open_project 原本只负责原生部分
print(f"Failed to connect to databases for {network}: {str(e)}")
# 如果数据库连接是必须的,可以抛出异常:
# raise HTTPException(status_code=500, detail=f"Database connection failed: {str(e)}")
return network
@router.post("/closeproject/", summary="关闭项目", description="将指定项目从内存中卸载,释放资源。")
async def close_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)")
):
"""
关闭项目
- **network**: 管网名称(或数据库名称)
"""
close_project(network)
return True
@router.post("/copyproject/", summary="复制项目", description="将现有项目复制为新项目。")
async def copy_project_endpoint(
source: str = Query(..., description="管网名称(或数据库名称)"),
target: str = Query(..., description="管网名称(或数据库名称)")
):
"""
复制项目
- **source**: 管网名称(或数据库名称)
- **target**: 管网名称(或数据库名称)
"""
copy_project(source, target)
return True
@router.post("/importinp/", summary="导入 INP 文件内容", description="将 INP 格式的文本内容导入到指定项目中。")
async def import_inp_endpoint(
req: Request,
network: str = Query(..., description="管网名称(或数据库名称)")
):
"""
导入 INP 文件内容
- **network**: 管网名称(或数据库名称)
- **req**: 请求体,需包含 `{"inp": "..."}` 结构
"""
jo_root = await req.json()
inp_text = jo_root["inp"]
ps = {"inp": inp_text}
ret = import_inp(network, ChangeSet(ps))
print(ret)
return ret
@router.get("/exportinp/", response_model=None, summary="导出项目为 ChangeSet", description="导出项目的变更集 (ChangeSet),包含顶点、SCADA 元素、DMA、SA、VD 等信息。")
async def export_inp_endpoint(
@router.get("/projects/current/exports/change-set", response_model=None, summary="导出项目为 ChangeSet", description="导出项目的变更集 (ChangeSet),包含顶点、SCADA 元素、DMA、SA、VD 等信息")
def export_inp_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
version: str = Query(..., description="版本号 (通常用于增量更新)")
) -> ChangeSet:
@@ -224,364 +51,7 @@ async def export_inp_endpoint(
"""
cs = export_inp(network, version)
op = cs.operations[0]
open_project(network)
op["vertex"] = json.dumps(get_all_vertices(network))
op["scada"] = json.dumps(get_all_scada_elements(network))
op["dma"] = json.dumps(get_all_district_metering_areas(network))
op["sa"] = json.dumps(get_all_service_areas(network))
op["vd"] = json.dumps(get_all_virtual_districts(network))
op["legend"] = get_extension_data(network, "legend")
db = get_extension_data(network, "scada_db")
print(db)
scada_db = ""
if db:
scada_db = db
print(scada_db)
op["scada_db"] = scada_db
close_project(network)
return cs
@router.post("/readinp/", summary="读取 INP 文件到项目", description="从服务器文件系统中读取指定的 INP 文件并加载到项目中。")
async def read_inp_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
inp: str = Query(..., description="INP 文件名 (不包含路径)")
) -> bool:
"""
读取 INP 文件到项目
- **network**: 管网名称(或数据库名称)
- **inp**: INP 文件名
"""
read_inp(network, inp)
return True
@router.get("/dumpinp/", summary="导出项目到 INP 文件", description="将项目当前状态保存为 INP 文件到服务器文件系统。")
async def dump_inp_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
inp: str = Query(..., description="目标文件名")
) -> bool:
"""
导出项目到 INP 文件
- **network**: 管网名称(或数据库名称)
- **inp**: 目标文件名
"""
dump_inp(network, inp)
return True
@router.get("/isprojectlocked/", summary="检查项目是否被锁定", description="检查指定项目是否处于锁定状态。")
async def is_project_locked_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
检查项目是否被锁定
- **network**: 管网名称(或数据库名称)
"""
return network in lockedPrjs.keys()
@router.get("/isprojectlockedbyme/", summary="检查项目是否被当前用户锁定", description="检查指定项目是否被当前客户端 (IP) 锁定。")
async def is_project_locked_by_me_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
检查项目是否被当前用户锁定
- **network**: 管网名称(或数据库名称)
"""
client_host = req.client.host
return lockedPrjs.get(network) == client_host
# 0 successfully locked
# 1 already locked by you
# 2 locked by others
@router.post("/lockproject/", summary="锁定项目", description="锁定指定项目以防止并发修改。")
async def lock_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
锁定项目
返回值:
- **0**: 锁定成功
- **1**: 已被当前用户锁定
- **2**: 已被其他用户锁定
"""
client_host = req.client.host
if not network in lockedPrjs.keys():
lockedPrjs[network] = client_host
return 0
else:
if lockedPrjs.get(network) == client_host:
return 1
else:
return 2
@router.post("/unlockproject/", summary="解锁项目", description="释放对项目的锁定。")
def unlock_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
解锁项目
只有锁定者才能解锁。
"""
client_host = req.client.host
if lockedPrjs.get(network) == client_host:
print("delete key")
del lockedPrjs[network]
return True
return False
# inp file operations
@router.post("/uploadinp/", status_code=status.HTTP_200_OK, summary="上传 INP 文件", description="上传 INP 文件到服务器数据目录。")
async def fastapi_upload_inp(
afile: bytes = Body(..., description="文件二进制内容"),
name: str = Query(..., description="保存的文件名")
):
"""
上传 INP 文件
- **afile**: 文件内容
- **name**: 文件名
"""
if not os.path.exists(inpDir):
os.makedirs(inpDir, exist_ok=True)
filePath = inpDir + str(name)
with open(filePath, "wb") as f:
f.write(afile)
return True
@router.get("/downloadinp/", status_code=status.HTTP_200_OK, summary="下载 INP 文件", description="从服务器数据目录下载指定的 INP 文件。")
async def fastapi_download_inp(
name: str = Query(..., description="文件名"),
response: Response = None
):
"""
下载 INP 文件
- **name**: 文件名
"""
filePath = inpDir + name
if os.path.exists(filePath):
return FileResponse(
filePath, media_type="application/octet-stream", filename="inp.inp"
)
else:
response.status_code = status.HTTP_400_BAD_REQUEST
return True
# DingZQ, 2024-12-28, convert v3 to v2
@router.get("/convertv3tov2/", response_model=None, summary="转换 INP V3 为 V2", description="将 EPANET 3.0 格式的 INP 内容转换为 2.x 格式。")
async def fastapi_convert_v3_to_v2(
req: Request
) -> ChangeSet:
"""
转换 INP V3 为 V2
- **req**: 请求体,需包含 `{"inp": "..."}` 结构
"""
network = "v3Tov2"
jo_root = await req.json()
inp = jo_root["inp"]
cs = convert_inp_v3_to_v2(inp)
op = cs.operations[0]
open_project(network)
op["vertex"] = json.dumps(get_all_vertices(network))
op["scada"] = json.dumps(get_all_scada_elements(network))
op["dma"] = json.dumps(get_all_district_metering_areas(network))
op["sa"] = json.dumps(get_all_service_areas(network))
op["vd"] = json.dumps(get_all_virtual_districts(network))
op["legend"] = get_extension_data(network, "legend")
db = get_extension_data(network, "scada_db")
print(db)
scada_db = ""
if db:
scada_db = db
print(scada_db)
op["scada_db"] = scada_db
close_project(network)
return cs
@router.post("/readinp/", summary="读取 INP 文件到项目", description="从服务器文件系统中读取指定的 INP 文件并加载到项目中。")
async def read_inp_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
inp: str = Query(..., description="INP 文件名 (不包含路径)")
) -> bool:
"""
读取 INP 文件到项目
- **network**: 管网名称(或数据库名称)
- **inp**: INP 文件名
"""
read_inp(network, inp)
return True
@router.get("/dumpinp/", summary="导出项目到 INP 文件", description="将项目当前状态保存为 INP 文件到服务器文件系统。")
async def dump_inp_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
inp: str = Query(..., description="目标文件名")
) -> bool:
"""
导出项目到 INP 文件
- **network**: 管网名称(或数据库名称)
- **inp**: 目标文件名
"""
dump_inp(network, inp)
return True
@router.get("/isprojectlocked/", summary="检查项目是否被锁定", description="检查指定项目是否处于锁定状态。")
async def is_project_locked_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
检查项目是否被锁定
- **network**: 管网名称(或数据库名称)
"""
return network in lockedPrjs.keys()
@router.get("/isprojectlockedbyme/", summary="检查项目是否被当前用户锁定", description="检查指定项目是否被当前客户端 (IP) 锁定。")
async def is_project_locked_by_me_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
检查项目是否被当前用户锁定
- **network**: 管网名称(或数据库名称)
"""
client_host = req.client.host
return lockedPrjs.get(network) == client_host
# 0 successfully locked
# 1 already locked by you
# 2 locked by others
@router.post("/lockproject/", summary="锁定项目", description="锁定指定项目以防止并发修改。")
async def lock_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
锁定项目
返回值:
- **0**: 锁定成功
- **1**: 已被当前用户锁定
- **2**: 已被其他用户锁定
"""
client_host = req.client.host
if not network in lockedPrjs.keys():
lockedPrjs[network] = client_host
return 0
else:
if lockedPrjs.get(network) == client_host:
return 1
else:
return 2
@router.post("/unlockproject/", summary="解锁项目", description="释放对项目的锁定。")
def unlock_project_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
):
"""
解锁项目
只有锁定者才能解锁。
"""
client_host = req.client.host
if lockedPrjs.get(network) == client_host:
print("delete key")
del lockedPrjs[network]
return True
return False
# inp file operations
@router.post("/uploadinp/", status_code=status.HTTP_200_OK, summary="上传 INP 文件", description="上传 INP 文件到服务器数据目录。")
async def fastapi_upload_inp(
afile: bytes = Body(..., description="文件二进制内容"),
name: str = Query(..., description="保存的文件名")
):
"""
上传 INP 文件
- **afile**: 文件内容
- **name**: 文件名
"""
if not os.path.exists(inpDir):
os.makedirs(inpDir, exist_ok=True)
filePath = inpDir + str(name)
with open(filePath, "wb") as f:
f.write(afile)
return True
@router.get("/downloadinp/", status_code=status.HTTP_200_OK, summary="下载 INP 文件", description="从服务器数据目录下载指定的 INP 文件。")
async def fastapi_download_inp(
name: str = Query(..., description="文件名"),
response: Response = None
):
"""
下载 INP 文件
- **name**: 文件名
"""
filePath = inpDir + name
if os.path.exists(filePath):
return FileResponse(
filePath, media_type="application/octet-stream", filename="inp.inp"
)
else:
response.status_code = status.HTTP_400_BAD_REQUEST
return True
# DingZQ, 2024-12-28, convert v3 to v2
@router.get("/convertv3tov2/", response_model=None, summary="转换 INP V3 为 V2", description="将 EPANET 3.0 格式的 INP 内容转换为 2.x 格式。")
async def fastapi_convert_v3_to_v2(
req: Request
) -> ChangeSet:
"""
转换 INP V3 为 V2
- **req**: 请求体,需包含 `{"inp": "..."}` 结构
"""
network = "v3Tov2"
jo_root = await req.json()
inp = jo_root["inp"]
cs = convert_inp_v3_to_v2(inp)
op = cs.operations[0]
open_project(network)
op["vertex"] = json.dumps(get_all_vertices(network))
op["scada"] = json.dumps(get_all_scada_elements(network))
op["dma"] = json.dumps(get_all_district_metering_areas(network))
op["sa"] = json.dumps(get_all_service_areas(network))
op["vd"] = json.dumps(get_all_virtual_districts(network))
op["legend"] = get_extension_data(network, "legend")
db = get_extension_data(network, "scada_db")
print(db)
scada_db = ""
if db:
scada_db = db
print(scada_db)
op["scada_db"] = scada_db
close_project(network)
op["scada"] = json.dumps(get_all_scada_info(network))
return cs
+25 -41
View File
@@ -1,10 +1,10 @@
from fastapi import APIRouter, Depends, HTTPException, Path, Query
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Query
from psycopg import AsyncConnection
import app.native.wndb as wndb
from app.infra.db.postgresql.scheme import SchemeRepository
from app.infra.db.postgresql.analysis import AnalysisRepository
from app.auth.project_dependencies import get_project_pg_connection
from app.services import project_info
router = APIRouter()
@@ -16,28 +16,8 @@ async def get_database_connection(
yield conn
@router.get("/scada-info", summary="获取SCADA信息", description="使用连接池查询所有SCADA信息")
async def get_scada_info_with_connection(
conn: AsyncConnection = Depends(get_database_connection),
):
"""
获取所有SCADA信息
返回项目中所有的SCADA设备信息
"""
try:
_ = conn
network_name = project_info.name
scada_data = wndb.get_all_scada_info(network_name) if network_name else []
return {"success": True, "data": scada_data, "count": len(scada_data)}
except Exception as e:
raise HTTPException(
status_code=500, detail=f"查询SCADA信息时发生错误: {str(e)}"
)
@router.get("/scheme-list", summary="获取方案列表", description="使用连接池查询所有方案信息")
async def get_scheme_list_with_connection(
@router.get("/analysis/runs", summary="获取分析运行列表")
async def get_analysis_runs(
conn: AsyncConnection = Depends(get_database_connection),
):
"""
@@ -46,14 +26,15 @@ async def get_scheme_list_with_connection(
返回项目中所有方案的详细信息
"""
try:
scheme_data = await SchemeRepository.get_schemes(conn)
return {"success": True, "data": scheme_data, "count": len(scheme_data)}
runs = await AnalysisRepository.list_runs(conn)
return {"success": True, "data": runs, "count": len(runs)}
except Exception as e:
raise HTTPException(status_code=500, detail=f"查询方案信息时发生错误: {str(e)}")
raise HTTPException(status_code=500, detail=f"查询分析运行时发生错误: {str(e)}")
@router.get("/burst-locate-result", summary="获取爆管定位结果", description="使用连接池查询所有爆管定位结果")
async def get_burst_locate_result_with_connection(
@router.get("/analysis/runs/{run_id}", summary="获取分析运行")
async def get_analysis_run(
run_id: UUID,
conn: AsyncConnection = Depends(get_database_connection),
):
"""
@@ -62,17 +43,22 @@ async def get_burst_locate_result_with_connection(
返回项目中所有的爆管定位分析结果
"""
try:
burst_data = await SchemeRepository.get_burst_locate_results(conn)
return {"success": True, "data": burst_data, "count": len(burst_data)}
run = await AnalysisRepository.get_run(conn, run_id)
if run is None:
raise HTTPException(status_code=404, detail="分析运行不存在")
return run
except HTTPException:
raise
except Exception as e:
raise HTTPException(
status_code=500, detail=f"查询爆管定位结果时发生错误: {str(e)}"
status_code=500, detail=f"查询分析运行时发生错误: {str(e)}"
)
@router.get("/burst-locate-result/{burst_incident}", summary="按事件查询爆管定位结果", description="根据爆管事件ID查询对应的爆管定位结果")
async def get_burst_locate_result_by_incident(
burst_incident: str = Path(..., description="爆管事件ID"),
@router.get("/analysis/runs/{run_id}/results", summary="获取分析结果")
async def get_analysis_results(
run_id: UUID,
result_type: str | None = Query(default=None, description="结果类型"),
conn: AsyncConnection = Depends(get_database_connection),
):
"""
@@ -82,11 +68,9 @@ async def get_burst_locate_result_by_incident(
burst_incident: 爆管事件的唯一标识符
"""
try:
return await SchemeRepository.get_burst_locate_result_by_incident(
conn, burst_incident
)
return await AnalysisRepository.list_results(conn, run_id, result_type)
except Exception as e:
raise HTTPException(
status_code=500,
detail=f"根据 burst_incident 查询爆管定位结果时发生错误: {str(e)}",
detail=f"查询分析结果时发生错误: {str(e)}",
)
-127
View File
@@ -1,127 +0,0 @@
from typing import Any, List, Dict
from fastapi import APIRouter, Query, Path
from app.services.tjnetwork import (
get_pipe_risk_probability_now,
get_pipe_risk_probability,
get_pipes_risk_probability,
get_network_pipe_risk_probability_now,
get_pipe_risk_probability_geometries,
)
router = APIRouter()
@router.get(
"/getpiperiskprobabilitynow/",
summary="获取管道当前风险概率",
description="获取指定管道当前时刻的风险概率值"
)
async def fastapi_get_pipe_risk_probability_now(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe_id: str = Query(..., description="管道ID")
) -> dict[str, Any]:
"""
获取管道当前风险概率。
查询指定管道在当前时刻的风险概率值。
Args:
network: 管网名称(或数据库名称)
pipe_id: 管道ID
Returns:
包含风险概率信息的字典
"""
return get_pipe_risk_probability_now(network, pipe_id)
@router.get(
"/getpiperiskprobability/",
summary="获取管道风险概率历史",
description="获取指定管道的风险概率历史数据"
)
async def fastapi_get_pipe_risk_probability(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe_id: str = Query(..., description="管道ID")
) -> dict[str, Any]:
"""
获取管道风险概率历史。
查询指定管道的历史风险概率数据。
Args:
network: 管网名称(或数据库名称)
pipe_id: 管道ID
Returns:
包含风险概率历史的字典
"""
return get_pipe_risk_probability(network, pipe_id)
@router.get(
"/getpipesriskprobability/",
summary="批量获取多条管道风险概率",
description="批量获取多条管道的风险概率值"
)
async def fastapi_get_pipes_risk_probability(
network: str = Query(..., description="管网名称(或数据库名称)"),
pipe_ids: str = Query(..., description="逗号分隔的管道ID列表")
) -> list[dict[str, Any]]:
"""
批量获取多条管道风险概率。
查询多条指定管道的风险概率值。
Args:
network: 管网名称(或数据库名称)
pipe_ids: 逗号分隔的管道ID列表(例如:pipe1,pipe2,pipe3
Returns:
包含多条管道风险概率的列表
"""
pipeids = pipe_ids.split(",")
return get_pipes_risk_probability(network, pipeids)
@router.get(
"/getnetworkpiperiskprobabilitynow/",
summary="获取整个网络的管道风险概率",
description="获取指定网络中所有管道的当前风险概率值"
)
async def fastapi_get_network_pipe_risk_probability_now(
network: str = Query(..., description="管网名称(或数据库名称)"),
) -> list[dict[str, Any]]:
"""
获取整个网络的管道风险概率。
查询指定网络中所有管道在当前时刻的风险概率值。
Args:
network: 管网名称(或数据库名称)
Returns:
包含网络内所有管道风险概率的列表
"""
return get_network_pipe_risk_probability_now(network)
@router.get(
"/getpiperiskprobabilitygeometries/",
summary="获取管道风险几何信息",
description="获取指定网络中管道的风险相关几何数据"
)
async def fastapi_get_pipe_risk_probability_geometries(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, Any]:
"""
获取管道风险几何信息。
查询指定网络中管道的地理和风险相关的几何数据。
Args:
network: 管网名称(或数据库名称)
Returns:
包含几何信息和风险数据的字典
"""
return get_pipe_risk_probability_geometries(network)
+33 -518
View File
@@ -1,527 +1,42 @@
from typing import Any
from fastapi import APIRouter, Request, Query
from app.services.tjnetwork import (
ChangeSet,
get_scada_info,
get_all_scada_info,
get_scada_device_schema,
get_scada_device,
set_scada_device,
add_scada_device,
delete_scada_device,
clean_scada_device,
get_all_scada_device_ids,
get_all_scada_devices,
get_scada_device_data_schema,
get_scada_device_data,
set_scada_device_data,
add_scada_device_data,
delete_scada_device_data,
clean_scada_device_data,
get_scada_element_schema,
get_scada_element,
set_scada_element,
add_scada_element,
delete_scada_element,
clean_scada_element,
get_all_scada_elements,
get_scada_element_schema,
get_scada_info_schema,
)
from fastapi import APIRouter, Depends, HTTPException
from psycopg import AsyncConnection
from app.auth.project_dependencies import get_project_pg_connection
from app.domain.schemas.scada import ScadaDeviceResponse
from app.infra.db.postgresql.scada import ScadaInfoRepository, get_scada_info_schema
router = APIRouter()
@router.get("/getscadaproperties/", summary="获取SCADA属性", tags=["SCADA基础"])
async def fast_get_scada_properties(
network: str = Query(..., description="管网名称(或数据库名称)"),
scada: str = Query(..., description="SCADA设备ID")
) -> dict[str, Any]:
"""
获取单个SCADA设备的属性信息
根据管网名称和SCADA设备ID获取该设备的完整属性。
Args:
network: 管网名称(或数据库名称)
scada: SCADA设备ID
Returns:
SCADA设备的属性字典
"""
return get_scada_info(network, scada)
@router.get("/getallscadaproperties/", summary="获取所有SCADA属性", tags=["SCADA基础"])
async def fast_get_all_scada_properties(
network: str = Query(..., description="管网名称(或数据库名称)")
@router.get("/network-schemas/scada-device", summary="获取 SCADA 设备结构")
def get_scada_device_schema() -> dict[str, dict[str, Any]]:
return get_scada_info_schema("")
@router.get(
"/scada-devices",
summary="获取 SCADA 设备列表",
response_model=list[ScadaDeviceResponse],
)
async def get_scada_devices(
conn: AsyncConnection = Depends(get_project_pg_connection),
) -> list[dict[str, Any]]:
"""
获取指定管网所有SCADA设备的属性信息
查询该管网下所有已配置的SCADA设备的属性列表。
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA设备属性列表
"""
return get_all_scada_info(network)
return await ScadaInfoRepository.get_scadas(conn)
############################################################
# scada_device 设备管理
############################################################
@router.get("/getscadadeviceschema/", summary="获取SCADA设备架构", tags=["SCADA设备"])
async def fastapi_get_scada_device_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""
获取SCADA设备的数据架构
返回SCADA设备表的字段定义和类型信息。
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA设备的字段架构信息
"""
return get_scada_device_schema(network)
@router.get("/getscadadevice/", summary="获取SCADA设备", tags=["SCADA设备"])
async def fastapi_get_scada_device(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="SCADA设备ID")
@router.get(
"/scada-devices/{device_id}",
summary="获取 SCADA 设备",
response_model=ScadaDeviceResponse,
)
async def get_scada_device(
device_id: str,
conn: AsyncConnection = Depends(get_project_pg_connection),
) -> dict[str, Any]:
"""
获取单个SCADA设备的信息
根据设备ID查询该设备的详细信息。
Args:
network: 管网名称(或数据库名称)
id: SCADA设备ID
Returns:
SCADA设备信息
"""
return get_scada_device(network, id)
@router.post("/setscadadevice/", response_model=None, summary="更新SCADA设备", tags=["SCADA设备"])
async def fastapi_set_scada_device(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
更新SCADA设备信息
修改指定SCADA设备的属性。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含要更新的设备属性
Returns:
变更集合信息
"""
props = await req.json()
return set_scada_device(network, ChangeSet(props))
@router.post("/addscadadevice/", response_model=None, summary="添加SCADA设备", tags=["SCADA设备"])
async def fastapi_add_scada_device(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
添加新的SCADA设备
在指定管网中添加一个新的SCADA设备。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含新设备的属性
Returns:
变更集合信息
"""
props = await req.json()
return add_scada_device(network, ChangeSet(props))
@router.post("/deletescadadevice/", response_model=None, summary="删除SCADA设备", tags=["SCADA设备"])
async def fastapi_delete_scada_device(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
删除SCADA设备
从指定管网中删除一个SCADA设备。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含要删除的设备ID
Returns:
变更集合信息
"""
props = await req.json()
return delete_scada_device(network, ChangeSet(props))
@router.post("/cleanscadadevice/", response_model=None, summary="清空SCADA设备表", tags=["SCADA设备"])
async def fastapi_clean_scada_device(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> ChangeSet:
"""
清空SCADA设备表
删除指定管网中所有的SCADA设备。
Args:
network: 管网名称(或数据库名称)
Returns:
变更集合信息
"""
return clean_scada_device(network)
@router.get("/getallscadadeviceids/", summary="获取所有SCADA设备ID", tags=["SCADA设备"])
async def fastapi_get_all_scada_device_ids(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[str]:
"""
获取指定管网所有SCADA设备的ID列表
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA设备ID列表
"""
return get_all_scada_device_ids(network)
@router.get("/getallscadadevices/", summary="获取所有SCADA设备", tags=["SCADA设备"])
async def fastapi_get_all_scada_devices(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
获取指定管网所有SCADA设备的完整信息
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA设备信息列表
"""
return get_all_scada_devices(network)
############################################################
# scada_device_data 设备数据管理
############################################################
@router.get("/getscadadevicedataschema/", summary="获取SCADA设备数据架构", tags=["SCADA设备数据"])
async def fastapi_get_scada_device_data_schema(
network: str = Query(..., description="管网名称(或数据库名称)"),
) -> dict[str, dict[str, Any]]:
"""
获取SCADA设备数据的表结构
返回SCADA设备数据表的字段定义和类型信息。
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA设备数据的字段架构信息
"""
return get_scada_device_data_schema(network)
@router.get("/getscadadevicedata/", summary="获取SCADA设备数据", tags=["SCADA设备数据"])
async def fastapi_get_scada_device_data(
network: str = Query(..., description="管网名称(或数据库名称)"),
device_id: str = Query(..., description="SCADA设备ID")
) -> dict[str, Any]:
"""
获取单个SCADA设备的数据
查询指定设备的监测数据或配置数据。
Args:
network: 管网名称(或数据库名称)
device_id: SCADA设备ID
Returns:
SCADA设备数据
"""
return get_scada_device_data(network, device_id)
@router.post("/setscadadevicedata/", response_model=None, summary="更新SCADA设备数据", tags=["SCADA设备数据"])
async def fastapi_set_scada_device_data(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
更新SCADA设备数据
修改指定SCADA设备的数据。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含要更新的数据
Returns:
变更集合信息
"""
props = await req.json()
return set_scada_device_data(network, ChangeSet(props))
@router.post("/addscadadevicedata/", response_model=None, summary="添加SCADA设备数据", tags=["SCADA设备数据"])
async def fastapi_add_scada_device_data(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
添加新的SCADA设备数据
为指定SCADA设备添加新的数据记录。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含新数据的内容
Returns:
变更集合信息
"""
props = await req.json()
return add_scada_device_data(network, ChangeSet(props))
@router.post("/deletescadadevicedata/", response_model=None, summary="删除SCADA设备数据", tags=["SCADA设备数据"])
async def fastapi_delete_scada_device_data(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
删除SCADA设备数据
删除指定SCADA设备的数据记录。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含要删除的数据ID
Returns:
变更集合信息
"""
props = await req.json()
return delete_scada_device_data(network, ChangeSet(props))
@router.post("/cleanscadadevicedata/", response_model=None, summary="清空SCADA设备数据表", tags=["SCADA设备数据"])
async def fastapi_clean_scada_device_data(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> ChangeSet:
"""
清空SCADA设备数据表
删除指定管网中所有SCADA设备的数据。
Args:
network: 管网名称(或数据库名称)
Returns:
变更集合信息
"""
return clean_scada_device_data(network)
############################################################
# scada_element SCADA元素映射
############################################################
@router.get("/getscadaelementschema/", summary="获取SCADA元素架构", tags=["SCADA元素映射"])
async def fastapi_get_scada_element_schema(
network: str = Query(..., description="管网名称(或数据库名称)"),
) -> dict[str, dict[str, Any]]:
"""
获取SCADA元素映射的表结构
返回SCADA元素映射表的字段定义和类型信息。
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA元素映射的字段架构信息
"""
return get_scada_element_schema(network)
@router.get("/getscadaelements/", summary="获取所有SCADA元素映射", tags=["SCADA元素映射"])
async def fastapi_get_scada_elements(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
获取指定管网所有SCADA元素映射
查询所有SCADA设备与管网元素(节点/管道)的映射关系。
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA元素映射列表
"""
return get_all_scada_elements(network)
@router.get("/getscadaelement/", summary="获取单个SCADA元素映射", tags=["SCADA元素映射"])
async def fastapi_get_scada_element(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="SCADA元素映射ID")
) -> dict[str, Any]:
"""
获取单个SCADA元素映射的信息
根据ID查询特定的SCADA设备与管网元素的映射关系。
Args:
network: 管网名称(或数据库名称)
id: SCADA元素映射ID
Returns:
SCADA元素映射信息
"""
return get_scada_element(network, id)
@router.post("/setscadaelement/", response_model=None, summary="更新SCADA元素映射", tags=["SCADA元素映射"])
async def fastapi_set_scada_element(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
更新SCADA元素映射
修改SCADA设备与管网元素的映射关系。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含要更新的映射信息
Returns:
变更集合信息
"""
props = await req.json()
return set_scada_element(network, ChangeSet(props))
@router.post("/addscadaelement/", response_model=None, summary="添加SCADA元素映射", tags=["SCADA元素映射"])
async def fastapi_add_scada_element(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
添加新的SCADA元素映射
创建SCADA设备与管网元素的新映射关系。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含新映射的信息
Returns:
变更集合信息
"""
props = await req.json()
return add_scada_element(network, ChangeSet(props))
@router.post("/deletescadaelement/", response_model=None, summary="删除SCADA元素映射", tags=["SCADA元素映射"])
async def fastapi_delete_scada_element(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
删除SCADA元素映射
移除SCADA设备与管网元素的映射关系。
Args:
network: 管网名称(或数据库名称)
req: 请求体,包含要删除的映射ID
Returns:
变更集合信息
"""
props = await req.json()
return delete_scada_element(network, ChangeSet(props))
@router.post("/cleanscadaelement/", response_model=None, summary="清空SCADA元素映射表", tags=["SCADA元素映射"])
async def fastapi_clean_scada_element(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> ChangeSet:
"""
清空SCADA元素映射表
删除指定管网中所有的SCADA元素映射。
Args:
network: 管网名称(或数据库名称)
Returns:
变更集合信息
"""
return clean_scada_element(network)
############################################################
# scada_info SCADA信息
############################################################
@router.get("/getscadainfoschema/", summary="获取SCADA信息架构", tags=["SCADA信息"])
async def fastapi_get_scada_info_schema(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> dict[str, dict[str, Any]]:
"""
获取SCADA信息表的结构
返回SCADA信息表的字段定义和类型信息。
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA信息的字段架构信息
"""
return get_scada_info_schema(network)
@router.get("/getscadainfo/", summary="获取SCADA信息", tags=["SCADA信息"])
async def fastapi_get_scada_info(
network: str = Query(..., description="管网名称(或数据库名称)"),
id: str = Query(..., description="SCADA信息ID")
) -> dict[str, Any]:
"""
获取单个SCADA信息
根据ID查询SCADA的详细配置信息。
Args:
network: 管网名称(或数据库名称)
id: SCADA信息ID
Returns:
SCADA信息详情
"""
return get_scada_info(network, id)
@router.get("/getallscadainfo/", summary="获取所有SCADA信息", tags=["SCADA信息"])
async def fastapi_get_all_scada_info(
network: str = Query(..., description="管网名称(或数据库名称)")
) -> list[dict[str, Any]]:
"""
获取指定管网所有SCADA的信息
查询该管网下所有已配置的SCADA的完整信息。
Args:
network: 管网名称(或数据库名称)
Returns:
SCADA信息列表
"""
return get_all_scada_info(network)
device = await ScadaInfoRepository.get_scada(conn, device_id)
if device is None:
raise HTTPException(status_code=404, detail="SCADA 设备不存在")
return device
-32
View File
@@ -1,32 +0,0 @@
from fastapi import APIRouter, Query
from typing import Any, List, Dict
from app.services.tjnetwork import get_scheme_schema, get_scheme, get_all_schemes
router = APIRouter()
@router.get("/getschemeschema/", summary="获取方案模式", description="获取指定网络的方案模式定义")
async def fastapi_get_scheme_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[Any, Any]]:
"""
获取方案模式定义
返回指定网络的方案模式结构定义
"""
return get_scheme_schema(network)
@router.get("/getscheme/", summary="获取单个方案", description="根据名称获取指定的方案信息")
async def fastapi_get_scheme(network: str = Query(..., description="管网名称(或数据库名称)"), schema_name: str = Query(..., description="方案名称")) -> dict[Any, Any]:
"""
获取单个方案详情
返回指定网络中指定名称的方案详细信息
"""
return get_scheme(network, schema_name)
@router.get("/getallschemes/", summary="获取所有方案", description="获取指定网络的所有方案信息")
async def fastapi_get_all_schemes(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[dict[Any, Any]]:
"""
获取所有方案列表
返回指定网络中所有可用的方案
"""
return get_all_schemes(network)
+286
View File
@@ -0,0 +1,286 @@
import logging
from typing import Any
from urllib.parse import quote
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Path, Query, status
from fastapi.responses import StreamingResponse
from starlette.concurrency import run_in_threadpool
from app.auth.metadata_dependencies import get_current_metadata_user
from app.auth.project_dependencies import (
ProjectContext,
get_project_context,
use_project_business_routing,
)
from app.infra.db.project_routing import ActiveProjectRouting
from app.domain.schemas.sensor_placement import (
SensorPointResponse,
SensorPlacementExportRequest,
SensorPlacementOptimizeRequest,
SensorPlacementSchemeResponse,
SensorPlacementUpdateRequest,
)
from app.services.sensor_placement import (
SensorPlacementConflictError,
SensorPlacementNotFoundError,
SensorPlacementValidationError,
build_sensor_placement_workbook,
can_edit_sensor_placement,
get_sensor_placement_candidate,
get_sensor_placement_run,
list_sensor_placement_runs,
optimize_sensor_placement_by_kmeans,
optimize_sensor_placement_by_sensitivity,
update_sensor_placement_run,
)
router = APIRouter()
logger = logging.getLogger(__name__)
def _project_network(network: str, project_context: ProjectContext) -> str:
if network != project_context.project_code:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="请求的管网不属于当前项目",
)
return project_context.project_code
def _can_modify_project(project_context: ProjectContext) -> bool:
return project_context.project_role == "member"
def _require_project_write(
project_context: ProjectContext,
) -> None:
if not _can_modify_project(project_context):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="当前项目角色为只读,不能修改监测点方案",
)
def _service_http_error(exc: Exception) -> HTTPException:
if isinstance(exc, SensorPlacementNotFoundError):
return HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=str(exc),
)
if isinstance(exc, SensorPlacementConflictError):
return HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=str(exc),
)
return HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail=str(exc),
)
def _get_run_response(
network: str,
run_id: UUID,
current_user: Any,
project_context: ProjectContext,
) -> dict[str, Any]:
try:
run = get_sensor_placement_run(network, run_id)
return {
**run,
"can_edit": (
_can_modify_project(project_context)
and can_edit_sensor_placement(current_user, run)
),
}
except (
SensorPlacementNotFoundError,
SensorPlacementValidationError,
) as exc:
raise _service_http_error(exc) from exc
@router.get(
"/sensor-placement-candidates/{node_id}",
response_model=SensorPointResponse,
summary="获取监测点候选节点详情",
)
async def get_sensor_placement_candidate_detail(
node_id: str = Path(..., min_length=1, max_length=32),
project_context: ProjectContext = Depends(get_project_context),
_routing: ActiveProjectRouting = Depends(use_project_business_routing),
) -> dict[str, Any]:
try:
return await run_in_threadpool(
get_sensor_placement_candidate,
project_context.project_code,
node_id,
)
except SensorPlacementValidationError as exc:
raise _service_http_error(exc) from exc
@router.post(
"/sensor-placement-runs",
response_model=SensorPlacementSchemeResponse,
summary="创建并返回监测点优化方案",
)
async def optimize_sensor_placement_scheme(
payload: SensorPlacementOptimizeRequest,
project_context: ProjectContext = Depends(get_project_context),
current_user=Depends(get_current_metadata_user),
) -> dict[str, Any]:
network = _project_network(payload.network, project_context)
_require_project_write(project_context)
optimizer = (
optimize_sensor_placement_by_sensitivity
if payload.method == "sensitivity"
else optimize_sensor_placement_by_kmeans
)
try:
created = await run_in_threadpool(
optimizer,
project_code=network,
run_name=payload.run_name,
sensor_count=payload.sensor_count,
min_diameter=payload.min_diameter,
created_by=current_user.username,
)
run = get_sensor_placement_run(network, created["run_id"])
return {**run, "can_edit": True}
except (
SensorPlacementConflictError,
SensorPlacementValidationError,
ValueError,
) as exc:
raise _service_http_error(exc) from exc
except Exception as exc:
logger.exception("Sensor placement optimization failed")
raise HTTPException(
status_code=500,
detail="监测点优化失败,请稍后重试",
) from exc
@router.get(
"/sensor-placement-runs",
response_model=list[SensorPlacementSchemeResponse],
summary="获取监测点优化运行",
)
async def get_sensor_placement_runs(
project_context: ProjectContext = Depends(get_project_context),
_routing: ActiveProjectRouting = Depends(use_project_business_routing),
) -> list[dict[str, Any]]:
return await run_in_threadpool(
list_sensor_placement_runs,
project_context.project_code,
)
@router.get(
"/sensor-placement-runs/{run_id}",
response_model=SensorPlacementSchemeResponse,
summary="获取监测点方案详情",
)
def get_sensor_placement_run_detail(
run_id: UUID,
network: str = Query(..., min_length=1),
project_context: ProjectContext = Depends(get_project_context),
current_user=Depends(get_current_metadata_user),
) -> dict[str, Any]:
return _get_run_response(
_project_network(network, project_context),
run_id,
current_user,
project_context,
)
@router.put(
"/sensor-placement-runs/{run_id}",
response_model=SensorPlacementSchemeResponse,
summary="覆盖保存监测点方案",
)
def overwrite_sensor_placement_run(
run_id: UUID,
payload: SensorPlacementUpdateRequest,
network: str = Query(..., min_length=1),
project_context: ProjectContext = Depends(get_project_context),
current_user=Depends(get_current_metadata_user),
) -> dict[str, Any]:
network = _project_network(network, project_context)
_require_project_write(project_context)
run = _get_run_response(
network,
run_id,
current_user,
project_context,
)
if not run["can_edit"]:
raise HTTPException(status_code=403, detail="无权修改该监测点优化运行")
try:
updated = update_sensor_placement_run(
network,
run_id,
expected_sensor_locations=payload.expected_sensor_locations,
sensor_locations=payload.sensor_locations,
)
return {**updated, "can_edit": True}
except (
SensorPlacementConflictError,
SensorPlacementNotFoundError,
SensorPlacementValidationError,
) as exc:
raise _service_http_error(exc) from exc
@router.post(
"/sensor-placement-runs/{run_id}/exports/excel",
summary="导出监测点工程清单",
)
async def export_sensor_placement_excel(
run_id: UUID,
payload: SensorPlacementExportRequest,
network: str = Query(..., min_length=1),
project_context: ProjectContext = Depends(get_project_context),
current_user=Depends(get_current_metadata_user),
) -> StreamingResponse:
network = _project_network(network, project_context)
run = _get_run_response(
network,
run_id,
current_user,
project_context,
)
if (
payload.sensor_locations != run["sensor_locations"]
and not run["can_edit"]
):
raise HTTPException(status_code=403, detail="无权导出该方案的未保存草稿")
try:
workbook = await run_in_threadpool(
build_sensor_placement_workbook,
network=network,
scheme=run,
sensor_location=payload.sensor_locations,
adjustment_status=payload.adjustment_status,
)
except SensorPlacementValidationError as exc:
raise _service_http_error(exc) from exc
filename = f"{run['name']}_监测点清单.xlsx"
encoded_filename = quote(filename)
return StreamingResponse(
workbook,
media_type=(
"application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
),
headers={
"Content-Disposition": (
f"attachment; filename*=UTF-8''{encoded_filename}"
)
},
)
File diff suppressed because it is too large Load Diff
-204
View File
@@ -1,204 +0,0 @@
from fastapi import APIRouter, Request, Query
from app.services.tjnetwork import (
ChangeSet,
get_current_operation,
execute_undo,
execute_redo,
list_snapshot,
have_snapshot,
have_snapshot_for_operation,
have_snapshot_for_current_operation,
take_snapshot_for_operation,
take_snapshot_for_current_operation,
take_snapshot,
pick_snapshot,
pick_operation,
sync_with_server,
execute_batch_commands,
execute_batch_command,
get_restore_operation,
set_restore_operation,
)
router = APIRouter()
@router.get("/getcurrentoperationid/", summary="获取当前操作ID", description="获取网络当前的操作ID")
async def get_current_operation_id_endpoint(network: str = Query(..., description="管网名称(或数据库名称)")) -> int:
"""
获取当前操作ID
返回网络当前正在执行的操作ID
"""
return get_current_operation(network)
@router.post("/undo/", summary="撤销操作", description="撤销网络上最后的一个操作")
async def undo_endpoint(network: str = Query(..., description="管网名称(或数据库名称)")):
"""
撤销操作
撤销网络上最近执行的一个操作
"""
return execute_undo(network)
@router.post("/redo/", summary="重做操作", description="重做网络上被撤销的操作")
async def redo_endpoint(network: str = Query(..., description="管网名称(或数据库名称)")):
"""
重做操作
重做网络上被撤销的操作
"""
return execute_redo(network)
@router.get("/getsnapshots/", summary="获取快照列表", description="获取网络中的所有快照")
async def list_snapshot_endpoint(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[tuple[int, str]]:
"""
获取快照列表
返回网络中所有可用的快照及其信息
"""
return list_snapshot(network)
@router.get("/havesnapshot/", summary="检查快照是否存在", description="检查指定标签的快照是否存在")
async def have_snapshot_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), tag: str = Query(..., description="快照标签")) -> bool:
"""
检查快照是否存在
返回指定标签的快照是否存在
"""
return have_snapshot(network, tag)
@router.get("/havesnapshotforoperation/", summary="检查操作快照是否存在", description="检查指定操作ID的快照是否存在")
async def have_snapshot_for_operation_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), operation: int = Query(..., description="操作ID")) -> bool:
"""
检查操作快照是否存在
返回指定操作ID的快照是否存在
"""
return have_snapshot_for_operation(network, operation)
@router.get("/havesnapshotforcurrentoperation/", summary="检查当前操作快照是否存在", description="检查当前操作的快照是否存在")
async def have_snapshot_for_current_operation_endpoint(network: str = Query(..., description="管网名称(或数据库名称)")) -> bool:
"""
检查当前操作快照是否存在
返回当前操作的快照是否存在
"""
return have_snapshot_for_current_operation(network)
@router.post("/takesnapshotforoperation/", summary="为操作创建快照", description="为指定的操作创建快照")
async def take_snapshot_for_operation_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
operation: int = Query(..., description="操作ID"),
tag: str = Query(..., description="快照标签")
) -> None:
"""
为操作创建快照
为指定操作创建一个带标签的快照
"""
return take_snapshot_for_operation(network, operation, tag)
@router.post("/takesnapshotforcurrentoperation", summary="为当前操作创建快照", description="为当前操作创建快照")
async def take_snapshot_for_current_operation_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), tag: str = Query(..., description="快照标签")) -> None:
"""
为当前操作创建快照
为网络当前操作创建一个快照
"""
return take_snapshot_for_current_operation(network, tag)
# 兼容旧拼写: takenapshotforcurrentoperation
@router.post("/takenapshotforcurrentoperation", summary="为当前操作创建快照(兼容模式)", description="为当前操作创建快照(兼容旧的API路径)")
async def take_snapshot_for_current_operation_legacy_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), tag: str = Query(..., description="快照标签")) -> None:
"""
为当前操作创建快照(兼容模式)
兼容旧的API路径,为网络当前操作创建一个快照
"""
return take_snapshot_for_current_operation(network, tag)
@router.post("/takesnapshot/", summary="创建快照", description="为网络创建一个快照")
async def take_snapshot_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), tag: str = Query(..., description="快照标签")) -> None:
"""
创建快照
为网络创建一个带标签的快照
"""
return take_snapshot(network, tag)
@router.post("/picksnapshot/", summary="选择快照", description="选择并恢复到指定的快照", response_model=None)
async def pick_snapshot_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), tag: str = Query(..., description="快照标签"), discard: bool = Query(False, description="是否丢弃当前更改")) -> ChangeSet:
"""
选择快照
选择并恢复到指定的快照
"""
return pick_snapshot(network, tag, discard)
@router.post("/pickoperation/", summary="选择操作", description="选择并恢复到指定的操作", response_model=None)
async def pick_operation_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
operation: int = Query(..., description="操作ID"),
discard: bool = Query(False, description="是否丢弃当前更改")
) -> ChangeSet:
"""
选择操作
选择并恢复到指定的操作
"""
return pick_operation(network, operation, discard)
@router.get("/syncwithserver/", summary="与服务器同步", description="将网络与服务器同步到指定操作", response_model=None)
async def sync_with_server_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), operation: int = Query(..., description="目标操作ID")) -> ChangeSet:
"""
与服务器同步
将网络与服务器同步到指定的操作
"""
return sync_with_server(network, operation)
@router.post("/batch/", summary="执行批量命令", description="执行多个网络操作命令", response_model=None)
async def execute_batch_commands_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), req: Request = None) -> ChangeSet:
"""
执行批量命令
在网络上执行多个操作命令
"""
jo_root = await req.json()
cs: ChangeSet = ChangeSet()
cs.operations = jo_root["operations"]
rcs = execute_batch_commands(network, cs)
return rcs
@router.post("/compressedbatch/", summary="执行压缩批量命令", description="执行压缩的批量命令", response_model=None)
async def execute_compressed_batch_commands_endpoint(
network: str = Query(..., description="管网名称(或数据库名称)"),
req: Request = None
) -> ChangeSet:
"""
执行压缩批量命令
执行压缩格式的批量命令
"""
jo_root = await req.json()
cs: ChangeSet = ChangeSet()
cs.operations = jo_root["operations"]
return execute_batch_command(network, cs)
@router.get("/getrestoreoperation/", summary="获取恢复操作ID", description="获取网络的恢复操作ID")
async def get_restore_operation_endpoint(network: str = Query(..., description="管网名称(或数据库名称)")) -> int:
"""
获取恢复操作ID
返回网络的恢复操作ID
"""
return get_restore_operation(network)
@router.post("/setrestoreoperation/", summary="设置恢复操作ID", description="设置网络的恢复操作ID")
async def set_restore_operation_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), operation: int = Query(..., description="操作ID")) -> None:
"""
设置恢复操作ID
设置网络的恢复操作ID
"""
return set_restore_operation(network, operation)
@@ -0,0 +1,61 @@
from datetime import datetime
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Query
from psycopg import AsyncConnection
from app.infra.db.timescaledb.repositories.analysis import AnalysisResultsRepository
from .dependencies import get_timescale_connection
router = APIRouter()
@router.get("/timeseries/analysis/runs/{run_id}/nodes/{node_id}")
async def get_analysis_node_series(
run_id: UUID,
node_id: str,
start_time: datetime = Query(...),
end_time: datetime = Query(...),
field: str = Query(...),
conn: AsyncConnection = Depends(get_timescale_connection),
):
try:
return await AnalysisResultsRepository.get_node_series(
conn, run_id, node_id, start_time, end_time, field
)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.get("/timeseries/analysis/runs/{run_id}/links/{link_id}")
async def get_analysis_link_series(
run_id: UUID,
link_id: str,
start_time: datetime = Query(...),
end_time: datetime = Query(...),
field: str = Query(...),
conn: AsyncConnection = Depends(get_timescale_connection),
):
try:
return await AnalysisResultsRepository.get_link_series(
conn, run_id, link_id, start_time, end_time, field
)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@router.get("/timeseries/analysis/runs/{run_id}/values")
async def get_analysis_values_at_time(
run_id: UUID,
result_time: datetime = Query(...),
element_type: str = Query(..., pattern="^(node|link)$"),
field: str = Query(...),
conn: AsyncConnection = Depends(get_timescale_connection),
):
try:
return await AnalysisResultsRepository.get_values_at_time(
conn, run_id, element_type, result_time, field
)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
+24 -184
View File
@@ -2,194 +2,36 @@ from fastapi import APIRouter, Depends, HTTPException, Query
from datetime import datetime
from psycopg import AsyncConnection
from app.infra.db.timescaledb.composite_queries import CompositeQueries
from app.services.timeseries_analysis import TimeseriesAnalysisService
from app.domain.schemas.timeseries_history import (
ElementHistoryQuery,
ElementHistoryResponse,
)
from app.services.timeseries_history import TimeseriesHistoryService
from .dependencies import get_timescale_connection, get_postgres_connection
router = APIRouter()
@router.get("/composite/scada-simulation", summary="获取SCADA关联的模拟数据",
tags=["复合查询"])
async def get_scada_associated_simulation_data(
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
device_ids: str = Query(..., description="SCADA设备ID列表,逗号分隔"),
scheme_type: str = Query(None, description="方案类型,若为空则查询实时数据"),
scheme_name: str = Query(None, description="方案名称,若为空则查询实时数据"),
@router.post(
"/timeseries/views/element-history/query",
summary="批量查询管网元素历史数据",
response_model=ElementHistoryResponse,
)
async def query_element_history(
payload: ElementHistoryQuery,
timescale_conn: AsyncConnection = Depends(get_timescale_connection),
postgres_conn: AsyncConnection = Depends(get_postgres_connection),
):
"""
获取SCADA关联的link/node模拟值
根据传入的SCADA device_ids,找到关联的link/node
并根据对应的type,查询对应的模拟数据。支持查询实时或方案数据。
Args:
start_time: 查询开始时间
end_time: 查询结束时间
device_ids: SCADA设备ID列表,用逗号分隔
scheme_type: 方案类型,若为空则查询实时数据
scheme_name: 方案名称,若为空则查询实时数据
timescale_conn: TimescaleDB连接
postgres_conn: PostgreSQL连接
Returns:
SCADA关联的模拟数据
Raises:
HTTPException: 当查询参数无效时返回400错误,未找到数据时返回404错误
"""
) -> ElementHistoryResponse:
try:
device_ids_list = (
[id.strip() for id in device_ids.split(",") if id.strip()]
if device_ids
else []
return await TimeseriesHistoryService.query(
timescale_conn, postgres_conn, payload
)
if scheme_type and scheme_name:
result = await CompositeQueries.get_scada_associated_scheme_simulation_data(
timescale_conn,
postgres_conn,
device_ids_list,
start_time,
end_time,
scheme_type,
scheme_name,
)
else:
result = (
await CompositeQueries.get_scada_associated_realtime_simulation_data(
timescale_conn,
postgres_conn,
device_ids_list,
start_time,
end_time,
)
)
if result is None:
raise HTTPException(status_code=404, detail="No simulation data found")
return result
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
@router.get("/composite/element-simulation", summary="获取管网元素的模拟数据",
tags=["复合查询"])
async def get_feature_simulation_data(
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
feature_infos: str = Query(
..., description="特征信息,格式: id1:type1,id2:type2type为pipe(管道)或junction(节点)"
),
scheme_type: str = Query(None, description="方案类型,若为空则查询实时数据"),
scheme_name: str = Query(None, description="方案名称,若为空则查询实时数据"),
timescale_conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
获取link/node模拟值
根据传入的featureInfos,找到关联的link/node
并根据对应的type,查询对应的模拟数据。支持查询实时或方案数据。
Args:
start_time: 查询开始时间
end_time: 查询结束时间
feature_infos: 格式为 "element_id1:type1,element_id2:type2"
例如: "P1:pipe,J1:junction"
scheme_type: 方案类型,若为空则查询实时数据
scheme_name: 方案名称,若为空则查询实时数据
timescale_conn: TimescaleDB连接
Returns:
管网元素的模拟数据
Raises:
HTTPException: 当feature_infos为空返回400错误,未找到数据返回404错误,其他错误返回400错误
"""
try:
feature_infos_list = []
if feature_infos:
for item in feature_infos.split(","):
item = item.strip()
if ":" in item:
element_id, element_type = item.split(":", 1)
feature_infos_list.append(
(element_id.strip(), element_type.strip())
)
if not feature_infos_list:
raise HTTPException(status_code=400, detail="feature_infos cannot be empty")
if scheme_type and scheme_name:
result = await CompositeQueries.get_scheme_simulation_data(
timescale_conn,
feature_infos_list,
start_time,
end_time,
scheme_type,
scheme_name,
)
else:
result = await CompositeQueries.get_realtime_simulation_data(
timescale_conn,
feature_infos_list,
start_time,
end_time,
)
if result is None:
raise HTTPException(status_code=404, detail="No simulation data found")
return result
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.get("/composite/element-scada", summary="获取管网元素关联的SCADA监测数据",
tags=["复合查询"])
async def get_element_associated_scada_data(
element_id: str = Query(..., description="管网元素ID(管道或节点)"),
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
use_cleaned: bool = Query(False, description="是否使用清洗后的数据"),
timescale_conn: AsyncConnection = Depends(get_timescale_connection),
postgres_conn: AsyncConnection = Depends(get_postgres_connection),
):
"""
获取link/node关联的SCADA监测值
根据传入的link/node id,匹配SCADA信息,
如果存在关联的SCADA device_id,获取实际的监测数据。
Args:
element_id: 管网元素ID
start_time: 查询开始时间
end_time: 查询结束时间
use_cleaned: 是否使用清洗后的数据,默认为False使用原始数据
timescale_conn: TimescaleDB连接
postgres_conn: PostgreSQL连接
Returns:
管网元素关联的SCADA监测数据
Raises:
HTTPException: 当查询参数无效时返回400错误,未找到关联数据返回404错误
"""
try:
result = await CompositeQueries.get_element_associated_scada_data(
timescale_conn, postgres_conn, element_id, start_time, end_time, use_cleaned
)
if result is None:
raise HTTPException(
status_code=404, detail="No associated SCADA data found"
)
return result
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.post("/composite/clean-scada", summary="清洗SCADA监测数据",
tags=["复合查询"])
@router.post("/timeseries/scada-cleaning-runs", summary="清洗SCADA监测数据")
async def clean_scada_data(
device_ids: str = Query(..., description="设备ID列表或 'all' 表示清洗所有设备"),
start_time: datetime = Query(..., description="清洗数据的开始时间"),
@@ -225,19 +67,18 @@ async def clean_scada_data(
if device_ids
else []
)
return await CompositeQueries.clean_scada_data(
return await TimeseriesAnalysisService.clean_scada_data(
timescale_conn, postgres_conn, device_ids_list, start_time, end_time
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.get("/composite/pipeline-health-prediction", summary="预测管道健康状况",
tags=["复合查询"])
@router.get("/pipeline-health-predictions", summary="预测管道健康状况")
async def predict_pipeline_health(
query_time: datetime = Query(..., description="查询时间"),
network_name: str = Query(..., description="管网名称(或数据库名称)"),
timescale_conn: AsyncConnection = Depends(get_timescale_connection),
postgres_conn: AsyncConnection = Depends(get_postgres_connection),
):
"""
预测管道健康状况
@@ -247,7 +88,6 @@ async def predict_pipeline_health(
Args:
query_time: 查询时间
network_name: 管网名称(或数据库名称)
timescale_conn: TimescaleDB连接
Returns:
@@ -257,8 +97,8 @@ async def predict_pipeline_health(
HTTPException: 当模型文件不存在返回404错误,其他错误返回400或500错误
"""
try:
return await CompositeQueries.predict_pipeline_health(
timescale_conn, network_name, query_time
return await TimeseriesAnalysisService.predict_pipeline_health(
timescale_conn, postgres_conn, query_time
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
+85 -34
View File
@@ -2,17 +2,43 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body
from typing import List
from datetime import datetime
from psycopg import AsyncConnection
from pydantic import BaseModel
from app.infra.db.timescaledb.repositories.realtime import RealtimeRepository
from .dependencies import get_timescale_connection
router = APIRouter()
TIME_WITH_TZ_DESC = "ISO 8601 / RFC 3339 时间,必须显式带时区;可直接传 UTC+8,服务端会先转换为 UTC 再处理。"
TIME_RANGE_START_DESC = f"时间范围开始时间。{TIME_WITH_TZ_DESC}"
TIME_RANGE_END_DESC = f"时间范围结束时间。{TIME_WITH_TZ_DESC}"
@router.post("/realtime/links/batch", status_code=201, summary="批量插入实时管道数据",
tags=["时间序列-实时数据"])
class RealtimeLinkBatchItem(BaseModel):
time: datetime
id: str
flow: float | None = None
friction: float | None = None
headloss: float | None = None
quality: float | None = None
reaction: float | None = None
setting: float | None = None
status: float | None = None
velocity: float | None = None
class RealtimeNodeBatchItem(BaseModel):
time: datetime
id: str
actual_demand: float | None = None
total_head: float | None = None
pressure: float | None = None
quality: float | None = None
@router.post("/timeseries/realtime/links/batches", status_code=201, summary="批量插入实时管道数据")
async def insert_realtime_links(
data: List[dict] = Body(..., description="管道数据列表,每项包含管道ID、时间戳等信息"),
data: List[RealtimeLinkBatchItem] = Body(..., description="同一时间点的管道快照数据"),
conn: AsyncConnection = Depends(get_timescale_connection)
):
"""
@@ -26,20 +52,27 @@ async def insert_realtime_links(
Returns:
插入成功的记录数
"""
await RealtimeRepository.insert_links_batch(conn, data)
await RealtimeRepository.insert_links_batch(
conn, [item.model_dump() for item in data]
)
return {"message": f"Inserted {len(data)} records"}
@router.get("/realtime/links", summary="查询实时管道数据", tags=["时间序列-实时数据"])
@router.get(
"/timeseries/realtime/links",
summary="查询实时管道数据",
description="按时间范围查询实时管道数据。start_time 和 end_time 必须显式带时区;允许传 UTC+8,服务端会先归一化为 UTC 再执行查询。",
)
async def get_realtime_links(
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
start_time: datetime = Query(..., description=TIME_RANGE_START_DESC),
end_time: datetime = Query(..., description=TIME_RANGE_END_DESC),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
查询指定时间范围内的实时管道数据
根据时间范围查询所有实时管道的监测值。
根据时间范围查询所有实时管道的监测值。传入时间必须显式包含时区,
可以直接使用 UTC+8,服务端会先统一转换为 UTC 再参与数据库查询。
Args:
start_time: 查询开始时间
@@ -51,10 +84,14 @@ async def get_realtime_links(
return await RealtimeRepository.get_links_by_time_range(conn, start_time, end_time)
@router.delete("/realtime/links", summary="删除实时管道数据", tags=["时间序列-实时数据"])
@router.delete(
"/timeseries/realtime/links",
summary="删除实时管道数据",
description="按时间范围删除实时管道数据。start_time 和 end_time 必须显式带时区;允许传 UTC+8,服务端按请求中的绝对时间删除对应 UTC 数据。",
)
async def delete_realtime_links(
start_time: datetime = Query(..., description="删除开始时间"),
end_time: datetime = Query(..., description="删除结束时间"),
start_time: datetime = Query(..., description=TIME_RANGE_START_DESC),
end_time: datetime = Query(..., description=TIME_RANGE_END_DESC),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
@@ -73,11 +110,10 @@ async def delete_realtime_links(
return {"message": "Deleted successfully"}
@router.patch("/realtime/links/{link_id}/field", summary="更新实时管道字段",
tags=["时间序列-实时数据"])
@router.patch("/timeseries/realtime/links/{link_id}/field", summary="更新实时管道字段")
async def update_realtime_link_field(
link_id: str = Path(..., description="管道ID"),
time: datetime = Query(..., description="更新数据的时间戳"),
time: datetime = Query(..., description=f"更新记录的时间戳{TIME_WITH_TZ_DESC}"),
field: str = Query(..., description="要更新的字段名称"),
value: float = Query(..., description="更新的字段值"),
conn: AsyncConnection = Depends(get_timescale_connection),
@@ -106,10 +142,9 @@ async def update_realtime_link_field(
raise HTTPException(status_code=400, detail=str(e))
@router.post("/realtime/nodes/batch", status_code=201, summary="批量插入实时节点数据",
tags=["时间序列-实时数据"])
@router.post("/timeseries/realtime/nodes/batches", status_code=201, summary="批量插入实时节点数据")
async def insert_realtime_nodes(
data: List[dict] = Body(..., description="节点数据列表,每项包含节点ID、时间戳等信息"),
data: List[RealtimeNodeBatchItem] = Body(..., description="同一时间点的节点快照数据"),
conn: AsyncConnection = Depends(get_timescale_connection)
):
"""
@@ -123,20 +158,27 @@ async def insert_realtime_nodes(
Returns:
插入成功的记录数
"""
await RealtimeRepository.insert_nodes_batch(conn, data)
await RealtimeRepository.insert_nodes_batch(
conn, [item.model_dump() for item in data]
)
return {"message": f"Inserted {len(data)} records"}
@router.get("/realtime/nodes", summary="查询实时节点数据", tags=["时间序列-实时数据"])
@router.get(
"/timeseries/realtime/nodes",
summary="查询实时节点数据",
description="按时间范围查询实时节点数据。start_time 和 end_time 必须显式带时区;允许传 UTC+8,服务端会先归一化为 UTC 再执行查询。",
)
async def get_realtime_nodes(
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
start_time: datetime = Query(..., description=TIME_RANGE_START_DESC),
end_time: datetime = Query(..., description=TIME_RANGE_END_DESC),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
查询指定时间范围内的实时节点数据
根据时间范围查询所有实时节点的监测值。
根据时间范围查询所有实时节点的监测值。传入时间必须显式包含时区,
可以直接使用 UTC+8,服务端会先统一转换为 UTC 再参与数据库查询。
Args:
start_time: 查询开始时间
@@ -148,10 +190,14 @@ async def get_realtime_nodes(
return await RealtimeRepository.get_nodes_by_time_range(conn, start_time, end_time)
@router.delete("/realtime/nodes", summary="删除实时节点数据", tags=["时间序列-实时数据"])
@router.delete(
"/timeseries/realtime/nodes",
summary="删除实时节点数据",
description="按时间范围删除实时节点数据。start_time 和 end_time 必须显式带时区;允许传 UTC+8,服务端按请求中的绝对时间删除对应 UTC 数据。",
)
async def delete_realtime_nodes(
start_time: datetime = Query(..., description="删除开始时间"),
end_time: datetime = Query(..., description="删除结束时间"),
start_time: datetime = Query(..., description=TIME_RANGE_START_DESC),
end_time: datetime = Query(..., description=TIME_RANGE_END_DESC),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
@@ -172,12 +218,11 @@ async def delete_realtime_nodes(
@router.post("/realtime/simulation/store", status_code=201, summary="存储实时模拟结果",
tags=["时间序列-实时数据"])
@router.post("/timeseries/realtime/simulation-results", status_code=201, summary="存储实时模拟结果")
async def store_realtime_simulation_result(
node_result_list: List[dict] = Body(..., description="节点模拟结果列表"),
link_result_list: List[dict] = Body(..., description="管道模拟结果列表"),
result_start_time: str = Query(..., description="模拟结果开始时间"),
result_start_time: str = Query(..., description=f"模拟结果开始时间{TIME_WITH_TZ_DESC}"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
@@ -199,10 +244,13 @@ async def store_realtime_simulation_result(
return {"message": "Simulation results stored successfully"}
@router.get("/realtime/query/by-time-property", summary="按时间和属性查询实时数据",
tags=["时间序列-实时数据"])
@router.get(
"/timeseries/realtime/records",
summary="按时间和属性查询实时数据",
description="查询指定时间点的实时属性值。query_time 必须显式带时区;允许传 UTC+8,服务端会先归一化为 UTC 再执行查询。",
)
async def query_realtime_records_by_time_property(
query_time: str = Query(..., description="查询时间"),
query_time: str = Query(..., description=f"查询时间{TIME_WITH_TZ_DESC}"),
type: str = Query(..., description="数据类型,pipe(管道)或 junction(节点)"),
property: str = Query(..., description="要查询的属性名称"),
conn: AsyncConnection = Depends(get_timescale_connection),
@@ -232,12 +280,15 @@ async def query_realtime_records_by_time_property(
raise HTTPException(status_code=400, detail=str(e))
@router.get("/realtime/query/by-id-time", summary="按ID和时间查询实时模拟数据",
tags=["时间序列-实时数据"])
@router.get(
"/timeseries/realtime/simulation-results",
summary="按ID和时间查询实时模拟数据",
description="查询指定元素在某一时间点的实时模拟结果。query_time 必须显式带时区;允许传 UTC+8,服务端会先归一化为 UTC 再执行查询。",
)
async def query_realtime_simulation_by_id_time(
id: str = Query(..., description="元素ID(管道ID或节点ID"),
type: str = Query(..., description="元素类型,pipe(管道)或 junction(节点)"),
query_time: str = Query(..., description="查询时间"),
query_time: str = Query(..., description=f"查询时间{TIME_WITH_TZ_DESC}"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
+71 -33
View File
@@ -2,52 +2,89 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body
from typing import List
from datetime import datetime
from psycopg import AsyncConnection
from pydantic import BaseModel, Field, field_validator
from app.infra.db.postgresql.scada import ScadaInfoRepository
from app.infra.db.timescaledb.repositories.scada import ScadaRepository
from .dependencies import get_timescale_connection
from .dependencies import get_postgres_connection, get_timescale_connection
router = APIRouter()
SCADA_BATCH_MAX_ITEMS = 10_000
@router.post("/scada/batch", status_code=201, summary="批量插入SCADA监测数据",
tags=["时间序列-监测数据"])
class ScadaReadingBatchItem(BaseModel):
time: datetime
device_id: str = Field(min_length=1)
monitored_value: float | None = None
cleaned_value: float | None = None
@field_validator("device_id")
@classmethod
def normalize_device_id(cls, value: str) -> str:
normalized = value.strip()
if not normalized:
raise ValueError("device_id must not be blank")
return normalized
@router.post("/timeseries/scada-readings/batches", status_code=201, summary="批量插入SCADA监测数据")
async def insert_scada_data(
data: List[dict] = Body(..., description="SCADA设备监测数据列表"),
conn: AsyncConnection = Depends(get_timescale_connection)
data: List[ScadaReadingBatchItem] = Body(
...,
min_length=1,
max_length=SCADA_BATCH_MAX_ITEMS,
description="SCADA设备监测数据列表",
),
conn: AsyncConnection = Depends(get_timescale_connection),
postgres_conn: AsyncConnection = Depends(get_postgres_connection),
):
"""
批量插入SCADA监测数据
将多个设备的实时监测数据批量插入时间序列数据库。
Args:
data: SCADA设备监测数据列表,每项包含device_id、时间戳和监测值等信息
Returns:
插入成功的记录数
"""
await ScadaRepository.insert_scada_batch(conn, data)
rows = [item.model_dump() for item in data]
requested_ids = list(dict.fromkeys(item["device_id"] for item in rows))
existing_ids = await ScadaInfoRepository.get_existing_device_ids(
postgres_conn, requested_ids
)
missing_ids = [
device_id for device_id in requested_ids if device_id not in existing_ids
]
if missing_ids:
raise HTTPException(
status_code=422,
detail=f"SCADA devices do not exist in BizDB: {', '.join(missing_ids)}",
)
await ScadaRepository.insert_scada_batch(conn, rows)
return {"message": f"Inserted {len(data)} records"}
@router.get("/scada/by-ids-time-range", summary="按设备ID和时间范围查询SCADA数据",
tags=["时间序列-监测数据"])
@router.get("/timeseries/scada-readings", summary="按设备ID和时间范围查询SCADA数据")
async def get_scada_by_ids_time_range(
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
device_ids: str = Query(..., description="设备ID列,逗号分隔,如 'device1,device2,device3'"),
device_ids: str = Query(
..., description="设备ID列表,逗号分隔,如 'device1,device2,device3'"
),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
按设备ID和时间范围查询SCADA监测数据
查询多个设备在指定时间范围内的所有监测数据。
Args:
start_time: 查询开始时间
end_time: 查询结束时间
device_ids: 设备ID列表,用逗号分隔
Returns:
SCADA监测数据列表
"""
@@ -59,29 +96,32 @@ async def get_scada_by_ids_time_range(
)
@router.get("/scada/by-ids-field-time-range", summary="按设备ID、字段和时间范围查询SCADA数据",
tags=["时间序列-监测数据"])
@router.get(
"/timeseries/scada-readings/fields", summary="按设备ID、字段和时间范围查询SCADA数据"
)
async def get_scada_field_by_ids_time_range(
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
field: str = Query(..., description="要查询的字段名称"),
device_ids: str = Query(..., description="设备ID列表,逗号分隔,如 'device1,device2,device3'"),
device_ids: str = Query(
..., description="设备ID列表,逗号分隔,如 'device1,device2,device3'"
),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
按设备ID、字段和时间范围查询特定SCADA数据
查询多个设备在指定时间范围内的特定字段监测数据。
Args:
start_time: 查询开始时间
end_time: 查询结束时间
field: 字段名称
device_ids: 设备ID列表,用逗号分隔
Returns:
SCADA字段数据列表
Raises:
HTTPException: 当字段不存在或查询参数无效时返回400错误
"""
@@ -98,8 +138,7 @@ async def get_scada_field_by_ids_time_range(
raise HTTPException(status_code=400, detail=str(e))
@router.patch("/scada/{device_id}/field", summary="更新SCADA设备字段",
tags=["时间序列-监测数据"])
@router.patch("/timeseries/scada-readings/{device_id}/field", summary="更新SCADA设备字段")
async def update_scada_field(
device_id: str = Path(..., description="设备ID"),
time: datetime = Query(..., description="更新数据的时间戳"),
@@ -109,18 +148,18 @@ async def update_scada_field(
):
"""
更新指定设备的字段值
更新SCADA设备在特定时间的某个字段监测数据。
Args:
device_id: 设备ID
time: 数据时间戳
field: 字段名称
value: 字段新值
Returns:
更新结果信息
Raises:
HTTPException: 当字段不存在或更新失败时返回400错误
"""
@@ -131,8 +170,7 @@ async def update_scada_field(
raise HTTPException(status_code=400, detail=str(e))
@router.delete("/scada/by-id-time-range", summary="按设备ID和时间范围删除SCADA数据",
tags=["时间序列-监测数据"])
@router.delete("/timeseries/scada-readings", summary="按设备ID和时间范围删除SCADA数据")
async def delete_scada_data(
device_id: str = Query(..., description="设备ID"),
start_time: datetime = Query(..., description="删除开始时间"),
@@ -141,14 +179,14 @@ async def delete_scada_data(
):
"""
删除指定设备和时间范围内的SCADA数据
删除在指定时间范围内的特定设备监测数据。
Args:
device_id: 设备ID
start_time: 删除开始时间
end_time: 删除结束时间
Returns:
删除结果信息
"""
-398
View File
@@ -1,398 +0,0 @@
from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body
from typing import List
from datetime import datetime
from psycopg import AsyncConnection
from app.infra.db.timescaledb.repositories.scheme import SchemeRepository
from .dependencies import get_timescale_connection
router = APIRouter()
@router.post("/scheme/links/batch", status_code=201, summary="批量插入方案管道数据",
tags=["时间序列-方案数据"])
async def insert_scheme_links(
data: List[dict] = Body(..., description="方案管道数据列表"),
conn: AsyncConnection = Depends(get_timescale_connection)
):
"""
批量插入方案管道数据
将特定方案的管道模拟数据批量插入时间序列数据库。
Args:
data: 方案管道数据列表
Returns:
插入成功的记录数
"""
await SchemeRepository.insert_links_batch(conn, data)
return {"message": f"Inserted {len(data)} records"}
@router.get("/scheme/links", summary="查询方案管道数据", tags=["时间序列-方案数据"])
async def get_scheme_links(
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
查询指定方案和时间范围内的管道数据
根据方案和时间范围查询管道的模拟值。
Args:
scheme_type: 方案类型
scheme_name: 方案名称
start_time: 查询开始时间
end_time: 查询结束时间
Returns:
方案管道数据列表
"""
return await SchemeRepository.get_links_by_scheme_and_time_range(
conn, scheme_type, scheme_name, start_time, end_time
)
@router.get("/scheme/links/{link_id}/field", summary="查询方案管道字段数据",
tags=["时间序列-方案数据"])
async def get_scheme_link_field(
link_id: str = Path(..., description="管道ID"),
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
field: str = Query(..., description="要查询的字段名称"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
查询指定方案管道的特定字段数据
查询特定方案中指定管道在时间范围内的特定字段值。
Args:
link_id: 管道ID
scheme_type: 方案类型
scheme_name: 方案名称
start_time: 查询开始时间
end_time: 查询结束时间
field: 字段名称
Returns:
字段数据列表
Raises:
HTTPException: 当查询参数无效时返回400错误
"""
try:
return await SchemeRepository.get_link_field_by_scheme_and_time_range(
conn, scheme_type, scheme_name, start_time, end_time, link_id, field
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.patch("/scheme/links/{link_id}/field", summary="更新方案管道字段",
tags=["时间序列-方案数据"])
async def update_scheme_link_field(
link_id: str = Path(..., description="管道ID"),
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
time: datetime = Query(..., description="更新数据的时间戳"),
field: str = Query(..., description="要更新的字段名称"),
value: float = Query(..., description="更新的字段值"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
更新指定方案管道的字段值
更新特定方案中指定管道在某个时间的字段数据。
Args:
link_id: 管道ID
scheme_type: 方案类型
scheme_name: 方案名称
time: 数据时间戳
field: 字段名称
value: 字段新值
Returns:
更新结果信息
Raises:
HTTPException: 当字段不存在或更新失败时返回400错误
"""
try:
await SchemeRepository.update_link_field(
conn, time, scheme_type, scheme_name, link_id, field, value
)
return {"message": "Updated successfully"}
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.delete("/scheme/links", summary="删除方案管道数据", tags=["时间序列-方案数据"])
async def delete_scheme_links(
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
start_time: datetime = Query(..., description="删除开始时间"),
end_time: datetime = Query(..., description="删除结束时间"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
删除指定方案和时间范围内的管道数据
删除在指定方案和时间范围内的所有管道模拟数据。
Args:
scheme_type: 方案类型
scheme_name: 方案名称
start_time: 删除开始时间
end_time: 删除结束时间
Returns:
删除结果信息
"""
await SchemeRepository.delete_links_by_scheme_and_time_range(
conn, scheme_type, scheme_name, start_time, end_time
)
return {"message": "Deleted successfully"}
@router.post("/scheme/nodes/batch", status_code=201, summary="批量插入方案节点数据",
tags=["时间序列-方案数据"])
async def insert_scheme_nodes(
data: List[dict] = Body(..., description="方案节点数据列表"),
conn: AsyncConnection = Depends(get_timescale_connection)
):
"""
批量插入方案节点数据
将特定方案的节点模拟数据批量插入时间序列数据库。
Args:
data: 方案节点数据列表
Returns:
插入成功的记录数
"""
await SchemeRepository.insert_nodes_batch(conn, data)
return {"message": f"Inserted {len(data)} records"}
@router.get("/scheme/nodes/{node_id}/field", summary="查询方案节点字段数据",
tags=["时间序列-方案数据"])
async def get_scheme_node_field(
node_id: str = Path(..., description="节点ID"),
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
start_time: datetime = Query(..., description="查询开始时间"),
end_time: datetime = Query(..., description="查询结束时间"),
field: str = Query(..., description="要查询的字段名称"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
查询指定方案节点的特定字段数据
查询特定方案中指定节点在时间范围内的特定字段值。
Args:
node_id: 节点ID
scheme_type: 方案类型
scheme_name: 方案名称
start_time: 查询开始时间
end_time: 查询结束时间
field: 字段名称
Returns:
字段数据列表
Raises:
HTTPException: 当查询参数无效时返回400错误
"""
try:
return await SchemeRepository.get_node_field_by_scheme_and_time_range(
conn, scheme_type, scheme_name, start_time, end_time, node_id, field
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.patch("/scheme/nodes/{node_id}/field", summary="更新方案节点字段",
tags=["时间序列-方案数据"])
async def update_scheme_node_field(
node_id: str = Path(..., description="节点ID"),
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
time: datetime = Query(..., description="更新数据的时间戳"),
field: str = Query(..., description="要更新的字段名称"),
value: float = Query(..., description="更新的字段值"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
更新指定方案节点的字段值
更新特定方案中指定节点在某个时间的字段数据。
Args:
node_id: 节点ID
scheme_type: 方案类型
scheme_name: 方案名称
time: 数据时间戳
field: 字段名称
value: 字段新值
Returns:
更新结果信息
Raises:
HTTPException: 当字段不存在或更新失败时返回400错误
"""
try:
await SchemeRepository.update_node_field(
conn, time, scheme_type, scheme_name, node_id, field, value
)
return {"message": "Updated successfully"}
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.delete("/scheme/nodes", summary="删除方案节点数据", tags=["时间序列-方案数据"])
async def delete_scheme_nodes(
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
start_time: datetime = Query(..., description="删除开始时间"),
end_time: datetime = Query(..., description="删除结束时间"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
删除指定方案和时间范围内的节点数据
删除在指定方案和时间范围内的所有节点模拟数据。
Args:
scheme_type: 方案类型
scheme_name: 方案名称
start_time: 删除开始时间
end_time: 删除结束时间
Returns:
删除结果信息
"""
await SchemeRepository.delete_nodes_by_scheme_and_time_range(
conn, scheme_type, scheme_name, start_time, end_time
)
return {"message": "Deleted successfully"}
@router.post("/scheme/simulation/store", status_code=201, summary="存储方案模拟结果",
tags=["时间序列-方案数据"])
async def store_scheme_simulation_result(
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
node_result_list: List[dict] = Body(..., description="节点模拟结果列表"),
link_result_list: List[dict] = Body(..., description="管道模拟结果列表"),
result_start_time: str = Query(..., description="模拟结果开始时间"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
存储方案模拟结果到时间序列数据库
将特定方案的节点和管道模拟计算结果批量存储到TimescaleDB数据库。
Args:
scheme_type: 方案类型
scheme_name: 方案名称
node_result_list: 节点模拟结果列表
link_result_list: 管道模拟结果列表
result_start_time: 模拟结果对应的起始时间
Returns:
存储结果信息
"""
await SchemeRepository.store_scheme_simulation_result(
conn,
scheme_type,
scheme_name,
node_result_list,
link_result_list,
result_start_time,
)
return {"message": "Scheme simulation results stored successfully"}
@router.get("/scheme/query/by-scheme-time-property", summary="按方案、时间和属性查询数据",
tags=["时间序列-方案数据"])
async def query_scheme_records_by_scheme_time_property(
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
query_time: str = Query(..., description="查询时间"),
type: str = Query(..., description="元素类型,pipe(管道)或 junction(节点)"),
property: str = Query(..., description="要查询的属性名称"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
按指定方案、时间和属性查询所有方案数据
查询在特定方案和时间点,所有指定类型元素的特定属性值。
Args:
scheme_type: 方案类型
scheme_name: 方案名称
query_time: 查询时间
type: 元素类型(pipe或junction
property: 属性名称
Returns:
查询结果列表
Raises:
HTTPException: 当查询参数无效时返回400错误
"""
try:
results = await SchemeRepository.query_all_record_by_scheme_time_property(
conn, scheme_type, scheme_name, query_time, type, property
)
return {"results": results}
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
@router.get("/scheme/query/by-id-time", summary="按ID和时间查询方案模拟数据",
tags=["时间序列-方案数据"])
async def query_scheme_simulation_by_id_time(
scheme_type: str = Query(..., description="方案类型"),
scheme_name: str = Query(..., description="方案名称"),
id: str = Query(..., description="元素ID(管道ID或节点ID"),
type: str = Query(..., description="元素类型,pipe(管道)或 junction(节点)"),
query_time: str = Query(..., description="查询时间"),
conn: AsyncConnection = Depends(get_timescale_connection),
):
"""
按指定ID和时间查询方案模拟结果
查询特定方案中的元素在某一时间点的模拟数据。
Args:
scheme_type: 方案类型
scheme_name: 方案名称
id: 元素ID
type: 元素类型(pipe或junction
query_time: 查询时间
Returns:
模拟结果数据
Raises:
HTTPException: 当查询参数无效时返回400错误
"""
try:
result = await SchemeRepository.query_scheme_simulation_result_by_id_time(
conn, scheme_type, scheme_name, id, type, query_time
)
return {"result": result}
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
-215
View File
@@ -1,215 +0,0 @@
"""
用户管理 API 接口
演示权限控制的使用
"""
from typing import List
from fastapi import APIRouter, Depends, HTTPException, status, Path, Query
from app.domain.schemas.user import UserResponse, UserUpdate, UserCreate
from app.domain.models.role import UserRole
from app.domain.schemas.user import UserInDB
from app.infra.db.metadb.repositories.user_repository import UserRepository
from app.auth.dependencies import get_user_repository, get_current_active_user
from app.auth.permissions import get_current_admin, require_role, check_resource_owner
router = APIRouter()
@router.get(
"/",
summary="列出所有用户",
description="获取用户列表(仅管理员)",
response_model=List[UserResponse],
)
async def list_users(
skip: int = Query(0, ge=0, description="跳过的用户数"),
limit: int = Query(100, ge=1, le=1000, description="返回的最大用户数"),
current_user: UserInDB = Depends(require_role(UserRole.ADMIN)),
user_repo: UserRepository = Depends(get_user_repository),
) -> List[UserResponse]:
"""
获取用户列表
获取系统中所有的用户信息(需要管理员权限)
"""
users = await user_repo.get_all_users(skip=skip, limit=limit)
return [UserResponse.model_validate(user) for user in users]
@router.get(
"/{user_id}",
summary="获取用户详情",
description="获取指定用户的详细信息",
response_model=UserResponse,
)
async def get_user(
user_id: int = Path(..., gt=0, description="用户ID"),
current_user: UserInDB = Depends(get_current_active_user),
user_repo: UserRepository = Depends(get_user_repository),
) -> UserResponse:
"""
获取用户详情
管理员可查看所有用户,普通用户只能查看自己
"""
# 检查权限
if not check_resource_owner(user_id, current_user):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="You don't have permission to view this user",
)
user = await user_repo.get_user_by_id(user_id)
if not user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
)
return UserResponse.model_validate(user)
@router.put(
"/{user_id}",
summary="更新用户信息",
description="更新指定用户的信息",
response_model=UserResponse,
)
async def update_user(
user_id: int = Path(..., gt=0, description="用户ID"),
user_update: UserUpdate = None,
current_user: UserInDB = Depends(get_current_active_user),
user_repo: UserRepository = Depends(get_user_repository),
) -> UserResponse:
"""
更新用户信息
管理员可更新所有用户,普通用户只能更新自己(且不能修改角色)
"""
# 检查用户是否存在
target_user = await user_repo.get_user_by_id(user_id)
if not target_user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
)
# 权限检查
is_owner = current_user.id == user_id
is_admin = UserRole(current_user.role).has_permission(UserRole.ADMIN)
if not is_owner and not is_admin:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="You don't have permission to update this user",
)
# 非管理员不能修改角色和激活状态
if not is_admin:
if user_update.role is not None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only admins can change user roles",
)
if user_update.is_active is not None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only admins can change user active status",
)
# 更新用户
updated_user = await user_repo.update_user(user_id, user_update)
if not updated_user:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to update user",
)
return UserResponse.model_validate(updated_user)
@router.delete("/{user_id}", summary="删除用户", description="删除指定用户(仅管理员)")
async def delete_user(
user_id: int = Path(..., gt=0, description="用户ID"),
current_user: UserInDB = Depends(get_current_admin),
user_repo: UserRepository = Depends(get_user_repository),
) -> dict:
"""
删除用户
删除指定用户(需要管理员权限,不能删除自己)
"""
# 不能删除自己
if current_user.id == user_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="You cannot delete your own account",
)
success = await user_repo.delete_user(user_id)
if not success:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
)
return {"message": "User deleted successfully"}
@router.post(
"/{user_id}/activate",
summary="激活用户",
description="激活指定用户账户(仅管理员)",
response_model=UserResponse,
)
async def activate_user(
user_id: int = Path(..., gt=0, description="用户ID"),
current_user: UserInDB = Depends(get_current_admin),
user_repo: UserRepository = Depends(get_user_repository),
) -> UserResponse:
"""
激活用户
激活指定用户的账户(需要管理员权限)
"""
user_update = UserUpdate(is_active=True)
updated_user = await user_repo.update_user(user_id, user_update)
if not updated_user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
)
return UserResponse.model_validate(updated_user)
@router.post(
"/{user_id}/deactivate",
summary="停用用户",
description="停用指定用户账户(仅管理员)",
response_model=UserResponse,
)
async def deactivate_user(
user_id: int = Path(..., gt=0, description="用户ID"),
current_user: UserInDB = Depends(get_current_admin),
user_repo: UserRepository = Depends(get_user_repository),
) -> UserResponse:
"""
停用用户
停用指定用户的账户(需要管理员权限,不能停用自己)
"""
# 不能停用自己
if current_user.id == user_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="You cannot deactivate your own account",
)
user_update = UserUpdate(is_active=False)
updated_user = await user_repo.update_user(user_id, user_update)
if not updated_user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="User not found"
)
return UserResponse.model_validate(updated_user)
-36
View File
@@ -1,36 +0,0 @@
from fastapi import APIRouter, Request, Query
from typing import Any, List, Dict, Union
from app.services.tjnetwork import Any, get_all_users, get_user, get_user_schema
router = APIRouter()
###########################################################
# user 39
###########################################################
@router.get("/getuserschema/", summary="获取用户模式", description="获取指定网络的用户模式定义")
async def fastapi_get_user_schema(network: str = Query(..., description="管网名称(或数据库名称)")) -> dict[str, dict[Any, Any]]:
"""
获取用户模式定义
返回指定网络的用户模式结构定义
"""
return get_user_schema(network)
@router.get("/getuser/", summary="获取单个用户", description="获取指定网络中的单个用户信息")
async def fastapi_get_user(network: str = Query(..., description="管网名称(或数据库名称)"), user_name: str = Query(..., description="用户名")) -> dict[Any, Any]:
"""
获取用户信息
返回指定网络中指定用户名的详细信息
"""
return get_user(network, user_name)
@router.get("/getallusers/", summary="获取所有用户", description="获取指定网络的所有用户列表")
async def fastapi_get_all_users(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[dict[Any, Any]]:
"""
获取所有用户列表
返回指定网络中所有用户的信息
"""
return get_all_users(network)
+29
View File
@@ -0,0 +1,29 @@
from typing import Any
from fastapi import APIRouter, HTTPException, status
from app.services.web_search import (
BochaSearchAPIError,
BochaSearchConfigError,
WebSearchRequest,
search_bocha_web,
)
router = APIRouter()
@router.post(
"/web-searches",
summary="Web Search",
description="调用 Bocha Web Search API 获取实时网页搜索结果",
)
async def web_search(request: WebSearchRequest) -> dict[str, Any]:
try:
return await search_bocha_web(request)
except BochaSearchConfigError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=str(exc),
) from exc
except BochaSearchAPIError as exc:
raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc

Some files were not shown because too many files have changed in this diff Show More