diff --git a/EvoScientist/internal_service.py b/EvoScientist/internal_service.py new file mode 100644 index 0000000..d917567 --- /dev/null +++ b/EvoScientist/internal_service.py @@ -0,0 +1,17 @@ +"""Authentication headers for LangGraph-to-Gateway internal calls.""" + +from __future__ import annotations + +import os + + +def internal_service_token() -> str: + return ( + os.environ.get("EVOSCIENTIST_BACKEND_SERVICE_TOKEN", "").strip() + or os.environ.get("AI4SCI_EVO_RUNTIME_GRANT_SECRET", "").strip() + ) + + +def internal_service_headers() -> dict[str, str]: + token = internal_service_token() + return {"X-Ai4Sci-Service-Token": token} if token else {} diff --git a/EvoScientist/llm/gateway_proxy.py b/EvoScientist/llm/gateway_proxy.py index c28bb62..dbd42ed 100644 --- a/EvoScientist/llm/gateway_proxy.py +++ b/EvoScientist/llm/gateway_proxy.py @@ -14,6 +14,8 @@ from langchain_core.tools import BaseTool from langchain_core.utils.function_calling import convert_to_openai_tool from pydantic import Field +from EvoScientist.internal_service import internal_service_headers + from .contracts import EvoRuntimeError @@ -93,6 +95,7 @@ class GatewayProxyChatModel(BaseChatModel): "tools": self.bound_tools, "tool_choice": self.bound_tool_choice, }, + headers=internal_service_headers(), ) response.raise_for_status() value = response.json() @@ -126,6 +129,7 @@ class GatewayProxyChatModel(BaseChatModel): "POST", f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/stream", json=payload, + headers=internal_service_headers(), ) as response: response.raise_for_status() saw_done = False diff --git a/EvoScientist/middleware/dynamic_review.py b/EvoScientist/middleware/dynamic_review.py index 90f38ff..7ff83d4 100644 --- a/EvoScientist/middleware/dynamic_review.py +++ b/EvoScientist/middleware/dynamic_review.py @@ -7,6 +7,8 @@ from collections.abc import Mapping from typing import Annotated, Any, NotRequired import httpx + +from EvoScientist.internal_service import internal_service_headers from langchain.agents.middleware import HumanInTheLoopMiddleware from langchain.agents.middleware.types import AgentState, OmitFromSchema from langgraph.config import get_config @@ -62,8 +64,7 @@ def _request_payload( def _service_headers() -> dict[str, str]: - token = os.environ.get("EVOSCIENTIST_BACKEND_SERVICE_TOKEN", "").strip() - return {"X-Ai4Sci-Service-Token": token} if token else {} + return internal_service_headers() def _validated_auto_state( diff --git a/EvoScientist/middleware/recoverable_metering.py b/EvoScientist/middleware/recoverable_metering.py index 2fcc2e3..f13ea13 100644 --- a/EvoScientist/middleware/recoverable_metering.py +++ b/EvoScientist/middleware/recoverable_metering.py @@ -16,6 +16,9 @@ from langchain.agents.middleware.types import ( ) from langchain_core.callbacks import AsyncCallbackHandler +from EvoScientist.internal_service import internal_service_headers +from EvoScientist.llm.errors import AgentControlError + logger = logging.getLogger(__name__) _clients: dict[str, httpx.AsyncClient] = {} @@ -154,15 +157,24 @@ async def _post(config: Mapping[str, str], phase: str, payload: dict[str, Any]) if client is None: client = httpx.AsyncClient(timeout=httpx.Timeout(15.0, connect=3.0)) _clients[base_url] = client - response = await client.post( - f"{base_url}/api/internal/recoverable-runs/metering/{phase}", - json={ - **payload, - "run_id": config["run_id"], - "envelope_signature": config["envelope_signature"], - }, - ) - response.raise_for_status() + try: + response = await client.post( + f"{base_url}/api/internal/recoverable-runs/metering/{phase}", + json={ + **payload, + "run_id": config["run_id"], + "envelope_signature": config["envelope_signature"], + }, + headers=internal_service_headers(), + ) + response.raise_for_status() + except httpx.HTTPError as exc: + raise AgentControlError( + "BILLING_UNAVAILABLE", + "The metering service is unavailable.", + status_code=503, + retryable=True, + ) from exc class RecoverableMeteringCallback(AsyncCallbackHandler): @@ -244,10 +256,17 @@ class RecoverableMeteringCallback(AsyncCallbackHandler): ) self._attempts.discard(run_id) return + except AgentControlError: + raise except Exception as exc: last_error = exc await asyncio.sleep(0.1 * (attempt + 1)) - raise RuntimeError("AI4SCI_METERING_TERMINAL_FAILED") from last_error + raise AgentControlError( + "BILLING_UNAVAILABLE", + "The metering service is unavailable.", + status_code=503, + retryable=True, + ) from last_error class RecoverableMeteringMiddleware(AgentMiddleware): diff --git a/EvoScientist/middleware/recoverable_tools.py b/EvoScientist/middleware/recoverable_tools.py index 9ff3a8e..d1ae029 100644 --- a/EvoScientist/middleware/recoverable_tools.py +++ b/EvoScientist/middleware/recoverable_tools.py @@ -13,6 +13,7 @@ from langchain_core.messages import ToolMessage, message_to_dict, messages_from_ from langgraph.types import Command from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.internal_service import internal_service_headers if TYPE_CHECKING: from langchain.agents.middleware.types import ToolCallRequest @@ -92,6 +93,7 @@ async def _post(proxy: Mapping[str, str], phase: str, payload: dict[str, Any]) - "attempt_id": proxy["run_id"], "envelope_signature": proxy["envelope_signature"], }, + headers=internal_service_headers(), ) if response.is_error: try: diff --git a/tests/test_gateway_proxy.py b/tests/test_gateway_proxy.py index 72961bb..3b51bcc 100644 --- a/tests/test_gateway_proxy.py +++ b/tests/test_gateway_proxy.py @@ -37,10 +37,11 @@ class _FakeClient: async def __aexit__(self, *exc): return False - def stream(self, method, url, json=None): + def stream(self, method, url, json=None, headers=None): assert method == "POST" assert url.endswith("/api/internal/recoverable-runs/model/stream") self.sent = json + self.headers = headers return _FakeStream(self._lines) @@ -58,6 +59,7 @@ def test_runtime_error_repr_preserves_only_stable_code(): @pytest.mark.anyio async def test_astream_yields_chunks_from_sse(monkeypatch): + monkeypatch.setenv("AI4SCI_EVO_RUNTIME_GRANT_SECRET", "runtime-service-secret") model = GatewayProxyChatModel( gateway_url="http://gw", run_id="run-1", @@ -75,6 +77,7 @@ async def test_astream_yields_chunks_from_sse(monkeypatch): chunks = [c async for c in model._astream([HumanMessage(content="hi")])] assert fake.sent["stream"] is True + assert fake.headers == {"X-Ai4Sci-Service-Token": "runtime-service-secret"} assert len(chunks) == 2 assert chunks[0].message.content == "hello" diff --git a/tests/test_host_metering_extensions.py b/tests/test_host_metering_extensions.py index afbaa07..4ef5ffd 100644 --- a/tests/test_host_metering_extensions.py +++ b/tests/test_host_metering_extensions.py @@ -1,5 +1,11 @@ from EvoScientist.llm.errors import AgentControlError from EvoScientist.middleware.model_fallback import _is_non_fallbackable +import uuid + +import httpx +import pytest + +from EvoScientist.middleware import recoverable_metering from EvoScientist.middleware.recoverable_metering import _metering_config, _source_type @@ -33,3 +39,89 @@ def test_recoverable_metering_reads_explicit_evomemory_scope(): assert _source_type({"metering_scope": "evomemory_subagent_worker"}, []) == ( "evomemory_subagent_worker" ) + + +@pytest.mark.anyio +async def test_recoverable_metering_sends_internal_service_identity(monkeypatch): + captured = {} + + class Response: + def raise_for_status(self): + return None + + class Client: + async def post(self, url, *, json, headers): + captured.update(url=url, json=json, headers=headers) + return Response() + + monkeypatch.setenv("AI4SCI_EVO_RUNTIME_GRANT_SECRET", "runtime-service-secret") + monkeypatch.delenv("EVOSCIENTIST_BACKEND_SERVICE_TOKEN", raising=False) + monkeypatch.setattr(recoverable_metering, "_clients", {"http://gateway": Client()}) + + await recoverable_metering._post( + { + "gateway_url": "http://gateway", + "run_id": "run-1", + "envelope_signature": "signature", + }, + "start", + {}, + ) + + assert captured["headers"] == { + "X-Ai4Sci-Service-Token": "runtime-service-secret" + } + + +@pytest.mark.anyio +async def test_recoverable_metering_reports_gateway_identity_failure(monkeypatch): + class Client: + async def post(self, url, *, json, headers): + return httpx.Response(401, request=httpx.Request("POST", url)) + + monkeypatch.setattr(recoverable_metering, "_clients", {"http://gateway": Client()}) + + with pytest.raises(AgentControlError) as exc_info: + await recoverable_metering._post( + { + "gateway_url": "http://gateway", + "run_id": "run-1", + "envelope_signature": "signature", + }, + "start", + {}, + ) + + assert exc_info.value.code == "BILLING_UNAVAILABLE" + + +async def _fake_sleep(_delay: float) -> None: + return None + + +@pytest.mark.anyio +async def test_terminal_metering_preserves_billing_unavailable_and_skips_retry_on_401(monkeypatch): + calls = {"count": 0} + + class Client: + async def post(self, url, *, json, headers): + calls["count"] += 1 + return httpx.Response(401, request=httpx.Request("POST", url)) + + monkeypatch.setattr(recoverable_metering, "_clients", {"http://gateway": Client()}) + monkeypatch.setattr(recoverable_metering.asyncio, "sleep", _fake_sleep) + + callback = recoverable_metering.RecoverableMeteringCallback( + { + "gateway_url": "http://gateway", + "run_id": "run-1", + "envelope_signature": "signature", + } + ) + callback._attempts.add(uuid.UUID(int=1)) + + with pytest.raises(AgentControlError) as exc_info: + await callback._terminal(uuid.UUID(int=1), "succeeded", None) + + assert exc_info.value.code == "BILLING_UNAVAILABLE" + assert calls["count"] == 1 diff --git a/tests/test_recoverable_tools.py b/tests/test_recoverable_tools.py index d6ca130..6c57750 100644 --- a/tests/test_recoverable_tools.py +++ b/tests/test_recoverable_tools.py @@ -59,6 +59,8 @@ def test_evomemory_never_falls_back_to_parent_model_proxy(monkeypatch): @pytest.mark.asyncio async def test_tool_effect_gateway_error_preserves_machine_code(monkeypatch): + captured = {} + class Client: async def __aenter__(self): return self @@ -66,13 +68,15 @@ async def test_tool_effect_gateway_error_preserves_machine_code(monkeypatch): async def __aexit__(self, *_): return None - async def post(self, url, json): + async def post(self, url, json, headers): + captured["headers"] = headers return httpx.Response( 409, json={"detail": {"code": "RUN_FENCE_LOST"}}, request=httpx.Request("POST", url), ) + monkeypatch.setenv("AI4SCI_EVO_RUNTIME_GRANT_SECRET", "runtime-service-secret") monkeypatch.setattr(recoverable_tools.httpx, "AsyncClient", lambda **_: Client()) with pytest.raises(EvoRuntimeError) as exc_info: @@ -87,4 +91,7 @@ async def test_tool_effect_gateway_error_preserves_machine_code(monkeypatch): ) assert exc_info.value.code == "RUN_FENCE_LOST" + assert captured["headers"] == { + "X-Ai4Sci-Service-Token": "runtime-service-secret" + } assert repr(exc_info.value) == "EvoRuntimeError(code='RUN_FENCE_LOST')"