#!/usr/bin/env python3
"""Phase 3 (hardening): Encryption Key Rotation — test suite.

Tests:
  - Dual key resolution (V1/V2 env vars)
  - Encrypt always uses current key
  - Decrypt tries current → legacy fallback
  - Migration re-encrypts all data under V2
  - No legacy key → no fallback
  - Cache clearing after rotation

Run: python -m pytest tests/test_crypto_rotation.py -v
"""

import os
import sys
import unittest

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from app.crypto import (
    encrypt_value, decrypt_value, encrypt_user_value, decrypt_user_value,
    encrypt_entity, decrypt_entity,
    master_key_is_rotation_active, clear_cache, reset_master_key_cache,
    _resolve_master_keys, key_is_configured
)


def _set_env(key, value):
    """Set env var, or delete if value is None."""
    if value is None:
        os.environ.pop(key, None)
    else:
        os.environ[key] = value


class TestKeyResolution(unittest.TestCase):
    """Test dual key resolution logic."""

    def setUp(self):
        self._orig_v1 = os.environ.get("ENCRYPTION_KEY")
        self._orig_v2 = os.environ.get("ENCRYPTION_KEY_V2")
        reset_master_key_cache()

    def tearDown(self):
        _set_env("ENCRYPTION_KEY", self._orig_v1)
        _set_env("ENCRYPTION_KEY_V2", self._orig_v2)
        reset_master_key_cache()

    def test_single_key_no_rotation(self):
        """Only ENCRYPTION_KEY set → current=V1, legacy=None."""
        _set_env("ENCRYPTION_KEY", "test-key-v1-only")
        _set_env("ENCRYPTION_KEY_V2", None)
        reset_master_key_cache()

        current, legacy = _resolve_master_keys()
        self.assertIsNotNone(current)
        self.assertIsNone(legacy)
        self.assertFalse(master_key_is_rotation_active())

    def test_dual_key_rotation_mode(self):
        """Both V1 and V2 set → current=V2, legacy=V1."""
        _set_env("ENCRYPTION_KEY", "test-key-v1-legacy")
        _set_env("ENCRYPTION_KEY_V2", "test-key-v2-current")
        reset_master_key_cache()

        current, legacy = _resolve_master_keys()
        self.assertIsNotNone(current)
        self.assertIsNotNone(legacy)
        self.assertNotEqual(current, legacy)
        self.assertTrue(master_key_is_rotation_active())

    def test_v2_only(self):
        """Only V2 set → current=V2, legacy=None."""
        _set_env("ENCRYPTION_KEY", None)
        _set_env("ENCRYPTION_KEY_V2", "test-key-v2-only")
        reset_master_key_cache()

        current, legacy = _resolve_master_keys()
        self.assertIsNotNone(current)
        self.assertIsNone(legacy)
        self.assertFalse(master_key_is_rotation_active())

    def test_same_key_no_rotation(self):
        """V1 == V2 → current=V1, legacy=None (no rotation needed)."""
        _set_env("ENCRYPTION_KEY", "same-key")
        _set_env("ENCRYPTION_KEY_V2", "same-key")
        reset_master_key_cache()

        current, legacy = _resolve_master_keys()
        self.assertIsNotNone(current)
        self.assertIsNone(legacy)


class TestSiteScopedRotation(unittest.TestCase):
    """Test site-scoped encrypt/decrypt with key rotation."""

    def setUp(self):
        self._orig_v1 = os.environ.get("ENCRYPTION_KEY")
        self._orig_v2 = os.environ.get("ENCRYPTION_KEY_V2")
        _set_env("ENCRYPTION_KEY", "rotation-test-v1")
        _set_env("ENCRYPTION_KEY_V2", "rotation-test-v2")
        reset_master_key_cache()

    def tearDown(self):
        _set_env("ENCRYPTION_KEY", self._orig_v1)
        _set_env("ENCRYPTION_KEY_V2", self._orig_v2)
        reset_master_key_cache()

    def test_encrypt_decrypt_current_key(self):
        """Encrypt with V2 → decrypt with V2 (happy path)."""
        encrypted = encrypt_value(1, "secret data")
        self.assertIsNotNone(encrypted)
        decrypted = decrypt_value(1, encrypted)
        self.assertEqual(decrypted, "secret data")

    def test_decrypt_fallback_to_legacy(self):
        """Data encrypted with V1 → decrypt falls back to V1."""
        # Temporarily use V1-only mode to encrypt
        _set_env("ENCRYPTION_KEY", "rotation-test-v1")
        _set_env("ENCRYPTION_KEY_V2", None)
        reset_master_key_cache()

        v1_encrypted = encrypt_value(1, "legacy secret")

        # Restore V2 mode
        _set_env("ENCRYPTION_KEY", "rotation-test-v1")
        _set_env("ENCRYPTION_KEY_V2", "rotation-test-v2")
        reset_master_key_cache()

        # Decrypt should fall back to V1
        decrypted = decrypt_value(1, v1_encrypted)
        self.assertEqual(decrypted, "legacy secret")

    def test_different_keys_produce_different_ciphertext(self):
        """Same plaintext with different keys → different ciphertext."""
        # Encrypt with V2 (current)
        v2_encrypted = encrypt_value(1, "test data")

        # Encrypt with V1 only
        _set_env("ENCRYPTION_KEY", "rotation-test-v1")
        _set_env("ENCRYPTION_KEY_V2", None)
        reset_master_key_cache()
        v1_encrypted = encrypt_value(1, "test data")

        # Ciphertext should differ
        self.assertNotEqual(v2_encrypted, v1_encrypted)


class TestUserScopedRotation(unittest.TestCase):
    """Test user-scoped encrypt/decrypt with key rotation."""

    def setUp(self):
        self._orig_v1 = os.environ.get("ENCRYPTION_KEY")
        self._orig_v2 = os.environ.get("ENCRYPTION_KEY_V2")
        _set_env("ENCRYPTION_KEY", "user-rotation-v1")
        _set_env("ENCRYPTION_KEY_V2", "user-rotation-v2")
        reset_master_key_cache()

    def tearDown(self):
        _set_env("ENCRYPTION_KEY", self._orig_v1)
        _set_env("ENCRYPTION_KEY_V2", self._orig_v2)
        reset_master_key_cache()

    def test_user_encrypt_decrypt(self):
        """User-scoped encrypt/decrypt works with V2."""
        encrypted = encrypt_user_value(42, "user email")
        self.assertIsNotNone(encrypted)
        decrypted = decrypt_user_value(42, encrypted)
        self.assertEqual(decrypted, "user email")

    def test_user_fallback_to_legacy(self):
        """User data encrypted with V1 → decrypt falls back."""
        # Encrypt with V1 only
        _set_env("ENCRYPTION_KEY", "user-rotation-v1")
        _set_env("ENCRYPTION_KEY_V2", None)
        reset_master_key_cache()
        v1_encrypted = encrypt_user_value(42, "legacy user email")

        # Restore V2
        _set_env("ENCRYPTION_KEY", "user-rotation-v1")
        _set_env("ENCRYPTION_KEY_V2", "user-rotation-v2")
        reset_master_key_cache()

        decrypted = decrypt_user_value(42, v1_encrypted)
        self.assertEqual(decrypted, "legacy user email")


class TestEntityScopedRotation(unittest.TestCase):
    """Test entity-scoped encrypt/decrypt with key rotation."""

    def setUp(self):
        self._orig_v1 = os.environ.get("ENCRYPTION_KEY")
        self._orig_v2 = os.environ.get("ENCRYPTION_KEY_V2")
        _set_env("ENCRYPTION_KEY", "entity-rotation-v1")
        _set_env("ENCRYPTION_KEY_V2", "entity-rotation-v2")
        reset_master_key_cache()

    def tearDown(self):
        _set_env("ENCRYPTION_KEY", self._orig_v1)
        _set_env("ENCRYPTION_KEY_V2", self._orig_v2)
        reset_master_key_cache()

    def test_entity_encrypt_decrypt(self):
        """Entity-scoped encrypt/decrypt works with V2."""
        encrypted = encrypt_entity("webhook_dest", 7, "https://example.com/hook")
        self.assertIsNotNone(encrypted)
        decrypted = decrypt_entity("webhook_dest", 7, encrypted)
        self.assertEqual(decrypted, "https://example.com/hook")

    def test_entity_fallback_to_legacy(self):
        """Entity data encrypted with V1 → decrypt falls back."""
        _set_env("ENCRYPTION_KEY", "entity-rotation-v1")
        _set_env("ENCRYPTION_KEY_V2", None)
        reset_master_key_cache()
        v1_encrypted = encrypt_entity("webhook_dest", 7, "legacy hook url")

        _set_env("ENCRYPTION_KEY", "entity-rotation-v1")
        _set_env("ENCRYPTION_KEY_V2", "entity-rotation-v2")
        reset_master_key_cache()

        decrypted = decrypt_entity("webhook_dest", 7, v1_encrypted)
        self.assertEqual(decrypted, "legacy hook url")


class TestNoFallback(unittest.TestCase):
    """Test behavior when no legacy key is configured."""

    def setUp(self):
        self._orig_v1 = os.environ.get("ENCRYPTION_KEY")
        self._orig_v2 = os.environ.get("ENCRYPTION_KEY_V2")

    def tearDown(self):
        _set_env("ENCRYPTION_KEY", self._orig_v1)
        _set_env("ENCRYPTION_KEY_V2", self._orig_v2)
        reset_master_key_cache()

    def test_no_legacy_no_fallback(self):
        """Data from unknown key → decrypt returns None (no fallback)."""
        _set_env("ENCRYPTION_KEY", "new-key-only")
        _set_env("ENCRYPTION_KEY_V2", None)
        reset_master_key_cache()

        # Fake encrypted data (not actually encrypted with this key)
        fake_encrypted = "dGVzdC1kYXRh"  # base64 of "test-data" (too short for valid ciphertext)
        result = decrypt_value(1, fake_encrypted)
        self.assertIsNone(result)


class TestCacheClearing(unittest.TestCase):
    """Test cache clearing after key changes."""

    def setUp(self):
        self._orig_v1 = os.environ.get("ENCRYPTION_KEY")
        self._orig_v2 = os.environ.get("ENCRYPTION_KEY_V2")
        _set_env("ENCRYPTION_KEY", "cache-test-v1")
        _set_env("ENCRYPTION_KEY_V2", "cache-test-v2")
        reset_master_key_cache()

    def tearDown(self):
        _set_env("ENCRYPTION_KEY", self._orig_v1)
        _set_env("ENCRYPTION_KEY_V2", self._orig_v2)
        reset_master_key_cache()

    def test_clear_cache_resolves_new_keys(self):
        """clear_cache() re-resolves master keys."""
        enc1 = encrypt_value(1, "before change")
        self.assertIsNotNone(enc1)

        # Change keys
        _set_env("ENCRYPTION_KEY", "new-v1-after-clear")
        _set_env("ENCRYPTION_KEY_V2", "new-v2-after-clear")
        clear_cache()

        enc2 = encrypt_value(1, "after change")
        self.assertIsNotNone(enc2)
        self.assertNotEqual(enc1, enc2)

        dec2 = decrypt_value(1, enc2)
        self.assertEqual(dec2, "after change")


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