"""共享测试 fixtures.""" from collections.abc import AsyncGenerator from typing import Any import pytest from src.ai.errors import AILLMUnavailableError from src.ai.providers import LLMProvider, LLMResponse, LLMStreamChunk class MockProvider(LLMProvider): """可控的 LLM Provider mock(用于测试 FailoverChain + Service).""" def __init__( self, name: str = "mock", available: bool = True, response_content: str = "mock response", fail: bool = False, stream_chunks: list[str] | None = None, ) -> None: self._name = name self._available = available self._response_content = response_content self._fail = fail self._stream_chunks = stream_chunks or ["hello", " world"] @property def name(self) -> str: return self._name def is_available(self) -> bool: return self._available async def chat( self, messages: list[dict[str, str]], model: str, temperature: float = 0.7, **kwargs: Any, ) -> LLMResponse: if self._fail: raise AILLMUnavailableError(f"{self._name} mock failure") return LLMResponse( content=self._response_content, model=model, usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, provider=self._name, ) async def stream_chat( self, messages: list[dict[str, str]], model: str, temperature: float = 0.7, **kwargs: Any, ) -> AsyncGenerator[LLMStreamChunk, None]: if self._fail: raise AILLMUnavailableError(f"{self._name} mock stream failure") for i, chunk in enumerate(self._stream_chunks): finish = "stop" if i == len(self._stream_chunks) - 1 else None yield LLMStreamChunk( delta=chunk, model=model, finish_reason=finish, provider=self._name, ) @pytest.fixture def mock_provider() -> MockProvider: """默认可用的 mock provider.""" return MockProvider() @pytest.fixture def failing_provider() -> MockProvider: """总是失败的 mock provider.""" return MockProvider(name="failing", fail=True) @pytest.fixture def unavailable_provider() -> MockProvider: """未配置的 mock provider.""" return MockProvider(name="unconfigured", available=False)