test_token_refresh.py 3.2 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495
  1. import unittest
  2. from datetime import datetime
  3. from unittest.mock import Mock, call, patch
  4. import requests
  5. import main
  6. class QueryHourDataWithTokenRefreshTests(unittest.TestCase):
  7. def setUp(self):
  8. self.energy_config = {"base_url": "http://example.test"}
  9. self.token_manager = Mock(spec=main.TokenManager)
  10. self.token_manager.get_token.return_value = "old-token"
  11. self.token_manager.refresh_token.return_value = "new-token"
  12. self.tag_ids = ["tag-1"]
  13. self.start_time = datetime(2026, 7, 16, 16, 0, 0)
  14. self.end_time = datetime(2026, 7, 16, 16, 59, 59)
  15. @staticmethod
  16. def http_error(status_code):
  17. response = requests.Response()
  18. response.status_code = status_code
  19. response.url = "http://example.test/query"
  20. return requests.HTTPError(response=response)
  21. def test_success_does_not_refresh_token(self):
  22. result = {"code": 0}
  23. with patch("main.query_hour_data", return_value=result):
  24. actual = main.query_hour_data_with_token_refresh(
  25. self.energy_config,
  26. self.token_manager,
  27. self.tag_ids,
  28. self.start_time,
  29. self.end_time,
  30. )
  31. self.assertIs(actual, result)
  32. self.token_manager.refresh_token.assert_not_called()
  33. def test_401_refreshes_token_and_retries_once(self):
  34. result = {"code": 0}
  35. with patch(
  36. "main.query_hour_data",
  37. side_effect=[self.http_error(401), result],
  38. ) as query:
  39. actual = main.query_hour_data_with_token_refresh(
  40. self.energy_config,
  41. self.token_manager,
  42. self.tag_ids,
  43. self.start_time,
  44. self.end_time,
  45. )
  46. self.assertIs(actual, result)
  47. self.token_manager.refresh_token.assert_called_once_with()
  48. query.assert_has_calls(
  49. [
  50. call(self.energy_config, "old-token", self.tag_ids, self.start_time, self.end_time),
  51. call(self.energy_config, "new-token", self.tag_ids, self.start_time, self.end_time),
  52. ]
  53. )
  54. def test_second_401_is_raised_without_another_refresh(self):
  55. with patch(
  56. "main.query_hour_data",
  57. side_effect=[self.http_error(401), self.http_error(401)],
  58. ):
  59. with self.assertRaises(requests.HTTPError):
  60. main.query_hour_data_with_token_refresh(
  61. self.energy_config,
  62. self.token_manager,
  63. self.tag_ids,
  64. self.start_time,
  65. self.end_time,
  66. )
  67. self.token_manager.refresh_token.assert_called_once_with()
  68. def test_non_401_error_does_not_refresh_token(self):
  69. with patch("main.query_hour_data", side_effect=self.http_error(500)):
  70. with self.assertRaises(requests.HTTPError):
  71. main.query_hour_data_with_token_refresh(
  72. self.energy_config,
  73. self.token_manager,
  74. self.tag_ids,
  75. self.start_time,
  76. self.end_time,
  77. )
  78. self.token_manager.refresh_token.assert_not_called()
  79. if __name__ == "__main__":
  80. unittest.main()