| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192 |
- """Unit tests for the Step0 connection checks; no real network is used."""
- from __future__ import annotations
- import unittest
- from types import SimpleNamespace
- from step0_pre_prepare.connection_checks import (
- ConnectionCheckError,
- Step0ConnectionSettings,
- check_deepseek_connection,
- check_redis_connection,
- )
- class FakeCompletions:
- def __init__(self) -> None:
- self.kwargs = None
- def create(self, **kwargs):
- self.kwargs = kwargs
- return SimpleNamespace(choices=[SimpleNamespace()])
- class FakeOpenAI:
- instance = None
- def __init__(self, **kwargs) -> None:
- self.client_kwargs = kwargs
- self.chat = SimpleNamespace(completions=FakeCompletions())
- type(self).instance = self
- class FakeRedis:
- def __init__(self) -> None:
- self.ping_called = False
- self.closed = False
- def ping(self) -> bool:
- self.ping_called = True
- return True
- def close(self) -> None:
- self.closed = True
- class Step0ConnectionChecksTest(unittest.TestCase):
- def setUp(self) -> None:
- self.settings = Step0ConnectionSettings(
- deepseek_api_key="test-key",
- deepseek_base_url="https://example.invalid/v1",
- deepseek_model="test-model",
- redis_url="redis://example.invalid:6379/0",
- )
- def test_deepseek_uses_a_minimal_request(self) -> None:
- check_deepseek_connection(self.settings, client_factory=FakeOpenAI)
- client = FakeOpenAI.instance
- self.assertEqual(client.chat.completions.kwargs["max_tokens"], 1)
- self.assertFalse(client.chat.completions.kwargs["stream"])
- self.assertEqual(client.chat.completions.kwargs["model"], "test-model")
- def test_redis_only_pings_and_closes(self) -> None:
- client = FakeRedis()
- factory_kwargs = {}
- def factory(url, **kwargs):
- factory_kwargs["url"] = url
- factory_kwargs.update(kwargs)
- return client
- check_redis_connection(self.settings, redis_factory=factory)
- self.assertTrue(client.ping_called)
- self.assertTrue(client.closed)
- self.assertEqual(factory_kwargs["url"], self.settings.redis_url)
- def test_failure_message_does_not_include_exception_details(self) -> None:
- class BrokenOpenAI:
- def __init__(self, **kwargs) -> None:
- raise RuntimeError("secret value must not leak")
- with self.assertRaises(ConnectionCheckError) as caught:
- check_deepseek_connection(self.settings, client_factory=BrokenOpenAI)
- self.assertNotIn("secret value", str(caught.exception))
- self.assertIn("RuntimeError", str(caught.exception))
- if __name__ == "__main__":
- unittest.main()
|