feat(relay): pass API request metadata to Relay tracking
Signed-off-by: David Gardner <dagardner@nvidia.com>
This commit is contained in:
@@ -922,7 +922,14 @@ class RelaySessionCoordinator:
|
||||
return host.register_subagent(event, metadata=metadata)
|
||||
return host.ensure_session({"session_id": session_id}, metadata=metadata)
|
||||
|
||||
def begin_turn(self, lease: ConversationLease, *, turn_id: str, task_id: str) -> RelayTurnContext:
|
||||
def begin_turn(
|
||||
self,
|
||||
lease: ConversationLease,
|
||||
*,
|
||||
turn_id: str,
|
||||
task_id: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> RelayTurnContext:
|
||||
if lease.released:
|
||||
raise RuntimeError("Hermes Relay conversation lease is released")
|
||||
turn = RelayTurnContext(lease=lease, turn_id=turn_id, task_id=task_id)
|
||||
@@ -942,10 +949,17 @@ class RelaySessionCoordinator:
|
||||
if host is not None:
|
||||
# Rotation happens HERE: no live turn scope on the stack, so the session scope can close/reopen LIFO.
|
||||
_warn_on_error("segment rotation", self._maybe_rotate_segment, host, lease.session)
|
||||
turn_metadata = dict(metadata or {})
|
||||
turn_metadata.update(
|
||||
runtime_metadata(
|
||||
host.runtime_id,
|
||||
**{"hermes.execution_surface": lease.platform or "unknown"},
|
||||
)
|
||||
)
|
||||
turn.handle = _warn_on_error(
|
||||
"turn initialization", host.run_in_session, lease.session, host.relay.scope.push,
|
||||
TURN_SCOPE, host.relay.ScopeType.Function, handle=lease.session.handle, input={},
|
||||
metadata=runtime_metadata(host.runtime_id, **{"hermes.execution_surface": lease.platform or "unknown"}),
|
||||
metadata=turn_metadata,
|
||||
timeout=_SCOPE_OP_TIMEOUT,
|
||||
)
|
||||
turn._previous_turn = _CURRENT_TURN.get()
|
||||
|
||||
@@ -27,6 +27,7 @@ class TurnFacadeMixin:
|
||||
persist_user_display_metadata: Optional[Dict[str, Any]]=None,
|
||||
persist_user_platform_id: Optional[str]=None, moa_config: Optional[dict[str, Any]]=None,
|
||||
turn_author: Optional[Dict[str, Any]] = None,
|
||||
relay_metadata: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Forwarder — see ``agent.conversation_loop.run_conversation``."""
|
||||
# A review shares this session_id for cache parity: fence review startup or interrupt
|
||||
@@ -94,8 +95,14 @@ class TurnFacadeMixin:
|
||||
parent_session_id=relay_parent_session_id,
|
||||
model=str(getattr(self, "model", None) or ""),
|
||||
)
|
||||
relay_turn_kwargs: Dict[str, Any] = {
|
||||
"turn_id": relay_turn_id,
|
||||
"task_id": effective_task_id,
|
||||
}
|
||||
if relay_metadata:
|
||||
relay_turn_kwargs["metadata"] = relay_metadata
|
||||
relay_turn = relay_runtime.SESSION_COORDINATOR.begin_turn(
|
||||
relay_lease, turn_id=relay_turn_id, task_id=effective_task_id
|
||||
relay_lease, **relay_turn_kwargs
|
||||
)
|
||||
# Minimal relay-runtime shims may lack the opt-out flag: default enabled.
|
||||
if getattr(relay_turn, "relay_enabled", True):
|
||||
|
||||
@@ -345,6 +345,16 @@ def _request_agent_overrides(
|
||||
return overrides
|
||||
|
||||
|
||||
def _request_relay_metadata(body: Any) -> Dict[str, Any]:
|
||||
"""Extract Relay metadata from an OpenAI request body."""
|
||||
if not isinstance(body, dict):
|
||||
return {}
|
||||
metadata = body.get("metadata")
|
||||
if not isinstance(metadata, dict):
|
||||
return {}
|
||||
return dict(metadata)
|
||||
|
||||
|
||||
def _is_compressed_summary_message(message: Any) -> bool:
|
||||
"""Recognize every compaction carrier shape via the compressor's own classifier
|
||||
(SessionDB drops the in-process marker; a prefix scan misses merge-into-tail carriers)."""
|
||||
@@ -3641,7 +3651,8 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter):
|
||||
route: Optional[Dict[str, Any]] = None, session_model: Optional[str] = None,
|
||||
requested_runtime: Optional[Dict[str, Any]] = None, route_source: str = "global",
|
||||
confirmed_runtime_lock: bool = False, bind_declared_conversation: bool = False,
|
||||
session_history_delivery: str = "", turn_author: Optional[Dict[str, Any]] = None) -> tuple:
|
||||
session_history_delivery: str = "", turn_author: Optional[Dict[str, Any]] = None,
|
||||
relay_metadata: Optional[Dict[str, Any]] = None) -> tuple:
|
||||
"""Create an agent and run one turn in a thread executor -> ``(result, usage)``.
|
||||
``agent_ref[0]`` receives the agent so SSE writers can interrupt it; ``active_run_id``
|
||||
registers it in ``_active_run_agents``. Under a confirmed model lock the actual
|
||||
@@ -3695,9 +3706,15 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter):
|
||||
self._shutdown_interruptible_agents[id(agent)] = agent
|
||||
# Passed only when set: a human turn keeps today's call shape.
|
||||
author_kwargs = {"turn_author": turn_author} if turn_author is not None else {}
|
||||
result = agent.run_conversation(
|
||||
user_message=user_message, conversation_history=conversation_history,
|
||||
task_id=effective_task_id, **author_kwargs)
|
||||
conversation_kwargs = dict(
|
||||
user_message=user_message,
|
||||
conversation_history=conversation_history,
|
||||
task_id=effective_task_id,
|
||||
**author_kwargs,
|
||||
)
|
||||
if relay_metadata:
|
||||
conversation_kwargs["relay_metadata"] = relay_metadata
|
||||
result = agent.run_conversation(**conversation_kwargs)
|
||||
return self._finish_turn_result(
|
||||
agent, result, session_id, route=route, requested_runtime=requested_runtime,
|
||||
route_source=route_source, confirmed_runtime_lock=confirmed_runtime_lock)
|
||||
|
||||
@@ -421,6 +421,8 @@ class OpenAICompatRoutesMixin:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
return _error_response("Invalid JSON in request body", 400)
|
||||
from gateway.platforms.api_server import _request_relay_metadata
|
||||
relay_metadata = _request_relay_metadata(body)
|
||||
messages = body.get("messages")
|
||||
if not messages or not isinstance(messages, list):
|
||||
return _invalid_request("Missing or invalid 'messages' field")
|
||||
@@ -502,6 +504,7 @@ class OpenAICompatRoutesMixin:
|
||||
user_message=user_message, conversation_history=history,
|
||||
ephemeral_system_prompt=system_prompt, session_id=session_id,
|
||||
gateway_session_key=gateway_session_key, **agent_overrides, route=route,
|
||||
relay_metadata=relay_metadata,
|
||||
# #98619: only an explicitly provided X-Hermes-Session-Id is wake-capable (the
|
||||
# header is 403-gated on API_SERVER_KEY, so the wake self-post can authenticate
|
||||
# and the client can resume the session by sending it again). A fingerprint-derived
|
||||
@@ -766,6 +769,8 @@ class OpenAICompatRoutesMixin:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
return _invalid_request("Invalid JSON in request body")
|
||||
from gateway.platforms.api_server import _request_relay_metadata
|
||||
relay_metadata = _request_relay_metadata(body)
|
||||
raw_input = body.get("input")
|
||||
if raw_input is None:
|
||||
return _error_response("Missing 'input' field", 400)
|
||||
@@ -846,7 +851,7 @@ class OpenAICompatRoutesMixin:
|
||||
user_message=user_message, conversation_history=conversation_history,
|
||||
ephemeral_system_prompt=instructions, session_id=session_id,
|
||||
gateway_session_key=gateway_session_key, bind_declared_conversation=_declared_selected,
|
||||
**agent_overrides, route=route)
|
||||
**agent_overrides, route=route, relay_metadata=relay_metadata)
|
||||
if stream:
|
||||
_stream_q = ThreadSafeAsyncQueue()
|
||||
|
||||
|
||||
@@ -167,6 +167,36 @@ def _run_turn(coordinator, lease, turn_id):
|
||||
return turn
|
||||
|
||||
|
||||
class TestTurnMetadata:
|
||||
def test_includes_request_metadata_without_overriding_runtime_fields(
|
||||
self, coordinator
|
||||
):
|
||||
fake = _FakeRelay()
|
||||
runtime = _make_runtime(fake)
|
||||
lease = _acquire(coordinator, runtime)
|
||||
|
||||
turn = coordinator.begin_turn(
|
||||
lease,
|
||||
turn_id="t1",
|
||||
task_id="task1",
|
||||
metadata={
|
||||
"request_id": "req-123",
|
||||
"context": {"tenant": "example"},
|
||||
relay_runtime.RUNTIME_INSTANCE_KEY: "caller-supplied",
|
||||
},
|
||||
)
|
||||
|
||||
turn_metadata = [
|
||||
push
|
||||
for push in fake.scope.pushes
|
||||
if push["name"] == relay_runtime.TURN_SCOPE
|
||||
][-1]["metadata"]
|
||||
assert turn_metadata["request_id"] == "req-123"
|
||||
assert turn_metadata["context"] == {"tenant": "example"}
|
||||
assert turn_metadata[relay_runtime.RUNTIME_INSTANCE_KEY] == runtime.runtime_id
|
||||
coordinator.end_turn(turn, outcome="success")
|
||||
|
||||
|
||||
class TestDefaultsNeverRotate:
|
||||
def test_no_rotation_across_many_turns_and_compactions(self, coordinator):
|
||||
fake = _FakeRelay()
|
||||
|
||||
@@ -35,6 +35,7 @@ from gateway.platforms.api_server import (
|
||||
_hermes_version,
|
||||
_redact_api_error_text,
|
||||
_request_agent_overrides,
|
||||
_request_relay_metadata,
|
||||
check_api_server_requirements,
|
||||
cors_middleware,
|
||||
security_headers_middleware,
|
||||
@@ -429,6 +430,36 @@ class TestAgentExecution:
|
||||
assert mock_agent._gateway_turn_process_baseline == frozenset()
|
||||
|
||||
|
||||
class TestRelayMetadataForwarding:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("endpoint", "payload"),
|
||||
[
|
||||
(
|
||||
"/v1/chat/completions",
|
||||
{"messages": [{"role": "user", "content": "hi"}]},
|
||||
),
|
||||
("/v1/responses", {"input": "hi"}),
|
||||
],
|
||||
)
|
||||
async def test_openai_requests_forward_metadata_to_relay(
|
||||
self, adapter, endpoint, payload
|
||||
):
|
||||
app = _create_app(adapter)
|
||||
metadata = {"request_id": "req-123", "context": {"tenant": "example"}}
|
||||
with patch.object(adapter, "_run_agent", new_callable=AsyncMock) as mock_run:
|
||||
mock_run.return_value = (
|
||||
{"final_response": "ok", "messages": [], "api_calls": 1},
|
||||
{"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
|
||||
)
|
||||
async with TestClient(TestServer(app)) as cli:
|
||||
response = await cli.post(endpoint, json={**payload, "metadata": metadata})
|
||||
|
||||
assert response.status == 200
|
||||
assert mock_run.call_args.kwargs["relay_metadata"] == metadata
|
||||
assert mock_run.call_args.kwargs["relay_metadata"] is not metadata
|
||||
|
||||
|
||||
class TestDisconnectedAgentReap:
|
||||
"""#76188 review: SSE disconnect handlers must reap only the background
|
||||
processes the disconnected turn created, and must no-op when no turn
|
||||
@@ -2844,6 +2875,30 @@ class TestKeyRejectionSetsNonRetryableFatalError:
|
||||
await self._assert_key_rejection_is_fatal(adapter)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Relay metadata extraction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRequestRelayMetadata:
|
||||
def test_copies_all_metadata_fields(self):
|
||||
metadata = {
|
||||
"request_id": "req-123",
|
||||
"attempt": 2,
|
||||
"tags": ["batch", "evaluation"],
|
||||
"context": {"tenant": "example"},
|
||||
}
|
||||
|
||||
extracted = _request_relay_metadata({"metadata": metadata})
|
||||
|
||||
assert extracted == metadata
|
||||
assert extracted is not metadata
|
||||
|
||||
@pytest.mark.parametrize("body", [None, [], {}, {"metadata": "invalid"}])
|
||||
def test_ignores_non_object_metadata(self, body):
|
||||
assert _request_relay_metadata(body) == {}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bare-model opt-in gate (direct_model_requests) for _request_agent_overrides
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user