Jack1808/Claude_Code
0
1import asyncio2import contextlib3import logging4import os5 6import pytest7 8# Set mock environment BEFORE any imports that use Settings9os.environ.setdefault("NVIDIA_NIM_API_KEY", "test_key")10os.environ.setdefault("MODEL", "nvidia_nim/test-model")11os.environ["PTB_TIMEDELTA"] = "1"12# Ensure tests don't pick up a server API key from the repo .env13# (tests expect endpoints to be unauthenticated by default)14os.environ["ANTHROPIC_AUTH_TOKEN"] = ""15 16from typing import Any17from unittest.mock import AsyncMock, MagicMock18 19from config.nim import NimSettings20from messaging.models import IncomingMessage21from messaging.platforms.base import (22 CLISession,23 MessagingPlatform,24 SessionManagerInterface,25)26from messaging.session import SessionStore27from providers.base import ProviderConfig28from providers.nvidia_nim import NvidiaNimProvider29 30 31@pytest.fixture(autouse=True)32def _isolate_from_dotenv(monkeypatch):33 """Prevent Pydantic BaseSettings from reading the .env file during tests."""34 from config.settings import Settings35 36 monkeypatch.setattr(37 Settings, "model_config", {**Settings.model_config, "env_file": None}38 )39 40 41@pytest.fixture42def provider_config():43 return ProviderConfig(44 api_key="test_key",45 base_url="https://test.api.nvidia.com/v1",46 rate_limit=10,47 rate_window=60,48 )49 50 51@pytest.fixture52def nim_provider(provider_config):53 return NvidiaNimProvider(provider_config, nim_settings=NimSettings())54 55 56@pytest.fixture57def open_router_provider(provider_config):58 from providers.open_router import OpenRouterProvider59 60 return OpenRouterProvider(provider_config)61 62 63@pytest.fixture64def lmstudio_provider(provider_config):65 from providers.lmstudio import LMStudioProvider66 67 lmstudio_config = ProviderConfig(68 api_key="lm-studio",69 base_url="http://localhost:1234/v1",70 rate_limit=provider_config.rate_limit,71 rate_window=provider_config.rate_window,72 )73 return LMStudioProvider(lmstudio_config)74 75 76@pytest.fixture77def llamacpp_provider(provider_config):78 from providers.llamacpp import LlamaCppProvider79 80 llamacpp_config = ProviderConfig(81 api_key="llamacpp",82 base_url="http://localhost:8080/v1",83 rate_limit=10,84 rate_window=60,85 )86 return LlamaCppProvider(llamacpp_config)87 88 89@pytest.fixture90def mock_cli_session():91 session = MagicMock(spec=CLISession)92 session.start_task = MagicMock() # This will return an async generator93 session.is_busy = False94 return session95 96 97@pytest.fixture98def mock_cli_manager():99 manager = MagicMock(spec=SessionManagerInterface)100 manager.get_or_create_session = AsyncMock()101 manager.register_real_session_id = AsyncMock(return_value=True)102 manager.stop_all = AsyncMock()103 manager.remove_session = AsyncMock(return_value=True)104 manager.get_stats = MagicMock(return_value={"active_sessions": 0})105 return manager106 107 108@pytest.fixture109def mock_platform():110 platform = MagicMock(spec=MessagingPlatform)111 platform.send_message = AsyncMock(return_value="msg_123")112 platform.edit_message = AsyncMock()113 platform.delete_message = AsyncMock()114 platform.queue_send_message = AsyncMock(return_value="msg_123")115 platform.queue_edit_message = AsyncMock()116 platform.queue_delete_message = AsyncMock()117 118 def _fire_and_forget(task):119 if asyncio.iscoroutine(task):120 # Create a task to avoid "coroutine was never awaited" warning121 return asyncio.create_task(task)122 return None123 124 platform.fire_and_forget = MagicMock(side_effect=_fire_and_forget)125 return platform126 127 128@pytest.fixture129def mock_session_store():130 store = MagicMock(spec=SessionStore)131 store.save_tree = MagicMock()132 store.get_tree = MagicMock(return_value=None)133 store.register_node = MagicMock()134 store.clear_all = MagicMock()135 store.record_message_id = MagicMock()136 store.get_message_ids_for_chat = MagicMock(return_value=[])137 return store138 139 140@pytest.fixture141def incoming_message_factory():142 _valid_keys = frozenset(143 {144 "text",145 "chat_id",146 "user_id",147 "message_id",148 "platform",149 "reply_to_message_id",150 "message_thread_id",151 "username",152 "timestamp",153 "raw_event",154 "status_message_id",155 }156 )157 158 def _create(**kwargs):159 defaults: dict[str, Any] = {160 "text": "hello",161 "chat_id": "chat_1",162 "user_id": "user_1",163 "message_id": "msg_1",164 "platform": "telegram",165 }166 defaults.update(kwargs)167 if "timestamp" in defaults and isinstance(defaults["timestamp"], str):168 from datetime import datetime169 170 defaults["timestamp"] = datetime.fromisoformat(defaults["timestamp"])171 filtered = {k: v for k, v in defaults.items() if k in _valid_keys}172 return IncomingMessage(**filtered)173 174 return _create175 176 177@pytest.fixture(autouse=True)178def _propagate_loguru_to_caplog():179 """Route loguru logs to stdlib logging so pytest caplog captures them."""180 from loguru import logger as loguru_logger181 182 class _PropagateHandler:183 def write(self, message):184 record = message.record185 level = record["level"].no186 stdlib_level = min(level, logging.CRITICAL)187 py_logger = logging.getLogger(record["name"])188 py_logger.log(stdlib_level, record["message"])189 190 handler_id = loguru_logger.add(_PropagateHandler(), format="{message}")191 yield192 with contextlib.suppress(ValueError):193 loguru_logger.remove(194 handler_id195 ) # Handler already removed (e.g. by test_logging_config)196 