From 94c9ff90e83b8e434f8c449d148027c31fd1dda9 Mon Sep 17 00:00:00 2001 From: adeshmukh-ks Date: Tue, 11 Aug 2026 11:31:55 +0530 Subject: [PATCH] Handle throttle api response with retry mechanism --- .../src/keepersdk/authentication/endpoint.py | 283 +++++++++++++----- .../keepersdk/authentication/keeper_auth.py | 35 ++- .../unit_tests/test_throttle_handling.py | 250 ++++++++++++++++ 3 files changed, 491 insertions(+), 77 deletions(-) create mode 100644 keepersdk-package/unit_tests/test_throttle_handling.py diff --git a/keepersdk-package/src/keepersdk/authentication/endpoint.py b/keepersdk-package/src/keepersdk/authentication/endpoint.py index f351b3e7..ce562bf7 100644 --- a/keepersdk-package/src/keepersdk/authentication/endpoint.py +++ b/keepersdk-package/src/keepersdk/authentication/endpoint.py @@ -3,9 +3,12 @@ import locale import logging import os +import re import time import warnings -from typing import Optional, Dict, Any, Type, TypeVar +from datetime import datetime, timezone +from email.utils import parsedate_to_datetime +from typing import Optional, Dict, Any, Type, TypeVar, Mapping, Tuple from urllib.parse import urlunparse, urlparse import requests @@ -22,6 +25,90 @@ TRQ = TypeVar('TRQ', bound=Message) TRS = TypeVar('TRS', bound=Message) +DEFAULT_TIMEOUT = (15, 120) +MAX_THROTTLE_RETRIES = 3 +MAX_KEY_RETRIES = 3 +MAX_THROTTLE_WAIT_SECONDS = 300 +DEFAULT_THROTTLE_WAIT_SECONDS = 60 + + +def _parse_retry_after(value: str) -> Optional[int]: + """Parse a ``Retry-After`` value: delta-seconds or an HTTP date.""" + value = (value or '').strip() + if not value: + return None + try: + return max(int(value), 0) + except ValueError: + pass + try: + retry_at = parsedate_to_datetime(value) + except (TypeError, ValueError): + return None + if retry_at is None: + return None + if retry_at.tzinfo is None: + retry_at = retry_at.replace(tzinfo=timezone.utc) + return max(int((retry_at - datetime.now(timezone.utc)).total_seconds()), 0) + + +def parse_throttle_wait_seconds(message: str = '', + headers: Optional[Mapping[str, str]] = None) -> int: + """Resolve how long to wait after a throttle / rate-limit response. + + Prefers the ``Retry-After`` header when present, otherwise parses a + human-readable duration from the server message. Result is capped at + ``MAX_THROTTLE_WAIT_SECONDS``. + """ + if headers: + retry_after = _parse_retry_after(headers.get('Retry-After') or '') + if retry_after is not None: + return min(retry_after, MAX_THROTTLE_WAIT_SECONDS) + + wait_seconds = DEFAULT_THROTTLE_WAIT_SECONDS + wait_match = re.search(r'(\d+)\s*(second|minute)', message or '', re.IGNORECASE) + if wait_match: + wait_val = int(wait_match.group(1)) + if 'minute' in wait_match.group(2).lower(): + wait_seconds = wait_val * 60 + else: + wait_seconds = wait_val + return min(wait_seconds, MAX_THROTTLE_WAIT_SECONDS) + + +def throttle_backoff_seconds(throttle_retries: int, wait_seconds: int) -> int: + """Exponential backoff floored at the server-suggested wait.""" + return max(wait_seconds, 30 * (2 ** (throttle_retries - 1))) + + +def parse_error_response(response: requests.Response) -> Tuple[Optional[Dict[str, Any]], str, str]: + """Extract ``(body, error_code, message)`` from a failed Keeper response. + + The body is ``None`` unless the server returned parsable JSON, so callers can + tell a structured API failure apart from a gateway/proxy error page. + """ + content_type = response.headers.get('Content-Type') or '' + if not content_type.startswith('application/json'): + return None, '', '' + try: + body = response.json() + except ValueError: + return None, '', '' + if not isinstance(body, dict): + return None, '', '' + + error_code = body.get('error') or '' + message = body.get('message') or '' + additional_info = body.get('additional_info') + if additional_info: + message += f'({additional_info})' + return body, error_code, message + + +def is_throttle_response(status_code: int, error_code: str) -> bool: + """Keeper signals throttling as 403 + ``error=throttled``; 429 is the HTTP standard.""" + return status_code == 429 or error_code == 'throttled' + _proxies: Optional[Dict] = None _certificate_check: bool = True @@ -150,29 +237,44 @@ def execute_router_rest(self, endpoint: str, *, session_token: bytes, payload: O logger.debug('>>> [ROUTER] POST Request: [%s]', url) if payload is not None: payload = crypto.encrypt_aes_v2(payload, transmission_key) - response = requests.post(url, headers=headers, data=payload, verify=get_certificate_check()) - logger.debug('<<< [ROUTER] Response Code: [%d]', response.status_code) - - if response.status_code == 200: - rs_body = response.content - if rs_body: - router_response = router_pb2.RouterResponse() - router_response.ParseFromString(rs_body) - if router_response.responseCode == router_pb2.RouterResponseCode.RRC_OK: - if router_response.encryptedPayload: - return crypto.decrypt_aes_v2(router_response.encryptedPayload, transmission_key) - else: - if router_response.responseCode == router_pb2.RouterResponseCode.RRC_BAD_REQUEST: - code = 'bad_request' - elif router_response.responseCode == router_pb2.RouterResponseCode.RRC_NOT_ALLOWED: - code = 'not_allowed' + + throttle_retries = 0 + while True: + response = requests.post(url, headers=headers, data=payload, verify=get_certificate_check(), + timeout=DEFAULT_TIMEOUT) + logger.debug('<<< [ROUTER] Response Code: [%d]', response.status_code) + + if response.status_code == 200: + rs_body = response.content + if rs_body: + router_response = router_pb2.RouterResponse() + router_response.ParseFromString(rs_body) + if router_response.responseCode == router_pb2.RouterResponseCode.RRC_OK: + if router_response.encryptedPayload: + return crypto.decrypt_aes_v2(router_response.encryptedPayload, transmission_key) else: - code = 'router_error' - raise errors.KeeperApiError(code, router_response.errorMessage) - return None - else: - message = response.reason - raise errors.KeeperApiError('router_error', f'{message}: {response.status_code}') + if router_response.responseCode == router_pb2.RouterResponseCode.RRC_BAD_REQUEST: + code = 'bad_request' + elif router_response.responseCode == router_pb2.RouterResponseCode.RRC_NOT_ALLOWED: + code = 'not_allowed' + else: + code = 'router_error' + raise errors.KeeperApiError(code, router_response.errorMessage) + return None + + _, error_code, error_message = parse_error_response(response) + if is_throttle_response(response.status_code, error_code): + throttle_retries = self._handle_throttle_and_retry( + status_code=response.status_code, + error_code=error_code, + error_message=error_message or response.reason or '', + headers=response.headers, + throttle_retries=throttle_retries) + continue + + raise errors.KeeperApiError( + error_code or 'router_error', + error_message or f'{response.reason}: {response.status_code}') def execute_router_bi(self, encryption_key: bytes, endpoint: str, request: Optional[TRQ], *, response_type: Type[TRS]) -> Optional[TRS]: @@ -196,21 +298,62 @@ def execute_router_bi(self, encryption_key: bytes, endpoint: str, request: Optio payload = crypto.encrypt_aes_v2(request.SerializeToString(), encryption_key) rq.payload = payload - response = requests.post(url, data=rq.SerializeToString()) - if response.status_code == 200: - rs_body = response.content - payload = crypto.decrypt_aes_v2(rs_body, encryption_key) - router_response = response_type() - router_response.ParseFromString(payload) - if logger.level <= logging.DEBUG: - js = MessageToJson(router_response) if router_response else '' - logger.debug('>>> [RS] \"%s\": %s', endpoint, js) - - return router_response - else: - message = response.reason - raise errors.KeeperApiError('router_error', f'{message}: {response.status_code}') - + throttle_retries = 0 + while True: + response = requests.post(url, data=rq.SerializeToString(), timeout=DEFAULT_TIMEOUT) + if response.status_code == 200: + rs_body = response.content + payload = crypto.decrypt_aes_v2(rs_body, encryption_key) + router_response = response_type() + router_response.ParseFromString(payload) + if logger.level <= logging.DEBUG: + js = MessageToJson(router_response) if router_response else '' + logger.debug('>>> [RS] \"%s\": %s', endpoint, js) + + return router_response + + _, error_code, error_message = parse_error_response(response) + if is_throttle_response(response.status_code, error_code): + throttle_retries = self._handle_throttle_and_retry( + status_code=response.status_code, + error_code=error_code, + error_message=error_message or response.reason or '', + headers=response.headers, + throttle_retries=throttle_retries) + continue + + raise errors.KeeperApiError( + error_code or 'router_error', + error_message or f'{response.reason}: {response.status_code}') + + + def _handle_throttle_and_retry(self, *, + status_code: int, + error_code: str, + error_message: str, + headers: Mapping[str, str], + throttle_retries: int) -> int: + """Sleep with backoff for a throttle / rate-limit response. + + Returns the updated retry count. Raises ``KeeperApiError`` when + ``fail_on_throttle`` is set or ``MAX_THROTTLE_RETRIES`` is exhausted. + """ + logger = utils.get_logger() + error_code = error_code or 'throttled' + error_message = error_message or 'Request was throttled' + if self.fail_on_throttle: + raise errors.KeeperApiError(error_code, error_message) + + throttle_retries += 1 + if throttle_retries > MAX_THROTTLE_RETRIES: + raise errors.KeeperApiError(error_code, error_message) + wait_seconds = parse_throttle_wait_seconds(error_message, headers) + backoff = throttle_backoff_seconds(throttle_retries, wait_seconds) + logger.warning( + 'Throttled (HTTP %d, attempt %d/%d), retrying in %d seconds: %s', + status_code, throttle_retries, MAX_THROTTLE_RETRIES, backoff, error_message) + time.sleep(backoff) + return throttle_retries def _communicate_keeper(self, endpoint: str, payload: Optional[bytes], @@ -219,10 +362,9 @@ def _communicate_keeper(self, endpoint: str, logger = utils.get_logger() transmission_key = utils.generate_aes_key() key_id = self.server_key_id - attempt = 0 - while attempt < 3: - attempt += 1 - + throttle_retries = 0 + key_retries = 0 + while True: api_request = prepare_api_request(key_id, transmission_key, payload, session_token=session_token, keeper_locale=self.locale, @@ -240,10 +382,9 @@ def _communicate_keeper(self, endpoint: str, 'User-Agent': 'KeeperSDK.Python/' + self.client_version } rs = requests.post(url, data=api_request.SerializeToString(), headers=headers, proxies=get_proxies(), - verify=get_certificate_check()) + verify=get_certificate_check(), timeout=DEFAULT_TIMEOUT) logger.debug('<<< Response Code: [%d]', rs.status_code) - content_type = rs.headers.get('Content-Type') or '' if rs.status_code == 200: if key_id != self._server_key_id: self._server_key_id = key_id @@ -255,31 +396,39 @@ def _communicate_keeper(self, endpoint: str, rs_body = rs.content return crypto.decrypt_aes_v2(rs_body, transmission_key) if rs_body else None - elif content_type.startswith('application/json'): - error_rs = rs.json() - if 'error' in error_rs: - error_code = error_rs['error'] - error_message = error_rs.get('message') or '' - additional_info = error_rs.get('additional_info') - if additional_info: - error_message += f'({additional_info})' - if error_code == 'key': - if 'key_id' in error_rs: - key_id = error_rs['key_id'] - continue - elif error_code == 'region_redirect': - raise errors.RegionRedirectError(error_rs['region_host'], error_message) - elif error_code == 'device_not_registered': - raise errors.InvalidDeviceTokenError(error_message) - elif error_code == 'throttled' and not self.fail_on_throttle: - logger.info('Throttled. sleeping for 10 seconds') - time.sleep(10) - continue - raise errors.KeeperApiError(error_code, error_message) - raise errors.KeeperApiError('http_error', f'{rs.reason}: {rs.status_code}') + error_rs, error_code, error_message = parse_error_response(rs) + if error_rs is not None: + logger.debug('<<< Response Error: [%s]', error_rs) + if error_code == 'key' and 'key_id' in error_rs: + key_retries += 1 + if key_retries > MAX_KEY_RETRIES: + raise errors.KeeperApiError(error_code, error_message or 'Invalid server public key') + key_id = error_rs['key_id'] + continue + if error_code == 'region_redirect': + raise errors.RegionRedirectError(error_rs.get('region_host') or '', error_message) + if error_code == 'device_not_registered': + raise errors.InvalidDeviceTokenError(error_message) + + if is_throttle_response(rs.status_code, error_code): + throttle_retries = self._handle_throttle_and_retry( + status_code=rs.status_code, + error_code=error_code, + error_message=error_message or rs.reason or '', + headers=rs.headers, + throttle_retries=throttle_retries) + continue + + if error_code: + raise errors.KeeperApiError(error_code, error_message) - raise errors.KeeperError('Failed to execute Keeper API request') + if logger.level <= logging.DEBUG: + if rs.text: + logger.debug('<<< Response Content: [%s]', rs.text) + else: + logger.debug('<<< HTTP Status: [%s] Reason: [%s]', rs.status_code, rs.reason) + raise errors.KeeperApiError('http_error', f'{rs.reason}: {rs.status_code}') def execute_rest(self, rest_endpoint: str, request: Optional[TRQ], diff --git a/keepersdk-package/src/keepersdk/authentication/keeper_auth.py b/keepersdk-package/src/keepersdk/authentication/keeper_auth.py index 5fbad726..0cc82a4a 100644 --- a/keepersdk-package/src/keepersdk/authentication/keeper_auth.py +++ b/keepersdk-package/src/keepersdk/authentication/keeper_auth.py @@ -187,6 +187,7 @@ def execute_batch(self, requests: List[Dict[str, Any]]) -> List[Dict[str, Any]]: sleep_interval = 0 chunk_size = 200 queue = requests.copy() + throttle_retries = 0 while len(queue) > 0: if sleep_interval > 0: time.sleep(sleep_interval) @@ -200,16 +201,30 @@ def execute_batch(self, requests: List[Dict[str, Any]]) -> List[Dict[str, Any]]: } rs = self.execute_auth_command(rq) results = rs.get('results') - if isinstance(results, list) and len(results) > 0: - error_status = results[-1] - throttled = error_status.get('result') != 'success' and error_status.get('result_code') == 'throttled' - if throttled: - sleep_interval = 10 - results.pop() - responses.extend(results) - - if len(results) < len(chunk): - queue = chunk[len(results):] + queue + if not isinstance(results, list) or len(results) == 0: + # No result for any queued request: re-queueing would spin forever + raise errors.KeeperApiError('server_error', 'Batch execution returned no results') + + error_status = results[-1] + throttled = error_status.get('result') != 'success' and error_status.get('result_code') == 'throttled' + if throttled: + throttle_retries += 1 + if self.keeper_endpoint.fail_on_throttle or throttle_retries > endpoint.MAX_THROTTLE_RETRIES: + raise errors.KeeperApiError( + error_status.get('result_code') or 'throttled', + error_status.get('message') or 'Request was throttled') + wait_seconds = endpoint.parse_throttle_wait_seconds(error_status.get('message') or '') + sleep_interval = endpoint.throttle_backoff_seconds(throttle_retries, wait_seconds) + utils.get_logger().warning( + 'Batch throttled (attempt %d/%d), retrying in %d seconds', + throttle_retries, endpoint.MAX_THROTTLE_RETRIES, sleep_interval) + results.pop() + else: + throttle_retries = 0 + responses.extend(results) + + if len(results) < len(chunk): + queue = chunk[len(results):] + queue return responses def execute_router(self, path: str, request: Optional[endpoint.TRQ], *, diff --git a/keepersdk-package/unit_tests/test_throttle_handling.py b/keepersdk-package/unit_tests/test_throttle_handling.py new file mode 100644 index 00000000..43f657f2 --- /dev/null +++ b/keepersdk-package/unit_tests/test_throttle_handling.py @@ -0,0 +1,250 @@ +import unittest +from datetime import datetime, timedelta, timezone +from email.utils import format_datetime +from unittest.mock import MagicMock, patch + +from keepersdk import errors +from keepersdk.authentication import endpoint +from keepersdk.authentication.keeper_auth import KeeperAuth, AuthContext + + +class TestThrottleHelpers(unittest.TestCase): + def test_parse_wait_from_message_seconds(self): + self.assertEqual(endpoint.parse_throttle_wait_seconds('Please wait 45 seconds'), 45) + + def test_parse_wait_from_message_minutes(self): + self.assertEqual(endpoint.parse_throttle_wait_seconds('Retry after 2 minutes'), 120) + + def test_parse_wait_defaults_when_message_has_no_duration(self): + self.assertEqual( + endpoint.parse_throttle_wait_seconds('Too many requests'), + endpoint.DEFAULT_THROTTLE_WAIT_SECONDS) + + def test_parse_wait_prefers_retry_after_header(self): + self.assertEqual( + endpoint.parse_throttle_wait_seconds('wait 10 seconds', {'Retry-After': '90'}), + 90) + + def test_parse_wait_accepts_retry_after_http_date(self): + retry_at = datetime.now(timezone.utc) + timedelta(seconds=120) + wait = endpoint.parse_throttle_wait_seconds('', {'Retry-After': format_datetime(retry_at)}) + self.assertGreaterEqual(wait, 110) + self.assertLessEqual(wait, 120) + + def test_parse_wait_ignores_malformed_retry_after(self): + self.assertEqual( + endpoint.parse_throttle_wait_seconds('wait 15 seconds', {'Retry-After': 'soon'}), + 15) + + def test_parse_wait_caps_at_max(self): + self.assertEqual( + endpoint.parse_throttle_wait_seconds('wait 10 minutes'), + endpoint.MAX_THROTTLE_WAIT_SECONDS) + self.assertEqual( + endpoint.parse_throttle_wait_seconds('', {'Retry-After': '99999'}), + endpoint.MAX_THROTTLE_WAIT_SECONDS) + + def test_backoff_grows_exponentially(self): + self.assertEqual(endpoint.throttle_backoff_seconds(1, 10), 30) + self.assertEqual(endpoint.throttle_backoff_seconds(2, 10), 60) + self.assertEqual(endpoint.throttle_backoff_seconds(3, 10), 120) + self.assertEqual(endpoint.throttle_backoff_seconds(1, 90), 90) + + def test_is_throttle_response(self): + self.assertTrue(endpoint.is_throttle_response(403, 'throttled')) + self.assertTrue(endpoint.is_throttle_response(429, '')) + self.assertFalse(endpoint.is_throttle_response(403, 'access_denied')) + self.assertFalse(endpoint.is_throttle_response(400, '')) + + +class TestCommunicateKeeperThrottle(unittest.TestCase): + def _make_endpoint(self): + storage = MagicMock() + storage.get.return_value = MagicMock(last_server='keepersecurity.com', servers=MagicMock(return_value={})) + ep = endpoint.KeeperEndpoint(storage, keeper_server='keepersecurity.com') + ep.fail_on_throttle = False + return ep + + def _json_response(self, status_code, body, headers=None): + rs = MagicMock() + rs.status_code = status_code + rs.headers = {'Content-Type': 'application/json', **(headers or {})} + rs.json.return_value = body + rs.reason = 'Forbidden' if status_code == 403 else 'Too Many Requests' + rs.content = b'' + rs.text = '' + return rs + + def _raw_response(self, status_code, reason, text='', headers=None): + rs = MagicMock() + rs.status_code = status_code + rs.headers = headers or {} + rs.reason = reason + rs.text = text + rs.content = b'' + rs.json.side_effect = ValueError('no json') + return rs + + @patch('keepersdk.authentication.endpoint.time.sleep') + @patch('keepersdk.authentication.endpoint.requests.post') + @patch('keepersdk.authentication.endpoint.prepare_api_request') + def test_403_throttled_retries_then_raises(self, mock_prepare, mock_post, mock_sleep): + mock_prepare.return_value = MagicMock(SerializeToString=MagicMock(return_value=b'rq')) + throttle_body = {'error': 'throttled', 'message': 'Please wait 1 second'} + ok = MagicMock() + ok.status_code = 200 + ok.headers = {} + ok.content = b'' + + # Fail twice with throttle, succeed on third + mock_post.side_effect = [ + self._json_response(403, throttle_body), + self._json_response(403, throttle_body), + ok, + ] + ep = self._make_endpoint() + result = ep._communicate_keeper('vault/test', b'payload') + self.assertIsNone(result) + self.assertEqual(mock_post.call_count, 3) + self.assertEqual(mock_sleep.call_count, 2) + + @patch('keepersdk.authentication.endpoint.time.sleep') + @patch('keepersdk.authentication.endpoint.requests.post') + @patch('keepersdk.authentication.endpoint.prepare_api_request') + def test_429_retries_then_raises(self, mock_prepare, mock_post, mock_sleep): + mock_prepare.return_value = MagicMock(SerializeToString=MagicMock(return_value=b'rq')) + throttle_rs = MagicMock() + throttle_rs.status_code = 429 + throttle_rs.headers = {'Retry-After': '1'} + throttle_rs.reason = 'Too Many Requests' + throttle_rs.content = b'' + throttle_rs.text = '' + + mock_post.side_effect = [throttle_rs] * (endpoint.MAX_THROTTLE_RETRIES + 1) + ep = self._make_endpoint() + with self.assertRaises(errors.KeeperApiError) as ctx: + ep._communicate_keeper('vault/test', b'payload') + self.assertEqual(ctx.exception.result_code, 'throttled') + self.assertEqual(mock_sleep.call_count, endpoint.MAX_THROTTLE_RETRIES) + + @patch('keepersdk.authentication.endpoint.requests.post') + @patch('keepersdk.authentication.endpoint.prepare_api_request') + def test_fail_on_throttle_raises_immediately(self, mock_prepare, mock_post): + mock_prepare.return_value = MagicMock(SerializeToString=MagicMock(return_value=b'rq')) + mock_post.return_value = self._json_response(403, {'error': 'throttled', 'message': 'slow down'}) + ep = self._make_endpoint() + ep.fail_on_throttle = True + with self.assertRaises(errors.KeeperApiError) as ctx: + ep._communicate_keeper('vault/test', b'payload') + self.assertEqual(ctx.exception.result_code, 'throttled') + self.assertEqual(mock_post.call_count, 1) + + @patch('keepersdk.authentication.endpoint.requests.post') + @patch('keepersdk.authentication.endpoint.prepare_api_request') + def test_403_non_throttle_error_is_not_retried(self, mock_prepare, mock_post): + mock_prepare.return_value = MagicMock(SerializeToString=MagicMock(return_value=b'rq')) + mock_post.return_value = self._json_response( + 403, {'error': 'access_denied', 'message': 'Not permitted'}) + ep = self._make_endpoint() + with self.assertRaises(errors.KeeperApiError) as ctx: + ep._communicate_keeper('vault/test', b'payload') + self.assertEqual(ctx.exception.result_code, 'access_denied') + self.assertEqual(ctx.exception.message, 'Not permitted') + self.assertEqual(mock_post.call_count, 1) + + @patch('keepersdk.authentication.endpoint.time.sleep') + @patch('keepersdk.authentication.endpoint.requests.post') + @patch('keepersdk.authentication.endpoint.prepare_api_request') + def test_429_without_json_body_is_retried(self, mock_prepare, mock_post, mock_sleep): + mock_prepare.return_value = MagicMock(SerializeToString=MagicMock(return_value=b'rq')) + ok = MagicMock() + ok.status_code = 200 + ok.headers = {} + ok.content = b'' + + mock_post.side_effect = [ + self._raw_response(429, 'Too Many Requests', 'rate limited'), + ok, + ] + ep = self._make_endpoint() + self.assertIsNone(ep._communicate_keeper('vault/test', b'payload')) + self.assertEqual(mock_sleep.call_count, 1) + + @patch('keepersdk.authentication.endpoint.requests.post') + @patch('keepersdk.authentication.endpoint.prepare_api_request') + def test_malformed_json_error_body_raises_http_error(self, mock_prepare, mock_post): + mock_prepare.return_value = MagicMock(SerializeToString=MagicMock(return_value=b'rq')) + rs = self._raw_response(500, 'Internal Server Error', 'not json', + headers={'Content-Type': 'application/json'}) + mock_post.return_value = rs + ep = self._make_endpoint() + with self.assertRaises(errors.KeeperApiError) as ctx: + ep._communicate_keeper('vault/test', b'payload') + self.assertEqual(ctx.exception.result_code, 'http_error') + + @patch('keepersdk.authentication.endpoint.requests.post') + @patch('keepersdk.authentication.endpoint.prepare_api_request') + def test_key_rotation_is_bounded(self, mock_prepare, mock_post): + mock_prepare.return_value = MagicMock(SerializeToString=MagicMock(return_value=b'rq')) + mock_post.return_value = self._json_response(401, {'error': 'key', 'key_id': 8}) + ep = self._make_endpoint() + with self.assertRaises(errors.KeeperApiError) as ctx: + ep._communicate_keeper('vault/test', b'payload') + self.assertEqual(ctx.exception.result_code, 'key') + self.assertEqual(mock_post.call_count, endpoint.MAX_KEY_RETRIES + 1) + + +class TestExecuteBatchThrottle(unittest.TestCase): + @staticmethod + def _make_auth(): + keeper_endpoint = MagicMock() + keeper_endpoint.fail_on_throttle = False + return KeeperAuth(keeper_endpoint, AuthContext()) + + @patch('keepersdk.authentication.keeper_auth.time.sleep') + def test_batch_throttled_retries_with_backoff(self, mock_sleep): + auth = self._make_auth() + throttle = {'result': 'fail', 'result_code': 'throttled', 'message': 'wait 1 second'} + success = {'result': 'success'} + + # First chunk: one success + trailing throttle; retry returns remaining success + auth.execute_auth_command = MagicMock(side_effect=[ + {'results': [success, throttle]}, + {'results': [success]}, + ]) + responses = auth.execute_batch([{'command': 'a'}, {'command': 'b'}]) + self.assertEqual(len(responses), 2) + self.assertEqual(mock_sleep.call_count, 1) + + @staticmethod + def _always_throttled(*_args, **_kwargs): + return {'results': [{'result': 'fail', 'result_code': 'throttled', 'message': 'slow down'}]} + + @patch('keepersdk.authentication.keeper_auth.time.sleep') + def test_batch_raises_after_max_throttle_retries(self, mock_sleep): + auth = self._make_auth() + auth.execute_auth_command = MagicMock(side_effect=self._always_throttled) + with self.assertRaises(errors.KeeperApiError) as ctx: + auth.execute_batch([{'command': 'a'}]) + self.assertEqual(ctx.exception.result_code, 'throttled') + self.assertEqual(mock_sleep.call_count, endpoint.MAX_THROTTLE_RETRIES) + + def test_batch_fail_on_throttle_raises_immediately(self): + auth = self._make_auth() + auth.keeper_endpoint.fail_on_throttle = True + auth.execute_auth_command = MagicMock(side_effect=self._always_throttled) + with self.assertRaises(errors.KeeperApiError) as ctx: + auth.execute_batch([{'command': 'a'}]) + self.assertEqual(ctx.exception.result_code, 'throttled') + self.assertEqual(auth.execute_auth_command.call_count, 1) + + def test_batch_empty_results_raises_instead_of_dropping_requests(self): + auth = self._make_auth() + auth.execute_auth_command = MagicMock(return_value={'results': []}) + with self.assertRaises(errors.KeeperApiError) as ctx: + auth.execute_batch([{'command': 'a'}]) + self.assertEqual(ctx.exception.result_code, 'server_error') + + +if __name__ == '__main__': + unittest.main()