"""Database boundary for messaging state and concurrency guarantees."""

from __future__ import annotations

from datetime import datetime, timedelta

from sqlalchemy import func, or_, select
from sqlalchemy.exc import IntegrityError

from config.control_plane import get_control_plane_session
from .models import MessagingDelivery, MessagingDeliveryAttempt, MessagingPermission, MessagingProvider, MessagingTemplate
from .templates import BUILTIN_TEMPLATES, TemplateDefinition


class MessagingRepository:
    def __init__(self, session_factory=get_control_plane_session) -> None:
        self.session_factory = session_factory

    def permission(self, tenant_id: str, requesting_app: str) -> MessagingPermission | None:
        with self.session_factory() as session:
            return session.execute(select(MessagingPermission).where(
                MessagingPermission.tenant_id == tenant_id,
                MessagingPermission.requesting_app == requesting_app,
                MessagingPermission.enabled.is_(True),
            )).scalar_one_or_none()

    def template(self, tenant_id: str, key: str, channel: str) -> TemplateDefinition | None:
        with self.session_factory() as session:
            record = session.execute(select(MessagingTemplate).where(
                or_(MessagingTemplate.tenant_id == tenant_id, MessagingTemplate.tenant_id.is_(None)),
                MessagingTemplate.template_key == key, MessagingTemplate.channel == channel,
                MessagingTemplate.enabled.is_(True),
            ).order_by(MessagingTemplate.tenant_id.desc(), MessagingTemplate.version.desc())).scalars().first()
            if record:
                return TemplateDefinition(record.template_key, record.channel, record.version,
                                          tuple(record.allowed_variables_json or ()), record.text_body,
                                          record.subject, record.html_body)
        return BUILTIN_TEMPLATES.get((key, channel))

    def providers(self, tenant_id: str, channel: str) -> list[MessagingProvider]:
        with self.session_factory() as session:
            return list(session.execute(select(MessagingProvider).where(
                MessagingProvider.tenant_id == tenant_id, MessagingProvider.channel == channel,
                MessagingProvider.enabled.is_(True),
                MessagingProvider.health_status.in_(("active", "degraded")),
            ).order_by(MessagingProvider.priority, MessagingProvider.created_at)).scalars())

    def provider(self, tenant_id: str, provider_id: str) -> MessagingProvider | None:
        with self.session_factory() as session:
            return session.execute(select(MessagingProvider).where(
                MessagingProvider.tenant_id == tenant_id,
                MessagingProvider.provider_id == provider_id,
            )).scalar_one_or_none()

    def delivery(self, tenant_id: str, delivery_id: str) -> MessagingDelivery | None:
        with self.session_factory() as session:
            return session.execute(select(MessagingDelivery).where(
                MessagingDelivery.tenant_id == tenant_id, MessagingDelivery.delivery_id == delivery_id,
            )).scalar_one_or_none()

    def existing(self, tenant_id: str, requesting_app: str, operation: str, key: str) -> MessagingDelivery | None:
        with self.session_factory() as session:
            return session.execute(select(MessagingDelivery).where(
                MessagingDelivery.tenant_id == tenant_id, MessagingDelivery.requesting_app == requesting_app,
                MessagingDelivery.operation == operation, MessagingDelivery.idempotency_key == key,
            )).scalar_one_or_none()

    def create_delivery(self, **values) -> tuple[MessagingDelivery, bool]:
        with self.session_factory() as session:
            delivery = MessagingDelivery(**values)
            session.add(delivery)
            try:
                session.commit()
                session.refresh(delivery)
                return delivery, True
            except IntegrityError:
                session.rollback()
                existing = session.execute(select(MessagingDelivery).where(
                    MessagingDelivery.tenant_id == values["tenant_id"],
                    MessagingDelivery.requesting_app == values["requesting_app"],
                    MessagingDelivery.operation == values["operation"],
                    MessagingDelivery.idempotency_key == values["idempotency_key"],
                )).scalar_one()
                return existing, False

    def complete_attempt(self, delivery_id: str, provider_id: str, attempt_number: int, result) -> MessagingDelivery:
        now = datetime.utcnow()
        with self.session_factory() as session:
            delivery = session.get(MessagingDelivery, delivery_id)
            session.add(MessagingDeliveryAttempt(
                delivery_id=delivery_id, provider_id=provider_id, attempt_number=attempt_number,
                status=result.status, error_category=result.error_category, provider_code=result.provider_code,
                retryable=result.retryable, latency_ms=result.latency_ms, completed_at=now,
            ))
            delivery.provider_id = provider_id
            delivery.status = result.status
            delivery.error_category = result.error_category
            delivery.provider_code = result.provider_code
            provider = session.get(MessagingProvider, provider_id)
            if result.accepted:
                provider.last_success_at, provider.last_error, provider.consecutive_failures = now, None, 0
            else:
                provider.last_error = result.error_category or "provider_error"
                provider.consecutive_failures += 1
            session.commit()
            session.refresh(delivery)
            return delivery

    def count_since(self, *, tenant_id: str, since: datetime, provider_id: str | None = None,
                    recipient: str | None = None, purpose: str | None = None) -> int:
        with self.session_factory() as session:
            filters = [MessagingDelivery.tenant_id == tenant_id, MessagingDelivery.created_at >= since]
            if provider_id: filters.append(MessagingDelivery.provider_id == provider_id)
            if recipient: filters.append(MessagingDelivery.recipient_normalized == recipient)
            if purpose: filters.append(MessagingDelivery.purpose == purpose)
            return int(session.execute(select(func.count()).select_from(MessagingDelivery).where(*filters)).scalar_one())

    def update_health(self, provider_id: str, result) -> None:
        with self.session_factory() as session:
            provider = session.get(MessagingProvider, provider_id)
            provider.last_health_check_at = datetime.utcnow()
            provider.health_status = result.status
            provider.last_error = result.error_category
            provider.consecutive_failures = 0 if result.accepted else provider.consecutive_failures + 1
            session.commit()

    def save_provider(self, tenant_id: str, values: dict) -> MessagingProvider:
        with self.session_factory() as session:
            provider_id = values.get("provider_id")
            provider = session.get(MessagingProvider, provider_id) if provider_id else None
            if provider and provider.tenant_id != tenant_id:
                raise PermissionError("Provider belongs to another tenant")
            if provider is None:
                provider = MessagingProvider(tenant_id=tenant_id)
                session.add(provider)
            for field in ("name", "driver", "channel", "enabled", "priority", "config_json", "secret_ref",
                          "daily_limit", "per_minute_limit", "max_retries"):
                if field in values: setattr(provider, field, values[field])
            session.commit(); session.refresh(provider); return provider

    def save_permission(self, tenant_id: str, values: dict) -> MessagingPermission:
        with self.session_factory() as session:
            app = str(values["requesting_app"])
            record = session.execute(select(MessagingPermission).where(
                MessagingPermission.tenant_id == tenant_id, MessagingPermission.requesting_app == app,
            )).scalar_one_or_none() or MessagingPermission(tenant_id=tenant_id, requesting_app=app)
            session.add(record)
            for field in ("enabled", "channels_json", "templates_json", "purposes_json", "allow_fallback"):
                if field in values: setattr(record, field, values[field])
            session.commit(); session.refresh(record); return record
