fix: propagate streamed-model attempt id via shared configurable

The legacy BaseChatModel.astream path does not forward run_manager to
_astream, so the per-call tracing run_id was unreachable and the stream
fell back to the conversation run_id, which the gateway's
verify_model_attempt rejected (401 RUN_ATTEMPT_NOT_ACCEPTED). Publish the
tracing run_id into the shared configurable dict from
on_chat_model_start and read it back in _attempt_id.
This commit is contained in:
m4
2026-08-20 15:32:18 +08:00
parent 386c8130ea
commit 3683cbfc13
3 changed files with 50 additions and 3 deletions
+22 -2
View File
@@ -45,6 +45,26 @@ class GatewayProxyChatModel(BaseChatModel):
update={"bound_tools": serialized, "bound_tool_choice": tool_choice} update={"bound_tools": serialized, "bound_tool_choice": tool_choice}
) )
def _attempt_id(self, run_manager: Any) -> str:
# The metering callback publishes the per-call tracing run_id into the
# shared configurable dict (see RecoverableMeteringCallback). The legacy
# BaseChatModel.astream path does NOT forward run_manager to _astream, so
# run_manager.run_id is unavailable here; the configurable channel is the
# only reliable source of the id the gateway registered under.
try:
from langgraph.config import get_config
config = get_config()
except Exception:
config = None
if isinstance(config, dict):
configurable = config.get("configurable")
if isinstance(configurable, dict):
current = configurable.get("ai4sci_attempt_id")
if current:
return str(current)
return str(getattr(run_manager, "run_id", None) or self.run_id)
def _generate(self, *args: Any, **kwargs: Any) -> ChatResult: def _generate(self, *args: Any, **kwargs: Any) -> ChatResult:
del args, kwargs del args, kwargs
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH") raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH")
@@ -57,7 +77,7 @@ class GatewayProxyChatModel(BaseChatModel):
**kwargs: Any, **kwargs: Any,
) -> ChatResult: ) -> ChatResult:
del stop, kwargs del stop, kwargs
attempt_id = str(getattr(run_manager, "run_id", None) or self.run_id) attempt_id = self._attempt_id(run_manager)
async with httpx.AsyncClient(timeout=httpx.Timeout(660.0, connect=5.0)) as client: async with httpx.AsyncClient(timeout=httpx.Timeout(660.0, connect=5.0)) as client:
response = await client.post( response = await client.post(
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke", f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke",
@@ -85,7 +105,7 @@ class GatewayProxyChatModel(BaseChatModel):
**kwargs: Any, **kwargs: Any,
): ):
del stop, kwargs del stop, kwargs
attempt_id = str(getattr(run_manager, "run_id", None) or self.run_id) attempt_id = self._attempt_id(run_manager)
payload = { payload = {
"run_id": self.run_id, "run_id": self.run_id,
"attempt_id": attempt_id, "attempt_id": attempt_id,
+4
View File
@@ -1798,6 +1798,10 @@ class _EvoWebRun:
except asyncio.CancelledError: except asyncio.CancelledError:
return return
except Exception as exc: except Exception as exc:
logger.exception(
"evo run %s failed before model provider error was reported",
getattr(self._admission, "request_id", "?"),
)
error_code = _safe_error_code(exc, fallback="AGENT_RUNTIME_ERROR") error_code = _safe_error_code(exc, fallback="AGENT_RUNTIME_ERROR")
error_details = _run_failure_details(exc, error_code) error_details = _run_failure_details(exc, error_code)
if self._last_model_failure is not None: if self._last_model_failure is not None:
@@ -47,6 +47,23 @@ def _metering_config(config: Mapping[str, Any] | None) -> dict[str, str] | None:
return normalized if all(normalized[name] for name in required) else None return normalized if all(normalized[name] for name in required) else None
def _publish_attempt_id(tracing_run_id: str) -> None:
"""Publish the per-call tracing run_id into the shared configurable dict.
LangChain dispatches ``on_chat_model_start`` inside a ``copy_context()``
task, so a ContextVar set here never reaches the model's ``_astream`` (which
runs in the parent context). The configurable dict is the same mutable
object in both contexts, so writing here lets ``_astream`` read the id the
gateway registered as this attempt's id.
"""
config = _config()
if not isinstance(config, dict):
return
configurable = config.get("configurable")
if isinstance(configurable, dict):
configurable["ai4sci_attempt_id"] = tracing_run_id
def _first(mapping: Mapping[str, Any], *keys: str) -> Any: def _first(mapping: Mapping[str, Any], *keys: str) -> Any:
for key in keys: for key in keys:
value = mapping.get(key) value = mapping.get(key)
@@ -169,6 +186,7 @@ class RecoverableMeteringCallback(AsyncCallbackHandler):
del messages del messages
metadata = dict(metadata or {}) metadata = dict(metadata or {})
invocation = dict(kwargs.get("invocation_params") or {}) invocation = dict(kwargs.get("invocation_params") or {})
_publish_attempt_id(str(run_id))
serialized_kwargs = dict(serialized.get("kwargs") or {}) serialized_kwargs = dict(serialized.get("kwargs") or {})
provider = self.config.get("provider_id") or str( provider = self.config.get("provider_id") or str(
_first(metadata, "route_provider", "ls_provider", "provider") _first(metadata, "route_provider", "ls_provider", "provider")
@@ -203,7 +221,12 @@ class RecoverableMeteringCallback(AsyncCallbackHandler):
async def on_llm_error( async def on_llm_error(
self, error: BaseException, *, run_id: uuid.UUID, **kwargs: Any self, error: BaseException, *, run_id: uuid.UUID, **kwargs: Any
) -> None: ) -> None:
del error, kwargs del kwargs
logger.error(
"recoverable metering: model error for attempt %s",
run_id,
exc_info=(type(error), error, error.__traceback__),
)
await self._terminal(run_id, "failed", None) await self._terminal(run_id, "failed", None)
async def _terminal( async def _terminal(