| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495 |
- import unittest
- from datetime import datetime
- from unittest.mock import Mock, call, patch
- import requests
- import main
- class QueryHourDataWithTokenRefreshTests(unittest.TestCase):
- def setUp(self):
- self.energy_config = {"base_url": "http://example.test"}
- self.token_manager = Mock(spec=main.TokenManager)
- self.token_manager.get_token.return_value = "old-token"
- self.token_manager.refresh_token.return_value = "new-token"
- self.tag_ids = ["tag-1"]
- self.start_time = datetime(2026, 7, 16, 16, 0, 0)
- self.end_time = datetime(2026, 7, 16, 16, 59, 59)
- @staticmethod
- def http_error(status_code):
- response = requests.Response()
- response.status_code = status_code
- response.url = "http://example.test/query"
- return requests.HTTPError(response=response)
- def test_success_does_not_refresh_token(self):
- result = {"code": 0}
- with patch("main.query_hour_data", return_value=result):
- actual = main.query_hour_data_with_token_refresh(
- self.energy_config,
- self.token_manager,
- self.tag_ids,
- self.start_time,
- self.end_time,
- )
- self.assertIs(actual, result)
- self.token_manager.refresh_token.assert_not_called()
- def test_401_refreshes_token_and_retries_once(self):
- result = {"code": 0}
- with patch(
- "main.query_hour_data",
- side_effect=[self.http_error(401), result],
- ) as query:
- actual = main.query_hour_data_with_token_refresh(
- self.energy_config,
- self.token_manager,
- self.tag_ids,
- self.start_time,
- self.end_time,
- )
- self.assertIs(actual, result)
- self.token_manager.refresh_token.assert_called_once_with()
- query.assert_has_calls(
- [
- call(self.energy_config, "old-token", self.tag_ids, self.start_time, self.end_time),
- call(self.energy_config, "new-token", self.tag_ids, self.start_time, self.end_time),
- ]
- )
- def test_second_401_is_raised_without_another_refresh(self):
- with patch(
- "main.query_hour_data",
- side_effect=[self.http_error(401), self.http_error(401)],
- ):
- with self.assertRaises(requests.HTTPError):
- main.query_hour_data_with_token_refresh(
- self.energy_config,
- self.token_manager,
- self.tag_ids,
- self.start_time,
- self.end_time,
- )
- self.token_manager.refresh_token.assert_called_once_with()
- def test_non_401_error_does_not_refresh_token(self):
- with patch("main.query_hour_data", side_effect=self.http_error(500)):
- with self.assertRaises(requests.HTTPError):
- main.query_hour_data_with_token_refresh(
- self.energy_config,
- self.token_manager,
- self.tag_ids,
- self.start_time,
- self.end_time,
- )
- self.token_manager.refresh_token.assert_not_called()
- if __name__ == "__main__":
- unittest.main()
|