"""Production secret-loading and signer-boundary regression tests."""

from __future__ import annotations

import base64
import os
from pathlib import Path
import tempfile
import unittest

from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey

from licensing_server.app.config import (
    DatabaseWorkerSettings,
    NotificationWorkerSettings,
    ServerSettings,
    SubscriptionWorkerSettings,
)
from licensing_server.app.security import Ed25519SnapshotSigner


REPOSITORY_ROOT = Path(__file__).resolve().parents[2]
COMPOSE_PATH = REPOSITORY_ROOT / "licensing_server/compose.production.example.yml"
UNRELATED_WORKER_SETTINGS = (
    "LICENSING_SIGNING_KEY_ID",
    "LICENSING_SIGNING_PRIVATE_KEY_PATH",
    "LICENSING_SERIAL_PEPPER_BASE64_FILE",
    "LICENSING_FINGERPRINT_PEPPER_BASE64_FILE",
    "LICENSING_OIDC_ISSUER",
    "LICENSING_OIDC_AUDIENCE",
    "LICENSING_OIDC_JWKS_URL",
)


class LicensingSecurityConfigTests(unittest.TestCase):
    def test_compose_workers_receive_only_component_owned_secrets(self) -> None:
        document = COMPOSE_PATH.read_text(encoding="utf-8")
        subscription = document.split("\n  subscription_worker:\n", 1)[1].split(
            "\n  notification_worker:\n",
            1,
        )[0]
        notification = document.split("\n  notification_worker:\n", 1)[1].split(
            "\nvolumes:\n",
            1,
        )[0]
        for setting in UNRELATED_WORKER_SETTINGS:
            self.assertNotIn(setting, subscription)
            self.assertNotIn(setting, notification)
        self.assertIn("LICENSING_STRIPE_SECRET_KEY_FILE", subscription)
        self.assertNotIn("LICENSING_SMTP_PASSWORD_FILE", subscription)
        self.assertIn("LICENSING_SMTP_PASSWORD_FILE", notification)
        self.assertIn("LICENSING_NOTIFICATION_PEPPER_BASE64_FILE", notification)
        self.assertNotIn("LICENSING_NOTIFICATION_PEPPER_BASE64_FILE", subscription)
        self.assertNotIn("LICENSING_STRIPE_SECRET_KEY_FILE", notification)
        self.assertNotIn("- serial_pepper_base64", subscription)
        self.assertNotIn("- fingerprint_pepper_base64", subscription)
        self.assertNotIn("- serial_pepper_base64", notification)
        self.assertNotIn("- fingerprint_pepper_base64", notification)
        self.assertIn(
            '"--component", "subscription_worker"',
            subscription,
        )
        self.assertIn(
            '"--component", "notification_worker"',
            notification,
        )
        self.assertNotIn("127.0.0.1:8080/v1/health", subscription)
        self.assertNotIn("127.0.0.1:8080/v1/health", notification)

    def test_worker_settings_ignore_and_do_not_expose_unrelated_secrets(self) -> None:
        poison_path = "/definitely/missing/unrelated-secret"
        common = {
            "LICENSING_DATABASE_URL": "sqlite+pysqlite:///:memory:",
            "LICENSING_ENVIRONMENT": "development",
            "LICENSING_SIGNING_KEY_ID": "license.test.must_not_load",
            "LICENSING_SIGNING_PRIVATE_KEY_PATH": poison_path,
            "LICENSING_SERIAL_PEPPER_BASE64_FILE": poison_path,
            "LICENSING_FINGERPRINT_PEPPER_BASE64_FILE": poison_path,
            "LICENSING_OIDC_ISSUER": "not-used",
            "LICENSING_OIDC_AUDIENCE": "not-used",
            "LICENSING_OIDC_JWKS_URL": "not-used",
            "LICENSING_STRIPE_WEBHOOK_SECRET_FILE": poison_path,
        }
        database = DatabaseWorkerSettings.from_environment(common)
        self.assertEqual(database.database_url, "sqlite+pysqlite:///:memory:")
        self.assertFalse(hasattr(database, "serial_pepper"))
        self.assertFalse(hasattr(database, "oidc_issuer"))

        subscription = SubscriptionWorkerSettings.from_environment(
            {
                **common,
                "LICENSING_STRIPE_SECRET_KEY": "sk_test_worker_owned",
                "LICENSING_STRIPE_PRICE_SKU_MAP_JSON": (
                    '{"price_test_ai":"subscription.ai_assist"}'
                ),
            }
        )
        self.assertEqual(subscription.stripe_secret_key, "sk_test_worker_owned")
        self.assertEqual(
            dict(subscription.stripe_price_sku_map),
            {"price_test_ai": "subscription.ai_assist"},
        )
        self.assertFalse(hasattr(subscription, "signing_key_id"))
        self.assertFalse(hasattr(subscription, "fingerprint_pepper"))

        notification = NotificationWorkerSettings.from_environment(
            {
                **common,
                "LICENSING_STRIPE_SECRET_KEY_FILE": poison_path,
                "LICENSING_EMAIL_NOTIFICATIONS_ENABLED": "true",
                "LICENSING_SMTP_HOST": "smtp.example.test",
                "LICENSING_SMTP_PORT": "465",
                "LICENSING_SMTP_TLS_MODE": "implicit",
                "LICENSING_SMTP_USERNAME": "smtp-user",
                "LICENSING_SMTP_PASSWORD": "smtp-password",
                "LICENSING_NOTIFICATION_PEPPER_BASE64": base64.b64encode(
                    b"n" * 32
                ).decode("ascii"),
                "LICENSING_EMAIL_FROM_ADDRESS": "licensing@example.test",
                "LICENSING_ACCOUNT_PORTAL_URL": (
                    "https://account.example.test/licenses"
                ),
            }
        )
        self.assertTrue(notification.email_notifications_enabled)
        self.assertEqual(notification.smtp_password, "smtp-password")
        self.assertEqual(notification.notification_pepper, b"n" * 32)
        self.assertFalse(hasattr(notification, "stripe_secret_key"))
        self.assertFalse(hasattr(notification, "serial_pepper"))
        missing_notification_pepper = {
            **common,
            "LICENSING_EMAIL_NOTIFICATIONS_ENABLED": "true",
            "LICENSING_SMTP_HOST": "smtp.example.test",
            "LICENSING_SMTP_USERNAME": "smtp-user",
            "LICENSING_SMTP_PASSWORD": "smtp-password",
            "LICENSING_EMAIL_FROM_ADDRESS": "licensing@example.test",
            "LICENSING_ACCOUNT_PORTAL_URL": (
                "https://account.example.test/licenses"
            ),
        }
        with self.assertRaisesRegex(ValueError, "notification pepper"):
            NotificationWorkerSettings.from_environment(
                missing_notification_pepper
            )

        with self.assertRaisesRegex(ValueError, "requires PostgreSQL"):
            DatabaseWorkerSettings(
                database_url="sqlite:///production.db",
                environment_name="production",
            )

    def test_environment_secret_files_load_without_direct_secret_values(self) -> None:
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            database = root / "database-url"
            serial_pepper = root / "serial-pepper"
            fingerprint_pepper = root / "fingerprint-pepper"
            webhook = root / "webhook"
            stripe_api_key = root / "stripe-api-key"
            smtp_password = root / "smtp-password"
            notification_pepper = root / "notification-pepper"
            database.write_text("postgresql+psycopg://test:test@db/test\n", encoding="utf-8")
            serial_pepper.write_text(base64.b64encode(b"s" * 32).decode("ascii"), encoding="utf-8")
            fingerprint_pepper.write_text(
                base64.b64encode(b"f" * 32).decode("ascii"),
                encoding="utf-8",
            )
            webhook.write_text("whsec_file_test\n", encoding="utf-8")
            stripe_api_key.write_text("sk_test_file_owned\n", encoding="utf-8")
            smtp_password.write_text("smtp-password-file-owned\n", encoding="utf-8")
            notification_pepper.write_text(
                base64.b64encode(b"n" * 32).decode("ascii"),
                encoding="utf-8",
            )
            environment = {
                "LICENSING_DATABASE_URL_FILE": str(database),
                "LICENSING_ENVIRONMENT": "production",
                "LICENSING_SIGNING_KEY_ID": "license.test.config",
                "LICENSING_SIGNING_PRIVATE_KEY_PATH": str(root / "signing-key"),
                "LICENSING_SERIAL_PEPPER_BASE64_FILE": str(serial_pepper),
                "LICENSING_FINGERPRINT_PEPPER_BASE64_FILE": str(fingerprint_pepper),
                "LICENSING_OIDC_ISSUER": "https://identity.test",
                "LICENSING_OIDC_AUDIENCE": "licensing-api",
                "LICENSING_OIDC_JWKS_URL": "https://identity.test/jwks.json",
                "LICENSING_STRIPE_WEBHOOK_SECRET_FILE": str(webhook),
                "LICENSING_STRIPE_SECRET_KEY_FILE": str(stripe_api_key),
                "LICENSING_STRIPE_COMMERCE_ENABLED": "true",
                "LICENSING_STRIPE_PRICE_SKU_MAP_JSON": (
                    '{"price_test_ai_assist":"subscription.ai_assist"}'
                ),
                "LICENSING_STRIPE_CHECKOUT_SUCCESS_URL": (
                    "https://account.example.test/checkout/success"
                ),
                "LICENSING_STRIPE_CHECKOUT_CANCEL_URL": (
                    "https://account.example.test/checkout/cancel"
                ),
                "LICENSING_STRIPE_BILLING_PORTAL_RETURN_URL": (
                    "https://account.example.test/licenses"
                ),
                "LICENSING_EMAIL_NOTIFICATIONS_ENABLED": "true",
                "LICENSING_SMTP_HOST": "smtp.example.test",
                "LICENSING_SMTP_PORT": "465",
                "LICENSING_SMTP_TLS_MODE": "implicit",
                "LICENSING_SMTP_USERNAME": "smtp-user",
                "LICENSING_SMTP_PASSWORD_FILE": str(smtp_password),
                "LICENSING_NOTIFICATION_PEPPER_BASE64_FILE": str(
                    notification_pepper
                ),
                "LICENSING_EMAIL_FROM_ADDRESS": "licensing@example.test",
                "LICENSING_ACCOUNT_PORTAL_URL": "https://account.example.test/licenses",
                "LICENSING_ALLOW_PRODUCTION_FILE_SIGNER": "true",
                "LICENSING_TRUSTED_PROXY_IPS": "127.0.0.1,10.20.0.0/16",
            }
            settings = ServerSettings.from_environment(environment)
            self.assertEqual(settings.database_url, "postgresql+psycopg://test:test@db/test")
            self.assertEqual(settings.serial_pepper, b"s" * 32)
            self.assertEqual(settings.fingerprint_pepper, b"f" * 32)
            self.assertEqual(settings.stripe_webhook_secret, "whsec_file_test")
            self.assertEqual(settings.stripe_secret_key, "sk_test_file_owned")
            self.assertTrue(settings.stripe_commerce_enabled)
            self.assertTrue(settings.email_notifications_enabled)
            self.assertEqual(settings.smtp_password, "smtp-password-file-owned")
            self.assertEqual(settings.email_from_address, "licensing@example.test")
            self.assertTrue(settings.allow_production_file_signer)
            self.assertEqual(
                settings.trusted_proxy_ips,
                ("127.0.0.1/32", "10.20.0.0/16"),
            )
            notification_settings = NotificationWorkerSettings.from_environment(
                environment
            )
            self.assertEqual(notification_settings.notification_pepper, b"n" * 32)

            environment["LICENSING_DATABASE_URL"] = "postgresql://duplicate"
            with self.assertRaises(ValueError):
                ServerSettings.from_environment(environment)

            invalid_database = dict(environment)
            invalid_database.pop("LICENSING_DATABASE_URL")
            invalid_database.pop("LICENSING_DATABASE_URL_FILE")
            invalid_database["LICENSING_DATABASE_URL"] = "sqlite:///production.db"
            with self.assertRaisesRegex(ValueError, "requires PostgreSQL"):
                ServerSettings.from_environment(invalid_database)

            invalid_identity = dict(environment)
            invalid_identity.pop("LICENSING_DATABASE_URL")
            invalid_identity["LICENSING_OIDC_JWKS_URL"] = "http://identity.test/jwks.json"
            with self.assertRaisesRegex(ValueError, "HTTPS URL"):
                ServerSettings.from_environment(invalid_identity)

            incomplete_commerce = dict(environment)
            incomplete_commerce.pop("LICENSING_DATABASE_URL", None)
            incomplete_commerce.pop("LICENSING_STRIPE_CHECKOUT_CANCEL_URL")
            with self.assertRaisesRegex(ValueError, "enabled Stripe commerce requires"):
                ServerSettings.from_environment(incomplete_commerce)

            duplicate_sku_prices = dict(environment)
            duplicate_sku_prices.pop("LICENSING_DATABASE_URL", None)
            duplicate_sku_prices["LICENSING_STRIPE_PRICE_SKU_MAP_JSON"] = (
                '{"price_monthly":"subscription.ai_assist",'
                '"price_yearly":"subscription.ai_assist"}'
            )
            with self.assertRaisesRegex(ValueError, "exactly one Stripe price"):
                ServerSettings.from_environment(duplicate_sku_prices)

            incomplete_notifications = dict(environment)
            incomplete_notifications.pop("LICENSING_DATABASE_URL", None)
            incomplete_notifications.pop("LICENSING_EMAIL_FROM_ADDRESS")
            with self.assertRaisesRegex(ValueError, "enabled email notifications require"):
                ServerSettings.from_environment(incomplete_notifications)

            plaintext_notifications = dict(environment)
            plaintext_notifications.pop("LICENSING_DATABASE_URL", None)
            plaintext_notifications["LICENSING_SMTP_TLS_MODE"] = "none"
            with self.assertRaisesRegex(ValueError, "implicit or starttls"):
                ServerSettings.from_environment(plaintext_notifications)

            invalid_proxy = dict(environment)
            invalid_proxy.pop("LICENSING_DATABASE_URL", None)
            invalid_proxy["LICENSING_TRUSTED_PROXY_IPS"] = "not-an-ip"
            with self.assertRaisesRegex(ValueError, "invalid IP or network"):
                ServerSettings.from_environment(invalid_proxy)

    @unittest.skipIf(os.name == "nt", "POSIX owner-only mode is not available")
    def test_production_signer_requires_owner_only_key_file(self) -> None:
        with tempfile.TemporaryDirectory() as directory:
            path = Path(directory) / "signing-key"
            private_key = Ed25519PrivateKey.generate()
            raw_key = private_key.private_bytes(
                encoding=serialization.Encoding.Raw,
                format=serialization.PrivateFormat.Raw,
                encryption_algorithm=serialization.NoEncryption(),
            )
            path.write_bytes(base64.b64encode(raw_key))
            path.chmod(0o644)
            with self.assertRaises(PermissionError):
                Ed25519SnapshotSigner.from_private_key_file(
                    "license.test.secure",
                    path,
                    require_owner_only=True,
                )

            path.chmod(0o600)
            signer = Ed25519SnapshotSigner.from_private_key_file(
                "license.test.secure",
                path,
                require_owner_only=True,
            )
            self.assertEqual(signer.key_id, "license.test.secure")
            self.assertEqual(len(signer.public_key_bytes()), 32)


if __name__ == "__main__":
    unittest.main()
