test_auth_client.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116
  1. import os
  2. import os
  3. import tempfile
  4. import unittest
  5. from config import GatewayConfig
  6. from services.auth_client import AuthClient
  7. from services.token_store import InMemoryTokenStore
  8. class DummyTransport:
  9. def __init__(self):
  10. self.calls = []
  11. def post_json(self, url, payload, headers, timeout):
  12. self.calls.append(
  13. {
  14. 'url': url,
  15. 'payload': payload,
  16. 'headers': headers,
  17. 'timeout': timeout,
  18. }
  19. )
  20. if url.endswith('/mcp/auth/exchange'):
  21. return {
  22. 'code': 'MCP_0000',
  23. 'msg': 'success',
  24. 'data': {
  25. 'mcp_token': 'MT_exchange',
  26. 'expire_time': '2099-01-01T00:00:00',
  27. },
  28. }
  29. if url.endswith('/mcp/auth/refresh'):
  30. return {
  31. 'code': 'MCP_0000',
  32. 'msg': 'success',
  33. 'data': {
  34. 'mcp_token': 'MT_refresh',
  35. 'expire_time': '2099-01-02T00:00:00',
  36. },
  37. }
  38. return {
  39. 'code': 'MCP_0000',
  40. 'msg': 'success',
  41. 'data': {},
  42. }
  43. class AuthClientTest(unittest.TestCase):
  44. def test_gateway_config_reads_expected_env_keys(self):
  45. env = {
  46. 'MCP_AUTH_BASE_URL': 'http://base.example.test',
  47. 'MCP_TOOLS_BASE_URL': 'http://tools.example.test',
  48. 'MCP_CLIENT_TYPE': 'workbuddy',
  49. 'MCP_TIMEOUT_SECONDS': '9',
  50. 'MCP_REFRESH_SKEW_SECONDS': '120',
  51. }
  52. config = GatewayConfig.from_env(env, dotenv_path=os.path.join(tempfile.gettempdir(), 'missing-auth-client.env'))
  53. self.assertEqual('http://base.example.test', config.auth_base_url)
  54. self.assertEqual('http://tools.example.test', config.tools_base_url)
  55. self.assertEqual('workbuddy', config.client_type)
  56. self.assertEqual(9, config.timeout_seconds)
  57. self.assertEqual(120, config.refresh_skew_seconds)
  58. def test_auth_client_uses_auth_routes_and_updates_token_store(self):
  59. transport = DummyTransport()
  60. store = InMemoryTokenStore(refresh_skew_seconds=60)
  61. client = AuthClient(
  62. base_url='http://base.example.test',
  63. client_type='workbuddy',
  64. token_store=store,
  65. transport=transport,
  66. timeout=9,
  67. )
  68. exchange = client.exchange('AUTH123')
  69. refresh = client.refresh('MT_exchange')
  70. self.assertEqual('MT_refresh', store.get()['token'])
  71. revoke = client.revoke('MT_refresh')
  72. self.assertEqual('MT_exchange', exchange['data']['mcp_token'])
  73. self.assertEqual('MT_refresh', refresh['data']['mcp_token'])
  74. self.assertIsNone(store.get())
  75. self.assertEqual({}, revoke['data'])
  76. self.assertEqual('http://base.example.test/mcp/auth/exchange', transport.calls[0]['url'])
  77. self.assertEqual('http://base.example.test/mcp/auth/refresh', transport.calls[1]['url'])
  78. self.assertEqual('http://base.example.test/mcp/auth/revoke', transport.calls[2]['url'])
  79. self.assertEqual({'auth_code': 'AUTH123', 'client_type': 'workbuddy'}, transport.calls[0]['payload'])
  80. self.assertEqual({'mcp_token': 'MT_exchange'}, transport.calls[1]['payload'])
  81. self.assertEqual({'mcp_token': 'MT_refresh'}, transport.calls[2]['payload'])
  82. def test_refresh_and_revoke_send_bearer_header_and_revoke_clears_token(self):
  83. transport = DummyTransport()
  84. store = InMemoryTokenStore(refresh_skew_seconds=60)
  85. store.save('MT_exchange', '2099-01-01T00:00:00')
  86. client = AuthClient(
  87. base_url='http://base.example.test',
  88. client_type='workbuddy',
  89. token_store=store,
  90. transport=transport,
  91. timeout=9,
  92. )
  93. client.refresh('MT_exchange')
  94. client.revoke('MT_refresh')
  95. self.assertEqual('Bearer MT_exchange', transport.calls[0]['headers']['Authorization'])
  96. self.assertEqual('Bearer MT_refresh', transport.calls[1]['headers']['Authorization'])
  97. self.assertIsNone(store.get())
  98. if __name__ == '__main__':
  99. unittest.main()