diff --git a/parsedmarc/mail/graph.py b/parsedmarc/mail/graph.py index e87ac7a3..b8c02bb5 100644 --- a/parsedmarc/mail/graph.py +++ b/parsedmarc/mail/graph.py @@ -16,10 +16,14 @@ from azure.identity import ( AuthenticationRecord, ) from msgraph.core import GraphClient +from requests.exceptions import RequestException from parsedmarc.log import logger from parsedmarc.mail.mailbox_connection import MailboxConnection +GRAPH_REQUEST_RETRY_ATTEMPTS = 3 +GRAPH_REQUEST_RETRY_DELAY_SECONDS = 5 + class AuthMethod(Enum): DeviceCode = 1 @@ -129,6 +133,23 @@ class MSGraphConnection(MailboxConnection): self._client = GraphClient(**client_params) self.mailbox_name = mailbox + def _request_with_retries(self, method_name: str, *args, **kwargs): + for attempt in range(1, GRAPH_REQUEST_RETRY_ATTEMPTS + 1): + try: + return getattr(self._client, method_name)(*args, **kwargs) + except RequestException as error: + if attempt == GRAPH_REQUEST_RETRY_ATTEMPTS: + raise + logger.warning( + "Transient MS Graph %s error on attempt %s/%s: %s", + method_name.upper(), + attempt, + GRAPH_REQUEST_RETRY_ATTEMPTS, + error, + ) + sleep(GRAPH_REQUEST_RETRY_DELAY_SECONDS) + raise RuntimeError("no retry attempts configured") + def create_folder(self, folder_name: str): sub_url = "" path_parts = folder_name.split("/") @@ -143,7 +164,7 @@ class MSGraphConnection(MailboxConnection): request_body = {"displayName": folder_name} request_url = f"/users/{self.mailbox_name}/mailFolders{sub_url}" - resp = self._client.post(request_url, json=request_body) + resp = self._request_with_retries("post", request_url, json=request_body) if resp.status_code == 409: logger.debug(f"Folder {folder_name} already exists, skipping creation") elif resp.status_code == 201: @@ -173,7 +194,7 @@ class MSGraphConnection(MailboxConnection): params["$top"] = batch_size else: params["$top"] = 100 - result = self._client.get(url, params=params) + result = self._request_with_retries("get", url, params=params) if result.status_code != 200: raise RuntimeError(f"Failed to fetch messages {result.text}") messages = result.json()["value"] @@ -181,7 +202,7 @@ class MSGraphConnection(MailboxConnection): while "@odata.nextLink" in result.json() and ( since is not None or (batch_size == 0 or batch_size - len(messages) > 0) ): - result = self._client.get(result.json()["@odata.nextLink"]) + result = self._request_with_retries("get", result.json()["@odata.nextLink"]) if result.status_code != 200: raise RuntimeError(f"Failed to fetch messages {result.text}") messages.extend(result.json()["value"]) @@ -190,7 +211,7 @@ class MSGraphConnection(MailboxConnection): def mark_message_read(self, message_id: str): """Marks a message as read""" url = f"/users/{self.mailbox_name}/messages/{message_id}" - resp = self._client.patch(url, json={"isRead": "true"}) + resp = self._request_with_retries("patch", url, json={"isRead": "true"}) if resp.status_code != 200: raise RuntimeWarning( f"Failed to mark message read{resp.status_code}: {resp.json()}" @@ -198,7 +219,7 @@ class MSGraphConnection(MailboxConnection): def fetch_message(self, message_id: str, **kwargs): url = f"/users/{self.mailbox_name}/messages/{message_id}/$value" - result = self._client.get(url) + result = self._request_with_retries("get", url) if result.status_code != 200: raise RuntimeWarning( f"Failed to fetch message{result.status_code}: {result.json()}" @@ -210,7 +231,7 @@ class MSGraphConnection(MailboxConnection): def delete_message(self, message_id: str): url = f"/users/{self.mailbox_name}/messages/{message_id}" - resp = self._client.delete(url) + resp = self._request_with_retries("delete", url) if resp.status_code != 204: raise RuntimeWarning( f"Failed to delete message {resp.status_code}: {resp.json()}" @@ -220,7 +241,7 @@ class MSGraphConnection(MailboxConnection): folder_id = self._find_folder_id_from_folder_path(folder_name) request_body = {"destinationId": folder_id} url = f"/users/{self.mailbox_name}/messages/{message_id}/move" - resp = self._client.post(url, json=request_body) + resp = self._request_with_retries("post", url, json=request_body) if resp.status_code != 201: raise RuntimeWarning( f"Failed to move message {resp.status_code}: {resp.json()}" @@ -256,7 +277,7 @@ class MSGraphConnection(MailboxConnection): sub_url = f"/{parent_folder_id}/childFolders" url = f"/users/{self.mailbox_name}/mailFolders{sub_url}" filter = f"?$filter=displayName eq '{folder_name}'" - folders_resp = self._client.get(url + filter) + folders_resp = self._request_with_retries("get", url + filter) if folders_resp.status_code != 200: raise RuntimeWarning(f"Failed to list folders.{folders_resp.json()}") folders: list = folders_resp.json()["value"] diff --git a/tests.py b/tests.py index ecfbc883..7f54b601 100755 --- a/tests.py +++ b/tests.py @@ -628,6 +628,31 @@ class TestGraphConnection(unittest.TestCase): with self.assertRaises(RuntimeError): connection._get_all_messages("/url", batch_size=0, since=None) + def testGetAllMessagesRetriesTransientRequestErrors(self): + connection = MSGraphConnection.__new__(MSGraphConnection) + connection._client = MagicMock() + connection._client.get.side_effect = [ + graph_module.RequestException("connection reset"), + _FakeGraphResponse(200, {"value": [{"id": "1"}]}), + ] + with patch.object(graph_module, "sleep") as mocked_sleep: + messages = connection._get_all_messages("/url", batch_size=0, since=None) + self.assertEqual([msg["id"] for msg in messages], ["1"]) + mocked_sleep.assert_called_once_with(graph_module.GRAPH_REQUEST_RETRY_DELAY_SECONDS) + + def testGetAllMessagesRaisesAfterRetryExhaustion(self): + connection = MSGraphConnection.__new__(MSGraphConnection) + connection._client = MagicMock() + connection._client.get.side_effect = graph_module.RequestException( + "connection reset" + ) + with patch.object(graph_module, "sleep") as mocked_sleep: + with self.assertRaises(graph_module.RequestException): + connection._get_all_messages("/url", batch_size=0, since=None) + self.assertEqual( + mocked_sleep.call_count, graph_module.GRAPH_REQUEST_RETRY_ATTEMPTS - 1 + ) + def testGetAllMessagesNextPageFailure(self): connection = MSGraphConnection.__new__(MSGraphConnection) first_response = _FakeGraphResponse(