Skip to content

Commit 97c8201

Browse files
fix(hosting): isolate OBS cache keys and bound refresh locks
- Key the cache by an (agent_id, tenant_id) tuple instead of "agent_id:tenant_id", so IDs that contain a colon cannot share another identity's cached token. make_key (new in this PR) becomes private. - Keep each refresh lock on its cache entry, so capacity eviction and invalidate_all() reclaim it. Eviction skips entries with a refresh in flight. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 5cbf5f6b-cc40-4b7e-a591-65848db73a12
1 parent 5d40056 commit 97c8201

2 files changed

Lines changed: 82 additions & 25 deletions

File tree

‎libraries/microsoft-agents-a365-observability-hosting/microsoft_agents_a365/observability/hosting/token_cache_helpers/agent_token_cache.py‎

Lines changed: 23 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
import logging
1212
import time
1313
from collections.abc import Awaitable, Callable, Sequence
14-
from dataclasses import dataclass
14+
from dataclasses import dataclass, field
1515
from inspect import isawaitable
1616
from threading import Lock
1717

@@ -59,6 +59,7 @@ class _Entry:
5959
token: str | None = None
6060
expires_on_ms: float | None = None
6161
acquired_on_ms: float | None = None
62+
lock: asyncio.Lock = field(default_factory=asyncio.Lock, repr=False, compare=False)
6263

6364
_default_refresh_skew_ms = 60_000
6465
_default_max_token_age_ms = 3_600_000
@@ -67,18 +68,17 @@ class _Entry:
6768

6869
def __init__(self, observability_scopes: Sequence[str] | None = None) -> None:
6970
"""Initialize the token cache."""
70-
self._map: dict[str, AgenticTokenCache._Entry] = {}
71-
self._key_locks: dict[str, asyncio.Lock] = {}
71+
self._map: dict[tuple[str, str], AgenticTokenCache._Entry] = {}
7272
self._lock = Lock()
7373
self._observability_scopes = (
7474
None if observability_scopes is None else tuple(observability_scopes)
7575
)
7676
self._removed_registration_logged = False
7777

7878
@staticmethod
79-
def make_key(agent_id: str, tenant_id: str) -> str:
80-
"""Create a cache key for an agent and tenant."""
81-
return f"{agent_id}:{tenant_id}"
79+
def _make_key(agent_id: str, tenant_id: str) -> tuple[str, str]:
80+
# A tuple key keeps identities apart even when an ID contains a separator character.
81+
return (agent_id, tenant_id)
8282

8383
def register_observability(
8484
self,
@@ -118,10 +118,9 @@ async def refresh_observability_token(
118118
if not agent_id or not agent_id.strip() or not tenant_id or not tenant_id.strip():
119119
raise ValueError("[AgenticTokenCache] Agent and tenant IDs are required")
120120

121-
key = self.make_key(agent_id, tenant_id)
122-
lock = self._get_key_lock(key)
123-
async with lock:
124-
entry = self._get_or_create_entry(key)
121+
key = self._make_key(agent_id, tenant_id)
122+
entry = self._get_or_create_entry(key)
123+
async with entry.lock:
125124
if entry.token is not None and not self._is_expired(entry):
126125
return entry.token
127126

@@ -133,7 +132,7 @@ async def get_observability_token(self, agent_id: str, tenant_id: str) -> str |
133132
This method is a pure cache read. It never acquires a token and never
134133
calls delegated token exchange.
135134
"""
136-
key = self.make_key(agent_id, tenant_id)
135+
key = self._make_key(agent_id, tenant_id)
137136
with self._lock:
138137
entry = self._map.get(key)
139138

@@ -147,7 +146,7 @@ async def get_observability_token(self, agent_id: str, tenant_id: str) -> str |
147146

148147
def invalidate_token(self, agent_id: str, tenant_id: str) -> None:
149148
"""Invalidate one cached token."""
150-
key = self.make_key(agent_id, tenant_id)
149+
key = self._make_key(agent_id, tenant_id)
151150
with self._lock:
152151
entry = self._map.get(key)
153152
if entry is not None:
@@ -158,7 +157,7 @@ def invalidate_all(self) -> None:
158157
with self._lock:
159158
self._map.clear()
160159

161-
def _get_or_create_entry(self, key: str) -> _Entry:
160+
def _get_or_create_entry(self, key: tuple[str, str]) -> _Entry:
162161
with self._lock:
163162
entry = self._map.get(key)
164163
if entry is not None:
@@ -171,22 +170,22 @@ def _get_or_create_entry(self, key: str) -> _Entry:
171170
raise ValueError("[AgenticTokenCache] No valid scopes")
172171

173172
if len(self._map) >= self._max_cache_size:
174-
oldest_key = next(iter(self._map), None)
175-
if oldest_key is not None:
176-
del self._map[oldest_key]
173+
# Evict the oldest idle entry; an entry with a refresh in flight keeps its lock.
174+
idle_key = next(
175+
(
176+
existing_key
177+
for existing_key, existing in self._map.items()
178+
if not existing.lock.locked()
179+
),
180+
None,
181+
)
182+
if idle_key is not None:
183+
del self._map[idle_key]
177184

178185
entry = AgenticTokenCache._Entry(scopes=scopes)
179186
self._map[key] = entry
180187
return entry
181188

182-
def _get_key_lock(self, key: str) -> asyncio.Lock:
183-
with self._lock:
184-
lock = self._key_locks.get(key)
185-
if lock is None:
186-
lock = asyncio.Lock()
187-
self._key_locks[key] = lock
188-
return lock
189-
190189
def _get_effective_scopes(self) -> tuple[str, ...]:
191190
scopes = self._observability_scopes
192191
if scopes is None:

‎tests/observability/hosting/token_cache_helpers/test_agent_token_cache.py‎

Lines changed: 59 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -253,7 +253,7 @@ async def test_opaque_token_uses_fresh_fallback_ttl(token_cache):
253253
await token_cache.refresh_observability_token("agent", "tenant", lambda *_: "opaque-token")
254254

255255
assert await token_cache.get_observability_token("agent", "tenant") == "opaque-token"
256-
entry = token_cache._map[AgenticTokenCache.make_key("agent", "tenant")]
256+
entry = token_cache._map[AgenticTokenCache._make_key("agent", "tenant")]
257257
entry.acquired_on_ms = (time.time() * 1000) - token_cache._default_max_token_age_ms - 1
258258
assert await token_cache.get_observability_token("agent", "tenant") is None
259259

@@ -307,3 +307,61 @@ async def test_cache_evicts_oldest_entry_when_capacity_is_reached(token_cache):
307307
assert await token_cache.get_observability_token("one", "tenant") is None
308308
assert await token_cache.get_observability_token("two", "tenant") == token
309309
assert await token_cache.get_observability_token("three", "tenant") == token
310+
311+
312+
@pytest.mark.asyncio
313+
async def test_separator_bearing_ids_do_not_share_a_cache_entry(token_cache):
314+
"""IDs that contain a separator character never alias another identity."""
315+
first = MagicMock(return_value="token-for-first")
316+
second = MagicMock(return_value="token-for-second")
317+
318+
assert await token_cache.refresh_observability_token("a:b", "c", first) == "token-for-first"
319+
assert await token_cache.refresh_observability_token("a", "b:c", second) == "token-for-second"
320+
321+
second.assert_called_once()
322+
assert await token_cache.get_observability_token("a:b", "c") == "token-for-first"
323+
assert await token_cache.get_observability_token("a", "b:c") == "token-for-second"
324+
325+
326+
@pytest.mark.asyncio
327+
async def test_identity_churn_keeps_every_per_identity_registry_bounded(token_cache):
328+
"""Refresh locks live and die with cache entries, so identity churn stays bounded."""
329+
token_cache._max_cache_size = 2
330+
token = make_jwt(300)
331+
for agent_id in ("one", "two", "three", "four"):
332+
await token_cache.refresh_observability_token(agent_id, "tenant", lambda *_: token)
333+
334+
registries = {
335+
name: len(value) for name, value in vars(token_cache).items() if isinstance(value, dict)
336+
}
337+
assert registries == {"_map": 2}
338+
339+
token_cache.invalidate_all()
340+
registries = {
341+
name: len(value) for name, value in vars(token_cache).items() if isinstance(value, dict)
342+
}
343+
assert registries == {"_map": 0}
344+
345+
346+
@pytest.mark.asyncio
347+
async def test_eviction_keeps_entry_with_refresh_in_flight(token_cache):
348+
"""Capacity eviction skips an identity whose refresh is still running."""
349+
token_cache._max_cache_size = 1
350+
release = asyncio.Event()
351+
352+
async def slow_resolver(agent_id: str, tenant_id: str, scopes: list[str]) -> str:
353+
await release.wait()
354+
return "token-for-slow"
355+
356+
in_flight = asyncio.create_task(
357+
token_cache.refresh_observability_token("slow", "tenant", slow_resolver)
358+
)
359+
await asyncio.sleep(0)
360+
fast = await token_cache.refresh_observability_token(
361+
"fast", "tenant", lambda *_: "token-for-fast"
362+
)
363+
364+
release.set()
365+
assert fast == "token-for-fast"
366+
assert await in_flight == "token-for-slow"
367+
assert await token_cache.get_observability_token("slow", "tenant") == "token-for-slow"

0 commit comments

Comments
 (0)