|
|
@@ -0,0 +1,95 @@
|
|
|
+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()
|