fix(security): encode JWT exp as NumericDate so PyJWT validates expiry (#1016)

JWTParams.exp was typed str and filled with an ISO-8601 string, copied
verbatim into the reserved RFC 7519 exp claim. PyJWT requires exp to be a
NumericDate and raised DecodeError (subclass of PyJWTError), which the
blanket handler turned into 'Invalid JWT' — so every token minted with
--expires or expires_at failed auth. The module's own ISO comparison was
dead code because decode always failed first.

- JWTParams.exp: datetime | None; create_jwt unchanged (PyJWT converts)
- verify_jwt: drop the dead ISO comparison; catch jwt.ExpiredSignatureError
  before PyJWTError so expired tokens report 'JWT expired'
- keys.py: pass expires_at through directly; drop unused format import
- generate_jwt.py: keep ISO string for display, datetime for the claim
- Tests: TestJWTExpiry (future/past/ISO-string/tampered) + rewritten
  test_create_key_with_expires_at verifying minted keys both ways

Signed-off-by: ChethanUK <chethanuk@outlook.com>
This commit is contained in:
ChethanUK 2026-08-16 18:15:14 +02:00
parent 444897975c
commit 2d6206f068
No known key found for this signature in database
5 changed files with 82 additions and 26 deletions

View File

@ -108,17 +108,16 @@ def main():
if (args.peer or args.session) and not args.workspace:
parser.error("--peer and --session require --workspace")
exp_str: str | None = None
if args.expires:
expiry: datetime.datetime | None = None
if args.expires is not None:
expiry = datetime.datetime.now(datetime.timezone.utc) + args.expires
exp_str = format_datetime_utc(expiry)
params = JWTParams(
ad=True if args.admin else None,
w=args.workspace,
p=args.peer,
s=args.session,
exp=exp_str,
exp=expiry,
)
token = create_jwt(params)
@ -137,8 +136,8 @@ def main():
scope_parts.append(f"session={args.session}")
print(f"Scope: {', '.join(scope_parts)}")
if exp_str:
print(f"Expires: {exp_str}")
if expiry:
print(f"Expires: {format_datetime_utc(expiry)}")
else:
print("Expires: never")
print(f"Token: {token}")

View File

@ -11,7 +11,6 @@ from src.security import (
require_auth,
scope_requires_workspace,
)
from src.utils.formatting import format_datetime_utc
logger = logging.getLogger(__name__)
@ -56,7 +55,7 @@ async def create_key(
key_str = create_jwt(
JWTParams(
exp=format_datetime_utc(expires_at) if expires_at else None,
exp=expires_at,
w=workspace_id,
p=peer_id,
s=session_id,

View File

@ -8,7 +8,7 @@ from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, Field
from src.config import settings
from src.utils.formatting import parse_datetime_iso, utc_now_iso
from src.utils.formatting import utc_now_iso
from .exceptions import AuthenticationException
@ -48,7 +48,7 @@ class JWTParams(BaseModel):
Fields (all optional other than `t`):
`t`: a string timestamp of when the JWT was created
`exp`: a string timestamp of when the JWT expires (optional)
`exp`: when the JWT expires (optional) a standard NumericDate claim
`ad`: a boolean flag indicating if the JWT is an admin JWT
`w`: (string) workspace name
`p`: (string) peer name
@ -56,7 +56,7 @@ class JWTParams(BaseModel):
"""
t: str = Field(default_factory=utc_now_iso)
exp: str | None = None
exp: datetime.datetime | None = None
ad: bool | None = None
w: str | None = None
p: str | None = None
@ -114,13 +114,6 @@ def verify_jwt(token: str) -> JWTParams:
)
if "t" in decoded:
params.t = decoded["t"]
if "exp" in decoded:
params.exp = decoded["exp"]
if params.exp:
exp_time = parse_datetime_iso(params.exp)
current_time = datetime.datetime.now(datetime.timezone.utc)
if exp_time < current_time:
raise AuthenticationException("JWT expired")
if "ad" in decoded:
params.ad = decoded["ad"]
# Normalize empty-string scope claims to None so a blank `w`/`p`/`s`
@ -141,6 +134,8 @@ def verify_jwt(token: str) -> JWTParams:
"Invalid JWT scope: peer/session token missing workspace"
)
return params
except jwt.ExpiredSignatureError:
raise AuthenticationException("JWT expired") from None
except jwt.PyJWTError:
raise AuthenticationException("Invalid JWT") from None

View File

@ -1,4 +1,10 @@
import datetime
import pytest
from src.exceptions import AuthenticationException
from src.models import Peer, Workspace
from src.security import verify_jwt
from tests.conftest import AuthClient
@ -59,17 +65,30 @@ def test_create_key_with_expires_at(
response = auth_client.post("/v3/keys", params={"expires_at": "2025-01-01"})
# Only admin JWT should be allowed
if auth_client.auth_type == "admin":
# key with no params should fail
assert response.status_code == 422
return
else:
if auth_client.auth_type != "admin":
assert response.status_code == 401
return
# key with no params should fail
assert response.status_code == 422
test_workspace, _ = sample_data
# assert that the key is expired
# Future expiry: mint succeeds and the token verifies (NumericDate exp, #1016)
future = datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(days=1)
response = auth_client.post(
"/v3/keys", params={"workspace_id": test_workspace.name}
"/v3/keys",
params={"workspace_id": test_workspace.name, "expires_at": future.isoformat()},
)
assert response.status_code == 401
assert response.status_code == 200
verify_jwt(response.json()["key"]) # must not raise "Invalid JWT"
# Past expiry: mint succeeds but verification reports expired
past = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta(days=1)
response = auth_client.post(
"/v3/keys",
params={"workspace_id": test_workspace.name, "expires_at": past.isoformat()},
)
assert response.status_code == 200
with pytest.raises(AuthenticationException, match="JWT expired"):
verify_jwt(response.json()["key"])

View File

@ -5,6 +5,7 @@ through to a workspace check, so a `{w, p}` token authorized any peer in `w`.
The contract now is: authorize by the token's narrowest claim, never widen.
"""
import datetime
from contextlib import asynccontextmanager
import jwt as pyjwt
@ -285,3 +286,46 @@ class TestAuthAdminAndUnscoped:
creds = _bearer(create_jwt(JWTParams()))
with pytest.raises(AuthenticationException):
await auth(credentials=creds, workspace_name="ws-a")
SCOPES = [
{"ad": True},
{"w": "ws-a"},
{"w": "ws-a", "p": "alice"},
{"w": "ws-a", "s": "sess-1"},
]
class TestJWTExpiry:
"""#1016: exp was an ISO string in the reserved NumericDate claim."""
def _exp(self, **delta):
return datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(
**delta
)
@pytest.mark.parametrize("scope", SCOPES)
def test_future_expiry_verifies(self, scope):
token = create_jwt(JWTParams(exp=self._exp(days=30), **scope))
params = verify_jwt(token)
assert (params.ad, params.w, params.p, params.s) == (
scope.get("ad"),
scope.get("w"),
scope.get("p"),
scope.get("s"),
)
@pytest.mark.parametrize("scope", [{"ad": True}, {"w": "ws-a"}])
def test_past_expiry_reports_expired(self, scope):
token = create_jwt(JWTParams(exp=self._exp(days=-1), **scope))
with pytest.raises(AuthenticationException, match="JWT expired"):
verify_jwt(token)
def test_iso_string_exp_still_accepted_by_the_model(self):
token = create_jwt(JWTParams(ad=True, exp="2099-01-01T00:00:00Z"))
assert verify_jwt(token).ad is True
def test_tampered_token_still_reports_invalid(self):
token = create_jwt(JWTParams(ad=True, exp=self._exp(days=30)))
with pytest.raises(AuthenticationException, match="Invalid JWT"):
verify_jwt(token + "x")