| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390 |
- """
- 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()
|