import datetime import logging import os from tempfile import NamedTemporaryFile from psycopg import sql from ..core.projects import ( copy_project, delete_project, have_project, temporary_project_name, temporary_template_database, ) from app.infra.db.project_routing import get_schema_template_database_name from ..core.connection import project_transaction from ..core.model_replace import replace_project_model from ..core.database import ( ChangeSet, refresh_materialized_views_after_commit, sql_literal, write, ) from .sections import ( BACKDROP, BOUND, CONTROLS, COORDINATES, CURVES, DEMANDS, EMITTERS, ENERGY, JUNCTIONS, LABELS, MIXING, OPTIONS, PATTERNS, PIPES, PUMPS, QUALITY, REACTIONS, REGION, REGION_NODES, REPORT, RESERVOIRS, RULES, SOURCES, STATUS, TAGS, TANKS, TIMES, TITLE, VALVES, VERTICES, section_name, ) from ..model.title import inp_in_title from ..model.junctions import inp_in_junction from ..model.reservoirs import inp_in_reservoir from ..model.tanks import inp_in_tank from ..model.pipes import inp_in_pipe from ..model.pumps import inp_in_pump from ..model.valves import inp_in_valve from ..model.tags import inp_in_tag from ..model.demands import inp_in_demand from ..model.status import inp_in_status from ..model.patterns import pattern_v3_types, inp_in_pattern from ..model.curves import curve_types, inp_in_curve from ..model.controls import inp_in_control from ..model.rules import inp_in_rule from ..model.energy import inp_in_energy from ..model.emitters import inp_in_emitter from ..model.quality import inp_in_quality from ..model.sources import inp_in_source from ..model.reactions import inp_in_reaction from ..model.mixing import inp_in_mixing from ..model.times import inp_in_time from ..model.reports import inp_in_report from ..model.options_v2 import inp_in_option_v2 from ..model.options_v3 import inp_in_option_v3 from ..gis.coordinates import inp_in_coord from ..gis.vertices import inp_in_vertex from ..gis.labels import inp_in_label from ..gis.backdrop import inp_in_backdrop from ..gis.regions import inp_in_region, inp_in_bound, inp_in_regionnodes from ..gis.region_geometry import to_postgis_polygon # DingZQ, 2024-12-28, export inp from .exporter import export_inp _S = "S" _L = "L" logger = logging.getLogger(__name__) def _inp_in_option(section: list[str], version: str = "3") -> str: return inp_in_option_v3(section) if version == "3" else inp_in_option_v2(section) _handler = { TITLE: (_S, inp_in_title), JUNCTIONS: (_L, inp_in_junction), # line, demand_outside RESERVOIRS: (_L, inp_in_reservoir), TANKS: (_L, inp_in_tank), PIPES: (_L, inp_in_pipe), PUMPS: (_L, inp_in_pump), VALVES: (_L, inp_in_valve), TAGS: (_L, inp_in_tag), DEMANDS: (_L, inp_in_demand), STATUS: (_L, inp_in_status), PATTERNS: (_L, inp_in_pattern), # line, fixed CURVES: (_L, inp_in_curve), CONTROLS: (_L, inp_in_control), RULES: (_L, inp_in_rule), ENERGY: (_L, inp_in_energy), EMITTERS: (_L, inp_in_emitter), QUALITY: (_L, inp_in_quality), SOURCES: (_L, inp_in_source), REACTIONS: (_L, inp_in_reaction), MIXING: (_L, inp_in_mixing), TIMES: (_S, inp_in_time), REPORT: (_S, inp_in_report), OPTIONS: (_S, _inp_in_option), # line, version COORDINATES: (_L, inp_in_coord), VERTICES: (_L, inp_in_vertex), REGION: (_L, inp_in_region), BOUND: (_L, inp_in_bound), REGION_NODES: (_L, inp_in_regionnodes), LABELS: (_L, inp_in_label), BACKDROP: (_S, inp_in_backdrop), # END : 'END', } _level_1 = [ TITLE, PATTERNS, CURVES, CONTROLS, RULES, TIMES, REPORT, OPTIONS, BACKDROP, ] _level_2 = [ JUNCTIONS, RESERVOIRS, TANKS, ] _level_3 = [ PIPES, PUMPS, VALVES, DEMANDS, EMITTERS, QUALITY, SOURCES, MIXING, COORDINATES, LABELS, ] _level_4 = [ TAGS, STATUS, ENERGY, REACTIONS, VERTICES, REGION, BOUND, REGION_NODES, ] map_regiontype = { # map the region types from desktop to server "DISTRIBUTION": "WDA", "DMA": "DMA", "PMA": "PMA", "VD": "VD", "SA": "SA", } class SQLBatch: def __init__(self, project: str, count: int = 100) -> None: self.batch: list[str] = [] self.project = project self.count = count def add(self, sql: str) -> None: self.batch.append(sql) if len(self.batch) == self.count: self.flush() def flush(self) -> None: write(self.project, "".join(self.batch)) self.batch.clear() def _print_time(desc: str) -> datetime.datetime: now = datetime.datetime.now() time = now.strftime("%Y-%m-%d %H:%M:%S") print(f"{time}: {desc}") return now def _get_file_offset(inp: str) -> tuple[dict[str, list[int]], bool]: offset: dict[str, list[int]] = {} current = "" demand_outside = False with open(inp, encoding="utf-8") as f: while True: line = f.readline() if not line: break line = line.strip() if line.startswith("["): for s in section_name: if line.startswith(f"[{s}"): if s not in offset: offset[s] = [] offset[s].append(f.tell()) current = s break elif line != "" and line.startswith(";") == False: if current == DEMANDS: demand_outside = True return (offset, demand_outside) def parse_file(project: str, inp: str, version: str = "3") -> None: start = _print_time(f'Start reading file "{inp}"...') _print_time("First scan...") offset, demand_outside = _get_file_offset(inp) levels = _level_1 + _level_2 + _level_3 + _level_4 # parse the whole section rather than line sections: dict[str, list[str]] = {} for [s, t] in _handler.items(): if t[0] == _S: sections[s] = [] variable_patterns = [] current_pattern = None current_curve = None curve_type_desc_line = None current_region = None current_bound = [] current_bound.clear() region_list = {} sql_batch = SQLBatch(project) _print_time("Second scan...") with open(inp, encoding="utf-8") as f: for s in levels: if s not in offset: continue if s == DEMANDS and demand_outside == False: continue _print_time(f"[{s}]") is_s = _handler[s][0] == _S handler = _handler[s][1] for ptr in offset[s]: f.seek(ptr) while True: line = f.readline() if not line: break line = line.strip() if line.startswith("["): break elif line == "": continue if is_s: sections[s].append(line) else: if line.startswith(";"): if version != "3": # v2 line = line.removeprefix(";") if s == PATTERNS: # ;desc pass elif s == CURVES: # ;type: desc curve_type_desc_line = line continue if s == PATTERNS: tokens = line.split() if tokens[1].upper() in pattern_v3_types: # v3 sql_batch.add( f"insert into network.patterns (id) values ({sql_literal(tokens[0])});" ) current_pattern = tokens[0] if tokens[1].upper() == "VARIABLE": variable_patterns.append(tokens[0]) continue if current_pattern != tokens[0]: sql_batch.add( f"insert into network.patterns (id) values ({sql_literal(tokens[0])});" ) current_pattern = tokens[0] elif s == CURVES: tokens = line.split() if tokens[1].upper() in curve_types: # v3 sql_batch.add( f"insert into network.curves (id, curve_type) values ({sql_literal(tokens[0])}, {sql_literal(tokens[1].upper())});" ) current_curve = tokens[0] continue if current_curve != tokens[0]: type = curve_types[0] if curve_type_desc_line != None: type = curve_type_desc_line.split(":")[0].strip() sql_batch.add( f"insert into network.curves (id, curve_type) values ({sql_literal(tokens[0])}, {sql_literal(type)});" ) current_curve = tokens[0] curve_type_desc_line = None elif s == REGION: tokens = line.split() region_list[tokens[0]] = tokens[1] continue elif s == BOUND: tokens = line.split() if tokens[0] != current_region and len(current_bound) > 0: current_bound.append(current_bound[0]) current_geometry = to_postgis_polygon(current_bound) region_type = map_regiontype.get( region_list[current_region], region_list[current_region], ) sql_batch.add( "insert into gis.regions(id, region_type, boundary) " f"values ({sql_literal(current_region)}, {sql_literal(region_type)}, " f"st_geomfromtext({sql_literal(current_geometry)}, 900914));" ) current_bound.clear() vertex_point = (float(tokens[1]), float(tokens[2])) current_bound.append(vertex_point) current_region = tokens[0] if s == JUNCTIONS: sql_batch.add(handler(line, demand_outside)) elif s == PATTERNS: sql_batch.add( handler(line, current_pattern not in variable_patterns) ) elif s == BOUND: continue else: sql_batch.add(handler(line)) f.seek(0) if is_s: if s == OPTIONS: sql_batch.add(handler(sections[s], version)) else: sql_batch.add(handler(sections[s])) # need to insert the last region into database if len(current_bound) > 0: current_bound.append(current_bound[0]) current_geometry = to_postgis_polygon(current_bound) region_type = map_regiontype.get( region_list[current_region], region_list[current_region], ) sql_batch.add( "insert into gis.regions(id, region_type, boundary) " f"values ({sql_literal(current_region)}, {sql_literal(region_type)}, " f"st_geomfromtext({sql_literal(current_geometry)}, 900914));" ) sql_batch.flush() end = _print_time(f'End reading file "{inp}"') print(f"Total (in second): {(end-start).seconds}(s)") def read_inp(project: str, inp: str, version: str = "3") -> bool: if version != "3" and version != "2": version = "2" if not have_project(project): raise ValueError(f"Project database {project!r} does not exist") staging_project = temporary_project_name(project, "model_import") replacement_committed = False try: copy_project( get_schema_template_database_name(), staging_project, allow_template_source=True, ) with project_transaction(staging_project): parse_file(staging_project, inp, version) replace_project_model(project, staging_project) replacement_committed = True finally: try: if have_project(staging_project): delete_project(staging_project) except Exception: logger.exception( "Failed to remove model-import staging database %s", staging_project, ) if replacement_committed: refresh_materialized_views_after_commit(project) return True # DingZQ, 2024-12-28, convert v3 to v2 def convert_inp_v3_to_v2(inp: str) -> ChangeSet: temp_path: str | None = None with temporary_template_database("conversion", "v3_to_v2") as project: try: with NamedTemporaryFile( mode="w", suffix=".inp", encoding="utf-8", delete=False ) as temp_file: temp_file.write(inp) temp_path = temp_file.name with project_transaction(project): parse_file(project, temp_path, "3") return export_inp(project, "2") finally: if temp_path is not None: os.remove(temp_path) def import_inp(project: str, cs: ChangeSet, version: str = "3") -> bool: if version != "3" and version != "2": version = "2" if "inp" not in cs.operations[0]: return False temp_path: str | None = None try: with NamedTemporaryFile( mode="w", suffix=".inp", prefix="tjwater_import_", encoding="utf-8", delete=False, ) as temp_file: temp_file.write(str(cs.operations[0]["inp"])) temp_path = temp_file.name return read_inp(project, temp_path, version) finally: if temp_path is not None: try: os.remove(temp_path) except FileNotFoundError: pass