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:
@@ -45,6 +45,26 @@ class GatewayProxyChatModel(BaseChatModel):
|
||||
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:
|
||||
del args, kwargs
|
||||
raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH")
|
||||
@@ -57,7 +77,7 @@ class GatewayProxyChatModel(BaseChatModel):
|
||||
**kwargs: Any,
|
||||
) -> ChatResult:
|
||||
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:
|
||||
response = await client.post(
|
||||
f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke",
|
||||
@@ -85,7 +105,7 @@ class GatewayProxyChatModel(BaseChatModel):
|
||||
**kwargs: Any,
|
||||
):
|
||||
del stop, kwargs
|
||||
attempt_id = str(getattr(run_manager, "run_id", None) or self.run_id)
|
||||
attempt_id = self._attempt_id(run_manager)
|
||||
payload = {
|
||||
"run_id": self.run_id,
|
||||
"attempt_id": attempt_id,
|
||||
|
||||
@@ -1798,6 +1798,10 @@ class _EvoWebRun:
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
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_details = _run_failure_details(exc, error_code)
|
||||
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
|
||||
|
||||
|
||||
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:
|
||||
for key in keys:
|
||||
value = mapping.get(key)
|
||||
@@ -169,6 +186,7 @@ class RecoverableMeteringCallback(AsyncCallbackHandler):
|
||||
del messages
|
||||
metadata = dict(metadata or {})
|
||||
invocation = dict(kwargs.get("invocation_params") or {})
|
||||
_publish_attempt_id(str(run_id))
|
||||
serialized_kwargs = dict(serialized.get("kwargs") or {})
|
||||
provider = self.config.get("provider_id") or str(
|
||||
_first(metadata, "route_provider", "ls_provider", "provider")
|
||||
@@ -203,7 +221,12 @@ class RecoverableMeteringCallback(AsyncCallbackHandler):
|
||||
async def on_llm_error(
|
||||
self, error: BaseException, *, run_id: uuid.UUID, **kwargs: Any
|
||||
) -> 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)
|
||||
|
||||
async def _terminal(
|
||||
|
||||
Reference in New Issue
Block a user