"""
Tests for the Amazon Cognito token provider

These tests verify:
1. The client-credentials request is shaped the way Cognito expects
2. Tokens are cached, refreshed ahead of expiry, and refetched after
   invalidate(); concurrent callers share a single fetch
3. Cognito/transport failures surface as AmazonAuthError with the right
   retryable flag and never leave a stale token cached
"""

import asyncio
import base64
import sys
import unittest
from pathlib import Path
from unittest.mock import patch

sys.path.insert(0, str(Path(__file__).resolve().parents[2]))

import app.src.amazon_auth as amazon_auth_mod
from app.src.amazon_auth import AmazonAuthError, CognitoTokenProvider

TOKEN_URL = "https://cognito.example.com/oauth2/token"
CLIENT_ID = "client-id"
CLIENT_SECRET = "client-secret"
SCOPE = "shuttle-service-gamma-api/consultation"


class FakeClock:
    def __init__(self, now=1000.0):
        self.now = now

    def __call__(self):
        return self.now

    def advance(self, seconds):
        self.now += seconds


class FakeResponse:
    """Fake aiohttp response usable as an async context manager"""

    def __init__(self, status, json_body=None, text_body="", delay=0.0):
        self.status = status
        self._json = json_body
        self._text = text_body
        self._delay = delay

    async def json(self):
        return self._json

    async def text(self):
        return self._text

    async def __aenter__(self):
        if self._delay:
            await asyncio.sleep(self._delay)
        return self

    async def __aexit__(self, *args):
        return False


class FailingPost:
    def __init__(self, exc):
        self.exc = exc

    async def __aenter__(self):
        raise self.exc

    async def __aexit__(self, *args):
        return False


class FakeSession:
    """Fake aiohttp session returning queued post outcomes and recording calls"""

    def __init__(self, outcomes):
        self.outcomes = list(outcomes)
        self.calls = []

    def post(self, url, **kwargs):
        self.calls.append((url, kwargs))
        outcome = self.outcomes.pop(0)
        if isinstance(outcome, Exception):
            return FailingPost(outcome)
        return outcome

    async def __aenter__(self):
        return self

    async def __aexit__(self, *args):
        return False


def token_response(token="tok-1", expires_in=3600, delay=0.0):
    return FakeResponse(
        200,
        json_body={"access_token": token, "token_type": "Bearer", "expires_in": expires_in},
        delay=delay,
    )


class CognitoTestCase(unittest.TestCase):
    def setUp(self):
        self.clock = FakeClock()
        self.provider = CognitoTokenProvider(
            TOKEN_URL, CLIENT_ID, CLIENT_SECRET, SCOPE, clock=self.clock
        )

    def run_with_session(self, coro_factory, outcomes):
        session = FakeSession(outcomes)
        with patch.object(
            amazon_auth_mod.aiohttp, "ClientSession", return_value=session
        ):
            result = asyncio.run(coro_factory())
        return result, session

    def get_token(self, outcomes):
        return self.run_with_session(self.provider.get_token, outcomes)


class TestTokenRequest(CognitoTestCase):
    def test_sends_client_credentials_grant_with_basic_auth(self):
        token, session = self.get_token([token_response()])

        self.assertEqual(token, "tok-1")
        self.assertEqual(len(session.calls), 1)
        url, kwargs = session.calls[0]
        self.assertEqual(url, TOKEN_URL)
        expected_basic = base64.b64encode(
            f"{CLIENT_ID}:{CLIENT_SECRET}".encode()
        ).decode()
        self.assertEqual(kwargs["headers"]["Authorization"], f"Basic {expected_basic}")
        self.assertEqual(
            kwargs["headers"]["Content-Type"], "application/x-www-form-urlencoded"
        )
        self.assertEqual(
            kwargs["data"], {"grant_type": "client_credentials", "scope": SCOPE}
        )

    def test_rejects_incomplete_configuration(self):
        with self.assertRaises(ValueError):
            CognitoTokenProvider(TOKEN_URL, CLIENT_ID, "", SCOPE)

    def test_from_config_reads_amazon_settings(self):
        class Cfg:
            amazon_cognito_token_url = TOKEN_URL
            amazon_client_id = CLIENT_ID
            amazon_client_secret = CLIENT_SECRET
            amazon_oauth_scope = SCOPE

        provider = CognitoTokenProvider.from_config(Cfg())

        self.assertFalse(provider.has_token)


class TestCaching(CognitoTestCase):
    def test_reuses_cached_token(self):
        async def twice():
            first = await self.provider.get_token()
            second = await self.provider.get_token()
            return first, second

        (first, second), session = self.run_with_session(twice, [token_response()])

        self.assertEqual((first, second), ("tok-1", "tok-1"))
        self.assertEqual(len(session.calls), 1)

    def test_keeps_token_until_refresh_margin(self):
        self.get_token([token_response("tok-1", expires_in=3600)])
        self.clock.advance(2999)

        token, session = self.get_token([])

        self.assertEqual(token, "tok-1")
        self.assertEqual(session.calls, [])

    def test_refreshes_ten_minutes_before_expiry(self):
        self.get_token([token_response("tok-1", expires_in=3600)])
        self.clock.advance(3000)

        token, session = self.get_token([token_response("tok-2")])

        self.assertEqual(token, "tok-2")
        self.assertEqual(len(session.calls), 1)

    def test_short_lived_token_refreshes_at_half_life_not_every_call(self):
        self.get_token([token_response("tok-1", expires_in=300)])

        self.clock.advance(100)
        token, session = self.get_token([])
        self.assertEqual(token, "tok-1")
        self.assertEqual(session.calls, [])

        self.clock.advance(60)
        token, session = self.get_token([token_response("tok-2", expires_in=300)])
        self.assertEqual(token, "tok-2")
        self.assertEqual(len(session.calls), 1)

    def test_missing_expires_in_defaults_to_one_hour(self):
        self.get_token([FakeResponse(200, json_body={"access_token": "tok-1"})])
        self.clock.advance(2999)

        token, session = self.get_token([])

        self.assertEqual(token, "tok-1")
        self.assertEqual(session.calls, [])

    def test_invalidate_forces_refetch(self):
        self.get_token([token_response("tok-1")])
        self.assertTrue(self.provider.has_token)

        self.provider.invalidate()
        self.assertFalse(self.provider.has_token)

        token, session = self.get_token([token_response("tok-2")])
        self.assertEqual(token, "tok-2")
        self.assertEqual(len(session.calls), 1)

    def test_concurrent_callers_share_one_fetch(self):
        async def burst():
            return await asyncio.gather(
                *(self.provider.get_token() for _ in range(5))
            )

        tokens, session = self.run_with_session(
            burst, [token_response("tok-1", delay=0.01)]
        )

        self.assertEqual(tokens, ["tok-1"] * 5)
        self.assertEqual(len(session.calls), 1)


class TestFailures(CognitoTestCase):
    def assert_auth_error(self, outcomes, *, status, retryable):
        with self.assertRaises(AmazonAuthError) as ctx:
            self.get_token(outcomes)
        self.assertEqual(ctx.exception.status, status)
        self.assertEqual(ctx.exception.retryable, retryable)
        self.assertFalse(self.provider.has_token)

    def test_invalid_client_is_not_retryable(self):
        self.assert_auth_error(
            [FakeResponse(401, text_body='{"error":"invalid_client"}')],
            status=401,
            retryable=False,
        )

    def test_bad_request_is_not_retryable(self):
        self.assert_auth_error(
            [FakeResponse(400, text_body="missing scope")], status=400, retryable=False
        )

    def test_server_error_is_retryable(self):
        self.assert_auth_error(
            [FakeResponse(503, text_body="unavailable")], status=503, retryable=True
        )

    def test_throttling_is_retryable(self):
        self.assert_auth_error(
            [FakeResponse(429, text_body="slow down")], status=429, retryable=True
        )

    def test_connection_error_is_retryable(self):
        self.assert_auth_error(
            [amazon_auth_mod.aiohttp.ClientConnectionError("refused")],
            status=None,
            retryable=True,
        )

    def test_timeout_is_retryable(self):
        self.assert_auth_error([asyncio.TimeoutError()], status=None, retryable=True)

    def test_response_without_access_token_is_an_error(self):
        self.assert_auth_error(
            [FakeResponse(200, json_body={"token_type": "Bearer"})],
            status=None,
            retryable=False,
        )

    def test_failed_early_refresh_falls_back_to_still_valid_token(self):
        self.get_token([token_response("tok-1", expires_in=3600)])
        self.clock.advance(3000)

        token, session = self.get_token([FakeResponse(503, text_body="unavailable")])

        self.assertEqual(token, "tok-1")
        self.assertEqual(len(session.calls), 1)

    def test_failed_refresh_after_real_expiry_raises(self):
        self.get_token([token_response("tok-1", expires_in=3600)])
        self.clock.advance(3601)

        with self.assertRaises(AmazonAuthError):
            self.get_token([FakeResponse(503, text_body="unavailable")])

    def test_next_call_after_fallback_retries_the_refresh(self):
        self.get_token([token_response("tok-1", expires_in=3600)])
        self.clock.advance(3000)
        self.get_token([FakeResponse(503, text_body="unavailable")])

        token, session = self.get_token([token_response("tok-2")])

        self.assertEqual(token, "tok-2")
        self.assertEqual(len(session.calls), 1)


if __name__ == "__main__":
    unittest.main()
