test_token_store_coverage.py 15 KB

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