refactor(agent/adapters): simplify azure identity, response guards, error surface, retry utils (-603 LOC)

Drop dead read_error_body_or_default / LAYER_RUNTIME; inline single-use
_build_default_credential/_safe_close; compact incident narratives to their
invariants in the guards. Byte-cap and deadline semantics of
read_streaming_error_body verified identical.
This commit is contained in:
Teknium
2026-09-02 11:21:12 -07:00
parent 63abd4d174
commit 6a88ec4eb3
7 changed files with 420 additions and 940 deletions
+176 -404
View File
@@ -1,32 +1,19 @@
"""Microsoft Entra ID adapter for Microsoft Foundry.
Provides keyless authentication for Microsoft Foundry deployments using the
`azure-identity` SDK's `DefaultAzureCredential` chain (env service principal
→ workload identity → managed identity → VS Code → Azure CLI → azd →
PowerShell → broker).
Keyless auth via the `azure-identity` ``DefaultAzureCredential`` chain (env
service principal → workload identity → managed identity → VS Code → Azure
CLI → azd → PowerShell → broker). Mirrors ``agent/bedrock_adapter.py``:
Architecture mirrors `agent/bedrock_adapter.py`:
* Lazy import. `azure-identity` is only loaded when ``model.auth_mode =
entra_id`` is selected. Users who stick with `AZURE_FOUNDRY_API_KEY`
never pay the import cost.
* SDK-callable contract. The public entry point ``build_token_provider``
returns a zero-arg callable produced by ``get_bearer_token_provider`` —
this is exactly the value Microsoft's documented sample plugs into
``OpenAI(api_key=token_provider, base_url=...)``. The OpenAI SDK calls
it before every request, so token refresh is transparent.
* Three explicit consumer-side helpers (display / cache / http-bearer)
rather than one generic "materialize" function — splitting them by
purpose prevents accidental token-minting in logging paths or token
leakage into cache keys / dashboard JSON.
* No persisted JWT. ``azure-identity`` caches in-process and (where
available) in the OS keychain or ``~/.IdentityService``. Hermes does
not duplicate that storage in ``auth.json``.
* Lazy import: `azure-identity` loads only when ``model.auth_mode = entra_id``.
* ``build_token_provider`` returns the zero-arg callable Microsoft's sample
plugs into ``OpenAI(api_key=token_provider, ...)``; the SDK calls it before
every request, so refresh is transparent.
* Consumer helpers are split by purpose (display / cache / http-bearer) so
logging paths never mint tokens and tokens never leak into cache keys.
* No persisted JWT: azure-identity caches in-process / OS keychain; Hermes
does not duplicate that in ``auth.json``.
Reference: https://learn.microsoft.com/azure/ai-foundry/foundry-models/how-to/configure-entra-id
Requires: ``azure-identity`` (optional dependency — only needed when
``model.auth_mode = entra_id``).
"""
from __future__ import annotations
@@ -40,28 +27,23 @@ from typing import Any, Callable, Dict, Optional
logger = logging.getLogger(__name__)
# Microsoft-documented scope for Foundry inference auth. Both the new
# Foundry portal and the legacy Azure OpenAI managed-identity docs use
# this scope for ALL Foundry endpoint shapes (*.openai.azure.com,
# *.services.ai.azure.com, *.ai.azure.com). The older control-plane
# scope ``https://cognitiveservices.azure.com/.default`` is for ARM
# resource management and is rejected for inference by newer
# resources — users with that requirement override via
# ``model.entra.scope`` in config.yaml.
# Microsoft-documented Foundry inference scope for ALL endpoint shapes
# (*.openai.azure.com, *.services.ai.azure.com, *.ai.azure.com). The older
# ``https://cognitiveservices.azure.com/.default`` is an ARM control-plane
# scope rejected for inference by newer resources; override via
# ``model.entra.scope`` if required.
SCOPE_AI_AZURE_DEFAULT = "https://ai.azure.com/.default"
# ---------------------------------------------------------------------------
# Lazy SDK import — only loaded when the Entra path is actually used.
# ---------------------------------------------------------------------------
_AZURE_IDENTITY_FEATURE = "provider.azure_identity"
_INSTALL_MSG = "The 'azure-identity' package is required for Azure AI Foundry Entra ID authentication. "
_LAZY_INSTALL_HINT = (
"pip install azure-identity manually, or enable lazy installs (security.allow_lazy_installs: true in config.yaml)."
)
_AUTH_HEADERS = ("Authorization", "authorization", "Api-Key", "api-key", "X-Api-Key", "x-api-key")
def has_azure_identity_installed() -> bool:
"""Return True if `azure-identity` can be imported right now.
Cheap check — does not walk the credential chain.
"""
"""Cheap importability check — does not walk the credential chain."""
try:
import azure.identity # noqa: F401
return True
@@ -70,11 +52,7 @@ def has_azure_identity_installed() -> bool:
def _require_azure_identity():
"""Import ``azure.identity``, lazy-installing it if allowed.
Raises ``ImportError`` with a clear actionable message when the
package is missing and lazy installs are disabled.
"""
"""Import ``azure.identity``, lazy-installing if allowed; ImportError with an actionable message otherwise."""
try:
import azure.identity as _ai
return _ai
@@ -82,68 +60,35 @@ def _require_azure_identity():
try:
from tools.lazy_deps import ensure, FeatureUnavailable
except ImportError as exc:
raise ImportError(
"The 'azure-identity' package is required for Azure AI "
"Foundry Entra ID authentication. Install it with: "
"pip install azure-identity"
) from exc
raise ImportError(_INSTALL_MSG + "Install it with: pip install azure-identity") from exc
try:
ensure(_AZURE_IDENTITY_FEATURE, prompt=False)
except FeatureUnavailable as exc:
raise ImportError(
"The 'azure-identity' package is required for Azure AI "
"Foundry Entra ID authentication. " + str(exc)
) from exc
raise ImportError(_INSTALL_MSG + str(exc)) from exc
# Retry import after lazy install.
import azure.identity as _ai # noqa: WPS440
import azure.identity as _ai # noqa: WPS440 — retry after lazy install
return _ai
def reset_credential_cache() -> None:
"""Clear the cached ``DefaultAzureCredential``. Used by tests and
profile switches.
Defensive against tests that ``monkeypatch.setattr`` over
``build_credential`` with a plain (non-lru-cached) function — those
won't expose ``cache_clear()`` until pytest reverts the patch.
"""
"""Clear the cached ``DefaultAzureCredential`` (tests, profile switches). Tolerates
tests that monkeypatch ``build_credential`` with a plain function lacking ``cache_clear``."""
cache_clear = getattr(build_credential, "cache_clear", None)
if callable(cache_clear):
cache_clear()
# ---------------------------------------------------------------------------
# Token-provider construction
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class EntraIdentityConfig:
"""Serializable Entra ID config.
"""Hermes-managed Entra knobs. Everything else (tenant, SP secret,
federated token file, sovereign authority, ``AZURE_CLIENT_ID``...) flows
through azure-identity's standard ``AZURE_*`` env vars.
Captures the Hermes-managed Entra knobs we need outside Azure SDK
environment configuration. Everything else
(tenant ID, service principal secret, federated token file, sovereign
cloud authority, etc.) flows through azure-identity's standard
``AZURE_*`` env vars — see the Bedrock pattern in
``hermes_cli/runtime_provider.py:1310-1377`` for the analogous
"let the SDK read env" approach.
``scope`` is Microsoft's documented Foundry inference audience. Almost
everyone uses the default; sovereign-cloud / non-standard tenants can
override via ``model.entra.scope``. Identity selection (user-assigned
managed identity, workload identity, service principal, tenant, authority)
stays in the standard Azure SDK env vars such as ``AZURE_CLIENT_ID``.
``exclude_interactive_browser`` is kept as an internal constructor knob
so probes stay non-interactive by default. It is not written by the setup
wizard.
The dataclass is frozen so it's hashable for ``functools.lru_cache``
keying, and serializable across multiprocessing boundaries (workers
rebuild the credential inside their own process).
``exclude_interactive_browser`` is an internal knob keeping probes
non-interactive; the setup wizard never writes it. Frozen so it is
hashable for ``lru_cache`` and serializable across multiprocessing
(workers rebuild the credential in their own process).
"""
scope: str = SCOPE_AI_AZURE_DEFAULT
@@ -154,261 +99,156 @@ class EntraIdentityConfig:
object.__setattr__(self, "scope", scope)
def to_dict(self) -> Dict[str, Any]:
return {
"scope": self.scope,
"exclude_interactive_browser": self.exclude_interactive_browser,
}
return {"scope": self.scope, "exclude_interactive_browser": self.exclude_interactive_browser}
@classmethod
def from_dict(cls, data: Optional[Dict[str, Any]],
*, default_scope: Optional[str] = None) -> "EntraIdentityConfig":
def from_dict(cls, data: Optional[Dict[str, Any]], *, default_scope: Optional[str] = None) -> "EntraIdentityConfig":
data = data or {}
scope = str(data.get("scope") or "").strip() or default_scope or SCOPE_AI_AZURE_DEFAULT
exclude_browser = bool(data.get("exclude_interactive_browser", True))
return cls(
scope=scope,
exclude_interactive_browser=exclude_browser,
scope=str(data.get("scope") or "").strip() or default_scope or SCOPE_AI_AZURE_DEFAULT,
exclude_interactive_browser=bool(data.get("exclude_interactive_browser", True)),
)
def _build_default_credential(config: EntraIdentityConfig) -> Any:
"""Construct a ``DefaultAzureCredential`` for ``config``.
Only Hermes-selected knobs are passed as kwargs. Everything else
(tenant, service principal secret, federated token file, sovereign
cloud authority, etc.) is read by ``azure-identity`` from the
standard ``AZURE_*`` environment variables — see Microsoft's
documented credential resolution chain. Users configure those in
``~/.hermes/.env`` or the deployment environment.
"""
@functools.lru_cache(maxsize=1)
def build_credential(config: EntraIdentityConfig) -> Any:
"""Cached ``DefaultAzureCredential``. ``maxsize=1`` is intentional: a process uses one
``model.entra.*`` block at a time (a second config just evicts the first). Only Hermes
knobs are passed as kwargs; the rest comes from ``AZURE_*`` env vars."""
ai = _require_azure_identity()
kwargs: Dict[str, Any] = {}
# SDK default is True (browser excluded); only pass when the user
# explicitly opts in to interactive browser auth.
# SDK default already excludes the browser; only pass when opting in.
if not config.exclude_interactive_browser:
kwargs["exclude_interactive_browser_credential"] = False
return ai.DefaultAzureCredential(**kwargs)
@functools.lru_cache(maxsize=1)
def build_credential(config: EntraIdentityConfig) -> Any:
"""Return the cached ``DefaultAzureCredential`` for ``config``.
Hermes processes use exactly one Entra config at a time (the
``model.entra.*`` block in config.yaml drives every aux task,
subagent, and credential probe in the session). ``maxsize=1`` is
intentional: it reflects the actual usage pattern and keeps the
cache trivially small.
``EntraIdentityConfig`` is a frozen dataclass, so it's hashable and
safe as an LRU-cache key. ``functools.lru_cache`` is thread-safe in
CPython.
If two distinct configs are ever passed (tests do this; production
rarely), the LRU eviction handles it correctly — each call still
returns a credential matching its config; only one is cached at a
time. Use :func:`reset_credential_cache` to clear (e.g. in tests).
"""
return _build_default_credential(config)
def _resolve_config(config: Optional[EntraIdentityConfig], scope: Optional[str],
**overrides: Any) -> EntraIdentityConfig:
if config is not None:
return config
return EntraIdentityConfig(scope=(scope or "").strip() or SCOPE_AI_AZURE_DEFAULT, **overrides)
def build_token_provider(scope: Optional[str] = None,
*,
config: Optional[EntraIdentityConfig] = None,
base_url: Optional[str] = None,
exclude_interactive_browser: bool = True,
def _install_failure(allow_install: bool) -> Optional[Dict[str, Any]]:
"""None when ``azure.identity`` is importable (lazy-installing if allowed), else ``{"error", "hint"}``."""
if has_azure_identity_installed():
return None
if not allow_install:
return {"error": "azure-identity not installed",
"hint": "pip install azure-identity (or rely on lazy install at first use)"}
try:
_require_azure_identity()
except ImportError as exc:
return {"error": str(exc) or "azure-identity not installed", "hint": _LAZY_INSTALL_HINT, "exc": exc}
return None
def build_token_provider(scope: Optional[str] = None, *, config: Optional[EntraIdentityConfig] = None,
base_url: Optional[str] = None, exclude_interactive_browser: bool = True,
) -> Callable[[], str]:
"""Return a zero-arg callable that mints a fresh Entra bearer JWT.
"""Zero-arg callable minting a fresh Entra bearer JWT — pass as ``OpenAI(api_key=...)``.
The returned callable is exactly what Microsoft's documented Foundry
sample expects::
from openai import OpenAI
client = OpenAI(
base_url="https://my-resource.openai.azure.com/openai/v1/",
api_key=build_token_provider(),
)
Scope resolution order:
1. ``config.scope`` when a config object is supplied
2. explicit ``scope`` kwarg
3. ``SCOPE_AI_AZURE_DEFAULT`` (Microsoft's documented Foundry scope)
``base_url`` is unused today and kept for back-compat. Tenant /
service-principal / sovereign-cloud configuration flows through
``azure-identity``'s standard ``AZURE_*`` environment variables —
see :func:`_build_default_credential` for the rationale.
NOT serializable across process boundaries. For multiprocessing
workers, serialize the ``EntraIdentityConfig`` and rebuild the
provider inside the worker.
Scope precedence: ``config.scope`` > ``scope`` kwarg > default. ``base_url`` is unused
(back-compat). Not picklable: ship the ``EntraIdentityConfig`` and rebuild in the worker.
"""
ai = _require_azure_identity()
if config is None:
config = EntraIdentityConfig(
scope=scope or SCOPE_AI_AZURE_DEFAULT,
exclude_interactive_browser=exclude_interactive_browser,
)
config = _resolve_config(config, scope, exclude_interactive_browser=exclude_interactive_browser)
credential = build_credential(config)
return ai.get_bearer_token_provider(credential, config.scope)
# ---------------------------------------------------------------------------
# Credential probing
# ---------------------------------------------------------------------------
def has_azure_identity_credentials(scope: Optional[str] = None,
*,
config: Optional[EntraIdentityConfig] = None,
timeout_seconds: float = 10.0,
allow_install: bool = True,
**overrides: Any) -> bool:
"""Best-effort probe: can `DefaultAzureCredential` mint a token now?
Runs ``credential.get_token(scope)`` under a thread-based timeout so
a slow token service can't hang the caller. Returns False on any
error — never raises. Use for ``hermes doctor`` /
``hermes auth status`` / wizard preflight.
``allow_install``: when True (default) and ``azure-identity`` is not
importable, the adapter triggers the standard lazy-install path
(subject to ``security.allow_lazy_installs``) before probing. Set
False to make this strictly an "is installed?" check — used on hot
paths like CLI startup where we never want pip to run.
NOT used by ``is_provider_configured()`` — that path is structural
only (no token mint), so CLI startup doesn't pay this latency.
"""
if not has_azure_identity_installed():
if not allow_install:
return False
try:
_require_azure_identity()
except ImportError as exc:
logger.debug("azure-identity lazy install unavailable: %s", exc)
return False
if config is None:
effective_scope = (scope or "").strip() or SCOPE_AI_AZURE_DEFAULT
config = EntraIdentityConfig(scope=effective_scope, **overrides)
result = {"ok": False}
def _probe() -> None:
try:
credential = build_credential(config)
tok = credential.get_token(config.scope)
result["ok"] = bool(getattr(tok, "token", None))
except Exception as exc:
logger.debug("Entra credential probe failed: %s", exc)
result["ok"] = False
thread = threading.Thread(target=_probe, daemon=True)
thread.start()
thread.join(timeout=max(0.01, timeout_seconds))
if thread.is_alive():
logger.debug("Entra token service probe timed out after %ss", timeout_seconds)
return False
return bool(result.get("ok"))
def describe_active_credential(config: Optional[EntraIdentityConfig] = None,
*,
scope: Optional[str] = None,
timeout_seconds: float = 10.0,
allow_install: bool = True,
**overrides: Any) -> Dict[str, Any]:
"""Return diagnostic info about the active credential chain.
Best-effort: runs ``get_token()`` and inspects what came back.
Designed for ``hermes doctor`` and the wizard preflight — never
raises, returns ``{"ok": False, "error": ...}`` on failure.
``allow_install``: when True (default) and ``azure-identity`` is not
importable, the adapter triggers the standard lazy-install path
(subject to ``security.allow_lazy_installs``) before probing. The
install failure is surfaced as the diagnostic error when it fails.
Set False for hot CLI paths that should never trigger pip.
``azure-identity`` doesn't expose the winning inner credential as
a public field, so we report a coarse picture (env vars present,
token expiry, claims-derived tenant) rather than the credential
class name. Users wanting the precise class can run with
``AZURE_LOG_LEVEL=DEBUG``.
"""
info: Dict[str, Any] = {"ok": False}
if not has_azure_identity_installed():
if not allow_install:
info["error"] = "azure-identity not installed"
info["hint"] = (
"pip install azure-identity (or rely on lazy install at "
"first use)"
)
return info
try:
_require_azure_identity()
except ImportError as exc:
info["error"] = str(exc) or "azure-identity not installed"
info["hint"] = (
"pip install azure-identity manually, or enable lazy "
"installs (security.allow_lazy_installs: true in "
"config.yaml)."
)
return info
if config is None:
effective_scope = (scope or "").strip() or SCOPE_AI_AZURE_DEFAULT
config = EntraIdentityConfig(scope=effective_scope, **overrides)
info["scope"] = config.scope
# Tenant / authority / service-principal config flow through the
# standard ``AZURE_*`` env vars; surface them below.
if os.environ.get("AZURE_TENANT_ID", "").strip():
info["tenant_id_env"] = os.environ["AZURE_TENANT_ID"].strip()
# Surface which env-var sources are present without minting yet.
# Credential-bearing vars (AZURE_CLIENT_SECRET, AZURE_FEDERATED_TOKEN_FILE)
# are read through the profile secret scope so a multiplexed profile's
# diagnostics don't report another profile's env-bridged credentials;
# unscoped CLI probes keep the legacy env read (Slack pattern).
def _scoped_env(name: str) -> str:
try:
from agent.secret_scope import UnscopedSecretError, get_secret
try:
return (get_secret(name) or "").strip()
except UnscopedSecretError:
pass
except Exception:
pass
return os.environ.get(name, "").strip()
env_sources = []
if _scoped_env("AZURE_FEDERATED_TOKEN_FILE"):
env_sources.append("WorkloadIdentityCredential (AZURE_FEDERATED_TOKEN_FILE)")
if (os.environ.get("AZURE_CLIENT_ID", "").strip()
and _scoped_env("AZURE_CLIENT_SECRET")
and os.environ.get("AZURE_TENANT_ID", "").strip()):
env_sources.append("EnvironmentCredential (client secret)")
if os.environ.get("IDENTITY_ENDPOINT", "").strip() or os.environ.get("MSI_ENDPOINT", "").strip():
env_sources.append("ManagedIdentityCredential (IDENTITY_ENDPOINT)")
info["env_sources"] = env_sources
# Now try minting.
def _probe_token(config: EntraIdentityConfig, timeout_seconds: float) -> Optional[Dict[str, Any]]:
"""``get_token`` on a daemon thread under a hard deadline → ``{"token"}`` / ``{"error"}`` / None on timeout."""
result: Dict[str, Any] = {}
def _probe() -> None:
try:
credential = build_credential(config)
tok = credential.get_token(config.scope)
result["token"] = tok
result["token"] = build_credential(config).get_token(config.scope)
except Exception as exc:
result["error"] = str(exc)
thread = threading.Thread(target=_probe, daemon=True)
thread.start()
thread.join(timeout=max(0.01, timeout_seconds))
if thread.is_alive():
return None if thread.is_alive() else result
def has_azure_identity_credentials(scope: Optional[str] = None, *, config: Optional[EntraIdentityConfig] = None,
timeout_seconds: float = 10.0, allow_install: bool = True,
**overrides: Any) -> bool:
"""Timeout-bounded probe: can the chain mint a token now? Never raises.
``allow_install=False`` makes it a strict "is installed?" check for hot paths (CLI startup)
where pip must never run. NOT used by ``is_provider_configured()`` (structural, no mint).
"""
failure = _install_failure(allow_install)
if failure is not None:
if "exc" in failure:
logger.debug("azure-identity lazy install unavailable: %s", failure["exc"])
return False
config = _resolve_config(config, scope, **overrides)
result = _probe_token(config, timeout_seconds)
if result is None:
logger.debug("Entra token service probe timed out after %ss", timeout_seconds)
return False
if "error" in result:
logger.debug("Entra credential probe failed: %s", result["error"])
return False
return bool(getattr(result.get("token"), "token", None))
def _env(name: str) -> str:
return os.environ.get(name, "").strip()
def _scoped_env(name: str) -> str:
"""Credential-bearing env read via the profile secret scope so a multiplexed profile never
reports another profile's env-bridged credentials; unscoped CLI probes fall back to plain env."""
try:
from agent.secret_scope import UnscopedSecretError, get_secret
try:
return (get_secret(name) or "").strip()
except UnscopedSecretError:
pass
except Exception:
pass
return _env(name)
# (label, predicate) for env-var-driven credential sources, in chain order.
_ENV_SOURCE_CHECKS = (
("WorkloadIdentityCredential (AZURE_FEDERATED_TOKEN_FILE)", lambda: _scoped_env("AZURE_FEDERATED_TOKEN_FILE")),
("EnvironmentCredential (client secret)",
lambda: _env("AZURE_CLIENT_ID") and _scoped_env("AZURE_CLIENT_SECRET") and _env("AZURE_TENANT_ID")),
("ManagedIdentityCredential (IDENTITY_ENDPOINT)", lambda: _env("IDENTITY_ENDPOINT") or _env("MSI_ENDPOINT")),
)
def describe_active_credential(config: Optional[EntraIdentityConfig] = None, *, scope: Optional[str] = None,
timeout_seconds: float = 10.0, allow_install: bool = True,
**overrides: Any) -> Dict[str, Any]:
"""Doctor / preflight diagnostics. Never raises; ``{"ok": False, "error": ...}`` on failure.
azure-identity hides the winning inner credential, so this reports a coarse picture (env
sources, token expiry) rather than a class name; ``AZURE_LOG_LEVEL=DEBUG`` shows the chain.
"""
info: Dict[str, Any] = {"ok": False}
failure = _install_failure(allow_install)
if failure is not None:
info["error"], info["hint"] = failure["error"], failure["hint"]
return info
config = _resolve_config(config, scope, **overrides)
info["scope"] = config.scope
if _env("AZURE_TENANT_ID"):
info["tenant_id_env"] = _env("AZURE_TENANT_ID")
info["env_sources"] = [label for label, present in _ENV_SOURCE_CHECKS if present()]
result = _probe_token(config, timeout_seconds)
if result is None:
info["error"] = f"Token probe timed out after {timeout_seconds:.0f}s"
info["hint"] = (
"DefaultAzureCredential can be slow when the token service is unreachable "
@@ -416,54 +256,30 @@ def describe_active_credential(config: Optional[EntraIdentityConfig] = None,
"AZURE_CLIENT_ID / AZURE_TENANT_ID / AZURE_CLIENT_SECRET."
)
return info
if "error" in result:
info["error"] = result["error"]
return info
token = result.get("token")
if token is None:
info["error"] = "credential chain exhausted"
return info
info["ok"] = True
info["expires_on"] = getattr(token, "expires_on", None)
return info
# ---------------------------------------------------------------------------
# Consumer-side helpers — split by purpose to prevent accidental token
# minting in logging / cache-key / dashboard paths.
# ---------------------------------------------------------------------------
# Consumer-side helpers — split by purpose so logging / cache-key / dashboard paths never mint tokens.
def is_token_provider(value: Any) -> bool:
"""Return True when ``value`` is a callable Entra token provider.
Used at the seams where a consumer must decide between
string-API-key semantics and bearer-callable semantics.
"""
"""True when ``value`` is a callable Entra token provider (vs. a string API key)."""
return callable(value) and not isinstance(value, str)
def materialize_bearer_for_http(value: Any) -> str:
"""Return a fresh Bearer JWT for a manual HTTP request.
"""Mint a fresh Bearer JWT for a manual HTTP request (calls the provider once).
Only call this at sites that must construct an ``Authorization``
header outside the OpenAI SDK (e.g. ``hermes_cli/azure_detect.py``).
Calls the callable exactly once and returns the resulting token.
**Anthropic SDK integration:** the Anthropic Python SDK does not
accept a ``Callable[[], str]`` for ``auth_token``. Instead,
:func:`build_bearer_http_client` returns an ``httpx.Client`` whose
request event hook calls this function and rewrites the
``Authorization`` header per request — and that client is passed to
the Anthropic SDK via ``http_client=...``. See
:func:`agent.anthropic_adapter.build_anthropic_client` for the
consumer.
Raises ``ValueError`` if ``value`` is not a callable token provider
or non-empty string.
Only for sites building ``Authorization`` outside the OpenAI SDK; the Anthropic SDK can't take
a callable, so :func:`build_bearer_http_client` calls this from an httpx hook. ``ValueError``
on an unusable value or empty token.
"""
if is_token_provider(value):
token = value()
@@ -475,41 +291,20 @@ def materialize_bearer_for_http(value: Any) -> str:
raise ValueError("no usable api_key / token provider")
def _strip_auth_headers(request: Any) -> None:
for header_name in _AUTH_HEADERS:
request.headers.pop(header_name, None)
def build_bearer_http_client(token_provider: Callable[[], str], **httpx_kwargs: Any) -> Any:
"""Return an ``httpx.Client`` that mints a fresh Entra bearer JWT
per outbound request.
"""``httpx.Client`` minting a fresh Entra bearer JWT per outbound request.
The Anthropic SDK (≤ 0.86.0 at the time of writing) stores
``api_key`` / ``auth_token`` as static strings and computes the
``Authorization`` header at construction time. To get per-request
token refresh (the Microsoft-recommended Foundry pattern for
callable bearer providers), we install an httpx ``request`` event
hook on a custom client and pass that client to the SDK via
``http_client=...``. The hook:
1. Calls :func:`materialize_bearer_for_http` to mint a fresh JWT
(azure-identity caches internally — this is cheap when the
cached token is still valid).
2. Strips any pre-set ``Authorization`` / ``api-key`` /
``x-api-key`` headers the SDK may have added (avoids
conflicting auth values).
3. Sets ``Authorization: Bearer <fresh-jwt>``.
``token_provider`` must be a zero-arg callable returning a string —
typically the result of :func:`build_token_provider`.
``httpx_kwargs`` are forwarded verbatim to ``httpx.Client(...)`` so
callers can attach a ``timeout``, ``transport``, ``proxy``, etc.
Raises ``ImportError`` if ``httpx`` is not installed (it is a
transitive dependency of both ``openai`` and ``anthropic`` SDKs, so
in practice always available when this helper is reached).
The Anthropic SDK computes ``Authorization`` once at construction, so per-request refresh needs
a ``request`` hook: mint (cheap — azure-identity caches), strip pre-set auth headers, set
``Authorization: Bearer``. ``httpx_kwargs`` are forwarded verbatim (``timeout``, ``transport``...).
"""
if not is_token_provider(token_provider):
raise ValueError(
"build_bearer_http_client requires a zero-arg callable "
"token provider"
)
raise ValueError("build_bearer_http_client requires a zero-arg callable token provider")
try:
import httpx
@@ -524,48 +319,25 @@ def build_bearer_http_client(token_provider: Callable[[], str], **httpx_kwargs:
try:
token = materialize_bearer_for_http(token_provider)
except ValueError as exc:
# Token provider failed (chain exhausted, token service unreachable,
# az login expired, etc.). Strip any auth headers the SDK
# may have set — including our own placeholder sentinel
# ``entra-id-bearer-via-http-hook`` from
# ``_build_anthropic_client_with_bearer_hook`` — so the
# outbound request hits Azure with NO Authorization rather
# than with the placeholder. Azure returns a clean 401
# "missing auth" that is easier to diagnose than a 401
# against the sentinel string, and the sentinel never
# appears in upstream access logs.
#
# Log at WARNING (not DEBUG) so the misconfiguration is
# visible at default log levels.
# Chain exhausted / az login expired: strip ALL auth headers (incl. the anthropic_adapter
# placeholder sentinel) so Azure returns a clean "missing auth" 401 and the sentinel never
# reaches upstream logs. WARNING so the misconfiguration is visible at default levels.
logger.warning(
"Bearer hook: Entra ID token provider returned empty (%s) "
"— stripping Authorization headers. Azure will respond 401. "
"Run `hermes doctor` or `az login` to recover.",
exc,
)
for header_name in ("Authorization", "authorization", "Api-Key", "api-key", "X-Api-Key", "x-api-key"):
request.headers.pop(header_name, None)
_strip_auth_headers(request)
return
for header_name in ("Authorization", "authorization", "Api-Key", "api-key", "X-Api-Key", "x-api-key"):
request.headers.pop(header_name, None)
_strip_auth_headers(request)
request.headers["Authorization"] = f"Bearer {token}"
return httpx.Client(
event_hooks={"request": [_inject_bearer]},
**httpx_kwargs,
)
return httpx.Client(event_hooks={"request": [_inject_bearer]}, **httpx_kwargs)
__all__ = [
"EntraIdentityConfig",
"SCOPE_AI_AZURE_DEFAULT",
"build_bearer_http_client",
"build_credential",
"build_token_provider",
"describe_active_credential",
"has_azure_identity_credentials",
"has_azure_identity_installed",
"is_token_provider",
"materialize_bearer_for_http",
"reset_credential_cache",
"EntraIdentityConfig", "SCOPE_AI_AZURE_DEFAULT", "build_bearer_http_client", "build_credential",
"build_token_provider", "describe_active_credential", "has_azure_identity_credentials",
"has_azure_identity_installed", "is_token_provider", "materialize_bearer_for_http", "reset_credential_cache",
]
+33 -87
View File
@@ -1,55 +1,38 @@
"""Bounded reads of HTTP error response bodies.
When a provider returns a non-OK status on a *streaming* request, Hermes reads
the response body to build a useful diagnostic error. A bare ``response.read()``
on a streaming httpx response is unbounded in two dangerous ways:
On a non-OK *streaming* response Hermes reads the body for a diagnostic. A
bare ``response.read()`` is unbounded two ways: a server can stream an
arbitrarily large body (memory), or open the body and stall forever (hang).
The diagnostic is only ever shown truncated to a few hundred chars, so
``read_streaming_error_body`` caps bytes and enforces a hard wall-clock
deadline; callers feed the returned text to their error builders instead of
touching ``response.text`` (unbounded / raises after a partial stream read).
1. A server can declare (or stream) an arbitrarily large body, so the read can
balloon memory.
2. A server can open the body and then stall forever (no ``Content-Length``,
no further bytes), so the read hangs the agent indefinitely.
Subtlety: ``httpx.iter_bytes()`` blocks *inside* the socket read, so a
deadline checked only between chunks can't interrupt a mid-chunk stall until
httpx's own (30s+) read timeout fires. The read therefore runs on a daemon
thread; on timeout we close the response (unblocking the read) and return the
partial bytes collected so far.
Both are realistic against a misbehaving proxy, a hijacked endpoint, or a
provider having a bad day. The diagnostic body is only ever shown to the user
truncated to a few hundred characters, so reading megabytes — or blocking
forever — buys nothing.
``read_streaming_error_body`` bounds the read to a byte cap and enforces a
hard wall-clock deadline, returning the decoded text snippet. Callers pass the
returned text into their existing error builders instead of touching
``response.text`` (which would be unbounded / would raise after a partial
stream read).
A subtlety the implementation must respect: ``httpx``'s ``iter_bytes()`` blocks
*inside* the C/socket read while waiting for the next chunk. A wall-clock check
placed only between yielded chunks cannot interrupt a server that opens the
body and then stalls mid-chunk — control never returns to Python until httpx's
own (often 30s+) read timeout fires. To guarantee a bounded stop regardless of
socket behavior, the read runs on a daemon worker thread and the caller waits
on it with a hard deadline; on timeout we close the response (which unblocks /
cancels the read) and return whatever partial bytes were collected.
Ported and adapted from openclaw/openclaw#95108 ("bound Anthropic error
streams"), generalized to cover Hermes's three streaming error-body sites
(native Gemini, Gemini Cloud Code, Antigravity Cloud Code).
Covers the three streaming error-body sites: native Gemini, Gemini Cloud
Code, Antigravity Cloud Code.
"""
from __future__ import annotations
import logging
import threading
from typing import List, Optional
from typing import List
import httpx
logger = logging.getLogger(__name__)
# Defaults chosen to comfortably hold any real provider error envelope (Google
# RPC error JSON, Anthropic error JSON) while rejecting pathological bodies.
# Comfortably holds any real provider error envelope (Google RPC / Anthropic
# error JSON) while rejecting pathological bodies.
DEFAULT_ERROR_BODY_MAX_BYTES = 64 * 1024
# Hard wall-clock deadline for the whole bounded read. A streaming error body
# that does not finish within this window is abandoned and the connection is
# closed; we keep whatever partial bytes arrived.
# Hard deadline for the whole read; past it the connection is closed and the
# partial bytes are kept.
DEFAULT_ERROR_BODY_TIMEOUT_S = 10.0
@@ -59,17 +42,11 @@ def read_streaming_error_body(
max_bytes: int = DEFAULT_ERROR_BODY_MAX_BYTES,
timeout_s: float = DEFAULT_ERROR_BODY_TIMEOUT_S,
) -> str:
"""Read a non-OK streaming response body with a byte cap and a hard deadline.
"""Read a non-OK streaming body with a byte cap and a hard deadline.
Returns the decoded body text (UTF-8, errors replaced), truncated to
``max_bytes``. Never raises: any transport error, stall, or oversize
condition is swallowed and the best-effort partial text (or an empty
string) is returned, because this runs on the error path and must not
mask the original HTTP failure with a read error.
The byte cap protects against huge bodies; the wall-clock deadline (enforced
via a worker thread so it can interrupt a socket read that stalls mid-chunk)
protects against bodies that open and then hang.
Returns UTF-8 text (errors replaced) truncated to ``max_bytes``. Never
raises: transport errors, stalls and oversize bodies yield best-effort
partial text (or ""), so a read error can't mask the original failure.
"""
chunks: List[bytes] = []
state = {"truncated": False}
@@ -82,12 +59,9 @@ def read_streaming_error_body(
if not chunk:
continue
remaining = max_bytes - total
if remaining <= 0:
state["truncated"] = True
break
if len(chunk) > remaining:
chunks.append(chunk[:remaining])
total += remaining
if remaining > 0:
chunks.append(chunk[:remaining])
state["truncated"] = True
break
chunks.append(chunk)
@@ -97,24 +71,19 @@ def read_streaming_error_body(
finally:
done.set()
worker = threading.Thread(
target=_drain, name="bounded-error-body-read", daemon=True
)
worker.start()
finished = done.wait(timeout=timeout_s)
if not finished:
threading.Thread(target=_drain, name="bounded-error-body-read", daemon=True).start()
if not done.wait(timeout=timeout_s):
logger.debug(
"bounded error-body read: hard timeout after %.1fs (%d bytes so far)",
timeout_s,
sum(len(c) for c in chunks),
)
# Closing the response cancels the in-flight socket read, letting the
# worker thread unwind. We do not join (it is a daemon and may be
# blocked in C); the partial `chunks` collected so far are returned.
_safe_close(response)
else:
_safe_close(response)
# Closing cancels any in-flight socket read so the worker unwinds. We do
# not join (it is a daemon and may be blocked in C).
try:
response.close()
except Exception: # noqa: BLE001
pass
if state["truncated"]:
logger.debug(
@@ -123,26 +92,3 @@ def read_streaming_error_body(
max_bytes,
)
return b"".join(chunks).decode("utf-8", errors="replace")
def _safe_close(response: httpx.Response) -> None:
try:
response.close()
except Exception: # noqa: BLE001
pass
def read_error_body_or_default(
response: httpx.Response,
*,
max_bytes: int = DEFAULT_ERROR_BODY_MAX_BYTES,
timeout_s: float = DEFAULT_ERROR_BODY_TIMEOUT_S,
) -> Optional[str]:
"""Like ``read_streaming_error_body`` but returns ``None`` on empty body.
Convenience for callers that distinguish "no body" from "empty string".
"""
text = read_streaming_error_body(
response, max_bytes=max_bytes, timeout_s=timeout_s
)
return text or None
+70 -133
View File
@@ -1,47 +1,30 @@
"""Deterministic-empty detection and cost-aware retry budgets (NS-503).
"""Deterministic-empty detection and cost-aware retry budgets.
When a provider returns an empty completion, the agent loop retries up to
3 times and then walks the fallback chain. Every attempt re-sends the full
conversation input — at large context on paid routes this bills the user
repeatedly for a turn that produces no text (the "charged ~$2.33 for an
empty answer" incident class).
On an empty completion the loop retries up to 3 times, then walks the
fallback chain — each attempt re-bills the full input. Signaled refusals
(``content_filter``, Anthropic ``refusal``, Bedrock guardrails) are already
terminal; this module handles *unsignaled* empties (success, zero output
tokens, generic finish reason — typical of portal-proxied refusals).
Signaled refusals (``finish_reason="content_filter"``, Anthropic
``stop_reason="refusal"``, Bedrock guardrails) are already terminal and
never reach the empty-retry loop. This module addresses the *unsignaled*
empties: the provider reports a successful completion with zero output
tokens and a generic finish reason (portal-proxied refusals commonly look
like this).
Two independent guards, both failing OPEN to legacy behaviour:
Two independent guards, both failing OPEN to today's behaviour:
1. Deterministic-empty: two consecutive empties, both with usage present and
``output_tokens == 0``, from the same (model, provider, finish_reason) →
the same prompt will keep producing the same empty, so skip remaining
retries and go straight to the fallback chain (a different model may
behave differently). Missing usage or ``output_tokens > 0`` (think-block
stripping, whitespace, flaky decoding) never classifies as deterministic.
2. Cost-aware budget: when one empty attempt's estimated input cost exceeds
the threshold (default $0.25), the retry budget drops from 3 to 1.
Unknown pricing / missing usage / included routes leave it untouched.
1. **Deterministic-empty detection** — two consecutive empty attempts,
both with usage present and ``output_tokens == 0``, from the same
(model, provider, finish_reason), are treated as deterministic: the
same prompt will keep producing the same empty. Remaining retries are
skipped and the loop proceeds straight to the fallback chain (a
different model may behave differently). Attempts with missing usage
or ``output_tokens > 0`` (model generated *something* — think-block
stripping, whitespace, flaky decoding) never classify as deterministic
and keep the full retry budget.
2. **Cost-aware retry budget** — when the estimated input cost of a
single empty attempt exceeds the configured threshold (default
$0.25), the empty-retry budget for this streak drops from 3 to 1.
Unknown pricing, missing usage, or included/subscription routes
leave the budget untouched.
Configured via the additive ``agent.empty_response_guard`` section in
``config.yaml`` (resolved once at agent init by ``agent_init``)::
Config (``config.yaml``, resolved once by ``agent_init`` and stashed on the
agent so the hot loop never re-reads config — no env vars)::
agent:
empty_response_guard:
enabled: true # false = legacy fixed 3-retry behaviour
cost_threshold_usd: 0.25 # per-attempt cost that halves the budget
Per project policy, no ``HERMES_*`` environment variables are involved —
``.env`` is reserved for credentials; behavioural settings live in
``config.yaml``.
"""
from __future__ import annotations
@@ -58,11 +41,10 @@ REDUCED_EMPTY_RETRY_BUDGET = 1
DEFAULT_COST_THRESHOLD_USD = Decimal("0.25")
DEFAULT_GUARD_ENABLED = True
# Attribute names stashed on the agent object. State is scoped to one
# consecutive empty streak: it is cleared whenever a streak starts
# (``_empty_content_retries == 0`` at record time), which transparently
# honours every existing reset site (turn start, compaction, tool
# success, fallback activation) without touching them.
# Agent-object attribute names. State is scoped to one consecutive empty
# streak: cleared whenever ``_empty_content_retries == 0`` at record time, so
# every existing counter-reset site (turn start, compaction, tool success,
# fallback activation) is honoured without touching it.
_ATTEMPTS_ATTR = "_empty_attempt_history"
_STREAK_COST_ATTR = "_empty_streak_cost_usd"
_ENABLED_ATTR = "_empty_guard_enabled"
@@ -85,23 +67,14 @@ class EmptyAttempt:
def resolve_guard_settings(section: Any) -> Tuple[bool, Decimal]:
"""Resolve ``agent.empty_response_guard`` config into (enabled, threshold).
Tolerant of malformed input: anything that isn't a well-formed dict
(or well-formed values within it) falls back to the schema defaults.
Called once per agent at init; the resolved values are stashed on the
agent object so the hot loop never re-reads config.
"""
"""Resolve ``agent.empty_response_guard`` into (enabled, threshold); malformed input → schema defaults."""
if not isinstance(section, dict):
return (DEFAULT_GUARD_ENABLED, DEFAULT_COST_THRESHOLD_USD)
enabled_raw = section.get("enabled", DEFAULT_GUARD_ENABLED)
if isinstance(enabled_raw, bool):
enabled = enabled_raw
elif isinstance(enabled_raw, str):
# YAML quoting can turn true/false into strings.
enabled = enabled_raw.strip().lower() not in ("0", "false", "no", "off")
else:
enabled = section.get("enabled", DEFAULT_GUARD_ENABLED)
if isinstance(enabled, str): # YAML quoting can turn true/false into strings.
enabled = enabled.strip().lower() not in ("0", "false", "no", "off")
elif not isinstance(enabled, bool):
enabled = DEFAULT_GUARD_ENABLED
threshold = DEFAULT_COST_THRESHOLD_USD
@@ -112,28 +85,19 @@ def resolve_guard_settings(section: Any) -> Tuple[bool, Decimal]:
if candidate > 0:
threshold = candidate
except Exception: # noqa: BLE001 — malformed config must not break init
logger.debug(
"empty-guard: invalid cost_threshold_usd %r, using default",
threshold_raw,
)
logger.debug("empty-guard: invalid cost_threshold_usd %r, using default", threshold_raw)
return (enabled, threshold)
def guard_enabled(agent: Any) -> bool:
"""Whether the guard is enabled for this agent (config-resolved).
Agents built before the config was threaded through (tests, embedded
callers) simply get the default: enabled.
"""
"""Config-resolved enabled flag; agents built without config default to enabled."""
value = getattr(agent, _ENABLED_ATTR, DEFAULT_GUARD_ENABLED)
return value if isinstance(value, bool) else DEFAULT_GUARD_ENABLED
def _cost_threshold_usd(agent: Any) -> Decimal:
value = getattr(agent, _THRESHOLD_ATTR, None)
if isinstance(value, Decimal) and value > 0:
return value
return DEFAULT_COST_THRESHOLD_USD
return value if isinstance(value, Decimal) and value > 0 else DEFAULT_COST_THRESHOLD_USD
def _attempts(agent: Any) -> List[EmptyAttempt]:
@@ -144,25 +108,32 @@ def _attempts(agent: Any) -> List[EmptyAttempt]:
return attempts
def _estimate_attempt_cost(agent: Any, response: Any) -> Optional[Decimal]:
"""Best-effort USD estimate for one attempt. None when unknown."""
def _normalized_usage(agent: Any, response: Any, what: str) -> Any:
"""Canonical usage for ``response`` or None (no usage / normalization failed)."""
raw_usage = getattr(response, "usage", None)
if not raw_usage:
return None
try:
from agent.usage_pricing import estimate_usage_cost, normalize_usage
from agent.usage_pricing import normalize_usage
return normalize_usage(raw_usage, provider=getattr(agent, "provider", None),
api_mode=getattr(agent, "api_mode", None))
except Exception: # noqa: BLE001 — pricing must never break the loop
logger.debug("empty-guard: %s failed", what, exc_info=True)
return None
def _estimate_attempt_cost(agent: Any, response: Any) -> Optional[Decimal]:
"""Best-effort USD estimate for one attempt. None when unknown."""
canonical = _normalized_usage(agent, response, "cost estimation")
if canonical is None:
return None
try:
from agent.usage_pricing import estimate_usage_cost
canonical = normalize_usage(
raw_usage,
provider=getattr(agent, "provider", None),
api_mode=getattr(agent, "api_mode", None),
)
result = estimate_usage_cost(
getattr(agent, "model", "") or "",
canonical,
provider=getattr(agent, "provider", None),
base_url=getattr(agent, "base_url", None),
api_key=getattr(agent, "api_key", None),
getattr(agent, "model", "") or "", canonical, provider=getattr(agent, "provider", None),
base_url=getattr(agent, "base_url", None), api_key=getattr(agent, "api_key", None),
)
except Exception: # noqa: BLE001 — pricing must never break the loop
logger.debug("empty-guard: cost estimation failed", exc_info=True)
@@ -172,31 +143,16 @@ def _estimate_attempt_cost(agent: Any, response: Any) -> Optional[Decimal]:
def _zero_output(agent: Any, response: Any) -> tuple:
"""Return (usage_present, zero_output) for a response, failing open."""
raw_usage = getattr(response, "usage", None)
if not raw_usage:
return (False, False)
try:
from agent.usage_pricing import normalize_usage
canonical = normalize_usage(
raw_usage,
provider=getattr(agent, "provider", None),
api_mode=getattr(agent, "api_mode", None),
)
except Exception: # noqa: BLE001
logger.debug("empty-guard: usage normalization failed", exc_info=True)
canonical = _normalized_usage(agent, response, "usage normalization")
if canonical is None:
return (False, False)
output = getattr(canonical, "output_tokens", None)
if output is None:
# A present-but-empty usage object (some proxies) normalizes to all zeros;
# a genuine completion always has input tokens — no evidence, fail open.
if output is None or getattr(canonical, "prompt_tokens", 0) <= 0:
return (False, False)
# A present-but-empty usage object (some proxies emit usage with no
# fields) normalizes to all zeros. A genuine completion always has
# input tokens — without them the usage is not evidence, fail open.
if getattr(canonical, "prompt_tokens", 0) <= 0:
return (False, False)
# Reasoning tokens count as real generation — a reasoning-only
# response is NOT a deterministic empty (the prefill-continuation
# path upstream owns that case).
# Reasoning tokens are real generation: a reasoning-only response is NOT
# a deterministic empty (the prefill-continuation path owns that case).
reasoning = getattr(canonical, "reasoning_tokens", 0) or 0
return (True, (output + reasoning) == 0)
@@ -204,10 +160,8 @@ def _zero_output(agent: Any, response: Any) -> tuple:
def record_empty_attempt(agent: Any, *, finish_reason: str, response: Any) -> None:
"""Record one empty completion in the current streak.
Must be called before ``_empty_content_retries`` is incremented for
this attempt: a counter of 0 marks the start of a new streak and
clears prior history (this transparently follows every existing
counter-reset site).
Call BEFORE ``_empty_content_retries`` is incremented: a counter of 0
marks a new streak and clears prior history.
"""
attempts = _attempts(agent)
if getattr(agent, "_empty_content_retries", 0) == 0:
@@ -215,15 +169,10 @@ def record_empty_attempt(agent: Any, *, finish_reason: str, response: Any) -> No
setattr(agent, _STREAK_COST_ATTR, Decimal("0"))
usage_present, zero_output = _zero_output(agent, response)
attempts.append(
EmptyAttempt(
model=str(getattr(agent, "model", "") or ""),
provider=str(getattr(agent, "provider", "") or ""),
finish_reason=str(finish_reason or ""),
usage_present=usage_present,
zero_output=zero_output,
)
)
attempts.append(EmptyAttempt(
model=str(getattr(agent, "model", "") or ""), provider=str(getattr(agent, "provider", "") or ""),
finish_reason=str(finish_reason or ""), usage_present=usage_present, zero_output=zero_output,
))
cost = _estimate_attempt_cost(agent, response)
if cost is not None and cost > 0:
@@ -232,22 +181,14 @@ def record_empty_attempt(agent: Any, *, finish_reason: str, response: Any) -> No
def deterministic_empty(agent: Any) -> bool:
"""True when the current streak looks deterministic.
Requires >= 2 consecutive attempts, ALL with usage present, zero
output tokens, and an identical (model, provider, finish_reason)
signature. Any attempt with missing usage or non-zero output keeps
this False (fail open — transients deserve their retries).
"""
"""True when >= 2 consecutive attempts ALL have usage present, zero output
and an identical signature. Any missing-usage / non-zero attempt → False
(fail open — transients deserve their retries)."""
if not guard_enabled(agent):
return False
attempts = getattr(agent, _ATTEMPTS_ATTR, None) or []
if len(attempts) < 2:
return False
first = attempts[0]
return all(
a.usage_present and a.zero_output and a.signature == first.signature
for a in attempts
return len(attempts) >= 2 and all(
a.usage_present and a.zero_output and a.signature == attempts[0].signature for a in attempts
)
@@ -257,9 +198,7 @@ def empty_retry_budget(agent: Any, response: Any) -> int:
if not guard_enabled(agent):
return DEFAULT_EMPTY_RETRY_BUDGET
cost = _estimate_attempt_cost(agent, response)
if cost is None:
return DEFAULT_EMPTY_RETRY_BUDGET
if cost >= _cost_threshold_usd(agent):
if cost is not None and cost >= _cost_threshold_usd(agent):
return REDUCED_EMPTY_RETRY_BUDGET
return DEFAULT_EMPTY_RETRY_BUDGET
@@ -267,6 +206,4 @@ def empty_retry_budget(agent: Any, response: Any) -> int:
def streak_cost_usd(agent: Any) -> Optional[Decimal]:
"""Accumulated estimated cost of the current empty streak, if known."""
cost = getattr(agent, _STREAK_COST_ATTR, None)
if cost is None or cost <= 0:
return None
return cost
return cost if cost is not None and cost > 0 else None
+60 -138
View File
@@ -1,29 +1,20 @@
"""Structured error-surface descriptors for UI clients (Desktop/TUI).
Maps the internal failure taxonomy (``agent.error_classifier.FailoverReason``
values carried in turn results as ``failure_reason``, or raw exceptions from
the turn dispatcher) onto a small, stable wire descriptor:
Maps the internal failure taxonomy (``FailoverReason`` values carried in turn
results as ``failure_reason``, or raw exceptions from the turn dispatcher)
onto a small, stable wire descriptor::
{"layer": <ui layer>, "code": <specific code>, "retryable": <bool>}
The *layer* names which part of the stack failed, so clients can say
"Provider error" / "Gateway error" instead of toasting an opaque string and
leaving the user to guess whether the model, the gateway, or the app froze:
Layers (wire values): provider (model API rejected/failed the call), endpoint
(user-configured custom/local endpoint transport failure), streaming (SSE
dropped mid-turn), auth, billing (fallback signal; clients usually have a
richer ``billing_block``), gateway (local runtime errored), disk (disk full /
persistence failure).
provider — the model/provider API rejected or failed the call
endpoint — a user-configured custom/local endpoint failed (transport)
streaming — the provider's SSE/stream connection dropped mid-turn
auth — authentication/authorization failed
billing — credits/quota wall (clients usually have a richer
billing_block descriptor; this is the fallback signal)
gateway — the local gateway/agent runtime itself errored
runtime — agent initialization / local environment failure
disk — local disk full / persistence failure
This module is intentionally dependency-light and NEVER raises: surfacing
diagnostics must not be able to break the error path it describes. Clients
treat the descriptor as advisory — an absent or partial descriptor falls
back to today's string-sniffing behavior (older backends keep working).
Dependency-light and NEVER raises: surfacing diagnostics must not break the
error path it describes. Descriptors are advisory — clients fall back to
string sniffing when absent or partial.
"""
from __future__ import annotations
@@ -33,19 +24,16 @@ from typing import Any, Optional
logger = logging.getLogger(__name__)
# UI layers (wire values — stable contract with desktop/TUI clients).
LAYER_PROVIDER = "provider"
LAYER_ENDPOINT = "endpoint"
LAYER_STREAMING = "streaming"
LAYER_AUTH = "auth"
LAYER_BILLING = "billing"
LAYER_GATEWAY = "gateway"
LAYER_RUNTIME = "runtime"
LAYER_DISK = "disk"
# failure_reason (FailoverReason.value) → UI layer. Reasons not listed fall
# back to LAYER_PROVIDER: every FailoverReason is produced by classifying a
# provider API call, so "the provider call failed" is the honest default.
# failure_reason → UI layer. Unlisted reasons fall back to LAYER_PROVIDER:
# every FailoverReason comes from classifying a provider call.
_REASON_TO_LAYER = {
"auth": LAYER_AUTH,
"auth_permanent": LAYER_AUTH,
@@ -53,75 +41,34 @@ _REASON_TO_LAYER = {
"billing_unverified": LAYER_BILLING,
}
# Transport-ish reasons: the failure is between us and the base_url, not a
# verdict the provider returned. On a custom/local endpoint these point at
# the user's endpoint config, so they surface as LAYER_ENDPOINT there.
_TRANSPORT_REASONS = {
"timeout",
"ssl_cert_verification",
}
# Failures between us and the base_url (not a provider verdict); on a
# custom/local endpoint they point at the user's endpoint config.
_TRANSPORT_REASONS = {"timeout", "ssl_cert_verification"}
# Reasons that are deterministic for the request — a bare "Retry" repeats the
# same failure, so clients shouldn't lead with it. Fallback only: results
# from current backends carry the classifier's own verdict in
# ``failure_retryable`` and never consult this set. Kept in sync with
# ``classify_api_error``'s retryable=False verdicts.
# Deterministic for the request — a bare "Retry" repeats the failure. Fallback
# only: current backends stamp the classifier's verdict in ``failure_retryable``.
# Kept in sync with ``classify_api_error``'s retryable=False verdicts.
_NON_RETRYABLE_REASONS = {
"auth",
"auth_permanent",
"billing",
"billing_unverified",
"content_policy_blocked",
"provider_policy_blocked",
"model_not_found",
"format_error",
"ssl_cert_verification",
"auth", "auth_permanent", "billing", "billing_unverified", "content_policy_blocked",
"provider_policy_blocked", "model_not_found", "format_error", "ssl_cert_verification",
}
# Providers whose base_url is user-supplied rather than a known vendor —
# transport failures against these are endpoint-config problems.
_CUSTOM_ENDPOINT_PROVIDERS = {
"custom",
"local",
"llama.cpp",
"llamacpp",
"ollama",
"lmstudio",
"vllm",
}
# Providers whose base_url is user-supplied rather than a known vendor.
_CUSTOM_ENDPOINT_PROVIDERS = {"custom", "local", "llama.cpp", "llamacpp", "ollama", "lmstudio", "vllm"}
# Message fragments that mark a mid-stream connection drop. Deliberately
# narrow: these strings come from our own retry-exhaustion summaries and the
# OpenAI SDK's stream-abort errors.
# Mid-stream drop markers. Deliberately narrow: our own retry-exhaustion
# summaries plus the OpenAI SDK's stream-abort errors.
_STREAM_DROP_FRAGMENTS = (
"stream connection",
"peer closed connection",
"incomplete chunked read",
"connection broken",
"stream ended prematurely",
"sse",
"mid-stream",
"stream connection", "peer closed connection", "incomplete chunked read",
"connection broken", "stream ended prematurely", "sse", "mid-stream",
)
# Exception modules that indicate the failure came from an API/transport call
# (vs. a bug in our own dispatcher code, which is a gateway-layer failure).
# Covers every SDK family our provider adapters raise from: OpenAI-compatible
# (openai/httpx/httpcore), Anthropic, Bedrock (botocore/boto3), Google
# (google.*/grpc), plus raw transports (requests/aiohttp/ssl/socket/urllib).
# Exception top-level modules that mean "API/transport call failed" (vs. a bug
# in our dispatcher = gateway layer): every SDK family our adapters raise from
# plus raw transports.
_API_EXC_MODULE_PREFIXES = (
"openai",
"httpx",
"httpcore",
"anthropic",
"botocore",
"boto3",
"google",
"grpc",
"requests",
"aiohttp",
"ssl",
"socket",
"urllib",
"openai", "httpx", "httpcore", "anthropic", "botocore", "boto3", "google",
"grpc", "requests", "aiohttp", "ssl", "socket", "urllib",
)
@@ -135,17 +82,10 @@ def _looks_like_stream_drop(message: str) -> bool:
return any(fragment in msg for fragment in _STREAM_DROP_FRAGMENTS)
def _surface(
layer: str,
code: str,
retryable: bool,
provider: str = "",
model: str = "",
) -> dict:
def _surface(layer: str, code: str, retryable: bool, provider: str = "", model: str = "") -> dict:
out = {"layer": layer, "code": code, "retryable": bool(retryable)}
# The failing session's identity, captured at classification time so
# clients report the model/provider that actually failed — not whatever
# the foreground composer points at when a button is clicked later.
# Identity captured at classification time, so clients report the session
# that actually failed — not whatever the composer points at later.
if provider:
out["provider"] = provider
if model:
@@ -153,14 +93,20 @@ def _surface(
return out
def build_error_surface_from_result(
result: Any, provider: str = "", model: str = ""
) -> Optional[dict]:
def _disk_full(candidate: Any) -> bool:
try:
from hermes_state import is_disk_full_error
return bool(is_disk_full_error(candidate))
except Exception: # pragma: no cover - defensive import guard
return False
def build_error_surface_from_result(result: Any, provider: str = "", model: str = "") -> Optional[dict]:
"""Descriptor for a returned-error turn result (``failed=True`` dicts).
Reads the ``failure_reason`` the conversation loop already stamps
(a ``FailoverReason.value``) plus the error text, and maps them onto a
UI layer. Returns None when the result carries no failure signal.
Uses the stamped ``failure_reason`` plus error text. None when the result
carries no failure signal.
"""
try:
if not isinstance(result, dict):
@@ -171,14 +117,9 @@ def build_error_surface_from_result(
return None
# Disk-full wins outright: the fix (free space) is unrelated to the
# provider stack, and hermes_state owns the pattern list.
try:
from hermes_state import is_disk_full_error
if error_text and is_disk_full_error(error_text):
return _surface(LAYER_DISK, "disk_full", False, provider, model)
except Exception: # pragma: no cover - defensive import guard
pass
# provider stack; hermes_state owns the pattern list.
if error_text and _disk_full(error_text):
return _surface(LAYER_DISK, "disk_full", False, provider, model)
if result.get("billing_block") or reason in ("billing", "billing_unverified"):
return _surface(LAYER_BILLING, reason or "billing", False, provider, model)
@@ -197,9 +138,8 @@ def build_error_surface_from_result(
layer = LAYER_STREAMING
else:
layer = LAYER_PROVIDER
# Prefer the classifier's own retry verdict when the result carries it
# (conversation_loop stamps ``failure_retryable`` next to
# ``failure_reason``); the reason-set fallback covers older results.
# Prefer the classifier's own verdict (``failure_retryable``); the
# reason-set fallback covers older results.
retryable = result.get("failure_retryable")
if not isinstance(retryable, bool):
retryable = reason not in _NON_RETRYABLE_REASONS
@@ -209,31 +149,20 @@ def build_error_surface_from_result(
return None
def build_error_surface_from_exception(
exc: BaseException, provider: str = "", model: str = ""
) -> Optional[dict]:
def build_error_surface_from_exception(exc: BaseException, provider: str = "", model: str = "") -> Optional[dict]:
"""Descriptor for an exception that escaped the turn dispatcher.
API/transport exceptions are classified through the real
``classify_api_error`` pipeline (same taxonomy as the retry loop);
anything else is a gateway-layer failure — a bug or environment problem
in our own dispatcher, not a provider verdict.
API/transport exceptions go through ``classify_api_error`` (same taxonomy
as the retry loop); anything else is a gateway-layer failure.
"""
try:
message = str(exc) or type(exc).__name__
try:
from hermes_state import is_disk_full_error
if is_disk_full_error(exc):
return _surface(LAYER_DISK, "disk_full", False, provider, model)
except Exception: # pragma: no cover - defensive import guard
pass
if _disk_full(exc):
return _surface(LAYER_DISK, "disk_full", False, provider, model)
exc_module = type(exc).__module__ or ""
api_like = exc_module.split(".")[0] in _API_EXC_MODULE_PREFIXES or hasattr(
exc, "status_code"
)
api_like = exc_module.split(".")[0] in _API_EXC_MODULE_PREFIXES or hasattr(exc, "status_code")
if not api_like or not isinstance(exc, Exception):
return _surface(LAYER_GATEWAY, type(exc).__name__, True, provider, model)
@@ -241,15 +170,8 @@ def build_error_surface_from_exception(
from agent.error_classifier import classify_api_error
classified = classify_api_error(exc, provider=provider, model=model)
reason = classified.reason.value
synthetic = {
"error": classified.message or message,
"failure_reason": reason,
}
surface = build_error_surface_from_result(
synthetic, provider=provider, model=model
)
synthetic = {"error": classified.message or message, "failure_reason": classified.reason.value}
surface = build_error_surface_from_result(synthetic, provider=provider, model=model)
if surface is not None:
surface["retryable"] = bool(classified.retryable)
return surface
+30 -48
View File
@@ -1,57 +1,42 @@
"""Cheap content-sanity checks for the truncated-response continuation path.
Issue #86581: a model in a degenerate repetition loop can spend its ENTIRE
output budget echoing one fragment. The ``finish_reason=length``
continuation path in ``conversation_loop.py`` would then retry with a
"continue, don't repeat" nudge — stitching a pathological fragment into the
final response with no content-sanity check. In the incident behind #86581
a single turn produced a 60,698-char response delivered as 31 Discord
messages.
A model in a degenerate repetition loop can spend its ENTIRE output budget
echoing one fragment; the ``finish_reason=length`` continuation path would
then stitch that fragment into the final response with a "continue" nudge
(one incident: a 60k-char turn delivered as 31 Discord messages). These
helpers detect repetition-dominated fragments BEFORE the nudge so the turn
aborts with a clear error (like the ``_thinking_exhausted`` guard).
These helpers detect repetition-dominated fragments BEFORE the continuation
nudge is appended so the turn can abort with a clear user-facing error
(mirroring the existing ``_thinking_exhausted`` guard) instead of flooding.
The detection is deliberately conservative: only LONG verbatim repeats
(60+ chars) whose occurrences cover a majority of the fragment trip the
guard, so ordinary truncated responses (a sentence cut mid-word, a heading
repeated, code with similar-looking lines) are never blocked.
Deliberately conservative: only LONG verbatim repeats (60+ chars) covering a
majority of the fragment trip the guard, so ordinary truncations (a sentence
cut mid-word, a repeated heading, similar code lines) are never blocked.
"""
from __future__ import annotations
import math
# A fragment must be at least this long before the repetition check runs at
# all. Short truncations (a sentence cut mid-word) can trivially contain
# repeated tokens and are legitimately continued.
# Below this length the check doesn't run: short truncations trivially
# contain repeated tokens and are legitimately continued.
MIN_FRAGMENT_LENGTH = 400
# Length of the exact-repeat window. A verbatim repeat of this many chars
# is far beyond ordinary phrasing reuse (citations, headings, similar code).
# Exact-repeat window; a verbatim repeat this long is far beyond ordinary
# phrasing reuse (citations, headings, similar code).
_REPEAT_WINDOW = 60
# A window that repeats at least this many times is a repetition signal,
# even for short fragments.
# A window repeating at least this often is a signal even for short fragments.
_MIN_REPEAT_COUNT = 5
# A fragment is "repetition-dominated" when repeated windows account for at
# least this fraction of its characters.
# "Repetition-dominated" = repeated windows cover at least this fraction.
_DOMINANCE_RATIO = 0.5
def is_repetition_dominated(text: str) -> bool:
"""True when ``text`` is dominated by verbatim repeated fragments.
"""True when a single 60+ char substring recurs often enough to cover at
least half of ``text`` — the signature of a repetition loop, where
continuing would only stitch in more repeats.
A truncated response is "repetition-dominated" when a single 60+ char
substring appears often enough that its occurrences cover at least half
of the fragment. That shape is the signature of a model repetition
loop (issue #86581), and continuing such a fragment is pointless — the
continuation nudge would just stitch more repeated text into the final
response.
Returns False for non-string / empty / short inputs (fail-open: never
blocks a continuation the guard cannot confidently judge).
Fail-open: False for non-string / empty / short inputs.
"""
if not isinstance(text, str):
return False
@@ -59,17 +44,15 @@ def is_repetition_dominated(text: str) -> bool:
if n < MIN_FRAGMENT_LENGTH:
return False
# Fast path: one normalized line duplicated often enough to cover half
# the fragment (the most common echo shape — a repeated paragraph or
# sentence on its own line). Cheap, no big allocations.
# Fast path: one normalized line duplicated enough to cover half the
# fragment (the most common echo shape). Cheap, no big allocations.
if _line_repetition_dominated(text, n):
return True
# General path: fixed-size exact-repeat windows, sliding one char at a
# time. Catches repetition loops that do not align to line boundaries.
# General path: fixed-size windows sliding one char at a time, catching
# loops that don't align to line boundaries. A window must appear
# ``needed`` times to cover >= _DOMINANCE_RATIO (and >= _MIN_REPEAT_COUNT).
window = _REPEAT_WINDOW
# A window must appear this many times for its occurrences to cover
# >= DOMINANCE_RATIO of the fragment (and at least _MIN_REPEAT_COUNT).
needed = max(_MIN_REPEAT_COUNT, math.ceil(n * _DOMINANCE_RATIO / window))
counts: dict[str, int] = {}
for i in range(n - window + 1):
@@ -86,10 +69,9 @@ def _line_repetition_dominated(text: str, n: int) -> bool:
counts: dict[str, int] = {}
for line in text.splitlines():
norm = line.strip()
if not norm:
continue
counts[norm] = counts.get(norm, 0) + 1
for line, c in counts.items():
if c >= _MIN_REPEAT_COUNT and c * len(line) >= n * _DOMINANCE_RATIO:
return True
return False
if norm:
counts[norm] = counts.get(norm, 0) + 1
return any(
c >= _MIN_REPEAT_COUNT and c * len(line) >= n * _DOMINANCE_RATIO
for line, c in counts.items()
)
+50 -112
View File
@@ -1,8 +1,7 @@
"""Retry utilities — jittered backoff for decorrelated retries.
Replaces fixed exponential backoff with jittered delays to prevent
thundering-herd retry spikes when multiple sessions hit the same
rate-limited provider concurrently.
Jittered delays (vs. fixed exponential) prevent thundering-herd retry spikes
when many sessions hit the same rate-limited provider concurrently.
"""
import random
@@ -12,59 +11,43 @@ from datetime import datetime, timezone
from email.utils import parsedate_to_datetime
from typing import Any, Optional
# Monotonic counter for jitter seed uniqueness within the same process.
# Protected by a lock to avoid race conditions in concurrent retry paths
# (e.g. multiple gateway sessions retrying simultaneously).
# Monotonic counter for jitter-seed uniqueness within a process; locked
# because concurrent gateway sessions retry simultaneously.
_jitter_counter = 0
_jitter_lock = threading.Lock()
# Z.AI Coding Plan's GLM-5.2 endpoint often returns HTTP 429 code 1305
# ("The service may be temporarily overloaded...") for otherwise valid
# Hermes requests. Short retries tend to hammer the same overloaded window;
# after a few normal retries, progressively widen the wait window. Keep the
# cap interactive-friendly: a simple TUI message should fail visibly in minutes,
# not sit silent for 20+ minutes.
# Z.AI Coding Plan's GLM-5.2 endpoint often returns 429 code 1305 ("service
# may be temporarily overloaded"). Short retries hammer the same window, so
# after a few normal retries the wait widens progressively. Cap stays
# interactive-friendly: a TUI message should fail visibly in minutes.
_ZAI_CODING_OVERLOAD_LONG_BACKOFF = (30.0, 60.0, 90.0, 120.0)
# Number of initial short retries before the adaptive long-backoff tier kicks
# in. Shared by ``adaptive_rate_limit_backoff`` (which walks the long table
# starting at attempt ``short_attempts + 1``) and
# ``zai_coding_overload_retry_ceiling`` (which sizes the retry loop so every
# long-tier entry is reachable). Keeping it a single module constant prevents
# the two from silently desyncing if the short-retry count is ever tuned.
# Short retries before the long tier. Shared by ``adaptive_rate_limit_backoff``
# (walks the long table from attempt ``short_attempts + 1``) and
# ``zai_coding_overload_retry_ceiling`` (sizes the loop so every long entry is
# reachable) so the two cannot silently desync.
_ZAI_CODING_OVERLOAD_SHORT_ATTEMPTS = 3
def parse_retry_after_seconds(value_or_headers: Any) -> Optional[float]:
"""Parse a ``Retry-After`` value into non-negative seconds.
Accepts either a raw header value (numeric string / HTTP-date / number)
or a headers mapping, in which case the ``Retry-After`` key is looked up
case-insensitively (``.get`` on dict-like objects tries both common
casings; real HTTP header containers like httpx/requests are already
case-insensitive).
Returns:
Seconds as a ``float`` (negative deltas clamped to ``0.0``), or
``None`` when the header is absent or unparseable.
"""Parse a ``Retry-After`` value (numeric / HTTP-date / number) or a headers
mapping into seconds. Both casings are tried for plain dicts (real header
containers are already case-insensitive). Float clamped at 0.0, or None
when absent / unparseable.
"""
raw = value_or_headers
if raw is not None and not isinstance(raw, (str, int, float)):
# Looks like a headers mapping — pull the header out of it.
getter = getattr(raw, "get", None)
if callable(getter):
try:
value = getter("Retry-After")
if value is None:
value = getter("retry-after")
except Exception:
return None
raw = value
else:
if not callable(getter):
return None
if raw is None:
return None
if isinstance(raw, bool):
try:
value = getter("Retry-After")
if value is None:
value = getter("retry-after")
except Exception:
return None
raw = value
if raw is None or isinstance(raw, bool):
return None
if isinstance(raw, (int, float)):
return max(0.0, float(raw))
@@ -80,35 +63,17 @@ def parse_retry_after_seconds(value_or_headers: Any) -> Optional[float]:
when = parsedate_to_datetime(text)
except (TypeError, ValueError):
return None
if when is None:
if when is None: # older stdlib returns None instead of raising
return None
if when.tzinfo is None:
when = when.replace(tzinfo=timezone.utc)
return max(0.0, (when - datetime.now(timezone.utc)).total_seconds())
def jittered_backoff(
attempt: int,
*,
base_delay: float = 5.0,
max_delay: float = 120.0,
jitter_ratio: float = 0.5,
) -> float:
"""Compute a jittered exponential backoff delay.
Args:
attempt: 1-based retry attempt number.
base_delay: Base delay in seconds for attempt 1.
max_delay: Maximum delay cap in seconds.
jitter_ratio: Fraction of computed delay to use as random jitter
range. 0.5 means jitter is uniform in [0, 0.5 * delay].
Returns:
Delay in seconds: min(base * 2^(attempt-1), max_delay) + jitter.
The jitter decorrelates concurrent retries so multiple sessions
hitting the same provider don't all retry at the same instant.
"""
def jittered_backoff(attempt: int, *, base_delay: float = 5.0, max_delay: float = 120.0,
jitter_ratio: float = 0.5) -> float:
"""min(base * 2^(attempt-1), max_delay) + uniform jitter in
[0, jitter_ratio * delay]. ``attempt`` is 1-based."""
global _jitter_counter
with _jitter_lock:
_jitter_counter += 1
@@ -120,41 +85,25 @@ def jittered_backoff(
else:
delay = min(base_delay * (2 ** exponent), max_delay)
# Seed from time + counter for decorrelation even with coarse clocks.
# Seed from time + counter so coarse clocks still decorrelate.
seed = (time.time_ns() ^ (tick * 0x9E3779B9)) & 0xFFFFFFFF
rng = random.Random(seed)
jitter = rng.uniform(0, jitter_ratio * delay)
return delay + jitter
return delay + random.Random(seed).uniform(0, jitter_ratio * delay)
def _error_text(error: Any) -> str:
"""Best-effort flattened provider error text for retry classification."""
parts = [
error,
getattr(error, "message", None),
getattr(error, "body", None),
getattr(error, "response", None),
]
parts = [error, getattr(error, "message", None), getattr(error, "body", None), getattr(error, "response", None)]
return " ".join(str(part) for part in parts if part is not None).lower()
def is_zai_coding_overload_error(*, base_url: str | None, model: str | None, error: Any) -> bool:
"""Return True for Z.AI Coding Plan transient overload 429s.
The coding-plan endpoint reports overload as HTTP 429 with body code 1305
and message "The service may be temporarily overloaded...". Treat only
that narrow shape specially so ordinary quota/billing 429s still fail fast
through the existing classifier.
"""
base = (base_url or "").lower()
model_name = (model or "").lower()
status = getattr(error, "status_code", None)
"""True only for the narrow Z.AI Coding Plan overload shape (429 + code
1305 / "temporarily overloaded"), so ordinary quota 429s still fail fast."""
text = _error_text(error)
return (
status == 429
and "api.z.ai/api/coding/paas/v4" in base
and "glm-5.2" in model_name
getattr(error, "status_code", None) == 429
and "api.z.ai/api/coding/paas/v4" in (base_url or "").lower()
and "glm-5.2" in (model or "").lower()
and ("1305" in text or "temporarily overloaded" in text)
)
@@ -168,16 +117,11 @@ def adaptive_rate_limit_backoff(
default_wait: float,
short_attempts: int = _ZAI_CODING_OVERLOAD_SHORT_ATTEMPTS,
) -> tuple[float, str | None]:
"""Provider-aware rate-limit backoff.
"""Provider-aware rate-limit backoff → ``(wait_seconds, reason_label)``.
For most providers this returns ``default_wait`` unchanged. For Z.AI
Coding Plan GLM-5.2 overloads, keep the first ``short_attempts`` retries on
the normal short exponential schedule, then switch to progressively longer
waits (30s → 60s → 90s → 120s, capped) plus light jitter.
``attempt`` is 1-based, matching the retry loop's logged attempt number.
Returns ``(wait_seconds, reason_label)`` where ``reason_label`` is suitable
for status/log decoration when a provider-specific policy fired.
Most providers get ``default_wait`` unchanged. Z.AI Coding GLM-5.2
overloads keep ``short_attempts`` short retries, then 30→60→90→120s
(capped) with light jitter. ``attempt`` is 1-based like the loop's log.
"""
if not is_zai_coding_overload_error(base_url=base_url, model=model, error=error):
return default_wait, None
@@ -186,23 +130,17 @@ def adaptive_rate_limit_backoff(
idx = min(attempt - short_attempts - 1, len(_ZAI_CODING_OVERLOAD_LONG_BACKOFF) - 1)
base_delay = _ZAI_CODING_OVERLOAD_LONG_BACKOFF[idx]
# A smaller jitter ratio keeps long waits readable while still avoiding
# synchronized retry storms across concurrent Hermes sessions.
return jittered_backoff(1, base_delay=base_delay, max_delay=base_delay, jitter_ratio=0.2), "zai_coding_overload_long"
# Smaller jitter keeps long waits readable while still de-synchronizing.
wait = jittered_backoff(1, base_delay=base_delay, max_delay=base_delay, jitter_ratio=0.2)
return wait, "zai_coding_overload_long"
def zai_coding_overload_retry_ceiling(short_attempts: int = _ZAI_CODING_OVERLOAD_SHORT_ATTEMPTS) -> int:
"""Retry-loop ceiling needed for the full Z.AI overload backoff schedule.
"""Retry-loop ceiling for the full Z.AI overload schedule.
The adaptive policy runs ``short_attempts`` short retries, then walks the
long-backoff table one entry per subsequent attempt. The retry loop gives
up as soon as ``retry_count >= ceiling`` — and that check runs *before* the
attempt's backoff is computed — so the ceiling must sit one past the final
long-backoff entry for every long tier to actually execute.
With the default ``api_max_retries`` (3) equal to ``short_attempts`` (3),
the loop always gave up before reaching the long tier, leaving the whole
long-backoff schedule as dead code. Callers extend the ceiling to this
value for Z.AI Coding overload 429s so the 30/60/90/120s waits run.
The loop gives up when ``retry_count >= ceiling`` *before* computing the
attempt's backoff, so the ceiling must sit one past the last long entry
or the long tier never runs (the default ``api_max_retries`` of 3 equals
``short_attempts``).
"""
return short_attempts + len(_ZAI_CODING_OVERLOAD_LONG_BACKOFF) + 1
+1 -18
View File
@@ -20,10 +20,7 @@ import time
import httpx
import pytest
from agent.bounded_response import (
read_error_body_or_default,
read_streaming_error_body,
)
from agent.bounded_response import read_streaming_error_body
class _ThreadingServer(socketserver.ThreadingTCPServer):
@@ -114,17 +111,3 @@ def test_oversize_body_is_capped(server_base, client):
assert 0 < len(text) <= 64 * 1024
# Capping must return promptly, not after draining the whole body.
assert elapsed < 9.0
def test_or_default_returns_text_when_present(server_base, client):
with client.stream("POST", server_base + "/normal") as response:
result = read_error_body_or_default(response)
assert result is not None and "RESOURCE_EXHAUSTED" in result