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:
@@ -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),
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user