fix: fail-closed internal identity and preserve billing error semantics
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
- Gateway internal identity: when a service token is configured, reject wrong/missing tokens even from loopback (closes SSRF/local bypass). - Terminal metering: classified AgentControlError propagates without retry; exhausted retries raise BILLING_UNAVAILABLE instead of a generic RuntimeError, keeping error attribution accurate.
This commit is contained in:
@@ -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 {}
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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')"
|
||||
|
||||
Reference in New Issue
Block a user