mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-03 06:55:13 +02:00
Compare commits
70
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
52bd5b13a7 | ||
|
|
e9d62944d3 | ||
|
|
cbd09abe58 | ||
|
|
4798443e31 | ||
|
|
ce900864fa | ||
|
|
577f95bd50 | ||
|
|
59a11c63b0 | ||
|
|
08098688d4 | ||
|
|
687ee02509 | ||
|
|
451bc038b6 | ||
|
|
c865e8c070 | ||
|
|
d74ec2c2de | ||
|
|
f70bfc6d87 | ||
|
|
c86f0af107 | ||
|
|
c6ee807de5 | ||
|
|
7256752f48 | ||
|
|
dac84951aa | ||
|
|
3aaa3e38a0 | ||
|
|
400d83708a | ||
|
|
1d9c7ef461 | ||
|
|
01e5ecedfd | ||
|
|
76199701b0 | ||
|
|
2766fccb5b | ||
|
|
6aef3e0117 | ||
|
|
c5023ba147 | ||
|
|
9a9fe2fdec | ||
|
|
effddca494 | ||
|
|
2ab59840e7 | ||
|
|
0fb65f6e67 | ||
|
|
0ecd23eec6 | ||
|
|
e137dabf22 | ||
|
|
fdc1e47aa1 | ||
|
|
866780b477 | ||
|
|
18d3fa2e15 | ||
|
|
4b0c53fb5c | ||
|
|
5183484322 | ||
|
|
8213e4719b | ||
|
|
f993dfcfcb | ||
|
|
9e31b82d8d | ||
|
|
a3c5b8fc37 | ||
|
|
1e0aebc3ec | ||
|
|
056f581342 | ||
|
|
a0d7323bec | ||
|
|
44ee0199fd | ||
|
|
1de61f6fce | ||
|
|
f33db6cec4 | ||
|
|
c72107177b | ||
|
|
f520a38d30 | ||
|
|
d0278f520c | ||
|
|
506539ac9d | ||
|
|
d6c6516f16 | ||
|
|
6f5d6d9993 | ||
|
|
8dcd058404 | ||
|
|
e3050b3a3e | ||
|
|
9e767afad7 | ||
|
|
62b35277ec | ||
|
|
931419909b | ||
|
|
e849c869cc | ||
|
|
5a580ae5ec | ||
|
|
47a0e09513 | ||
|
|
f37486efe2 | ||
|
|
fa61be9fbc | ||
|
|
08097a78bd | ||
|
|
43f610e9a6 | ||
|
|
aa1ddee67e | ||
|
|
12b46e8a69 | ||
|
|
d90f69105a | ||
|
|
fbd3b67183 | ||
|
|
6d8be543e7 | ||
|
|
1c1772f7ec |
@@ -6,6 +6,9 @@ import re
|
||||
from typing import List, Literal, Optional
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
|
||||
from functools import lru_cache
|
||||
|
||||
import nbformat
|
||||
from nbconvert.preprocessors import Preprocessor
|
||||
|
||||
@@ -47,6 +50,8 @@ MANUAL_API_REFERENCES_LANGGRAPH = [
|
||||
(["langgraph.graph"], "langgraph.constants", "END", "constants"),
|
||||
(["langgraph.constants"], "langgraph.types", "Send", "types"),
|
||||
(["langgraph.constants"], "langgraph.types", "Interrupt", "types"),
|
||||
(["langgraph.constants"], "langgraph.types", "interrupt", "types"),
|
||||
(["langgraph.constants"], "langgraph.types", "Command", "types"),
|
||||
([], "langgraph.types", "RetryPolicy", "types"),
|
||||
([], "langgraph.checkpoint.base", "Checkpoint", "checkpoints"),
|
||||
([], "langgraph.checkpoint.base", "CheckpointMetadata", "checkpoints"),
|
||||
@@ -83,8 +88,11 @@ _IMPORT_LANGCHAIN_RE = _make_regular_expression("langchain")
|
||||
_IMPORT_LANGGRAPH_RE = _make_regular_expression("langgraph")
|
||||
|
||||
|
||||
def _get_full_module_name(module_path, class_name) -> Optional[str]:
|
||||
"""Get full module name using inspect"""
|
||||
|
||||
|
||||
@lru_cache(maxsize=10_000)
|
||||
def _get_full_module_name(module_path: str, class_name: str) -> Optional[str]:
|
||||
"""Get full module name using inspect, with LRU cache to memoize results."""
|
||||
try:
|
||||
module = importlib.import_module(module_path)
|
||||
class_ = getattr(module, class_name)
|
||||
@@ -95,13 +103,12 @@ def _get_full_module_name(module_path, class_name) -> Optional[str]:
|
||||
return module_path
|
||||
return module.__name__
|
||||
except AttributeError as e:
|
||||
logger.warning(f"Could not find module for {class_name}, {e}")
|
||||
logger.warning(f"API Reference: Could not find module for {class_name}, {e}")
|
||||
return None
|
||||
except ImportError as e:
|
||||
logger.warning(f"Failed to load for class {class_name}, {e}")
|
||||
logger.warning(f"API Reference: Failed to load for class {class_name}, {e}")
|
||||
return None
|
||||
|
||||
|
||||
def _get_doc_title(data: str, file_name: str) -> str:
|
||||
try:
|
||||
return re.findall(r"^#\s*(.*)", data, re.MULTILINE)[0]
|
||||
@@ -115,10 +122,10 @@ def _get_doc_title(data: str, file_name: str) -> str:
|
||||
|
||||
|
||||
class ImportInformation(TypedDict):
|
||||
imported: str # imported class name
|
||||
source: str # module path
|
||||
docs: str # URL to the documentation
|
||||
title: str # Title of the document
|
||||
imported: str # The name of the class that was imported.
|
||||
source: str # The full module path from which the class was imported.
|
||||
docs: str # The URL pointing to the class's documentation.
|
||||
title: str # The title of the document where the import is used.
|
||||
|
||||
|
||||
def _get_imports(
|
||||
@@ -211,36 +218,73 @@ def _get_imports(
|
||||
return imports
|
||||
|
||||
|
||||
class ImportPreprocessor(Preprocessor):
|
||||
"""A preprocessor to replace imports in each Python code cell with links to their
|
||||
documentation and append the import info in a comment."""
|
||||
def get_imports(code: str, doc_title: str) -> List[ImportInformation]:
|
||||
"""Retrieve all import references from the given code for specified ecosystems.
|
||||
|
||||
def preprocess(self, nb, resources):
|
||||
self.all_imports = []
|
||||
file_name = os.path.basename(resources.get("metadata", {}).get("name", ""))
|
||||
_DOC_TITLE = _get_doc_title(nb.cells[0].source, file_name)
|
||||
Args:
|
||||
code: The source code from which to extract import references.
|
||||
doc_title: The documentation title associated with the code.
|
||||
|
||||
cells = []
|
||||
for cell in nb.cells:
|
||||
if cell.cell_type == "code":
|
||||
cells.append(cell)
|
||||
imports = _get_imports(
|
||||
cell.source, _DOC_TITLE, "langchain"
|
||||
) + _get_imports(cell.source, _DOC_TITLE, "langgraph")
|
||||
if not imports:
|
||||
continue
|
||||
Returns:
|
||||
A list of import information for each import found.
|
||||
"""
|
||||
ecosystems = ["langchain", "langgraph"]
|
||||
all_imports = []
|
||||
for package_ecosystem in ecosystems:
|
||||
all_imports.extend(_get_imports(code, doc_title, package_ecosystem))
|
||||
return all_imports
|
||||
|
||||
cells.append(
|
||||
nbformat.v4.new_markdown_cell(
|
||||
source=f"""
|
||||
<div>
|
||||
<b>API Reference:</b>
|
||||
{' | '.join(f'<a href="{imp["docs"]}">{imp["imported"]}</a>' for imp in imports)}
|
||||
</div>
|
||||
"""
|
||||
)
|
||||
)
|
||||
else:
|
||||
cells.append(cell)
|
||||
nb.cells = cells
|
||||
return nb, resources
|
||||
|
||||
def update_markdown_with_imports(markdown: str) -> str:
|
||||
"""Update markdown to include API reference links for imports in Python code blocks.
|
||||
|
||||
This function scans the markdown content for Python code blocks, extracts any imports, and appends links to their API documentation.
|
||||
|
||||
Args:
|
||||
markdown: The markdown content to process.
|
||||
|
||||
Returns:
|
||||
Updated markdown with API reference links appended to Python code blocks.
|
||||
|
||||
Example:
|
||||
Given a markdown with a Python code block:
|
||||
|
||||
```python
|
||||
from langchain.nlp import TextGenerator
|
||||
```
|
||||
This function will append an API reference link to the `TextGenerator` class from the `langchain.nlp` module if it's recognized.
|
||||
"""
|
||||
code_block_pattern = re.compile(
|
||||
r'(?P<indent>[ \t]*)```(?P<language>python|py)\n(?P<code>.*?)\n(?P=indent)```', re.DOTALL
|
||||
)
|
||||
|
||||
def replace_code_block(match: re.Match) -> str:
|
||||
"""Replace the matched code block with additional API reference links if imports are found.
|
||||
|
||||
Args:
|
||||
match (re.Match): The regex match object containing the code block.
|
||||
|
||||
Returns:
|
||||
str: The modified code block with API reference links appended if applicable.
|
||||
"""
|
||||
indent = match.group('indent')
|
||||
code_block = match.group('code')
|
||||
language = match.group('language') # Preserve the language from the regex match
|
||||
# Retrieve import information from the code block
|
||||
imports = get_imports(code_block, "__unused__")
|
||||
|
||||
original_code_block = match.group(0)
|
||||
# If no imports are found, return the original code block
|
||||
if not imports:
|
||||
return original_code_block
|
||||
|
||||
# Generate API reference links for each import
|
||||
api_links = ' | '.join(
|
||||
f'<a href="{imp["docs"]}">{imp["imported"]}</a>' for imp in imports
|
||||
)
|
||||
# Return the code block with appended API reference links
|
||||
return f'{original_code_block}\n\n{indent}API Reference: {api_links}'
|
||||
|
||||
# Apply the replace_code_block function to all matches in the markdown
|
||||
updated_markdown = code_block_pattern.sub(replace_code_block, markdown)
|
||||
return updated_markdown
|
||||
@@ -6,8 +6,6 @@ import nbformat
|
||||
from nbconvert.exporters import MarkdownExporter
|
||||
from nbconvert.preprocessors import Preprocessor
|
||||
|
||||
from generate_api_reference_links import ImportPreprocessor
|
||||
|
||||
|
||||
class EscapePreprocessor(Preprocessor):
|
||||
def preprocess_cell(self, cell, resources, cell_index):
|
||||
@@ -107,7 +105,6 @@ exporter = MarkdownExporter(
|
||||
preprocessors=[
|
||||
EscapePreprocessor,
|
||||
ExtractAttachmentsPreprocessor,
|
||||
ImportPreprocessor,
|
||||
],
|
||||
template_name="mdoutput",
|
||||
extra_template_basedirs=[
|
||||
|
||||
@@ -1,10 +1,13 @@
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Dict
|
||||
|
||||
from mkdocs.structure.pages import Page
|
||||
from mkdocs.structure.files import Files, File
|
||||
from mkdocs.structure.pages import Page
|
||||
|
||||
from notebook_convert import convert_notebook
|
||||
from generate_api_reference_links import update_markdown_with_imports
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logging.basicConfig()
|
||||
@@ -35,12 +38,83 @@ def on_files(files: Files, **kwargs: Dict[str, Any]):
|
||||
return new_files
|
||||
|
||||
|
||||
def _highlight_code_blocks(markdown: str) -> str:
|
||||
"""Find code blocks with highlight comments and add hl_lines attribute.
|
||||
|
||||
Args:
|
||||
markdown: The markdown content to process.
|
||||
|
||||
Returns:
|
||||
updated Markdown code with code blocks containing highlight comments
|
||||
updated to use the hl_lines attribute.
|
||||
"""
|
||||
# Pattern to find code blocks with highlight comments and without
|
||||
# existing hl_lines for Python and JavaScript
|
||||
# Pattern to find code blocks with highlight comments, handling optional indentation
|
||||
code_block_pattern = re.compile(
|
||||
r"(?P<indent>[ \t]*)```(?P<language>py|python|js|javascript)(?!\s+hl_lines=)\n"
|
||||
r"(?P<code>((?:.*\n)*?))" # Capture the code inside the block using named group
|
||||
r"(?P=indent)```" # Match closing backticks with the same indentation
|
||||
)
|
||||
|
||||
def replace_highlight_comments(match: re.Match) -> str:
|
||||
indent = match.group("indent")
|
||||
language = match.group("language")
|
||||
code_block = match.group("code")
|
||||
lines = code_block.split("\n")
|
||||
highlighted_lines = []
|
||||
|
||||
# Skip initial empty lines
|
||||
while lines and not lines[0].strip():
|
||||
lines.pop(0)
|
||||
|
||||
lines_to_keep = []
|
||||
|
||||
comment_syntax = (
|
||||
"# highlight-next-line"
|
||||
if language in ["py", "python"]
|
||||
else "// highlight-next-line"
|
||||
)
|
||||
|
||||
for line in lines:
|
||||
if comment_syntax in line:
|
||||
count = len(lines_to_keep) + 1
|
||||
highlighted_lines.append(str(count))
|
||||
else:
|
||||
lines_to_keep.append(line)
|
||||
|
||||
# Reconstruct the new code block
|
||||
new_code_block = "\n".join(lines_to_keep)
|
||||
|
||||
if highlighted_lines:
|
||||
return (
|
||||
f'{indent}```{language} hl_lines="{" ".join(highlighted_lines)}"\n'
|
||||
# The indent and terminating \n is already included in the code block
|
||||
f'{new_code_block}'
|
||||
f'{indent}```'
|
||||
)
|
||||
else:
|
||||
return (
|
||||
f"{indent}```{language}\n"
|
||||
# The indent and terminating \n is already included in the code block
|
||||
f"{new_code_block}"
|
||||
f"{indent}```"
|
||||
)
|
||||
|
||||
# Replace all code blocks in the markdown
|
||||
markdown = code_block_pattern.sub(replace_highlight_comments, markdown)
|
||||
return markdown
|
||||
|
||||
|
||||
def on_page_markdown(markdown: str, page: Page, **kwargs: Dict[str, Any]):
|
||||
if DISABLED:
|
||||
return markdown
|
||||
if page.file.src_path.endswith(".ipynb"):
|
||||
logger.info("Processing Jupyter notebook: %s", page.file.src_path)
|
||||
body = convert_notebook(page.file.abs_src_path)
|
||||
return body
|
||||
markdown = convert_notebook(page.file.abs_src_path)
|
||||
|
||||
# Append API reference links to code blocks
|
||||
markdown = update_markdown_with_imports(markdown)
|
||||
# Apply highlight comments to code blocks
|
||||
markdown = _highlight_code_blocks(markdown)
|
||||
return markdown
|
||||
|
||||
@@ -5,7 +5,7 @@ LangGraph Cloud is available within <a href="https://www.langchain.com/langsmith
|
||||
## Prerequisites
|
||||
|
||||
1. LangGraph Cloud applications are deployed from GitHub repositories. Configure and upload a LangGraph Cloud application to a GitHub repository in order to deploy it to LangGraph Cloud.
|
||||
1. [Verify that the LangGraph API runs locally](test_locally.md). If the API does not build and run successfully (i.e. `langgraph up`), deploying to LangGraph Cloud will fail as well.
|
||||
1. [Verify that the LangGraph API runs locally](test_locally.md). If the API does not run successfully (i.e. `langgraph dev`), deploying to LangGraph Cloud will fail as well.
|
||||
|
||||
## Create New Deployment
|
||||
|
||||
|
||||
@@ -6,17 +6,11 @@ Testing locally ensures that there are no errors or conflicts with Python depend
|
||||
|
||||
## Setup
|
||||
|
||||
Install the proper packages:
|
||||
Install the LangGraph CLI package:
|
||||
|
||||
|
||||
=== "pip"
|
||||
```bash
|
||||
pip install -U langgraph-cli
|
||||
```
|
||||
=== "Homebrew (macOS only)"
|
||||
```bash
|
||||
brew install langgraph-cli
|
||||
```
|
||||
```bash
|
||||
pip install -U "langgraph-cli[inmem]"
|
||||
```
|
||||
|
||||
Ensure you have an API key, which you can create from the [LangSmith UI](https://smith.langchain.com) (Settings > API Keys). This is required to authenticate that you have LangGraph Cloud access. After you have saved the key to a safe place, place the following line in your `.env` file:
|
||||
|
||||
@@ -29,16 +23,26 @@ LANGSMITH_API_KEY = *********
|
||||
Once you have installed the CLI, you can run the following command to start the API server for local testing:
|
||||
|
||||
```shell
|
||||
langgraph up
|
||||
langgraph dev
|
||||
```
|
||||
|
||||
This will start up the LangGraph API server locally. If this runs successfully, you should see something like:
|
||||
|
||||
```shell
|
||||
Ready!
|
||||
- API: http://localhost:8123
|
||||
2024-06-26 19:20:41,056:INFO:uvicorn.access 127.0.0.1:44138 - "GET /ok HTTP/1.1" 200
|
||||
```
|
||||
> Ready!
|
||||
>
|
||||
> - API: [http://localhost:2024](http://localhost:2024/)
|
||||
>
|
||||
> - Docs: http://localhost:2024/docs
|
||||
>
|
||||
> - LangGraph Studio Web UI: https://smith.langchain.com/studio/?baseUrl=http://127.0.0.1:2024
|
||||
|
||||
!!! note "In-Memory Mode"
|
||||
|
||||
The `langgraph dev` command starts LangGraph Server in an in-memory mode. This mode is suitable for development and testing purposes. For production use, you should deploy LangGraph Server with access to a persistent storage backend.
|
||||
|
||||
If you want to test your application with a persistent storage backend, you can use the `langgraph up` command instead of `langgraph dev`. You will
|
||||
need to have `docker` installed on your machine to use this command.
|
||||
|
||||
|
||||
### Interact with the server
|
||||
|
||||
@@ -53,7 +57,7 @@ You can either initialize by passing authentication or by setting an environment
|
||||
```python
|
||||
from langgraph_sdk import get_client
|
||||
|
||||
# only pass the url argument to get_client() if you changed the default port when calling langgraph up
|
||||
# only pass the url argument to get_client() if you changed the default port when calling langgraph dev
|
||||
client = get_client(url=<DEPLOYMENT_URL>,api_key=<LANGSMITH_API_KEY>)
|
||||
# Using the graph deployed with the name "agent"
|
||||
assistant_id = "agent"
|
||||
@@ -65,7 +69,7 @@ You can either initialize by passing authentication or by setting an environment
|
||||
```js
|
||||
import { Client } from "@langchain/langgraph-sdk";
|
||||
|
||||
// only set the apiUrl if you changed the default port when calling langgraph up
|
||||
// only set the apiUrl if you changed the default port when calling langgraph dev
|
||||
const client = new Client({ apiUrl: <DEPLOYMENT_URL>, apiKey: <LANGSMITH_API_KEY> });
|
||||
// Using the graph deployed with the name "agent"
|
||||
const assistantId = "agent";
|
||||
@@ -91,7 +95,7 @@ If you have a `LANGSMITH_API_KEY` set in your environment, you do not need to ex
|
||||
```python
|
||||
from langgraph_sdk import get_client
|
||||
|
||||
# only pass the url argument to get_client() if you changed the default port when calling langgraph up
|
||||
# only pass the url argument to get_client() if you changed the default port when calling langgraph dev
|
||||
client = get_client()
|
||||
# Using the graph deployed with the name "agent"
|
||||
assistant_id = "agent"
|
||||
@@ -103,7 +107,7 @@ If you have a `LANGSMITH_API_KEY` set in your environment, you do not need to ex
|
||||
```js
|
||||
import { Client } from "@langchain/langgraph-sdk";
|
||||
|
||||
// only set the apiUrl if you changed the default port when calling langgraph up
|
||||
// only set the apiUrl if you changed the default port when calling langgraph dev
|
||||
const client = new Client();
|
||||
// Using the graph deployed with the name "agent"
|
||||
const assistantId = "agent";
|
||||
|
||||
@@ -7,17 +7,21 @@
|
||||
|
||||
Make sure you have setup your app correctly, by creating a compiled graph, a `.env` file with any environment variables, and a `langgraph.json` config file that points to your environment file and compiled graph. See [here](https://langchain-ai.github.io/langgraph/cloud/deployment/setup/) for more detailed instructions.
|
||||
|
||||
After you have your app setup, head into the directory with your `langgraph.json` file and call `langgraph up -c langgraph.json --watch` to start the API server in watch mode which means it will restart on code changes, which is ideal for local testing. If the API server start correctly you should see logs that look something like this:
|
||||
After you have your app setup, head into the directory with your `langgraph.json` file and call `langgraph dev` to start the API server in watch mode which means it will restart on code changes, which is ideal for local testing. If the API server start correctly you should see logs that look something like this:
|
||||
|
||||
Ready!
|
||||
- API: http://localhost:8123
|
||||
2024-06-26 19:20:41,056:INFO:uvicorn.access 127.0.0.1:44138 - "GET /ok HTTP/1.1" 200
|
||||
> Ready!
|
||||
>
|
||||
> - API: [http://localhost:2024](http://localhost:2024/)
|
||||
>
|
||||
> - Docs: http://localhost:2024/docs
|
||||
>
|
||||
> - LangGraph Studio Web UI: https://smith.langchain.com/studio/?baseUrl=http://127.0.0.1:2024
|
||||
|
||||
Read this [reference](https://langchain-ai.github.io/langgraph/cloud/reference/cli/#up) to learn about all the options for starting the API server.
|
||||
|
||||
## Access Studio
|
||||
|
||||
Once you have successfully started the API server, you can access the studio by going to the following URL: `https://smith.langchain.com/studio/?baseUrl=http://127.0.0.1:8123` (see warning above if using Safari).
|
||||
Once you have successfully started the API server, you can access the studio by going to the following URL: `https://smith.langchain.com/studio/?baseUrl=http://127.0.0.1:2024` (see warning above if using Safari).
|
||||
|
||||
If everything is working correctly you should see the studio show up looking something like this (with your graph diagram on the left hand side):
|
||||
|
||||
|
||||
@@ -208,7 +208,6 @@ export LANGSMITH_API_KEY=...
|
||||
```js
|
||||
const { Client } = await import("@langchain/langgraph-sdk");
|
||||
|
||||
// only set the apiUrl if you changed the default port when calling langgraph up
|
||||
const client = new Client({ apiUrl: "your-deployment-url", apiKey: "your-langsmith-api-key" });
|
||||
|
||||
const streamResponse = client.runs.stream(
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
<!doctype html>
|
||||
<html>
|
||||
<head>
|
||||
<title>LangGraph Cloud API Reference</title>
|
||||
<meta charset="utf-8" />
|
||||
<meta
|
||||
name="viewport"
|
||||
content="width=device-width, initial-scale=1" />
|
||||
</head>
|
||||
<body>
|
||||
<script id="api-reference" data-url="./openapi_control_plane.json"></script>
|
||||
<script>
|
||||
var configuration = {}
|
||||
document.getElementById('api-reference').dataset.configuration =
|
||||
JSON.stringify(configuration)
|
||||
</script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@scalar/api-reference"></script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,695 @@
|
||||
{
|
||||
"openapi": "3.1.0",
|
||||
"info": {
|
||||
"title": "LangGraph Control Plane API (Beta)",
|
||||
"version": "0.0.1",
|
||||
"description": "The LangGraph Control Plane API is used to programmatically create and manage LangGraph Server deployments. For example, the APIs can be orchestrated to create custom CI/CD workflows.\n\n### Beta\nThis API is currently in beta and may change or break without notice. This API documentation may not be up-to-date with actual API functionality.\n### Host\nhttps://api.host.langchain.com/\n\n### Authentication\nTo authenticate with the LangGraph Control Plane API, set the `X-Api-Key` header to a valid LangSmith API key for each request.\n\n### Versioning\nEach endpoint path is prefixed with a version (e.g. `v1`).\n\n### Quick Start\n\n1. Call `GET /{version}/projects` to retrieve the `Project` `id`. The `Project` `id` is needed in subsequent API calls.\n2. Call `POST /{version}/projects/{project_id}/revisions` to create a new `Revision` for the `Project`.\n3. Call `GET /{version}/projects/{project_id}/revisions` to get the latest `Revision` (first element in returned list). Get the `Revision` `id`.\n4. Poll for `Revision` `status` until `status` is `DEPLOYED` by calling `GET /{version}/projects/{project_id}/revisions/{revision_id}`."
|
||||
},
|
||||
"servers": [
|
||||
{
|
||||
"url": "https://api.host.langchain.com"
|
||||
}
|
||||
],
|
||||
"tags": [
|
||||
{
|
||||
"name": "Projects (v1)",
|
||||
"description": "A project corresponds to a LangGraph Server deployment and the associated LangSmith tracing project.\n\nCreating a project via API is not currently supported/documented."
|
||||
},
|
||||
{
|
||||
"name": "Revisions (v1)",
|
||||
"description": "A revision is a version of a LangGraph Server deployment. Different revisions may contain different code and/or environment variables. A project can have many revisions."
|
||||
}
|
||||
],
|
||||
"paths": {
|
||||
"/v1/projects": {
|
||||
"get": {
|
||||
"tags": ["Projects (v1)"],
|
||||
"summary": "List Projects",
|
||||
"description": "List all projects.",
|
||||
"operationId": "list_projects_projects_get",
|
||||
"parameters": [
|
||||
{
|
||||
"required": false,
|
||||
"schema": {
|
||||
"type": "integer",
|
||||
"title": "Limit",
|
||||
"description": "Maximum number of results to return. Minimum: 1. Maximum: 100.",
|
||||
"default": 20
|
||||
},
|
||||
"name": "limit",
|
||||
"in": "query"
|
||||
},
|
||||
{
|
||||
"required": false,
|
||||
"schema": {
|
||||
"type": "integer",
|
||||
"title": "Offset",
|
||||
"description": "Pagination offset value. Pass this value in subsequent requests to retrieve the next page of results. Minimum: 0.",
|
||||
"default": 0
|
||||
},
|
||||
"name": "offset",
|
||||
"in": "query"
|
||||
},
|
||||
{
|
||||
"required": false,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"title": "Name Contains",
|
||||
"description": "Filter string to filter projects by `name`."
|
||||
},
|
||||
"name": "name_contains",
|
||||
"in": "query"
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/components/schemas/Project"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/projects/{project_id}": {
|
||||
"get": {
|
||||
"tags": ["Projects (v1)"],
|
||||
"summary": "Get Project",
|
||||
"description": "Get project by ID.",
|
||||
"operationId": "get_project_projects__project_id__get",
|
||||
"parameters": [
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Project ID"
|
||||
},
|
||||
"name": "project_id",
|
||||
"in": "path"
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/Project"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"delete": {
|
||||
"tags": ["Projects (v1)"],
|
||||
"summary": "Delete Project",
|
||||
"description": "Delete project by ID.",
|
||||
"operationId": "delete_project_projects__project_id__delete",
|
||||
"parameters": [
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Project ID"
|
||||
},
|
||||
"name": "project_id",
|
||||
"in": "path"
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/Project"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/projects/{project_id}/revisions": {
|
||||
"get": {
|
||||
"tags": ["Revisions (v1)"],
|
||||
"summary": "List Revisions",
|
||||
"description": "List revisions of a project.",
|
||||
"operationId": "list_revisions_projects__project_id__revisions_get",
|
||||
"parameters": [
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Project ID"
|
||||
},
|
||||
"name": "project_id",
|
||||
"in": "path"
|
||||
},
|
||||
{
|
||||
"required": false,
|
||||
"schema": {
|
||||
"type": "integer",
|
||||
"title": "Limit",
|
||||
"description": "Maximum number of results to return. Minimum: 1. Maximum: 100.",
|
||||
"default": 20
|
||||
},
|
||||
"name": "limit",
|
||||
"in": "query"
|
||||
},
|
||||
{
|
||||
"required": false,
|
||||
"schema": {
|
||||
"type": "integer",
|
||||
"title": "Offset",
|
||||
"description": "Pagination offset value. Pass this value in subsequent requests to retrieve the next page of results. Minimum: 0.",
|
||||
"default": 0
|
||||
},
|
||||
"name": "offset",
|
||||
"in": "query"
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/components/schemas/Revision"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"post": {
|
||||
"tags": ["Revisions (v1)"],
|
||||
"summary": "Create Revision",
|
||||
"description": "Create a new revision for a project.",
|
||||
"operationId": "create_revision_projects__project_id__revisions_post",
|
||||
"parameters": [
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Project ID"
|
||||
},
|
||||
"name": "project_id",
|
||||
"in": "path"
|
||||
}
|
||||
],
|
||||
"requestBody": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/CreateRevisionRequest"
|
||||
}
|
||||
}
|
||||
},
|
||||
"required": true
|
||||
},
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/Project"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/projects/{project_id}/revisions/{revision_id}": {
|
||||
"get": {
|
||||
"tags": ["Revisions (v1)"],
|
||||
"summary": "Get Revision",
|
||||
"description": "Get revision by ID.",
|
||||
"operationId": "get_revision_projects__project_id__revisions__revision_id__get",
|
||||
"parameters": [
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Project ID"
|
||||
},
|
||||
"name": "project_id",
|
||||
"in": "path"
|
||||
},
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Revision ID"
|
||||
},
|
||||
"name": "revision_id",
|
||||
"in": "path"
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Success",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/Revision"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/projects/{project_id}/revisions/{revision_id}/interrupt": {
|
||||
"post": {
|
||||
"tags": ["Revisions (v1)"],
|
||||
"summary": "Interrupt Revision",
|
||||
"description": "Interrupt revision by ID.\n\nIf the deployment of a revision appears \"stuck\", the revision may need to be interrupted. A new revision cannot be created if the latest revision is in a non-terminal `status`. In this scenario, the revision may need to be interrupted.",
|
||||
"operationId": "interrupt_revision_projects__project_id__revisions__revision_id__interrupt_post",
|
||||
"parameters": [
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Project ID"
|
||||
},
|
||||
"name": "project_id",
|
||||
"in": "path"
|
||||
},
|
||||
{
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"title": "Revision ID"
|
||||
},
|
||||
"name": "revision_id",
|
||||
"in": "path"
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
},
|
||||
"components": {
|
||||
"securitySchemes": {
|
||||
"apiKeyAuth": {
|
||||
"type": "apiKey",
|
||||
"in": "header",
|
||||
"name": "X-Api-Key"
|
||||
}
|
||||
},
|
||||
"schemas": {
|
||||
"EnvVar": {
|
||||
"type": "object",
|
||||
"description": "An environment variable or secret.",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Environment variable or secret name.",
|
||||
"required": true
|
||||
},
|
||||
"value": {
|
||||
"type": "string",
|
||||
"description": "Environment variable or secret value.",
|
||||
"required": true
|
||||
},
|
||||
"type": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
"default",
|
||||
"secret"
|
||||
],
|
||||
"description": "Field to designate type of the environment variable (default) or secret.",
|
||||
"required": true
|
||||
}
|
||||
}
|
||||
},
|
||||
"ContainerSpec": {
|
||||
"type": "object",
|
||||
"description": "Container specification for a revision's deployment.\n\nIf any field is omitted or set to `null`, the internal default value is used depending on the deployment type (`dev` or `prod`).",
|
||||
"properties": {
|
||||
"min_scale": {
|
||||
"type": ["integer", "null"],
|
||||
"description": "Minimum number of replicas in deployment.",
|
||||
"default": "null"
|
||||
},
|
||||
"max_scale": {
|
||||
"type": ["integer", "null"],
|
||||
"description": "Maximum number of replicas in deployment.",
|
||||
"default": "null"
|
||||
},
|
||||
"cpu": {
|
||||
"type": ["integer", "null"],
|
||||
"description": "Number of vCPU cores per replica.",
|
||||
"default": "null"
|
||||
},
|
||||
"memory_mb": {
|
||||
"type": ["integer", "null"],
|
||||
"description": "Amount of memory in MB per replica.",
|
||||
"default": "null"
|
||||
}
|
||||
}
|
||||
},
|
||||
"CreateRevisionRequest": {
|
||||
"type": "object",
|
||||
"description": "Object for creating a new revision.",
|
||||
"properties": {
|
||||
"image_path": {
|
||||
"type": ["string", "null"],
|
||||
"description": "URI of the Docker image to deploy.\n\nIf this field is omitted or set to `null`, the previous revision's `image_path` value is used. Set this field for BYOC deployments. Omit this field if creating a new revision from a GitHub repository.",
|
||||
"default": "null"
|
||||
},
|
||||
"repo_path": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Path to `langgraph.json` configuration file. For example, `langgraph.json` or `src/langgraph.json`.\n\nIf this field is omitted or set to `null`, the previous revision's `repo_path` value is used. Set this field for deployments from a GitHub repository. Omit this field if creating a new revision from a Docker image.",
|
||||
"default": "null"
|
||||
},
|
||||
"env_vars": {
|
||||
"type": "array",
|
||||
"description": "List of environment variables or secrets.\n\nIf this field is omitted or set to `null`, the previous revision's `env_vars` value is used.",
|
||||
"items": {
|
||||
"$ref": "#/components/schemas/EnvVar"
|
||||
},
|
||||
"default": "null"
|
||||
},
|
||||
"shareable": {
|
||||
"type": ["boolean", "null"],
|
||||
"description": "Boolean flag to configure if a deployment is shareable through LangGraph Studio.\n\nIf this field is omitted or set to `null`, the previous revision's `shareable` value is used. This field does not apply to BYOC deployments.",
|
||||
"default": "null"
|
||||
},
|
||||
"container_spec": {
|
||||
"description": "If this field is omitted or set to `null`, the previous revision's `container_spec` value is used.",
|
||||
"$ref": "#/components/schemas/ContainerSpec",
|
||||
"default": "null"
|
||||
}
|
||||
}
|
||||
},
|
||||
"Project": {
|
||||
"type": "object",
|
||||
"description": "A project corresponds to a LangGraph Server deployment and the associated LangSmith tracing project.",
|
||||
"properties": {
|
||||
"id": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"description": "ID of the project.",
|
||||
"required": true
|
||||
},
|
||||
"tool_name": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Do not use."
|
||||
},
|
||||
"display_name": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Do not use."
|
||||
},
|
||||
"description": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Do not use."
|
||||
},
|
||||
"example_input": {
|
||||
"type": ["object", "null"],
|
||||
"description": "Do not use."
|
||||
},
|
||||
"tenant_id": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"description": "ID of the tenant/workspace of the project.",
|
||||
"required": true
|
||||
},
|
||||
"created_at": {
|
||||
"type": "string",
|
||||
"format": "date-time",
|
||||
"description": "Timestamp of when the project was created.",
|
||||
"required": true
|
||||
},
|
||||
"updated_at": {
|
||||
"type": "string",
|
||||
"format": "date-time",
|
||||
"description": "Timestamp of when the project was updated.",
|
||||
"required": true
|
||||
},
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Name of the project.\n\nThis is also the name of the LangSmith tracing project for the LangGraph deployment.",
|
||||
"required": true
|
||||
},
|
||||
"lc_hosted": {
|
||||
"type": "boolean",
|
||||
"description": "Boolean flag to indicate if the deployment is hosted in LangChain's cloud or an external cloud (e.g. BYOC).",
|
||||
"required": true
|
||||
},
|
||||
"repo_url": {
|
||||
"type": ["string", "null"],
|
||||
"description": "URL of the GitHub repository.\n\nThis field is not used for deployments from a Docker image."
|
||||
},
|
||||
"repo_branch": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Branch of the GitHub repository.\n\nThis field is not used for deployments from a Docker image."
|
||||
},
|
||||
"tracer_session_id": {
|
||||
"type": ["string", "null"],
|
||||
"format": "uuid",
|
||||
"description": "Do not use."
|
||||
},
|
||||
"api_key_id": {
|
||||
"type": ["string", "null"],
|
||||
"format": "uuid",
|
||||
"description": "Do not use."
|
||||
},
|
||||
"build_on_push": {
|
||||
"type": "boolean",
|
||||
"description": "Boolean flag to indicate if a new revision is automatically created on push to GitHub branch (`repo_branch`).\n\nThis field does not apply for BYOC deployments."
|
||||
},
|
||||
"input_json_schemas": {
|
||||
"type": ["object", "null"],
|
||||
"description": "Do not use."
|
||||
},
|
||||
"output_json_schemas": {
|
||||
"type": ["object", "null"],
|
||||
"description": "Do not use."
|
||||
},
|
||||
"host_integration_id": {
|
||||
"type": ["string", "null"],
|
||||
"format": "uuid",
|
||||
"description": "Do not use."
|
||||
},
|
||||
"metadata": {
|
||||
"$ref": "#/components/schemas/ProjectMetadata"
|
||||
},
|
||||
"resource": {
|
||||
"$ref": "#/components/schemas/ResourceService"
|
||||
}
|
||||
}
|
||||
},
|
||||
"ProjectMetadata": {
|
||||
"type": "object",
|
||||
"description": "Metadata associated with a `Project`.",
|
||||
"properties": {
|
||||
"deployment_type": {
|
||||
"type": "string",
|
||||
"description": "Development (`dev`) or Production (`prod`) type deployment.",
|
||||
"enum": [
|
||||
"dev",
|
||||
"prod"
|
||||
]
|
||||
},
|
||||
"image_source": {
|
||||
"type": "string",
|
||||
"description": "Do not use.",
|
||||
"enum": [
|
||||
"github",
|
||||
"internal_docker",
|
||||
"external_docker"
|
||||
]
|
||||
},
|
||||
"shareable": {
|
||||
"type": "boolean",
|
||||
"description": "Boolean flag to configure if a deployment is shareable through LangGraph Studio.\n\nThis field does not apply to BYOC deployments."
|
||||
},
|
||||
"region": {
|
||||
"type": "string",
|
||||
"description": "Region of deployment.\n\nRegion value is cloud provider specific."
|
||||
},
|
||||
"aws_account_id": {
|
||||
"type": "string",
|
||||
"description": "AWS account ID of BYOC deployment.\n\nThis field does not apply to non-BYOC deployments."
|
||||
},
|
||||
"aws_external_id": {
|
||||
"type": "string",
|
||||
"description": "Do not use."
|
||||
}
|
||||
}
|
||||
},
|
||||
"ResourceId": {
|
||||
"type": "object",
|
||||
"description": "Internal identifier for a `ResourceRevision` or `ResourceService`.",
|
||||
"properties": {
|
||||
"type": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
"revisions",
|
||||
"services"
|
||||
]
|
||||
},
|
||||
"name": {
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
},
|
||||
"ResourceRevision": {
|
||||
"type": "object",
|
||||
"description": "Internal revision resource for a `ResourceService`.",
|
||||
"properties": {
|
||||
"id": {
|
||||
"$ref": "#/components/schemas/ResourceId"
|
||||
},
|
||||
"env_vars": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/components/schemas/EnvVar"
|
||||
}
|
||||
},
|
||||
"hosted_langserve_revision_id": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"description": "References `id` of a `Revision`."
|
||||
}
|
||||
}
|
||||
},
|
||||
"ResourceService": {
|
||||
"type": "object",
|
||||
"description": "Internal service resource for a `Project`.",
|
||||
"properties": {
|
||||
"id": {
|
||||
"$ref": "#/components/schemas/ResourceId"
|
||||
},
|
||||
"url": {
|
||||
"type": ["string", "null"],
|
||||
"description": "URL of LangGraph Server deployment."
|
||||
},
|
||||
"latest_revision": {
|
||||
"description": "References latest `ResourceRevision`.\n\nThe latest `ResourceRevision` may not be active if it's currently being deployed.",
|
||||
"$ref": "#/components/schemas/ResourceRevision"
|
||||
},
|
||||
"latest_active_revision": {
|
||||
"description": "References latest active `ResourceRevision`.\n\nThe latest active `ResourceRevision` is not always the latest `ResourceRevision`.",
|
||||
"$ref": "#/components/schemas/ResourceRevision"
|
||||
}
|
||||
}
|
||||
},
|
||||
"Revision": {
|
||||
"type": "object",
|
||||
"description": "A revision is a version of a LangGraph Server deployment.\n\nDifferent revisions may contain different code and/or environment variables. A project can have many revisions.",
|
||||
"properties": {
|
||||
"id": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"description": "ID of the revision.",
|
||||
"required": true
|
||||
},
|
||||
"project_id": {
|
||||
"type": "string",
|
||||
"format": "uuid",
|
||||
"description": "References `id` of `Project`.",
|
||||
"required": true
|
||||
},
|
||||
"created_at": {
|
||||
"type": "string",
|
||||
"format": "date-time",
|
||||
"description": "Timestamp of when the revision was created.",
|
||||
"required": true
|
||||
},
|
||||
"updated_at": {
|
||||
"type": "string",
|
||||
"format": "date-time",
|
||||
"description": "Timestamp of when the revision was updated.",
|
||||
"required": true
|
||||
},
|
||||
"repo_path": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Path to `langgraph.json` configuration file. For example, `langgraph.json` or `src/langgraph.json`.\n\nThis field only applies to deployments from a GitHub repository.",
|
||||
"default": "null"
|
||||
},
|
||||
"repo_commit": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Git branch name of deployment.\n\nThis field only applies to deployments from a GitHub repository.",
|
||||
"default": "null"
|
||||
},
|
||||
"status": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
"CREATING",
|
||||
"AWAITING_BUILD",
|
||||
"BUILDING",
|
||||
"AWAITING_DEPLOY",
|
||||
"DEPLOYING",
|
||||
"CREATE_FAILED",
|
||||
"BUILD_FAILED",
|
||||
"DEPLOY_FAILED",
|
||||
"DEPLOYED",
|
||||
"INTERRUPTED",
|
||||
"UNKNOWN"
|
||||
],
|
||||
"description": "Deployment status of the revision.\n\nNon-terminal statuses: `CREATING`, `AWAITING_BUILD`, `BUILDING`, `AWAITING_DEPLOY`, `DEPLOYING`. All other statuses are terminal."
|
||||
},
|
||||
"status_message": {
|
||||
"type": "string",
|
||||
"description": "Message associated with the `status`."
|
||||
},
|
||||
"gcp_build_name": {
|
||||
"type": ["string", "null"],
|
||||
"description": "Do not use."
|
||||
},
|
||||
"metadata": {
|
||||
"$ref": "#/components/schemas/RevisionMetadata"
|
||||
},
|
||||
"image_path": {
|
||||
"type": ["string", "null"],
|
||||
"description": "URI of the Docker image to deploy.\n\nThis field does not apply to deployments from a GitHub repository.",
|
||||
"default": "null"
|
||||
},
|
||||
"container_spec": {
|
||||
"$ref": "#/components/schemas/ContainerSpec"
|
||||
},
|
||||
"resource": {
|
||||
"$ref": "#/components/schemas/ResourceRevision"
|
||||
}
|
||||
}
|
||||
},
|
||||
"RevisionMetadata": {
|
||||
"type": "object",
|
||||
"description": "Metadata associated with a `Revision`.",
|
||||
"properties": {
|
||||
"created_by": {
|
||||
"type": "object",
|
||||
"description": "Do not use."
|
||||
},
|
||||
"repo_commit_sha": {
|
||||
"type": "string",
|
||||
"description": "Git commit SHA of the deployment.\n\nThis field only applies to deployments from a GitHub repository."
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -61,6 +61,7 @@ The LangGraph CLI requires a JSON configuration file with the following keys:
|
||||
All deployments come with a DB-backed BaseStore. Adding an "index" configuration to your `langgraph.json` will enable [semantic search](../deployment/semantic_search.md) within the BaseStore of your deployment.
|
||||
|
||||
The `fields` configuration determines which parts of your documents to embed:
|
||||
|
||||
- If omitted or set to `["$"]`, the entire document will be embedded
|
||||
- To embed specific fields, use JSON path notation: `["metadata.title", "content.text"]`
|
||||
- Documents missing specified fields will still be stored but won't have embeddings for those fields
|
||||
|
||||
@@ -1,26 +1,26 @@
|
||||
# Agent architectures
|
||||
|
||||
Many LLM applications implement a particular control flow of steps before and / or after LLM calls. As an example, [RAG](https://github.com/langchain-ai/rag-from-scratch) performs retrieval of relevant documents to a question, and passes those documents to an LLM in order to ground the model's response.
|
||||
Many LLM applications implement a particular control flow of steps before and / or after LLM calls. As an example, [RAG](https://github.com/langchain-ai/rag-from-scratch) performs retrieval of documents relevant to a user question, and passes those documents to an LLM in order to ground the model's response in the provided document context.
|
||||
|
||||
Instead of hard-coding a fixed control flow, we sometimes want LLM systems that can pick its own control flow to solve more complex problems! This is one definition of an [agent](https://blog.langchain.dev/what-is-an-agent/): *an agent is a system that uses an LLM to decide the control flow of an application.* There are many ways that an LLM can control application:
|
||||
Instead of hard-coding a fixed control flow, we sometimes want LLM systems that can pick their own control flow to solve more complex problems! This is one definition of an [agent](https://blog.langchain.dev/what-is-an-agent/): *an agent is a system that uses an LLM to decide the control flow of an application.* There are many ways that an LLM can control application:
|
||||
|
||||
- An LLM can route between two potential paths
|
||||
- An LLM can decide which of many tools to call
|
||||
- An LLM can decide whether the generated answer is sufficient or more work is needed
|
||||
|
||||
As a result, there are many different types of [agent architectures](https://blog.langchain.dev/what-is-a-cognitive-architecture/), which given an LLM varying levels of control.
|
||||
As a result, there are many different types of [agent architectures](https://blog.langchain.dev/what-is-a-cognitive-architecture/), which give an LLM varying levels of control.
|
||||
|
||||

|
||||
|
||||
## Router
|
||||
|
||||
A router allows an LLM to select a single step from a specified set of options. This is an agent architecture that exhibits a relatively limited level of control because the LLM usually governs a single decision and can return a narrow set of outputs. Routers typically employ a few different concepts to achieve this.
|
||||
A router allows an LLM to select a single step from a specified set of options. This is an agent architecture that exhibits a relatively limited level of control because the LLM usually focuses on making a single decision and produces a specific output from limited set of pre-defined options. Routers typically employ a few different concepts to achieve this.
|
||||
|
||||
### Structured Output
|
||||
|
||||
Structured outputs with LLMs work by providing a specific format or schema that the LLM should follow in its response. This is similar to tool calling, but more general. While tool calling typically involves selecting and using predefined functions, structured outputs can be used for any type of formatted response. Common methods to achieve structured outputs include:
|
||||
|
||||
1. Prompt engineering: Instructing the LLM to respond in a specific format.
|
||||
1. Prompt engineering: Instructing the LLM to respond in a specific format via the system prompt.
|
||||
2. Output parsers: Using post-processing to extract structured data from LLM responses.
|
||||
3. Tool calling: Leveraging built-in tool calling capabilities of some LLMs to generate structured outputs.
|
||||
|
||||
@@ -30,7 +30,7 @@ Structured outputs are crucial for routing as they ensure the LLM's decision can
|
||||
|
||||
While a router allows an LLM to make a single decision, more complex agent architectures expand the LLM's control in two key ways:
|
||||
|
||||
1. Multi-step decision making: The LLM can control a sequence of decisions rather than just one.
|
||||
1. Multi-step decision making: The LLM can make a series of decisions, one after another, instead of just one.
|
||||
2. Tool access: The LLM can choose from and use a variety of tools to accomplish tasks.
|
||||
|
||||
[ReAct](https://arxiv.org/abs/2210.03629) is a popular general purpose agent architecture that combines these expansions, integrating three core concepts.
|
||||
@@ -39,13 +39,13 @@ While a router allows an LLM to make a single decision, more complex agent archi
|
||||
2. `Memory`: Enabling the agent to retain and use information from previous steps.
|
||||
3. `Planning`: Empowering the LLM to create and follow multi-step plans to achieve goals.
|
||||
|
||||
This architecture allows for more complex and flexible agent behaviors, going beyond simple routing to enable dynamic problem-solving across multiple steps. You can use it with [`create_react_agent`][langgraph.prebuilt.chat_agent_executor.create_react_agent].
|
||||
This architecture allows for more complex and flexible agent behaviors, going beyond simple routing to enable dynamic problem-solving with multiple steps. You can use it with [`create_react_agent`][langgraph.prebuilt.chat_agent_executor.create_react_agent].
|
||||
|
||||
### Tool calling
|
||||
|
||||
Tools are useful whenever you want an agent to interact with external systems. External systems (e.g., APIs) often require a particular input schema or payload, rather than natural language. When we bind an API, for example, as a tool we given the model awareness of the required input schema. The model will choose to call a tool based upon the natural language input from the user and it will return an output that adheres to the tool's schema.
|
||||
Tools are useful whenever you want an agent to interact with external systems. External systems (e.g., APIs) often require a particular input schema or payload, rather than natural language. When we bind an API, for example, as a tool, we give the model awareness of the required input schema. The model will choose to call a tool based upon the natural language input from the user and it will return an output that adheres to the tool's required schema.
|
||||
|
||||
[Many LLM providers support tool calling](https://python.langchain.com/v0.1/docs/integrations/chat/) and [tool calling interface](https://blog.langchain.dev/improving-core-tool-interfaces-and-docs-in-langchain/) in LangChain is simple: you can simply pass any Python `function` into `ChatModel.bind_tools(function)`.
|
||||
[Many LLM providers support tool calling](https://python.langchain.com/docs/integrations/chat/) and [tool calling interface](https://blog.langchain.dev/improving-core-tool-interfaces-and-docs-in-langchain/) in LangChain is simple: you can simply pass any Python `function` into `ChatModel.bind_tools(function)`.
|
||||
|
||||

|
||||
|
||||
@@ -67,11 +67,11 @@ Effective memory management enhances an agent's ability to maintain context, lea
|
||||
|
||||
### Planning
|
||||
|
||||
In the ReAct architecture, an LLM is called repeatedly in a while-loop. At each step the agent decides which tools to call, and what the inputs to those tools should be. Those tools are then executed, and the outputs are fed back into the LLM as observations. The while-loop terminates when the agent decides it is not worth calling any more tools.
|
||||
In the ReAct architecture, an LLM is called repeatedly in a while-loop. At each step the agent decides which tools to call, and what the inputs to those tools should be. Those tools are then executed, and the outputs are fed back into the LLM as observations. The while-loop terminates when the agent decides it has enough information to solve the user request and it is not worth calling any more tools.
|
||||
|
||||
### ReAct implementation
|
||||
|
||||
There are several differences between this paper and the pre-built [`create_react_agent`][langgraph.prebuilt.chat_agent_executor.create_react_agent] implementation:
|
||||
There are several differences between [this](https://arxiv.org/abs/2210.03629) paper and the pre-built [`create_react_agent`][langgraph.prebuilt.chat_agent_executor.create_react_agent] implementation:
|
||||
|
||||
- First, we use [tool-calling](#tool-calling) to have LLMs call tools, whereas the paper used prompting + parsing of raw output. This is because tool calling did not exist when the paper was written, but is generally better and more reliable.
|
||||
- Second, we use messages to prompt the LLM, whereas the paper used string formatting. This is because at the time of writing, LLMs didn't even expose a message-based interface, whereas now that's the only interface they expose.
|
||||
|
||||
+36
-38
@@ -17,6 +17,22 @@ While often used interchangeably, these terms represent distinct security concep
|
||||
|
||||
In LangGraph Platform, authentication is handled by your [`@auth.authenticate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.Auth.authenticate) handler, and authorization is handled by your [`@auth.on`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.Auth.on) handlers.
|
||||
|
||||
## Default Security Models
|
||||
|
||||
LangGraph Platform provides different security defaults:
|
||||
|
||||
### LangGraph Cloud
|
||||
|
||||
- Uses LangSmith API keys by default
|
||||
- Requires valid API key in `x-api-key` header
|
||||
- Can be customized with your auth handler
|
||||
|
||||
### Self-Hosted
|
||||
|
||||
- No default authentication
|
||||
- Complete flexibility to implement your security model
|
||||
- You control all aspects of authentication and authorization
|
||||
|
||||
## System Architecture
|
||||
|
||||
A typical authentication setup involves three main components:
|
||||
@@ -123,7 +139,7 @@ The returned user information is available:
|
||||
|
||||
After authentication, LangGraph calls your [`@auth.on`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.Auth.on) handlers to control access to specific resources (e.g., threads, assistants, crons). These handlers can:
|
||||
|
||||
1. Add metadata to be saved during resource creation by mutating the `value["metadata"]` dictionary directly.
|
||||
1. Add metadata to be saved during resource creation by mutating the `value["metadata"]` dictionary directly. See the [supported actions table](##supported-actions) for the list of types the value can take for each action.
|
||||
2. Filter resources by metadata during search/list or read operations by returning a [filter dictionary](#filter-operations).
|
||||
3. Raise an HTTP exception if access is denied.
|
||||
|
||||
@@ -342,10 +358,6 @@ async def rbac_create(ctx: Auth.types.AuthContext, value: dict):
|
||||
|
||||
## Supported Resources
|
||||
|
||||
LangGraph provides authorization handlers for the following resource types:
|
||||
|
||||
## Supported Resources
|
||||
|
||||
LangGraph provides three levels of authorization handlers, from most general to most specific:
|
||||
|
||||
1. **Global Handler** (`@auth.on`): Matches all resources and actions
|
||||
@@ -381,46 +393,32 @@ If a more specific handler is registered, the more general handler will not be c
|
||||
```
|
||||
More specific handlers provide better type hints since they handle fewer action types.
|
||||
|
||||
#### Supported actions and types {#supported-actions}
|
||||
Here are all the supported action handlers:
|
||||
|
||||
| Resource | Handler | Description |
|
||||
|----------|---------|-------------|
|
||||
| **Threads** | `@auth.on.threads.create` | Thread creation |
|
||||
| | `@auth.on.threads.read` | Thread retrieval |
|
||||
| | `@auth.on.threads.update` | Thread updates |
|
||||
| | `@auth.on.threads.delete` | Thread deletion |
|
||||
| | `@auth.on.threads.search` | Listing threads |
|
||||
| | `@auth.on.threads.create_run` | Creating or updating a run |
|
||||
| **Assistants** | `@auth.on.assistants.create` | Assistant creation |
|
||||
| | `@auth.on.assistants.read` | Assistant retrieval |
|
||||
| | `@auth.on.assistants.update` | Assistant updates |
|
||||
| | `@auth.on.assistants.delete` | Assistant deletion |
|
||||
| | `@auth.on.assistants.search` | Listing assistants |
|
||||
| **Crons** | `@auth.on.crons.create` | Cron job creation |
|
||||
| | `@auth.on.crons.read` | Cron job retrieval |
|
||||
| | `@auth.on.crons.update` | Cron job updates |
|
||||
| | `@auth.on.crons.delete` | Cron job deletion |
|
||||
| | `@auth.on.crons.search` | Listing cron jobs |
|
||||
| Resource | Handler | Description | Value Type |
|
||||
|----------|---------|-------------|------------|
|
||||
| **Threads** | `@auth.on.threads.create` | Thread creation | [`ThreadsCreate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.ThreadsCreate) |
|
||||
| | `@auth.on.threads.read` | Thread retrieval | [`ThreadsRead`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.ThreadsRead) |
|
||||
| | `@auth.on.threads.update` | Thread updates | [`ThreadsUpdate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.ThreadsUpdate) |
|
||||
| | `@auth.on.threads.delete` | Thread deletion | [`ThreadsDelete`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.ThreadsDelete) |
|
||||
| | `@auth.on.threads.search` | Listing threads | [`ThreadsSearch`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.ThreadsSearch) |
|
||||
| | `@auth.on.threads.create_run` | Creating or updating a run | [`RunsCreate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.RunsCreate) |
|
||||
| **Assistants** | `@auth.on.assistants.create` | Assistant creation | [`AssistantsCreate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.AssistantsCreate) |
|
||||
| | `@auth.on.assistants.read` | Assistant retrieval | [`AssistantsRead`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.AssistantsRead) |
|
||||
| | `@auth.on.assistants.update` | Assistant updates | [`AssistantsUpdate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.AssistantsUpdate) |
|
||||
| | `@auth.on.assistants.delete` | Assistant deletion | [`AssistantsDelete`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.AssistantsDelete) |
|
||||
| | `@auth.on.assistants.search` | Listing assistants | [`AssistantsSearch`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.AssistantsSearch) |
|
||||
| **Crons** | `@auth.on.crons.create` | Cron job creation | [`CronsCreate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.CronsCreate) |
|
||||
| | `@auth.on.crons.read` | Cron job retrieval | [`CronsRead`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.CronsRead) |
|
||||
| | `@auth.on.crons.update` | Cron job updates | [`CronsUpdate`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.CronsUpdate) |
|
||||
| | `@auth.on.crons.delete` | Cron job deletion | [`CronsDelete`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.CronsDelete) |
|
||||
| | `@auth.on.crons.search` | Listing cron jobs | [`CronsSearch`](../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.types.CronsSearch) |
|
||||
|
||||
???+ note "About Runs"
|
||||
Runs are scoped to their parent thread for access control. This means permissions are typically inherited from the thread, reflecting the conversational nature of the data model. All run operations (reading, listing) except creation are controlled by the thread's handlers.
|
||||
There is a specific `create_run` handler for creating new runs because it had more arguments that you can view in the handler.
|
||||
|
||||
## Default Security Models
|
||||
|
||||
LangGraph Platform provides different security defaults:
|
||||
|
||||
### LangGraph Cloud
|
||||
|
||||
- Uses LangSmith API keys by default
|
||||
- Requires valid API key in `x-api-key` header
|
||||
- Can be customized with your auth handler
|
||||
|
||||
### Self-Hosted
|
||||
|
||||
- No default authentication
|
||||
- Complete flexibility to implement your security model
|
||||
- You control all aspects of authentication and authorization
|
||||
|
||||
## Next Steps
|
||||
|
||||
|
||||
@@ -19,7 +19,7 @@ See the [how-to guide](../cloud/deployment/cloud.md#create-new-deployment) for c
|
||||
| **Deployment Type** | **CPU** | **Memory** | **Scaling** |
|
||||
|---------------------|---------|------------|---------------------|
|
||||
| Development | 1 CPU | 1 GB | Up to 1 container |
|
||||
| Production | 1 CPU | 2 GB | Up to 10 containers |
|
||||
| Production | 2 CPU | 2 GB | Up to 10 containers |
|
||||
|
||||
## Autoscaling
|
||||
`Production` type deployments automatically scale up to 10 containers. Scaling is based on the current request load for a single container. Specifically, the autoscaling implementation scales the deployment so that each container is processing about 10 concurrent requests. For example...
|
||||
|
||||
@@ -191,7 +191,7 @@ class State(MessagesState):
|
||||
|
||||
## Nodes
|
||||
|
||||
In LangGraph, nodes are typically python functions (sync or `async`) where the **first** positional argument is the [state](#state), and (optionally), the **second** positional argument is a "config", containing optional [configurable parameters](#configuration) (such as a `thread_id`).
|
||||
In LangGraph, nodes are typically python functions (sync or async) where the **first** positional argument is the [state](#state), and (optionally), the **second** positional argument is a "config", containing optional [configurable parameters](#configuration) (such as a `thread_id`).
|
||||
|
||||
Similar to `NetworkX`, you add these nodes to a graph using the [add_node][langgraph.graph.StateGraph.add_node] method:
|
||||
|
||||
|
||||
@@ -516,9 +516,7 @@
|
||||
"\n",
|
||||
" tool_response = tool_.invoke(tool_call)\n",
|
||||
" if isinstance(tool_response, ToolMessage):\n",
|
||||
" results.append(\n",
|
||||
" Command(goto=\"call_model\", update={\"messages\": [tool_response]})\n",
|
||||
" )\n",
|
||||
" results.append(Command(update={\"messages\": [tool_response]}))\n",
|
||||
"\n",
|
||||
" # handle tools that return Command directly\n",
|
||||
" elif isinstance(tool_response, Command):\n",
|
||||
@@ -531,6 +529,7 @@
|
||||
" graph.add_node(call_model)\n",
|
||||
" graph.add_node(call_tools)\n",
|
||||
" graph.add_edge(START, \"call_model\")\n",
|
||||
" graph.add_edge(\"call_tools\", \"call_model\")\n",
|
||||
"\n",
|
||||
" return graph.compile()"
|
||||
]
|
||||
|
||||
@@ -82,7 +82,7 @@ Assuming you are using JWT token authentication, you could access your deploymen
|
||||
url="http://localhost:2024",
|
||||
headers={"Authorization": f"Bearer {my_token}"}
|
||||
)
|
||||
threads = await client.threads.list()
|
||||
threads = await client.threads.search()
|
||||
```
|
||||
|
||||
=== "Python RemoteGraph"
|
||||
@@ -96,7 +96,7 @@ Assuming you are using JWT token authentication, you could access your deploymen
|
||||
url="http://localhost:2024",
|
||||
headers={"Authorization": f"Bearer {my_token}"}
|
||||
)
|
||||
threads = await remote_graph.threads.list()
|
||||
threads = await remote_graph.ainvoke(...)
|
||||
```
|
||||
|
||||
=== "JavaScript Client"
|
||||
@@ -109,7 +109,7 @@ Assuming you are using JWT token authentication, you could access your deploymen
|
||||
apiUrl: "http://localhost:2024",
|
||||
headers: { Authorization: `Bearer ${my_token}` },
|
||||
});
|
||||
const threads = await client.threads.list();
|
||||
const threads = await client.threads.search();
|
||||
```
|
||||
|
||||
=== "JavaScript RemoteGraph"
|
||||
@@ -123,7 +123,7 @@ Assuming you are using JWT token authentication, you could access your deploymen
|
||||
url: "http://localhost:2024",
|
||||
headers: { Authorization: `Bearer ${my_token}` },
|
||||
});
|
||||
const threads = await remoteGraph.threads.list();
|
||||
const threads = await remoteGraph.invoke(...);
|
||||
```
|
||||
|
||||
=== "CURL"
|
||||
|
||||
@@ -66,13 +66,10 @@ Since we're using Supabase for this, we can do this in the Supabase dashboard:
|
||||
```shell
|
||||
echo "SUPABASE_URL=your-project-url" >> .env
|
||||
```
|
||||
|
||||
3. Next, copy your service role secret key and add it to your `.env` file
|
||||
|
||||
```shell
|
||||
echo "SUPABASE_SERVICE_KEY=your-service-role-key" >> .env
|
||||
```
|
||||
|
||||
4. Finally, copy your "anon public" key and note it down. This will be used later when we set up our client code.
|
||||
|
||||
```bash
|
||||
@@ -96,7 +93,7 @@ And we'll keep our existing resource authorization logic unchanged
|
||||
|
||||
Let's update `src/security/auth.py` to implement this:
|
||||
|
||||
```python
|
||||
```python hl_lines="8-9 20-30" title="src/security/auth.py"
|
||||
import os
|
||||
import httpx
|
||||
from langgraph_sdk import Auth
|
||||
@@ -135,6 +132,7 @@ async def get_current_user(authorization: str | None):
|
||||
except Exception as e:
|
||||
raise Auth.exceptions.HTTPException(status_code=401, detail=str(e))
|
||||
|
||||
# ... the rest is the same as before
|
||||
|
||||
# Keep our resource authorization from the previous tutorial
|
||||
@auth.on
|
||||
@@ -153,9 +151,10 @@ Let's test this with a real user account!
|
||||
## Testing Authentication Flow
|
||||
|
||||
Let's test out our new authentication flow. You can run the following code in a file or notebook. You will need to provide:
|
||||
|
||||
- A valid email address
|
||||
- A Supabase project URL (from [above](#setup-auth-provider))
|
||||
- A Supabase service role key (also from [above](#setup-auth-provider))
|
||||
- A Supabase anon **public key** (also from [above](#setup-auth-provider))
|
||||
|
||||
```python
|
||||
import os
|
||||
@@ -174,10 +173,12 @@ email2 = f"{base_email[0]}+2@{base_email[1]}"
|
||||
SUPABASE_URL = os.environ.get("SUPABASE_URL")
|
||||
if not SUPABASE_URL:
|
||||
SUPABASE_URL = getpass("Enter your Supabase project URL: ")
|
||||
|
||||
SUPABASE_SERVICE_KEY = os.environ.get("SUPABASE_SERVICE_KEY")
|
||||
if not SUPABASE_SERVICE_KEY:
|
||||
SUPABASE_SERVICE_KEY = getpass("Enter your Supabase service role key: ")
|
||||
|
||||
# This is your PUBLIC anon key (which is safe to use client-side)
|
||||
# Do NOT mistake this for the secret service role key
|
||||
SUPABASE_ANON_KEY = os.environ.get("SUPABASE_ANON_KEY")
|
||||
if not SUPABASE_ANON_KEY:
|
||||
SUPABASE_ANON_KEY = getpass("Enter your public Supabase anon key: ")
|
||||
|
||||
|
||||
async def sign_up(email: str, password: str):
|
||||
@@ -186,7 +187,7 @@ async def sign_up(email: str, password: str):
|
||||
response = await client.post(
|
||||
f"{SUPABASE_URL}/auth/v1/signup",
|
||||
json={"email": email, "password": password},
|
||||
headers={"apiKey": SUPABASE_SERVICE_KEY},
|
||||
headers={"apiKey": SUPABASE_ANON_KEY},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
return response.json()
|
||||
@@ -207,15 +208,6 @@ Then run the code.
|
||||
Now let's test that users can only see their own data. Make sure the server is running (run `langgraph dev`) before proceeding. The following snippet requires the "anon public" key that you copied from the Supabase dashboard while [setting up the auth provider](#setup-auth-provider) previously.
|
||||
|
||||
```python
|
||||
import os
|
||||
import httpx
|
||||
|
||||
from langgraph_sdk import get_client
|
||||
|
||||
SUPABASE_ANON_KEY = os.environ.get("SUPABASE_ANON_KEY")
|
||||
if not SUPABASE_ANON_KEY:
|
||||
SUPABASE_ANON_KEY = getpass("Enter your Supabase anon key: ")
|
||||
|
||||
async def login(email: str, password: str):
|
||||
"""Get an access token for an existing user."""
|
||||
async with httpx.AsyncClient() as client:
|
||||
@@ -230,10 +222,8 @@ async def login(email: str, password: str):
|
||||
"Content-Type": "application/json"
|
||||
},
|
||||
)
|
||||
if response.status_code == 200:
|
||||
return response.json()["access_token"]
|
||||
else:
|
||||
raise ValueError(f"Login failed: {response.status_code} - {response.text}")
|
||||
assert response.status_code == 200
|
||||
return response.json()["access_token"]
|
||||
|
||||
|
||||
# Log in as user 1
|
||||
@@ -268,10 +258,11 @@ except Exception as e:
|
||||
```
|
||||
The output should look like this:
|
||||
|
||||
> ➜ custom-auth SUPABASE_ANON_KEY=eyJh... python test_oauth.py CHANGEME@example.com
|
||||
> ✅ User 1 created thread: d6af3754-95df-4176-aa10-dbd8dca40f1a
|
||||
> ✅ Unauthenticated access blocked: Client error '403 Forbidden' for url 'http://localhost:2024/threads'
|
||||
> ✅ User 2 blocked from User 1's thread: Client error '404 Not Found' for url 'http://localhost:2024/threads/d6af3754-95df-4176-aa10-dbd8dca40f1a'
|
||||
```shell
|
||||
✅ User 1 created thread: d6af3754-95df-4176-aa10-dbd8dca40f1a
|
||||
✅ Unauthenticated access blocked: Client error '403 Forbidden' for url 'http://localhost:2024/threads'
|
||||
✅ User 2 blocked from User 1's thread: Client error '404 Not Found' for url 'http://localhost:2024/threads/d6af3754-95df-4176-aa10-dbd8dca40f1a'
|
||||
```
|
||||
|
||||
Perfect! Our authentication and authorization are working together:
|
||||
1. Users must log in to access the bot
|
||||
|
||||
@@ -6,6 +6,17 @@
|
||||
2. [Resource Authorization](resource_auth.md) - Let users have private conversations
|
||||
3. [Production Auth](add_auth_server.md) - Add real user accounts and validate using OAuth2
|
||||
|
||||
!!! tip "Prerequisites"
|
||||
|
||||
This guide assumes basic familiarity with the following concepts:
|
||||
|
||||
* [**Authentication & Access Control**](../../concepts/auth.md)
|
||||
* [**LangGraph Platform**](../../concepts/index.md#langgraph-platform)
|
||||
|
||||
!!! note "Python only"
|
||||
|
||||
We currently only support custom authentication and authorization in Python deployments with `langgraph-api>=0.0.11`. Support for LangGraph.JS will be added soon.
|
||||
|
||||
In this tutorial, we will build a chatbot that only lets specific users access it. We'll start with the LangGraph template and add token-based security step by step. By the end, you'll have a working chatbot that checks for valid tokens before allowing access.
|
||||
|
||||
## Setting up our project
|
||||
@@ -32,8 +43,17 @@ If everything works, the server should start and open the studio in your browser
|
||||
> This in-memory server is designed for development and testing.
|
||||
> For production use, please use LangGraph Cloud.
|
||||
|
||||
Now that we've seen the base LangGraph app, let's add authentication to it! In part 1, we will start with a hard-coded token for illustration purposes.
|
||||
We will get to a "production-ready" authentication scheme in part 3, after mastering the basics.
|
||||
The graph should run, and if you were to self-host this on the public internet, anyone could access it!
|
||||
|
||||

|
||||
|
||||
Now that we've seen the base LangGraph app, let's add authentication to it!
|
||||
|
||||
???+ tip "Placeholder token"
|
||||
|
||||
In part 1, we will start with a hard-coded token for illustration purposes.
|
||||
We will get to a "production-ready" authentication scheme in part 3, after mastering the basics.
|
||||
|
||||
|
||||
## Adding Authentication
|
||||
|
||||
@@ -41,10 +61,10 @@ The [`Auth`](../../cloud/reference/sdk/python_sdk_ref.md#langgraph_sdk.auth.Auth
|
||||
|
||||
Create a new file `src/security/auth.py`. This is where our code will live to check if users are allowed to access our bot:
|
||||
|
||||
```python
|
||||
```python hl_lines="10 15-16" title="src/security/auth.py"
|
||||
from langgraph_sdk import Auth
|
||||
|
||||
# This is our toy user database
|
||||
# This is our toy user database. Do not do this in production
|
||||
VALID_TOKENS = {
|
||||
"user1-token": {"id": "user1", "name": "Alice"},
|
||||
"user2-token": {"id": "user2", "name": "Bob"},
|
||||
@@ -80,8 +100,13 @@ Notice that our [authentication](../../cloud/reference/sdk/python_sdk_ref.md#lan
|
||||
|
||||
Now tell LangGraph to use our authentication by adding the following to the [`langgraph.json`](../../cloud/reference/cli.md#configuration-file) configuration:
|
||||
|
||||
```json
|
||||
```json hl_lines="7-9" title="langgraph.json"
|
||||
{
|
||||
"dependencies": ["."],
|
||||
"graphs": {
|
||||
"agent": "./src/agent/graph.py:graph"
|
||||
},
|
||||
"env": ".env",
|
||||
"auth": {
|
||||
"path": "src/security/auth.py:auth"
|
||||
}
|
||||
@@ -109,7 +134,11 @@ langgraph dev --no-browser
|
||||
}
|
||||
```
|
||||
|
||||
Now let's try to chat with our bot. Run the following code in a file or notebook:
|
||||
Now let's try to chat with our bot. If we've implemented authentication correctly, we should only be able to access the bot if we provide a valid token in the request header. Users will still, however, be able to access each other's resources until we add [resource authorization handlers](../../concepts/auth.md#resource-authorization) in the next section of our tutorial.
|
||||
|
||||

|
||||
|
||||
Run the following code in a file or notebook:
|
||||
|
||||
```python
|
||||
from langgraph_sdk import get_client
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 614 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 545 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 293 KiB |
@@ -8,6 +8,13 @@
|
||||
|
||||
In this tutorial, we will extend our chatbot to give each user their own private conversations. We'll add [resource-level access control](../../concepts/auth.md#resource-level-access-control) so users can only see their own threads.
|
||||
|
||||

|
||||
|
||||
???+ tip "Placeholder token"
|
||||
|
||||
As we did in [part 1](getting_started.md), for this section, we will use a hard-coded token for illustration purposes.
|
||||
We will get to a "production-ready" authentication scheme in part 3, after mastering the basics.
|
||||
|
||||
## Understanding Resource Authorization
|
||||
|
||||
In the last tutorial, we controlled who could access our bot. But right now, any authenticated user can see everyone else's conversations! Let's fix that by adding [resource authorization](../../concepts/auth.md#resource-authorization).
|
||||
@@ -32,7 +39,7 @@ Authorization handlers are functions that run **after** authentication succeeds.
|
||||
|
||||
Let's update our `src/security/auth.py` and add one authorization handler that is run on every request:
|
||||
|
||||
```python hl_lines="29-39"
|
||||
```python hl_lines="29-39" title="src/security/auth.py"
|
||||
from langgraph_sdk import Auth
|
||||
|
||||
# Keep our test users from the previous tutorial
|
||||
@@ -66,7 +73,51 @@ async def add_owner(
|
||||
value: dict, # The resource being created/accessed
|
||||
):
|
||||
"""Make resources private to their creator."""
|
||||
# Add owner when creating resources
|
||||
# Examples:
|
||||
# ctx: AuthContext(
|
||||
# permissions=[],
|
||||
# user=ProxyUser(
|
||||
# identity='user1',
|
||||
# is_authenticated=True,
|
||||
# display_name='user1'
|
||||
# ),
|
||||
# resource='threads',
|
||||
# action='create_run'
|
||||
# )
|
||||
# value:
|
||||
# {
|
||||
# 'thread_id': UUID('1e1b2733-303f-4dcd-9620-02d370287d72'),
|
||||
# 'assistant_id': UUID('fe096781-5601-53d2-b2f6-0d3403f7e9ca'),
|
||||
# 'run_id': UUID('1efbe268-1627-66d4-aa8d-b956b0f02a41'),
|
||||
# 'status': 'pending',
|
||||
# 'metadata': {},
|
||||
# 'prevent_insert_if_inflight': True,
|
||||
# 'multitask_strategy': 'reject',
|
||||
# 'if_not_exists': 'reject',
|
||||
# 'after_seconds': 0,
|
||||
# 'kwargs': {
|
||||
# 'input': {'messages': [{'role': 'user', 'content': 'Hello!'}]},
|
||||
# 'command': None,
|
||||
# 'config': {
|
||||
# 'configurable': {
|
||||
# 'langgraph_auth_user': ... Your user object...
|
||||
# 'langgraph_auth_user_id': 'user1'
|
||||
# }
|
||||
# },
|
||||
# 'stream_mode': ['values'],
|
||||
# 'interrupt_before': None,
|
||||
# 'interrupt_after': None,
|
||||
# 'webhook': None,
|
||||
# 'feedback_keys': None,
|
||||
# 'temporary': False,
|
||||
# 'subgraphs': False
|
||||
# }
|
||||
# }
|
||||
|
||||
# Do 2 things:
|
||||
# 1. Add the user's ID to the resource's metadata. Each LangGraph resource has a `metadata` dict that persists with the resource.
|
||||
# this metadata is useful for filtering in read and update operations
|
||||
# 2. Return a filter that lets users only see their own resources
|
||||
filters = {"owner": ctx.user.identity}
|
||||
metadata = value.setdefault("metadata", {})
|
||||
metadata.update(filters)
|
||||
@@ -103,6 +154,10 @@ bob = get_client(
|
||||
headers={"Authorization": "Bearer user2-token"}
|
||||
)
|
||||
|
||||
# Alice creates an assistant
|
||||
alice_assistant = await alice.assistants.create()
|
||||
print(f"✅ Alice created assistant: {alice_assistant['assistant_id']}")
|
||||
|
||||
# Alice creates a thread and chats
|
||||
alice_thread = await alice.threads.create()
|
||||
print(f"✅ Alice created thread: {alice_thread['thread_id']}")
|
||||
@@ -130,16 +185,16 @@ await bob.runs.create(
|
||||
print(f"✅ Bob created his own thread: {bob_thread['thread_id']}")
|
||||
|
||||
# List threads - each user only sees their own
|
||||
alice_threads = await alice.threads.list()
|
||||
bob_threads = await bob.threads.list()
|
||||
alice_threads = await alice.threads.search()
|
||||
bob_threads = await bob.threads.search()
|
||||
print(f"✅ Alice sees {len(alice_threads)} thread")
|
||||
print(f"✅ Bob sees {len(bob_threads)} thread")
|
||||
|
||||
```
|
||||
|
||||
Run the test code and you should see output like this:
|
||||
|
||||
```bash
|
||||
✅ Alice created assistant: fc50fb08-78da-45a9-93cc-1d3928a3fc37
|
||||
✅ Alice created thread: 533179b7-05bc-4d48-b47a-a83cbdb5781d
|
||||
✅ Bob correctly denied access: Client error '404 Not Found' for url 'http://localhost:2024/threads/533179b7-05bc-4d48-b47a-a83cbdb5781d'
|
||||
For more information check: https://developer.mozilla.org/en-US/docs/Web/HTTP/Status/404
|
||||
@@ -176,11 +231,15 @@ async def on_thread_create(
|
||||
1. Sets metadata on the thread being created to track ownership
|
||||
2. Returns a filter that ensures only the creator can access it
|
||||
"""
|
||||
# Example value:
|
||||
# {'thread_id': UUID('99b045bc-b90b-41a8-b882-dabc541cf740'), 'metadata': {}, 'if_exists': 'raise'}
|
||||
|
||||
# Add owner metadata to the thread being created
|
||||
# This metadata is stored with the thread and persists
|
||||
metadata = value.setdefault("metadata", {})
|
||||
metadata["owner"] = ctx.user.identity
|
||||
|
||||
|
||||
# Return filter to restrict access to just the creator
|
||||
return {"owner": ctx.user.identity}
|
||||
|
||||
@@ -197,53 +256,61 @@ async def on_thread_read(
|
||||
"""
|
||||
return {"owner": ctx.user.identity}
|
||||
|
||||
@auth.on.threads.create_run
|
||||
async def on_run_create(
|
||||
@auth.on.assistants
|
||||
async def on_assistants(
|
||||
ctx: Auth.types.AuthContext,
|
||||
value: Auth.types.on.threads.create_run.value,
|
||||
value: Auth.types.on.assistants.value,
|
||||
):
|
||||
"""Only let thread owners create runs.
|
||||
|
||||
This handler runs when creating runs on a thread. The filter
|
||||
applies to the parent thread, not the run being created.
|
||||
This ensures only thread owners can create runs on their threads.
|
||||
"""
|
||||
return {"owner": ctx.user.identity}
|
||||
# For illustration purposes, we will deny all requests
|
||||
# that touch the assistants resource
|
||||
# Example value:
|
||||
# {
|
||||
# 'assistant_id': UUID('63ba56c3-b074-4212-96e2-cc333bbc4eb4'),
|
||||
# 'graph_id': 'agent',
|
||||
# 'config': {},
|
||||
# 'metadata': {},
|
||||
# 'name': 'Untitled'
|
||||
# }
|
||||
raise Auth.exceptions.HTTPException(
|
||||
status_code=403,
|
||||
detail="User lacks the required permissions.",
|
||||
)
|
||||
```
|
||||
|
||||
Notice that instead of one global handler, we now have specific handlers for:
|
||||
|
||||
1. Creating threads
|
||||
2. Reading threads
|
||||
3. Creating runs
|
||||
4. Accessing assistants
|
||||
3. Accessing assistants
|
||||
|
||||
The first three of these match specific **actions** on each resource (see [resource actions](../../concepts/auth.md#resource-actions)), while the last one (`@auth.on.assistants`) matches _any_ action on the `assistants` resource. For each request, LangGraph will run the most specific handler that matches the resource and action being accessed. This means that the four handlers above will run rather than the broad "@auth.on" handler.
|
||||
The first three of these match specific **actions** on each resource (see [resource actions](../../concepts/auth.md#resource-actions)), while the last one (`@auth.on.assistants`) matches _any_ action on the `assistants` resource. For each request, LangGraph will run the most specific handler that matches the resource and action being accessed. This means that the four handlers above will run rather than the broadly scoped "`@auth.on`" handler.
|
||||
|
||||
Try adding the following test code to `test_private.py`:
|
||||
Try adding the following test code to your test file:
|
||||
|
||||
```python
|
||||
async def test_private():
|
||||
# ... Same as before
|
||||
# Try creating an assistant. This should fail
|
||||
try:
|
||||
await alice.assistants.create("agent")
|
||||
print("❌ Alice shouldn't be able to create assistants!")
|
||||
except Exception as e:
|
||||
print("✅ Alice correctly denied access:", e)
|
||||
# ... Same as before
|
||||
# Try creating an assistant. This should fail
|
||||
try:
|
||||
await alice.assistants.create("agent")
|
||||
print("❌ Alice shouldn't be able to create assistants!")
|
||||
except Exception as e:
|
||||
print("✅ Alice correctly denied access:", e)
|
||||
|
||||
# Try searching for assistants. This also should fail
|
||||
try:
|
||||
await alice.assistants.search()
|
||||
print("❌ Alice shouldn't be able to search assistants!")
|
||||
except Exception as e:
|
||||
print("✅ Alice correctly denied access to searching assistants:", e)
|
||||
# Try searching for assistants. This also should fail
|
||||
try:
|
||||
await alice.assistants.search()
|
||||
print("❌ Alice shouldn't be able to search assistants!")
|
||||
except Exception as e:
|
||||
print("✅ Alice correctly denied access to searching assistants:", e)
|
||||
|
||||
# Alice can still create threads
|
||||
alice_thread = await alice.threads.create()
|
||||
print(f"✅ Alice created thread: {alice_thread['thread_id']}")
|
||||
```
|
||||
|
||||
And then run the test code again:
|
||||
|
||||
```bash
|
||||
> python test_private.py
|
||||
✅ Alice created thread: dcea5cd8-eb70-4a01-a4b6-643b14e8f754
|
||||
✅ Bob correctly denied access: Client error '404 Not Found' for url 'http://localhost:2024/threads/dcea5cd8-eb70-4a01-a4b6-643b14e8f754'
|
||||
For more information check: https://developer.mozilla.org/en-US/docs/Web/HTTP/Status/404
|
||||
|
||||
@@ -553,9 +553,6 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from typing import Literal\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def route_tools(\n",
|
||||
" state: State,\n",
|
||||
"):\n",
|
||||
@@ -2653,7 +2650,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from typing import Annotated, Literal\n",
|
||||
"from typing import Annotated\n",
|
||||
"\n",
|
||||
"from langchain_anthropic import ChatAnthropic\n",
|
||||
"from langchain_community.tools.tavily_search import TavilySearchResults\n",
|
||||
|
||||
@@ -180,7 +180,7 @@ LangGraph Studio Web is a specialized UI that you can connect to LangGraph API s
|
||||
```js
|
||||
const { Client } = await import("@langchain/langgraph-sdk");
|
||||
|
||||
// only set the apiUrl if you changed the default port when calling langgraph up
|
||||
// only set the apiUrl if you changed the default port when calling langgraph dev
|
||||
const client = new Client({ apiUrl: "http://localhost:2024"});
|
||||
|
||||
const streamResponse = client.runs.stream(
|
||||
|
||||
@@ -870,7 +870,7 @@
|
||||
" # No deps or all deps satisfied\n",
|
||||
" # can schedule now\n",
|
||||
" schedule_task.invoke(dict(task=task, observations=observations))\n",
|
||||
" # futures.append(executor.submit(schedule_task.invoke dict(task=task, observations=observations)))\n",
|
||||
" # futures.append(executor.submit(schedule_task.invoke, dict(task=task, observations=observations)))\n",
|
||||
"\n",
|
||||
" # All tasks have been submitted or enqueued\n",
|
||||
" # Wait for them to complete\n",
|
||||
|
||||
@@ -2,6 +2,7 @@ site_name: ""
|
||||
site_description: Build language agents as graphs
|
||||
site_url: https://langchain-ai.github.io/langgraph/
|
||||
repo_url: https://github.com/langchain-ai/langgraph
|
||||
edit_uri: edit/main/docs/docs/
|
||||
theme:
|
||||
name: material
|
||||
custom_dir: overrides
|
||||
@@ -16,6 +17,7 @@ theme:
|
||||
- content.code.copy
|
||||
- content.code.select
|
||||
- content.tabs.link
|
||||
- content.action.edit
|
||||
- content.tooltips
|
||||
- header.autohide
|
||||
- navigation.expand
|
||||
|
||||
@@ -19,6 +19,7 @@ from langgraph.checkpoint.base import (
|
||||
)
|
||||
from langgraph.checkpoint.postgres import _internal
|
||||
from langgraph.checkpoint.postgres.base import BasePostgresSaver
|
||||
from langgraph.checkpoint.postgres.shallow import ShallowPostgresSaver
|
||||
from langgraph.checkpoint.serde.base import SerializerProtocol
|
||||
|
||||
Conn = _internal.Conn # For backward compatibility
|
||||
@@ -396,4 +397,4 @@ class PostgresSaver(BasePostgresSaver):
|
||||
yield cur
|
||||
|
||||
|
||||
__all__ = ["PostgresSaver", "BasePostgresSaver", "Conn"]
|
||||
__all__ = ["PostgresSaver", "BasePostgresSaver", "ShallowPostgresSaver", "Conn"]
|
||||
|
||||
@@ -19,6 +19,7 @@ from langgraph.checkpoint.base import (
|
||||
)
|
||||
from langgraph.checkpoint.postgres import _ainternal
|
||||
from langgraph.checkpoint.postgres.base import BasePostgresSaver
|
||||
from langgraph.checkpoint.postgres.shallow import AsyncShallowPostgresSaver
|
||||
from langgraph.checkpoint.serde.base import SerializerProtocol
|
||||
|
||||
Conn = _ainternal.Conn # For backward compatibility
|
||||
@@ -464,4 +465,4 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
).result()
|
||||
|
||||
|
||||
__all__ = ["AsyncPostgresSaver", "Conn"]
|
||||
__all__ = ["AsyncPostgresSaver", "AsyncShallowPostgresSaver", "Conn"]
|
||||
|
||||
@@ -0,0 +1,918 @@
|
||||
import asyncio
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator, Sequence
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from typing import Any, Optional
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from psycopg import (
|
||||
AsyncConnection,
|
||||
AsyncCursor,
|
||||
AsyncPipeline,
|
||||
Capabilities,
|
||||
Connection,
|
||||
Cursor,
|
||||
Pipeline,
|
||||
)
|
||||
from psycopg.rows import DictRow, dict_row
|
||||
from psycopg.types.json import Jsonb
|
||||
from psycopg_pool import AsyncConnectionPool, ConnectionPool
|
||||
|
||||
from langgraph.checkpoint.base import (
|
||||
WRITES_IDX_MAP,
|
||||
ChannelVersions,
|
||||
Checkpoint,
|
||||
CheckpointMetadata,
|
||||
CheckpointTuple,
|
||||
)
|
||||
from langgraph.checkpoint.postgres import _ainternal, _internal
|
||||
from langgraph.checkpoint.postgres.base import BasePostgresSaver
|
||||
from langgraph.checkpoint.serde.base import SerializerProtocol
|
||||
from langgraph.checkpoint.serde.types import TASKS
|
||||
|
||||
"""
|
||||
To add a new migration, add a new string to the MIGRATIONS list.
|
||||
The position of the migration in the list is the version number.
|
||||
"""
|
||||
MIGRATIONS = [
|
||||
"""CREATE TABLE IF NOT EXISTS checkpoint_migrations (
|
||||
v INTEGER PRIMARY KEY
|
||||
);""",
|
||||
"""CREATE TABLE IF NOT EXISTS checkpoints (
|
||||
thread_id TEXT NOT NULL,
|
||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||
type TEXT,
|
||||
checkpoint JSONB NOT NULL,
|
||||
metadata JSONB NOT NULL DEFAULT '{}',
|
||||
PRIMARY KEY (thread_id, checkpoint_ns)
|
||||
);""",
|
||||
"""CREATE TABLE IF NOT EXISTS checkpoint_blobs (
|
||||
thread_id TEXT NOT NULL,
|
||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT NOT NULL,
|
||||
blob BYTEA,
|
||||
PRIMARY KEY (thread_id, checkpoint_ns, channel)
|
||||
);""",
|
||||
"""CREATE TABLE IF NOT EXISTS checkpoint_writes (
|
||||
thread_id TEXT NOT NULL,
|
||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||
checkpoint_id TEXT NOT NULL,
|
||||
task_id TEXT NOT NULL,
|
||||
idx INTEGER NOT NULL,
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT,
|
||||
blob BYTEA NOT NULL,
|
||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
||||
);""",
|
||||
"""
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoints_thread_id_idx ON checkpoints(thread_id);
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoint_blobs_thread_id_idx ON checkpoint_blobs(thread_id);
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoint_writes_thread_id_idx ON checkpoint_writes(thread_id);
|
||||
""",
|
||||
]
|
||||
|
||||
SELECT_SQL = f"""
|
||||
select
|
||||
thread_id,
|
||||
checkpoint,
|
||||
checkpoint_ns,
|
||||
metadata,
|
||||
(
|
||||
select array_agg(array[bl.channel::bytea, bl.type::bytea, bl.blob])
|
||||
from jsonb_each_text(checkpoint -> 'channel_versions')
|
||||
inner join checkpoint_blobs bl
|
||||
on bl.thread_id = checkpoints.thread_id
|
||||
and bl.checkpoint_ns = checkpoints.checkpoint_ns
|
||||
and bl.channel = jsonb_each_text.key
|
||||
) as channel_values,
|
||||
(
|
||||
select
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
||||
from checkpoint_writes cw
|
||||
where cw.thread_id = checkpoints.thread_id
|
||||
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
||||
and cw.checkpoint_id = (checkpoint->>'id')
|
||||
) as pending_writes,
|
||||
(
|
||||
select array_agg(array[cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
||||
from checkpoint_writes cw
|
||||
where cw.thread_id = checkpoints.thread_id
|
||||
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
||||
and cw.channel = '{TASKS}'
|
||||
) as pending_sends
|
||||
from checkpoints """
|
||||
|
||||
UPSERT_CHECKPOINT_BLOBS_SQL = """
|
||||
INSERT INTO checkpoint_blobs (thread_id, checkpoint_ns, channel, type, blob)
|
||||
VALUES (%s, %s, %s, %s, %s)
|
||||
ON CONFLICT (thread_id, checkpoint_ns, channel) DO UPDATE SET
|
||||
type = EXCLUDED.type,
|
||||
blob = EXCLUDED.blob;
|
||||
"""
|
||||
|
||||
UPSERT_CHECKPOINTS_SQL = """
|
||||
INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint, metadata)
|
||||
VALUES (%s, %s, %s, %s)
|
||||
ON CONFLICT (thread_id, checkpoint_ns)
|
||||
DO UPDATE SET
|
||||
checkpoint = EXCLUDED.checkpoint,
|
||||
metadata = EXCLUDED.metadata;
|
||||
"""
|
||||
|
||||
UPSERT_CHECKPOINT_WRITES_SQL = """
|
||||
INSERT INTO checkpoint_writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob)
|
||||
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) DO UPDATE SET
|
||||
channel = EXCLUDED.channel,
|
||||
type = EXCLUDED.type,
|
||||
blob = EXCLUDED.blob;
|
||||
"""
|
||||
|
||||
INSERT_CHECKPOINT_WRITES_SQL = """
|
||||
INSERT INTO checkpoint_writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob)
|
||||
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) DO NOTHING
|
||||
"""
|
||||
|
||||
|
||||
def _dump_blobs(
|
||||
serde: SerializerProtocol,
|
||||
thread_id: str,
|
||||
checkpoint_ns: str,
|
||||
values: dict[str, Any],
|
||||
versions: ChannelVersions,
|
||||
) -> list[tuple[str, str, str, str, str, Optional[bytes]]]:
|
||||
if not versions:
|
||||
return []
|
||||
|
||||
return [
|
||||
(
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
k,
|
||||
*(serde.dumps_typed(values[k]) if k in values else ("empty", None)),
|
||||
)
|
||||
for k in versions
|
||||
]
|
||||
|
||||
|
||||
class ShallowPostgresSaver(BasePostgresSaver):
|
||||
"""A checkpoint saver that uses Postgres to store checkpoints.
|
||||
|
||||
This checkpointer ONLY stores the most recent checkpoint and does NOT retain any history.
|
||||
It is meant to be a light-weight drop-in replacement for the PostgresSaver that
|
||||
supports most of the LangGraph persistence functionality with the exception of time travel.
|
||||
"""
|
||||
|
||||
SELECT_SQL = SELECT_SQL
|
||||
MIGRATIONS = MIGRATIONS
|
||||
UPSERT_CHECKPOINT_BLOBS_SQL = UPSERT_CHECKPOINT_BLOBS_SQL
|
||||
UPSERT_CHECKPOINTS_SQL = UPSERT_CHECKPOINTS_SQL
|
||||
UPSERT_CHECKPOINT_WRITES_SQL = UPSERT_CHECKPOINT_WRITES_SQL
|
||||
INSERT_CHECKPOINT_WRITES_SQL = INSERT_CHECKPOINT_WRITES_SQL
|
||||
|
||||
lock: threading.Lock
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
conn: _internal.Conn,
|
||||
pipe: Optional[Pipeline] = None,
|
||||
serde: Optional[SerializerProtocol] = None,
|
||||
) -> None:
|
||||
super().__init__(serde=serde)
|
||||
if isinstance(conn, ConnectionPool) and pipe is not None:
|
||||
raise ValueError(
|
||||
"Pipeline should be used only with a single Connection, not ConnectionPool."
|
||||
)
|
||||
|
||||
self.conn = conn
|
||||
self.pipe = pipe
|
||||
self.lock = threading.Lock()
|
||||
self.supports_pipeline = Capabilities().has_pipeline()
|
||||
|
||||
@classmethod
|
||||
@contextmanager
|
||||
def from_conn_string(
|
||||
cls, conn_string: str, *, pipeline: bool = False
|
||||
) -> Iterator["ShallowPostgresSaver"]:
|
||||
"""Create a new ShallowPostgresSaver instance from a connection string.
|
||||
|
||||
Args:
|
||||
conn_string (str): The Postgres connection info string.
|
||||
pipeline (bool): whether to use Pipeline
|
||||
|
||||
Returns:
|
||||
ShallowPostgresSaver: A new ShallowPostgresSaver instance.
|
||||
"""
|
||||
with Connection.connect(
|
||||
conn_string, autocommit=True, prepare_threshold=0, row_factory=dict_row
|
||||
) as conn:
|
||||
if pipeline:
|
||||
with conn.pipeline() as pipe:
|
||||
yield cls(conn, pipe)
|
||||
else:
|
||||
yield cls(conn)
|
||||
|
||||
def setup(self) -> None:
|
||||
"""Set up the checkpoint database asynchronously.
|
||||
|
||||
This method creates the necessary tables in the Postgres database if they don't
|
||||
already exist and runs database migrations. It MUST be called directly by the user
|
||||
the first time checkpointer is used.
|
||||
"""
|
||||
with self._cursor() as cur:
|
||||
cur.execute(self.MIGRATIONS[0])
|
||||
results = cur.execute(
|
||||
"SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1"
|
||||
)
|
||||
row = results.fetchone()
|
||||
if row is None:
|
||||
version = -1
|
||||
else:
|
||||
version = row["v"]
|
||||
for v, migration in zip(
|
||||
range(version + 1, len(self.MIGRATIONS)),
|
||||
self.MIGRATIONS[version + 1 :],
|
||||
):
|
||||
cur.execute(migration)
|
||||
cur.execute(f"INSERT INTO checkpoint_migrations (v) VALUES ({v})")
|
||||
if self.pipe:
|
||||
self.pipe.sync()
|
||||
|
||||
def list(
|
||||
self,
|
||||
config: Optional[RunnableConfig],
|
||||
*,
|
||||
filter: Optional[dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> Iterator[CheckpointTuple]:
|
||||
"""List checkpoints from the database.
|
||||
|
||||
This method retrieves a list of checkpoint tuples from the Postgres database based
|
||||
on the provided config. For ShallowPostgresSaver, this method returns a list with
|
||||
ONLY the most recent checkpoint.
|
||||
"""
|
||||
where, args = self._search_where(config, filter, before)
|
||||
query = self.SELECT_SQL + where
|
||||
if limit:
|
||||
query += f" LIMIT {limit}"
|
||||
with self._cursor() as cur:
|
||||
cur.execute(self.SELECT_SQL + where, args, binary=True)
|
||||
for value in cur:
|
||||
checkpoint = self._load_checkpoint(
|
||||
value["checkpoint"],
|
||||
value["channel_values"],
|
||||
value["pending_sends"],
|
||||
)
|
||||
yield CheckpointTuple(
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": value["thread_id"],
|
||||
"checkpoint_ns": value["checkpoint_ns"],
|
||||
"checkpoint_id": checkpoint["id"],
|
||||
}
|
||||
},
|
||||
checkpoint=checkpoint,
|
||||
metadata=self._load_metadata(value["metadata"]),
|
||||
pending_writes=self._load_writes(value["pending_writes"]),
|
||||
)
|
||||
|
||||
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||
"""Get a checkpoint tuple from the database.
|
||||
|
||||
This method retrieves a checkpoint tuple from the Postgres database based on the
|
||||
provided config (matching the thread ID in the config).
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): The config to use for retrieving the checkpoint.
|
||||
|
||||
Returns:
|
||||
Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found.
|
||||
|
||||
Examples:
|
||||
|
||||
Basic:
|
||||
>>> config = {"configurable": {"thread_id": "1"}}
|
||||
>>> checkpoint_tuple = memory.get_tuple(config)
|
||||
>>> print(checkpoint_tuple)
|
||||
CheckpointTuple(...)
|
||||
|
||||
With timestamp:
|
||||
|
||||
>>> config = {
|
||||
... "configurable": {
|
||||
... "thread_id": "1",
|
||||
... "checkpoint_ns": "",
|
||||
... "checkpoint_id": "1ef4f797-8335-6428-8001-8a1503f9b875",
|
||||
... }
|
||||
... }
|
||||
>>> checkpoint_tuple = memory.get_tuple(config)
|
||||
>>> print(checkpoint_tuple)
|
||||
CheckpointTuple(...)
|
||||
""" # noqa
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
args = (thread_id, checkpoint_ns)
|
||||
where = "WHERE thread_id = %s AND checkpoint_ns = %s"
|
||||
|
||||
with self._cursor() as cur:
|
||||
cur.execute(
|
||||
self.SELECT_SQL + where,
|
||||
args,
|
||||
binary=True,
|
||||
)
|
||||
|
||||
for value in cur:
|
||||
checkpoint = self._load_checkpoint(
|
||||
value["checkpoint"],
|
||||
value["channel_values"],
|
||||
value["pending_sends"],
|
||||
)
|
||||
return CheckpointTuple(
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"checkpoint_id": checkpoint["id"],
|
||||
}
|
||||
},
|
||||
checkpoint=checkpoint,
|
||||
metadata=self._load_metadata(value["metadata"]),
|
||||
pending_writes=self._load_writes(value["pending_writes"]),
|
||||
)
|
||||
|
||||
def put(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
"""Save a checkpoint to the database.
|
||||
|
||||
This method saves a checkpoint to the Postgres database. The checkpoint is associated
|
||||
with the provided config. For ShallowPostgresSaver, this method saves ONLY the most recent
|
||||
checkpoint and overwrites a previous checkpoint, if it exists.
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): The config to associate with the checkpoint.
|
||||
checkpoint (Checkpoint): The checkpoint to save.
|
||||
metadata (CheckpointMetadata): Additional metadata to save with the checkpoint.
|
||||
new_versions (ChannelVersions): New channel versions as of this write.
|
||||
|
||||
Returns:
|
||||
RunnableConfig: Updated configuration after storing the checkpoint.
|
||||
|
||||
Examples:
|
||||
|
||||
>>> from langgraph.checkpoint.postgres import ShallowPostgresSaver
|
||||
>>> DB_URI = "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable"
|
||||
>>> with ShallowPostgresSaver.from_conn_string(DB_URI) as memory:
|
||||
>>> config = {"configurable": {"thread_id": "1", "checkpoint_ns": ""}}
|
||||
>>> checkpoint = {"ts": "2024-05-04T06:32:42.235444+00:00", "id": "1ef4f797-8335-6428-8001-8a1503f9b875", "channel_values": {"key": "value"}}
|
||||
>>> saved_config = memory.put(config, checkpoint, {"source": "input", "step": 1, "writes": {"key": "value"}}, {})
|
||||
>>> print(saved_config)
|
||||
{'configurable': {'thread_id': '1', 'checkpoint_ns': '', 'checkpoint_id': '1ef4f797-8335-6428-8001-8a1503f9b875'}}
|
||||
"""
|
||||
configurable = config["configurable"].copy()
|
||||
thread_id = configurable.pop("thread_id")
|
||||
checkpoint_ns = configurable.pop("checkpoint_ns")
|
||||
|
||||
copy = checkpoint.copy()
|
||||
next_config = {
|
||||
"configurable": {
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"checkpoint_id": checkpoint["id"],
|
||||
}
|
||||
}
|
||||
|
||||
with self._cursor(pipeline=True) as cur:
|
||||
cur.execute(
|
||||
"""DELETE FROM checkpoint_writes
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s AND checkpoint_id NOT IN (%s, %s)""",
|
||||
(
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
checkpoint["id"],
|
||||
configurable.get("checkpoint_id", ""),
|
||||
),
|
||||
)
|
||||
cur.executemany(
|
||||
self.UPSERT_CHECKPOINT_BLOBS_SQL,
|
||||
_dump_blobs(
|
||||
self.serde,
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
copy.pop("channel_values"), # type: ignore[misc]
|
||||
new_versions,
|
||||
),
|
||||
)
|
||||
cur.execute(
|
||||
self.UPSERT_CHECKPOINTS_SQL,
|
||||
(
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
Jsonb(self._dump_checkpoint(copy)),
|
||||
self._dump_metadata(metadata),
|
||||
),
|
||||
)
|
||||
return next_config
|
||||
|
||||
def put_writes(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
writes: Sequence[tuple[str, Any]],
|
||||
task_id: str,
|
||||
) -> None:
|
||||
"""Store intermediate writes linked to a checkpoint.
|
||||
|
||||
This method saves intermediate writes associated with a checkpoint to the Postgres database.
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): Configuration of the related checkpoint.
|
||||
writes (List[Tuple[str, Any]]): List of writes to store.
|
||||
task_id (str): Identifier for the task creating the writes.
|
||||
"""
|
||||
query = (
|
||||
self.UPSERT_CHECKPOINT_WRITES_SQL
|
||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||
else self.INSERT_CHECKPOINT_WRITES_SQL
|
||||
)
|
||||
with self._cursor(pipeline=True) as cur:
|
||||
cur.executemany(
|
||||
query,
|
||||
self._dump_writes(
|
||||
config["configurable"]["thread_id"],
|
||||
config["configurable"]["checkpoint_ns"],
|
||||
config["configurable"]["checkpoint_id"],
|
||||
task_id,
|
||||
writes,
|
||||
),
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
def _cursor(self, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]:
|
||||
"""Create a database cursor as a context manager.
|
||||
|
||||
Args:
|
||||
pipeline (bool): whether to use pipeline for the DB operations inside the context manager.
|
||||
Will be applied regardless of whether the ShallowPostgresSaver instance was initialized with a pipeline.
|
||||
If pipeline mode is not supported, will fall back to using transaction context manager.
|
||||
"""
|
||||
with _internal.get_connection(self.conn) as conn:
|
||||
if self.pipe:
|
||||
# a connection in pipeline mode can be used concurrently
|
||||
# in multiple threads/coroutines, but only one cursor can be
|
||||
# used at a time
|
||||
try:
|
||||
with conn.cursor(binary=True, row_factory=dict_row) as cur:
|
||||
yield cur
|
||||
finally:
|
||||
if pipeline:
|
||||
self.pipe.sync()
|
||||
elif pipeline:
|
||||
# a connection not in pipeline mode can only be used by one
|
||||
# thread/coroutine at a time, so we acquire a lock
|
||||
if self.supports_pipeline:
|
||||
with (
|
||||
self.lock,
|
||||
conn.pipeline(),
|
||||
conn.cursor(binary=True, row_factory=dict_row) as cur,
|
||||
):
|
||||
yield cur
|
||||
else:
|
||||
# Use connection's transaction context manager when pipeline mode not supported
|
||||
with (
|
||||
self.lock,
|
||||
conn.transaction(),
|
||||
conn.cursor(binary=True, row_factory=dict_row) as cur,
|
||||
):
|
||||
yield cur
|
||||
else:
|
||||
with self.lock, conn.cursor(binary=True, row_factory=dict_row) as cur:
|
||||
yield cur
|
||||
|
||||
|
||||
class AsyncShallowPostgresSaver(BasePostgresSaver):
|
||||
"""A checkpoint saver that uses Postgres to store checkpoints asynchronously.
|
||||
|
||||
This checkpointer ONLY stores the most recent checkpoint and does NOT retain any history.
|
||||
It is meant to be a light-weight drop-in replacement for the AsyncPostgresSaver that
|
||||
supports most of the LangGraph persistence functionality with the exception of time travel.
|
||||
"""
|
||||
|
||||
SELECT_SQL = SELECT_SQL
|
||||
MIGRATIONS = MIGRATIONS
|
||||
UPSERT_CHECKPOINT_BLOBS_SQL = UPSERT_CHECKPOINT_BLOBS_SQL
|
||||
UPSERT_CHECKPOINTS_SQL = UPSERT_CHECKPOINTS_SQL
|
||||
UPSERT_CHECKPOINT_WRITES_SQL = UPSERT_CHECKPOINT_WRITES_SQL
|
||||
INSERT_CHECKPOINT_WRITES_SQL = INSERT_CHECKPOINT_WRITES_SQL
|
||||
lock: asyncio.Lock
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
conn: _ainternal.Conn,
|
||||
pipe: Optional[AsyncPipeline] = None,
|
||||
serde: Optional[SerializerProtocol] = None,
|
||||
) -> None:
|
||||
super().__init__(serde=serde)
|
||||
if isinstance(conn, AsyncConnectionPool) and pipe is not None:
|
||||
raise ValueError(
|
||||
"Pipeline should be used only with a single AsyncConnection, not AsyncConnectionPool."
|
||||
)
|
||||
|
||||
self.conn = conn
|
||||
self.pipe = pipe
|
||||
self.lock = asyncio.Lock()
|
||||
self.loop = asyncio.get_running_loop()
|
||||
self.supports_pipeline = Capabilities().has_pipeline()
|
||||
|
||||
@classmethod
|
||||
@asynccontextmanager
|
||||
async def from_conn_string(
|
||||
cls,
|
||||
conn_string: str,
|
||||
*,
|
||||
pipeline: bool = False,
|
||||
serde: Optional[SerializerProtocol] = None,
|
||||
) -> AsyncIterator["AsyncShallowPostgresSaver"]:
|
||||
"""Create a new AsyncShallowPostgresSaver instance from a connection string.
|
||||
|
||||
Args:
|
||||
conn_string (str): The Postgres connection info string.
|
||||
pipeline (bool): whether to use AsyncPipeline
|
||||
|
||||
Returns:
|
||||
AsyncShallowPostgresSaver: A new AsyncShallowPostgresSaver instance.
|
||||
"""
|
||||
async with await AsyncConnection.connect(
|
||||
conn_string, autocommit=True, prepare_threshold=0, row_factory=dict_row
|
||||
) as conn:
|
||||
if pipeline:
|
||||
async with conn.pipeline() as pipe:
|
||||
yield cls(conn=conn, pipe=pipe, serde=serde)
|
||||
else:
|
||||
yield cls(conn=conn, serde=serde)
|
||||
|
||||
async def setup(self) -> None:
|
||||
"""Set up the checkpoint database asynchronously.
|
||||
|
||||
This method creates the necessary tables in the Postgres database if they don't
|
||||
already exist and runs database migrations. It MUST be called directly by the user
|
||||
the first time checkpointer is used.
|
||||
"""
|
||||
async with self._cursor() as cur:
|
||||
await cur.execute(self.MIGRATIONS[0])
|
||||
results = await cur.execute(
|
||||
"SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1"
|
||||
)
|
||||
row = await results.fetchone()
|
||||
if row is None:
|
||||
version = -1
|
||||
else:
|
||||
version = row["v"]
|
||||
for v, migration in zip(
|
||||
range(version + 1, len(self.MIGRATIONS)),
|
||||
self.MIGRATIONS[version + 1 :],
|
||||
):
|
||||
await cur.execute(migration)
|
||||
await cur.execute(f"INSERT INTO checkpoint_migrations (v) VALUES ({v})")
|
||||
if self.pipe:
|
||||
await self.pipe.sync()
|
||||
|
||||
async def alist(
|
||||
self,
|
||||
config: Optional[RunnableConfig],
|
||||
*,
|
||||
filter: Optional[dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> AsyncIterator[CheckpointTuple]:
|
||||
"""List checkpoints from the database asynchronously.
|
||||
|
||||
This method retrieves a list of checkpoint tuples from the Postgres database based
|
||||
on the provided config. For ShallowPostgresSaver, this method returns a list with
|
||||
ONLY the most recent checkpoint.
|
||||
"""
|
||||
where, args = self._search_where(config, filter, before)
|
||||
query = self.SELECT_SQL + where
|
||||
if limit:
|
||||
query += f" LIMIT {limit}"
|
||||
async with self._cursor() as cur:
|
||||
await cur.execute(self.SELECT_SQL + where, args, binary=True)
|
||||
async for value in cur:
|
||||
checkpoint = await asyncio.to_thread(
|
||||
self._load_checkpoint,
|
||||
value["checkpoint"],
|
||||
value["channel_values"],
|
||||
value["pending_sends"],
|
||||
)
|
||||
yield CheckpointTuple(
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": value["thread_id"],
|
||||
"checkpoint_ns": value["checkpoint_ns"],
|
||||
"checkpoint_id": checkpoint["id"],
|
||||
}
|
||||
},
|
||||
checkpoint=checkpoint,
|
||||
metadata=self._load_metadata(value["metadata"]),
|
||||
pending_writes=await asyncio.to_thread(
|
||||
self._load_writes, value["pending_writes"]
|
||||
),
|
||||
)
|
||||
|
||||
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||
"""Get a checkpoint tuple from the database asynchronously.
|
||||
|
||||
This method retrieves a checkpoint tuple from the Postgres database based on the
|
||||
provided config (matching the thread ID in the config).
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): The config to use for retrieving the checkpoint.
|
||||
|
||||
Returns:
|
||||
Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found.
|
||||
"""
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
args = (thread_id, checkpoint_ns)
|
||||
where = "WHERE thread_id = %s AND checkpoint_ns = %s"
|
||||
|
||||
async with self._cursor() as cur:
|
||||
await cur.execute(
|
||||
self.SELECT_SQL + where,
|
||||
args,
|
||||
binary=True,
|
||||
)
|
||||
|
||||
async for value in cur:
|
||||
checkpoint = await asyncio.to_thread(
|
||||
self._load_checkpoint,
|
||||
value["checkpoint"],
|
||||
value["channel_values"],
|
||||
value["pending_sends"],
|
||||
)
|
||||
return CheckpointTuple(
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"checkpoint_id": checkpoint["id"],
|
||||
}
|
||||
},
|
||||
checkpoint=checkpoint,
|
||||
metadata=self._load_metadata(value["metadata"]),
|
||||
pending_writes=await asyncio.to_thread(
|
||||
self._load_writes, value["pending_writes"]
|
||||
),
|
||||
)
|
||||
|
||||
async def aput(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
"""Save a checkpoint to the database asynchronously.
|
||||
|
||||
This method saves a checkpoint to the Postgres database. The checkpoint is associated
|
||||
with the provided config. For AsyncShallowPostgresSaver, this method saves ONLY the most recent
|
||||
checkpoint and overwrites a previous checkpoint, if it exists.
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): The config to associate with the checkpoint.
|
||||
checkpoint (Checkpoint): The checkpoint to save.
|
||||
metadata (CheckpointMetadata): Additional metadata to save with the checkpoint.
|
||||
new_versions (ChannelVersions): New channel versions as of this write.
|
||||
|
||||
Returns:
|
||||
RunnableConfig: Updated configuration after storing the checkpoint.
|
||||
"""
|
||||
configurable = config["configurable"].copy()
|
||||
thread_id = configurable.pop("thread_id")
|
||||
checkpoint_ns = configurable.pop("checkpoint_ns")
|
||||
|
||||
copy = checkpoint.copy()
|
||||
next_config = {
|
||||
"configurable": {
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"checkpoint_id": checkpoint["id"],
|
||||
}
|
||||
}
|
||||
|
||||
async with self._cursor(pipeline=True) as cur:
|
||||
await cur.execute(
|
||||
"""DELETE FROM checkpoint_writes
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s AND checkpoint_id NOT IN (%s, %s)""",
|
||||
(
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
checkpoint["id"],
|
||||
configurable.get("checkpoint_id", ""),
|
||||
),
|
||||
)
|
||||
await cur.executemany(
|
||||
self.UPSERT_CHECKPOINT_BLOBS_SQL,
|
||||
_dump_blobs(
|
||||
self.serde,
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
copy.pop("channel_values"), # type: ignore[misc]
|
||||
new_versions,
|
||||
),
|
||||
)
|
||||
await cur.execute(
|
||||
self.UPSERT_CHECKPOINTS_SQL,
|
||||
(
|
||||
thread_id,
|
||||
checkpoint_ns,
|
||||
Jsonb(self._dump_checkpoint(copy)),
|
||||
self._dump_metadata(metadata),
|
||||
),
|
||||
)
|
||||
return next_config
|
||||
|
||||
async def aput_writes(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
writes: Sequence[tuple[str, Any]],
|
||||
task_id: str,
|
||||
) -> None:
|
||||
"""Store intermediate writes linked to a checkpoint asynchronously.
|
||||
|
||||
This method saves intermediate writes associated with a checkpoint to the database.
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): Configuration of the related checkpoint.
|
||||
writes (Sequence[Tuple[str, Any]]): List of writes to store, each as (channel, value) pair.
|
||||
task_id (str): Identifier for the task creating the writes.
|
||||
"""
|
||||
query = (
|
||||
self.UPSERT_CHECKPOINT_WRITES_SQL
|
||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||
else self.INSERT_CHECKPOINT_WRITES_SQL
|
||||
)
|
||||
params = await asyncio.to_thread(
|
||||
self._dump_writes,
|
||||
config["configurable"]["thread_id"],
|
||||
config["configurable"]["checkpoint_ns"],
|
||||
config["configurable"]["checkpoint_id"],
|
||||
task_id,
|
||||
writes,
|
||||
)
|
||||
async with self._cursor(pipeline=True) as cur:
|
||||
await cur.executemany(query, params)
|
||||
|
||||
@asynccontextmanager
|
||||
async def _cursor(
|
||||
self, *, pipeline: bool = False
|
||||
) -> AsyncIterator[AsyncCursor[DictRow]]:
|
||||
"""Create a database cursor as a context manager.
|
||||
|
||||
Args:
|
||||
pipeline (bool): whether to use pipeline for the DB operations inside the context manager.
|
||||
Will be applied regardless of whether the AsyncShallowPostgresSaver instance was initialized with a pipeline.
|
||||
If pipeline mode is not supported, will fall back to using transaction context manager.
|
||||
"""
|
||||
async with _ainternal.get_connection(self.conn) as conn:
|
||||
if self.pipe:
|
||||
# a connection in pipeline mode can be used concurrently
|
||||
# in multiple threads/coroutines, but only one cursor can be
|
||||
# used at a time
|
||||
try:
|
||||
async with conn.cursor(binary=True, row_factory=dict_row) as cur:
|
||||
yield cur
|
||||
finally:
|
||||
if pipeline:
|
||||
await self.pipe.sync()
|
||||
elif pipeline:
|
||||
# a connection not in pipeline mode can only be used by one
|
||||
# thread/coroutine at a time, so we acquire a lock
|
||||
if self.supports_pipeline:
|
||||
async with (
|
||||
self.lock,
|
||||
conn.pipeline(),
|
||||
conn.cursor(binary=True, row_factory=dict_row) as cur,
|
||||
):
|
||||
yield cur
|
||||
else:
|
||||
# Use connection's transaction context manager when pipeline mode not supported
|
||||
async with (
|
||||
self.lock,
|
||||
conn.transaction(),
|
||||
conn.cursor(binary=True, row_factory=dict_row) as cur,
|
||||
):
|
||||
yield cur
|
||||
else:
|
||||
async with (
|
||||
self.lock,
|
||||
conn.cursor(binary=True, row_factory=dict_row) as cur,
|
||||
):
|
||||
yield cur
|
||||
|
||||
def list(
|
||||
self,
|
||||
config: Optional[RunnableConfig],
|
||||
*,
|
||||
filter: Optional[dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> Iterator[CheckpointTuple]:
|
||||
"""List checkpoints from the database.
|
||||
|
||||
This method retrieves a list of checkpoint tuples from the Postgres database based
|
||||
on the provided config. For ShallowPostgresSaver, this method returns a list with
|
||||
ONLY the most recent checkpoint.
|
||||
"""
|
||||
aiter_ = self.alist(config, filter=filter, before=before, limit=limit)
|
||||
while True:
|
||||
try:
|
||||
yield asyncio.run_coroutine_threadsafe(
|
||||
anext(aiter_), # noqa: F821
|
||||
self.loop,
|
||||
).result()
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
|
||||
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||
"""Get a checkpoint tuple from the database.
|
||||
|
||||
This method retrieves a checkpoint tuple from the Postgres database based on the
|
||||
provided config (matching the thread ID in the config).
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): The config to use for retrieving the checkpoint.
|
||||
|
||||
Returns:
|
||||
Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found.
|
||||
"""
|
||||
try:
|
||||
# check if we are in the main thread, only bg threads can block
|
||||
# we don't check in other methods to avoid the overhead
|
||||
if asyncio.get_running_loop() is self.loop:
|
||||
raise asyncio.InvalidStateError(
|
||||
"Synchronous calls to AsyncShallowPostgresSaver are only allowed from a "
|
||||
"different thread. From the main thread, use the async interface."
|
||||
"For example, use `await checkpointer.aget_tuple(...)` or `await "
|
||||
"graph.ainvoke(...)`."
|
||||
)
|
||||
except RuntimeError:
|
||||
pass
|
||||
return asyncio.run_coroutine_threadsafe(
|
||||
self.aget_tuple(config), self.loop
|
||||
).result()
|
||||
|
||||
def put(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
"""Save a checkpoint to the database.
|
||||
|
||||
This method saves a checkpoint to the Postgres database. The checkpoint is associated
|
||||
with the provided config. For AsyncShallowPostgresSaver, this method saves ONLY the most recent
|
||||
checkpoint and overwrites a previous checkpoint, if it exists.
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): The config to associate with the checkpoint.
|
||||
checkpoint (Checkpoint): The checkpoint to save.
|
||||
metadata (CheckpointMetadata): Additional metadata to save with the checkpoint.
|
||||
new_versions (ChannelVersions): New channel versions as of this write.
|
||||
|
||||
Returns:
|
||||
RunnableConfig: Updated configuration after storing the checkpoint.
|
||||
"""
|
||||
return asyncio.run_coroutine_threadsafe(
|
||||
self.aput(config, checkpoint, metadata, new_versions), self.loop
|
||||
).result()
|
||||
|
||||
def put_writes(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
writes: Sequence[tuple[str, Any]],
|
||||
task_id: str,
|
||||
) -> None:
|
||||
"""Store intermediate writes linked to a checkpoint.
|
||||
|
||||
This method saves intermediate writes associated with a checkpoint to the database.
|
||||
|
||||
Args:
|
||||
config (RunnableConfig): Configuration of the related checkpoint.
|
||||
writes (Sequence[Tuple[str, Any]]): List of writes to store, each as (channel, value) pair.
|
||||
task_id (str): Identifier for the task creating the writes.
|
||||
"""
|
||||
return asyncio.run_coroutine_threadsafe(
|
||||
self.aput_writes(config, writes, task_id), self.loop
|
||||
).result()
|
||||
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "langgraph-checkpoint-postgres"
|
||||
version = "2.0.8"
|
||||
version = "2.0.9"
|
||||
description = "Library with a Postgres implementation of LangGraph checkpoint saver."
|
||||
authors = []
|
||||
license = "MIT"
|
||||
|
||||
@@ -16,7 +16,10 @@ from langgraph.checkpoint.base import (
|
||||
create_checkpoint,
|
||||
empty_checkpoint,
|
||||
)
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
from langgraph.checkpoint.postgres.aio import (
|
||||
AsyncPostgresSaver,
|
||||
AsyncShallowPostgresSaver,
|
||||
)
|
||||
from tests.conftest import DEFAULT_POSTGRES_URI
|
||||
|
||||
|
||||
@@ -103,11 +106,41 @@ async def _base_saver():
|
||||
await conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _shallow_saver():
|
||||
"""Fixture for shallow connection mode testing."""
|
||||
database = f"test_{uuid4().hex[:16]}"
|
||||
# create unique db
|
||||
async with await AsyncConnection.connect(
|
||||
DEFAULT_POSTGRES_URI, autocommit=True
|
||||
) as conn:
|
||||
await conn.execute(f"CREATE DATABASE {database}")
|
||||
try:
|
||||
async with await AsyncConnection.connect(
|
||||
DEFAULT_POSTGRES_URI + database,
|
||||
autocommit=True,
|
||||
prepare_threshold=0,
|
||||
row_factory=dict_row,
|
||||
) as conn:
|
||||
checkpointer = AsyncShallowPostgresSaver(conn)
|
||||
await checkpointer.setup()
|
||||
yield checkpointer
|
||||
finally:
|
||||
# drop unique db
|
||||
async with await AsyncConnection.connect(
|
||||
DEFAULT_POSTGRES_URI, autocommit=True
|
||||
) as conn:
|
||||
await conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _saver(name: str):
|
||||
if name == "base":
|
||||
async with _base_saver() as saver:
|
||||
yield saver
|
||||
elif name == "shallow":
|
||||
async with _shallow_saver() as saver:
|
||||
yield saver
|
||||
elif name == "pool":
|
||||
async with _pool_saver() as saver:
|
||||
yield saver
|
||||
@@ -167,7 +200,7 @@ def test_data():
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe", "shallow"])
|
||||
async def test_asearch(request, saver_name: str, test_data) -> None:
|
||||
async with _saver(saver_name) as saver:
|
||||
configs = test_data["configs"]
|
||||
@@ -212,7 +245,7 @@ async def test_asearch(request, saver_name: str, test_data) -> None:
|
||||
} == {"", "inner"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe", "shallow"])
|
||||
async def test_null_chars(request, saver_name: str, test_data) -> None:
|
||||
async with _saver(saver_name) as saver:
|
||||
config = await saver.aput(
|
||||
|
||||
@@ -16,7 +16,7 @@ from langgraph.checkpoint.base import (
|
||||
create_checkpoint,
|
||||
empty_checkpoint,
|
||||
)
|
||||
from langgraph.checkpoint.postgres import PostgresSaver
|
||||
from langgraph.checkpoint.postgres import PostgresSaver, ShallowPostgresSaver
|
||||
from tests.conftest import DEFAULT_POSTGRES_URI
|
||||
|
||||
|
||||
@@ -91,11 +91,37 @@ def _base_saver():
|
||||
conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _shallow_saver():
|
||||
"""Fixture for regular connection mode testing with a shallow checkpointer."""
|
||||
database = f"test_{uuid4().hex[:16]}"
|
||||
# create unique db
|
||||
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
|
||||
conn.execute(f"CREATE DATABASE {database}")
|
||||
try:
|
||||
with Connection.connect(
|
||||
DEFAULT_POSTGRES_URI + database,
|
||||
autocommit=True,
|
||||
prepare_threshold=0,
|
||||
row_factory=dict_row,
|
||||
) as conn:
|
||||
checkpointer = ShallowPostgresSaver(conn)
|
||||
checkpointer.setup()
|
||||
yield checkpointer
|
||||
finally:
|
||||
# drop unique db
|
||||
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
|
||||
conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _saver(name: str):
|
||||
if name == "base":
|
||||
with _base_saver() as saver:
|
||||
yield saver
|
||||
elif name == "shallow":
|
||||
with _shallow_saver() as saver:
|
||||
yield saver
|
||||
elif name == "pool":
|
||||
with _pool_saver() as saver:
|
||||
yield saver
|
||||
@@ -155,7 +181,7 @@ def test_data():
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe", "shallow"])
|
||||
def test_search(saver_name: str, test_data) -> None:
|
||||
with _saver(saver_name) as saver:
|
||||
configs = test_data["configs"]
|
||||
@@ -198,7 +224,7 @@ def test_search(saver_name: str, test_data) -> None:
|
||||
} == {"", "inner"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
|
||||
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe", "shallow"])
|
||||
def test_null_chars(saver_name: str, test_data) -> None:
|
||||
with _saver(saver_name) as saver:
|
||||
config = saver.put(
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import operator
|
||||
from typing import Annotated, TypedDict
|
||||
from typing import Annotated
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.constants import END, START, Send
|
||||
from langgraph.graph.state import StateGraph
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import asyncio
|
||||
import concurrent
|
||||
import concurrent.futures
|
||||
import functools
|
||||
import inspect
|
||||
import types
|
||||
from functools import partial, update_wrapper
|
||||
from typing import (
|
||||
Any,
|
||||
Awaitable,
|
||||
@@ -33,17 +33,17 @@ T = TypeVar("T")
|
||||
|
||||
|
||||
def call(
|
||||
func: Callable[[P1], T],
|
||||
input: P1,
|
||||
*,
|
||||
func: Callable[P, T],
|
||||
*args: Any,
|
||||
retry: Optional[RetryPolicy] = None,
|
||||
**kwargs: Any,
|
||||
) -> concurrent.futures.Future[T]:
|
||||
from langgraph.constants import CONFIG_KEY_CALL
|
||||
from langgraph.utils.config import get_configurable
|
||||
|
||||
conf = get_configurable()
|
||||
impl = conf[CONFIG_KEY_CALL]
|
||||
fut = impl(func, input, retry=retry)
|
||||
fut = impl(func, (args, kwargs), retry=retry)
|
||||
return fut
|
||||
|
||||
|
||||
@@ -59,16 +59,51 @@ def task( # type: ignore[overload-cannot-match]
|
||||
) -> Callable[[Callable[P, T]], Callable[P, concurrent.futures.Future[T]]]: ...
|
||||
|
||||
|
||||
@overload
|
||||
def task(
|
||||
*, retry: Optional[RetryPolicy] = None
|
||||
__func_or_none__: Callable[P, T],
|
||||
) -> Callable[P, concurrent.futures.Future[T]]: ...
|
||||
|
||||
|
||||
@overload
|
||||
def task(
|
||||
__func_or_none__: Callable[P, Awaitable[T]],
|
||||
) -> Callable[P, asyncio.Future[T]]: ...
|
||||
|
||||
|
||||
def task(
|
||||
__func_or_none__: Optional[Union[Callable[P, T], Callable[P, Awaitable[T]]]] = None,
|
||||
*,
|
||||
retry: Optional[RetryPolicy] = None,
|
||||
) -> Union[
|
||||
Callable[[Callable[P, Awaitable[T]]], Callable[P, asyncio.Future[T]]],
|
||||
Callable[[Callable[P, T]], Callable[P, concurrent.futures.Future[T]]],
|
||||
Callable[P, asyncio.Future[T]],
|
||||
Callable[P, concurrent.futures.Future[T]],
|
||||
]:
|
||||
def _task(func: Callable[P, T]) -> Callable[P, concurrent.futures.Future[T]]:
|
||||
return update_wrapper(partial(call, func, retry=retry), func)
|
||||
def decorator(
|
||||
func: Union[Callable[P, Awaitable[T]], Callable[P, T]],
|
||||
) -> Callable[P, concurrent.futures.Future[T]]:
|
||||
if asyncio.iscoroutinefunction(func):
|
||||
|
||||
return _task
|
||||
@functools.wraps(func)
|
||||
async def _tick(__allargs__: tuple) -> T:
|
||||
return await func(*__allargs__[0], **__allargs__[1])
|
||||
|
||||
else:
|
||||
|
||||
@functools.wraps(func)
|
||||
def _tick(__allargs__: tuple) -> T:
|
||||
return func(*__allargs__[0], **__allargs__[1])
|
||||
|
||||
return functools.update_wrapper(
|
||||
functools.partial(call, _tick, retry=retry), func
|
||||
)
|
||||
|
||||
if __func_or_none__ is not None:
|
||||
return decorator(__func_or_none__)
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def entrypoint(
|
||||
|
||||
@@ -8,7 +8,6 @@ from typing import (
|
||||
Literal,
|
||||
Optional,
|
||||
Sequence,
|
||||
TypedDict,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
@@ -22,6 +21,7 @@ from langchain_core.messages import (
|
||||
convert_to_messages,
|
||||
message_chunk_to_message,
|
||||
)
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.graph.state import StateGraph
|
||||
|
||||
|
||||
@@ -381,7 +381,7 @@ def create_react_agent(
|
||||
Add complex prompt with custom graph state:
|
||||
|
||||
```pycon
|
||||
>>> from typing import TypedDict
|
||||
>>> from typing_extensions import TypedDict
|
||||
>>>
|
||||
>>> from langgraph.managed import IsLastStep
|
||||
>>> prompt = ChatPromptTemplate.from_messages(
|
||||
|
||||
@@ -601,7 +601,8 @@ def tools_condition(
|
||||
>>> from langgraph.prebuilt import ToolNode, tools_condition
|
||||
>>> from langgraph.graph.message import add_messages
|
||||
...
|
||||
>>> from typing import TypedDict, Annotated
|
||||
>>> from typing import Annotated
|
||||
>>> from typing_extensions import TypedDict
|
||||
...
|
||||
>>> @tool
|
||||
>>> def divide(a: float, b: float) -> int:
|
||||
|
||||
@@ -74,7 +74,8 @@ class ValidationNode(RunnableCallable):
|
||||
|
||||
Examples:
|
||||
Example usage for re-prompting the model to generate a valid response:
|
||||
>>> from typing import Literal, Annotated, TypedDict
|
||||
>>> from typing import Literal, Annotated
|
||||
>>> from typing_extensions import TypedDict
|
||||
...
|
||||
>>> from langchain_anthropic import ChatAnthropic
|
||||
>>> from pydantic import BaseModel, validator
|
||||
|
||||
@@ -1139,6 +1139,10 @@ class Pregel(PregelProtocol):
|
||||
values: dict[str, Any] | Any,
|
||||
as_node: Optional[str] = None,
|
||||
) -> RunnableConfig:
|
||||
"""Update the state of the graph asynchronously with the given values, as if they came from
|
||||
node `as_node`. If `as_node` is not provided, it will be set to the last node
|
||||
that updated the state, if not ambiguous.
|
||||
"""
|
||||
checkpointer: Optional[BaseCheckpointSaver] = ensure_config(config)[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINTER, self.checkpointer
|
||||
)
|
||||
|
||||
@@ -10,13 +10,13 @@ from typing import (
|
||||
Mapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
TypedDict,
|
||||
Union,
|
||||
)
|
||||
from uuid import UUID
|
||||
|
||||
from langchain_core.runnables.config import RunnableConfig
|
||||
from langchain_core.utils.input import get_bolded_text, get_colored_text
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.checkpoint.base import Checkpoint, CheckpointMetadata, PendingWrite
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from collections import Counter
|
||||
from typing import Any, Iterator, Literal, Mapping, Optional, Sequence, TypeVar, Union
|
||||
from uuid import UUID
|
||||
|
||||
@@ -181,12 +182,27 @@ def map_output_updates(
|
||||
(task.name, value) for chan, value in writes if chan == output_channels
|
||||
)
|
||||
elif any(chan in output_channels for chan, _ in writes):
|
||||
updated.append(
|
||||
(
|
||||
task.name,
|
||||
{chan: value for chan, value in writes if chan in output_channels},
|
||||
counts = Counter(chan for chan, _ in writes)
|
||||
if any(counts[chan] > 1 for chan in output_channels):
|
||||
updated.extend(
|
||||
(
|
||||
task.name,
|
||||
{chan: value},
|
||||
)
|
||||
for chan, value in writes
|
||||
if chan in output_channels
|
||||
)
|
||||
else:
|
||||
updated.append(
|
||||
(
|
||||
task.name,
|
||||
{
|
||||
chan: value
|
||||
for chan, value in writes
|
||||
if chan in output_channels
|
||||
},
|
||||
)
|
||||
)
|
||||
)
|
||||
grouped: dict[str, list[Any]] = {t.name: [] for t, _ in output_tasks}
|
||||
for node, value in updated:
|
||||
grouped[node].append(value)
|
||||
|
||||
@@ -1032,6 +1032,13 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
# unwind stack
|
||||
return await asyncio.shield(
|
||||
exit_task = asyncio.create_task(
|
||||
self.stack.__aexit__(exc_type, exc_value, traceback)
|
||||
)
|
||||
try:
|
||||
return await exit_task
|
||||
except asyncio.CancelledError as e:
|
||||
# Bubble up the exit task upon cancellation to permit the API
|
||||
# consumer to await it before e.g., re-using the DB connection.
|
||||
e.args = (*e.args, exit_task)
|
||||
raise
|
||||
|
||||
@@ -13,14 +13,13 @@ from typing import (
|
||||
Optional,
|
||||
Sequence,
|
||||
Type,
|
||||
TypedDict,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from typing_extensions import Self
|
||||
from typing_extensions import Self, TypedDict
|
||||
|
||||
from langgraph.checkpoint.base import (
|
||||
BaseCheckpointSaver,
|
||||
@@ -373,7 +372,8 @@ def interrupt(value: Any) -> Any:
|
||||
Example:
|
||||
```python
|
||||
import uuid
|
||||
from typing import TypedDict, Optional
|
||||
from typing import Optional
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.checkpoint.memory import MemorySaver
|
||||
from langgraph.constants import START
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "langgraph"
|
||||
version = "0.2.60"
|
||||
version = "0.2.61"
|
||||
description = "Building stateful, multi-actor applications with LLMs"
|
||||
authors = []
|
||||
license = "MIT"
|
||||
@@ -38,7 +38,7 @@ py-spy = "^0.3.14"
|
||||
types-requests = "^2.32.0.20240914"
|
||||
|
||||
[tool.ruff]
|
||||
lint.select = [ "E", "F", "I" ]
|
||||
lint.select = [ "E", "F", "I", "TID251" ]
|
||||
lint.ignore = [ "E501" ]
|
||||
line-length = 88
|
||||
indent-width = 4
|
||||
@@ -52,6 +52,9 @@ line-ending = "auto"
|
||||
docstring-code-format = false
|
||||
docstring-code-line-length = "dynamic"
|
||||
|
||||
[tool.ruff.lint.flake8-tidy-imports.banned-api]
|
||||
"typing.TypedDict".msg = "Use typing_extensions.TypedDict instead."
|
||||
|
||||
[tool.mypy]
|
||||
# https://mypy.readthedocs.io/en/stable/config_file.html
|
||||
disallow_untyped_defs = "True"
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1,151 @@
|
||||
# serializer version: 1
|
||||
# name: test_weather_subgraph[memory]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
graph TD;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
router_node(router_node)
|
||||
normal_llm_node(normal_llm_node)
|
||||
weather_graph_model_node(model_node)
|
||||
weather_graph_weather_node(weather_node<hr/><small><em>__interrupt = before</em></small>)
|
||||
__end__([<p>__end__</p>]):::last
|
||||
__start__ --> router_node;
|
||||
normal_llm_node --> __end__;
|
||||
weather_graph_weather_node --> __end__;
|
||||
router_node -.-> normal_llm_node;
|
||||
router_node -.-> weather_graph_model_node;
|
||||
router_node -.-> __end__;
|
||||
subgraph weather_graph
|
||||
weather_graph_model_node --> weather_graph_weather_node;
|
||||
end
|
||||
classDef default fill:#f2f0ff,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#bfb6fc
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_weather_subgraph[postgres_aio]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
graph TD;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
router_node(router_node)
|
||||
normal_llm_node(normal_llm_node)
|
||||
weather_graph_model_node(model_node)
|
||||
weather_graph_weather_node(weather_node<hr/><small><em>__interrupt = before</em></small>)
|
||||
__end__([<p>__end__</p>]):::last
|
||||
__start__ --> router_node;
|
||||
normal_llm_node --> __end__;
|
||||
weather_graph_weather_node --> __end__;
|
||||
router_node -.-> normal_llm_node;
|
||||
router_node -.-> weather_graph_model_node;
|
||||
router_node -.-> __end__;
|
||||
subgraph weather_graph
|
||||
weather_graph_model_node --> weather_graph_weather_node;
|
||||
end
|
||||
classDef default fill:#f2f0ff,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#bfb6fc
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_weather_subgraph[postgres_aio_pipe]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
graph TD;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
router_node(router_node)
|
||||
normal_llm_node(normal_llm_node)
|
||||
weather_graph_model_node(model_node)
|
||||
weather_graph_weather_node(weather_node<hr/><small><em>__interrupt = before</em></small>)
|
||||
__end__([<p>__end__</p>]):::last
|
||||
__start__ --> router_node;
|
||||
normal_llm_node --> __end__;
|
||||
weather_graph_weather_node --> __end__;
|
||||
router_node -.-> normal_llm_node;
|
||||
router_node -.-> weather_graph_model_node;
|
||||
router_node -.-> __end__;
|
||||
subgraph weather_graph
|
||||
weather_graph_model_node --> weather_graph_weather_node;
|
||||
end
|
||||
classDef default fill:#f2f0ff,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#bfb6fc
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_weather_subgraph[postgres_aio_pool]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
graph TD;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
router_node(router_node)
|
||||
normal_llm_node(normal_llm_node)
|
||||
weather_graph_model_node(model_node)
|
||||
weather_graph_weather_node(weather_node<hr/><small><em>__interrupt = before</em></small>)
|
||||
__end__([<p>__end__</p>]):::last
|
||||
__start__ --> router_node;
|
||||
normal_llm_node --> __end__;
|
||||
weather_graph_weather_node --> __end__;
|
||||
router_node -.-> normal_llm_node;
|
||||
router_node -.-> weather_graph_model_node;
|
||||
router_node -.-> __end__;
|
||||
subgraph weather_graph
|
||||
weather_graph_model_node --> weather_graph_weather_node;
|
||||
end
|
||||
classDef default fill:#f2f0ff,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#bfb6fc
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_weather_subgraph[postgres_aio_shallow]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
graph TD;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
router_node(router_node)
|
||||
normal_llm_node(normal_llm_node)
|
||||
weather_graph_model_node(model_node)
|
||||
weather_graph_weather_node(weather_node<hr/><small><em>__interrupt = before</em></small>)
|
||||
__end__([<p>__end__</p>]):::last
|
||||
__start__ --> router_node;
|
||||
normal_llm_node --> __end__;
|
||||
weather_graph_weather_node --> __end__;
|
||||
router_node -.-> normal_llm_node;
|
||||
router_node -.-> weather_graph_model_node;
|
||||
router_node -.-> __end__;
|
||||
subgraph weather_graph
|
||||
weather_graph_model_node --> weather_graph_weather_node;
|
||||
end
|
||||
classDef default fill:#f2f0ff,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#bfb6fc
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_weather_subgraph[sqlite_aio]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
graph TD;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
router_node(router_node)
|
||||
normal_llm_node(normal_llm_node)
|
||||
weather_graph_model_node(model_node)
|
||||
weather_graph_weather_node(weather_node<hr/><small><em>__interrupt = before</em></small>)
|
||||
__end__([<p>__end__</p>]):::last
|
||||
__start__ --> router_node;
|
||||
normal_llm_node --> __end__;
|
||||
weather_graph_weather_node --> __end__;
|
||||
router_node -.-> normal_llm_node;
|
||||
router_node -.-> weather_graph_model_node;
|
||||
router_node -.-> __end__;
|
||||
subgraph weather_graph
|
||||
weather_graph_model_node --> weather_graph_weather_node;
|
||||
end
|
||||
classDef default fill:#f2f0ff,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#bfb6fc
|
||||
|
||||
'''
|
||||
# ---
|
||||
@@ -2878,6 +2878,19 @@
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge[postgres_shallow]
|
||||
'''
|
||||
graph TD;
|
||||
__start__ --> rewrite_query;
|
||||
analyzer_one --> retriever_one;
|
||||
qa --> __end__;
|
||||
retriever_one --> qa;
|
||||
retriever_two --> qa;
|
||||
rewrite_query --> analyzer_one;
|
||||
rewrite_query --> retriever_two;
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge[sqlite]
|
||||
'''
|
||||
graph TD;
|
||||
@@ -3311,6 +3324,76 @@
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[postgres_shallow]
|
||||
'''
|
||||
graph TD;
|
||||
__start__ --> rewrite_query;
|
||||
analyzer_one --> retriever_one;
|
||||
qa --> __end__;
|
||||
retriever_one --> qa;
|
||||
retriever_two --> qa;
|
||||
rewrite_query --> analyzer_one;
|
||||
rewrite_query -.-> retriever_two;
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[postgres_shallow].1
|
||||
dict({
|
||||
'definitions': dict({
|
||||
'InnerObject': dict({
|
||||
'properties': dict({
|
||||
'yo': dict({
|
||||
'title': 'Yo',
|
||||
'type': 'integer',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'yo',
|
||||
]),
|
||||
'title': 'InnerObject',
|
||||
'type': 'object',
|
||||
}),
|
||||
}),
|
||||
'properties': dict({
|
||||
'inner': dict({
|
||||
'$ref': '#/definitions/InnerObject',
|
||||
}),
|
||||
'query': dict({
|
||||
'title': 'Query',
|
||||
'type': 'string',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'query',
|
||||
'inner',
|
||||
]),
|
||||
'title': 'Input',
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[postgres_shallow].2
|
||||
dict({
|
||||
'properties': dict({
|
||||
'answer': dict({
|
||||
'title': 'Answer',
|
||||
'type': 'string',
|
||||
}),
|
||||
'docs': dict({
|
||||
'items': dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
'title': 'Docs',
|
||||
'type': 'array',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'answer',
|
||||
'docs',
|
||||
]),
|
||||
'title': 'Output',
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[sqlite]
|
||||
'''
|
||||
graph TD;
|
||||
@@ -3788,6 +3871,76 @@
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[postgres_shallow]
|
||||
'''
|
||||
graph TD;
|
||||
__start__ --> rewrite_query;
|
||||
analyzer_one --> retriever_one;
|
||||
qa --> __end__;
|
||||
retriever_one --> qa;
|
||||
retriever_two --> qa;
|
||||
rewrite_query --> analyzer_one;
|
||||
rewrite_query -.-> retriever_two;
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[postgres_shallow].1
|
||||
dict({
|
||||
'$defs': dict({
|
||||
'InnerObject': dict({
|
||||
'properties': dict({
|
||||
'yo': dict({
|
||||
'title': 'Yo',
|
||||
'type': 'integer',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'yo',
|
||||
]),
|
||||
'title': 'InnerObject',
|
||||
'type': 'object',
|
||||
}),
|
||||
}),
|
||||
'properties': dict({
|
||||
'inner': dict({
|
||||
'$ref': '#/$defs/InnerObject',
|
||||
}),
|
||||
'query': dict({
|
||||
'title': 'Query',
|
||||
'type': 'string',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'query',
|
||||
'inner',
|
||||
]),
|
||||
'title': 'Input',
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[postgres_shallow].2
|
||||
dict({
|
||||
'properties': dict({
|
||||
'answer': dict({
|
||||
'title': 'Answer',
|
||||
'type': 'string',
|
||||
}),
|
||||
'docs': dict({
|
||||
'items': dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
'title': 'Docs',
|
||||
'type': 'array',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'answer',
|
||||
'docs',
|
||||
]),
|
||||
'title': 'Output',
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[sqlite]
|
||||
'''
|
||||
graph TD;
|
||||
@@ -3923,6 +4076,19 @@
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_via_branch[postgres_shallow]
|
||||
'''
|
||||
graph TD;
|
||||
__start__ --> rewrite_query;
|
||||
analyzer_one --> retriever_one;
|
||||
qa --> __end__;
|
||||
retriever_one --> qa;
|
||||
retriever_two --> qa;
|
||||
rewrite_query --> analyzer_one;
|
||||
rewrite_query -.-> retriever_two;
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_via_branch[sqlite]
|
||||
'''
|
||||
graph TD;
|
||||
|
||||
@@ -934,6 +934,127 @@
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[postgres_aio_shallow]
|
||||
'''
|
||||
graph TD;
|
||||
__start__ --> rewrite_query;
|
||||
analyzer_one --> retriever_one;
|
||||
qa --> __end__;
|
||||
retriever_one --> qa;
|
||||
retriever_two --> qa;
|
||||
rewrite_query --> analyzer_one;
|
||||
rewrite_query -.-> retriever_two;
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[postgres_aio_shallow].1
|
||||
dict({
|
||||
'$defs': dict({
|
||||
'InnerObject': dict({
|
||||
'properties': dict({
|
||||
'yo': dict({
|
||||
'title': 'Yo',
|
||||
'type': 'integer',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'yo',
|
||||
]),
|
||||
'title': 'InnerObject',
|
||||
'type': 'object',
|
||||
}),
|
||||
}),
|
||||
'properties': dict({
|
||||
'answer': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Answer',
|
||||
}),
|
||||
'docs': dict({
|
||||
'items': dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
'title': 'Docs',
|
||||
'type': 'array',
|
||||
}),
|
||||
'inner': dict({
|
||||
'$ref': '#/$defs/InnerObject',
|
||||
}),
|
||||
'query': dict({
|
||||
'title': 'Query',
|
||||
'type': 'string',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'query',
|
||||
'inner',
|
||||
'docs',
|
||||
]),
|
||||
'title': 'State',
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[postgres_aio_shallow].2
|
||||
dict({
|
||||
'$defs': dict({
|
||||
'InnerObject': dict({
|
||||
'properties': dict({
|
||||
'yo': dict({
|
||||
'title': 'Yo',
|
||||
'type': 'integer',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'yo',
|
||||
]),
|
||||
'title': 'InnerObject',
|
||||
'type': 'object',
|
||||
}),
|
||||
}),
|
||||
'properties': dict({
|
||||
'answer': dict({
|
||||
'anyOf': list([
|
||||
dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
dict({
|
||||
'type': 'null',
|
||||
}),
|
||||
]),
|
||||
'default': None,
|
||||
'title': 'Answer',
|
||||
}),
|
||||
'docs': dict({
|
||||
'items': dict({
|
||||
'type': 'string',
|
||||
}),
|
||||
'title': 'Docs',
|
||||
'type': 'array',
|
||||
}),
|
||||
'inner': dict({
|
||||
'$ref': '#/$defs/InnerObject',
|
||||
}),
|
||||
'query': dict({
|
||||
'title': 'Query',
|
||||
'type': 'string',
|
||||
}),
|
||||
}),
|
||||
'required': list([
|
||||
'query',
|
||||
'inner',
|
||||
'docs',
|
||||
]),
|
||||
'title': 'State',
|
||||
'type': 'object',
|
||||
})
|
||||
# ---
|
||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[sqlite_aio]
|
||||
'''
|
||||
graph TD;
|
||||
@@ -1362,6 +1483,21 @@
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_send_react_interrupt_control[postgres_aio_shallow]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
graph TD;
|
||||
__start__([<p>__start__</p>]):::first
|
||||
agent(agent)
|
||||
foo([foo]):::last
|
||||
__start__ --> agent;
|
||||
agent -.-> foo;
|
||||
classDef default fill:#f2f0ff,line-height:1.2
|
||||
classDef first fill-opacity:0
|
||||
classDef last fill:#bfb6fc
|
||||
|
||||
'''
|
||||
# ---
|
||||
# name: test_send_react_interrupt_control[sqlite_aio]
|
||||
'''
|
||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||
|
||||
@@ -13,8 +13,11 @@ from pytest_mock import MockerFixture
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.checkpoint.duckdb import DuckDBSaver
|
||||
from langgraph.checkpoint.duckdb.aio import AsyncDuckDBSaver
|
||||
from langgraph.checkpoint.postgres import PostgresSaver
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
from langgraph.checkpoint.postgres import PostgresSaver, ShallowPostgresSaver
|
||||
from langgraph.checkpoint.postgres.aio import (
|
||||
AsyncPostgresSaver,
|
||||
AsyncShallowPostgresSaver,
|
||||
)
|
||||
from langgraph.checkpoint.sqlite import SqliteSaver
|
||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
||||
from langgraph.store.base import BaseStore
|
||||
@@ -100,6 +103,25 @@ def checkpointer_postgres():
|
||||
conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def checkpointer_postgres_shallow():
|
||||
database = f"test_{uuid4().hex[:16]}"
|
||||
# create unique db
|
||||
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
|
||||
conn.execute(f"CREATE DATABASE {database}")
|
||||
try:
|
||||
# yield checkpointer
|
||||
with ShallowPostgresSaver.from_conn_string(
|
||||
DEFAULT_POSTGRES_URI + database
|
||||
) as checkpointer:
|
||||
checkpointer.setup()
|
||||
yield checkpointer
|
||||
finally:
|
||||
# drop unique db
|
||||
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
|
||||
conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def checkpointer_postgres_pipe():
|
||||
database = f"test_{uuid4().hex[:16]}"
|
||||
@@ -167,6 +189,31 @@ async def _checkpointer_postgres_aio():
|
||||
await conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _checkpointer_postgres_aio_shallow():
|
||||
if sys.version_info < (3, 10):
|
||||
pytest.skip("Async Postgres tests require Python 3.10+")
|
||||
database = f"test_{uuid4().hex[:16]}"
|
||||
# create unique db
|
||||
async with await AsyncConnection.connect(
|
||||
DEFAULT_POSTGRES_URI, autocommit=True
|
||||
) as conn:
|
||||
await conn.execute(f"CREATE DATABASE {database}")
|
||||
try:
|
||||
# yield checkpointer
|
||||
async with AsyncShallowPostgresSaver.from_conn_string(
|
||||
DEFAULT_POSTGRES_URI + database
|
||||
) as checkpointer:
|
||||
await checkpointer.setup()
|
||||
yield checkpointer
|
||||
finally:
|
||||
# drop unique db
|
||||
async with await AsyncConnection.connect(
|
||||
DEFAULT_POSTGRES_URI, autocommit=True
|
||||
) as conn:
|
||||
await conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _checkpointer_postgres_aio_pipe():
|
||||
if sys.version_info < (3, 10):
|
||||
@@ -240,6 +287,9 @@ async def awith_checkpointer(
|
||||
elif checkpointer_name == "postgres_aio":
|
||||
async with _checkpointer_postgres_aio() as checkpointer:
|
||||
yield checkpointer
|
||||
elif checkpointer_name == "postgres_aio_shallow":
|
||||
async with _checkpointer_postgres_aio_shallow() as checkpointer:
|
||||
yield checkpointer
|
||||
elif checkpointer_name == "postgres_aio_pipe":
|
||||
async with _checkpointer_postgres_aio_pipe() as checkpointer:
|
||||
yield checkpointer
|
||||
@@ -417,20 +467,30 @@ async def awith_store(store_name: Optional[str]) -> AsyncIterator[BaseStore]:
|
||||
raise NotImplementedError(f"Unknown store {store_name}")
|
||||
|
||||
|
||||
ALL_CHECKPOINTERS_SYNC = [
|
||||
SHALLOW_CHECKPOINTERS_SYNC = ["postgres_shallow"]
|
||||
REGULAR_CHECKPOINTERS_SYNC = [
|
||||
"memory",
|
||||
"sqlite",
|
||||
"postgres",
|
||||
"postgres_pipe",
|
||||
"postgres_pool",
|
||||
]
|
||||
ALL_CHECKPOINTERS_ASYNC = [
|
||||
ALL_CHECKPOINTERS_SYNC = [
|
||||
*REGULAR_CHECKPOINTERS_SYNC,
|
||||
*SHALLOW_CHECKPOINTERS_SYNC,
|
||||
]
|
||||
SHALLOW_CHECKPOINTERS_ASYNC = ["postgres_aio_shallow"]
|
||||
REGULAR_CHECKPOINTERS_ASYNC = [
|
||||
"memory",
|
||||
"sqlite_aio",
|
||||
"postgres_aio",
|
||||
"postgres_aio_pipe",
|
||||
"postgres_aio_pool",
|
||||
]
|
||||
ALL_CHECKPOINTERS_ASYNC = [
|
||||
*REGULAR_CHECKPOINTERS_ASYNC,
|
||||
*SHALLOW_CHECKPOINTERS_ASYNC,
|
||||
]
|
||||
ALL_CHECKPOINTERS_ASYNC_PLUS_NONE = [
|
||||
*ALL_CHECKPOINTERS_ASYNC,
|
||||
None,
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from typing import TypedDict
|
||||
|
||||
import pytest
|
||||
from pytest_mock import MockerFixture
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from tests.conftest import (
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -21,7 +21,6 @@ from typing import (
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
TypedDict,
|
||||
Union,
|
||||
get_type_hints,
|
||||
)
|
||||
@@ -36,6 +35,7 @@ from langchain_core.runnables import (
|
||||
from langsmith import traceable
|
||||
from pytest_mock import MockerFixture
|
||||
from syrupy import SnapshotAssertion
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||
@@ -78,6 +78,7 @@ from tests.any_str import AnyStr, AnyVersion, FloatBetween, UnsortedSequence
|
||||
from tests.conftest import (
|
||||
ALL_CHECKPOINTERS_SYNC,
|
||||
ALL_STORES_SYNC,
|
||||
REGULAR_CHECKPOINTERS_SYNC,
|
||||
SHOULD_CHECK_SNAPSHOTS,
|
||||
)
|
||||
from tests.memory_assert import MemorySaverAssertCheckpointMetadata
|
||||
@@ -624,7 +625,7 @@ def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
|
||||
assert step == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_SYNC)
|
||||
def test_run_from_checkpoint_id_retains_previous_writes(
|
||||
request: pytest.FixtureRequest, checkpointer_name: str, mocker: MockerFixture
|
||||
) -> None:
|
||||
@@ -1157,6 +1158,10 @@ def test_pending_writes_resume(
|
||||
# both the pending write and the new write were applied, 1 + 2 + 3 = 6
|
||||
assert graph.invoke(None, thread1) == {"value": 6}
|
||||
|
||||
if "shallow" in checkpointer_name:
|
||||
assert len(list(checkpointer.list(thread1))) == 1
|
||||
return
|
||||
|
||||
# check all final checkpoints
|
||||
checkpoints = [c for c in checkpointer.list(thread1)]
|
||||
# we should have 3
|
||||
@@ -1510,27 +1515,32 @@ def test_imp_stream_order(
|
||||
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
|
||||
|
||||
@task()
|
||||
def foo(state: dict) -> dict:
|
||||
return {"a": state["a"] + "foo", "b": "bar"}
|
||||
def foo(state: dict) -> tuple:
|
||||
return state["a"] + "foo", "bar"
|
||||
|
||||
@task()
|
||||
def bar(state: dict) -> dict:
|
||||
return {"a": state["a"] + state["b"], "c": "bark"}
|
||||
@task
|
||||
def bar(a: str, b: str, c: Optional[str] = None) -> dict:
|
||||
return {"a": a + b, "c": (c or "") + "bark"}
|
||||
|
||||
@task()
|
||||
@task
|
||||
def baz(state: dict) -> dict:
|
||||
return {"a": state["a"] + "baz", "c": "something else"}
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
def graph(state: dict) -> dict:
|
||||
fut_foo = foo(state)
|
||||
fut_bar = bar(fut_foo.result())
|
||||
fut_bar = bar(*fut_foo.result())
|
||||
fut_baz = baz(fut_bar.result())
|
||||
return fut_baz.result()
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert [c for c in graph.stream({"a": "0"}, thread1)] == [
|
||||
{"foo": {"a": "0foo", "b": "bar"}},
|
||||
{
|
||||
"foo": (
|
||||
"0foo",
|
||||
"bar",
|
||||
)
|
||||
},
|
||||
{"bar": {"a": "0foobar", "c": "bark"}},
|
||||
{"baz": {"a": "0foobarbaz", "c": "something else"}},
|
||||
{"graph": {"a": "0foobarbaz", "c": "something else"}},
|
||||
@@ -1618,6 +1628,9 @@ def test_invoke_checkpoint_three(
|
||||
assert state.values.get("total") == 5
|
||||
assert state.next == ()
|
||||
|
||||
if "shallow" in checkpointer_name:
|
||||
return
|
||||
|
||||
assert len(list(app.get_state_history(thread_1, limit=1))) == 1
|
||||
# list all checkpoints for thread 1
|
||||
thread_1_history = [c for c in app.get_state_history(thread_1)]
|
||||
@@ -2270,6 +2283,11 @@ def test_in_one_fan_out_state_graph_waiting_edge(
|
||||
]
|
||||
|
||||
app_w_interrupt.update_state(config, {"docs": ["doc5"]})
|
||||
expected_parent_config = (
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else list(app_w_interrupt.checkpointer.list(config, limit=2))[-1].config
|
||||
)
|
||||
assert app_w_interrupt.get_state(config) == StateSnapshot(
|
||||
values={
|
||||
"query": "analyzed: query: what is weather in sf",
|
||||
@@ -2277,8 +2295,14 @@ def test_in_one_fan_out_state_graph_waiting_edge(
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "qa", (PULL, "qa")),),
|
||||
next=("qa",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
created_at=app_w_interrupt.checkpointer.get_tuple(config).checkpoint["ts"],
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
metadata={
|
||||
"parents": {},
|
||||
"source": "update",
|
||||
@@ -2286,7 +2310,7 @@ def test_in_one_fan_out_state_graph_waiting_edge(
|
||||
"writes": {"retriever_one": {"docs": ["doc5"]}},
|
||||
"thread_id": "2",
|
||||
},
|
||||
parent_config=[*app_w_interrupt.checkpointer.list(config, limit=2)][-1].config,
|
||||
parent_config=expected_parent_config,
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config, debug=1)] == [
|
||||
@@ -4149,10 +4173,12 @@ def test_store_injected(
|
||||
def __call__(self, inputs: State, config: RunnableConfig, store: BaseStore):
|
||||
assert isinstance(store, BaseStore)
|
||||
store.put(
|
||||
namespace
|
||||
if self.i is not None
|
||||
and config["configurable"]["thread_id"] in (thread_1, thread_2)
|
||||
else (f"foo_{self.i}", "bar"),
|
||||
(
|
||||
namespace
|
||||
if self.i is not None
|
||||
and config["configurable"]["thread_id"] in (thread_1, thread_2)
|
||||
else (f"foo_{self.i}", "bar")
|
||||
),
|
||||
doc_id,
|
||||
{
|
||||
**doc,
|
||||
@@ -4670,13 +4696,17 @@ def test_parent_command(request: pytest.FixtureRequest, checkpointer_name: str)
|
||||
"parents": {},
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(),
|
||||
)
|
||||
|
||||
@@ -5203,11 +5233,15 @@ def test_checkpoint_recovery(request: pytest.FixtureRequest, checkpointer_name:
|
||||
assert state is not None
|
||||
assert state.values == {"steps": ["start"], "attempt": 1} # input state saved
|
||||
assert state.next == ("node1",) # Should retry failed node
|
||||
assert "RuntimeError('Simulated failure')" in state.tasks[0].error
|
||||
|
||||
# Retry with updated attempt count
|
||||
result = graph.invoke({"steps": [], "attempt": 2}, config)
|
||||
assert result == {"steps": ["start", "node1", "node2"], "attempt": 2}
|
||||
|
||||
if "shallow" in checkpointer_name:
|
||||
return
|
||||
|
||||
# Verify checkpoint history shows both attempts
|
||||
history = list(graph.get_state_history(config))
|
||||
assert len(history) == 6 # Initial + failed attempt + successful attempt
|
||||
@@ -5215,3 +5249,54 @@ def test_checkpoint_recovery(request: pytest.FixtureRequest, checkpointer_name:
|
||||
# Verify the error was recorded in checkpoint
|
||||
failed_checkpoint = next(c for c in history if c.tasks and c.tasks[0].error)
|
||||
assert "RuntimeError('Simulated failure')" in failed_checkpoint.tasks[0].error
|
||||
|
||||
|
||||
def test_multiple_updates_root() -> None:
|
||||
def node_a(state):
|
||||
return [Command(update="a1"), Command(update="a2")]
|
||||
|
||||
def node_b(state):
|
||||
return "b"
|
||||
|
||||
graph = (
|
||||
StateGraph(Annotated[str, operator.add])
|
||||
.add_sequence([node_a, node_b])
|
||||
.add_edge(START, "node_a")
|
||||
.compile()
|
||||
)
|
||||
|
||||
assert graph.invoke("") == "a1a2b"
|
||||
|
||||
# only streams the last update from node_a
|
||||
assert [c for c in graph.stream("", stream_mode="updates")] == [
|
||||
{"node_a": ["a1", "a2"]},
|
||||
{"node_b": "b"},
|
||||
]
|
||||
|
||||
|
||||
def test_multiple_updates() -> None:
|
||||
class State(TypedDict):
|
||||
foo: Annotated[str, operator.add]
|
||||
|
||||
def node_a(state):
|
||||
return [Command(update={"foo": "a1"}), Command(update={"foo": "a2"})]
|
||||
|
||||
def node_b(state):
|
||||
return {"foo": "b"}
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_sequence([node_a, node_b])
|
||||
.add_edge(START, "node_a")
|
||||
.compile()
|
||||
)
|
||||
|
||||
assert graph.invoke({"foo": ""}) == {
|
||||
"foo": "a1a2b",
|
||||
}
|
||||
|
||||
# only streams the last update from node_a
|
||||
assert [c for c in graph.stream({"foo": ""}, stream_mode="updates")] == [
|
||||
{"node_a": [{"foo": "a1"}, {"foo": "a2"}]},
|
||||
{"node_b": {"foo": "b"}},
|
||||
]
|
||||
|
||||
@@ -19,7 +19,6 @@ from typing import (
|
||||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
TypedDict,
|
||||
Union,
|
||||
)
|
||||
from uuid import UUID
|
||||
@@ -34,6 +33,7 @@ from langchain_core.runnables import (
|
||||
from langchain_core.utils.aiter import aclosing
|
||||
from pytest_mock import MockerFixture
|
||||
from syrupy import SnapshotAssertion
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||
@@ -75,6 +75,7 @@ from tests.conftest import (
|
||||
ALL_CHECKPOINTERS_ASYNC,
|
||||
ALL_CHECKPOINTERS_ASYNC_PLUS_NONE,
|
||||
ALL_STORES_ASYNC,
|
||||
REGULAR_CHECKPOINTERS_ASYNC,
|
||||
SHOULD_CHECK_SNAPSHOTS,
|
||||
awith_checkpointer,
|
||||
awith_store,
|
||||
@@ -179,6 +180,262 @@ async def test_checkpoint_errors() -> None:
|
||||
pass
|
||||
|
||||
|
||||
async def test_py_async_with_cancel_behavior() -> None:
|
||||
"""This test confirms that in all versions of Python we support, __aexit__
|
||||
is not cancelled when the coroutine containing the async with block is cancelled."""
|
||||
|
||||
logs: list[str] = []
|
||||
|
||||
class MyContextManager:
|
||||
async def __aenter__(self):
|
||||
logs.append("Entering")
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
logs.append("Starting exit")
|
||||
try:
|
||||
# Simulate some cleanup work
|
||||
await asyncio.sleep(2)
|
||||
logs.append("Cleanup completed")
|
||||
except asyncio.CancelledError:
|
||||
logs.append("Cleanup was cancelled!")
|
||||
raise
|
||||
logs.append("Exit finished")
|
||||
|
||||
async def main():
|
||||
try:
|
||||
async with MyContextManager():
|
||||
logs.append("In context")
|
||||
await asyncio.sleep(1)
|
||||
logs.append("This won't print if cancelled")
|
||||
except asyncio.CancelledError:
|
||||
logs.append("Context was cancelled")
|
||||
raise
|
||||
|
||||
# create task
|
||||
t = asyncio.create_task(main())
|
||||
# cancel after 0.2 seconds
|
||||
await asyncio.sleep(0.2)
|
||||
t.cancel()
|
||||
# check logs before cancellation is handled
|
||||
assert logs == [
|
||||
"Entering",
|
||||
"In context",
|
||||
], "Cancelled before cleanup started"
|
||||
# wait for task to finish
|
||||
try:
|
||||
await t
|
||||
except asyncio.CancelledError:
|
||||
# check logs after cancellation is handled
|
||||
assert logs == [
|
||||
"Entering",
|
||||
"In context",
|
||||
"Starting exit",
|
||||
"Cleanup completed",
|
||||
"Exit finished",
|
||||
"Context was cancelled",
|
||||
], "Cleanup started and finished after cancellation"
|
||||
else:
|
||||
assert False, "Task should be cancelled"
|
||||
|
||||
|
||||
async def test_checkpoint_put_after_cancellation() -> None:
|
||||
logs: list[str] = []
|
||||
|
||||
class LongPutCheckpointer(MemorySaver):
|
||||
async def aput(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
logs.append("checkpoint.aput.start")
|
||||
try:
|
||||
await asyncio.sleep(1)
|
||||
return await super().aput(config, checkpoint, metadata, new_versions)
|
||||
finally:
|
||||
logs.append("checkpoint.aput.end")
|
||||
|
||||
inner_task_cancelled = False
|
||||
|
||||
async def awhile(input: Any) -> None:
|
||||
logs.append("awhile.start")
|
||||
try:
|
||||
await asyncio.sleep(1)
|
||||
except asyncio.CancelledError:
|
||||
nonlocal inner_task_cancelled
|
||||
inner_task_cancelled = True
|
||||
raise
|
||||
finally:
|
||||
logs.append("awhile.end")
|
||||
|
||||
builder = Graph()
|
||||
builder.add_node("agent", awhile)
|
||||
builder.set_entry_point("agent")
|
||||
builder.set_finish_point("agent")
|
||||
|
||||
graph = builder.compile(checkpointer=LongPutCheckpointer())
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# start the task
|
||||
t = asyncio.create_task(graph.ainvoke(1, thread1))
|
||||
# cancel after 0.2 seconds
|
||||
await asyncio.sleep(0.2)
|
||||
t.cancel()
|
||||
# check logs before cancellation is handled
|
||||
assert sorted(logs) == [
|
||||
"awhile.start",
|
||||
"checkpoint.aput.start",
|
||||
], "Cancelled before checkpoint put started"
|
||||
# wait for task to finish
|
||||
try:
|
||||
await t
|
||||
except asyncio.CancelledError:
|
||||
# check logs after cancellation is handled
|
||||
assert sorted(logs) == [
|
||||
"awhile.end",
|
||||
"awhile.start",
|
||||
"checkpoint.aput.end",
|
||||
"checkpoint.aput.start",
|
||||
], "Checkpoint put is not cancelled"
|
||||
else:
|
||||
assert False, "Task should be cancelled"
|
||||
|
||||
|
||||
async def test_checkpoint_put_after_cancellation_stream_anext() -> None:
|
||||
logs: list[str] = []
|
||||
|
||||
class LongPutCheckpointer(MemorySaver):
|
||||
async def aput(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
logs.append("checkpoint.aput.start")
|
||||
try:
|
||||
await asyncio.sleep(1)
|
||||
return await super().aput(config, checkpoint, metadata, new_versions)
|
||||
finally:
|
||||
logs.append("checkpoint.aput.end")
|
||||
|
||||
inner_task_cancelled = False
|
||||
|
||||
async def awhile(input: Any) -> None:
|
||||
logs.append("awhile.start")
|
||||
try:
|
||||
await asyncio.sleep(1)
|
||||
except asyncio.CancelledError:
|
||||
nonlocal inner_task_cancelled
|
||||
inner_task_cancelled = True
|
||||
raise
|
||||
finally:
|
||||
logs.append("awhile.end")
|
||||
|
||||
builder = Graph()
|
||||
builder.add_node("agent", awhile)
|
||||
builder.set_entry_point("agent")
|
||||
builder.set_finish_point("agent")
|
||||
|
||||
graph = builder.compile(checkpointer=LongPutCheckpointer())
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# start the task
|
||||
s = graph.astream(1, thread1)
|
||||
t = asyncio.create_task(s.__anext__())
|
||||
# cancel after 0.2 seconds
|
||||
await asyncio.sleep(0.2)
|
||||
t.cancel()
|
||||
# check logs before cancellation is handled
|
||||
assert sorted(logs) == [
|
||||
"awhile.start",
|
||||
"checkpoint.aput.start",
|
||||
], "Cancelled before checkpoint put started"
|
||||
# wait for task to finish
|
||||
try:
|
||||
await t
|
||||
except asyncio.CancelledError:
|
||||
# check logs after cancellation is handled
|
||||
assert sorted(logs) == [
|
||||
"awhile.end",
|
||||
"awhile.start",
|
||||
"checkpoint.aput.end",
|
||||
"checkpoint.aput.start",
|
||||
], "Checkpoint put is not cancelled"
|
||||
else:
|
||||
assert False, "Task should be cancelled"
|
||||
|
||||
|
||||
async def test_checkpoint_put_after_cancellation_stream_events_anext() -> None:
|
||||
logs: list[str] = []
|
||||
|
||||
class LongPutCheckpointer(MemorySaver):
|
||||
async def aput(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
logs.append("checkpoint.aput.start")
|
||||
try:
|
||||
await asyncio.sleep(1)
|
||||
return await super().aput(config, checkpoint, metadata, new_versions)
|
||||
finally:
|
||||
logs.append("checkpoint.aput.end")
|
||||
|
||||
inner_task_cancelled = False
|
||||
|
||||
async def awhile(input: Any) -> None:
|
||||
logs.append("awhile.start")
|
||||
try:
|
||||
await asyncio.sleep(1)
|
||||
except asyncio.CancelledError:
|
||||
nonlocal inner_task_cancelled
|
||||
inner_task_cancelled = True
|
||||
raise
|
||||
finally:
|
||||
logs.append("awhile.end")
|
||||
|
||||
builder = Graph()
|
||||
builder.add_node("agent", awhile)
|
||||
builder.set_entry_point("agent")
|
||||
builder.set_finish_point("agent")
|
||||
|
||||
graph = builder.compile(checkpointer=LongPutCheckpointer())
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# start the task
|
||||
s = graph.astream_events(1, thread1, version="v2", include_names=["LangGraph"])
|
||||
# skip first event (happens right away)
|
||||
await s.__anext__()
|
||||
# start the task for 2nd event
|
||||
t = asyncio.create_task(s.__anext__())
|
||||
# cancel after 0.2 seconds
|
||||
await asyncio.sleep(0.2)
|
||||
t.cancel()
|
||||
# check logs before cancellation is handled
|
||||
assert logs == [
|
||||
"checkpoint.aput.start",
|
||||
"awhile.start",
|
||||
], "Cancelled before checkpoint put started"
|
||||
# wait for task to finish
|
||||
try:
|
||||
await t
|
||||
except asyncio.CancelledError:
|
||||
# check logs after cancellation is handled
|
||||
assert logs == [
|
||||
"checkpoint.aput.start",
|
||||
"awhile.start",
|
||||
"awhile.end",
|
||||
"checkpoint.aput.end",
|
||||
], "Checkpoint put is not cancelled"
|
||||
else:
|
||||
assert False, "Task should be cancelled"
|
||||
|
||||
|
||||
async def test_node_cancellation_on_external_cancel() -> None:
|
||||
inner_task_cancelled = False
|
||||
|
||||
@@ -347,22 +604,23 @@ async def test_dynamic_interrupt(checkpointer_name: str) -> None:
|
||||
)
|
||||
},
|
||||
]
|
||||
assert [c.metadata async for c in tool_two.checkpointer.alist(thread1)] == [
|
||||
{
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 0,
|
||||
"writes": None,
|
||||
"thread_id": "1",
|
||||
},
|
||||
{
|
||||
"parents": {},
|
||||
"source": "input",
|
||||
"step": -1,
|
||||
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
]
|
||||
if "shallow" not in checkpointer_name:
|
||||
assert [c.metadata async for c in tool_two.checkpointer.alist(thread1)] == [
|
||||
{
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 0,
|
||||
"writes": None,
|
||||
"thread_id": "1",
|
||||
},
|
||||
{
|
||||
"parents": {},
|
||||
"source": "input",
|
||||
"step": -1,
|
||||
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
]
|
||||
tup = await tool_two.checkpointer.aget_tuple(thread1)
|
||||
assert await tool_two.aget_state(thread1) == StateSnapshot(
|
||||
values={"my_key": "value ⛰️", "market": "DE"},
|
||||
@@ -390,9 +648,13 @@ async def test_dynamic_interrupt(checkpointer_name: str) -> None:
|
||||
"writes": None,
|
||||
"thread_id": "1",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
|
||||
][-1].config,
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else [c async for c in tool_two.checkpointer.alist(thread1, limit=2)][
|
||||
-1
|
||||
].config
|
||||
),
|
||||
)
|
||||
|
||||
# clear the interrupt and next tasks
|
||||
@@ -412,9 +674,13 @@ async def test_dynamic_interrupt(checkpointer_name: str) -> None:
|
||||
"writes": {},
|
||||
"thread_id": "1",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
|
||||
][-1].config,
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else [c async for c in tool_two.checkpointer.alist(thread1, limit=2)][
|
||||
-1
|
||||
].config
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -524,22 +790,25 @@ async def test_dynamic_interrupt_subgraph(checkpointer_name: str) -> None:
|
||||
)
|
||||
},
|
||||
]
|
||||
assert [c.metadata async for c in tool_two.checkpointer.alist(thread1root)] == [
|
||||
{
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 0,
|
||||
"writes": None,
|
||||
"thread_id": "1",
|
||||
},
|
||||
{
|
||||
"parents": {},
|
||||
"source": "input",
|
||||
"step": -1,
|
||||
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
]
|
||||
if "shallow" not in checkpointer_name:
|
||||
assert [
|
||||
c.metadata async for c in tool_two.checkpointer.alist(thread1root)
|
||||
] == [
|
||||
{
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 0,
|
||||
"writes": None,
|
||||
"thread_id": "1",
|
||||
},
|
||||
{
|
||||
"parents": {},
|
||||
"source": "input",
|
||||
"step": -1,
|
||||
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
]
|
||||
tup = await tool_two.checkpointer.aget_tuple(thread1)
|
||||
assert await tool_two.aget_state(thread1) == StateSnapshot(
|
||||
values={"my_key": "value ⛰️", "market": "DE"},
|
||||
@@ -573,9 +842,13 @@ async def test_dynamic_interrupt_subgraph(checkpointer_name: str) -> None:
|
||||
"writes": None,
|
||||
"thread_id": "1",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in tool_two.checkpointer.alist(thread1root, limit=2)
|
||||
][-1].config,
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else [
|
||||
c async for c in tool_two.checkpointer.alist(thread1root, limit=2)
|
||||
][-1].config
|
||||
),
|
||||
)
|
||||
|
||||
# clear the interrupt and next tasks
|
||||
@@ -595,9 +868,13 @@ async def test_dynamic_interrupt_subgraph(checkpointer_name: str) -> None:
|
||||
"writes": {},
|
||||
"thread_id": "1",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in tool_two.checkpointer.alist(thread1root, limit=2)
|
||||
][-1].config,
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else [
|
||||
c async for c in tool_two.checkpointer.alist(thread1root, limit=2)
|
||||
][-1].config
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -699,22 +976,25 @@ async def test_copy_checkpoint(checkpointer_name: str) -> None:
|
||||
"my_key": "value ⛰️ one",
|
||||
"market": "DE",
|
||||
}
|
||||
assert [c.metadata async for c in tool_two.checkpointer.alist(thread1)] == [
|
||||
{
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 0,
|
||||
"writes": {"tool_one": {"my_key": " one"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
{
|
||||
"parents": {},
|
||||
"source": "input",
|
||||
"step": -1,
|
||||
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
]
|
||||
|
||||
if "shallow" not in checkpointer_name:
|
||||
assert [c.metadata async for c in tool_two.checkpointer.alist(thread1)] == [
|
||||
{
|
||||
"parents": {},
|
||||
"source": "loop",
|
||||
"step": 0,
|
||||
"writes": {"tool_one": {"my_key": " one"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
{
|
||||
"parents": {},
|
||||
"source": "input",
|
||||
"step": -1,
|
||||
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
]
|
||||
|
||||
tup = await tool_two.checkpointer.aget_tuple(thread1)
|
||||
assert await tool_two.aget_state(thread1) == StateSnapshot(
|
||||
values={"my_key": "value ⛰️ one", "market": "DE"},
|
||||
@@ -742,9 +1022,13 @@ async def test_copy_checkpoint(checkpointer_name: str) -> None:
|
||||
"writes": {"tool_one": {"my_key": " one"}},
|
||||
"thread_id": "1",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
|
||||
][-1].config,
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else [c async for c in tool_two.checkpointer.alist(thread1, limit=2)][
|
||||
-1
|
||||
].config
|
||||
),
|
||||
)
|
||||
# clear the interrupt and next tasks
|
||||
await tool_two.aupdate_state(thread1, None)
|
||||
@@ -770,9 +1054,13 @@ async def test_copy_checkpoint(checkpointer_name: str) -> None:
|
||||
"writes": {},
|
||||
"thread_id": "1",
|
||||
},
|
||||
parent_config=[
|
||||
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
|
||||
][-1].config,
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else [c async for c in tool_two.checkpointer.alist(thread1, limit=2)][
|
||||
-1
|
||||
].config
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -1754,6 +2042,10 @@ async def test_pending_writes_resume(
|
||||
# both the pending write and the new write were applied, 1 + 2 + 3 = 6
|
||||
assert await graph.ainvoke(None, thread1) == {"value": 6}
|
||||
|
||||
if "shallow" in checkpointer_name:
|
||||
assert len([c async for c in checkpointer.alist(thread1)]) == 1
|
||||
return
|
||||
|
||||
# check all final checkpoints
|
||||
checkpoints = [c async for c in checkpointer.alist(thread1)]
|
||||
# we should have 3
|
||||
@@ -1911,7 +2203,7 @@ async def test_pending_writes_resume(
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC)
|
||||
async def test_run_from_checkpoint_id_retains_previous_writes(
|
||||
request: pytest.FixtureRequest, checkpointer_name: str, mocker: MockerFixture
|
||||
) -> None:
|
||||
@@ -2279,9 +2571,9 @@ async def test_imp_sync_from_async(checkpointer_name: str) -> None:
|
||||
def foo(state: dict) -> dict:
|
||||
return {"a": state["a"] + "foo", "b": "bar"}
|
||||
|
||||
@task()
|
||||
def bar(state: dict) -> dict:
|
||||
return {"a": state["a"] + state["b"], "c": "bark"}
|
||||
@task
|
||||
def bar(a: str, b: str, c: Optional[str] = None) -> dict:
|
||||
return {"a": a + b, "c": (c or "") + "bark"}
|
||||
|
||||
@task()
|
||||
def baz(state: dict) -> dict:
|
||||
@@ -2289,8 +2581,8 @@ async def test_imp_sync_from_async(checkpointer_name: str) -> None:
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
def graph(state: dict) -> dict:
|
||||
fut_foo = foo(state)
|
||||
fut_bar = bar(fut_foo.result())
|
||||
foo_result = foo(state).result()
|
||||
fut_bar = bar(foo_result["a"], foo_result["b"])
|
||||
fut_baz = baz(fut_bar.result())
|
||||
return fut_baz.result()
|
||||
|
||||
@@ -2315,9 +2607,9 @@ async def test_imp_stream_order(checkpointer_name: str) -> None:
|
||||
async def foo(state: dict) -> dict:
|
||||
return {"a": state["a"] + "foo", "b": "bar"}
|
||||
|
||||
@task()
|
||||
async def bar(state: dict) -> dict:
|
||||
return {"a": state["a"] + state["b"], "c": "bark"}
|
||||
@task
|
||||
async def bar(a: str, b: str, c: Optional[str] = None) -> dict:
|
||||
return {"a": a + b, "c": (c or "") + "bark"}
|
||||
|
||||
@task()
|
||||
async def baz(state: dict) -> dict:
|
||||
@@ -2325,8 +2617,9 @@ async def test_imp_stream_order(checkpointer_name: str) -> None:
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
async def graph(state: dict) -> dict:
|
||||
fut_foo = foo(state)
|
||||
fut_bar = bar(await fut_foo)
|
||||
foo_res = await foo(state)
|
||||
|
||||
fut_bar = bar(foo_res["a"], foo_res["b"])
|
||||
fut_baz = baz(await fut_bar)
|
||||
return await fut_baz
|
||||
|
||||
@@ -2339,7 +2632,7 @@ async def test_imp_stream_order(checkpointer_name: str) -> None:
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC)
|
||||
async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
|
||||
if not FF_SEND_V2:
|
||||
pytest.skip("Send deduplication is only available in Send V2")
|
||||
@@ -2792,13 +3085,17 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
"thread_id": "2",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
@@ -2875,13 +3172,17 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
"thread_id": "2",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(),
|
||||
)
|
||||
|
||||
@@ -2951,13 +3252,17 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
"thread_id": "3",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
@@ -3060,13 +3365,17 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
"thread_id": "3",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "3",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
@@ -3260,13 +3569,17 @@ async def test_send_react_interrupt_control(
|
||||
"thread_id": "2",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
@@ -3343,13 +3656,17 @@ async def test_send_react_interrupt_control(
|
||||
"thread_id": "2",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "2",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(),
|
||||
)
|
||||
|
||||
@@ -3572,6 +3889,9 @@ async def test_invoke_checkpoint_three(
|
||||
assert state.values.get("total") == 5
|
||||
assert state.next == ()
|
||||
|
||||
if "shallow" in checkpointer_name:
|
||||
return
|
||||
|
||||
assert len([c async for c in app.aget_state_history(thread_1, limit=1)]) == 1
|
||||
# list all checkpoints for thread 1
|
||||
thread_1_history = [c async for c in app.aget_state_history(thread_1)]
|
||||
@@ -4279,13 +4599,17 @@ async def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class(
|
||||
"thread_id": "1",
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
async with assert_ctx_once():
|
||||
@@ -5758,13 +6082,17 @@ async def test_parent_command(checkpointer_name: str) -> None:
|
||||
"parents": {},
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
parent_config=(
|
||||
None
|
||||
if "shallow" in checkpointer_name
|
||||
else {
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
}
|
||||
},
|
||||
),
|
||||
tasks=(),
|
||||
)
|
||||
|
||||
@@ -6281,6 +6609,9 @@ async def test_checkpoint_recovery_async(checkpointer_name: str):
|
||||
result = await graph.ainvoke({"steps": [], "attempt": 2}, config)
|
||||
assert result == {"steps": ["start", "node1", "node2"], "attempt": 2}
|
||||
|
||||
if "shallow" in checkpointer_name:
|
||||
return
|
||||
|
||||
# Verify checkpoint history shows both attempts
|
||||
history = [c async for c in graph.aget_state_history(config)]
|
||||
assert len(history) == 6 # Initial + failed attempt + successful attempt
|
||||
@@ -6288,3 +6619,54 @@ async def test_checkpoint_recovery_async(checkpointer_name: str):
|
||||
# Verify the error was recorded in checkpoint
|
||||
failed_checkpoint = next(c for c in history if c.tasks and c.tasks[0].error)
|
||||
assert "RuntimeError('Simulated failure')" in failed_checkpoint.tasks[0].error
|
||||
|
||||
|
||||
async def test_multiple_updates_root() -> None:
|
||||
def node_a(state):
|
||||
return [Command(update="a1"), Command(update="a2")]
|
||||
|
||||
def node_b(state):
|
||||
return "b"
|
||||
|
||||
graph = (
|
||||
StateGraph(Annotated[str, operator.add])
|
||||
.add_sequence([node_a, node_b])
|
||||
.add_edge(START, "node_a")
|
||||
.compile()
|
||||
)
|
||||
|
||||
assert await graph.ainvoke("") == "a1a2b"
|
||||
|
||||
# only streams the last update from node_a
|
||||
assert [c async for c in graph.astream("", stream_mode="updates")] == [
|
||||
{"node_a": ["a1", "a2"]},
|
||||
{"node_b": "b"},
|
||||
]
|
||||
|
||||
|
||||
async def test_multiple_updates() -> None:
|
||||
class State(TypedDict):
|
||||
foo: Annotated[str, operator.add]
|
||||
|
||||
def node_a(state):
|
||||
return [Command(update={"foo": "a1"}), Command(update={"foo": "a2"})]
|
||||
|
||||
def node_b(state):
|
||||
return {"foo": "b"}
|
||||
|
||||
graph = (
|
||||
StateGraph(State)
|
||||
.add_sequence([node_a, node_b])
|
||||
.add_edge(START, "node_a")
|
||||
.compile()
|
||||
)
|
||||
|
||||
assert await graph.ainvoke({"foo": ""}) == {
|
||||
"foo": "a1a2b",
|
||||
}
|
||||
|
||||
# only streams the last update from node_a
|
||||
assert [c async for c in graph.astream({"foo": ""}, stream_mode="updates")] == [
|
||||
{"node_a": [{"foo": "a1"}, {"foo": "a2"}]},
|
||||
{"node_b": {"foo": "b"}},
|
||||
]
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
from typing import Any, Callable, Tuple, TypedDict, TypeVar
|
||||
from typing import Any, Callable, Tuple, TypeVar
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import langsmith as ls
|
||||
import pytest
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langchain_core.tracers import LangChainTracer
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.graph import StateGraph
|
||||
|
||||
|
||||
@@ -9,7 +9,6 @@ from typing import (
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
TypedDict,
|
||||
TypeVar,
|
||||
Union,
|
||||
)
|
||||
@@ -17,7 +16,7 @@ from unittest.mock import patch
|
||||
|
||||
import langsmith
|
||||
import pytest
|
||||
from typing_extensions import Annotated, NotRequired, Required
|
||||
from typing_extensions import Annotated, NotRequired, Required, TypedDict
|
||||
|
||||
from langgraph.graph import END, StateGraph
|
||||
from langgraph.graph.graph import CompiledGraph
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@langchain/langgraph-sdk",
|
||||
"version": "0.0.32",
|
||||
"version": "0.0.33",
|
||||
"description": "Client library for interacting with the LangGraph API",
|
||||
"type": "module",
|
||||
"packageManager": "yarn@1.22.19",
|
||||
|
||||
@@ -32,7 +32,7 @@ export interface Command {
|
||||
/**
|
||||
* An object to update the thread state with.
|
||||
*/
|
||||
update?: Record<string, unknown>;
|
||||
update?: Record<string, unknown> | [string, unknown][];
|
||||
|
||||
/**
|
||||
* The value to return from an `interrupt` function call.
|
||||
|
||||
@@ -461,6 +461,73 @@ class _CronsOn(
|
||||
Search = types.CronsSearch
|
||||
|
||||
|
||||
class _StoreOn:
|
||||
def __init__(self, auth: Auth) -> None:
|
||||
self._auth = auth
|
||||
|
||||
@typing.overload
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
actions: typing.Optional[
|
||||
typing.Union[
|
||||
typing.Literal["put", "get", "search", "list_namespaces", "delete"],
|
||||
Sequence[
|
||||
typing.Literal["put", "get", "search", "list_namespaces", "delete"]
|
||||
],
|
||||
]
|
||||
] = None,
|
||||
) -> Callable[[AHO], AHO]: ...
|
||||
|
||||
@typing.overload
|
||||
def __call__(self, fn: AHO) -> AHO: ...
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
fn: typing.Optional[AHO] = None,
|
||||
*,
|
||||
actions: typing.Optional[
|
||||
typing.Union[
|
||||
typing.Literal["put", "get", "search", "list_namespaces", "delete"],
|
||||
Sequence[
|
||||
typing.Literal["put", "get", "search", "list_namespaces", "delete"]
|
||||
],
|
||||
]
|
||||
] = None,
|
||||
) -> typing.Union[AHO, Callable[[AHO], AHO]]:
|
||||
"""Register a handler for specific resources and actions.
|
||||
|
||||
Can be used as a decorator or with explicit resource/action parameters:
|
||||
|
||||
@auth.on.store
|
||||
async def handler(): ... # Handle all store ops
|
||||
|
||||
@auth.on.store(actions=("put", "get", "search", "delete"))
|
||||
async def handler(): ... # Handle specific store ops
|
||||
|
||||
@auth.on.store.put
|
||||
async def handler(): ... # Handle store.put ops
|
||||
"""
|
||||
if fn is not None:
|
||||
# Used as a plain decorator
|
||||
_register_handler(self._auth, None, None, fn)
|
||||
return fn
|
||||
|
||||
# Used with parameters, return a decorator
|
||||
def decorator(
|
||||
handler: AHO,
|
||||
) -> AHO:
|
||||
if isinstance(actions, str):
|
||||
action_list = [actions]
|
||||
else:
|
||||
action_list = list(actions) if actions is not None else ["*"]
|
||||
for action in action_list:
|
||||
_register_handler(self._auth, "store", action, handler)
|
||||
return handler
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
AHO = typing.TypeVar("AHO", bound=_ActionHandler[dict[str, typing.Any]])
|
||||
|
||||
|
||||
@@ -524,6 +591,7 @@ class _On:
|
||||
"threads",
|
||||
"runs",
|
||||
"crons",
|
||||
"store",
|
||||
"value",
|
||||
)
|
||||
|
||||
@@ -532,6 +600,7 @@ class _On:
|
||||
self.assistants = _AssistantsOn(auth, "assistants")
|
||||
self.threads = _ThreadsOn(auth, "threads")
|
||||
self.crons = _CronsOn(auth, "crons")
|
||||
self.store = _StoreOn(auth)
|
||||
self.value = dict[str, typing.Any]
|
||||
|
||||
@typing.overload
|
||||
|
||||
@@ -5,7 +5,7 @@ request handling in LangGraph. It includes user protocols, authentication contex
|
||||
and typed dictionaries for various API operations.
|
||||
|
||||
Note:
|
||||
All typing.TypedDict classes use total=False to make all fields optional by default.
|
||||
All typing.TypedDict classes use total=False to make all fields typing.Optional by default.
|
||||
"""
|
||||
|
||||
import functools
|
||||
@@ -157,7 +157,7 @@ class MinimalUserDict(typing.TypedDict, total=False):
|
||||
identity: typing_extensions.Required[str]
|
||||
"""The required unique identifier for the user."""
|
||||
display_name: str
|
||||
"""The optional display name for the user."""
|
||||
"""The typing.Optional display name for the user."""
|
||||
is_authenticated: bool
|
||||
"""Whether the user is authenticated. Defaults to True."""
|
||||
permissions: Sequence[str]
|
||||
@@ -358,11 +358,34 @@ class AuthContext(BaseAuthContext):
|
||||
allowing for fine-grained access control decisions.
|
||||
"""
|
||||
|
||||
resource: typing.Literal["runs", "threads", "crons", "assistants"]
|
||||
resource: typing.Literal["runs", "threads", "crons", "assistants", "store"]
|
||||
"""The resource being accessed."""
|
||||
|
||||
action: typing.Literal["create", "read", "update", "delete", "search", "create_run"]
|
||||
"""The action being performed on the resource."""
|
||||
action: typing.Literal[
|
||||
"create",
|
||||
"read",
|
||||
"update",
|
||||
"delete",
|
||||
"search",
|
||||
"create_run",
|
||||
"put",
|
||||
"get",
|
||||
"list_namespaces",
|
||||
]
|
||||
"""The action being performed on the resource.
|
||||
|
||||
Most resources support the following actions:
|
||||
- create: Create a new resource
|
||||
- read: Read information about a resource
|
||||
- update: Update an existing resource
|
||||
- delete: Delete a resource
|
||||
- search: Search for resources
|
||||
|
||||
The store supports the following actions:
|
||||
- put: Add or update a document in the store
|
||||
- get: Get a document from the store
|
||||
- list_namespaces: List the namespaces in the store
|
||||
"""
|
||||
|
||||
|
||||
class ThreadsCreate(typing.TypedDict, total=False):
|
||||
@@ -759,6 +782,93 @@ class CronsSearch(typing.TypedDict, total=False):
|
||||
"""Offset for pagination."""
|
||||
|
||||
|
||||
class StoreGet(typing.TypedDict):
|
||||
"""Operation to retrieve a specific item by its namespace and key."""
|
||||
|
||||
namespace: tuple[str, ...]
|
||||
"""Hierarchical path that uniquely identifies the item's location."""
|
||||
|
||||
key: str
|
||||
"""Unique identifier for the item within its specific namespace."""
|
||||
|
||||
|
||||
class StoreSearch(typing.TypedDict):
|
||||
"""Operation to search for items within a specified namespace hierarchy."""
|
||||
|
||||
namespace_prefix: tuple[str, ...]
|
||||
"""Hierarchical path prefix defining the search scope.
|
||||
|
||||
???+ example "Examples"
|
||||
|
||||
```python
|
||||
() # Search entire store
|
||||
("documents",) # Search all documents
|
||||
("users", "content") # Search within user content
|
||||
```
|
||||
"""
|
||||
|
||||
filter: typing.Optional[dict[str, typing.Any]]
|
||||
"""Key-value pairs for filtering results based on exact matches or comparison operators."""
|
||||
|
||||
limit: int
|
||||
"""Maximum number of items to return in the search results."""
|
||||
|
||||
offset: int
|
||||
"""Number of matching items to skip for pagination."""
|
||||
|
||||
query: typing.Optional[str]
|
||||
"""Naturalj language search query for semantic search capabilities."""
|
||||
|
||||
|
||||
class StoreListNamespaces(typing.TypedDict):
|
||||
"""Operation to list and filter namespaces in the store."""
|
||||
|
||||
prefix: typing.Optional[tuple[str, ...]]
|
||||
"""Optional conditions for filtering namespaces."""
|
||||
|
||||
suffix: typing.Optional[tuple[str, ...]]
|
||||
"""Optional conditions for filtering namespaces."""
|
||||
|
||||
max_depth: typing.Optional[int]
|
||||
"""Maximum depth of namespace hierarchy to return.
|
||||
|
||||
Note:
|
||||
Namespaces deeper than this level will be truncated.
|
||||
"""
|
||||
|
||||
limit: int
|
||||
"""Maximum number of namespaces to return."""
|
||||
|
||||
offset: int
|
||||
"""Number of namespaces to skip for pagination."""
|
||||
|
||||
|
||||
class StorePut(typing.TypedDict):
|
||||
"""Operation to store, update, or delete an item in the store."""
|
||||
|
||||
namespace: tuple[str, ...]
|
||||
"""Hierarchical path that identifies the location of the item."""
|
||||
|
||||
key: str
|
||||
"""Unique identifier for the item within its namespace."""
|
||||
|
||||
value: typing.Optional[dict[str, typing.Any]]
|
||||
"""The data to store, or None to mark the item for deletion."""
|
||||
|
||||
index: typing.Optional[typing.Union[typing.Literal[False], list[str]]]
|
||||
"""Optional index configuration for full-text search."""
|
||||
|
||||
|
||||
class StoreDelete(typing.TypedDict):
|
||||
"""Operation to delete an item from the store."""
|
||||
|
||||
namespace: tuple[str, ...]
|
||||
"""Hierarchical path that uniquely identifies the item's location."""
|
||||
|
||||
key: str
|
||||
"""Unique identifier for the item within its specific namespace."""
|
||||
|
||||
|
||||
class on:
|
||||
"""Namespace for type definitions of different API operations.
|
||||
|
||||
@@ -894,6 +1004,38 @@ class on:
|
||||
|
||||
value = CronsSearch
|
||||
|
||||
class store:
|
||||
"""Types for store-related operations."""
|
||||
|
||||
value = typing.Union[
|
||||
StoreGet, StoreSearch, StoreListNamespaces, StorePut, StoreDelete
|
||||
]
|
||||
|
||||
class put:
|
||||
"""Type for store put parameters."""
|
||||
|
||||
value = StorePut
|
||||
|
||||
class get:
|
||||
"""Type for store get parameters."""
|
||||
|
||||
value = StoreGet
|
||||
|
||||
class search:
|
||||
"""Type for store search parameters."""
|
||||
|
||||
value = StoreSearch
|
||||
|
||||
class delete:
|
||||
"""Type for store delete parameters."""
|
||||
|
||||
value = StoreDelete
|
||||
|
||||
class list_namespaces:
|
||||
"""Type for store list namespaces parameters."""
|
||||
|
||||
value = StoreListNamespaces
|
||||
|
||||
|
||||
__all__ = [
|
||||
"on",
|
||||
@@ -909,4 +1051,9 @@ __all__ = [
|
||||
"AssistantsUpdate",
|
||||
"AssistantsDelete",
|
||||
"AssistantsSearch",
|
||||
"StoreGet",
|
||||
"StoreSearch",
|
||||
"StoreListNamespaces",
|
||||
"StorePut",
|
||||
"StoreDelete",
|
||||
]
|
||||
|
||||
@@ -1779,7 +1779,7 @@ class RunsClient:
|
||||
|
||||
Args:
|
||||
thread_id: The thread ID to cancel.
|
||||
run_id: The run ID to cancek.
|
||||
run_id: The run ID to cancel.
|
||||
wait: Whether to wait until run has completed.
|
||||
action: Action to take when cancelling the run. Possible values
|
||||
are `interrupt` or `rollback`. Default is `interrupt`.
|
||||
@@ -3917,7 +3917,7 @@ class SyncRunsClient:
|
||||
|
||||
Args:
|
||||
thread_id: The thread ID to cancel.
|
||||
run_id: The run ID to cancek.
|
||||
run_id: The run ID to cancel.
|
||||
wait: Whether to wait until run has completed.
|
||||
action: Action to take when cancelling the run. Possible values
|
||||
are `interrupt` or `rollback`. Default is `interrupt`.
|
||||
|
||||
@@ -1,7 +1,17 @@
|
||||
"""Data models for interacting with the LangGraph API."""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Literal, NamedTuple, Optional, Sequence, TypedDict, Union
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
TypedDict,
|
||||
Union,
|
||||
)
|
||||
|
||||
Json = Optional[dict[str, Any]]
|
||||
"""Represents a JSON-like structure, which can be None or a dictionary with string keys and any values."""
|
||||
@@ -374,5 +384,5 @@ class Send(TypedDict):
|
||||
|
||||
class Command(TypedDict, total=False):
|
||||
goto: Union[Send, str, Sequence[Union[Send, str]]]
|
||||
update: dict[str, Any]
|
||||
update: Union[dict[str, Any], Sequence[Tuple[str, Any]]]
|
||||
resume: Any
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "langgraph-sdk"
|
||||
version = "0.1.48"
|
||||
version = "0.1.50"
|
||||
description = "SDK for interacting with LangGraph API"
|
||||
authors = []
|
||||
license = "MIT"
|
||||
|
||||
Reference in New Issue
Block a user