import concurrent.futures
import tempfile
import unittest
from pathlib import Path

import jwt
from cryptography.hazmat.primitives.asymmetric import ec
from verification import NonceStore, verify_with_key


class VerificationTest(unittest.TestCase):
    def setUp(self):
        self.directory = tempfile.TemporaryDirectory()
        self.addCleanup(self.directory.cleanup)
        self.store = NonceStore(str(Path(self.directory.name) / "nonces.sqlite"))
        self.key = ec.generate_private_key(ec.SECP256R1())
        self.now = 1_800_000_000_000
        self.session = "server-generated-cookie-with-32-random-bytes"
        self.claims = dict(iss="https://ims.koee.app", aud="com.example.shop", sub="a" * 43,
                           iat=self.now // 1000, exp=self.now // 1000 + 300, jti="unique-id",
                           origin="https://shop.example.com", scope="profile", name="Person", locale="en",
                           nonce=self.store.issue(self.session, self.now))

    def sign(self, changes=None, header=None):
        return jwt.encode({**self.claims, **(changes or {})}, self.key, algorithm="ES256",
                          headers={"typ": "kit-launch+jwt", "kid": "configured-key", **(header or {})})

    def verify(self, token, offset=0, session=None):
        return verify_with_key(token, self.key.public_key(), "https://ims.koee.app", "com.example.shop",
                               {"https://shop.example.com"}, self.store, session or self.session, self.now + offset)

    def test_binding_and_replay(self):
        token = self.sign()
        with self.assertRaises(ValueError):
            self.verify(token, session="another-server-generated-pre-session-cookie")
        self.verify(token)
        with self.assertRaises(ValueError):
            self.verify(token)

    def test_atomic_concurrent_consumption(self):
        nonce = self.claims["nonce"]
        def consume(_):
            return self.store.consume(nonce, self.session, self.now)
        with concurrent.futures.ThreadPoolExecutor(4) as executor:
            self.assertEqual(sum(executor.map(consume, range(8))), 1)

    def test_rejected_claims(self):
        cases = [({}, 300000), ({}, 300001), ({"exp": self.now // 1000 + 301}, 0),
                 ({"iat": self.now // 1000 + 61}, 0), ({"iat": self.now // 1000 + .5}, 0),
                 ({"iat": True}, 0), ({"exp": self.now // 1000 + 299.5}, 0),
                 ({"iss": "https://attacker.example"}, 0), ({"aud": "com.other.kit"}, 0),
                 ({"aud": ["com.example.shop"]}, 0), ({"origin": "https://attacker.example"}, 0),
                 ({"nonce": "b" * 43}, 0), ({"scope": "handle"}, 0), ({"nbf": self.now // 1000 + 1}, 0)]
        for claims, offset in cases:
            with self.subTest(claims=claims, offset=offset), self.assertRaises((ValueError, jwt.PyJWTError)):
                self.verify(self.sign(claims), offset)

    def test_headers(self):
        for header in ({"typ": "JWT"}, {"kid": ""}, {"jku": "https://attacker.example/jwks"}):
            with self.subTest(header=header), self.assertRaises(ValueError):
                self.verify(self.sign(header=header))

    def test_iat_tolerance_only(self):
        self.verify(self.sign({"iat": self.now // 1000 + 60}))

    def test_invalid_signature_does_not_consume_nonce(self):
        token = jwt.encode(self.claims, ec.generate_private_key(ec.SECP256R1()), algorithm="ES256",
                           headers={"typ": "kit-launch+jwt", "kid": "configured-key"})
        with self.assertRaises(jwt.PyJWTError):
            self.verify(token)
        self.verify(self.sign())


if __name__ == "__main__":
    unittest.main()
