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]