diff --git a/EvoScientist/llm/gateway_proxy.py b/EvoScientist/llm/gateway_proxy.py index 75531d9..6fb9a2e 100644 --- a/EvoScientist/llm/gateway_proxy.py +++ b/EvoScientist/llm/gateway_proxy.py @@ -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, diff --git a/EvoScientist/llm/runtime.py b/EvoScientist/llm/runtime.py index e24f432..64a7361 100644 --- a/EvoScientist/llm/runtime.py +++ b/EvoScientist/llm/runtime.py @@ -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: diff --git a/EvoScientist/middleware/recoverable_metering.py b/EvoScientist/middleware/recoverable_metering.py index 7418213..2fcc2e3 100644 --- a/EvoScientist/middleware/recoverable_metering.py +++ b/EvoScientist/middleware/recoverable_metering.py @@ -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(