|
21 | 21 | from google.adk.events.event import Event |
22 | 22 | from google.adk.models.cache_metadata import CacheMetadata |
23 | 23 | from google.adk.sessions.base_session_service import BaseSessionService |
| 24 | +from google.adk.sessions.in_memory_session_service import InMemorySessionService |
24 | 25 | from google.adk.sessions.session import Session |
25 | 26 | from google.adk.utils.cache_performance_analyzer import CachePerformanceAnalyzer |
26 | 27 | from google.genai import types |
27 | 28 | import pytest |
28 | 29 |
|
29 | 30 |
|
| 31 | +@pytest.mark.parametrize( |
| 32 | + "snapshots,expected_total,expected_average", |
| 33 | + [ |
| 34 | + ([("cache1", 1), ("cache1", 2), ("cache1", 3)], 3, 3.0), |
| 35 | + ([("cache1", 1), ("cache1", 1), ("cache1", 2)], 2, 2.0), |
| 36 | + ([("cache1", 1), ("cache1", 2), ("cache2", 1)], 3, 1.5), |
| 37 | + ([("cache1", 3)], 3, 3.0), |
| 38 | + ([("cache1", 3), ("cache1", 2)], 3, 3.0), |
| 39 | + ([(None, None), ("cache1", 2), (None, None)], 2, 2.0), |
| 40 | + ([(None, None)], 0, 0), |
| 41 | + ([("cache1", 0), ("cache1", 0)], 0, 0), |
| 42 | + ], |
| 43 | + ids=[ |
| 44 | + "successive-invocations", |
| 45 | + "multiple-responses-per-invocation", |
| 46 | + "cache-refresh-with-same-fingerprint", |
| 47 | + "single-snapshot", |
| 48 | + "out-of-order-snapshots", |
| 49 | + "fingerprint-only-around-active-cache", |
| 50 | + "fingerprint-only", |
| 51 | + "unused-cache", |
| 52 | + ], |
| 53 | +) |
| 54 | +async def test_cache_invocation_totals_count_each_cache_once( |
| 55 | + snapshots: list[tuple[str | None, int | None]], |
| 56 | + expected_total: int, |
| 57 | + expected_average: float, |
| 58 | +) -> None: |
| 59 | + """Cumulative cache counters are not added once per response.""" |
| 60 | + service = InMemorySessionService() |
| 61 | + session = await service.create_session(app_name="app", user_id="user") |
| 62 | + for cache_name, count in snapshots: |
| 63 | + await service.append_event( |
| 64 | + session=session, |
| 65 | + event=Event( |
| 66 | + author="agent", |
| 67 | + invocation_id=f"invocation{count}", |
| 68 | + cache_metadata=CacheMetadata( |
| 69 | + cache_name=cache_name, |
| 70 | + invocations_used=count, |
| 71 | + expire_time=2_000_000_000 if cache_name is not None else None, |
| 72 | + fingerprint="prefix", |
| 73 | + contents_count=2, |
| 74 | + ), |
| 75 | + usage_metadata=types.GenerateContentResponseUsageMetadata( |
| 76 | + prompt_token_count=100, |
| 77 | + cached_content_token_count=80 if cache_name is not None else 0, |
| 78 | + ), |
| 79 | + ), |
| 80 | + ) |
| 81 | + |
| 82 | + result = await CachePerformanceAnalyzer( |
| 83 | + service |
| 84 | + ).analyze_agent_cache_performance(session.id, "user", "app", "agent") |
| 85 | + |
| 86 | + assert result["total_invocations"] == expected_total |
| 87 | + assert result["avg_invocations_used"] == expected_average |
| 88 | + assert result["cache_refreshes"] == len( |
| 89 | + {name for name, _ in snapshots if name is not None} |
| 90 | + ) |
| 91 | + assert result["total_requests"] == len(snapshots) |
| 92 | + assert result["total_prompt_tokens"] == 100 * len(snapshots) |
| 93 | + assert result["total_cached_tokens"] == 80 * sum( |
| 94 | + name is not None for name, _ in snapshots |
| 95 | + ) |
| 96 | + |
| 97 | + |
30 | 98 | class TestCachePerformanceAnalyzer: |
31 | 99 | """Test suite for CachePerformanceAnalyzer.""" |
32 | 100 |
|
|
0 commit comments