test_step0_connections.py 2.7 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192
  1. """Unit tests for the Step0 connection checks; no real network is used."""
  2. from __future__ import annotations
  3. import unittest
  4. from types import SimpleNamespace
  5. from step0_pre_prepare.connection_checks import (
  6. ConnectionCheckError,
  7. Step0ConnectionSettings,
  8. check_deepseek_connection,
  9. check_redis_connection,
  10. )
  11. class FakeCompletions:
  12. def __init__(self) -> None:
  13. self.kwargs = None
  14. def create(self, **kwargs):
  15. self.kwargs = kwargs
  16. return SimpleNamespace(choices=[SimpleNamespace()])
  17. class FakeOpenAI:
  18. instance = None
  19. def __init__(self, **kwargs) -> None:
  20. self.client_kwargs = kwargs
  21. self.chat = SimpleNamespace(completions=FakeCompletions())
  22. type(self).instance = self
  23. class FakeRedis:
  24. def __init__(self) -> None:
  25. self.ping_called = False
  26. self.closed = False
  27. def ping(self) -> bool:
  28. self.ping_called = True
  29. return True
  30. def close(self) -> None:
  31. self.closed = True
  32. class Step0ConnectionChecksTest(unittest.TestCase):
  33. def setUp(self) -> None:
  34. self.settings = Step0ConnectionSettings(
  35. deepseek_api_key="test-key",
  36. deepseek_base_url="https://example.invalid/v1",
  37. deepseek_model="test-model",
  38. redis_url="redis://example.invalid:6379/0",
  39. )
  40. def test_deepseek_uses_a_minimal_request(self) -> None:
  41. check_deepseek_connection(self.settings, client_factory=FakeOpenAI)
  42. client = FakeOpenAI.instance
  43. self.assertEqual(client.chat.completions.kwargs["max_tokens"], 1)
  44. self.assertFalse(client.chat.completions.kwargs["stream"])
  45. self.assertEqual(client.chat.completions.kwargs["model"], "test-model")
  46. def test_redis_only_pings_and_closes(self) -> None:
  47. client = FakeRedis()
  48. factory_kwargs = {}
  49. def factory(url, **kwargs):
  50. factory_kwargs["url"] = url
  51. factory_kwargs.update(kwargs)
  52. return client
  53. check_redis_connection(self.settings, redis_factory=factory)
  54. self.assertTrue(client.ping_called)
  55. self.assertTrue(client.closed)
  56. self.assertEqual(factory_kwargs["url"], self.settings.redis_url)
  57. def test_failure_message_does_not_include_exception_details(self) -> None:
  58. class BrokenOpenAI:
  59. def __init__(self, **kwargs) -> None:
  60. raise RuntimeError("secret value must not leak")
  61. with self.assertRaises(ConnectionCheckError) as caught:
  62. check_deepseek_connection(self.settings, client_factory=BrokenOpenAI)
  63. self.assertNotIn("secret value", str(caught.exception))
  64. self.assertIn("RuntimeError", str(caught.exception))
  65. if __name__ == "__main__":
  66. unittest.main()