diff --git a/src/shared/shared/side_effects.py b/src/shared/shared/side_effects.py index 8df4dbe..3231f1e 100644 --- a/src/shared/shared/side_effects.py +++ b/src/shared/shared/side_effects.py @@ -86,23 +86,17 @@ def maybe_repoint_today(date_str: str, shift_type: str, extension: str) -> bool: return False -_cached_3cx: ThreeCXClient | None = None - - def make_3cx_client() -> ThreeCXClient | None: - global _cached_3cx - if _cached_3cx is not None: - return _cached_3cx + """Return the process OAuth client, refreshing it when the token or secret changed.""" secret_prefix = os.environ.get("TCX_SECRET_PREFIX") if not secret_prefix: logger.warning("3CX env vars not set — skipping 3CX call") return None - _cached_3cx = oauth_client( + return oauth_client( domain=get_secret(f"{secret_prefix}domain"), client_id=get_secret(f"{secret_prefix}client-id"), client_secret=get_secret(f"{secret_prefix}client-secret"), ) - return _cached_3cx def set_holiday_queue_agents(schedule, date_str: str) -> None: diff --git a/src/shared/shared/three_cx_client.py b/src/shared/shared/three_cx_client.py index 37f30d2..c579fc5 100644 --- a/src/shared/shared/three_cx_client.py +++ b/src/shared/shared/three_cx_client.py @@ -1,13 +1,23 @@ import logging +import threading +import time + import requests logger = logging.getLogger(__name__) _oauth_clients: dict[tuple[str, str], "ThreeCXClient"] = {} +# Refresh this long before 3CX's expires_in, so a job does not start on a token +# that dies mid-call. 3CX client-credentials tokens last 3600 seconds. +_TOKEN_SKEW_SECONDS = 60 def oauth_client(domain: str, client_id: str, client_secret: str) -> "ThreeCXClient": - """Reuse one OAuth client per (domain, client_id) in this process.""" + """Reuse one OAuth client per (domain, client_id) in this process. + + Re-authenticates when the access token is near expiry, and immediately when + ``client_secret`` differs from the one the cached client logged in with. + """ key = (domain, client_id) client = _oauth_clients.get(key) if client is None: @@ -18,6 +28,8 @@ def oauth_client(domain: str, client_id: str, client_secret: str) -> "ThreeCXCli client_secret=client_secret, ) _oauth_clients[key] = client + return client + client.use_client_secret(client_secret) return client @@ -32,7 +44,14 @@ class ThreeCXClient: auth_kwargs: credentials — see _authenticate_user / _authenticate_oauth """ self.base_url = f"https://{domain}" + self._auth_mode = auth_mode + self._client_id = auth_kwargs.get("client_id") + self._client_secret = auth_kwargs.get("client_secret") + self._token_expires_at = 0.0 + self._auth_lock = threading.RLock() self.session = requests.Session() + self._raw_request = self.session.request + self.session.request = self._request self.session.headers.update( { "OData-Version": "4.0", @@ -41,12 +60,40 @@ class ThreeCXClient: ) if auth_mode == "oauth": - self._authenticate_oauth( - auth_kwargs["client_id"], auth_kwargs["client_secret"] - ) + self._authenticate_oauth(self._client_id, self._client_secret) else: 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: + self._client_secret = client_secret + self._token_expires_at = 0.0 + self.ensure_fresh_token() + + def ensure_fresh_token(self) -> None: + """Fetch a new access token when the current one is missing or near expiry.""" + if self._auth_mode != "oauth": + return + 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) + + def _request(self, method, url, **kwargs): + if self._auth_mode == "oauth": + self.ensure_fresh_token() + response = self._raw_request(method, url, **kwargs) + if self._auth_mode == "oauth" and response.status_code == 401: + with self._auth_lock: + self._token_expires_at = 0.0 + self._authenticate_oauth(self._client_id, self._client_secret) + response = self._raw_request(method, url, **kwargs) + return response + def _authenticate_user(self, username: str, password: str): """Authenticate via extension/user credentials (any license tier).""" resp = self.session.post( @@ -64,7 +111,8 @@ class ThreeCXClient: def _authenticate_oauth(self, client_id: str, client_secret: str): """Authenticate via OAuth2 client credentials (Enterprise license required). API client must be created in 3CX Admin > Integrations > API.""" - resp = self.session.post( + resp = self._raw_request( + "POST", f"{self.base_url}/connect/token", data={ "client_id": client_id, @@ -74,7 +122,13 @@ class ThreeCXClient: headers={"Content-Type": "application/x-www-form-urlencoded"}, ) resp.raise_for_status() - token = resp.json()["access_token"] + body = resp.json() + token = body.get("access_token") + if not token: + raise ValueError("Failed to get access token from 3CX OAuth response") + expires_in = int(body.get("expires_in") or 3600) + skew = min(_TOKEN_SKEW_SECONDS, expires_in // 10) + self._token_expires_at = time.monotonic() + max(expires_in - skew, 0) self.session.headers.update({"Authorization": f"Bearer {token}"}) logger.info("Authenticated to 3CX via OAuth2 client credentials") diff --git a/tests/shared/test_three_cx_client.py b/tests/shared/test_three_cx_client.py index 994841c..c1595cb 100644 --- a/tests/shared/test_three_cx_client.py +++ b/tests/shared/test_three_cx_client.py @@ -218,6 +218,10 @@ def test_extract_ivr_routes_missing_key0_is_none(): assert routes == {"key0": None, "timeout": None} +def _token_posts(): + return [c for c in responses.calls if c.request.url.endswith("/connect/token")] + + @responses.activate def test_oauth_client_reuses_process_cache(): from shared import three_cx_client as tcx @@ -227,8 +231,92 @@ def test_oauth_client_reuses_process_cache(): first = tcx.oauth_client("test.3cx.us", "cid", "secret") second = tcx.oauth_client("test.3cx.us", "cid", "secret") assert first is second - token_posts = [ - c for c in responses.calls if c.request.url.endswith("/connect/token") - ] - assert len(token_posts) == 1 + assert len(_token_posts()) == 1 tcx._oauth_clients.clear() + + +@responses.activate +def test_oauth_client_refreshes_expired_token(): + from shared import three_cx_client as tcx + + tcx._oauth_clients.clear() + responses.add( + responses.POST, + f"{BASE}/connect/token", + json={"access_token": "tok-1", "expires_in": 3600}, + status=200, + ) + responses.add( + responses.POST, + f"{BASE}/connect/token", + json={"access_token": "tok-2", "expires_in": 3600}, + status=200, + ) + responses.add( + responses.GET, + f"{BASE}/xapi/v1/Queues/Pbx.GetByNumber(number='801')", + json={"Id": 83}, + status=200, + ) + client = tcx.oauth_client("test.3cx.us", "cid", "secret") + client._token_expires_at = 0 + queue = client.get_queue("801") + assert queue["Id"] == 83 + assert client.session.headers["Authorization"] == "Bearer tok-2" + assert len(_token_posts()) == 2 + tcx._oauth_clients.clear() + + +@responses.activate +def test_oauth_client_reauths_when_client_secret_changes(): + from shared import three_cx_client as tcx + + tcx._oauth_clients.clear() + responses.add( + responses.POST, + f"{BASE}/connect/token", + json={"access_token": "tok-old", "expires_in": 3600}, + status=200, + ) + responses.add( + responses.POST, + f"{BASE}/connect/token", + json={"access_token": "tok-new", "expires_in": 3600}, + status=200, + ) + 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" + posts = _token_posts() + assert len(posts) == 2 + assert "client_secret=new-secret" in posts[1].request.body + tcx._oauth_clients.clear() + + +@responses.activate +def test_oauth_request_retries_once_after_401(): + _stub_oauth() + responses.add( + responses.POST, + f"{BASE}/connect/token", + json={"access_token": "tok-refreshed", "expires_in": 3600}, + status=200, + ) + responses.add( + responses.GET, + f"{BASE}/xapi/v1/Queues/Pbx.GetByNumber(number='801')", + status=401, + ) + responses.add( + responses.GET, + f"{BASE}/xapi/v1/Queues/Pbx.GetByNumber(number='801')", + json={"Id": 83}, + status=200, + ) + client = ThreeCXClient( + domain="test.3cx.us", auth_mode="oauth", client_id="c", client_secret="s" + ) + assert client.get_queue("801")["Id"] == 83 + assert client.session.headers["Authorization"] == "Bearer tok-refreshed" + assert len(_token_posts()) == 2