"""Transactional command dispatch for WNDB model mutations.""" from collections.abc import Callable from ..core.database import ( API_ADD, API_DELETE, API_UPDATE, ChangeSet, changes_affect_materialized_views, model_mutation_transaction, refresh_materialized_views_after_commit, ) from ..gis.backdrop import set_backdrop from ..gis.labels import add_label, delete_label, set_label from ..gis.regions import add_region, delete_region, set_region from ..gis.vertices import add_vertex, delete_vertex, set_vertex from ..model.controls import set_control from ..model.curves import add_curve, delete_curve, set_curve from ..model.demands import set_demand from ..model.emitters import set_emitter from ..model.energy import set_energy, set_pump_energy from ..model.junctions import add_junction, delete_junction, set_junction from ..model.mixing import add_mixing, delete_mixing, set_mixing from ..model.options import set_option, set_option_v3 from ..model.patterns import add_pattern, delete_pattern, set_pattern from ..model.pipes import add_pipe, delete_pipe, set_pipe from ..model.pumps import add_pump, delete_pump, set_pump from ..model.quality import set_quality from ..model.reactions import ( set_pipe_reaction, set_reaction, set_tank_reaction, ) from ..model.reservoirs import add_reservoir, delete_reservoir, set_reservoir from ..model.rules import set_rule from ..model.sources import add_source, delete_source, set_source from ..model.status import set_status from ..model.tags import set_tag from ..model.tanks import add_tank, delete_tank, set_tank from ..model.times import set_time from ..model.title import set_title from ..model.valves import add_valve, delete_valve, set_valve from .cascade import expand_command CommandHandler = Callable[[str, ChangeSet], ChangeSet] _ADD_HANDLERS: dict[str, CommandHandler] = { "junction": add_junction, "reservoir": add_reservoir, "tank": add_tank, "pipe": add_pipe, "pump": add_pump, "valve": add_valve, "pattern": add_pattern, "curve": add_curve, "source": add_source, "mixing": add_mixing, "vertex": add_vertex, "label": add_label, "region": add_region, } _UPDATE_HANDLERS: dict[str, CommandHandler] = { "title": set_title, "junction": set_junction, "reservoir": set_reservoir, "tank": set_tank, "pipe": set_pipe, "pump": set_pump, "valve": set_valve, "tag": set_tag, "demand": set_demand, "status": set_status, "pattern": set_pattern, "curve": set_curve, "control": set_control, "rule": set_rule, "energy": set_energy, "pump_energy": set_pump_energy, "emitter": set_emitter, "quality": set_quality, "source": set_source, "reaction": set_reaction, "pipe_reaction": set_pipe_reaction, "tank_reaction": set_tank_reaction, "mixing": set_mixing, "time": set_time, "option": set_option, "option_v3": set_option_v3, "vertex": set_vertex, "label": set_label, "backdrop": set_backdrop, "region": set_region, } _DELETE_HANDLERS: dict[str, CommandHandler] = { "junction": delete_junction, "reservoir": delete_reservoir, "tank": delete_tank, "pipe": delete_pipe, "pump": delete_pump, "valve": delete_valve, "pattern": delete_pattern, "curve": delete_curve, "source": delete_source, "mixing": delete_mixing, "vertex": delete_vertex, "label": delete_label, "region": delete_region, } def _dispatch( handlers: dict[str, CommandHandler], name: str, change_set: ChangeSet ) -> ChangeSet: element_type = change_set.operations[0]["type"] handler = handlers.get(element_type) return handler(name, change_set) if handler else ChangeSet() def _execute_add_command(name: str, change_set: ChangeSet) -> ChangeSet: return _dispatch(_ADD_HANDLERS, name, change_set) def _execute_update_command(name: str, change_set: ChangeSet) -> ChangeSet: return _dispatch(_UPDATE_HANDLERS, name, change_set) def _execute_delete_command(name: str, change_set: ChangeSet) -> ChangeSet: return _dispatch(_DELETE_HANDLERS, name, change_set) def execute_batch_commands(name: str, change_set: ChangeSet) -> ChangeSet: with model_mutation_transaction(name): rewritten = ChangeSet() for operation in change_set.operations: rewritten.merge(expand_command(name, ChangeSet(operation))) result = ChangeSet() for operation in rewritten.operations: operation_type = operation["operation"] if operation_type == API_ADD: result.merge(_execute_add_command(name, ChangeSet(operation))) elif operation_type == API_UPDATE: result.merge(_execute_update_command(name, ChangeSet(operation))) elif operation_type == API_DELETE: result.merge(_execute_delete_command(name, ChangeSet(operation))) if changes_affect_materialized_views(rewritten): refresh_materialized_views_after_commit(name) return result def execute_batch_command(name: str, change_set: ChangeSet) -> ChangeSet: return execute_batch_commands(name, change_set)