"""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()