Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
268 changes: 268 additions & 0 deletions libs/code/deepagents_code/_js_cost.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,268 @@
"""Accounting-only transport for JavaScript subagents."""

from __future__ import annotations

import asyncio
import dataclasses
import hashlib
import itertools
import json
import math
from collections.abc import Mapping
from typing import TYPE_CHECKING, override

from langchain_core.tools import StructuredTool
from langchain_quickjs import CodeInterpreterMiddleware
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
CheckpointMetadata,
empty_checkpoint,
)
from langgraph.config import get_config
from langgraph.types import Command

from deepagents_code.cost_tracking import (
_empty_cost_breakdown,
_merge_cost_breakdowns,
_parent_checkpoint_scope,
)

if TYPE_CHECKING:
from collections.abc import Awaitable, Callable

from langchain_core.messages import ToolMessage
from langchain_core.runnables import RunnableConfig
from langgraph.prebuilt.tool_node import ToolCallRequest, ToolRuntime

from deepagents_code.cost_tracking import CostBreakdown

_ACCOUNTING_KEY = "__deepagents_js_cost_owner"
_RESPONSE_FORMAT_KEY = "__deepagents_subagent_response_format"


class _ReceiptMetadata(CheckpointMetadata):
js_cost_owner: str


def _digest(value: str) -> str:
"""Return a hashed identity without storing prompts in checkpoint names."""
return hashlib.sha256(value.encode()).hexdigest()


class CostAwareCodeInterpreterMiddleware(CodeInterpreterMiddleware):
"""Transport cost receipts without caching JavaScript or forwarding state."""

@override
def wrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], ToolMessage | Command],
) -> ToolMessage | Command:
"""Preserve synchronous execution; receipt transport is async-only.

Args:
request: Tool call and its runtime context.
handler: Synchronous tool executor.

Returns:
The tool result without modification.
"""
return handler(request)

async def awrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
) -> ToolMessage | Command:
"""Return owned receipts after settling all dispatched children."""
if request.tool not in self.tools:
return await handler(request)
configurable = request.runtime.config.get("configurable", {})
scope = configurable.get("checkpoint_ns", "")
owner = configurable.get(_ACCOUNTING_KEY)
accounting = (
owner or f"{scope}|js_cost:{_digest(request.runtime.tool_call_id or '')}"
)
active: set[asyncio.Task] = set()
tools = [
_cost_task(tool, request.runtime, accounting, active)
if isinstance(tool, StructuredTool) and tool.name == "task"
else tool
for tool in request.runtime.tools
]
runtime = dataclasses.replace(request.runtime, tools=tools)
try:
result = await handler(request.override(runtime=runtime))
finally:
for invocation in list(active):
if not invocation.cancelling():
invocation.cancel()
await asyncio.gather(*active, return_exceptions=True)
if owner:
return result
receipt = await _receipt_total(runtime.config, accounting)
if receipt is None:
return result
total, breakdown = receipt
transfers = {
accounting: {
"owner_scope": _parent_checkpoint_scope(scope),
"cost_usd": total,
"breakdown": breakdown,
}
}
if isinstance(result, Command):
update = dict(result.update) if isinstance(result.update, Mapping) else {}
update["_session_cost_transfers"] = transfers
return dataclasses.replace(result, update=update)
return Command(
update={"messages": [result], "_session_cost_transfers": transfers}
)


def _cost_task(
tool: StructuredTool,
outer_runtime: ToolRuntime,
owner: str,
active: set[asyncio.Task],
) -> StructuredTool:
"""Return a task proxy with request-isolated child graph checkpoints."""
occurrences: dict[str, int] = {}

async def invoke(
description: str, subagent_type: str, runtime: ToolRuntime
) -> object:
configurable = dict(runtime.config.get("configurable", {}))
response_format = configurable.get(_RESPONSE_FORMAT_KEY)
fingerprint = _digest(
json.dumps(
[description, subagent_type, getattr(response_format, "schema", None)],
sort_keys=True,
)
)
occurrence = occurrences.get(fingerprint, 0)
occurrences[fingerprint] = occurrence + 1
scope = outer_runtime.config.get("configurable", {}).get("checkpoint_ns", "")
configurable["checkpoint_ns"] = (
f"{scope}|js_dispatch:{fingerprint}_{occurrence}"
)
configurable[_ACCOUNTING_KEY] = owner
# The SDK's task closure owns the actual child ainvoke, including
# dynamic response schemas. Public durability="sync" belongs there;
# until that boundary exposes it, inherit LangGraph 1.2's config key.
configurable["__pregel_durability"] = "sync"
configurable["__deepagents_js_cost_loop"] = asyncio.get_running_loop()
scratchpad = configurable.get("__pregel_scratchpad")
if scratchpad is not None:
configurable["__pregel_scratchpad"] = dataclasses.replace(
scratchpad, subgraph_counter=itertools.count().__next__
)
config: RunnableConfig = {**runtime.config, "configurable": configurable}
invocation = asyncio.current_task()
if invocation is not None:
active.add(invocation)
try:
return await tool.arun(
{
"description": description,
"subagent_type": subagent_type,
"runtime": dataclasses.replace(runtime, config=config),
},
config=config,
tool_call_id=runtime.tool_call_id,
)
finally:
if invocation is not None:
active.discard(invocation)

return tool.model_copy(update={"coroutine": invoke})


def record_cost_receipt(
amount: float, *, breakdown: CostBreakdown | None = None
) -> None:
"""Persist one local node delta before its graph update returns.

Args:
amount: Dollars priced here, excluding claimed descendant transfers.
breakdown: Local requests, including free and unpriceable usage.
Omitted detail is preserved as incomplete legacy accounting.
"""
configurable = get_config().get("configurable", {})
owner = configurable.get(_ACCOUNTING_KEY)
saver = configurable.get("__pregel_checkpointer")
if not owner or not isinstance(saver, BaseCheckpointSaver):
return
if amount <= 0 and (breakdown is None or breakdown["request_count"] == 0):
return
namespace = configurable.get("checkpoint_ns", "")
config: RunnableConfig = {
"configurable": {
"thread_id": configurable["thread_id"],
"checkpoint_ns": f"{owner}|receipt:{_digest(namespace)}",
}
}
loop = configurable["__deepagents_js_cost_loop"]
asyncio.run_coroutine_threadsafe(
_put_receipt(saver, config, owner, amount, breakdown), loop
).result()


async def _put_receipt(
saver: BaseCheckpointSaver,
config: RunnableConfig,
owner: str,
amount: float,
breakdown: CostBreakdown | None,
) -> None:
"""Persist a node's first usage delta without replacing it on replay."""
if await saver.aget_tuple(config) is not None:
return
checkpoint = empty_checkpoint()
checkpoint["channel_values"] = {"cost_usd": amount}
if breakdown is not None:
checkpoint["channel_values"]["breakdown"] = breakdown
checkpoint["channel_versions"] = dict.fromkeys(
checkpoint["channel_values"], checkpoint["id"]
)
metadata: _ReceiptMetadata = {"js_cost_owner": owner}
await saver.aput(config, checkpoint, metadata, checkpoint["channel_versions"])


async def _receipt_total(
config: RunnableConfig, owner: str
) -> tuple[float, CostBreakdown] | None:
"""Return the sum of the latest local receipts owned by this eval."""
configurable = config.get("configurable", {})
saver = configurable.get("__pregel_checkpointer")
if not isinstance(saver, BaseCheckpointSaver):
return None
seen: set[str] = set()
total = 0.0
breakdown = _empty_cost_breakdown()
async for receipt in saver.alist(
{"configurable": {"thread_id": configurable["thread_id"]}},
filter={"js_cost_owner": owner},
):
namespace = receipt.config["configurable"]["checkpoint_ns"]
if namespace in seen:
continue
seen.add(namespace)
values = receipt.checkpoint["channel_values"]
amount = values.get("cost_usd")
if (
isinstance(amount, int | float)
and not isinstance(amount, bool)
and math.isfinite(amount)
and amount >= 0
):
total += amount
detail = values.get("breakdown")
if not isinstance(detail, Mapping):
# Keep legacy dollars without inventing requests or attribution.
detail = _empty_cost_breakdown(historical_complete=False)
detail["total_cost_usd"] = amount
breakdown = _merge_cost_breakdowns(breakdown, detail)
else:
breakdown["historical_complete"] = False
return (total, breakdown) if seen else None
6 changes: 4 additions & 2 deletions libs/code/deepagents_code/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from langchain.messages import ToolCall
from langchain_core.language_models import BaseChatModel
from langchain_core.messages import ToolMessage
from langchain_quickjs import PTCOption
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.prebuilt.tool_node import ToolCallRequest
from langgraph.pregel import Pregel
Expand Down Expand Up @@ -3153,7 +3154,8 @@ def _subagent_cli_middleware(
from langchain_core._api import ( # noqa: PLC2701 # re-exported in _api.__all__
suppress_langchain_beta_warning,
)
from langchain_quickjs import CodeInterpreterMiddleware, PTCOption

from deepagents_code._js_cost import CostAwareCodeInterpreterMiddleware

interpreter = interpreter_config or InterpreterConfig.from_resolver()
ptc_names = _resolve_ptc_option(
Expand All @@ -3170,7 +3172,7 @@ def _subagent_cli_middleware(
# and the warning is not actionable for users, so suppress it.
with suppress_langchain_beta_warning():
agent_middleware.append(
CodeInterpreterMiddleware(
CostAwareCodeInterpreterMiddleware(
tool_name="js_eval",
timeout=interpreter.timeout_seconds,
memory_limit=interpreter.memory_limit_mb * 1024 * 1024,
Expand Down
43 changes: 43 additions & 0 deletions libs/code/deepagents_code/cost_tracking.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@

from __future__ import annotations

import asyncio
import dataclasses
import errno
import json
Expand Down Expand Up @@ -2945,6 +2946,32 @@ def after_model( # ty: ignore[invalid-method-override]
logger.warning("Cost tracking failed to charge a model step", exc_info=True)
return None

async def aafter_model(
self, state: CostState, runtime: Runtime[ContextT]
) -> dict[str, Any] | None:
"""Return model cost updates after off-loop pricing and receipt writes."""
return await self._arun_cost_hook(self.after_model, state, runtime)

async def aafter_agent(
self, state: CostState, runtime: Runtime[ContextT]
) -> dict[str, Any] | None:
"""Return final cost updates after off-loop pricing and receipt writes."""
return await self._arun_cost_hook(self.after_agent, state, runtime)

@staticmethod
async def _arun_cost_hook(
hook: Callable[[CostState, Runtime[ContextT]], dict[str, Any] | None],
state: CostState,
runtime: Runtime[ContextT],
) -> dict[str, Any] | None:
"""Return the cost update only after any in-flight receipt write settles."""
invocation = asyncio.create_task(asyncio.to_thread(hook, state, runtime))
try:
return await asyncio.shield(invocation)
finally:
if not invocation.done():
await invocation

def after_agent( # ty: ignore[invalid-method-override]
self,
state: CostState,
Expand Down Expand Up @@ -3087,6 +3114,12 @@ def _charge(
)
remaining_transfers.pop(source_scope, None)
claimed_transfer = True
transferred_usd = delta_usd
transferred_breakdown = breakdown
# Receipts describe only this node's requests. Descendants have their
# own receipts, and completeness flags cannot be subtracted after merge.
delta_usd = 0.0
breakdown = _empty_cost_breakdown()
represented_message_ids: set[str] = set()
represented_count = 0
pricing_attempted = False
Expand Down Expand Up @@ -3187,6 +3220,16 @@ def _charge(
if estimate is not None:
delta_usd += estimate.total_cost_usd

if (
ensure_config()
.get("configurable", {})
.get("__deepagents_js_cost_owner")
):
from deepagents_code._js_cost import record_cost_receipt

record_cost_receipt(delta_usd, breakdown=breakdown)
delta_usd += transferred_usd
breakdown = _merge_cost_breakdowns(transferred_breakdown, breakdown)
has_breakdown = breakdown["request_count"] > 0
if not self._nested and (
delta_usd > 0 or pricing_attempted or has_breakdown
Expand Down
10 changes: 6 additions & 4 deletions libs/code/tests/unit_tests/test_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -5424,8 +5424,9 @@ def test_appends_interpreter_middleware_when_enabled(self, tmp_path: Path) -> No
)

_, kwargs = mock_create.call_args
middleware_types = [type(m) for m in kwargs["middleware"]]
assert CodeInterpreterMiddleware in middleware_types
assert any(
isinstance(m, CodeInterpreterMiddleware) for m in kwargs["middleware"]
)

def test_no_interpreter_middleware_when_disabled(self, tmp_path: Path) -> None:
from langchain_quickjs import CodeInterpreterMiddleware
Expand Down Expand Up @@ -5457,8 +5458,9 @@ def test_no_interpreter_middleware_when_disabled(self, tmp_path: Path) -> None:
)

_, kwargs = mock_create.call_args
middleware_types = [type(m) for m in kwargs["middleware"]]
assert CodeInterpreterMiddleware not in middleware_types
assert not any(
isinstance(m, CodeInterpreterMiddleware) for m in kwargs["middleware"]
)

def test_raises_when_sandbox_present(self, tmp_path: Path) -> None:
mock_settings = self._build_mock_settings(tmp_path)
Expand Down
Loading
Loading