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
11 changes: 11 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,17 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2.

### Fixed

- 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
the record `regenerate_session_id()` had just deleted, so the user was logged
straight back out. Header clients got no token at all. The deprecated
`SessionMiddleware` could overwrite the new token with a sliding-renewed one
for the old ID. Both middlewares now notice the changed ID and send the new
token through the request's transport. `SessionManager.issue_token(session)`
is the one place tokens are signed.
([#103](https://github.com/allen0099/FastAPI-CacheX/issues/103))

- `ttl` means the same thing on every backend. Zero or negative TTLs now raise
`ValueError` from `set`, `set_if_absent` and `increment` on all built-in
backends, from the base-class fallbacks, and from `CacheManager` and
Expand Down
28 changes: 23 additions & 5 deletions docs/SESSION.md
Original file line number Diff line number Diff line change
Expand Up @@ -598,14 +598,32 @@ match the bound one (or has no address) is treated as having no session.
Prevents session fixation attacks:

```python
# After a successful login
session, _renewed_token = await session_manager.get_session(current_token)
session, new_token = await session_manager.regenerate_session_id(session)
# Hand new_token to the client; the old token no longer resolves to a session.
from fastapi_cachex.session.dependencies import SessionDep, SessionManagerDep


@app.post("/login")
async def login(request: Request, session: SessionDep, manager: SessionManagerDep):
... # verify the credentials
await manager.regenerate_session_id(session)
request.session["user_id"] = "123"
return {"ok": True}
```

`regenerate_session_id()` deletes the backend record under the old ID and saves the session under
a new ID, keeping its data, user, `created_at` and expiry.
a new ID, keeping its data, user, `created_at` and expiry. When the session is the request's own
(from `SessionDep`, `get_session` and friends), either middleware sees the new ID and sends a token
for it through the transport the request used: `Set-Cookie` for a cookie, the response header for a
header token. After that the old token no longer resolves to a session.

Outside a middleware, load the session with the same bindings the middleware would pass, and hand
the returned token to the client yourself:

```python
session, _ = await manager.get_session(
current_token, ip_address=client_ip, user_agent=user_agent
)
session, new_token = await manager.regenerate_session_id(session)
```

## SessionManager at a glance

Expand Down
36 changes: 26 additions & 10 deletions fastapi_cachex/session/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,8 +153,7 @@ async def _create_session(
# Store in backend
await self._save_session(session)

# Generate signed token (pass expires_at so JWT exp reflects sliding expiration)
token = self._create_token(session.session_id, expires_at=session.expires_at)
token = self.issue_token(session)
logger.debug(
"Session created; id=%s ttl=%s ip=%s ua=%s",
session.session_id,
Expand All @@ -163,7 +162,7 @@ async def _create_session(
session.user_agent,
)

return session, self._serializer.to_string(token)
return session, token

async def get_session(
self,
Expand Down Expand Up @@ -276,10 +275,7 @@ async def get_session(

if time_remaining < threshold:
session.renew(self.config.session_ttl)
token = self._create_token(
session.session_id, expires_at=session.expires_at
)
renewed_token = self._serializer.to_string(token)
renewed_token = self.issue_token(session)
logger.debug(
"Session renewed (sliding expiration); id=%s ttl=%s",
session.session_id,
Expand Down Expand Up @@ -326,6 +322,11 @@ async def regenerate_session_id(
) -> tuple[Session, str]:
"""Regenerate session ID (after login for security).

``session`` is changed in place. When it is the request's session under
either middleware, the middleware notices the new ID and sends the new
token through the request's transport, so a handler that only needs
the cookie or header updated can ignore the returned token.

Args:
session: Session to regenerate

Expand All @@ -342,13 +343,28 @@ async def regenerate_session_id(
# Save with new ID
await self._save_session(session)

# Create new token (pass expires_at so JWT exp reflects current expiry)
token = self._create_token(session.session_id, expires_at=session.expires_at)
logger.debug(
"Session ID regenerated; old_id=%s new_id=%s", old_id, session.session_id
)

return session, self._serializer.to_string(token)
return session, self.issue_token(session)

def issue_token(self, session: Session) -> str:
"""Sign a token string for ``session``'s current ID and expiry.

This is the token ``create_session``, sliding renewal and
``regenerate_session_id`` hand out; the middleware also uses it to
send a fresh token once a handler has regenerated the session ID.

Args:
session: Session to issue a token for

Returns:
Serialized session token
"""
# expires_at is passed so a JWT's exp matches the session's expiry.
token = self._create_token(session.session_id, expires_at=session.expires_at)
return self._serializer.to_string(token)

async def delete_user_sessions(self, user_id: str) -> int:
"""Delete all sessions for a user.
Expand Down
50 changes: 42 additions & 8 deletions fastapi_cachex/session/middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -228,11 +228,19 @@ async def dispatch(
# Store session in request state
setattr(request.state, "__fastapi_cachex_session", session)

loaded_session_id = session.session_id if session is not None else None

# Process request
response: Response = await call_next(request)

# Propagate renewed token to client so its JWT exp stays in sync
if renewed_token is not None:
if session is not None and session.session_id != loaded_session_id:
# The handler regenerated the session ID; a renewed token would
# name the deleted record, so send a token for the new ID.
response.headers[self.config.header_name] = (
self.session_manager.issue_token(session)
)
elif renewed_token is not None:
# Propagate renewed token to client so its JWT exp stays in sync
response.headers[self.config.header_name] = renewed_token

return response
Expand Down Expand Up @@ -321,6 +329,7 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
loaded_token: str | None = None
renewed_token: str | None = None
backend_session: Session | None = None
loaded_session_id: str | None = None

# Resolve the incoming session token: prefer the header/bearer transport
# (e.g. X-Session-Token, as used by SessionMiddleware) and fall back to
Expand All @@ -340,6 +349,7 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
user_agent=user_agent,
)
loaded_token = renewed_token or token_value
loaded_session_id = backend_session.session_id
scope["session"] = StarletteSession(backend_session.data)
initial_session_was_empty = not backend_session.data
except SessionError:
Expand All @@ -363,6 +373,10 @@ async def send_wrapper(message: Message) -> None:
session: StarletteSession = scope["session"]
headers = MutableHeaders(scope=message)

current_token, fresh_token = self._response_tokens(
backend_session, loaded_session_id, loaded_token, renewed_token
)

if session.accessed:
headers.add_vary_header("Cookie")

Expand All @@ -371,8 +385,8 @@ async def send_wrapper(message: Message) -> None:
session,
connection,
backend_session,
loaded_token,
renewed_token,
current_token,
fresh_token,
)
# Header clients only need a genuinely new/renewed token (an
# unchanged one is already held); cookie clients always get a
Expand All @@ -394,15 +408,35 @@ async def send_wrapper(message: Message) -> None:
# Cookie transport: expire the cookie. A header-based client
# simply drops its now-dangling token (record is deleted).
headers.append("Set-Cookie", self._build_clear_cookie_header())
elif renewed_token is not None:
# Sliding expiration renewed the token even though the dict
# itself was untouched; propagate it via the same transport.
self._emit_token(headers, renewed_token, from_header=from_header)
elif fresh_token is not None:
# Sliding expiration renewed the token, or the ID was
# regenerated, even though the dict itself was untouched;
# propagate it via the same transport.
self._emit_token(headers, fresh_token, from_header=from_header)

await send(message)

await self.app(scope, receive, send_wrapper)

def _response_tokens(
self,
backend_session: "Session | None",
loaded_session_id: str | None,
loaded_token: str | None,
renewed_token: str | None,
) -> tuple[str | None, str | None]:
"""The ``(current, fresh)`` tokens to answer with.

Normally the loaded token and the sliding-renewed one (if any). If the
handler regenerated the session ID (at login, say), the loaded token
names a deleted record, so both become a token for the new ID and the
client gets it on either transport.
"""
if backend_session is None or backend_session.session_id == loaded_session_id:
return loaded_token, renewed_token
token = self.session_manager.issue_token(backend_session)
return token, token

def _emit_token(
self, headers: MutableHeaders, token: str, *, from_header: bool
) -> None:
Expand Down
110 changes: 110 additions & 0 deletions tests/session/test_starlette_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -726,3 +726,113 @@ async def logout_route(request: Request) -> dict[str, bool]:
assert "set-cookie" not in response.headers
with pytest.raises(SessionNotFoundError):
await manager.get_session(token)


def _regenerating_app(
manager: SessionManager,
config: SessionConfig,
*,
write_data: bool,
middleware: Any = FastAPICacheXSessionMiddleware,
) -> FastAPI:
"""An app whose /login regenerates the request's session ID, as docs advise."""
app = FastAPI()
app.add_middleware(middleware, session_manager=manager, config=config)

@app.post("/login")
async def login(request: Request, session=Depends(get_session)):
await manager.regenerate_session_id(session)
if write_data:
request.session["logged_in"] = True
return {"ok": True}

return app


async def _shorten_expiry(manager: SessionManager, token: str) -> None:
"""Push the session under the sliding threshold so the load renews it."""
session, _ = await manager.get_session(token)
session.expires_at = datetime.now(timezone.utc) + timedelta(seconds=100)
await manager._save_session(session)


@pytest.mark.parametrize("write_data", [True, False])
@pytest.mark.parametrize("sliding", [True, False])
@pytest.mark.asyncio
async def test_regenerated_session_id_is_sent_as_cookie(
manager: SessionManager, config: SessionConfig, write_data: bool, sliding: bool
) -> None:
"""After regenerate_session_id() the cookie must carry a token for the new ID.

The middleware used to re-send the loaded token, which named the deleted
record, so logging in with the documented fixation defence logged the user
straight back out.
"""
_session, old_token = await manager.create_session(
user=SessionUser(user_id="u"), a=1
)
if sliding:
await _shorten_expiry(manager, old_token)
client = TestClient(_regenerating_app(manager, config, write_data=write_data))
client.cookies.set(config.cookie_name, old_token)

response = client.post("/login")

assert response.status_code == 200
new_token = _extract_cookie_token(
response.headers["set-cookie"], config.cookie_name
)
assert new_token != old_token
session, _ = await manager.get_session(new_token)
expected = {"a": 1, "logged_in": True} if write_data else {"a": 1}
assert session.data == expected
with pytest.raises(SessionNotFoundError):
await manager.get_session(old_token)


@pytest.mark.parametrize("write_data", [True, False])
@pytest.mark.parametrize("sliding", [True, False])
@pytest.mark.asyncio
async def test_regenerated_session_id_is_sent_in_the_header(
manager: SessionManager, config: SessionConfig, write_data: bool, sliding: bool
) -> None:
"""A header client gets the new ID's token back in the response header."""
_session, old_token = await manager.create_session(user=SessionUser(user_id="u"))
if sliding:
await _shorten_expiry(manager, old_token)
client = TestClient(_regenerating_app(manager, config, write_data=write_data))

response = client.post("/login", headers={config.header_name: old_token})

assert response.status_code == 200
assert "set-cookie" not in response.headers
new_token = response.headers[config.header_name]
assert new_token != old_token
await manager.get_session(new_token)
with pytest.raises(SessionNotFoundError):
await manager.get_session(old_token)


@pytest.mark.filterwarnings("ignore::DeprecationWarning")
@pytest.mark.parametrize("sliding", [True, False])
@pytest.mark.asyncio
async def test_deprecated_middleware_sends_regenerated_token(
manager: SessionManager, config: SessionConfig, sliding: bool
) -> None:
"""SessionMiddleware must not overwrite the new ID's token with a renewed old one."""
_session, old_token = await manager.create_session(user=SessionUser(user_id="u"))
if sliding:
await _shorten_expiry(manager, old_token)
client = TestClient(
_regenerating_app(
manager, config, write_data=False, middleware=SessionMiddleware
)
)

response = client.post("/login", headers={config.header_name: old_token})

assert response.status_code == 200
new_token = response.headers[config.header_name]
await manager.get_session(new_token)
with pytest.raises(SessionNotFoundError):
await manager.get_session(old_token)
Loading