v0.163.0
  1from __future__ import annotations
  2
  3import datetime
  4import os
  5import pathlib
  6import socket
  7import subprocess
  8import sys
  9import tempfile
 10import time
 11from typing import TYPE_CHECKING
 12from urllib.parse import urlparse
 13
 14from cryptography import x509
 15from cryptography.hazmat.primitives import hashes, serialization
 16from cryptography.hazmat.primitives.asymmetric import rsa
 17from cryptography.x509.oid import NameOID
 18from plain.test import Client
 19
 20if TYPE_CHECKING:
 21    from playwright.sync_api import Browser, Page  # ty: ignore[unresolved-import]
 22
 23
 24class TestBrowser:
 25    def __init__(self, browser: Browser, database_url: str):
 26        self.browser = browser
 27
 28        self.database_url = database_url
 29        self.protocol = "https"
 30        self.host = "localhost"
 31        self.port = _get_available_port()
 32        self.base_url = f"{self.protocol}://{self.host}:{self.port}"
 33        self.server_process: subprocess.Popen | None = None
 34        self.tmpdir = tempfile.TemporaryDirectory()
 35
 36        # Set the initial browser context
 37        self.reset_context()
 38
 39    def force_login(self, user: object) -> None:
 40        # Make sure existing session cookies are cleared
 41        self.context.clear_cookies()
 42
 43        client = Client()
 44        client.force_login(user)
 45
 46        cookies = []
 47
 48        for morsel in client.cookies.values():
 49            cookie = {
 50                "name": morsel.key,
 51                "value": morsel.value,
 52                # Set this by default because playwright needs url or domain/path pair
 53                # (Plain does this in response, but this isn't going through a response)
 54                "domain": self.host,
 55            }
 56            # These fields are all optional
 57            if url := morsel.get("url"):
 58                cookie["url"] = url
 59            if domain := morsel.get("domain"):
 60                cookie["domain"] = domain
 61            if path := morsel.get("path"):
 62                cookie["path"] = path
 63            if expires := morsel.get("expires"):
 64                cookie["expires"] = expires
 65            if httponly := morsel.get("httponly"):
 66                cookie["httpOnly"] = httponly
 67            if secure := morsel.get("secure"):
 68                cookie["secure"] = secure
 69            if samesite := morsel.get("samesite"):
 70                cookie["sameSite"] = samesite
 71
 72            cookies.append(cookie)
 73
 74        self.context.add_cookies(cookies)
 75
 76    def logout(self) -> None:
 77        self.context.clear_cookies()
 78
 79    def reset_context(self) -> None:
 80        """Create a new browser context with the base URL and ignore HTTPS errors."""
 81        self.context = self.browser.new_context(
 82            base_url=self.base_url,
 83            ignore_https_errors=True,
 84        )
 85
 86    def new_page(self) -> Page:
 87        """Create a new page in the current context."""
 88        return self.context.new_page()
 89
 90    def discover_urls(self, urls: list[str]) -> list[str]:
 91        """Recursively discover all URLs on the page and related pages until we don't see anything new"""
 92
 93        def relative_url(url: str) -> str:
 94            """Convert a URL to a relative URL based on the base URL."""
 95            return url.removeprefix(self.base_url)
 96
 97        # Start with the initial URLs
 98        to_visit = {relative_url(url) for url in urls}
 99        visited = set()
100
101        # Create a new page to use for all crawling
102        page = self.context.new_page()
103
104        while to_visit:
105            # Move the url from to_visit to visited
106            url = to_visit.pop()
107
108            response = page.goto(url)
109
110            visited.add(url)
111
112            # Don't process links that aren't on our site
113            if not response.url.startswith(self.base_url):
114                continue
115
116            # Get the current page's path for resolving relative URLs
117            current_page_path = response.url.removeprefix(self.base_url)
118
119            # Find all <a> links on the page
120            for link in page.query_selector_all("a"):
121                if href := link.get_attribute("href"):
122                    # Remove fragments
123                    href = href.split("#")[0]
124                    if not href:
125                        # Empty URL, skip it
126                        continue
127
128                    parsed = urlparse(href)
129                    # Skip non-http(s) links (mailto:, tel:, javascript:, etc.)
130                    if parsed.scheme and parsed.scheme not in ("http", "https"):
131                        continue
132
133                    # Skip external HTTP links
134                    if parsed.scheme in ("http", "https") and not href.startswith(
135                        self.base_url
136                    ):
137                        continue
138
139                    # Handle query-only URLs (e.g., "?stage=approved")
140                    if href.startswith("?"):
141                        href = current_page_path.split("?")[0] + href
142
143                    visit_url = relative_url(href)
144                    if visit_url not in visited:
145                        to_visit.add(visit_url)
146
147        page.close()
148
149        return list(visited)
150
151    def generate_certificates(self) -> tuple[str, str]:
152        """Generate self-signed certificates for HTTPS."""
153
154        # Generate private key
155        private_key = rsa.generate_private_key(
156            public_exponent=65537,
157            key_size=2048,
158        )
159
160        # Create certificate
161        subject = issuer = x509.Name(
162            [
163                x509.NameAttribute(NameOID.COMMON_NAME, self.host),
164            ]
165        )
166
167        cert = (
168            x509.CertificateBuilder()
169            .subject_name(subject)
170            .issuer_name(issuer)
171            .public_key(private_key.public_key())
172            .serial_number(x509.random_serial_number())
173            .not_valid_before(datetime.datetime.now(datetime.UTC))
174            .not_valid_after(
175                datetime.datetime.now(datetime.UTC) + datetime.timedelta(days=365)
176            )
177            .add_extension(
178                x509.SubjectAlternativeName(
179                    [
180                        x509.DNSName(self.host),
181                    ]
182                ),
183                critical=False,
184            )
185            .sign(private_key, hashes.SHA256())
186        )
187
188        # Write certificate and key to files
189        cert_file = pathlib.Path(self.tmpdir.name) / "cert.pem"
190        key_file = pathlib.Path(self.tmpdir.name) / "key.pem"
191
192        with open(cert_file, "wb") as f:
193            f.write(cert.public_bytes(serialization.Encoding.PEM))
194
195        with open(key_file, "wb") as f:
196            f.write(
197                private_key.private_bytes(
198                    encoding=serialization.Encoding.PEM,
199                    format=serialization.PrivateFormat.PKCS8,
200                    encryption_algorithm=serialization.NoEncryption(),
201                )
202            )
203
204        return str(cert_file), str(key_file)
205
206    def run_server(self) -> None:
207        cert_file, key_file = self.generate_certificates()
208
209        env = os.environ.copy()
210
211        if self.database_url:
212            env["DATABASE_URL"] = self.database_url
213
214        self.server_process = subprocess.Popen(
215            [
216                sys.executable,
217                "-m",
218                "plain",
219                "server",
220                "--bind",
221                f"{self.host}:{self.port}",
222                "--certfile",
223                cert_file,
224                "--keyfile",
225                key_file,
226                "--workers",
227                "2",
228                "--timeout",
229                "10",
230            ],
231            env=env,
232        )
233
234        self._wait_for_server()
235
236    def _wait_for_server(self, timeout: float = 10.0, interval: float = 0.1) -> None:
237        """Wait until the server is accepting connections."""
238        deadline = time.monotonic() + timeout
239        while time.monotonic() < deadline:
240            # Check that the server process hasn't crashed
241            if self.server_process and self.server_process.poll() is not None:
242                raise RuntimeError(
243                    f"Server process exited with code {self.server_process.returncode}"
244                )
245            try:
246                with socket.create_connection((self.host, self.port), timeout=interval):
247                    return
248            except OSError:
249                time.sleep(interval)
250        raise RuntimeError(
251            f"Server did not start within {timeout}s at {self.host}:{self.port}"
252        )
253
254    def cleanup_server(self) -> None:
255        if self.server_process:
256            self.server_process.terminate()
257            self.server_process.wait()
258            self.server_process = None
259
260        self.tmpdir.cleanup()
261
262
263def _get_available_port() -> int:
264    """Get a randomly available port."""
265    with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
266        s.bind(("", 0))
267        s.listen(1)
268        return s.getsockname()[1]