Spaces:
Running
Running
| import pytest | |
| import asyncio | |
| from typing import Tuple, Any, Dict, List | |
| from config.schemas import ModelTier, AgentConfig, AgentModelsConfig | |
| from config.loader import load_agent_config, AgentConfigResolver | |
| from config.settings import Settings, parse_comma_separated_keys | |
| from llm.errors import ErrorCategory, ErrorClassifier | |
| from llm.key_state import KeyState, KeyMetadata, MemoryKeyStateStore, hash_key | |
| from llm.key_pool import APIKeyPool | |
| from llm.telemetry import LLMTelemetryRecord, LLMTelemetry | |
| from agents.runtime import AgentRuntime | |
| def test_parse_comma_separated_keys(): | |
| raw = "key1, key2 , 'key3', \"key4\"" | |
| keys = parse_comma_separated_keys(raw) | |
| assert keys == ["key1", "key2", "key3", "key4"] | |
| def test_agent_models_config_schema_validation(): | |
| tier1 = ModelTier(model="gemini/gemini-3.5-flash-lite", max_attempts=1, reasoning_effort="low") | |
| tier2 = ModelTier(model="gemini/gemini-3.5-flash", max_attempts=1, reasoning_effort="medium") | |
| agent = AgentConfig( | |
| name="test_agent", | |
| description="Testing agent schema", | |
| tiers=[tier1, tier2], | |
| temperature=0.1, | |
| max_tokens=4096, | |
| timeout_seconds=60, | |
| ) | |
| config = AgentModelsConfig(version=2, agents={"test_agent": agent}) | |
| assert config.version == 2 | |
| assert "test_agent" in config.agents | |
| assert len(config.agents["test_agent"].tiers) == 2 | |
| assert config.agents["test_agent"].tiers[0].reasoning_effort == "low" | |
| assert config.agents["test_agent"].tiers[1].reasoning_effort == "medium" | |
| def test_error_classifier(): | |
| assert ErrorClassifier.classify(Exception("429 Too Many Requests")) == ErrorCategory.RATE_LIMIT | |
| assert ErrorClassifier.classify(Exception("API_KEY_INVALID: User not authorized")) == ErrorCategory.AUTH_ERROR | |
| assert ErrorClassifier.classify(Exception("Daily quota exceeded for project")) == ErrorCategory.QUOTA_EXHAUSTED | |
| assert ErrorClassifier.classify(Exception("Connection reset by peer")) == ErrorCategory.NETWORK | |
| assert ErrorClassifier.classify(Exception("Internal Server Error 500")) == ErrorCategory.SERVER_ERROR | |
| assert ErrorClassifier.classify(Exception("Request timed out")) == ErrorCategory.TIMEOUT | |
| async def test_key_pool_round_robin_and_cooldown(): | |
| store = MemoryKeyStateStore() | |
| pool = APIKeyPool(state_store=store) | |
| custom_prov = "test_custom_prov" | |
| pool.register_keys(custom_prov, ["key_alpha", "key_beta", "key_gamma"]) | |
| # First rotation | |
| k1, h1 = await pool.get_next_key(custom_prov) | |
| k2, h2 = await pool.get_next_key(custom_prov) | |
| k3, h3 = await pool.get_next_key(custom_prov) | |
| assert [k1, k2, k3] == ["key_alpha", "key_beta", "key_gamma"] | |
| # Put key_alpha on cooldown | |
| await pool.mark_cooldown("key_alpha", retry_after=120) | |
| # Next key should skip key_alpha | |
| k_next, _ = await pool.get_next_key(custom_prov) | |
| assert k_next in ("key_beta", "key_gamma") | |
| async def test_agent_runtime_validator_cascade(): | |
| """Simulates Tier 1 (3.5-flash-lite, low) failing validation and Tier 2 (3.5-flash, medium) succeeding validation on geometry_parser.""" | |
| call_history = [] | |
| reasoning_efforts = [] | |
| class MockLLMService: | |
| async def acomplete(self, model: str, messages: list, reasoning_effort: str = None, **kwargs) -> str: | |
| call_history.append(model) | |
| reasoning_efforts.append(reasoning_effort) | |
| if "lite" in model: | |
| return "INVALID_OUTPUT_FROM_TIER_1" | |
| return '{"type": "pyramid", "analysis": "Valid analysis from Tier 2"}' | |
| runtime = AgentRuntime(llm_service=MockLLMService()) | |
| def mock_validator(raw_output: str) -> Tuple[bool, Any]: | |
| if "INVALID" in raw_output: | |
| return False, "Malformed analysis output" | |
| return True, {"valid": True, "raw": raw_output} | |
| messages = [{"role": "user", "content": "Analyze problem"}] | |
| res = await runtime.run( | |
| agent="geometry_parser", | |
| messages=messages, | |
| validator=mock_validator, | |
| ) | |
| assert res["valid"] is True | |
| # Verify that Tier 1 (lite) was attempted with reasoning_effort='low' and escalated to Tier 2 with reasoning_effort='medium' | |
| assert any("lite" in m for m in call_history) | |
| assert any("3.5-flash" in m and "lite" not in m for m in call_history) | |
| assert reasoning_efforts == ["low", "medium"] | |