from decimal import Decimal
import base64
from io import BytesIO
from unittest.mock import patch

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

from accounts.models import ShipperProfile, TransporterProfile, UserRole, Vehicle
from core.models import Shipment, Zone

User = get_user_model()


def seed_zones(*country_codes, allow_return_trips=False, radius_km=Decimal('50')):
    for code in country_codes:
        Zone.objects.get_or_create(
            country_code=code,
            defaults={'allow_return_trips': allow_return_trips, 'radius_km': radius_km},
        )


_LOAD_TYPE = {'local': True, 'country_to_country': False}


class ShipmentPublishCountryCodeTests(TestCase):
    def setUp(self):
        self.shipper = User.objects.create_user('pubship@test.com', email='pubship@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)
        token, _ = Token.objects.get_or_create(user=self.shipper)
        self.client = APIClient()
        self.client.credentials(HTTP_AUTHORIZATION=f'Bearer {token.key}')

        self.shipment = Shipment.objects.create(
            shipper=self.shipper,
            pickup_address='Pickup',
            delivery_address='Delivery',
            cargo_type='General',
            weight='1 ton',
            vehicle_type_required='Flatbed',
            local=True,
            country_to_country=False,
            status=Shipment.Status.DRAFT,
        )

    def test_publish_requires_country_codes(self):
        response = self.client.patch(
            f'/api/shipper/shipments/{self.shipment.id}/',
            {
                **_LOAD_TYPE,
                'status': Shipment.Status.PUBLISHED},
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        self.assertIn('pickup_country_code', response.json()['error'])
        self.assertIn('delivery_country_code', response.json()['error'])
        self.shipment.refresh_from_db()
        self.assertEqual(self.shipment.status, Shipment.Status.DRAFT)

    def test_publish_succeeds_with_valid_country_codes(self):
        seed_zones('AE', 'SA')
        response = self.client.patch(
            f'/api/shipper/shipments/{self.shipment.id}/',
            {
                **_LOAD_TYPE,

                'status': Shipment.Status.PUBLISHED,
                'pickup_country_code': 'AE',
                'delivery_country_code': 'SA',
            },
            format='json',
        )
        self.assertEqual(response.status_code, 200)
        self.shipment.refresh_from_db()
        self.assertEqual(self.shipment.status, Shipment.Status.PUBLISHED)
        self.assertEqual(self.shipment.pickup_country_code, 'AE')
        self.assertEqual(self.shipment.delivery_country_code, 'SA')

    @patch('core.geocoding.reverse_geocode_country_code')
    def test_publish_auto_fills_country_codes_from_coordinates(self, mock_reverse):
        seed_zones('AE', 'SA')
        mock_reverse.side_effect = lambda lat, lon: 'AE' if float(lon) > 60 else 'SA'
        self.shipment.pickup_lat = Decimal('24.8607000')
        self.shipment.pickup_lon = Decimal('67.0011000')
        self.shipment.delivery_lat = Decimal('24.9000000')
        self.shipment.delivery_lon = Decimal('46.7000000')
        self.shipment.save()

        response = self.client.patch(
            f'/api/shipper/shipments/{self.shipment.id}/',
            {
                **_LOAD_TYPE,
                'status': Shipment.Status.PUBLISHED},
            format='json',
        )
        self.assertEqual(response.status_code, 200)
        self.shipment.refresh_from_db()
        self.assertEqual(self.shipment.pickup_country_code, 'AE')
        self.assertEqual(self.shipment.delivery_country_code, 'SA')

    def test_publish_rejects_country_outside_zones(self):
        seed_zones('AE')
        response = self.client.patch(
            f'/api/shipper/shipments/{self.shipment.id}/',
            {
                **_LOAD_TYPE,

                'status': Shipment.Status.PUBLISHED,
                'pickup_country_code': 'AE',
                'delivery_country_code': 'ZA',
            },
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        body = response.json()
        self.assertEqual(body['message'], 'Validation failed.')
        self.assertEqual(body['error'], 'We are not operating in South Africa.')
        self.assertIsNone(body['data'])


class ShipmentCreateUpdateCountryCodeTests(TestCase):
    def setUp(self):
        self.shipper = User.objects.create_user('crship@test.com', email='crship@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)
        token, _ = Token.objects.get_or_create(user=self.shipper)
        self.client = APIClient()
        self.client.credentials(HTTP_AUTHORIZATION=f'Bearer {token.key}')

    def _shipment_payload(self, **overrides):
        payload = {
            'pickup_address': 'Pickup',
            'delivery_address': 'Delivery',
            'cargo_type': 'General',
            'weight': '1 ton',
            'vehicle_type_required': 'Flatbed',
            'local': True,
            'country_to_country': False,
        }
        payload.update(overrides)
        return payload

    def test_create_saves_pickup_and_dropoff_company(self):
        seed_zones('AE', 'SA')
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_country_code='AE',
                delivery_country_code='SA',
                pickup_company='Acme Warehouse',
                dropoff_company='Lahore Depot',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 201)
        data = response.json()['data']
        self.assertEqual(data['pickup_company'], 'Acme Warehouse')
        self.assertEqual(data['dropoff_company'], 'Lahore Depot')
        shipment = Shipment.objects.get(pk=data['id'])
        self.assertEqual(shipment.pickup_company, 'Acme Warehouse')
        self.assertEqual(shipment.dropoff_company, 'Lahore Depot')

    def test_create_accepts_country_codes(self):
        seed_zones('AE', 'SA')
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_country_code='AE',
                delivery_country_code='SA',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 201)
        data = response.json()['data']
        self.assertEqual(data['pickup_country_code'], 'AE')
        self.assertEqual(data['delivery_country_code'], 'SA')

    def test_create_accepts_luggage_image(self):
        seed_zones('AE', 'SA')
        image_buffer = BytesIO()
        Image.new('RGB', (8, 8), color='red').save(image_buffer, format='JPEG')
        encoded = base64.b64encode(image_buffer.getvalue()).decode()
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_country_code='AE',
                delivery_country_code='SA',
                luggage_image=f'data:image/jpeg;base64,{encoded}',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 201)
        data = response.json()['data']
        self.assertTrue(data['luggage_image_url'])
        self.assertTrue(data['luggage_image_url'].startswith('http'))
        shipment = Shipment.objects.get(pk=data['id'])
        self.assertEqual(shipment.luggage_image_url, data['luggage_image_url'])

    def test_create_rejects_invalid_country_codes(self):
        seed_zones('SA')
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_country_code='INVALID',
                delivery_country_code='SA',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        self.assertIn('pickup_country_code', response.json()['error'])

    @patch('core.geocoding.reverse_geocode_country_code')
    def test_create_auto_fills_country_codes_from_coordinates(self, mock_reverse):
        seed_zones('AE', 'SA')
        mock_reverse.side_effect = lambda lat, lon: 'AE' if float(lon) > 60 else 'SA'
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_lat='24.8607000',
                pickup_lon='67.0011000',
                delivery_lat='24.9000000',
                delivery_lon='46.7000000',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 201)
        data = response.json()['data']
        self.assertEqual(data['pickup_country_code'], 'AE')
        self.assertEqual(data['delivery_country_code'], 'SA')

    def test_edit_accepts_country_codes(self):
        seed_zones('OM', 'QA')
        shipment = Shipment.objects.create(
            shipper=self.shipper,
            pickup_address='Pickup',
            delivery_address='Delivery',
            cargo_type='General',
            weight='1 ton',
            vehicle_type_required='Flatbed',
            local=True,
            country_to_country=False,
            status=Shipment.Status.DRAFT,
        )
        response = self.client.patch(
            f'/api/shipper/shipments/{shipment.id}/',
            {
                **_LOAD_TYPE,

                'pickup_country_code': 'OM',
                'delivery_country_code': 'QA',
            },
            format='json',
        )
        self.assertEqual(response.status_code, 200)
        data = response.json()['data']
        self.assertEqual(data['pickup_country_code'], 'OM')
        self.assertEqual(data['delivery_country_code'], 'QA')
        shipment.refresh_from_db()
        self.assertEqual(shipment.pickup_country_code, 'OM')
        self.assertEqual(shipment.delivery_country_code, 'QA')

    def test_create_rejects_country_outside_zones(self):
        seed_zones('AE')
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_country_code='AE',
                delivery_country_code='ZA',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        body = response.json()
        self.assertEqual(body['message'], 'Validation failed.')
        self.assertEqual(body['error'], 'We are not operating in South Africa.')
        self.assertIsNone(body['data'])

    def test_create_rejects_pickup_outside_zones(self):
        seed_zones('SA')
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_country_code='ZA',
                delivery_country_code='SA',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        self.assertEqual(
            response.json()['error'],
            'We are not operating in South Africa.',
        )

    @patch('core.geocoding.reverse_geocode_country_code')
    def test_create_rejects_geocoded_country_outside_zones(self, mock_reverse):
        seed_zones('AE')
        mock_reverse.side_effect = lambda lat, lon: 'AE' if float(lon) > 60 else 'ZA'
        response = self.client.post(
            '/api/shipper/shipments/',
            self._shipment_payload(
                pickup_lat='24.8607000',
                pickup_lon='67.0011000',
                delivery_lat='-26.2041000',
                delivery_lon='28.0473000',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        self.assertEqual(
            response.json()['error'],
            'We are not operating in South Africa.',
        )


class ShipmentReturnTripTests(TestCase):
    def setUp(self):
        self.shipper = User.objects.create_user('retcr@test.com', email='retcr@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)
        token, _ = Token.objects.get_or_create(user=self.shipper)
        self.client = APIClient()
        self.client.credentials(HTTP_AUTHORIZATION=f'Bearer {token.key}')

    def _payload(self, **overrides):
        payload = {
            'pickup_address': 'Pickup',
            'delivery_address': 'Delivery',
            'cargo_type': 'General',
            'weight': '1 ton',
            'vehicle_type_required': 'Flatbed',
            'local': True,
            'country_to_country': False,
        }
        payload.update(overrides)
        return payload

    def test_create_return_trip_when_zone_allows_pickup_country(self):
        seed_zones('AE', 'SA', allow_return_trips=True)
        response = self.client.post(
            '/api/shipper/shipments/',
            self._payload(
                is_return=True,
                pickup_country_code='AE',
                delivery_country_code='SA',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 201)
        self.assertTrue(response.json()['data']['is_return'])

    def test_create_return_trip_rejected_when_zone_disallows_pickup_country(self):
        seed_zones('AE', 'SA')
        Zone.objects.filter(country_code='AE').update(allow_return_trips=False)
        response = self.client.post(
            '/api/shipper/shipments/',
            self._payload(
                is_return=True,
                pickup_country_code='AE',
                delivery_country_code='SA',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        self.assertEqual(
            response.json()['error'],
            'Return trips are not allowed from United Arab Emirates.',
        )

    def test_create_non_return_trip_skips_zone_check(self):
        seed_zones('AE', 'SA')
        response = self.client.post(
            '/api/shipper/shipments/',
            self._payload(
                is_return=False,
                pickup_country_code='AE',
                delivery_country_code='SA',
            ),
            format='json',
        )
        self.assertEqual(response.status_code, 201)
        self.assertFalse(response.json()['data']['is_return'])

    def test_edit_return_trip_rejected_when_zone_disallows(self):
        seed_zones('AE', 'SA')
        Zone.objects.filter(country_code='AE').update(allow_return_trips=False)
        shipment = Shipment.objects.create(
            shipper=self.shipper,
            pickup_address='Pickup',
            delivery_address='Delivery',
            cargo_type='General',
            weight='1 ton',
            vehicle_type_required='Flatbed',
            pickup_country_code='AE',
            delivery_country_code='SA',
            local=True,
            country_to_country=False,
            status=Shipment.Status.DRAFT,
        )
        response = self.client.patch(
            f'/api/shipper/shipments/{shipment.id}/',
            {
                **_LOAD_TYPE,
'is_return': True},
            format='json',
        )
        self.assertEqual(response.status_code, 400)
        self.assertEqual(
            response.json()['error'],
            'Return trips are not allowed from United Arab Emirates.',
        )


class ReturnLoadsDiscoveryTests(TestCase):
    def setUp(self):
        self.shipper = User.objects.create_user('retship@test.com', email='retship@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('rettr@test.com', email='rettr@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,
            tc_id='222',
        )
        Vehicle.objects.create(
            owner=self.transporter,
            vehicle_type='Flatbed',
            registration_number='RET-222',
            load_capacity=Decimal('5000'),
            is_verified=True,
            is_active=True,
        )

        self.shipment = Shipment.objects.create(
            shipper=self.shipper,
            pickup_address='A',
            pickup_lat=Decimal('24.8607000'),
            pickup_lon=Decimal('67.0011000'),
            pickup_country_code='AE',
            delivery_address='B',
            delivery_lat=Decimal('24.9000000'),
            delivery_lon=Decimal('67.0500000'),
            delivery_country_code='AE',
            cargo_type='General',
            weight='1 ton',
            vehicle_type_required='Flatbed',
            is_return=True,
            local=True,
            country_to_country=False,
            status=Shipment.Status.PUBLISHED,
        )

        token, _ = Token.objects.get_or_create(user=self.transporter)
        self.client = APIClient()
        self.client.credentials(HTTP_AUTHORIZATION=f'Bearer {token.key}')

    @patch('api.views.get_latest_device_position')
    def test_return_shipment_included_in_load_discovery(self, mock_position):
        Zone.objects.create(country_code='AE', radius_km=Decimal('50'), allow_return_trips=True)
        mock_position.return_value = {
            'latitude': 24.8610,
            'longitude': 67.0020,
            'raw': {'id': 10},
        }

        response = self.client.get('/api/transporter/load-discovery/?country_code=AE')
        self.assertEqual(response.status_code, 200)
        data = response.json()['data']
        load_ids = {item['id'] for item in data['loads']}
        self.assertIn(self.shipment.id, load_ids)
