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
« 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
10import requests
12from benefits.enrollment.api import BaseDataClass
14logger = logging.getLogger(__name__)
17@dataclass
18class Registration(BaseDataClass):
19 regId: str
20 gtwUrl: str
23class RegistrationMode(Enum):
24 REGISTER = "register"
25 IDENTIFY = "identify"
28class EshopResponseMode(Enum):
29 FRAGMENT = "fragment"
30 QUERY = "query"
31 FORM_POST = "form_post"
32 POST_MESSAGE = "post_message"
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
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
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)
72 key.write(self.private_key)
73 key.seek(0)
75 ca.write(self.ca_certificate)
76 ca.seek(0)
78 # request using temp file paths
79 return request_func(verify=ca.name, cert=(cert.name, key.name))
82class TokenizationClient(Client):
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
98 def _signature_input_string(self, timestamp: str, method: str, request_path: str, body: str = None):
99 if body is None:
100 body = ""
102 return f"{timestamp}{method}{request_path}{body}"
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)
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()
112 return stp_signature
114 def _get_headers(self, method, request_path, request_body: dict = None):
115 timestamp = str(int(datetime.now().timestamp()))
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 }
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 }
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 )
153 response.raise_for_status()
155 return Registration.from_kwargs(**response.json())
157 def get_registration_status(self, registration_id, timeout=5) -> RegistrationStatus:
158 request_path = f"/api/v1/registration/{registration_id}"
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 )
170 response.raise_for_status()
172 return RegistrationStatus.from_kwargs(**response.json())
175@dataclass
176class Group(BaseDataClass):
177 id: int
178 operatorId: int
179 name: str
180 code: str
181 value: int
184@dataclass
185class GroupExpiry(BaseDataClass):
186 group: str
187 expiresAt: datetime | None = None
189 def __post_init__(self):
190 """Parses any date parameters into aware Python datetime objects.
192 For @dataclasses with a generated __init__ function, this function is called automatically.
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
204class EnrollmentClient(Client):
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
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"
230 def _get_headers(self):
231 return {"Authorization": self.authorization_header_value}
233 def healthcheck(self, timeout=5):
234 request_path = "/api/external/discount/echo"
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 )
246 response.raise_for_status()
248 return response.text
250 def get_groups(self, pto_id, timeout=5):
251 request_path = f"/api/external/discount/{pto_id}/groups"
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 )
263 response.raise_for_status()
265 return [Group.from_kwargs(**discount_group) for discount_group in response.json()]
267 def get_groups_for_token(self, pto_id, token, timeout=5):
268 request_path = f"/api/external/discount/{pto_id}/token/{token}"
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 )
280 response.raise_for_status()
282 return [GroupExpiry.from_kwargs(**group_expiry) for group_expiry in response.json()]
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"
287 request_body = {"group": group_id}
288 if expiry:
289 request_body["expiresAt"] = self._format_expiry(expiry)
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 )
302 response.raise_for_status()
304 return response.text
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"
309 request_body = {"group": group_id}
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 )
322 response.raise_for_status()
324 return response.text