Files
tavolo/server/tests/test_config.py
T
woggioni e0896b4a95
CI / Build and push docker image (push) Successful in 3m24s
Configure CORS headers from environment variables
2026-09-18 19:18:27 +08:00

135 lines
4.9 KiB
Python

"""Unit tests for the database DSN assembly in :mod:`tavolo.config`.
``Settings.from_env`` is called directly with a fully replaced
``os.environ`` so no test leaks its ``DATABASE_*`` overrides into the
suite (``tests/__init__.py`` sets ``DATABASE_URL=sqlite://:memory:``
globally for the application tests).
"""
from __future__ import annotations
import os
import unittest
from unittest.mock import patch
from tavolo.config import Settings
def _settings(env: dict) -> Settings:
with patch.dict(os.environ, env, clear=True):
return Settings.from_env()
class DatabaseUrlTests(unittest.TestCase):
def test_defaults_assemble_from_parts(self):
settings = _settings({})
self.assertEqual(
settings.database_url,
"postgres://tavolo:password@localhost/tavolo",
)
def test_components_override_defaults(self):
settings = _settings({
"DATABASE_ENGINE": "postgres",
"DATABASE_HOST": "db.internal",
"DATABASE_PORT": "5433",
"DATABASE_NAME": "cards",
"DATABASE_USER": "scopa",
"DATABASE_PASSWORD": "s3cret",
})
self.assertEqual(
settings.database_url,
"postgres://scopa:s3cret@db.internal:5433/cards",
)
def test_options_are_appended_as_query_string(self):
settings = _settings({"DATABASE_OPTIONS": "ssl=require"})
self.assertEqual(
settings.database_url,
"postgres://tavolo:password@localhost/tavolo?ssl=require",
)
def test_options_leading_question_mark_is_stripped(self):
settings = _settings({"DATABASE_OPTIONS": "?ssl=require"})
self.assertEqual(
settings.database_url,
"postgres://tavolo:password@localhost/tavolo?ssl=require",
)
def test_credentials_are_percent_encoded(self):
settings = _settings({
"DATABASE_USER": "u@x",
"DATABASE_PASSWORD": "p@ss/word:1",
})
self.assertEqual(
settings.database_url,
"postgres://u%40x:p%40ss%2Fword%3A1@localhost/tavolo",
)
def test_database_url_takes_precedence_over_parts(self):
settings = _settings({
"DATABASE_URL": "sqlite://:memory:",
"DATABASE_HOST": "db.internal",
"DATABASE_PASSWORD": "ignored",
})
self.assertEqual(settings.database_url, "sqlite://:memory:")
def test_empty_database_url_falls_back_to_parts(self):
settings = _settings({"DATABASE_URL": ""})
self.assertEqual(
settings.database_url,
"postgres://tavolo:password@localhost/tavolo",
)
class CorsSettingsTests(unittest.TestCase):
def test_cors_disabled_by_default(self):
settings = _settings({})
self.assertIsNone(settings.cors_allow_origins)
self.assertIsNone(settings.cors_allow_origin_regex)
self.assertIsNone(settings.cors_allow_methods)
self.assertIsNone(settings.cors_allow_headers)
self.assertFalse(settings.cors_allow_credentials)
self.assertIsNone(settings.cors_expose_headers)
self.assertEqual(600, settings.cors_max_age)
def test_allow_origins_parses_comma_separated_list(self):
settings = _settings({
"CORS_ALLOW_ORIGINS": "https://a.example, https://b.example ,,https://c.example",
})
self.assertEqual(
("https://a.example", "https://b.example", "https://c.example"),
settings.cors_allow_origins,
)
def test_allow_origins_star_is_passed_through(self):
settings = _settings({"CORS_ALLOW_ORIGINS": "*"})
self.assertEqual(("*",), settings.cors_allow_origins)
def test_allow_origin_regex_is_passed_through(self):
settings = _settings({"CORS_ALLOW_ORIGIN_REGEX": r"https://.*\.example\.com"})
self.assertEqual(r"https://.*\.example\.com", settings.cors_allow_origin_regex)
def test_allow_methods_and_headers_parse_as_lists(self):
settings = _settings({
"CORS_ALLOW_METHODS": "GET,POST",
"CORS_ALLOW_HEADERS": "Authorization, X-Custom-Header",
"CORS_EXPOSE_HEADERS": "X-Total-Count",
})
self.assertEqual(("GET", "POST"), settings.cors_allow_methods)
self.assertEqual(("Authorization", "X-Custom-Header"), settings.cors_allow_headers)
self.assertEqual(("X-Total-Count",), settings.cors_expose_headers)
def test_allow_credentials_parses_boolean(self):
for value in ("1", "true", "TRUE", "yes", "on"):
self.assertTrue(_settings({"CORS_ALLOW_CREDENTIALS": value}).cors_allow_credentials)
for value in ("0", "false", "no", "off", "anything-else"):
self.assertFalse(_settings({"CORS_ALLOW_CREDENTIALS": value}).cors_allow_credentials)
def test_max_age_parses_int(self):
settings = _settings({"CORS_MAX_AGE": "3600"})
self.assertEqual(3600, settings.cors_max_age)
if __name__ == "__main__":
unittest.main()