from __future__ import annotations from typing import Any from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from app.core.database import SessionLocal from app.models.entities import FixtureDefinition, FixtureInstance from app.models.schemas import PatchPayload, PatchValidationPayload class PatchService: def __init__(self, session_factory: async_sessionmaker[AsyncSession] = SessionLocal) -> None: self._session_factory = session_factory async def list_instances(self) -> list[dict[str, object]]: async with self._session_factory() as session: result = await session.execute(select(FixtureInstance).order_by(FixtureInstance.universe, FixtureInstance.start_address)) instances = result.scalars().all() fixtures = await self._definitions_by_ids(session, {instance.definition_id for instance in instances}) return [self._serialize_instance(instance, fixtures.get(instance.definition_id)) for instance in instances] async def validate(self, payload: PatchValidationPayload) -> dict[str, object]: async with self._session_factory() as session: definition = await session.get(FixtureDefinition, payload.definition_id) if definition is None: raise LookupError("Fixture not found") mode = self._resolve_mode(definition, payload.mode_key) if mode is None: raise LookupError("Fixture mode not found") channel_count = int(mode["channel_count"]) end_address = payload.start_address + channel_count - 1 conflicts: list[dict[str, object]] = [] if end_address > 512: conflicts.append( { "type": "range", "message": f"Patch {payload.start_address}-{end_address} overskrider DMX kanal 512.", } ) result = await session.execute( select(FixtureInstance).where(FixtureInstance.universe == payload.universe).order_by(FixtureInstance.start_address) ) instances = result.scalars().all() fixtures = await self._definitions_by_ids(session, {instance.definition_id for instance in instances}) for instance in instances: if payload.exclude_patch_id is not None and instance.id == payload.exclude_patch_id: continue instance_end = instance.start_address + instance.channel_count - 1 overlaps = not (instance_end < payload.start_address or instance.start_address > end_address) if overlaps: definition_for_instance = fixtures.get(instance.definition_id) conflicts.append( { "type": "overlap", "patch_id": instance.id, "name": instance.name, "range": f"{instance.start_address}-{instance_end}", "fixture": self._definition_label(definition_for_instance), "message": f"{instance.name} optager allerede {instance.start_address}-{instance_end}.", } ) occupied = sorted( { address for instance in instances if payload.exclude_patch_id is None or instance.id != payload.exclude_patch_id for address in range(instance.start_address, instance.start_address + instance.channel_count) } ) return { "valid": not conflicts, "range": f"{payload.start_address}-{end_address}", "end_address": end_address, "channel_count": channel_count, "occupied_channels": occupied, "conflicts": conflicts, } async def create_instance(self, payload: PatchPayload) -> dict[str, object]: validation = await self.validate( PatchValidationPayload( universe=payload.universe, definition_id=payload.definition_id, mode_key=payload.mode_key, start_address=payload.start_address, enabled=payload.enabled, ) ) if not validation["valid"]: raise ValueError("Patch overlapper eller overskrider kanalrangen") async with self._session_factory() as session: definition = await session.get(FixtureDefinition, payload.definition_id) if definition is None: raise LookupError("Fixture not found") mode = self._resolve_mode(definition, payload.mode_key) if mode is None: raise LookupError("Fixture mode not found") instance = FixtureInstance( universe=payload.universe, name=payload.name, definition_id=payload.definition_id, mode_key=payload.mode_key, start_address=payload.start_address, channel_count=int(mode["channel_count"]), enabled=payload.enabled, group_names=self._sanitize_groups(payload.group_names), position=payload.position.model_dump(), tags=[], ) session.add(instance) await session.commit() await session.refresh(instance) return self._serialize_instance(instance, definition) async def update_instance(self, patch_id: int, payload: PatchPayload) -> dict[str, object]: validation = await self.validate( PatchValidationPayload( universe=payload.universe, definition_id=payload.definition_id, mode_key=payload.mode_key, start_address=payload.start_address, enabled=payload.enabled, exclude_patch_id=patch_id, ) ) if not validation["valid"]: raise ValueError("Patch overlapper eller overskrider kanalrangen") async with self._session_factory() as session: instance = await session.get(FixtureInstance, patch_id) if instance is None: raise LookupError("Patch not found") definition = await session.get(FixtureDefinition, payload.definition_id) if definition is None: raise LookupError("Fixture not found") mode = self._resolve_mode(definition, payload.mode_key) if mode is None: raise LookupError("Fixture mode not found") instance.universe = payload.universe instance.name = payload.name instance.definition_id = payload.definition_id instance.mode_key = payload.mode_key instance.start_address = payload.start_address instance.channel_count = int(mode["channel_count"]) instance.enabled = payload.enabled instance.group_names = self._sanitize_groups(payload.group_names) instance.position = payload.position.model_dump() await session.commit() await session.refresh(instance) return self._serialize_instance(instance, definition) async def delete_instance(self, patch_id: int) -> dict[str, object]: async with self._session_factory() as session: instance = await session.get(FixtureInstance, patch_id) if instance is None: raise LookupError("Patch not found") deleted_name = instance.name await session.delete(instance) await session.commit() return {"deleted": deleted_name, "id": patch_id} async def _definitions_by_ids(self, session: AsyncSession, ids: set[int]) -> dict[int, FixtureDefinition]: if not ids: return {} result = await session.execute(select(FixtureDefinition).where(FixtureDefinition.id.in_(ids))) definitions = result.scalars().all() return {definition.id: definition for definition in definitions} def _resolve_mode(self, definition: FixtureDefinition, mode_key: str) -> dict[str, Any] | None: modes = definition.normalized_data.get("modes", []) if not isinstance(modes, list): return None for mode in modes: if isinstance(mode, dict) and str(mode.get("key")) == mode_key: return mode return None def _serialize_instance( self, instance: FixtureInstance, definition: FixtureDefinition | None, ) -> dict[str, object]: end_address = instance.start_address + instance.channel_count - 1 mode = self._resolve_mode(definition, instance.mode_key) if definition is not None else None return { "id": instance.id, "universe": instance.universe, "name": instance.name, "definition_id": instance.definition_id, "fixture_slug": definition.slug if definition is not None else None, "manufacturer": definition.manufacturer if definition is not None else None, "model": definition.model if definition is not None else None, "mode_key": instance.mode_key, "start_address": instance.start_address, "end_address": end_address, "channel_count": instance.channel_count, "enabled": instance.enabled, "group_names": self._sanitize_groups(instance.group_names), "position": self._serialize_position(instance.position), "channels": mode.get("channels", []) if isinstance(mode, dict) else [], } def _definition_label(self, definition: FixtureDefinition | None) -> str: if definition is None: return "Ukendt fixture" return f"{definition.manufacturer} {definition.model}" def _sanitize_groups(self, groups: list[str]) -> list[str]: unique: list[str] = [] seen: set[str] = set() for group in groups: normalized = str(group).strip() key = normalized.casefold() if not normalized or key in seen: continue unique.append(normalized) seen.add(key) return unique def _serialize_position(self, payload: object) -> dict[str, float]: if not isinstance(payload, dict): payload = {} return { "x": float(payload.get("x", 50.0)), "y": float(payload.get("y", 50.0)), "z": float(payload.get("z", 0.0)), "rotation": float(payload.get("rotation", 0.0)), }