fix(a2a/wxo): evict expired token cache entries and raise asyncio.TimeoutError on poll timeout
Some checks failed
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled
Unit Tests: Proxy DB Operations / auth-checks (push) Has been cancelled
Unit Tests: Proxy DB Operations / budgets (push) Has been cancelled
Unit Tests: Proxy DB Operations / custom-logging (push) Has been cancelled
Unit Tests: Proxy DB Operations / db-and-spend (push) Has been cancelled
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Has been cancelled
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Has been cancelled
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Has been cancelled
Unit Tests: Proxy DB Operations / key-generation (push) Has been cancelled
Unit Tests: Proxy DB Operations / logging-misc (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-runtime (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-server-core (push) Has been cancelled
Unit Tests: Proxy DB Operations / schema-migration (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-utils (push) Has been cancelled

This commit is contained in:
mateo-berri 2026-06-02 06:25:32 +00:00
parent 65a06a57e2
commit 26672902a8
No known key found for this signature in database
2 changed files with 99 additions and 1 deletions

View file

@ -118,6 +118,9 @@ class WatsonxOrchestrateHandler:
expires_at = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 0)
_token_cache[cache_key] = (token, expires_at)
for stale_key, (_, stale_expires_at) in list(_token_cache.items()):
if stale_expires_at <= now:
del _token_cache[stale_key]
return token
@staticmethod
@ -143,7 +146,7 @@ class WatsonxOrchestrateHandler:
if status in WatsonxOrchestrateTransformation.TERMINAL_STATES:
return result
raise TimeoutError(
raise asyncio.TimeoutError(
f"WXO run '{run_id}' did not reach a terminal state after "
f"{max_attempts * interval_s:.0f}s"
)

View file

@ -1,6 +1,8 @@
import asyncio
import json
import os
import sys
import time
from pathlib import Path
import httpx
@ -9,6 +11,7 @@ import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager
from litellm.a2a_protocol.providers.watsonx_orchestrate import handler as wxo_handler
from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import (
WatsonxOrchestrateHandler,
)
@ -251,6 +254,98 @@ async def test_short_lived_tokens_are_not_served_from_cache():
assert client.calls == 2
class _CP4DAuthClient:
def __init__(self, expiration):
self.expiration = expiration
self.calls = []
async def post(self, url, **kwargs):
self.calls.append((url, kwargs))
return _JsonResponse({"token": "cp4d-token", "expiration": self.expiration})
@pytest.mark.asyncio
async def test_cp4d_auth_posts_to_authorize_and_caches_token():
client = _CP4DAuthClient(int(time.time()) + 3600)
token_1 = await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host="https://cpd.example.com/",
auth_mode="cp4d",
api_key="cp4d-e2e-cache-key",
username="cp4d-user",
client=client,
)
token_2 = await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host="https://cpd.example.com/",
auth_mode="cp4d",
api_key="cp4d-e2e-cache-key",
username="cp4d-user",
client=client,
)
assert token_1 == "cp4d-token"
assert token_2 == "cp4d-token"
assert len(client.calls) == 1
url, kwargs = client.calls[0]
assert url == "https://cpd.example.com/icp4d-api/v1/authorize"
assert kwargs["json"] == {"username": "cp4d-user", "api_key": "cp4d-e2e-cache-key"}
@pytest.mark.asyncio
async def test_cp4d_auth_requires_username():
client = _CP4DAuthClient(int(time.time()) + 3600)
with pytest.raises(ValueError, match="username"):
await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host="https://cpd.example.com",
auth_mode="cp4d",
api_key="cp4d-missing-username-key",
username=None,
client=client,
)
assert client.calls == []
@pytest.mark.asyncio
async def test_expired_token_cache_entries_are_evicted():
stale_key = "wxo-stale-cache-entry"
wxo_handler._token_cache[stale_key] = ("stale-token", time.monotonic() - 1)
class _FreshTokenClient:
async def post(self, *args, **kwargs):
return _JsonResponse({"access_token": "fresh", "expires_in": 3600})
await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host="https://cpd.example.com",
auth_mode="ibm_cloud",
api_key="wxo-eviction-trigger-key",
client=_FreshTokenClient(),
)
assert stale_key not in wxo_handler._token_cache
@pytest.mark.asyncio
async def test_poll_run_raises_asyncio_timeout_when_never_terminal():
class _NeverTerminalClient:
def __init__(self):
self.get_calls = 0
async def get(self, url, headers=None):
self.get_calls += 1
return _JsonResponse({"status": "running"})
client = _NeverTerminalClient()
with pytest.raises(asyncio.TimeoutError):
await WatsonxOrchestrateHandler._poll_run(
base_url="https://cpd.example.com/orchestrate/cpd/instances/i",
run_id="run-1",
auth_headers={},
client=client,
max_attempts=2,
interval_s=0,
)
assert client.get_calls == 2
@pytest.mark.asyncio
async def test_handle_streaming_polls_non_sse_json_until_complete(monkeypatch):
client = _JsonStreamClient({"status": "running", "run_id": "run-1"})