"""Fault injection tests for Model Gateway transient errors, 429 backoff, 5xx, and failovers.""" from __future__ import annotations import asyncio from typing import Any, Dict, List import httpx from src.runtime.core.config import ( ModelRoleConfig, RuntimeConfig, RuntimeLimits, RuntimeObservabilityConfig, RuntimePricing, RuntimeStoragePaths, ) from src.runtime.gateway.adapters import ProviderAdapter from src.runtime.gateway.client import ModelGatewayClient class FaultyMockAdapter(ProviderAdapter): def __init__(self, responses: List[Any]): super().__init__("faulty_provider") self.responses = list(responses) self.call_count = 0 async def execute_call( self, model: str, messages: List[Dict[str, str]], temperature: float = 0.0, timeout_seconds: int = 30, response_format: Any = None, ) -> Dict[str, Any]: self.call_count += 1 if not self.responses: raise IOError("No more fault injection responses configured") curr = self.responses.pop(0) if isinstance(curr, Exception): raise curr return curr def create_fault_test_config() -> RuntimeConfig: return RuntimeConfig( config_version="1.0.0", paths=RuntimeStoragePaths(), roles={ "runtime_primary": ModelRoleConfig( role_config_version="1.0.0", provider="groq", model="llama-3.1-8b-instant", endpoint_url="https://api.groq.com/openai/v1", timeout_seconds=2.0, max_retries=3, parameters={"temperature": 0.0}, ), "runtime_fallback": ModelRoleConfig( role_config_version="1.0.0", provider="deepseek", model="deepseek-chat", endpoint_url="https://api.deepseek.com/v1", timeout_seconds=2.0, max_retries=3, parameters={"temperature": 0.0}, ), }, prompts={}, ecp={}, limits=RuntimeLimits(), pricing=RuntimePricing(), langfuse=RuntimeObservabilityConfig(), sqlite_busy_timeout_ms=2000, raw_config_bytes_sha256="abc", ) def test_gateway_fault_http_429_rate_limit_and_fallback(): async def _test(): config = create_fault_test_config() client = ModelGatewayClient(config) # Primary repeatedly throws 429 req = httpx.Request("POST", "https://api.groq.com/openai/v1/chat/completions") resp_429 = httpx.Response(429, request=req) primary_mock = FaultyMockAdapter( [ httpx.HTTPStatusError("Rate limit exceeded", request=req, response=resp_429), httpx.HTTPStatusError("Rate limit exceeded", request=req, response=resp_429), httpx.HTTPStatusError("Rate limit exceeded", request=req, response=resp_429), ] ) # Fallback recovers successfully fallback_mock = FaultyMockAdapter( [ { "choices": [{"message": {"content": '{"status": "recovered_by_fallback"}'}}], "usage": {"prompt_tokens": 50, "completion_tokens": 10}, } ] ) client.register_adapter("groq", primary_mock) client.register_adapter("deepseek", fallback_mock) resp = await client.execute_structured_call( messages=[{"role": "user", "content": "test"}], schema_dict={"type": "object"}, ) assert resp.status == "success" assert resp.effective_role == "runtime_fallback" assert resp.used_fallback is True assert resp.content_json == {"status": "recovered_by_fallback"} assert primary_mock.call_count == 3 assert fallback_mock.call_count == 1 asyncio.run(_test()) def test_gateway_fault_http_500_server_error_and_fallback(): async def _test(): config = create_fault_test_config() client = ModelGatewayClient(config) req = httpx.Request("POST", "https://api.groq.com/openai/v1/chat/completions") resp_500 = httpx.Response(500, request=req) primary_mock = FaultyMockAdapter( [ httpx.HTTPStatusError("Internal Server Error", request=req, response=resp_500), httpx.HTTPStatusError("Internal Server Error", request=req, response=resp_500), httpx.HTTPStatusError("Internal Server Error", request=req, response=resp_500), ] ) fallback_mock = FaultyMockAdapter( [ { "choices": [{"message": {"content": '{"status": "recovered_from_500"}'}}], "usage": {"prompt_tokens": 60, "completion_tokens": 15}, } ] ) client.register_adapter("groq", primary_mock) client.register_adapter("deepseek", fallback_mock) resp = await client.execute_structured_call( messages=[{"role": "user", "content": "test"}], schema_dict={"type": "object"}, ) assert resp.status == "success" assert resp.effective_role == "runtime_fallback" assert resp.used_fallback is True assert resp.content_json == {"status": "recovered_from_500"} asyncio.run(_test())