|
|
@@ -0,0 +1,340 @@
|
|
|
+"""
|
|
|
+Coverage补充测试:services/token_store.py
|
|
|
+
|
|
|
+目标路径:
|
|
|
+- InMemoryTokenStore.is_expiring — 无session、无expire_time、Z后缀、过期中、未过期
|
|
|
+- InMemoryTokenStore.require_token — 无session、token为空、正常返回
|
|
|
+- RedisTokenStore — 初始化校验、save(含/不含TTL)、get(bytes/str/空串/None)、clear、is_expiring
|
|
|
+- RedisTokenStore._ttl_seconds — None、Z后缀、naive datetime
|
|
|
+- RedisSocketClient._read_response — +/-/:/$/$(nil)/*(含空数组)/unsupported/empty
|
|
|
+- RedisSocketClient._execute — 带 password+db、不带 password+db
|
|
|
+- RedisSocketClient.set/get/delete(mocked socket)
|
|
|
+"""
|
|
|
+import io
|
|
|
+import json
|
|
|
+import unittest
|
|
|
+from datetime import datetime, timedelta, timezone
|
|
|
+from unittest.mock import MagicMock, patch
|
|
|
+
|
|
|
+from services.token_store import (
|
|
|
+ FileTokenStore,
|
|
|
+ InMemoryTokenStore,
|
|
|
+ RedisSocketClient,
|
|
|
+ RedisTokenStore,
|
|
|
+)
|
|
|
+
|
|
|
+
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+# InMemoryTokenStore
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+
|
|
|
+class InMemoryIsExpiringTest(unittest.TestCase):
|
|
|
+ def test_no_session_is_expiring(self):
|
|
|
+ store = InMemoryTokenStore()
|
|
|
+ self.assertTrue(store.is_expiring())
|
|
|
+
|
|
|
+ def test_session_with_none_expire_time_is_expiring(self):
|
|
|
+ store = InMemoryTokenStore()
|
|
|
+ store._session = {'token': 'T', 'expire_time': None}
|
|
|
+ self.assertTrue(store.is_expiring())
|
|
|
+
|
|
|
+ def test_session_with_empty_expire_time_is_expiring(self):
|
|
|
+ store = InMemoryTokenStore()
|
|
|
+ store._session = {'token': 'T', 'expire_time': ''}
|
|
|
+ self.assertTrue(store.is_expiring())
|
|
|
+
|
|
|
+ def test_expire_soon_within_skew_is_expiring(self):
|
|
|
+ store = InMemoryTokenStore(refresh_skew_seconds=120)
|
|
|
+ # 60 秒后到期,在 120s 的 skew 窗口内 → 算作 expiring
|
|
|
+ soon = (datetime.now(timezone.utc) + timedelta(seconds=60)).strftime('%Y-%m-%dT%H:%M:%SZ')
|
|
|
+ store._session = {'token': 'T', 'expire_time': soon}
|
|
|
+ self.assertTrue(store.is_expiring())
|
|
|
+
|
|
|
+ def test_expire_far_future_not_expiring(self):
|
|
|
+ store = InMemoryTokenStore(refresh_skew_seconds=120)
|
|
|
+ store._session = {'token': 'T', 'expire_time': '2099-01-01T00:00:00Z'}
|
|
|
+ self.assertFalse(store.is_expiring())
|
|
|
+
|
|
|
+ def test_expire_time_z_suffix_parsed_correctly(self):
|
|
|
+ store = InMemoryTokenStore(refresh_skew_seconds=0)
|
|
|
+ # 已过期(1秒前)
|
|
|
+ past = (datetime.now(timezone.utc) - timedelta(seconds=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
|
|
|
+ store._session = {'token': 'T', 'expire_time': past}
|
|
|
+ self.assertTrue(store.is_expiring())
|
|
|
+
|
|
|
+ def test_expire_time_without_z_suffix(self):
|
|
|
+ store = InMemoryTokenStore(refresh_skew_seconds=0)
|
|
|
+ # 无时区后缀的 ISO 格式(naive datetime)
|
|
|
+ past = (datetime.now() - timedelta(seconds=1)).strftime('%Y-%m-%dT%H:%M:%S')
|
|
|
+ store._session = {'token': 'T', 'expire_time': past}
|
|
|
+ self.assertTrue(store.is_expiring())
|
|
|
+
|
|
|
+
|
|
|
+class InMemoryRequireTokenTest(unittest.TestCase):
|
|
|
+ def test_raises_when_no_session(self):
|
|
|
+ store = InMemoryTokenStore()
|
|
|
+ with self.assertRaises(RuntimeError) as ctx:
|
|
|
+ store.require_token()
|
|
|
+ self.assertIn('mcp token missing', str(ctx.exception))
|
|
|
+
|
|
|
+ def test_raises_when_token_is_empty_string(self):
|
|
|
+ store = InMemoryTokenStore()
|
|
|
+ store._session = {'token': '', 'expire_time': '2099-01-01T00:00:00'}
|
|
|
+ with self.assertRaises(RuntimeError):
|
|
|
+ store.require_token()
|
|
|
+
|
|
|
+ def test_returns_token_when_present(self):
|
|
|
+ store = InMemoryTokenStore()
|
|
|
+ store.save('MT_abc123', '2099-01-01T00:00:00')
|
|
|
+ self.assertEqual('MT_abc123', store.require_token())
|
|
|
+
|
|
|
+
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+# RedisTokenStore(使用 MagicMock client 隔离真实 Redis)
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+
|
|
|
+class RedisTokenStoreMockedTest(unittest.TestCase):
|
|
|
+ def _make_store(self, raw_get=None, prefix='test:', session_key='session_1'):
|
|
|
+ client = MagicMock()
|
|
|
+ client.get.return_value = raw_get
|
|
|
+ store = RedisTokenStore(client, prefix=prefix, session_key=session_key)
|
|
|
+ return store, client
|
|
|
+
|
|
|
+ # 初始化
|
|
|
+ def test_empty_session_key_raises_value_error(self):
|
|
|
+ client = MagicMock()
|
|
|
+ with self.assertRaises(ValueError) as ctx:
|
|
|
+ RedisTokenStore(client, session_key='')
|
|
|
+ self.assertIn('session_key is required', str(ctx.exception))
|
|
|
+
|
|
|
+ def test_whitespace_session_key_raises_value_error(self):
|
|
|
+ client = MagicMock()
|
|
|
+ with self.assertRaises(ValueError):
|
|
|
+ RedisTokenStore(client, session_key=' ')
|
|
|
+
|
|
|
+ def test_key_composed_from_prefix_and_session_key(self):
|
|
|
+ store, _ = self._make_store(prefix='ns:', session_key='abc')
|
|
|
+ self.assertEqual('ns:abc', store.key)
|
|
|
+
|
|
|
+ # save
|
|
|
+ def test_save_calls_client_set_with_positive_ttl(self):
|
|
|
+ store, client = self._make_store()
|
|
|
+ store.save('MT_001', '2099-01-01T00:00:00Z')
|
|
|
+ self.assertTrue(client.set.called)
|
|
|
+ args, kwargs = client.set.call_args
|
|
|
+ self.assertEqual('test:session_1', args[0])
|
|
|
+ parsed = json.loads(args[1])
|
|
|
+ self.assertEqual('MT_001', parsed['token'])
|
|
|
+ self.assertIsNotNone(kwargs.get('ex'))
|
|
|
+ self.assertGreater(kwargs['ex'], 0)
|
|
|
+
|
|
|
+ def test_save_with_none_expire_time_passes_no_ex(self):
|
|
|
+ store, client = self._make_store()
|
|
|
+ store.save('MT_001', None)
|
|
|
+ _, kwargs = client.set.call_args
|
|
|
+ self.assertIsNone(kwargs.get('ex'))
|
|
|
+
|
|
|
+ # get
|
|
|
+ def test_get_returns_none_when_redis_returns_none(self):
|
|
|
+ store, _ = self._make_store(raw_get=None)
|
|
|
+ self.assertIsNone(store.get())
|
|
|
+
|
|
|
+ def test_get_returns_none_on_empty_string(self):
|
|
|
+ store, _ = self._make_store(raw_get='')
|
|
|
+ self.assertIsNone(store.get())
|
|
|
+
|
|
|
+ def test_get_decodes_bytes(self):
|
|
|
+ data = {'token': 'MT_bytes', 'expire_time': '2099-01-01T00:00:00'}
|
|
|
+ store, _ = self._make_store(raw_get=json.dumps(data).encode('utf-8'))
|
|
|
+ result = store.get()
|
|
|
+ self.assertEqual('MT_bytes', result['token'])
|
|
|
+
|
|
|
+ def test_get_handles_string_response(self):
|
|
|
+ data = {'token': 'MT_str', 'expire_time': '2099-01-01T00:00:00'}
|
|
|
+ store, _ = self._make_store(raw_get=json.dumps(data))
|
|
|
+ result = store.get()
|
|
|
+ self.assertEqual('MT_str', result['token'])
|
|
|
+
|
|
|
+ # clear
|
|
|
+ def test_clear_calls_client_delete_with_correct_key(self):
|
|
|
+ store, client = self._make_store()
|
|
|
+ store.clear()
|
|
|
+ client.delete.assert_called_once_with('test:session_1')
|
|
|
+ self.assertIsNone(store._session)
|
|
|
+
|
|
|
+ # is_expiring — refreshes from Redis first
|
|
|
+ def test_is_expiring_reads_from_redis_first(self):
|
|
|
+ client = MagicMock()
|
|
|
+ data = {'token': 'MT_live', 'expire_time': '2099-01-01T00:00:00'}
|
|
|
+ client.get.return_value = json.dumps(data)
|
|
|
+ store = RedisTokenStore(client, prefix='t:', session_key='s')
|
|
|
+ self.assertFalse(store.is_expiring())
|
|
|
+ self.assertTrue(client.get.called)
|
|
|
+
|
|
|
+ def test_is_expiring_true_when_redis_has_no_data(self):
|
|
|
+ store, _ = self._make_store(raw_get=None)
|
|
|
+ self.assertTrue(store.is_expiring())
|
|
|
+
|
|
|
+
|
|
|
+class RedisTtlSecondsTest(unittest.TestCase):
|
|
|
+ def test_none_returns_none(self):
|
|
|
+ self.assertIsNone(RedisTokenStore._ttl_seconds(None))
|
|
|
+
|
|
|
+ def test_empty_string_returns_none(self):
|
|
|
+ self.assertIsNone(RedisTokenStore._ttl_seconds(''))
|
|
|
+
|
|
|
+ def test_z_suffix_parsed_as_utc(self):
|
|
|
+ future = (datetime.now(timezone.utc) + timedelta(hours=2)).strftime('%Y-%m-%dT%H:%M:%SZ')
|
|
|
+ ttl = RedisTokenStore._ttl_seconds(future)
|
|
|
+ self.assertGreater(ttl, 3600)
|
|
|
+
|
|
|
+ def test_naive_datetime_uses_local_time(self):
|
|
|
+ future = (datetime.now() + timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M:%S')
|
|
|
+ ttl = RedisTokenStore._ttl_seconds(future)
|
|
|
+ self.assertIsNotNone(ttl)
|
|
|
+ self.assertGreater(ttl, 0)
|
|
|
+
|
|
|
+ def test_already_past_returns_1_at_minimum(self):
|
|
|
+ # 已过期的时间仍返回 max(1, ...)
|
|
|
+ past = (datetime.now(timezone.utc) - timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M:%SZ')
|
|
|
+ ttl = RedisTokenStore._ttl_seconds(past)
|
|
|
+ self.assertEqual(1, ttl)
|
|
|
+
|
|
|
+
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+# RedisSocketClient._read_response(直接传入 BytesIO 流,不需要真实 socket)
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+
|
|
|
+class RedisReadResponseTest(unittest.TestCase):
|
|
|
+ """_read_response 是实例方法,通过实例调用。"""
|
|
|
+
|
|
|
+ def _client(self):
|
|
|
+ return RedisSocketClient()
|
|
|
+
|
|
|
+ def _stream(self, data: bytes) -> io.BytesIO:
|
|
|
+ return io.BytesIO(data)
|
|
|
+
|
|
|
+ def test_simple_string_ok(self):
|
|
|
+ result = self._client()._read_response(self._stream(b'+OK\r\n'))
|
|
|
+ self.assertEqual('OK', result)
|
|
|
+
|
|
|
+ def test_simple_string_pong(self):
|
|
|
+ result = self._client()._read_response(self._stream(b'+PONG\r\n'))
|
|
|
+ self.assertEqual('PONG', result)
|
|
|
+
|
|
|
+ def test_error_raises_runtime_error(self):
|
|
|
+ with self.assertRaises(RuntimeError) as ctx:
|
|
|
+ self._client()._read_response(self._stream(b'-ERR unknown command\r\n'))
|
|
|
+ self.assertIn('ERR unknown command', str(ctx.exception))
|
|
|
+
|
|
|
+ def test_integer(self):
|
|
|
+ result = self._client()._read_response(self._stream(b':7\r\n'))
|
|
|
+ self.assertEqual(7, result)
|
|
|
+
|
|
|
+ def test_integer_zero(self):
|
|
|
+ result = self._client()._read_response(self._stream(b':0\r\n'))
|
|
|
+ self.assertEqual(0, result)
|
|
|
+
|
|
|
+ def test_bulk_string(self):
|
|
|
+ result = self._client()._read_response(self._stream(b'$5\r\nhello\r\n'))
|
|
|
+ self.assertEqual('hello', result)
|
|
|
+
|
|
|
+ def test_bulk_string_nil(self):
|
|
|
+ result = self._client()._read_response(self._stream(b'$-1\r\n'))
|
|
|
+ self.assertIsNone(result)
|
|
|
+
|
|
|
+ def test_array_two_elements(self):
|
|
|
+ result = self._client()._read_response(self._stream(b'*2\r\n+alpha\r\n+beta\r\n'))
|
|
|
+ self.assertEqual(['alpha', 'beta'], result)
|
|
|
+
|
|
|
+ def test_array_empty(self):
|
|
|
+ result = self._client()._read_response(self._stream(b'*0\r\n'))
|
|
|
+ self.assertEqual([], result)
|
|
|
+
|
|
|
+ def test_array_with_bulk_strings(self):
|
|
|
+ result = self._client()._read_response(self._stream(b'*2\r\n$3\r\nfoo\r\n$3\r\nbar\r\n'))
|
|
|
+ self.assertEqual(['foo', 'bar'], result)
|
|
|
+
|
|
|
+ def test_unsupported_prefix_raises(self):
|
|
|
+ with self.assertRaises(RuntimeError) as ctx:
|
|
|
+ self._client()._read_response(self._stream(b'?nope\r\n'))
|
|
|
+ self.assertIn('unsupported redis response', str(ctx.exception))
|
|
|
+
|
|
|
+ def test_empty_response_raises(self):
|
|
|
+ with self.assertRaises(RuntimeError) as ctx:
|
|
|
+ self._client()._read_response(self._stream(b''))
|
|
|
+ self.assertIn('empty redis response', str(ctx.exception))
|
|
|
+
|
|
|
+
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+# RedisSocketClient._execute(用 patch mock socket,覆盖 AUTH/SELECT 分支)
|
|
|
+# ---------------------------------------------------------------------------
|
|
|
+
|
|
|
+class RedisSocketClientExecuteTest(unittest.TestCase):
|
|
|
+ """通过 mock socket 验证 _execute 中 AUTH/SELECT 的条件分支。"""
|
|
|
+
|
|
|
+ def _mock_sock(self, responses):
|
|
|
+ """
|
|
|
+ responses: list[bytes],每个是一次 _read_response 调用的完整 RESP 响应。
|
|
|
+ """
|
|
|
+ stream = io.BytesIO(b''.join(responses))
|
|
|
+ sock = MagicMock()
|
|
|
+ sock.makefile.return_value = stream
|
|
|
+ sock.__enter__ = lambda s: s
|
|
|
+ sock.__exit__ = MagicMock(return_value=False)
|
|
|
+ return sock
|
|
|
+
|
|
|
+ def test_execute_with_password_sends_auth(self):
|
|
|
+ """password 非空 → 发送 AUTH 命令。"""
|
|
|
+ client = RedisSocketClient(password='secret', db=0)
|
|
|
+ # AUTH → OK;GET → nil
|
|
|
+ mock_sock = self._mock_sock([b'+OK\r\n', b'$-1\r\n'])
|
|
|
+ with patch('socket.create_connection', return_value=mock_sock):
|
|
|
+ result = client.get('k')
|
|
|
+ self.assertIsNone(result)
|
|
|
+ # sendall 应被调用多次(AUTH + GET)
|
|
|
+ self.assertGreater(mock_sock.sendall.call_count, 0)
|
|
|
+
|
|
|
+ def test_execute_with_db_sends_select(self):
|
|
|
+ """db 非0 → 发送 SELECT 命令。"""
|
|
|
+ client = RedisSocketClient(password='', db=3)
|
|
|
+ # SELECT → OK;GET → value
|
|
|
+ mock_sock = self._mock_sock([b'+OK\r\n', b'$5\r\nvalue\r\n'])
|
|
|
+ with patch('socket.create_connection', return_value=mock_sock):
|
|
|
+ result = client.get('k')
|
|
|
+ self.assertEqual('value', result)
|
|
|
+
|
|
|
+ def test_execute_with_password_and_db(self):
|
|
|
+ """password 非空 且 db 非0 → 发送 AUTH + SELECT。"""
|
|
|
+ client = RedisSocketClient(password='pass', db=2)
|
|
|
+ # AUTH → OK;SELECT → OK;SET → OK
|
|
|
+ mock_sock = self._mock_sock([b'+OK\r\n', b'+OK\r\n', b'+OK\r\n'])
|
|
|
+ with patch('socket.create_connection', return_value=mock_sock):
|
|
|
+ ok = client.set('k', 'v')
|
|
|
+ self.assertTrue(ok)
|
|
|
+
|
|
|
+ def test_execute_without_password_or_db(self):
|
|
|
+ """password 为空、db 为0 → 直接执行命令,不发 AUTH/SELECT。"""
|
|
|
+ client = RedisSocketClient(password='', db=0)
|
|
|
+ mock_sock = self._mock_sock([b'+OK\r\n'])
|
|
|
+ with patch('socket.create_connection', return_value=mock_sock):
|
|
|
+ ok = client.set('k', 'v', ex=300)
|
|
|
+ self.assertTrue(ok)
|
|
|
+
|
|
|
+ def test_delete_returns_integer(self):
|
|
|
+ client = RedisSocketClient()
|
|
|
+ mock_sock = self._mock_sock([b':1\r\n'])
|
|
|
+ with patch('socket.create_connection', return_value=mock_sock):
|
|
|
+ result = client.delete('k')
|
|
|
+ self.assertEqual(1, result)
|
|
|
+
|
|
|
+ def test_set_without_ex(self):
|
|
|
+ client = RedisSocketClient()
|
|
|
+ mock_sock = self._mock_sock([b'+OK\r\n'])
|
|
|
+ with patch('socket.create_connection', return_value=mock_sock):
|
|
|
+ ok = client.set('k', 'v') # no ex= argument
|
|
|
+ self.assertTrue(ok)
|
|
|
+
|
|
|
+
|
|
|
+if __name__ == '__main__':
|
|
|
+ unittest.main()
|