diff --git a/src/shared/shared/three_cx_client.py b/src/shared/shared/three_cx_client.py index c579fc5..d463f69 100644 --- a/src/shared/shared/three_cx_client.py +++ b/src/shared/shared/three_cx_client.py @@ -47,6 +47,7 @@ class ThreeCXClient: self._auth_mode = auth_mode self._client_id = auth_kwargs.get("client_id") self._client_secret = auth_kwargs.get("client_secret") + self._retired_secrets: set[str] = set() self._token_expires_at = 0.0 self._auth_lock = threading.RLock() self.session = requests.Session() @@ -65,12 +66,22 @@ class ThreeCXClient: self._authenticate_user(auth_kwargs["username"], auth_kwargs["password"]) def use_client_secret(self, client_secret: str) -> None: - """Point this client at ``client_secret`` and log in again if needed.""" - if client_secret != self._client_secret: - with self._auth_lock: + """Point this client at ``client_secret`` and log in again if needed. + + The compare-and-set runs under ``_auth_lock``. A secret this client has + already replaced is ignored, so a slower caller still holding the + pre-rotation secret cannot write it back. + """ + with self._auth_lock: + if ( + client_secret != self._client_secret + and client_secret not in self._retired_secrets + ): + if self._client_secret: + self._retired_secrets.add(self._client_secret) self._client_secret = client_secret self._token_expires_at = 0.0 - self.ensure_fresh_token() + self._refresh_expired_token() def ensure_fresh_token(self) -> None: """Fetch a new access token when the current one is missing or near expiry.""" @@ -79,9 +90,15 @@ class ThreeCXClient: if time.monotonic() < self._token_expires_at: return with self._auth_lock: - if time.monotonic() < self._token_expires_at: - return - self._authenticate_oauth(self._client_id, self._client_secret) + self._refresh_expired_token() + + def _refresh_expired_token(self) -> None: + """Log in again when the token is due. Caller holds ``_auth_lock``.""" + if self._auth_mode != "oauth": + return + if time.monotonic() < self._token_expires_at: + return + self._authenticate_oauth(self._client_id, self._client_secret) def _request(self, method, url, **kwargs): if self._auth_mode == "oauth": diff --git a/tests/shared/test_three_cx_client.py b/tests/shared/test_three_cx_client.py index c1595cb..0a4ff5a 100644 --- a/tests/shared/test_three_cx_client.py +++ b/tests/shared/test_three_cx_client.py @@ -286,8 +286,9 @@ def test_oauth_client_reauths_when_client_secret_changes(): ) first = tcx.oauth_client("test.3cx.us", "cid", "old-secret") second = tcx.oauth_client("test.3cx.us", "cid", "new-secret") - assert first is second - assert second.session.headers["Authorization"] == "Bearer tok-new" + stale = tcx.oauth_client("test.3cx.us", "cid", "old-secret") + assert first is second is stale + assert stale.session.headers["Authorization"] == "Bearer tok-new" posts = _token_posts() assert len(posts) == 2 assert "client_secret=new-secret" in posts[1].request.body