test_token_store_coverage.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341
  1. """
  2. Coverage补充测试:services/token_store.py
  3. 目标路径:
  4. - InMemoryTokenStore.is_expiring — 无session、无expire_time、Z后缀、过期中、未过期
  5. - InMemoryTokenStore.require_token — 无session、token为空、正常返回
  6. - RedisTokenStore — 初始化校验、save(含/不含TTL)、get(bytes/str/空串/None)、clear、is_expiring
  7. - RedisTokenStore._ttl_seconds — None、Z后缀、naive datetime
  8. - RedisSocketClient._read_response — +/-/:/$/$(nil)/*(含空数组)/unsupported/empty
  9. - RedisSocketClient._execute — 带 password+db、不带 password+db
  10. - RedisSocketClient.set/get/delete(mocked socket)
  11. """
  12. import io
  13. import json
  14. import unittest
  15. from datetime import datetime, timedelta, timezone
  16. from unittest.mock import MagicMock, patch
  17. from services.token_store import (
  18. FileTokenStore,
  19. InMemoryTokenStore,
  20. RedisSocketClient,
  21. RedisTokenStore,
  22. )
  23. # ---------------------------------------------------------------------------
  24. # InMemoryTokenStore
  25. # ---------------------------------------------------------------------------
  26. class InMemoryIsExpiringTest(unittest.TestCase):
  27. def test_no_session_is_expiring(self):
  28. store = InMemoryTokenStore()
  29. self.assertTrue(store.is_expiring())
  30. def test_session_with_none_expire_time_is_expiring(self):
  31. store = InMemoryTokenStore()
  32. store._session = {'token': 'T', 'expire_time': None}
  33. self.assertTrue(store.is_expiring())
  34. def test_session_with_empty_expire_time_is_expiring(self):
  35. store = InMemoryTokenStore()
  36. store._session = {'token': 'T', 'expire_time': ''}
  37. self.assertTrue(store.is_expiring())
  38. def test_expire_soon_within_skew_is_expiring(self):
  39. store = InMemoryTokenStore(refresh_skew_seconds=120)
  40. # 60 秒后到期,在 120s 的 skew 窗口内 → 算作 expiring
  41. soon = (datetime.now(timezone.utc) + timedelta(seconds=60)).strftime('%Y-%m-%dT%H:%M:%SZ')
  42. store._session = {'token': 'T', 'expire_time': soon}
  43. self.assertTrue(store.is_expiring())
  44. def test_expire_far_future_not_expiring(self):
  45. store = InMemoryTokenStore(refresh_skew_seconds=120)
  46. store._session = {'token': 'T', 'expire_time': '2099-01-01T00:00:00Z'}
  47. self.assertFalse(store.is_expiring())
  48. def test_expire_time_z_suffix_parsed_correctly(self):
  49. store = InMemoryTokenStore(refresh_skew_seconds=0)
  50. # 已过期(1秒前)
  51. past = (datetime.now(timezone.utc) - timedelta(seconds=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
  52. store._session = {'token': 'T', 'expire_time': past}
  53. self.assertTrue(store.is_expiring())
  54. def test_expire_time_without_z_suffix(self):
  55. store = InMemoryTokenStore(refresh_skew_seconds=0)
  56. # 无时区后缀的 ISO 格式(naive datetime)
  57. past = (datetime.now() - timedelta(seconds=1)).strftime('%Y-%m-%dT%H:%M:%S')
  58. store._session = {'token': 'T', 'expire_time': past}
  59. self.assertTrue(store.is_expiring())
  60. class InMemoryRequireTokenTest(unittest.TestCase):
  61. def test_raises_when_no_session(self):
  62. store = InMemoryTokenStore()
  63. with self.assertRaises(RuntimeError) as ctx:
  64. store.require_token()
  65. self.assertIn('mcp token missing', str(ctx.exception))
  66. def test_raises_when_token_is_empty_string(self):
  67. store = InMemoryTokenStore()
  68. store._session = {'token': '', 'expire_time': '2099-01-01T00:00:00'}
  69. with self.assertRaises(RuntimeError):
  70. store.require_token()
  71. def test_returns_token_when_present(self):
  72. store = InMemoryTokenStore()
  73. store.save('MT_abc123', '2099-01-01T00:00:00')
  74. self.assertEqual('MT_abc123', store.require_token())
  75. # ---------------------------------------------------------------------------
  76. # RedisTokenStore(使用 MagicMock client 隔离真实 Redis)
  77. # ---------------------------------------------------------------------------
  78. class RedisTokenStoreMockedTest(unittest.TestCase):
  79. def _make_store(self, raw_get=None, prefix='test:', session_key='session_1'):
  80. client = MagicMock()
  81. client.get.return_value = raw_get
  82. store = RedisTokenStore(client, prefix=prefix, session_key=session_key)
  83. return store, client
  84. # 初始化
  85. def test_empty_session_key_raises_value_error(self):
  86. client = MagicMock()
  87. with self.assertRaises(ValueError) as ctx:
  88. RedisTokenStore(client, session_key='')
  89. self.assertIn('session_key is required', str(ctx.exception))
  90. def test_whitespace_session_key_raises_value_error(self):
  91. client = MagicMock()
  92. with self.assertRaises(ValueError):
  93. RedisTokenStore(client, session_key=' ')
  94. def test_key_composed_from_prefix_and_session_key(self):
  95. store, _ = self._make_store(prefix='ns:', session_key='abc')
  96. self.assertEqual('ns:abc', store.key)
  97. # save
  98. def test_save_calls_client_set_with_positive_ttl(self):
  99. store, client = self._make_store()
  100. store.save('MT_001', '2099-01-01T00:00:00Z')
  101. self.assertTrue(client.set.called)
  102. args, kwargs = client.set.call_args
  103. self.assertEqual('test:session_1', args[0])
  104. parsed = json.loads(args[1])
  105. self.assertEqual('MT_001', parsed['token'])
  106. self.assertIsNotNone(kwargs.get('ex'))
  107. self.assertGreater(kwargs['ex'], 0)
  108. def test_save_with_none_expire_time_passes_no_ex(self):
  109. store, client = self._make_store()
  110. store.save('MT_001', None)
  111. _, kwargs = client.set.call_args
  112. self.assertIsNone(kwargs.get('ex'))
  113. # get
  114. def test_get_returns_none_when_redis_returns_none(self):
  115. store, _ = self._make_store(raw_get=None)
  116. self.assertIsNone(store.get())
  117. def test_get_returns_none_on_empty_string(self):
  118. store, _ = self._make_store(raw_get='')
  119. self.assertIsNone(store.get())
  120. def test_get_decodes_bytes(self):
  121. data = {'token': 'MT_bytes', 'expire_time': '2099-01-01T00:00:00'}
  122. store, _ = self._make_store(raw_get=json.dumps(data).encode('utf-8'))
  123. result = store.get()
  124. self.assertEqual('MT_bytes', result['token'])
  125. def test_get_handles_string_response(self):
  126. data = {'token': 'MT_str', 'expire_time': '2099-01-01T00:00:00'}
  127. store, _ = self._make_store(raw_get=json.dumps(data))
  128. result = store.get()
  129. self.assertEqual('MT_str', result['token'])
  130. # clear
  131. def test_clear_calls_client_delete_with_correct_key(self):
  132. store, client = self._make_store()
  133. store.clear()
  134. client.delete.assert_called_once_with('test:session_1')
  135. self.assertIsNone(store._session)
  136. # is_expiring — refreshes from Redis first
  137. def test_is_expiring_reads_from_redis_first(self):
  138. client = MagicMock()
  139. data = {'token': 'MT_live', 'expire_time': '2099-01-01T00:00:00'}
  140. client.get.return_value = json.dumps(data)
  141. store = RedisTokenStore(client, prefix='t:', session_key='s')
  142. self.assertFalse(store.is_expiring())
  143. self.assertTrue(client.get.called)
  144. def test_is_expiring_true_when_redis_has_no_data(self):
  145. store, _ = self._make_store(raw_get=None)
  146. self.assertTrue(store.is_expiring())
  147. class RedisTtlSecondsTest(unittest.TestCase):
  148. def test_none_returns_none(self):
  149. self.assertIsNone(RedisTokenStore._ttl_seconds(None))
  150. def test_empty_string_returns_none(self):
  151. self.assertIsNone(RedisTokenStore._ttl_seconds(''))
  152. def test_z_suffix_parsed_as_utc(self):
  153. future = (datetime.now(timezone.utc) + timedelta(hours=2)).strftime('%Y-%m-%dT%H:%M:%SZ')
  154. ttl = RedisTokenStore._ttl_seconds(future)
  155. self.assertGreater(ttl, 3600)
  156. def test_naive_datetime_uses_local_time(self):
  157. future = (datetime.now() + timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M:%S')
  158. ttl = RedisTokenStore._ttl_seconds(future)
  159. self.assertIsNotNone(ttl)
  160. self.assertGreater(ttl, 0)
  161. def test_already_past_returns_1_at_minimum(self):
  162. # 已过期的时间仍返回 max(1, ...)
  163. past = (datetime.now(timezone.utc) - timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
  164. ttl = RedisTokenStore._ttl_seconds(past)
  165. self.assertEqual(1, ttl)
  166. # ---------------------------------------------------------------------------
  167. # RedisSocketClient._read_response(直接传入 BytesIO 流,不需要真实 socket)
  168. # ---------------------------------------------------------------------------
  169. class RedisReadResponseTest(unittest.TestCase):
  170. """_read_response 是实例方法,通过实例调用。"""
  171. def _client(self):
  172. return RedisSocketClient()
  173. def _stream(self, data: bytes) -> io.BytesIO:
  174. return io.BytesIO(data)
  175. def test_simple_string_ok(self):
  176. result = self._client()._read_response(self._stream(b'+OK\r\n'))
  177. self.assertEqual('OK', result)
  178. def test_simple_string_pong(self):
  179. result = self._client()._read_response(self._stream(b'+PONG\r\n'))
  180. self.assertEqual('PONG', result)
  181. def test_error_raises_runtime_error(self):
  182. with self.assertRaises(RuntimeError) as ctx:
  183. self._client()._read_response(self._stream(b'-ERR unknown command\r\n'))
  184. self.assertIn('ERR unknown command', str(ctx.exception))
  185. def test_integer(self):
  186. result = self._client()._read_response(self._stream(b':7\r\n'))
  187. self.assertEqual(7, result)
  188. def test_integer_zero(self):
  189. result = self._client()._read_response(self._stream(b':0\r\n'))
  190. self.assertEqual(0, result)
  191. def test_bulk_string(self):
  192. result = self._client()._read_response(self._stream(b'$5\r\nhello\r\n'))
  193. self.assertEqual('hello', result)
  194. def test_bulk_string_nil(self):
  195. result = self._client()._read_response(self._stream(b'$-1\r\n'))
  196. self.assertIsNone(result)
  197. def test_array_two_elements(self):
  198. result = self._client()._read_response(self._stream(b'*2\r\n+alpha\r\n+beta\r\n'))
  199. self.assertEqual(['alpha', 'beta'], result)
  200. def test_array_empty(self):
  201. result = self._client()._read_response(self._stream(b'*0\r\n'))
  202. self.assertEqual([], result)
  203. def test_array_with_bulk_strings(self):
  204. result = self._client()._read_response(self._stream(b'*2\r\n$3\r\nfoo\r\n$3\r\nbar\r\n'))
  205. self.assertEqual(['foo', 'bar'], result)
  206. def test_unsupported_prefix_raises(self):
  207. with self.assertRaises(RuntimeError) as ctx:
  208. self._client()._read_response(self._stream(b'?nope\r\n'))
  209. self.assertIn('unsupported redis response', str(ctx.exception))
  210. def test_empty_response_raises(self):
  211. with self.assertRaises(RuntimeError) as ctx:
  212. self._client()._read_response(self._stream(b''))
  213. self.assertIn('empty redis response', str(ctx.exception))
  214. # ---------------------------------------------------------------------------
  215. # RedisSocketClient._execute(用 patch mock socket,覆盖 AUTH/SELECT 分支)
  216. # ---------------------------------------------------------------------------
  217. class RedisSocketClientExecuteTest(unittest.TestCase):
  218. """通过 mock socket 验证 _execute 中 AUTH/SELECT 的条件分支。"""
  219. def _mock_sock(self, responses):
  220. """
  221. responses: list[bytes],每个是一次 _read_response 调用的完整 RESP 响应。
  222. """
  223. stream = io.BytesIO(b''.join(responses))
  224. sock = MagicMock()
  225. sock.makefile.return_value = stream
  226. sock.__enter__ = lambda s: s
  227. sock.__exit__ = MagicMock(return_value=False)
  228. return sock
  229. def test_execute_with_password_sends_auth(self):
  230. """password 非空 → 发送 AUTH 命令。"""
  231. client = RedisSocketClient(password='secret', db=0)
  232. # AUTH → OK;GET → nil
  233. mock_sock = self._mock_sock([b'+OK\r\n', b'$-1\r\n'])
  234. with patch('socket.create_connection', return_value=mock_sock):
  235. result = client.get('k')
  236. self.assertIsNone(result)
  237. # sendall 应被调用多次(AUTH + GET)
  238. self.assertGreater(mock_sock.sendall.call_count, 0)
  239. def test_execute_with_db_sends_select(self):
  240. """db 非0 → 发送 SELECT 命令。"""
  241. client = RedisSocketClient(password='', db=3)
  242. # SELECT → OK;GET → value
  243. mock_sock = self._mock_sock([b'+OK\r\n', b'$5\r\nvalue\r\n'])
  244. with patch('socket.create_connection', return_value=mock_sock):
  245. result = client.get('k')
  246. self.assertEqual('value', result)
  247. def test_execute_with_password_and_db(self):
  248. """password 非空 且 db 非0 → 发送 AUTH + SELECT。"""
  249. client = RedisSocketClient(password='pass', db=2)
  250. # AUTH → OK;SELECT → OK;SET → OK
  251. mock_sock = self._mock_sock([b'+OK\r\n', b'+OK\r\n', b'+OK\r\n'])
  252. with patch('socket.create_connection', return_value=mock_sock):
  253. ok = client.set('k', 'v')
  254. self.assertTrue(ok)
  255. def test_execute_without_password_or_db(self):
  256. """password 为空、db 为0 → 直接执行命令,不发 AUTH/SELECT。"""
  257. client = RedisSocketClient(password='', db=0)
  258. mock_sock = self._mock_sock([b'+OK\r\n'])
  259. with patch('socket.create_connection', return_value=mock_sock):
  260. ok = client.set('k', 'v', ex=300)
  261. self.assertTrue(ok)
  262. def test_delete_returns_integer(self):
  263. client = RedisSocketClient()
  264. mock_sock = self._mock_sock([b':1\r\n'])
  265. with patch('socket.create_connection', return_value=mock_sock):
  266. result = client.delete('k')
  267. self.assertEqual(1, result)
  268. def test_set_without_ex(self):
  269. client = RedisSocketClient()
  270. mock_sock = self._mock_sock([b'+OK\r\n'])
  271. with patch('socket.create_connection', return_value=mock_sock):
  272. ok = client.set('k', 'v') # no ex= argument
  273. self.assertTrue(ok)
  274. if __name__ == '__main__':
  275. unittest.main()