mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 13:53:15 +00:00
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com> (cherry picked from commit 1ea0e600dcc234fa5a333c6f4b80b90e2e6679d3) Co-authored-by: Johannes du Plessis <johannes@langchain.dev>
216 lines
7.7 KiB
Python
216 lines
7.7 KiB
Python
"""GitHub token lookup utilities."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
from collections.abc import Mapping
|
|
from datetime import UTC, datetime, timedelta
|
|
from typing import Any
|
|
|
|
from langgraph.config import get_config
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# Treat tokens with <= this many seconds remaining as expired so we re-auth
|
|
# before kicking off long agent runs.
|
|
_GITHUB_TOKEN_EXPIRY_SKEW_SECONDS = 60
|
|
# Hard cap on how long an entry stays cached regardless of the token's own
|
|
# expiry, so entries for threads that are never read again don't accumulate.
|
|
_GITHUB_TOKEN_MAX_TTL = timedelta(hours=24)
|
|
_BOT_PRINCIPAL = "bot"
|
|
# (thread_id, principal) -> (token, token_expires_at, cached_at, bound_repo).
|
|
# ``principal`` binds user tokens to the user they were resolved for so a
|
|
# cached token can never be served to a different user on the same thread.
|
|
# ``bound_repo`` ("owner/name") binds the entry to the repo it was resolved for
|
|
# so a colliding thread_id originating from a different repo cannot reuse
|
|
# another repo's token.
|
|
_GITHUB_TOKEN_CACHE: dict[tuple[str, str], tuple[str, str | None, datetime, str | None]] = {}
|
|
|
|
|
|
def github_token_principal(*, login: str | None = None, email: str | None = None) -> str | None:
|
|
"""Return the normalized principal used to isolate cached user tokens."""
|
|
if isinstance(login, str) and login.strip():
|
|
return f"login:{login.strip().casefold()}"
|
|
if isinstance(email, str) and email.strip():
|
|
return f"email:{email.strip().casefold()}"
|
|
return None
|
|
|
|
|
|
class GitHubAuthError(Exception):
|
|
"""Raised when a GitHub call returns 401, signalling a stale/revoked token."""
|
|
|
|
|
|
def repo_cache_key(repo: Any) -> str | None:
|
|
"""Normalize a repo dict/string to a casefolded ``owner/name`` (None if unknown).
|
|
|
|
GitHub owner/name are case-insensitive, so the key is casefolded on both the
|
|
write (binding) and read (compare) sides to keep ``Org/Repo`` and ``org/repo``
|
|
a single repo and avoid spurious cross-repo mismatches.
|
|
"""
|
|
if isinstance(repo, str):
|
|
cleaned = repo.strip()
|
|
return cleaned.casefold() or None
|
|
if isinstance(repo, Mapping):
|
|
owner = repo.get("owner")
|
|
name = repo.get("name")
|
|
if isinstance(owner, str) and isinstance(name, str) and owner and name:
|
|
return f"{owner}/{name}".casefold()
|
|
return None
|
|
|
|
|
|
def cache_github_token_for_thread(
|
|
thread_id: str,
|
|
token: str,
|
|
expires_at: str | None = None,
|
|
*,
|
|
repo: Any = None,
|
|
principal: str | None = None,
|
|
is_bot_token: bool = False,
|
|
) -> None:
|
|
"""Cache a GitHub token in process for the current thread and principal."""
|
|
if not thread_id or not token:
|
|
return
|
|
cache_principal = _BOT_PRINCIPAL if is_bot_token else principal
|
|
if not cache_principal:
|
|
logger.warning("Refusing to cache an unbound user GitHub token for thread %s", thread_id)
|
|
return
|
|
now = datetime.now(UTC)
|
|
_GITHUB_TOKEN_CACHE[(thread_id, cache_principal)] = (
|
|
token,
|
|
expires_at,
|
|
now,
|
|
repo_cache_key(repo),
|
|
)
|
|
_evict_expired(now=now)
|
|
|
|
|
|
def _is_expired(expires_at: Any, *, now: datetime | None = None) -> bool:
|
|
"""Return True when ``expires_at`` is past (or close to) ``now``."""
|
|
if expires_at is None:
|
|
return False
|
|
|
|
parsed: datetime | None = None
|
|
if isinstance(expires_at, int | float):
|
|
try:
|
|
parsed = datetime.fromtimestamp(float(expires_at), tz=UTC)
|
|
except (OverflowError, OSError, ValueError):
|
|
return False
|
|
elif isinstance(expires_at, str):
|
|
raw = expires_at.strip()
|
|
if not raw:
|
|
return False
|
|
if raw.endswith("Z"):
|
|
raw = raw[:-1] + "+00:00"
|
|
try:
|
|
parsed = datetime.fromisoformat(raw)
|
|
except ValueError:
|
|
return False
|
|
|
|
if parsed is None:
|
|
return False
|
|
if parsed.tzinfo is None:
|
|
parsed = parsed.replace(tzinfo=UTC)
|
|
|
|
current = (now or datetime.now(UTC)).astimezone(UTC)
|
|
return (parsed - current).total_seconds() <= _GITHUB_TOKEN_EXPIRY_SKEW_SECONDS
|
|
|
|
|
|
def _entry_expired(expires_at: str | None, cached_at: datetime, *, now: datetime) -> bool:
|
|
"""Expired when past the token's own expiry or the 24h cache cap."""
|
|
if now - cached_at >= _GITHUB_TOKEN_MAX_TTL:
|
|
return True
|
|
return _is_expired(expires_at, now=now)
|
|
|
|
|
|
def _evict_expired(*, now: datetime | None = None) -> None:
|
|
current = now or datetime.now(UTC)
|
|
stale = [
|
|
key
|
|
for key, (_token, expires_at, cached_at, _repo) in _GITHUB_TOKEN_CACHE.items()
|
|
if _entry_expired(expires_at, cached_at, now=current)
|
|
]
|
|
for key in stale:
|
|
_GITHUB_TOKEN_CACHE.pop(key, None)
|
|
|
|
|
|
def _cached_token_if_fresh(
|
|
thread_id: str | None, principal: str | None, *, expected_repo: Any = None
|
|
) -> tuple[str | None, str | None]:
|
|
if not thread_id:
|
|
return None, None
|
|
keys = []
|
|
if principal:
|
|
keys.append((thread_id, principal))
|
|
keys.append((thread_id, _BOT_PRINCIPAL))
|
|
expected = repo_cache_key(expected_repo)
|
|
for key in keys:
|
|
cached = _GITHUB_TOKEN_CACHE.get(key)
|
|
if not cached:
|
|
continue
|
|
token, expires_at, cached_at, bound_repo = cached
|
|
if _entry_expired(expires_at, cached_at, now=datetime.now(UTC)):
|
|
_GITHUB_TOKEN_CACHE.pop(key, None)
|
|
logger.info("Cached GitHub token for thread %s has expired; re-resolving", thread_id)
|
|
continue
|
|
if expected and bound_repo and expected != bound_repo:
|
|
_GITHUB_TOKEN_CACHE.pop(key, None)
|
|
logger.warning(
|
|
"Cached GitHub token for thread %s is bound to repo %s, not %s; "
|
|
"refusing cross-repo reuse",
|
|
thread_id,
|
|
bound_repo,
|
|
expected,
|
|
)
|
|
continue
|
|
return token, expires_at
|
|
return None, None
|
|
|
|
|
|
def _thread_id_from_config(run_config: Mapping[str, Any]) -> str | None:
|
|
configurable = run_config.get("configurable", {})
|
|
if not isinstance(configurable, Mapping):
|
|
return None
|
|
thread_id = configurable.get("thread_id")
|
|
return thread_id if isinstance(thread_id, str) and thread_id else None
|
|
|
|
|
|
def _repo_from_config(run_config: Mapping[str, Any]) -> Any:
|
|
configurable = run_config.get("configurable", {})
|
|
if not isinstance(configurable, Mapping):
|
|
return None
|
|
return configurable.get("repo")
|
|
|
|
|
|
def _principal_from_config(run_config: Mapping[str, Any]) -> str | None:
|
|
configurable = run_config.get("configurable", {})
|
|
if not isinstance(configurable, Mapping):
|
|
return None
|
|
return github_token_principal(
|
|
login=configurable.get("github_login"),
|
|
email=configurable.get("user_email"),
|
|
)
|
|
|
|
|
|
def get_github_token(run_config: Mapping[str, Any] | None = None) -> str | None:
|
|
"""Resolve the current thread's GitHub token from process memory."""
|
|
resolved = run_config if run_config is not None else get_config()
|
|
token, _expires_at = _cached_token_if_fresh(
|
|
_thread_id_from_config(resolved),
|
|
_principal_from_config(resolved),
|
|
expected_repo=_repo_from_config(resolved),
|
|
)
|
|
return token
|
|
|
|
|
|
async def get_github_token_from_thread(
|
|
thread_id: str, *, principal: str | None = None, expected_repo: Any = None
|
|
) -> tuple[str | None, str | None]:
|
|
"""Resolve the current process's cached GitHub token for a thread and principal."""
|
|
return _cached_token_if_fresh(thread_id, principal, expected_repo=expected_repo)
|
|
|
|
|
|
async def invalidate_cached_github_token(thread_id: str) -> None:
|
|
"""Clear every cached GitHub token for a thread."""
|
|
for key in [key for key in _GITHUB_TOKEN_CACHE if key[0] == thread_id]:
|
|
_GITHUB_TOKEN_CACHE.pop(key, None)
|
|
logger.info("Invalidated cached GitHub token for thread %s", thread_id)
|