"""LLM Provider 适配器测试. 使用 httpx.MockTransport mock HTTP 响应,验证 4 个 Provider 的 chat / embed / 流式解析行为,以及 create_failover_chain 工厂与 ProviderFailoverChain.embed 故障切换。 """ import contextlib from collections.abc import AsyncGenerator, Callable, Iterator from typing import Any from unittest.mock import patch import httpx import pytest from src.ai.config import Settings from src.ai.errors import AILLMUnavailableError from src.ai.providers import ( LLMProvider, LLMResponse, LLMStreamChunk, ProviderFailoverChain, create_failover_chain, ) from src.ai.providers.anthropic_provider import AnthropicProvider from src.ai.providers.baichuan_provider import BaichuanProvider from src.ai.providers.circuit_breaker import CircuitBreaker from src.ai.providers.ollama_provider import LocalOllamaProvider from src.ai.providers.openai_provider import OpenAIProvider from .conftest import MockProvider # Capture the real AsyncClient before any patching to avoid recursion # (patch replaces httpx.AsyncClient globally via the shared module object). _REAL_ASYNC_CLIENT = httpx.AsyncClient # --------------------------------------------------------------------------- # Mock transport handlers # --------------------------------------------------------------------------- def _openai_chat_handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, json={ "choices": [{"message": {"content": "hello world"}}], "usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}, "model": "gpt-4o", }, ) def _openai_embed_handler(request: httpx.Request) -> httpx.Response: return httpx.Response(200, json={"data": [{"embedding": [0.1, 0.2, 0.3]}]}) def _empty_embed_handler(request: httpx.Request) -> httpx.Response: return httpx.Response(200, json={"data": []}) def _anthropic_chat_handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, json={ "content": [{"type": "text", "text": "hello from claude"}], "usage": {"input_tokens": 5, "output_tokens": 10}, "model": "claude-3-5-sonnet-20241022", }, ) def _baichuan_chat_handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, json={ "choices": [{"message": {"content": "baichuan response"}}], "usage": {"prompt_tokens": 3, "completion_tokens": 7, "total_tokens": 10}, "model": "Baichuan2-53B", }, ) def _ollama_chat_handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, json={ "message": {"role": "assistant", "content": "ollama response"}, "prompt_eval_count": 4, "eval_count": 8, "model": "llama3", }, ) def _ollama_embed_handler(request: httpx.Request) -> httpx.Response: return httpx.Response(200, json={"embedding": [0.5, 0.6]}) def _server_error_handler(request: httpx.Request) -> httpx.Response: return httpx.Response(500, text="internal server error") def _network_error_handler(request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("connection refused", request=request) @contextlib.contextmanager def _mock_http( module_path: str, handler: Callable[[httpx.Request], httpx.Response], ) -> Iterator[None]: """Patch httpx.AsyncClient in a provider module to use a MockTransport. Providers create ``httpx.AsyncClient`` internally; patching the shared ``httpx`` module attribute makes the mock transport take effect. The real class is captured up-front to avoid infinite recursion. """ transport = httpx.MockTransport(handler) with patch( module_path, lambda **kwargs: _REAL_ASYNC_CLIENT(transport=transport, **kwargs), ): yield # --------------------------------------------------------------------------- # Helper provider for failover-chain embed tests # --------------------------------------------------------------------------- class _EmbeddableProvider(LLMProvider): """Provider with embed support for failover chain tests.""" def __init__( self, name: str, embedding: list[float], fail: bool = False, ) -> None: self._name = name self._embedding = embedding self._fail = fail @property def name(self) -> str: return self._name def is_available(self) -> bool: return True async def chat( self, messages: list[dict[str, str]], model: str, temperature: float = 0.7, **kwargs: Any, ) -> LLMResponse: return LLMResponse(content="", model=model, provider=self._name) async def stream_chat( self, messages: list[dict[str, str]], model: str, temperature: float = 0.7, **kwargs: Any, ) -> AsyncGenerator[LLMStreamChunk, None]: yield LLMStreamChunk(delta="", model=model, provider=self._name) async def embed(self, text: str, model: str) -> list[float]: if self._fail: raise AILLMUnavailableError(f"{self._name} embed failed") return self._embedding # --------------------------------------------------------------------------- # OpenAIProvider # --------------------------------------------------------------------------- class TestOpenAIProvider: _MODULE = "src.ai.providers.openai_provider.httpx.AsyncClient" def test_name(self) -> None: assert OpenAIProvider(api_key="test-key").name == "openai" def test_is_available_true(self) -> None: assert OpenAIProvider(api_key="test-key").is_available() is True def test_is_available_false(self) -> None: assert OpenAIProvider(api_key="").is_available() is False async def test_chat_success(self) -> None: with _mock_http(self._MODULE, _openai_chat_handler): provider = OpenAIProvider(api_key="test-key") result = await provider.chat( [{"role": "user", "content": "hi"}], "gpt-4o", ) assert result.content == "hello world" assert result.provider == "openai" assert result.model == "gpt-4o" assert result.usage == { "prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15, } async def test_chat_not_configured_raises(self) -> None: provider = OpenAIProvider(api_key="") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "gpt-4o") async def test_chat_http_error_raises(self) -> None: with _mock_http(self._MODULE, _server_error_handler): provider = OpenAIProvider(api_key="test-key") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "gpt-4o") async def test_chat_network_error_raises(self) -> None: with _mock_http(self._MODULE, _network_error_handler): provider = OpenAIProvider(api_key="test-key") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "gpt-4o") def test_parse_sse_line_empty(self) -> None: assert OpenAIProvider._parse_sse_line("", "gpt-4o") is None def test_parse_sse_line_non_data(self) -> None: assert OpenAIProvider._parse_sse_line("event: ping", "gpt-4o") is None def test_parse_sse_line_done(self) -> None: chunk = OpenAIProvider._parse_sse_line("data: [DONE]", "gpt-4o") assert chunk is not None assert chunk.finish_reason == "stop" assert chunk.delta == "" def test_parse_sse_line_content(self) -> None: line = 'data: {"choices":[{"delta":{"content":"hello"}}]}' chunk = OpenAIProvider._parse_sse_line(line, "gpt-4o") assert chunk is not None assert chunk.delta == "hello" assert chunk.finish_reason is None assert chunk.provider == "openai" def test_parse_sse_line_finish_reason(self) -> None: line = 'data: {"choices":[{"delta":{},"finish_reason":"stop"}]}' chunk = OpenAIProvider._parse_sse_line(line, "gpt-4o") assert chunk is not None assert chunk.finish_reason == "stop" def test_parse_sse_line_invalid_json(self) -> None: assert OpenAIProvider._parse_sse_line("data: {invalid}", "gpt-4o") is None def test_parse_sse_line_no_choices(self) -> None: assert OpenAIProvider._parse_sse_line('data: {"choices":[]}', "gpt-4o") is None async def test_embed_success(self) -> None: with _mock_http(self._MODULE, _openai_embed_handler): provider = OpenAIProvider(api_key="test-key") result = await provider.embed("hello", "text-embedding-3-small") assert result == [0.1, 0.2, 0.3] async def test_embed_not_configured_raises(self) -> None: provider = OpenAIProvider(api_key="") with pytest.raises(AILLMUnavailableError): await provider.embed("hello", "text-embedding-3-small") async def test_embed_empty_response(self) -> None: with _mock_http(self._MODULE, _empty_embed_handler): provider = OpenAIProvider(api_key="test-key") result = await provider.embed("hello", "text-embedding-3-small") assert result == [] async def test_embed_http_error_raises(self) -> None: with _mock_http(self._MODULE, _server_error_handler): provider = OpenAIProvider(api_key="test-key") with pytest.raises(AILLMUnavailableError): await provider.embed("hello", "text-embedding-3-small") # --------------------------------------------------------------------------- # AnthropicProvider # --------------------------------------------------------------------------- class TestAnthropicProvider: _MODULE = "src.ai.providers.anthropic_provider.httpx.AsyncClient" def test_name(self) -> None: assert AnthropicProvider(api_key="test-key").name == "anthropic" def test_is_available_true(self) -> None: assert AnthropicProvider(api_key="test-key").is_available() is True def test_is_available_false(self) -> None: assert AnthropicProvider(api_key="").is_available() is False async def test_chat_success(self) -> None: with _mock_http(self._MODULE, _anthropic_chat_handler): provider = AnthropicProvider(api_key="test-key") result = await provider.chat( [{"role": "user", "content": "hi"}], "claude-3-5-sonnet-20241022", ) assert result.content == "hello from claude" assert result.provider == "anthropic" assert result.model == "claude-3-5-sonnet-20241022" assert result.usage == { "prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15, } async def test_chat_not_configured_raises(self) -> None: provider = AnthropicProvider(api_key="") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "claude-3") async def test_chat_http_error_raises(self) -> None: with _mock_http(self._MODULE, _server_error_handler): provider = AnthropicProvider(api_key="test-key") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "claude-3") async def test_chat_network_error_raises(self) -> None: with _mock_http(self._MODULE, _network_error_handler): provider = AnthropicProvider(api_key="test-key") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "claude-3") def test_parse_sse_line_empty(self) -> None: assert AnthropicProvider._parse_sse_line("", "claude") is None def test_parse_sse_line_non_data(self) -> None: assert AnthropicProvider._parse_sse_line( "event: content_block_delta", "claude", ) is None def test_parse_sse_line_content_block_delta(self) -> None: line = ( 'data: {"type":"content_block_delta",' '"delta":{"type":"text_delta","text":"hello"}}' ) chunk = AnthropicProvider._parse_sse_line(line, "claude") assert chunk is not None assert chunk.delta == "hello" assert chunk.provider == "anthropic" def test_parse_sse_line_message_stop(self) -> None: chunk = AnthropicProvider._parse_sse_line( 'data: {"type":"message_stop"}', "claude", ) assert chunk is not None assert chunk.finish_reason == "end_turn" def test_parse_sse_line_invalid_json(self) -> None: assert AnthropicProvider._parse_sse_line("data: {bad}", "claude") is None def test_convert_messages(self) -> None: provider = AnthropicProvider(api_key="test-key") messages = [ {"role": "system", "content": "You are helpful."}, {"role": "system", "content": "Be concise."}, {"role": "user", "content": "Hi"}, {"role": "assistant", "content": "Hello!"}, ] system_prompt, converted = provider._convert_messages(messages) assert system_prompt == "You are helpful.\n\nBe concise." assert len(converted) == 2 assert converted[0] == {"role": "user", "content": "Hi"} assert converted[1] == {"role": "assistant", "content": "Hello!"} def test_convert_messages_no_system(self) -> None: provider = AnthropicProvider(api_key="test-key") system_prompt, converted = provider._convert_messages( [{"role": "user", "content": "Hi"}], ) assert system_prompt == "" assert len(converted) == 1 assert converted[0] == {"role": "user", "content": "Hi"} # --------------------------------------------------------------------------- # BaichuanProvider # --------------------------------------------------------------------------- class TestBaichuanProvider: _MODULE = "src.ai.providers.baichuan_provider.httpx.AsyncClient" def test_name(self) -> None: assert BaichuanProvider(api_key="test-key").name == "baichuan" def test_is_available_true(self) -> None: assert BaichuanProvider(api_key="test-key").is_available() is True def test_is_available_false(self) -> None: assert BaichuanProvider(api_key="").is_available() is False async def test_chat_success(self) -> None: with _mock_http(self._MODULE, _baichuan_chat_handler): provider = BaichuanProvider(api_key="test-key") result = await provider.chat( [{"role": "user", "content": "hi"}], "Baichuan2-53B", ) assert result.content == "baichuan response" assert result.provider == "baichuan" assert result.usage == { "prompt_tokens": 3, "completion_tokens": 7, "total_tokens": 10, } async def test_chat_not_configured_raises(self) -> None: provider = BaichuanProvider(api_key="") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "Baichuan2") async def test_chat_http_error_raises(self) -> None: with _mock_http(self._MODULE, _server_error_handler): provider = BaichuanProvider(api_key="test-key") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "Baichuan2") async def test_chat_network_error_raises(self) -> None: with _mock_http(self._MODULE, _network_error_handler): provider = BaichuanProvider(api_key="test-key") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "Baichuan2") def test_parse_sse_line_done(self) -> None: # Baichuan reuses OpenAI SSE parser (compatible format) chunk = OpenAIProvider._parse_sse_line("data: [DONE]", "Baichuan2") assert chunk is not None assert chunk.finish_reason == "stop" def test_parse_sse_line_content(self) -> None: line = 'data: {"choices":[{"delta":{"content":"hi"}}]}' chunk = OpenAIProvider._parse_sse_line(line, "Baichuan2") assert chunk is not None assert chunk.delta == "hi" async def test_embed_not_implemented(self) -> None: provider = BaichuanProvider(api_key="test-key") with pytest.raises(NotImplementedError): await provider.embed("hello", "any-model") # --------------------------------------------------------------------------- # LocalOllamaProvider # --------------------------------------------------------------------------- class TestOllamaProvider: _MODULE = "src.ai.providers.ollama_provider.httpx.AsyncClient" def test_name(self) -> None: assert LocalOllamaProvider().name == "local_ollama" def test_is_available_true(self) -> None: provider = LocalOllamaProvider(base_url="http://localhost:11434") assert provider.is_available() is True def test_is_available_false(self) -> None: assert LocalOllamaProvider(base_url="").is_available() is False async def test_chat_success(self) -> None: with _mock_http(self._MODULE, _ollama_chat_handler): provider = LocalOllamaProvider() result = await provider.chat( [{"role": "user", "content": "hi"}], "llama3", ) assert result.content == "ollama response" assert result.provider == "local_ollama" assert result.model == "llama3" assert result.usage == { "prompt_tokens": 4, "completion_tokens": 8, "total_tokens": 12, } async def test_chat_not_configured_raises(self) -> None: provider = LocalOllamaProvider(base_url="") with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "llama3") async def test_chat_http_error_raises(self) -> None: with _mock_http(self._MODULE, _server_error_handler): provider = LocalOllamaProvider() with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "llama3") async def test_chat_network_error_raises(self) -> None: with _mock_http(self._MODULE, _network_error_handler): provider = LocalOllamaProvider() with pytest.raises(AILLMUnavailableError): await provider.chat([{"role": "user", "content": "hi"}], "llama3") def test_parse_ndjson_line_empty(self) -> None: assert LocalOllamaProvider._parse_ndjson_line("", "llama3") is None def test_parse_ndjson_line_done_true(self) -> None: line = ( '{"model":"llama3","message":{"role":"assistant","content":""},' '"done":true}' ) chunk = LocalOllamaProvider._parse_ndjson_line(line, "llama3") assert chunk is not None assert chunk.finish_reason == "stop" def test_parse_ndjson_line_done_false(self) -> None: line = ( '{"model":"llama3","message":{"role":"assistant","content":"hi"},' '"done":false}' ) chunk = LocalOllamaProvider._parse_ndjson_line(line, "llama3") assert chunk is not None assert chunk.delta == "hi" assert chunk.finish_reason is None assert chunk.provider == "local_ollama" def test_parse_ndjson_line_invalid_json(self) -> None: assert LocalOllamaProvider._parse_ndjson_line("{invalid}", "llama3") is None async def test_embed_success(self) -> None: with _mock_http(self._MODULE, _ollama_embed_handler): provider = LocalOllamaProvider() result = await provider.embed("hello", "nomic-embed-text") assert result == [0.5, 0.6] async def test_embed_not_configured_raises(self) -> None: provider = LocalOllamaProvider(base_url="") with pytest.raises(AILLMUnavailableError): await provider.embed("hello", "nomic-embed-text") # --------------------------------------------------------------------------- # create_failover_chain factory # --------------------------------------------------------------------------- class TestCreateFailoverChain: def test_creates_chain_with_configured_providers(self) -> None: settings = Settings( _env_file=None, openai_api_key="sk-openai", anthropic_api_key="sk-anthropic", baichuan_api_key="sk-baichuan", ollama_base_url="http://localhost:11434", llm_provider_priority="openai,anthropic,baichuan,local_ollama", ) chain = create_failover_chain(settings) names = [p.name for p in chain.providers] assert names == ["openai", "anthropic", "baichuan", "local_ollama"] def test_falls_back_to_openai_when_no_providers(self) -> None: settings = Settings(_env_file=None, llm_provider_priority="") chain = create_failover_chain(settings) assert len(chain.providers) == 1 assert chain.providers[0].name == "openai" def test_priority_order_respected(self) -> None: settings = Settings( _env_file=None, openai_api_key="sk-openai", anthropic_api_key="sk-anthropic", llm_provider_priority="anthropic,openai", ) chain = create_failover_chain(settings) names = [p.name for p in chain.providers] assert names == ["anthropic", "openai"] # --------------------------------------------------------------------------- # ProviderFailoverChain.embed # --------------------------------------------------------------------------- class TestFailoverChainEmbed: async def test_embed_success(self) -> None: provider = _EmbeddableProvider(name="p1", embedding=[0.5, 0.6]) chain = ProviderFailoverChain([provider], CircuitBreaker()) embedding, provider_name = await chain.embed("hello", "model") assert embedding == [0.5, 0.6] assert provider_name == "p1" async def test_embed_no_provider_raises(self) -> None: # MockProvider does not implement embed → NotImplementedError, skipped provider = MockProvider(name="p1") chain = ProviderFailoverChain([provider], CircuitBreaker()) with pytest.raises(AILLMUnavailableError): await chain.embed("hello", "model")