"""Transactional notification privacy, retry, recipient, and SMTP tests."""

from __future__ import annotations

from dataclasses import replace
from datetime import datetime, timedelta, timezone
import unittest
from unittest.mock import MagicMock, patch

from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool

from licensing_shared.catalog import LicensingCatalog, load_builtin_catalog
from licensing_shared.constants import PRODUCT_ID
from licensing_server.app.database import Base
from licensing_server.app.models import (
    AuditEvent,
    License,
    Membership,
    NotificationDelivery,
    NotificationFeedbackEvent,
    NotificationSuppression,
    Organization,
    OutboxEvent,
    Subscription,
    SubscriptionItem,
    User,
)
from licensing_server.app.notifications import (
    CustomerNotificationKind,
    NotificationFeedbackType,
    NotificationProviderFailure,
    NotificationSendResult,
    SmtpCustomerNotificationProvider,
    billing_hold_notification_kind,
    clear_notification_suppression,
    enqueue_subscription_notification,
    process_customer_notification_outbox,
    purge_notification_feedback_events,
    record_notification_feedback,
    recover_stale_customer_notification_claims,
    requeue_failed_customer_notifications,
    subscription_transition_notification_kinds,
)


TEST_NOW = datetime(2026, 8, 15, 21, 0, tzinfo=timezone.utc)
TEST_PORTAL_URL = "https://account.example.test/licenses"
TEST_NOTIFICATION_PEPPER = b"notification-test-pepper-value-32"


class RecordingNotificationProvider:
    def __init__(self, failures: int = 0) -> None:
        self.failures = failures
        self.messages = []

    def send(self, message):
        self.messages.append(message)
        if self.failures > 0:
            self.failures -= 1
            raise NotificationProviderFailure("smtp_delivery_failed")
        return NotificationSendResult(
            provider_message_id=f"provider-message-{len(self.messages)}"
        )


class CustomerNotificationTests(unittest.TestCase):
    def setUp(self) -> None:
        self.engine = create_engine(
            "sqlite+pysqlite:///:memory:",
            connect_args={"check_same_thread": False},
            poolclass=StaticPool,
        )
        Base.metadata.create_all(self.engine)
        original = load_builtin_catalog()
        subscription_sku = replace(
            original.skus["subscription.ai_assist"],
            active=True,
        )
        self.catalog = LicensingCatalog(
            product_id=original.product_id,
            revision=original.revision,
            entitlements=original.entitlements,
            device_policies=original.device_policies,
            skus={**original.skus, subscription_sku.sku_id: subscription_sku},
        )
        with Session(self.engine) as session, session.begin():
            session.add(
                User(
                    id="user.test.notification.owner",
                    external_issuer="https://identity.test",
                    external_subject="notification-owner",
                    verified_email="owner@example.test",
                    status="active",
                    is_server_admin=False,
                )
            )
            session.add(
                License(
                    id="license.test.notification",
                    product_id=PRODUCT_ID,
                    owner_user_id="user.test.notification.owner",
                    owner_organization_id=None,
                    anonymous_subject_digest=None,
                    subject_type="user",
                    status="active",
                    device_policy_id="policy.named_user.standard",
                    revocation_generation=0,
                )
            )
            session.add(
                Subscription(
                    id="subscription.test.notification",
                    license_id="license.test.notification",
                    provider="stripe",
                    provider_subscription_id="sub_test_notification",
                    provider_customer_id="cus_test_notification",
                    status="active",
                    current_period_start=TEST_NOW,
                    current_period_end=TEST_NOW + timedelta(days=30),
                    cancel_at_period_end=False,
                    grace_ends_at=None,
                    provider_event_created_at=TEST_NOW,
                    billing_hold_status=None,
                    billing_hold_provider_id=None,
                    billing_hold_event_created_at=None,
                )
            )
            session.add(
                SubscriptionItem(
                    id="subscription_item.test.notification",
                    subscription_id="subscription.test.notification",
                    sku_id="subscription.ai_assist",
                    quantity=1,
                    provider_item_id="si_test_notification",
                )
            )

    def tearDown(self) -> None:
        self.engine.dispose()

    def _enqueue(self, kind: CustomerNotificationKind, event_id: str) -> bool:
        with Session(self.engine) as session, session.begin():
            return enqueue_subscription_notification(
                session,
                "subscription.test.notification",
                kind,
                event_id,
                now=TEST_NOW,
            )

    def _process(self, provider, *, now: datetime = TEST_NOW) -> int:
        with Session(self.engine) as session:
            return process_customer_notification_outbox(
                session,
                self.catalog,
                provider,
                TEST_PORTAL_URL,
                now=now,
                notification_pepper=TEST_NOTIFICATION_PEPPER,
            )

    def test_personal_notification_is_idempotent_private_and_delivered_once(self) -> None:
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.PAYMENT_FAILED,
                "evt_test_notification_payment_failed",
            )
        )
        self.assertFalse(
            self._enqueue(
                CustomerNotificationKind.PAYMENT_FAILED,
                "evt_test_notification_payment_failed",
            )
        )
        provider = RecordingNotificationProvider()
        self.assertEqual(self._process(provider), 1)
        self.assertEqual(self._process(provider), 0)
        self.assertEqual(len(provider.messages), 1)
        message = provider.messages[0]
        self.assertEqual(message.recipient, "owner@example.test")
        self.assertIn("payment needs attention", message.body)
        self.assertIn(TEST_PORTAL_URL, message.body)
        with Session(self.engine) as session:
            outbox = session.scalar(select(OutboxEvent))
            delivery = session.scalar(select(NotificationDelivery))
            audit = session.scalar(
                select(AuditEvent).where(AuditEvent.action == "notification.delivered")
            )
            self.assertEqual(outbox.status, "processed")
            self.assertEqual(delivery.status, "processed")
            self.assertNotIn("owner@example.test", str(outbox.payload_json))
            self.assertNotIn("owner@example.test", str(audit.metadata_json))

    def test_failure_requeues_only_failed_delivery_with_stable_idempotency(self) -> None:
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.DISPUTE_OPEN,
                "evt_test_notification_dispute",
            )
        )
        provider = RecordingNotificationProvider(failures=1)
        self.assertEqual(self._process(provider), 0)
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(OutboxEvent)).status, "failed")
            self.assertEqual(session.scalar(select(NotificationDelivery)).status, "failed")
        with Session(self.engine) as session:
            self.assertEqual(
                requeue_failed_customer_notifications(
                    session,
                    now=TEST_NOW + timedelta(minutes=1),
                ),
                1,
            )
        self.assertEqual(
            self._process(provider, now=TEST_NOW + timedelta(minutes=1)),
            1,
        )
        self.assertEqual(len(provider.messages), 2)
        self.assertEqual(
            provider.messages[0].idempotency_key,
            provider.messages[1].idempotency_key,
        )
        with Session(self.engine) as session:
            delivery = session.scalar(select(NotificationDelivery))
            self.assertEqual(delivery.attempts, 2)
            self.assertEqual(delivery.status, "processed")

    def test_delayed_hard_bounce_suppresses_only_the_delivered_address(self) -> None:
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.PAYMENT_FAILED,
                "evt_test_notification_before_bounce",
            )
        )
        provider = RecordingNotificationProvider()
        self.assertEqual(self._process(provider), 1)
        feedback_time = datetime.now(timezone.utc)
        with Session(self.engine) as session:
            delivery = session.scalar(select(NotificationDelivery))
            self.assertIsNotNone(delivery.recipient_address_digest)
            delivery_id = delivery.id
            provider_message_id = delivery.provider_message_id
        with Session(self.engine) as session:
            self.assertTrue(
                record_notification_feedback(
                    session,
                    "smtp.test",
                    "feedback.test.hard_bounce",
                    NotificationFeedbackType.HARD_BOUNCE,
                    occurred_at=feedback_time,
                    now=feedback_time,
                    delivery_id=delivery_id,
                    provider_message_id=provider_message_id,
                )
            )
        with Session(self.engine) as session:
            self.assertFalse(
                record_notification_feedback(
                    session,
                    "smtp.test",
                    "feedback.test.hard_bounce",
                    NotificationFeedbackType.HARD_BOUNCE,
                    occurred_at=feedback_time,
                    now=feedback_time,
                    delivery_id=delivery_id,
                    provider_message_id=provider_message_id,
                )
            )
        with Session(self.engine) as session:
            with self.assertRaisesRegex(ValueError, "new content"):
                record_notification_feedback(
                    session,
                    "smtp.test",
                    "feedback.test.hard_bounce",
                    NotificationFeedbackType.COMPLAINT,
                    occurred_at=feedback_time,
                    now=feedback_time,
                    delivery_id=delivery_id,
                    provider_message_id=provider_message_id,
                )
        with Session(self.engine) as session, session.begin():
            session.scalar(select(NotificationFeedbackEvent)).delivery_id = None
        with Session(self.engine) as session:
            self.assertFalse(
                record_notification_feedback(
                    session,
                    "smtp.test",
                    "feedback.test.hard_bounce",
                    NotificationFeedbackType.HARD_BOUNCE,
                    occurred_at=feedback_time,
                    now=feedback_time,
                    delivery_id=delivery_id,
                )
            )

        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.ACCESS_SUSPENDED,
                "evt_test_notification_suppressed",
            )
        )
        self.assertEqual(self._process(provider), 1)
        self.assertEqual(len(provider.messages), 1)
        with Session(self.engine) as session:
            suppression = session.scalar(select(NotificationSuppression))
            feedback = session.scalar(select(NotificationFeedbackEvent))
            skipped = session.scalar(
                select(NotificationDelivery).where(
                    NotificationDelivery.last_error_code
                    == "notification_recipient_suppressed"
                )
            )
            audits = tuple(session.scalars(select(AuditEvent)))
            self.assertTrue(suppression.active)
            self.assertEqual(suppression.reason, "hard_bounce")
            self.assertEqual(feedback.feedback_type, "hard_bounce")
            self.assertIsNotNone(skipped)
            self.assertNotIn("owner@example.test", str(suppression.__dict__))
            self.assertNotIn("owner@example.test", str(feedback.__dict__))
            self.assertNotIn(
                "owner@example.test",
                str([value.metadata_json for value in audits]),
            )

        with Session(self.engine) as session, session.begin():
            session.get(User, "user.test.notification.owner").verified_email = (
                "replacement@example.test"
            )
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.PAYMENT_RECOVERED,
                "evt_test_notification_new_address",
            )
        )
        self.assertEqual(self._process(provider), 1)
        self.assertEqual(len(provider.messages), 2)
        self.assertEqual(provider.messages[-1].recipient, "replacement@example.test")
        with Session(self.engine) as session:
            self.assertEqual(
                purge_notification_feedback_events(
                    session,
                    now=feedback_time + timedelta(days=366),
                    retention_days=365,
                ),
                1,
            )
        with Session(self.engine) as session:
            self.assertEqual(
                session.scalar(select(NotificationFeedbackEvent)),
                None,
            )
            suppression = session.scalar(select(NotificationSuppression))
            self.assertTrue(suppression.active)
            self.assertIsNone(suppression.source_feedback_id)

    def test_complaint_can_be_cleared_by_admin_and_soft_bounce_does_not_suppress(self) -> None:
        with Session(self.engine) as session, session.begin():
            session.add(
                User(
                    id="user.test.notification.admin",
                    external_issuer="https://identity.test",
                    external_subject="notification-admin",
                    verified_email="admin@example.test",
                    status="active",
                    is_server_admin=True,
                )
            )
        provider = RecordingNotificationProvider()
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.DISPUTE_OPEN,
                "evt_test_notification_before_complaint",
            )
        )
        self.assertEqual(self._process(provider), 1)
        with Session(self.engine) as session:
            delivery = session.scalar(select(NotificationDelivery))
            delivery_id = delivery.id
        feedback_time = datetime.now(timezone.utc)
        with Session(self.engine) as session:
            self.assertTrue(
                record_notification_feedback(
                    session,
                    "smtp.test",
                    "feedback.test.complaint",
                    NotificationFeedbackType.COMPLAINT,
                    occurred_at=feedback_time,
                    now=feedback_time,
                    delivery_id=delivery_id,
                )
            )
        with Session(self.engine) as session:
            with self.assertRaisesRegex(PermissionError, "administrator"):
                clear_notification_suppression(
                    session,
                    "user.test.notification.owner",
                    "user.test.notification.owner",
                    "customer_request",
                    TEST_NOTIFICATION_PEPPER,
                    now=feedback_time + timedelta(seconds=1),
                )
        with Session(self.engine) as session:
            self.assertTrue(
                clear_notification_suppression(
                    session,
                    "user.test.notification.owner",
                    "user.test.notification.admin",
                    "customer_request",
                    TEST_NOTIFICATION_PEPPER,
                    now=feedback_time + timedelta(seconds=1),
                )
            )
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.DISPUTE_RESOLVED,
                "evt_test_notification_after_clear",
            )
        )
        self.assertEqual(self._process(provider), 1)
        latest_delivery_id = provider.messages[-1].idempotency_key
        soft_bounce_time = datetime.now(timezone.utc)
        with Session(self.engine) as session:
            self.assertTrue(
                record_notification_feedback(
                    session,
                    "smtp.test",
                    "feedback.test.soft_bounce",
                    NotificationFeedbackType.SOFT_BOUNCE,
                    occurred_at=soft_bounce_time,
                    now=soft_bounce_time,
                    delivery_id=latest_delivery_id,
                )
            )
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.PAYMENT_RECOVERED,
                "evt_test_notification_after_soft_bounce",
            )
        )
        self.assertEqual(self._process(provider), 1)
        self.assertEqual(len(provider.messages), 3)
        with Session(self.engine) as session:
            suppression = session.scalar(select(NotificationSuppression))
            self.assertFalse(suppression.active)
            self.assertIsNotNone(suppression.cleared_at)
            audit = session.scalar(
                select(AuditEvent).where(
                    AuditEvent.action == "notification.suppression_cleared"
                )
            )
            self.assertEqual(audit.actor_id, "user.test.notification.admin")

    def test_organization_notification_targets_only_active_billing_managers(self) -> None:
        with Session(self.engine) as session, session.begin():
            personal_license = session.get(License, "license.test.notification")
            personal_license.owner_user_id = None
            personal_license.owner_organization_id = "organization.test.notification"
            personal_license.subject_type = "organization"
            session.add(
                Organization(
                    id="organization.test.notification",
                    display_name="Notification Team",
                    status="active",
                )
            )
            session.add(
                Membership(
                    id="membership.test.notification.owner",
                    organization_id="organization.test.notification",
                    user_id="user.test.notification.owner",
                    role="owner",
                    status="active",
                    valid_until=None,
                )
            )
            for suffix, role, status in (
                ("admin", "admin", "active"),
                ("member", "member", "active"),
                ("inactive", "admin", "inactive"),
            ):
                user_id = f"user.test.notification.{suffix}"
                session.add(
                    User(
                        id=user_id,
                        external_issuer="https://identity.test",
                        external_subject=f"notification-{suffix}",
                        verified_email=f"{suffix}@example.test",
                        status="active",
                        is_server_admin=False,
                    )
                )
                session.add(
                    Membership(
                        id=f"membership.test.notification.{suffix}",
                        organization_id="organization.test.notification",
                        user_id=user_id,
                        role=role,
                        status=status,
                        valid_until=None,
                    )
                )
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.REFUNDED,
                "evt_test_notification_refund",
            )
        )
        provider = RecordingNotificationProvider()
        self.assertEqual(self._process(provider), 1)
        self.assertEqual(
            {message.recipient for message in provider.messages},
            {"owner@example.test", "admin@example.test"},
        )
        with Session(self.engine) as session:
            delivery_user_ids = set(
                session.scalars(select(NotificationDelivery.user_id))
            )
        self.assertEqual(
            delivery_user_ids,
            {
                "user.test.notification.owner",
                "user.test.notification.admin",
            },
        )

        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.PAYMENT_FAILED,
                "evt_test_notification_role_retry",
            )
        )
        retry_provider = RecordingNotificationProvider(failures=2)
        self.assertEqual(self._process(retry_provider), 0)
        with Session(self.engine) as session, session.begin():
            membership = session.get(Membership, "membership.test.notification.admin")
            self.assertIsNotNone(membership)
            membership.status = "removed"
        retry_provider.failures = 0
        with Session(self.engine) as session:
            self.assertEqual(
                requeue_failed_customer_notifications(
                    session,
                    now=TEST_NOW + timedelta(minutes=1),
                ),
                1,
            )
        self.assertEqual(
            self._process(
                retry_provider,
                now=TEST_NOW + timedelta(minutes=1),
            ),
            1,
        )
        self.assertEqual(
            sum(
                message.recipient == "admin@example.test"
                for message in retry_provider.messages
            ),
            1,
        )
        with Session(self.engine) as session:
            admin_delivery_statuses = set(
                session.scalars(
                    select(NotificationDelivery.status).where(
                        NotificationDelivery.user_id
                        == "user.test.notification.admin"
                    )
                )
            )
        self.assertEqual(admin_delivery_statuses, {"processed", "skipped"})

    def test_transition_mapping_stale_recovery_and_tls_smtp_message(self) -> None:
        self.assertEqual(
            subscription_transition_notification_kinds(
                "past_due",
                "active",
                False,
                False,
            ),
            (CustomerNotificationKind.PAYMENT_RECOVERED,),
        )
        self.assertEqual(
            billing_hold_notification_kind("dispute_open", None),
            CustomerNotificationKind.DISPUTE_RESOLVED,
        )
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.SUBSCRIPTION_STARTED,
                "evt_test_notification_stale",
            )
        )
        with Session(self.engine) as session, session.begin():
            outbox = session.scalar(select(OutboxEvent))
            outbox.status = "processing"
            outbox.updated_at = TEST_NOW - timedelta(hours=1)
            session.add(
                NotificationDelivery(
                    id="notification_delivery.test.stale",
                    outbox_event_id=outbox.id,
                    user_id="user.test.notification.owner",
                    status="sending",
                    attempts=1,
                    provider_message_id=None,
                    last_error_code=None,
                    delivered_at=None,
                    updated_at=TEST_NOW - timedelta(hours=1),
                )
            )
        with Session(self.engine) as session:
            self.assertEqual(
                recover_stale_customer_notification_claims(
                    session,
                    now=TEST_NOW,
                ),
                1,
            )
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(OutboxEvent)).status, "pending")
            self.assertEqual(session.scalar(select(NotificationDelivery)).status, "pending")

        smtp_client = MagicMock()
        smtp_connection = smtp_client.__enter__.return_value
        smtp_connection.send_message.return_value = {}
        provider = SmtpCustomerNotificationProvider(
            "smtp.example.test",
            465,
            "implicit",
            "smtp-user",
            "smtp-password",
            "licensing@example.test",
        )
        with patch(
            "licensing_server.app.notifications.smtplib.SMTP_SSL",
            return_value=smtp_client,
        ) as smtp_ssl:
            result = provider.send(
                self._message_for_smtp()
            )
        self.assertTrue(result.provider_message_id.startswith("<apolon-"))
        smtp_ssl.assert_called_once()
        smtp_connection.login.assert_called_once_with("smtp-user", "smtp-password")
        sent_message = smtp_connection.send_message.call_args.args[0]
        self.assertEqual(sent_message["Auto-Submitted"], "auto-generated")
        self.assertEqual(sent_message["Message-ID"], result.provider_message_id)
        self.assertEqual(
            sent_message["X-Apolon-Delivery-ID"],
            "notification_delivery.test.smtp",
        )

    def test_invalid_outbox_payload_is_skipped_without_blocking_valid_work(self) -> None:
        with Session(self.engine) as session, session.begin():
            session.add(
                OutboxEvent(
                    id="notification.test.invalid_payload",
                    event_type="customer.notification.send",
                    aggregate_type="subscription_notification",
                    aggregate_id="subscription.test.notification",
                    payload_json={"schema": "invalid"},
                    status="pending",
                    available_at=TEST_NOW,
                    attempts=0,
                    processed_at=None,
                    last_error_code=None,
                )
            )
        self.assertTrue(
            self._enqueue(
                CustomerNotificationKind.PAYMENT_FAILED,
                "evt_test_notification_after_invalid",
            )
        )
        provider = RecordingNotificationProvider()
        self.assertEqual(self._process(provider), 2)
        self.assertEqual(len(provider.messages), 1)
        with Session(self.engine) as session:
            invalid = session.get(OutboxEvent, "notification.test.invalid_payload")
            self.assertIsNotNone(invalid)
            self.assertEqual(invalid.status, "processed")
            self.assertEqual(
                invalid.last_error_code,
                "notification_payload_invalid",
            )
            audit = session.scalar(
                select(AuditEvent).where(
                    AuditEvent.target_id == "notification.test.invalid_payload"
                )
            )
            self.assertIsNotNone(audit)
            self.assertEqual(audit.metadata_json["notificationKind"], "invalid")

    @staticmethod
    def _message_for_smtp():
        from licensing_server.app.notifications import CustomerNotificationMessage

        return CustomerNotificationMessage(
            recipient="owner@example.test",
            subject="Apolon licensing status",
            body="Review your licensing status.",
            idempotency_key="notification_delivery.test.smtp",
            notification_kind=CustomerNotificationKind.PAYMENT_FAILED,
        )


if __name__ == "__main__":
    unittest.main()
