""" 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 os import tempfile 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()) class FileTokenStoreBoundaryTest(unittest.TestCase): def test_save_without_parent_directory(self): with tempfile.TemporaryDirectory() as tmp_dir: old_cwd = os.getcwd() try: os.chdir(tmp_dir) store = FileTokenStore('token.json') store.save('MT_test', '2099-01-01T00:00:00') self.assertTrue(os.path.exists('token.json')) finally: os.chdir(old_cwd) def test_get_reloads_when_session_is_missing_in_memory(self): with tempfile.TemporaryDirectory() as tmp_dir: path = os.path.join(tmp_dir, 'token.json') store = FileTokenStore(path) with open(path, 'w', encoding='utf-8') as handle: json.dump({ 'token': 'MT_disk', 'expire_time': '2099-01-01T00:00:00', }, handle) self.assertEqual('MT_disk', store.get()['token']) def test_clear_when_file_is_already_missing(self): with tempfile.TemporaryDirectory() as tmp_dir: path = os.path.join(tmp_dir, 'missing-token.json') store = FileTokenStore(path) store.clear() self.assertIsNone(store.get()) # --------------------------------------------------------------------------- # 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)) def test_read_line_rejects_missing_crlf(self): with self.assertRaisesRegex(RuntimeError, 'invalid redis line'): RedisSocketClient._read_line(self._stream(b'invalid')) class RedisSendTest(unittest.TestCase): def test_send_accepts_bytes(self): sock = MagicMock() RedisSocketClient._send(sock, b'PING') self.assertIn(b'PING', sock.sendall.call_args.args[0]) # --------------------------------------------------------------------------- # 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()