feat(openai): support async tools (#40208)

This commit is contained in:
ccurme authored and GitHub committed 2026-09-04 15:49:14 -04:00
1 parent 1a3d81756f
commit 3f212e7a88
5 files changed
+217 -5

No files matched your search

@@ -735,7 +735,7 @@ def _convert_to_v1_from_responses(message: AIMessage) -> list[types.ContentBlock
tool_call_block["extras"]["item_id"] = block["id"]
if "index" in block:
tool_call_block["index"] = f"lc_tc_{block['index']}"
for extra_key in ("status", "namespace"):
for extra_key in ("status", "namespace", "async"):
if extra_key in block:
if "extras" not in tool_call_block:
tool_call_block["extras"] = {}
@@ -603,3 +603,53 @@ def test_convert_to_openai_data_block() -> None:
expected = {"type": "input_file", "file_id": "file-abc123"}
result = convert_to_openai_data_block(block, api="responses")
assert result == expected
def test_convert_to_v1_from_responses_async_tool_call() -> None:
"""Test that the `async` flag on a function call reaches `tool_call` extras."""
message = AIMessage(
[
{
"type": "function_call",
"call_id": "call_A",
"id": "fc_1",
"name": "lookup_price",
"arguments": '{"sku": "WIDGET"}',
"async": True,
},
{
"type": "function_call",
"call_id": "call_B",
"id": "fc_2",
"name": "get_time",
"arguments": "{}",
},
],
tool_calls=[
{
"type": "tool_call",
"id": "call_A",
"name": "lookup_price",
"args": {"sku": "WIDGET"},
},
{"type": "tool_call", "id": "call_B", "name": "get_time", "args": {}},
],
response_metadata={"model_provider": "openai"},
)
expected_content: list[types.ContentBlock] = [
{
"type": "tool_call",
"id": "call_A",
"name": "lookup_price",
"args": {"sku": "WIDGET"},
"extras": {"item_id": "fc_1", "async": True},
},
{
"type": "tool_call",
"id": "call_B",
"name": "get_time",
"args": {},
"extras": {"item_id": "fc_2"},
},
]
assert message.content_blocks == expected_content
@@ -452,7 +452,7 @@ def _convert_from_v1_to_responses(
tool_call["args"], separators=(",", ":")
)
if "extras" in block:
for extra_key in ("status", "namespace"):
for extra_key in ("status", "namespace", "async"):
if extra_key in block["extras"]:
new_block[extra_key] = block["extras"][extra_key]
new_content.append(new_block)
@@ -211,6 +211,8 @@ WellKnownTools = (
"apply_patch",
)
_TOOL_EXTRAS_PASSTHROUGH = ("defer_loading", "async")
def _convert_dict_to_message(_dict: Mapping[str, Any]) -> BaseMessage:
"""Convert a dictionary to a LangChain message.
@@ -2462,9 +2464,10 @@ class BaseChatOpenAI(BaseChatModel):
isinstance(original, BaseTool)
and hasattr(original, "extras")
and isinstance(original.extras, dict)
and "defer_loading" in original.extras
):
formatted["defer_loading"] = original.extras["defer_loading"]
for key in _TOOL_EXTRAS_PASSTHROUGH:
if key in original.extras:
formatted[key] = original.extras[key]
tool_names = []
for tool in formatted_tools:
if "function" in tool:
@@ -5093,7 +5096,12 @@ def _construct_lc_result_from_responses_api(
refusal_block["phase"] = phase
content_blocks.append(refusal_block)
elif output.type == "function_call":
content_blocks.append(output.model_dump(exclude_none=True, mode="json"))
function_call_block = output.model_dump(exclude_none=True, mode="json")
# The SDK names the reserved word `async_`; content blocks carry the
# wire name so it round-trips on the next request.
if "async_" in function_call_block:
function_call_block["async"] = function_call_block.pop("async_")
content_blocks.append(function_call_block)
try:
args = json.loads(output.arguments, strict=False)
error = None
@@ -5375,6 +5383,12 @@ def _convert_responses_chunk_to_generation_chunk(
}
if getattr(chunk.item, "namespace", None) is not None:
function_call_content["namespace"] = chunk.item.namespace
# SDKs expose the reserved word as `async_`; older ones keep it in extras.
async_flag = getattr(chunk.item, "async_", None)
if async_flag is None:
async_flag = getattr(chunk.item, "async", None)
if async_flag is not None:
function_call_content["async"] = async_flag
content.append(function_call_content)
elif chunk.type == "response.output_item.done" and chunk.item.type in (
"compaction",
@@ -5157,6 +5157,154 @@ def test_defer_loading_in_responses_api_payload() -> None:
assert {"type": "tool_search"} in result["tools"]
def test__construct_lc_result_from_responses_api_async_tool_call() -> None:
"""Test that `async` on a `function_call` item reaches `tool_call` extras."""
response = Response(
id="resp_123",
created_at=1234567890,
model=OPENAI_TEST_MODEL,
object="response",
parallel_tool_calls=True,
tools=[],
tool_choice="auto",
output=[
ResponseFunctionToolCall.model_validate(
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_A",
"name": "lookup_price",
"arguments": '{"sku": "WIDGET"}',
"async": True,
}
),
ResponseFunctionToolCall.model_validate(
{
"type": "function_call",
"id": "fc_2",
"call_id": "call_B",
"name": "get_time",
"arguments": "{}",
}
),
],
)
message = cast(
AIMessage,
_construct_lc_result_from_responses_api(response).generations[0].message,
)
extras: dict[Any, dict[str, Any]] = {
block.get("id"): cast(dict[str, Any], block.get("extras") or {})
for block in message.content_blocks
}
assert extras["call_A"]["async"] is True
assert "async" not in extras["call_B"]
# The flag is metadata only; both remain ordinary tool calls.
assert [tc["id"] for tc in message.tool_calls] == ["call_A", "call_B"]
def test_async_tool_call_round_trips_to_next_request() -> None:
"""Test that `async` survives response -> message -> next request."""
from langchain_core.tools import tool
@tool(extras={"async": True})
def lookup_price(sku: str) -> str:
"""Look up a price."""
return "1200"
response = Response(
id="resp_123",
created_at=1234567890,
model=OPENAI_TEST_MODEL,
object="response",
parallel_tool_calls=True,
tools=[],
tool_choice="auto",
output=[
ResponseFunctionToolCall.model_validate(
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_A",
"name": "lookup_price",
"arguments": '{"sku": "WIDGET"}',
"async": True,
}
)
],
)
for output_version in ("responses/v1", "v1"):
message = (
_construct_lc_result_from_responses_api(
response, output_version=output_version
)
.generations[0]
.message
)
llm = ChatOpenAI(
model=OPENAI_TEST_MODEL,
use_responses_api=True,
output_version=output_version,
)
bound = llm.bind_tools([lookup_price])
payload = bound._get_request_payload( # type: ignore[attr-defined]
[HumanMessage("price?"), message, HumanMessage("anything else?")],
**bound.kwargs, # type: ignore[attr-defined]
)
function_calls = [
item for item in payload["input"] if item.get("type") == "function_call"
]
# Without the flag the Responses API rejects the turn for a missing output.
assert function_calls[0]["async"] is True, output_version
def test_async_tool_from_extras_in_payload() -> None:
"""Test that `async` from `BaseTool.extras` reaches the Responses tool def."""
from langchain_core.tools import tool
@tool(extras={"async": True})
def lookup_price(sku: str) -> str:
"""Look up a price."""
return "1200"
@tool
def get_time() -> str:
"""Get the current time."""
return "14:05"
llm = ChatOpenAI(model=OPENAI_TEST_MODEL, use_responses_api=True)
bound = llm.bind_tools([lookup_price, get_time])
payload = bound._get_request_payload( # type: ignore[attr-defined]
"test",
**bound.kwargs, # type: ignore[attr-defined]
)
tools_by_name = {t["name"]: t for t in payload["tools"]}
assert tools_by_name["lookup_price"]["async"] is True
# Tools that don't opt in are unaffected.
assert "async" not in tools_by_name["get_time"]
def test_async_tool_raw_dict_passthrough() -> None:
"""Test that `async` on a raw tool dict is preserved."""
llm = ChatOpenAI(model=OPENAI_TEST_MODEL, use_responses_api=True)
raw_tool = {
"type": "function",
"name": "lookup_price",
"description": "Look up a price.",
"async": True,
"parameters": {
"type": "object",
"properties": {"sku": {"type": "string"}},
},
}
bound = llm.bind_tools([raw_tool])
payload = bound._get_request_payload( # type: ignore[attr-defined]
"test",
**bound.kwargs, # type: ignore[attr-defined]
)
assert payload["tools"][0]["async"] is True
def test_langsmith_gateway_true(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LANGSMITH_GATEWAY", "true")
llm = ChatOpenAI(model=OPENAI_TEST_MODEL, api_key=SecretStr("test"))