mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 17:35:28 +03:00
release: v1.0.0 (#32567)
Co-authored-by: Mohammad Mohtashim <45242107+keenborder786@users.noreply.github.com> Co-authored-by: Caspar Broekhuizen <caspar@langchain.dev> Co-authored-by: ccurme <chester.curme@gmail.com> Co-authored-by: Christophe Bornet <cbornet@hotmail.com> Co-authored-by: Eugene Yurtsev <eyurtsev@gmail.com> Co-authored-by: Sadra Barikbin <sadraqazvin1@yahoo.com> Co-authored-by: Vadym Barda <vadim.barda@gmail.com>
This commit is contained in:
188 files changed
+23644
-17479
No files matched your search
@@ -136,14 +136,14 @@ def _get_configs_for_single_dir(job: str, dir_: str) -> List[Dict[str, str]]:
|
||||
if job == "codspeed":
|
||||
py_versions = ["3.12"] # 3.13 is not yet supported
|
||||
elif dir_ == "libs/core":
|
||||
py_versions = ["3.9", "3.10", "3.11", "3.12", "3.13"]
|
||||
py_versions = ["3.10", "3.11", "3.12", "3.13"]
|
||||
# custom logic for specific directories
|
||||
|
||||
elif dir_ in PY_312_MAX_PACKAGES:
|
||||
py_versions = ["3.9", "3.12"]
|
||||
py_versions = ["3.10", "3.12"]
|
||||
|
||||
elif dir_ == "libs/langchain" and job == "extended-tests":
|
||||
py_versions = ["3.9", "3.13"]
|
||||
py_versions = ["3.10", "3.13"]
|
||||
elif dir_ == "libs/langchain_v1":
|
||||
py_versions = ["3.10", "3.13"]
|
||||
elif dir_ in {"libs/cli"}:
|
||||
@@ -151,9 +151,9 @@ def _get_configs_for_single_dir(job: str, dir_: str) -> List[Dict[str, str]]:
|
||||
|
||||
elif dir_ == ".":
|
||||
# unable to install with 3.13 because tokenizers doesn't support 3.13 yet
|
||||
py_versions = ["3.9", "3.12"]
|
||||
py_versions = ["3.10", "3.12"]
|
||||
else:
|
||||
py_versions = ["3.9", "3.13"]
|
||||
py_versions = ["3.10", "3.13"]
|
||||
|
||||
return [{"working-directory": dir_, "python-version": py_v} for py_v in py_versions]
|
||||
|
||||
|
||||
@@ -431,7 +431,7 @@ jobs:
|
||||
git ls-remote --tags origin "langchain-${{ matrix.partner }}*" \
|
||||
| awk '{print $2}' \
|
||||
| sed 's|refs/tags/||' \
|
||||
| grep -E '[0-9]+\.[0-9]+\.[0-9]+$' \
|
||||
| grep -E '[0-9]+\.[0-9]+\.[0-9]+([a-zA-Z]+[0-9]+)?$' \
|
||||
| sort -Vr \
|
||||
| head -n 1
|
||||
)"
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
# Runs daily. Can also be triggered manually for immediate updates.
|
||||
|
||||
name: '⏰ Scheduled Integration Tests'
|
||||
run-name: "Run Integration Tests - ${{ inputs.working-directory-force || 'all libs' }} (Python ${{ inputs.python-version-force || '3.9, 3.11' }})"
|
||||
run-name: "Run Integration Tests - ${{ inputs.working-directory-force || 'all libs' }} (Python ${{ inputs.python-version-force || '3.10, 3.13' }})"
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
@@ -15,7 +15,7 @@ on:
|
||||
description: "From which folder this pipeline executes - defaults to all in matrix - example value: libs/partners/anthropic"
|
||||
python-version-force:
|
||||
type: string
|
||||
description: "Python version to use - defaults to 3.9 and 3.11 in matrix - example value: 3.9"
|
||||
description: "Python version to use - defaults to 3.10 and 3.13 in matrix - example value: 3.11"
|
||||
schedule:
|
||||
- cron: '0 13 * * *' # Runs daily at 1PM UTC (9AM EDT/6AM PDT)
|
||||
|
||||
@@ -46,9 +46,9 @@ jobs:
|
||||
PYTHON_VERSION_FORCE: ${{ github.event.inputs.python-version-force || '' }}
|
||||
run: |
|
||||
# echo "matrix=..." where matrix is a json formatted str with keys python-version and working-directory
|
||||
# python-version should default to 3.9 and 3.11, but is overridden to [PYTHON_VERSION_FORCE] if set
|
||||
# python-version should default to 3.10 and 3.13, but is overridden to [PYTHON_VERSION_FORCE] if set
|
||||
# working-directory should default to DEFAULT_LIBS, but is overridden to [WORKING_DIRECTORY_FORCE] if set
|
||||
python_version='["3.9", "3.11"]'
|
||||
python_version='["3.10", "3.13"]'
|
||||
working_directory="$DEFAULT_LIBS"
|
||||
if [ -n "$PYTHON_VERSION_FORCE" ]; then
|
||||
python_version="[\"$PYTHON_VERSION_FORCE\"]"
|
||||
|
||||
@@ -58,7 +58,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"execution_count": null,
|
||||
"id": "1fcf7b27-1cc3-420a-b920-0420b5892e20",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -102,7 +102,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -133,7 +133,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"execution_count": null,
|
||||
"id": "99d27f8f-ae78-48bc-9bf2-3cef35213ec7",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -163,7 +163,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -176,7 +176,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"execution_count": null,
|
||||
"id": "325fb4ca",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -198,7 +198,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -234,7 +234,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"execution_count": null,
|
||||
"id": "6c1455a9-699a-4702-a7e0-7f6eaec76a21",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -284,7 +284,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -312,7 +312,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"execution_count": null,
|
||||
"id": "55e1d937-3b22-4deb-b9f0-9e688f0609dc",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -342,7 +342,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -417,7 +417,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -443,7 +443,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"execution_count": null,
|
||||
"id": "83593b9d-a8d3-4c99-9dac-64e0a9d397cb",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -488,13 +488,13 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())\n",
|
||||
"print(response.text)\n",
|
||||
"response.usage_metadata"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"execution_count": null,
|
||||
"id": "9bbf578e-794a-4dc0-a469-78c876ccd4a3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -530,7 +530,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message, response, next_message])\n",
|
||||
"print(response.text())\n",
|
||||
"print(response.text)\n",
|
||||
"response.usage_metadata"
|
||||
]
|
||||
},
|
||||
@@ -600,7 +600,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"execution_count": null,
|
||||
"id": "ae076c9b-ff8f-461d-9349-250f396c9a25",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -641,7 +641,7 @@
|
||||
" ],\n",
|
||||
"}\n",
|
||||
"response = llm.invoke([message])\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -54,7 +54,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"execution_count": null,
|
||||
"id": "5df2e558-321d-4cf7-994e-2815ac37e704",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -75,7 +75,7 @@
|
||||
"\n",
|
||||
"chain = prompt | llm\n",
|
||||
"response = chain.invoke({\"image_url\": url})\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -117,7 +117,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"execution_count": null,
|
||||
"id": "25e4829e-0073-49a8-9669-9f43e5778383",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -144,7 +144,7 @@
|
||||
" \"cache_type\": \"ephemeral\",\n",
|
||||
" }\n",
|
||||
")\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -1593,7 +1593,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"execution_count": null,
|
||||
"id": "30a0af36-2327-4b1d-9ba5-e47cb72db0be",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1629,7 +1629,7 @@
|
||||
"response = llm_with_tools.invoke(\n",
|
||||
" \"There's a syntax error in my primes.py file. Can you help me fix it?\"\n",
|
||||
")\n",
|
||||
"print(response.text())\n",
|
||||
"print(response.text)\n",
|
||||
"response.tool_calls"
|
||||
]
|
||||
},
|
||||
|
||||
@@ -243,12 +243,12 @@
|
||||
"id": "0ef05abb-9c04-4dc3-995e-f857779644d5",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"You can filter to text using the [.text()](https://python.langchain.com/api_reference/core/messages/langchain_core.messages.ai.AIMessage.html#langchain_core.messages.ai.AIMessage.text) method on the output:"
|
||||
"You can filter to text using the [.text](https://python.langchain.com/api_reference/core/messages/langchain_core.messages.ai.AIMessage.html#langchain_core.messages.ai.AIMessage.text) property on the output:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"execution_count": null,
|
||||
"id": "2a4e743f-ea7d-4e5a-9b12-f9992362de8b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -262,7 +262,7 @@
|
||||
],
|
||||
"source": [
|
||||
"for chunk in llm.stream(messages):\n",
|
||||
" print(chunk.text(), end=\"|\")"
|
||||
" print(chunk.text, end=\"|\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -261,7 +261,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"execution_count": null,
|
||||
"id": "c5fac0e9-05a4-4fc1-a3b3-e5bbb24b971b",
|
||||
"metadata": {
|
||||
"colab": {
|
||||
@@ -286,7 +286,7 @@
|
||||
],
|
||||
"source": [
|
||||
"async for token in llm.astream(\"Hello, please explain how antibiotics work\"):\n",
|
||||
" print(token.text(), end=\"\")"
|
||||
" print(token.text, end=\"\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -814,7 +814,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"execution_count": null,
|
||||
"id": "1f758726-33ef-4c04-8a54-49adb783bbb3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -860,7 +860,7 @@
|
||||
"llm_with_tools = llm.bind_tools([tool])\n",
|
||||
"\n",
|
||||
"response = llm_with_tools.invoke(\"What is deep research by OpenAI?\")\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1151,7 +1151,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"execution_count": null,
|
||||
"id": "073f6010-6b0e-4db6-b2d3-7427c8dec95b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1167,7 +1167,7 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"response_2.text()"
|
||||
"response_2.text"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1198,7 +1198,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"execution_count": null,
|
||||
"id": "b6da5bd6-a44a-4c64-970b-30da26b003d6",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1214,7 +1214,7 @@
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"response_2.text()"
|
||||
"response_2.text"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1404,7 +1404,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"execution_count": null,
|
||||
"id": "51d3e4d3-ea78-426c-9205-aecb0937fca7",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1428,13 +1428,13 @@
|
||||
"messages = [{\"role\": \"user\", \"content\": first_query}]\n",
|
||||
"\n",
|
||||
"response = llm_with_tools.invoke(messages)\n",
|
||||
"response_text = response.text()\n",
|
||||
"response_text = response.text\n",
|
||||
"print(f\"{response_text[:100]}... {response_text[-100:]}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"execution_count": null,
|
||||
"id": "b248bedf-2050-4c17-a90e-3a26eeb1b055",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1460,7 +1460,7 @@
|
||||
" ]\n",
|
||||
")\n",
|
||||
"second_response = llm_with_tools.invoke(messages)\n",
|
||||
"print(second_response.text())"
|
||||
"print(second_response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1482,7 +1482,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"execution_count": null,
|
||||
"id": "009e541a-b372-410e-b9dd-608a8052ce09",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1502,12 +1502,12 @@
|
||||
" output_version=\"responses/v1\",\n",
|
||||
")\n",
|
||||
"response = llm.invoke(\"Hi, I'm Bob.\")\n",
|
||||
"print(response.text())"
|
||||
"print(response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"execution_count": null,
|
||||
"id": "393a443a-4c5f-4a07-bc0e-c76e529b35e3",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1524,7 +1524,7 @@
|
||||
" \"What is my name?\",\n",
|
||||
" previous_response_id=response.response_metadata[\"id\"],\n",
|
||||
")\n",
|
||||
"print(second_response.text())"
|
||||
"print(second_response.text)"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1589,7 +1589,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"execution_count": null,
|
||||
"id": "8d322f3a-0732-45ab-ac95-dfd4596e0d85",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -1616,7 +1616,7 @@
|
||||
"response = llm.invoke(\"What is 3^3?\")\n",
|
||||
"\n",
|
||||
"# Output\n",
|
||||
"response.text()"
|
||||
"response.text"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -302,7 +302,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 7,
|
||||
"execution_count": null,
|
||||
"id": "c96c960b",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -320,7 +320,7 @@
|
||||
"source": [
|
||||
"query = \"Hi!\"\n",
|
||||
"response = model.invoke([{\"role\": \"user\", \"content\": query}])\n",
|
||||
"response.text()"
|
||||
"response.text"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -351,7 +351,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"execution_count": null,
|
||||
"id": "b6a7e925",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -371,7 +371,7 @@
|
||||
"query = \"Hi!\"\n",
|
||||
"response = model_with_tools.invoke([{\"role\": \"user\", \"content\": query}])\n",
|
||||
"\n",
|
||||
"print(f\"Message content: {response.text()}\\n\")\n",
|
||||
"print(f\"Message content: {response.text}\\n\")\n",
|
||||
"print(f\"Tool calls: {response.tool_calls}\")"
|
||||
]
|
||||
},
|
||||
@@ -385,7 +385,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 16,
|
||||
"execution_count": null,
|
||||
"id": "688b465d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -403,7 +403,7 @@
|
||||
"query = \"Search for the weather in SF\"\n",
|
||||
"response = model_with_tools.invoke([{\"role\": \"user\", \"content\": query}])\n",
|
||||
"\n",
|
||||
"print(f\"Message content: {response.text()}\\n\")\n",
|
||||
"print(f\"Message content: {response.text}\\n\")\n",
|
||||
"print(f\"Tool calls: {response.tool_calls}\")"
|
||||
]
|
||||
},
|
||||
@@ -615,19 +615,12 @@
|
||||
"## Streaming tokens\n",
|
||||
"\n",
|
||||
"In addition to streaming back messages, it is also useful to stream back tokens.\n",
|
||||
"We can do this by specifying `stream_mode=\"messages\"`.\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"::: note\n",
|
||||
"\n",
|
||||
"Below we use `message.text()`, which requires `langchain-core>=0.3.37`.\n",
|
||||
"\n",
|
||||
":::"
|
||||
"We can do this by specifying `stream_mode=\"messages\"`."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 18,
|
||||
"execution_count": null,
|
||||
"id": "63198158-380e-43a3-a2ad-d4288949c1d4",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -651,7 +644,7 @@
|
||||
"for step, metadata in agent_executor.stream(\n",
|
||||
" {\"messages\": [input_message]}, config, stream_mode=\"messages\"\n",
|
||||
"):\n",
|
||||
" if metadata[\"langgraph_node\"] == \"agent\" and (text := step.text()):\n",
|
||||
" if metadata[\"langgraph_node\"] == \"agent\" and (text := step.text):\n",
|
||||
" print(text, end=\"|\")"
|
||||
]
|
||||
},
|
||||
|
||||
@@ -102,7 +102,7 @@ def _is_relevant_import(module: str) -> bool:
|
||||
"langchain",
|
||||
"langchain_core",
|
||||
"langchain_community",
|
||||
"langchain_experimental",
|
||||
# "langchain_experimental",
|
||||
"langchain_text_splitters",
|
||||
]
|
||||
return module.split(".")[0] in recognized_packages
|
||||
|
||||
Generated
+17
-18
@@ -277,6 +277,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/2a/39/e50c7c3a983047577ee07d2a9e53faf5a69493943ec3f6a384bdc792deb2/httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad", size = 73517, upload-time = "2024-12-06T15:37:21.509Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "httpx-sse"
|
||||
version = "0.4.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/6e/fa/66bd985dd0b7c109a3bcb89272ee0bfb7e2b4d06309ad7b38ff866734b2a/httpx_sse-0.4.1.tar.gz", hash = "sha256:8f44d34414bc7b21bf3602713005c5df4917884f76072479b21f68befa4ea26e", size = 12998, upload-time = "2025-06-24T13:21:05.71Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/25/0a/6269e3473b09aed2dab8aa1a600c70f31f00ae1349bee30658f7e358a159/httpx_sse-0.4.1-py3-none-any.whl", hash = "sha256:cba42174344c3a5b06f255ce65b350880f962d99ead85e776f23c6618a377a37", size = 8054, upload-time = "2025-06-24T13:21:04.772Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "idna"
|
||||
version = "3.10"
|
||||
@@ -335,31 +344,20 @@ dependencies = [
|
||||
requires-dist = [
|
||||
{ name = "async-timeout", marker = "python_full_version < '3.11'", specifier = ">=4.0.0,<5.0.0" },
|
||||
{ name = "langchain-anthropic", marker = "extra == 'anthropic'" },
|
||||
{ name = "langchain-aws", marker = "extra == 'aws'" },
|
||||
{ name = "langchain-azure-ai", marker = "extra == 'azure-ai'" },
|
||||
{ name = "langchain-cohere", marker = "extra == 'cohere'" },
|
||||
{ name = "langchain-community", marker = "extra == 'community'" },
|
||||
{ name = "langchain-core", editable = "../core" },
|
||||
{ name = "langchain-deepseek", marker = "extra == 'deepseek'" },
|
||||
{ name = "langchain-fireworks", marker = "extra == 'fireworks'" },
|
||||
{ name = "langchain-google-genai", marker = "extra == 'google-genai'" },
|
||||
{ name = "langchain-google-vertexai", marker = "extra == 'google-vertexai'" },
|
||||
{ name = "langchain-groq", marker = "extra == 'groq'" },
|
||||
{ name = "langchain-huggingface", marker = "extra == 'huggingface'" },
|
||||
{ name = "langchain-mistralai", marker = "extra == 'mistralai'" },
|
||||
{ name = "langchain-ollama", marker = "extra == 'ollama'" },
|
||||
{ name = "langchain-openai", marker = "extra == 'openai'", editable = "../partners/openai" },
|
||||
{ name = "langchain-perplexity", marker = "extra == 'perplexity'" },
|
||||
{ name = "langchain-text-splitters", editable = "../text-splitters" },
|
||||
{ name = "langchain-together", marker = "extra == 'together'" },
|
||||
{ name = "langchain-xai", marker = "extra == 'xai'" },
|
||||
{ name = "langsmith", specifier = ">=0.1.17,<1.0.0" },
|
||||
{ name = "pydantic", specifier = ">=2.7.4,<3.0.0" },
|
||||
{ name = "pyyaml", specifier = ">=5.3.0,<7.0.0" },
|
||||
{ name = "requests", specifier = ">=2.0.0,<3.0.0" },
|
||||
{ name = "sqlalchemy", specifier = ">=1.4.0,<3.0.0" },
|
||||
]
|
||||
provides-extras = ["community", "anthropic", "openai", "azure-ai", "cohere", "google-vertexai", "google-genai", "fireworks", "ollama", "together", "mistralai", "huggingface", "groq", "aws", "deepseek", "xai", "perplexity"]
|
||||
provides-extras = ["community", "anthropic", "openai", "google-vertexai", "google-genai", "together"]
|
||||
|
||||
[package.metadata.requires-dev]
|
||||
dev = [
|
||||
@@ -485,7 +483,7 @@ typing = [{ name = "langchain", editable = "../langchain" }]
|
||||
|
||||
[[package]]
|
||||
name = "langchain-core"
|
||||
version = "0.3.76"
|
||||
version = "1.0.0a5"
|
||||
source = { editable = "../core" }
|
||||
dependencies = [
|
||||
{ name = "jsonpatch" },
|
||||
@@ -586,28 +584,29 @@ typing = [
|
||||
{ name = "beautifulsoup4", specifier = ">=4.13.5,<5.0.0" },
|
||||
{ name = "lxml-stubs", specifier = ">=0.5.1,<1.0.0" },
|
||||
{ name = "mypy", specifier = ">=1.18.1,<1.19.0" },
|
||||
{ name = "tiktoken", specifier = ">=0.11.0,<1.0.0" },
|
||||
{ name = "tiktoken", specifier = ">=0.8.0,<1.0.0" },
|
||||
{ name = "types-requests", specifier = ">=2.31.0.20240218,<3.0.0.0" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "langserve"
|
||||
version = "0.3.2"
|
||||
version = "0.0.51"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "httpx" },
|
||||
{ name = "langchain-core" },
|
||||
{ name = "langchain" },
|
||||
{ name = "orjson" },
|
||||
{ name = "pydantic" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/5a/fb/86e1f5049fb3593743f0fb049c4991f4984020cda00b830ae31f2c47b46b/langserve-0.3.2.tar.gz", hash = "sha256:134b78b1d897c6bcd1fb8a6258e30cf0fb318294505e4ea59c2bea72fa152129", size = 1141270, upload-time = "2025-09-17T20:01:22.183Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/06/af/243c8a6ad0efee30186fba5a05a68b4bd9553d3662f946e2e8302cb4a141/langserve-0.0.51.tar.gz", hash = "sha256:036c0104c512bcc2c2406ae089ef9e7e718c32c39ebf6dcb2212f168c7d09816", size = 1135441, upload-time = "2024-03-12T06:16:32.374Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/bd/f0/193c34bf61e1dee8bd637dbeddcc644c46d14e8b03068792ca60b1909bc1/langserve-0.3.2-py3-none-any.whl", hash = "sha256:d9c4cd19d12f6362b82ceecb10357b339b3640a858b9bc30643d5f8a0a036bce", size = 1173213, upload-time = "2025-09-17T20:01:20.603Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/b3/49/5b407071f7ea5a861b3f4c3ed2f034cdafb75db1554bbe1a256a092d2669/langserve-0.0.51-py3-none-any.whl", hash = "sha256:e735eef2b6fde7e1514f4be8234b9f0727283e639822ca9c25e8ccc2d24e8492", size = 1167759, upload-time = "2024-03-12T06:16:30.099Z" },
|
||||
]
|
||||
|
||||
[package.optional-dependencies]
|
||||
all = [
|
||||
{ name = "fastapi" },
|
||||
{ name = "httpx-sse" },
|
||||
{ name = "sse-starlette" },
|
||||
]
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ from collections.abc import Generator
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
ParamSpec,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
@@ -25,7 +26,6 @@ from typing import (
|
||||
|
||||
from pydantic.fields import FieldInfo
|
||||
from pydantic.v1.fields import FieldInfo as FieldInfoV1
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
from langchain_core._api.internal import is_caller_internal
|
||||
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""Some **beta** features that are not yet ready for production."""
|
||||
@@ -1 +0,0 @@
|
||||
"""Runnables."""
|
||||
@@ -1,447 +0,0 @@
|
||||
"""Context management for runnables."""
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
from collections import defaultdict
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from functools import partial
|
||||
from itertools import groupby
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Optional,
|
||||
TypeVar,
|
||||
Union,
|
||||
)
|
||||
|
||||
from pydantic import ConfigDict
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain_core._api.beta_decorator import beta
|
||||
from langchain_core.runnables.base import (
|
||||
Runnable,
|
||||
RunnableSerializable,
|
||||
coerce_to_runnable,
|
||||
)
|
||||
from langchain_core.runnables.config import RunnableConfig, ensure_config, patch_config
|
||||
from langchain_core.runnables.utils import ConfigurableFieldSpec, Input, Output
|
||||
|
||||
T = TypeVar("T")
|
||||
Values = dict[Union[asyncio.Event, threading.Event], Any]
|
||||
CONTEXT_CONFIG_PREFIX = "__context__/"
|
||||
CONTEXT_CONFIG_SUFFIX_GET = "/get"
|
||||
CONTEXT_CONFIG_SUFFIX_SET = "/set"
|
||||
|
||||
|
||||
async def _asetter(done: asyncio.Event, values: Values, value: T) -> T:
|
||||
values[done] = value
|
||||
done.set()
|
||||
return value
|
||||
|
||||
|
||||
async def _agetter(done: asyncio.Event, values: Values) -> Any:
|
||||
await done.wait()
|
||||
return values[done]
|
||||
|
||||
|
||||
def _setter(done: threading.Event, values: Values, value: T) -> T:
|
||||
values[done] = value
|
||||
done.set()
|
||||
return value
|
||||
|
||||
|
||||
def _getter(done: threading.Event, values: Values) -> Any:
|
||||
done.wait()
|
||||
return values[done]
|
||||
|
||||
|
||||
def _key_from_id(id_: str) -> str:
|
||||
wout_prefix = id_.split(CONTEXT_CONFIG_PREFIX, maxsplit=1)[1]
|
||||
if wout_prefix.endswith(CONTEXT_CONFIG_SUFFIX_GET):
|
||||
return wout_prefix[: -len(CONTEXT_CONFIG_SUFFIX_GET)]
|
||||
if wout_prefix.endswith(CONTEXT_CONFIG_SUFFIX_SET):
|
||||
return wout_prefix[: -len(CONTEXT_CONFIG_SUFFIX_SET)]
|
||||
msg = f"Invalid context config id {id_}"
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def _config_with_context(
|
||||
config: RunnableConfig,
|
||||
steps: list[Runnable],
|
||||
setter: Callable,
|
||||
getter: Callable,
|
||||
event_cls: Union[type[threading.Event], type[asyncio.Event]],
|
||||
) -> RunnableConfig:
|
||||
if any(k.startswith(CONTEXT_CONFIG_PREFIX) for k in config.get("configurable", {})):
|
||||
return config
|
||||
|
||||
context_specs = [
|
||||
(spec, i)
|
||||
for i, step in enumerate(steps)
|
||||
for spec in step.config_specs
|
||||
if spec.id.startswith(CONTEXT_CONFIG_PREFIX)
|
||||
]
|
||||
grouped_by_key = {
|
||||
key: list(group)
|
||||
for key, group in groupby(
|
||||
sorted(context_specs, key=lambda s: s[0].id),
|
||||
key=lambda s: _key_from_id(s[0].id),
|
||||
)
|
||||
}
|
||||
deps_by_key = {
|
||||
key: {
|
||||
_key_from_id(dep) for spec in group for dep in (spec[0].dependencies or [])
|
||||
}
|
||||
for key, group in grouped_by_key.items()
|
||||
}
|
||||
|
||||
values: Values = {}
|
||||
events: defaultdict[str, Union[asyncio.Event, threading.Event]] = defaultdict(
|
||||
event_cls
|
||||
)
|
||||
context_funcs: dict[str, Callable[[], Any]] = {}
|
||||
for key, group in grouped_by_key.items():
|
||||
getters = [s for s in group if s[0].id.endswith(CONTEXT_CONFIG_SUFFIX_GET)]
|
||||
setters = [s for s in group if s[0].id.endswith(CONTEXT_CONFIG_SUFFIX_SET)]
|
||||
|
||||
for dep in deps_by_key[key]:
|
||||
if key in deps_by_key[dep]:
|
||||
msg = f"Deadlock detected between context keys {key} and {dep}"
|
||||
raise ValueError(msg)
|
||||
if len(setters) != 1:
|
||||
msg = f"Expected exactly one setter for context key {key}"
|
||||
raise ValueError(msg)
|
||||
setter_idx = setters[0][1]
|
||||
if any(getter_idx < setter_idx for _, getter_idx in getters):
|
||||
msg = f"Context setter for key {key} must be defined after all getters."
|
||||
raise ValueError(msg)
|
||||
|
||||
if getters:
|
||||
context_funcs[getters[0][0].id] = partial(getter, events[key], values)
|
||||
context_funcs[setters[0][0].id] = partial(setter, events[key], values)
|
||||
|
||||
return patch_config(config, configurable=context_funcs)
|
||||
|
||||
|
||||
def aconfig_with_context(
|
||||
config: RunnableConfig,
|
||||
steps: list[Runnable],
|
||||
) -> RunnableConfig:
|
||||
"""Asynchronously patch a runnable config with context getters and setters.
|
||||
|
||||
Args:
|
||||
config: The runnable config.
|
||||
steps: The runnable steps.
|
||||
|
||||
Returns:
|
||||
The patched runnable config.
|
||||
"""
|
||||
return _config_with_context(config, steps, _asetter, _agetter, asyncio.Event)
|
||||
|
||||
|
||||
def config_with_context(
|
||||
config: RunnableConfig,
|
||||
steps: list[Runnable],
|
||||
) -> RunnableConfig:
|
||||
"""Patch a runnable config with context getters and setters.
|
||||
|
||||
Args:
|
||||
config: The runnable config.
|
||||
steps: The runnable steps.
|
||||
|
||||
Returns:
|
||||
The patched runnable config.
|
||||
"""
|
||||
return _config_with_context(config, steps, _setter, _getter, threading.Event)
|
||||
|
||||
|
||||
@beta()
|
||||
class ContextGet(RunnableSerializable):
|
||||
"""Get a context value."""
|
||||
|
||||
prefix: str = ""
|
||||
|
||||
key: Union[str, list[str]]
|
||||
|
||||
@override
|
||||
def __str__(self) -> str:
|
||||
return f"ContextGet({_print_keys(self.key)})"
|
||||
|
||||
@property
|
||||
def ids(self) -> list[str]:
|
||||
"""The context getter ids."""
|
||||
prefix = self.prefix + "/" if self.prefix else ""
|
||||
keys = self.key if isinstance(self.key, list) else [self.key]
|
||||
return [
|
||||
f"{CONTEXT_CONFIG_PREFIX}{prefix}{k}{CONTEXT_CONFIG_SUFFIX_GET}"
|
||||
for k in keys
|
||||
]
|
||||
|
||||
@property
|
||||
@override
|
||||
def config_specs(self) -> list[ConfigurableFieldSpec]:
|
||||
return super().config_specs + [
|
||||
ConfigurableFieldSpec(
|
||||
id=id_,
|
||||
annotation=Callable[[], Any],
|
||||
)
|
||||
for id_ in self.ids
|
||||
]
|
||||
|
||||
@override
|
||||
def invoke(
|
||||
self, input: Any, config: Optional[RunnableConfig] = None, **kwargs: Any
|
||||
) -> Any:
|
||||
config = ensure_config(config)
|
||||
configurable = config.get("configurable", {})
|
||||
if isinstance(self.key, list):
|
||||
return {key: configurable[id_]() for key, id_ in zip(self.key, self.ids)}
|
||||
return configurable[self.ids[0]]()
|
||||
|
||||
@override
|
||||
async def ainvoke(
|
||||
self, input: Any, config: Optional[RunnableConfig] = None, **kwargs: Any
|
||||
) -> Any:
|
||||
config = ensure_config(config)
|
||||
configurable = config.get("configurable", {})
|
||||
if isinstance(self.key, list):
|
||||
values = await asyncio.gather(*(configurable[id_]() for id_ in self.ids))
|
||||
return dict(zip(self.key, values))
|
||||
return await configurable[self.ids[0]]()
|
||||
|
||||
|
||||
SetValue = Union[
|
||||
Runnable[Input, Output],
|
||||
Callable[[Input], Output],
|
||||
Callable[[Input], Awaitable[Output]],
|
||||
Any,
|
||||
]
|
||||
|
||||
|
||||
def _coerce_set_value(value: SetValue) -> Runnable[Input, Output]:
|
||||
if not isinstance(value, Runnable) and not callable(value):
|
||||
return coerce_to_runnable(lambda _: value)
|
||||
return coerce_to_runnable(value)
|
||||
|
||||
|
||||
@beta()
|
||||
class ContextSet(RunnableSerializable):
|
||||
"""Set a context value."""
|
||||
|
||||
prefix: str = ""
|
||||
|
||||
keys: Mapping[str, Optional[Runnable]]
|
||||
|
||||
model_config = ConfigDict(
|
||||
arbitrary_types_allowed=True,
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
key: Optional[str] = None,
|
||||
value: Optional[SetValue] = None,
|
||||
prefix: str = "",
|
||||
**kwargs: SetValue,
|
||||
):
|
||||
"""Create a context setter.
|
||||
|
||||
Args:
|
||||
key: The context setter key.
|
||||
value: The context setter value.
|
||||
prefix: The context setter prefix.
|
||||
**kwargs: Additional context setter key-value pairs.
|
||||
"""
|
||||
if key is not None:
|
||||
kwargs[key] = value
|
||||
super().__init__(
|
||||
keys={
|
||||
k: _coerce_set_value(v) if v is not None else None
|
||||
for k, v in kwargs.items()
|
||||
},
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
@override
|
||||
def __str__(self) -> str:
|
||||
return f"ContextSet({_print_keys(list(self.keys.keys()))})"
|
||||
|
||||
@property
|
||||
def ids(self) -> list[str]:
|
||||
"""The context setter ids."""
|
||||
prefix = self.prefix + "/" if self.prefix else ""
|
||||
return [
|
||||
f"{CONTEXT_CONFIG_PREFIX}{prefix}{key}{CONTEXT_CONFIG_SUFFIX_SET}"
|
||||
for key in self.keys
|
||||
]
|
||||
|
||||
@property
|
||||
@override
|
||||
def config_specs(self) -> list[ConfigurableFieldSpec]:
|
||||
mapper_config_specs = [
|
||||
s
|
||||
for mapper in self.keys.values()
|
||||
if mapper is not None
|
||||
for s in mapper.config_specs
|
||||
]
|
||||
for spec in mapper_config_specs:
|
||||
if spec.id.endswith(CONTEXT_CONFIG_SUFFIX_GET):
|
||||
getter_key = spec.id.split("/")[1]
|
||||
if getter_key in self.keys:
|
||||
msg = f"Circular reference in context setter for key {getter_key}"
|
||||
raise ValueError(msg)
|
||||
return super().config_specs + [
|
||||
ConfigurableFieldSpec(
|
||||
id=id_,
|
||||
annotation=Callable[[], Any],
|
||||
)
|
||||
for id_ in self.ids
|
||||
]
|
||||
|
||||
@override
|
||||
def invoke(
|
||||
self, input: Any, config: Optional[RunnableConfig] = None, **kwargs: Any
|
||||
) -> Any:
|
||||
config = ensure_config(config)
|
||||
configurable = config.get("configurable", {})
|
||||
for id_, mapper in zip(self.ids, self.keys.values()):
|
||||
if mapper is not None:
|
||||
configurable[id_](mapper.invoke(input, config))
|
||||
else:
|
||||
configurable[id_](input)
|
||||
return input
|
||||
|
||||
@override
|
||||
async def ainvoke(
|
||||
self, input: Any, config: Optional[RunnableConfig] = None, **kwargs: Any
|
||||
) -> Any:
|
||||
config = ensure_config(config)
|
||||
configurable = config.get("configurable", {})
|
||||
for id_, mapper in zip(self.ids, self.keys.values()):
|
||||
if mapper is not None:
|
||||
await configurable[id_](await mapper.ainvoke(input, config))
|
||||
else:
|
||||
await configurable[id_](input)
|
||||
return input
|
||||
|
||||
|
||||
class Context:
|
||||
"""Context for a runnable.
|
||||
|
||||
The `Context` class provides methods for creating context scopes,
|
||||
getters, and setters within a runnable. It allows for managing
|
||||
and accessing contextual information throughout the execution
|
||||
of a program.
|
||||
|
||||
Example:
|
||||
.. code-block:: python
|
||||
|
||||
from langchain_core.beta.runnables.context import Context
|
||||
from langchain_core.runnables.passthrough import RunnablePassthrough
|
||||
from langchain_core.prompts.prompt import PromptTemplate
|
||||
from langchain_core.output_parsers.string import StrOutputParser
|
||||
from tests.unit_tests.fake.llm import FakeListLLM
|
||||
|
||||
chain = (
|
||||
Context.setter("input")
|
||||
| {
|
||||
"context": RunnablePassthrough() | Context.setter("context"),
|
||||
"question": RunnablePassthrough(),
|
||||
}
|
||||
| PromptTemplate.from_template("{context} {question}")
|
||||
| FakeListLLM(responses=["hello"])
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"result": RunnablePassthrough(),
|
||||
"context": Context.getter("context"),
|
||||
"input": Context.getter("input"),
|
||||
}
|
||||
)
|
||||
|
||||
# Use the chain
|
||||
output = chain.invoke("What's your name?")
|
||||
print(output["result"]) # Output: "hello"
|
||||
print(output["context"]) # Output: "What's your name?"
|
||||
print(output["input"]) # Output: "What's your name?
|
||||
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def create_scope(scope: str, /) -> "PrefixContext":
|
||||
"""Create a context scope.
|
||||
|
||||
Args:
|
||||
scope: The scope.
|
||||
|
||||
Returns:
|
||||
The context scope.
|
||||
"""
|
||||
return PrefixContext(prefix=scope)
|
||||
|
||||
@staticmethod
|
||||
def getter(key: Union[str, list[str]], /) -> ContextGet:
|
||||
"""Return a context getter.
|
||||
|
||||
Args:
|
||||
key: The context getter key.
|
||||
"""
|
||||
return ContextGet(key=key)
|
||||
|
||||
@staticmethod
|
||||
def setter(
|
||||
_key: Optional[str] = None,
|
||||
_value: Optional[SetValue] = None,
|
||||
/,
|
||||
**kwargs: SetValue,
|
||||
) -> ContextSet:
|
||||
"""Return a context setter.
|
||||
|
||||
Args:
|
||||
_key: The context setter key.
|
||||
_value: The context setter value.
|
||||
**kwargs: Additional context setter key-value pairs.
|
||||
"""
|
||||
return ContextSet(_key, _value, prefix="", **kwargs)
|
||||
|
||||
|
||||
class PrefixContext:
|
||||
"""Context for a runnable with a prefix."""
|
||||
|
||||
prefix: str = ""
|
||||
|
||||
def __init__(self, prefix: str = ""):
|
||||
"""Create a prefix context.
|
||||
|
||||
Args:
|
||||
prefix: The prefix.
|
||||
"""
|
||||
self.prefix = prefix
|
||||
|
||||
def getter(self, key: Union[str, list[str]], /) -> ContextGet:
|
||||
"""Return a prefixed context getter.
|
||||
|
||||
Args:
|
||||
key: The context getter key.
|
||||
"""
|
||||
return ContextGet(key=key, prefix=self.prefix)
|
||||
|
||||
def setter(
|
||||
self,
|
||||
_key: Optional[str] = None,
|
||||
_value: Optional[SetValue] = None,
|
||||
/,
|
||||
**kwargs: SetValue,
|
||||
) -> ContextSet:
|
||||
"""Return a prefixed context setter.
|
||||
|
||||
Args:
|
||||
_key: The context setter key.
|
||||
_value: The context setter value.
|
||||
**kwargs: Additional context setter key-value pairs.
|
||||
"""
|
||||
return ContextSet(_key, _value, prefix=self.prefix, **kwargs)
|
||||
|
||||
|
||||
def _print_keys(keys: Union[str, Sequence[str]]) -> str:
|
||||
if isinstance(keys, str):
|
||||
return f"'{keys}'"
|
||||
return ", ".join(f"'{k}'" for k in keys)
|
||||
@@ -45,6 +45,7 @@ https://python.langchain.com/docs/how_to/custom_llm/
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from langchain_core._import_utils import import_attr
|
||||
from langchain_core.language_models._utils import is_openai_data_block
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.language_models.base import (
|
||||
@@ -85,6 +86,7 @@ __all__ = (
|
||||
"ParrotFakeChatModel",
|
||||
"SimpleChatModel",
|
||||
"get_tokenizer",
|
||||
"is_openai_data_block",
|
||||
)
|
||||
|
||||
_dynamic_imports = {
|
||||
@@ -104,6 +106,7 @@ _dynamic_imports = {
|
||||
"ParrotFakeChatModel": "fake_chat_models",
|
||||
"LLM": "llms",
|
||||
"BaseLLM": "llms",
|
||||
"is_openai_data_block": "_utils",
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -1,13 +1,49 @@
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
from typing import Optional
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Literal,
|
||||
Optional,
|
||||
TypedDict,
|
||||
TypeVar,
|
||||
Union,
|
||||
)
|
||||
|
||||
from langchain_core.messages import BaseMessage
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.messages import BaseMessage
|
||||
from langchain_core.messages.content import (
|
||||
ContentBlock,
|
||||
)
|
||||
|
||||
|
||||
def _is_openai_data_block(block: dict) -> bool:
|
||||
"""Check if the block contains multimodal data in OpenAI Chat Completions format."""
|
||||
def is_openai_data_block(
|
||||
block: dict, filter_: Union[Literal["image", "audio", "file"], None] = None
|
||||
) -> bool:
|
||||
"""Check whether a block contains multimodal data in OpenAI Chat Completions format.
|
||||
|
||||
Supports both data and ID-style blocks (e.g. ``'file_data'`` and ``'file_id'``)
|
||||
|
||||
If additional keys are present, they are ignored / will not affect outcome as long
|
||||
as the required keys are present and valid.
|
||||
|
||||
Args:
|
||||
block: The content block to check.
|
||||
filter_: If provided, only return True for blocks matching this specific type.
|
||||
- "image": Only match image_url blocks
|
||||
- "audio": Only match input_audio blocks
|
||||
- "file": Only match file blocks
|
||||
If None, match any valid OpenAI data block type. Note that this means that
|
||||
if the block has a valid OpenAI data type but the filter_ is set to a
|
||||
different type, this function will return False.
|
||||
|
||||
Returns:
|
||||
True if the block is a valid OpenAI data block and matches the filter_
|
||||
(if provided).
|
||||
|
||||
"""
|
||||
if block.get("type") == "image_url":
|
||||
if filter_ is not None and filter_ != "image":
|
||||
return False
|
||||
if (
|
||||
(set(block.keys()) <= {"type", "image_url", "detail"})
|
||||
and (image_url := block.get("image_url"))
|
||||
@@ -15,29 +51,47 @@ def _is_openai_data_block(block: dict) -> bool:
|
||||
):
|
||||
url = image_url.get("url")
|
||||
if isinstance(url, str):
|
||||
# Required per OpenAI spec
|
||||
return True
|
||||
# Ignore `'detail'` since it's optional and specific to OpenAI
|
||||
|
||||
elif block.get("type") == "input_audio":
|
||||
if filter_ is not None and filter_ != "audio":
|
||||
return False
|
||||
if (audio := block.get("input_audio")) and isinstance(audio, dict):
|
||||
audio_data = audio.get("data")
|
||||
audio_format = audio.get("format")
|
||||
# Both required per OpenAI spec
|
||||
if isinstance(audio_data, str) and isinstance(audio_format, str):
|
||||
return True
|
||||
|
||||
elif block.get("type") == "file":
|
||||
if filter_ is not None and filter_ != "file":
|
||||
return False
|
||||
if (file := block.get("file")) and isinstance(file, dict):
|
||||
file_data = file.get("file_data")
|
||||
if isinstance(file_data, str):
|
||||
return True
|
||||
|
||||
elif block.get("type") == "input_audio":
|
||||
if (input_audio := block.get("input_audio")) and isinstance(input_audio, dict):
|
||||
audio_data = input_audio.get("data")
|
||||
audio_format = input_audio.get("format")
|
||||
if isinstance(audio_data, str) and isinstance(audio_format, str):
|
||||
file_id = file.get("file_id")
|
||||
# Files can be either base64-encoded or pre-uploaded with an ID
|
||||
if isinstance(file_data, str) or isinstance(file_id, str):
|
||||
return True
|
||||
|
||||
else:
|
||||
return False
|
||||
|
||||
# Has no `'type'` key
|
||||
return False
|
||||
|
||||
|
||||
def _parse_data_uri(uri: str) -> Optional[dict]:
|
||||
"""Parse a data URI into its components. If parsing fails, return None.
|
||||
class ParsedDataUri(TypedDict):
|
||||
source_type: Literal["base64"]
|
||||
data: str
|
||||
mime_type: str
|
||||
|
||||
|
||||
def _parse_data_uri(uri: str) -> Optional[ParsedDataUri]:
|
||||
"""Parse a data URI into its components.
|
||||
|
||||
If parsing fails, return None. If either MIME type or data is missing, return None.
|
||||
|
||||
Example:
|
||||
|
||||
@@ -57,84 +111,219 @@ def _parse_data_uri(uri: str) -> Optional[dict]:
|
||||
match = re.match(regex, uri)
|
||||
if match is None:
|
||||
return None
|
||||
|
||||
mime_type = match.group("mime_type")
|
||||
data = match.group("data")
|
||||
if not mime_type or not data:
|
||||
return None
|
||||
|
||||
return {
|
||||
"source_type": "base64",
|
||||
"data": match.group("data"),
|
||||
"mime_type": match.group("mime_type"),
|
||||
"data": data,
|
||||
"mime_type": mime_type,
|
||||
}
|
||||
|
||||
|
||||
def _convert_openai_format_to_data_block(block: dict) -> dict:
|
||||
"""Convert OpenAI image content block to standard data content block.
|
||||
def _normalize_messages(
|
||||
messages: Sequence["BaseMessage"],
|
||||
) -> list["BaseMessage"]:
|
||||
"""Normalize message formats to LangChain v1 standard content blocks.
|
||||
|
||||
If parsing fails, pass-through.
|
||||
Chat models already implement support for:
|
||||
- Images in OpenAI Chat Completions format
|
||||
These will be passed through unchanged
|
||||
- LangChain v1 standard content blocks
|
||||
|
||||
Args:
|
||||
block: The OpenAI image content block to convert.
|
||||
This function extends support to:
|
||||
- `Audio <https://platform.openai.com/docs/api-reference/chat/create>`__ and
|
||||
`file <https://platform.openai.com/docs/api-reference/files>`__ data in OpenAI
|
||||
Chat Completions format
|
||||
- Images are technically supported but we expect chat models to handle them
|
||||
directly; this may change in the future
|
||||
- LangChain v0 standard content blocks for backward compatibility
|
||||
|
||||
Returns:
|
||||
The converted standard data content block.
|
||||
"""
|
||||
if block["type"] == "image_url":
|
||||
parsed = _parse_data_uri(block["image_url"]["url"])
|
||||
if parsed is not None:
|
||||
parsed["type"] = "image"
|
||||
return parsed
|
||||
return block
|
||||
.. versionchanged:: 1.0.0
|
||||
In previous versions, this function returned messages in LangChain v0 format.
|
||||
Now, it returns messages in LangChain v1 format, which upgraded chat models now
|
||||
expect to receive when passing back in message history. For backward
|
||||
compatibility, this function will convert v0 message content to v1 format.
|
||||
|
||||
if block["type"] == "file":
|
||||
parsed = _parse_data_uri(block["file"]["file_data"])
|
||||
if parsed is not None:
|
||||
parsed["type"] = "file"
|
||||
if filename := block["file"].get("filename"):
|
||||
parsed["filename"] = filename
|
||||
return parsed
|
||||
return block
|
||||
.. dropdown:: v0 Content Block Schemas
|
||||
|
||||
if block["type"] == "input_audio":
|
||||
data = block["input_audio"].get("data")
|
||||
audio_format = block["input_audio"].get("format")
|
||||
if data and audio_format:
|
||||
return {
|
||||
"type": "audio",
|
||||
"source_type": "base64",
|
||||
"data": data,
|
||||
"mime_type": f"audio/{audio_format}",
|
||||
``URLContentBlock``:
|
||||
|
||||
.. codeblock::
|
||||
|
||||
{
|
||||
mime_type: NotRequired[str]
|
||||
type: Literal['image', 'audio', 'file'],
|
||||
source_type: Literal['url'],
|
||||
url: str,
|
||||
}
|
||||
return block
|
||||
|
||||
return block
|
||||
``Base64ContentBlock``:
|
||||
|
||||
.. codeblock::
|
||||
|
||||
def _normalize_messages(messages: Sequence[BaseMessage]) -> list[BaseMessage]:
|
||||
"""Extend support for message formats.
|
||||
{
|
||||
mime_type: NotRequired[str]
|
||||
type: Literal['image', 'audio', 'file'],
|
||||
source_type: Literal['base64'],
|
||||
data: str,
|
||||
}
|
||||
|
||||
``IDContentBlock``:
|
||||
|
||||
(In practice, this was never used)
|
||||
|
||||
.. codeblock::
|
||||
|
||||
{
|
||||
type: Literal['image', 'audio', 'file'],
|
||||
source_type: Literal['id'],
|
||||
id: str,
|
||||
}
|
||||
|
||||
``PlainTextContentBlock``:
|
||||
|
||||
.. codeblock::
|
||||
|
||||
{
|
||||
mime_type: NotRequired[str]
|
||||
type: Literal['file'],
|
||||
source_type: Literal['text'],
|
||||
url: str,
|
||||
}
|
||||
|
||||
If a v1 message is passed in, it will be returned as-is, meaning it is safe to
|
||||
always pass in v1 messages to this function for assurance.
|
||||
|
||||
For posterity, here are the OpenAI Chat Completions schemas we expect:
|
||||
|
||||
Chat Completions image. Can be URL-based or base64-encoded. Supports MIME types
|
||||
png, jpeg/jpg, webp, static gif:
|
||||
{
|
||||
"type": Literal['image_url'],
|
||||
"image_url": {
|
||||
"url": Union["data:$MIME_TYPE;base64,$BASE64_ENCODED_IMAGE", "$IMAGE_URL"],
|
||||
"detail": Literal['low', 'high', 'auto'] = 'auto', # Supported by OpenAI
|
||||
}
|
||||
}
|
||||
|
||||
Chat Completions audio:
|
||||
{
|
||||
"type": Literal['input_audio'],
|
||||
"input_audio": {
|
||||
"format": Literal['wav', 'mp3'],
|
||||
"data": str = "$BASE64_ENCODED_AUDIO",
|
||||
},
|
||||
}
|
||||
|
||||
Chat Completions files: either base64 or pre-uploaded file ID
|
||||
{
|
||||
"type": Literal['file'],
|
||||
"file": Union[
|
||||
{
|
||||
"filename": Optional[str] = "$FILENAME",
|
||||
"file_data": str = "$BASE64_ENCODED_FILE",
|
||||
},
|
||||
{
|
||||
"file_id": str = "$FILE_ID", # For pre-uploaded files to OpenAI
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
Chat models implement support for images in OpenAI Chat Completions format, as well
|
||||
as other multimodal data as standard data blocks. This function extends support to
|
||||
audio and file data in OpenAI Chat Completions format by converting them to standard
|
||||
data blocks.
|
||||
"""
|
||||
from langchain_core.messages.block_translators.langchain_v0 import ( # noqa: PLC0415
|
||||
_convert_legacy_v0_content_block_to_v1,
|
||||
)
|
||||
from langchain_core.messages.block_translators.openai import ( # noqa: PLC0415
|
||||
_convert_openai_format_to_data_block,
|
||||
)
|
||||
|
||||
formatted_messages = []
|
||||
for message in messages:
|
||||
# We preserve input messages - the caller may reuse them elsewhere and expects
|
||||
# them to remain unchanged. We only create a copy if we need to translate.
|
||||
formatted_message = message
|
||||
|
||||
if isinstance(message.content, list):
|
||||
for idx, block in enumerate(message.content):
|
||||
# OpenAI Chat Completions multimodal data blocks to v1 standard
|
||||
if (
|
||||
isinstance(block, dict)
|
||||
# Subset to (PDF) files and audio, as most relevant chat models
|
||||
# support images in OAI format (and some may not yet support the
|
||||
# standard data block format)
|
||||
and block.get("type") in {"file", "input_audio"}
|
||||
and _is_openai_data_block(block)
|
||||
and block.get("type") in {"input_audio", "file"}
|
||||
# Discriminate between OpenAI/LC format since they share `'type'`
|
||||
and is_openai_data_block(block)
|
||||
):
|
||||
if formatted_message is message:
|
||||
formatted_message = message.model_copy()
|
||||
# Also shallow-copy content
|
||||
formatted_message.content = list(formatted_message.content)
|
||||
formatted_message = _ensure_message_copy(message, formatted_message)
|
||||
|
||||
converted_block = _convert_openai_format_to_data_block(block)
|
||||
_update_content_block(formatted_message, idx, converted_block)
|
||||
|
||||
# Convert multimodal LangChain v0 to v1 standard content blocks
|
||||
elif (
|
||||
isinstance(block, dict)
|
||||
and block.get("type")
|
||||
in {
|
||||
"image",
|
||||
"audio",
|
||||
"file",
|
||||
}
|
||||
and block.get("source_type") # v1 doesn't have `source_type`
|
||||
in {
|
||||
"url",
|
||||
"base64",
|
||||
"id",
|
||||
"text",
|
||||
}
|
||||
):
|
||||
formatted_message = _ensure_message_copy(message, formatted_message)
|
||||
|
||||
converted_block = _convert_legacy_v0_content_block_to_v1(block)
|
||||
_update_content_block(formatted_message, idx, converted_block)
|
||||
continue
|
||||
|
||||
# else, pass through blocks that look like they have v1 format unchanged
|
||||
|
||||
formatted_message.content[idx] = ( # type: ignore[index] # mypy confused by .model_copy
|
||||
_convert_openai_format_to_data_block(block)
|
||||
)
|
||||
formatted_messages.append(formatted_message)
|
||||
|
||||
return formatted_messages
|
||||
|
||||
|
||||
T = TypeVar("T", bound="BaseMessage")
|
||||
|
||||
|
||||
def _ensure_message_copy(message: T, formatted_message: T) -> T:
|
||||
"""Create a copy of the message if it hasn't been copied yet."""
|
||||
if formatted_message is message:
|
||||
formatted_message = message.model_copy()
|
||||
# Shallow-copy content list to allow modifications
|
||||
formatted_message.content = list(formatted_message.content)
|
||||
return formatted_message
|
||||
|
||||
|
||||
def _update_content_block(
|
||||
formatted_message: "BaseMessage", idx: int, new_block: Union[ContentBlock, dict]
|
||||
) -> None:
|
||||
"""Update a content block at the given index, handling type issues."""
|
||||
# Type ignore needed because:
|
||||
# - `BaseMessage.content` is typed as `Union[str, list[Union[str, dict]]]`
|
||||
# - When content is str, indexing fails (index error)
|
||||
# - When content is list, the items are `Union[str, dict]` but we're assigning
|
||||
# `Union[ContentBlock, dict]` where ContentBlock is richer than dict
|
||||
# - This is safe because we only call this when we've verified content is a list and
|
||||
# we're doing content block conversions
|
||||
formatted_message.content[idx] = new_block # type: ignore[index, assignment]
|
||||
|
||||
|
||||
def _update_message_content_to_blocks(message: T, output_version: str) -> T:
|
||||
return message.model_copy(
|
||||
update={
|
||||
"content": message.content_blocks,
|
||||
"response_metadata": {
|
||||
**message.response_metadata,
|
||||
"output_version": output_version,
|
||||
},
|
||||
}
|
||||
)
|
||||
@@ -12,18 +12,20 @@ from typing import (
|
||||
Callable,
|
||||
Literal,
|
||||
Optional,
|
||||
TypeAlias,
|
||||
TypeVar,
|
||||
Union,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
from typing_extensions import TypeAlias, TypedDict, override
|
||||
from typing_extensions import TypedDict, override
|
||||
|
||||
from langchain_core._api import deprecated
|
||||
from langchain_core.caches import BaseCache
|
||||
from langchain_core.callbacks import Callbacks
|
||||
from langchain_core.globals import get_verbose
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AnyMessage,
|
||||
BaseMessage,
|
||||
MessageLikeRepresentation,
|
||||
@@ -101,7 +103,7 @@ def _get_token_ids_default_method(text: str) -> list[int]:
|
||||
LanguageModelInput = Union[PromptValue, str, Sequence[MessageLikeRepresentation]]
|
||||
LanguageModelOutput = Union[BaseMessage, str]
|
||||
LanguageModelLike = Runnable[LanguageModelInput, LanguageModelOutput]
|
||||
LanguageModelOutputVar = TypeVar("LanguageModelOutputVar", BaseMessage, str)
|
||||
LanguageModelOutputVar = TypeVar("LanguageModelOutputVar", AIMessage, str)
|
||||
|
||||
|
||||
def _get_verbosity() -> bool:
|
||||
|
||||
@@ -27,7 +27,10 @@ from langchain_core.callbacks import (
|
||||
Callbacks,
|
||||
)
|
||||
from langchain_core.globals import get_llm_cache
|
||||
from langchain_core.language_models._utils import _normalize_messages
|
||||
from langchain_core.language_models._utils import (
|
||||
_normalize_messages,
|
||||
_update_message_content_to_blocks,
|
||||
)
|
||||
from langchain_core.language_models.base import (
|
||||
BaseLanguageModel,
|
||||
LangSmithParams,
|
||||
@@ -36,16 +39,17 @@ from langchain_core.language_models.base import (
|
||||
from langchain_core.load import dumpd, dumps
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AIMessageChunk,
|
||||
AnyMessage,
|
||||
BaseMessage,
|
||||
BaseMessageChunk,
|
||||
HumanMessage,
|
||||
convert_to_messages,
|
||||
convert_to_openai_image_block,
|
||||
is_data_content_block,
|
||||
message_chunk_to_message,
|
||||
)
|
||||
from langchain_core.messages.ai import _LC_ID_PREFIX
|
||||
from langchain_core.messages.block_translators.openai import (
|
||||
convert_to_openai_image_block,
|
||||
)
|
||||
from langchain_core.output_parsers.openai_tools import (
|
||||
JsonOutputKeyToolsParser,
|
||||
PydanticToolsParser,
|
||||
@@ -69,6 +73,7 @@ from langchain_core.utils.function_calling import (
|
||||
convert_to_openai_tool,
|
||||
)
|
||||
from langchain_core.utils.pydantic import TypeBaseModel, is_basemodel_subclass
|
||||
from langchain_core.utils.utils import LC_ID_PREFIX, from_env
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import uuid
|
||||
@@ -129,7 +134,7 @@ def _format_for_tracing(messages: list[BaseMessage]) -> list[BaseMessage]:
|
||||
if (
|
||||
block.get("type") == "image"
|
||||
and is_data_content_block(block)
|
||||
and block.get("source_type") != "id"
|
||||
and not ("file_id" in block or block.get("source_type") == "id")
|
||||
):
|
||||
if message_to_trace is message:
|
||||
# Shallow copy
|
||||
@@ -139,6 +144,22 @@ def _format_for_tracing(messages: list[BaseMessage]) -> list[BaseMessage]:
|
||||
message_to_trace.content[idx] = ( # type: ignore[index] # mypy confused by .model_copy
|
||||
convert_to_openai_image_block(block)
|
||||
)
|
||||
elif (
|
||||
block.get("type") == "file"
|
||||
and is_data_content_block(block) # v0 (image/audio/file) or v1
|
||||
and "base64" in block
|
||||
# Backward compat: convert v1 base64 blocks to v0
|
||||
):
|
||||
if message_to_trace is message:
|
||||
# Shallow copy
|
||||
message_to_trace = message.model_copy()
|
||||
message_to_trace.content = list(message_to_trace.content)
|
||||
|
||||
message_to_trace.content[idx] = { # type: ignore[index]
|
||||
**{k: v for k, v in block.items() if k != "base64"},
|
||||
"data": block["base64"],
|
||||
"source_type": "base64",
|
||||
}
|
||||
elif len(block) == 1 and "type" not in block:
|
||||
# Tracing assumes all content blocks have a "type" key. Here
|
||||
# we add this key if it is missing, and there's an obvious
|
||||
@@ -221,7 +242,7 @@ def _format_ls_structured_output(ls_structured_output_format: Optional[dict]) ->
|
||||
return ls_structured_output_format_dict
|
||||
|
||||
|
||||
class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
class BaseChatModel(BaseLanguageModel[AIMessage], ABC):
|
||||
"""Base class for chat models.
|
||||
|
||||
Key imperative methods:
|
||||
@@ -330,6 +351,28 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
|
||||
"""
|
||||
|
||||
output_version: Optional[str] = Field(
|
||||
default_factory=from_env("LC_OUTPUT_VERSION", default=None)
|
||||
)
|
||||
"""Version of ``AIMessage`` output format to store in message content.
|
||||
|
||||
``AIMessage.content_blocks`` will lazily parse the contents of ``content`` into a
|
||||
standard format. This flag can be used to additionally store the standard format
|
||||
in message content, e.g., for serialization purposes.
|
||||
|
||||
Supported values:
|
||||
|
||||
- ``"v0"``: provider-specific format in content (can lazily-parse with
|
||||
``.content_blocks``)
|
||||
- ``"v1"``: standardized format in content (consistent with ``.content_blocks``)
|
||||
|
||||
Partner packages (e.g., ``langchain-openai``) can also use this field to roll out
|
||||
new content formats in a backward-compatible way.
|
||||
|
||||
.. versionadded:: 1.0
|
||||
|
||||
"""
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def raise_deprecation(cls, values: dict) -> Any:
|
||||
@@ -388,21 +431,24 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
*,
|
||||
stop: Optional[list[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> BaseMessage:
|
||||
) -> AIMessage:
|
||||
config = ensure_config(config)
|
||||
return cast(
|
||||
"ChatGeneration",
|
||||
self.generate_prompt(
|
||||
[self._convert_input(input)],
|
||||
stop=stop,
|
||||
callbacks=config.get("callbacks"),
|
||||
tags=config.get("tags"),
|
||||
metadata=config.get("metadata"),
|
||||
run_name=config.get("run_name"),
|
||||
run_id=config.pop("run_id", None),
|
||||
**kwargs,
|
||||
).generations[0][0],
|
||||
).message
|
||||
"AIMessage",
|
||||
cast(
|
||||
"ChatGeneration",
|
||||
self.generate_prompt(
|
||||
[self._convert_input(input)],
|
||||
stop=stop,
|
||||
callbacks=config.get("callbacks"),
|
||||
tags=config.get("tags"),
|
||||
metadata=config.get("metadata"),
|
||||
run_name=config.get("run_name"),
|
||||
run_id=config.pop("run_id", None),
|
||||
**kwargs,
|
||||
).generations[0][0],
|
||||
).message,
|
||||
)
|
||||
|
||||
@override
|
||||
async def ainvoke(
|
||||
@@ -412,7 +458,7 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
*,
|
||||
stop: Optional[list[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> BaseMessage:
|
||||
) -> AIMessage:
|
||||
config = ensure_config(config)
|
||||
llm_result = await self.agenerate_prompt(
|
||||
[self._convert_input(input)],
|
||||
@@ -424,7 +470,9 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
run_id=config.pop("run_id", None),
|
||||
**kwargs,
|
||||
)
|
||||
return cast("ChatGeneration", llm_result.generations[0][0]).message
|
||||
return cast(
|
||||
"AIMessage", cast("ChatGeneration", llm_result.generations[0][0]).message
|
||||
)
|
||||
|
||||
def _should_stream(
|
||||
self,
|
||||
@@ -469,11 +517,11 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
*,
|
||||
stop: Optional[list[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterator[BaseMessageChunk]:
|
||||
) -> Iterator[AIMessageChunk]:
|
||||
if not self._should_stream(async_api=False, **{**kwargs, "stream": True}):
|
||||
# Model doesn't implement streaming, so use default implementation
|
||||
yield cast(
|
||||
"BaseMessageChunk",
|
||||
"AIMessageChunk",
|
||||
self.invoke(input, config=config, stop=stop, **kwargs),
|
||||
)
|
||||
else:
|
||||
@@ -518,16 +566,41 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
|
||||
try:
|
||||
input_messages = _normalize_messages(messages)
|
||||
run_id = "-".join((_LC_ID_PREFIX, str(run_manager.run_id)))
|
||||
run_id = "-".join((LC_ID_PREFIX, str(run_manager.run_id)))
|
||||
yielded = False
|
||||
for chunk in self._stream(input_messages, stop=stop, **kwargs):
|
||||
if chunk.message.id is None:
|
||||
chunk.message.id = run_id
|
||||
chunk.message.response_metadata = _gen_info_and_msg_metadata(chunk)
|
||||
if self.output_version == "v1":
|
||||
# Overwrite .content with .content_blocks
|
||||
chunk.message = _update_message_content_to_blocks(
|
||||
chunk.message, "v1"
|
||||
)
|
||||
run_manager.on_llm_new_token(
|
||||
cast("str", chunk.message.content), chunk=chunk
|
||||
)
|
||||
chunks.append(chunk)
|
||||
yield chunk.message
|
||||
yield cast("AIMessageChunk", chunk.message)
|
||||
yielded = True
|
||||
|
||||
# Yield a final empty chunk with chunk_position="last" if not yet
|
||||
# yielded
|
||||
if (
|
||||
yielded
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: Union[str, list] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
msg_chunk = AIMessageChunk(
|
||||
content=empty_content, chunk_position="last", id=run_id
|
||||
)
|
||||
run_manager.on_llm_new_token(
|
||||
"", chunk=ChatGenerationChunk(message=msg_chunk)
|
||||
)
|
||||
yield msg_chunk
|
||||
except BaseException as e:
|
||||
generations_with_error_metadata = _generate_response_from_error(e)
|
||||
chat_generation_chunk = merge_chat_generation_chunks(chunks)
|
||||
@@ -560,11 +633,11 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
*,
|
||||
stop: Optional[list[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[BaseMessageChunk]:
|
||||
) -> AsyncIterator[AIMessageChunk]:
|
||||
if not self._should_stream(async_api=True, **{**kwargs, "stream": True}):
|
||||
# No async or sync stream is implemented, so fall back to ainvoke
|
||||
yield cast(
|
||||
"BaseMessageChunk",
|
||||
"AIMessageChunk",
|
||||
await self.ainvoke(input, config=config, stop=stop, **kwargs),
|
||||
)
|
||||
return
|
||||
@@ -611,7 +684,8 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
|
||||
try:
|
||||
input_messages = _normalize_messages(messages)
|
||||
run_id = "-".join((_LC_ID_PREFIX, str(run_manager.run_id)))
|
||||
run_id = "-".join((LC_ID_PREFIX, str(run_manager.run_id)))
|
||||
yielded = False
|
||||
async for chunk in self._astream(
|
||||
input_messages,
|
||||
stop=stop,
|
||||
@@ -620,11 +694,34 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
if chunk.message.id is None:
|
||||
chunk.message.id = run_id
|
||||
chunk.message.response_metadata = _gen_info_and_msg_metadata(chunk)
|
||||
if self.output_version == "v1":
|
||||
# Overwrite .content with .content_blocks
|
||||
chunk.message = _update_message_content_to_blocks(
|
||||
chunk.message, "v1"
|
||||
)
|
||||
await run_manager.on_llm_new_token(
|
||||
cast("str", chunk.message.content), chunk=chunk
|
||||
)
|
||||
chunks.append(chunk)
|
||||
yield chunk.message
|
||||
yield cast("AIMessageChunk", chunk.message)
|
||||
yielded = True
|
||||
|
||||
# Yield a final empty chunk with chunk_position="last" if not yet yielded
|
||||
if (
|
||||
yielded
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: Union[str, list] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
msg_chunk = AIMessageChunk(
|
||||
content=empty_content, chunk_position="last", id=run_id
|
||||
)
|
||||
await run_manager.on_llm_new_token(
|
||||
"", chunk=ChatGenerationChunk(message=msg_chunk)
|
||||
)
|
||||
yield msg_chunk
|
||||
except BaseException as e:
|
||||
generations_with_error_metadata = _generate_response_from_error(e)
|
||||
chat_generation_chunk = merge_chat_generation_chunks(chunks)
|
||||
@@ -1077,15 +1174,43 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
**kwargs,
|
||||
):
|
||||
chunks: list[ChatGenerationChunk] = []
|
||||
run_id: Optional[str] = (
|
||||
f"{LC_ID_PREFIX}-{run_manager.run_id}" if run_manager else None
|
||||
)
|
||||
yielded = False
|
||||
for chunk in self._stream(messages, stop=stop, **kwargs):
|
||||
chunk.message.response_metadata = _gen_info_and_msg_metadata(chunk)
|
||||
if self.output_version == "v1":
|
||||
# Overwrite .content with .content_blocks
|
||||
chunk.message = _update_message_content_to_blocks(
|
||||
chunk.message, "v1"
|
||||
)
|
||||
if run_manager:
|
||||
if chunk.message.id is None:
|
||||
chunk.message.id = f"{_LC_ID_PREFIX}-{run_manager.run_id}"
|
||||
chunk.message.id = run_id
|
||||
run_manager.on_llm_new_token(
|
||||
cast("str", chunk.message.content), chunk=chunk
|
||||
)
|
||||
chunks.append(chunk)
|
||||
yielded = True
|
||||
|
||||
# Yield a final empty chunk with chunk_position="last" if not yet yielded
|
||||
if (
|
||||
yielded
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: Union[str, list] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
chunk = ChatGenerationChunk(
|
||||
message=AIMessageChunk(
|
||||
content=empty_content, chunk_position="last", id=run_id
|
||||
)
|
||||
)
|
||||
if run_manager:
|
||||
run_manager.on_llm_new_token("", chunk=chunk)
|
||||
chunks.append(chunk)
|
||||
result = generate_from_stream(iter(chunks))
|
||||
elif inspect.signature(self._generate).parameters.get("run_manager"):
|
||||
result = self._generate(
|
||||
@@ -1094,10 +1219,17 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
else:
|
||||
result = self._generate(messages, stop=stop, **kwargs)
|
||||
|
||||
if self.output_version == "v1":
|
||||
# Overwrite .content with .content_blocks
|
||||
for generation in result.generations:
|
||||
generation.message = _update_message_content_to_blocks(
|
||||
generation.message, "v1"
|
||||
)
|
||||
|
||||
# Add response metadata to each generation
|
||||
for idx, generation in enumerate(result.generations):
|
||||
if run_manager and generation.message.id is None:
|
||||
generation.message.id = f"{_LC_ID_PREFIX}-{run_manager.run_id}-{idx}"
|
||||
generation.message.id = f"{LC_ID_PREFIX}-{run_manager.run_id}-{idx}"
|
||||
generation.message.response_metadata = _gen_info_and_msg_metadata(
|
||||
generation
|
||||
)
|
||||
@@ -1150,15 +1282,43 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
**kwargs,
|
||||
):
|
||||
chunks: list[ChatGenerationChunk] = []
|
||||
run_id: Optional[str] = (
|
||||
f"{LC_ID_PREFIX}-{run_manager.run_id}" if run_manager else None
|
||||
)
|
||||
yielded = False
|
||||
async for chunk in self._astream(messages, stop=stop, **kwargs):
|
||||
chunk.message.response_metadata = _gen_info_and_msg_metadata(chunk)
|
||||
if self.output_version == "v1":
|
||||
# Overwrite .content with .content_blocks
|
||||
chunk.message = _update_message_content_to_blocks(
|
||||
chunk.message, "v1"
|
||||
)
|
||||
if run_manager:
|
||||
if chunk.message.id is None:
|
||||
chunk.message.id = f"{_LC_ID_PREFIX}-{run_manager.run_id}"
|
||||
chunk.message.id = run_id
|
||||
await run_manager.on_llm_new_token(
|
||||
cast("str", chunk.message.content), chunk=chunk
|
||||
)
|
||||
chunks.append(chunk)
|
||||
yielded = True
|
||||
|
||||
# Yield a final empty chunk with chunk_position="last" if not yet yielded
|
||||
if (
|
||||
yielded
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not chunk.message.chunk_position
|
||||
):
|
||||
empty_content: Union[str, list] = (
|
||||
"" if isinstance(chunk.message.content, str) else []
|
||||
)
|
||||
chunk = ChatGenerationChunk(
|
||||
message=AIMessageChunk(
|
||||
content=empty_content, chunk_position="last", id=run_id
|
||||
)
|
||||
)
|
||||
if run_manager:
|
||||
await run_manager.on_llm_new_token("", chunk=chunk)
|
||||
chunks.append(chunk)
|
||||
result = generate_from_stream(iter(chunks))
|
||||
elif inspect.signature(self._agenerate).parameters.get("run_manager"):
|
||||
result = await self._agenerate(
|
||||
@@ -1167,10 +1327,17 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
else:
|
||||
result = await self._agenerate(messages, stop=stop, **kwargs)
|
||||
|
||||
if self.output_version == "v1":
|
||||
# Overwrite .content with .content_blocks
|
||||
for generation in result.generations:
|
||||
generation.message = _update_message_content_to_blocks(
|
||||
generation.message, "v1"
|
||||
)
|
||||
|
||||
# Add response metadata to each generation
|
||||
for idx, generation in enumerate(result.generations):
|
||||
if run_manager and generation.message.id is None:
|
||||
generation.message.id = f"{_LC_ID_PREFIX}-{run_manager.run_id}-{idx}"
|
||||
generation.message.id = f"{LC_ID_PREFIX}-{run_manager.run_id}-{idx}"
|
||||
generation.message.response_metadata = _gen_info_and_msg_metadata(
|
||||
generation
|
||||
)
|
||||
@@ -1443,7 +1610,7 @@ class BaseChatModel(BaseLanguageModel[BaseMessage], ABC):
|
||||
*,
|
||||
tool_choice: Optional[Union[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Runnable[LanguageModelInput, BaseMessage]:
|
||||
) -> Runnable[LanguageModelInput, AIMessage]:
|
||||
"""Bind tools to the model.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -4,7 +4,7 @@ import asyncio
|
||||
import re
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Optional, Union, cast
|
||||
from typing import Any, Literal, Optional, Union, cast
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
@@ -113,7 +113,12 @@ class FakeListChatModel(SimpleChatModel):
|
||||
):
|
||||
raise FakeListChatModelError
|
||||
|
||||
yield ChatGenerationChunk(message=AIMessageChunk(content=c))
|
||||
chunk_position: Optional[Literal["last"]] = (
|
||||
"last" if i_c == len(response) - 1 else None
|
||||
)
|
||||
yield ChatGenerationChunk(
|
||||
message=AIMessageChunk(content=c, chunk_position=chunk_position)
|
||||
)
|
||||
|
||||
@override
|
||||
async def _astream(
|
||||
@@ -136,7 +141,12 @@ class FakeListChatModel(SimpleChatModel):
|
||||
and i_c == self.error_on_chunk_number
|
||||
):
|
||||
raise FakeListChatModelError
|
||||
yield ChatGenerationChunk(message=AIMessageChunk(content=c))
|
||||
chunk_position: Optional[Literal["last"]] = (
|
||||
"last" if i_c == len(response) - 1 else None
|
||||
)
|
||||
yield ChatGenerationChunk(
|
||||
message=AIMessageChunk(content=c, chunk_position=chunk_position)
|
||||
)
|
||||
|
||||
@property
|
||||
@override
|
||||
@@ -152,7 +162,7 @@ class FakeListChatModel(SimpleChatModel):
|
||||
*,
|
||||
return_exceptions: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> list[BaseMessage]:
|
||||
) -> list[AIMessage]:
|
||||
if isinstance(config, list):
|
||||
return [self.invoke(m, c, **kwargs) for m, c in zip(inputs, config)]
|
||||
return [self.invoke(m, config, **kwargs) for m in inputs]
|
||||
@@ -165,7 +175,7 @@ class FakeListChatModel(SimpleChatModel):
|
||||
*,
|
||||
return_exceptions: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> list[BaseMessage]:
|
||||
) -> list[AIMessage]:
|
||||
if isinstance(config, list):
|
||||
# do Not use an async iterator here because need explicit ordering
|
||||
return [await self.ainvoke(m, c, **kwargs) for m, c in zip(inputs, config)]
|
||||
@@ -284,10 +294,16 @@ class GenericFakeChatModel(BaseChatModel):
|
||||
|
||||
content_chunks = cast("list[str]", re.split(r"(\s)", content))
|
||||
|
||||
for token in content_chunks:
|
||||
for idx, token in enumerate(content_chunks):
|
||||
chunk = ChatGenerationChunk(
|
||||
message=AIMessageChunk(content=token, id=message.id)
|
||||
)
|
||||
if (
|
||||
idx == len(content_chunks) - 1
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not message.additional_kwargs
|
||||
):
|
||||
chunk.message.chunk_position = "last"
|
||||
if run_manager:
|
||||
run_manager.on_llm_new_token(token, chunk=chunk)
|
||||
yield chunk
|
||||
|
||||
@@ -1466,10 +1466,10 @@ class BaseLLM(BaseLanguageModel[str], ABC):
|
||||
prompt_dict = self.dict()
|
||||
|
||||
if save_path.suffix == ".json":
|
||||
with save_path.open("w") as f:
|
||||
with save_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(prompt_dict, f, indent=4)
|
||||
elif save_path.suffix.endswith((".yaml", ".yml")):
|
||||
with save_path.open("w") as f:
|
||||
with save_path.open("w", encoding="utf-8") as f:
|
||||
yaml.dump(prompt_dict, f, default_flow_style=False)
|
||||
else:
|
||||
msg = f"{save_path} must be json or yaml"
|
||||
|
||||
@@ -18,6 +18,7 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from langchain_core._import_utils import import_attr
|
||||
from langchain_core.utils.utils import LC_AUTO_PREFIX, LC_ID_PREFIX, ensure_id
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.messages.ai import (
|
||||
@@ -31,10 +32,29 @@ if TYPE_CHECKING:
|
||||
message_to_dict,
|
||||
messages_to_dict,
|
||||
)
|
||||
from langchain_core.messages.chat import ChatMessage, ChatMessageChunk
|
||||
from langchain_core.messages.content_blocks import (
|
||||
from langchain_core.messages.block_translators.openai import (
|
||||
convert_to_openai_data_block,
|
||||
convert_to_openai_image_block,
|
||||
)
|
||||
from langchain_core.messages.chat import ChatMessage, ChatMessageChunk
|
||||
from langchain_core.messages.content import (
|
||||
Annotation,
|
||||
AudioContentBlock,
|
||||
Citation,
|
||||
ContentBlock,
|
||||
DataContentBlock,
|
||||
FileContentBlock,
|
||||
ImageContentBlock,
|
||||
InvalidToolCall,
|
||||
NonStandardAnnotation,
|
||||
NonStandardContentBlock,
|
||||
PlainTextContentBlock,
|
||||
ReasoningContentBlock,
|
||||
ServerToolCall,
|
||||
ServerToolCallChunk,
|
||||
ServerToolResult,
|
||||
TextContentBlock,
|
||||
VideoContentBlock,
|
||||
is_data_content_block,
|
||||
)
|
||||
from langchain_core.messages.function import FunctionMessage, FunctionMessageChunk
|
||||
@@ -42,7 +62,6 @@ if TYPE_CHECKING:
|
||||
from langchain_core.messages.modifier import RemoveMessage
|
||||
from langchain_core.messages.system import SystemMessage, SystemMessageChunk
|
||||
from langchain_core.messages.tool import (
|
||||
InvalidToolCall,
|
||||
ToolCall,
|
||||
ToolCallChunk,
|
||||
ToolMessage,
|
||||
@@ -63,31 +82,50 @@ if TYPE_CHECKING:
|
||||
)
|
||||
|
||||
__all__ = (
|
||||
"LC_AUTO_PREFIX",
|
||||
"LC_ID_PREFIX",
|
||||
"AIMessage",
|
||||
"AIMessageChunk",
|
||||
"Annotation",
|
||||
"AnyMessage",
|
||||
"AudioContentBlock",
|
||||
"BaseMessage",
|
||||
"BaseMessageChunk",
|
||||
"ChatMessage",
|
||||
"ChatMessageChunk",
|
||||
"Citation",
|
||||
"ContentBlock",
|
||||
"DataContentBlock",
|
||||
"FileContentBlock",
|
||||
"FunctionMessage",
|
||||
"FunctionMessageChunk",
|
||||
"HumanMessage",
|
||||
"HumanMessageChunk",
|
||||
"ImageContentBlock",
|
||||
"InvalidToolCall",
|
||||
"MessageLikeRepresentation",
|
||||
"NonStandardAnnotation",
|
||||
"NonStandardContentBlock",
|
||||
"PlainTextContentBlock",
|
||||
"ReasoningContentBlock",
|
||||
"RemoveMessage",
|
||||
"ServerToolCall",
|
||||
"ServerToolCallChunk",
|
||||
"ServerToolResult",
|
||||
"SystemMessage",
|
||||
"SystemMessageChunk",
|
||||
"TextContentBlock",
|
||||
"ToolCall",
|
||||
"ToolCallChunk",
|
||||
"ToolMessage",
|
||||
"ToolMessageChunk",
|
||||
"VideoContentBlock",
|
||||
"_message_from_dict",
|
||||
"convert_to_messages",
|
||||
"convert_to_openai_data_block",
|
||||
"convert_to_openai_image_block",
|
||||
"convert_to_openai_messages",
|
||||
"ensure_id",
|
||||
"filter_messages",
|
||||
"get_buffer_string",
|
||||
"is_data_content_block",
|
||||
@@ -103,35 +141,51 @@ __all__ = (
|
||||
_dynamic_imports = {
|
||||
"AIMessage": "ai",
|
||||
"AIMessageChunk": "ai",
|
||||
"Annotation": "content",
|
||||
"AudioContentBlock": "content",
|
||||
"BaseMessage": "base",
|
||||
"BaseMessageChunk": "base",
|
||||
"merge_content": "base",
|
||||
"message_to_dict": "base",
|
||||
"messages_to_dict": "base",
|
||||
"Citation": "content",
|
||||
"ContentBlock": "content",
|
||||
"ChatMessage": "chat",
|
||||
"ChatMessageChunk": "chat",
|
||||
"DataContentBlock": "content",
|
||||
"FileContentBlock": "content",
|
||||
"FunctionMessage": "function",
|
||||
"FunctionMessageChunk": "function",
|
||||
"HumanMessage": "human",
|
||||
"HumanMessageChunk": "human",
|
||||
"NonStandardAnnotation": "content",
|
||||
"NonStandardContentBlock": "content",
|
||||
"PlainTextContentBlock": "content",
|
||||
"ReasoningContentBlock": "content",
|
||||
"RemoveMessage": "modifier",
|
||||
"ServerToolCall": "content",
|
||||
"ServerToolCallChunk": "content",
|
||||
"ServerToolResult": "content",
|
||||
"SystemMessage": "system",
|
||||
"SystemMessageChunk": "system",
|
||||
"ImageContentBlock": "content",
|
||||
"InvalidToolCall": "tool",
|
||||
"TextContentBlock": "content",
|
||||
"ToolCall": "tool",
|
||||
"ToolCallChunk": "tool",
|
||||
"ToolMessage": "tool",
|
||||
"ToolMessageChunk": "tool",
|
||||
"VideoContentBlock": "content",
|
||||
"AnyMessage": "utils",
|
||||
"MessageLikeRepresentation": "utils",
|
||||
"_message_from_dict": "utils",
|
||||
"convert_to_messages": "utils",
|
||||
"convert_to_openai_data_block": "content_blocks",
|
||||
"convert_to_openai_image_block": "content_blocks",
|
||||
"convert_to_openai_data_block": "block_translators.openai",
|
||||
"convert_to_openai_image_block": "block_translators.openai",
|
||||
"convert_to_openai_messages": "utils",
|
||||
"filter_messages": "utils",
|
||||
"get_buffer_string": "utils",
|
||||
"is_data_content_block": "content_blocks",
|
||||
"is_data_content_block": "content",
|
||||
"merge_message_runs": "utils",
|
||||
"message_chunk_to_message": "utils",
|
||||
"messages_from_dict": "utils",
|
||||
|
||||
@@ -3,42 +3,37 @@
|
||||
import json
|
||||
import logging
|
||||
import operator
|
||||
from typing import Any, Literal, Optional, Union, cast
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, Literal, Optional, Union, cast, overload
|
||||
|
||||
from pydantic import model_validator
|
||||
from typing_extensions import NotRequired, Self, TypedDict, override
|
||||
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.base import (
|
||||
BaseMessage,
|
||||
BaseMessageChunk,
|
||||
_extract_reasoning_from_additional_kwargs,
|
||||
merge_content,
|
||||
)
|
||||
from langchain_core.messages.content import InvalidToolCall
|
||||
from langchain_core.messages.tool import (
|
||||
InvalidToolCall,
|
||||
ToolCall,
|
||||
ToolCallChunk,
|
||||
default_tool_chunk_parser,
|
||||
default_tool_parser,
|
||||
)
|
||||
from langchain_core.messages.tool import (
|
||||
invalid_tool_call as create_invalid_tool_call,
|
||||
)
|
||||
from langchain_core.messages.tool import (
|
||||
tool_call as create_tool_call,
|
||||
)
|
||||
from langchain_core.messages.tool import (
|
||||
tool_call_chunk as create_tool_call_chunk,
|
||||
)
|
||||
from langchain_core.messages.tool import invalid_tool_call as create_invalid_tool_call
|
||||
from langchain_core.messages.tool import tool_call as create_tool_call
|
||||
from langchain_core.messages.tool import tool_call_chunk as create_tool_call_chunk
|
||||
from langchain_core.utils._merge import merge_dicts, merge_lists
|
||||
from langchain_core.utils.json import parse_partial_json
|
||||
from langchain_core.utils.usage import _dict_int_op
|
||||
from langchain_core.utils.utils import LC_AUTO_PREFIX, LC_ID_PREFIX
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_LC_ID_PREFIX = "run-"
|
||||
|
||||
|
||||
class InputTokenDetails(TypedDict, total=False):
|
||||
"""Breakdown of input token counts.
|
||||
|
||||
@@ -162,13 +157,6 @@ class AIMessage(BaseMessage):
|
||||
|
||||
"""
|
||||
|
||||
example: bool = False
|
||||
"""Use to denote that a message is part of an example conversation.
|
||||
|
||||
At the moment, this is ignored by most models. Usage is discouraged.
|
||||
|
||||
"""
|
||||
|
||||
tool_calls: list[ToolCall] = []
|
||||
"""If provided, tool calls associated with the message."""
|
||||
invalid_tool_calls: list[InvalidToolCall] = []
|
||||
@@ -181,20 +169,52 @@ class AIMessage(BaseMessage):
|
||||
"""
|
||||
|
||||
type: Literal["ai"] = "ai"
|
||||
"""The type of the message (used for deserialization). Defaults to ``'ai'``."""
|
||||
"""The type of the message (used for deserialization). Defaults to "ai"."""
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Union[str, list[Union[str, dict]]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize ``AIMessage``.
|
||||
|
||||
Specify ``content`` as positional arg or ``content_blocks`` for typing.
|
||||
|
||||
Args:
|
||||
content: The content of the message.
|
||||
content_blocks: Typed standard content.
|
||||
kwargs: Additional arguments to pass to the parent class.
|
||||
"""
|
||||
super().__init__(content=content, **kwargs)
|
||||
if content_blocks is not None:
|
||||
# If there are tool calls in content_blocks, but not in tool_calls, add them
|
||||
content_tool_calls = [
|
||||
block for block in content_blocks if block.get("type") == "tool_call"
|
||||
]
|
||||
if content_tool_calls and "tool_calls" not in kwargs:
|
||||
kwargs["tool_calls"] = content_tool_calls
|
||||
|
||||
super().__init__(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
super().__init__(content=content, **kwargs)
|
||||
|
||||
@property
|
||||
def lc_attributes(self) -> dict:
|
||||
@@ -204,6 +224,65 @@ class AIMessage(BaseMessage):
|
||||
"invalid_tool_calls": self.invalid_tool_calls,
|
||||
}
|
||||
|
||||
@property
|
||||
def content_blocks(self) -> list[types.ContentBlock]:
|
||||
"""Return content blocks of the message.
|
||||
|
||||
If the message has a known model provider, use the provider-specific translator
|
||||
first before falling back to best-effort parsing. For details, see the property
|
||||
on ``BaseMessage``.
|
||||
"""
|
||||
if self.response_metadata.get("output_version") == "v1":
|
||||
return cast("list[types.ContentBlock]", self.content)
|
||||
|
||||
model_provider = self.response_metadata.get("model_provider")
|
||||
if model_provider:
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
get_translator,
|
||||
)
|
||||
|
||||
translator = get_translator(model_provider)
|
||||
if translator:
|
||||
try:
|
||||
return translator["translate_content"](self)
|
||||
except NotImplementedError:
|
||||
pass
|
||||
|
||||
# Otherwise, use best-effort parsing
|
||||
blocks = super().content_blocks
|
||||
|
||||
if self.tool_calls:
|
||||
# Add from tool_calls if missing from content
|
||||
content_tool_call_ids = {
|
||||
block.get("id")
|
||||
for block in self.content
|
||||
if isinstance(block, dict) and block.get("type") == "tool_call"
|
||||
}
|
||||
for tool_call in self.tool_calls:
|
||||
if (id_ := tool_call.get("id")) and id_ not in content_tool_call_ids:
|
||||
tool_call_block: types.ToolCall = {
|
||||
"type": "tool_call",
|
||||
"id": id_,
|
||||
"name": tool_call["name"],
|
||||
"args": tool_call["args"],
|
||||
}
|
||||
if "index" in tool_call:
|
||||
tool_call_block["index"] = tool_call["index"] # type: ignore[typeddict-item]
|
||||
if "extras" in tool_call:
|
||||
tool_call_block["extras"] = tool_call["extras"] # type: ignore[typeddict-item]
|
||||
blocks.append(tool_call_block)
|
||||
|
||||
# Best-effort reasoning extraction from additional_kwargs
|
||||
# Only add reasoning if not already present
|
||||
# Insert before all other blocks to keep reasoning at the start
|
||||
has_reasoning = any(block.get("type") == "reasoning" for block in blocks)
|
||||
if not has_reasoning and (
|
||||
reasoning_block := _extract_reasoning_from_additional_kwargs(self)
|
||||
):
|
||||
blocks.insert(0, reasoning_block)
|
||||
|
||||
return blocks
|
||||
|
||||
# TODO: remove this logic if possible, reducing breaking nature of changes
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
@@ -232,7 +311,9 @@ class AIMessage(BaseMessage):
|
||||
# Ensure "type" is properly set on all tool call-like dicts.
|
||||
if tool_calls := values.get("tool_calls"):
|
||||
values["tool_calls"] = [
|
||||
create_tool_call(**{k: v for k, v in tc.items() if k != "type"})
|
||||
create_tool_call(
|
||||
**{k: v for k, v in tc.items() if k not in ("type", "extras")}
|
||||
)
|
||||
for tc in tool_calls
|
||||
]
|
||||
if invalid_tool_calls := values.get("invalid_tool_calls"):
|
||||
@@ -307,6 +388,13 @@ class AIMessageChunk(AIMessage, BaseMessageChunk):
|
||||
tool_call_chunks: list[ToolCallChunk] = []
|
||||
"""If provided, tool call chunks associated with the message."""
|
||||
|
||||
chunk_position: Optional[Literal["last"]] = None
|
||||
"""Optional span represented by an aggregated AIMessageChunk.
|
||||
|
||||
If a chunk with ``chunk_position="last"`` is aggregated into a stream,
|
||||
``tool_call_chunks`` in message content will be parsed into ``tool_calls``.
|
||||
"""
|
||||
|
||||
@property
|
||||
def lc_attributes(self) -> dict:
|
||||
"""Attrs to be serialized even if they are derived from other init args."""
|
||||
@@ -315,6 +403,60 @@ class AIMessageChunk(AIMessage, BaseMessageChunk):
|
||||
"invalid_tool_calls": self.invalid_tool_calls,
|
||||
}
|
||||
|
||||
@property
|
||||
def content_blocks(self) -> list[types.ContentBlock]:
|
||||
"""Return content blocks of the message."""
|
||||
if self.response_metadata.get("output_version") == "v1":
|
||||
return cast("list[types.ContentBlock]", self.content)
|
||||
|
||||
model_provider = self.response_metadata.get("model_provider")
|
||||
if model_provider:
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
get_translator,
|
||||
)
|
||||
|
||||
translator = get_translator(model_provider)
|
||||
if translator:
|
||||
try:
|
||||
return translator["translate_content_chunk"](self)
|
||||
except NotImplementedError:
|
||||
pass
|
||||
|
||||
# Otherwise, use best-effort parsing
|
||||
blocks = super().content_blocks
|
||||
|
||||
if (
|
||||
self.tool_call_chunks
|
||||
and not self.content
|
||||
and self.chunk_position != "last" # keep tool_calls if aggregated
|
||||
):
|
||||
blocks = [
|
||||
block
|
||||
for block in blocks
|
||||
if block["type"] not in ("tool_call", "invalid_tool_call")
|
||||
]
|
||||
for tool_call_chunk in self.tool_call_chunks:
|
||||
tc: types.ToolCallChunk = {
|
||||
"type": "tool_call_chunk",
|
||||
"id": tool_call_chunk.get("id"),
|
||||
"name": tool_call_chunk.get("name"),
|
||||
"args": tool_call_chunk.get("args"),
|
||||
}
|
||||
if (idx := tool_call_chunk.get("index")) is not None:
|
||||
tc["index"] = idx
|
||||
blocks.append(tc)
|
||||
|
||||
# Best-effort reasoning extraction from additional_kwargs
|
||||
# Only add reasoning if not already present
|
||||
# Insert before all other blocks to keep reasoning at the start
|
||||
has_reasoning = any(block.get("type") == "reasoning" for block in blocks)
|
||||
if not has_reasoning and (
|
||||
reasoning_block := _extract_reasoning_from_additional_kwargs(self)
|
||||
):
|
||||
blocks.insert(0, reasoning_block)
|
||||
|
||||
return blocks
|
||||
|
||||
@model_validator(mode="after")
|
||||
def init_tool_calls(self) -> Self:
|
||||
"""Initialize tool calls from tool call chunks.
|
||||
@@ -379,10 +521,70 @@ class AIMessageChunk(AIMessage, BaseMessageChunk):
|
||||
add_chunk_to_invalid_tool_calls(chunk)
|
||||
self.tool_calls = tool_calls
|
||||
self.invalid_tool_calls = invalid_tool_calls
|
||||
|
||||
if (
|
||||
self.chunk_position == "last"
|
||||
and self.tool_call_chunks
|
||||
and self.response_metadata.get("output_version") == "v1"
|
||||
and isinstance(self.content, list)
|
||||
):
|
||||
id_to_tc: dict[str, types.ToolCall] = {
|
||||
cast("str", tc.get("id")): {
|
||||
"type": "tool_call",
|
||||
"name": tc["name"],
|
||||
"args": tc["args"],
|
||||
"id": tc.get("id"),
|
||||
}
|
||||
for tc in self.tool_calls
|
||||
if "id" in tc
|
||||
}
|
||||
for idx, block in enumerate(self.content):
|
||||
if (
|
||||
isinstance(block, dict)
|
||||
and block.get("type") == "tool_call_chunk"
|
||||
and (call_id := block.get("id"))
|
||||
and call_id in id_to_tc
|
||||
):
|
||||
self.content[idx] = cast("dict[str, Any]", id_to_tc[call_id])
|
||||
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def init_server_tool_calls(self) -> Self:
|
||||
"""Parse server_tool_call_chunks."""
|
||||
if (
|
||||
self.chunk_position == "last"
|
||||
and self.response_metadata.get("output_version") == "v1"
|
||||
and isinstance(self.content, list)
|
||||
):
|
||||
for idx, block in enumerate(self.content):
|
||||
if (
|
||||
isinstance(block, dict)
|
||||
and block.get("type")
|
||||
in ("server_tool_call", "server_tool_call_chunk")
|
||||
and (args_str := block.get("args"))
|
||||
and isinstance(args_str, str)
|
||||
):
|
||||
try:
|
||||
args = json.loads(args_str)
|
||||
if isinstance(args, dict):
|
||||
self.content[idx]["type"] = "server_tool_call" # type: ignore[index]
|
||||
self.content[idx]["args"] = args # type: ignore[index]
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
return self
|
||||
|
||||
@overload # type: ignore[override] # summing BaseMessages gives ChatPromptTemplate
|
||||
def __add__(self, other: "AIMessageChunk") -> "AIMessageChunk": ...
|
||||
|
||||
@overload
|
||||
def __add__(self, other: Sequence["AIMessageChunk"]) -> "AIMessageChunk": ...
|
||||
|
||||
@overload
|
||||
def __add__(self, other: Any) -> BaseMessageChunk: ...
|
||||
|
||||
@override
|
||||
def __add__(self, other: Any) -> BaseMessageChunk: # type: ignore[override]
|
||||
def __add__(self, other: Any) -> BaseMessageChunk:
|
||||
if isinstance(other, AIMessageChunk):
|
||||
return add_ai_message_chunks(self, other)
|
||||
if isinstance(other, (list, tuple)) and all(
|
||||
@@ -401,17 +603,10 @@ def add_ai_message_chunks(
|
||||
left: The first ``AIMessageChunk``.
|
||||
*others: Other ``AIMessageChunk``s to add.
|
||||
|
||||
Raises:
|
||||
ValueError: If the example values of the chunks are not the same.
|
||||
|
||||
Returns:
|
||||
The resulting ``AIMessageChunk``.
|
||||
|
||||
"""
|
||||
if any(left.example != o.example for o in others):
|
||||
msg = "Cannot concatenate AIMessageChunks with different example values."
|
||||
raise ValueError(msg)
|
||||
|
||||
content = merge_content(left.content, *(o.content for o in others))
|
||||
additional_kwargs = merge_dicts(
|
||||
left.additional_kwargs, *(o.additional_kwargs for o in others)
|
||||
@@ -446,26 +641,40 @@ def add_ai_message_chunks(
|
||||
|
||||
chunk_id = None
|
||||
candidates = [left.id] + [o.id for o in others]
|
||||
# first pass: pick the first non-run-* id
|
||||
# first pass: pick the first provider-assigned id (non-run-* and non-lc_*)
|
||||
for id_ in candidates:
|
||||
if id_ and not id_.startswith(_LC_ID_PREFIX):
|
||||
if (
|
||||
id_
|
||||
and not id_.startswith(LC_ID_PREFIX)
|
||||
and not id_.startswith(LC_AUTO_PREFIX)
|
||||
):
|
||||
chunk_id = id_
|
||||
break
|
||||
else:
|
||||
# second pass: no provider-assigned id found, just take the first non-null
|
||||
# second pass: prefer lc_run-* ids over lc_* ids
|
||||
for id_ in candidates:
|
||||
if id_:
|
||||
if id_ and id_.startswith(LC_ID_PREFIX):
|
||||
chunk_id = id_
|
||||
break
|
||||
else:
|
||||
# third pass: take any remaining id (auto-generated lc_* ids)
|
||||
for id_ in candidates:
|
||||
if id_:
|
||||
chunk_id = id_
|
||||
break
|
||||
|
||||
chunk_position: Optional[Literal["last"]] = (
|
||||
"last" if any(x.chunk_position == "last" for x in [left, *others]) else None
|
||||
)
|
||||
|
||||
return left.__class__(
|
||||
example=left.example,
|
||||
content=content,
|
||||
additional_kwargs=additional_kwargs,
|
||||
tool_call_chunks=tool_call_chunks,
|
||||
response_metadata=response_metadata,
|
||||
usage_metadata=usage_metadata,
|
||||
id=chunk_id,
|
||||
chunk_position=chunk_position,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -2,11 +2,14 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Optional, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Optional, Union, cast, overload
|
||||
|
||||
from pydantic import ConfigDict, Field
|
||||
from typing_extensions import Self
|
||||
|
||||
from langchain_core._api.deprecation import warn_deprecated
|
||||
from langchain_core.load.serializable import Serializable
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.utils import get_bolded_text
|
||||
from langchain_core.utils._merge import merge_dicts, merge_lists
|
||||
from langchain_core.utils.interactive_env import is_interactive_env
|
||||
@@ -17,10 +20,79 @@ if TYPE_CHECKING:
|
||||
from langchain_core.prompts.chat import ChatPromptTemplate
|
||||
|
||||
|
||||
def _extract_reasoning_from_additional_kwargs(
|
||||
message: BaseMessage,
|
||||
) -> Optional[types.ReasoningContentBlock]:
|
||||
"""Extract `reasoning_content` from `additional_kwargs`.
|
||||
|
||||
Handles reasoning content stored in various formats:
|
||||
- `additional_kwargs["reasoning_content"]` (string) - Ollama, DeepSeek, XAI, Groq
|
||||
|
||||
Args:
|
||||
message: The message to extract reasoning from.
|
||||
|
||||
Returns:
|
||||
A `ReasoningContentBlock` if reasoning content is found, None otherwise.
|
||||
"""
|
||||
additional_kwargs = getattr(message, "additional_kwargs", {})
|
||||
|
||||
reasoning_content = additional_kwargs.get("reasoning_content")
|
||||
if reasoning_content is not None and isinstance(reasoning_content, str):
|
||||
return {"type": "reasoning", "reasoning": reasoning_content}
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class TextAccessor(str):
|
||||
"""String-like object that supports both property and method access patterns.
|
||||
|
||||
Exists to maintain backward compatibility while transitioning from method-based to
|
||||
property-based text access in message objects. In LangChain <v1.0, message text was
|
||||
accessed via ``.text()`` method calls. In v1.0=<, the preferred pattern is property
|
||||
access via ``.text``.
|
||||
|
||||
Rather than breaking existing code immediately, ``TextAccessor`` allows both
|
||||
patterns:
|
||||
- Modern property access: ``message.text`` (returns string directly)
|
||||
- Legacy method access: ``message.text()`` (callable, emits deprecation warning)
|
||||
|
||||
"""
|
||||
|
||||
__slots__ = ()
|
||||
|
||||
def __new__(cls, value: str) -> Self:
|
||||
"""Create new TextAccessor instance."""
|
||||
return str.__new__(cls, value)
|
||||
|
||||
def __call__(self) -> str:
|
||||
"""Enable method-style text access for backward compatibility.
|
||||
|
||||
This method exists solely to support legacy code that calls ``.text()``
|
||||
as a method. New code should use property access (``.text``) instead.
|
||||
|
||||
.. deprecated:: 1.0.0
|
||||
Calling ``.text()`` as a method is deprecated. Use ``.text`` as a property
|
||||
instead. This method will be removed in 2.0.0.
|
||||
|
||||
Returns:
|
||||
The string content, identical to property access.
|
||||
|
||||
"""
|
||||
warn_deprecated(
|
||||
since="1.0.0",
|
||||
message=(
|
||||
"Calling .text() as a method is deprecated. "
|
||||
"Use .text as a property instead (e.g., message.text)."
|
||||
),
|
||||
removal="2.0.0",
|
||||
)
|
||||
return str(self)
|
||||
|
||||
|
||||
class BaseMessage(Serializable):
|
||||
"""Base abstract message class.
|
||||
|
||||
Messages are the inputs and outputs of ``ChatModel``s.
|
||||
Messages are the inputs and outputs of a ``ChatModel``.
|
||||
"""
|
||||
|
||||
content: Union[str, list[Union[str, dict]]]
|
||||
@@ -66,17 +138,40 @@ class BaseMessage(Serializable):
|
||||
extra="allow",
|
||||
)
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Union[str, list[Union[str, dict]]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize ``BaseMessage``.
|
||||
|
||||
Specify ``content`` as positional arg or ``content_blocks`` for typing.
|
||||
|
||||
Args:
|
||||
content: The string contents of the message.
|
||||
content_blocks: Typed standard content.
|
||||
kwargs: Additional arguments to pass to the parent class.
|
||||
"""
|
||||
super().__init__(content=content, **kwargs)
|
||||
if content_blocks is not None:
|
||||
super().__init__(content=content_blocks, **kwargs)
|
||||
else:
|
||||
super().__init__(content=content, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def is_lc_serializable(cls) -> bool:
|
||||
@@ -96,26 +191,98 @@ class BaseMessage(Serializable):
|
||||
"""
|
||||
return ["langchain", "schema", "messages"]
|
||||
|
||||
def text(self) -> str:
|
||||
"""Get the text ``content`` of the message.
|
||||
@property
|
||||
def content_blocks(self) -> list[types.ContentBlock]:
|
||||
r"""Load content blocks from the message content.
|
||||
|
||||
.. versionadded:: 1.0.0
|
||||
|
||||
"""
|
||||
# Needed here to avoid circular import, as these classes import BaseMessages
|
||||
from langchain_core.messages import content as types # noqa: PLC0415
|
||||
from langchain_core.messages.block_translators.anthropic import ( # noqa: PLC0415
|
||||
_convert_to_v1_from_anthropic_input,
|
||||
)
|
||||
from langchain_core.messages.block_translators.bedrock_converse import ( # noqa: PLC0415
|
||||
_convert_to_v1_from_converse_input,
|
||||
)
|
||||
from langchain_core.messages.block_translators.google_genai import ( # noqa: PLC0415
|
||||
_convert_to_v1_from_genai_input,
|
||||
)
|
||||
from langchain_core.messages.block_translators.langchain_v0 import ( # noqa: PLC0415
|
||||
_convert_v0_multimodal_input_to_v1,
|
||||
)
|
||||
from langchain_core.messages.block_translators.openai import ( # noqa: PLC0415
|
||||
_convert_to_v1_from_chat_completions_input,
|
||||
)
|
||||
|
||||
blocks: list[types.ContentBlock] = []
|
||||
content = (
|
||||
# Transpose string content to list, otherwise assumed to be list
|
||||
[self.content]
|
||||
if isinstance(self.content, str) and self.content
|
||||
else self.content
|
||||
)
|
||||
for item in content:
|
||||
if isinstance(item, str):
|
||||
# Plain string content is treated as a text block
|
||||
blocks.append({"type": "text", "text": item})
|
||||
elif isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
if item_type not in types.KNOWN_BLOCK_TYPES:
|
||||
# Handle all provider-specific or None type blocks as non-standard -
|
||||
# we'll come back to these later
|
||||
blocks.append({"type": "non_standard", "value": item})
|
||||
else:
|
||||
# Guard against v0 blocks that share the same `type` keys
|
||||
if "source_type" in item:
|
||||
blocks.append({"type": "non_standard", "value": item})
|
||||
continue
|
||||
|
||||
# This can't be a v0 block (since they require `source_type`),
|
||||
# so it's a known v1 block type
|
||||
blocks.append(cast("types.ContentBlock", item))
|
||||
|
||||
# Subsequent passes: attempt to unpack non-standard blocks.
|
||||
# This is the last stop - if we can't parse it here, it is left as non-standard
|
||||
for parsing_step in [
|
||||
_convert_v0_multimodal_input_to_v1,
|
||||
_convert_to_v1_from_chat_completions_input,
|
||||
_convert_to_v1_from_anthropic_input,
|
||||
_convert_to_v1_from_genai_input,
|
||||
_convert_to_v1_from_converse_input,
|
||||
]:
|
||||
blocks = parsing_step(blocks)
|
||||
return blocks
|
||||
|
||||
@property
|
||||
def text(self) -> TextAccessor:
|
||||
"""Get the text content of the message as a string.
|
||||
|
||||
Can be used as both property (``message.text``) and method (``message.text()``).
|
||||
|
||||
.. deprecated:: 1.0.0
|
||||
Calling ``.text()`` as a method is deprecated. Use ``.text`` as a property
|
||||
instead. This method will be removed in 2.0.0.
|
||||
|
||||
Returns:
|
||||
The text content of the message.
|
||||
|
||||
"""
|
||||
if isinstance(self.content, str):
|
||||
return self.content
|
||||
|
||||
# must be a list
|
||||
blocks = [
|
||||
block
|
||||
for block in self.content
|
||||
if isinstance(block, str)
|
||||
or (block.get("type") == "text" and isinstance(block.get("text"), str))
|
||||
]
|
||||
return "".join(
|
||||
block if isinstance(block, str) else block["text"] for block in blocks
|
||||
)
|
||||
text_value = self.content
|
||||
else:
|
||||
# must be a list
|
||||
blocks = [
|
||||
block
|
||||
for block in self.content
|
||||
if isinstance(block, str)
|
||||
or (block.get("type") == "text" and isinstance(block.get("text"), str))
|
||||
]
|
||||
text_value = "".join(
|
||||
block if isinstance(block, str) else block["text"] for block in blocks
|
||||
)
|
||||
return TextAccessor(text_value)
|
||||
|
||||
def __add__(self, other: Any) -> ChatPromptTemplate:
|
||||
"""Concatenate this message with another message.
|
||||
@@ -171,7 +338,9 @@ def merge_content(
|
||||
The merged content.
|
||||
|
||||
"""
|
||||
merged = first_content
|
||||
merged: Union[str, list[Union[str, dict]]]
|
||||
merged = "" if first_content is None else first_content
|
||||
|
||||
for content in contents:
|
||||
# If current is a string
|
||||
if isinstance(merged, str):
|
||||
@@ -190,8 +359,10 @@ def merge_content(
|
||||
elif merged and isinstance(merged[-1], str):
|
||||
merged[-1] += content
|
||||
# If second content is an empty string, treat as a no-op
|
||||
elif content:
|
||||
# Otherwise, add the second content as a new element of the list
|
||||
elif content == "":
|
||||
pass
|
||||
# Otherwise, add the second content as a new element of the list
|
||||
elif merged:
|
||||
merged.append(content)
|
||||
return merged
|
||||
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
"""Derivations of standard content blocks from provider content.
|
||||
|
||||
``AIMessage`` will first attempt to use a provider-specific translator if
|
||||
``model_provider`` is set in ``response_metadata`` on the message. Consequently, each
|
||||
provider translator must handle all possible content response types from the provider,
|
||||
including text.
|
||||
|
||||
If no provider is set, or if the provider does not have a registered translator,
|
||||
``AIMessage`` will fall back to best-effort parsing of the content into blocks using
|
||||
the implementation in ``BaseMessage``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
# Provider to translator mapping
|
||||
PROVIDER_TRANSLATORS: dict[str, dict[str, Callable[..., list[types.ContentBlock]]]] = {}
|
||||
"""Map model provider names to translator functions.
|
||||
|
||||
The dictionary maps provider names (e.g. ``'openai'``, ``'anthropic'``) to another
|
||||
dictionary with two keys:
|
||||
- ``'translate_content'``: Function to translate ``AIMessage`` content.
|
||||
- ``'translate_content_chunk'``: Function to translate ``AIMessageChunk`` content.
|
||||
|
||||
When calling `.content_blocks` on an ``AIMessage`` or ``AIMessageChunk``, if
|
||||
``model_provider`` is set in ``response_metadata``, the corresponding translator
|
||||
functions will be used to parse the content into blocks. Otherwise, best-effort parsing
|
||||
in ``BaseMessage`` will be used.
|
||||
"""
|
||||
|
||||
|
||||
def register_translator(
|
||||
provider: str,
|
||||
translate_content: Callable[[AIMessage], list[types.ContentBlock]],
|
||||
translate_content_chunk: Callable[[AIMessageChunk], list[types.ContentBlock]],
|
||||
) -> None:
|
||||
"""Register content translators for a provider in `PROVIDER_TRANSLATORS`.
|
||||
|
||||
Args:
|
||||
provider: The model provider name (e.g. ``'openai'``, ``'anthropic'``).
|
||||
translate_content: Function to translate ``AIMessage`` content.
|
||||
translate_content_chunk: Function to translate ``AIMessageChunk`` content.
|
||||
"""
|
||||
PROVIDER_TRANSLATORS[provider] = {
|
||||
"translate_content": translate_content,
|
||||
"translate_content_chunk": translate_content_chunk,
|
||||
}
|
||||
|
||||
|
||||
def get_translator(
|
||||
provider: str,
|
||||
) -> dict[str, Callable[..., list[types.ContentBlock]]] | None:
|
||||
"""Get the translator functions for a provider.
|
||||
|
||||
Args:
|
||||
provider: The model provider name.
|
||||
|
||||
Returns:
|
||||
Dictionary with ``'translate_content'`` and ``'translate_content_chunk'``
|
||||
functions, or None if no translator is registered for the provider. In such
|
||||
case, best-effort parsing in ``BaseMessage`` will be used.
|
||||
"""
|
||||
return PROVIDER_TRANSLATORS.get(provider)
|
||||
|
||||
|
||||
def _register_translators() -> None:
|
||||
"""Register all translators in langchain-core.
|
||||
|
||||
A unit test ensures all modules in ``block_translators`` are represented here.
|
||||
|
||||
For translators implemented outside langchain-core, they can be registered by
|
||||
calling ``register_translator`` from within the integration package.
|
||||
"""
|
||||
from langchain_core.messages.block_translators.anthropic import ( # noqa: PLC0415
|
||||
_register_anthropic_translator,
|
||||
)
|
||||
from langchain_core.messages.block_translators.bedrock import ( # noqa: PLC0415
|
||||
_register_bedrock_translator,
|
||||
)
|
||||
from langchain_core.messages.block_translators.bedrock_converse import ( # noqa: PLC0415
|
||||
_register_bedrock_converse_translator,
|
||||
)
|
||||
from langchain_core.messages.block_translators.google_genai import ( # noqa: PLC0415
|
||||
_register_google_genai_translator,
|
||||
)
|
||||
from langchain_core.messages.block_translators.google_vertexai import ( # noqa: PLC0415
|
||||
_register_google_vertexai_translator,
|
||||
)
|
||||
from langchain_core.messages.block_translators.groq import ( # noqa: PLC0415
|
||||
_register_groq_translator,
|
||||
)
|
||||
from langchain_core.messages.block_translators.openai import ( # noqa: PLC0415
|
||||
_register_openai_translator,
|
||||
)
|
||||
|
||||
_register_bedrock_translator()
|
||||
_register_bedrock_converse_translator()
|
||||
_register_anthropic_translator()
|
||||
_register_google_genai_translator()
|
||||
_register_google_vertexai_translator()
|
||||
_register_groq_translator()
|
||||
_register_openai_translator()
|
||||
|
||||
|
||||
_register_translators()
|
||||
@@ -0,0 +1,470 @@
|
||||
"""Derivations of standard content blocks from Anthropic content."""
|
||||
|
||||
import json
|
||||
from collections.abc import Iterable
|
||||
from typing import Any, Optional, Union, cast
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
|
||||
def _populate_extras(
|
||||
standard_block: types.ContentBlock, block: dict[str, Any], known_fields: set[str]
|
||||
) -> types.ContentBlock:
|
||||
"""Mutate a block, populating extras."""
|
||||
if standard_block.get("type") == "non_standard":
|
||||
return standard_block
|
||||
|
||||
for key, value in block.items():
|
||||
if key not in known_fields:
|
||||
if "extras" not in standard_block:
|
||||
# Below type-ignores are because mypy thinks a non-standard block can
|
||||
# get here, although we exclude them above.
|
||||
standard_block["extras"] = {} # type: ignore[typeddict-unknown-key]
|
||||
standard_block["extras"][key] = value # type: ignore[typeddict-item]
|
||||
|
||||
return standard_block
|
||||
|
||||
|
||||
def _convert_to_v1_from_anthropic_input(
|
||||
content: list[types.ContentBlock],
|
||||
) -> list[types.ContentBlock]:
|
||||
"""Convert Anthropic format blocks to v1 format.
|
||||
|
||||
During the `.content_blocks` parsing process, we wrap blocks not recognized as a v1
|
||||
block as a ``'non_standard'`` block with the original block stored in the ``value``
|
||||
field. This function attempts to unpack those blocks and convert any blocks that
|
||||
might be Anthropic format to v1 ContentBlocks.
|
||||
|
||||
If conversion fails, the block is left as a ``'non_standard'`` block.
|
||||
|
||||
Args:
|
||||
content: List of content blocks to process.
|
||||
|
||||
Returns:
|
||||
Updated list with Anthropic blocks converted to v1 format.
|
||||
"""
|
||||
|
||||
def _iter_blocks() -> Iterable[types.ContentBlock]:
|
||||
blocks: list[dict[str, Any]] = [
|
||||
cast("dict[str, Any]", block)
|
||||
if block.get("type") != "non_standard"
|
||||
else block["value"] # type: ignore[typeddict-item] # this is only non-standard blocks
|
||||
for block in content
|
||||
]
|
||||
for block in blocks:
|
||||
block_type = block.get("type")
|
||||
|
||||
if (
|
||||
block_type == "document"
|
||||
and "source" in block
|
||||
and "type" in block["source"]
|
||||
):
|
||||
if block["source"]["type"] == "base64":
|
||||
file_block: types.FileContentBlock = {
|
||||
"type": "file",
|
||||
"base64": block["source"]["data"],
|
||||
"mime_type": block["source"]["media_type"],
|
||||
}
|
||||
_populate_extras(file_block, block, {"type", "source"})
|
||||
yield file_block
|
||||
|
||||
elif block["source"]["type"] == "url":
|
||||
file_block = {
|
||||
"type": "file",
|
||||
"url": block["source"]["url"],
|
||||
}
|
||||
_populate_extras(file_block, block, {"type", "source"})
|
||||
yield file_block
|
||||
|
||||
elif block["source"]["type"] == "file":
|
||||
file_block = {
|
||||
"type": "file",
|
||||
"id": block["source"]["file_id"],
|
||||
}
|
||||
_populate_extras(file_block, block, {"type", "source"})
|
||||
yield file_block
|
||||
|
||||
elif block["source"]["type"] == "text":
|
||||
plain_text_block: types.PlainTextContentBlock = {
|
||||
"type": "text-plain",
|
||||
"text": block["source"]["data"],
|
||||
"mime_type": block.get("media_type", "text/plain"),
|
||||
}
|
||||
_populate_extras(plain_text_block, block, {"type", "source"})
|
||||
yield plain_text_block
|
||||
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
elif (
|
||||
block_type == "image"
|
||||
and "source" in block
|
||||
and "type" in block["source"]
|
||||
):
|
||||
if block["source"]["type"] == "base64":
|
||||
image_block: types.ImageContentBlock = {
|
||||
"type": "image",
|
||||
"base64": block["source"]["data"],
|
||||
"mime_type": block["source"]["media_type"],
|
||||
}
|
||||
_populate_extras(image_block, block, {"type", "source"})
|
||||
yield image_block
|
||||
|
||||
elif block["source"]["type"] == "url":
|
||||
image_block = {
|
||||
"type": "image",
|
||||
"url": block["source"]["url"],
|
||||
}
|
||||
_populate_extras(image_block, block, {"type", "source"})
|
||||
yield image_block
|
||||
|
||||
elif block["source"]["type"] == "file":
|
||||
image_block = {
|
||||
"type": "image",
|
||||
"id": block["source"]["file_id"],
|
||||
}
|
||||
_populate_extras(image_block, block, {"type", "source"})
|
||||
yield image_block
|
||||
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
elif block_type in types.KNOWN_BLOCK_TYPES:
|
||||
yield cast("types.ContentBlock", block)
|
||||
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
return list(_iter_blocks())
|
||||
|
||||
|
||||
def _convert_citation_to_v1(citation: dict[str, Any]) -> types.Annotation:
|
||||
citation_type = citation.get("type")
|
||||
|
||||
if citation_type == "web_search_result_location":
|
||||
url_citation: types.Citation = {
|
||||
"type": "citation",
|
||||
"cited_text": citation["cited_text"],
|
||||
"url": citation["url"],
|
||||
}
|
||||
if title := citation.get("title"):
|
||||
url_citation["title"] = title
|
||||
known_fields = {"type", "cited_text", "url", "title", "index", "extras"}
|
||||
for key, value in citation.items():
|
||||
if key not in known_fields:
|
||||
if "extras" not in url_citation:
|
||||
url_citation["extras"] = {}
|
||||
url_citation["extras"][key] = value
|
||||
|
||||
return url_citation
|
||||
|
||||
if citation_type in (
|
||||
"char_location",
|
||||
"content_block_location",
|
||||
"page_location",
|
||||
"search_result_location",
|
||||
):
|
||||
document_citation: types.Citation = {
|
||||
"type": "citation",
|
||||
"cited_text": citation["cited_text"],
|
||||
}
|
||||
if "document_title" in citation:
|
||||
document_citation["title"] = citation["document_title"]
|
||||
elif title := citation.get("title"):
|
||||
document_citation["title"] = title
|
||||
else:
|
||||
pass
|
||||
known_fields = {
|
||||
"type",
|
||||
"cited_text",
|
||||
"document_title",
|
||||
"title",
|
||||
"index",
|
||||
"extras",
|
||||
}
|
||||
for key, value in citation.items():
|
||||
if key not in known_fields:
|
||||
if "extras" not in document_citation:
|
||||
document_citation["extras"] = {}
|
||||
document_citation["extras"][key] = value
|
||||
|
||||
return document_citation
|
||||
|
||||
return {
|
||||
"type": "non_standard_annotation",
|
||||
"value": citation,
|
||||
}
|
||||
|
||||
|
||||
def _convert_to_v1_from_anthropic(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Convert Anthropic message content to v1 format."""
|
||||
if isinstance(message.content, str):
|
||||
content: list[Union[str, dict]] = [{"type": "text", "text": message.content}]
|
||||
else:
|
||||
content = message.content
|
||||
|
||||
def _iter_blocks() -> Iterable[types.ContentBlock]:
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
block_type = block.get("type")
|
||||
|
||||
if block_type == "text":
|
||||
if citations := block.get("citations"):
|
||||
text_block: types.TextContentBlock = {
|
||||
"type": "text",
|
||||
"text": block.get("text", ""),
|
||||
"annotations": [_convert_citation_to_v1(a) for a in citations],
|
||||
}
|
||||
else:
|
||||
text_block = {"type": "text", "text": block["text"]}
|
||||
if "index" in block:
|
||||
text_block["index"] = block["index"]
|
||||
yield text_block
|
||||
|
||||
elif block_type == "thinking":
|
||||
reasoning_block: types.ReasoningContentBlock = {
|
||||
"type": "reasoning",
|
||||
"reasoning": block.get("thinking", ""),
|
||||
}
|
||||
if "index" in block:
|
||||
reasoning_block["index"] = block["index"]
|
||||
known_fields = {"type", "thinking", "index", "extras"}
|
||||
for key in block:
|
||||
if key not in known_fields:
|
||||
if "extras" not in reasoning_block:
|
||||
reasoning_block["extras"] = {}
|
||||
reasoning_block["extras"][key] = block[key]
|
||||
yield reasoning_block
|
||||
|
||||
elif block_type == "tool_use":
|
||||
if (
|
||||
isinstance(message, AIMessageChunk)
|
||||
and len(message.tool_call_chunks) == 1
|
||||
and message.chunk_position != "last"
|
||||
):
|
||||
# Isolated chunk
|
||||
tool_call_chunk: types.ToolCallChunk = (
|
||||
message.tool_call_chunks[0].copy() # type: ignore[assignment]
|
||||
)
|
||||
if "type" not in tool_call_chunk:
|
||||
tool_call_chunk["type"] = "tool_call_chunk"
|
||||
yield tool_call_chunk
|
||||
else:
|
||||
tool_call_block: Optional[types.ToolCall] = None
|
||||
# Non-streaming or gathered chunk
|
||||
if len(message.tool_calls) == 1:
|
||||
tool_call_block = {
|
||||
"type": "tool_call",
|
||||
"name": message.tool_calls[0]["name"],
|
||||
"args": message.tool_calls[0]["args"],
|
||||
"id": message.tool_calls[0].get("id"),
|
||||
}
|
||||
elif call_id := block.get("id"):
|
||||
for tc in message.tool_calls:
|
||||
if tc.get("id") == call_id:
|
||||
tool_call_block = {
|
||||
"type": "tool_call",
|
||||
"name": tc["name"],
|
||||
"args": tc["args"],
|
||||
"id": tc.get("id"),
|
||||
}
|
||||
break
|
||||
else:
|
||||
pass
|
||||
if not tool_call_block:
|
||||
tool_call_block = {
|
||||
"type": "tool_call",
|
||||
"name": block.get("name", ""),
|
||||
"args": block.get("input", {}),
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
if "index" in block:
|
||||
tool_call_block["index"] = block["index"]
|
||||
yield tool_call_block
|
||||
|
||||
elif block_type == "input_json_delta" and isinstance(
|
||||
message, AIMessageChunk
|
||||
):
|
||||
if len(message.tool_call_chunks) == 1:
|
||||
tool_call_chunk = (
|
||||
message.tool_call_chunks[0].copy() # type: ignore[assignment]
|
||||
)
|
||||
if "type" not in tool_call_chunk:
|
||||
tool_call_chunk["type"] = "tool_call_chunk"
|
||||
yield tool_call_chunk
|
||||
|
||||
else:
|
||||
server_tool_call_chunk: types.ServerToolCallChunk = {
|
||||
"type": "server_tool_call_chunk",
|
||||
"args": block.get("partial_json", ""),
|
||||
}
|
||||
if "index" in block:
|
||||
server_tool_call_chunk["index"] = block["index"]
|
||||
yield server_tool_call_chunk
|
||||
|
||||
elif block_type == "server_tool_use":
|
||||
if block.get("name") == "code_execution":
|
||||
server_tool_use_name = "code_interpreter"
|
||||
else:
|
||||
server_tool_use_name = block.get("name", "")
|
||||
if (
|
||||
isinstance(message, AIMessageChunk)
|
||||
and block.get("input") == {}
|
||||
and "partial_json" not in block
|
||||
and message.chunk_position != "last"
|
||||
):
|
||||
# First chunk in a stream
|
||||
server_tool_call_chunk = {
|
||||
"type": "server_tool_call_chunk",
|
||||
"name": server_tool_use_name,
|
||||
"args": "",
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
if "index" in block:
|
||||
server_tool_call_chunk["index"] = block["index"]
|
||||
known_fields = {"type", "name", "input", "id", "index"}
|
||||
_populate_extras(server_tool_call_chunk, block, known_fields)
|
||||
yield server_tool_call_chunk
|
||||
else:
|
||||
server_tool_call: types.ServerToolCall = {
|
||||
"type": "server_tool_call",
|
||||
"name": server_tool_use_name,
|
||||
"args": block.get("input", {}),
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
|
||||
if block.get("input") == {} and "partial_json" in block:
|
||||
try:
|
||||
input_ = json.loads(block["partial_json"])
|
||||
if isinstance(input_, dict):
|
||||
server_tool_call["args"] = input_
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
if "index" in block:
|
||||
server_tool_call["index"] = block["index"]
|
||||
known_fields = {
|
||||
"type",
|
||||
"name",
|
||||
"input",
|
||||
"partial_json",
|
||||
"id",
|
||||
"index",
|
||||
}
|
||||
_populate_extras(server_tool_call, block, known_fields)
|
||||
|
||||
yield server_tool_call
|
||||
|
||||
elif block_type == "mcp_tool_use":
|
||||
if (
|
||||
isinstance(message, AIMessageChunk)
|
||||
and block.get("input") == {}
|
||||
and "partial_json" not in block
|
||||
and message.chunk_position != "last"
|
||||
):
|
||||
# First chunk in a stream
|
||||
server_tool_call_chunk = {
|
||||
"type": "server_tool_call_chunk",
|
||||
"name": "remote_mcp",
|
||||
"args": "",
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
if "name" in block:
|
||||
server_tool_call_chunk["extras"] = {"tool_name": block["name"]}
|
||||
known_fields = {"type", "name", "input", "id", "index"}
|
||||
_populate_extras(server_tool_call_chunk, block, known_fields)
|
||||
if "index" in block:
|
||||
server_tool_call_chunk["index"] = block["index"]
|
||||
yield server_tool_call_chunk
|
||||
else:
|
||||
server_tool_call = {
|
||||
"type": "server_tool_call",
|
||||
"name": "remote_mcp",
|
||||
"args": block.get("input", {}),
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
|
||||
if block.get("input") == {} and "partial_json" in block:
|
||||
try:
|
||||
input_ = json.loads(block["partial_json"])
|
||||
if isinstance(input_, dict):
|
||||
server_tool_call["args"] = input_
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
if "name" in block:
|
||||
server_tool_call["extras"] = {"tool_name": block["name"]}
|
||||
known_fields = {
|
||||
"type",
|
||||
"name",
|
||||
"input",
|
||||
"partial_json",
|
||||
"id",
|
||||
"index",
|
||||
}
|
||||
_populate_extras(server_tool_call, block, known_fields)
|
||||
if "index" in block:
|
||||
server_tool_call["index"] = block["index"]
|
||||
|
||||
yield server_tool_call
|
||||
|
||||
elif block_type and block_type.endswith("_tool_result"):
|
||||
server_tool_result: types.ServerToolResult = {
|
||||
"type": "server_tool_result",
|
||||
"tool_call_id": block.get("tool_use_id", ""),
|
||||
"status": "success",
|
||||
"extras": {"block_type": block_type},
|
||||
}
|
||||
if output := block.get("content", []):
|
||||
server_tool_result["output"] = output
|
||||
if isinstance(output, dict) and output.get(
|
||||
"error_code" # web_search, code_interpreter
|
||||
):
|
||||
server_tool_result["status"] = "error"
|
||||
if block.get("is_error"): # mcp_tool_result
|
||||
server_tool_result["status"] = "error"
|
||||
if "index" in block:
|
||||
server_tool_result["index"] = block["index"]
|
||||
|
||||
known_fields = {"type", "tool_use_id", "content", "is_error", "index"}
|
||||
_populate_extras(server_tool_result, block, known_fields)
|
||||
|
||||
yield server_tool_result
|
||||
|
||||
else:
|
||||
new_block: types.NonStandardContentBlock = {
|
||||
"type": "non_standard",
|
||||
"value": block,
|
||||
}
|
||||
if "index" in new_block["value"]:
|
||||
new_block["index"] = new_block["value"].pop("index")
|
||||
yield new_block
|
||||
|
||||
return list(_iter_blocks())
|
||||
|
||||
|
||||
def translate_content(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a message with Anthropic content."""
|
||||
return _convert_to_v1_from_anthropic(message)
|
||||
|
||||
|
||||
def translate_content_chunk(message: AIMessageChunk) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a message chunk with Anthropic content."""
|
||||
return _convert_to_v1_from_anthropic(message)
|
||||
|
||||
|
||||
def _register_anthropic_translator() -> None:
|
||||
"""Register the Anthropic translator with the central registry.
|
||||
|
||||
Run automatically when the module is imported.
|
||||
"""
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
register_translator,
|
||||
)
|
||||
|
||||
register_translator("anthropic", translate_content, translate_content_chunk)
|
||||
|
||||
|
||||
_register_anthropic_translator()
|
||||
@@ -0,0 +1,94 @@
|
||||
"""Derivations of standard content blocks from Bedrock content."""
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.block_translators.anthropic import (
|
||||
_convert_to_v1_from_anthropic,
|
||||
)
|
||||
|
||||
|
||||
def _convert_to_v1_from_bedrock(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Convert bedrock message content to v1 format."""
|
||||
out = _convert_to_v1_from_anthropic(message)
|
||||
|
||||
content_tool_call_ids = {
|
||||
block.get("id")
|
||||
for block in out
|
||||
if isinstance(block, dict) and block.get("type") == "tool_call"
|
||||
}
|
||||
for tool_call in message.tool_calls:
|
||||
if (id_ := tool_call.get("id")) and id_ not in content_tool_call_ids:
|
||||
tool_call_block: types.ToolCall = {
|
||||
"type": "tool_call",
|
||||
"id": id_,
|
||||
"name": tool_call["name"],
|
||||
"args": tool_call["args"],
|
||||
}
|
||||
if "index" in tool_call:
|
||||
tool_call_block["index"] = tool_call["index"] # type: ignore[typeddict-item]
|
||||
if "extras" in tool_call:
|
||||
tool_call_block["extras"] = tool_call["extras"] # type: ignore[typeddict-item]
|
||||
out.append(tool_call_block)
|
||||
return out
|
||||
|
||||
|
||||
def _convert_to_v1_from_bedrock_chunk(
|
||||
message: AIMessageChunk,
|
||||
) -> list[types.ContentBlock]:
|
||||
"""Convert bedrock message chunk content to v1 format."""
|
||||
if (
|
||||
message.content == ""
|
||||
and not message.additional_kwargs
|
||||
and not message.tool_calls
|
||||
):
|
||||
# Bedrock outputs multiple chunks containing response metadata
|
||||
return []
|
||||
|
||||
out = _convert_to_v1_from_anthropic(message)
|
||||
|
||||
if (
|
||||
message.tool_call_chunks
|
||||
and not message.content
|
||||
and message.chunk_position != "last" # keep tool_calls if aggregated
|
||||
):
|
||||
for tool_call_chunk in message.tool_call_chunks:
|
||||
tc: types.ToolCallChunk = {
|
||||
"type": "tool_call_chunk",
|
||||
"id": tool_call_chunk.get("id"),
|
||||
"name": tool_call_chunk.get("name"),
|
||||
"args": tool_call_chunk.get("args"),
|
||||
}
|
||||
if (idx := tool_call_chunk.get("index")) is not None:
|
||||
tc["index"] = idx
|
||||
out.append(tc)
|
||||
return out
|
||||
|
||||
|
||||
def translate_content(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a message with Bedrock content."""
|
||||
if "claude" not in message.response_metadata.get("model_name", "").lower():
|
||||
raise NotImplementedError # fall back to best-effort parsing
|
||||
return _convert_to_v1_from_bedrock(message)
|
||||
|
||||
|
||||
def translate_content_chunk(message: AIMessageChunk) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a message chunk with Bedrock content."""
|
||||
# TODO: add model_name to all Bedrock chunks and update core merging logic
|
||||
# to not append during aggregation. Then raise NotImplementedError here if
|
||||
# not an Anthropic model to fall back to best-effort parsing.
|
||||
return _convert_to_v1_from_bedrock_chunk(message)
|
||||
|
||||
|
||||
def _register_bedrock_translator() -> None:
|
||||
"""Register the bedrock translator with the central registry.
|
||||
|
||||
Run automatically when the module is imported.
|
||||
"""
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
register_translator,
|
||||
)
|
||||
|
||||
register_translator("bedrock", translate_content, translate_content_chunk)
|
||||
|
||||
|
||||
_register_bedrock_translator()
|
||||
@@ -0,0 +1,297 @@
|
||||
"""Derivations of standard content blocks from Amazon (Bedrock Converse) content."""
|
||||
|
||||
import base64
|
||||
from collections.abc import Iterable
|
||||
from typing import Any, Optional, cast
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
|
||||
def _bytes_to_b64_str(bytes_: bytes) -> str:
|
||||
return base64.b64encode(bytes_).decode("utf-8")
|
||||
|
||||
|
||||
def _populate_extras(
|
||||
standard_block: types.ContentBlock, block: dict[str, Any], known_fields: set[str]
|
||||
) -> types.ContentBlock:
|
||||
"""Mutate a block, populating extras."""
|
||||
if standard_block.get("type") == "non_standard":
|
||||
return standard_block
|
||||
|
||||
for key, value in block.items():
|
||||
if key not in known_fields:
|
||||
if "extras" not in standard_block:
|
||||
# Below type-ignores are because mypy thinks a non-standard block can
|
||||
# get here, although we exclude them above.
|
||||
standard_block["extras"] = {} # type: ignore[typeddict-unknown-key]
|
||||
standard_block["extras"][key] = value # type: ignore[typeddict-item]
|
||||
|
||||
return standard_block
|
||||
|
||||
|
||||
def _convert_to_v1_from_converse_input(
|
||||
content: list[types.ContentBlock],
|
||||
) -> list[types.ContentBlock]:
|
||||
"""Convert Bedrock Converse format blocks to v1 format.
|
||||
|
||||
During the `.content_blocks` parsing process, we wrap blocks not recognized as a v1
|
||||
block as a ``'non_standard'`` block with the original block stored in the ``value``
|
||||
field. This function attempts to unpack those blocks and convert any blocks that
|
||||
might be Converse format to v1 ContentBlocks.
|
||||
|
||||
If conversion fails, the block is left as a ``'non_standard'`` block.
|
||||
|
||||
Args:
|
||||
content: List of content blocks to process.
|
||||
|
||||
Returns:
|
||||
Updated list with Converse blocks converted to v1 format.
|
||||
"""
|
||||
|
||||
def _iter_blocks() -> Iterable[types.ContentBlock]:
|
||||
blocks: list[dict[str, Any]] = [
|
||||
cast("dict[str, Any]", block)
|
||||
if block.get("type") != "non_standard"
|
||||
else block["value"] # type: ignore[typeddict-item] # this is only non-standard blocks
|
||||
for block in content
|
||||
]
|
||||
for block in blocks:
|
||||
num_keys = len(block)
|
||||
|
||||
if num_keys == 1 and (text := block.get("text")):
|
||||
yield {"type": "text", "text": text}
|
||||
|
||||
elif (
|
||||
num_keys == 1
|
||||
and (document := block.get("document"))
|
||||
and isinstance(document, dict)
|
||||
and "format" in document
|
||||
):
|
||||
if document.get("format") == "pdf":
|
||||
if "bytes" in document.get("source", {}):
|
||||
file_block: types.FileContentBlock = {
|
||||
"type": "file",
|
||||
"base64": _bytes_to_b64_str(document["source"]["bytes"]),
|
||||
"mime_type": "application/pdf",
|
||||
}
|
||||
_populate_extras(file_block, document, {"format", "source"})
|
||||
yield file_block
|
||||
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
elif document["format"] == "txt":
|
||||
if "text" in document.get("source", {}):
|
||||
plain_text_block: types.PlainTextContentBlock = {
|
||||
"type": "text-plain",
|
||||
"text": document["source"]["text"],
|
||||
"mime_type": "text/plain",
|
||||
}
|
||||
_populate_extras(
|
||||
plain_text_block, document, {"format", "source"}
|
||||
)
|
||||
yield plain_text_block
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
elif (
|
||||
num_keys == 1
|
||||
and (image := block.get("image"))
|
||||
and isinstance(image, dict)
|
||||
and "format" in image
|
||||
):
|
||||
if "bytes" in image.get("source", {}):
|
||||
image_block: types.ImageContentBlock = {
|
||||
"type": "image",
|
||||
"base64": _bytes_to_b64_str(image["source"]["bytes"]),
|
||||
"mime_type": f"image/{image['format']}",
|
||||
}
|
||||
_populate_extras(image_block, image, {"format", "source"})
|
||||
yield image_block
|
||||
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
elif block.get("type") in types.KNOWN_BLOCK_TYPES:
|
||||
yield cast("types.ContentBlock", block)
|
||||
|
||||
else:
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
return list(_iter_blocks())
|
||||
|
||||
|
||||
def _convert_citation_to_v1(citation: dict[str, Any]) -> types.Annotation:
|
||||
standard_citation: types.Citation = {"type": "citation"}
|
||||
if "title" in citation:
|
||||
standard_citation["title"] = citation["title"]
|
||||
if (
|
||||
(source_content := citation.get("source_content"))
|
||||
and isinstance(source_content, list)
|
||||
and all(isinstance(item, dict) for item in source_content)
|
||||
):
|
||||
standard_citation["cited_text"] = "".join(
|
||||
item.get("text", "") for item in source_content
|
||||
)
|
||||
|
||||
known_fields = {"type", "source_content", "title", "index", "extras"}
|
||||
|
||||
for key, value in citation.items():
|
||||
if key not in known_fields:
|
||||
if "extras" not in standard_citation:
|
||||
standard_citation["extras"] = {}
|
||||
standard_citation["extras"][key] = value
|
||||
|
||||
return standard_citation
|
||||
|
||||
|
||||
def _convert_to_v1_from_converse(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Convert Bedrock Converse message content to v1 format."""
|
||||
if (
|
||||
message.content == ""
|
||||
and not message.additional_kwargs
|
||||
and not message.tool_calls
|
||||
):
|
||||
# Converse outputs multiple chunks containing response metadata
|
||||
return []
|
||||
|
||||
if isinstance(message.content, str):
|
||||
message.content = [{"type": "text", "text": message.content}]
|
||||
|
||||
def _iter_blocks() -> Iterable[types.ContentBlock]:
|
||||
for block in message.content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
block_type = block.get("type")
|
||||
|
||||
if block_type == "text":
|
||||
if citations := block.get("citations"):
|
||||
text_block: types.TextContentBlock = {
|
||||
"type": "text",
|
||||
"text": block.get("text", ""),
|
||||
"annotations": [_convert_citation_to_v1(a) for a in citations],
|
||||
}
|
||||
else:
|
||||
text_block = {"type": "text", "text": block["text"]}
|
||||
if "index" in block:
|
||||
text_block["index"] = block["index"]
|
||||
yield text_block
|
||||
|
||||
elif block_type == "reasoning_content":
|
||||
reasoning_block: types.ReasoningContentBlock = {"type": "reasoning"}
|
||||
if reasoning_content := block.get("reasoning_content"):
|
||||
if reasoning := reasoning_content.get("text"):
|
||||
reasoning_block["reasoning"] = reasoning
|
||||
if signature := reasoning_content.get("signature"):
|
||||
if "extras" not in reasoning_block:
|
||||
reasoning_block["extras"] = {}
|
||||
reasoning_block["extras"]["signature"] = signature
|
||||
|
||||
if "index" in block:
|
||||
reasoning_block["index"] = block["index"]
|
||||
|
||||
known_fields = {"type", "reasoning_content", "index", "extras"}
|
||||
for key in block:
|
||||
if key not in known_fields:
|
||||
if "extras" not in reasoning_block:
|
||||
reasoning_block["extras"] = {}
|
||||
reasoning_block["extras"][key] = block[key]
|
||||
yield reasoning_block
|
||||
|
||||
elif block_type == "tool_use":
|
||||
if (
|
||||
isinstance(message, AIMessageChunk)
|
||||
and len(message.tool_call_chunks) == 1
|
||||
and message.chunk_position != "last"
|
||||
):
|
||||
# Isolated chunk
|
||||
tool_call_chunk: types.ToolCallChunk = (
|
||||
message.tool_call_chunks[0].copy() # type: ignore[assignment]
|
||||
)
|
||||
if "type" not in tool_call_chunk:
|
||||
tool_call_chunk["type"] = "tool_call_chunk"
|
||||
yield tool_call_chunk
|
||||
else:
|
||||
tool_call_block: Optional[types.ToolCall] = None
|
||||
# Non-streaming or gathered chunk
|
||||
if len(message.tool_calls) == 1:
|
||||
tool_call_block = {
|
||||
"type": "tool_call",
|
||||
"name": message.tool_calls[0]["name"],
|
||||
"args": message.tool_calls[0]["args"],
|
||||
"id": message.tool_calls[0].get("id"),
|
||||
}
|
||||
elif call_id := block.get("id"):
|
||||
for tc in message.tool_calls:
|
||||
if tc.get("id") == call_id:
|
||||
tool_call_block = {
|
||||
"type": "tool_call",
|
||||
"name": tc["name"],
|
||||
"args": tc["args"],
|
||||
"id": tc.get("id"),
|
||||
}
|
||||
break
|
||||
else:
|
||||
pass
|
||||
if not tool_call_block:
|
||||
tool_call_block = {
|
||||
"type": "tool_call",
|
||||
"name": block.get("name", ""),
|
||||
"args": block.get("input", {}),
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
if "index" in block:
|
||||
tool_call_block["index"] = block["index"]
|
||||
yield tool_call_block
|
||||
|
||||
elif (
|
||||
block_type == "input_json_delta"
|
||||
and isinstance(message, AIMessageChunk)
|
||||
and len(message.tool_call_chunks) == 1
|
||||
):
|
||||
tool_call_chunk = (
|
||||
message.tool_call_chunks[0].copy() # type: ignore[assignment]
|
||||
)
|
||||
if "type" not in tool_call_chunk:
|
||||
tool_call_chunk["type"] = "tool_call_chunk"
|
||||
yield tool_call_chunk
|
||||
|
||||
else:
|
||||
new_block: types.NonStandardContentBlock = {
|
||||
"type": "non_standard",
|
||||
"value": block,
|
||||
}
|
||||
if "index" in new_block["value"]:
|
||||
new_block["index"] = new_block["value"].pop("index")
|
||||
yield new_block
|
||||
|
||||
return list(_iter_blocks())
|
||||
|
||||
|
||||
def translate_content(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a message with Bedrock Converse content."""
|
||||
return _convert_to_v1_from_converse(message)
|
||||
|
||||
|
||||
def translate_content_chunk(message: AIMessageChunk) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a chunk with Bedrock Converse content."""
|
||||
return _convert_to_v1_from_converse(message)
|
||||
|
||||
|
||||
def _register_bedrock_converse_translator() -> None:
|
||||
"""Register the Bedrock Converse translator with the central registry.
|
||||
|
||||
Run automatically when the module is imported.
|
||||
"""
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
register_translator,
|
||||
)
|
||||
|
||||
register_translator("bedrock_converse", translate_content, translate_content_chunk)
|
||||
|
||||
|
||||
_register_bedrock_converse_translator()
|
||||
@@ -0,0 +1,529 @@
|
||||
"""Derivations of standard content blocks from Google (GenAI) content."""
|
||||
|
||||
import base64
|
||||
import re
|
||||
from collections.abc import Iterable
|
||||
from typing import Any, cast
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.content import Citation, create_citation
|
||||
|
||||
|
||||
def _bytes_to_b64_str(bytes_: bytes) -> str:
|
||||
"""Convert bytes to base64 encoded string."""
|
||||
return base64.b64encode(bytes_).decode("utf-8")
|
||||
|
||||
|
||||
def translate_grounding_metadata_to_citations(
|
||||
grounding_metadata: dict[str, Any],
|
||||
) -> list[Citation]:
|
||||
"""Translate Google AI grounding metadata to LangChain Citations.
|
||||
|
||||
Args:
|
||||
grounding_metadata: Google AI grounding metadata containing web search
|
||||
queries, grounding chunks, and grounding supports.
|
||||
|
||||
Returns:
|
||||
List of Citation content blocks derived from the grounding metadata.
|
||||
|
||||
Example:
|
||||
>>> metadata = {
|
||||
... "web_search_queries": ["UEFA Euro 2024 winner"],
|
||||
... "grounding_chunks": [
|
||||
... {
|
||||
... "web": {
|
||||
... "uri": "https://uefa.com/euro2024",
|
||||
... "title": "UEFA Euro 2024 Results",
|
||||
... }
|
||||
... }
|
||||
... ],
|
||||
... "grounding_supports": [
|
||||
... {
|
||||
... "segment": {
|
||||
... "start_index": 0,
|
||||
... "end_index": 47,
|
||||
... "text": "Spain won the UEFA Euro 2024 championship",
|
||||
... },
|
||||
... "grounding_chunk_indices": [0],
|
||||
... }
|
||||
... ],
|
||||
... }
|
||||
>>> citations = translate_grounding_metadata_to_citations(metadata)
|
||||
>>> len(citations)
|
||||
1
|
||||
>>> citations[0]["url"]
|
||||
'https://uefa.com/euro2024'
|
||||
"""
|
||||
if not grounding_metadata:
|
||||
return []
|
||||
|
||||
grounding_chunks = grounding_metadata.get("grounding_chunks", [])
|
||||
grounding_supports = grounding_metadata.get("grounding_supports", [])
|
||||
web_search_queries = grounding_metadata.get("web_search_queries", [])
|
||||
|
||||
citations: list[Citation] = []
|
||||
|
||||
for support in grounding_supports:
|
||||
segment = support.get("segment", {})
|
||||
chunk_indices = support.get("grounding_chunk_indices", [])
|
||||
|
||||
start_index = segment.get("start_index")
|
||||
end_index = segment.get("end_index")
|
||||
cited_text = segment.get("text")
|
||||
|
||||
# Create a citation for each referenced chunk
|
||||
for chunk_index in chunk_indices:
|
||||
if chunk_index < len(grounding_chunks):
|
||||
chunk = grounding_chunks[chunk_index]
|
||||
web_info = chunk.get("web", {})
|
||||
|
||||
citation = create_citation(
|
||||
url=web_info.get("uri"),
|
||||
title=web_info.get("title"),
|
||||
start_index=start_index,
|
||||
end_index=end_index,
|
||||
cited_text=cited_text,
|
||||
extras={
|
||||
"google_ai_metadata": {
|
||||
"web_search_queries": web_search_queries,
|
||||
"grounding_chunk_index": chunk_index,
|
||||
"confidence_scores": support.get("confidence_scores", []),
|
||||
}
|
||||
},
|
||||
)
|
||||
citations.append(citation)
|
||||
|
||||
return citations
|
||||
|
||||
|
||||
def _convert_to_v1_from_genai_input(
|
||||
content: list[types.ContentBlock],
|
||||
) -> list[types.ContentBlock]:
|
||||
"""Convert Google GenAI format blocks to v1 format.
|
||||
|
||||
Called when message isn't an `AIMessage` or `model_provider` isn't set on
|
||||
`response_metadata`.
|
||||
|
||||
During the `.content_blocks` parsing process, we wrap blocks not recognized as a v1
|
||||
block as a ``'non_standard'`` block with the original block stored in the ``value``
|
||||
field. This function attempts to unpack those blocks and convert any blocks that
|
||||
might be GenAI format to v1 ContentBlocks.
|
||||
|
||||
If conversion fails, the block is left as a ``'non_standard'`` block.
|
||||
|
||||
Args:
|
||||
content: List of content blocks to process.
|
||||
|
||||
Returns:
|
||||
Updated list with GenAI blocks converted to v1 format.
|
||||
"""
|
||||
|
||||
def _iter_blocks() -> Iterable[types.ContentBlock]:
|
||||
blocks: list[dict[str, Any]] = [
|
||||
cast("dict[str, Any]", block)
|
||||
if block.get("type") != "non_standard"
|
||||
else block["value"] # type: ignore[typeddict-item] # this is only non-standard blocks
|
||||
for block in content
|
||||
]
|
||||
for block in blocks:
|
||||
num_keys = len(block)
|
||||
block_type = block.get("type")
|
||||
|
||||
if num_keys == 1 and (text := block.get("text")):
|
||||
# This is probably a TextContentBlock
|
||||
yield {"type": "text", "text": text}
|
||||
|
||||
elif (
|
||||
num_keys == 1
|
||||
and (document := block.get("document"))
|
||||
and isinstance(document, dict)
|
||||
and "format" in document
|
||||
):
|
||||
# Handle document format conversion
|
||||
doc_format = document.get("format")
|
||||
source = document.get("source", {})
|
||||
|
||||
if doc_format == "pdf" and "bytes" in source:
|
||||
# PDF document with byte data
|
||||
file_block: types.FileContentBlock = {
|
||||
"type": "file",
|
||||
"base64": source["bytes"]
|
||||
if isinstance(source["bytes"], str)
|
||||
else _bytes_to_b64_str(source["bytes"]),
|
||||
"mime_type": "application/pdf",
|
||||
}
|
||||
# Preserve extra fields
|
||||
extras = {
|
||||
key: value
|
||||
for key, value in document.items()
|
||||
if key not in {"format", "source"}
|
||||
}
|
||||
if extras:
|
||||
file_block["extras"] = extras
|
||||
yield file_block
|
||||
|
||||
elif doc_format == "txt" and "text" in source:
|
||||
# Text document
|
||||
plain_text_block: types.PlainTextContentBlock = {
|
||||
"type": "text-plain",
|
||||
"text": source["text"],
|
||||
"mime_type": "text/plain",
|
||||
}
|
||||
# Preserve extra fields
|
||||
extras = {
|
||||
key: value
|
||||
for key, value in document.items()
|
||||
if key not in {"format", "source"}
|
||||
}
|
||||
if extras:
|
||||
plain_text_block["extras"] = extras
|
||||
yield plain_text_block
|
||||
|
||||
else:
|
||||
# Unknown document format
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
elif (
|
||||
num_keys == 1
|
||||
and (image := block.get("image"))
|
||||
and isinstance(image, dict)
|
||||
and "format" in image
|
||||
):
|
||||
# Handle image format conversion
|
||||
img_format = image.get("format")
|
||||
source = image.get("source", {})
|
||||
|
||||
if "bytes" in source:
|
||||
# Image with byte data
|
||||
image_block: types.ImageContentBlock = {
|
||||
"type": "image",
|
||||
"base64": source["bytes"]
|
||||
if isinstance(source["bytes"], str)
|
||||
else _bytes_to_b64_str(source["bytes"]),
|
||||
"mime_type": f"image/{img_format}",
|
||||
}
|
||||
# Preserve extra fields
|
||||
extras = {}
|
||||
for key, value in image.items():
|
||||
if key not in {"format", "source"}:
|
||||
extras[key] = value
|
||||
if extras:
|
||||
image_block["extras"] = extras
|
||||
yield image_block
|
||||
|
||||
else:
|
||||
# Image without byte data
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
elif block_type == "file_data" and "file_uri" in block:
|
||||
# Handle FileData URI-based content
|
||||
uri_file_block: types.FileContentBlock = {
|
||||
"type": "file",
|
||||
"url": block["file_uri"],
|
||||
}
|
||||
if mime_type := block.get("mime_type"):
|
||||
uri_file_block["mime_type"] = mime_type
|
||||
yield uri_file_block
|
||||
|
||||
elif block_type == "function_call" and "name" in block:
|
||||
# Handle function calls
|
||||
tool_call_block: types.ToolCall = {
|
||||
"type": "tool_call",
|
||||
"name": block["name"],
|
||||
"args": block.get("args", {}),
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
yield tool_call_block
|
||||
|
||||
elif block_type == "executable_code":
|
||||
server_tool_call_input: types.ServerToolCall = {
|
||||
"type": "server_tool_call",
|
||||
"name": "code_interpreter",
|
||||
"args": {
|
||||
"code": block.get("executable_code", ""),
|
||||
"language": block.get("language", "python"),
|
||||
},
|
||||
"id": block.get("id", ""),
|
||||
}
|
||||
yield server_tool_call_input
|
||||
|
||||
elif block_type == "code_execution_result":
|
||||
outcome = block.get("outcome", 1)
|
||||
status = "success" if outcome == 1 else "error"
|
||||
server_tool_result_input: types.ServerToolResult = {
|
||||
"type": "server_tool_result",
|
||||
"tool_call_id": block.get("tool_call_id", ""),
|
||||
"status": status, # type: ignore[typeddict-item]
|
||||
"output": block.get("code_execution_result", ""),
|
||||
}
|
||||
if outcome is not None:
|
||||
server_tool_result_input["extras"] = {"outcome": outcome}
|
||||
yield server_tool_result_input
|
||||
|
||||
elif block.get("type") in types.KNOWN_BLOCK_TYPES:
|
||||
# We see a standard block type, so we just cast it, even if
|
||||
# we don't fully understand it. This may be dangerous, but
|
||||
# it's better than losing information.
|
||||
yield cast("types.ContentBlock", block)
|
||||
|
||||
else:
|
||||
# We don't understand this block at all.
|
||||
yield {"type": "non_standard", "value": block}
|
||||
|
||||
return list(_iter_blocks())
|
||||
|
||||
|
||||
def _convert_to_v1_from_genai(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Convert Google GenAI message content to v1 format.
|
||||
|
||||
Calling `.content_blocks` on an `AIMessage` where `response_metadata.model_provider`
|
||||
is set to `'google_genai'` will invoke this function to parse the content into
|
||||
standard content blocks for returning.
|
||||
|
||||
Args:
|
||||
message: The AIMessage or AIMessageChunk to convert.
|
||||
|
||||
Returns:
|
||||
List of standard content blocks derived from the message content.
|
||||
"""
|
||||
if isinstance(message.content, str):
|
||||
# String content -> TextContentBlock (only add if non-empty in case of audio)
|
||||
string_blocks: list[types.ContentBlock] = []
|
||||
if message.content:
|
||||
string_blocks.append({"type": "text", "text": message.content})
|
||||
|
||||
# Add any missing tool calls from message.tool_calls field
|
||||
content_tool_call_ids = {
|
||||
block.get("id")
|
||||
for block in string_blocks
|
||||
if isinstance(block, dict) and block.get("type") == "tool_call"
|
||||
}
|
||||
for tool_call in message.tool_calls:
|
||||
id_ = tool_call.get("id")
|
||||
if id_ and id_ not in content_tool_call_ids:
|
||||
string_tool_call_block: types.ToolCall = {
|
||||
"type": "tool_call",
|
||||
"id": id_,
|
||||
"name": tool_call["name"],
|
||||
"args": tool_call["args"],
|
||||
}
|
||||
string_blocks.append(string_tool_call_block)
|
||||
|
||||
# Handle audio from additional_kwargs if present (for empty content cases)
|
||||
audio_data = message.additional_kwargs.get("audio")
|
||||
if audio_data and isinstance(audio_data, bytes):
|
||||
audio_block: types.AudioContentBlock = {
|
||||
"type": "audio",
|
||||
"base64": _bytes_to_b64_str(audio_data),
|
||||
"mime_type": "audio/wav", # Default to WAV for Google GenAI
|
||||
}
|
||||
string_blocks.append(audio_block)
|
||||
|
||||
grounding_metadata = message.response_metadata.get("grounding_metadata")
|
||||
if grounding_metadata:
|
||||
citations = translate_grounding_metadata_to_citations(grounding_metadata)
|
||||
|
||||
for block in string_blocks:
|
||||
if block["type"] == "text" and citations:
|
||||
# Add citations to the first text block only
|
||||
block["annotations"] = cast("list[types.Annotation]", citations)
|
||||
break
|
||||
|
||||
return string_blocks
|
||||
|
||||
if not isinstance(message.content, list):
|
||||
# Unexpected content type, attempt to represent as text
|
||||
return [{"type": "text", "text": str(message.content)}]
|
||||
|
||||
converted_blocks: list[types.ContentBlock] = []
|
||||
|
||||
for item in message.content:
|
||||
if isinstance(item, str):
|
||||
# Conversation history strings
|
||||
|
||||
# Citations are handled below after all blocks are converted
|
||||
converted_blocks.append({"type": "text", "text": item}) # TextContentBlock
|
||||
|
||||
elif isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
if item_type == "image_url":
|
||||
# Convert image_url to standard image block (base64)
|
||||
# (since the original implementation returned as url-base64 CC style)
|
||||
image_url = item.get("image_url", {})
|
||||
url = image_url.get("url", "")
|
||||
if url:
|
||||
# Extract base64 data
|
||||
match = re.match(r"data:([^;]+);base64,(.+)", url)
|
||||
if match:
|
||||
# Data URI provided
|
||||
mime_type, base64_data = match.groups()
|
||||
converted_blocks.append(
|
||||
{
|
||||
"type": "image",
|
||||
"base64": base64_data,
|
||||
"mime_type": mime_type,
|
||||
}
|
||||
)
|
||||
else:
|
||||
# Assume it's raw base64 without data URI
|
||||
try:
|
||||
# Validate base64 and decode for mime type detection
|
||||
decoded_bytes = base64.b64decode(url, validate=True)
|
||||
|
||||
image_url_b64_block = {
|
||||
"type": "image",
|
||||
"base64": url,
|
||||
}
|
||||
|
||||
try:
|
||||
import filetype # type: ignore[import-not-found] # noqa: PLC0415
|
||||
|
||||
# Guess mime type based on file bytes
|
||||
mime_type = None
|
||||
kind = filetype.guess(decoded_bytes)
|
||||
if kind:
|
||||
mime_type = kind.mime
|
||||
if mime_type:
|
||||
image_url_b64_block["mime_type"] = mime_type
|
||||
except ImportError:
|
||||
# filetype library not available, skip type detection
|
||||
pass
|
||||
|
||||
converted_blocks.append(
|
||||
cast("types.ImageContentBlock", image_url_b64_block)
|
||||
)
|
||||
except Exception:
|
||||
# Not valid base64, treat as non-standard
|
||||
converted_blocks.append(
|
||||
{"type": "non_standard", "value": item}
|
||||
)
|
||||
else:
|
||||
# This likely won't be reached according to previous implementations
|
||||
converted_blocks.append({"type": "non_standard", "value": item})
|
||||
msg = "Image URL not a data URI; appending as non-standard block."
|
||||
raise ValueError(msg)
|
||||
elif item_type == "function_call":
|
||||
# Handle Google GenAI function calls
|
||||
function_call_block: types.ToolCall = {
|
||||
"type": "tool_call",
|
||||
"name": item.get("name", ""),
|
||||
"args": item.get("args", {}),
|
||||
"id": item.get("id", ""),
|
||||
}
|
||||
converted_blocks.append(function_call_block)
|
||||
elif item_type == "file_data":
|
||||
# Handle FileData URI-based content
|
||||
file_block: types.FileContentBlock = {
|
||||
"type": "file",
|
||||
"url": item.get("file_uri", ""),
|
||||
}
|
||||
if mime_type := item.get("mime_type"):
|
||||
file_block["mime_type"] = mime_type
|
||||
converted_blocks.append(file_block)
|
||||
elif item_type == "thinking":
|
||||
# Handling for the 'thinking' type we package thoughts as
|
||||
reasoning_block: types.ReasoningContentBlock = {
|
||||
"type": "reasoning",
|
||||
"reasoning": item.get("thinking", ""),
|
||||
}
|
||||
if signature := item.get("signature"):
|
||||
reasoning_block["extras"] = {"signature": signature}
|
||||
|
||||
converted_blocks.append(reasoning_block)
|
||||
elif item_type == "executable_code":
|
||||
# Convert to standard server tool call block at the moment
|
||||
server_tool_call_block: types.ServerToolCall = {
|
||||
"type": "server_tool_call",
|
||||
"name": "code_interpreter",
|
||||
"args": {
|
||||
"code": item.get("executable_code", ""),
|
||||
"language": item.get("language", "python"), # Default to python
|
||||
},
|
||||
"id": item.get("id", ""),
|
||||
}
|
||||
converted_blocks.append(server_tool_call_block)
|
||||
elif item_type == "code_execution_result":
|
||||
# Map outcome to status: OUTCOME_OK (1) → success, else → error
|
||||
outcome = item.get("outcome", 1)
|
||||
status = "success" if outcome == 1 else "error"
|
||||
server_tool_result_block: types.ServerToolResult = {
|
||||
"type": "server_tool_result",
|
||||
"tool_call_id": item.get("tool_call_id", ""),
|
||||
"status": status, # type: ignore[typeddict-item]
|
||||
"output": item.get("code_execution_result", ""),
|
||||
}
|
||||
# Preserve original outcome in extras
|
||||
if outcome is not None:
|
||||
server_tool_result_block["extras"] = {"outcome": outcome}
|
||||
converted_blocks.append(server_tool_result_block)
|
||||
else:
|
||||
# Unknown type, preserve as non-standard
|
||||
converted_blocks.append({"type": "non_standard", "value": item})
|
||||
else:
|
||||
# Non-dict, non-string content
|
||||
converted_blocks.append({"type": "non_standard", "value": item})
|
||||
|
||||
grounding_metadata = message.response_metadata.get("grounding_metadata")
|
||||
if grounding_metadata:
|
||||
citations = translate_grounding_metadata_to_citations(grounding_metadata)
|
||||
|
||||
for block in converted_blocks:
|
||||
if block["type"] == "text" and citations:
|
||||
# Add citations to text blocks (only the first text block)
|
||||
block["annotations"] = cast("list[types.Annotation]", citations)
|
||||
break
|
||||
|
||||
# Audio is stored on the message.additional_kwargs
|
||||
audio_data = message.additional_kwargs.get("audio")
|
||||
if audio_data and isinstance(audio_data, bytes):
|
||||
audio_block_kwargs: types.AudioContentBlock = {
|
||||
"type": "audio",
|
||||
"base64": _bytes_to_b64_str(audio_data),
|
||||
"mime_type": "audio/wav", # Default to WAV for Google GenAI
|
||||
}
|
||||
converted_blocks.append(audio_block_kwargs)
|
||||
|
||||
# Add any missing tool calls from message.tool_calls field
|
||||
content_tool_call_ids = {
|
||||
block.get("id")
|
||||
for block in converted_blocks
|
||||
if isinstance(block, dict) and block.get("type") == "tool_call"
|
||||
}
|
||||
for tool_call in message.tool_calls:
|
||||
id_ = tool_call.get("id")
|
||||
if id_ and id_ not in content_tool_call_ids:
|
||||
missing_tool_call_block: types.ToolCall = {
|
||||
"type": "tool_call",
|
||||
"id": id_,
|
||||
"name": tool_call["name"],
|
||||
"args": tool_call["args"],
|
||||
}
|
||||
converted_blocks.append(missing_tool_call_block)
|
||||
|
||||
return converted_blocks
|
||||
|
||||
|
||||
def translate_content(message: AIMessage) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a message with Google (GenAI) content."""
|
||||
return _convert_to_v1_from_genai(message)
|
||||
|
||||
|
||||
def translate_content_chunk(message: AIMessageChunk) -> list[types.ContentBlock]:
|
||||
"""Derive standard content blocks from a chunk with Google (GenAI) content."""
|
||||
return _convert_to_v1_from_genai(message)
|
||||
|
||||
|
||||
def _register_google_genai_translator() -> None:
|
||||
"""Register the Google (GenAI) translator with the central registry.
|
||||
|
||||
Run automatically when the module is imported.
|
||||
"""
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
register_translator,
|
||||
)
|
||||
|
||||
register_translator("google_genai", translate_content, translate_content_chunk)
|
||||
|
||||
|
||||
_register_google_genai_translator()
|
||||
@@ -0,0 +1,49 @@
|
||||
"""Derivations of standard content blocks from Google (VertexAI) content."""
|
||||
|
||||
import warnings
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
WARNED = False
|
||||
|
||||
|
||||
def translate_content(message: AIMessage) -> list[types.ContentBlock]: # noqa: ARG001
|
||||
"""Derive standard content blocks from a message with Google (VertexAI) content."""
|
||||
global WARNED # noqa: PLW0603
|
||||
if not WARNED:
|
||||
warning_message = (
|
||||
"Content block standardization is not yet fully supported for Google "
|
||||
"VertexAI."
|
||||
)
|
||||
warnings.warn(warning_message, stacklevel=2)
|
||||
WARNED = True
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def translate_content_chunk(message: AIMessageChunk) -> list[types.ContentBlock]: # noqa: ARG001
|
||||
"""Derive standard content blocks from a chunk with Google (VertexAI) content."""
|
||||
global WARNED # noqa: PLW0603
|
||||
if not WARNED:
|
||||
warning_message = (
|
||||
"Content block standardization is not yet fully supported for Google "
|
||||
"VertexAI."
|
||||
)
|
||||
warnings.warn(warning_message, stacklevel=2)
|
||||
WARNED = True
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def _register_google_vertexai_translator() -> None:
|
||||
"""Register the Google (VertexAI) translator with the central registry.
|
||||
|
||||
Run automatically when the module is imported.
|
||||
"""
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
register_translator,
|
||||
)
|
||||
|
||||
register_translator("google_vertexai", translate_content, translate_content_chunk)
|
||||
|
||||
|
||||
_register_google_vertexai_translator()
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Derivations of standard content blocks from Groq content."""
|
||||
|
||||
import warnings
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
WARNED = False
|
||||
|
||||
|
||||
def translate_content(message: AIMessage) -> list[types.ContentBlock]: # noqa: ARG001
|
||||
"""Derive standard content blocks from a message with Groq content."""
|
||||
global WARNED # noqa: PLW0603
|
||||
if not WARNED:
|
||||
warning_message = (
|
||||
"Content block standardization is not yet fully supported for Groq."
|
||||
)
|
||||
warnings.warn(warning_message, stacklevel=2)
|
||||
WARNED = True
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def translate_content_chunk(message: AIMessageChunk) -> list[types.ContentBlock]: # noqa: ARG001
|
||||
"""Derive standard content blocks from a message chunk with Groq content."""
|
||||
global WARNED # noqa: PLW0603
|
||||
if not WARNED:
|
||||
warning_message = (
|
||||
"Content block standardization is not yet fully supported for Groq."
|
||||
)
|
||||
warnings.warn(warning_message, stacklevel=2)
|
||||
WARNED = True
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def _register_groq_translator() -> None:
|
||||
"""Register the Groq translator with the central registry.
|
||||
|
||||
Run automatically when the module is imported.
|
||||
"""
|
||||
from langchain_core.messages.block_translators import ( # noqa: PLC0415
|
||||
register_translator,
|
||||
)
|
||||
|
||||
register_translator("groq", translate_content, translate_content_chunk)
|
||||
|
||||
|
||||
_register_groq_translator()
|
||||
@@ -0,0 +1,301 @@
|
||||
"""Derivations of standard content blocks from LangChain v0 multimodal content."""
|
||||
|
||||
from typing import Any, Union, cast
|
||||
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
|
||||
def _convert_v0_multimodal_input_to_v1(
|
||||
content: list[types.ContentBlock],
|
||||
) -> list[types.ContentBlock]:
|
||||
"""Convert v0 multimodal blocks to v1 format.
|
||||
|
||||
During the `.content_blocks` parsing process, we wrap blocks not recognized as a v1
|
||||
block as a ``'non_standard'`` block with the original block stored in the ``value``
|
||||
field. This function attempts to unpack those blocks and convert any v0 format
|
||||
blocks to v1 format.
|
||||
|
||||
If conversion fails, the block is left as a ``'non_standard'`` block.
|
||||
|
||||
Args:
|
||||
content: List of content blocks to process.
|
||||
|
||||
Returns:
|
||||
v1 content blocks.
|
||||
"""
|
||||
converted_blocks = []
|
||||
unpacked_blocks: list[dict[str, Any]] = [
|
||||
cast("dict[str, Any]", block)
|
||||
if block.get("type") != "non_standard"
|
||||
else block["value"] # type: ignore[typeddict-item] # this is only non-standard blocks
|
||||
for block in content
|
||||
]
|
||||
for block in unpacked_blocks:
|
||||
if block.get("type") in {"image", "audio", "file"} and "source_type" in block:
|
||||
converted_block = _convert_legacy_v0_content_block_to_v1(block)
|
||||
converted_blocks.append(cast("types.ContentBlock", converted_block))
|
||||
elif block.get("type") in types.KNOWN_BLOCK_TYPES:
|
||||
# Guard in case this function is used outside of the .content_blocks flow
|
||||
converted_blocks.append(cast("types.ContentBlock", block))
|
||||
else:
|
||||
converted_blocks.append({"type": "non_standard", "value": block})
|
||||
|
||||
return converted_blocks
|
||||
|
||||
|
||||
def _convert_legacy_v0_content_block_to_v1(
|
||||
block: dict,
|
||||
) -> Union[types.ContentBlock, dict]:
|
||||
"""Convert a LangChain v0 content block to v1 format.
|
||||
|
||||
Preserves unknown keys as extras to avoid data loss.
|
||||
|
||||
Returns the original block unchanged if it's not in v0 format.
|
||||
"""
|
||||
|
||||
def _extract_v0_extras(block_dict: dict, known_keys: set[str]) -> dict[str, Any]:
|
||||
"""Extract unknown keys from v0 block to preserve as extras.
|
||||
|
||||
Args:
|
||||
block_dict: The original v0 block dictionary.
|
||||
known_keys: Set of keys known to be part of the v0 format for this block.
|
||||
|
||||
Returns:
|
||||
A dictionary of extra keys not part of the known v0 format.
|
||||
"""
|
||||
return {k: v for k, v in block_dict.items() if k not in known_keys}
|
||||
|
||||
# Check if this is actually a v0 format block
|
||||
block_type = block.get("type")
|
||||
if block_type not in {"image", "audio", "file"} or "source_type" not in block:
|
||||
# Not a v0 format block, return unchanged
|
||||
return block
|
||||
|
||||
if block.get("type") == "image":
|
||||
source_type = block.get("source_type")
|
||||
if source_type == "url":
|
||||
# image-url
|
||||
known_keys = {"mime_type", "type", "source_type", "url"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
if "id" in block:
|
||||
return types.create_image_block(
|
||||
url=block["url"],
|
||||
mime_type=block.get("mime_type"),
|
||||
id=block["id"],
|
||||
**extras,
|
||||
)
|
||||
|
||||
# Don't construct with an ID if not present in original block
|
||||
v1_image_url = types.ImageContentBlock(type="image", url=block["url"])
|
||||
if block.get("mime_type"):
|
||||
v1_image_url["mime_type"] = block["mime_type"]
|
||||
|
||||
v1_image_url["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_image_url["extras"][key] = value
|
||||
if v1_image_url["extras"] == {}:
|
||||
del v1_image_url["extras"]
|
||||
|
||||
return v1_image_url
|
||||
if source_type == "base64":
|
||||
# image-base64
|
||||
known_keys = {"mime_type", "type", "source_type", "data"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
if "id" in block:
|
||||
return types.create_image_block(
|
||||
base64=block["data"],
|
||||
mime_type=block.get("mime_type"),
|
||||
id=block["id"],
|
||||
**extras,
|
||||
)
|
||||
|
||||
v1_image_base64 = types.ImageContentBlock(
|
||||
type="image", base64=block["data"]
|
||||
)
|
||||
if block.get("mime_type"):
|
||||
v1_image_base64["mime_type"] = block["mime_type"]
|
||||
|
||||
v1_image_base64["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_image_base64["extras"][key] = value
|
||||
if v1_image_base64["extras"] == {}:
|
||||
del v1_image_base64["extras"]
|
||||
|
||||
return v1_image_base64
|
||||
if source_type == "id":
|
||||
# image-id
|
||||
known_keys = {"type", "source_type", "id"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
# For id `source_type`, `id` is the file reference, not block ID
|
||||
v1_image_id = types.ImageContentBlock(type="image", file_id=block["id"])
|
||||
|
||||
v1_image_id["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_image_id["extras"][key] = value
|
||||
if v1_image_id["extras"] == {}:
|
||||
del v1_image_id["extras"]
|
||||
|
||||
return v1_image_id
|
||||
elif block.get("type") == "audio":
|
||||
source_type = block.get("source_type")
|
||||
if source_type == "url":
|
||||
# audio-url
|
||||
known_keys = {"mime_type", "type", "source_type", "url"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
if "id" in block:
|
||||
return types.create_audio_block(
|
||||
url=block["url"],
|
||||
mime_type=block.get("mime_type"),
|
||||
id=block["id"],
|
||||
**extras,
|
||||
)
|
||||
|
||||
# Don't construct with an ID if not present in original block
|
||||
v1_audio_url: types.AudioContentBlock = types.AudioContentBlock(
|
||||
type="audio", url=block["url"]
|
||||
)
|
||||
if block.get("mime_type"):
|
||||
v1_audio_url["mime_type"] = block["mime_type"]
|
||||
|
||||
v1_audio_url["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_audio_url["extras"][key] = value
|
||||
if v1_audio_url["extras"] == {}:
|
||||
del v1_audio_url["extras"]
|
||||
|
||||
return v1_audio_url
|
||||
if source_type == "base64":
|
||||
# audio-base64
|
||||
known_keys = {"mime_type", "type", "source_type", "data"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
if "id" in block:
|
||||
return types.create_audio_block(
|
||||
base64=block["data"],
|
||||
mime_type=block.get("mime_type"),
|
||||
id=block["id"],
|
||||
**extras,
|
||||
)
|
||||
|
||||
v1_audio_base64: types.AudioContentBlock = types.AudioContentBlock(
|
||||
type="audio", base64=block["data"]
|
||||
)
|
||||
if block.get("mime_type"):
|
||||
v1_audio_base64["mime_type"] = block["mime_type"]
|
||||
|
||||
v1_audio_base64["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_audio_base64["extras"][key] = value
|
||||
if v1_audio_base64["extras"] == {}:
|
||||
del v1_audio_base64["extras"]
|
||||
|
||||
return v1_audio_base64
|
||||
if source_type == "id":
|
||||
# audio-id
|
||||
known_keys = {"type", "source_type", "id"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
v1_audio_id: types.AudioContentBlock = types.AudioContentBlock(
|
||||
type="audio", file_id=block["id"]
|
||||
)
|
||||
|
||||
v1_audio_id["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_audio_id["extras"][key] = value
|
||||
if v1_audio_id["extras"] == {}:
|
||||
del v1_audio_id["extras"]
|
||||
|
||||
return v1_audio_id
|
||||
elif block.get("type") == "file":
|
||||
source_type = block.get("source_type")
|
||||
if source_type == "url":
|
||||
# file-url
|
||||
known_keys = {"mime_type", "type", "source_type", "url"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
if "id" in block:
|
||||
return types.create_file_block(
|
||||
url=block["url"],
|
||||
mime_type=block.get("mime_type"),
|
||||
id=block["id"],
|
||||
**extras,
|
||||
)
|
||||
|
||||
v1_file_url: types.FileContentBlock = types.FileContentBlock(
|
||||
type="file", url=block["url"]
|
||||
)
|
||||
if block.get("mime_type"):
|
||||
v1_file_url["mime_type"] = block["mime_type"]
|
||||
|
||||
v1_file_url["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_file_url["extras"][key] = value
|
||||
if v1_file_url["extras"] == {}:
|
||||
del v1_file_url["extras"]
|
||||
|
||||
return v1_file_url
|
||||
if source_type == "base64":
|
||||
# file-base64
|
||||
known_keys = {"mime_type", "type", "source_type", "data"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
if "id" in block:
|
||||
return types.create_file_block(
|
||||
base64=block["data"],
|
||||
mime_type=block.get("mime_type"),
|
||||
id=block["id"],
|
||||
**extras,
|
||||
)
|
||||
|
||||
v1_file_base64: types.FileContentBlock = types.FileContentBlock(
|
||||
type="file", base64=block["data"]
|
||||
)
|
||||
if block.get("mime_type"):
|
||||
v1_file_base64["mime_type"] = block["mime_type"]
|
||||
|
||||
v1_file_base64["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_file_base64["extras"][key] = value
|
||||
if v1_file_base64["extras"] == {}:
|
||||
del v1_file_base64["extras"]
|
||||
|
||||
return v1_file_base64
|
||||
if source_type == "id":
|
||||
# file-id
|
||||
known_keys = {"type", "source_type", "id"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
return types.create_file_block(file_id=block["id"], **extras)
|
||||
if source_type == "text":
|
||||
# file-text
|
||||
known_keys = {"mime_type", "type", "source_type", "url"}
|
||||
extras = _extract_v0_extras(block, known_keys)
|
||||
if "id" in block:
|
||||
return types.create_plaintext_block(
|
||||
# In v0, URL points to the text file content
|
||||
# TODO: attribute this claim
|
||||
text=block["url"],
|
||||
id=block["id"],
|
||||
**extras,
|
||||
)
|
||||
|
||||
v1_file_text: types.PlainTextContentBlock = types.PlainTextContentBlock(
|
||||
type="text-plain", text=block["url"], mime_type="text/plain"
|
||||
)
|
||||
if block.get("mime_type"):
|
||||
v1_file_text["mime_type"] = block["mime_type"]
|
||||
|
||||
v1_file_text["extras"] = {}
|
||||
for key, value in extras.items():
|
||||
if value is not None:
|
||||
v1_file_text["extras"][key] = value
|
||||
if v1_file_text["extras"] == {}:
|
||||
del v1_file_text["extras"]
|
||||
|
||||
return v1_file_text
|
||||
|
||||
# If we can't convert, return the block unchanged
|
||||
return block
|
||||
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -1,176 +0,0 @@
|
||||
"""Types for content blocks."""
|
||||
|
||||
import warnings
|
||||
from typing import Any, Literal, Union
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
|
||||
class BaseDataContentBlock(TypedDict, total=False):
|
||||
"""Base class for data content blocks."""
|
||||
|
||||
mime_type: NotRequired[str]
|
||||
"""MIME type of the content block (if needed)."""
|
||||
|
||||
|
||||
class URLContentBlock(BaseDataContentBlock):
|
||||
"""Content block for data from a URL."""
|
||||
|
||||
type: Literal["image", "audio", "file"]
|
||||
"""Type of the content block."""
|
||||
source_type: Literal["url"]
|
||||
"""Source type (url)."""
|
||||
url: str
|
||||
"""URL for data."""
|
||||
|
||||
|
||||
class Base64ContentBlock(BaseDataContentBlock):
|
||||
"""Content block for inline data from a base64 string."""
|
||||
|
||||
type: Literal["image", "audio", "file"]
|
||||
"""Type of the content block."""
|
||||
source_type: Literal["base64"]
|
||||
"""Source type (base64)."""
|
||||
data: str
|
||||
"""Data as a base64 string."""
|
||||
|
||||
|
||||
class PlainTextContentBlock(BaseDataContentBlock):
|
||||
"""Content block for plain text data (e.g., from a document)."""
|
||||
|
||||
type: Literal["file"]
|
||||
"""Type of the content block."""
|
||||
source_type: Literal["text"]
|
||||
"""Source type (text)."""
|
||||
text: str
|
||||
"""Text data."""
|
||||
|
||||
|
||||
class IDContentBlock(TypedDict):
|
||||
"""Content block for data specified by an identifier."""
|
||||
|
||||
type: Literal["image", "audio", "file"]
|
||||
"""Type of the content block."""
|
||||
source_type: Literal["id"]
|
||||
"""Source type (id)."""
|
||||
id: str
|
||||
"""Identifier for data source."""
|
||||
|
||||
|
||||
DataContentBlock = Union[
|
||||
URLContentBlock,
|
||||
Base64ContentBlock,
|
||||
PlainTextContentBlock,
|
||||
IDContentBlock,
|
||||
]
|
||||
|
||||
_DataContentBlockAdapter: TypeAdapter[DataContentBlock] = TypeAdapter(DataContentBlock)
|
||||
|
||||
|
||||
def is_data_content_block(
|
||||
content_block: dict,
|
||||
) -> bool:
|
||||
"""Check if the content block is a standard data content block.
|
||||
|
||||
Args:
|
||||
content_block: The content block to check.
|
||||
|
||||
Returns:
|
||||
True if the content block is a data content block, False otherwise.
|
||||
"""
|
||||
try:
|
||||
_ = _DataContentBlockAdapter.validate_python(content_block)
|
||||
except ValidationError:
|
||||
return False
|
||||
else:
|
||||
return True
|
||||
|
||||
|
||||
def convert_to_openai_image_block(content_block: dict[str, Any]) -> dict:
|
||||
"""Convert image content block to format expected by OpenAI Chat Completions API.
|
||||
|
||||
Args:
|
||||
content_block: The content block to convert.
|
||||
|
||||
Raises:
|
||||
ValueError: If the source type is not supported or if ``mime_type`` is missing
|
||||
for base64 data.
|
||||
|
||||
Returns:
|
||||
A dictionary formatted for OpenAI's API.
|
||||
"""
|
||||
if content_block["source_type"] == "url":
|
||||
return {
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": content_block["url"],
|
||||
},
|
||||
}
|
||||
if content_block["source_type"] == "base64":
|
||||
if "mime_type" not in content_block:
|
||||
error_message = "mime_type key is required for base64 data."
|
||||
raise ValueError(error_message)
|
||||
mime_type = content_block["mime_type"]
|
||||
return {
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:{mime_type};base64,{content_block['data']}",
|
||||
},
|
||||
}
|
||||
error_message = "Unsupported source type. Only 'url' and 'base64' are supported."
|
||||
raise ValueError(error_message)
|
||||
|
||||
|
||||
def convert_to_openai_data_block(block: dict) -> dict:
|
||||
"""Format standard data content block to format expected by OpenAI.
|
||||
|
||||
Args:
|
||||
block: A data content block.
|
||||
|
||||
Raises:
|
||||
ValueError: If the block type or source type is not supported.
|
||||
|
||||
Returns:
|
||||
A dictionary formatted for OpenAI's API.
|
||||
"""
|
||||
if block["type"] == "image":
|
||||
formatted_block = convert_to_openai_image_block(block)
|
||||
|
||||
elif block["type"] == "file":
|
||||
if block["source_type"] == "base64":
|
||||
file = {"file_data": f"data:{block['mime_type']};base64,{block['data']}"}
|
||||
if filename := block.get("filename"):
|
||||
file["filename"] = filename
|
||||
elif (metadata := block.get("metadata")) and ("filename" in metadata):
|
||||
file["filename"] = metadata["filename"]
|
||||
else:
|
||||
warnings.warn(
|
||||
"OpenAI may require a filename for file inputs. Specify a filename "
|
||||
"in the content block: {'type': 'file', 'source_type': 'base64', "
|
||||
"'mime_type': 'application/pdf', 'data': '...', "
|
||||
"'filename': 'my-pdf'}",
|
||||
stacklevel=1,
|
||||
)
|
||||
formatted_block = {"type": "file", "file": file}
|
||||
elif block["source_type"] == "id":
|
||||
formatted_block = {"type": "file", "file": {"file_id": block["id"]}}
|
||||
else:
|
||||
error_msg = "source_type base64 or id is required for file blocks."
|
||||
raise ValueError(error_msg)
|
||||
|
||||
elif block["type"] == "audio":
|
||||
if block["source_type"] == "base64":
|
||||
audio_format = block["mime_type"].split("/")[-1]
|
||||
formatted_block = {
|
||||
"type": "input_audio",
|
||||
"input_audio": {"data": block["data"], "format": audio_format},
|
||||
}
|
||||
else:
|
||||
error_msg = "source_type base64 is required for audio blocks."
|
||||
raise ValueError(error_msg)
|
||||
else:
|
||||
error_msg = f"Block of type {block['type']} is not supported."
|
||||
raise ValueError(error_msg)
|
||||
|
||||
return formatted_block
|
||||
@@ -1,7 +1,8 @@
|
||||
"""Human message."""
|
||||
|
||||
from typing import Any, Literal, Union
|
||||
from typing import Any, Literal, Optional, Union, cast, overload
|
||||
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.base import BaseMessage, BaseMessageChunk
|
||||
|
||||
|
||||
@@ -27,14 +28,6 @@ class HumanMessage(BaseMessage):
|
||||
|
||||
"""
|
||||
|
||||
example: bool = False
|
||||
"""Use to denote that a message is part of an example conversation.
|
||||
|
||||
At the moment, this is ignored by most models. Usage is discouraged.
|
||||
Defaults to False.
|
||||
|
||||
"""
|
||||
|
||||
type: Literal["human"] = "human"
|
||||
"""The type of the message (used for serialization).
|
||||
|
||||
@@ -42,18 +35,35 @@ class HumanMessage(BaseMessage):
|
||||
|
||||
"""
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Union[str, list[Union[str, dict]]],
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize ``HumanMessage``.
|
||||
) -> None: ...
|
||||
|
||||
Args:
|
||||
content: The string contents of the message.
|
||||
kwargs: Additional fields to pass to the message.
|
||||
"""
|
||||
super().__init__(content=content, **kwargs)
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Specify ``content`` as positional arg or ``content_blocks`` for typing."""
|
||||
if content_blocks is not None:
|
||||
super().__init__(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
super().__init__(content=content, **kwargs)
|
||||
|
||||
|
||||
class HumanMessageChunk(HumanMessage, BaseMessageChunk):
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
"""System message."""
|
||||
|
||||
from typing import Any, Literal, Union
|
||||
from typing import Any, Literal, Optional, Union, cast, overload
|
||||
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.base import BaseMessage, BaseMessageChunk
|
||||
|
||||
|
||||
@@ -34,16 +35,35 @@ class SystemMessage(BaseMessage):
|
||||
|
||||
"""
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self, content: Union[str, list[Union[str, dict]]], **kwargs: Any
|
||||
) -> None:
|
||||
"""Pass in content as positional arg.
|
||||
self,
|
||||
content: Union[str, list[Union[str, dict]]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
Args:
|
||||
content: The string contents of the message.
|
||||
kwargs: Additional fields to pass to the message.
|
||||
"""
|
||||
super().__init__(content=content, **kwargs)
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Specify ``content`` as positional arg or ``content_blocks`` for typing."""
|
||||
if content_blocks is not None:
|
||||
super().__init__(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
super().__init__(content=content, **kwargs)
|
||||
|
||||
|
||||
class SystemMessageChunk(SystemMessage, BaseMessageChunk):
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
"""Messages for tools."""
|
||||
|
||||
import json
|
||||
from typing import Any, Literal, Optional, Union
|
||||
from typing import Any, Literal, Optional, Union, cast, overload
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import Field, model_validator
|
||||
from typing_extensions import NotRequired, TypedDict, override
|
||||
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.base import BaseMessage, BaseMessageChunk, merge_content
|
||||
from langchain_core.messages.content import InvalidToolCall
|
||||
from langchain_core.utils._merge import merge_dicts, merge_obj
|
||||
|
||||
|
||||
@@ -142,18 +144,43 @@ class ToolMessage(BaseMessage, ToolOutputMixin):
|
||||
values["tool_call_id"] = str(tool_call_id)
|
||||
return values
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Union[str, list[Union[str, dict]]],
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None: ...
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
content: Optional[Union[str, list[Union[str, dict]]]] = None,
|
||||
content_blocks: Optional[list[types.ContentBlock]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize ``ToolMessage``.
|
||||
|
||||
Specify ``content`` as positional arg or ``content_blocks`` for typing.
|
||||
|
||||
Args:
|
||||
content: The string contents of the message.
|
||||
content_blocks: Typed standard content.
|
||||
**kwargs: Additional fields.
|
||||
"""
|
||||
super().__init__(content=content, **kwargs)
|
||||
if content_blocks is not None:
|
||||
super().__init__(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", content_blocks),
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
super().__init__(content=content, **kwargs)
|
||||
|
||||
|
||||
class ToolMessageChunk(ToolMessage, BaseMessageChunk):
|
||||
@@ -290,24 +317,6 @@ def tool_call_chunk(
|
||||
)
|
||||
|
||||
|
||||
class InvalidToolCall(TypedDict):
|
||||
"""Allowance for errors made by LLM.
|
||||
|
||||
Here we add an ``error`` key to surface errors made during generation
|
||||
(e.g., invalid JSON arguments.)
|
||||
"""
|
||||
|
||||
name: Optional[str]
|
||||
"""The name of the tool to be called."""
|
||||
args: Optional[str]
|
||||
"""The arguments to the tool call."""
|
||||
id: Optional[str]
|
||||
"""An identifier associated with the tool call."""
|
||||
error: Optional[str]
|
||||
"""An error message associated with the tool call."""
|
||||
type: NotRequired[Literal["invalid_tool_call"]]
|
||||
|
||||
|
||||
def invalid_tool_call(
|
||||
*,
|
||||
name: Optional[str] = None,
|
||||
|
||||
@@ -32,10 +32,15 @@ from typing import (
|
||||
from pydantic import Discriminator, Field, Tag
|
||||
|
||||
from langchain_core.exceptions import ErrorCode, create_message
|
||||
from langchain_core.messages import convert_to_openai_data_block, is_data_content_block
|
||||
from langchain_core.messages.ai import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages.base import BaseMessage, BaseMessageChunk
|
||||
from langchain_core.messages.block_translators.openai import (
|
||||
convert_to_openai_data_block,
|
||||
)
|
||||
from langchain_core.messages.chat import ChatMessage, ChatMessageChunk
|
||||
from langchain_core.messages.content import (
|
||||
is_data_content_block,
|
||||
)
|
||||
from langchain_core.messages.function import FunctionMessage, FunctionMessageChunk
|
||||
from langchain_core.messages.human import HumanMessage, HumanMessageChunk
|
||||
from langchain_core.messages.modifier import RemoveMessage
|
||||
@@ -137,7 +142,7 @@ def get_buffer_string(
|
||||
else:
|
||||
msg = f"Got unsupported message type: {m}"
|
||||
raise ValueError(msg) # noqa: TRY004
|
||||
message = f"{role}: {m.text()}"
|
||||
message = f"{role}: {m.text}"
|
||||
if isinstance(m, AIMessage) and "function_call" in m.additional_kwargs:
|
||||
message += f"{m.additional_kwargs['function_call']}"
|
||||
string_messages.append(message)
|
||||
@@ -204,7 +209,7 @@ def message_chunk_to_message(chunk: BaseMessage) -> BaseMessage:
|
||||
# chunk classes always have the equivalent non-chunk class as their first parent
|
||||
ignore_keys = ["type"]
|
||||
if isinstance(chunk, AIMessageChunk):
|
||||
ignore_keys.append("tool_call_chunks")
|
||||
ignore_keys.extend(["tool_call_chunks", "chunk_position"])
|
||||
return chunk.__class__.__mro__[1](
|
||||
**{k: v for k, v in chunk.__dict__.items() if k not in ignore_keys}
|
||||
)
|
||||
@@ -1617,11 +1622,15 @@ def _msg_to_chunk(message: BaseMessage) -> BaseMessageChunk:
|
||||
def _chunk_to_msg(chunk: BaseMessageChunk) -> BaseMessage:
|
||||
if chunk.__class__ in _CHUNK_MSG_MAP:
|
||||
return _CHUNK_MSG_MAP[chunk.__class__](
|
||||
**chunk.model_dump(exclude={"type", "tool_call_chunks"})
|
||||
**chunk.model_dump(exclude={"type", "tool_call_chunks", "chunk_position"})
|
||||
)
|
||||
for chunk_cls, msg_cls in _CHUNK_MSG_MAP.items():
|
||||
if isinstance(chunk, chunk_cls):
|
||||
return msg_cls(**chunk.model_dump(exclude={"type", "tool_call_chunks"}))
|
||||
return msg_cls(
|
||||
**chunk.model_dump(
|
||||
exclude={"type", "tool_call_chunks", "chunk_position"}
|
||||
)
|
||||
)
|
||||
|
||||
msg = (
|
||||
f"Unrecognized message chunk class {chunk.__class__}. Supported classes are "
|
||||
|
||||
@@ -133,7 +133,7 @@ class ImagePromptValue(PromptValue):
|
||||
|
||||
def to_string(self) -> str:
|
||||
"""Return prompt (image URL) as string."""
|
||||
return self.image_url["url"]
|
||||
return self.image_url.get("url", "")
|
||||
|
||||
def to_messages(self) -> list[BaseMessage]:
|
||||
"""Return prompt (image URL) as messages."""
|
||||
|
||||
@@ -379,10 +379,10 @@ class BasePromptTemplate(
|
||||
directory_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if save_path.suffix == ".json":
|
||||
with save_path.open("w") as f:
|
||||
with save_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(prompt_dict, f, indent=4)
|
||||
elif save_path.suffix.endswith((".yaml", ".yml")):
|
||||
with save_path.open("w") as f:
|
||||
with save_path.open("w", encoding="utf-8") as f:
|
||||
yaml.dump(prompt_dict, f, default_flow_style=False)
|
||||
else:
|
||||
msg = f"{save_path} must be json or yaml"
|
||||
|
||||
@@ -543,8 +543,7 @@ class _StringImageMessagePromptTemplate(BaseMessagePromptTemplate):
|
||||
Returns:
|
||||
A new instance of this class.
|
||||
"""
|
||||
template = Path(template_file).read_text()
|
||||
# TODO: .read_text(encoding="utf-8") for v0.4
|
||||
template = Path(template_file).read_text(encoding="utf-8")
|
||||
return cls.from_template(template, input_variables=input_variables, **kwargs)
|
||||
|
||||
def format_messages(self, **kwargs: Any) -> list[BaseMessage]:
|
||||
|
||||
@@ -53,7 +53,7 @@ def _load_template(var_name: str, config: dict) -> dict:
|
||||
template_path = Path(config.pop(f"{var_name}_path"))
|
||||
# Load the template.
|
||||
if template_path.suffix == ".txt":
|
||||
template = template_path.read_text()
|
||||
template = template_path.read_text(encoding="utf-8")
|
||||
else:
|
||||
raise ValueError
|
||||
# Set the template variable to the extracted variable.
|
||||
@@ -67,7 +67,7 @@ def _load_examples(config: dict) -> dict:
|
||||
pass
|
||||
elif isinstance(config["examples"], str):
|
||||
path = Path(config["examples"])
|
||||
with path.open() as f:
|
||||
with path.open(encoding="utf-8") as f:
|
||||
if path.suffix == ".json":
|
||||
examples = json.load(f)
|
||||
elif path.suffix in {".yaml", ".yml"}:
|
||||
|
||||
@@ -20,7 +20,7 @@ from collections.abc import (
|
||||
)
|
||||
from concurrent.futures import FIRST_COMPLETED, wait
|
||||
from functools import wraps
|
||||
from itertools import groupby, tee
|
||||
from itertools import tee
|
||||
from operator import itemgetter
|
||||
from types import GenericAlias
|
||||
from typing import (
|
||||
@@ -28,17 +28,19 @@ from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Generic,
|
||||
Literal,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
get_args,
|
||||
get_type_hints,
|
||||
overload,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, RootModel
|
||||
from typing_extensions import Literal, get_args, override
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain_core._api import beta_decorator
|
||||
from langchain_core.callbacks.manager import AsyncCallbackManager, CallbackManager
|
||||
@@ -1653,7 +1655,7 @@ class Runnable(ABC, Generic[Input, Output]):
|
||||
from langchain_ollama import ChatOllama
|
||||
from langchain_core.output_parsers import StrOutputParser
|
||||
|
||||
llm = ChatOllama(model="llama2")
|
||||
llm = ChatOllama(model="llama3.1")
|
||||
|
||||
# Without bind.
|
||||
chain = llm | StrOutputParser()
|
||||
@@ -3068,50 +3070,10 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
|
||||
"""
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.beta.runnables.context import ( # noqa: PLC0415
|
||||
CONTEXT_CONFIG_PREFIX,
|
||||
_key_from_id,
|
||||
return get_unique_config_specs(
|
||||
[spec for step in self.steps for spec in step.config_specs]
|
||||
)
|
||||
|
||||
# get all specs
|
||||
all_specs = [
|
||||
(spec, idx)
|
||||
for idx, step in enumerate(self.steps)
|
||||
for spec in step.config_specs
|
||||
]
|
||||
# calculate context dependencies
|
||||
specs_by_pos = groupby(
|
||||
[tup for tup in all_specs if tup[0].id.startswith(CONTEXT_CONFIG_PREFIX)],
|
||||
itemgetter(1),
|
||||
)
|
||||
next_deps: set[str] = set()
|
||||
deps_by_pos: dict[int, set[str]] = {}
|
||||
for pos, specs in specs_by_pos:
|
||||
deps_by_pos[pos] = next_deps
|
||||
next_deps = next_deps | {spec[0].id for spec in specs}
|
||||
# assign context dependencies
|
||||
for pos, (spec, idx) in enumerate(all_specs):
|
||||
if spec.id.startswith(CONTEXT_CONFIG_PREFIX):
|
||||
all_specs[pos] = (
|
||||
ConfigurableFieldSpec(
|
||||
id=spec.id,
|
||||
annotation=spec.annotation,
|
||||
name=spec.name,
|
||||
default=spec.default,
|
||||
description=spec.description,
|
||||
is_shared=spec.is_shared,
|
||||
dependencies=[
|
||||
d
|
||||
for d in deps_by_pos[idx]
|
||||
if _key_from_id(d) != _key_from_id(spec.id)
|
||||
]
|
||||
+ (spec.dependencies or []),
|
||||
),
|
||||
idx,
|
||||
)
|
||||
|
||||
return get_unique_config_specs(spec for spec, _ in all_specs)
|
||||
|
||||
@override
|
||||
def get_graph(self, config: Optional[RunnableConfig] = None) -> Graph:
|
||||
"""Get the graph representation of the ``Runnable``.
|
||||
@@ -3215,13 +3177,8 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
def invoke(
|
||||
self, input: Input, config: Optional[RunnableConfig] = None, **kwargs: Any
|
||||
) -> Output:
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.beta.runnables.context import ( # noqa: PLC0415
|
||||
config_with_context,
|
||||
)
|
||||
|
||||
# setup callbacks and context
|
||||
config = config_with_context(ensure_config(config), self.steps)
|
||||
config = ensure_config(config)
|
||||
callback_manager = get_callback_manager_for_config(config)
|
||||
# start the root run
|
||||
run_manager = callback_manager.on_chain_start(
|
||||
@@ -3259,13 +3216,8 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
config: Optional[RunnableConfig] = None,
|
||||
**kwargs: Optional[Any],
|
||||
) -> Output:
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.beta.runnables.context import ( # noqa: PLC0415
|
||||
aconfig_with_context,
|
||||
)
|
||||
|
||||
# setup callbacks and context
|
||||
config = aconfig_with_context(ensure_config(config), self.steps)
|
||||
config = ensure_config(config)
|
||||
callback_manager = get_async_callback_manager_for_config(config)
|
||||
# start the root run
|
||||
run_manager = await callback_manager.on_chain_start(
|
||||
@@ -3306,19 +3258,11 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
return_exceptions: bool = False,
|
||||
**kwargs: Optional[Any],
|
||||
) -> list[Output]:
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.beta.runnables.context import ( # noqa: PLC0415
|
||||
config_with_context,
|
||||
)
|
||||
|
||||
if not inputs:
|
||||
return []
|
||||
|
||||
# setup callbacks and context
|
||||
configs = [
|
||||
config_with_context(c, self.steps)
|
||||
for c in get_config_list(config, len(inputs))
|
||||
]
|
||||
configs = get_config_list(config, len(inputs))
|
||||
callback_managers = [
|
||||
CallbackManager.configure(
|
||||
inheritable_callbacks=config.get("callbacks"),
|
||||
@@ -3438,19 +3382,11 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
return_exceptions: bool = False,
|
||||
**kwargs: Optional[Any],
|
||||
) -> list[Output]:
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.beta.runnables.context import ( # noqa: PLC0415
|
||||
aconfig_with_context,
|
||||
)
|
||||
|
||||
if not inputs:
|
||||
return []
|
||||
|
||||
# setup callbacks and context
|
||||
configs = [
|
||||
aconfig_with_context(c, self.steps)
|
||||
for c in get_config_list(config, len(inputs))
|
||||
]
|
||||
configs = get_config_list(config, len(inputs))
|
||||
callback_managers = [
|
||||
AsyncCallbackManager.configure(
|
||||
inheritable_callbacks=config.get("callbacks"),
|
||||
@@ -3571,14 +3507,7 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
config: RunnableConfig,
|
||||
**kwargs: Any,
|
||||
) -> Iterator[Output]:
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.beta.runnables.context import ( # noqa: PLC0415
|
||||
config_with_context,
|
||||
)
|
||||
|
||||
steps = [self.first, *self.middle, self.last]
|
||||
config = config_with_context(config, self.steps)
|
||||
|
||||
# transform the input stream of each step with the next
|
||||
# steps that don't natively support transforming an input stream will
|
||||
# buffer input in memory until all available, and then start emitting output
|
||||
@@ -3601,14 +3530,7 @@ class RunnableSequence(RunnableSerializable[Input, Output]):
|
||||
config: RunnableConfig,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[Output]:
|
||||
# Import locally to prevent circular import
|
||||
from langchain_core.beta.runnables.context import ( # noqa: PLC0415
|
||||
aconfig_with_context,
|
||||
)
|
||||
|
||||
steps = [self.first, *self.middle, self.last]
|
||||
config = aconfig_with_context(config, self.steps)
|
||||
|
||||
# stream the last steps
|
||||
# transform the input stream of each step with the next
|
||||
# steps that don't natively support transforming an input stream will
|
||||
|
||||
@@ -12,10 +12,6 @@ from typing import (
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain_core.beta.runnables.context import (
|
||||
CONTEXT_CONFIG_PREFIX,
|
||||
CONTEXT_CONFIG_SUFFIX_SET,
|
||||
)
|
||||
from langchain_core.runnables.base import (
|
||||
Runnable,
|
||||
RunnableLike,
|
||||
@@ -181,7 +177,7 @@ class RunnableBranch(RunnableSerializable[Input, Output]):
|
||||
@property
|
||||
@override
|
||||
def config_specs(self) -> list[ConfigurableFieldSpec]:
|
||||
specs = get_unique_config_specs(
|
||||
return get_unique_config_specs(
|
||||
spec
|
||||
for step in (
|
||||
[self.default]
|
||||
@@ -190,14 +186,6 @@ class RunnableBranch(RunnableSerializable[Input, Output]):
|
||||
)
|
||||
for spec in step.config_specs
|
||||
)
|
||||
if any(
|
||||
s.id.startswith(CONTEXT_CONFIG_PREFIX)
|
||||
and s.id.endswith(CONTEXT_CONFIG_SUFFIX_SET)
|
||||
for s in specs
|
||||
):
|
||||
msg = "RunnableBranch cannot contain context setters."
|
||||
raise ValueError(msg)
|
||||
return specs
|
||||
|
||||
@override
|
||||
def invoke(
|
||||
|
||||
@@ -10,10 +10,19 @@ from concurrent.futures import Executor, Future, ThreadPoolExecutor
|
||||
from contextlib import contextmanager
|
||||
from contextvars import Context, ContextVar, Token, copy_context
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any, Callable, Optional, TypeVar, Union, cast
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Optional,
|
||||
ParamSpec,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
from langsmith.run_helpers import _set_tracing_context, get_tracing_context
|
||||
from typing_extensions import ParamSpec, TypedDict
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langchain_core.callbacks.manager import AsyncCallbackManager, CallbackManager
|
||||
from langchain_core.runnables.utils import (
|
||||
|
||||
@@ -23,6 +23,13 @@ class EventData(TypedDict, total=False):
|
||||
won't be known until the *END* of the Runnable when it has finished streaming
|
||||
its inputs.
|
||||
"""
|
||||
error: NotRequired[BaseException]
|
||||
"""The error that occurred during the execution of the Runnable.
|
||||
|
||||
This field is only available if the Runnable raised an exception.
|
||||
|
||||
.. versionadded:: 1.0.0
|
||||
"""
|
||||
output: Any
|
||||
"""The output of the Runnable that generated the event.
|
||||
|
||||
|
||||
@@ -18,11 +18,12 @@ from typing import (
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypeGuard,
|
||||
TypeVar,
|
||||
Union,
|
||||
)
|
||||
|
||||
from typing_extensions import TypeGuard, override
|
||||
from typing_extensions import override
|
||||
|
||||
# Re-export create-model for backwards compatibility
|
||||
from langchain_core.utils.pydantic import create_model # noqa: F401
|
||||
|
||||
@@ -1272,7 +1272,7 @@ class InjectedToolCallId(InjectedToolArg):
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from typing_extensions import Annotated
|
||||
from typing import Annotated
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langchain_core.tools import tool, InjectedToolCallId
|
||||
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sys
|
||||
import traceback
|
||||
from abc import ABC, abstractmethod
|
||||
from datetime import datetime, timezone
|
||||
@@ -98,12 +97,7 @@ class _TracerCore(ABC):
|
||||
"""Get the stacktrace of the parent error."""
|
||||
msg = repr(error)
|
||||
try:
|
||||
if sys.version_info < (3, 10):
|
||||
tb = traceback.format_exception(
|
||||
error.__class__, error, error.__traceback__
|
||||
)
|
||||
else:
|
||||
tb = traceback.format_exception(error)
|
||||
tb = traceback.format_exception(error)
|
||||
return (msg + "\n\n".join(tb)).strip()
|
||||
except: # noqa: E722
|
||||
return msg
|
||||
|
||||
@@ -610,6 +610,28 @@ class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHand
|
||||
run_type,
|
||||
)
|
||||
|
||||
def _get_tool_run_info_with_inputs(self, run_id: UUID) -> tuple[RunInfo, Any]:
|
||||
"""Get run info for a tool and extract inputs, with validation.
|
||||
|
||||
Args:
|
||||
run_id: The run ID of the tool.
|
||||
|
||||
Returns:
|
||||
A tuple of (run_info, inputs).
|
||||
|
||||
Raises:
|
||||
AssertionError: If the run ID is a tool call and does not have inputs.
|
||||
"""
|
||||
run_info = self.run_map.pop(run_id)
|
||||
if "inputs" not in run_info:
|
||||
msg = (
|
||||
f"Run ID {run_id} is a tool call and is expected to have "
|
||||
f"inputs associated with it."
|
||||
)
|
||||
raise AssertionError(msg)
|
||||
inputs = run_info["inputs"]
|
||||
return run_info, inputs
|
||||
|
||||
@override
|
||||
async def on_tool_start(
|
||||
self,
|
||||
@@ -652,6 +674,35 @@ class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHand
|
||||
"tool",
|
||||
)
|
||||
|
||||
@override
|
||||
async def on_tool_error(
|
||||
self,
|
||||
error: BaseException,
|
||||
*,
|
||||
run_id: UUID,
|
||||
parent_run_id: Optional[UUID] = None,
|
||||
tags: Optional[list[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Run when tool errors."""
|
||||
run_info, inputs = self._get_tool_run_info_with_inputs(run_id)
|
||||
|
||||
self._send(
|
||||
{
|
||||
"event": "on_tool_error",
|
||||
"data": {
|
||||
"error": error,
|
||||
"input": inputs,
|
||||
},
|
||||
"run_id": str(run_id),
|
||||
"name": run_info["name"],
|
||||
"tags": run_info["tags"],
|
||||
"metadata": run_info["metadata"],
|
||||
"parent_ids": self._get_parent_ids(run_id),
|
||||
},
|
||||
"tool",
|
||||
)
|
||||
|
||||
@override
|
||||
async def on_tool_end(self, output: Any, *, run_id: UUID, **kwargs: Any) -> None:
|
||||
"""End a trace for a tool run.
|
||||
@@ -659,14 +710,7 @@ class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHand
|
||||
Raises:
|
||||
AssertionError: If the run ID is a tool call and does not have inputs
|
||||
"""
|
||||
run_info = self.run_map.pop(run_id)
|
||||
if "inputs" not in run_info:
|
||||
msg = (
|
||||
f"Run ID {run_id} is a tool call and is expected to have "
|
||||
f"inputs associated with it."
|
||||
)
|
||||
raise AssertionError(msg)
|
||||
inputs = run_info["inputs"]
|
||||
run_info, inputs = self._get_tool_run_info_with_inputs(run_id)
|
||||
|
||||
self._send(
|
||||
{
|
||||
|
||||
@@ -57,6 +57,11 @@ def merge_dicts(left: dict[str, Any], *others: dict[str, Any]) -> dict[str, Any]
|
||||
# "should either occur once or have the same value across "
|
||||
# "all dicts."
|
||||
# )
|
||||
if (right_k == "index" and merged[right_k].startswith("lc_")) or (
|
||||
right_k in ("id", "output_version", "model_provider")
|
||||
and merged[right_k] == right_v
|
||||
):
|
||||
continue
|
||||
merged[right_k] += right_v
|
||||
elif isinstance(merged[right_k], dict):
|
||||
merged[right_k] = merge_dicts(merged[right_k], right_v)
|
||||
@@ -93,7 +98,16 @@ def merge_lists(left: Optional[list], *others: Optional[list]) -> Optional[list]
|
||||
merged = other.copy()
|
||||
else:
|
||||
for e in other:
|
||||
if isinstance(e, dict) and "index" in e and isinstance(e["index"], int):
|
||||
if (
|
||||
isinstance(e, dict)
|
||||
and "index" in e
|
||||
and (
|
||||
isinstance(e["index"], int)
|
||||
or (
|
||||
isinstance(e["index"], str) and e["index"].startswith("lc_")
|
||||
)
|
||||
)
|
||||
):
|
||||
to_merge = [
|
||||
i
|
||||
for i, e_left in enumerate(merged)
|
||||
@@ -102,11 +116,35 @@ def merge_lists(left: Optional[list], *others: Optional[list]) -> Optional[list]
|
||||
if to_merge:
|
||||
# TODO: Remove this once merge_dict is updated with special
|
||||
# handling for 'type'.
|
||||
new_e = (
|
||||
{k: v for k, v in e.items() if k != "type"}
|
||||
if "type" in e
|
||||
else e
|
||||
)
|
||||
if (left_type := merged[to_merge[0]].get("type")) and (
|
||||
e.get("type") == "non_standard" and "value" in e
|
||||
):
|
||||
if left_type != "non_standard":
|
||||
# standard + non_standard
|
||||
new_e: dict[str, Any] = {
|
||||
"extras": {
|
||||
k: v
|
||||
for k, v in e["value"].items()
|
||||
if k != "type"
|
||||
}
|
||||
}
|
||||
else:
|
||||
# non_standard + non_standard
|
||||
new_e = {
|
||||
"value": {
|
||||
k: v
|
||||
for k, v in e["value"].items()
|
||||
if k != "type"
|
||||
}
|
||||
}
|
||||
if "index" in e:
|
||||
new_e["index"] = e["index"]
|
||||
else:
|
||||
new_e = (
|
||||
{k: v for k, v in e.items() if k != "type"}
|
||||
if "type" in e
|
||||
else e
|
||||
)
|
||||
merged[to_merge[0]] = merge_dicts(merged[to_merge[0]], new_e)
|
||||
else:
|
||||
merged.append(e)
|
||||
|
||||
@@ -17,12 +17,15 @@ from typing import (
|
||||
Optional,
|
||||
Union,
|
||||
cast,
|
||||
get_args,
|
||||
get_origin,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic.v1 import BaseModel as BaseModelV1
|
||||
from pydantic.v1 import Field, create_model
|
||||
from typing_extensions import TypedDict, get_args, get_origin, is_typeddict
|
||||
from pydantic.v1 import Field as Field_v1
|
||||
from pydantic.v1 import create_model as create_model_v1
|
||||
from typing_extensions import TypedDict, is_typeddict
|
||||
|
||||
import langchain_core
|
||||
from langchain_core._api import beta, deprecated
|
||||
@@ -294,7 +297,7 @@ def _convert_any_typed_dicts_to_pydantic(
|
||||
raise ValueError(msg)
|
||||
if arg_desc := arg_descriptions.get(arg):
|
||||
field_kwargs["description"] = arg_desc
|
||||
fields[arg] = (new_arg_type, Field(**field_kwargs))
|
||||
fields[arg] = (new_arg_type, Field_v1(**field_kwargs))
|
||||
else:
|
||||
new_arg_type = _convert_any_typed_dicts_to_pydantic(
|
||||
arg_type, depth=depth + 1, visited=visited
|
||||
@@ -302,8 +305,8 @@ def _convert_any_typed_dicts_to_pydantic(
|
||||
field_kwargs = {"default": ...}
|
||||
if arg_desc := arg_descriptions.get(arg):
|
||||
field_kwargs["description"] = arg_desc
|
||||
fields[arg] = (new_arg_type, Field(**field_kwargs))
|
||||
model = create_model(typed_dict.__name__, **fields)
|
||||
fields[arg] = (new_arg_type, Field_v1(**field_kwargs))
|
||||
model = create_model_v1(typed_dict.__name__, **fields)
|
||||
model.__doc__ = description
|
||||
visited[typed_dict] = model
|
||||
return model
|
||||
|
||||
@@ -8,14 +8,13 @@ from types import TracebackType
|
||||
from typing import (
|
||||
Any,
|
||||
Generic,
|
||||
Literal,
|
||||
Optional,
|
||||
TypeVar,
|
||||
Union,
|
||||
overload,
|
||||
)
|
||||
|
||||
from typing_extensions import Literal
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from typing import (
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from typing_extensions import TypeAlias
|
||||
from typing import TypeAlias
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ import warnings
|
||||
from collections.abc import Iterator, Sequence
|
||||
from importlib.metadata import version
|
||||
from typing import Any, Callable, Optional, Union, overload
|
||||
from uuid import uuid4
|
||||
|
||||
from packaging.version import parse
|
||||
from pydantic import SecretStr
|
||||
@@ -482,3 +483,31 @@ def secret_from_env(
|
||||
raise ValueError(msg)
|
||||
|
||||
return get_secret_from_env
|
||||
|
||||
|
||||
LC_AUTO_PREFIX = "lc_"
|
||||
"""LangChain auto-generated ID prefix for messages and content blocks."""
|
||||
|
||||
LC_ID_PREFIX = "lc_run-"
|
||||
"""Internal tracing/callback system identifier.
|
||||
|
||||
Used for:
|
||||
- Tracing. Every LangChain operation (LLM call, chain execution, tool use, etc.)
|
||||
gets a unique run_id (UUID)
|
||||
- Enables tracking parent-child relationships between operations
|
||||
"""
|
||||
|
||||
|
||||
def ensure_id(id_val: Optional[str]) -> str:
|
||||
"""Ensure the ID is a valid string, generating a new UUID if not provided.
|
||||
|
||||
Auto-generated UUIDs are prefixed by ``'lc_'`` to indicate they are
|
||||
LangChain-generated IDs.
|
||||
|
||||
Args:
|
||||
id_val: Optional string ID value to validate.
|
||||
|
||||
Returns:
|
||||
A string ID, either the validated provided value or a newly generated UUID4.
|
||||
"""
|
||||
return id_val or str(f"{LC_AUTO_PREFIX}{uuid4()}")
|
||||
@@ -603,7 +603,7 @@ class InMemoryVectorStore(VectorStore):
|
||||
A VectorStore object.
|
||||
"""
|
||||
path_: Path = Path(path)
|
||||
with path_.open("r") as f:
|
||||
with path_.open("r", encoding="utf-8") as f:
|
||||
store = load(json.load(f))
|
||||
vectorstore = cls(embedding=embedding, **kwargs)
|
||||
vectorstore.store = store
|
||||
@@ -617,5 +617,5 @@ class InMemoryVectorStore(VectorStore):
|
||||
"""
|
||||
path_: Path = Path(path)
|
||||
path_.parent.mkdir(exist_ok=True, parents=True)
|
||||
with path_.open("w") as f:
|
||||
with path_.open("w", encoding="utf-8") as f:
|
||||
json.dump(dumpd(self.store), f, indent=2)
|
||||
@@ -1,3 +1,3 @@
|
||||
"""langchain-core version information and utilities."""
|
||||
|
||||
VERSION = "0.3.77"
|
||||
VERSION = "1.0.0a6"
|
||||
@@ -5,7 +5,7 @@ build-backend = "pdm.backend"
|
||||
[project]
|
||||
authors = []
|
||||
license = {text = "MIT"}
|
||||
requires-python = ">=3.9.0,<4.0.0"
|
||||
requires-python = ">=3.10.0,<4.0.0"
|
||||
dependencies = [
|
||||
"langsmith>=0.3.45,<1.0.0",
|
||||
"tenacity!=8.4.0,>=8.1.0,<10.0.0",
|
||||
@@ -16,7 +16,7 @@ dependencies = [
|
||||
"pydantic>=2.7.4,<3.0.0",
|
||||
]
|
||||
name = "langchain-core"
|
||||
version = "0.3.77"
|
||||
version = "1.0.0a6"
|
||||
description = "Building applications with LLMs through composability"
|
||||
readme = "README.md"
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
import time
|
||||
from itertools import cycle
|
||||
from typing import Any, Optional, Union
|
||||
from typing import Any, Optional, Union, cast
|
||||
from uuid import UUID
|
||||
|
||||
from typing_extensions import override
|
||||
@@ -59,7 +59,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
assert chunks == [
|
||||
_any_id_ai_message_chunk(content="hello"),
|
||||
_any_id_ai_message_chunk(content=" "),
|
||||
_any_id_ai_message_chunk(content="goodbye"),
|
||||
_any_id_ai_message_chunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in chunks}) == 1
|
||||
|
||||
@@ -67,7 +67,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
assert chunks == [
|
||||
_any_id_ai_message_chunk(content="hello"),
|
||||
_any_id_ai_message_chunk(content=" "),
|
||||
_any_id_ai_message_chunk(content="goodbye"),
|
||||
_any_id_ai_message_chunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in chunks}) == 1
|
||||
|
||||
@@ -79,6 +79,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
assert chunks == [
|
||||
_any_id_ai_message_chunk(content="", additional_kwargs={"foo": 42}),
|
||||
_any_id_ai_message_chunk(content="", additional_kwargs={"bar": 24}),
|
||||
_any_id_ai_message_chunk(content="", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in chunks}) == 1
|
||||
|
||||
@@ -97,7 +98,8 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
|
||||
assert chunks == [
|
||||
_any_id_ai_message_chunk(
|
||||
content="", additional_kwargs={"function_call": {"name": "move_file"}}
|
||||
content="",
|
||||
additional_kwargs={"function_call": {"name": "move_file"}},
|
||||
),
|
||||
_any_id_ai_message_chunk(
|
||||
content="",
|
||||
@@ -114,6 +116,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
"function_call": {"arguments": '\n "destination_path": "bar"\n}'},
|
||||
},
|
||||
),
|
||||
_any_id_ai_message_chunk(content="", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in chunks}) == 1
|
||||
|
||||
@@ -134,6 +137,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
}
|
||||
},
|
||||
id=chunks[0].id,
|
||||
chunk_position="last",
|
||||
)
|
||||
|
||||
|
||||
@@ -148,7 +152,7 @@ async def test_generic_fake_chat_model_astream_log() -> None:
|
||||
assert final.state["streamed_output"] == [
|
||||
_any_id_ai_message_chunk(content="hello"),
|
||||
_any_id_ai_message_chunk(content=" "),
|
||||
_any_id_ai_message_chunk(content="goodbye"),
|
||||
_any_id_ai_message_chunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in final.state["streamed_output"]}) == 1
|
||||
|
||||
@@ -205,7 +209,7 @@ async def test_callback_handlers() -> None:
|
||||
assert results == [
|
||||
_any_id_ai_message_chunk(content="hello"),
|
||||
_any_id_ai_message_chunk(content=" "),
|
||||
_any_id_ai_message_chunk(content="goodbye"),
|
||||
_any_id_ai_message_chunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
assert tokens == ["hello", " ", "goodbye"]
|
||||
assert len({chunk.id for chunk in results}) == 1
|
||||
@@ -214,7 +218,9 @@ async def test_callback_handlers() -> None:
|
||||
def test_chat_model_inputs() -> None:
|
||||
fake = ParrotFakeChatModel()
|
||||
|
||||
assert fake.invoke("hello") == _any_id_human_message(content="hello")
|
||||
assert cast("HumanMessage", fake.invoke("hello")) == _any_id_human_message(
|
||||
content="hello"
|
||||
)
|
||||
assert fake.invoke([("ai", "blah")]) == _any_id_ai_message(content="blah")
|
||||
assert fake.invoke([AIMessage(content="blah")]) == _any_id_ai_message(
|
||||
content="blah"
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""Test base chat model."""
|
||||
|
||||
import uuid
|
||||
import warnings
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional, Union
|
||||
|
||||
@@ -14,11 +15,15 @@ from langchain_core.language_models import (
|
||||
ParrotFakeChatModel,
|
||||
)
|
||||
from langchain_core.language_models._utils import _normalize_messages
|
||||
from langchain_core.language_models.fake_chat_models import FakeListChatModelError
|
||||
from langchain_core.language_models.fake_chat_models import (
|
||||
FakeListChatModelError,
|
||||
GenericFakeChatModel,
|
||||
)
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AIMessageChunk,
|
||||
BaseMessage,
|
||||
BaseMessageChunk,
|
||||
HumanMessage,
|
||||
SystemMessage,
|
||||
)
|
||||
@@ -40,6 +45,37 @@ if TYPE_CHECKING:
|
||||
from langchain_core.outputs.llm_result import LLMResult
|
||||
|
||||
|
||||
def _content_blocks_equal_ignore_id(
|
||||
actual: Union[str, list[Any]], expected: Union[str, list[Any]]
|
||||
) -> bool:
|
||||
"""Compare content blocks, ignoring auto-generated `id` fields.
|
||||
|
||||
Args:
|
||||
actual: Actual content from response (string or list of content blocks).
|
||||
expected: Expected content to compare against (string or list of blocks).
|
||||
|
||||
Returns:
|
||||
True if content matches (excluding `id` fields), False otherwise.
|
||||
|
||||
"""
|
||||
if isinstance(actual, str) or isinstance(expected, str):
|
||||
return actual == expected
|
||||
|
||||
if len(actual) != len(expected):
|
||||
return False
|
||||
for actual_block, expected_block in zip(actual, expected):
|
||||
actual_without_id = (
|
||||
{k: v for k, v in actual_block.items() if k != "id"}
|
||||
if isinstance(actual_block, dict) and "id" in actual_block
|
||||
else actual_block
|
||||
)
|
||||
|
||||
if actual_without_id != expected_block:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def messages() -> list:
|
||||
return [
|
||||
@@ -141,7 +177,7 @@ async def test_stream_error_callback() -> None:
|
||||
|
||||
|
||||
async def test_astream_fallback_to_ainvoke() -> None:
|
||||
"""Test astream uses appropriate implementation."""
|
||||
"""Test `astream()` uses appropriate implementation."""
|
||||
|
||||
class ModelWithGenerate(BaseChatModel):
|
||||
@override
|
||||
@@ -168,10 +204,10 @@ async def test_astream_fallback_to_ainvoke() -> None:
|
||||
# is not strictly correct.
|
||||
# LangChain documents a pattern of adding BaseMessageChunks to accumulate a stream.
|
||||
# This may be better done with `reduce(operator.add, chunks)`.
|
||||
assert chunks == [_any_id_ai_message(content="hello")] # type: ignore[comparison-overlap]
|
||||
assert chunks == [_any_id_ai_message(content="hello")]
|
||||
|
||||
chunks = [chunk async for chunk in model.astream("anything")]
|
||||
assert chunks == [_any_id_ai_message(content="hello")] # type: ignore[comparison-overlap]
|
||||
assert chunks == [_any_id_ai_message(content="hello")]
|
||||
|
||||
|
||||
async def test_astream_implementation_fallback_to_stream() -> None:
|
||||
@@ -198,7 +234,9 @@ async def test_astream_implementation_fallback_to_stream() -> None:
|
||||
) -> Iterator[ChatGenerationChunk]:
|
||||
"""Stream the output of the model."""
|
||||
yield ChatGenerationChunk(message=AIMessageChunk(content="a"))
|
||||
yield ChatGenerationChunk(message=AIMessageChunk(content="b"))
|
||||
yield ChatGenerationChunk(
|
||||
message=AIMessageChunk(content="b", chunk_position="last")
|
||||
)
|
||||
|
||||
@property
|
||||
def _llm_type(self) -> str:
|
||||
@@ -207,15 +245,19 @@ async def test_astream_implementation_fallback_to_stream() -> None:
|
||||
model = ModelWithSyncStream()
|
||||
chunks = list(model.stream("anything"))
|
||||
assert chunks == [
|
||||
_any_id_ai_message_chunk(content="a"),
|
||||
_any_id_ai_message_chunk(content="b"),
|
||||
_any_id_ai_message_chunk(
|
||||
content="a",
|
||||
),
|
||||
_any_id_ai_message_chunk(content="b", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in chunks}) == 1
|
||||
assert type(model)._astream == BaseChatModel._astream
|
||||
astream_chunks = [chunk async for chunk in model.astream("anything")]
|
||||
assert astream_chunks == [
|
||||
_any_id_ai_message_chunk(content="a"),
|
||||
_any_id_ai_message_chunk(content="b"),
|
||||
_any_id_ai_message_chunk(
|
||||
content="a",
|
||||
),
|
||||
_any_id_ai_message_chunk(content="b", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in astream_chunks}) == 1
|
||||
|
||||
@@ -244,7 +286,9 @@ async def test_astream_implementation_uses_astream() -> None:
|
||||
) -> AsyncIterator[ChatGenerationChunk]:
|
||||
"""Stream the output of the model."""
|
||||
yield ChatGenerationChunk(message=AIMessageChunk(content="a"))
|
||||
yield ChatGenerationChunk(message=AIMessageChunk(content="b"))
|
||||
yield ChatGenerationChunk(
|
||||
message=AIMessageChunk(content="b", chunk_position="last")
|
||||
)
|
||||
|
||||
@property
|
||||
def _llm_type(self) -> str:
|
||||
@@ -253,8 +297,10 @@ async def test_astream_implementation_uses_astream() -> None:
|
||||
model = ModelWithAsyncStream()
|
||||
chunks = [chunk async for chunk in model.astream("anything")]
|
||||
assert chunks == [
|
||||
_any_id_ai_message_chunk(content="a"),
|
||||
_any_id_ai_message_chunk(content="b"),
|
||||
_any_id_ai_message_chunk(
|
||||
content="a",
|
||||
),
|
||||
_any_id_ai_message_chunk(content="b", chunk_position="last"),
|
||||
]
|
||||
assert len({chunk.id for chunk in chunks}) == 1
|
||||
|
||||
@@ -427,11 +473,12 @@ class FakeChatModelStartTracer(FakeTracer):
|
||||
|
||||
|
||||
def test_trace_images_in_openai_format() -> None:
|
||||
"""Test that images are traced in OpenAI format."""
|
||||
"""Test that images are traced in OpenAI Chat Completions format."""
|
||||
llm = ParrotFakeChatModel()
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
# v0 format
|
||||
"content": [
|
||||
{
|
||||
"type": "image",
|
||||
@@ -442,7 +489,7 @@ def test_trace_images_in_openai_format() -> None:
|
||||
}
|
||||
]
|
||||
tracer = FakeChatModelStartTracer()
|
||||
response = llm.invoke(messages, config={"callbacks": [tracer]})
|
||||
llm.invoke(messages, config={"callbacks": [tracer]})
|
||||
assert tracer.messages == [
|
||||
[
|
||||
[
|
||||
@@ -457,19 +504,90 @@ def test_trace_images_in_openai_format() -> None:
|
||||
]
|
||||
]
|
||||
]
|
||||
# Test no mutation
|
||||
assert response.content == [
|
||||
|
||||
|
||||
def test_trace_pdfs() -> None:
|
||||
# For backward compat
|
||||
llm = ParrotFakeChatModel()
|
||||
messages = [
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://example.com/image.png",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "file",
|
||||
"mime_type": "application/pdf",
|
||||
"base64": "<base64 string>",
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
tracer = FakeChatModelStartTracer()
|
||||
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("error")
|
||||
llm.invoke(messages, config={"callbacks": [tracer]})
|
||||
|
||||
assert tracer.messages == [
|
||||
[
|
||||
[
|
||||
HumanMessage(
|
||||
content=[
|
||||
{
|
||||
"type": "file",
|
||||
"mime_type": "application/pdf",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 string>",
|
||||
}
|
||||
]
|
||||
)
|
||||
]
|
||||
]
|
||||
]
|
||||
|
||||
|
||||
def test_trace_content_blocks_with_no_type_key() -> None:
|
||||
"""Test that we add a ``type`` key to certain content blocks that don't have one."""
|
||||
llm = ParrotFakeChatModel()
|
||||
def test_content_block_transformation_v0_to_v1_image() -> None:
|
||||
"""Test that v0 format image content blocks are transformed to v1 format."""
|
||||
# Create a message with v0 format image content
|
||||
image_message = AIMessage(
|
||||
content=[
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://example.com/image.png",
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
llm = GenericFakeChatModel(messages=iter([image_message]), output_version="v1")
|
||||
response = llm.invoke("test")
|
||||
|
||||
# With v1 output_version, .content should be transformed
|
||||
# Check structure, ignoring auto-generated IDs
|
||||
assert len(response.content) == 1
|
||||
content_block = response.content[0]
|
||||
if isinstance(content_block, dict) and "id" in content_block:
|
||||
# Remove auto-generated id for comparison
|
||||
content_without_id = {k: v for k, v in content_block.items() if k != "id"}
|
||||
expected_content = {
|
||||
"type": "image",
|
||||
"url": "https://example.com/image.png",
|
||||
}
|
||||
assert content_without_id == expected_content
|
||||
else:
|
||||
assert content_block == {
|
||||
"type": "image",
|
||||
"url": "https://example.com/image.png",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("output_version", ["v0", "v1"])
|
||||
def test_trace_content_blocks_with_no_type_key(output_version: str) -> None:
|
||||
"""Test behavior of content blocks that don't have a `type` key.
|
||||
|
||||
Only for blocks with one key, in which case, the name of the key is used as `type`.
|
||||
|
||||
"""
|
||||
llm = ParrotFakeChatModel(output_version=output_version)
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
@@ -504,155 +622,381 @@ def test_trace_content_blocks_with_no_type_key() -> None:
|
||||
]
|
||||
]
|
||||
]
|
||||
# Test no mutation
|
||||
assert response.content == [
|
||||
|
||||
if output_version == "v0":
|
||||
assert response.content == [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Hello",
|
||||
},
|
||||
{
|
||||
"cachePoint": {"type": "default"},
|
||||
},
|
||||
]
|
||||
else:
|
||||
assert response.content == [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Hello",
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {
|
||||
"cachePoint": {"type": "default"},
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
assert response.content_blocks == [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Hello",
|
||||
},
|
||||
{
|
||||
"cachePoint": {"type": "default"},
|
||||
"type": "non_standard",
|
||||
"value": {
|
||||
"cachePoint": {"type": "default"},
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def test_extend_support_to_openai_multimodal_formats() -> None:
|
||||
"""Test that chat models normalize OpenAI file and audio inputs."""
|
||||
llm = ParrotFakeChatModel()
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Hello"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/image.png"},
|
||||
"""Test normalizing OpenAI audio, image, and file inputs to v1."""
|
||||
# Audio and file only (chat model default)
|
||||
messages = HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Hello"},
|
||||
{ # audio-base64
|
||||
"type": "input_audio",
|
||||
"input_audio": {
|
||||
"format": "wav",
|
||||
"data": "<base64 string>",
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQSkZJRg..."},
|
||||
},
|
||||
{ # file-base64
|
||||
"type": "file",
|
||||
"file": {
|
||||
"filename": "draconomicon.pdf",
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {
|
||||
"filename": "draconomicon.pdf",
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {"file_id": "<file id>"},
|
||||
},
|
||||
{
|
||||
"type": "input_audio",
|
||||
"input_audio": {"data": "<base64 data>", "format": "wav"},
|
||||
},
|
||||
],
|
||||
},
|
||||
]
|
||||
expected_content = [
|
||||
{"type": "text", "text": "Hello"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/image.png"},
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQSkZJRg..."},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
"filename": "draconomicon.pdf",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {"file_id": "<file id>"},
|
||||
},
|
||||
{
|
||||
"type": "audio",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"mime_type": "audio/wav",
|
||||
},
|
||||
]
|
||||
response = llm.invoke(messages)
|
||||
assert response.content == expected_content
|
||||
},
|
||||
{ # file-id
|
||||
"type": "file",
|
||||
"file": {"file_id": "<file id>"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
# Test no mutation
|
||||
assert messages[0]["content"] == [
|
||||
{"type": "text", "text": "Hello"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/image.png"},
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQSkZJRg..."},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {
|
||||
"filename": "draconomicon.pdf",
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
expected_content_messages = HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Hello"}, # TextContentBlock
|
||||
{ # AudioContentBlock
|
||||
"type": "audio",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "audio/wav",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
{ # FileContentBlock
|
||||
"type": "file",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
"extras": {"filename": "draconomicon.pdf"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {"file_id": "<file id>"},
|
||||
},
|
||||
{
|
||||
"type": "input_audio",
|
||||
"input_audio": {"data": "<base64 data>", "format": "wav"},
|
||||
},
|
||||
]
|
||||
{ # ...
|
||||
"type": "file",
|
||||
"file_id": "<file id>",
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
normalized_content = _normalize_messages([messages])
|
||||
|
||||
# Check structure, ignoring auto-generated IDs
|
||||
assert len(normalized_content) == 1
|
||||
normalized_message = normalized_content[0]
|
||||
assert len(normalized_message.content) == len(expected_content_messages.content)
|
||||
|
||||
assert _content_blocks_equal_ignore_id(
|
||||
normalized_message.content, expected_content_messages.content
|
||||
)
|
||||
|
||||
messages = HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Hello"},
|
||||
{ # image-url
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/image.png"},
|
||||
},
|
||||
{ # image-base64
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQSkZJRg..."},
|
||||
},
|
||||
{ # audio-base64
|
||||
"type": "input_audio",
|
||||
"input_audio": {
|
||||
"format": "wav",
|
||||
"data": "<base64 string>",
|
||||
},
|
||||
},
|
||||
{ # file-base64
|
||||
"type": "file",
|
||||
"file": {
|
||||
"filename": "draconomicon.pdf",
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
},
|
||||
},
|
||||
{ # file-id
|
||||
"type": "file",
|
||||
"file": {"file_id": "<file id>"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
expected_content_messages = HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Hello"}, # TextContentBlock
|
||||
{ # image-url passes through
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/image.png"},
|
||||
},
|
||||
{ # image-url passes through with inline data
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQSkZJRg..."},
|
||||
},
|
||||
{ # AudioContentBlock
|
||||
"type": "audio",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "audio/wav",
|
||||
},
|
||||
{ # FileContentBlock
|
||||
"type": "file",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
"extras": {"filename": "draconomicon.pdf"},
|
||||
},
|
||||
{ # ...
|
||||
"type": "file",
|
||||
"file_id": "<file id>",
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
normalized_content = _normalize_messages([messages])
|
||||
|
||||
# Check structure, ignoring auto-generated IDs
|
||||
assert len(normalized_content) == 1
|
||||
normalized_message = normalized_content[0]
|
||||
assert len(normalized_message.content) == len(expected_content_messages.content)
|
||||
|
||||
assert _content_blocks_equal_ignore_id(
|
||||
normalized_message.content, expected_content_messages.content
|
||||
)
|
||||
|
||||
|
||||
def test_normalize_messages_edge_cases() -> None:
|
||||
# Test some blocks that should pass through
|
||||
# Test behavior of malformed/unrecognized content blocks
|
||||
|
||||
messages = [
|
||||
HumanMessage(
|
||||
content=[
|
||||
{
|
||||
"type": "file",
|
||||
"file": "uri",
|
||||
"type": "input_image", # Responses API type; not handled
|
||||
"image_url": "uri",
|
||||
},
|
||||
{
|
||||
"type": "input_file",
|
||||
# Standard OpenAI Chat Completions type but malformed structure
|
||||
"type": "input_audio",
|
||||
"input_audio": "uri", # Should be nested in `audio`
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": "uri", # `file` should be a dict for Chat Completions
|
||||
},
|
||||
{
|
||||
"type": "input_file", # Responses API type; not handled
|
||||
"file_data": "uri",
|
||||
"filename": "file-name",
|
||||
},
|
||||
{
|
||||
"type": "input_audio",
|
||||
"input_audio": "uri",
|
||||
},
|
||||
{
|
||||
"type": "input_image",
|
||||
"image_url": "uri",
|
||||
},
|
||||
]
|
||||
)
|
||||
]
|
||||
|
||||
assert messages == _normalize_messages(messages)
|
||||
|
||||
|
||||
def test_normalize_messages_v1_content_blocks_unchanged() -> None:
|
||||
"""Test passing v1 content blocks to `_normalize_messages()` leaves unchanged."""
|
||||
input_messages = [
|
||||
HumanMessage(
|
||||
content=[
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Hello world",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"url": "https://example.com/image.png",
|
||||
"mime_type": "image/png",
|
||||
},
|
||||
{
|
||||
"type": "audio",
|
||||
"base64": "base64encodedaudiodata",
|
||||
"mime_type": "audio/wav",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"id": "file_123",
|
||||
},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"reasoning": "Let me think about this...",
|
||||
},
|
||||
]
|
||||
)
|
||||
]
|
||||
|
||||
result = _normalize_messages(input_messages)
|
||||
|
||||
# Verify the result is identical to the input (message should not be copied)
|
||||
assert len(result) == 1
|
||||
assert result[0] is input_messages[0]
|
||||
assert result[0].content == input_messages[0].content
|
||||
|
||||
|
||||
def test_output_version_invoke(monkeypatch: Any) -> None:
|
||||
messages = [AIMessage("hello")]
|
||||
|
||||
llm = GenericFakeChatModel(messages=iter(messages), output_version="v1")
|
||||
response = llm.invoke("hello")
|
||||
assert response.content == [{"type": "text", "text": "hello"}]
|
||||
assert response.response_metadata["output_version"] == "v1"
|
||||
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
response = llm.invoke("hello")
|
||||
assert response.content == "hello"
|
||||
|
||||
monkeypatch.setenv("LC_OUTPUT_VERSION", "v1")
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
response = llm.invoke("hello")
|
||||
assert response.content == [{"type": "text", "text": "hello"}]
|
||||
assert response.response_metadata["output_version"] == "v1"
|
||||
|
||||
|
||||
# -- v1 output version tests --
|
||||
|
||||
|
||||
async def test_output_version_ainvoke(monkeypatch: Any) -> None:
|
||||
messages = [AIMessage("hello")]
|
||||
|
||||
# v0
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
response = await llm.ainvoke("hello")
|
||||
assert response.content == "hello"
|
||||
|
||||
# v1
|
||||
llm = GenericFakeChatModel(messages=iter(messages), output_version="v1")
|
||||
response = await llm.ainvoke("hello")
|
||||
assert response.content == [{"type": "text", "text": "hello"}]
|
||||
assert response.response_metadata["output_version"] == "v1"
|
||||
|
||||
# v1 from env var
|
||||
monkeypatch.setenv("LC_OUTPUT_VERSION", "v1")
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
response = await llm.ainvoke("hello")
|
||||
assert response.content == [{"type": "text", "text": "hello"}]
|
||||
assert response.response_metadata["output_version"] == "v1"
|
||||
|
||||
|
||||
def test_output_version_stream(monkeypatch: Any) -> None:
|
||||
messages = [AIMessage("foo bar")]
|
||||
|
||||
# v0
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
full = None
|
||||
for chunk in llm.stream("hello"):
|
||||
assert isinstance(chunk, AIMessageChunk)
|
||||
assert isinstance(chunk.content, str)
|
||||
assert chunk.content
|
||||
full = chunk if full is None else full + chunk
|
||||
assert isinstance(full, AIMessageChunk)
|
||||
assert full.content == "foo bar"
|
||||
|
||||
# v1
|
||||
llm = GenericFakeChatModel(messages=iter(messages), output_version="v1")
|
||||
full_v1: Optional[BaseMessageChunk] = None
|
||||
for chunk in llm.stream("hello"):
|
||||
assert isinstance(chunk, AIMessageChunk)
|
||||
assert isinstance(chunk.content, list)
|
||||
assert len(chunk.content) == 1
|
||||
block = chunk.content[0]
|
||||
assert isinstance(block, dict)
|
||||
assert block["type"] == "text"
|
||||
assert block["text"]
|
||||
full_v1 = chunk if full_v1 is None else full_v1 + chunk
|
||||
assert isinstance(full_v1, AIMessageChunk)
|
||||
assert full_v1.response_metadata["output_version"] == "v1"
|
||||
|
||||
# v1 from env var
|
||||
monkeypatch.setenv("LC_OUTPUT_VERSION", "v1")
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
full_env = None
|
||||
for chunk in llm.stream("hello"):
|
||||
assert isinstance(chunk, AIMessageChunk)
|
||||
assert isinstance(chunk.content, list)
|
||||
assert len(chunk.content) == 1
|
||||
block = chunk.content[0]
|
||||
assert isinstance(block, dict)
|
||||
assert block["type"] == "text"
|
||||
assert block["text"]
|
||||
full_env = chunk if full_env is None else full_env + chunk
|
||||
assert isinstance(full_env, AIMessageChunk)
|
||||
assert full_env.response_metadata["output_version"] == "v1"
|
||||
|
||||
|
||||
async def test_output_version_astream(monkeypatch: Any) -> None:
|
||||
messages = [AIMessage("foo bar")]
|
||||
|
||||
# v0
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
full = None
|
||||
async for chunk in llm.astream("hello"):
|
||||
assert isinstance(chunk, AIMessageChunk)
|
||||
assert isinstance(chunk.content, str)
|
||||
assert chunk.content
|
||||
full = chunk if full is None else full + chunk
|
||||
assert isinstance(full, AIMessageChunk)
|
||||
assert full.content == "foo bar"
|
||||
|
||||
# v1
|
||||
llm = GenericFakeChatModel(messages=iter(messages), output_version="v1")
|
||||
full_v1: Optional[BaseMessageChunk] = None
|
||||
async for chunk in llm.astream("hello"):
|
||||
assert isinstance(chunk, AIMessageChunk)
|
||||
assert isinstance(chunk.content, list)
|
||||
assert len(chunk.content) == 1
|
||||
block = chunk.content[0]
|
||||
assert isinstance(block, dict)
|
||||
assert block["type"] == "text"
|
||||
assert block["text"]
|
||||
full_v1 = chunk if full_v1 is None else full_v1 + chunk
|
||||
assert isinstance(full_v1, AIMessageChunk)
|
||||
assert full_v1.response_metadata["output_version"] == "v1"
|
||||
|
||||
# v1 from env var
|
||||
monkeypatch.setenv("LC_OUTPUT_VERSION", "v1")
|
||||
llm = GenericFakeChatModel(messages=iter(messages))
|
||||
full_env = None
|
||||
async for chunk in llm.astream("hello"):
|
||||
assert isinstance(chunk, AIMessageChunk)
|
||||
assert isinstance(chunk.content, list)
|
||||
assert len(chunk.content) == 1
|
||||
block = chunk.content[0]
|
||||
assert isinstance(block, dict)
|
||||
assert block["type"] == "text"
|
||||
assert block["text"]
|
||||
full_env = chunk if full_env is None else full_env + chunk
|
||||
assert isinstance(full_env, AIMessageChunk)
|
||||
assert full_env.response_metadata["output_version"] == "v1"
|
||||
assert messages == _normalize_messages(messages)
|
||||
|
||||
|
||||
|
||||
@@ -215,8 +215,8 @@ def test_rate_limit_skips_cache() -> None:
|
||||
assert list(cache._cache) == [
|
||||
(
|
||||
'[{"lc": 1, "type": "constructor", "id": ["langchain", "schema", '
|
||||
'"messages", '
|
||||
'"HumanMessage"], "kwargs": {"content": "foo", "type": "human"}}]',
|
||||
'"messages", "HumanMessage"], "kwargs": {"content": "foo", '
|
||||
'"type": "human"}}]',
|
||||
"[('_type', 'generic-fake-chat-model'), ('stop', None)]",
|
||||
)
|
||||
]
|
||||
@@ -240,7 +240,8 @@ def test_serialization_with_rate_limiter() -> None:
|
||||
assert InMemoryRateLimiter.__name__ not in serialized_model
|
||||
|
||||
|
||||
async def test_rate_limit_skips_cache_async() -> None:
|
||||
@pytest.mark.parametrize("output_version", ["v0", "v1"])
|
||||
async def test_rate_limit_skips_cache_async(output_version: str) -> None:
|
||||
"""Test that rate limiting does not rate limit cache look ups."""
|
||||
cache = InMemoryCache()
|
||||
model = GenericFakeChatModel(
|
||||
@@ -249,6 +250,7 @@ async def test_rate_limit_skips_cache_async() -> None:
|
||||
requests_per_second=20, check_every_n_seconds=0.1, max_bucket_size=1
|
||||
),
|
||||
cache=cache,
|
||||
output_version=output_version,
|
||||
)
|
||||
|
||||
tic = time.time()
|
||||
|
||||
@@ -18,6 +18,7 @@ EXPECTED_ALL = [
|
||||
"FakeStreamingListLLM",
|
||||
"FakeListLLM",
|
||||
"ParrotFakeChatModel",
|
||||
"is_openai_data_block",
|
||||
]
|
||||
|
||||
|
||||
|
||||
Whitespace-only changes.
@@ -0,0 +1,489 @@
|
||||
from typing import Optional
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk, HumanMessage
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
|
||||
def test_convert_to_v1_from_anthropic() -> None:
|
||||
message = AIMessage(
|
||||
[
|
||||
{"type": "thinking", "thinking": "foo", "signature": "foo_signature"},
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"input": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "It's sunny.",
|
||||
"citations": [
|
||||
{
|
||||
"type": "search_result_location",
|
||||
"cited_text": "The weather is sunny.",
|
||||
"source": "source_123",
|
||||
"title": "Document Title",
|
||||
"search_result_index": 1,
|
||||
"start_block_index": 0,
|
||||
"end_block_index": 2,
|
||||
},
|
||||
{"bar": "baz"},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "server_tool_use",
|
||||
"name": "web_search",
|
||||
"input": {"query": "web search query"},
|
||||
"id": "srvtoolu_abc123",
|
||||
},
|
||||
{
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_abc123",
|
||||
"content": [
|
||||
{
|
||||
"type": "web_search_result",
|
||||
"title": "Page Title 1",
|
||||
"url": "<page url 1>",
|
||||
"page_age": "January 1, 2025",
|
||||
"encrypted_content": "<encrypted content 1>",
|
||||
},
|
||||
{
|
||||
"type": "web_search_result",
|
||||
"title": "Page Title 2",
|
||||
"url": "<page url 2>",
|
||||
"page_age": "January 2, 2025",
|
||||
"encrypted_content": "<encrypted content 2>",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_def456",
|
||||
"name": "code_execution",
|
||||
"input": {"code": "import numpy as np..."},
|
||||
},
|
||||
{
|
||||
"type": "code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_def456",
|
||||
"content": {
|
||||
"type": "code_execution_result",
|
||||
"stdout": "Mean: 5.5\nStandard deviation...",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
{"type": "something_else", "foo": "bar"},
|
||||
],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
)
|
||||
expected_content: list[types.ContentBlock] = [
|
||||
{
|
||||
"type": "reasoning",
|
||||
"reasoning": "foo",
|
||||
"extras": {"signature": "foo_signature"},
|
||||
},
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "It's sunny.",
|
||||
"annotations": [
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "Document Title",
|
||||
"cited_text": "The weather is sunny.",
|
||||
"extras": {
|
||||
"source": "source_123",
|
||||
"search_result_index": 1,
|
||||
"start_block_index": 0,
|
||||
"end_block_index": 2,
|
||||
},
|
||||
},
|
||||
{"type": "non_standard_annotation", "value": {"bar": "baz"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "web_search",
|
||||
"id": "srvtoolu_abc123",
|
||||
"args": {"query": "web search query"},
|
||||
},
|
||||
{
|
||||
"type": "server_tool_result",
|
||||
"tool_call_id": "srvtoolu_abc123",
|
||||
"output": [
|
||||
{
|
||||
"type": "web_search_result",
|
||||
"title": "Page Title 1",
|
||||
"url": "<page url 1>",
|
||||
"page_age": "January 1, 2025",
|
||||
"encrypted_content": "<encrypted content 1>",
|
||||
},
|
||||
{
|
||||
"type": "web_search_result",
|
||||
"title": "Page Title 2",
|
||||
"url": "<page url 2>",
|
||||
"page_age": "January 2, 2025",
|
||||
"encrypted_content": "<encrypted content 2>",
|
||||
},
|
||||
],
|
||||
"status": "success",
|
||||
"extras": {"block_type": "web_search_tool_result"},
|
||||
},
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "code_interpreter",
|
||||
"id": "srvtoolu_def456",
|
||||
"args": {"code": "import numpy as np..."},
|
||||
},
|
||||
{
|
||||
"type": "server_tool_result",
|
||||
"tool_call_id": "srvtoolu_def456",
|
||||
"output": {
|
||||
"type": "code_execution_result",
|
||||
"return_code": 0,
|
||||
"stdout": "Mean: 5.5\nStandard deviation...",
|
||||
"stderr": "",
|
||||
},
|
||||
"status": "success",
|
||||
"extras": {"block_type": "code_execution_tool_result"},
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {"type": "something_else", "foo": "bar"},
|
||||
},
|
||||
]
|
||||
assert message.content_blocks == expected_content
|
||||
|
||||
# Check no mutation
|
||||
assert message.content != expected_content
|
||||
|
||||
message = AIMessage("Hello", response_metadata={"model_provider": "anthropic"})
|
||||
expected_content = [{"type": "text", "text": "Hello"}]
|
||||
assert message.content_blocks == expected_content
|
||||
assert message.content != expected_content # check no mutation
|
||||
|
||||
|
||||
def test_convert_to_v1_from_anthropic_chunk() -> None:
|
||||
chunks = [
|
||||
AIMessageChunk(
|
||||
content=[{"text": "Looking ", "type": "text", "index": 0}],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[{"text": "now.", "type": "text", "index": 0}],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "tool_use",
|
||||
"name": "get_weather",
|
||||
"input": {},
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": "",
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[{"type": "input_json_delta", "partial_json": "", "index": 1}],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": "",
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": '{"loca', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": '{"loca',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": 'tion": "San ', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": 'tion": "San ',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": 'Francisco"}', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": 'Francisco"}',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
),
|
||||
]
|
||||
expected_contents: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "Looking ", "index": 0},
|
||||
{"type": "text", "text": "now.", "index": 0},
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": "",
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
{"name": None, "args": "", "id": None, "index": 1, "type": "tool_call_chunk"},
|
||||
{
|
||||
"name": None,
|
||||
"args": '{"loca',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
{
|
||||
"name": None,
|
||||
"args": 'tion": "San ',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
{
|
||||
"name": None,
|
||||
"args": 'Francisco"}',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
]
|
||||
for chunk, expected in zip(chunks, expected_contents):
|
||||
assert chunk.content_blocks == [expected]
|
||||
|
||||
full: Optional[AIMessageChunk] = None
|
||||
for chunk in chunks:
|
||||
full = chunk if full is None else full + chunk
|
||||
assert isinstance(full, AIMessageChunk)
|
||||
|
||||
expected_content = [
|
||||
{"type": "text", "text": "Looking now.", "index": 0},
|
||||
{
|
||||
"type": "tool_use",
|
||||
"name": "get_weather",
|
||||
"partial_json": '{"location": "San Francisco"}',
|
||||
"input": {},
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content == expected_content
|
||||
|
||||
expected_content_blocks = [
|
||||
{"type": "text", "text": "Looking now.", "index": 0},
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": '{"location": "San Francisco"}',
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content_blocks == expected_content_blocks
|
||||
|
||||
# Test parse partial json
|
||||
full = AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"id": "srvtoolu_abc123",
|
||||
"input": {},
|
||||
"name": "web_fetch",
|
||||
"type": "server_tool_use",
|
||||
"index": 0,
|
||||
"partial_json": '{"url": "https://docs.langchain.com"}',
|
||||
},
|
||||
{
|
||||
"id": "mcptoolu_abc123",
|
||||
"input": {},
|
||||
"name": "ask_question",
|
||||
"server_name": "<my server name>",
|
||||
"type": "mcp_tool_use",
|
||||
"index": 1,
|
||||
"partial_json": '{"repoName": "<my repo>", "question": "<my query>"}',
|
||||
},
|
||||
],
|
||||
response_metadata={"model_provider": "anthropic"},
|
||||
chunk_position="last",
|
||||
)
|
||||
expected_content_blocks = [
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "web_fetch",
|
||||
"id": "srvtoolu_abc123",
|
||||
"args": {"url": "https://docs.langchain.com"},
|
||||
"index": 0,
|
||||
},
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "remote_mcp",
|
||||
"id": "mcptoolu_abc123",
|
||||
"args": {"repoName": "<my repo>", "question": "<my query>"},
|
||||
"extras": {"tool_name": "ask_question", "server_name": "<my server name>"},
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content_blocks == expected_content_blocks
|
||||
|
||||
|
||||
def test_convert_to_v1_from_anthropic_input() -> None:
|
||||
message = HumanMessage(
|
||||
[
|
||||
{"type": "text", "text": "foo"},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"media_type": "application/pdf",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "url",
|
||||
"url": "<document url>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "content",
|
||||
"content": [
|
||||
{"type": "text", "text": "The grass is green"},
|
||||
{"type": "text", "text": "The sky is blue"},
|
||||
],
|
||||
},
|
||||
"citations": {"enabled": True},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "text",
|
||||
"data": "<plain text data>",
|
||||
"media_type": "text/plain",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"media_type": "image/jpeg",
|
||||
"data": "<base64 image data>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "url",
|
||||
"url": "<image url>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "file",
|
||||
"file_id": "<image file id>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {"type": "file", "file_id": "<pdf file id>"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
expected: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "foo"},
|
||||
{
|
||||
"type": "file",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "application/pdf",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"url": "<document url>",
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "content",
|
||||
"content": [
|
||||
{"type": "text", "text": "The grass is green"},
|
||||
{"type": "text", "text": "The sky is blue"},
|
||||
],
|
||||
},
|
||||
"citations": {"enabled": True},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "text-plain",
|
||||
"text": "<plain text data>",
|
||||
"mime_type": "text/plain",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"base64": "<base64 image data>",
|
||||
"mime_type": "image/jpeg",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"url": "<image url>",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"id": "<image file id>",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"id": "<pdf file id>",
|
||||
},
|
||||
]
|
||||
|
||||
assert message.content_blocks == expected
|
||||
@@ -0,0 +1,407 @@
|
||||
from typing import Optional
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk, HumanMessage
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
|
||||
def test_convert_to_v1_from_bedrock() -> None:
|
||||
message = AIMessage(
|
||||
[
|
||||
{"type": "thinking", "thinking": "foo", "signature": "foo_signature"},
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"input": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "It's sunny.",
|
||||
"citations": [
|
||||
{
|
||||
"type": "search_result_location",
|
||||
"cited_text": "The weather is sunny.",
|
||||
"source": "source_123",
|
||||
"title": "Document Title",
|
||||
"search_result_index": 1,
|
||||
"start_block_index": 0,
|
||||
"end_block_index": 2,
|
||||
},
|
||||
{"bar": "baz"},
|
||||
],
|
||||
},
|
||||
{"type": "something_else", "foo": "bar"},
|
||||
],
|
||||
tool_calls=[
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_234",
|
||||
"name": "another_tool",
|
||||
"args": {"arg_1": "value_1"},
|
||||
},
|
||||
],
|
||||
response_metadata={
|
||||
"model_provider": "bedrock",
|
||||
"model_name": "us.anthropic.claude-sonnet-4-20250514-v1:0",
|
||||
},
|
||||
)
|
||||
expected_content: list[types.ContentBlock] = [
|
||||
{
|
||||
"type": "reasoning",
|
||||
"reasoning": "foo",
|
||||
"extras": {"signature": "foo_signature"},
|
||||
},
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "It's sunny.",
|
||||
"annotations": [
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "Document Title",
|
||||
"cited_text": "The weather is sunny.",
|
||||
"extras": {
|
||||
"source": "source_123",
|
||||
"search_result_index": 1,
|
||||
"start_block_index": 0,
|
||||
"end_block_index": 2,
|
||||
},
|
||||
},
|
||||
{"type": "non_standard_annotation", "value": {"bar": "baz"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {"type": "something_else", "foo": "bar"},
|
||||
},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_234",
|
||||
"name": "another_tool",
|
||||
"args": {"arg_1": "value_1"},
|
||||
},
|
||||
]
|
||||
assert message.content_blocks == expected_content
|
||||
|
||||
# Check no mutation
|
||||
assert message.content != expected_content
|
||||
|
||||
# Test with a non-Anthropic message
|
||||
message = AIMessage(
|
||||
[
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{"type": "something_else", "foo": "bar"},
|
||||
],
|
||||
tool_calls=[
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
)
|
||||
expected_content = [
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {"type": "something_else", "foo": "bar"},
|
||||
},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
]
|
||||
assert message.content_blocks == expected_content
|
||||
|
||||
|
||||
def test_convert_to_v1_from_bedrock_chunk() -> None:
|
||||
chunks = [
|
||||
AIMessageChunk(
|
||||
content=[{"text": "Looking ", "type": "text", "index": 0}],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[{"text": "now.", "type": "text", "index": 0}],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "tool_use",
|
||||
"name": "get_weather",
|
||||
"input": {},
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": "",
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[{"type": "input_json_delta", "partial_json": "", "index": 1}],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": "",
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": '{"loca', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": '{"loca',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": 'tion": "San ', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": 'tion": "San ',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": 'Francisco"}', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": 'Francisco"}',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock"},
|
||||
),
|
||||
]
|
||||
expected_contents: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "Looking ", "index": 0},
|
||||
{"type": "text", "text": "now.", "index": 0},
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": "",
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
{"name": None, "args": "", "id": None, "index": 1, "type": "tool_call_chunk"},
|
||||
{
|
||||
"name": None,
|
||||
"args": '{"loca',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
{
|
||||
"name": None,
|
||||
"args": 'tion": "San ',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
{
|
||||
"name": None,
|
||||
"args": 'Francisco"}',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
]
|
||||
for chunk, expected in zip(chunks, expected_contents):
|
||||
assert chunk.content_blocks == [expected]
|
||||
|
||||
full: Optional[AIMessageChunk] = None
|
||||
for chunk in chunks:
|
||||
full = chunk if full is None else full + chunk
|
||||
assert isinstance(full, AIMessageChunk)
|
||||
|
||||
expected_content = [
|
||||
{"type": "text", "text": "Looking now.", "index": 0},
|
||||
{
|
||||
"type": "tool_use",
|
||||
"name": "get_weather",
|
||||
"partial_json": '{"location": "San Francisco"}',
|
||||
"input": {},
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content == expected_content
|
||||
|
||||
expected_content_blocks = [
|
||||
{"type": "text", "text": "Looking now.", "index": 0},
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": '{"location": "San Francisco"}',
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content_blocks == expected_content_blocks
|
||||
|
||||
|
||||
def test_convert_to_v1_from_bedrock_input() -> None:
|
||||
message = HumanMessage(
|
||||
[
|
||||
{"type": "text", "text": "foo"},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"media_type": "application/pdf",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "url",
|
||||
"url": "<document url>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "content",
|
||||
"content": [
|
||||
{"type": "text", "text": "The grass is green"},
|
||||
{"type": "text", "text": "The sky is blue"},
|
||||
],
|
||||
},
|
||||
"citations": {"enabled": True},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "text",
|
||||
"data": "<plain text data>",
|
||||
"media_type": "text/plain",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"media_type": "image/jpeg",
|
||||
"data": "<base64 image data>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "url",
|
||||
"url": "<image url>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "file",
|
||||
"file_id": "<image file id>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "document",
|
||||
"source": {"type": "file", "file_id": "<pdf file id>"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
expected: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "foo"},
|
||||
{
|
||||
"type": "file",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "application/pdf",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"url": "<document url>",
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {
|
||||
"type": "document",
|
||||
"source": {
|
||||
"type": "content",
|
||||
"content": [
|
||||
{"type": "text", "text": "The grass is green"},
|
||||
{"type": "text", "text": "The sky is blue"},
|
||||
],
|
||||
},
|
||||
"citations": {"enabled": True},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "text-plain",
|
||||
"text": "<plain text data>",
|
||||
"mime_type": "text/plain",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"base64": "<base64 image data>",
|
||||
"mime_type": "image/jpeg",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"url": "<image url>",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"id": "<image file id>",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"id": "<pdf file id>",
|
||||
},
|
||||
]
|
||||
|
||||
assert message.content_blocks == expected
|
||||
@@ -0,0 +1,381 @@
|
||||
from typing import Optional
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk, HumanMessage
|
||||
from langchain_core.messages import content as types
|
||||
|
||||
|
||||
def test_convert_to_v1_from_bedrock_converse() -> None:
|
||||
message = AIMessage(
|
||||
[
|
||||
{
|
||||
"type": "reasoning_content",
|
||||
"reasoning_content": {"text": "foo", "signature": "foo_signature"},
|
||||
},
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"input": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "It's sunny.",
|
||||
"citations": [
|
||||
{
|
||||
"title": "Document Title",
|
||||
"source_content": [{"text": "The weather is sunny."}],
|
||||
"location": {
|
||||
"document_char": {
|
||||
"document_index": 0,
|
||||
"start": 58,
|
||||
"end": 96,
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
"title": "Document Title",
|
||||
"source_content": [{"text": "The weather is sunny."}],
|
||||
"location": {
|
||||
"document_page": {"document_index": 0, "start": 1, "end": 2}
|
||||
},
|
||||
},
|
||||
{
|
||||
"title": "Document Title",
|
||||
"source_content": [{"text": "The weather is sunny."}],
|
||||
"location": {
|
||||
"document_chunk": {
|
||||
"document_index": 0,
|
||||
"start": 1,
|
||||
"end": 2,
|
||||
}
|
||||
},
|
||||
},
|
||||
{"bar": "baz"},
|
||||
],
|
||||
},
|
||||
{"type": "something_else", "foo": "bar"},
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
)
|
||||
expected_content: list[types.ContentBlock] = [
|
||||
{
|
||||
"type": "reasoning",
|
||||
"reasoning": "foo",
|
||||
"extras": {"signature": "foo_signature"},
|
||||
},
|
||||
{"type": "text", "text": "Let's call a tool."},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "abc_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "It's sunny.",
|
||||
"annotations": [
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "Document Title",
|
||||
"cited_text": "The weather is sunny.",
|
||||
"extras": {
|
||||
"location": {
|
||||
"document_char": {
|
||||
"document_index": 0,
|
||||
"start": 58,
|
||||
"end": 96,
|
||||
}
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "Document Title",
|
||||
"cited_text": "The weather is sunny.",
|
||||
"extras": {
|
||||
"location": {
|
||||
"document_page": {"document_index": 0, "start": 1, "end": 2}
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "Document Title",
|
||||
"cited_text": "The weather is sunny.",
|
||||
"extras": {
|
||||
"location": {
|
||||
"document_chunk": {
|
||||
"document_index": 0,
|
||||
"start": 1,
|
||||
"end": 2,
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
{"type": "citation", "extras": {"bar": "baz"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {"type": "something_else", "foo": "bar"},
|
||||
},
|
||||
]
|
||||
assert message.content_blocks == expected_content
|
||||
|
||||
# Check no mutation
|
||||
assert message.content != expected_content
|
||||
|
||||
|
||||
def test_convert_to_v1_from_converse_chunk() -> None:
|
||||
chunks = [
|
||||
AIMessageChunk(
|
||||
content=[{"text": "Looking ", "type": "text", "index": 0}],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[{"text": "now.", "type": "text", "index": 0}],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "tool_use",
|
||||
"name": "get_weather",
|
||||
"input": {},
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": "",
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[{"type": "input_json_delta", "partial_json": "", "index": 1}],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": "",
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": '{"loca', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": '{"loca',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": 'tion": "San ', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": 'tion": "San ',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{"type": "input_json_delta", "partial_json": 'Francisco"}', "index": 1}
|
||||
],
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"name": None,
|
||||
"args": 'Francisco"}',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "bedrock_converse"},
|
||||
),
|
||||
]
|
||||
expected_contents: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "Looking ", "index": 0},
|
||||
{"type": "text", "text": "now.", "index": 0},
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": "",
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
{"name": None, "args": "", "id": None, "index": 1, "type": "tool_call_chunk"},
|
||||
{
|
||||
"name": None,
|
||||
"args": '{"loca',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
{
|
||||
"name": None,
|
||||
"args": 'tion": "San ',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
{
|
||||
"name": None,
|
||||
"args": 'Francisco"}',
|
||||
"id": None,
|
||||
"index": 1,
|
||||
"type": "tool_call_chunk",
|
||||
},
|
||||
]
|
||||
for chunk, expected in zip(chunks, expected_contents):
|
||||
assert chunk.content_blocks == [expected]
|
||||
|
||||
full: Optional[AIMessageChunk] = None
|
||||
for chunk in chunks:
|
||||
full = chunk if full is None else full + chunk
|
||||
assert isinstance(full, AIMessageChunk)
|
||||
|
||||
expected_content = [
|
||||
{"type": "text", "text": "Looking now.", "index": 0},
|
||||
{
|
||||
"type": "tool_use",
|
||||
"name": "get_weather",
|
||||
"partial_json": '{"location": "San Francisco"}',
|
||||
"input": {},
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content == expected_content
|
||||
|
||||
expected_content_blocks = [
|
||||
{"type": "text", "text": "Looking now.", "index": 0},
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "get_weather",
|
||||
"args": '{"location": "San Francisco"}',
|
||||
"id": "toolu_abc123",
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content_blocks == expected_content_blocks
|
||||
|
||||
|
||||
def test_convert_to_v1_from_converse_input() -> None:
|
||||
message = HumanMessage(
|
||||
[
|
||||
{"text": "foo"},
|
||||
{
|
||||
"document": {
|
||||
"format": "txt",
|
||||
"name": "doc_name_1",
|
||||
"source": {"text": "doc_text_1"},
|
||||
"context": "doc_context_1",
|
||||
"citations": {"enabled": True},
|
||||
},
|
||||
},
|
||||
{
|
||||
"document": {
|
||||
"format": "pdf",
|
||||
"name": "doc_name_2",
|
||||
"source": {"bytes": b"doc_text_2"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"document": {
|
||||
"format": "txt",
|
||||
"name": "doc_name_3",
|
||||
"source": {"content": [{"text": "doc_text"}, {"text": "_3"}]},
|
||||
"context": "doc_context_3",
|
||||
},
|
||||
},
|
||||
{
|
||||
"image": {
|
||||
"format": "jpeg",
|
||||
"source": {"bytes": b"image_bytes"},
|
||||
}
|
||||
},
|
||||
{
|
||||
"document": {
|
||||
"format": "pdf",
|
||||
"name": "doc_name_4",
|
||||
"source": {
|
||||
"s3Location": {"uri": "s3://bla", "bucketOwner": "owner"}
|
||||
},
|
||||
},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
expected: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "foo"},
|
||||
{
|
||||
"type": "text-plain",
|
||||
"mime_type": "text/plain",
|
||||
"text": "doc_text_1",
|
||||
"extras": {
|
||||
"name": "doc_name_1",
|
||||
"context": "doc_context_1",
|
||||
"citations": {"enabled": True},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"mime_type": "application/pdf",
|
||||
"base64": "ZG9jX3RleHRfMg==",
|
||||
"extras": {"name": "doc_name_2"},
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {
|
||||
"document": {
|
||||
"format": "txt",
|
||||
"name": "doc_name_3",
|
||||
"source": {"content": [{"text": "doc_text"}, {"text": "_3"}]},
|
||||
"context": "doc_context_3",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"base64": "aW1hZ2VfYnl0ZXM=",
|
||||
"mime_type": "image/jpeg",
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {
|
||||
"document": {
|
||||
"format": "pdf",
|
||||
"name": "doc_name_4",
|
||||
"source": {
|
||||
"s3Location": {"uri": "s3://bla", "bucketOwner": "owner"}
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
assert message.content_blocks == expected
|
||||
@@ -0,0 +1,113 @@
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.block_translators.langchain_v0 import (
|
||||
_convert_legacy_v0_content_block_to_v1,
|
||||
)
|
||||
from tests.unit_tests.language_models.chat_models.test_base import (
|
||||
_content_blocks_equal_ignore_id,
|
||||
)
|
||||
|
||||
|
||||
def test_convert_to_v1_from_openai_input() -> None:
|
||||
message = HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Hello"},
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://example.com/image.png",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"mime_type": "image/png",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"source_type": "url",
|
||||
"url": "<document url>",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"mime_type": "application/pdf",
|
||||
},
|
||||
{
|
||||
"type": "audio",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"mime_type": "audio/mpeg",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"source_type": "id",
|
||||
"id": "<file id>",
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
expected: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "Hello"},
|
||||
{
|
||||
"type": "image",
|
||||
"url": "https://example.com/image.png",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "image/png",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"url": "<document url>",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "application/pdf",
|
||||
},
|
||||
{
|
||||
"type": "audio",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "audio/mpeg",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file_id": "<file id>",
|
||||
},
|
||||
]
|
||||
|
||||
assert _content_blocks_equal_ignore_id(message.content_blocks, expected)
|
||||
|
||||
|
||||
def test_convert_with_extras_on_v0_block() -> None:
|
||||
"""Test that extras on old-style blocks are preserved in conversion.
|
||||
|
||||
Refer to `_extract_v0_extras` for details.
|
||||
"""
|
||||
block = {
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://example.com/image.png",
|
||||
# extras follow
|
||||
"alt_text": "An example image",
|
||||
"caption": "Example caption",
|
||||
"name": "example_image",
|
||||
"description": None,
|
||||
"attribution": None,
|
||||
}
|
||||
expected_output = {
|
||||
"type": "image",
|
||||
"url": "https://example.com/image.png",
|
||||
"extras": {
|
||||
"alt_text": "An example image",
|
||||
"caption": "Example caption",
|
||||
"name": "example_image",
|
||||
# "description": None, # These are filtered out
|
||||
# "attribution": None,
|
||||
},
|
||||
}
|
||||
|
||||
assert _convert_legacy_v0_content_block_to_v1(block) == expected_output
|
||||
@@ -0,0 +1,606 @@
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk, HumanMessage
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.block_translators.openai import (
|
||||
convert_to_openai_data_block,
|
||||
)
|
||||
from tests.unit_tests.language_models.chat_models.test_base import (
|
||||
_content_blocks_equal_ignore_id,
|
||||
)
|
||||
|
||||
|
||||
def test_convert_to_v1_from_responses() -> None:
|
||||
message = AIMessage(
|
||||
[
|
||||
{"type": "reasoning", "id": "abc123", "summary": []},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"summary": [
|
||||
{"type": "summary_text", "text": "foo bar"},
|
||||
{"type": "summary_text", "text": "baz"},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": "call_123",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "San Francisco"}',
|
||||
},
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": "call_234",
|
||||
"name": "get_weather_2",
|
||||
"arguments": '{"location": "New York"}',
|
||||
"id": "fc_123",
|
||||
},
|
||||
{"type": "text", "text": "Hello "},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "world",
|
||||
"annotations": [
|
||||
{"type": "url_citation", "url": "https://example.com"},
|
||||
{
|
||||
"type": "file_citation",
|
||||
"filename": "my doc",
|
||||
"index": 1,
|
||||
"file_id": "file_123",
|
||||
},
|
||||
{"bar": "baz"},
|
||||
],
|
||||
},
|
||||
{"type": "image_generation_call", "id": "ig_123", "result": "..."},
|
||||
{
|
||||
"type": "file_search_call",
|
||||
"id": "fs_123",
|
||||
"queries": ["query for file search"],
|
||||
"results": [{"file_id": "file-123"}],
|
||||
"status": "completed",
|
||||
},
|
||||
{"type": "something_else", "foo": "bar"},
|
||||
],
|
||||
tool_calls=[
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "call_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "call_234",
|
||||
"name": "get_weather_2",
|
||||
"args": {"location": "New York"},
|
||||
},
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
)
|
||||
expected_content: list[types.ContentBlock] = [
|
||||
{"type": "reasoning", "id": "abc123"},
|
||||
{"type": "reasoning", "id": "abc234", "reasoning": "foo bar"},
|
||||
{"type": "reasoning", "id": "abc234", "reasoning": "baz"},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "call_123",
|
||||
"name": "get_weather",
|
||||
"args": {"location": "San Francisco"},
|
||||
},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "call_234",
|
||||
"name": "get_weather_2",
|
||||
"args": {"location": "New York"},
|
||||
"extras": {"item_id": "fc_123"},
|
||||
},
|
||||
{"type": "text", "text": "Hello "},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "world",
|
||||
"annotations": [
|
||||
{"type": "citation", "url": "https://example.com"},
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "my doc",
|
||||
"extras": {"file_id": "file_123", "index": 1},
|
||||
},
|
||||
{"type": "non_standard_annotation", "value": {"bar": "baz"}},
|
||||
],
|
||||
},
|
||||
{"type": "image", "base64": "...", "id": "ig_123"},
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "file_search",
|
||||
"id": "fs_123",
|
||||
"args": {"queries": ["query for file search"]},
|
||||
},
|
||||
{
|
||||
"type": "server_tool_result",
|
||||
"tool_call_id": "fs_123",
|
||||
"output": [{"file_id": "file-123"}],
|
||||
"status": "success",
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {"type": "something_else", "foo": "bar"},
|
||||
},
|
||||
]
|
||||
assert message.content_blocks == expected_content
|
||||
|
||||
# Check no mutation
|
||||
assert message.content != expected_content
|
||||
|
||||
|
||||
def test_convert_to_v1_from_responses_chunk() -> None:
|
||||
chunks = [
|
||||
AIMessageChunk(
|
||||
content=[{"type": "reasoning", "id": "abc123", "summary": [], "index": 0}],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"summary": [
|
||||
{"type": "summary_text", "text": "foo ", "index": 0},
|
||||
],
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"summary": [
|
||||
{"type": "summary_text", "text": "bar", "index": 0},
|
||||
],
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"summary": [
|
||||
{"type": "summary_text", "text": "baz", "index": 1},
|
||||
],
|
||||
"index": 1,
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
]
|
||||
expected_chunks = [
|
||||
AIMessageChunk(
|
||||
content=[{"type": "reasoning", "id": "abc123", "index": "lc_rs_305f30"}],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"reasoning": "foo ",
|
||||
"index": "lc_rs_315f30",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"reasoning": "bar",
|
||||
"index": "lc_rs_315f30",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"reasoning": "baz",
|
||||
"index": "lc_rs_315f31",
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
),
|
||||
]
|
||||
for chunk, expected in zip(chunks, expected_chunks):
|
||||
assert chunk.content_blocks == expected.content_blocks
|
||||
|
||||
full: Optional[AIMessageChunk] = None
|
||||
for chunk in chunks:
|
||||
full = chunk if full is None else full + chunk
|
||||
assert isinstance(full, AIMessageChunk)
|
||||
|
||||
expected_content = [
|
||||
{"type": "reasoning", "id": "abc123", "summary": [], "index": 0},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"summary": [
|
||||
{"type": "summary_text", "text": "foo bar", "index": 0},
|
||||
{"type": "summary_text", "text": "baz", "index": 1},
|
||||
],
|
||||
"index": 1,
|
||||
},
|
||||
]
|
||||
assert full.content == expected_content
|
||||
|
||||
expected_content_blocks = [
|
||||
{"type": "reasoning", "id": "abc123", "index": "lc_rs_305f30"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"reasoning": "foo bar",
|
||||
"index": "lc_rs_315f30",
|
||||
},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "abc234",
|
||||
"reasoning": "baz",
|
||||
"index": "lc_rs_315f31",
|
||||
},
|
||||
]
|
||||
assert full.content_blocks == expected_content_blocks
|
||||
|
||||
|
||||
def test_convert_to_v1_from_openai_input() -> None:
|
||||
message = HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Hello"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/image.png"},
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQSkZJRg..."},
|
||||
},
|
||||
{
|
||||
"type": "input_audio",
|
||||
"input_audio": {
|
||||
"format": "wav",
|
||||
"data": "<base64 string>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {
|
||||
"filename": "draconomicon.pdf",
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {"file_id": "<file id>"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
expected: list[types.ContentBlock] = [
|
||||
{"type": "text", "text": "Hello"},
|
||||
{
|
||||
"type": "image",
|
||||
"url": "https://example.com/image.png",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"base64": "/9j/4AAQSkZJRg...",
|
||||
"mime_type": "image/jpeg",
|
||||
},
|
||||
{
|
||||
"type": "audio",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "audio/wav",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
"extras": {"filename": "draconomicon.pdf"},
|
||||
},
|
||||
{"type": "file", "file_id": "<file id>"},
|
||||
]
|
||||
|
||||
assert _content_blocks_equal_ignore_id(message.content_blocks, expected)
|
||||
|
||||
|
||||
def test_compat_responses_v03() -> None:
|
||||
# Check compatibility with v0.3 legacy message format
|
||||
message_v03 = AIMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "Hello, world!", "annotations": [{"type": "foo"}]}
|
||||
],
|
||||
additional_kwargs={
|
||||
"reasoning": {
|
||||
"type": "reasoning",
|
||||
"id": "rs_123",
|
||||
"summary": [
|
||||
{"type": "summary_text", "text": "summary 1"},
|
||||
{"type": "summary_text", "text": "summary 2"},
|
||||
],
|
||||
},
|
||||
"tool_outputs": [
|
||||
{
|
||||
"type": "web_search_call",
|
||||
"id": "websearch_123",
|
||||
"status": "completed",
|
||||
}
|
||||
],
|
||||
"refusal": "I cannot assist with that.",
|
||||
"__openai_function_call_ids__": {"call_abc": "fc_abc"},
|
||||
},
|
||||
tool_calls=[
|
||||
{"type": "tool_call", "name": "my_tool", "args": {"x": 3}, "id": "call_abc"}
|
||||
],
|
||||
response_metadata={"id": "resp_123", "model_provider": "openai"},
|
||||
id="msg_123",
|
||||
)
|
||||
|
||||
expected_content: list[types.ContentBlock] = [
|
||||
{"type": "reasoning", "id": "rs_123", "reasoning": "summary 1"},
|
||||
{"type": "reasoning", "id": "rs_123", "reasoning": "summary 2"},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Hello, world!",
|
||||
"annotations": [
|
||||
{"type": "non_standard_annotation", "value": {"type": "foo"}}
|
||||
],
|
||||
"id": "msg_123",
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {"type": "refusal", "refusal": "I cannot assist with that."},
|
||||
},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"name": "my_tool",
|
||||
"args": {"x": 3},
|
||||
"id": "call_abc",
|
||||
"extras": {"item_id": "fc_abc"},
|
||||
},
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "web_search",
|
||||
"args": {},
|
||||
"id": "websearch_123",
|
||||
},
|
||||
{
|
||||
"type": "server_tool_result",
|
||||
"tool_call_id": "websearch_123",
|
||||
"status": "success",
|
||||
},
|
||||
]
|
||||
assert message_v03.content_blocks == expected_content
|
||||
|
||||
# Test chunks
|
||||
## Tool calls
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=[],
|
||||
additional_kwargs={"__openai_function_call_ids__": {"call_abc": "fc_abc"}},
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "my_tool",
|
||||
"args": "",
|
||||
"id": "call_abc",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
response_metadata={"model_provider": "openai"},
|
||||
)
|
||||
expected_content = [
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "my_tool",
|
||||
"args": "",
|
||||
"id": "call_abc",
|
||||
"index": 0,
|
||||
"extras": {"item_id": "fc_abc"},
|
||||
}
|
||||
]
|
||||
assert chunk_1.content_blocks == expected_content
|
||||
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=[],
|
||||
additional_kwargs={"__openai_function_call_ids__": {}},
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": None,
|
||||
"args": "{",
|
||||
"id": None,
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
)
|
||||
expected_content = [
|
||||
{"type": "tool_call_chunk", "name": None, "args": "{", "id": None, "index": 0}
|
||||
]
|
||||
|
||||
chunk = chunk_1 + chunk_2
|
||||
expected_content = [
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "my_tool",
|
||||
"args": "{",
|
||||
"id": "call_abc",
|
||||
"index": 0,
|
||||
"extras": {"item_id": "fc_abc"},
|
||||
}
|
||||
]
|
||||
assert chunk.content_blocks == expected_content
|
||||
|
||||
## Reasoning
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=[],
|
||||
additional_kwargs={
|
||||
"reasoning": {"id": "rs_abc", "summary": [], "type": "reasoning"}
|
||||
},
|
||||
response_metadata={"model_provider": "openai"},
|
||||
)
|
||||
expected_content = [{"type": "reasoning", "id": "rs_abc"}]
|
||||
assert chunk_1.content_blocks == expected_content
|
||||
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=[],
|
||||
additional_kwargs={
|
||||
"reasoning": {
|
||||
"summary": [
|
||||
{"index": 0, "type": "summary_text", "text": "reasoning text"}
|
||||
]
|
||||
}
|
||||
},
|
||||
response_metadata={"model_provider": "openai"},
|
||||
)
|
||||
expected_content = [{"type": "reasoning", "reasoning": "reasoning text"}]
|
||||
assert chunk_2.content_blocks == expected_content
|
||||
|
||||
chunk = chunk_1 + chunk_2
|
||||
expected_content = [
|
||||
{"type": "reasoning", "reasoning": "reasoning text", "id": "rs_abc"}
|
||||
]
|
||||
assert chunk.content_blocks == expected_content
|
||||
|
||||
|
||||
def test_convert_to_openai_data_block() -> None:
|
||||
# Chat completions
|
||||
## Image / url
|
||||
block = {
|
||||
"type": "image",
|
||||
"url": "https://example.com/test.png",
|
||||
}
|
||||
expected = {
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://example.com/test.png"},
|
||||
}
|
||||
result = convert_to_openai_data_block(block)
|
||||
assert result == expected
|
||||
|
||||
## Image / base64
|
||||
block = {
|
||||
"type": "image",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "image/png",
|
||||
}
|
||||
expected = {
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,<base64 string>"},
|
||||
}
|
||||
result = convert_to_openai_data_block(block)
|
||||
assert result == expected
|
||||
|
||||
## File / url
|
||||
block = {
|
||||
"type": "file",
|
||||
"url": "https://example.com/test.pdf",
|
||||
}
|
||||
with pytest.raises(ValueError, match="does not support"):
|
||||
result = convert_to_openai_data_block(block)
|
||||
|
||||
## File / base64
|
||||
block = {
|
||||
"type": "file",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
"filename": "test.pdf",
|
||||
}
|
||||
expected = {
|
||||
"type": "file",
|
||||
"file": {
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
"filename": "test.pdf",
|
||||
},
|
||||
}
|
||||
result = convert_to_openai_data_block(block)
|
||||
assert result == expected
|
||||
|
||||
## File / file ID
|
||||
block = {
|
||||
"type": "file",
|
||||
"file_id": "file-abc123",
|
||||
}
|
||||
expected = {"type": "file", "file": {"file_id": "file-abc123"}}
|
||||
result = convert_to_openai_data_block(block)
|
||||
assert result == expected
|
||||
|
||||
## Audio / base64
|
||||
block = {
|
||||
"type": "audio",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "audio/wav",
|
||||
}
|
||||
expected = {
|
||||
"type": "input_audio",
|
||||
"input_audio": {"data": "<base64 string>", "format": "wav"},
|
||||
}
|
||||
result = convert_to_openai_data_block(block)
|
||||
assert result == expected
|
||||
|
||||
# Responses
|
||||
## Image / url
|
||||
block = {
|
||||
"type": "image",
|
||||
"url": "https://example.com/test.png",
|
||||
}
|
||||
expected = {"type": "input_image", "image_url": "https://example.com/test.png"}
|
||||
result = convert_to_openai_data_block(block, api="responses")
|
||||
assert result == expected
|
||||
|
||||
## Image / base64
|
||||
block = {
|
||||
"type": "image",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "image/png",
|
||||
}
|
||||
expected = {
|
||||
"type": "input_image",
|
||||
"image_url": "data:image/png;base64,<base64 string>",
|
||||
}
|
||||
result = convert_to_openai_data_block(block, api="responses")
|
||||
assert result == expected
|
||||
|
||||
## File / url
|
||||
block = {
|
||||
"type": "file",
|
||||
"url": "https://example.com/test.pdf",
|
||||
}
|
||||
expected = {"type": "input_file", "file_url": "https://example.com/test.pdf"}
|
||||
|
||||
## File / base64
|
||||
block = {
|
||||
"type": "file",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
"filename": "test.pdf",
|
||||
}
|
||||
expected = {
|
||||
"type": "input_file",
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
"filename": "test.pdf",
|
||||
}
|
||||
result = convert_to_openai_data_block(block, api="responses")
|
||||
assert result == expected
|
||||
|
||||
## File / file ID
|
||||
block = {
|
||||
"type": "file",
|
||||
"file_id": "file-abc123",
|
||||
}
|
||||
expected = {"type": "input_file", "file_id": "file-abc123"}
|
||||
result = convert_to_openai_data_block(block, api="responses")
|
||||
assert result == expected
|
||||
@@ -0,0 +1,29 @@
|
||||
import pkgutil
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from langchain_core.messages.block_translators import PROVIDER_TRANSLATORS
|
||||
|
||||
|
||||
def test_all_providers_registered() -> None:
|
||||
"""Test that all block translators implemented in langchain-core are registered.
|
||||
|
||||
If this test fails, it is likely that a block translator is implemented but not
|
||||
registered on import. Check that the provider is included in
|
||||
``langchain_core.messages.block_translators.__init__._register_translators``.
|
||||
"""
|
||||
package_path = (
|
||||
Path(__file__).parents[4] / "langchain_core" / "messages" / "block_translators"
|
||||
)
|
||||
|
||||
for module_info in pkgutil.iter_modules([str(package_path)]):
|
||||
module_name = module_info.name
|
||||
|
||||
# Skip the __init__ module, any private modules, and ``langchain_v0``, which is
|
||||
# only used to parse v0 multimodal inputs.
|
||||
if module_name.startswith("_") or module_name == "langchain_v0":
|
||||
continue
|
||||
|
||||
if module_name not in PROVIDER_TRANSLATORS:
|
||||
pytest.fail(f"Block translator not registered: {module_name}")
|
||||
@@ -1,5 +1,10 @@
|
||||
from typing import Union, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from langchain_core.load import dumpd, load
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
from langchain_core.messages import content as types
|
||||
from langchain_core.messages.ai import (
|
||||
InputTokenDetails,
|
||||
OutputTokenDetails,
|
||||
@@ -196,3 +201,323 @@ def test_add_ai_message_chunks_usage() -> None:
|
||||
output_token_details=OutputTokenDetails(audio=1, reasoning=2),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_init_tool_calls() -> None:
|
||||
# Test we add "type" key on init
|
||||
msg = AIMessage("", tool_calls=[{"name": "foo", "args": {"a": "b"}, "id": "abc"}])
|
||||
assert len(msg.tool_calls) == 1
|
||||
assert msg.tool_calls[0]["type"] == "tool_call"
|
||||
|
||||
# Test we can assign without adding type key
|
||||
msg.tool_calls = [{"name": "bar", "args": {"c": "d"}, "id": "def"}]
|
||||
|
||||
|
||||
def test_content_blocks() -> None:
|
||||
message = AIMessage(
|
||||
"",
|
||||
tool_calls=[
|
||||
{"type": "tool_call", "name": "foo", "args": {"a": "b"}, "id": "abc_123"}
|
||||
],
|
||||
)
|
||||
assert len(message.content_blocks) == 1
|
||||
assert message.content_blocks[0]["type"] == "tool_call"
|
||||
assert message.content_blocks == [
|
||||
{"type": "tool_call", "id": "abc_123", "name": "foo", "args": {"a": "b"}}
|
||||
]
|
||||
assert message.content == ""
|
||||
|
||||
message = AIMessage(
|
||||
"foo",
|
||||
tool_calls=[
|
||||
{"type": "tool_call", "name": "foo", "args": {"a": "b"}, "id": "abc_123"}
|
||||
],
|
||||
)
|
||||
assert len(message.content_blocks) == 2
|
||||
assert message.content_blocks[0]["type"] == "text"
|
||||
assert message.content_blocks[1]["type"] == "tool_call"
|
||||
assert message.content_blocks == [
|
||||
{"type": "text", "text": "foo"},
|
||||
{"type": "tool_call", "id": "abc_123", "name": "foo", "args": {"a": "b"}},
|
||||
]
|
||||
assert message.content == "foo"
|
||||
|
||||
# With standard blocks
|
||||
standard_content: list[types.ContentBlock] = [
|
||||
{"type": "reasoning", "reasoning": "foo"},
|
||||
{"type": "text", "text": "bar"},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "baz",
|
||||
"annotations": [{"type": "citation", "url": "http://example.com"}],
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"url": "http://example.com/image.png",
|
||||
"extras": {"foo": "bar"},
|
||||
},
|
||||
{
|
||||
"type": "non_standard",
|
||||
"value": {"custom_key": "custom_value", "another_key": 123},
|
||||
},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"name": "foo",
|
||||
"args": {"a": "b"},
|
||||
"id": "abc_123",
|
||||
},
|
||||
]
|
||||
missing_tool_call: types.ToolCall = {
|
||||
"type": "tool_call",
|
||||
"name": "bar",
|
||||
"args": {"c": "d"},
|
||||
"id": "abc_234",
|
||||
}
|
||||
message = AIMessage(
|
||||
content_blocks=standard_content,
|
||||
tool_calls=[
|
||||
{"type": "tool_call", "name": "foo", "args": {"a": "b"}, "id": "abc_123"},
|
||||
missing_tool_call,
|
||||
],
|
||||
)
|
||||
assert message.content_blocks == [*standard_content, missing_tool_call]
|
||||
|
||||
# Check we auto-populate tool_calls
|
||||
standard_content = [
|
||||
{"type": "text", "text": "foo"},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"name": "foo",
|
||||
"args": {"a": "b"},
|
||||
"id": "abc_123",
|
||||
},
|
||||
missing_tool_call,
|
||||
]
|
||||
message = AIMessage(content_blocks=standard_content)
|
||||
assert message.tool_calls == [
|
||||
{"type": "tool_call", "name": "foo", "args": {"a": "b"}, "id": "abc_123"},
|
||||
missing_tool_call,
|
||||
]
|
||||
|
||||
# Chunks
|
||||
message = AIMessageChunk(
|
||||
content="",
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "foo",
|
||||
"args": "",
|
||||
"id": "abc_123",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
)
|
||||
assert len(message.content_blocks) == 1
|
||||
assert message.content_blocks[0]["type"] == "tool_call_chunk"
|
||||
assert message.content_blocks == [
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "foo",
|
||||
"args": "",
|
||||
"id": "abc_123",
|
||||
"index": 0,
|
||||
}
|
||||
]
|
||||
assert message.content == ""
|
||||
|
||||
# Test we parse tool call chunks into tool calls for v1 content
|
||||
chunk_1 = AIMessageChunk(
|
||||
content="",
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "foo",
|
||||
"args": '{"foo": "b',
|
||||
"id": "abc_123",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
chunk_2 = AIMessageChunk(
|
||||
content="",
|
||||
tool_call_chunks=[
|
||||
{
|
||||
"type": "tool_call_chunk",
|
||||
"name": "",
|
||||
"args": 'ar"}',
|
||||
"id": "abc_123",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
)
|
||||
chunk_3 = AIMessageChunk(content="", chunk_position="last")
|
||||
chunk = chunk_1 + chunk_2 + chunk_3
|
||||
assert chunk.content == ""
|
||||
assert chunk.content_blocks == chunk.tool_calls
|
||||
|
||||
# test v1 content
|
||||
chunk_1.content = cast("Union[str, list[Union[str, dict]]]", chunk_1.content_blocks)
|
||||
chunk_1.response_metadata["output_version"] = "v1"
|
||||
chunk_2.content = cast("Union[str, list[Union[str, dict]]]", chunk_2.content_blocks)
|
||||
|
||||
chunk = chunk_1 + chunk_2 + chunk_3
|
||||
assert chunk.content == [
|
||||
{
|
||||
"type": "tool_call",
|
||||
"name": "foo",
|
||||
"args": {"foo": "bar"},
|
||||
"id": "abc_123",
|
||||
}
|
||||
]
|
||||
|
||||
# Non-standard
|
||||
standard_content_1: list[types.ContentBlock] = [
|
||||
{"type": "non_standard", "index": 0, "value": {"foo": "bar "}}
|
||||
]
|
||||
standard_content_2: list[types.ContentBlock] = [
|
||||
{"type": "non_standard", "index": 0, "value": {"foo": "baz"}}
|
||||
]
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", standard_content_1)
|
||||
)
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", standard_content_2)
|
||||
)
|
||||
merged_chunk = chunk_1 + chunk_2
|
||||
assert merged_chunk.content == [
|
||||
{"type": "non_standard", "index": 0, "value": {"foo": "bar baz"}},
|
||||
]
|
||||
|
||||
# Test server_tool_call_chunks
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "server_tool_call_chunk",
|
||||
"index": 0,
|
||||
"name": "foo",
|
||||
}
|
||||
]
|
||||
)
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=[{"type": "server_tool_call_chunk", "index": 0, "args": '{"a'}]
|
||||
)
|
||||
chunk_3 = AIMessageChunk(
|
||||
content=[{"type": "server_tool_call_chunk", "index": 0, "args": '": 1}'}]
|
||||
)
|
||||
merged_chunk = chunk_1 + chunk_2 + chunk_3
|
||||
assert merged_chunk.content == [
|
||||
{
|
||||
"type": "server_tool_call_chunk",
|
||||
"name": "foo",
|
||||
"index": 0,
|
||||
"args": '{"a": 1}',
|
||||
}
|
||||
]
|
||||
|
||||
full_chunk = merged_chunk + AIMessageChunk(
|
||||
content=[], chunk_position="last", response_metadata={"output_version": "v1"}
|
||||
)
|
||||
assert full_chunk.content == [
|
||||
{"type": "server_tool_call", "name": "foo", "index": 0, "args": {"a": 1}}
|
||||
]
|
||||
|
||||
# Test non-standard + non-standard
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "non_standard",
|
||||
"index": 0,
|
||||
"value": {"type": "non_standard_tool", "foo": "bar"},
|
||||
}
|
||||
]
|
||||
)
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "non_standard",
|
||||
"index": 0,
|
||||
"value": {"type": "input_json_delta", "partial_json": "a"},
|
||||
}
|
||||
]
|
||||
)
|
||||
chunk_3 = AIMessageChunk(
|
||||
content=[
|
||||
{
|
||||
"type": "non_standard",
|
||||
"index": 0,
|
||||
"value": {"type": "input_json_delta", "partial_json": "b"},
|
||||
}
|
||||
]
|
||||
)
|
||||
merged_chunk = chunk_1 + chunk_2 + chunk_3
|
||||
assert merged_chunk.content == [
|
||||
{
|
||||
"type": "non_standard",
|
||||
"index": 0,
|
||||
"value": {"type": "non_standard_tool", "foo": "bar", "partial_json": "ab"},
|
||||
}
|
||||
]
|
||||
|
||||
# Test standard + non-standard with same index
|
||||
standard_content_1 = [
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "web_search",
|
||||
"id": "ws_123",
|
||||
"args": {"query": "web query"},
|
||||
"index": 0,
|
||||
}
|
||||
]
|
||||
standard_content_2 = [{"type": "non_standard", "value": {"foo": "bar"}, "index": 0}]
|
||||
chunk_1 = AIMessageChunk(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", standard_content_1)
|
||||
)
|
||||
chunk_2 = AIMessageChunk(
|
||||
content=cast("Union[str, list[Union[str, dict]]]", standard_content_2)
|
||||
)
|
||||
merged_chunk = chunk_1 + chunk_2
|
||||
assert merged_chunk.content == [
|
||||
{
|
||||
"type": "server_tool_call",
|
||||
"name": "web_search",
|
||||
"id": "ws_123",
|
||||
"args": {"query": "web query"},
|
||||
"index": 0,
|
||||
"extras": {"foo": "bar"},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_provider_warns() -> None:
|
||||
# Test that major providers warn if content block standardization is not yet
|
||||
# implemented.
|
||||
# This test should be removed when all major providers support content block
|
||||
# standardization.
|
||||
message = AIMessage("Hello.", response_metadata={"model_provider": "groq"})
|
||||
with pytest.warns(match="not yet fully supported for Groq"):
|
||||
content_blocks = message.content_blocks
|
||||
|
||||
assert content_blocks == [{"type": "text", "text": "Hello."}]
|
||||
|
||||
|
||||
def test_content_blocks_reasoning_extraction() -> None:
|
||||
"""Test best-effort reasoning extraction from `additional_kwargs`."""
|
||||
message = AIMessage(
|
||||
content="The answer is 42.",
|
||||
additional_kwargs={"reasoning_content": "Let me think about this problem..."},
|
||||
)
|
||||
content_blocks = message.content_blocks
|
||||
assert len(content_blocks) == 2
|
||||
assert content_blocks[0]["type"] == "reasoning"
|
||||
assert content_blocks[0].get("reasoning") == "Let me think about this problem..."
|
||||
assert content_blocks[1]["type"] == "text"
|
||||
assert content_blocks[1]["text"] == "The answer is 42."
|
||||
|
||||
# Test no reasoning extraction when no reasoning content
|
||||
message = AIMessage(
|
||||
content="The answer is 42.", additional_kwargs={"other_field": "some value"}
|
||||
)
|
||||
content_blocks = message.content_blocks
|
||||
assert len(content_blocks) == 1
|
||||
assert content_blocks[0]["type"] == "text"
|
||||
@@ -5,24 +5,43 @@ EXPECTED_ALL = [
|
||||
"_message_from_dict",
|
||||
"AIMessage",
|
||||
"AIMessageChunk",
|
||||
"Annotation",
|
||||
"AnyMessage",
|
||||
"AudioContentBlock",
|
||||
"BaseMessage",
|
||||
"BaseMessageChunk",
|
||||
"ContentBlock",
|
||||
"ChatMessage",
|
||||
"ChatMessageChunk",
|
||||
"Citation",
|
||||
"DataContentBlock",
|
||||
"FileContentBlock",
|
||||
"FunctionMessage",
|
||||
"FunctionMessageChunk",
|
||||
"HumanMessage",
|
||||
"HumanMessageChunk",
|
||||
"ImageContentBlock",
|
||||
"InvalidToolCall",
|
||||
"LC_AUTO_PREFIX",
|
||||
"LC_ID_PREFIX",
|
||||
"NonStandardAnnotation",
|
||||
"NonStandardContentBlock",
|
||||
"PlainTextContentBlock",
|
||||
"ServerToolCall",
|
||||
"ServerToolCallChunk",
|
||||
"ServerToolResult",
|
||||
"SystemMessage",
|
||||
"SystemMessageChunk",
|
||||
"TextContentBlock",
|
||||
"ToolCall",
|
||||
"ToolCallChunk",
|
||||
"ToolMessage",
|
||||
"ToolMessageChunk",
|
||||
"VideoContentBlock",
|
||||
"ReasoningContentBlock",
|
||||
"RemoveMessage",
|
||||
"convert_to_messages",
|
||||
"ensure_id",
|
||||
"get_buffer_string",
|
||||
"is_data_content_block",
|
||||
"merge_content",
|
||||
|
||||
@@ -1215,13 +1215,14 @@ def test_convert_to_openai_messages_developer() -> None:
|
||||
|
||||
|
||||
def test_convert_to_openai_messages_multimodal() -> None:
|
||||
"""v0 and v1 content to OpenAI messages conversion."""
|
||||
messages = [
|
||||
HumanMessage(
|
||||
content=[
|
||||
# Prior v0 blocks
|
||||
{"type": "text", "text": "Text message"},
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://example.com/test.png",
|
||||
},
|
||||
{
|
||||
@@ -1238,6 +1239,7 @@ def test_convert_to_openai_messages_multimodal() -> None:
|
||||
"filename": "test.pdf",
|
||||
},
|
||||
{
|
||||
# OpenAI Chat Completions file format
|
||||
"type": "file",
|
||||
"file": {
|
||||
"filename": "draconomicon.pdf",
|
||||
@@ -1262,22 +1264,47 @@ def test_convert_to_openai_messages_multimodal() -> None:
|
||||
"format": "wav",
|
||||
},
|
||||
},
|
||||
# v1 Additions
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "url", # backward compatibility v0 block field
|
||||
"url": "https://example.com/test.png",
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "image/png",
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
"filename": "test.pdf", # backward compatibility v0 block field
|
||||
},
|
||||
{
|
||||
"type": "file",
|
||||
"file_id": "file-abc123",
|
||||
},
|
||||
{
|
||||
"type": "audio",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "audio/wav",
|
||||
},
|
||||
]
|
||||
)
|
||||
]
|
||||
result = convert_to_openai_messages(messages, text_format="block")
|
||||
assert len(result) == 1
|
||||
message = result[0]
|
||||
assert len(message["content"]) == 8
|
||||
assert len(message["content"]) == 13
|
||||
|
||||
# Test adding filename
|
||||
# Test auto-adding filename
|
||||
messages = [
|
||||
HumanMessage(
|
||||
content=[
|
||||
{
|
||||
"type": "file",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 string>",
|
||||
"base64": "<base64 string>",
|
||||
"mime_type": "application/pdf",
|
||||
},
|
||||
]
|
||||
@@ -1290,6 +1317,7 @@ def test_convert_to_openai_messages_multimodal() -> None:
|
||||
assert len(message["content"]) == 1
|
||||
block = message["content"][0]
|
||||
assert block == {
|
||||
# OpenAI Chat Completions file format
|
||||
"type": "file",
|
||||
"file": {
|
||||
"file_data": "data:application/pdf;base64,<base64 string>",
|
||||
|
||||
@@ -39,11 +39,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -122,6 +117,19 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
'type': 'string',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -143,11 +151,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -560,11 +563,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -635,11 +633,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -747,6 +740,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -758,6 +755,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -776,9 +784,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -1457,11 +1466,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -1540,6 +1544,19 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
'type': 'string',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -1561,11 +1578,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -1978,11 +1990,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2053,11 +2060,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2165,6 +2167,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2176,6 +2182,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2194,9 +2211,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
|
||||
@@ -260,13 +260,9 @@ async def test_chat_prompt_template_from_messages_using_role_strings() -> None:
|
||||
SystemMessage(
|
||||
content="You are a helpful AI bot. Your name is Bob.", additional_kwargs={}
|
||||
),
|
||||
HumanMessage(
|
||||
content="Hello, how are you doing?", additional_kwargs={}, example=False
|
||||
),
|
||||
AIMessage(
|
||||
content="I'm doing well, thanks!", additional_kwargs={}, example=False
|
||||
),
|
||||
HumanMessage(content="What is your name?", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="Hello, how are you doing?", additional_kwargs={}),
|
||||
AIMessage(content="I'm doing well, thanks!", additional_kwargs={}),
|
||||
HumanMessage(content="What is your name?", additional_kwargs={}),
|
||||
]
|
||||
|
||||
messages = template.format_messages(name="Bob", user_input="What is your name?")
|
||||
@@ -296,13 +292,9 @@ def test_chat_prompt_template_from_messages_mustache() -> None:
|
||||
SystemMessage(
|
||||
content="You are a helpful AI bot. Your name is Bob.", additional_kwargs={}
|
||||
),
|
||||
HumanMessage(
|
||||
content="Hello, how are you doing?", additional_kwargs={}, example=False
|
||||
),
|
||||
AIMessage(
|
||||
content="I'm doing well, thanks!", additional_kwargs={}, example=False
|
||||
),
|
||||
HumanMessage(content="What is your name?", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="Hello, how are you doing?", additional_kwargs={}),
|
||||
AIMessage(content="I'm doing well, thanks!", additional_kwargs={}),
|
||||
HumanMessage(content="What is your name?", additional_kwargs={}),
|
||||
]
|
||||
|
||||
|
||||
@@ -324,13 +316,9 @@ def test_chat_prompt_template_from_messages_jinja2() -> None:
|
||||
SystemMessage(
|
||||
content="You are a helpful AI bot. Your name is Bob.", additional_kwargs={}
|
||||
),
|
||||
HumanMessage(
|
||||
content="Hello, how are you doing?", additional_kwargs={}, example=False
|
||||
),
|
||||
AIMessage(
|
||||
content="I'm doing well, thanks!", additional_kwargs={}, example=False
|
||||
),
|
||||
HumanMessage(content="What is your name?", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="Hello, how are you doing?", additional_kwargs={}),
|
||||
AIMessage(content="I'm doing well, thanks!", additional_kwargs={}),
|
||||
HumanMessage(content="What is your name?", additional_kwargs={}),
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -357,11 +357,11 @@ async def test_few_shot_chat_message_prompt_template() -> None:
|
||||
|
||||
expected = [
|
||||
SystemMessage(content="You are a helpful AI Assistant", additional_kwargs={}),
|
||||
HumanMessage(content="2+2", additional_kwargs={}, example=False),
|
||||
AIMessage(content="4", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="2+3", additional_kwargs={}, example=False),
|
||||
AIMessage(content="5", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="100 + 1", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="2+2", additional_kwargs={}),
|
||||
AIMessage(content="4", additional_kwargs={}),
|
||||
HumanMessage(content="2+3", additional_kwargs={}),
|
||||
AIMessage(content="5", additional_kwargs={}),
|
||||
HumanMessage(content="100 + 1", additional_kwargs={}),
|
||||
]
|
||||
|
||||
messages = final_prompt.format_messages(input="100 + 1")
|
||||
@@ -432,11 +432,11 @@ def test_few_shot_chat_message_prompt_template_with_selector() -> None:
|
||||
)
|
||||
expected = [
|
||||
SystemMessage(content="You are a helpful AI Assistant", additional_kwargs={}),
|
||||
HumanMessage(content="2+2", additional_kwargs={}, example=False),
|
||||
AIMessage(content="4", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="2+3", additional_kwargs={}, example=False),
|
||||
AIMessage(content="5", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="100 + 1", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="2+2", additional_kwargs={}),
|
||||
AIMessage(content="4", additional_kwargs={}),
|
||||
HumanMessage(content="2+3", additional_kwargs={}),
|
||||
AIMessage(content="5", additional_kwargs={}),
|
||||
HumanMessage(content="100 + 1", additional_kwargs={}),
|
||||
]
|
||||
messages = final_prompt.format_messages(input="100 + 1")
|
||||
assert messages == expected
|
||||
@@ -531,11 +531,11 @@ async def test_few_shot_chat_message_prompt_template_with_selector_async() -> No
|
||||
)
|
||||
expected = [
|
||||
SystemMessage(content="You are a helpful AI Assistant", additional_kwargs={}),
|
||||
HumanMessage(content="2+2", additional_kwargs={}, example=False),
|
||||
AIMessage(content="4", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="2+3", additional_kwargs={}, example=False),
|
||||
AIMessage(content="5", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="100 + 1", additional_kwargs={}, example=False),
|
||||
HumanMessage(content="2+2", additional_kwargs={}),
|
||||
AIMessage(content="4", additional_kwargs={}),
|
||||
HumanMessage(content="2+3", additional_kwargs={}),
|
||||
AIMessage(content="5", additional_kwargs={}),
|
||||
HumanMessage(content="100 + 1", additional_kwargs={}),
|
||||
]
|
||||
messages = await final_prompt.aformat_messages(input="100 + 1")
|
||||
assert messages == expected
|
||||
@@ -463,11 +463,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -546,6 +541,19 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
'type': 'string',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -567,11 +575,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -984,11 +987,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -1059,11 +1057,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -1171,6 +1164,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -1182,6 +1179,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -1200,9 +1208,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
|
||||
@@ -1995,11 +1995,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2077,6 +2072,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2098,11 +2105,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2510,11 +2512,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2584,11 +2581,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2695,6 +2687,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2706,6 +2702,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -2723,9 +2730,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -3394,11 +3402,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -3476,6 +3479,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -3497,11 +3512,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -3972,11 +3982,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -4046,11 +4051,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -4157,6 +4157,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -4168,6 +4172,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -4185,9 +4200,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -4887,11 +4903,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -4969,6 +4980,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -4990,11 +5013,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -5465,11 +5483,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -5539,11 +5552,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -5650,6 +5658,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -5661,6 +5673,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -5678,9 +5701,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -6318,11 +6342,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -6400,6 +6419,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -6421,11 +6452,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -6833,11 +6859,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -6907,11 +6928,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -7018,6 +7034,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -7029,6 +7049,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -7046,9 +7077,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -7759,11 +7791,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -7841,6 +7868,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -7862,11 +7901,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -8337,11 +8371,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -8411,11 +8440,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -8522,6 +8546,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -8533,6 +8561,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -8550,9 +8589,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -9235,11 +9275,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -9317,6 +9352,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -9338,11 +9385,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -9750,11 +9792,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -9824,11 +9861,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -9935,6 +9967,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -9946,6 +9982,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -9963,9 +10010,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -10584,11 +10632,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -10666,6 +10709,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -10687,11 +10742,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -11162,11 +11212,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -11236,11 +11281,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -11347,6 +11387,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -11358,6 +11402,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -11375,9 +11430,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
@@ -12038,11 +12094,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -12120,6 +12171,18 @@
|
||||
'title': 'Additional Kwargs',
|
||||
'type': 'object',
|
||||
}),
|
||||
'chunk_position': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'const': 'last',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Chunk Position',
|
||||
}),
|
||||
'content': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -12141,11 +12204,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -12616,11 +12674,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -12690,11 +12743,6 @@
|
||||
]),
|
||||
'title': 'Content',
|
||||
}),
|
||||
'example': dict({
|
||||
'default': False,
|
||||
'title': 'Example',
|
||||
'type': 'boolean',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -12801,6 +12849,10 @@
|
||||
]),
|
||||
'title': 'Error',
|
||||
}),
|
||||
'extras': dict({
|
||||
'title': 'Extras',
|
||||
'type': 'object',
|
||||
}),
|
||||
'id': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -12812,6 +12864,17 @@
|
||||
]),
|
||||
'title': 'Id',
|
||||
}),
|
||||
'index': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'integer',
|
||||
}),
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
]),
|
||||
'title': 'Index',
|
||||
}),
|
||||
'name': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
@@ -12829,9 +12892,10 @@
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'type',
|
||||
'id',
|
||||
'name',
|
||||
'args',
|
||||
'id',
|
||||
'error',
|
||||
]),
|
||||
'title': 'InvalidToolCall',
|
||||
|
||||
@@ -1,427 +0,0 @@
|
||||
import asyncio
|
||||
from typing import Any, Callable, NamedTuple, Union
|
||||
|
||||
import pytest
|
||||
|
||||
from langchain_core.beta.runnables.context import Context
|
||||
from langchain_core.language_models import FakeListLLM, FakeStreamingListLLM
|
||||
from langchain_core.output_parsers.string import StrOutputParser
|
||||
from langchain_core.prompt_values import StringPromptValue
|
||||
from langchain_core.prompts.prompt import PromptTemplate
|
||||
from langchain_core.runnables.base import Runnable, RunnableLambda
|
||||
from langchain_core.runnables.passthrough import RunnablePassthrough
|
||||
from langchain_core.runnables.utils import aadd, add
|
||||
|
||||
|
||||
class _TestCase(NamedTuple):
|
||||
input: Any
|
||||
output: Any
|
||||
|
||||
|
||||
def seq_naive_rag() -> Runnable:
|
||||
context = [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
]
|
||||
|
||||
retriever = RunnableLambda(lambda _: context)
|
||||
prompt = PromptTemplate.from_template("{context} {question}")
|
||||
llm = FakeListLLM(responses=["hello"])
|
||||
|
||||
return (
|
||||
Context.setter("input")
|
||||
| {
|
||||
"context": retriever | Context.setter("context"),
|
||||
"question": RunnablePassthrough(),
|
||||
}
|
||||
| prompt
|
||||
| llm
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"result": RunnablePassthrough(),
|
||||
"context": Context.getter("context"),
|
||||
"input": Context.getter("input"),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def seq_naive_rag_alt() -> Runnable:
|
||||
context = [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
]
|
||||
|
||||
retriever = RunnableLambda(lambda _: context)
|
||||
prompt = PromptTemplate.from_template("{context} {question}")
|
||||
llm = FakeListLLM(responses=["hello"])
|
||||
|
||||
return (
|
||||
Context.setter("input")
|
||||
| {
|
||||
"context": retriever | Context.setter("context"),
|
||||
"question": RunnablePassthrough(),
|
||||
}
|
||||
| prompt
|
||||
| llm
|
||||
| StrOutputParser()
|
||||
| Context.setter("result")
|
||||
| Context.getter(["context", "input", "result"])
|
||||
)
|
||||
|
||||
|
||||
def seq_naive_rag_scoped() -> Runnable:
|
||||
context = [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
]
|
||||
|
||||
retriever = RunnableLambda(lambda _: context)
|
||||
prompt = PromptTemplate.from_template("{context} {question}")
|
||||
llm = FakeListLLM(responses=["hello"])
|
||||
|
||||
scoped = Context.create_scope("a_scope")
|
||||
|
||||
return (
|
||||
Context.setter("input")
|
||||
| {
|
||||
"context": retriever | Context.setter("context"),
|
||||
"question": RunnablePassthrough(),
|
||||
"scoped": scoped.setter("context") | scoped.getter("context"),
|
||||
}
|
||||
| prompt
|
||||
| llm
|
||||
| StrOutputParser()
|
||||
| Context.setter("result")
|
||||
| Context.getter(["context", "input", "result"])
|
||||
)
|
||||
|
||||
|
||||
test_cases = [
|
||||
(
|
||||
Context.setter("foo") | Context.getter("foo"),
|
||||
(
|
||||
_TestCase("foo", "foo"),
|
||||
_TestCase("bar", "bar"),
|
||||
),
|
||||
),
|
||||
(
|
||||
Context.setter("input") | {"bar": Context.getter("input")},
|
||||
(
|
||||
_TestCase("foo", {"bar": "foo"}),
|
||||
_TestCase("bar", {"bar": "bar"}),
|
||||
),
|
||||
),
|
||||
(
|
||||
{"bar": Context.setter("input")} | Context.getter("input"),
|
||||
(
|
||||
_TestCase("foo", "foo"),
|
||||
_TestCase("bar", "bar"),
|
||||
),
|
||||
),
|
||||
(
|
||||
(
|
||||
PromptTemplate.from_template("{foo} {bar}")
|
||||
| Context.setter("prompt")
|
||||
| FakeListLLM(responses=["hello"])
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"response": RunnablePassthrough(),
|
||||
"prompt": Context.getter("prompt"),
|
||||
}
|
||||
),
|
||||
(
|
||||
_TestCase(
|
||||
{"foo": "foo", "bar": "bar"},
|
||||
{"response": "hello", "prompt": StringPromptValue(text="foo bar")},
|
||||
),
|
||||
_TestCase(
|
||||
{"foo": "bar", "bar": "foo"},
|
||||
{"response": "hello", "prompt": StringPromptValue(text="bar foo")},
|
||||
),
|
||||
),
|
||||
),
|
||||
(
|
||||
(
|
||||
PromptTemplate.from_template("{foo} {bar}")
|
||||
| Context.setter("prompt", prompt_str=lambda x: x.to_string())
|
||||
| FakeListLLM(responses=["hello"])
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"response": RunnablePassthrough(),
|
||||
"prompt": Context.getter("prompt"),
|
||||
"prompt_str": Context.getter("prompt_str"),
|
||||
}
|
||||
),
|
||||
(
|
||||
_TestCase(
|
||||
{"foo": "foo", "bar": "bar"},
|
||||
{
|
||||
"response": "hello",
|
||||
"prompt": StringPromptValue(text="foo bar"),
|
||||
"prompt_str": "foo bar",
|
||||
},
|
||||
),
|
||||
_TestCase(
|
||||
{"foo": "bar", "bar": "foo"},
|
||||
{
|
||||
"response": "hello",
|
||||
"prompt": StringPromptValue(text="bar foo"),
|
||||
"prompt_str": "bar foo",
|
||||
},
|
||||
),
|
||||
),
|
||||
),
|
||||
(
|
||||
(
|
||||
PromptTemplate.from_template("{foo} {bar}")
|
||||
| Context.setter(prompt_str=lambda x: x.to_string())
|
||||
| FakeListLLM(responses=["hello"])
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"response": RunnablePassthrough(),
|
||||
"prompt_str": Context.getter("prompt_str"),
|
||||
}
|
||||
),
|
||||
(
|
||||
_TestCase(
|
||||
{"foo": "foo", "bar": "bar"},
|
||||
{"response": "hello", "prompt_str": "foo bar"},
|
||||
),
|
||||
_TestCase(
|
||||
{"foo": "bar", "bar": "foo"},
|
||||
{"response": "hello", "prompt_str": "bar foo"},
|
||||
),
|
||||
),
|
||||
),
|
||||
(
|
||||
(
|
||||
PromptTemplate.from_template("{foo} {bar}")
|
||||
| Context.setter("prompt_str", lambda x: x.to_string())
|
||||
| FakeListLLM(responses=["hello"])
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"response": RunnablePassthrough(),
|
||||
"prompt_str": Context.getter("prompt_str"),
|
||||
}
|
||||
),
|
||||
(
|
||||
_TestCase(
|
||||
{"foo": "foo", "bar": "bar"},
|
||||
{"response": "hello", "prompt_str": "foo bar"},
|
||||
),
|
||||
_TestCase(
|
||||
{"foo": "bar", "bar": "foo"},
|
||||
{"response": "hello", "prompt_str": "bar foo"},
|
||||
),
|
||||
),
|
||||
),
|
||||
(
|
||||
(
|
||||
PromptTemplate.from_template("{foo} {bar}")
|
||||
| Context.setter("prompt")
|
||||
| FakeStreamingListLLM(responses=["hello"])
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"response": RunnablePassthrough(),
|
||||
"prompt": Context.getter("prompt"),
|
||||
}
|
||||
),
|
||||
(
|
||||
_TestCase(
|
||||
{"foo": "foo", "bar": "bar"},
|
||||
{"response": "hello", "prompt": StringPromptValue(text="foo bar")},
|
||||
),
|
||||
_TestCase(
|
||||
{"foo": "bar", "bar": "foo"},
|
||||
{"response": "hello", "prompt": StringPromptValue(text="bar foo")},
|
||||
),
|
||||
),
|
||||
),
|
||||
(
|
||||
seq_naive_rag,
|
||||
(
|
||||
_TestCase(
|
||||
"What up",
|
||||
{
|
||||
"result": "hello",
|
||||
"context": [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
],
|
||||
"input": "What up",
|
||||
},
|
||||
),
|
||||
_TestCase(
|
||||
"Howdy",
|
||||
{
|
||||
"result": "hello",
|
||||
"context": [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
],
|
||||
"input": "Howdy",
|
||||
},
|
||||
),
|
||||
),
|
||||
),
|
||||
(
|
||||
seq_naive_rag_alt,
|
||||
(
|
||||
_TestCase(
|
||||
"What up",
|
||||
{
|
||||
"result": "hello",
|
||||
"context": [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
],
|
||||
"input": "What up",
|
||||
},
|
||||
),
|
||||
_TestCase(
|
||||
"Howdy",
|
||||
{
|
||||
"result": "hello",
|
||||
"context": [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
],
|
||||
"input": "Howdy",
|
||||
},
|
||||
),
|
||||
),
|
||||
),
|
||||
(
|
||||
seq_naive_rag_scoped,
|
||||
(
|
||||
_TestCase(
|
||||
"What up",
|
||||
{
|
||||
"result": "hello",
|
||||
"context": [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
],
|
||||
"input": "What up",
|
||||
},
|
||||
),
|
||||
_TestCase(
|
||||
"Howdy",
|
||||
{
|
||||
"result": "hello",
|
||||
"context": [
|
||||
"Hi there!",
|
||||
"How are you?",
|
||||
"What's your name?",
|
||||
],
|
||||
"input": "Howdy",
|
||||
},
|
||||
),
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("runnable", "cases"), test_cases)
|
||||
def test_context_runnables(
|
||||
runnable: Union[Runnable, Callable[[], Runnable]], cases: list[_TestCase]
|
||||
) -> None:
|
||||
runnable = runnable if isinstance(runnable, Runnable) else runnable()
|
||||
assert runnable.invoke(cases[0].input) == cases[0].output
|
||||
assert runnable.batch([case.input for case in cases]) == [
|
||||
case.output for case in cases
|
||||
]
|
||||
assert add(runnable.stream(cases[0].input)) == cases[0].output
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("runnable", "cases"), test_cases)
|
||||
async def test_context_runnables_async(
|
||||
runnable: Union[Runnable, Callable[[], Runnable]], cases: list[_TestCase]
|
||||
) -> None:
|
||||
runnable = runnable if isinstance(runnable, Runnable) else runnable()
|
||||
assert await runnable.ainvoke(cases[1].input) == cases[1].output
|
||||
assert await runnable.abatch([case.input for case in cases]) == [
|
||||
case.output for case in cases
|
||||
]
|
||||
assert await aadd(runnable.astream(cases[1].input)) == cases[1].output
|
||||
|
||||
|
||||
def test_runnable_context_seq_key_not_found() -> None:
|
||||
seq: Runnable = {"bar": Context.setter("input")} | Context.getter("foo")
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="Expected exactly one setter for context key foo"
|
||||
):
|
||||
seq.invoke("foo")
|
||||
|
||||
|
||||
def test_runnable_context_seq_key_order() -> None:
|
||||
seq: Runnable = {"bar": Context.getter("foo")} | Context.setter("foo")
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="Context setter for key foo must be defined after all getters",
|
||||
):
|
||||
seq.invoke("foo")
|
||||
|
||||
|
||||
def test_runnable_context_deadlock() -> None:
|
||||
seq: Runnable = {
|
||||
"bar": Context.setter("input") | Context.getter("foo"),
|
||||
"foo": Context.setter("foo") | Context.getter("input"),
|
||||
} | RunnablePassthrough()
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="Deadlock detected between context keys foo and input"
|
||||
):
|
||||
seq.invoke("foo")
|
||||
|
||||
|
||||
def test_runnable_context_seq_key_circular_ref() -> None:
|
||||
seq: Runnable = {
|
||||
"bar": Context.setter(input=Context.getter("input"))
|
||||
} | Context.getter("foo")
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="Circular reference in context setter for key input"
|
||||
):
|
||||
seq.invoke("foo")
|
||||
|
||||
|
||||
async def test_runnable_seq_streaming_chunks() -> None:
|
||||
chain: Runnable = (
|
||||
PromptTemplate.from_template("{foo} {bar}")
|
||||
| Context.setter("prompt")
|
||||
| FakeStreamingListLLM(responses=["hello"])
|
||||
| StrOutputParser()
|
||||
| {
|
||||
"response": RunnablePassthrough(),
|
||||
"prompt": Context.getter("prompt"),
|
||||
}
|
||||
)
|
||||
chunks = await asyncio.to_thread(list, chain.stream({"foo": "foo", "bar": "bar"}))
|
||||
achunks = [c async for c in chain.astream({"foo": "foo", "bar": "bar"})]
|
||||
for c in chunks:
|
||||
assert c in achunks
|
||||
for c in achunks:
|
||||
assert c in chunks
|
||||
|
||||
assert len(chunks) == 6
|
||||
assert [c for c in chunks if c.get("response")] == [
|
||||
{"response": "h"},
|
||||
{"response": "e"},
|
||||
{"response": "l"},
|
||||
{"response": "l"},
|
||||
{"response": "o"},
|
||||
]
|
||||
assert [c for c in chunks if c.get("prompt")] == [
|
||||
{"prompt": StringPromptValue(text="foo bar")},
|
||||
]
|
||||
@@ -18,7 +18,7 @@ from langchain_core.language_models import (
|
||||
LanguageModelInput,
|
||||
)
|
||||
from langchain_core.load import dumps
|
||||
from langchain_core.messages import BaseMessage
|
||||
from langchain_core.messages import AIMessage, BaseMessage
|
||||
from langchain_core.outputs import ChatResult
|
||||
from langchain_core.prompts import PromptTemplate
|
||||
from langchain_core.runnables import (
|
||||
@@ -340,7 +340,7 @@ class FakeStructuredOutputModel(BaseChatModel):
|
||||
self,
|
||||
tools: Sequence[Union[dict[str, Any], type[BaseModel], Callable, BaseTool]],
|
||||
**kwargs: Any,
|
||||
) -> Runnable[LanguageModelInput, BaseMessage]:
|
||||
) -> Runnable[LanguageModelInput, AIMessage]:
|
||||
return self.bind(tools=tools)
|
||||
|
||||
@override
|
||||
@@ -373,7 +373,7 @@ class FakeModel(BaseChatModel):
|
||||
self,
|
||||
tools: Sequence[Union[dict[str, Any], type[BaseModel], Callable, BaseTool]],
|
||||
**kwargs: Any,
|
||||
) -> Runnable[LanguageModelInput, BaseMessage]:
|
||||
) -> Runnable[LanguageModelInput, AIMessage]:
|
||||
return self.bind(tools=tools)
|
||||
|
||||
@property
|
||||
|
||||
@@ -531,9 +531,6 @@ def test_passthrough_assign_schema() -> None:
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 9), reason="Requires python version >= 3.9 to run."
|
||||
)
|
||||
def test_lambda_schemas(snapshot: SnapshotAssertion) -> None:
|
||||
first_lambda = lambda x: x["hello"] # noqa: E731
|
||||
assert RunnableLambda(first_lambda).get_input_jsonschema() == {
|
||||
@@ -1856,7 +1853,7 @@ def test_prompt_with_chat_model(
|
||||
] == [
|
||||
_any_id_ai_message_chunk(content="f"),
|
||||
_any_id_ai_message_chunk(content="o"),
|
||||
_any_id_ai_message_chunk(content="o"),
|
||||
_any_id_ai_message_chunk(content="o", chunk_position="last"),
|
||||
]
|
||||
assert prompt_spy.call_args.args[1] == {"question": "What is your name?"}
|
||||
assert chat_spy.call_args.args[1] == ChatPromptValue(
|
||||
@@ -1965,7 +1962,7 @@ async def test_prompt_with_chat_model_async(
|
||||
] == [
|
||||
_any_id_ai_message_chunk(content="f"),
|
||||
_any_id_ai_message_chunk(content="o"),
|
||||
_any_id_ai_message_chunk(content="o"),
|
||||
_any_id_ai_message_chunk(content="o", chunk_position="last"),
|
||||
]
|
||||
assert prompt_spy.call_args.args[1] == {"question": "What is your name?"}
|
||||
assert chat_spy.call_args.args[1] == ChatPromptValue(
|
||||
@@ -4790,9 +4787,6 @@ async def test_runnable_branch_astream_with_callbacks() -> None:
|
||||
assert tracer.runs[2].outputs == {"output": "bye"}
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 9), reason="Requires python version >= 3.9 to run."
|
||||
)
|
||||
def test_representation_of_runnables() -> None:
|
||||
"""Test representation of runnables."""
|
||||
runnable = RunnableLambda(lambda x: x * 2)
|
||||
|
||||
@@ -503,7 +503,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="hello")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="hello",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {"a": "b"},
|
||||
"name": "my_model",
|
||||
@@ -521,7 +525,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="world!")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="world!", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {"a": "b"},
|
||||
"name": "my_model",
|
||||
@@ -530,7 +538,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"output": _any_id_ai_message_chunk(content="hello world!")},
|
||||
"data": {
|
||||
"output": _any_id_ai_message_chunk(
|
||||
content="hello world!", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_end",
|
||||
"metadata": {"a": "b"},
|
||||
"name": "my_model",
|
||||
@@ -574,7 +586,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="hello")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="hello",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -600,7 +616,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="world!")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="world!", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -698,7 +718,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="hello")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="hello",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -724,7 +748,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="world!")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="world!", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -891,7 +919,12 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="hello", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="hello",
|
||||
id="ai1",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -905,7 +938,12 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="hello", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="hello",
|
||||
id="ai1",
|
||||
)
|
||||
},
|
||||
"event": "on_chain_stream",
|
||||
"metadata": {"foo": "bar"},
|
||||
"name": "my_chain",
|
||||
@@ -937,7 +975,11 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="world!", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="world!", id="ai1", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -951,7 +993,11 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="world!", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="world!", id="ai1", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chain_stream",
|
||||
"metadata": {"foo": "bar"},
|
||||
"name": "my_chain",
|
||||
@@ -975,7 +1021,9 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
{
|
||||
"generation_info": None,
|
||||
"message": AIMessageChunk(
|
||||
content="hello world!", id="ai1"
|
||||
content="hello world!",
|
||||
id="ai1",
|
||||
chunk_position="last",
|
||||
),
|
||||
"text": "hello world!",
|
||||
"type": "ChatGenerationChunk",
|
||||
@@ -1000,7 +1048,11 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"output": AIMessageChunk(content="hello world!", id="ai1")},
|
||||
"data": {
|
||||
"output": AIMessageChunk(
|
||||
content="hello world!", id="ai1", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chain_end",
|
||||
"metadata": {"foo": "bar"},
|
||||
"name": "my_chain",
|
||||
@@ -1851,7 +1903,12 @@ async def test_events_astream_config() -> None:
|
||||
"tags": [],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="Goodbye", id="ai2")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="Goodbye",
|
||||
id="ai2",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {},
|
||||
"name": "RunnableConfigurableFields",
|
||||
@@ -1869,7 +1926,11 @@ async def test_events_astream_config() -> None:
|
||||
"tags": [],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="world", id="ai2")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="world", id="ai2", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {},
|
||||
"name": "RunnableConfigurableFields",
|
||||
@@ -1878,7 +1939,11 @@ async def test_events_astream_config() -> None:
|
||||
"tags": [],
|
||||
},
|
||||
{
|
||||
"data": {"output": AIMessageChunk(content="Goodbye world", id="ai2")},
|
||||
"data": {
|
||||
"output": AIMessageChunk(
|
||||
content="Goodbye world", id="ai2", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_end",
|
||||
"metadata": {},
|
||||
"name": "RunnableConfigurableFields",
|
||||
|
||||
@@ -540,7 +540,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="hello")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="hello",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -566,7 +570,11 @@ async def test_astream_events_from_model() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="world!")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="world!", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -580,7 +588,9 @@ async def test_astream_events_from_model() -> None:
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"output": _any_id_ai_message_chunk(content="hello world!"),
|
||||
"output": _any_id_ai_message_chunk(
|
||||
content="hello world!", chunk_position="last"
|
||||
),
|
||||
},
|
||||
"event": "on_chat_model_end",
|
||||
"metadata": {
|
||||
@@ -646,7 +656,11 @@ async def test_astream_with_model_in_chain() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="hello")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="hello",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -672,7 +686,11 @@ async def test_astream_with_model_in_chain() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="world!")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="world!", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -754,7 +772,11 @@ async def test_astream_with_model_in_chain() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="hello")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="hello",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -780,7 +802,11 @@ async def test_astream_with_model_in_chain() -> None:
|
||||
"tags": ["my_model"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": _any_id_ai_message_chunk(content="world!")},
|
||||
"data": {
|
||||
"chunk": _any_id_ai_message_chunk(
|
||||
content="world!", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -931,7 +957,12 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="hello", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="hello",
|
||||
id="ai1",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -945,7 +976,12 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="hello", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="hello",
|
||||
id="ai1",
|
||||
)
|
||||
},
|
||||
"event": "on_chain_stream",
|
||||
"metadata": {"foo": "bar"},
|
||||
"name": "my_chain",
|
||||
@@ -977,7 +1013,11 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="world!", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="world!", id="ai1", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {
|
||||
"a": "b",
|
||||
@@ -991,7 +1031,11 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="world!", id="ai1")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="world!", id="ai1", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chain_stream",
|
||||
"metadata": {"foo": "bar"},
|
||||
"name": "my_chain",
|
||||
@@ -1009,7 +1053,9 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
]
|
||||
]
|
||||
},
|
||||
"output": AIMessageChunk(content="hello world!", id="ai1"),
|
||||
"output": AIMessageChunk(
|
||||
content="hello world!", id="ai1", chunk_position="last"
|
||||
),
|
||||
},
|
||||
"event": "on_chat_model_end",
|
||||
"metadata": {
|
||||
@@ -1024,7 +1070,11 @@ async def test_event_stream_with_simple_chain() -> None:
|
||||
"tags": ["my_chain", "my_model", "seq:step:2"],
|
||||
},
|
||||
{
|
||||
"data": {"output": AIMessageChunk(content="hello world!", id="ai1")},
|
||||
"data": {
|
||||
"output": AIMessageChunk(
|
||||
content="hello world!", id="ai1", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chain_end",
|
||||
"metadata": {"foo": "bar"},
|
||||
"name": "my_chain",
|
||||
@@ -1806,7 +1856,12 @@ async def test_events_astream_config() -> None:
|
||||
"tags": [],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="Goodbye", id="ai2")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="Goodbye",
|
||||
id="ai2",
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {"ls_model_type": "chat"},
|
||||
"name": "GenericFakeChatModel",
|
||||
@@ -1824,7 +1879,11 @@ async def test_events_astream_config() -> None:
|
||||
"tags": [],
|
||||
},
|
||||
{
|
||||
"data": {"chunk": AIMessageChunk(content="world", id="ai2")},
|
||||
"data": {
|
||||
"chunk": AIMessageChunk(
|
||||
content="world", id="ai2", chunk_position="last"
|
||||
)
|
||||
},
|
||||
"event": "on_chat_model_stream",
|
||||
"metadata": {"ls_model_type": "chat"},
|
||||
"name": "GenericFakeChatModel",
|
||||
@@ -1834,7 +1893,9 @@ async def test_events_astream_config() -> None:
|
||||
},
|
||||
{
|
||||
"data": {
|
||||
"output": AIMessageChunk(content="Goodbye world", id="ai2"),
|
||||
"output": AIMessageChunk(
|
||||
content="Goodbye world", id="ai2", chunk_position="last"
|
||||
),
|
||||
},
|
||||
"event": "on_chat_model_end",
|
||||
"metadata": {"ls_model_type": "chat"},
|
||||
|
||||
@@ -5,7 +5,7 @@ import sys
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, Coroutine, Generator
|
||||
from inspect import isasyncgenfunction
|
||||
from typing import Any, Callable, Optional
|
||||
from typing import Any, Callable, Literal, Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -13,7 +13,6 @@ from langsmith import Client, get_current_run_tree, traceable
|
||||
from langsmith.run_helpers import tracing_context
|
||||
from langsmith.run_trees import RunTree
|
||||
from langsmith.utils import get_env_var
|
||||
from typing_extensions import Literal
|
||||
|
||||
from langchain_core.callbacks import BaseCallbackHandler
|
||||
from langchain_core.runnables.base import RunnableLambda, RunnableParallel
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import sys
|
||||
from typing import Callable
|
||||
|
||||
import pytest
|
||||
@@ -11,9 +10,6 @@ from langchain_core.runnables.utils import (
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 9), reason="Requires python version >= 3.9 to run."
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("func", "expected_source"),
|
||||
[
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import uuid
|
||||
from typing import Optional, Union
|
||||
from typing import Optional, Union, get_args
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -29,6 +29,7 @@ from langchain_core.messages import (
|
||||
messages_from_dict,
|
||||
messages_to_dict,
|
||||
)
|
||||
from langchain_core.messages.content import KNOWN_BLOCK_TYPES, ContentBlock
|
||||
from langchain_core.messages.tool import invalid_tool_call as create_invalid_tool_call
|
||||
from langchain_core.messages.tool import tool_call as create_tool_call
|
||||
from langchain_core.messages.tool import tool_call_chunk as create_tool_call_chunk
|
||||
@@ -177,21 +178,23 @@ def test_message_chunks() -> None:
|
||||
assert AIMessageChunk(content="") + left == left
|
||||
assert right + AIMessageChunk(content="") == right
|
||||
|
||||
default_id = "lc_run--abc123"
|
||||
meaningful_id = "msg_def456"
|
||||
|
||||
# Test ID order of precedence
|
||||
null_id = AIMessageChunk(content="", id=None)
|
||||
default_id = AIMessageChunk(
|
||||
content="", id="run-abc123"
|
||||
null_id_chunk = AIMessageChunk(content="", id=None)
|
||||
default_id_chunk = AIMessageChunk(
|
||||
content="", id=default_id
|
||||
) # LangChain-assigned run ID
|
||||
meaningful_id = AIMessageChunk(content="", id="msg_def456") # provider-assigned ID
|
||||
provider_chunk = AIMessageChunk(
|
||||
content="", id=meaningful_id
|
||||
) # provided ID (either by user or provider)
|
||||
|
||||
assert (null_id + default_id).id == "run-abc123"
|
||||
assert (default_id + null_id).id == "run-abc123"
|
||||
assert (null_id_chunk + default_id_chunk).id == default_id
|
||||
assert (null_id_chunk + provider_chunk).id == meaningful_id
|
||||
|
||||
assert (null_id + meaningful_id).id == "msg_def456"
|
||||
assert (meaningful_id + null_id).id == "msg_def456"
|
||||
|
||||
assert (default_id + meaningful_id).id == "msg_def456"
|
||||
assert (meaningful_id + default_id).id == "msg_def456"
|
||||
# Provider assigned IDs have highest precedence
|
||||
assert (default_id_chunk + provider_chunk).id == meaningful_id
|
||||
|
||||
|
||||
def test_chat_message_chunks() -> None:
|
||||
@@ -319,20 +322,12 @@ def test_function_message_chunks() -> None:
|
||||
|
||||
|
||||
def test_ai_message_chunks() -> None:
|
||||
assert AIMessageChunk(example=True, content="I am") + AIMessageChunk(
|
||||
example=True, content=" indeed."
|
||||
) == AIMessageChunk(example=True, content="I am indeed."), (
|
||||
assert AIMessageChunk(content="I am") + AIMessageChunk(
|
||||
content=" indeed."
|
||||
) == AIMessageChunk(content="I am indeed."), (
|
||||
"AIMessageChunk + AIMessageChunk should be a AIMessageChunk"
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="Cannot concatenate AIMessageChunks with different example values",
|
||||
):
|
||||
AIMessageChunk(example=True, content="I am") + AIMessageChunk(
|
||||
example=False, content=" indeed."
|
||||
)
|
||||
|
||||
|
||||
class TestGetBufferString:
|
||||
_HUMAN_MSG = HumanMessage(content="human")
|
||||
@@ -750,7 +745,6 @@ def test_convert_to_messages() -> None:
|
||||
"type": "human",
|
||||
"name": None,
|
||||
"id": "1",
|
||||
"example": False,
|
||||
},
|
||||
]
|
||||
)
|
||||
@@ -786,7 +780,6 @@ def test_convert_to_messages() -> None:
|
||||
additional_kwargs={"metadata": {"speaker_name": "Presenter"}},
|
||||
response_metadata={},
|
||||
id="1",
|
||||
example=False,
|
||||
),
|
||||
]
|
||||
assert expected == actual
|
||||
@@ -1034,12 +1027,13 @@ def test_tool_message_content() -> None:
|
||||
ToolMessage(["foo"], tool_call_id="1")
|
||||
ToolMessage([{"foo": "bar"}], tool_call_id="1")
|
||||
|
||||
assert ToolMessage(("a", "b", "c"), tool_call_id="1").content == ["a", "b", "c"] # type: ignore[arg-type]
|
||||
assert ToolMessage(5, tool_call_id="1").content == "5" # type: ignore[arg-type]
|
||||
assert ToolMessage(5.1, tool_call_id="1").content == "5.1" # type: ignore[arg-type]
|
||||
assert ToolMessage({"foo": "bar"}, tool_call_id="1").content == "{'foo': 'bar'}" # type: ignore[arg-type]
|
||||
# Ignoring since we're testing that tuples get converted to lists in `coerce_args`
|
||||
assert ToolMessage(("a", "b", "c"), tool_call_id="1").content == ["a", "b", "c"] # type: ignore[call-overload]
|
||||
assert ToolMessage(5, tool_call_id="1").content == "5" # type: ignore[call-overload]
|
||||
assert ToolMessage(5.1, tool_call_id="1").content == "5.1" # type: ignore[call-overload]
|
||||
assert ToolMessage({"foo": "bar"}, tool_call_id="1").content == "{'foo': 'bar'}" # type: ignore[call-overload]
|
||||
assert (
|
||||
ToolMessage(Document("foo"), tool_call_id="1").content == "page_content='foo'" # type: ignore[arg-type]
|
||||
ToolMessage(Document("foo"), tool_call_id="1").content == "page_content='foo'" # type: ignore[call-overload]
|
||||
)
|
||||
|
||||
|
||||
@@ -1057,9 +1051,9 @@ def test_message_text() -> None:
|
||||
# content: [empty], [single element], [multiple elements]
|
||||
# content dict types: [text], [not text], [no type]
|
||||
|
||||
assert HumanMessage(content="foo").text() == "foo"
|
||||
assert not AIMessage(content=[]).text()
|
||||
assert AIMessage(content=["foo", "bar"]).text() == "foobar"
|
||||
assert HumanMessage(content="foo").text == "foo"
|
||||
assert AIMessage(content=[]).text == ""
|
||||
assert AIMessage(content=["foo", "bar"]).text == "foobar"
|
||||
assert (
|
||||
AIMessage(
|
||||
content=[
|
||||
@@ -1071,12 +1065,11 @@ def test_message_text() -> None:
|
||||
"input": {"location": "San Francisco, CA"},
|
||||
},
|
||||
]
|
||||
).text()
|
||||
).text
|
||||
== "<thinking>thinking...</thinking>"
|
||||
)
|
||||
assert (
|
||||
SystemMessage(content=[{"type": "text", "text": "foo"}, "bar"]).text()
|
||||
== "foobar"
|
||||
SystemMessage(content=[{"type": "text", "text": "foo"}, "bar"]).text == "foobar"
|
||||
)
|
||||
assert (
|
||||
ToolMessage(
|
||||
@@ -1092,40 +1085,62 @@ def test_message_text() -> None:
|
||||
},
|
||||
],
|
||||
tool_call_id="1",
|
||||
).text()
|
||||
).text
|
||||
== "15 degrees"
|
||||
)
|
||||
assert (
|
||||
AIMessage(content=[{"text": "hi there"}, "hi"]).text() == "hi"
|
||||
AIMessage(content=[{"text": "hi there"}, "hi"]).text == "hi"
|
||||
) # missing type: text
|
||||
assert not AIMessage(content=[{"type": "nottext", "text": "hi"}]).text()
|
||||
assert not AIMessage(content=[]).text()
|
||||
assert not AIMessage(
|
||||
content="", tool_calls=[create_tool_call(name="a", args={"b": 1}, id=None)]
|
||||
).text()
|
||||
assert AIMessage(content=[{"type": "nottext", "text": "hi"}]).text == ""
|
||||
assert AIMessage(content=[]).text == ""
|
||||
assert (
|
||||
AIMessage(
|
||||
content="", tool_calls=[create_tool_call(name="a", args={"b": 1}, id=None)]
|
||||
).text
|
||||
== ""
|
||||
)
|
||||
|
||||
|
||||
def test_is_data_content_block() -> None:
|
||||
# Test all DataContentBlock types with various data fields
|
||||
|
||||
# Image blocks
|
||||
assert is_data_content_block({"type": "image", "url": "https://..."})
|
||||
assert is_data_content_block(
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://...",
|
||||
}
|
||||
{"type": "image", "base64": "<base64 data>", "mime_type": "image/jpeg"}
|
||||
)
|
||||
|
||||
# Video blocks
|
||||
assert is_data_content_block({"type": "video", "url": "https://video.mp4"})
|
||||
assert is_data_content_block(
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"mime_type": "image/jpeg",
|
||||
}
|
||||
{"type": "video", "base64": "<base64 video>", "mime_type": "video/mp4"}
|
||||
)
|
||||
assert is_data_content_block({"type": "video", "file_id": "vid_123"})
|
||||
|
||||
# Audio blocks
|
||||
assert is_data_content_block({"type": "audio", "url": "https://audio.mp3"})
|
||||
assert is_data_content_block(
|
||||
{"type": "audio", "base64": "<base64 audio>", "mime_type": "audio/mp3"}
|
||||
)
|
||||
assert is_data_content_block({"type": "audio", "file_id": "aud_123"})
|
||||
|
||||
# Plain text blocks
|
||||
assert is_data_content_block({"type": "text-plain", "text": "document content"})
|
||||
assert is_data_content_block({"type": "text-plain", "url": "https://doc.txt"})
|
||||
assert is_data_content_block({"type": "text-plain", "file_id": "txt_123"})
|
||||
|
||||
# File blocks
|
||||
assert is_data_content_block({"type": "file", "url": "https://file.pdf"})
|
||||
assert is_data_content_block(
|
||||
{"type": "file", "base64": "<base64 file>", "mime_type": "application/pdf"}
|
||||
)
|
||||
assert is_data_content_block({"type": "file", "file_id": "file_123"})
|
||||
|
||||
# Blocks with additional metadata (should still be valid)
|
||||
assert is_data_content_block(
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "image/jpeg",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
@@ -1133,65 +1148,191 @@ def test_is_data_content_block() -> None:
|
||||
assert is_data_content_block(
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "image/jpeg",
|
||||
"metadata": {"cache_control": {"type": "ephemeral"}},
|
||||
}
|
||||
)
|
||||
|
||||
assert not is_data_content_block(
|
||||
assert is_data_content_block(
|
||||
{
|
||||
"type": "text",
|
||||
"text": "foo",
|
||||
"type": "image",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "image/jpeg",
|
||||
"extras": "hi",
|
||||
}
|
||||
)
|
||||
|
||||
# Invalid cases - wrong type
|
||||
assert not is_data_content_block({"type": "text", "text": "foo"})
|
||||
assert not is_data_content_block(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://..."},
|
||||
}
|
||||
)
|
||||
assert not is_data_content_block(
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "base64",
|
||||
}
|
||||
)
|
||||
assert not is_data_content_block(
|
||||
{
|
||||
"type": "image",
|
||||
"source": "<base64 data>",
|
||||
}
|
||||
} # This is OpenAI Chat Completions
|
||||
)
|
||||
assert not is_data_content_block({"type": "tool_call", "name": "func", "args": {}})
|
||||
assert not is_data_content_block({"type": "invalid", "url": "something"})
|
||||
|
||||
# Invalid cases - valid type but no data or `source_type` fields
|
||||
assert not is_data_content_block({"type": "image"})
|
||||
assert not is_data_content_block({"type": "video", "mime_type": "video/mp4"})
|
||||
assert not is_data_content_block({"type": "audio", "extras": {"key": "value"}})
|
||||
|
||||
# Invalid cases - valid type but wrong data field name
|
||||
assert not is_data_content_block({"type": "image", "source": "<base64 data>"})
|
||||
assert not is_data_content_block({"type": "video", "data": "video_data"})
|
||||
|
||||
# Edge cases - empty or missing values
|
||||
assert not is_data_content_block({})
|
||||
assert not is_data_content_block({"url": "https://..."}) # missing type
|
||||
|
||||
|
||||
def test_convert_to_openai_image_block() -> None:
|
||||
input_block = {
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://...",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
expected = {
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://..."},
|
||||
}
|
||||
result = convert_to_openai_image_block(input_block)
|
||||
assert result == expected
|
||||
|
||||
input_block = {
|
||||
"type": "image",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"mime_type": "image/jpeg",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
expected = {
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "data:image/jpeg;base64,<base64 data>",
|
||||
for input_block in [
|
||||
{
|
||||
"type": "image",
|
||||
"url": "https://...",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "url",
|
||||
"url": "https://...",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
]:
|
||||
expected = {
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://..."},
|
||||
}
|
||||
result = convert_to_openai_image_block(input_block)
|
||||
assert result == expected
|
||||
|
||||
for input_block in [
|
||||
{
|
||||
"type": "image",
|
||||
"base64": "<base64 data>",
|
||||
"mime_type": "image/jpeg",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
{
|
||||
"type": "image",
|
||||
"source_type": "base64",
|
||||
"data": "<base64 data>",
|
||||
"mime_type": "image/jpeg",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
]:
|
||||
expected = {
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "data:image/jpeg;base64,<base64 data>",
|
||||
},
|
||||
}
|
||||
result = convert_to_openai_image_block(input_block)
|
||||
assert result == expected
|
||||
|
||||
|
||||
def test_known_block_types() -> None:
|
||||
expected = {
|
||||
bt
|
||||
for bt in get_args(ContentBlock)
|
||||
for bt in get_args(bt.__annotations__["type"])
|
||||
}
|
||||
result = convert_to_openai_image_block(input_block)
|
||||
assert result == expected
|
||||
# Normalize any Literal[...] types in block types to their string values.
|
||||
# This ensures all entries are plain strings, not Literal objects.
|
||||
expected = {
|
||||
t
|
||||
if isinstance(t, str)
|
||||
else t.__args__[0]
|
||||
if hasattr(t, "__args__") and len(t.__args__) == 1
|
||||
else t
|
||||
for t in expected
|
||||
}
|
||||
assert expected == KNOWN_BLOCK_TYPES
|
||||
|
||||
|
||||
def test_typed_init() -> None:
|
||||
ai_message = AIMessage(content_blocks=[{"type": "text", "text": "Hello"}])
|
||||
assert ai_message.content == [{"type": "text", "text": "Hello"}]
|
||||
assert ai_message.content_blocks == ai_message.content
|
||||
|
||||
human_message = HumanMessage(content_blocks=[{"type": "text", "text": "Hello"}])
|
||||
assert human_message.content == [{"type": "text", "text": "Hello"}]
|
||||
assert human_message.content_blocks == human_message.content
|
||||
|
||||
system_message = SystemMessage(content_blocks=[{"type": "text", "text": "Hello"}])
|
||||
assert system_message.content == [{"type": "text", "text": "Hello"}]
|
||||
assert system_message.content_blocks == system_message.content
|
||||
|
||||
tool_message = ToolMessage(
|
||||
content_blocks=[{"type": "text", "text": "Hello"}],
|
||||
tool_call_id="abc123",
|
||||
)
|
||||
assert tool_message.content == [{"type": "text", "text": "Hello"}]
|
||||
assert tool_message.content_blocks == tool_message.content
|
||||
|
||||
for message_class in [AIMessage, HumanMessage, SystemMessage]:
|
||||
message = message_class("Hello")
|
||||
assert message.content == "Hello"
|
||||
assert message.content_blocks == [{"type": "text", "text": "Hello"}]
|
||||
|
||||
message = message_class(content="Hello")
|
||||
assert message.content == "Hello"
|
||||
assert message.content_blocks == [{"type": "text", "text": "Hello"}]
|
||||
|
||||
# Test we get type errors for malformed blocks (type checker will complain if
|
||||
# below type-ignores are unused).
|
||||
_ = AIMessage(content_blocks=[{"type": "text", "bad": "Hello"}]) # type: ignore[list-item]
|
||||
_ = HumanMessage(content_blocks=[{"type": "text", "bad": "Hello"}]) # type: ignore[list-item]
|
||||
_ = SystemMessage(content_blocks=[{"type": "text", "bad": "Hello"}]) # type: ignore[list-item]
|
||||
_ = ToolMessage(
|
||||
content_blocks=[{"type": "text", "bad": "Hello"}], # type: ignore[list-item]
|
||||
tool_call_id="abc123",
|
||||
)
|
||||
|
||||
|
||||
def test_text_accessor() -> None:
|
||||
"""Test that `message.text` property and `.text()` method return the same value."""
|
||||
human_msg = HumanMessage(content="Hello world")
|
||||
assert human_msg.text == "Hello world"
|
||||
assert human_msg.text == "Hello world"
|
||||
assert str(human_msg.text) == str(human_msg.text)
|
||||
|
||||
system_msg = SystemMessage(content="You are a helpful assistant")
|
||||
assert system_msg.text == "You are a helpful assistant"
|
||||
assert system_msg.text == "You are a helpful assistant"
|
||||
assert str(system_msg.text) == str(system_msg.text)
|
||||
|
||||
ai_msg = AIMessage(content="I can help you with that")
|
||||
assert ai_msg.text == "I can help you with that"
|
||||
assert ai_msg.text == "I can help you with that"
|
||||
assert str(ai_msg.text) == str(ai_msg.text)
|
||||
|
||||
tool_msg = ToolMessage(content="Task completed", tool_call_id="tool_1")
|
||||
assert tool_msg.text == "Task completed"
|
||||
assert tool_msg.text == "Task completed"
|
||||
assert str(tool_msg.text) == str(tool_msg.text)
|
||||
|
||||
complex_msg = HumanMessage(
|
||||
content=[{"type": "text", "text": "Hello "}, {"type": "text", "text": "world"}]
|
||||
)
|
||||
assert complex_msg.text == "Hello world"
|
||||
assert complex_msg.text == "Hello world"
|
||||
assert str(complex_msg.text) == str(complex_msg.text)
|
||||
|
||||
mixed_msg = AIMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "The answer is "},
|
||||
{"type": "tool_use", "name": "calculate", "input": {"x": 2}, "id": "1"},
|
||||
{"type": "text", "text": "42"},
|
||||
]
|
||||
)
|
||||
assert mixed_msg.text == "The answer is 42"
|
||||
assert mixed_msg.text == "The answer is 42"
|
||||
assert str(mixed_msg.text) == str(mixed_msg.text)
|
||||
|
||||
empty_msg = HumanMessage(content=[])
|
||||
assert empty_msg.text == ""
|
||||
assert empty_msg.text == ""
|
||||
assert str(empty_msg.text) == str(empty_msg.text)
|
||||
@@ -2321,7 +2321,7 @@ def test_tool_injected_tool_call_id() -> None:
|
||||
@tool
|
||||
def foo(x: int, tool_call_id: Annotated[str, InjectedToolCallId]) -> ToolMessage:
|
||||
"""Foo."""
|
||||
return ToolMessage(x, tool_call_id=tool_call_id) # type: ignore[arg-type]
|
||||
return ToolMessage(x, tool_call_id=tool_call_id) # type: ignore[call-overload]
|
||||
|
||||
assert foo.invoke(
|
||||
{
|
||||
@@ -2330,7 +2330,7 @@ def test_tool_injected_tool_call_id() -> None:
|
||||
"name": "foo",
|
||||
"id": "bar",
|
||||
}
|
||||
) == ToolMessage(0, tool_call_id="bar") # type: ignore[arg-type]
|
||||
) == ToolMessage(0, tool_call_id="bar") # type: ignore[call-overload]
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
@@ -2342,7 +2342,7 @@ def test_tool_injected_tool_call_id() -> None:
|
||||
@tool
|
||||
def foo2(x: int, tool_call_id: Annotated[str, InjectedToolCallId()]) -> ToolMessage:
|
||||
"""Foo."""
|
||||
return ToolMessage(x, tool_call_id=tool_call_id) # type: ignore[arg-type]
|
||||
return ToolMessage(x, tool_call_id=tool_call_id) # type: ignore[call-overload]
|
||||
|
||||
assert foo2.invoke(
|
||||
{
|
||||
@@ -2351,7 +2351,7 @@ def test_tool_injected_tool_call_id() -> None:
|
||||
"name": "foo",
|
||||
"id": "bar",
|
||||
}
|
||||
) == ToolMessage(0, tool_call_id="bar") # type: ignore[arg-type]
|
||||
) == ToolMessage(0, tool_call_id="bar") # type: ignore[call-overload]
|
||||
|
||||
|
||||
def test_tool_injected_tool_call_id_override_llm_generated() -> None:
|
||||
@@ -2360,7 +2360,7 @@ def test_tool_injected_tool_call_id_override_llm_generated() -> None:
|
||||
@tool
|
||||
def foo(x: int, tool_call_id: Annotated[str, InjectedToolCallId]) -> ToolMessage:
|
||||
"""Foo."""
|
||||
return ToolMessage(x, tool_call_id=tool_call_id) # type: ignore[arg-type]
|
||||
return ToolMessage(str(x), tool_call_id=tool_call_id)
|
||||
|
||||
# Test that when LLM generates the tool_call_id, it gets overridden
|
||||
result = foo.invoke(
|
||||
@@ -2373,14 +2373,14 @@ def test_tool_injected_tool_call_id_override_llm_generated() -> None:
|
||||
)
|
||||
|
||||
# The tool should receive the real tool call ID, not the LLM-generated one
|
||||
assert result == ToolMessage(0, tool_call_id="real_tool_call_id") # type: ignore[arg-type]
|
||||
assert result == ToolMessage("0", tool_call_id="real_tool_call_id")
|
||||
|
||||
|
||||
def test_tool_uninjected_tool_call_id() -> None:
|
||||
@tool
|
||||
def foo(x: int, tool_call_id: str) -> ToolMessage:
|
||||
"""Foo."""
|
||||
return ToolMessage(x, tool_call_id=tool_call_id) # type: ignore[arg-type]
|
||||
return ToolMessage(str(x), tool_call_id=tool_call_id)
|
||||
|
||||
with pytest.raises(ValueError, match="1 validation error for foo"):
|
||||
foo.invoke({"type": "tool_call", "args": {"x": 0}, "name": "foo", "id": "bar"})
|
||||
@@ -2392,7 +2392,7 @@ def test_tool_uninjected_tool_call_id() -> None:
|
||||
"name": "foo",
|
||||
"id": "bar",
|
||||
}
|
||||
) == ToolMessage(0, tool_call_id="zap") # type: ignore[arg-type]
|
||||
) == ToolMessage(0, tool_call_id="zap") # type: ignore[call-overload]
|
||||
|
||||
|
||||
def test_tool_return_output_mixin() -> None:
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import sys
|
||||
import typing
|
||||
from collections.abc import Iterable, Mapping, MutableMapping, Sequence
|
||||
from typing import Annotated as ExtensionsAnnotated
|
||||
@@ -1051,12 +1050,9 @@ def test__convert_typed_dict_to_openai_function_fail(typed_dict: type) -> None:
|
||||
_convert_typed_dict_to_openai_function(Tool)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 10), reason="Requires python version >= 3.10 to run."
|
||||
)
|
||||
def test_convert_union_type_py_39() -> None:
|
||||
def test_convert_union_type() -> None:
|
||||
@tool
|
||||
def magic_function(value: int | str) -> str: # type: ignore[syntax,unused-ignore] # noqa: ARG001,FA102
|
||||
def magic_function(value: int | str) -> str: # noqa: ARG001,FA102
|
||||
"""Compute a magic function."""
|
||||
return ""
|
||||
|
||||
|
||||
Generated
+1090
-1194
File diff suppressed because it is too large.
Load diff
@@ -47,7 +47,12 @@ def parse_ai_message_to_tool_action(
|
||||
try:
|
||||
args = json.loads(function["arguments"] or "{}")
|
||||
tool_calls.append(
|
||||
ToolCall(name=function_name, args=args, id=tool_call["id"]),
|
||||
ToolCall(
|
||||
type="tool_call",
|
||||
name=function_name,
|
||||
args=args,
|
||||
id=tool_call["id"],
|
||||
),
|
||||
)
|
||||
except JSONDecodeError as e:
|
||||
msg = (
|
||||
|
||||
@@ -14,7 +14,7 @@ from langchain_core.language_models.chat_models import (
|
||||
agenerate_from_stream,
|
||||
generate_from_stream,
|
||||
)
|
||||
from langchain_core.messages import AnyMessage, BaseMessage
|
||||
from langchain_core.messages import AIMessage, AnyMessage
|
||||
from langchain_core.runnables import Runnable, RunnableConfig, ensure_config
|
||||
from langchain_core.runnables.schema import StreamEvent
|
||||
from langchain_core.tools import BaseTool
|
||||
@@ -948,7 +948,7 @@ class _ConfigurableModel(Runnable[LanguageModelInput, Any]):
|
||||
self,
|
||||
tools: Sequence[Union[dict[str, Any], type[BaseModel], Callable, BaseTool]],
|
||||
**kwargs: Any,
|
||||
) -> Runnable[LanguageModelInput, BaseMessage]:
|
||||
) -> Runnable[LanguageModelInput, AIMessage]:
|
||||
return self.__getattr__("bind_tools")(tools, **kwargs)
|
||||
|
||||
# Explicitly added to satisfy downstream linters.
|
||||
|
||||
@@ -5,7 +5,7 @@ build-backend = "pdm.backend"
|
||||
[project]
|
||||
authors = []
|
||||
license = { text = "MIT" }
|
||||
requires-python = ">=3.9.0,<4.0.0"
|
||||
requires-python = ">=3.10.0,<4.0.0"
|
||||
dependencies = [
|
||||
"langchain-core>=0.3.72,<1.0.0",
|
||||
"langchain-text-splitters>=0.3.9,<1.0.0",
|
||||
@@ -25,20 +25,20 @@ readme = "README.md"
|
||||
community = ["langchain-community"]
|
||||
anthropic = ["langchain-anthropic"]
|
||||
openai = ["langchain-openai"]
|
||||
azure-ai = ["langchain-azure-ai"]
|
||||
cohere = ["langchain-cohere"]
|
||||
# azure-ai = ["langchain-azure-ai"]
|
||||
# cohere = ["langchain-cohere"]
|
||||
google-vertexai = ["langchain-google-vertexai"]
|
||||
google-genai = ["langchain-google-genai"]
|
||||
fireworks = ["langchain-fireworks"]
|
||||
ollama = ["langchain-ollama"]
|
||||
# fireworks = ["langchain-fireworks"]
|
||||
# ollama = ["langchain-ollama"]
|
||||
together = ["langchain-together"]
|
||||
mistralai = ["langchain-mistralai"]
|
||||
huggingface = ["langchain-huggingface"]
|
||||
groq = ["langchain-groq"]
|
||||
aws = ["langchain-aws"]
|
||||
deepseek = ["langchain-deepseek"]
|
||||
xai = ["langchain-xai"]
|
||||
perplexity = ["langchain-perplexity"]
|
||||
# mistralai = ["langchain-mistralai"]
|
||||
# huggingface = ["langchain-huggingface"]
|
||||
# groq = ["langchain-groq"]
|
||||
# aws = ["langchain-aws"]
|
||||
# deepseek = ["langchain-deepseek"]
|
||||
# xai = ["langchain-xai"]
|
||||
# perplexity = ["langchain-perplexity"]
|
||||
|
||||
[project.urls]
|
||||
"Source Code" = "https://github.com/langchain-ai/langchain/tree/master/libs/langchain"
|
||||
@@ -128,7 +128,6 @@ strict = "True"
|
||||
strict_bytes = "True"
|
||||
ignore_missing_imports = "True"
|
||||
enable_error_code = "deprecated"
|
||||
report_deprecated_as_note = "True"
|
||||
warn_unreachable = "True"
|
||||
|
||||
# TODO: activate for 'strict' checking
|
||||
|
||||
@@ -907,8 +907,8 @@ async def test_openai_agent_with_streaming() -> None:
|
||||
"name": "find_pet",
|
||||
},
|
||||
},
|
||||
"chunk_position": "last",
|
||||
"content": "",
|
||||
"example": False,
|
||||
"invalid_tool_calls": [],
|
||||
"name": None,
|
||||
"response_metadata": {},
|
||||
@@ -946,7 +946,6 @@ async def test_openai_agent_with_streaming() -> None:
|
||||
{
|
||||
"additional_kwargs": {},
|
||||
"content": "The cat is spying from under the bed.",
|
||||
"example": False,
|
||||
"invalid_tool_calls": [],
|
||||
"name": None,
|
||||
"response_metadata": {},
|
||||
@@ -1111,6 +1110,7 @@ async def test_openai_agent_tools_agent() -> None:
|
||||
},
|
||||
],
|
||||
},
|
||||
chunk_position="last",
|
||||
),
|
||||
],
|
||||
tool_call_id="0",
|
||||
@@ -1137,6 +1137,7 @@ async def test_openai_agent_tools_agent() -> None:
|
||||
},
|
||||
],
|
||||
},
|
||||
chunk_position="last",
|
||||
),
|
||||
],
|
||||
},
|
||||
@@ -1167,6 +1168,7 @@ async def test_openai_agent_tools_agent() -> None:
|
||||
},
|
||||
],
|
||||
},
|
||||
chunk_position="last",
|
||||
),
|
||||
],
|
||||
tool_call_id="1",
|
||||
@@ -1193,6 +1195,7 @@ async def test_openai_agent_tools_agent() -> None:
|
||||
},
|
||||
],
|
||||
},
|
||||
chunk_position="last",
|
||||
),
|
||||
],
|
||||
},
|
||||
@@ -1230,6 +1233,7 @@ async def test_openai_agent_tools_agent() -> None:
|
||||
},
|
||||
],
|
||||
},
|
||||
chunk_position="last",
|
||||
),
|
||||
],
|
||||
tool_call_id="0",
|
||||
@@ -1273,6 +1277,7 @@ async def test_openai_agent_tools_agent() -> None:
|
||||
},
|
||||
],
|
||||
},
|
||||
chunk_position="last",
|
||||
),
|
||||
],
|
||||
tool_call_id="1",
|
||||
|
||||
@@ -142,7 +142,7 @@ def test_configurable() -> None:
|
||||
"openai_api_base": None,
|
||||
"openai_organization": None,
|
||||
"openai_proxy": None,
|
||||
"output_version": "v0",
|
||||
"output_version": None,
|
||||
"request_timeout": None,
|
||||
"max_retries": None,
|
||||
"presence_penalty": None,
|
||||
@@ -260,7 +260,7 @@ def test_configurable_with_default() -> None:
|
||||
"disable_streaming": False,
|
||||
"model": "claude-3-7-sonnet-20250219",
|
||||
"mcp_servers": None,
|
||||
"max_tokens": 1024,
|
||||
"max_tokens": 64000,
|
||||
"temperature": None,
|
||||
"thinking": None,
|
||||
"top_k": None,
|
||||
@@ -277,6 +277,7 @@ def test_configurable_with_default() -> None:
|
||||
"model_kwargs": {},
|
||||
"streaming": False,
|
||||
"stream_usage": True,
|
||||
"output_version": None,
|
||||
},
|
||||
"kwargs": {
|
||||
"tools": [{"name": "foo", "description": "foo", "input_schema": {}}],
|
||||
|
||||
@@ -130,10 +130,16 @@ class GenericFakeChatModel(BaseChatModel):
|
||||
assert isinstance(content, str)
|
||||
content_chunks = cast("list[str]", re.split(r"(\s)", content))
|
||||
|
||||
for token in content_chunks:
|
||||
for idx, token in enumerate(content_chunks):
|
||||
chunk = ChatGenerationChunk(
|
||||
message=AIMessageChunk(id=message.id, content=token),
|
||||
)
|
||||
if (
|
||||
idx == len(content_chunks) - 1
|
||||
and isinstance(chunk.message, AIMessageChunk)
|
||||
and not message.additional_kwargs
|
||||
):
|
||||
chunk.message.chunk_position = "last"
|
||||
if run_manager:
|
||||
run_manager.on_llm_new_token(token, chunk=chunk)
|
||||
yield chunk
|
||||
|
||||
@@ -49,14 +49,14 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
assert chunks == [
|
||||
_AnyIdAIMessageChunk(content="hello"),
|
||||
_AnyIdAIMessageChunk(content=" "),
|
||||
_AnyIdAIMessageChunk(content="goodbye"),
|
||||
_AnyIdAIMessageChunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
|
||||
chunks = list(model.stream("meow"))
|
||||
assert chunks == [
|
||||
_AnyIdAIMessageChunk(content="hello"),
|
||||
_AnyIdAIMessageChunk(content=" "),
|
||||
_AnyIdAIMessageChunk(content="goodbye"),
|
||||
_AnyIdAIMessageChunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
|
||||
# Test streaming of additional kwargs.
|
||||
@@ -67,6 +67,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
assert chunks == [
|
||||
_AnyIdAIMessageChunk(content="", additional_kwargs={"foo": 42}),
|
||||
_AnyIdAIMessageChunk(content="", additional_kwargs={"bar": 24}),
|
||||
_AnyIdAIMessageChunk(content="", chunk_position="last"),
|
||||
]
|
||||
|
||||
message = AIMessage(
|
||||
@@ -108,6 +109,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
"function_call": {"arguments": '\n "destination_path": "bar"\n}'},
|
||||
},
|
||||
),
|
||||
_AnyIdAIMessageChunk(content="", chunk_position="last"),
|
||||
]
|
||||
|
||||
accumulate_chunks = None
|
||||
@@ -127,6 +129,7 @@ async def test_generic_fake_chat_model_stream() -> None:
|
||||
'destination_path": "bar"\n}',
|
||||
},
|
||||
},
|
||||
chunk_position="last",
|
||||
)
|
||||
|
||||
|
||||
@@ -141,7 +144,7 @@ async def test_generic_fake_chat_model_astream_log() -> None:
|
||||
assert final.state["streamed_output"] == [
|
||||
_AnyIdAIMessageChunk(content="hello"),
|
||||
_AnyIdAIMessageChunk(content=" "),
|
||||
_AnyIdAIMessageChunk(content="goodbye"),
|
||||
_AnyIdAIMessageChunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
|
||||
|
||||
@@ -198,6 +201,6 @@ async def test_callback_handlers() -> None:
|
||||
assert results == [
|
||||
_AnyIdAIMessageChunk(content="hello"),
|
||||
_AnyIdAIMessageChunk(content=" "),
|
||||
_AnyIdAIMessageChunk(content="goodbye"),
|
||||
_AnyIdAIMessageChunk(content="goodbye", chunk_position="last"),
|
||||
]
|
||||
assert tokens == ["hello", " ", "goodbye"]
|
||||
@@ -118,7 +118,7 @@ def extract_deprecated_lookup(file_path: str) -> Optional[dict[str, Any]]:
|
||||
Returns:
|
||||
dict or None: The value of DEPRECATED_LOOKUP if it exists, None otherwise.
|
||||
"""
|
||||
tree = ast.parse(Path(file_path).read_text(), filename=file_path)
|
||||
tree = ast.parse(Path(file_path).read_text(encoding="utf-8"), filename=file_path)
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Assign):
|
||||
|
||||
Generated
+2489
-2928
File diff suppressed because it is too large.
Load diff
@@ -1,4 +1,8 @@
|
||||
"""Main entrypoint into LangChain."""
|
||||
"""Main entrypoint into LangChain.
|
||||
|
||||
AKA `version.py` in CI's `check_version_equality`.
|
||||
`langchain_v1 versions in pyproject.toml and __init__.py do not match!`
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
Loaded 100 of 188 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user