from __future__ import annotations

import logging
from datetime import datetime
from enum import Enum
from typing import Literal

from pydantic import BaseModel
from sqlalchemy import and_, func, select

from app.db import NotificationReadRow, NotificationRow, SessionLocal, UserRow

logger = logging.getLogger(__name__)

def resolve_client_id(platform: str, email: str, requested: str | None = None) -> str | None:
    """Resolve which ledger client a user may access.

    Client-supplied ``requested`` ids are ignored. Authorization always
    comes from the user row in the database.
    """
    del requested
    if platform == "firm":
        return None
    email_n = email.strip().lower()
    with SessionLocal() as session:
        row = session.scalar(
            select(UserRow).where(UserRow.platform == platform, UserRow.email == email_n)
        )
        if row and row.client_id:
            return row.client_id
    return None


class NotificationType(str, Enum):
    DOCUMENT_UPLOADED = "DOCUMENT_UPLOADED"
    DOCUMENT_PROCESSING = "DOCUMENT_PROCESSING"
    DOCUMENT_READY = "DOCUMENT_READY"
    DOCUMENT_NEEDS_REVIEW = "DOCUMENT_NEEDS_REVIEW"
    ACCOUNTANT_QUESTION = "ACCOUNTANT_QUESTION"
    QUESTION_ANSWERED = "QUESTION_ANSWERED"
    QUESTION_RESOLVED = "QUESTION_RESOLVED"
    FIRM_QUERY_REPLY = "FIRM_QUERY_REPLY"
    INVOICE_APPROVED = "INVOICE_APPROVED"
    INVOICE_REJECTED = "INVOICE_REJECTED"
    SNLSTART_SYNC_SUCCESS = "SNLSTART_SYNC_SUCCESS"
    SNLSTART_SYNC_FAILED = "SNLSTART_SYNC_FAILED"
    CLIENT_REMINDER = "CLIENT_REMINDER"
    DOCUMENT_DUPLICATE = "DOCUMENT_DUPLICATE"
    FINANCIAL_REPORT_READY = "FINANCIAL_REPORT_READY"
    AI_ADVICE_REQUIRES_APPROVAL = "AI_ADVICE_REQUIRES_APPROVAL"


class NotificationOut(BaseModel):
    id: int
    recipient_platform: str
    client_id: str | None = None
    type: str
    title: str
    message: str
    reference_type: str | None = None
    reference_id: str | None = None
    priority: str = "normal"
    channel: str = "in_app"
    is_read: bool = False
    created_at: datetime
    read_at: datetime | None = None


class NotificationListOut(BaseModel):
    items: list[NotificationOut]
    unread_count: int
    total: int


Channel = Literal["in_app", "email", "sms"]


def create_notification(
    *,
    recipient_platform: Literal["firm", "client"],
    type: NotificationType | str,
    title: str,
    message: str,
    client_id: str | None = None,
    reference_type: str | None = None,
    reference_id: str | None = None,
    priority: Literal["normal", "high"] = "normal",
    channel: Channel = "in_app",
) -> NotificationRow | None:
    """
    Persist a notification. Never raises to callers — failures are logged so
    primary business operations keep succeeding.
    """
    try:
        if recipient_platform == "client" and not client_id:
            logger.warning("Skipping client notification without client_id type=%s", type)
            return None

        type_value = type.value if isinstance(type, NotificationType) else str(type)

        with SessionLocal() as session:
            row = NotificationRow(
                recipient_platform=recipient_platform,
                client_id=client_id if recipient_platform == "client" else client_id,
                type=type_value,
                title=title.strip(),
                message=message.strip(),
                reference_type=reference_type,
                reference_id=reference_id,
                priority=priority,
                channel=channel,
                created_at=datetime.utcnow(),
            )
            # Firm-wide notifications keep optional client_id for context (which client)
            session.add(row)
            session.commit()
            session.refresh(row)
            session.expunge(row)
            return row
    except Exception:
        logger.exception("Failed to create notification type=%s", type)
        return None


def notify_firm(
    *,
    type: NotificationType | str,
    title: str,
    message: str,
    client_id: str | None = None,
    reference_type: str | None = None,
    reference_id: str | None = None,
    priority: Literal["normal", "high"] = "normal",
) -> NotificationRow | None:
    return create_notification(
        recipient_platform="firm",
        type=type,
        title=title,
        message=message,
        client_id=client_id,
        reference_type=reference_type,
        reference_id=reference_id,
        priority=priority,
    )


def notify_client(
    *,
    client_id: str,
    type: NotificationType | str,
    title: str,
    message: str,
    reference_type: str | None = None,
    reference_id: str | None = None,
    priority: Literal["normal", "high"] = "normal",
) -> NotificationRow | None:
    return create_notification(
        recipient_platform="client",
        type=type,
        title=title,
        message=message,
        client_id=client_id,
        reference_type=reference_type,
        reference_id=reference_id,
        priority=priority,
    )


def _authorized_filter(platform: str, client_id: str | None):
    if platform == "firm":
        return NotificationRow.recipient_platform == "firm"
    return and_(
        NotificationRow.recipient_platform == "client",
        NotificationRow.client_id == client_id,
    )


def list_notifications(
    *,
    platform: str,
    email: str,
    client_id: str | None = None,
    unread_only: bool = False,
    limit: int = 50,
    offset: int = 0,
) -> NotificationListOut:
    email_n = email.strip().lower()
    resolved_client = resolve_client_id(platform, email_n, client_id)
    auth_filter = _authorized_filter(platform, resolved_client)

    with SessionLocal() as session:
        total = session.scalar(select(func.count()).select_from(NotificationRow).where(auth_filter)) or 0

        read_subq = (
            select(NotificationReadRow.notification_id, NotificationReadRow.read_at)
            .where(
                NotificationReadRow.platform == platform,
                NotificationReadRow.email == email_n,
            )
            .subquery()
        )

        stmt = (
            select(NotificationRow, read_subq.c.read_at)
            .outerjoin(read_subq, read_subq.c.notification_id == NotificationRow.id)
            .where(auth_filter)
            .order_by(NotificationRow.created_at.desc())
        )
        if unread_only:
            stmt = stmt.where(read_subq.c.read_at.is_(None))

        rows = session.execute(stmt.offset(offset).limit(limit)).all()

        unread_count = session.scalar(
            select(func.count())
            .select_from(NotificationRow)
            .outerjoin(
                read_subq,
                read_subq.c.notification_id == NotificationRow.id,
            )
            .where(auth_filter, read_subq.c.read_at.is_(None))
        ) or 0

        items: list[NotificationOut] = []
        for row, read_at in rows:
            items.append(
                NotificationOut(
                    id=row.id,
                    recipient_platform=row.recipient_platform,
                    client_id=row.client_id,
                    type=row.type,
                    title=row.title,
                    message=row.message,
                    reference_type=row.reference_type,
                    reference_id=row.reference_id,
                    priority=row.priority,
                    channel=row.channel,
                    is_read=read_at is not None,
                    created_at=row.created_at,
                    read_at=read_at,
                )
            )

        return NotificationListOut(items=items, unread_count=unread_count, total=total)


def unread_count(*, platform: str, email: str, client_id: str | None = None) -> int:
    return list_notifications(
        platform=platform,
        email=email,
        client_id=client_id,
        limit=1,
    ).unread_count


def mark_read(*, platform: str, email: str, notification_id: int, client_id: str | None = None) -> NotificationOut | None:
    email_n = email.strip().lower()
    resolved_client = resolve_client_id(platform, email_n, client_id)
    auth_filter = _authorized_filter(platform, resolved_client)

    with SessionLocal() as session:
        row = session.scalar(select(NotificationRow).where(NotificationRow.id == notification_id, auth_filter))
        if not row:
            return None

        existing = session.scalar(
            select(NotificationReadRow).where(
                NotificationReadRow.notification_id == notification_id,
                NotificationReadRow.platform == platform,
                NotificationReadRow.email == email_n,
            )
        )
        now = datetime.utcnow()
        if existing:
            read_at = existing.read_at
        else:
            session.add(
                NotificationReadRow(
                    notification_id=notification_id,
                    platform=platform,
                    email=email_n,
                    read_at=now,
                )
            )
            session.commit()
            read_at = now

        return NotificationOut(
            id=row.id,
            recipient_platform=row.recipient_platform,
            client_id=row.client_id,
            type=row.type,
            title=row.title,
            message=row.message,
            reference_type=row.reference_type,
            reference_id=row.reference_id,
            priority=row.priority,
            channel=row.channel,
            is_read=True,
            created_at=row.created_at,
            read_at=read_at,
        )


def mark_all_read(*, platform: str, email: str, client_id: str | None = None) -> int:
    email_n = email.strip().lower()
    resolved_client = resolve_client_id(platform, email_n, client_id)
    auth_filter = _authorized_filter(platform, resolved_client)

    with SessionLocal() as session:
        read_ids = set(
            session.scalars(
                select(NotificationReadRow.notification_id).where(
                    NotificationReadRow.platform == platform,
                    NotificationReadRow.email == email_n,
                )
            ).all()
        )
        unread_ids = [
            nid
            for nid in session.scalars(select(NotificationRow.id).where(auth_filter)).all()
            if nid not in read_ids
        ]
        now = datetime.utcnow()
        for nid in unread_ids:
            session.add(
                NotificationReadRow(
                    notification_id=nid,
                    platform=platform,
                    email=email_n,
                    read_at=now,
                )
            )
        session.commit()
        return len(unread_ids)


# --- Domain helpers (safe wrappers for store hooks) ---


def on_document_uploaded(*, client_id: str, client_name: str, invoice_id: str, file_name: str) -> None:
    notify_firm(
        type=NotificationType.DOCUMENT_UPLOADED,
        title="New document uploaded",
        message=f"{client_name} uploaded {file_name}.",
        client_id=client_id,
        reference_type="invoice",
        reference_id=invoice_id,
        priority="normal",
    )


def on_document_needs_review(
    *,
    client_id: str,
    invoice_id: str,
    supplier: str,
    total: float,
    file_name: str | None = None,
) -> None:
    label = (supplier or file_name or "A document").strip() or "A document"
    if total and total > 0:
        message = f"{label} · €{total:,.2f} requires your review."
    else:
        message = f"{label} requires your review."
    notify_firm(
        type=NotificationType.DOCUMENT_NEEDS_REVIEW,
        title="New document requires review",
        message=message,
        client_id=client_id,
        reference_type="invoice",
        reference_id=invoice_id,
        priority="high",
    )


def on_accountant_question(*, client_id: str, query_id: str, invoice_id: str | None, supplier: str | None, question: str) -> None:
    about = supplier or "one of your documents"
    notify_client(
        client_id=client_id,
        type=NotificationType.ACCOUNTANT_QUESTION,
        title="Your accountant has a question",
        message=f"Your accountant needs some information about your {about} invoice. “{question[:120]}”",
        reference_type="query",
        reference_id=query_id,
        priority="high",
    )


def on_question_answered(*, client_id: str, client_name: str, query_id: str, invoice_id: str | None, supplier: str | None) -> None:
    about = supplier or "a document"
    notify_firm(
        type=NotificationType.QUESTION_ANSWERED,
        title="Client answered your question",
        message=f"{client_name} answered your question about {about}.",
        client_id=client_id,
        reference_type="query",
        reference_id=query_id,
        priority="high",
    )


def on_firm_query_reply(
    *,
    client_id: str,
    query_id: str,
    invoice_id: str | None,
    supplier: str | None,
    message: str,
) -> None:
    about = supplier or "your documents"
    notify_client(
        client_id=client_id,
        type=NotificationType.FIRM_QUERY_REPLY,
        title="Your accountant replied",
        message=f"Your accountant sent a follow-up about {about}. “{message[:100]}”",
        reference_type="query",
        reference_id=query_id,
        priority="normal",
    )


def on_question_resolved(*, client_id: str, query_id: str, invoice_id: str | None) -> None:
    notify_client(
        client_id=client_id,
        type=NotificationType.QUESTION_RESOLVED,
        title="Question resolved",
        message="Your accountant resolved a question.",
        reference_type="query",
        reference_id=query_id,
        priority="normal",
    )


def on_invoice_approved(*, client_id: str, invoice_id: str, supplier: str) -> None:
    notify_client(
        client_id=client_id,
        type=NotificationType.INVOICE_APPROVED,
        title="Document approved",
        message=f"Your document “{supplier}” has been approved.",
        reference_type="invoice",
        reference_id=invoice_id,
        priority="normal",
    )


def on_invoice_rejected(
    *,
    client_id: str,
    invoice_id: str,
    supplier: str,
    reason: str | None = None,
) -> None:
    reason_n = (reason or "").strip()
    if reason_n:
        message = f"Your document was rejected: {reason_n}"
    else:
        message = f"Your document “{supplier}” was rejected. Please check with your accountant."
    notify_client(
        client_id=client_id,
        type=NotificationType.INVOICE_REJECTED,
        title="Document rejected",
        message=message,
        reference_type="invoice",
        reference_id=invoice_id,
        priority="high",
    )


def on_sync_success(*, client_id: str, invoice_id: str, number: str) -> None:
    notify_firm(
        type=NotificationType.SNLSTART_SYNC_SUCCESS,
        title="Invoice synced",
        message=f"Invoice {number} was successfully synced to SnelStart.",
        client_id=client_id,
        reference_type="invoice",
        reference_id=invoice_id,
        priority="normal",
    )


def on_sync_failed(*, client_id: str, invoice_id: str, number: str) -> None:
    notify_firm(
        type=NotificationType.SNLSTART_SYNC_FAILED,
        title="SnelStart sync failed",
        message=f"Invoice {number} could not be synchronized.",
        client_id=client_id,
        reference_type="invoice",
        reference_id=invoice_id,
        priority="high",
    )


def on_document_duplicate_confirmed(*, client_id: str, invoice_id: str, file_name: str) -> None:
    notify_client(
        client_id=client_id,
        type=NotificationType.DOCUMENT_DUPLICATE,
        title="Duplicate document noted",
        message=f"Your accountant marked “{file_name}” as a duplicate.",
        reference_type="invoice",
        reference_id=invoice_id,
        priority="normal",
    )


def on_upload_reminder(*, client_id: str, client_name: str) -> None:
    notify_client(
        client_id=client_id,
        type=NotificationType.CLIENT_REMINDER,
        title="Reminder to upload invoices",
        message=f"Hi {client_name}, please upload your recent invoices so your accountant can keep your books up to date.",
        reference_type="client",
        reference_id=client_id,
        priority="normal",
    )
