diff --git a/agent/azure_identity_adapter.py b/agent/azure_identity_adapter.py index dd0f62ab97..a6cf20ed58 100644 --- a/agent/azure_identity_adapter.py +++ b/agent/azure_identity_adapter.py @@ -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 ``. - - ``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", ] diff --git a/agent/bounded_response.py b/agent/bounded_response.py index e5177bc8a2..4ee32dfcb0 100644 --- a/agent/bounded_response.py +++ b/agent/bounded_response.py @@ -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 diff --git a/agent/empty_response_guard.py b/agent/empty_response_guard.py index fbde4b58b0..9f163d1ed7 100644 --- a/agent/empty_response_guard.py +++ b/agent/empty_response_guard.py @@ -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 diff --git a/agent/error_surface.py b/agent/error_surface.py index 4016a9fb3c..0e041b4791 100644 --- a/agent/error_surface.py +++ b/agent/error_surface.py @@ -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": , "code": , "retryable": } -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 diff --git a/agent/repetition_guard.py b/agent/repetition_guard.py index 6a1d466e92..c84f5dec37 100644 --- a/agent/repetition_guard.py +++ b/agent/repetition_guard.py @@ -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() + ) diff --git a/agent/retry_utils.py b/agent/retry_utils.py index 58c6231b60..a0c5a62bb2 100644 --- a/agent/retry_utils.py +++ b/agent/retry_utils.py @@ -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 diff --git a/tests/agent/test_bounded_response.py b/tests/agent/test_bounded_response.py index db5b30c8e2..46a18b9ee5 100644 --- a/tests/agent/test_bounded_response.py +++ b/tests/agent/test_bounded_response.py @@ -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