Coverage for benefits/enrollment_switchio/api.py: 99%

135 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-10-08 19:50 +0000

1import hashlib 

2import hmac 

3import json 

4import logging 

5from dataclasses import dataclass 

6from datetime import datetime, timezone 

7from enum import Enum 

8from tempfile import NamedTemporaryFile 

9 

10import requests 

11 

12from benefits.enrollment.api import BaseDataClass 

13 

14logger = logging.getLogger(__name__) 

15 

16 

17@dataclass 

18class Registration(BaseDataClass): 

19 regId: str 

20 gtwUrl: str 

21 

22 

23class RegistrationMode(Enum): 

24 REGISTER = "register" 

25 IDENTIFY = "identify" 

26 

27 

28class EshopResponseMode(Enum): 

29 FRAGMENT = "fragment" 

30 QUERY = "query" 

31 FORM_POST = "form_post" 

32 POST_MESSAGE = "post_message" 

33 

34 

35@dataclass 

36class RegistrationStatus(BaseDataClass): 

37 regState: str 

38 created: datetime 

39 mode: str 

40 tokens: list[dict] 

41 eshopResponseMode: str 

42 identType: str = None 

43 maskCln: str = None 

44 cardExp: str = None 

45 

46 

47class Client: 

48 def __init__( 

49 self, 

50 private_key, 

51 client_certificate, 

52 ca_certificate, 

53 ): 

54 self.private_key = private_key 

55 self.client_certificate = client_certificate 

56 self.ca_certificate = ca_certificate 

57 

58 # see https://github.com/cal-itp/benefits/issues/2848 for more context about this 

59 def _cert_request(self, request_func): 

60 """ 

61 Creates named (on-disk) temp files for client cert auth. 

62 * request_func: curried callable from `requests` library (e.g. `requests.get`). 

63 """ 

64 # requests library reads temp files from file path 

65 # The "with" context destroys temp files when response comes back 

66 with NamedTemporaryFile("w+") as cert, NamedTemporaryFile("w+") as key, NamedTemporaryFile("w+") as ca: 

67 # write client cert data to temp files 

68 # resetting so they can be read again by requests 

69 cert.write(self.client_certificate) 

70 cert.seek(0) 

71 

72 key.write(self.private_key) 

73 key.seek(0) 

74 

75 ca.write(self.ca_certificate) 

76 ca.seek(0) 

77 

78 # request using temp file paths 

79 return request_func(verify=ca.name, cert=(cert.name, key.name)) 

80 

81 

82class TokenizationClient(Client): 

83 

84 def __init__( 

85 self, 

86 api_url, 

87 api_key, 

88 api_secret, 

89 private_key, 

90 client_certificate, 

91 ca_certificate, 

92 ): 

93 super().__init__(private_key, client_certificate, ca_certificate) 

94 self.api_url = api_url.strip("/") 

95 self.api_key = api_key 

96 self.api_secret = api_secret 

97 

98 def _signature_input_string(self, timestamp: str, method: str, request_path: str, body: str = None): 

99 if body is None: 

100 body = "" 

101 

102 return f"{timestamp}{method}{request_path}{body}" 

103 

104 def _stp_signature(self, timestamp: str, method: str, request_path, body: str = None): 

105 input_string = self._signature_input_string(timestamp, method, request_path, body) 

106 

107 # must encode inputs for hashing, according to https://stackoverflow.com/a/66958131 

108 byte_key = self.api_secret.encode("utf-8") 

109 message = input_string.encode("utf-8") 

110 stp_signature = hmac.new(byte_key, message, hashlib.sha256).hexdigest() 

111 

112 return stp_signature 

113 

114 def _get_headers(self, method, request_path, request_body: dict = None): 

115 timestamp = str(int(datetime.now().timestamp())) 

116 

117 return { 

118 "STP-APIKEY": self.api_key, 

119 "STP-TIMESTAMP": timestamp, 

120 "STP-SIGNATURE": self._stp_signature( 

121 timestamp=timestamp, 

122 method=method, 

123 request_path=request_path, 

124 body=json.dumps(request_body) if request_body else None, 

125 ), 

126 } 

127 

128 def request_registration( 

129 self, 

130 eshopRedirectUrl: str, 

131 mode: RegistrationMode, 

132 eshopResponseMode: EshopResponseMode, 

133 timeout=5, 

134 ) -> Registration: 

135 registration_path = "/api/v1/registration" 

136 request_body = { 

137 "eshopRedirectUrl": eshopRedirectUrl, 

138 "mode": mode.value, 

139 "eshopResponseMode": eshopResponseMode.value, 

140 } 

141 

142 response = self._cert_request( 

143 lambda verify, cert: requests.post( 

144 self.api_url + registration_path, 

145 json=request_body, 

146 headers=self._get_headers(method="POST", request_path=registration_path, request_body=request_body), 

147 cert=cert, 

148 verify=verify, 

149 timeout=timeout, 

150 ) 

151 ) 

152 

153 response.raise_for_status() 

154 

155 return Registration.from_kwargs(**response.json()) 

156 

157 def get_registration_status(self, registration_id, timeout=5) -> RegistrationStatus: 

158 request_path = f"/api/v1/registration/{registration_id}" 

159 

160 response = self._cert_request( 

161 lambda verify, cert: requests.get( 

162 self.api_url + request_path, 

163 headers=self._get_headers(method="GET", request_path=request_path), 

164 cert=cert, 

165 verify=verify, 

166 timeout=timeout, 

167 ) 

168 ) 

169 

170 response.raise_for_status() 

171 

172 return RegistrationStatus.from_kwargs(**response.json()) 

173 

174 

175@dataclass 

176class Group(BaseDataClass): 

177 id: int 

178 operatorId: int 

179 name: str 

180 code: str 

181 value: int 

182 

183 

184@dataclass 

185class GroupExpiry(BaseDataClass): 

186 group: str 

187 expiresAt: datetime | None = None 

188 

189 def __post_init__(self): 

190 """Parses any date parameters into aware Python datetime objects. 

191 

192 For @dataclasses with a generated __init__ function, this function is called automatically. 

193 

194 https://docs.python.org/3.12/library/datetime.html#datetime.datetime.fromisoformat 

195 """ 

196 if self.expiresAt: 

197 # per the spec, we expect expiresAt to be in ISO format UTC 

198 # make the resulting datetime "aware" by replacing tzinfo explicitly 

199 self.expiresAt = datetime.fromisoformat(self.expiresAt).replace(tzinfo=timezone.utc) 

200 else: 

201 self.expiresAt = None 

202 

203 

204class EnrollmentClient(Client): 

205 

206 def __init__(self, api_url, authorization_header_value, private_key, client_certificate, ca_certificate): 

207 super().__init__(private_key, client_certificate, ca_certificate) 

208 self.api_url = api_url.strip("/") 

209 self.authorization_header_value = authorization_header_value 

210 

211 def _format_expiry(self, expiry: datetime) -> str: 

212 """Formats an expiry datetime into a string suitable for using in an API request body.""" 

213 if not isinstance(expiry, datetime): 213 ↛ 214line 213 didn't jump to line 214 because the condition on line 213 was never true

214 raise TypeError("expiry must be a Python datetime instance") 

215 # determine if expiry is an "aware" or "naive" datetime instance 

216 # https://docs.python.org/3/library/datetime.html#determining-if-an-object-is-aware-or-naive 

217 if expiry.tzinfo is not None and expiry.tzinfo.utcoffset(expiry) is not None: 

218 # expiry is an "aware" datetime instance, meaning it has associated time zone information 

219 # ensure this datetime instance is expressed in UTC 

220 expiry = expiry.astimezone(timezone.utc) 

221 else: 

222 # expiry is a "naive" datetime instance, meaning it has no associated time zone information 

223 # assume this datetime instance was provided in UTC 

224 expiry = expiry.replace(tzinfo=timezone.utc) 

225 # now expiry is an "aware" datetime instance in UTC format 

226 # datetime.isoformat() adds the UTC offset like +00:00 

227 # so keep everything but the last 6 characters and add the Z offset character 

228 return f"{expiry.isoformat(timespec='seconds')[:-6]}Z" 

229 

230 def _get_headers(self): 

231 return {"Authorization": self.authorization_header_value} 

232 

233 def healthcheck(self, timeout=5): 

234 request_path = "/api/external/discount/echo" 

235 

236 response = self._cert_request( 

237 lambda verify, cert: requests.get( 

238 self.api_url.strip("/") + request_path, 

239 headers=self._get_headers(), 

240 cert=cert, 

241 verify=verify, 

242 timeout=timeout, 

243 ) 

244 ) 

245 

246 response.raise_for_status() 

247 

248 return response.text 

249 

250 def get_groups(self, pto_id, timeout=5): 

251 request_path = f"/api/external/discount/{pto_id}/groups" 

252 

253 response = self._cert_request( 

254 lambda verify, cert: requests.get( 

255 self.api_url + request_path, 

256 headers=self._get_headers(), 

257 cert=cert, 

258 verify=verify, 

259 timeout=timeout, 

260 ) 

261 ) 

262 

263 response.raise_for_status() 

264 

265 return [Group.from_kwargs(**discount_group) for discount_group in response.json()] 

266 

267 def get_groups_for_token(self, pto_id, token, timeout=5): 

268 request_path = f"/api/external/discount/{pto_id}/token/{token}" 

269 

270 response = self._cert_request( 

271 lambda verify, cert: requests.get( 

272 self.api_url + request_path, 

273 headers=self._get_headers(), 

274 cert=cert, 

275 verify=verify, 

276 timeout=timeout, 

277 ) 

278 ) 

279 

280 response.raise_for_status() 

281 

282 return [GroupExpiry.from_kwargs(**group_expiry) for group_expiry in response.json()] 

283 

284 def add_group_to_token(self, pto_id, group_id, token, expiry: datetime = None, timeout=5): 

285 request_path = f"/api/external/discount/{pto_id}/token/{token}/add" 

286 

287 request_body = {"group": group_id} 

288 if expiry: 

289 request_body["expiresAt"] = self._format_expiry(expiry) 

290 

291 response = self._cert_request( 

292 lambda verify, cert: requests.post( 

293 self.api_url + request_path, 

294 json=request_body, 

295 headers=self._get_headers(), 

296 cert=cert, 

297 verify=verify, 

298 timeout=timeout, 

299 ) 

300 ) 

301 

302 response.raise_for_status() 

303 

304 return response.text 

305 

306 def remove_group_from_token(self, pto_id, group_id, token, timeout=5): 

307 request_path = f"/api/external/discount/{pto_id}/token/{token}/remove" 

308 

309 request_body = {"group": group_id} 

310 

311 response = self._cert_request( 

312 lambda verify, cert: requests.post( 

313 self.api_url + request_path, 

314 json=request_body, 

315 headers=self._get_headers(), 

316 cert=cert, 

317 verify=verify, 

318 timeout=timeout, 

319 ) 

320 ) 

321 

322 response.raise_for_status() 

323 

324 return response.text