diff --git a/CHANGELOG.md b/CHANGELOG.md index 40e806916..d636f2f84 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -52,6 +52,8 @@ to include examples, links to docs, or any other relevant information. ### Fixed +- `contrib.deepagents`: summarization middleware configured with a model name string now routes its LLM calls through Activities instead of running them in the Workflow. + - **Experimental**: External storage metrics now report the wall-clock time storage was in flight. Previously each batch's duration was summed, over-reporting the time whenever storage operations ran concurrently. diff --git a/temporalio/contrib/deepagents/_model.py b/temporalio/contrib/deepagents/_model.py index 2d97f992d..c9f2b9a09 100644 --- a/temporalio/contrib/deepagents/_model.py +++ b/temporalio/contrib/deepagents/_model.py @@ -263,31 +263,29 @@ def _stream( # create_deep_agent model patch (implicit wrapping) # --------------------------------------------------------------------------- # -# The durability seam is ``deepagents._models.resolve_model``, which -# ``create_deep_agent`` calls to turn a ``model=`` string (or instance) into a -# ``BaseChatModel`` — for both the top-level agent and every ``SubAgent`` -# (``graph.py`` lines 592 and 634). We patch it on the ``deepagents.graph`` -# module, where ``create_deep_agent``'s body resolves the ``resolve_model`` name -# at call time. +# The durability seam is ``resolve_model``, which turns a ``model=`` name +# string into a live ``BaseChatModel``. We patch every binding a string can +# reach in-workflow: # -# Patching *this* seam (not ``deepagents.create_deep_agent``) is what makes the -# rewrite survive the user's import style. A user who writes the idiomatic -# ``from deepagents import create_deep_agent`` binds the *original* function -# object into their module; rebinding the ``deepagents.create_deep_agent`` -# attribute would never be seen by that already-bound reference, so string -# models would reach the real provider inside the workflow (a hang / non- -# determinism). ``create_deep_agent``'s body, by contrast, always looks up -# ``resolve_model`` in the ``deepagents.graph`` globals afresh on each call, so -# rebinding it there is observed no matter how the caller imported the factory. -# It also preserves ``_model_spec`` (the original string), which the factory -# reads *before* calling ``resolve_model`` for harness-profile lookup. +# - ``deepagents.graph`` — read afresh by ``create_deep_agent`` for the main +# agent and every sub-agent (its module-top import froze its own binding). +# - ``deepagents._models`` — the definition site; covers call-time importers +# such as ``create_summarization_tool_middleware``. +# - ``SummarizationMiddleware.__init__`` — pre-resolves a string summarizer +# before the middleware delegates to LangChain's ``init_chat_model``. # -# ``create_deep_agent`` is still wrapped separately, best-effort, purely to fire -# the construction-time warnings that need the ``tools`` / ``checkpointer`` -# kwargs (those warnings are advisory and carry no durability weight). +# An unpatched binding means a real provider client constructed (and called) +# inside the workflow: nondeterministic and replay-unsafe. Every patched seam +# is deepagents-internal and covered by this package's deepagents version +# pin, so none is guarded: a missing seam is a broken install and fails the +# worker at startup rather than silently reverting. We do NOT rebind +# ``deepagents.create_deep_agent`` for durability — callers who already did +# ``from deepagents import create_deep_agent`` hold the original object — only +# a best-effort wrap to fire the advisory construction-time warnings. _original_create_deep_agent: Any = None _original_resolve_model: Any = None +_original_summarization_init: Any = None def _wrap_model_arg(model: Any) -> Any: @@ -317,21 +315,21 @@ def _wrap_model_arg(model: Any) -> Any: def install_model_patch() -> None: """Route Deep Agents' model resolution through :class:`TemporalModel`. - Patches ``deepagents.graph.resolve_model`` (the seam ``create_deep_agent`` - uses for the main agent *and* every sub-agent) so a bare ``model="..."`` - string becomes a durable :class:`TemporalModel`, and additionally wraps - ``deepagents.create_deep_agent`` to fire the advisory tool / checkpointer - warnings. Both only act when called inside a workflow, so importing - deepagents on a plain client / activity worker is unaffected. Idempotent. + Patches the model-resolution bindings listed above so a name string + becomes a durable :class:`TemporalModel` in-workflow, and wraps + ``deepagents.create_deep_agent`` for the advisory warnings. No effect + outside workflows. Idempotent. """ global _original_create_deep_agent, _original_resolve_model # importlib: `deepagents` is absent on Python 3.10 environments (its floor # is 3.11), so static imports here fail type-checking there. deepagents = importlib.import_module("deepagents") _graph = importlib.import_module("deepagents.graph") + _models = importlib.import_module("deepagents._models") if _original_resolve_model is None: - _original_resolve_model = _graph.resolve_model + # One original serves both bindings (same function object). + _original_resolve_model = _models.resolve_model def patched_resolve_model(model: Any) -> Any: if workflow.in_workflow(): @@ -339,6 +337,23 @@ def patched_resolve_model(model: Any) -> Any: return _original_resolve_model(model) setattr(_graph, "resolve_model", patched_resolve_model) + setattr(_models, "resolve_model", patched_resolve_model) + + global _original_summarization_init + if _original_summarization_init is None: + _da_sum = importlib.import_module("deepagents.middleware.summarization") + summarization_cls = _da_sum.SummarizationMiddleware + original_init = summarization_cls.__init__ + _original_summarization_init = original_init + + def patched_summarization_init( + self: Any, model: Any, *args: Any, **kwargs: Any + ) -> None: + if workflow.in_workflow() and isinstance(model, str): + model = _wrap_model_arg(model) + original_init(self, model, *args, **kwargs) + + summarization_cls.__init__ = patched_summarization_init if _original_create_deep_agent is None: _original_create_deep_agent = deepagents.create_deep_agent @@ -362,9 +377,17 @@ def uninstall_model_patch() -> None: global _original_create_deep_agent, _original_resolve_model if _original_resolve_model is not None: _graph = importlib.import_module("deepagents.graph") + _models = importlib.import_module("deepagents._models") setattr(_graph, "resolve_model", _original_resolve_model) + setattr(_models, "resolve_model", _original_resolve_model) _original_resolve_model = None + global _original_summarization_init + if _original_summarization_init is not None: + _da_sum = importlib.import_module("deepagents.middleware.summarization") + + _da_sum.SummarizationMiddleware.__init__ = _original_summarization_init + _original_summarization_init = None if _original_create_deep_agent is not None: deepagents = importlib.import_module("deepagents") diff --git a/tests/contrib/deepagents/test_summarization.py b/tests/contrib/deepagents/test_summarization.py new file mode 100644 index 000000000..a6c942952 --- /dev/null +++ b/tests/contrib/deepagents/test_summarization.py @@ -0,0 +1,97 @@ +"""String summarizer models must resolve to a durable ``TemporalModel``. + +``SummarizationMiddleware`` resolves a name string via LangChain's +``init_chat_model`` binding, and ``create_summarization_tool_middleware`` via +a call-time import of ``deepagents._models.resolve_model`` — neither reads +the patched ``deepagents.graph`` seam, so in-workflow both built a real +provider client (the default stack is unaffected: it receives the +already-resolved agent model). One pin per seam. +""" + +from __future__ import annotations + +import sys +import uuid + +import pytest + +from temporalio.testing import WorkflowEnvironment + +pytestmark = pytest.mark.skipif( + sys.version_info < (3, 11), reason="deepagents requires Python >= 3.11" +) +pytest.importorskip("deepagents") +pytest.importorskip("langchain_core") + +from temporalio import workflow +from temporalio.contrib.deepagents import DeepAgentsPlugin, TemporalModel +from temporalio.worker import Worker + +# Bind deepagents symbols off importorskip modules: static imports cannot +# resolve on Python 3.10 environments, where deepagents is absent. +_backends_mod = pytest.importorskip("deepagents.backends") +_middleware_mod = pytest.importorskip("deepagents.middleware") +_summarization_mod = pytest.importorskip("deepagents.middleware.summarization") +StateBackend = _backends_mod.StateBackend +SummarizationMiddleware = _middleware_mod.SummarizationMiddleware +create_summarization_tool_middleware = ( + _summarization_mod.create_summarization_tool_middleware +) + + +@workflow.defn +class SummarizerResolutionWorkflow: + @workflow.run + async def run(self) -> str: + middleware = SummarizationMiddleware( + "anthropic:claude-sonnet-4-5", + backend=StateBackend(), + ) + return type(middleware.model).__name__ + + +@pytest.mark.asyncio +async def test_summarizer_model_resolves_durable(env: WorkflowEnvironment) -> None: + plugin = DeepAgentsPlugin() + async with Worker( + env.client, + task_queue="da-summarizer-resolve", + workflows=[SummarizerResolutionWorkflow], + plugins=[plugin], + ): + out = await env.client.execute_workflow( + SummarizerResolutionWorkflow.run, + id=f"da-summarizer-resolve-{uuid.uuid4()}", + task_queue="da-summarizer-resolve", + ) + assert out == TemporalModel.__name__, out + + +@workflow.defn +class ToolMiddlewareResolutionWorkflow: + @workflow.run + async def run(self) -> str: + middleware = create_summarization_tool_middleware( + "anthropic:claude-sonnet-4-5", + StateBackend(), + ) + return type(middleware._summarization.model).__name__ + + +@pytest.mark.asyncio +async def test_tool_middleware_summarizer_resolves_durable( + env: WorkflowEnvironment, +) -> None: + plugin = DeepAgentsPlugin() + async with Worker( + env.client, + task_queue="da-summarizer-tool-resolve", + workflows=[ToolMiddlewareResolutionWorkflow], + plugins=[plugin], + ): + out = await env.client.execute_workflow( + ToolMiddlewareResolutionWorkflow.run, + id=f"da-summarizer-tool-resolve-{uuid.uuid4()}", + task_queue="da-summarizer-tool-resolve", + ) + assert out == TemporalModel.__name__, out