Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
192 changes: 192 additions & 0 deletions server_tests/test_login_and_hardening.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,192 @@
"""Tests for login rate limiting, request-size limits, and image validation.

Covers the server-side security hardening added to the authentication flow and
the AI message endpoints:

- Correct/incorrect PIN handling with constant-time comparison.
- Per-IP failed-attempt cooldown keyed on the real socket peer (X-Forwarded-For
spoofing must not reset or bypass the limit).
- A global failed-attempt ceiling that locks out even fresh IPs.
- Server-side attached-image validation (count and per-image size caps).
- MAX_CONTENT_LENGTH configuration and the JSON 413 error handler.
"""

from __future__ import annotations

import json
import os
import unittest
from typing import List, Optional
from unittest.mock import Mock, patch

from werkzeug.test import TestResponse

from static.app_manager import AppManager, MatHudFlask
from static.config import MAX_ATTACHED_IMAGES, MAX_CONTENT_LENGTH_BYTES, MAX_IMAGE_BASE64_BYTES
from static.routes import GLOBAL_FAILED_ATTEMPTS_LIMIT, reset_login_rate_limit_state


TEST_PIN = "123456"


class TestLoginRateLimiting(unittest.TestCase):
"""Exercise the /login handler's authentication and rate-limiting logic."""

def setUp(self) -> None:
self._original_require_auth: Optional[str] = os.environ.get("REQUIRE_AUTH")
self._original_auth_pin: Optional[str] = os.environ.get("AUTH_PIN")
os.environ["REQUIRE_AUTH"] = "true"
os.environ["AUTH_PIN"] = TEST_PIN

self.app: MatHudFlask = AppManager.create_app()
self.app.config["TESTING"] = True
self.client = self.app.test_client()

# Ensure no leaked rate-limit state from other tests.
reset_login_rate_limit_state()

def tearDown(self) -> None:
reset_login_rate_limit_state()
self._restore_env("REQUIRE_AUTH", self._original_require_auth)
self._restore_env("AUTH_PIN", self._original_auth_pin)

@staticmethod
def _restore_env(name: str, value: Optional[str]) -> None:
if value is not None:
os.environ[name] = value
else:
os.environ.pop(name, None)

def _post_login(self, pin: str, remote_addr: str = "127.0.0.1", xff: Optional[str] = None) -> TestResponse:
headers = {"X-Forwarded-For": xff} if xff is not None else None
return self.client.post(
"/login",
data={"pin": pin},
headers=headers,
environ_base={"REMOTE_ADDR": remote_addr},
)

def test_correct_pin_succeeds(self) -> None:
response = self._post_login(TEST_PIN)
self.assertEqual(response.status_code, 302) # redirect to index
with self.client.session_transaction() as sess:
self.assertTrue(sess.get("authenticated"))

def test_wrong_pin_fails(self) -> None:
response = self._post_login("000000")
self.assertEqual(response.status_code, 200) # re-render login, not a redirect
with self.client.session_transaction() as sess:
self.assertFalse(sess.get("authenticated"))

def test_per_ip_cooldown_triggers_after_failure(self) -> None:
# First failure records a timestamp and re-renders the login page.
first = self._post_login("000000")
self.assertEqual(first.status_code, 200)

# Second failure within the cooldown window is rejected with 429.
second = self._post_login("000000")
self.assertEqual(second.status_code, 429)

def test_x_forwarded_for_does_not_bypass_per_ip_limit(self) -> None:
# Same socket peer (REMOTE_ADDR), different spoofed XFF values on each request.
self.assertEqual(self._post_login("000000", xff="10.0.0.1").status_code, 200)

# Varying the client-controlled header must NOT reset the per-IP cooldown,
# because limiting keys on REMOTE_ADDR, not X-Forwarded-For.
self.assertEqual(self._post_login("000000", xff="10.0.0.2").status_code, 429)
self.assertEqual(self._post_login("000000", xff="203.0.113.9").status_code, 429)

def test_global_ceiling_locks_out_fresh_ip(self) -> None:
# Drive the global counter to its ceiling using a distinct IP per request so
# the per-IP cooldown never engages.
for i in range(GLOBAL_FAILED_ATTEMPTS_LIMIT):
response = self._post_login("000000", remote_addr=f"198.51.100.{i}")
self.assertEqual(response.status_code, 200)

# A brand-new, never-seen IP is now locked out purely by the global ceiling.
locked = self._post_login("000000", remote_addr="192.0.2.55")
self.assertEqual(locked.status_code, 429)

# Even a correct PIN is rejected while the global lockout is active.
locked_correct = self._post_login(TEST_PIN, remote_addr="192.0.2.77")
self.assertEqual(locked_correct.status_code, 429)


class TestImageAndSizeHardening(unittest.TestCase):
"""Exercise server-side image validation and request-size configuration."""

SAMPLE_PNG_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/w8AAwMB/aqVw0sAAAAASUVORK5CYII="

def setUp(self) -> None:
self._original_require_auth: Optional[str] = os.environ.get("REQUIRE_AUTH")
os.environ["REQUIRE_AUTH"] = "false"

self.app = AppManager.create_app()
self.app.config["TESTING"] = True
self.client = self.app.test_client()

def tearDown(self) -> None:
if self._original_require_auth is not None:
os.environ["REQUIRE_AUTH"] = self._original_require_auth
else:
os.environ.pop("REQUIRE_AUTH", None)

def _small_image(self) -> str:
return f"data:image/png;base64,{self.SAMPLE_PNG_BASE64}"

def _post_message(self, images: List[str]) -> TestResponse:
payload = {"message": json.dumps({"user_message": "hi", "attached_images": images})}
return self.client.post("/send_message", json=payload)

def test_too_many_images_rejected(self) -> None:
images = [self._small_image()] * (MAX_ATTACHED_IMAGES + 1)
response = self._post_message(images)
self.assertEqual(response.status_code, 400)
data = json.loads(response.data)
self.assertEqual(data["status"], "error")
self.assertIn("Too many attached images", data["message"])

def test_oversized_single_image_rejected(self) -> None:
oversized = "data:image/png;base64," + ("A" * (MAX_IMAGE_BASE64_BYTES + 1))
response = self._post_message([oversized])
self.assertEqual(response.status_code, 400)
data = json.loads(response.data)
self.assertEqual(data["status"], "error")
self.assertIn("maximum size", data["message"])

@patch("static.openai_completions_api.OpenAIChatCompletionsAPI.create_chat_completion")
def test_valid_image_payload_passes_validation(self, mock_chat: Mock) -> None:
class MockMessage:
content = "ok"
tool_calls = None

class MockResponse:
message = MockMessage()
finish_reason = "stop"

mock_chat.return_value = MockResponse()

response = self._post_message([self._small_image()] * MAX_ATTACHED_IMAGES)
self.assertEqual(response.status_code, 200)
data = json.loads(response.data)
self.assertEqual(data["status"], "success")

def test_max_content_length_configured(self) -> None:
self.assertEqual(self.app.config["MAX_CONTENT_LENGTH"], MAX_CONTENT_LENGTH_BYTES)

def test_413_handler_returns_json(self) -> None:
# Shrink the limit for this request so we don't have to send 80 MB.
self.app.config["MAX_CONTENT_LENGTH"] = 64
response = self.client.post(
"/send_message",
data="x" * 512,
content_type="application/json",
)
self.assertEqual(response.status_code, 413)
self.assertIn("application/json", response.content_type)
data = json.loads(response.data)
self.assertIn("error", data)


if __name__ == "__main__":
unittest.main()
76 changes: 75 additions & 1 deletion server_tests/test_markdown_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,10 +89,84 @@ def hello():
self.assertIn("</code></pre>", result)

def test_links(self) -> None:
"""Test link formatting."""
"""Test link formatting for a safe (allowed-scheme) URL."""
result = self.parser.parse("This is a [link](https://example.com) text")
self.assertIn('<a href="https://example.com">link</a>', result)

def test_raw_script_tag_is_escaped(self) -> None:
"""Raw <script> in AI output must be escaped, not executed."""
result = self.parser.parse("<script>alert(1)</script>")
self.assertNotIn("<script>", result)
self.assertIn("&lt;script&gt;", result)
self.assertIn("&lt;/script&gt;", result)

def test_img_onerror_is_escaped(self) -> None:
"""An <img onerror=...> payload must be escaped, not rendered."""
result = self.parser.parse("<img src=x onerror=alert(1)>")
self.assertNotIn("<img", result)
self.assertIn("&lt;img", result)

def test_javascript_link_is_neutralized(self) -> None:
"""javascript: links must not produce a live href."""
result = self.parser.parse("[x](javascript:alert(1))")
self.assertNotIn("href=\"javascript", result)
self.assertNotIn("javascript:alert", result)
# Link text is preserved as plain text.
self.assertIn("x", result)

def test_javascript_link_mixed_case_is_neutralized(self) -> None:
"""Mixed-case JaVaScRiPt: scheme must be rejected case-insensitively."""
result = self.parser.parse("[x](JaVaScRiPt:alert(1))")
self.assertNotIn("<a href", result)
self.assertNotIn("alert(1)", result)

def test_javascript_link_with_control_char_is_neutralized(self) -> None:
"""Embedded control chars (java\\tscript:) must not bypass sanitization."""
result = self.parser.parse("[x](java\tscript:alert(1))")
self.assertNotIn("<a href", result)
self.assertNotIn("alert(1)", result)

def test_data_uri_link_is_neutralized(self) -> None:
"""data: URIs must be rejected."""
result = self.parser.parse("[x](data:text/html,<script>alert(1)</script>)")
self.assertNotIn("<a href", result)

def test_safe_https_link_still_renders(self) -> None:
"""Allowed https links must still render as real anchors."""
result = self.parser.parse("[ok](https://example.com)")
self.assertIn('<a href="https://example.com">ok</a>', result)

def test_mailto_link_still_renders(self) -> None:
"""mailto: links are allowed."""
result = self.parser.parse("[mail](mailto:a@b.com)")
self.assertIn('<a href="mailto:a@b.com">mail</a>', result)

def test_relative_link_still_renders(self) -> None:
"""Relative URLs (no scheme) are allowed."""
result = self.parser.parse("[rel](/path/page)")
self.assertIn('<a href="/path/page">rel</a>', result)

def test_normal_markdown_still_renders_after_escaping(self) -> None:
"""Bold, code, and headers must still render once escaping is in place."""
result = self.parser.parse("# Title\n\nSome **bold** and `code` here")
self.assertIn("<h1>Title</h1>", result)
self.assertIn("<strong>bold</strong>", result)
self.assertIn("<code>code</code>", result)

def test_angle_brackets_in_prose_are_escaped(self) -> None:
"""Comparison operators in prose render as escaped entities."""
result = self.parser.parse("a < b and c > d")
self.assertIn("a &lt; b and c &gt; d", result)
self.assertNotIn("<b", result)

def test_code_block_with_script_is_escaped(self) -> None:
"""A code block containing <script> must be escaped inside <pre><code>."""
code_block = "```\n<script>alert(1)</script>\n```"
result = self.parser.parse(code_block)
self.assertIn("<pre><code>", result)
self.assertIn("&lt;script&gt;", result)
self.assertNotIn("<script>", result)

def test_unordered_lists(self) -> None:
"""Test unordered list formatting."""
markdown = """- Item 1
Expand Down
30 changes: 28 additions & 2 deletions static/app_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

from __future__ import annotations

import logging
import os
import secrets
from typing import TYPE_CHECKING, Dict, Optional, Tuple, TypedDict, Union
Expand All @@ -22,6 +23,7 @@
from flask import Flask, Response, jsonify
from flask_session import Session as FlaskSession

from static.config import MAX_CONTENT_LENGTH_BYTES
from static.env_config import load_env_files
from static.log_manager import LogManager
from static.openai_completions_api import OpenAIChatCompletionsAPI
Expand All @@ -35,6 +37,9 @@
from static.webdriver_manager import WebDriverManager


_logger = logging.getLogger(__name__)


JsonValue = Union[str, int, float, bool, None, Dict[str, "JsonValue"], list["JsonValue"]]


Expand Down Expand Up @@ -156,8 +161,29 @@ def create_app() -> MatHudFlask:
# Load environment variables from project .env and parent .env (API keys)
AppManager._load_env()

# Configure session management for authentication using modern CacheLib backend
app.secret_key = os.getenv("SECRET_KEY", secrets.token_hex(32))
# Configure session management for authentication using modern CacheLib backend.
# A persistent SECRET_KEY is required for sessions to survive process restarts.
# In deployed mode an ephemeral fallback silently logs everyone out on every
# restart, so warn prominently (but do not hard-fail).
secret_key_env = os.getenv("SECRET_KEY")
if secret_key_env:
app.secret_key = secret_key_env
else:
app.secret_key = secrets.token_hex(32)
if AppManager.is_deployed():
_logger.warning(
"SECRET_KEY is not set in a deployed environment. A random key was "
"generated for this process, so ALL user sessions will be invalidated "
"on every restart. Set SECRET_KEY to a stable secret to persist sessions."
)

# Cap total request body size and return clean JSON on overflow so clients
# do not receive Flask's default HTML 413 page.
app.config["MAX_CONTENT_LENGTH"] = MAX_CONTENT_LENGTH_BYTES

@app.errorhandler(413)
def _handle_request_entity_too_large(_error: Exception) -> Tuple[Response, int]:
return jsonify({"error": "Request payload too large"}), 413

# Create session directory if it doesn't exist
session_dir = os.path.join(os.getcwd(), "flask_session")
Expand Down
8 changes: 5 additions & 3 deletions static/client/chat_ui_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,10 +187,12 @@ def handler(event: Any) -> None:

except Exception as e:
print(f"Error creating message element: {e}")
# Fall back to simple paragraph
# Fall back to simple paragraph. Escape sender/content before
# interpolating into innerHTML so raw HTML never executes here.
if sender == "AI":
content = message.replace("\n", "<br>")
return html.P(f"<strong>{sender}:</strong> {content}", innerHTML=True)
content = self._escape_html(message).replace("\n", "<br>")
safe_sender = self._escape_html(sender)
return html.P(f"<strong>{safe_sender}:</strong> {content}", innerHTML=True)
else:
return html.P(f"<strong>{sender}:</strong> {message}")

Expand Down
Loading
Loading