Browse Source

修复查询数据时token失效问题

Lu Xianghui 2 weeks ago
parent
commit
04da4a4511
3 changed files with 121 additions and 3 deletions
  1. 1 1
      config.yaml
  2. 25 2
      main.py
  3. 95 0
      tests/test_token_refresh.py

+ 1 - 1
config.yaml

@@ -14,7 +14,7 @@ services:
   calcagg: "http://192.168.1.109:42941"
 
 schedule:
-  minute: 14
+  minute: 25
 
 logging:
   file: "logs/jdf-energy-collector.log"

+ 25 - 2
main.py

@@ -221,6 +221,9 @@ class TokenManager:
             LOGGER.info("复用登录 Token,下次登录时间: %s", format_api_time(self.last_login_at + self.refresh_interval))
             return self.token
 
+        return self.refresh_token()
+
+    def refresh_token(self) -> str:
         LOGGER.info("开始获取登录 Token")
         self.token = login(self.energy_config)
         self.last_login_at = datetime.now()
@@ -261,6 +264,25 @@ def query_hour_data(energy_config: dict, token: str, tag_ids: list[str], start_t
     raise RuntimeError("查询接口 timeout,已达到最大重试次数")
 
 
+def query_hour_data_with_token_refresh(
+    energy_config: dict,
+    token_manager: TokenManager,
+    tag_ids: list[str],
+    start_time: datetime,
+    end_time: datetime,
+) -> dict:
+    token = token_manager.get_token()
+    try:
+        return query_hour_data(energy_config, token, tag_ids, start_time, end_time)
+    except requests.HTTPError as exc:
+        if exc.response is None or exc.response.status_code != 401:
+            raise
+
+    LOGGER.warning("查询接口返回 401,重新登录获取 Token 后重试")
+    token = token_manager.refresh_token()
+    return query_hour_data(energy_config, token, tag_ids, start_time, end_time)
+
+
 def extract_records(result: dict) -> list[dict]:
     records = []
     data_root = result.get("data") or {}
@@ -351,8 +373,9 @@ def run_once(config: dict, token_manager: TokenManager, tag_ids: list[str], id_t
         end_time,
     )
 
-    token = token_manager.get_token()
-    result = query_hour_data(config["energy_chaowang"], token, tag_ids, start_time, end_time)
+    result = query_hour_data_with_token_refresh(
+        config["energy_chaowang"], token_manager, tag_ids, start_time, end_time
+    )
     records = extract_records(result)
 
     written_point_ids = call_addpointdatum(config["services"]["basedataportal"], records, id_to_point_id)

+ 95 - 0
tests/test_token_refresh.py

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