File size: 4,377 Bytes
0772b5a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
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


@pytest.mark.asyncio
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")


@pytest.mark.asyncio
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"]