test_gateway_session_store_unit.py 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124
  1. import json
  2. import unittest
  3. from unittest.mock import MagicMock
  4. from services.gateway_session_store import GatewaySessionStore
  5. from utils.security import hash_gateway_session_id
  6. class TestGatewaySessionStore(unittest.TestCase):
  7. def setUp(self):
  8. self.mock_redis = MagicMock()
  9. self.store = GatewaySessionStore(
  10. self.mock_redis,
  11. prefix='fms:mcp:gateway:',
  12. ttl_seconds=3600
  13. )
  14. def test_key_for_generates_correct_key(self):
  15. session_id = 'GWS_abc123'
  16. expected_hash = hash_gateway_session_id(session_id)
  17. expected_key = f'fms:mcp:gateway:session:{expected_hash}'
  18. key = self.store.key_for(session_id)
  19. self.assertEqual(key, expected_key)
  20. def test_save_stores_session_with_hash(self):
  21. session_id = 'GWS_test123'
  22. session_data = {
  23. 'mcp_token': 'MT_xyz',
  24. 'admin_id': 123,
  25. 'company_id': 100,
  26. }
  27. result = self.store.save(session_id, session_data)
  28. # Should add gateway_session_id_hash
  29. self.assertIn('gateway_session_id_hash', result)
  30. self.assertEqual(result['gateway_session_id_hash'], hash_gateway_session_id(session_id))
  31. # Should call Redis set with correct parameters
  32. self.mock_redis.set.assert_called_once()
  33. call_args = self.mock_redis.set.call_args
  34. key_arg = call_args[0][0]
  35. value_arg = call_args[0][1]
  36. ttl_arg = call_args[1]['ex']
  37. self.assertEqual(key_arg, self.store.key_for(session_id))
  38. self.assertEqual(ttl_arg, 3600)
  39. # Verify JSON contains expected fields
  40. saved_data = json.loads(value_arg)
  41. self.assertEqual(saved_data['mcp_token'], 'MT_xyz')
  42. self.assertEqual(saved_data['admin_id'], 123)
  43. self.assertEqual(saved_data['company_id'], 100)
  44. self.assertIn('gateway_session_id_hash', saved_data)
  45. def test_get_returns_session(self):
  46. session_id = 'GWS_test456'
  47. session_data = {
  48. 'mcp_token': 'MT_abc',
  49. 'admin_id': 456,
  50. 'company_id': 200,
  51. 'gateway_session_id_hash': hash_gateway_session_id(session_id),
  52. }
  53. self.mock_redis.get.return_value = json.dumps(session_data)
  54. result = self.store.get(session_id)
  55. self.assertEqual(result['mcp_token'], 'MT_abc')
  56. self.assertEqual(result['admin_id'], 456)
  57. self.assertEqual(result['company_id'], 200)
  58. self.mock_redis.get.assert_called_once_with(self.store.key_for(session_id))
  59. def test_get_returns_none_when_not_found(self):
  60. self.mock_redis.get.return_value = None
  61. result = self.store.get('GWS_nonexistent')
  62. self.assertIsNone(result)
  63. def test_get_handles_bytes(self):
  64. session_id = 'GWS_bytes_test'
  65. session_data = {'mcp_token': 'MT_test'}
  66. self.mock_redis.get.return_value = json.dumps(session_data).encode('utf-8')
  67. result = self.store.get(session_id)
  68. self.assertEqual(result['mcp_token'], 'MT_test')
  69. def test_delete_removes_session(self):
  70. session_id = 'GWS_delete_test'
  71. self.mock_redis.delete.return_value = 1
  72. result = self.store.delete(session_id)
  73. self.assertEqual(result, 1)
  74. self.mock_redis.delete.assert_called_once_with(self.store.key_for(session_id))
  75. def test_different_session_ids_have_different_keys(self):
  76. session_1 = 'GWS_session_1'
  77. session_2 = 'GWS_session_2'
  78. key_1 = self.store.key_for(session_1)
  79. key_2 = self.store.key_for(session_2)
  80. self.assertNotEqual(key_1, key_2)
  81. def test_prefix_is_applied_correctly(self):
  82. store_with_prefix = GatewaySessionStore(
  83. self.mock_redis,
  84. prefix='custom:prefix:',
  85. ttl_seconds=1800
  86. )
  87. session_id = 'GWS_test'
  88. key = store_with_prefix.key_for(session_id)
  89. self.assertTrue(key.startswith('custom:prefix:session:'))
  90. if __name__ == '__main__':
  91. unittest.main()