from __future__ import annotations import logging import unittest from .access_logging import ( UvicornAccessLogRedactionFilter, install_uvicorn_access_log_redaction, redact_access_log_target, ) UVICORN_ACCESS_FORMAT = '%s - "%s %s HTTP/%s" %d' def _uvicorn_access_record(target: object) -> logging.LogRecord: return logging.LogRecord( name="uvicorn.access", level=logging.INFO, pathname=__file__, lineno=1, msg=UVICORN_ACCESS_FORMAT, args=(("127.0.0.1", 12345), "GET", target, "1.1", 302), exc_info=None, ) class AccessLogRedactionTests(unittest.TestCase): def test_oauth_callback_query_is_removed_from_uvicorn_access_record(self) -> None: record = _uvicorn_access_record( "/auth/callback?code=oauth-code-secret&state=oauth-state-secret", ) self.assertTrue(UvicornAccessLogRedactionFilter().filter(record)) rendered = record.getMessage() self.assertIn('GET /auth/callback HTTP/1.1" 302', rendered) self.assertNotIn("oauth-code-secret", rendered) self.assertNotIn("oauth-state-secret", rendered) self.assertNotIn("?", rendered) def test_callback_redirect_variant_also_redacts_query(self) -> None: target = "/auth/callback/?code=oauth-code-secret&state=oauth-state-secret" self.assertEqual("/auth/callback/", redact_access_log_target(target)) def test_non_callback_query_is_not_changed(self) -> None: target = "/auth/login?provider=google&next=%2Flearn" record = _uvicorn_access_record(target) UvicornAccessLogRedactionFilter().filter(record) self.assertEqual(target, record.args[2]) self.assertEqual(target, redact_access_log_target(target)) def test_non_string_uvicorn_target_is_preserved(self) -> None: marker = object() record = _uvicorn_access_record(marker) UvicornAccessLogRedactionFilter().filter(record) self.assertIs(marker, record.args[2]) def test_install_is_idempotent(self) -> None: logger = logging.getLogger("uvicorn.access") original_filters = list(logger.filters) try: logger.filters = [ existing for existing in logger.filters if not isinstance(existing, UvicornAccessLogRedactionFilter) ] install_uvicorn_access_log_redaction() install_uvicorn_access_log_redaction() installed = [ existing for existing in logger.filters if isinstance(existing, UvicornAccessLogRedactionFilter) ] self.assertEqual(1, len(installed)) finally: logger.filters = original_filters def test_installed_filter_redacts_an_emitted_uvicorn_access_message(self) -> None: logger = logging.getLogger("uvicorn.access") original_filters = list(logger.filters) original_handlers = list(logger.handlers) original_level = logger.level original_propagate = logger.propagate rendered: list[str] = [] class CaptureHandler(logging.Handler): def emit(self, record: logging.LogRecord) -> None: rendered.append(record.getMessage()) try: logger.filters = [] logger.handlers = [CaptureHandler()] logger.setLevel(logging.INFO) logger.propagate = False install_uvicorn_access_log_redaction() logger.info( UVICORN_ACCESS_FORMAT, ("127.0.0.1", 12345), "GET", "/auth/callback?code=oauth-code-secret&state=oauth-state-secret", "1.1", 302, ) self.assertEqual(1, len(rendered)) self.assertIn("GET /auth/callback HTTP/1.1", rendered[0]) self.assertNotIn("oauth-code-secret", rendered[0]) self.assertNotIn("oauth-state-secret", rendered[0]) finally: logger.filters = original_filters logger.handlers = original_handlers logger.setLevel(original_level) logger.propagate = original_propagate if __name__ == "__main__": unittest.main()