Skip to content

Commit ebace76

Browse files
authored
Ensure that the refresh token is only used within the expiration window (bachya#266)
1 parent a8650ef commit ebace76

4 files changed

Lines changed: 57 additions & 4 deletions

File tree

‎simplipy/api.py‎

Lines changed: 24 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
11
"""Define functionality for interacting with the SimpliSafe API."""
22
from __future__ import annotations
33

4+
from datetime import datetime, timedelta
45
from json.decoder import JSONDecodeError
56
import sys
6-
from typing import Any, Callable
7+
from typing import TYPE_CHECKING, Any, Callable
78

89
from aiohttp import ClientSession
910
from aiohttp.client_exceptions import ClientResponseError
@@ -28,10 +29,21 @@
2829
API_URL_HOSTNAME = "api.simplisafe.com"
2930
API_URL_BASE = f"https://{API_URL_HOSTNAME}/v1"
3031

32+
DEFAULT_EXPIRATION_PADDING = 60
3133
DEFAULT_REQUEST_RETRIES = 4
3234
DEFAULT_TIMEOUT = 10
3335

3436

37+
def get_expiration_datetime(expires_in_seconds: int) -> datetime:
38+
"""Get a token expiration datetime as an offset of UTC now + a number of seconds.
39+
40+
Note that we pad the value to ensure the token doesn't expire without us knowing.
41+
"""
42+
return datetime.utcnow() + (
43+
timedelta(seconds=expires_in_seconds - DEFAULT_EXPIRATION_PADDING)
44+
)
45+
46+
3547
class API: # pylint: disable=too-many-instance-attributes
3648
"""An API object to interact with the SimpliSafe cloud.
3749
@@ -56,6 +68,7 @@ def __init__(
5668
self.session: ClientSession = session
5769

5870
# These will get filled in after initial authentication:
71+
self._access_token_expire_dt: datetime | None = None
5972
self.access_token: str | None = None
6073
self.refresh_token: str | None = None
6174
self.subscription_data: dict[int, Any] = {}
@@ -114,6 +127,7 @@ async def async_from_auth(
114127
raise InvalidCredentialsError("Invalid credentials") from err
115128
raise RequestError(err) from err
116129

130+
api._access_token_expire_dt = get_expiration_datetime(token_resp["expires_in"])
117131
api.access_token = token_resp["access_token"]
118132
api.refresh_token = token_resp["refresh_token"]
119133
await api._async_post_init()
@@ -149,8 +163,14 @@ async def _async_handle_on_backoff(self, _: dict[str, Any]) -> None:
149163
err: ClientResponseError = err_info[1].with_traceback(err_info[2]) # type: ignore
150164

151165
if err.status == 401 or err.status == 403:
152-
LOGGER.info("401 detected; attempting refresh token")
153-
await self._async_refresh_access_token()
166+
if TYPE_CHECKING:
167+
assert self._access_token_expire_dt
168+
if datetime.utcnow() >= self._access_token_expire_dt:
169+
# Since we might have multiple requests (each running their own retry
170+
# sequence) land here, we only refresh the access token if it hasn't
171+
# been refreshed within the expiration window:
172+
LOGGER.info("401 detected; attempting refresh token")
173+
await self._async_refresh_access_token()
154174

155175
async def _async_handle_on_giveup(self, _: dict[str, Any]) -> None:
156176
"""Handle a give up after retries are exhausted."""
@@ -183,6 +203,7 @@ async def _async_refresh_access_token(self) -> None:
183203
raise InvalidCredentialsError("Invalid refresh token") from err
184204
raise RequestError(err) from err
185205

206+
self._access_token_expire_dt = get_expiration_datetime(token_resp["expires_in"])
186207
self.access_token = token_resp["access_token"]
187208
self.refresh_token = token_resp["refresh_token"]
188209

‎tests/system/test_v3.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -232,6 +232,10 @@ async def test_no_state_change_on_failure(aresponses, v3_server):
232232
simplisafe = await API.async_from_auth(
233233
TEST_AUTHORIZATION_CODE, TEST_CODE_VERIFIER, session=session
234234
)
235+
236+
# Manually set the expiration datetime to force a refresh token flow:
237+
simplisafe._access_token_expire_dt = datetime.utcnow()
238+
235239
systems = await simplisafe.async_get_systems()
236240
system = systems[TEST_SYSTEM_ID]
237241
assert system.state == SystemStates.off

‎tests/test_api.py‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
"""Define tests for the System object."""
22
# pylint: disable=protected-access,too-many-arguments
3+
from datetime import datetime
34
from unittest.mock import Mock
45

56
import aiohttp
@@ -63,6 +64,10 @@ async def test_401_refresh_token_failure(
6364
simplisafe = await API.async_from_auth(
6465
TEST_AUTHORIZATION_CODE, TEST_CODE_VERIFIER, session=session
6566
)
67+
68+
# Manually set the expiration datetime to force a refresh token flow:
69+
simplisafe._access_token_expire_dt = datetime.utcnow()
70+
6671
await simplisafe.async_get_systems()
6772

6873
aresponses.assert_plan_strictly_followed()
@@ -113,6 +118,9 @@ async def test_401_refresh_token_success(
113118
TEST_AUTHORIZATION_CODE, TEST_CODE_VERIFIER, session=session
114119
)
115120

121+
# Manually set the expiration datetime to force a refresh token flow:
122+
simplisafe._access_token_expire_dt = datetime.utcnow()
123+
116124
# If this succeeds without throwing an exception, the retry is successful:
117125
await simplisafe.async_get_systems()
118126
assert simplisafe.access_token == "jjhhgg66"
@@ -200,17 +208,27 @@ async def test_client_async_from_refresh_token(
200208
async def test_refresh_token_listener_callback(
201209
api_token_response,
202210
aresponses,
211+
caplog,
203212
server,
204213
v2_settings_response,
205214
v2_subscriptions_response,
206215
):
207216
"""Test that listener callbacks are executed correctly."""
217+
import logging
218+
219+
caplog.set_level(logging.DEBUG)
208220
server.add(
209221
"api.simplisafe.com",
210222
f"/v1/users/{TEST_SUBSCRIPTION_ID}/subscriptions",
211223
"get",
212224
response=aresponses.Response(text="Unauthorized", status=401),
213225
)
226+
server.add(
227+
"api.simplisafe.com",
228+
f"/v1/subscriptions/{TEST_SUBSCRIPTION_ID}/settings",
229+
"get",
230+
response=aresponses.Response(text="Unauthorized", status=401),
231+
)
214232

215233
api_token_response["access_token"] = "jjhhgg66"
216234
api_token_response["refresh_token"] = "aabbcc11"
@@ -244,6 +262,9 @@ async def test_refresh_token_listener_callback(
244262
TEST_AUTHORIZATION_CODE, TEST_CODE_VERIFIER, session=session
245263
)
246264

265+
# Manually set the expiration datetime to force a refresh token flow:
266+
simplisafe._access_token_expire_dt = datetime.utcnow()
267+
247268
# We'll hang onto one listener callback:
248269
simplisafe.add_refresh_token_listener(mock_listener_1)
249270
assert mock_listener_1.call_count == 0
@@ -254,6 +275,7 @@ async def test_refresh_token_listener_callback(
254275

255276
await simplisafe.async_get_systems()
256277
mock_listener_1.assert_called_once_with("aabbcc11")
278+
assert mock_listener_1.call_count == 1
257279
assert mock_listener_2.call_count == 0
258280

259281

‎tests/test_lock.py‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
"""Define tests for the Lock objects."""
2-
# pylint: disable=unused-argument
2+
# pylint: disable=protected-access,unused-argument
3+
from datetime import datetime
4+
35
import aiohttp
46
import pytest
57

@@ -105,6 +107,10 @@ async def test_no_state_change_on_failure(
105107
TEST_CODE_VERIFIER,
106108
session=session,
107109
)
110+
111+
# Manually set the expiration datetime to force a refresh token flow:
112+
simplisafe._access_token_expire_dt = datetime.utcnow()
113+
108114
systems = await simplisafe.async_get_systems()
109115
system = systems[TEST_SYSTEM_ID]
110116
lock = system.locks[TEST_LOCK_ID]

0 commit comments

Comments
 (0)