From 1d6d87e2bef1a0f0db78896abd1239708654783b Mon Sep 17 00:00:00 2001 From: Alexandre Teixeira <111787685+alteixeira20@users.noreply.github.com> Date: Tue, 6 Oct 2026 03:16:54 +0100 Subject: [PATCH] fix(security): harden Python service boundaries --- mcp_servers/email_server.py | 19 +- routes/calendar_routes.py | 8 +- routes/chat_helpers.py | 12 +- routes/chat_routes.py | 22 +- routes/email/email_helpers.py | 169 +++++- routes/email/email_routes.py | 359 +++++++----- routes/hwfit_routes.py | 9 +- routes/search/search_routes.py | 21 +- routes/session_routes.py | 10 +- routes/skills_routes.py | 9 +- routes/task/task_routes.py | 8 +- services/memory/memory_extractor.py | 11 +- src/llm_core.py | 16 +- src/mail_attachment_paths.py | 32 ++ src/readiness.py | 11 +- src/text_helpers.py | 14 + src/tools/calendar.py | 10 +- src/tools/notes.py | 2 +- src/upload_handler.py | 79 ++- src/url_safety.py | 49 ++ .../test_calendar_reminder_minutes_parsing.py | 34 ++ tests/test_codeql_python_residual_closure.py | 522 ++++++++++++++++++ tests/test_codeql_python_residual_followup.py | 260 +++++++++ tests/test_codeql_python_residual_security.py | 164 ++++++ tests/test_email_imap_timeout.py | 12 +- tests/test_email_send_only_no_inbox.py | 2 +- tests/test_email_smtp_security.py | 18 +- tests/test_email_test_connection_oauth.py | 24 +- tests/test_imap_leak_fixes.py | 2 +- tests/test_notes_reminder_temporal_redos.py | 60 ++ tests/test_redos_residual_python.py | 183 ++++++ tests/test_security_regressions.py | 8 +- 32 files changed, 1911 insertions(+), 248 deletions(-) create mode 100644 src/mail_attachment_paths.py create mode 100644 tests/test_codeql_python_residual_closure.py create mode 100644 tests/test_codeql_python_residual_followup.py create mode 100644 tests/test_codeql_python_residual_security.py create mode 100644 tests/test_notes_reminder_temporal_redos.py create mode 100644 tests/test_redos_residual_python.py diff --git a/mcp_servers/email_server.py b/mcp_servers/email_server.py index af1005ef3..ed8d5acff 100644 --- a/mcp_servers/email_server.py +++ b/mcp_servers/email_server.py @@ -3504,6 +3504,21 @@ def _block_sender(sender=None, uids=None, folder="INBOX", account=None, reason=" } +def _attachment_dir(folder, uid, account_id) -> Path: + """Owner/account-scoped extraction directory (see src.mail_attachment_paths). + + `folder` comes from tool arguments and the IMAP server's own mailbox names + (`/` hierarchies, absolute or `..` segments), so it must not become a path; + the old `MAIL_ATTACHMENTS_DIR/{folder}_{uid}` wrote outside the root and + shared one directory between every owner's "INBOX 42". + """ + from src.mail_attachment_paths import attachment_scope_dir + + return attachment_scope_dir( + MAIL_ATTACHMENTS_DIR, folder, uid, owner=_current_owner(), account_id=account_id, + ) + + def _download_attachment(uid, index, folder="INBOX", account=None): """Extract a specific attachment to disk and return its local path.""" fixture = _fixture_attachment_source(uid, index, folder=folder, account=account) @@ -3512,7 +3527,7 @@ def _download_attachment(uid, index, folder="INBOX", account=None): filename = str(att.get("filename") or f"attachment-{index}.txt") safe_name = re.sub(r"[^\w\s\-.]", "_", filename).strip() or f"attachment-{index}.txt" content = str(att.get("content") or "") - target_dir = Path(MAIL_ATTACHMENTS_DIR) / re.sub(r"[^A-Za-z0-9._-]", "_", f"{folder}_{uid}") + target_dir = _attachment_dir(folder, uid, _row.get("account_id") or account) path = target_dir / safe_name size = len(content.encode("utf-8")) try: @@ -3545,7 +3560,7 @@ def _download_attachment(uid, index, folder="INBOX", account=None): raw = msg_data[0][1] msg = email.message_from_bytes(raw) - target_dir = Path(MAIL_ATTACHMENTS_DIR) / f"{folder}_{uid}" + target_dir = _attachment_dir(folder, uid, _load_config(account).get("account_id") or account) filepath = _extract_attachment_to_disk(msg, index, target_dir) if not filepath: return {"error": f"Attachment index {index} not found"} diff --git a/routes/calendar_routes.py b/routes/calendar_routes.py index ac7107ccb..aa84768fd 100644 --- a/routes/calendar_routes.py +++ b/routes/calendar_routes.py @@ -409,7 +409,7 @@ def parse_due_for_user(s: str) -> str: lower = s.lower().strip() def _parse_time(t): - t = _re.sub(r'\b([ap])\s*\.?\s*m\.?\b', r'\1m', t.strip(), flags=_re.IGNORECASE) + t = _re.sub(r'\b([ap])(?:\s*\.)?\s*m\.?\b', r'\1m', t.strip(), flags=_re.IGNORECASE) m = _re.match(r'^\s*(\d{1,2})(?::(\d{2}))?\s*(am|pm)?\s*$', t, _re.IGNORECASE) if not m: return None h = int(m.group(1)); mn = int(m.group(2) or 0); ampm = (m.group(3) or "").lower() @@ -433,7 +433,7 @@ def parse_due_for_user(s: str) -> str: return base.replace(hour=t[0], minute=t[1]).isoformat() # Time-first: "3pm today", "11pm today", "9am tomorrow" - m = _re.match(r'^(.+?)\s+(today|tonight|tomorrow|tmrw|yesterday)$', lower) + m = _re.match(r'^(.*\S)\s+(today|tonight|tomorrow|tmrw|yesterday)$', lower) if m: time_part, word = m.group(1).strip(), m.group(2) base = today @@ -530,7 +530,7 @@ def _parse_dt(s: str) -> datetime: def _parse_time(t: str): """Return (hour, minute) from '1pm', '1:30 PM', '13:00', etc., or None.""" - t = _re.sub(r'\b([ap])\s*\.?\s*m\.?\b', r'\1m', t.strip(), flags=_re.IGNORECASE) + t = _re.sub(r'\b([ap])(?:\s*\.)?\s*m\.?\b', r'\1m', t.strip(), flags=_re.IGNORECASE) m = _re.match(r'^\s*(\d{1,2})(?::(\d{2}))?\s*(am|pm)?\s*$', t, _re.IGNORECASE) if not m: return None @@ -562,7 +562,7 @@ def _parse_dt(s: str) -> datetime: # time-first: "3pm today", "9am tomorrow", "11pm tonight" # (parity with parse_due_for_user, which handles these via the same form) - m = _re.match(r'^(.+?)\s+(today|tonight|tomorrow|tmrw|yesterday)$', lower) + m = _re.match(r'^(.*\S)\s+(today|tonight|tomorrow|tmrw|yesterday)$', lower) if m: time_part, word = m.group(1).strip(), m.group(2) base = today diff --git a/routes/chat_helpers.py b/routes/chat_helpers.py index eb5fcf489..472128eaf 100644 --- a/routes/chat_helpers.py +++ b/routes/chat_helpers.py @@ -98,7 +98,9 @@ def clean_repeated_assistant_content(text: object) -> str: "", value, ).strip() - value = re.sub(r"(?is)\s*\s*$", "", value).strip() + # No leading `\s*`: .strip() removes that whitespace anyway, and scanning + # it from every offset of a long whitespace run was quadratic (ReDoS). + value = re.sub(r"(?is)\s*$", "", value).strip() return value _CASUAL_OPENING_RE = re.compile( @@ -1324,13 +1326,17 @@ def _normalize_thinking(text: str) -> str: # Handle garbled tags: reasoning text followed by as separator # e.g. "The user said...I should respond.\nHey! What's up?" + # Linear form of `^([\s\S]+?)\n*\s*([\s\S]*?)(?:)?\s*$`: + # the lookbehind stops the lazy prefix re-scanning a newline run from every + # offset, and the optional trailing closer is dropped after the match + # instead of being retried at every body offset (both were quadratic). garbled = re.match( - r'^([\s\S]+?)\n*\s*([\s\S]*?)(?:)?\s*$', + r'^([\s\S]+?)(?\s*([\s\S]*)$', text, re.IGNORECASE ) if garbled: before = garbled.group(1).strip() - after = garbled.group(2).strip() + after = re.sub(r'$', '', garbled.group(2).rstrip(), flags=re.IGNORECASE).strip() # Only treat as garbled if the part before looks like reasoning reasoning_starts = ( 'The user ', 'I need ', 'I should ', 'I will ', diff --git a/routes/chat_routes.py b/routes/chat_routes.py index 3bd0f01bf..e3f7bfa02 100644 --- a/routes/chat_routes.py +++ b/routes/chat_routes.py @@ -327,13 +327,31 @@ def _is_external_discovery_request(text: str) -> bool: )) +_URL_SCHEME_RE = re.compile(r"https?://", re.I) +_PDF_URL_TAIL_RE = re.compile(r"\.pdf\b|/pdf/", re.I) + + +def _mentions_pdf_url(value: str) -> bool: + """Same as ``re.search(r"https?://[^\\s]+(?:\\.pdf\\b|/pdf/)", value, re.I)``. + + Checked per whitespace-free token from its first scheme only (a later + scheme's tail is a suffix of the first's), so a token packed with + `http://` repeats is scanned once instead of once per repeat (ReDoS). + """ + for token in value.split(): + scheme = _URL_SCHEME_RE.search(token) + if scheme and _PDF_URL_TAIL_RE.search(token, scheme.end() + 1): + return True + return False + + def _prefers_structured_document_tools(text: str) -> bool: """Identify external paper/PDF extraction where shell is a bad source route.""" value = str(text or "") if re.search(r"(?:^|\s)(?:file://)?/workspace/[^\s`\"']+\.pdf\b", value, re.I): return False return bool( - re.search(r"https?://[^\s]+(?:\.pdf\b|/pdf/)", value, re.I) + _mentions_pdf_url(value) or re.search( r"\b(?:paper|report|study)\b[\s\S]{0,1200}?" r"\b(?:tables?|figures?|benchmarks?|scores?|metrics?)\b", @@ -2507,7 +2525,7 @@ def setup_chat_routes( _explicit_web_intent = ( _explicit_url_target or bool(re.search( - r"\b(search|look\s+(?:this|that|it|them|these|those)?\s*up|lookup|find\s*out|google|browse|web|online|latest|current|today|news|weather|forecast|rate|exchange\s+rate)\b", + r"\b(search|look\s+(?:(?:this|that|it|them|these|those)\s*)?up|lookup|find\s*out|google|browse|web|online|latest|current|today|news|weather|forecast|rate|exchange\s+rate)\b", _msg_l, )) or requires_external_web_verification(message) diff --git a/routes/email/email_helpers.py b/routes/email/email_helpers.py index 26ef3aa50..163132b85 100644 --- a/routes/email/email_helpers.py +++ b/routes/email/email_helpers.py @@ -17,6 +17,8 @@ import base64 import time import imaplib import smtplib +import socket +import ssl import email as email_mod import email.header import email.utils @@ -36,6 +38,7 @@ from typing import Optional, List from src.auth_helpers import _auth_disabled, get_current_user from src.secret_storage import decrypt as _decrypt +from src.url_safety import OutboundAddressBlocked, connect_outbound_tcp logger = logging.getLogger(__name__) @@ -160,6 +163,94 @@ def _smtp_security_mode(cfg: dict) -> str: return "ssl" +# Raised when a mail host resolves into a denied range (see _mail_private_blocked). +MailServerAddressBlocked = OutboundAddressBlocked + + +def _mail_private_blocked(owner: str | None) -> bool: + """Whether private/loopback/shared mail destinations are denied for *owner*. + + Link-local (cloud metadata), multicast, reserved and unspecified addresses + are always denied (src/url_safety.py). Private ranges are where a mail + tester or saved account becomes an internal-network probe, so they are + only allowed for an admin or in single-user mode (AUTH_ENABLED=false), the + same principals trusted with host-level tools. Operators of multi-user + deployments with a LAN mail server opt in with EMAIL_ALLOW_PRIVATE_IPS=true; + EMAIL_BLOCK_PRIVATE_IPS=true denies private ranges for everyone. + """ + if os.getenv("EMAIL_BLOCK_PRIVATE_IPS", "").strip().lower() == "true": + return True + if os.getenv("EMAIL_ALLOW_PRIVATE_IPS", "").strip().lower() == "true": + return False + from src.tool_security import owner_is_admin_or_single_user + + return not owner_is_admin_or_single_user((owner or "").strip() or None) + + +def _mail_socket_timeout(timeout): + return socket._GLOBAL_DEFAULT_TIMEOUT if timeout is None else timeout + + +class _PolicyIMAP4(imaplib.IMAP4): + """IMAP4 whose socket comes from connect_outbound_tcp. + + The mail host is resolved once and connected to only at the addresses the + address policy approved, so DNS rebinding between check and connect is + impossible. STARTTLS still verifies against ``self.host``. + """ + + def __init__(self, host, port, *, block_private: bool, timeout=None): + self._block_private = block_private + super().__init__(host, port, timeout=timeout) + + def _create_socket(self, timeout): + if timeout is not None and not timeout: + raise ValueError("Non-blocking socket (timeout=0) is not supported") + return connect_outbound_tcp( + self.host, self.port, timeout=_mail_socket_timeout(timeout), block_private=self._block_private, + ) + + +class _PolicyIMAP4_SSL(imaplib.IMAP4_SSL): + """IMAP4_SSL over a policy-checked socket; TLS SNI/verification use the hostname.""" + + def __init__(self, host, port, *, block_private: bool, timeout=None, ssl_context=None): + self._block_private = block_private + super().__init__(host, port, ssl_context=ssl_context, timeout=timeout) + + def _create_socket(self, timeout): + sock = _PolicyIMAP4._create_socket(self, timeout) + return self.ssl_context.wrap_socket(sock, server_hostname=self.host) + + +class _PolicySMTP(smtplib.SMTP): + """SMTP whose socket comes from connect_outbound_tcp (see _PolicyIMAP4).""" + + def __init__(self, host="", port=0, *, block_private: bool, **kwargs): + self._block_private = block_private + super().__init__(host, port, **kwargs) + + def _get_socket(self, host, port, timeout): + if timeout is not None and not timeout: + raise ValueError("Non-blocking socket (timeout=0) is not supported") + return connect_outbound_tcp( + host, port, timeout=timeout, block_private=self._block_private, + source_address=self.source_address, + ) + + +class _PolicySMTP_SSL(smtplib.SMTP_SSL): + """SMTP_SSL over a policy-checked socket; TLS SNI/verification use the hostname.""" + + def __init__(self, host="", port=0, *, block_private: bool, **kwargs): + self._block_private = block_private + super().__init__(host, port, **kwargs) + + def _get_socket(self, host, port, timeout): + sock = _PolicySMTP._get_socket(self, host, port, timeout) + return self.context.wrap_socket(sock, server_hostname=self._host) + + def _send_smtp_message(cfg: dict, from_addr: str, recipients: list[str], message: str | bytes, timeout: int = 30) -> None: """Send through SMTP using the configured transport security mode.""" host = cfg["smtp_host"] @@ -178,14 +269,15 @@ def _send_smtp_message(cfg: dict, from_addr: str, recipients: list[str], message smtp.login(user, password) security = _smtp_security_mode(cfg) + block_private = _mail_private_blocked(cfg.get("owner")) if security == "ssl": - with smtplib.SMTP_SSL(host, port, timeout=timeout) as smtp: + with _PolicySMTP_SSL(host, port, timeout=timeout, block_private=block_private) as smtp: _auth_smtp(smtp) smtp.sendmail(from_addr, recipients, message) return - with smtplib.SMTP(host, port, timeout=timeout) as smtp: + with _PolicySMTP(host, port, timeout=timeout, block_private=block_private) as smtp: if security == "starttls": smtp.starttls() _auth_smtp(smtp) @@ -224,6 +316,39 @@ def _friendly_email_auth_error(protocol: str, host: str, error: object) -> str: return raw[:200] +# Bound at import: callers (and tests) may swap out imaplib.IMAP4 itself. +_IMAP_ABORT = imaplib.IMAP4.abort + + +def _mail_connection_test_error(protocol: str, host: str, error: BaseException) -> str: + """User-facing result for a failed IMAP/SMTP connection test. + + Transport failures are reported by category, never with the peer's own + bytes: a non-mail service answering on the chosen host/port would + otherwise have its banner echoed back (an internal-service fingerprinting + oracle). Errors from a server that does speak the protocol (login/auth + rejections) keep their text via _friendly_email_auth_error, since that is + what the user needs to fix their settings. + """ + if isinstance(error, MailServerAddressBlocked): + return f"{protocol} server address is not allowed by this server's network policy" + if isinstance(error, (_IMAP_ABORT, smtplib.SMTPConnectError, smtplib.SMTPServerDisconnected)): + return f"{protocol} server did not respond like an {protocol} server" + if isinstance(error, ssl.SSLCertVerificationError): + return f"{protocol} TLS certificate verification failed" + if isinstance(error, ssl.SSLError): + return f"{protocol} TLS handshake failed; check the port and security setting" + if isinstance(error, (socket.timeout, TimeoutError)): + return f"{protocol} connection timed out" + if isinstance(error, ConnectionRefusedError): + return f"{protocol} connection refused" + if isinstance(error, socket.gaierror): + return f"{protocol} server name could not be resolved" + if isinstance(error, OSError) and not isinstance(error, smtplib.SMTPException): + return f"{protocol} connection failed" + return _friendly_email_auth_error(protocol, host, error) + + def _strip_think(text: str) -> str: """Email-flavored think strip — thin wrapper over the central helper. @@ -682,18 +807,21 @@ def _ensure_sender_signatures_table(conn): _lg.getLogger(__name__).warning(f"sender_signatures owner-migration skipped: {_mig_e}") -def attachment_extract_dir(folder: str, uid: str) -> Path: - """Containment-safe extraction directory for an attachment. +def attachment_extract_dir(folder: str, uid: str, *, owner: str, account_id: str | None) -> Path: + """Containment-safe extraction directory for one message's attachments. - `folder` and `uid` are user-controlled (query/path params). Flatten them to - a single safe path segment so a value like folder='../../tmp' can't escape - ATTACHMENTS_DIR, then assert containment as belt-and-suspenders.""" - key = re.sub(r"[^A-Za-z0-9._-]", "_", f"{folder}_{uid}") or "_" - target = (ATTACHMENTS_DIR / key).resolve() - base = ATTACHMENTS_DIR.resolve() - if target != base and base not in target.parents: + IMAP UIDs are small per-mailbox counters, so `folder_uid` alone put every + user's and every account's "INBOX 42" in one directory, where extractions + overwrote each other by filename (a path handed to one user's agent could + then hold another user's attachment). See src.mail_attachment_paths: the + directory is keyed by (owner, account, folder, uid) and contained in + ATTACHMENTS_DIR whatever `folder`/`uid` hold.""" + from src.mail_attachment_paths import attachment_scope_dir + + try: + return attachment_scope_dir(ATTACHMENTS_DIR, folder, uid, owner=owner, account_id=account_id) + except ValueError: raise HTTPException(400, "Invalid attachment location") - return target def _init_scheduled_db(): @@ -1074,6 +1202,8 @@ def _get_email_config(account_id: str | None = None, owner: str = "") -> dict: cfg = { "account_id": row.id, "account_name": row.name, + # Principal for the outbound mail address policy. + "owner": row.owner or owner or "", "smtp_host": row.smtp_host or "", "smtp_port": int(row.smtp_port or 465), "smtp_security": _smtp_security_mode({"smtp_security": getattr(row, "smtp_security", ""), "smtp_port": row.smtp_port}), @@ -1107,6 +1237,7 @@ def _get_email_config(account_id: str | None = None, owner: str = "") -> dict: cfg = { "account_id": resolved_id, "account_name": "legacy", + "owner": owner or "", "smtp_host": settings.get("smtp_host", os.environ.get("SMTP_HOST", "")), "smtp_port": int(settings.get("smtp_port", os.environ.get("SMTP_PORT", "465")) or 465), "smtp_security": _smtp_security_mode({ @@ -1170,11 +1301,16 @@ def _open_imap_connection( starttls: bool, timeout: int = _IMAP_TIMEOUT_SECONDS, ssl_context=None, + owner: str | None = None, ): - """Open an IMAP connection using the configured security mode.""" + """Open an IMAP connection using the configured security mode. + + `owner` is the principal the destination is judged for + (see _mail_private_blocked).""" port = int(port or 993) + block_private = _mail_private_blocked(owner) if starttls: - conn = imaplib.IMAP4(host, port, timeout=timeout) + conn = _PolicyIMAP4(host, port, timeout=timeout, block_private=block_private) try: if ssl_context: conn.starttls(ssl_context=ssl_context) @@ -1190,9 +1326,9 @@ def _open_imap_connection( raise elif port == 993: kwargs = {"ssl_context": ssl_context} if ssl_context else {} - conn = imaplib.IMAP4_SSL(host, port, timeout=timeout, **kwargs) + conn = _PolicyIMAP4_SSL(host, port, timeout=timeout, block_private=block_private, **kwargs) else: - conn = imaplib.IMAP4(host, port, timeout=timeout) + conn = _PolicyIMAP4(host, port, timeout=timeout, block_private=block_private) try: conn.sock.settimeout(timeout) except Exception: @@ -1232,6 +1368,7 @@ def _imap_connect(account_id: str | None = None, owner: str = "", cfg["imap_port"], starttls=bool(cfg.get("imap_starttls")), timeout=timeout, + owner=cfg.get("owner"), ) try: if cfg.get("oauth_provider") == "google": diff --git a/routes/email/email_routes.py b/routes/email/email_routes.py index e5b0d3224..0b71116d6 100644 --- a/routes/email/email_routes.py +++ b/routes/email/email_routes.py @@ -36,6 +36,7 @@ from pathlib import Path from email.mime.text import MIMEText from email.mime.multipart import MIMEMultipart +from starlette.concurrency import run_in_threadpool from fastapi import APIRouter, Query, UploadFile, File, BackgroundTasks, HTTPException, Depends, Request from fastapi.responses import FileResponse, StreamingResponse from src.constants import DATA_DIR @@ -60,6 +61,7 @@ from .email_helpers import ( _fetch_sender_thread_context, _pre_retrieve_context, _EMAIL_REPLY_SYS_PROMPT_BASE, _POOL_HOOKS, _friendly_email_auth_error, _email_summary_failure_log_detail, + _mail_connection_test_error, _mail_private_blocked, _PolicySMTP, _PolicySMTP_SSL, _generate_email_summary, EMAIL_SUMMARY_ERROR_CODE, EMAIL_SUMMARY_ERROR_MESSAGE, SendEmailRequest, ExtractStyleRequest, ATTACHMENTS_DIR, COMPOSE_UPLOADS_DIR, SCHEDULED_DB, @@ -1523,6 +1525,44 @@ def _envelope_recipients(*fields: str) -> list: return out +_MD_LINK_URL_RE = re.compile(r"\((https?://)[^)\s]*") + + +def _md_links_to_html(s: str) -> str: + """Linear ``re.sub(r"\\[([^\\]]+)\\]\\((https?://[^)\\s]+)\\)", '\\1', s)``. + + Every `[` before the next `]` shares that `]`, so only the first can match; + and when a URL runs into whitespace/end without `)`, every later `[` whose + `]` also lies before that stop fails the same way. Skipping those spans + keeps a `[[[` or `[a](http://x[a](http://x` flood O(n) instead of the + regex's O(n^2). + """ + out: list[str] = [] + emitted = scan = 0 + while True: + open_at = s.find("[", scan) + if open_at < 0: + break + close_at = s.find("]", open_at + 1) + if close_at < 0: + break + scan = close_at + 1 + if close_at == open_at + 1: + continue + url = _MD_LINK_URL_RE.match(s, close_at + 1) + if url is None: + continue + end = url.end() + if end < len(s) and s[end] == ")" and end > url.end(1): + out.append(s[emitted:open_at]) + out.append(f'{s[open_at + 1:close_at]}') + emitted = scan = end + 1 + elif end > url.end(1): + scan = max(scan, s.rfind("]", scan, end) + 1) + out.append(s[emitted:]) + return "".join(out) + + def _md_to_email_html(text: str) -> str: """Render the compose markdown body to a SAFE HTML fragment for the email's text/html part. Everything is HTML-escaped FIRST (so a pasted " + "" +) +_SOURCES = [ + {"title": "Example result", "url": "https://example.com/a", "snippet": "first"}, + {"title": "Script link", "url": "javascript:document.body.dataset.pwned='2'", "snippet": "second"}, +] + + +def _csp_app(): + from fastapi import FastAPI + from fastapi.responses import HTMLResponse + from core.middleware import SecurityHeadersMiddleware + from routes.search.search_routes import setup_search_routes + + app = FastAPI() + app.add_middleware(SecurityHeadersMiddleware) + + @app.post("/api/search") + async def _stub_search(): + return {"sources": _SOURCES} + + @app.get("/csp-control") + async def _control(): + return HTMLResponse("") + + app.include_router(setup_search_routes(SimpleNamespace())) + return app + + +@pytest.mark.parametrize("q", ["weather in lisbon", _PAYLOAD]) +def test_search_page_script_carries_the_response_csp_nonce(q): + from fastapi.testclient import TestClient + + response = TestClient(_csp_app()).get("/search/web", params={"q": q}) + csp = response.headers["content-security-policy"] + nonce = re.search(r"'nonce-([^']+)'", csp).group(1) + assert "unsafe-inline" not in csp.split("script-src", 1)[1].split(";", 1)[0] + scripts = re.findall(r"]*>", response.text, re.I) + assert scripts == [f'', + "", + "