-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_database.py
More file actions
211 lines (170 loc) · 8.12 KB
/
Copy pathtest_database.py
File metadata and controls
211 lines (170 loc) · 8.12 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
# -*- coding: UTF-8 -*-
"""
测试数据库写入和POST接口调用的独立脚本
用于调试从HDF5读取数据并写入数据库、调用POST接口的逻辑
"""
import os
import h5py
import numpy as np
import pandas as pd
import requests
from sqlserver_handler import SQLServerHandler
from config import *
from logger import logger
def load_section_mapping():
"""
加载断面ID和名称的映射关系
从断面id名称对应关系.xlsx文件中读取映射
:return: 字典,key为cross_sections_name(字节串),value为(SECTION_ID, SECTION_NAME)元组
"""
try:
import pandas as pd
import os
# 断面映射文件路径(相对于当前工作目录)
mapping_file = os.path.join(os.path.dirname(__file__), "断面id名称对应关系.xlsx")
if not os.path.exists(mapping_file):
logger.warning(f"断面映射文件不存在: {mapping_file}")
return {}
# 读取Excel文件
df = pd.read_excel(mapping_file, sheet_name='FLOODAREA')
# 建立映射字典
section_map = {}
for _, row in df.iterrows():
cross_sections_id = row['cross_sections_id']
# 跳过cross_sections_id为"无"的记录
if pd.isna(cross_sections_id) or str(cross_sections_id).strip() == "无":
continue
cross_sections_name = row['cross_sections_name']
section_id = int(row['SECTION_ID'])
section_name = str(row['SECTION_NAME'])
# 将cross_sections_name转为字节串作为key(与HDF5中的格式一致)
key = cross_sections_name.encode('utf-8') if isinstance(cross_sections_name, str) else cross_sections_name
section_map[key] = (section_id, section_name)
logger.info(f"成功加载{len(section_map)}个断面映射")
return section_map
except Exception as e:
logger.error(f"加载断面映射失败: {e}")
return {}
def test_database_and_post():
"""
测试数据库写入和POST接口调用
"""
# ========== 配置部分 - 请修改这些参数 ==========
scheme_name = "小流量调度方案仿真" # 方案名称
hdf5_file_path = "/root/results/小流量调度方案仿真/小流量调度方案仿真.hdf5" # HDF5文件路径
output_path = "/root/results/小流量调度方案仿真" # 输出路径
# ========== 初始化数据库连接 ==========
logger.info("初始化数据库连接...")
sqlserver_handler = SQLServerHandler(
SQLSERVER_HOST,
SQLSERVER_PORT,
SQLSERVER_USER,
SQLSERVER_PASSWORD,
SQLSERVER_DATABASE
)
# ========== 读取HDF5文件 ==========
try:
logger.info(f"开始读取HDF5文件: {hdf5_file_path}")
if not os.path.exists(hdf5_file_path):
logger.error(f"HDF5文件不存在: {hdf5_file_path}")
return
with h5py.File(hdf5_file_path, 'r') as hf:
# 读取断面数据
cross_sections_ws = hf['data']['CrossSections']['WaterSurface'][:]
cross_sections_name = hf['data']['CrossSections']['Name'][:]
cross_sections_flow = hf['data']['CrossSections']['Flow'][:]
time_date_stamp = hf['data']['TimeDateStamp'][:]
# 读取淹没面积数据
flooded_area = hf['data']['2DFlowAreas']['FloodedArea'][:]
logger.info(f"断面数量: {len(cross_sections_name)}")
logger.info(f"时间步数(time_date_stamp): {len(time_date_stamp)}")
logger.info(f"断面数据形状(WaterSurface): {cross_sections_ws.shape}")
logger.info(f"淹没面积数据长度: {len(flooded_area)}")
except Exception as e:
logger.error(f"读取HDF5文件失败: {e}")
import traceback
logger.error(traceback.format_exc())
return
# ========== 写入数据库 ==========
try:
logger.info("开始写入数据库...")
# 1. 更新FLOOD_REHEARSAL记录的STATUS为1和MAX_FLOOD_AREA
max_flood_area = int(np.max(flooded_area))
logger.info(f"最大淹没面积: {max_flood_area} km²")
success = sqlserver_handler.update_flood_rehearsal_status(
flood_dispatch_name=scheme_name,
status=1,
max_flood_area=max_flood_area
)
if not success:
logger.warning("FLOOD_REHEARSAL状态更新失败")
# 2. 准备并插入FLOOD_SECTION记录
logger.info("开始准备FLOOD_SECTION数据...")
# 加载断面映射
section_mapping = load_section_mapping()
# 获取断面数据的时间步数(使用cross_sections_ws的实际列数)
num_cross_section_timesteps = cross_sections_ws.shape[0] if len(cross_sections_ws.shape) > 1 else len(time_date_stamp)
logger.info(f"断面数据实际时间步数: {num_cross_section_timesteps}")
# 准备FLOOD_SECTION批量插入数据
section_records = []
for i, cs_name in enumerate(cross_sections_name):
# 查找映射
if cs_name in section_mapping:
section_id, section_name = section_mapping[cs_name]
logger.info(f"处理断面 [{i}] {section_name} (ID: {section_id})")
# 为每个时间步创建一条记录(使用断面数据实际的时间步数)
for j in range(num_cross_section_timesteps):
time_str = time_date_stamp[j].decode('utf-8') if isinstance(time_date_stamp[j], bytes) else str(time_date_stamp[j])
z_value = float(cross_sections_ws[j, i])
q_value = float(cross_sections_flow[j, i])
depth_value = 0 # DEPTH字段暂填0
section_records.append((
section_id,
section_name,
scheme_name,
time_str,
z_value,
depth_value,
q_value
))
else:
cs_name_str = cs_name.decode('utf-8') if isinstance(cs_name, bytes) else str(cs_name)
logger.warning(f"断面 [{i}] {cs_name_str} 未找到映射,跳过")
logger.info(f"准备插入{len(section_records)}条FLOOD_SECTION记录")
if section_records:
success = sqlserver_handler.insert_flood_section_batch(section_records)
if not success:
logger.error("FLOOD_SECTION批量写入失败")
else:
logger.info("FLOOD_SECTION批量写入成功")
# 3. 准备并插入FLOODAREA记录
logger.info("开始准备FLOODAREA数据...")
floodarea_records = []
for j in range(len(time_date_stamp)):
time_str = time_date_stamp[j].decode('utf-8') if isinstance(time_date_stamp[j], bytes) else str(time_date_stamp[j])
flooded_area_value = float(flooded_area[j])
floodarea_records.append((
time_str,
flooded_area_value,
scheme_name
))
logger.info(f"准备插入{len(floodarea_records)}条FLOODAREA记录")
if floodarea_records:
success = sqlserver_handler.insert_floodarea_batch(floodarea_records)
if not success:
logger.error("FLOODAREA批量写入失败")
else:
logger.info("FLOODAREA批量写入成功")
logger.info("数据库写入完成")
except Exception as e:
logger.error(f"写入数据库时出错: {e}")
import traceback
logger.error(traceback.format_exc())
if __name__ == '__main__':
logger.info("=" * 60)
logger.info("开始测试数据库写入和POST接口调用")
logger.info("=" * 60)
test_database_and_post()
logger.info("=" * 60)
logger.info("测试完成")
logger.info("=" * 60)