from decimal import Decimal

from django.contrib.auth import get_user_model
from django.test import TestCase
from rest_framework.authtoken.models import Token
from rest_framework.test import APIClient

from accounts.models import ShipperProfile, TransporterProfile, UserRole
from billing.models import Invoice, Payment, WalletAccount
from core.models import Bid, Shipment, Trip

User = get_user_model()


class PaymentFlowRequirementTests(TestCase):
    def setUp(self):
        self.shipper = User.objects.create_user('ship-pay@test.com', email='ship-pay@test.com', password='secret')
        UserRole.objects.create(user=self.shipper, role=UserRole.Role.SHIPPER, phone='', language='en')
        ShipperProfile.objects.create(user=self.shipper, account_type=ShipperProfile.AccountType.INDIVIDUAL)

        self.transporter = User.objects.create_user('tr-pay@test.com', email='tr-pay@test.com', password='secret')
        UserRole.objects.create(user=self.transporter, role=UserRole.Role.TRANSPORTER, phone='123', language='en')
        TransporterProfile.objects.create(
            user=self.transporter,
            account_type=TransporterProfile.AccountType.DRIVER,
            local=True,
            country_to_country=False,
            documents_verified=True,
        )

        self.shipment = Shipment.objects.create(
            shipper=self.shipper,
            pickup_address='A',
            delivery_address='B',
            cargo_type='General',
            weight='1 ton',
            vehicle_type_required='Flatbed',
            local=True,
            country_to_country=False,
            status=Shipment.Status.DELIVERED,
        )
        bid = Bid.objects.create(
            shipment=self.shipment,
            transporter=self.transporter,
            amount=Decimal('120.00'),
            status=Bid.Status.ACCEPTED,
        )
        self.trip = Trip.objects.create(
            shipment=self.shipment,
            accepted_bid=bid,
            transporter=self.transporter,
            status=Trip.Status.DELIVERED,
        )

    def _shipper_client(self):
        token, _ = Token.objects.get_or_create(user=self.shipper)
        c = APIClient()
        c.credentials(HTTP_AUTHORIZATION=f'Bearer {token.key}')
        return c

    def _transporter_client(self):
        token, _ = Token.objects.get_or_create(user=self.transporter)
        c = APIClient()
        c.credentials(HTTP_AUTHORIZATION=f'Bearer {token.key}')
        return c

    def test_payment_rejected_before_trip_completed(self):
        response = self._shipper_client().post(
            f'/api/shipper/trips/{self.trip.id}/pay/',
            {'method': 'WALLET'},
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        self.assertIn('COMPLETED', str(response.json()).upper())

    def test_invoice_auto_generated_when_trip_moves_to_completed(self):
        response = self._transporter_client().patch(
            f'/api/transporter/trips/{self.trip.id}/status/',
            {'status': 'COMPLETED'},
            format='json',
        )
        self.assertEqual(response.status_code, 200)
        invoice = Invoice.objects.get(trip=self.trip)
        self.assertEqual(invoice.status, Invoice.Status.ISSUED)
        self.assertIsNone(invoice.payment_id)

        detail = self._shipper_client().get(f'/api/shipper/invoices/{invoice.id}/')
        self.assertEqual(detail.status_code, 200)
        self.assertIn('currency', detail.json()['data'])
        self.assertEqual(detail.json()['data']['currency'], 'USD')

    def test_wallet_payment_marks_invoice_paid_and_archived(self):
        WalletAccount.objects.create(user=self.shipper, balance=Decimal('500.00'), currency='USD')
        self._transporter_client().patch(
            f'/api/transporter/trips/{self.trip.id}/status/',
            {'status': 'COMPLETED'},
            format='json',
        )
        response = self._shipper_client().post(
            f'/api/shipper/trips/{self.trip.id}/pay/',
            {'method': 'WALLET'},
            format='json',
        )
        self.assertEqual(response.status_code, 200)
        payload = response.json()['data']['invoice']
        self.assertEqual(payload['status'], 'PAID')
        self.assertTrue(payload['is_archived'])
        self.assertIsNotNone(payload['archived_at'])

    def test_cash_confirm_credits_transporter_without_shipper_wallet_debit(self):
        WalletAccount.objects.create(user=self.shipper, balance=Decimal('0.00'), currency='USD')
        WalletAccount.objects.create(user=self.transporter, balance=Decimal('0.00'), currency='USD')
        self._transporter_client().patch(
            f'/api/transporter/trips/{self.trip.id}/status/',
            {'status': 'COMPLETED'},
            format='json',
        )
        pay_resp = self._shipper_client().post(
            f'/api/shipper/trips/{self.trip.id}/pay/',
            {'method': 'CASH'},
            format='json',
        )
        self.assertEqual(pay_resp.status_code, 200)
        self.assertEqual(pay_resp.json()['data']['payment']['status'], 'PENDING_COD')

        confirm_resp = self._transporter_client().post(
            f'/api/transporter/trips/{self.trip.id}/confirm-cash/',
            format='json',
        )
        self.assertEqual(confirm_resp.status_code, 200)

        shipper_wallet = WalletAccount.objects.get(user=self.shipper)
        transporter_wallet = WalletAccount.objects.get(user=self.transporter)
        self.assertEqual(shipper_wallet.balance, Decimal('0.00'))
        self.assertEqual(transporter_wallet.balance, Decimal('120.00'))

        payment = Payment.objects.get(trip=self.trip)
        self.assertEqual(payment.status, Payment.Status.CAPTURED)
        invoice = Invoice.objects.get(trip=self.trip)
        self.assertEqual(invoice.status, Invoice.Status.PAID)
        self.assertTrue(invoice.is_archived)

    def test_card_pay_then_confirm_credits_transporter_without_shipper_wallet_debit(self):
        WalletAccount.objects.create(user=self.shipper, balance=Decimal('0.00'), currency='USD')
        WalletAccount.objects.create(user=self.transporter, balance=Decimal('0.00'), currency='USD')
        self._transporter_client().patch(
            f'/api/transporter/trips/{self.trip.id}/status/',
            {'status': 'COMPLETED'},
            format='json',
        )
        pay_resp = self._shipper_client().post(
            f'/api/shipper/trips/{self.trip.id}/pay/',
            {'method': 'CARD'},
            format='json',
        )
        self.assertEqual(pay_resp.status_code, 200)
        self.assertEqual(pay_resp.json()['data']['payment']['status'], 'REQUIRES_ACTION')
        self.assertEqual(pay_resp.json()['data']['invoice']['status'], 'ISSUED')

        confirm_resp = self._transporter_client().post(
            f'/api/transporter/trips/{self.trip.id}/confirm-card/',
            format='json',
        )
        self.assertEqual(confirm_resp.status_code, 200)

        shipper_wallet = WalletAccount.objects.get(user=self.shipper)
        transporter_wallet = WalletAccount.objects.get(user=self.transporter)
        self.assertEqual(shipper_wallet.balance, Decimal('0.00'))
        self.assertEqual(transporter_wallet.balance, Decimal('120.00'))

        payment = Payment.objects.get(trip=self.trip)
        self.assertEqual(payment.status, Payment.Status.CAPTURED)
        self.assertEqual(payment.method, Payment.Method.CARD)
        invoice = Invoice.objects.get(trip=self.trip)
        self.assertEqual(invoice.status, Invoice.Status.PAID)
        self.assertTrue(invoice.is_archived)

    def test_shipper_cannot_confirm_card_payment(self):
        self._transporter_client().patch(
            f'/api/transporter/trips/{self.trip.id}/status/',
            {'status': 'COMPLETED'},
            format='json',
        )
        self._shipper_client().post(
            f'/api/shipper/trips/{self.trip.id}/pay/',
            {'method': 'CARD'},
            format='json',
        )
        confirm_resp = self._shipper_client().post(
            f'/api/transporter/trips/{self.trip.id}/confirm-card/',
            format='json',
        )
        self.assertEqual(confirm_resp.status_code, 403)

