File size: 4,901 Bytes
1993d5c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
104
105
106
107
108
109
110
111
112
113
114
115
116
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(monkeypatch):
    """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 = []

    mock_agent_config = AgentConfig(
        name="geometry_parser",
        description="Testing cascade",
        tiers=[
            ModelTier(model="gemini/gemini-3.5-flash-lite", max_attempts=1, reasoning_effort="low"),
            ModelTier(model="gemini/gemini-3.5-flash", max_attempts=1, reasoning_effort="medium"),
        ],
        temperature=0.1,
        max_tokens=4096,
        timeout_seconds=60,
    )
    monkeypatch.setattr("agents.runtime.load_agent_config", lambda agent: mock_agent_config)

    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"]