from decimal import Decimal
from unittest.mock import patch

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

from accounts.models import PlatformSettings, ShipperProfile, TransporterProfile, UserRole, Vehicle
from core.models import Bid, RateRequest, Shipment, Trip, Zone
from core.rate_distribution import (
    dispatch_rate_requests_for_shipment,
    find_eligible_drivers_for_shipment,
    handle_rate_request_accept,
    handle_rate_request_reject,
)

User = get_user_model()


class RateDistributionTests(TestCase):
    def setUp(self):
        # 1. Admin
        self.admin = User.objects.create_user(
            username='admin@truckme.test',
            email='admin@truckme.test',
            password='adminpassword',
            first_name='Super',
            last_name='Admin',
            is_staff=True,
            is_superuser=True,
        )
        UserRole.objects.create(user=self.admin, role=UserRole.Role.ADMIN)
        self.admin_token = Token.objects.create(user=self.admin)

        # 2. Shipper
        self.shipper = User.objects.create_user(
            username='shipper@truckme.test',
            email='shipper@truckme.test',
            password='shipperpassword',
            first_name='Shipper',
            last_name='One',
        )
        UserRole.objects.create(user=self.shipper, role=UserRole.Role.SHIPPER)
        ShipperProfile.objects.create(user=self.shipper, account_type=ShipperProfile.AccountType.INDIVIDUAL)
        self.shipper_token = Token.objects.create(user=self.shipper)

        # 3. Drivers in AE (Dubai)
        self.driver1 = User.objects.create_user(
            username='driver1@truckme.test',
            email='driver1@truckme.test',
            password='driverpassword',
            first_name='Driver',
            last_name='One',
        )
        UserRole.objects.create(user=self.driver1, role=UserRole.Role.TRANSPORTER)
        TransporterProfile.objects.create(
            user=self.driver1,
            account_type=TransporterProfile.AccountType.DRIVER,
            documents_verified=True,
            local=True,
            country_to_country=False,
            tc_id='tc_driver_1',
        )
        Vehicle.objects.create(
            owner=self.driver1,
            registration_number='DXB-101',
            vehicle_type='Flatbed',
            load_capacity=Decimal('10000.00'),
            is_verified=True,
            is_active=True,
        )
        self.driver1_token = Token.objects.create(user=self.driver1)

        self.driver2 = User.objects.create_user(
            username='driver2@truckme.test',
            email='driver2@truckme.test',
            password='driverpassword',
            first_name='Driver',
            last_name='Two',
        )
        UserRole.objects.create(user=self.driver2, role=UserRole.Role.TRANSPORTER)
        TransporterProfile.objects.create(
            user=self.driver2,
            account_type=TransporterProfile.AccountType.DRIVER,
            documents_verified=True,
            local=True,
            country_to_country=False,
            tc_id='tc_driver_2',
        )
        Vehicle.objects.create(
            owner=self.driver2,
            registration_number='DXB-102',
            vehicle_type='Flatbed',
            load_capacity=Decimal('10000.00'),
            is_verified=True,
            is_active=True,
        )
        self.driver2_token = Token.objects.create(user=self.driver2)

        self.driver3 = User.objects.create_user(
            username='driver3@truckme.test',
            email='driver3@truckme.test',
            password='driverpassword',
            first_name='Driver',
            last_name='Three',
        )
        UserRole.objects.create(user=self.driver3, role=UserRole.Role.TRANSPORTER)
        TransporterProfile.objects.create(
            user=self.driver3,
            account_type=TransporterProfile.AccountType.DRIVER,
            documents_verified=True,
            local=True,
            country_to_country=False,
            tc_id='tc_driver_3',
        )
        Vehicle.objects.create(
            owner=self.driver3,
            registration_number='DXB-103',
            vehicle_type='Flatbed',
            load_capacity=Decimal('10000.00'),
            is_verified=True,
            is_active=True,
        )
        self.driver3_token = Token.objects.create(user=self.driver3)

        # Zone for AE
        Zone.objects.get_or_create(country_code='AE', defaults={'radius_km': Decimal('100.00'), 'rate_per_km': Decimal('5.00')})

        # Base Shipment in Dubai (25.2048, 55.2708)
        self.shipment = Shipment.objects.create(
            shipper=self.shipper,
            pickup_address='Dubai Marina',
            pickup_lat=Decimal('25.2048'),
            pickup_lon=Decimal('55.2708'),
            pickup_country_code='AE',
            delivery_address='Abu Dhabi Mall',
            delivery_lat=Decimal('24.4539'),
            delivery_lon=Decimal('54.3773'),
            delivery_country_code='AE',
            cargo_type='Construction Equipment',
            weight='10 tons',
            vehicle_type_required='Flatbed',
            local=True,
            country_to_country=False,
            suggested_price=Decimal('2500.00'),
            currency='AED',
            status=Shipment.Status.PUBLISHED,
        )

        self.client = Client()
        self.api_client = APIClient()

    def _mock_positions(self):
        return {
            'tc_driver_1': {'latitude': 25.2050, 'longitude': 55.2710},  # ~0.03 km
            'tc_driver_2': {'latitude': 25.2200, 'longitude': 55.2800},  # ~2 km
            'tc_driver_3': {'latitude': 25.2600, 'longitude': 55.3000},  # ~6.5 km
        }

    @patch('core.rate_distribution.get_latest_positions_map')
    @patch('core.rate_distribution.country_code_for_point', return_value='AE')
    def test_single_driver_mode_sequential_dispatch_and_rejection_cascade(self, mock_country, mock_positions):
        mock_positions.return_value = self._mock_positions()

        # Set PlatformSettings count = 1 (Single Driver Sequential Mode)
        ps = PlatformSettings.load()
        ps.bidding_rate_distribution_count = 1
        ps.save()

        # 1. First dispatch -> should create only 1 rate request for driver 1 (nearest)
        reqs = dispatch_rate_requests_for_shipment(self.shipment)
        self.assertEqual(len(reqs), 1)
        self.assertEqual(reqs[0].driver, self.driver1)
        self.assertEqual(reqs[0].batch_number, 1)
        self.assertEqual(reqs[0].status, RateRequest.Status.PENDING)

        # 2. Driver 1 rejects -> should automatically dispatch Batch #2 to driver 2
        handle_rate_request_reject(reqs[0], reason='Busy with other work')
        reqs[0].refresh_from_db()
        self.assertEqual(reqs[0].status, RateRequest.Status.REJECTED)

        batch2_reqs = RateRequest.objects.filter(shipment=self.shipment, batch_number=2)
        self.assertEqual(batch2_reqs.count(), 1)
        rr2 = batch2_reqs.first()
        self.assertEqual(rr2.driver, self.driver2)
        self.assertEqual(rr2.status, RateRequest.Status.PENDING)

        # 3. Driver 2 accepts -> should assign shipment, create trip, mark ACCEPTED
        trip, accepted_rr = handle_rate_request_accept(rr2)
        self.assertEqual(accepted_rr.status, RateRequest.Status.ACCEPTED)
        self.shipment.refresh_from_db()
        self.assertEqual(self.shipment.status, Shipment.Status.ASSIGNED)
        self.assertEqual(trip.status, Trip.Status.ASSIGNED)
        self.assertEqual(trip.transporter, self.driver2)
        self.assertEqual(trip.accepted_bid.amount, Decimal('2500.00'))

    @patch('core.rate_distribution.get_latest_positions_map')
    @patch('core.rate_distribution.country_code_for_point', return_value='AE')
    def test_multi_driver_mode_broadcast_and_acceptance_cancellation(self, mock_country, mock_positions):
        mock_positions.return_value = self._mock_positions()

        # Set PlatformSettings count = 3 (Broadcast Multi-Driver Mode)
        ps = PlatformSettings.load()
        ps.bidding_rate_distribution_count = 3
        ps.save()

        reqs = dispatch_rate_requests_for_shipment(self.shipment)
        self.assertEqual(len(reqs), 3)
        self.assertEqual([r.driver_id for r in reqs], [self.driver1.id, self.driver2.id, self.driver3.id])
        for r in reqs:
            self.assertEqual(r.batch_number, 1)
            self.assertEqual(r.status, RateRequest.Status.PENDING)

        # Driver 2 accepts first -> driver 1 & 3 should be CANCELLED
        trip, accepted_rr = handle_rate_request_accept(reqs[1])
        self.assertEqual(accepted_rr.status, RateRequest.Status.ACCEPTED)

        reqs[0].refresh_from_db()
        reqs[2].refresh_from_db()
        self.assertEqual(reqs[0].status, RateRequest.Status.CANCELLED)
        self.assertEqual(reqs[2].status, RateRequest.Status.CANCELLED)

    @patch('core.rate_distribution.get_latest_positions_map')
    @patch('core.rate_distribution.country_code_for_point', return_value='AE')
    def test_driver_rate_request_apis(self, mock_country, mock_positions):
        mock_positions.return_value = self._mock_positions()

        ps = PlatformSettings.load()
        ps.bidding_rate_distribution_count = 1
        ps.save()

        reqs = dispatch_rate_requests_for_shipment(self.shipment)
        rr = reqs[0]

        # 1. Driver 1 views rate requests
        self.api_client.credentials(HTTP_AUTHORIZATION=f'Token {self.driver1_token.key}')
        res = self.api_client.get('/api/transporter/rate-requests/')
        self.assertEqual(res.status_code, 200)
        self.assertEqual(res.data['data']['count'], 1)
        self.assertEqual(res.data['data']['results'][0]['id'], rr.id)

        # 2. Driver 1 accepts
        res_accept = self.api_client.post(f'/api/transporter/rate-requests/{rr.id}/accept/')
        self.assertEqual(res_accept.status_code, 200)
        self.assertIn('trip_id', res_accept.data['data'])

        self.shipment.refresh_from_db()
        self.assertEqual(self.shipment.status, Shipment.Status.ASSIGNED)

    @patch('core.rate_distribution.get_latest_positions_map')
    @patch('core.rate_distribution.country_code_for_point', return_value='AE')
    def test_admin_rate_requests_audit_and_manual_dispatch_api(self, mock_country, mock_positions):
        mock_positions.return_value = self._mock_positions()

        ps = PlatformSettings.load()
        ps.bidding_rate_distribution_count = 2
        ps.save()

        reqs = dispatch_rate_requests_for_shipment(self.shipment)

        self.api_client.credentials(HTTP_AUTHORIZATION=f'Token {self.admin_token.key}')

        # 1. Admin checks rate requests for shipment
        res = self.api_client.get(f'/api/admin/shipments/{self.shipment.id}/rate-requests/')
        self.assertEqual(res.status_code, 200)
        self.assertEqual(res.data['data']['count'], 2)

        # 2. Admin manually dispatches next batch
        res_dispatch = self.api_client.post(f'/api/admin/shipments/{self.shipment.id}/dispatch-next-batch/')
        self.assertEqual(res_dispatch.status_code, 200)
        self.assertEqual(res_dispatch.data['data']['dispatched_count'], 1)  # only driver 3 was remaining

    def test_platform_settings_form_and_view(self):
        self.client.force_login(self.admin)
        res = self.client.get('/settings')
        self.assertEqual(res.status_code, 200)
        self.assertContains(res, 'bidding_rate_distribution_count')

        # POST update count = 5
        post_data = {
            'default_wallet_currency': 'USD',
            'load_visibility_radius_km': '50.00',
            'load_visibility_max_radius_km': '500.00',
            'document_reminder_days': 30,
            'analytics_transporter_bid_window_days': 90,
            'on_time_delivery_grace_hours': 24,
            'nearby_load_notify_radius_km': '50.00',
            'nearby_publish_notify_max': 50,
            'navigation_naive_speed_kph': '45.00',
            'pod_max_delivery_distance_km': '0.00',
            'admin_stale_trip_hours': 48,
            'bidding_rate_distribution_count': 5,
        }
        res_post = self.client.post('/settings', data=post_data)
        self.assertEqual(res_post.status_code, 302)

        ps = PlatformSettings.load()
        self.assertEqual(ps.bidding_rate_distribution_count, 5)
