feat(auth): JWT TokenPayload accepts wid + role

TokenPayload gains optional wid (workspace_id) and role claims;
create_access_token accepts them as keyword-only args and only
encodes them when provided. Existing tokens and existing callers
keep working unchanged — the contract that protected requests
must carry wid is enforced by middleware (PR4 T4.6/T4.7), not by
the JWT type system.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
1445043649
2026-05-12 22:17:55 +08:00
parent 54762f491c
commit 54cb94c30f
2 changed files with 81 additions and 7 deletions
+29 -7
View File
@@ -1,6 +1,7 @@
"""JWT token creation and verification.""" """JWT token creation and verification."""
from datetime import UTC, datetime, timedelta from datetime import UTC, datetime, timedelta
from typing import Any
import jwt import jwt
from pydantic import BaseModel from pydantic import BaseModel
@@ -10,30 +11,51 @@ from app.gateway.auth.errors import TokenError
class TokenPayload(BaseModel): class TokenPayload(BaseModel):
"""JWT token payload.""" """JWT token payload.
`wid` / `role` were added in Stage 0 PR4 and are optional at the
model level so legacy 4-field tokens still parse into a value —
callers (middleware, decode_token) decide what "missing wid" means.
Post-PR4 production tokens always carry both fields.
"""
sub: str # user_id sub: str # user_id
wid: str | None = None # workspace_id (Stage 0 PR4)
role: str | None = None # owner / admin / member (Stage 0 PR4)
exp: datetime exp: datetime
iat: datetime | None = None iat: datetime | None = None
ver: int = 0 # token_version — must match User.token_version ver: int = 0 # token_version — must match User.token_version
def create_access_token(user_id: str, expires_delta: timedelta | None = None, token_version: int = 0) -> str: def create_access_token(
user_id: str,
expires_delta: timedelta | None = None,
token_version: int = 0,
*,
workspace_id: str | None = None,
role: str | None = None,
) -> str:
"""Create a JWT access token. """Create a JWT access token.
Args: Args:
user_id: The user's UUID as string user_id: The user's UUID as string.
expires_delta: Optional custom expiry, defaults to 7 days expires_delta: Optional custom expiry, defaults to 7 days.
token_version: User's current token_version for invalidation token_version: User's current token_version for invalidation.
workspace_id: Optional active workspace id; encoded as the ``wid`` claim.
role: Optional workspace role; encoded as the ``role`` claim.
Returns: Returns:
Encoded JWT string Encoded JWT string.
""" """
config = get_auth_config() config = get_auth_config()
expiry = expires_delta or timedelta(days=config.token_expiry_days) expiry = expires_delta or timedelta(days=config.token_expiry_days)
now = datetime.now(UTC) now = datetime.now(UTC)
payload = {"sub": user_id, "exp": now + expiry, "iat": now, "ver": token_version} payload: dict[str, Any] = {"sub": user_id, "exp": now + expiry, "iat": now, "ver": token_version}
if workspace_id is not None:
payload["wid"] = workspace_id
if role is not None:
payload["role"] = role
return jwt.encode(payload, config.jwt_secret, algorithm="HS256") return jwt.encode(payload, config.jwt_secret, algorithm="HS256")
+52
View File
@@ -0,0 +1,52 @@
"""JWT carries wid + role claims (Stage 0 PR4).
Tokens issued after PR4 must include `wid` (workspace_id) and `role`
(owner/admin/member) so the AuthMiddleware can resolve the active
workspace without a DB hit. Legacy token compatibility lives in
:mod:`test_legacy_token_compat`.
"""
from __future__ import annotations
from uuid import uuid4
import jwt
import pytest
from app.gateway.auth import create_access_token, decode_token
from app.gateway.auth.config import get_auth_config
@pytest.fixture(autouse=True)
def _stable_jwt_secret(monkeypatch):
"""Pin a deterministic JWT secret across tests in this module."""
monkeypatch.setenv("AUTH_JWT_SECRET", "test-secret-key-for-jwt-testing-minimum-32-chars")
yield
def test_jwt_includes_wid_and_role() -> None:
"""create_access_token records wid + role in the encoded payload."""
user_id = str(uuid4())
workspace_id = str(uuid4())
token = create_access_token(user_id, workspace_id=workspace_id, role="owner")
raw = jwt.decode(token, get_auth_config().jwt_secret, algorithms=["HS256"])
assert raw["sub"] == user_id
assert raw["wid"] == workspace_id
assert raw["role"] == "owner"
def test_decode_round_trip_keeps_wid_and_role() -> None:
"""decode_token returns a TokenPayload exposing wid + role attributes."""
user_id = str(uuid4())
workspace_id = str(uuid4())
token = create_access_token(user_id, workspace_id=workspace_id, role="member")
payload = decode_token(token)
# decode_token returns TokenError on failure — must be the success branch here.
assert hasattr(payload, "wid"), f"got {payload!r}"
assert payload.sub == user_id
assert payload.wid == workspace_id
assert payload.role == "member"