diff --git a/CHANGELOG.md b/CHANGELOG.md index 1f5ecc5..e4bcd31 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -40,6 +40,14 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2. ### Fixed +- `StateManager` no longer writes the raw OAuth state to its logs. The state + comes from the callback query string, so logging it leaked live tokens and let + a caller forge log lines with CR/LF. Log lines now carry `state_ref`, the + first 12 hex characters of the state's SHA-256. An unknown or expired state + in `consume_state()` is logged at INFO instead of WARNING, and malformed + stored data is logged once at WARNING without a traceback, instead of two + ERROR records with a traceback that echoed the stored state. + - Regenerating the session ID of the request's session (the documented defence against session fixation at login) now sends a token for the new ID. `FastAPICacheXSessionMiddleware` used to re-send the loaded token, which named diff --git a/docs/STATE.md b/docs/STATE.md index 0b08778..0a4a272 100644 --- a/docs/STATE.md +++ b/docs/STATE.md @@ -167,3 +167,7 @@ CacheXError sides. Override `get_and_delete()` in that case. - Do not store sensitive data in a state. `metadata` is stored as plain-text JSON in the cache backend. +- **Logs never contain the state itself.** Log lines from `fastapi_cachex.state.manager` + identify a state by `state_ref`, the first 12 hex characters of its SHA-256, which you can + compute from a known state to match it. An unknown or expired state is logged at INFO, + since it is routine client input; malformed stored data is logged once at WARNING. diff --git a/fastapi_cachex/state/manager.py b/fastapi_cachex/state/manager.py index 4885c19..fa77f8e 100644 --- a/fastapi_cachex/state/manager.py +++ b/fastapi_cachex/state/manager.py @@ -29,6 +29,23 @@ def _is_past(moment: datetime) -> bool: return datetime.now(timezone.utc) > moment +def _state_ref(state: str) -> str: + """Return a short digest that identifies a state in logs without revealing it. + + The state comes straight from the callback query string, so logging it raw + would leak live tokens and let a caller forge log lines with CR/LF. + """ + return hashlib.sha256(state.encode("utf-8", "surrogatepass")).hexdigest()[:12] + + +def _log_decode_failure(state: str) -> None: + # No traceback: pydantic validation errors echo the stored input, which + # includes the state itself. + logger.warning( + "Stored OAuth state data is malformed; state_ref=%s", _state_ref(state) + ) + + class StateManager: """Manages OAuth state and session state lifecycle and storage.""" @@ -56,12 +73,13 @@ def __init__( def _cache_key(self, state: str) -> str: return f"{self.key_prefix}{state}" - def _decode_state(self, cached: CacheEntry, state: str) -> StateData: + def _decode_state(self, cached: CacheEntry) -> StateData: """Turn a backend entry into a StateData model. + Nothing is logged here; the caller logs the failure once. + Args: cached: The CacheEntry retrieved from backend - state: State string (used for logging) Returns: StateData instance @@ -80,31 +98,29 @@ def _decode_state(self, cached: CacheEntry, state: str) -> StateData: state_dict: dict[str, Any] = json.loads(json_content) except json.JSONDecodeError as e: msg = f"Failed to parse state data: {e}" - logger.exception("Failed to parse state data; state=%s", state) raise StateDataError(msg) from e try: return StateData(**state_dict) except ValueError as e: msg = f"Invalid state data structure: {e}" - logger.exception("Failed to create StateData model; state=%s", state) raise StateDataError(msg) from e async def _peek_state(self, state: str) -> StateData | None: """Load a state without consuming it; None when missing, malformed or expired.""" cached = await self.backend.get(self._cache_key(state)) if cached is None: - logger.debug("State not found; state=%s", state) + logger.debug("State not found; state_ref=%s", _state_ref(state)) return None try: - state_data = self._decode_state(cached, state) + state_data = self._decode_state(cached) except StateDataError: - logger.exception("Failed to parse or validate state data; state=%s", state) + _log_decode_failure(state) return None if _is_past(state_data.expires_at): - logger.debug("State expired; state=%s", state) + logger.debug("State expired; state_ref=%s", _state_ref(state)) return None return state_data @@ -149,7 +165,9 @@ async def create_state( ) await self.backend.set(self._cache_key(state), entry, ttl=effective_ttl) - logger.debug("OAuth state created; state=%s ttl=%s", state, effective_ttl) + logger.debug( + "OAuth state created; state_ref=%s ttl=%s", _state_ref(state), effective_ttl + ) return state async def consume_state(self, state: str) -> StateData: @@ -172,18 +190,26 @@ async def consume_state(self, state: str) -> StateData: # out to be malformed or past its wall-clock expiry is gone as well. cached = await self.backend.get_and_delete(self._cache_key(state)) if cached is None: - logger.warning("OAuth state not found or expired; state=%s", state) + logger.info( + "OAuth state not found or expired; state_ref=%s", _state_ref(state) + ) msg = "Invalid or expired state" raise InvalidStateError(msg) - state_data = self._decode_state(cached, state) + try: + state_data = self._decode_state(cached) + except StateDataError: + _log_decode_failure(state) + raise if _is_past(state_data.expires_at): - logger.warning("OAuth state expired; state=%s", state) + logger.info("OAuth state expired; state_ref=%s", _state_ref(state)) msg = "State has expired" raise StateExpiredError(msg) - logger.debug("OAuth state consumed and deleted; state=%s", state) + logger.debug( + "OAuth state consumed and deleted; state_ref=%s", _state_ref(state) + ) return state_data async def validate_state(self, state: str) -> bool: @@ -219,7 +245,9 @@ async def delete_state(self, state: str) -> bool: True if state was deleted, False if it didn't exist """ if await self.backend.get_and_delete(self._cache_key(state)) is None: - logger.debug("OAuth state not found for deletion; state=%s", state) + logger.debug( + "OAuth state not found for deletion; state_ref=%s", _state_ref(state) + ) return False - logger.debug("OAuth state deleted; state=%s", state) + logger.debug("OAuth state deleted; state_ref=%s", _state_ref(state)) return True diff --git a/tests/state/test_manager.py b/tests/state/test_manager.py index 3815172..6a375fe 100644 --- a/tests/state/test_manager.py +++ b/tests/state/test_manager.py @@ -2,6 +2,7 @@ import hashlib import json +import logging from collections.abc import AsyncGenerator from datetime import datetime from datetime import timedelta @@ -16,6 +17,7 @@ from fastapi_cachex.proxy import BackendProxy from fastapi_cachex.state.exceptions import InvalidStateError from fastapi_cachex.state.exceptions import StateDataError +from fastapi_cachex.state.exceptions import StateExpiredError from fastapi_cachex.state.manager import StateManager from fastapi_cachex.state.models import StateData from fastapi_cachex.types import CacheEntry @@ -823,3 +825,92 @@ async def test_peeking_a_state_past_its_wall_clock_expiry_treats_it_as_gone( assert await state_manager.validate_state(state) is False assert await state_manager.get_state_metadata(state) is None + + +def _state_records(caplog: pytest.LogCaptureFixture) -> list[logging.LogRecord]: + return [r for r in caplog.records if r.name == "fastapi_cachex.state.manager"] + + +@pytest.mark.asyncio +async def test_logs_never_contain_the_raw_state( + memory_backend: MemoryBackend, caplog: pytest.LogCaptureFixture +) -> None: + """Every log line identifies a state by digest only, so CR/LF cannot forge lines.""" + caplog.set_level(logging.DEBUG, logger="fastapi_cachex.state.manager") + manager = StateManager(backend=memory_backend) + + state = await manager.create_state() + await manager.validate_state(state) + await manager.consume_state(state) + await manager.delete_state(state) + forged = "abc\r\nERROR forged entry" + with pytest.raises(InvalidStateError): + await manager.consume_state(forged) + await manager.validate_state(forged) + + records = _state_records(caplog) + assert records + assert state not in caplog.text + assert "forged entry" not in caplog.text + assert all("\n" not in r.getMessage() for r in records) + assert hashlib.sha256(state.encode()).hexdigest()[:12] in caplog.text + + +@pytest.mark.asyncio +async def test_unknown_state_is_not_logged_as_warning( + memory_backend: MemoryBackend, caplog: pytest.LogCaptureFixture +) -> None: + """A missing or expired state is routine client input, not an operator warning.""" + caplog.set_level(logging.DEBUG, logger="fastapi_cachex.state.manager") + manager = StateManager(backend=memory_backend) + + with pytest.raises(InvalidStateError): + await manager.consume_state("unknown") + + state = await manager.create_state(ttl=60) + stale = StateData( + state=state, expires_at=datetime.now(timezone.utc) - timedelta(seconds=1) + ) + content = stale.model_dump_json().encode() + await memory_backend.set( + f"{manager.key_prefix}{state}", + CacheEntry(fingerprint=hashlib.sha256(content).hexdigest(), content=content), + ttl=60, + ) + with pytest.raises(StateExpiredError): + await manager.consume_state(state) + + assert all(r.levelno < logging.WARNING for r in _state_records(caplog)) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["consume", "validate", "metadata"]) +async def test_malformed_state_data_is_logged_once_without_the_state( + memory_backend: MemoryBackend, + caplog: pytest.LogCaptureFixture, + operation: str, +) -> None: + """A decode failure produces one warning, with no traceback echoing the stored state.""" + caplog.set_level(logging.DEBUG, logger="fastapi_cachex.state.manager") + manager = StateManager(backend=memory_backend) + state = "secret-state-token" + content = json.dumps({"state": state, "expires_at": "not a date"}).encode() + await memory_backend.set( + f"{manager.key_prefix}{state}", + CacheEntry(fingerprint=hashlib.sha256(content).hexdigest(), content=content), + ttl=60, + ) + + if operation == "consume": + with pytest.raises(StateDataError): + await manager.consume_state(state) + elif operation == "validate": + assert await manager.validate_state(state) is False + else: + assert await manager.get_state_metadata(state) is None + + records = [r for r in _state_records(caplog) if r.levelno >= logging.WARNING] + assert len(records) == 1 + assert records[0].levelno == logging.WARNING + assert records[0].exc_info is None + assert state not in caplog.text