patx/projectpay

import asyncio
import json
import os
import unittest
from types import SimpleNamespace
from unittest.mock import patch

import micropie

with patch.dict(os.environ, {"APP_ENV": "test"}, clear=False):
    import app as invoice_app


class SecurityHardeningTests(unittest.TestCase):
    def make_app(self, **env):
        defaults = {
            "APP_ENV": "development",
            "MONGODB_URI": "mongodb://localhost:27017",
            "MONGODB_DB": "test_projectpay",
        }
        defaults.update(env)
        with patch.dict(os.environ, defaults, clear=True):
            return invoice_app.InvoiceApp()

    def bind_request(self, app, request):
        token = micropie.current_request.set(request)
        self.addCleanup(micropie.current_request.reset, token)
        return app

    async def asgi_get(self, app, path):
        messages = []

        async def receive():
            return {"type": "http.request", "body": b"", "more_body": False}

        async def send(message):
            messages.append(message)

        await app(
            {
                "type": "http",
                "method": "GET",
                "path": path,
                "query_string": b"",
                "headers": [],
                "client": ("127.0.0.1", 1234),
            },
            receive,
            send,
        )
        response_start = messages[0]
        body = b"".join(
            message.get("body", b"")
            for message in messages
            if message["type"] == "http.response.body"
        )
        headers = {
            key.decode("latin-1").lower(): value.decode("latin-1")
            for key, value in response_start["headers"]
        }
        return response_start["status"], headers, body

    def test_admin_login_starts_fresh_server_side_session(self):
        class FakeSessionBackend:
            def __init__(self):
                self.saved = None

            async def save(self, session_id, data, timeout):
                self.saved = (session_id, data, timeout)

        app = self.make_app()
        backend = FakeSessionBackend()
        app.session_backend = backend
        request = SimpleNamespace(
            session={"csrf_token": "old-token"},
            body_params={},
            headers={},
            scope={"client": ("127.0.0.1", 1234)},
            form=lambda name, default=None: default,
        )
        self.bind_request(app, request)

        response = asyncio.run(app._start_admin_session())

        self.assertEqual(response[0], 302)
        self.assertEqual(request.session, {})
        self.assertIsNotNone(backend.saved)
        session_id, data, timeout = backend.saved
        self.assertTrue(session_id)
        self.assertTrue(data["is_admin"])
        self.assertTrue(data["csrf_token"])
        self.assertGreater(timeout, 0)
        self.assertIn(f"{invoice_app.SESSION_COOKIE}={session_id}", response[2][1][1])

    def test_logout_expires_session_cookie(self):
        app = self.make_app()
        request = SimpleNamespace(
            session={"is_admin": True},
            body_params={},
            headers={},
            scope={"client": ("127.0.0.1", 1234)},
            form=lambda name, default=None: default,
        )
        self.bind_request(app, request)

        response = app.logout()

        self.assertEqual(request.session, {})
        self.assertEqual(response[0], 302)
        self.assertIn(f"{invoice_app.SESSION_COOKIE}=", response[2][1][1])
        self.assertIn("Max-Age=0", response[2][1][1])

    def test_production_requires_non_default_admin_password(self):
        with patch.dict(
            os.environ,
            {
                "APP_ENV": "production",
                "ADMIN_PASSWORD": "admin",
                "APP_BASE_URL": "https://example.com",
            },
            clear=True,
        ):
            with self.assertRaisesRegex(RuntimeError, "ADMIN_PASSWORD"):
                invoice_app.InvoiceApp()

    def test_app_env_must_be_explicit(self):
        with patch.dict(os.environ, {}, clear=True):
            with self.assertRaisesRegex(RuntimeError, "APP_ENV"):
                invoice_app.InvoiceApp()

    def test_production_requires_app_base_url(self):
        with patch.dict(
            os.environ,
            {
                "APP_ENV": "production",
                "ADMIN_PASSWORD": "strong-password",
                "APP_BASE_URL": "",
            },
            clear=True,
        ):
            with self.assertRaisesRegex(RuntimeError, "APP_BASE_URL"):
                invoice_app.InvoiceApp()

    def test_csrf_token_is_session_backed_and_validated(self):
        app = self.make_app()
        request = SimpleNamespace(
            session={},
            body_params={},
            headers={},
            scope={"client": ("127.0.0.1", 1234)},
            form=lambda name, default=None: default,
        )
        self.bind_request(app, request)

        csrf_token = app._csrf_token()
        self.assertTrue(csrf_token)
        self.assertFalse(app._valid_csrf())

        request.form = lambda name, default=None: {
            "csrf_token": csrf_token,
        }.get(name, default)
        self.assertTrue(app._valid_csrf())

    def test_malformed_csrf_token_is_rejected_without_error(self):
        app = self.make_app()
        request = SimpleNamespace(
            session={invoice_app.CSRF_SESSION_KEY: "ascii-token"},
            body_params={},
            headers={},
            scope={"client": ("127.0.0.1", 1234)},
            form=lambda name, default=None: {
                "csrf_token": "snowman-\u2603",
            }.get(name, default),
        )
        self.bind_request(app, request)

        self.assertFalse(app._valid_csrf())

    def test_direct_run_host_defaults_to_loopback(self):
        self.assertEqual(invoice_app.DEFAULT_HOST, "127.0.0.1")

    def test_manifest_json_uses_default_app_icon_url(self):
        app = self.make_app()

        status, headers, body = asyncio.run(self.asgi_get(app, "/manifest.json"))

        self.assertEqual(status, 200)
        self.assertEqual(headers["content-type"], "application/manifest+json")
        manifest = json.loads(body)
        self.assertEqual(manifest["name"], "ProjectPay")
        self.assertEqual(manifest["short_name"], "ProjectPay")
        self.assertEqual(manifest["start_url"], "/")
        self.assertEqual(manifest["scope"], "/")
        self.assertEqual(manifest["display"], "standalone")
        self.assertEqual(manifest["theme_color"], "#176b5b")
        self.assertEqual(manifest["icons"][0]["src"], invoice_app.DEFAULT_APP_ICON_URL)
        self.assertEqual(manifest["icons"][0]["sizes"], "512x512")
        self.assertEqual(manifest["icons"][0]["type"], "image/png")

    def test_manifest_json_uses_configured_app_icon_url(self):
        app = self.make_app(PUBLIC_APP_ICON_URL="https://example.com/app-icon.png")

        status, headers, body = asyncio.run(self.asgi_get(app, "/manifest.json"))

        self.assertEqual(status, 200)
        self.assertEqual(headers["content-type"], "application/manifest+json")
        manifest = json.loads(body)
        self.assertEqual(manifest["icons"][0]["src"], "https://example.com/app-icon.png")

    def test_public_dark_logo_url_defaults_to_public_logo_url(self):
        app = self.make_app(PUBLIC_LOGO_URL="https://example.com/logo.svg")

        self.assertEqual(app.public_dark_logo_url, "https://example.com/logo.svg")

    def test_public_dark_logo_url_can_be_configured(self):
        app = self.make_app(
            PUBLIC_LOGO_URL="https://example.com/logo.svg",
            PUBLIC_DARK_LOGO_URL="https://example.com/logo-dark.svg",
        )

        self.assertEqual(
            app.public_dark_logo_url, "https://example.com/logo-dark.svg"
        )

    def test_development_base_url_uses_request_headers_only_in_development(self):
        app = self.make_app(APP_BASE_URL="")
        request = SimpleNamespace(
            session={},
            body_params={},
            headers={"host": "local.test:8000", "x-forwarded-proto": "gopher"},
            scope={"client": ("127.0.0.1", 1234)},
            form=lambda name, default=None: default,
        )
        self.bind_request(app, request)

        self.assertEqual(app._base_url(), "http://local.test:8000")

    def test_heroku_client_ip_uses_last_forwarded_for_entry(self):
        app = self.make_app(DYNO="web.1")
        request = SimpleNamespace(
            session={},
            body_params={},
            headers={"x-forwarded-for": "spoofed, 203.0.113.10"},
            scope={"client": ("10.0.0.1", 1234)},
            form=lambda name, default=None: default,
        )
        self.bind_request(app, request)

        self.assertEqual(app._client_ip(), "203.0.113.10")

    def test_stripe_event_summary_excludes_full_customer_payload(self):
        app = self.make_app()
        summary = app._stripe_event_summary(
            {
                "id": "evt_123",
                "type": "checkout.session.completed",
                "created": 1234567890,
                "livemode": False,
                "data": {
                    "object": {
                        "id": "cs_123",
                        "object": "checkout.session",
                        "payment_status": "paid",
                        "amount_total": 2500,
                        "currency": "usd",
                        "customer_email": "[email protected]",
                        "metadata": {
                            "project_id": "507f1f77bcf86cd799439011",
                            "project_number": "PRJ-2026-0001",
                            "share_token": "share-token",
                        },
                    }
                },
            }
        )

        self.assertEqual(summary["id"], "evt_123")
        self.assertNotIn("customer_email", str(summary))
        self.assertEqual(summary["object"]["amount_total"], 2500)


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