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

+5 -5
View File
@@ -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]
+1 -1
View File
@@ -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
)"
+4 -4
View File
@@ -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\"]"
+17 -17
View File
@@ -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)"
]
},
{
+4 -4
View File
@@ -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)"
]
},
{
+2 -2
View File
@@ -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"
]
},
+3 -3
View File
@@ -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=\"|\")"
]
},
{
+2 -2
View File
@@ -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=\"\")"
]
},
{
+16 -16
View File
@@ -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"
]
},
{
+9 -16
View File
@@ -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=\"|\")"
]
},
+1 -1
View File
@@ -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
+17 -18
View File
@@ -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" },
]
+1 -1
View File
@@ -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"
+60 -6
View File
@@ -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",
+246 -37
View File
@@ -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,
)
+191 -20
View File
@@ -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
+26 -16
View File
@@ -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):
+29 -9
View File
@@ -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):
+29 -20
View File
@@ -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,
+14 -5
View File
@@ -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 "
+1 -1
View File
@@ -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."""
+2 -2
View File
@@ -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"
+1 -2
View File
@@ -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]:
+2 -2
View File
@@ -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"}:
+11 -89
View File
@@ -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
+1 -13
View File
@@ -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(
+11 -2
View File
@@ -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.
+2 -1
View File
@@ -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
+1 -1
View File
@@ -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
+1 -7
View File
@@ -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(
{
+44 -6
View File
@@ -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
+1 -2
View File
@@ -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")
+1 -1
View File
@@ -18,7 +18,7 @@ from typing import (
)
if TYPE_CHECKING:
from typing_extensions import TypeAlias
from typing import TypeAlias
logger = logging.getLogger(__name__)
+29
View File
@@ -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 -1
View File
@@ -1,3 +1,3 @@
"""langchain-core version information and utilities."""
VERSION = "0.3.77"
VERSION = "1.0.0a6"
+2 -2
View File
@@ -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"),
[
+242 -101
View File
@@ -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)
+8 -8
View File
@@ -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 ""
+1090 -1194
View File
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 = (
+2 -2
View File
@@ -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.
+12 -13
View File
@@ -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):
+2489 -2928
View File
File diff suppressed because it is too large. Load diff
+5 -1
View File
@@ -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