mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
feat(openai): support async tools (#40208)
This commit is contained in:
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"))
|
||||
|
||||
Reference in new issue
Block a user