diff --git a/CHANGELOG.md b/CHANGELOG.md index c1614a0..60c16f8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,18 @@ Versioning follows [docs/versioning.md](docs/versioning.md). ## [0.11.1] - 2026-09-24 +### Added + +- `bind_session_id(..., header_names=...)` and + `AgentLoopHooks.session_header_names` choose which request header(s) carry + the session id (default unchanged: `x-upstream-session-id`). Every listed + header gets the same value, so a host can keep the legacy header and add a + gateway's own affinity header (e.g. `X-Llmhub-Session`, which pins a session + to one upstream key in an account pool so prompt caches hit). An empty tuple + binds nothing; a bare `str` is rejected with `TypeError`. Ignored when + `AgentLoopHooks.bind_session` is supplied. `DEFAULT_SESSION_HEADER_NAMES` is + exported from `agent_core.runtime.loop`. + ### Fixed - `build_protocol_client` now forwards `cfg["default_headers"]` to diff --git a/agent_core/runtime/loop/__init__.py b/agent_core/runtime/loop/__init__.py index ce47795..adace2e 100644 --- a/agent_core/runtime/loop/__init__.py +++ b/agent_core/runtime/loop/__init__.py @@ -20,6 +20,7 @@ check_budget_exhausted, ) from agent_core.runtime.loop.llm_client import ( + DEFAULT_SESSION_HEADER_NAMES, RUNAWAY_STATE_KEY, TRUNCATION_CONTINUATION_GUIDANCE, LLMCallExhausted, @@ -79,6 +80,7 @@ "COMPACTION_TRIGGER_RATIO", "DEFAULT_DUPLICATE_THRESHOLDS", "DEFAULT_KEY_ALIASES", + "DEFAULT_SESSION_HEADER_NAMES", "DEFAULT_TRIGGER_RATIO", "LITERAL_CONTENT_KEYS", "RUNAWAY_STATE_KEY", diff --git a/agent_core/runtime/loop/_bind.py b/agent_core/runtime/loop/_bind.py index 17d6aa9..ede08e7 100644 --- a/agent_core/runtime/loop/_bind.py +++ b/agent_core/runtime/loop/_bind.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging -from collections.abc import Callable +from collections.abc import Callable, Sequence from dataclasses import dataclass, replace from typing import Any @@ -145,13 +145,23 @@ def bind_tools(llm: Any, tools: list[Any]) -> Any: return replace(_ensure_bound(llm), tools=schemas) +DEFAULT_SESSION_HEADER_NAMES: tuple[str, ...] = ("x-upstream-session-id",) + + def bind_session_id( llm: Any, task_id: str, *, sticky_session_enabled: Callable[[], bool] | None = None, + header_names: Sequence[str] = DEFAULT_SESSION_HEADER_NAMES, ) -> Any: - """Attach ``x-upstream-session-id: `` to every LLM request. + """Attach ``
: `` to every LLM request, per ``header_names``. + + ``header_names`` defaults to ``x-upstream-session-id``. Gateways that key + affinity on a different header (e.g. an account-pool proxy pinning a + session to one upstream key so prompt caches hit) pass their own name(s); + every listed header carries the same ``task_id``. An empty sequence binds + nothing. Pinning it at client-construction time is the obvious approach, but the LLM is per-profile-cached (one client shared across tasks), so the @@ -171,13 +181,16 @@ def bind_session_id( that attribution was retracted when the real cause turned out to be silent mid-stream stalls, now handled by the stall watchdog. """ - if not task_id or ( + if isinstance(header_names, str): + raise TypeError("header_names must be a sequence of names, not a str") + if not task_id or not header_names or ( sticky_session_enabled is not None and not sticky_session_enabled() ): return llm bound = _ensure_bound(llm) headers = dict(bound.extra_headers or {}) - headers["x-upstream-session-id"] = task_id + for name in header_names: + headers[name] = task_id return replace(bound, extra_headers=headers) diff --git a/agent_core/runtime/loop/agent_loop.py b/agent_core/runtime/loop/agent_loop.py index a685432..f6cb921 100644 --- a/agent_core/runtime/loop/agent_loop.py +++ b/agent_core/runtime/loop/agent_loop.py @@ -51,6 +51,7 @@ ) from agent_core.runtime.loop.image_attach import attach_images, evict_old_images from agent_core.runtime.loop.llm_client import ( + DEFAULT_SESSION_HEADER_NAMES, RUNAWAY_STATE_KEY, TRUNCATION_CONTINUATION_GUIDANCE, LLMCallExhausted, @@ -179,6 +180,12 @@ class AgentLoopHooks: # consult the flag a host binding is free to ignore. Supplying both is a # configuration error and is reported once at loop start. sticky_session_enabled: Callable[[], bool] | None = None + # Header(s) the built-in binding stamps with the session id. Ignored when + # ``bind_session`` is supplied, like ``sticky_session_enabled``. + # Keyword-only to preserve the existing positional hook arguments. + session_header_names: tuple[str, ...] = field( + default=DEFAULT_SESSION_HEADER_NAMES, kw_only=True, + ) bind_session: Callable[[LLMClient, str], LLMClient] | None = None wall_deadline_remaining: Callable[[], float | None] = _no_deadline chain_fallback_active: Callable[[], bool] = _false @@ -300,6 +307,7 @@ async def run_agent_loop( llm, llm_session_id, sticky_session_enabled=runtime.sticky_session_enabled, + header_names=runtime.session_header_names, ) llm_with_tools = bind_tools(llm_with_session, list(tools)) diff --git a/agent_core/runtime/loop/llm_client.py b/agent_core/runtime/loop/llm_client.py index 3bf6489..1a15f3f 100644 --- a/agent_core/runtime/loop/llm_client.py +++ b/agent_core/runtime/loop/llm_client.py @@ -16,14 +16,15 @@ LLMStreamStalled, ) from agent_core.runtime.loop._bind import ( - _ensure_bound as _ensure_bound, -) -from agent_core.runtime.loop._bind import ( + DEFAULT_SESSION_HEADER_NAMES, bind_max_tokens, bind_session_id, bind_temperature, bind_tools, ) +from agent_core.runtime.loop._bind import ( + _ensure_bound as _ensure_bound, +) from agent_core.runtime.loop._call import call_llm from agent_core.runtime.loop._response import ( extract_final_content, @@ -45,6 +46,7 @@ ) __all__ = [ + "DEFAULT_SESSION_HEADER_NAMES", "RUNAWAY_STATE_KEY", "TRUNCATION_CONTINUATION_GUIDANCE", "LLMCallExhausted", diff --git a/tests/test_agent_loop_engine.py b/tests/test_agent_loop_engine.py index 017f4c8..e656cbc 100644 --- a/tests/test_agent_loop_engine.py +++ b/tests/test_agent_loop_engine.py @@ -359,10 +359,21 @@ async def on_context_compacted(self, ctx) -> None: @pytest.mark.asyncio -async def test_host_can_override_session_binding() -> None: +@pytest.mark.parametrize("positional", [False, True]) +async def test_host_can_override_session_binding(positional: bool) -> None: bound: list[str] = [] llm = SequenceLLM([LLMResponse(content="done")]) + def bind_session(client, session_id): + bound.append(session_id) + return client + + hooks = ( + AgentLoopHooks(None, bind_session) + if positional + else AgentLoopHooks(bind_session=bind_session) + ) + result = await run_agent_loop( system_prompt="system", user_message="start", @@ -375,17 +386,43 @@ async def test_host_can_override_session_binding() -> None: loop_policy=LoopPolicy(no_tool_behavior="stop"), max_llm_retries=1, ), - runtime_hooks=AgentLoopHooks( - bind_session=lambda client, session_id: ( - bound.append(session_id) or client - ) - ), + runtime_hooks=hooks, ) assert result.final_content == "done" assert bound == ["gateway-session"] + +@pytest.mark.asyncio +async def test_session_header_names_hook_reaches_the_provider_call() -> None: + seen: list[dict[str, str]] = [] + + class HeaderLLM(SequenceLLM): + async def chat(self, messages, **kwargs) -> LLMResponse: + seen.append(dict(kwargs.get("extra_headers") or {})) + return await super().chat(messages, **kwargs) + + await run_agent_loop( + system_prompt="system", + user_message="start", + llm=HeaderLLM([LLMResponse(content="done")]), + tools=[], + config=LoopConfig( + max_turns=1, + task_id="runtime-task", + llm_session_id="gateway-session", + loop_policy=LoopPolicy(no_tool_behavior="stop"), + max_llm_retries=1, + ), + runtime_hooks=AgentLoopHooks( + session_header_names=("x-upstream-session-id", "X-Llmhub-Session"), + ), + ) + + assert seen and seen[0]["X-Llmhub-Session"] == "gateway-session" + assert seen[0]["x-upstream-session-id"] == "gateway-session" + def _orphan_tool_call_ids(messages: list[dict[str, Any]]) -> set[str]: """Ids an assistant message announces that no tool message answers. diff --git a/tests/test_llm_runtime_hooks.py b/tests/test_llm_runtime_hooks.py index 84cedbc..306591f 100644 --- a/tests/test_llm_runtime_hooks.py +++ b/tests/test_llm_runtime_hooks.py @@ -83,3 +83,22 @@ async def observe(mode: str) -> str: "disabled", ] assert current_thinking_retry_override() is None + + +def test_bind_session_id_stamps_every_configured_header_name() -> None: + llm = FakeLLM() + + bound = bind_session_id( + llm, "task-1", header_names=("x-upstream-session-id", "X-Llmhub-Session") + ) + assert bound.extra_headers == { + "x-upstream-session-id": "task-1", + "X-Llmhub-Session": "task-1", + } + assert bind_session_id(llm, "task-1", header_names=()) is llm + + +def test_bind_session_id_rejects_a_bare_string_header_name() -> None: + # A str is a Sequence[str]; iterating it would stamp one header per letter. + with pytest.raises(TypeError): + bind_session_id(FakeLLM(), "task-1", header_names="X-Llmhub-Session")