"""Stripe billing adapter parameter and response-boundary tests."""

from __future__ import annotations

from datetime import datetime, timezone
import unittest
from unittest.mock import patch

from licensing_server.app.billing import BillingProviderFailure, StripeBillingProvider


class StripeBillingProviderTests(unittest.TestCase):
    def setUp(self) -> None:
        self.provider = StripeBillingProvider("sk_test_billing_provider")

    def test_checkout_uses_subscription_metadata_and_server_redirects(self) -> None:
        with patch(
            "licensing_server.app.billing.stripe.checkout.Session.create",
            return_value={
                "id": "cs_test_provider",
                "url": "https://checkout.stripe.test/session/provider",
                "expires_at": int(datetime(2026, 8, 15, 19, 0, tzinfo=timezone.utc).timestamp()),
            },
        ) as create:
            result = self.provider.create_checkout_session(
                price_id="price_test_ai_assist",
                sku_id="subscription.ai_assist",
                license_id="license.test.provider",
                customer_id=None,
                customer_email="owner@example.test",
                success_url="https://account.example.test/checkout/success",
                cancel_url="https://account.example.test/checkout/cancel",
                idempotency_key="apolon_provider_checkout_test",
            )
        self.assertEqual(result.session_id, "cs_test_provider")
        parameters = create.call_args.kwargs
        self.assertEqual(parameters["mode"], "subscription")
        self.assertEqual(parameters["client_reference_id"], "license.test.provider")
        self.assertEqual(parameters["success_url"], "https://account.example.test/checkout/success")
        self.assertEqual(parameters["cancel_url"], "https://account.example.test/checkout/cancel")
        self.assertEqual(
            parameters["consent_collection"],
            {"terms_of_service": "required"},
        )
        self.assertEqual(
            parameters["subscription_data"]["metadata"]["license_id"],
            "license.test.provider",
        )
        self.assertEqual(parameters["customer_email"], "owner@example.test")
        self.assertNotIn("customer", parameters)

    def test_portal_uses_existing_customer_and_retrieval_is_recursive(self) -> None:
        with patch(
            "licensing_server.app.billing.stripe.billing_portal.Session.create",
            return_value={
                "id": "bps_test_provider",
                "url": "https://billing.stripe.test/session/provider",
            },
        ) as create:
            result = self.provider.create_portal_session(
                customer_id="cus_test_provider",
                return_url="https://account.example.test/licenses",
                idempotency_key="apolon_provider_portal_test",
            )
        self.assertEqual(result.session_id, "bps_test_provider")
        self.assertEqual(create.call_args.kwargs["customer"], "cus_test_provider")

        class RetrievedSubscription:
            def to_dict_recursive(self):
                return {"id": "sub_test_provider", "status": "active"}

        with patch(
            "licensing_server.app.billing.stripe.Subscription.retrieve",
            return_value=RetrievedSubscription(),
        ) as retrieve:
            subscription = self.provider.retrieve_subscription("sub_test_provider")
        self.assertEqual(subscription["id"], "sub_test_provider")
        self.assertEqual(retrieve.call_args.args, ("sub_test_provider",))

    def test_provider_response_must_return_bounded_https_url(self) -> None:
        with patch(
            "licensing_server.app.billing.stripe.checkout.Session.create",
            return_value={
                "id": "cs_test_invalid_provider",
                "url": "http://checkout.example.test/session",
                "expires_at": None,
            },
        ):
            with self.assertRaises(BillingProviderFailure):
                self.provider.create_checkout_session(
                    price_id="price_test_ai_assist",
                    sku_id="subscription.ai_assist",
                    license_id="license.test.provider",
                    customer_id="cus_test_provider",
                    customer_email=None,
                    success_url="https://account.example.test/checkout/success",
                    cancel_url="https://account.example.test/checkout/cancel",
                    idempotency_key="apolon_provider_invalid_test",
                )

    def test_billing_incidents_resolve_through_charge_and_invoice_authority(self) -> None:
        dispute_incident = {
            "providerEventType": "charge.dispute.updated",
            "incidentId": "dp_test_provider",
            "incidentStatus": "needs_response",
            "chargeId": "ch_test_provider",
        }
        with (
            patch(
                "licensing_server.app.billing.stripe.Dispute.retrieve",
                return_value={
                    "id": "dp_test_provider",
                    "charge": "ch_test_provider",
                    "status": "under_review",
                },
            ),
            patch(
                "licensing_server.app.billing.stripe.Charge.retrieve",
                return_value={
                    "id": "ch_test_provider",
                    "amount": 5000,
                    "amount_refunded": 0,
                    "invoice": "in_test_provider",
                },
            ),
            patch(
                "licensing_server.app.billing.stripe.Invoice.retrieve",
                return_value={
                    "id": "in_test_provider",
                    "parent": {
                        "subscription_details": {
                            "subscription": "sub_test_provider"
                        }
                    },
                },
            ),
        ):
            dispute = self.provider.resolve_billing_incident(dispute_incident)
        self.assertEqual(dispute.subscription_id, "sub_test_provider")
        self.assertTrue(dispute.apply_hold)
        self.assertEqual(dispute.hold_status, "dispute_open")

        refund_incident = {
            "providerEventType": "refund.created",
            "incidentId": "re_test_provider",
            "incidentStatus": "succeeded",
            "chargeId": "ch_test_provider",
        }
        with patch(
            "licensing_server.app.billing.stripe.Charge.retrieve",
            return_value={
                "id": "ch_test_provider",
                "amount": 5000,
                "amount_refunded": 5000,
                "invoice": {"id": "in_test_provider", "subscription": "sub_test_provider"},
            },
        ):
            refund = self.provider.resolve_billing_incident(refund_incident)
        self.assertTrue(refund.apply_hold)
        self.assertEqual(refund.hold_status, "refunded")

        with patch(
            "licensing_server.app.billing.stripe.Charge.retrieve",
            return_value={
                "id": "ch_test_provider",
                "amount": 5000,
                "amount_refunded": 1000,
                "invoice": {"id": "in_test_provider", "subscription": "sub_test_provider"},
            },
        ):
            partial_refund = self.provider.resolve_billing_incident(refund_incident)
        self.assertFalse(partial_refund.apply_hold)
        self.assertIsNone(partial_refund.hold_status)

        with patch(
            "licensing_server.app.billing.stripe.Charge.retrieve",
            return_value={
                "id": "ch_test_provider",
                "amount": 5000,
                "amount_refunded": 5001,
                "invoice": {"id": "in_test_provider", "subscription": "sub_test_provider"},
            },
        ):
            with self.assertRaises(BillingProviderFailure):
                self.provider.resolve_billing_incident(refund_incident)


if __name__ == "__main__":
    unittest.main()
