mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: keep session limiter migration state together
This commit is contained in:
parent
7bdfe7884c
commit
d21a5ec97f
5 changed files with 435 additions and 187 deletions
|
|
@ -43,16 +43,15 @@ else:
|
|||
# Cluster. Old proxy instances continue to update the aggregate key.
|
||||
MAX_BUDGET_SESSION_INCREMENT_SCRIPT: Final = """
|
||||
local legacy_key = KEYS[1]
|
||||
local total_new_key = KEYS[2]
|
||||
local agent_key = KEYS[3]
|
||||
local agent_scope_key = KEYS[2]
|
||||
local amount = tonumber(ARGV[1])
|
||||
local ttl = tonumber(ARGV[2])
|
||||
local agent_field = ARGV[3]
|
||||
if amount == nil or amount <= 0 or amount ~= amount or math.abs(amount) == math.huge then
|
||||
return redis.error_reply('agent session spend increment is invalid')
|
||||
end
|
||||
|
||||
if redis.call('EXISTS', total_new_key) == 0 then
|
||||
if redis.call('EXISTS', agent_key) == 1 then
|
||||
return redis.error_reply('agent session spend exists without migration total')
|
||||
end
|
||||
|
||||
if redis.call('EXISTS', agent_scope_key) == 0 then
|
||||
if redis.call('EXISTS', legacy_key) == 0 then
|
||||
redis.call('SET', legacy_key, '0')
|
||||
redis.call('PEXPIRE', legacy_key, ttl * 1000)
|
||||
|
|
@ -62,10 +61,13 @@ if redis.call('EXISTS', total_new_key) == 0 then
|
|||
if legacy_ttl == -2 then
|
||||
return redis.error_reply('legacy session spend expired during migration')
|
||||
end
|
||||
if legacy_ttl >= 0 and legacy_ttl <= 1 then
|
||||
return redis.error_reply('legacy session spend is expiring before agent scope')
|
||||
end
|
||||
|
||||
redis.call('SET', total_new_key, '0')
|
||||
redis.call('HSET', agent_scope_key, '__total_new', '0')
|
||||
if legacy_ttl >= 0 then
|
||||
redis.call('PEXPIRE', total_new_key, legacy_ttl + 1000)
|
||||
redis.call('PEXPIRE', agent_scope_key, legacy_ttl - 1)
|
||||
end
|
||||
end
|
||||
|
||||
|
|
@ -73,35 +75,60 @@ if redis.call('EXISTS', legacy_key) == 0 then
|
|||
return redis.error_reply('legacy session spend expired before agent scope')
|
||||
end
|
||||
|
||||
local legacy_value = tonumber(redis.call('INCRBYFLOAT', legacy_key, amount))
|
||||
local total_new_value = tonumber(redis.call('INCRBYFLOAT', total_new_key, amount))
|
||||
local agent_existed = redis.call('EXISTS', agent_key)
|
||||
local agent_value = tonumber(redis.call('INCRBYFLOAT', agent_key, amount))
|
||||
if agent_existed == 0 then
|
||||
local migration_ttl = redis.call('PTTL', total_new_key)
|
||||
if migration_ttl >= 0 then
|
||||
redis.call('PEXPIRE', agent_key, migration_ttl + 1000)
|
||||
end
|
||||
local total_new_raw = redis.call('HGET', agent_scope_key, '__total_new')
|
||||
if total_new_raw == false then
|
||||
return redis.error_reply('agent session scope is missing its migration total')
|
||||
end
|
||||
local total_new_value = tonumber(total_new_raw)
|
||||
if total_new_value == nil or total_new_value ~= total_new_value or math.abs(total_new_value) == math.huge then
|
||||
return redis.error_reply('agent session migration total is not numeric')
|
||||
end
|
||||
local agent_raw = redis.call('HGET', agent_scope_key, agent_field)
|
||||
local agent_value = tonumber(agent_raw or '0')
|
||||
if agent_value == nil or agent_value ~= agent_value or math.abs(agent_value) == math.huge then
|
||||
return redis.error_reply('agent session counter is not numeric')
|
||||
end
|
||||
if math.abs(total_new_value + amount) == math.huge or math.abs(agent_value + amount) == math.huge then
|
||||
return redis.error_reply('agent session spend increment is out of range')
|
||||
end
|
||||
|
||||
return tostring(math.max(legacy_value - total_new_value, 0) + agent_value)
|
||||
local legacy_value = tonumber(redis.call('INCRBYFLOAT', legacy_key, amount))
|
||||
local next_total_new = tonumber(redis.call('HINCRBYFLOAT', agent_scope_key, '__total_new', amount))
|
||||
local next_agent_value = tonumber(redis.call('HINCRBYFLOAT', agent_scope_key, agent_field, amount))
|
||||
|
||||
return tostring(math.max(legacy_value - next_total_new, 0) + next_agent_value)
|
||||
"""
|
||||
|
||||
MAX_BUDGET_SESSION_GET_AGENT_SPEND_SCRIPT: Final = """
|
||||
local legacy_value = tonumber(redis.call('GET', KEYS[1])) or 0
|
||||
local total_new_value = tonumber(redis.call('GET', KEYS[2]))
|
||||
local agent_value = tonumber(redis.call('GET', KEYS[3])) or 0
|
||||
local legacy_key = KEYS[1]
|
||||
local agent_scope_key = KEYS[2]
|
||||
local legacy_value = tonumber(redis.call('GET', legacy_key)) or 0
|
||||
|
||||
if total_new_value == nil then
|
||||
if redis.call('EXISTS', KEYS[3]) == 1 then
|
||||
return redis.error_reply('agent session spend exists without migration total')
|
||||
end
|
||||
return tostring(math.max(legacy_value, agent_value))
|
||||
if redis.call('EXISTS', agent_scope_key) == 0 then
|
||||
return tostring(legacy_value)
|
||||
end
|
||||
if redis.call('EXISTS', KEYS[1]) == 0 then
|
||||
|
||||
if redis.call('EXISTS', legacy_key) == 0 then
|
||||
return redis.error_reply('legacy session spend expired before agent scope')
|
||||
end
|
||||
|
||||
local total_new_raw = redis.call('HGET', agent_scope_key, '__total_new')
|
||||
if total_new_raw == false then
|
||||
return redis.error_reply('agent session scope is missing its migration total')
|
||||
end
|
||||
local total_new_value = tonumber(total_new_raw)
|
||||
if total_new_value == nil or total_new_value ~= total_new_value or math.abs(total_new_value) == math.huge then
|
||||
return redis.error_reply('agent session migration total is not numeric')
|
||||
end
|
||||
local agent_raw = redis.call('HGET', agent_scope_key, ARGV[1])
|
||||
local agent_value = 0
|
||||
if agent_raw ~= false then
|
||||
agent_value = tonumber(agent_raw)
|
||||
if agent_value == nil or agent_value ~= agent_value or math.abs(agent_value) == math.huge then
|
||||
return redis.error_reply('agent session counter is not numeric')
|
||||
end
|
||||
end
|
||||
|
||||
return tostring(math.max(legacy_value - total_new_value, 0) + agent_value)
|
||||
"""
|
||||
|
||||
|
|
@ -125,6 +152,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
self._local_lock = asyncio.Lock()
|
||||
self.increment_script: Callable[..., Awaitable[object]] | None = None
|
||||
self.get_agent_spend_script: Callable[..., Awaitable[object]] | None = None
|
||||
self._registered_redis_cache: object | None = None
|
||||
self.ttl = int(
|
||||
os.getenv(
|
||||
"LITELLM_MAX_BUDGET_PER_SESSION_TTL",
|
||||
|
|
@ -132,19 +160,27 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
)
|
||||
)
|
||||
|
||||
if self.internal_usage_cache.dual_cache.redis_cache is not None:
|
||||
self.increment_script = cast(
|
||||
Callable[..., Awaitable[object]],
|
||||
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
|
||||
MAX_BUDGET_SESSION_INCREMENT_SCRIPT
|
||||
),
|
||||
)
|
||||
self.get_agent_spend_script = cast(
|
||||
Callable[..., Awaitable[object]],
|
||||
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
|
||||
MAX_BUDGET_SESSION_GET_AGENT_SPEND_SCRIPT
|
||||
),
|
||||
)
|
||||
self._ensure_redis_scripts()
|
||||
|
||||
def _ensure_redis_scripts(self) -> None:
|
||||
redis_cache = self.internal_usage_cache.dual_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
self.increment_script = None
|
||||
self.get_agent_spend_script = None
|
||||
self._registered_redis_cache = None
|
||||
return
|
||||
if redis_cache is self._registered_redis_cache:
|
||||
return
|
||||
|
||||
self.increment_script = cast(
|
||||
Callable[..., Awaitable[object]],
|
||||
redis_cache.async_register_script(MAX_BUDGET_SESSION_INCREMENT_SCRIPT),
|
||||
)
|
||||
self.get_agent_spend_script = cast(
|
||||
Callable[..., Awaitable[object]],
|
||||
redis_cache.async_register_script(MAX_BUDGET_SESSION_GET_AGENT_SPEND_SCRIPT),
|
||||
)
|
||||
self._registered_redis_cache = redis_cache
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
|
|
@ -265,27 +301,52 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
return float(max_budget)
|
||||
return None
|
||||
|
||||
def _make_cache_key(self, session_id: str, agent_id: str) -> str:
|
||||
def _make_agent_scope_cache_key(self, session_id: str) -> str:
|
||||
return f"{{session_budget:{session_id}}}:agent-scope"
|
||||
|
||||
def _make_agent_scope_field(self, agent_id: str) -> str:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
stable_agent_id: Final = json.dumps(global_agent_registry.stable_agent_id(agent_id), separators=(",", ":"))
|
||||
return f"{{session_budget:{session_id}}}:agent:{stable_agent_id}:spend"
|
||||
return f"agent:{stable_agent_id}"
|
||||
|
||||
def _make_legacy_cache_key(self, session_id: str) -> str:
|
||||
return f"{{session_budget:{session_id}}}:spend"
|
||||
|
||||
def _make_total_new_cache_key(self, session_id: str) -> str:
|
||||
return f"{{session_budget:{session_id}}}:agent-scope-total"
|
||||
async def _get_local_scope(self, cache_key: str) -> dict[str, object] | None:
|
||||
result: Final[object | None] = cast(
|
||||
object | None,
|
||||
await self.internal_usage_cache.async_get_cache(
|
||||
key=cache_key,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
),
|
||||
)
|
||||
if result is None:
|
||||
return None
|
||||
if isinstance(result, dict) and all(isinstance(key, str) for key in result):
|
||||
return cast(dict[str, object], result)
|
||||
raise RuntimeError("Agent session scope cache has an invalid value")
|
||||
|
||||
@staticmethod
|
||||
def _get_scope_spend(scope: dict[str, object], field: str) -> float:
|
||||
value = scope.get(field)
|
||||
if value is None:
|
||||
return 0.0
|
||||
if isinstance(value, (int, float, str, bytes)):
|
||||
return float(value)
|
||||
raise RuntimeError("Agent session scope has an invalid counter")
|
||||
|
||||
async def _get_agent_spend(self, session_id: str, agent_id: str) -> float:
|
||||
legacy_key: Final = self._make_legacy_cache_key(session_id)
|
||||
total_new_key: Final = self._make_total_new_cache_key(session_id)
|
||||
agent_key: Final = self._make_cache_key(session_id, agent_id)
|
||||
agent_scope_key: Final = self._make_agent_scope_cache_key(session_id)
|
||||
agent_field: Final = self._make_agent_scope_field(agent_id)
|
||||
self._ensure_redis_scripts()
|
||||
if self.get_agent_spend_script is not None:
|
||||
try:
|
||||
result: Final[object] = await self.get_agent_spend_script(
|
||||
keys=[legacy_key, total_new_key, agent_key],
|
||||
args=[],
|
||||
keys=[legacy_key, agent_scope_key],
|
||||
args=[agent_field],
|
||||
)
|
||||
if isinstance(result, (int, float, str, bytes)):
|
||||
return float(result)
|
||||
|
|
@ -300,15 +361,19 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
raise
|
||||
|
||||
legacy_value: Final = await self._get_local_spend(legacy_key)
|
||||
total_new_value: Final = await self._get_local_spend(total_new_key)
|
||||
agent_value: Final = await self._get_local_spend(agent_key)
|
||||
agent_scope: Final = await self._get_local_scope(agent_scope_key)
|
||||
if legacy_value is None:
|
||||
if total_new_value is not None:
|
||||
if agent_scope is not None:
|
||||
raise RuntimeError("Legacy session spend expired before agent scope")
|
||||
return float(agent_value or 0.0)
|
||||
if total_new_value is None:
|
||||
return max(float(legacy_value), float(agent_value or 0.0))
|
||||
return max(float(legacy_value) - float(total_new_value), 0.0) + float(agent_value or 0.0)
|
||||
return 0.0
|
||||
if agent_scope is None:
|
||||
return float(legacy_value)
|
||||
raw_total_new = agent_scope.get("__total_new")
|
||||
if not isinstance(raw_total_new, (int, float, str, bytes)):
|
||||
raise RuntimeError("Agent session scope is missing its migration total")
|
||||
total_new_value = float(raw_total_new)
|
||||
agent_value = self._get_scope_spend(agent_scope, agent_field)
|
||||
return max(float(legacy_value) - total_new_value, 0.0) + agent_value
|
||||
|
||||
async def _get_local_spend(self, cache_key: str) -> float | None:
|
||||
result: Final[object | None] = cast(
|
||||
|
|
@ -325,13 +390,14 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
|
||||
async def _increment_agent_spend(self, session_id: str, agent_id: str, amount: float) -> float:
|
||||
legacy_key: Final = self._make_legacy_cache_key(session_id)
|
||||
total_new_key: Final = self._make_total_new_cache_key(session_id)
|
||||
agent_key: Final = self._make_cache_key(session_id, agent_id)
|
||||
agent_scope_key: Final = self._make_agent_scope_cache_key(session_id)
|
||||
agent_field: Final = self._make_agent_scope_field(agent_id)
|
||||
self._ensure_redis_scripts()
|
||||
if self.increment_script is not None:
|
||||
try:
|
||||
result: Final[object] = await self.increment_script(
|
||||
keys=[legacy_key, total_new_key, agent_key],
|
||||
args=[str(amount), self.ttl],
|
||||
keys=[legacy_key, agent_scope_key],
|
||||
args=[str(amount), self.ttl, agent_field],
|
||||
)
|
||||
if isinstance(result, (int, float, str, bytes)):
|
||||
return float(result)
|
||||
|
|
@ -347,35 +413,39 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
|
||||
async with self._local_lock:
|
||||
legacy_value = await self._get_local_spend(legacy_key)
|
||||
total_new_value = await self._get_local_spend(total_new_key)
|
||||
agent_value = await self._get_local_spend(agent_key)
|
||||
agent_scope = await self._get_local_scope(agent_scope_key)
|
||||
if legacy_value is None:
|
||||
if total_new_value is not None:
|
||||
if agent_scope is not None:
|
||||
raise RuntimeError("Legacy session spend expired before agent scope")
|
||||
legacy_value = 0.0
|
||||
total_new_value = total_new_value or 0.0
|
||||
agent_value = agent_value or 0.0
|
||||
if agent_scope is None:
|
||||
total_new_value = 0.0
|
||||
agent_value = 0.0
|
||||
else:
|
||||
raw_total_new = agent_scope.get("__total_new")
|
||||
if not isinstance(raw_total_new, (int, float, str, bytes)):
|
||||
raise RuntimeError("Agent session scope is missing its migration total")
|
||||
total_new_value = float(raw_total_new)
|
||||
agent_value = self._get_scope_spend(agent_scope, agent_field)
|
||||
new_legacy: Final = float(legacy_value) + amount
|
||||
new_total_new: Final = float(total_new_value) + amount
|
||||
new_agent: Final = float(agent_value) + amount
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=legacy_key,
|
||||
value=new_legacy,
|
||||
ttl=self.ttl if legacy_value == 0.0 else None,
|
||||
ttl=self.ttl,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=total_new_key,
|
||||
value=new_total_new,
|
||||
ttl=self.ttl + 1 if total_new_value == 0.0 else None,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=agent_key,
|
||||
value=new_agent,
|
||||
ttl=self.ttl + 1 if agent_value == 0.0 else None,
|
||||
key=agent_scope_key,
|
||||
value={
|
||||
**(agent_scope or {}),
|
||||
"__total_new": new_total_new,
|
||||
agent_field: new_agent,
|
||||
},
|
||||
# Keep this map on a shorter rolling TTL than the aggregate fallback.
|
||||
ttl=max(self.ttl - 1, 0),
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -45,14 +45,10 @@ if #KEYS == 1 then
|
|||
return current
|
||||
end
|
||||
|
||||
local total_new_key = KEYS[2]
|
||||
local agent_key = KEYS[3]
|
||||
local agent_scope_key = KEYS[2]
|
||||
local ttl = tonumber(ARGV[1])
|
||||
if redis.call('EXISTS', total_new_key) == 0 then
|
||||
if redis.call('EXISTS', agent_key) == 1 then
|
||||
return redis.error_reply('agent session count exists without migration total')
|
||||
end
|
||||
|
||||
local agent_field = ARGV[2]
|
||||
if redis.call('EXISTS', agent_scope_key) == 0 then
|
||||
if redis.call('EXISTS', legacy_key) == 0 then
|
||||
redis.call('SET', legacy_key, '0')
|
||||
redis.call('PEXPIRE', legacy_key, ttl * 1000)
|
||||
|
|
@ -62,10 +58,13 @@ if redis.call('EXISTS', total_new_key) == 0 then
|
|||
if legacy_ttl == -2 then
|
||||
return redis.error_reply('legacy session count expired during migration')
|
||||
end
|
||||
if legacy_ttl >= 0 and legacy_ttl <= 1 then
|
||||
return redis.error_reply('legacy session count is expiring before agent scope')
|
||||
end
|
||||
|
||||
redis.call('SET', total_new_key, '0')
|
||||
redis.call('HSET', agent_scope_key, '__total_new', '0')
|
||||
if legacy_ttl >= 0 then
|
||||
redis.call('PEXPIRE', total_new_key, legacy_ttl + 1000)
|
||||
redis.call('PEXPIRE', agent_scope_key, legacy_ttl - 1)
|
||||
end
|
||||
end
|
||||
|
||||
|
|
@ -73,16 +72,30 @@ if redis.call('EXISTS', legacy_key) == 0 then
|
|||
return redis.error_reply('legacy session count expired before agent scope')
|
||||
end
|
||||
|
||||
local legacy_value = redis.call('INCR', legacy_key)
|
||||
local total_new_value = redis.call('INCR', total_new_key)
|
||||
local agent_existed = redis.call('EXISTS', agent_key)
|
||||
local agent_value = redis.call('INCR', agent_key)
|
||||
if agent_existed == 0 then
|
||||
local migration_ttl = redis.call('PTTL', total_new_key)
|
||||
if migration_ttl >= 0 then
|
||||
redis.call('PEXPIRE', agent_key, migration_ttl + 1000)
|
||||
end
|
||||
local total_new_raw = redis.call('HGET', agent_scope_key, '__total_new')
|
||||
if total_new_raw == false then
|
||||
return redis.error_reply('agent session scope is missing its migration total')
|
||||
end
|
||||
if string.match(total_new_raw, '^%d+$') == nil then
|
||||
return redis.error_reply('agent session migration total is not an integer')
|
||||
end
|
||||
local validated_total_new = tonumber(total_new_raw)
|
||||
if validated_total_new == nil or validated_total_new >= 9223372036854774784 then
|
||||
return redis.error_reply('agent session migration total is out of range')
|
||||
end
|
||||
|
||||
local agent_raw = redis.call('HGET', agent_scope_key, agent_field)
|
||||
if agent_raw ~= false and string.match(agent_raw, '^%d+$') == nil then
|
||||
return redis.error_reply('agent session counter is not an integer')
|
||||
end
|
||||
local validated_agent_value = tonumber(agent_raw or '0')
|
||||
if validated_agent_value >= 9223372036854774784 then
|
||||
return redis.error_reply('agent session counter is out of range')
|
||||
end
|
||||
|
||||
local legacy_value = redis.call('INCR', legacy_key)
|
||||
local total_new_value = redis.call('HINCRBY', agent_scope_key, '__total_new', 1)
|
||||
local agent_value = redis.call('HINCRBY', agent_scope_key, agent_field, 1)
|
||||
|
||||
return math.max(legacy_value - total_new_value, 0) + agent_value
|
||||
"""
|
||||
|
|
@ -117,14 +130,25 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
self.internal_usage_cache = internal_usage_cache
|
||||
self._local_lock = asyncio.Lock()
|
||||
self.increment_script: Callable[..., Awaitable[object]] | None = None
|
||||
self._registered_redis_cache: object | None = None
|
||||
self.ttl = int(os.getenv("LITELLM_MAX_ITERATIONS_TTL", DEFAULT_MAX_ITERATIONS_TTL))
|
||||
|
||||
# Register Lua script with Redis if available (same pattern as v3 limiter)
|
||||
if self.internal_usage_cache.dual_cache.redis_cache is not None:
|
||||
self.increment_script = cast(
|
||||
Callable[..., Awaitable[object]],
|
||||
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(MAX_ITERATIONS_INCREMENT_SCRIPT),
|
||||
)
|
||||
self._ensure_redis_scripts()
|
||||
|
||||
def _ensure_redis_scripts(self) -> None:
|
||||
redis_cache = self.internal_usage_cache.dual_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
self.increment_script = None
|
||||
self._registered_redis_cache = None
|
||||
return
|
||||
if redis_cache is self._registered_redis_cache:
|
||||
return
|
||||
|
||||
self.increment_script = cast(
|
||||
Callable[..., Awaitable[object]],
|
||||
redis_cache.async_register_script(MAX_ITERATIONS_INCREMENT_SCRIPT),
|
||||
)
|
||||
self._registered_redis_cache = redis_cache
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
|
|
@ -219,25 +243,32 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
return int(max_iterations)
|
||||
return None
|
||||
|
||||
def _make_cache_key(self, session_id: str, agent_id: str | None = None) -> str:
|
||||
"""
|
||||
Create cache key for session iteration counter.
|
||||
|
||||
Agent-scoped counters share the legacy session hash tag so migration
|
||||
scripts can atomically update both scopes on Redis Cluster.
|
||||
"""
|
||||
if agent_id is None:
|
||||
return f"{{session_iterations:{session_id}}}:count"
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
stable_agent_id: Final = json.dumps(global_agent_registry.stable_agent_id(agent_id), separators=(",", ":"))
|
||||
return f"{{session_iterations:{session_id}}}:agent:{stable_agent_id}:count"
|
||||
|
||||
def _make_legacy_cache_key(self, session_id: str) -> str:
|
||||
return f"{{session_iterations:{session_id}}}:count"
|
||||
|
||||
def _make_total_new_cache_key(self, session_id: str) -> str:
|
||||
return f"{{session_iterations:{session_id}}}:agent-scope-total"
|
||||
def _make_agent_scope_cache_key(self, session_id: str) -> str:
|
||||
return f"{{session_iterations:{session_id}}}:agent-scope"
|
||||
|
||||
def _make_agent_scope_field(self, agent_id: str) -> str:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
stable_agent_id: Final = json.dumps(global_agent_registry.stable_agent_id(agent_id), separators=(",", ":"))
|
||||
return f"agent:{stable_agent_id}"
|
||||
|
||||
async def _get_local_scope(self, cache_key: str) -> dict[str, object] | None:
|
||||
result: Final[object | None] = cast(
|
||||
object | None,
|
||||
await self.internal_usage_cache.async_get_cache(
|
||||
key=cache_key,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
),
|
||||
)
|
||||
if result is None:
|
||||
return None
|
||||
if isinstance(result, dict) and all(isinstance(key, str) for key in result):
|
||||
return cast(dict[str, object], result)
|
||||
raise RuntimeError("Agent session scope cache has an invalid value")
|
||||
|
||||
async def _get_local_count(self, cache_key: str) -> int | None:
|
||||
local_result: Final[object | None] = cast(
|
||||
|
|
@ -253,6 +284,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
return None
|
||||
|
||||
async def _increment_legacy_and_get(self, cache_key: str) -> int:
|
||||
self._ensure_redis_scripts()
|
||||
if self.increment_script is not None:
|
||||
result: Final[object] = await self.increment_script(
|
||||
keys=[cache_key],
|
||||
|
|
@ -268,7 +300,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
await self.internal_usage_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=new_value,
|
||||
ttl=self.ttl if current is None else None,
|
||||
ttl=self.ttl,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
|
|
@ -276,13 +308,14 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
|
||||
async def _increment_agent_and_get(self, session_id: str, agent_id: str) -> int:
|
||||
legacy_key: Final = self._make_legacy_cache_key(session_id)
|
||||
total_new_key: Final = self._make_total_new_cache_key(session_id)
|
||||
agent_key: Final = self._make_cache_key(session_id, agent_id)
|
||||
agent_scope_key: Final = self._make_agent_scope_cache_key(session_id)
|
||||
agent_field: Final = self._make_agent_scope_field(agent_id)
|
||||
self._ensure_redis_scripts()
|
||||
if self.increment_script is not None:
|
||||
try:
|
||||
result: Final[object] = await self.increment_script(
|
||||
keys=[legacy_key, total_new_key, agent_key],
|
||||
args=[self.ttl],
|
||||
keys=[legacy_key, agent_scope_key],
|
||||
args=[self.ttl, agent_field],
|
||||
)
|
||||
if isinstance(result, (int, float, str, bytes)):
|
||||
return int(result)
|
||||
|
|
@ -296,35 +329,41 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
|
||||
async with self._local_lock:
|
||||
legacy_value = await self._get_local_count(legacy_key)
|
||||
total_new_value = await self._get_local_count(total_new_key)
|
||||
agent_value = await self._get_local_count(agent_key)
|
||||
agent_scope = await self._get_local_scope(agent_scope_key)
|
||||
if legacy_value is None:
|
||||
if total_new_value is not None:
|
||||
if agent_scope is not None:
|
||||
raise RuntimeError("Legacy session count expired before agent scope")
|
||||
legacy_value = 0
|
||||
total_new_value = total_new_value or 0
|
||||
agent_value = agent_value or 0
|
||||
total_new_value = 0
|
||||
agent_value = 0
|
||||
if agent_scope is not None:
|
||||
raw_total_new = agent_scope.get("__total_new")
|
||||
if not isinstance(raw_total_new, (int, float, str, bytes)):
|
||||
raise RuntimeError("Agent session scope is missing its migration total")
|
||||
raw_agent_value = agent_scope.get(agent_field, 0)
|
||||
if not isinstance(raw_agent_value, (int, float, str, bytes)):
|
||||
raise RuntimeError("Agent session scope has an invalid counter")
|
||||
total_new_value = int(raw_total_new)
|
||||
agent_value = int(raw_agent_value)
|
||||
new_legacy: Final = legacy_value + 1
|
||||
new_total_new: Final = total_new_value + 1
|
||||
new_agent: Final = agent_value + 1
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=legacy_key,
|
||||
value=new_legacy,
|
||||
ttl=self.ttl if legacy_value == 0 else None,
|
||||
ttl=self.ttl,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=total_new_key,
|
||||
value=new_total_new,
|
||||
ttl=self.ttl + 1 if total_new_value == 0 else None,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=agent_key,
|
||||
value=new_agent,
|
||||
ttl=self.ttl + 1 if agent_value == 0 else None,
|
||||
key=agent_scope_key,
|
||||
value={
|
||||
**(agent_scope or {}),
|
||||
"__total_new": new_total_new,
|
||||
agent_field: new_agent,
|
||||
},
|
||||
# Keep this map on a shorter rolling TTL than the aggregate fallback.
|
||||
ttl=max(self.ttl - 1, 0),
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -295,6 +295,71 @@ async def test_agent_budget_carries_existing_spend_into_its_counter() -> None:
|
|||
assert spend == pytest.approx(0.85)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_budget_recovers_conservatively_after_scope_eviction() -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
|
||||
session_id: Final = "scope-eviction-session"
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(_make_mock_agent(10.0, "agent-a"))
|
||||
registry.register_agent(_make_mock_agent(10.0, "agent-b"))
|
||||
|
||||
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
|
||||
await handler._increment_agent_spend(session_id, "agent-a", 0.4)
|
||||
await cache.async_delete_cache(key=handler._make_agent_scope_cache_key(session_id))
|
||||
# Lost per-agent detail falls back to the aggregate, rather than resetting spend.
|
||||
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(0.4)
|
||||
assert await handler._increment_agent_spend(session_id, "agent-a", 0.1) == pytest.approx(0.5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_budget_fails_closed_when_scope_loses_migration_total() -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
|
||||
session_id: Final = "inconsistent-scope-session"
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(_make_mock_agent(10.0, "agent-a"))
|
||||
scope_key: Final = handler._make_agent_scope_cache_key(session_id)
|
||||
await cache.async_set_cache(key=handler._make_legacy_cache_key(session_id), value=0.4)
|
||||
await cache.async_set_cache(key=scope_key, value={})
|
||||
|
||||
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
|
||||
with pytest.raises(RuntimeError, match="missing its migration total"):
|
||||
await handler._increment_agent_spend(session_id, "agent-a", 0.1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_budget_registers_redis_scripts_when_redis_is_attached_late() -> None:
|
||||
class FakeRedisCache:
|
||||
def __init__(self) -> None:
|
||||
self.scripts: list[str] = []
|
||||
self.calls: list[dict[str, object]] = []
|
||||
|
||||
def async_register_script(self, source: str):
|
||||
self.scripts.append(source)
|
||||
|
||||
async def call(**kwargs: object) -> object:
|
||||
self.calls.append(kwargs)
|
||||
return "0.1"
|
||||
|
||||
return call
|
||||
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
|
||||
redis_cache: Final = FakeRedisCache()
|
||||
cache.redis_cache = redis_cache
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(_make_mock_agent(10.0, "agent-a"))
|
||||
|
||||
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
|
||||
await handler._increment_agent_spend("late-redis", "agent-a", 0.1)
|
||||
assert await handler._get_agent_spend("late-redis", "agent-a") == pytest.approx(0.1)
|
||||
|
||||
assert len(redis_cache.scripts) == 2
|
||||
assert len(redis_cache.calls) == 2
|
||||
assert all(len(call["keys"]) == 2 for call in redis_cache.calls)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_budget_tracks_old_and_new_pods_without_mixing_new_agent_spend() -> None:
|
||||
cache: Final = DualCache()
|
||||
|
|
@ -358,7 +423,11 @@ async def test_agent_spend_is_recorded_before_a_budget_is_configured() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_budget_redis_errors_are_not_retried_or_read_locally() -> None:
|
||||
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(DualCache()))
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxBudgetPerSessionHandler(InternalUsageCache(cache))
|
||||
redis_cache: Final = object()
|
||||
cache.redis_cache = redis_cache
|
||||
handler._registered_redis_cache = redis_cache
|
||||
increment_calls = 0
|
||||
read_calls = 0
|
||||
|
||||
|
|
@ -405,10 +474,8 @@ async def test_redis_budget_migration_is_atomic_and_preserves_session_ttl() -> N
|
|||
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
|
||||
try:
|
||||
legacy_key: Final = handler._make_legacy_cache_key(session_id)
|
||||
total_new_key: Final = handler._make_total_new_cache_key(session_id)
|
||||
agent_a_key: Final = handler._make_cache_key(session_id, "agent-a")
|
||||
agent_b_key: Final = handler._make_cache_key(session_id, "agent-b")
|
||||
keys.extend((legacy_key, total_new_key, agent_a_key, agent_b_key))
|
||||
agent_scope_key: Final = handler._make_agent_scope_cache_key(session_id)
|
||||
keys.extend((legacy_key, agent_scope_key))
|
||||
await redis.async_set_cache(key=legacy_key, value=0.5, ttl=30)
|
||||
legacy_redis_key: Final = redis.check_and_fix_namespace(legacy_key)
|
||||
initial_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
|
|
@ -420,32 +487,32 @@ async def test_redis_budget_migration_is_atomic_and_preserves_session_ttl() -> N
|
|||
assert await handler._get_agent_spend(session_id, "agent-a") == pytest.approx(1.2)
|
||||
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(0.8)
|
||||
|
||||
hash_tags: Final = {key[key.index("{") : key.index("}") + 1] for key in keys[:4]}
|
||||
hash_tags: Final = {key[key.index("{") : key.index("}") + 1] for key in keys[:2]}
|
||||
assert len(hash_tags) == 1
|
||||
sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(total_new_key))
|
||||
agent_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_a_key))
|
||||
sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
|
||||
current_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
assert current_ttl <= sidecar_ttl <= current_ttl + 1100
|
||||
assert sidecar_ttl <= agent_ttl <= sidecar_ttl + 1100
|
||||
assert 0 <= sidecar_ttl <= current_ttl
|
||||
assert current_ttl <= initial_ttl
|
||||
await asyncio.sleep(0.25)
|
||||
before_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
before_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(total_new_key))
|
||||
before_increment_sidecar_ttl: Final = await redis_client.pttl(
|
||||
redis.check_and_fix_namespace(agent_scope_key)
|
||||
)
|
||||
await handler._increment_agent_spend(session_id, "agent-a", 0.01)
|
||||
after_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
after_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(total_new_key))
|
||||
after_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
|
||||
assert after_increment_ttl <= before_increment_ttl + 50
|
||||
assert after_increment_sidecar_ttl <= before_increment_sidecar_ttl + 50
|
||||
|
||||
# Losing the hash evicts migration and agent counters together; the aggregate remains a safe baseline.
|
||||
await redis_client.delete(redis.check_and_fix_namespace(agent_scope_key))
|
||||
assert await handler._get_agent_spend(session_id, "agent-b") == pytest.approx(1.31)
|
||||
await handler._increment_agent_spend(session_id, "agent-a", 0.01)
|
||||
assert await handler._get_agent_spend(session_id, "agent-a") == pytest.approx(1.32)
|
||||
|
||||
concurrent_legacy_key: Final = handler._make_legacy_cache_key(concurrent_session)
|
||||
keys.extend(
|
||||
(
|
||||
concurrent_legacy_key,
|
||||
handler._make_total_new_cache_key(concurrent_session),
|
||||
handler._make_cache_key(concurrent_session, "agent-a"),
|
||||
handler._make_cache_key(concurrent_session, "agent-b"),
|
||||
)
|
||||
)
|
||||
concurrent_scope_key: Final = handler._make_agent_scope_cache_key(concurrent_session)
|
||||
keys.extend((concurrent_legacy_key, concurrent_scope_key))
|
||||
await redis.async_set_cache(key=concurrent_legacy_key, value=5.0, ttl=30)
|
||||
await asyncio.gather(
|
||||
*(handler._increment_agent_spend(concurrent_session, "agent-a", 0.01) for _ in range(100)),
|
||||
|
|
@ -454,13 +521,18 @@ async def test_redis_budget_migration_is_atomic_and_preserves_session_ttl() -> N
|
|||
assert await handler._get_agent_spend(concurrent_session, "agent-a") == pytest.approx(6.0)
|
||||
assert await handler._get_agent_spend(concurrent_session, "agent-b") == pytest.approx(5.5)
|
||||
|
||||
saved_total_new: Final = await redis_client.hget(
|
||||
redis.check_and_fix_namespace(agent_scope_key), "__total_new"
|
||||
)
|
||||
await redis_client.hdel(redis.check_and_fix_namespace(agent_scope_key), "__total_new")
|
||||
with pytest.raises(Exception, match="agent session scope is missing its migration total"):
|
||||
await handler._get_agent_spend(session_id, "agent-a")
|
||||
assert saved_total_new is not None
|
||||
await redis_client.hset(redis.check_and_fix_namespace(agent_scope_key), "__total_new", saved_total_new)
|
||||
|
||||
await redis_client.pexpire(legacy_redis_key, 100)
|
||||
await asyncio.sleep(0.15)
|
||||
with pytest.raises(Exception, match="legacy session spend expired before agent scope"):
|
||||
await handler._get_agent_spend(session_id, "agent-a")
|
||||
|
||||
await redis_client.delete(redis.check_and_fix_namespace(total_new_key))
|
||||
with pytest.raises(Exception, match="agent session spend exists without migration total"):
|
||||
await handler._get_agent_spend(session_id, "agent-a")
|
||||
finally:
|
||||
await redis_client.delete(*(redis.check_and_fix_namespace(key) for key in keys))
|
||||
|
|
|
|||
|
|
@ -264,9 +264,71 @@ async def test_agent_iteration_limit_keeps_count_from_an_existing_session() -> N
|
|||
assert rejected.value.status_code == 429
|
||||
assert "Current count: 3" in str(rejected.value.detail)
|
||||
legacy_count: Final = await cache.async_get_cache(key=handler._make_legacy_cache_key(session_id))
|
||||
total_new_count: Final = await cache.async_get_cache(key=handler._make_total_new_cache_key(session_id))
|
||||
agent_count: Final = await cache.async_get_cache(key=handler._make_cache_key(session_id, "agent-test-123"))
|
||||
assert legacy_count - total_new_count + agent_count == 3
|
||||
scope: Final = await cache.async_get_cache(key=handler._make_agent_scope_cache_key(session_id))
|
||||
agent_field: Final = handler._make_agent_scope_field("agent-test-123")
|
||||
assert isinstance(scope, dict)
|
||||
assert legacy_count - scope["__total_new"] + scope[agent_field] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_iteration_recovers_conservatively_after_scope_eviction() -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
|
||||
session_id: Final = "scope-eviction-session"
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(_make_mock_agent(20, "agent-a"))
|
||||
registry.register_agent(_make_mock_agent(20, "agent-b"))
|
||||
|
||||
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
|
||||
assert await handler._increment_agent_and_get(session_id, "agent-a") == 1
|
||||
await cache.async_delete_cache(key=handler._make_agent_scope_cache_key(session_id))
|
||||
# With the migration scope gone, prior usage becomes a conservative shared baseline.
|
||||
assert await handler._increment_agent_and_get(session_id, "agent-b") == 2
|
||||
assert await handler._increment_agent_and_get(session_id, "agent-a") == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_iteration_fails_closed_when_scope_loses_migration_total() -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
|
||||
session_id: Final = "inconsistent-scope-session"
|
||||
scope_key: Final = handler._make_agent_scope_cache_key(session_id)
|
||||
await cache.async_set_cache(key=handler._make_legacy_cache_key(session_id), value=1)
|
||||
await cache.async_set_cache(key=scope_key, value={'agent:"agent-a"': 1})
|
||||
|
||||
with pytest.raises(RuntimeError, match="missing its migration total"):
|
||||
await handler._increment_agent_and_get(session_id, "agent-a")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_iteration_registers_redis_scripts_when_redis_is_attached_late() -> None:
|
||||
class FakeRedisCache:
|
||||
def __init__(self) -> None:
|
||||
self.scripts: list[str] = []
|
||||
self.calls: list[dict[str, object]] = []
|
||||
|
||||
def async_register_script(self, source: str):
|
||||
self.scripts.append(source)
|
||||
|
||||
async def call(**kwargs: object) -> object:
|
||||
self.calls.append(kwargs)
|
||||
return 1
|
||||
|
||||
return call
|
||||
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
|
||||
redis_cache: Final = FakeRedisCache()
|
||||
cache.redis_cache = redis_cache
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(_make_mock_agent(10, "agent-a"))
|
||||
|
||||
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
|
||||
assert await handler._increment_agent_and_get("late-redis", "agent-a") == 1
|
||||
|
||||
assert len(redis_cache.scripts) == 1
|
||||
assert len(redis_cache.calls) == 1
|
||||
assert len(redis_cache.calls[0]["keys"]) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -292,7 +354,11 @@ async def test_agent_iteration_counter_tracks_old_and_new_pods_independently() -
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_iteration_redis_error_is_not_retried_or_counted_locally() -> None:
|
||||
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(DualCache()))
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxIterationsHandler(InternalUsageCache(cache))
|
||||
redis_cache: Final = object()
|
||||
cache.redis_cache = redis_cache
|
||||
handler._registered_redis_cache = redis_cache
|
||||
calls = 0
|
||||
|
||||
async def fail_after_attempt(**_kwargs: object) -> object:
|
||||
|
|
@ -328,15 +394,8 @@ async def test_redis_iteration_migration_is_atomic_and_preserves_session_ttl() -
|
|||
with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry):
|
||||
try:
|
||||
legacy_key: Final = handler._make_legacy_cache_key(session_id)
|
||||
total_new_key: Final = handler._make_total_new_cache_key(session_id)
|
||||
keys.extend(
|
||||
(
|
||||
legacy_key,
|
||||
total_new_key,
|
||||
handler._make_cache_key(session_id, "agent-a"),
|
||||
handler._make_cache_key(session_id, "agent-b"),
|
||||
)
|
||||
)
|
||||
agent_scope_key: Final = handler._make_agent_scope_cache_key(session_id)
|
||||
keys.extend((legacy_key, agent_scope_key))
|
||||
await redis.async_set_cache(key=legacy_key, value=3, ttl=30)
|
||||
legacy_redis_key: Final = redis.check_and_fix_namespace(legacy_key)
|
||||
initial_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
|
|
@ -349,20 +408,18 @@ async def test_redis_iteration_migration_is_atomic_and_preserves_session_ttl() -
|
|||
assert await handler._increment_agent_and_get(session_id, "agent-a") == 6
|
||||
assert await handler._increment_agent_and_get(session_id, "agent-b") == 6
|
||||
|
||||
sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(total_new_key))
|
||||
agent_ttl: Final = await redis_client.pttl(
|
||||
redis.check_and_fix_namespace(handler._make_cache_key(session_id, "agent-a"))
|
||||
)
|
||||
sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
|
||||
current_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
assert current_ttl <= sidecar_ttl <= current_ttl + 1100
|
||||
assert sidecar_ttl <= agent_ttl <= sidecar_ttl + 1100
|
||||
assert 0 <= sidecar_ttl <= current_ttl
|
||||
assert current_ttl <= initial_ttl
|
||||
await asyncio.sleep(0.25)
|
||||
before_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
before_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(total_new_key))
|
||||
before_increment_sidecar_ttl: Final = await redis_client.pttl(
|
||||
redis.check_and_fix_namespace(agent_scope_key)
|
||||
)
|
||||
assert await handler._increment_agent_and_get(session_id, "agent-a") == 7
|
||||
after_increment_ttl: Final = await redis_client.pttl(legacy_redis_key)
|
||||
after_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(total_new_key))
|
||||
after_increment_sidecar_ttl: Final = await redis_client.pttl(redis.check_and_fix_namespace(agent_scope_key))
|
||||
assert after_increment_ttl <= before_increment_ttl + 50
|
||||
assert after_increment_sidecar_ttl <= before_increment_sidecar_ttl + 50
|
||||
|
||||
|
|
@ -371,13 +428,23 @@ async def test_redis_iteration_migration_is_atomic_and_preserves_session_ttl() -
|
|||
)
|
||||
assert sorted(values) == list(range(8, 108))
|
||||
|
||||
# Losing the hash evicts migration and agent counters together; the aggregate remains a safe baseline.
|
||||
await redis_client.delete(redis.check_and_fix_namespace(agent_scope_key))
|
||||
assert await handler._increment_agent_and_get(session_id, "agent-b") == 110
|
||||
assert await handler._increment_agent_and_get(session_id, "agent-a") == 110
|
||||
|
||||
saved_total_new: Final = await redis_client.hget(
|
||||
redis.check_and_fix_namespace(agent_scope_key), "__total_new"
|
||||
)
|
||||
await redis_client.hdel(redis.check_and_fix_namespace(agent_scope_key), "__total_new")
|
||||
with pytest.raises(Exception, match="agent session scope is missing its migration total"):
|
||||
await handler._increment_agent_and_get(session_id, "agent-a")
|
||||
assert saved_total_new is not None
|
||||
await redis_client.hset(redis.check_and_fix_namespace(agent_scope_key), "__total_new", saved_total_new)
|
||||
|
||||
await redis_client.pexpire(legacy_redis_key, 100)
|
||||
await asyncio.sleep(0.15)
|
||||
with pytest.raises(Exception, match="legacy session count expired before agent scope"):
|
||||
await handler._increment_agent_and_get(session_id, "agent-a")
|
||||
|
||||
await redis_client.delete(redis.check_and_fix_namespace(total_new_key))
|
||||
with pytest.raises(Exception, match="agent session count exists without migration total"):
|
||||
await handler._increment_agent_and_get(session_id, "agent-a")
|
||||
finally:
|
||||
await redis_client.delete(*(redis.check_and_fix_namespace(key) for key in keys))
|
||||
|
|
|
|||
|
|
@ -842,7 +842,7 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError:
|
|||
|
||||
internal_cache = MagicMock()
|
||||
internal_cache.dual_cache = DualCache()
|
||||
internal_cache.async_get_cache = AsyncMock(return_value=10.0)
|
||||
internal_cache.async_get_cache = AsyncMock(side_effect=[10.0, None])
|
||||
handler = _PROXY_MaxBudgetPerSessionHandler(
|
||||
internal_usage_cache=internal_cache,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue