import json import unittest from unittest.mock import MagicMock from services.gateway_session_store import GatewaySessionStore from utils.security import hash_gateway_session_id class TestGatewaySessionStore(unittest.TestCase): def setUp(self): self.mock_redis = MagicMock() self.store = GatewaySessionStore( self.mock_redis, prefix='fms:mcp:gateway:', ttl_seconds=3600 ) def test_key_for_generates_correct_key(self): session_id = 'GWS_abc123' expected_hash = hash_gateway_session_id(session_id) expected_key = f'fms:mcp:gateway:session:{expected_hash}' key = self.store.key_for(session_id) self.assertEqual(key, expected_key) def test_save_stores_session_with_hash(self): session_id = 'GWS_test123' session_data = { 'mcp_token': 'MT_xyz', 'admin_id': 123, 'company_id': 100, } result = self.store.save(session_id, session_data) # Should add gateway_session_id_hash self.assertIn('gateway_session_id_hash', result) self.assertEqual(result['gateway_session_id_hash'], hash_gateway_session_id(session_id)) # Should call Redis set with correct parameters self.mock_redis.set.assert_called_once() call_args = self.mock_redis.set.call_args key_arg = call_args[0][0] value_arg = call_args[0][1] ttl_arg = call_args[1]['ex'] self.assertEqual(key_arg, self.store.key_for(session_id)) self.assertEqual(ttl_arg, 3600) # Verify JSON contains expected fields saved_data = json.loads(value_arg) self.assertEqual(saved_data['mcp_token'], 'MT_xyz') self.assertEqual(saved_data['admin_id'], 123) self.assertEqual(saved_data['company_id'], 100) self.assertIn('gateway_session_id_hash', saved_data) def test_get_returns_session(self): session_id = 'GWS_test456' session_data = { 'mcp_token': 'MT_abc', 'admin_id': 456, 'company_id': 200, 'gateway_session_id_hash': hash_gateway_session_id(session_id), } self.mock_redis.get.return_value = json.dumps(session_data) result = self.store.get(session_id) self.assertEqual(result['mcp_token'], 'MT_abc') self.assertEqual(result['admin_id'], 456) self.assertEqual(result['company_id'], 200) self.mock_redis.get.assert_called_once_with(self.store.key_for(session_id)) def test_get_returns_none_when_not_found(self): self.mock_redis.get.return_value = None result = self.store.get('GWS_nonexistent') self.assertIsNone(result) def test_get_handles_bytes(self): session_id = 'GWS_bytes_test' session_data = {'mcp_token': 'MT_test'} self.mock_redis.get.return_value = json.dumps(session_data).encode('utf-8') result = self.store.get(session_id) self.assertEqual(result['mcp_token'], 'MT_test') def test_delete_removes_session(self): session_id = 'GWS_delete_test' self.mock_redis.delete.return_value = 1 result = self.store.delete(session_id) self.assertEqual(result, 1) self.mock_redis.delete.assert_called_once_with(self.store.key_for(session_id)) def test_different_session_ids_have_different_keys(self): session_1 = 'GWS_session_1' session_2 = 'GWS_session_2' key_1 = self.store.key_for(session_1) key_2 = self.store.key_for(session_2) self.assertNotEqual(key_1, key_2) def test_prefix_is_applied_correctly(self): store_with_prefix = GatewaySessionStore( self.mock_redis, prefix='custom:prefix:', ttl_seconds=1800 ) session_id = 'GWS_test' key = store_with_prefix.key_for(session_id) self.assertTrue(key.startswith('custom:prefix:session:')) if __name__ == '__main__': unittest.main()