mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-20 15:33:54 +00:00
fix(google-chat): cache callback token cert fetches
This commit is contained in:
parent
1305a690e0
commit
a7ec1b6e39
2 changed files with 105 additions and 2 deletions
|
|
@ -43,6 +43,8 @@ import logging
|
|||
import os
|
||||
import random
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path as _Path
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple
|
||||
|
||||
|
|
@ -71,18 +73,57 @@ HttpError: Any = Exception # type: ignore
|
|||
MediaFileUpload: Any = None # type: ignore
|
||||
|
||||
_google_modules_loaded: bool = False
|
||||
_GOOGLE_ID_TOKEN_CERTS_TTL_SECONDS = 300
|
||||
_google_id_token_request: Any = None
|
||||
_google_id_token_request_lock = threading.Lock()
|
||||
|
||||
|
||||
class _CachedGoogleAuthRequest:
|
||||
def __init__(self, request: Any, ttl_seconds: int = _GOOGLE_ID_TOKEN_CERTS_TTL_SECONDS) -> None:
|
||||
self._request = request
|
||||
self._ttl_seconds = ttl_seconds
|
||||
self._lock = threading.Lock()
|
||||
self._cache: Dict[Tuple[str, str], Tuple[float, Any]] = {}
|
||||
|
||||
def __call__(self, url: str, method: str = "GET", **kwargs: Any) -> Any:
|
||||
cache_key = (method.upper(), url)
|
||||
if cache_key[0] != "GET":
|
||||
return self._request(url=url, method=method, **kwargs)
|
||||
|
||||
now = time.monotonic()
|
||||
with self._lock:
|
||||
cached = self._cache.get(cache_key)
|
||||
if cached and cached[0] > now:
|
||||
return cached[1]
|
||||
|
||||
response = self._request(url=url, method=method, **kwargs)
|
||||
if getattr(response, "status", None) == 200:
|
||||
with self._lock:
|
||||
self._cache[cache_key] = (now + self._ttl_seconds, response)
|
||||
return response
|
||||
|
||||
|
||||
def _get_google_id_token_request() -> Any:
|
||||
global _google_id_token_request
|
||||
with _google_id_token_request_lock:
|
||||
if _google_id_token_request is None:
|
||||
try:
|
||||
from google.auth.transport import requests as google_requests
|
||||
except ImportError as exc:
|
||||
raise RuntimeError("google-auth is required for Google Chat HTTP callbacks") from exc
|
||||
_google_id_token_request = _CachedGoogleAuthRequest(google_requests.Request())
|
||||
return _google_id_token_request
|
||||
|
||||
|
||||
def _verify_google_id_token(token: str, audience: str) -> Dict[str, Any]:
|
||||
try:
|
||||
from google.auth.transport import requests as google_requests
|
||||
from google.oauth2 import id_token
|
||||
except ImportError as exc:
|
||||
raise RuntimeError("google-auth is required for Google Chat HTTP callbacks") from exc
|
||||
|
||||
return id_token.verify_oauth2_token(
|
||||
token,
|
||||
google_requests.Request(),
|
||||
_get_google_id_token_request(),
|
||||
audience,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -439,6 +439,68 @@ class TestValidateConfig:
|
|||
|
||||
|
||||
class TestHttpEventIngress:
|
||||
def test_cached_google_auth_request_reuses_successful_get_response(self, monkeypatch):
|
||||
now = [100.0]
|
||||
calls = []
|
||||
|
||||
class Response:
|
||||
status = 200
|
||||
|
||||
response = Response()
|
||||
|
||||
def raw_request(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(_gc_mod.time, "monotonic", lambda: now[0])
|
||||
request = _gc_mod._CachedGoogleAuthRequest(raw_request, ttl_seconds=300)
|
||||
|
||||
assert request("https://www.googleapis.com/oauth2/v1/certs") is response
|
||||
assert request("https://www.googleapis.com/oauth2/v1/certs") is response
|
||||
assert len(calls) == 1
|
||||
|
||||
now[0] += 301
|
||||
assert request("https://www.googleapis.com/oauth2/v1/certs") is response
|
||||
assert len(calls) == 2
|
||||
|
||||
def test_cached_google_auth_request_does_not_cache_post_response(self):
|
||||
calls = []
|
||||
|
||||
def raw_request(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return object()
|
||||
|
||||
request = _gc_mod._CachedGoogleAuthRequest(raw_request, ttl_seconds=300)
|
||||
|
||||
request("https://example.test/token", method="POST")
|
||||
request("https://example.test/token", method="POST")
|
||||
|
||||
assert len(calls) == 2
|
||||
|
||||
def test_verify_google_id_token_uses_cached_request_and_configured_audience(self, monkeypatch):
|
||||
request = object()
|
||||
captured = {}
|
||||
|
||||
monkeypatch.setattr(_gc_mod, "_get_google_id_token_request", lambda: request)
|
||||
|
||||
class FakeIdToken:
|
||||
@staticmethod
|
||||
def verify_oauth2_token(token, req, audience):
|
||||
captured.update(token=token, request=req, audience=audience)
|
||||
return {"email": "bot@example.test"}
|
||||
|
||||
monkeypatch.setitem(sys.modules, "google.oauth2.id_token", FakeIdToken)
|
||||
monkeypatch.setattr(sys.modules["google.oauth2"], "id_token", FakeIdToken, raising=False)
|
||||
|
||||
assert _gc_mod._verify_google_id_token("signed-token", "https://callback.example/events") == {
|
||||
"email": "bot@example.test"
|
||||
}
|
||||
assert captured == {
|
||||
"token": "signed-token",
|
||||
"request": request,
|
||||
"audience": "https://callback.example/events",
|
||||
}
|
||||
|
||||
def test_verify_http_event_request_accepts_expected_google_identity(self, monkeypatch):
|
||||
cfg = PlatformConfig(enabled=True)
|
||||
cfg.extra.update(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue