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
8 changes: 8 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 4 additions & 0 deletions docs/STATE.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
58 changes: 43 additions & 15 deletions fastapi_cachex/state/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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
91 changes: 91 additions & 0 deletions tests/state/test_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import hashlib
import json
import logging
from collections.abc import AsyncGenerator
from datetime import datetime
from datetime import timedelta
Expand All @@ -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
Expand Down Expand Up @@ -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
Loading