feat(agent): per-provider session_affinity_header (#104449)
This commit is contained in:
@@ -17,7 +17,7 @@ so the header cannot drift per code path.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
OPENCODE_SESSION_HEADER = "x-opencode-session"
|
||||
|
||||
@@ -43,6 +43,19 @@ def is_opencode_target(provider: Optional[str], base_url: Optional[str]) -> bool
|
||||
return False
|
||||
|
||||
|
||||
def resolve_affinity_key(session_id: Optional[str] = None) -> str:
|
||||
"""Return the normalized rotation-stable conversation affinity key."""
|
||||
try:
|
||||
from agent.portal_tags import get_affinity_scope, get_conversation_context
|
||||
from agent.transports.codex import _cache_scope_from_session_id
|
||||
|
||||
return _cache_scope_from_session_id(
|
||||
get_affinity_scope() or get_conversation_context() or session_id
|
||||
)
|
||||
except Exception:
|
||||
return str(session_id or "")
|
||||
|
||||
|
||||
def opencode_session_headers(
|
||||
provider: Optional[str],
|
||||
base_url: Optional[str],
|
||||
@@ -51,42 +64,50 @@ def opencode_session_headers(
|
||||
"""Return ``{"x-opencode-session": <key>}`` for OpenCode targets, else ``{}``."""
|
||||
if not is_opencode_target(provider, base_url):
|
||||
return {}
|
||||
try:
|
||||
from agent.portal_tags import get_affinity_scope, get_conversation_context
|
||||
from agent.transports.codex import _cache_scope_from_session_id
|
||||
|
||||
key = _cache_scope_from_session_id(
|
||||
# Top-level session_id → OpenRouter's sticky routing key. Per their prompt-caching docs it is
|
||||
# used directly as the routing key instead of hashing the opening messages, and it activates
|
||||
# stickiness on the first successful request rather than only after a cache hit. Resolve it from
|
||||
# the declared routing scope first (set only by a host that names its own conversation, #96811),
|
||||
# then the ambient conversation contextvar, with the explicit argument as fallback. The gap this
|
||||
# closes is the auxiliary call sites — compression, title generation, vision, web_extract,
|
||||
# session_search, MoA slots — which funnel through ``agent.auxiliary_client``. That module has
|
||||
# no session handle and passes no ``session_id``, so those calls sent NO sticky key at all and
|
||||
# each routed independently of the conversation it belonged to (#70820). Mirrors the Nous Portal
|
||||
# profile, which resolves the same way (f2f4df064d). The ambient value is the session-lineage
|
||||
# ROOT, so it also stays stable for installs that opt out of the default ``compression.in_place:
|
||||
# true`` and across delegate-subagent trees.
|
||||
get_affinity_scope() or get_conversation_context() or session_id
|
||||
)
|
||||
except Exception:
|
||||
key = str(session_id or "")
|
||||
key = resolve_affinity_key(session_id)
|
||||
return {OPENCODE_SESSION_HEADER: key} if key else {}
|
||||
|
||||
|
||||
def custom_provider_session_affinity_headers(
|
||||
provider: Optional[str],
|
||||
base_url: Optional[str],
|
||||
session_id: Optional[str] = None,
|
||||
custom_providers: Optional[List[Dict[str, Any]]] = None,
|
||||
) -> dict[str, str]:
|
||||
"""Return ``{<session_affinity_header>: <key>}`` for custom providers declaring one, else ``{}``."""
|
||||
try:
|
||||
from hermes_cli.config_providers import get_custom_provider_session_affinity_header
|
||||
|
||||
header = get_custom_provider_session_affinity_header(
|
||||
base_url=base_url,
|
||||
provider=provider,
|
||||
custom_providers=custom_providers,
|
||||
)
|
||||
if not header:
|
||||
return {}
|
||||
key = resolve_affinity_key(session_id)
|
||||
return {header: key} if key else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def merge_opencode_session_headers(
|
||||
kwargs: dict[str, Any],
|
||||
provider: Optional[str],
|
||||
base_url: Optional[str],
|
||||
session_id: Optional[str] = None,
|
||||
custom_providers: Optional[List[Dict[str, Any]]] = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Merge the affinity header into ``kwargs["extra_headers"]`` (in place).
|
||||
"""Merge OpenCode or custom provider affinity headers into ``kwargs["extra_headers"]`` (in place).
|
||||
|
||||
Existing per-request headers win, so a caller-pinned value is preserved.
|
||||
Non-OpenCode targets are left untouched.
|
||||
Non-OpenCode and non-affinity targets are left untouched.
|
||||
"""
|
||||
headers = opencode_session_headers(provider, base_url, session_id)
|
||||
if not headers:
|
||||
headers = custom_provider_session_affinity_headers(
|
||||
provider, base_url, session_id, custom_providers=custom_providers
|
||||
)
|
||||
if headers:
|
||||
existing = kwargs.get("extra_headers")
|
||||
merged = dict(existing) if isinstance(existing, dict) else {}
|
||||
@@ -94,3 +115,7 @@ def merge_opencode_session_headers(
|
||||
merged.setdefault(key, value)
|
||||
kwargs["extra_headers"] = merged
|
||||
return kwargs
|
||||
|
||||
|
||||
# Alias for callers naming the generalized capability
|
||||
merge_session_affinity_headers = merge_opencode_session_headers
|
||||
|
||||
@@ -710,6 +710,7 @@ from hermes_cli.config_providers import ( # noqa: E402,F401 (re-exported; call
|
||||
apply_custom_provider_tls_to_client_kwargs, coerce_provider_id, find_provider_entry,
|
||||
get_compatible_custom_providers, get_custom_provider_context_length,
|
||||
get_custom_provider_extra_headers, get_custom_provider_model_capability,
|
||||
get_custom_provider_session_affinity_header,
|
||||
get_custom_provider_tls_settings, is_provider_enabled, normalize_extra_headers,
|
||||
providers_dict_to_custom_providers, stringify_provider_map)
|
||||
# Back-compat re-exports — :mod:`hermes_cli.personality` owns personality/overlay semantics.
|
||||
|
||||
@@ -107,7 +107,8 @@ _CAMEL_ALIASES: Dict[str, str] = {
|
||||
"apiKeyEnv": "key_env", # OpenClaw-compatible + docs variant
|
||||
"defaultModel": "default_model",
|
||||
"contextLength": "context_length",
|
||||
"rateLimitDelay": "rate_limit_delay"}
|
||||
"rateLimitDelay": "rate_limit_delay",
|
||||
"sessionAffinityHeader": "session_affinity_header"}
|
||||
|
||||
|
||||
_KNOWN_PROVIDER_KEYS = {
|
||||
@@ -118,7 +119,7 @@ _KNOWN_PROVIDER_KEYS = {
|
||||
"api_mode", "transport", "model", "default_model", "models", "models_discovered",
|
||||
"context_length", "rate_limit_delay", "request_timeout_seconds", "stale_timeout_seconds",
|
||||
"discover_models", "extra_body", "extra_headers", "capabilities", "ssl_ca_cert", "ssl_verify",
|
||||
"catalog_provider"}
|
||||
"catalog_provider", "session_affinity_header"}
|
||||
|
||||
|
||||
def _pick_provider_base_url(entry: Dict[str, Any], provider_key: str) -> str:
|
||||
@@ -266,6 +267,7 @@ def _normalize_custom_provider_entry(
|
||||
|
||||
# Per-provider extra HTTP headers may carry credentials — never log them downstream.
|
||||
_put("extra_headers", normalize_extra_headers(entry.get("extra_headers")))
|
||||
_put("session_affinity_header", _stripped("session_affinity_header"))
|
||||
_put("ssl_ca_cert", _stripped("ssl_ca_cert"))
|
||||
|
||||
ssl_verify = entry.get("ssl_verify")
|
||||
@@ -288,7 +290,7 @@ def _custom_provider_entry_to_provider_config(
|
||||
for field in (
|
||||
"name", "api_key", "key_env", "key_cmd", "models", "models_discovered", "context_length",
|
||||
"rate_limit_delay", "discover_models", "extra_body", "extra_headers",
|
||||
"ssl_ca_cert", "ssl_verify", "catalog_provider"):
|
||||
"session_affinity_header", "ssl_ca_cert", "ssl_verify", "catalog_provider"):
|
||||
if field in normalized:
|
||||
provider_entry[field] = normalized[field]
|
||||
if "model" in normalized:
|
||||
@@ -488,6 +490,46 @@ def apply_custom_provider_extra_headers_to_client_kwargs(
|
||||
client_kwargs["default_headers"] = merged
|
||||
|
||||
|
||||
def get_custom_provider_session_affinity_header(
|
||||
base_url: Optional[str] = None,
|
||||
custom_providers: Optional[List[Dict[str, Any]]] = None,
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
provider: Optional[str] = None) -> Optional[str]:
|
||||
"""Return the declared ``session_affinity_header`` for a provider or route, or None."""
|
||||
from hermes_cli.config import get_compatible_custom_providers
|
||||
if custom_providers is None:
|
||||
try:
|
||||
custom_providers = get_compatible_custom_providers(config)
|
||||
except Exception:
|
||||
custom_providers = []
|
||||
if not isinstance(custom_providers, list):
|
||||
return None
|
||||
|
||||
want_p = str(provider or "").strip().lower()
|
||||
target_url = normalize_route_base_url(base_url) if base_url else ""
|
||||
|
||||
for entry in custom_providers:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
header = entry.get("session_affinity_header")
|
||||
if not (isinstance(header, str) and header.strip()):
|
||||
continue
|
||||
header_clean = header.strip()
|
||||
|
||||
if want_p:
|
||||
entry_p = str(entry.get("provider_key") or "").strip().lower()
|
||||
entry_n = str(entry.get("name") or "").strip().lower()
|
||||
if want_p == entry_p or want_p == entry_n:
|
||||
return header_clean
|
||||
|
||||
if target_url:
|
||||
entry_url = normalize_route_base_url(entry.get("base_url"))
|
||||
if entry_url and entry_url == target_url:
|
||||
return header_clean
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_custom_provider_context_length(
|
||||
model: str,
|
||||
base_url: str,
|
||||
|
||||
172
tests/agent/test_session_affinity_header.py
Normal file
172
tests/agent/test_session_affinity_header.py
Normal file
@@ -0,0 +1,172 @@
|
||||
"""Tests for per-provider session_affinity_header in agent and auxiliary requests.
|
||||
|
||||
Closes #104449: custom gateways (such as LiteLLM x-litellm-session-id) receive
|
||||
the turn's affinity scope to pin requests to warm prompt caches.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
import pytest
|
||||
|
||||
from agent import auxiliary_client as aux
|
||||
from agent.chat_completion_helpers import build_api_kwargs
|
||||
from hermes_cli.config_providers import (
|
||||
_normalize_custom_provider_entry,
|
||||
get_custom_provider_session_affinity_header,
|
||||
)
|
||||
from run_agent import AIAgent
|
||||
|
||||
_MSGS = [{"role": "user", "content": "hello"}]
|
||||
|
||||
|
||||
def _agent(provider, model, base_url, api_mode=None, session_id="sess-affinity-test-1"):
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
provider=provider,
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
session_id=session_id,
|
||||
)
|
||||
if api_mode:
|
||||
agent.api_mode = api_mode
|
||||
agent._transport = None
|
||||
agent._anthropic_base_url = base_url
|
||||
return agent
|
||||
|
||||
|
||||
def test_normalization_and_camel_case_alias():
|
||||
entry = _normalize_custom_provider_entry(
|
||||
{
|
||||
"name": "litellm-lan",
|
||||
"base_url": "http://localhost:4000/v1",
|
||||
"sessionAffinityHeader": "x-litellm-session-id",
|
||||
}
|
||||
)
|
||||
assert entry is not None
|
||||
assert entry.get("session_affinity_header") == "x-litellm-session-id"
|
||||
|
||||
|
||||
def test_lookup_by_provider_or_base_url():
|
||||
custom_providers = [
|
||||
{
|
||||
"name": "litellm-lan",
|
||||
"provider_key": "litellm-lan",
|
||||
"base_url": "http://localhost:4000/v1",
|
||||
"session_affinity_header": "x-litellm-session-id",
|
||||
}
|
||||
]
|
||||
# Matched by provider name
|
||||
assert (
|
||||
get_custom_provider_session_affinity_header(
|
||||
provider="litellm-lan", custom_providers=custom_providers
|
||||
)
|
||||
== "x-litellm-session-id"
|
||||
)
|
||||
# Matched by base_url
|
||||
assert (
|
||||
get_custom_provider_session_affinity_header(
|
||||
base_url="http://localhost:4000/v1", custom_providers=custom_providers
|
||||
)
|
||||
== "x-litellm-session-id"
|
||||
)
|
||||
# Unmatched
|
||||
assert (
|
||||
get_custom_provider_session_affinity_header(
|
||||
provider="unconfigured",
|
||||
base_url="http://other:8000/v1",
|
||||
custom_providers=custom_providers,
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_main_turn_sends_session_affinity_header():
|
||||
mock_providers = [
|
||||
{
|
||||
"name": "litellm-lan",
|
||||
"provider_key": "litellm-lan",
|
||||
"base_url": "http://localhost:4000/v1",
|
||||
"session_affinity_header": "x-litellm-session-id",
|
||||
}
|
||||
]
|
||||
agent = _agent("litellm-lan", "claude-3-5-sonnet", "http://localhost:4000/v1")
|
||||
with patch(
|
||||
"hermes_cli.config.get_compatible_custom_providers",
|
||||
return_value=mock_providers,
|
||||
):
|
||||
kwargs = build_api_kwargs(agent, _MSGS)
|
||||
extra = kwargs.get("extra_headers") or {}
|
||||
assert extra.get("x-litellm-session-id") == "sess-affinity-test-1"
|
||||
|
||||
|
||||
def test_auxiliary_call_inherits_session_affinity_header():
|
||||
mock_providers = [
|
||||
{
|
||||
"name": "litellm-lan",
|
||||
"provider_key": "litellm-lan",
|
||||
"base_url": "http://localhost:4000/v1",
|
||||
"session_affinity_header": "x-litellm-session-id",
|
||||
}
|
||||
]
|
||||
token = aux.set_runtime_main(
|
||||
"litellm-lan",
|
||||
"claude-3-5-sonnet",
|
||||
base_url="http://localhost:4000/v1",
|
||||
session_id="sess-affinity-test-1",
|
||||
)
|
||||
try:
|
||||
with patch(
|
||||
"hermes_cli.config.get_compatible_custom_providers",
|
||||
return_value=mock_providers,
|
||||
):
|
||||
kwargs = aux._build_call_kwargs(
|
||||
"litellm-lan",
|
||||
"claude-3-5-sonnet",
|
||||
_MSGS,
|
||||
base_url="http://localhost:4000/v1",
|
||||
)
|
||||
extra = kwargs.get("extra_headers") or {}
|
||||
assert extra.get("x-litellm-session-id") == "sess-affinity-test-1"
|
||||
finally:
|
||||
aux._RUNTIME_MAIN_CONTEXT.reset(token)
|
||||
|
||||
|
||||
def test_unconfigured_custom_provider_does_not_inject_header():
|
||||
mock_providers = [
|
||||
{
|
||||
"name": "plain-proxy",
|
||||
"base_url": "http://localhost:5000/v1",
|
||||
}
|
||||
]
|
||||
agent = _agent("plain-proxy", "model-a", "http://localhost:5000/v1")
|
||||
with patch(
|
||||
"hermes_cli.config.get_compatible_custom_providers",
|
||||
return_value=mock_providers,
|
||||
):
|
||||
kwargs = build_api_kwargs(agent, _MSGS)
|
||||
extra = kwargs.get("extra_headers") or {}
|
||||
assert "x-litellm-session-id" not in extra
|
||||
assert "x-opencode-session" not in extra
|
||||
|
||||
|
||||
def test_caller_pinned_header_wins():
|
||||
mock_providers = [
|
||||
{
|
||||
"name": "litellm-lan",
|
||||
"base_url": "http://localhost:4000/v1",
|
||||
"session_affinity_header": "x-litellm-session-id",
|
||||
}
|
||||
]
|
||||
agent = _agent("litellm-lan", "claude-3-5-sonnet", "http://localhost:4000/v1")
|
||||
agent.request_overrides = {"extra_headers": {"x-litellm-session-id": "pinned-by-caller"}}
|
||||
with patch(
|
||||
"hermes_cli.config.get_compatible_custom_providers",
|
||||
return_value=mock_providers,
|
||||
):
|
||||
kwargs = build_api_kwargs(agent, _MSGS)
|
||||
extra = kwargs.get("extra_headers") or {}
|
||||
assert extra.get("x-litellm-session-id") == "pinned-by-caller"
|
||||
Reference in New Issue
Block a user