feat: inherit the caller's model for async sub-agent launch and update (#446)

* feat: inherit the caller's model for async sub-agent launch and update

* docs: tighten middleware related docstrings
This commit is contained in:
jfilipiuk
2026-09-11 18:58:50 +02:00
committed by GitHub
parent a486b85851
commit 0410b40f57
4 changed files with 374 additions and 15 deletions
+50 -4
View File
@@ -36,6 +36,7 @@ from collections.abc import (
Mapping,
Sequence,
)
from contextvars import ContextVar
from typing import Any
from langchain_core.messages import AIMessage, BaseMessage
@@ -1086,6 +1087,41 @@ def _patch_anthropic_structured_output() -> None:
# ---------------------------------------------------------------------------
_model_passthrough_patched = False
# The launching run's per-run (model, provider), set by the async-task tools
# from ``runtime.config`` for the duration of a single ``runs.create`` call.
# ``_merge_runs_config_kwargs`` reads it as the highest-precedence source.
#
# Why a ContextVar rather than passing ``config=`` through each tool: the fix
# has to hold at every ``runs.create`` site (``start_async_task`` and
# ``update_async_task``), but only the tool functions can see the caller's
# per-run model (via ``runtime.config``, the config langgraph's ToolNode
# injects into tool calls). Threading it through a ContextVar
# lets the single merge point below inject it, so ``update_async_task`` (whose
# body we delegate to upstream unchanged) is covered without reimplementing it.
_caller_configurable: ContextVar[dict[str, str] | None] = ContextVar(
"_evo_caller_configurable", default=None
)
def _extract_caller_configurable(config: Any) -> dict[str, str]:
"""Pull ``(model, model_provider)`` out of a launching run's ``config``.
``config`` is the tool's ``runtime.config`` (a ``RunnableConfig``-shaped
dict). Returns a dict suitable for ``configurable``; empty when the caller
carries no model override, so the config-default fallback still applies.
``model_provider`` is only forwarded alongside a model — a bare provider
without a model is meaningless to the deployed graph's resolver.
"""
configurable = (config or {}).get("configurable") or {}
model = configurable.get("model")
provider = configurable.get("model_provider")
out: dict[str, str] = {}
if isinstance(model, str) and model:
out["model"] = model
if isinstance(provider, str) and provider:
out["model_provider"] = provider
return out
def _read_cfg_configurable() -> dict[str, str]:
"""Read live ``(model, provider)`` from EvoScientist config.
@@ -1114,11 +1150,21 @@ def _read_cfg_configurable() -> dict[str, str]:
def _merge_runs_config_kwargs(kwargs: dict) -> dict:
"""Merge the live model override into ``kwargs`` for ``runs.create``.
Preserves any caller-supplied ``config.configurable`` keys. EvoScientist's
keys take precedence on conflict (callers shouldn't be passing model
overrides — the CLI is the source of truth).
Model source, in increasing precedence:
1. ``_read_cfg_configurable()`` — the effective EvoScientist config. On the
in-process backend this is the CLI's live per-run model; on the
``langgraph_server`` backend this proxy runs inside the dev-server
process, where it reports the *server's* config-default model, not the
CLI's per-run choice.
2. ``_caller_configurable`` — the launching run's actual model, forwarded
by the async-task tools from ``runtime.config``. This wins so that a
sub-agent launched (or continued) while the caller runs on, say, a free
model does not silently fall back to a billed config-default.
Any unrelated caller-supplied ``config.configurable`` keys are preserved.
"""
overrides = _read_cfg_configurable()
overrides = {**_read_cfg_configurable(), **(_caller_configurable.get() or {})}
if not overrides:
return kwargs
existing = kwargs.get("config")
@@ -59,6 +59,7 @@ then called without ``runtime`` and raises ``TypeError``.
import asyncio
import logging
import threading
from contextlib import contextmanager
from datetime import UTC, datetime
from typing import Any, NotRequired
@@ -83,6 +84,74 @@ from langgraph.types import Command
_logger = logging.getLogger(__name__)
@contextmanager
def _caller_model_scope(runtime: ToolRuntime):
"""Forward the launching run's model to any ``runs.create`` in the block.
An async sub-agent is launched via a bare ``runs.create`` inside the
caller's run. Without help it falls back to the server's config-default
model rather than the model the caller is running on — so a run started on
a free model silently bills the config-default. Read the caller's per-run
model from ``runtime.config`` — the config langgraph's ToolNode injects
into every tool call, the same ``configurable`` channel that carries
``thread_id`` (``runtime`` is already injected into these tool
signatures, so it is the channel already in hand) — and publish it, for
the duration of the block,
to the contextvar the ``runs.create`` proxy reads. A sync ``with`` around an
``await`` is fine: the value is set before the await and reset after, and
contextvars propagate across awaits within the same task. Empty when the
caller has no override, preserving the default-model behaviour.
"""
from ..llm.patches import _caller_configurable, _extract_caller_configurable
token = _caller_configurable.set(
_extract_caller_configurable(getattr(runtime, "config", None))
)
try:
yield
finally:
_caller_configurable.reset(token)
def _build_expert_update_tool(
agent_map: dict[str, AsyncSubAgent],
clients: Any,
) -> StructuredTool:
"""``update_async_task`` wrapped to inherit the caller's model.
Delegates to upstream's tool body verbatim — preserving its
``multitask_strategy`` and task-envelope semantics — inside
``_caller_model_scope`` so the follow-up ``runs.create`` reaches the
sub-agent on the caller's model, not the config-default. The explicit
``runtime: ToolRuntime`` signature is required: langchain decides runtime
injection from ``inspect.signature``, so a ``*args`` wrapper would strip it.
"""
base = _build_update_tool(agent_map, clients)
orig_func = base.func
orig_coro = base.coroutine
def update_async_task(
task_id: str, message: str, runtime: ToolRuntime
) -> str | Command:
with _caller_model_scope(runtime):
return orig_func(task_id=task_id, message=message, runtime=runtime)
async def aupdate_async_task(
task_id: str, message: str, runtime: ToolRuntime
) -> str | Command:
with _caller_model_scope(runtime):
return await orig_coro(task_id=task_id, message=message, runtime=runtime)
return StructuredTool.from_function(
name=base.name,
func=update_async_task,
coroutine=aupdate_async_task,
description=base.description,
infer_schema=False,
args_schema=base.args_schema,
)
class ExpertAsyncSubAgent(AsyncSubAgent):
"""AsyncSubAgent spec extended with the expert-dispatch marker.
@@ -307,11 +376,12 @@ def _build_expert_start_tool(
try:
client = clients.get_sync(subagent_type)
thread = client.threads.create()
run = client.runs.create(
thread_id=thread["thread_id"],
assistant_id=spec["graph_id"],
input=input_dict,
)
with _caller_model_scope(runtime):
run = client.runs.create(
thread_id=thread["thread_id"],
assistant_id=spec["graph_id"],
input=input_dict,
)
except Exception as e:
_logger.warning(
"Failed to launch async subagent '%s': %s", subagent_type, e
@@ -349,11 +419,12 @@ def _build_expert_start_tool(
try:
client = clients.get_async(subagent_type)
thread = await client.threads.create()
run = await client.runs.create(
thread_id=thread["thread_id"],
assistant_id=spec["graph_id"],
input=input_dict,
)
with _caller_model_scope(runtime):
run = await client.runs.create(
thread_id=thread["thread_id"],
assistant_id=spec["graph_id"],
input=input_dict,
)
except Exception as e:
_logger.warning(
"Failed to launch async subagent '%s': %s", subagent_type, e
@@ -454,7 +525,7 @@ class EvoAsyncSubAgentMiddleware(AsyncSubAgentMiddleware):
self._resolve_lock,
),
_build_check_tool(clients),
_build_update_tool(agent_map, clients),
_build_expert_update_tool(agent_map, clients),
_build_cancel_tool(clients),
_build_list_tasks_tool(clients),
]