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()