Skip to content

Commit ed3d934

Browse files
authored
feat(core): strict UTC datetime validation for pydantic models (#477)
* feat: add strict UTC datetime validation to pydantic models Introduce UTCDatetime type that rejects naive and non-UTC datetimes at the pydantic validation boundary. Applied to all datetime fields in job and auth models. * test: add utc datetime validation enforcement to all pydantic models * refactor: convert InsertedJob and TokenPayload from TypedDict to BaseModel TypedDict fields are not validated by pydantic at runtime, so UTCDatetime annotations had no effect. Converting to BaseModel ensures datetime fields are actually validated on construction. * refactor: separate token creation from raw JWT signing Keep create_token strictly typed (TokenPayload only). Extract _sign_token_payload for tests that craft malformed JWTs from raw dicts (expired tokens, bad signatures, key testing).
1 parent bbca4b9 commit ed3d934

10 files changed

Lines changed: 217 additions & 72 deletions

File tree

diracx-core/src/diracx/core/models/auth.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,12 @@
11
from __future__ import annotations
22

3-
from datetime import datetime
43
from enum import StrEnum
54

65
from pydantic import BaseModel
76
from typing_extensions import TypedDict
87

8+
from .types import UTCDatetime
9+
910

1011
class UserInfo(BaseModel):
1112
sub: str # dirac generated vo:sub
@@ -55,9 +56,9 @@ class OpenIDConfiguration(TypedDict):
5556
code_challenge_methods_supported: list[str]
5657

5758

58-
class TokenPayload(TypedDict):
59+
class TokenPayload(BaseModel):
5960
jti: str
60-
exp: datetime
61+
exp: UTCDatetime
6162
dirac_policies: dict
6263

6364

diracx-core/src/diracx/core/models/job.py

Lines changed: 17 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -5,19 +5,19 @@
55

66
from __future__ import annotations
77

8-
from datetime import datetime
98
from enum import StrEnum
109
from typing import Literal
1110

1211
from pydantic import BaseModel, Field, field_validator
13-
from typing_extensions import TypedDict
1412

13+
from .types import UTCDatetime
1514

16-
class InsertedJob(TypedDict):
15+
16+
class InsertedJob(BaseModel):
1717
JobID: int
1818
Status: str
1919
MinorStatus: str
20-
TimeStamp: datetime
20+
TimeStamp: UTCDatetime
2121

2222

2323
class HeartbeatData(BaseModel, extra="forbid"):
@@ -39,7 +39,7 @@ class JobCommand(BaseModel):
3939
class JobParameters(BaseModel, populate_by_name=True, extra="allow"):
4040
"""Some of the most important parameters that can be set for a job."""
4141

42-
timestamp: datetime | None = None
42+
timestamp: UTCDatetime | None = None
4343
cpu_normalization_factor: int | None = Field(None, alias="CPUNormalizationFactor")
4444
norm_cpu_time_s: int | None = Field(None, alias="NormCPUTime(s)")
4545
total_cpu_time_s: int | None = Field(None, alias="TotalCPUTime(s)")
@@ -84,12 +84,12 @@ class JobAttributes(BaseModel, populate_by_name=True, extra="forbid"):
8484
owner: str | None = Field(None, alias="Owner")
8585
owner_group: str | None = Field(None, alias="OwnerGroup")
8686
vo: str | None = Field(None, alias="VO")
87-
submission_time: datetime | None = Field(None, alias="SubmissionTime")
88-
reschedule_time: datetime | None = Field(None, alias="RescheduleTime")
89-
last_update_time: datetime | None = Field(None, alias="LastUpdateTime")
90-
start_exec_time: datetime | None = Field(None, alias="StartExecTime")
91-
heart_beat_time: datetime | None = Field(None, alias="HeartBeatTime")
92-
end_exec_time: datetime | None = Field(None, alias="EndExecTime")
87+
submission_time: UTCDatetime | None = Field(None, alias="SubmissionTime")
88+
reschedule_time: UTCDatetime | None = Field(None, alias="RescheduleTime")
89+
last_update_time: UTCDatetime | None = Field(None, alias="LastUpdateTime")
90+
start_exec_time: UTCDatetime | None = Field(None, alias="StartExecTime")
91+
heart_beat_time: UTCDatetime | None = Field(None, alias="HeartBeatTime")
92+
end_exec_time: UTCDatetime | None = Field(None, alias="EndExecTime")
9393
status: str | None = Field(None, alias="Status")
9494
minor_status: str | None = Field(None, alias="MinorStatus")
9595
application_status: str | None = Field(None, alias="ApplicationStatus")
@@ -131,7 +131,7 @@ class JobLoggingRecord(BaseModel):
131131
status: JobStatus | Literal["idem"]
132132
minor_status: str
133133
application_status: str
134-
date: datetime
134+
date: UTCDatetime
135135
source: str
136136

137137

@@ -149,7 +149,7 @@ class LimitedJobStatusReturn(BaseModel):
149149

150150

151151
class JobStatusReturn(LimitedJobStatusReturn):
152-
StatusTime: datetime
152+
StatusTime: UTCDatetime
153153
Source: str
154154

155155

@@ -160,10 +160,10 @@ class SetJobStatusReturnSuccess(BaseModel):
160160
Status: JobStatus | None = None
161161
MinorStatus: str | None = None
162162
ApplicationStatus: str | None = None
163-
HeartBeatTime: datetime | None = None
164-
StartExecTime: datetime | None = None
165-
EndExecTime: datetime | None = None
166-
LastUpdateTime: datetime | None = None
163+
HeartBeatTime: UTCDatetime | None = None
164+
StartExecTime: UTCDatetime | None = None
165+
EndExecTime: UTCDatetime | None = None
166+
LastUpdateTime: UTCDatetime | None = None
167167

168168
success: dict[int, SetJobStatusReturnSuccess]
169169
failed: dict[int, dict[str, str]]
Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,21 @@
1+
"""Custom types for DiracX pydantic models."""
2+
3+
from __future__ import annotations
4+
5+
from datetime import UTC, datetime, timedelta
6+
7+
from pydantic import AfterValidator, AwareDatetime
8+
from typing_extensions import Annotated
9+
10+
11+
def _validate_utc(v: datetime) -> datetime:
12+
"""Reject aware datetimes that are not in UTC.
13+
14+
AwareDatetime already rejects naive datetimes before this runs.
15+
"""
16+
if v.utcoffset() != timedelta(0):
17+
raise ValueError(f"Datetime must be in UTC, got offset {v.utcoffset()}")
18+
return v.replace(tzinfo=UTC)
19+
20+
21+
UTCDatetime = Annotated[AwareDatetime, AfterValidator(_validate_utc)]
Lines changed: 118 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,118 @@
1+
"""Tests for UTCDatetime pydantic type validation."""
2+
3+
from __future__ import annotations
4+
5+
import importlib
6+
import inspect
7+
import pkgutil
8+
from datetime import UTC, datetime, timedelta, timezone
9+
10+
import pytest
11+
from pydantic import BaseModel, ValidationError
12+
13+
import diracx.core.models
14+
from diracx.core.models.types import UTCDatetime, _validate_utc
15+
16+
17+
class SampleModel(BaseModel):
18+
ts: UTCDatetime
19+
optional_ts: UTCDatetime | None = None
20+
21+
22+
class TestUTCDatetimeAcceptsUTC:
23+
def test_utc_timezone(self):
24+
dt = datetime(2024, 1, 1, 12, 0, 0, tzinfo=UTC)
25+
m = SampleModel(ts=dt)
26+
assert m.ts == dt
27+
assert m.ts.tzinfo is UTC
28+
29+
def test_timezone_utc(self):
30+
dt = datetime(2024, 1, 1, 12, 0, 0, tzinfo=timezone.utc)
31+
m = SampleModel(ts=dt)
32+
assert m.ts.utcoffset() == timedelta(0)
33+
34+
def test_iso_string_utc(self):
35+
m = SampleModel(ts="2024-01-01T12:00:00Z")
36+
assert m.ts.tzinfo is UTC
37+
38+
def test_iso_string_plus_zero(self):
39+
m = SampleModel(ts="2024-01-01T12:00:00+00:00")
40+
assert m.ts.utcoffset() == timedelta(0)
41+
42+
def test_optional_none(self):
43+
m = SampleModel(ts="2024-01-01T12:00:00Z", optional_ts=None)
44+
assert m.optional_ts is None
45+
46+
47+
class TestUTCDatetimeRejectsNonUTC:
48+
def test_naive_datetime(self):
49+
dt = datetime(2024, 1, 1, 12, 0, 0) # noqa: DTZ001
50+
with pytest.raises(ValidationError, match="timezone"):
51+
SampleModel(ts=dt)
52+
53+
def test_non_utc_timezone(self):
54+
cet = timezone(timedelta(hours=1))
55+
dt = datetime(2024, 1, 1, 12, 0, 0, tzinfo=cet)
56+
with pytest.raises(ValidationError, match="must be in UTC"):
57+
SampleModel(ts=dt)
58+
59+
def test_iso_string_non_utc(self):
60+
with pytest.raises(ValidationError, match="must be in UTC"):
61+
SampleModel(ts="2024-01-01T12:00:00+05:30")
62+
63+
def test_naive_iso_string(self):
64+
with pytest.raises(ValidationError):
65+
SampleModel(ts="2024-01-01T12:00:00")
66+
67+
68+
def _is_datetime_type(annotation: type) -> bool:
69+
"""Check if an annotation is datetime or a subclass of datetime."""
70+
try:
71+
return isinstance(annotation, type) and issubclass(annotation, datetime)
72+
except TypeError:
73+
return False
74+
75+
76+
def _collect_model_classes() -> list[type[BaseModel]]:
77+
"""Discover all BaseModel subclasses in diracx.core.models."""
78+
models = []
79+
package = diracx.core.models
80+
for _importer, modname, _ispkg in pkgutil.walk_packages(
81+
package.__path__, prefix=package.__name__ + "."
82+
):
83+
if modname.endswith(".types"):
84+
continue
85+
module = importlib.import_module(modname)
86+
for _name, obj in inspect.getmembers(module, inspect.isclass):
87+
if (
88+
issubclass(obj, BaseModel)
89+
and obj is not BaseModel
90+
and obj.__module__ == modname
91+
):
92+
models.append(obj)
93+
return models
94+
95+
96+
def _check_field_uses_utc_validator(model: type[BaseModel], field_name: str) -> bool:
97+
"""Check that a datetime field has the _validate_utc AfterValidator."""
98+
field_info = model.model_fields[field_name]
99+
return any(getattr(m, "func", None) is _validate_utc for m in field_info.metadata)
100+
101+
102+
def test_all_datetime_fields_use_utc_datetime():
103+
"""Ensure no pydantic model in diracx.core.models uses bare datetime.
104+
105+
Every datetime field must use UTCDatetime to enforce UTC validation.
106+
"""
107+
violations = []
108+
for model in _collect_model_classes():
109+
for field_name, field_info in model.model_fields.items():
110+
if not _is_datetime_type(field_info.annotation):
111+
continue
112+
if not _check_field_uses_utc_validator(model, field_name):
113+
violations.append(f"{model.__name__}.{field_name}")
114+
115+
assert not violations, (
116+
"The following fields use bare datetime instead of UTCDatetime:\n"
117+
+ "\n".join(f" - {v}" for v in violations)
118+
)

diracx-logic/src/diracx/logic/auth/token.py

Lines changed: 23 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -332,14 +332,14 @@ async def exchange_token(
332332
refresh_exp = uuid7_to_datetime(refresh_jti) + timedelta(
333333
minutes=refresh_token_expire_minutes
334334
)
335-
refresh_payload = {
336-
"jti": str(refresh_jti),
337-
"exp": refresh_exp,
335+
refresh_payload = RefreshTokenPayload(
336+
jti=str(refresh_jti),
337+
exp=refresh_exp,
338338
# legacy_exchange is used to indicate that the original refresh token
339339
# was obtained from the legacy_exchange endpoint
340-
"legacy_exchange": legacy_exchange,
341-
"dirac_policies": {},
342-
}
340+
legacy_exchange=legacy_exchange,
341+
dirac_policies={},
342+
)
343343

344344
# Generate access token payload
345345
# For now, the access token is only used to access DIRAC services,
@@ -348,23 +348,28 @@ async def exchange_token(
348348
access_exp = uuid7_to_datetime(access_jti) + timedelta(
349349
minutes=settings.access_token_expire_minutes
350350
)
351-
access_payload: AccessTokenPayload = {
352-
"sub": sub,
353-
"vo": vo,
354-
"iss": settings.token_issuer,
355-
"dirac_properties": list(properties),
356-
"jti": str(access_jti),
357-
"preferred_username": preferred_username,
358-
"dirac_group": dirac_group,
359-
"exp": access_exp,
360-
"dirac_policies": {},
361-
}
351+
access_payload = AccessTokenPayload(
352+
sub=sub,
353+
vo=vo,
354+
iss=settings.token_issuer,
355+
dirac_properties=list(properties),
356+
jti=str(access_jti),
357+
preferred_username=preferred_username,
358+
dirac_group=dirac_group,
359+
exp=access_exp,
360+
dirac_policies={},
361+
)
362362

363363
return access_payload, refresh_payload
364364

365365

366366
def create_token(payload: TokenPayload, settings: AuthSettings) -> str:
367367
"""Create a JWT token with the given payload and settings."""
368+
return _sign_token_payload(payload.model_dump(), settings)
369+
370+
371+
def _sign_token_payload(claims: dict, settings: AuthSettings) -> str:
372+
"""Sign a raw claims dict as a JWT. Used by create_token and tests."""
368373
signing_key = None
369374
for key in settings.token_keystore.jwks.keys:
370375
key_ops = key.get("key_ops")
@@ -379,7 +384,7 @@ def create_token(payload: TokenPayload, settings: AuthSettings) -> str:
379384

380385
return jwt.encode(
381386
header={"alg": signing_key.get("alg"), "kid": signing_key.get("kid")},
382-
claims=cast(Claims, payload),
387+
claims=cast(Claims, claims),
383388
key=settings.token_keystore.jwks,
384389
algorithms=settings.token_allowed_algorithms,
385390
)

diracx-routers/src/diracx/routers/auth/token.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,12 +61,12 @@ async def mint_token(
6161
dirac_refresh_policies[policy_name] = refresh_extra
6262

6363
# Create the access token
64-
access_payload["dirac_policies"] = dirac_access_policies
64+
access_payload.dirac_policies = dirac_access_policies
6565
access_token = create_token(access_payload, settings)
6666

6767
# Create the refresh token
6868
if refresh_payload:
69-
refresh_payload["dirac_policies"] = dirac_refresh_policies
69+
refresh_payload.dirac_policies = dirac_refresh_policies
7070
refresh_token = create_token(refresh_payload, settings)
7171
elif existing_refresh_token:
7272
refresh_token = existing_refresh_token

0 commit comments

Comments
 (0)