Compare commits

..
15 Commits
Author SHA1 Message Date
Nuno Campos ac472357b7 Type safety w generics 2025-03-02 19:35:32 -08:00
Nuno Campos e6cdd4a0af Code review 2025-03-02 13:09:22 -08:00
Nuno Campos 0894f3e21e Implement Channel and Pregel 2025-03-01 22:43:42 -08:00
Nuno Campos 196bcfe08d java langgraph-checkpoint 2025-03-01 18:51:45 -08:00
Nuno Campos c4275bdc32 Add spec 2025-03-01 18:51:13 -08:00
Nuno Campos 1b9b0a686e Remove more mentions of async 2025-03-01 17:25:58 -08:00
Nuno Campos 9e02a23682 Rm other mentions of stream_mode=messages 2025-03-01 14:20:15 -08:00
Nuno Campos 5e70e6f307 Rm docs 2025-03-01 14:10:02 -08:00
Nuno Campos eb57c06896 Remove features and dependencies
- rm langchain_core dependency
- replace callbacks w run tree
- rm Runnable dependency
- rm non-state Graph
- rm managed values
- rm entrypoint/task/call
- rm async methods
- rm shallow checkpointer
- rm messages stream mode
- rm debug flag
- rm remote graph
2025-03-01 13:53:27 -08:00
Nuno Campos 9284b57ba0 Remove prebuilt 2025-03-01 10:30:18 -08:00
Nuno Campos 25fea591b5 Remove sqlite 2025-03-01 10:14:10 -08:00
Nuno Campos b9fe53777f Remove cli 2025-03-01 10:13:58 -08:00
Nuno Campos 35c2e8a679 Remove kafka 2025-03-01 10:13:47 -08:00
Nuno Campos e0fb56c6a3 Remove sdks 2025-03-01 10:13:36 -08:00
Nuno Campos 3458a3cecb Remove examples 2025-03-01 10:13:24 -08:00
498 changed files with 42147 additions and 116054 deletions
-6
View File
@@ -1,6 +0,0 @@
# Contributing to LangGraph
Hi there! Thank you for even being interested in contributing to LangGraph.
As an open-source project in a rapidly developing field, we are extremely open to contributions, whether they involve new features, improved infrastructure, better documentation, or bug fixes.
To learn how to contribute to LangGraph, please follow the [contribution guide here](https://docs.langchain.com/oss/python/contributing).
+15 -14
View File
@@ -1,28 +1,29 @@
name: "\U0001F41B Bug Report"
description: Report a bug in LangGraph. To report a security issue, please instead use the security option below. For questions, please use the LangChain Forum at forum.langchain.com.
labels: [pending, bug]
description: Report a bug in LangGraph. To report a security issue, please instead use the security option below. For questions, please use the GitHub Discussions.
labels: ["02 Bug Report"]
body:
- type: markdown
attributes:
value: |
value: >
Thank you for taking the time to file a bug report.
Use this to report BUGS in LangGraph. For usage questions, feature requests and general design questions, please use the [LangChain Forum](https://forum.langchain.com/).
Use this to report BUGS in LangGraph. For usage questions, feature requests and general design questions, please use [GitHub Discussions](https://github.com/langchain-ai/langgraph/discussions).
Relevant links to check before filing a bug report to see if your issue has already been reported, fixed or
if there's another way to solve your problem:
* [LangChain Forum](https://forum.langchain.com/),
* [LangGraph Github Issues](https://github.com/langchain-ai/langgraph/issues),
* [LangChain documentation with the integrated search](https://docs.langchain.com/),
* [GitHub search](https://github.com/langchain-ai/langgraph),
[LangGraph Github Discussions](https://github.com/langchain-ai/langgraph/discussions),
[LangGraph Github Issues](https://github.com/langchain-ai/langgraph/issues),
[LangGraph how-to guides](https://langchain-ai.github.io/langgraph/how-tos/).
[LangChain documentation with the integrated search](https://python.langchain.com/docs/get_started/introduction),
[GitHub search](https://github.com/langchain-ai/langgraph),
- type: checkboxes
id: checks
attributes:
label: Checked other resources
description: Before submitting this issue, please confirm that you have completed all the steps below by checking each option. These steps help ensure your issue is well-defined, relevant, and actionable.
options:
- label: This is a bug, not a usage question. For questions, please use the LangChain Forum (https://forum.langchain.com/).
- label: This is a bug, not a usage question. For questions, please use GitHub Discussions.
required: true
- label: I added a clear and detailed title that summarizes the issue.
required: true
@@ -37,7 +38,7 @@ body:
attributes:
label: Example Code
description: |
Please add a self-contained, [minimal, reproducible, example](https://stackoverflow.com/help/minimal-reproducible-example) with your use case. Replace this code with your own!
Please add a self-contained, [minimal, reproducible, example](https://stackoverflow.com/help/minimal-reproducible-example) with your use case.
placeholder: |
from langgraph.graph import StateGraph
@@ -77,7 +78,7 @@ body:
attributes:
label: System Info
description: |
Run on your machine: `python -m langchain_core.sys_info`
python -m langchain_core.sys_info
placeholder: |
python -m langchain_core.sys_info
validations:
+12 -6
View File
@@ -1,9 +1,15 @@
blank_issues_enabled: false
version: 2.1
contact_links:
- name: Documentation
url: https://github.com/langchain-ai/docs/issues/new?template=langgraph.yml
about: Report an issue related to the LangGraph documentation
- name: LangChain Forum
url: https://forum.langchain.com/
about: General community discussions and support
- name: 🤔 Question or Problem
about: Ask a question or ask about a problem in GitHub Discussions.
url: https://github.com/langchain-ai/langgraph/discussions/categories/q-a
- name: Feature Request
url: https://github.com/langchain-ai/langgraph/discussions/categories/ideas
about: Suggest a feature or an idea
- name: Show and tell
about: Show what you built with LangChain
url: https://github.com/langchain-ai/langgraph/discussions/categories/show-and-tell
- name: Slack
url: https://www.langchain.com/join-community
about: General community discussions
+19
View File
@@ -0,0 +1,19 @@
name: Documentation
description: Report an issue related to the LangGraph documentation.
title: "DOC: <Please write a comprehensive title after the 'DOC: ' prefix>"
labels: [03 - Documentation]
body:
- type: textarea
attributes:
label: "Issue with current documentation:"
description: >
Please make sure to leave a reference to the document/code you're
referring to.
- type: textarea
attributes:
label: "Idea or request for content:"
description: >
Please describe as clearly as possible what topics you think are missing
from the current documentation.
+8 -12
View File
@@ -1,29 +1,25 @@
name: 🔒 Privileged
description: You are a LangGraph maintainer, or was asked directly by a maintainer to create an issue here. If not, check the other options.
description: You are a LangChain maintainer, or was asked directly by a maintainer to create an issue here. If not, check the other options.
body:
- type: markdown
attributes:
value: |
Thanks for your interest in LangGraph! 🚀
If you are not a LangGraph maintainer or were not asked directly by a maintainer to create an issue, then please start the conversation on the [LangChain Forum](https://forum.langchain.com/) instead.
You are a LangGraph maintainer if you maintain any of the packages inside of the LangGraph repository
or are a regular contributor to LangGraph with previous merged merged pull requests.
Thanks for your interest in LangChain! 🚀
If you are not a LangChain maintainer or were not asked directly by a maintainer to create an issue, then please start the conversation in a [Question in GitHub Discussions](https://github.com/langchain-ai/langchain/discussions/categories/q-a) instead.
You are a LangChain maintainer if you maintain any of the packages inside of the LangChain repository
or are a regular contributor to LangChain with previous merged merged pull requests.
- type: checkboxes
id: privileged
attributes:
label: Privileged issue
description: Confirm that you are allowed to create an issue here.
options:
- label: I am a LangGraph maintainer, or was asked directly by a LangGraph maintainer to create an issue here.
- label: I am a LangChain maintainer, or was asked directly by a LangChain maintainer to create an issue here.
required: true
- type: textarea
id: content
attributes:
label: Issue Content
description: Add the content of the issue here.
- type: markdown
attributes:
value: |
Community members should **NOT** work on Privileged issues unless these issues have been explicitly marked with a "help-wanted" tag.
-31
View File
@@ -1,31 +0,0 @@
Thank you for contributing to LangGraph! Follow these steps to mark your pull request as ready for review. **If any of these steps are not completed, your PR will not be considered for review.**
- [ ] **PR title**: Follows the format: {TYPE}({SCOPE}): {DESCRIPTION}
- Examples:
- feat(core): add multi-tenant support
- fix(cli): resolve flag parsing error
- docs(openai): update API usage examples
- Allowed `{TYPE}` values:
- feat, fix, docs, style, refactor, perf, test, build, ci, chore, revert, release
- Allowed `{SCOPE}` values (optional):
- langgraph, docs, cli, checkpoint, checkpoint-postgres, checkpoint-sqlite, prebuilt, scheduler-kafka, sdk-py
- Once you've written the title, please delete this checklist item; do not include it in the PR.
- [ ] **PR message**: ***Delete this entire checklist*** and replace with
- **Description:** a description of the change. Include a [closing keyword](https://docs.github.com/en/issues/tracking-your-work-with-issues/using-issues/linking-a-pull-request-to-an-issue#linking-a-pull-request-to-an-issue-using-a-keyword) if applicable.
- **Issue:** the issue # it fixes, if applicable
- **Dependencies:** any dependencies required for this change
- **Twitter handle:** if your PR gets announced, and you'd like a mention, we'll gladly shout you out!
- [ ] **Add tests and docs**: If you're adding a new integration, you must include:
1. A test for the integration, preferably unit tests that do not rely on network access,
2. An example notebook showing its use. It lives in `docs/docs/integrations` directory.
- [ ] **Lint and test**: Run `make format`, `make lint` and `make test` from the root of the package(s) you've modified. We will not consider a PR unless these three are passing in CI. See [contribution guidelines](https://github.com/langchain-ai/langgraph/blob/main/CONTRIBUTING.md) for more.
Additional guidelines:
- Make sure optional dependencies are imported within a function.
- Please do not add dependencies to `pyproject.toml` files (even optional ones) unless they are **required** for unit tests.
- Most PRs should not touch more than one package.
- Changes should be backwards compatible.
+88
View File
@@ -0,0 +1,88 @@
# An action for setting up poetry install with caching.
# Using a custom action since the default action does not
# take poetry install groups into account.
# Action code from:
# https://github.com/actions/setup-python/issues/505#issuecomment-1273013236
name: poetry-install-with-caching
description: Poetry install with support for caching of dependency groups.
inputs:
python-version:
description: Python version, supporting MAJOR.MINOR only
required: true
poetry-version:
description: Poetry version
required: true
cache-key:
description: Cache key to use for manual handling of caching
required: true
runs:
using: composite
steps:
- uses: actions/setup-python@v5
name: Setup python ${{ inputs.python-version }}
id: setup-python
with:
python-version: ${{ inputs.python-version }}
- uses: actions/cache@v3
id: cache-bin-poetry
name: Cache Poetry binary - Python ${{ inputs.python-version }}
env:
SEGMENT_DOWNLOAD_TIMEOUT_MIN: "1"
with:
path: |
/opt/pipx/venvs/poetry
# This step caches the poetry installation, so make sure it's keyed on the poetry version as well.
key: bin-poetry-${{ runner.os }}-${{ runner.arch }}-py-${{ inputs.python-version }}-${{ inputs.poetry-version }}
- name: Refresh shell hashtable and fixup softlinks
if: steps.cache-bin-poetry.outputs.cache-hit == 'true'
shell: bash
env:
POETRY_VERSION: ${{ inputs.poetry-version }}
PYTHON_VERSION: ${{ inputs.python-version }}
run: |
set -eux
# Refresh the shell hashtable, to ensure correct `which` output.
hash -r
# `actions/cache@v3` doesn't always seem able to correctly unpack softlinks.
# Delete and recreate the softlinks pipx expects to have.
rm /opt/pipx/venvs/poetry/bin/python
cd /opt/pipx/venvs/poetry/bin
ln -s "$(which "python$PYTHON_VERSION")" python
chmod +x python
cd /opt/pipx_bin/
ln -s /opt/pipx/venvs/poetry/bin/poetry poetry
chmod +x poetry
# Ensure everything got set up correctly.
/opt/pipx/venvs/poetry/bin/python --version
/opt/pipx_bin/poetry --version
- name: Install poetry
if: steps.cache-bin-poetry.outputs.cache-hit != 'true'
shell: bash
env:
POETRY_VERSION: ${{ inputs.poetry-version }}
PYTHON_VERSION: ${{ inputs.python-version }}
# Install poetry using the python version installed by setup-python step.
run: pipx install "poetry==$POETRY_VERSION" --python '${{ steps.setup-python.outputs.python-path }}' --verbose
- name: Restore pip and poetry cached dependencies
uses: actions/cache@v3
env:
SEGMENT_DOWNLOAD_TIMEOUT_MIN: "4"
with:
path: |
~/.cache/pip
~/.cache/pypoetry/virtualenvs
~/.cache/pypoetry/cache
~/.cache/pypoetry/artifacts
./.venv
key: py-deps-${{ runner.os }}-${{ runner.arch }}-py-${{ inputs.python-version }}-poetry-${{ inputs.poetry-version }}-${{ inputs.cache-key }}-${{ hashFiles('./poetry.lock') }}
-18
View File
@@ -1,18 +0,0 @@
version: 2
updates:
- package-ecosystem: "github-actions"
directory: "/"
schedule:
interval: "weekly"
- package-ecosystem: "pip"
directories:
- "libs/checkpoint"
- "libs/checkpoint-postgres"
- "libs/checkpoint-sqlite"
- "libs/cli"
- "libs/langgraph"
- "libs/prebuilt"
- "libs/sdk-py"
schedule:
interval: "weekly"
+3 -8
View File
@@ -1,15 +1,10 @@
import ast
import os
from itertools import filterfalse
from typing import Dict, List, Tuple
from typing import List, Tuple
ROOT_PATH = os.path.abspath(os.path.join(__file__, "..", "..", ".."))
CLIENT_PATH = os.path.join(ROOT_PATH, "libs", "sdk-py", "langgraph_sdk", "client.py")
ASYNC_TO_SYNC_METHOD_MAP: Dict[str, str] = {
"aclose": "close",
"__aenter__": "__enter__",
"__aexit__": "__exit__",
}
def get_class_methods(node: ast.ClassDef) -> List[str]:
@@ -27,7 +22,7 @@ def find_classes(tree: ast.AST) -> List[Tuple[str, List[str]]]:
def compare_sync_async_methods(sync_methods: List[str], async_methods: List[str]) -> List[str]:
sync_set = set(sync_methods)
async_set = {ASYNC_TO_SYNC_METHOD_MAP.get(async_method, async_method) for async_method in async_methods}
async_set = set(async_methods)
missing_in_sync = list(async_set - sync_set)
missing_in_async = list(sync_set - async_set)
return missing_in_sync + missing_in_async
@@ -38,7 +33,7 @@ def main():
tree = ast.parse(file.read())
classes = find_classes(tree)
def is_sync(class_spec: Tuple[str, List[str]]) -> bool:
return class_spec[0].startswith("Sync")
+84 -146
View File
@@ -1,164 +1,108 @@
import logging
import asyncio
import json
import os
import pathlib
import sys
import time
from urllib import error, request
import langgraph_cli
import langgraph_cli.config
import langgraph_cli.docker
from langgraph_cli.cli import prepare_args_and_stdin
from langgraph_cli.constants import DEFAULT_PORT
import langgraph_cli.config
from langgraph_cli.exec import Runner, subp_exec
from langgraph_cli.progress import Progress
logger = logging.getLogger(__name__)
logging.basicConfig(level=logging.INFO)
from langgraph_cli.constants import DEFAULT_PORT
def test(config: pathlib.Path, port: int, tag: str, verbose: bool):
"""Spin up API with Postgres/Redis via docker compose and wait until ready."""
logger.info("Starting test...")
def test(
config: pathlib.Path,
port: int,
tag: str,
verbose: bool,
):
with Runner() as runner, Progress(message="Pulling...") as set:
# Detect docker/compose capabilities
# check docker available
capabilities = langgraph_cli.docker.check_capabilities(runner)
# Validate config and prepare compose stdin/args using built image
# open config
config_json = langgraph_cli.config.validate_config_file(config)
args, stdin = prepare_args_and_stdin(
capabilities=capabilities,
config_path=config,
config=config_json,
docker_compose=None,
port=port,
watch=False,
debugger_port=None,
debugger_base_url=f"http://127.0.0.1:{port}",
postgres_uri=None,
api_version=None,
image=tag,
base_image=None,
)
# Compose up with wait (implies detach), similar to `langgraph up --wait`
args_up = [*args, "up", "--remove-orphans", "--wait"]
compose_cmd = ["docker", "compose"]
if capabilities.compose_type == "standalone":
compose_cmd = ["docker-compose"]
set("Starting...")
try:
runner.run(
subp_exec(
*compose_cmd,
*args_up,
input=stdin,
verbose=verbose,
)
set("Running...")
args = [
"run",
"--rm",
"-p",
f"{port}:8000",
]
if isinstance(config_json["env"], str):
args.extend(
[
"--env-file",
str(config.parent / config_json["env"]),
]
)
except Exception as e: # noqa: BLE001
# On failure, show diagnostics then ensure clean teardown
sys.stderr.write(f"docker compose up failed: {e}\n")
try:
sys.stderr.write("\n== docker compose ps ==\n")
runner.run(
subp_exec(*compose_cmd, *args, "ps", input=stdin, verbose=False)
)
except Exception:
pass
try:
sys.stderr.write("\n== docker compose logs (api) ==\n")
runner.run(
subp_exec(
*compose_cmd,
*args,
"logs",
"langgraph-api",
input=stdin,
verbose=False,
)
)
except Exception:
pass
finally:
try:
runner.run(
subp_exec(
*compose_cmd,
*args,
"down",
"-v",
"--remove-orphans",
input=stdin,
verbose=False,
)
)
finally:
raise
set("")
base_url = f"http://localhost:{port}"
ok_url = f"{base_url}/ok"
logger.info(f"Waiting for {ok_url} to respond with 200...")
deadline = time.time() + 30
last_err: Exception | None = None
while time.time() < deadline:
try:
with request.urlopen(ok_url, timeout=2) as resp:
if resp.status == 200:
sys.stdout.write(
f"""Ready!\n- API: {base_url}\n- /ok: 200 OK\n"""
)
sys.stdout.flush()
break
else:
last_err = RuntimeError(f"Unexpected status: {resp.status}")
logger.error(f"Unexpected status: {resp.status}")
except error.URLError as e:
logger.error(f"URLError: {e}")
last_err = e
except Exception as e: # noqa: BLE001
logger.error(f"Exception: {e}")
last_err = e
time.sleep(0.5)
else:
logger.error("Timeout waiting for /ok to return 200")
# Bring stack down before raising
args_down = [*args, "down", "-v", "--remove-orphans"]
try:
runner.run(
subp_exec(
*compose_cmd,
*args_down,
input=stdin,
verbose=verbose,
)
)
finally:
raise SystemExit(
f"/ok did not return 202 within timeout. Last error: {last_err}"
for k, v in config_json["env"].items():
args.extend(
[
"-e",
f"{k}={v}",
]
)
if capabilities.healthcheck_start_interval:
args.extend(
[
"--health-interval",
"5s",
"--health-retries",
"1",
"--health-start-period",
"10s",
"--health-start-interval",
"1s",
]
)
else:
args.extend(
[
"--health-interval",
"5s",
"--health-retries",
"2",
]
)
_task = None
def on_stdout(line: str):
nonlocal _task
if "GET /ok" in line or "Uvicorn running on" in line:
set("")
sys.stdout.write(
f"""Ready!
- API: http://localhost:{port}
"""
)
sys.stdout.flush()
_task.cancel()
return True
return False
async def subp_exec_task(*args, **kwargs):
nonlocal _task
_task = asyncio.create_task(subp_exec(*args, **kwargs))
await _task
# Clean up: bring compose stack down to free ports for next test
logger.info("Test succeeded. Bringing down compose stack...")
try:
args_down = [*args, "down", "-v", "--remove-orphans"]
runner.run(
subp_exec(
*compose_cmd,
*args_down,
input=stdin,
subp_exec_task(
"docker",
*args,
tag,
verbose=verbose,
on_stdout=on_stdout,
)
)
logger.info("Compose stack down. Finishing...")
except Exception:
logger.exception("Failed to bring down compose stack")
except asyncio.CancelledError:
pass
logger.info("Test finished")
if __name__ == "__main__":
import argparse
@@ -166,12 +110,6 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("-t", "--tag", type=str)
parser.add_argument("-c", "--config", type=str, default="./langgraph.json")
parser.add_argument("-p", "--port", type=int, default=DEFAULT_PORT)
parser.add_argument("-p", "--port", default=DEFAULT_PORT)
args = parser.parse_args()
try:
test(pathlib.Path(args.config), args.port, args.tag, verbose=True)
except BaseException:
logger.exception("Test failed")
raise
logger.info("Test execution finished")
test(pathlib.Path(args.config), args.port, args.tag, verbose=True)
+38 -82
View File
@@ -3,8 +3,8 @@ name: CLI integration test
on:
workflow_call:
permissions:
contents: read
env:
POETRY_VERSION: "1.7.1"
jobs:
build:
@@ -13,106 +13,62 @@ jobs:
matrix:
python-version:
- "3.10"
- "3.14"
example:
- name: A
workdir: libs/cli/examples
tag: langgraph-test-a
- name: B
workdir: libs/cli/examples/graphs
tag: langgraph-test-b
- name: C
workdir: libs/cli/examples/graphs_reqs_a
tag: langgraph-test-c
- name: D
workdir: libs/cli/examples/graphs_reqs_b
tag: langgraph-test-d
- "3.11"
name: "CLI integration test"
defaults:
run:
working-directory: libs/cli
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Get changed files
id: changed-files
uses: Ana06/get-changed-files@v2.3.0
with:
filter: "libs/cli/**"
- name: Set up Python ${{ matrix.python-version }}
- name: Set up Python ${{ matrix.python-version }} + Poetry ${{ env.POETRY_VERSION }}
if: steps.changed-files.outputs.all
uses: astral-sh/setup-uv@v7
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ matrix.python-version }}
enable-cache: true
cache-suffix: "cli-integration-test"
ignore-nothing-to-cache: true
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: integration-test-cli
- name: Setup env
if: steps.changed-files.outputs.all
working-directory: libs/cli/examples
run: cat .env.example > .env
- name: Install cli globally
if: steps.changed-files.outputs.all
run: pip install -e .
- name: Build and test service ${{ matrix.example.name }}
- name: Build and test service A
if: steps.changed-files.outputs.all
working-directory: ${{ matrix.example.workdir }}
env:
LANGSMITH_API_KEY: ${{ secrets.LANGSMITH_API_KEY }}
working-directory: libs/cli/examples
run: |
# Build the image for this example
langgraph build -t ${{ matrix.example.tag }}
# Prepare environment file from local or parent example directory
if [ -f .env.example ]; then cp .env.example .env; elif [ -f ../.env.example ]; then cp ../.env.example .env && cp ../.env.example ../.env; fi
if [ -n "${{ secrets.LANGSMITH_API_KEY }}" ]; then echo "LANGSMITH_API_KEY=${{ secrets.LANGSMITH_API_KEY }}" >> .env; if [ -f ../.env ]; then echo "LANGSMITH_API_KEY=${{ secrets.LANGSMITH_API_KEY }}" >> ../.env; fi; fi
# Run the integration test using the built tag
# Compute repo root to reference the shared script robustly
REPO_ROOT=$(git rev-parse --show-toplevel)
timeout 60 python "$REPO_ROOT/.github/scripts/run_langgraph_cli_test.py" -t ${{ matrix.example.tag }}
# The build-arg isn't used; just testing that we accept other args
langgraph build -t langgraph-test-a --base-image "langchain/langgraph-trial"
cp .env.example .envg
timeout 60 python ../../../.github/scripts/run_langgraph_cli_test.py -c langgraph.json -t langgraph-test-a
- name: Build and test service B
if: steps.changed-files.outputs.all
working-directory: libs/cli/examples/graphs
run: |
langgraph build -t langgraph-test-b --base-image "langchain/langgraph-trial"
timeout 60 python ../../../../.github/scripts/run_langgraph_cli_test.py -t langgraph-test-b
- name: Build and test service C
if: steps.changed-files.outputs.all
working-directory: libs/cli/examples/graphs_reqs_a
run: |
langgraph build -t langgraph-test-c --base-image "langchain/langgraph-trial"
timeout 60 python ../../../../.github/scripts/run_langgraph_cli_test.py -t langgraph-test-c
- name: Build and test service D
if: steps.changed-files.outputs.all
working-directory: libs/cli/examples/graphs_reqs_b
run: |
langgraph build -t langgraph-test-d --base-image "langchain/langgraph-trial"
timeout 60 python ../../../../.github/scripts/run_langgraph_cli_test.py -t langgraph-test-d
- name: Build JS service
if: ${{ steps.changed-files.outputs.all && matrix.example.name == 'A' }}
if: steps.changed-files.outputs.all
working-directory: libs/cli/js-examples
run: |
langgraph build -t langgraph-test-e
- name: Build JS monorepo service
if: ${{ steps.changed-files.outputs.all && matrix.example.name == 'A' }}
working-directory: libs/cli/js-monorepo-example
run: |
langgraph build -t langgraph-test-f -c apps/agent/langgraph.json --build-command "yarn run turbo build" --install-command "yarn install"
- name: Build Python monorepo service
if: ${{ steps.changed-files.outputs.all && matrix.example.name == 'A' }}
working-directory: libs/cli/python-monorepo-example
run: |
langgraph build -t langgraph-test-g -c apps/agent/langgraph.json
cp apps/agent/.env.example apps/agent/.env
if [ -n "${{ secrets.LANGSMITH_API_KEY }}" ]; then echo "LANGSMITH_API_KEY=${{ secrets.LANGSMITH_API_KEY }}" >> apps/agent/.env; fi
timeout 60 python ../../../.github/scripts/run_langgraph_cli_test.py -t langgraph-test-g -c apps/agent/langgraph.json
- name: Build and test prerelease reqs service
if: ${{ steps.changed-files.outputs.all && matrix.example.name == 'A' }}
working-directory: libs/cli/examples/graph_prerelease_reqs
run: |
langgraph build -t langgraph-test-h
cp ../.env.example .env
if [ -n "${{ secrets.LANGSMITH_API_KEY }}" ]; then echo "LANGSMITH_API_KEY=${{ secrets.LANGSMITH_API_KEY }}" >> .env; fi
timeout 60 python ../../../../.github/scripts/run_langgraph_cli_test.py -t langgraph-test-h
echo "Finished starting up langgraph-test-h"
LANGGRAPH_VERSION=$(docker run --rm --entrypoint "" langgraph-test-h python -c "import sys; from importlib.metadata import version; v = version('langgraph'); print(v);")
if [ "$LANGGRAPH_VERSION" != "1.0.2" ]; then
echo "LANGGRAPH_VERSION != 1.0.2; $LANGGRAPH_VERSION"
exit 1
fi
LANGCHAIN_OPENAI_VERSION=$(docker run --rm --entrypoint "" langgraph-test-h python -c "import sys; from importlib.metadata import version; v = version('langchain-openai'); print(v);")
if [ "$LANGCHAIN_OPENAI_VERSION" != "1.0.1" ]; then
echo "LANGCHAIN_OPENAI_VERSION != 1.0.1; $LANGCHAIN_OPENAI_VERSION"
exit 1
fi
LANGCHAIN_ANTHROPIC_VERSION=$(docker run --rm --entrypoint "" langgraph-test-h python -c "import sys; from importlib.metadata import version; v = version('langchain-anthropic'); print(v);")
if [ "$LANGCHAIN_ANTHROPIC_VERSION" != "1.0.0a5" ]; then
echo "LANGCHAIN_ANTHROPIC_VERSION != 1.0.0a5; $LANGCHAIN_ANTHROPIC_VERSION"
exit 1
fi
- name: Build and test prerelease reqs fail service
if: ${{ steps.changed-files.outputs.all && matrix.example.name == 'A' }}
working-directory: libs/cli/examples/graph_prerelease_reqs_fail
run: |
langgraph build -t langgraph-test-i || [ $? -eq 1 ]
+42 -14
View File
@@ -8,10 +8,9 @@ on:
type: string
description: "From which folder this pipeline executes"
permissions:
contents: read
env:
POETRY_VERSION: "1.7.1"
# This env var allows us to get inline annotations when ruff has complaints.
RUFF_OUTPUT_FORMAT: github
@@ -31,34 +30,54 @@ jobs:
- "3.12"
name: "lint #${{ matrix.python-version }}"
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Get changed files
id: changed-files
uses: Ana06/get-changed-files@v2.3.0
with:
filter: "${{ inputs.working-directory }}/**"
- name: Set up Python ${{ matrix.python-version }}
- name: Set up Python ${{ matrix.python-version }} + Poetry ${{ env.POETRY_VERSION }}
if: steps.changed-files.outputs.all
uses: astral-sh/setup-uv@v7
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ matrix.python-version }}
enable-cache: true
cache-suffix: lint-${{ inputs.working-directory }}
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: lint-${{ inputs.working-directory }}
- name: Check Poetry File
if: steps.changed-files.outputs.all
shell: bash
working-directory: ${{ inputs.working-directory }}
run: poetry check
- name: Check lock file
if: steps.changed-files.outputs.all
shell: bash
working-directory: ${{ inputs.working-directory }}
run: poetry lock --check
- name: Install dependencies
if: steps.changed-files.outputs.all
# Also installs dev/lint/test/typing dependencies, to ensure we have
# type hints for as many of our libraries as possible.
# This helps catch errors that require dependencies to be spotted, for example:
# https://github.com/langchain-ai/langchain/pull/10249/files#diff-935185cd488d015f026dcd9e19616ff62863e8cde8c0bee70318d3ccbca98341
#
# If you change this configuration, make sure to change the `cache-key`
# in the `poetry_setup` action above to stop using the old cache.
# It doesn't matter how you change it, any change will cause a cache-bust.
working-directory: ${{ inputs.working-directory }}
run: uv sync --frozen --group lint
run: poetry install --with dev
- name: Get .mypy_cache to speed up mypy
if: steps.changed-files.outputs.all
uses: actions/cache@v5
uses: actions/cache@v3
env:
SEGMENT_DOWNLOAD_TIMEOUT_MIN: "2"
with:
path: |
${{ inputs.working-directory }}/.mypy_cache
key: mypy-lint-${{ runner.os }}-${{ runner.arch }}-py${{ matrix.python-version }}-${{ inputs.working-directory }}-${{ hashFiles(format('{0}/uv.lock', inputs.working-directory)) }}
key: mypy-lint-${{ runner.os }}-${{ runner.arch }}-py${{ matrix.python-version }}-${{ inputs.working-directory }}-${{ hashFiles(format('{0}/poetry.lock', inputs.working-directory)) }}
- name: Analysing package code with our lint
if: steps.changed-files.outputs.all
@@ -73,18 +92,27 @@ jobs:
- name: Install test dependencies
if: steps.changed-files.outputs.all
# Also installs dev/lint/test/typing dependencies, to ensure we have
# type hints for as many of our libraries as possible.
# This helps catch errors that require dependencies to be spotted, for example:
# https://github.com/langchain-ai/langchain/pull/10249/files#diff-935185cd488d015f026dcd9e19616ff62863e8cde8c0bee70318d3ccbca98341
#
# If you change this configuration, make sure to change the `cache-key`
# in the `poetry_setup` action above to stop using the old cache.
# It doesn't matter how you change it, any change will cause a cache-bust.
working-directory: ${{ inputs.working-directory }}
run: uv sync --group lint
run: |
poetry install --with dev
- name: Get .mypy_cache_test to speed up mypy
if: steps.changed-files.outputs.all
uses: actions/cache@v5
uses: actions/cache@v3
env:
SEGMENT_DOWNLOAD_TIMEOUT_MIN: "2"
with:
path: |
${{ inputs.working-directory }}/.mypy_cache_test
key: mypy-test-${{ runner.os }}-${{ runner.arch }}-py${{ matrix.python-version }}-${{ inputs.working-directory }}-${{ hashFiles(format('{0}/uv.lock', inputs.working-directory)) }}
key: mypy-test-${{ runner.os }}-${{ runner.arch }}-py${{ matrix.python-version }}-${{ inputs.working-directory }}-${{ hashFiles(format('{0}/poetry.lock', inputs.working-directory)) }}
- name: Analysing tests with our lint
if: steps.changed-files.outputs.all
+12 -10
View File
@@ -8,8 +8,8 @@ on:
type: string
description: "From which folder this pipeline executes"
permissions:
contents: read
env:
POETRY_VERSION: "1.7.1"
jobs:
build:
@@ -17,21 +17,21 @@ jobs:
strategy:
matrix:
python-version:
- "3.9"
- "3.10"
- "3.11"
- "3.12"
- "3.13"
- "3.14"
name: "test #${{ matrix.python-version }}"
steps:
- uses: actions/checkout@v6
- name: Set up Python ${{ matrix.python-version }}
uses: astral-sh/setup-uv@v7
- uses: actions/checkout@v4
- name: Set up Python ${{ matrix.python-version }} + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ matrix.python-version }}
enable-cache: true
cache-suffix: test-${{ inputs.working-directory }}
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: test-${{ inputs.working-directory }}
- name: Login to Docker Hub
uses: docker/login-action@v3
if: ${{ !github.event.pull_request.head.repo.fork }}
@@ -42,12 +42,14 @@ jobs:
- name: Install dependencies
shell: bash
working-directory: ${{ inputs.working-directory }}
run: uv sync --frozen --group test --no-dev
run: |
poetry install --with dev
- name: Run tests
shell: bash
working-directory: ${{ inputs.working-directory }}
run: make test
run: |
make test
- name: Ensure the tests did not create any additional files
shell: bash
+12 -10
View File
@@ -3,8 +3,8 @@ name: test
on:
workflow_call:
permissions:
contents: read
env:
POETRY_VERSION: "1.7.1"
jobs:
build:
@@ -12,24 +12,24 @@ jobs:
strategy:
matrix:
python-version:
- "3.9"
- "3.10"
- "3.11"
- "3.12"
- "3.13"
- "3.14"
defaults:
run:
working-directory: libs/langgraph
name: "test #${{ matrix.python-version }}"
steps:
- uses: actions/checkout@v6
- name: Set up Python ${{ matrix.python-version }}
uses: astral-sh/setup-uv@v7
- uses: actions/checkout@v4
- name: Set up Python ${{ matrix.python-version }} + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ matrix.python-version }}
enable-cache: true
cache-suffix: "test-langgraph"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: test-langgraph
- name: Login to Docker Hub
uses: docker/login-action@v3
if: ${{ !github.event.pull_request.head.repo.fork }}
@@ -39,11 +39,13 @@ jobs:
- name: Install dependencies
shell: bash
run: uv sync --frozen --group test --no-dev
run: |
poetry install --with dev
- name: Run tests
shell: bash
run: make test_parallel
run: |
make test_parallel
- name: Ensure the tests did not create any additional files
shell: bash
+13 -14
View File
@@ -9,13 +9,12 @@ on:
description: "From which folder this pipeline executes"
env:
POETRY_VERSION: "1.7.1"
PYTHON_VERSION: "3.10"
permissions:
contents: read
jobs:
build:
if: github.ref == 'refs/heads/main'
runs-on: ubuntu-latest
outputs:
@@ -23,14 +22,14 @@ jobs:
version: ${{ steps.check-version.outputs.version }}
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Python $${ env.PYTHON_VERSION }}
uses: astral-sh/setup-uv@v7
- name: Set up Python + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ env.PYTHON_VERSION }}
enable-cache: true
cache-suffix: "release"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: release
# We want to keep this build stage *separate* from the release stage,
# so that there's no sharing of permissions between them.
@@ -44,11 +43,11 @@ jobs:
# > from the publish job.
# https://github.com/pypa/gh-action-pypi-publish#non-goals
- name: Build project for distribution
run: uv build
run: poetry build
working-directory: ${{ inputs.working-directory }}
- name: Upload build
uses: actions/upload-artifact@v6
uses: actions/upload-artifact@v4
with:
name: test-dist
path: ${{ inputs.working-directory }}/dist/
@@ -58,8 +57,8 @@ jobs:
shell: bash
working-directory: ${{ inputs.working-directory }}
run: |
echo pkg-name=$(grep -m 1 "^name = " pyproject.toml | cut -d '"' -f 2)
echo version=$(grep -m 1 "^version = " pyproject.toml | cut -d '"' -f 2)
echo pkg-name="$(poetry version | cut -d ' ' -f 1)" >> $GITHUB_OUTPUT
echo version="$(poetry version --short)" >> $GITHUB_OUTPUT
publish:
needs:
@@ -74,9 +73,9 @@ jobs:
id-token: write
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- uses: actions/download-artifact@v7
- uses: actions/download-artifact@v4
with:
name: test-dist
path: ${{ inputs.working-directory }}/dist/
@@ -0,0 +1,57 @@
name: test
on:
workflow_call:
env:
POETRY_VERSION: "1.7.1"
jobs:
build:
runs-on: ubuntu-latest
strategy:
matrix:
python-version:
- "3.11"
- "3.12"
defaults:
run:
working-directory: libs/scheduler-kafka
name: "test #${{ matrix.python-version }}"
steps:
- uses: actions/checkout@v4
- name: Set up Python ${{ matrix.python-version }} + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ matrix.python-version }}
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: test-scheduler-kafka
- name: Login to Docker Hub
uses: docker/login-action@v3
if: ${{ !github.event.pull_request.head.repo.fork }}
with:
username: ${{ secrets.DOCKERHUB_USERNAME }}
password: ${{ secrets.DOCKERHUB_RO_TOKEN }}
- name: Install dependencies
shell: bash
run: |
poetry install --with dev
- name: Run tests
shell: bash
run: |
make test
- name: Ensure the tests did not create any additional files
shell: bash
run: |
set -eu
STATUS="$(git status)"
echo "$STATUS"
# grep will exit non-zero if the target message isn't found,
# and `set -e` above will cause the step to fail.
echo "$STATUS" | grep 'nothing to commit, working tree clean'
+9 -9
View File
@@ -7,8 +7,8 @@ on:
paths:
- "libs/**"
permissions:
contents: read
env:
POETRY_VERSION: "1.7.1"
jobs:
benchmark:
@@ -17,20 +17,20 @@ jobs:
run:
working-directory: libs/langgraph
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- run: SHA=$(git rev-parse HEAD) && echo "SHA=$SHA" >> $GITHUB_ENV
- name: Set up Python 3.11
uses: astral-sh/setup-uv@v7
- name: Set up Python 3.11 + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: "3.11"
enable-cache: true
cache-suffix: "bench"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: bench
- name: Install dependencies
run: uv sync --group test
run: poetry install --with dev
- name: Run benchmarks
run: OUTPUT=out/benchmark-baseline.json make -s benchmark
- name: Save outputs
uses: actions/cache/save@v5
uses: actions/cache/save@v4
with:
key: ${{ runner.os }}-benchmark-baseline-${{ env.SHA }}
path: |
+12 -12
View File
@@ -5,8 +5,8 @@ on:
paths:
- "libs/**"
permissions:
contents: read
env:
POETRY_VERSION: "1.7.1"
jobs:
benchmark:
@@ -15,22 +15,22 @@ jobs:
run:
working-directory: libs/langgraph
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- id: files
name: Get changed files
uses: Ana06/get-changed-files@v2.3.0
with:
format: json
- name: Set up Python 3.11
uses: astral-sh/setup-uv@v7
- name: Set up Python 3.11 + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: "3.11"
enable-cache: true
cache-suffix: "bench"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: bench
- name: Install dependencies
run: uv sync --group test
run: poetry install --with dev
- name: Download baseline
uses: actions/cache/restore@v5
uses: actions/cache/restore@v4
with:
key: ${{ runner.os }}-benchmark-baseline
restore-keys: |
@@ -43,7 +43,7 @@ jobs:
run: |
{
echo 'OUTPUT<<EOF'
make -s benchmark-fast
make -s benchmark
echo EOF
} >> "$GITHUB_OUTPUT"
- name: Compare benchmarks
@@ -53,11 +53,11 @@ jobs:
echo 'OUTPUT<<EOF'
mv out/benchmark-baseline.json out/main.json
mv out/benchmark.json out/changes.json
uv run pyperf compare_to out/main.json out/changes.json --table --group-by-speed
poetry run pyperf compare_to out/main.json out/changes.json --table --group-by-speed
echo EOF
} >> "$GITHUB_OUTPUT"
- name: Annotation
uses: actions/github-script@v8
uses: actions/github-script@v7
with:
script: |
const file = JSON.parse(`${{ steps.files.outputs.added_modified_renamed }}`)[0]
+77 -56
View File
@@ -3,13 +3,9 @@ name: CI
on:
push:
branches:
- main
branches: [main]
pull_request:
permissions:
contents: read
# If another push to the same PR or branch happens while this workflow is still running,
# cancel the earlier run in favor of the next run.
#
@@ -20,14 +16,17 @@ concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true
env:
POETRY_VERSION: "1.7.1"
jobs:
changes:
runs-on: ubuntu-latest
outputs:
python: ${{ steps.filter.outputs.python }}
deps: ${{ steps.filter.outputs.deps }}
sdk-js: ${{ steps.filter.outputs.sdk-js }}
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
id: filter
with:
@@ -39,10 +38,10 @@ jobs:
- 'libs/checkpoint/**'
- 'libs/checkpoint-sqlite/**'
- 'libs/checkpoint-postgres/**'
- 'libs/scheduler-kafka/**'
- 'libs/prebuilt/**'
deps:
- '**/pyproject.toml'
- '**/uv.lock'
sdk-js:
- 'libs/sdk-js/**'
lint:
needs: changes
@@ -57,10 +56,10 @@ jobs:
"libs/checkpoint",
"libs/checkpoint-sqlite",
"libs/checkpoint-postgres",
"libs/scheduler-kafka",
"libs/prebuilt",
]
if: needs.changes.outputs.python == 'true' || needs.changes.outputs.deps == 'true'
if: needs.changes.outputs.python == 'true'
uses: ./.github/workflows/_lint.yml
with:
working-directory: ${{ matrix.working-directory }}
@@ -78,9 +77,8 @@ jobs:
"libs/checkpoint-sqlite",
"libs/checkpoint-postgres",
"libs/prebuilt",
"libs/sdk-py",
]
if: needs.changes.outputs.python == 'true' || needs.changes.outputs.deps == 'true'
if: needs.changes.outputs.python == 'true'
uses: ./.github/workflows/_test.yml
with:
working-directory: ${{ matrix.working-directory }}
@@ -89,78 +87,101 @@ jobs:
# NOTE: we're testing langgraph separately because it requires a different matrix
test-langgraph:
needs: changes
if: needs.changes.outputs.python == 'true' || needs.changes.outputs.deps == 'true'
if: needs.changes.outputs.python == 'true'
name: "cd libs/langgraph"
uses: ./.github/workflows/_test_langgraph.yml
secrets: inherit
# NOTE: we're testing scheduler-kafka separately because it requires a different matrix
test-scheduler-kafka:
needs: changes
if: needs.changes.outputs.python == 'true'
name: "cd libs/scheduler-kafka"
uses: ./.github/workflows/_test_scheduler_kafka.yml
secrets: inherit
check-sdk-methods:
needs: changes
if: needs.changes.outputs.python == 'true'
name: "Check SDK methods matching"
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v6
uses: actions/setup-python@v5
with:
python-version: "3.11"
- name: Run check_sdk_methods script
run: python .github/scripts/check_sdk_methods.py
check-schema:
needs: changes
if: needs.changes.outputs.python == 'true'
name: "Check CLI schema hasn't changed #${{ matrix.python-version }}"
runs-on: ubuntu-latest
strategy:
matrix:
python-version:
- "3.13"
steps:
- uses: actions/checkout@v6
- name: Set up Python ${{ matrix.python-version }}
uses: astral-sh/setup-uv@v7
with:
python-version: "3.13"
enable-cache: true
cache-suffix: "schema-check-cli"
- name: Install CLI dependencies
run: |
cd libs/cli
uv sync
- name: Generate schema and check for changes
run: |
cd libs/cli
# Create a temporary copy of the current schema
cp schemas/schema.json schemas/schema.current.json
# Generate new schema
uv run python generate_schema.py
# Compare the new schema with the original
if ! diff -q schemas/schema.json schemas/schema.current.json > /dev/null; then
echo "Error: Langgraph.json configuration schema has changed. Please run 'uv run python generate_schema.py' in the libs/cli directory and commit the changes."
diff schemas/schema.json schemas/schema.current.json
exit 1
fi
echo "Schema check passed - no changes detected"
integration-test:
needs: changes
if: needs.changes.outputs.python == 'true' || needs.changes.outputs.deps == 'true'
if: needs.changes.outputs.python == 'true'
name: CLI integration test
uses: ./.github/workflows/_integration_test.yml
secrets: inherit
lint-js:
needs: changes
if: needs.changes.outputs.sdk-js == 'true'
runs-on: ubuntu-latest
strategy:
matrix:
working-directory:
- "libs/sdk-js"
defaults:
run:
working-directory: ${{ matrix.working-directory }}
steps:
- uses: actions/checkout@v3
- name: Setup Node.js (LTS)
uses: actions/setup-node@v3
with:
node-version: "20"
cache: "yarn"
cache-dependency-path: ${{ matrix.working-directory }}/yarn.lock
- name: Install dependencies
run: yarn install
- name: Run lint
run: yarn lint
- name: Build
run: yarn build
test-js:
needs: changes
if: needs.changes.outputs.sdk-js == 'true'
runs-on: ubuntu-latest
strategy:
matrix:
working-directory:
- "libs/sdk-js"
defaults:
run:
working-directory: ${{ matrix.working-directory }}
steps:
- uses: actions/checkout@v3
- name: Setup Node.js (LTS)
uses: actions/setup-node@v3
with:
node-version: "20"
cache: "yarn"
cache-dependency-path: ${{ matrix.working-directory }}/yarn.lock
- name: Install dependencies
run: yarn install
- name: Run tests
run: yarn test
ci_success:
name: "CI Success"
needs:
[
lint,
lint-js,
test,
test-langgraph,
check-sdk-methods,
check-schema,
test-scheduler-kafka,
integration-test,
test-js,
]
if: |
always()
+43
View File
@@ -0,0 +1,43 @@
---
name: CI / cd . / make spell_check
on:
push:
branches: [main]
pull_request:
branches: [main]
permissions:
contents: read
defaults:
run:
working-directory: docs
jobs:
codespell:
name: (Check for spelling errors)
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Install Dependencies
run: |
pip install toml codespell==2.3.0 jupytext
- name: Extract Ignore Words List
run: |
# Use a Python script to extract the ignore words list from pyproject.toml
python ../.github/workflows/extract_ignored_words_list.py
id: extract_ignore_words
- name: Codespell
uses: codespell-project/actions-codespell@v2
with:
skip: '*.ambr,*.lock,*.ipynb,*.yaml,*.zlib,*.md'
ignore_words_list: ${{ steps.extract_ignore_words.outputs.ignore_words_list }}
# We do this to avoid spellchecking cell outputs
- name: Codespell Notebooks
run: make codespell
+170
View File
@@ -0,0 +1,170 @@
name: Deploy Docs
on:
push:
branches:
- main
pull_request:
branches:
- main
workflow_dispatch:
env:
POETRY_VERSION: "1.7.1"
permissions:
contents: read
pages: write
id-token: write
concurrency:
group: "pages"
cancel-in-progress: false
defaults:
run:
working-directory: docs
jobs:
get-changed-files:
runs-on: ubuntu-latest
outputs:
changed-files: ${{ steps.changed-files.outputs.added_modified }}
steps:
- uses: actions/checkout@v4
- name: Get changed files
id: changed-files
uses: Ana06/get-changed-files@v2.3.0
with:
filter: "docs/docs/**"
run-changed-notebooks:
needs: get-changed-files
uses: ./.github/workflows/run_notebooks.yml
secrets: inherit
with:
changed-files: ${{ needs.get-changed-files.outputs.changed-files }}
deploy:
# needs: run-changed-notebooks
runs-on: ubuntu-latest
timeout-minutes: 10 # Job will be cancelled if it runs for more than 10 minutes
env:
GITHUB_TOKEN: ${{ secrets.MKDOCS_GITHUB_TOKEN }}
steps:
- uses: actions/checkout@v4
with:
fetch-depth: 0
- name: Set up Python + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: "3.12"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: docs
- name: Use Node.js
uses: actions/setup-node@v3
with:
node-version: "22"
cache: "yarn"
cache-dependency-path: docs/yarn.lock
- name: Install dependencies
run: |
yarn
poetry install --with test --with docs --no-root
poetry run pip install -U \
pytest \
pytest-check-links \
GitPython \
"git+https://github.com/benjamincburns/markdown-exec.git@cc0d39d737e5ffd4b83d23cd8729d7ea16e363c8"
# we run this installation only for internal PRs
# as GITHUB_TOKEN is not available for PRs from outside contributors
if [ -n "${GITHUB_TOKEN}" ]; then
poetry run pip install "git+https://${GITHUB_TOKEN}@github.com/langchain-ai/mkdocs-material-insiders.git"
fi
poetry run jupyter kernelspec list
poetry run python3 -m ipykernel install --user --name=python3
npm install -g tslab
poetry run tslab install --python=python3
poetry run jupyter kernelspec list
- name: Run unit tests
# Run unit tests on the docs build pipeline
run: make tests
- name: Lint Docs
# This step lints the docs using the existing linting set up.
# It should be very fast and should not require any external services.
run: make lint-docs
- name: Build llms-text
run: make llms-text
- name: Build site
run: make build-docs
env:
MKDOCS_GIT_COMMITTERS_APIKEY: ${{ secrets.MKDOCS_GIT_COMMITTERS_APIKEY }}
OPENAI_API_KEY: sf-proj-1234567890 # fake placeholder, shouldn't actually be used
ANTHROPIC_API_KEY: sk-ant-api03-1234567890 # fake placeholder, shouldn't actually be used
- name: Check links in notebooks
env:
LANGCHAIN_API_KEY: test
run: |
if [ "${{ github.event_name }}" == "schedule" ] || [ "${{ github.event_name }}" == "workflow_dispatch" ] || ([ "${{ github.event_name }}" == "push" ] && [ "${{ github.ref }}" == "refs/heads/main" ]); then
echo "Running link check on all HTML files matching notebooks in docs directory..."
poetry run pytest -v \
--check-links-ignore "https://(api|web|docs)\.smith\.langchain\.com/.*" \
--check-links-ignore "https://academy\.langchain\.com/.*" \
--check-links-ignore "https://x.com/.*" \
--check-links-ignore "https://twitter.com/.*" \
--check-links-ignore "https://github\.com/.*" \
--check-links-ignore "http://localhost:8123/.*" \
--check-links-ignore "http://localhost:2024.*" \
--check-links-ignore "http://127.0.0.1:.*" \
--check-links-ignore "/.*\.(ipynb|html)$" \
--check-links-ignore "https://python\.langchain\.com/.*" \
--check-links-ignore "https://openai\.com/.*" \
--check-links-ignore "https://www\.uber\.com/.*" \
--check-links-ignore "https://pepy\.tech/.*" \
--check-links $(find site -name "index.html" | grep -v 'storm/index.html')
else
echo "Fetching changes from origin/main..."
git fetch origin main
echo "Checking for changed notebook files..."
CHANGED_FILES=$(git diff --name-only --diff-filter=d origin/main | grep 'docs/docs/.*\.ipynb$' | grep -v 'storm.ipynb' | sed -E 's|^docs/docs/|site/|; s/\.ipynb$/\/index.html/' || true)
echo "Changed files: ${CHANGED_FILES}"
if [ -n "${CHANGED_FILES}" ]; then
echo "Running link check on HTML files matching changed notebook files..."
poetry run pytest -v \
--check-links-ignore "https://(api|web|docs)\.smith\.langchain\.com/.*" \
--check-links-ignore "https://academy\.langchain\.com/.*" \
--check-links-ignore "http://localhost:8123/.*" \
--check-links-ignore "http://localhost:2024.*" \
--check-links-ignore "http://127.0.0.1:.*" \
--check-links-ignore "https://x.com/.*" \
--check-links-ignore "https://twitter.com/.*" \
--check-links-ignore "https://github\.com/.*" \
--check-links-ignore "/.*\.(ipynb|html)$" \
--check-links ${CHANGED_FILES} \
|| ([ $? = 5 ] && exit 0 || exit $?)
else
echo "No notebook files changed."
fi
fi
- name: Configure GitHub Pages
if: github.ref == 'refs/heads/main'
uses: actions/configure-pages@v4
- name: Upload Pages Artifact
# if: github.ref == 'refs/heads/main'
uses: actions/upload-pages-artifact@v3
with:
path: ./docs/site/
- name: Deploy to GitHub Pages
if: github.ref == 'refs/heads/main'
id: deployment
uses: actions/deploy-pages@v4
@@ -0,0 +1,10 @@
import toml
pyproject_toml = toml.load("pyproject.toml")
# Extract the ignore words list (adjust the key as per your TOML structure)
ignore_words_list = (
pyproject_toml.get("tool", {}).get("codespell", {}).get("ignore-words-list")
)
print(f"::set-output name=ignore_words_list::{ignore_words_list}") # noqa: T201
+49
View File
@@ -0,0 +1,49 @@
name: Check Docs & Links
on:
pull_request:
branches:
- main
push:
branches:
- main
schedule:
- cron: "0 5 * * *"
workflow_dispatch:
env:
POETRY_VERSION: "1.7.1"
jobs:
markdown-link-check:
runs-on: ubuntu-latest
steps:
- name: Checkout code
uses: actions/checkout@v4
with:
fetch-depth: 0
- name: Check links in Markdown files
uses: gaurav-nelson/github-action-markdown-link-check@v1
with:
folder-path: "docs/"
check-modified-files-only: ${{ github.event_name != 'schedule' }}
file-path: "./README.md"
config-file: "./.markdown-link-check.config.json"
check-readmes-synced:
# This checks that the repo README.md is identical to the libs/langgraph/README.md
runs-on: ubuntu-latest
steps:
- name: Checkout code
uses: actions/checkout@v4
with:
fetch-depth: 1
- name: Check README.md is in sync
run: |
if ! diff -q README.md libs/langgraph/README.md >/dev/null; then
echo "README.md is out of sync with libs/langgraph/README.md"
diff -C 3 README.md libs/langgraph/README.md
exit 1
fi
-46
View File
@@ -1,46 +0,0 @@
name: PR Title Lint
permissions:
pull-requests: read
on:
pull_request:
types: [opened, edited, synchronize]
jobs:
lint-pr-title:
runs-on: ubuntu-latest
steps:
- name: Validate PR Title
uses: amannn/action-semantic-pull-request@v6
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
with:
types: |
feat
fix
docs
style
refactor
perf
test
build
ci
chore
revert
release
scopes: |
checkpoint
checkpoint-postgres
checkpoint-sqlite
cli
langgraph
prebuilt
scheduler-kafka
sdk-py
docs
ci
deps
requireScope: false
ignoreLabels: |
ignore-lint-pr-title
+35 -45
View File
@@ -8,14 +8,13 @@ on:
type: string
default: "libs/langgraph"
permissions:
contents: read
env:
PYTHON_VERSION: "3.11"
POETRY_VERSION: "1.7.1"
jobs:
build:
if: github.ref == 'refs/heads/main'
runs-on: ubuntu-latest
outputs:
@@ -25,14 +24,14 @@ jobs:
tag: ${{ steps.check-version.outputs.tag }}
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Python
uses: astral-sh/setup-uv@v7
- name: Set up Python + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ env.PYTHON_VERSION }}
enable-cache: true
cache-suffix: "release"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: release
# We want to keep this build stage *separate* from the release stage,
# so that there's no sharing of permissions between them.
@@ -46,11 +45,11 @@ jobs:
# > from the publish job.
# https://github.com/pypa/gh-action-pypi-publish#non-goals
- name: Build project for distribution
run: uv build
run: poetry build
working-directory: ${{ inputs.working-directory }}
- name: Upload build
uses: actions/upload-artifact@v6
uses: actions/upload-artifact@v4
with:
name: dist
path: ${{ inputs.working-directory }}/dist/
@@ -60,14 +59,8 @@ jobs:
shell: bash
working-directory: ${{ inputs.working-directory }}
run: |
PKG_NAME=$(grep -m 1 "^name = " pyproject.toml | cut -d '"' -f 2)
if grep -q 'dynamic.*=.*\[.*"version".*\]' pyproject.toml; then
# handle dynamic versioning
DIR_NAME=$(echo "$PKG_NAME" | tr '-' '_')
VERSION=$(grep -m 1 '^__version__' "${DIR_NAME}/__init__.py" | cut -d '"' -f 2)
else
VERSION=$(grep -m 1 "^version = " pyproject.toml | cut -d '"' -f 2)
fi
PKG_NAME="$(poetry version | cut -d ' ' -f 1)"
VERSION="$(poetry version --short)"
SHORT_PKG_NAME="$(echo "$PKG_NAME" | sed -e 's/langgraph//g' -e 's/-//g')"
if [ -z $SHORT_PKG_NAME ]; then
TAG="$VERSION"
@@ -86,7 +79,7 @@ jobs:
outputs:
release-body: ${{ steps.generate-release-body.outputs.release-body }}
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
with:
repository: langchain-ai/langgraph
path: langgraph
@@ -142,9 +135,7 @@ jobs:
needs:
- build
- release-notes
permissions:
contents: read
id-token: write
permissions: write-all
uses: ./.github/workflows/_test_release.yml
with:
working-directory: ${{ inputs.working-directory }}
@@ -157,7 +148,7 @@ jobs:
- test-pypi-publish
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
# We explicitly *don't* set up caching here. This ensures our tests are
# maximally sensitive to catching breakage.
@@ -172,11 +163,11 @@ jobs:
# - The package is published, and it breaks on the missing dependency when
# used in the real world.
- name: Set up Python
uses: astral-sh/setup-uv@v7
- name: Set up Python + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ env.PYTHON_VERSION }}
enable-cache: true
poetry-version: ${{ env.POETRY_VERSION }}
- name: Import published package
shell: bash
@@ -194,18 +185,18 @@ jobs:
# - attempt install again after 5 seconds if it fails because there is
# sometimes a delay in availability on test pypi
run: |
uv run pip install \
poetry run pip install \
--extra-index-url https://test.pypi.org/simple/ \
"$PKG_NAME==$VERSION" || \
( \
sleep 5 && \
uv run pip install \
poetry run pip install \
--extra-index-url https://test.pypi.org/simple/ \
"$PKG_NAME==$VERSION" \
)
if [[ "$PKG_NAME" == *prebuilt* ]]; then
uv run pip install langgraph
poetry run pip install langgraph
fi
if [[ "$PKG_NAME" == *checkpoint* || "$PKG_NAME" == *prebuilt* ]]; then
@@ -218,10 +209,10 @@ jobs:
IMPORT_NAME="$(echo "$PKG_NAME" | sed s/-/_/g)"
fi
uv run python -c "import $IMPORT_NAME; print(dir($IMPORT_NAME))"
poetry run python -c "import $IMPORT_NAME; print(dir($IMPORT_NAME))"
- name: Import test dependencies
run: uv sync --group test
run: poetry install --with dev
working-directory: ${{ inputs.working-directory }}
# Overwrite the local version of the package with the test PyPI version.
@@ -232,7 +223,7 @@ jobs:
PKG_NAME: ${{ needs.build.outputs.pkg-name }}
VERSION: ${{ needs.build.outputs.version }}
run: |
uv run pip install \
poetry run pip install \
--extra-index-url https://test.pypi.org/simple/ \
"$PKG_NAME==$VERSION"
@@ -260,16 +251,16 @@ jobs:
working-directory: ${{ inputs.working-directory }}
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Python
uses: astral-sh/setup-uv@v7
- name: Set up Python + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ env.PYTHON_VERSION }}
enable-cache: true
cache-suffix: "release"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: release
- uses: actions/download-artifact@v7
- uses: actions/download-artifact@v4
with:
name: dist
path: ${{ inputs.working-directory }}/dist/
@@ -301,16 +292,16 @@ jobs:
working-directory: ${{ inputs.working-directory }}
steps:
- uses: actions/checkout@v6
- uses: actions/checkout@v4
- name: Set up Python
uses: astral-sh/setup-uv@v7
- name: Set up Python + Poetry ${{ env.POETRY_VERSION }}
uses: "./.github/actions/poetry_setup"
with:
python-version: ${{ env.PYTHON_VERSION }}
enable-cache: true
cache-suffix: "release"
poetry-version: ${{ env.POETRY_VERSION }}
cache-key: release
- uses: actions/download-artifact@v7
- uses: actions/download-artifact@v4
with:
name: dist
path: ${{ inputs.working-directory }}/dist/
@@ -322,6 +313,5 @@ jobs:
token: ${{ secrets.GITHUB_TOKEN }}
generateReleaseNotes: false
tag: ${{needs.build.outputs.tag}}
name: ${{ needs.build.outputs.pkg-name }}==${{ needs.build.outputs.version }}
body: ${{ needs.release-notes.outputs.release-body }}
commit: ${{ github.sha }}
+38
View File
@@ -0,0 +1,38 @@
name: JS Release
on:
workflow_dispatch:
jobs:
publish:
# Disallow publishing from branches that aren't `main`.
if: github.ref == 'refs/heads/main'
runs-on: ubuntu-latest
strategy:
matrix:
working-directory:
- "libs/sdk-js"
defaults:
run:
working-directory: ${{ matrix.working-directory }}
steps:
- uses: actions/checkout@v4
# JS Build
- name: Use Node.js
uses: actions/setup-node@v3
with:
node-version: "20"
cache: "yarn"
cache-dependency-path: ${{ matrix.working-directory }}/yarn.lock
- name: Install dependencies
run: yarn install
- name: Build
run: yarn build
- name: Publish package to NPM
run: |
echo "//registry.npmjs.org/:_authToken=${{ secrets.NPM_TOKEN }}" > .npmrc
npm publish
+82
View File
@@ -0,0 +1,82 @@
name: Run notebooks
on:
workflow_dispatch:
workflow_call:
inputs:
changed-files:
required: false
type: string
description: "JSON string of changed files"
schedule:
- cron: '0 13 * * *'
defaults:
run:
working-directory: docs
jobs:
build:
runs-on: ubuntu-latest
strategy:
matrix:
lib-version:
- "development"
- "latest"
steps:
- uses: actions/checkout@v4
- name: Set up Python + Poetry
uses: "./.github/actions/poetry_setup"
with:
python-version: 3.11
poetry-version: 1.7.1
cache-key: test-langgraph-notebooks
- name: Install dependencies
run: |
poetry install --with test
poetry run pip install jupyter
- name: Start services
run: make start-services
- name: Pre-download tiktoken files
run: |
poetry run python _scripts/download_tiktoken.py
- name: Prepare notebooks
run: |
if [ "${{ matrix.lib-version }}" = "development" ]; then
poetry run python _scripts/prepare_notebooks_for_ci.py --comment-install-cells
else
poetry run python _scripts/prepare_notebooks_for_ci.py
fi
- name: Run notebooks
env:
# these won't actually be used because of the VCR cassettes
# but need to set them to avoid triggering getpass()
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }}
TAVILY_API_KEY: ${{ secrets.TAVILY_API_KEY }}
LANGSMITH_API_KEY: ${{ secrets.LANGSMITH_API_KEY }}
NOMIC_API_KEY: ${{ secrets.NOMIC_API_KEY }}
COHERE_API_KEY: ${{ secrets.COHERE_API_KEY }}
FIREWORKS_API_KEY: ${{ secrets.FIREWORKS_API_KEY }}
run: |
if [ "${{ github.event_name }}" = "workflow_dispatch" ] || [ "${{ github.event_name }}" = "schedule" ]; then
echo "Running all notebooks"
./_scripts/execute_notebooks.sh
else
CHANGED_FILES=$(echo '${{ inputs.changed-files }}' | tr ' ' '\n' | sed 's|^docs/docs/|docs/|' | grep '\.ipynb$' || true)
if [ -n "$CHANGED_FILES" ]; then
echo "Running changed notebooks: $CHANGED_FILES"
./_scripts/execute_notebooks.sh $CHANGED_FILES
else
echo "No notebook files changed, skipping execution"
fi
fi
- name: Stop services
run: make stop-services
+29
View File
@@ -0,0 +1,29 @@
name: Check File Size
on:
push:
branches:
- main
pull_request:
branches:
- main
workflow_dispatch:
jobs:
file-size-check:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Get changed files
id: changed-files
uses: tj-actions/changed-files@v44
- name: Filter by size
# TODO: roll back the web voyager hack
run: |
large_added_files=$(find ${{ steps.changed-files.outputs.added_files }} -maxdepth 0 -size +1M | grep -v "web_voyager" || true)
if [ -n "$large_added_files" ]; then
echo "Large files added: $large_added_files"
echo "# Large files added:" >> $GITHUB_STEP_SUMMARY
echo "$large_added_files" >> $GITHUB_STEP_SUMMARY
exit 1
fi
-45
View File
@@ -1,45 +0,0 @@
name: UV Lock Upgrade
on:
schedule:
# run at midnight every Sunday
- cron: '0 0 * * 0'
# allow manual triggering
workflow_dispatch:
permissions:
contents: write
pull-requests: write
jobs:
upgrade-dependencies:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- name: Set up uv
uses: astral-sh/setup-uv@v7
with:
# use minimum supported Python version
python-version: "3.10"
enable-cache: true
cache-suffix: "uv-lock-upgrade"
- name: Run uv lock --upgrade in all Python packages
run: make lock-upgrade
- name: Create Pull Request
uses: peter-evans/create-pull-request@v8
with:
token: ${{ secrets.GITHUB_TOKEN }}
commit-message: "chore(deps): upgrade dependencies with `uv lock --upgrade`"
title: "chore(deps): upgrade dependencies with `uv lock --upgrade`"
body: |
This PR updates the dependencies in all Python packages using `uv lock --upgrade`.
This is an automated PR created by the UV Lock Upgrade workflow.
branch: deps/uv-lock-upgrade
delete-branch: true
labels: |
dependencies
+82 -2
View File
@@ -6,6 +6,9 @@ __pycache__/
*.py[cod]
*$py.class
# C extensions
*.so
# Distribution / packaging
.Python
build/
@@ -51,12 +54,27 @@ coverage.xml
.hypothesis/
.pytest_cache/
# Translations
*.mo
*.pot
# Django stuff:
*.log
local_settings.py
db.sqlite3
db.sqlite3-journal
# Flask stuff:
instance/
.webassets-cache
# Scrapy stuff:
.scrapy
# Sphinx documentation
docs/_build/
docs/docs/_build/
# PyBuilder
target/
@@ -71,9 +89,23 @@ ipython_config.py
# pyenv
.python-version
# pipenv
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
# However, in case of collaboration, if having platform-specific dependencies or dependencies
# having no cross-platform support, pipenv may install dependencies that don't work, or not
# install all needed dependencies.
#Pipfile.lock
# PEP 582; used by e.g. github.com/David-OConnor/pyflow
__pypackages__/
# Celery stuff
celerybeat-schedule
celerybeat.pid
# SageMath parsed files
*.sage.py
# Environments
.env
.envrc
@@ -85,6 +117,16 @@ ENV/
env.bak/
venv.bak/
# Spyder project settings
.spyderproject
.spyproject
# Rope project settings
.ropeproject
# mkdocs documentation
/site
# mypy
.mypy_cache/
.dmypy.json
@@ -96,7 +138,45 @@ dmypy.json
# macOS display setting files
.DS_Store
# Wandb directory
wandb/
# asdf tool versions
.tool-versions
/.ruff_cache/
*.pkl
*.bin
# integration test artifacts
data_map*
\[('_type', 'fake'), ('stop', None)]
# Replit files
*replit*
node_modules
docs/.yarn/
docs/node_modules/
docs/.docusaurus/
docs/.cache-loader/
docs/_dist
docs/api_reference/api_reference.rst
docs/api_reference/experimental_api_reference.rst
docs/api_reference/_build
docs/api_reference/*/
!docs/api_reference/_static/
!docs/api_reference/templates/
!docs/api_reference/themes/
docs/docs_skeleton/build
docs/docs_skeleton/node_modules
docs/docs_skeleton/yarn.lock
# Any new jupyter notebooks
# not intended for the repo
Untitled*.ipynb
Chinook.db
.vercel
.turbo
.editorconfig
.scratch
+4
View File
@@ -0,0 +1,4 @@
{
"aliveStatusCodes": [200, 206, 402],
"ignorePatterns": ["*dcbadge.vercel.app*"]
}
-55
View File
@@ -1,55 +0,0 @@
# AGENTS Instructions
This repository is a monorepo. Each library lives in a subdirectory under `libs/`.
When you modify code in any library, run the following commands in that library's directory before creating a pull request:
- `make format` – run code formatters
- `make lint` – run the linter
- `make test` – execute the test suite
To run a particular test file or to pass additional pytest options you can specify the `TEST` variable:
```txt
TEST=path/to/test.py make test
```
Other pytest arguments can also be supplied inside the `TEST` variable.
## Libraries
The repository contains several Python and JavaScript/TypeScript libraries.
Below is a high-level overview:
- **checkpoint** – base interfaces for LangGraph checkpointers.
- **checkpoint-postgres** – Postgres implementation of the checkpoint saver.
- **checkpoint-sqlite** – SQLite implementation of the checkpoint saver.
- **cli** – official command-line interface for LangGraph.
- **langgraph** – core framework for building stateful, multi-actor agents.
- **prebuilt** – high-level APIs for creating and running agents and tools.
- **sdk-js** – JS/TS SDK for interacting with the LangGraph REST API.
- **sdk-py** – Python SDK for the LangGraph Server API.
### Dependency map
The diagram below lists downstream libraries for each production dependency as
declared in that library's `pyproject.toml` (or `package.json`).
```text
checkpoint
├── checkpoint-postgres
├── checkpoint-sqlite
├── prebuilt
└── langgraph
prebuilt
└── langgraph
sdk-py
├── langgraph
└── cli
sdk-js (standalone)
```
Changes to a library may impact all of its dependents shown above.
+126 -39
View File
@@ -1,55 +1,142 @@
# AGENTS Instructions
# LangGraph Coding Guide
This repository is a monorepo. Each library lives in a subdirectory under `libs/`.
## Repository Structure
When you modify code in any library, run the following commands in that library's directory before creating a pull request:
LangGraph follows a monorepo organization, with the following structure:
- `make format` – run code formatters
- `make lint` – run the linter
- `make test` – execute the test suite
- `libs/langgraph` is the main Python library, published to pypi as `langgraph`. This contains the majority of the code for the framework, as well as the majority of the unit tests.
- `libs/checkpoint` , published to pypi as `langgraph-checkpoint` contains the base classes for the persistence layer of langgraph. The two main abstractions are BaseCheckpointSaver (base class for persistence of workflow runs step-by-step) and BaseStore (base class for "long-term memory" operations, offering a key-value interface combined with semantic search over documents, used for persisting information across distinct workflow runs). This library is a dependency of both the main langgraph library, as well as implementations of these storage interfaces for specific databases. This library also contains reference implementations
- `libs/checkpoint-postgres` published to pypi as langgraph-checkpoint-postgres, contains implementations of checkpoint and store backed by postgres. Majority of the test coverage is in `libs/langgraph` in the form of tests that run over all storage implementations in the repo.
- `langgraph-java` contains a Java implementation of the langgraph framework, which is in the early stages of development.
To run a particular test file or to pass additional pytest options you can specify the `TEST` variable:
## Feature Overview
```
TEST=path/to/test.py make test
```
langgraph is an orchestration framework (in the style of airflow or temporal) designed for LLM applications, with a focus on streaming output, cyclical and parallel workflows, and interrupt/resume capabilities. Applications built with langgraph are variously called workflows, graphs, cognitive architectures, agents. Key features:
Other pytest arguments can also be supplied inside the `TEST` variable.
1. **Graph-based Architecture**: Build directed computation graphs with nodes and edges
2. **State Management**: Type-safe state schema with custom reducers and transformations
3. **Human-in-the-loop**: Support for interrupts, checkpoints, and tool call review
4. **Persistence**: Save and resume execution with in-memory or database storage
5. **Streaming**: Multiple modes (values, updates, custom) for real-time feedback
6. **Multi-agent Patterns**: Support for network, supervisor, and hierarchical architectures
## Libraries
## Python Development
The repository contains several Python and JavaScript/TypeScript libraries.
Below is a high-level overview:
### Build/Test/Lint Commands
- **checkpoint** – base interfaces for LangGraph checkpointers.
- **checkpoint-postgres** – Postgres implementation of the checkpoint saver.
- **checkpoint-sqlite** – SQLite implementation of the checkpoint saver.
- **cli** – official command-line interface for LangGraph.
- **langgraph** – core framework for building stateful, multi-actor agents.
- **prebuilt** – high-level APIs for creating and running agents and tools.
- **sdk-js** – JS/TS SDK for interacting with the LangGraph REST API.
- **sdk-py** – Python SDK for the LangGraph Server API.
(in the respective subdirectory)
### Dependency map
- Run all tests: `make test`
- Run single test: `make test TEST=path/to/test_file.py::test_function`
- Watch mode tests: `make test_watch`
- Run tests in parallel: `make test_parallel`
- Generate coverage report: `make coverage`
- Format code: `make format`
- Lint code: `make lint`
- Check spelling: `make spell_check`
- Fix spelling: `make spell_fix`
- Build documentation: `make serve-docs` (from repo root)
- Run benchmarks: `make benchmark` or `make benchmark-fast`
The diagram below lists downstream libraries for each production dependency as
declared in that library's `pyproject.toml` (or `package.json`).
### Code Style Guidelines
```text
checkpoint
├── checkpoint-postgres
├── checkpoint-sqlite
├── prebuilt
└── langgraph
- Follow [ruff](https://github.com/astral-sh/ruff) formatting/linting rules
- Use [Google Python Style Guide](https://google.github.io/styleguide/pyguide.html) for docstrings
- Enforce type annotations with mypy (`disallow_untyped_defs = True`)
- Use double quotes for strings
- Maximum line length of 88 characters
- Follow imports sorting with `ruff`
- All functions/classes must have proper docstrings with args/returns
- Write comprehensive unit tests for new features
- Keep backward compatibility
- PR scope should be isolated (changes shouldn't affect multiple packages)
- Use descriptive variable names following Python conventions
- Error handling should use appropriate exception types and messaging
prebuilt
└── langgraph
## Java Development
sdk-py
├── langgraph
└── cli
(in the `langgraph-java` subdirectory)
sdk-js (standalone)
```
### Build/Test/Lint Commands
Changes to a library may impact all of its dependents shown above.
- Build the project: `./gradlew build`
- Run tests: `./gradlew test`
- Run a specific test: `./gradlew test --tests "com.langgraph.package.TestClass.testMethod"`
- Check formatting: `./gradlew spotlessCheck`
- Apply formatting: `./gradlew spotlessApply`
- Run all checks: `./gradlew check`
- Generate Javadoc: `./gradlew javadoc`
### Code Style Guidelines
- Follow standard Java code style (Google Java Style Guide)
- Use 4 spaces for indentation
- Maximum line length of 100 characters
- All public methods/classes must have proper Javadoc with @param/@return tags
- Use descriptive variable names following Java conventions (camelCase)
- Exception handling should use appropriate exception types with descriptive messages
- Favor composition over inheritance
- Use the Builder pattern for complex object creation
- Write comprehensive unit tests for new features
### Python Compatibility Guidelines
- When implementing features from the Python version:
- Maintain semantic equivalence with the Python implementation
- Preserve the same behavior for all public APIs
- Document any intentional differences in behavior with comments
- Pay special attention to collections handling (Python lists vs Java Lists)
- Ensure that iteration order and value handling match Python where relevant
- Use the same test cases as the Python version when possible
- Do not introduce Java-specific shortcuts that would break Python compatibility
- Never add test-specific code to source files - tests should adapt to implementation, not vice versa
### Implementation Mapping
- Always consult and update the `PYTHON_JAVA_MAPPING.md` file when:
- Adding new Java files or classes
- Updating existing Java implementations
- Fixing test failures in Java
- Implementing Python features in Java
- This mapping file documents:
- Where to find equivalent functionality in Python and Java
- Any intentional deviations between implementations
- Implementation status and compatibility notes
- When tests fail, check if the Java implementation matches Python behavior:
- Fix the implementation to match Python semantics whenever possible
- Update tests only if the Python version also differs
- Never create special cases or workarounds in Java just to make tests pass
- Document any implementation differences clearly in the mapping file
- For new features, implement the Python behavior first, then adapt to Java idioms
### Backward Compatibility and API Design
- LangGraph Java has not been released publicly, so there is no need to maintain backward compatibility
- When renaming methods, members, or classes:
- Use the clearest, most intuitive names that match Python semantics
- Remove old/deprecated methods completely rather than marking them as deprecated
- Update all tests and documentation to use the new names
- Do not leave deprecated methods or tests for backward compatibility
### API Design Principles
- Prefer a single, clear way to accomplish each task rather than multiple convenience methods
- Prefer builder patterns over static factory methods where appropriate
- For collections, prefer methods that operate on collections rather than having both single-item and collection variants
- Choose method names that clearly express their purpose and align with Java conventions
- Maintain consistent naming patterns across similar components
- Document the recommended usage pattern in JavaDoc
### Project Structure
- `langgraph-core`: Core functionality of the framework
- `langgraph-checkpoint`: Persistence layer for checkpoints and state management
- `langgraph-examples`: Example applications and usage patterns
### Error Handling
- Use runtime exceptions for unexpected errors
- Use checked exceptions for recoverable errors
- Provide clear error messages that include context about what went wrong
- Validate inputs early to prevent cascading errors
- Ensure all resources are properly closed even in error conditions
+293
View File
@@ -0,0 +1,293 @@
# Contributing to LangGraph
Thank you for being interested in contributing to LangGraph!
## General guidelines
Here are some things to keep in mind for all types of contributions:
- Follow the ["fork and pull request"](https://docs.github.com/en/get-started/exploring-projects-on-github/contributing-to-a-project) workflow.
- Fill out the checked-in pull request template when opening pull requests. Note related issues and tag relevant maintainers.
- Ensure your PR passes formatting, linting, and testing checks before requesting a review.
- If you would like comments or feedback, please open an issue or discussion and tag a maintainer.
- Backwards compatibility is key. Your changes must not be breaking, except in case of critical bug and security fixes.
- Look for duplicate PRs or issues that have already been opened before opening a new one.
- Keep scope as isolated as possible. As a general rule, your changes should not affect more than one package at a time.
### Bugfixes
For bug fixes, please open up an issue before proposing a fix to ensure the proposal properly addresses the underlying problem. In general, bug fixes should all have an accompanying unit test that fails before the fix.
### New features
For new features, please start a new [discussion](https://github.com/langchain-ai/langgraph/discussions), where the maintainers will help with scoping out the necessary changes.
## Contribute Documentation
Documentation is a vital part of LangGraph. We welcome both new documentation for new features and
community improvements to our current documentation. Please read the resources below before getting started:
- [Documentation style guide](#documentation-style-guide)
- [Documentation setup](#setup)
## Documentation Style Guide
As LangGraph continues to grow, the surface area of documentation required to cover it continues to grow too.
This page provides guidelines for anyone writing documentation for LangGraph, as well as some of our philosophies around organization and structure.
## Philosophy
LangGraph's documentation follows the [Diataxis framework](https://diataxis.fr).
Under this framework, all documentation falls under one of four categories: [Tutorials](#tutorials),
[How-to guides](#how-to-guides),
[References](#references), and [Explanations (aka conceptual guides)](#conceptual-guide).
### Tutorials
Tutorials are lessons that take the reader through a practical activity. Their purpose is to help the user
gain understanding of concepts and how they interact by showing one way to achieve some goal in a hands-on way.
They should **avoid** giving
multiple permutations of ways to achieve that goal in-depth. Choice is burdensome. Instead, they should guide a new user through a recommended path to accomplishing a concrete goal. While the end result of a tutorial does not necessarily need to
be completely production-ready, it should be useful and practically satisfy the goal that you clearly stated in the tutorial's introduction.
To quote the Diataxis website:
> A tutorial serves the user’s *acquisition* of skills and knowledge - their study. Its purpose is not to help the user get something done, but to help them learn.
In LangGraph, these are often higher level guides that show off end-to-end use cases.
Some examples include:
- [Build a Customer Support Bot](https://langchain-ai.github.io/langgraph/tutorials/customer-support/customer-support/)
- [Build a SQL Agent](https://langchain-ai.github.io/langgraph/tutorials/sql-agent/)
Here are some high-level tips on writing a good tutorial:
- Focus on guiding the user to get something done, but keep in mind the end-goal is more to impart principles than to create a perfect production system.
- Be specific, not abstract and follow one path.
- No need to go deeply into alternative approaches, but it’s ok to reference them, ideally with a link to an appropriate how-to guide.
- Get "a point on the board" as soon as possible - something the user can run that outputs something.
- You can iterate and expand afterwards.
- Try to frequently checkpoint at given steps where the user can run code and see progress.
- Focus on results, not technical explanation.
- Crosslink heavily to appropriate conceptual/reference pages
- The first time you mention a LangGraph concept, use its full name (e.g. "human-in-the-loop"), and link to its conceptual/other documentation page.
- It's also helpful to add a prerequisite callout that links to any pages with necessary background information.
- End with a recap/next steps section summarizing what the tutorial covered and future reading, such as related how-to guides.
- Use phrases like "Next we can run X & Y. We will expect Z.". Then afterwards, use language like "Notice Z" that recalls our expectations and directs the reader's attention to the topic we are trying to teach.
- Do not shy away from repetition.
### How-to guides
A how-to guide, as the name implies, demonstrates how to do something discrete and specific.
It should assume that the user is already familiar with underlying concepts, and is trying to solve an immediate problem, but
should still give some background or list the scenarios where the information contained within can be relevant.
They can and should discuss alternatives if one approach may be better than another in certain cases.
To quote the Diataxis website:
> A how-to guide serves the work of the already-competent user, whom you can assume to know what they want to do, and to be able to follow your instructions correctly.
Some examples include:
- [How to add persistence to your graph](https://langchain-ai.github.io/langgraph/how-tos/persistence/)
- [How to view and update past graph state](https://langchain-ai.github.io/langgraph/how-tos/human_in_the_loop/time-travel/)
Here are some high-level tips on writing a good how-to guide:
- Clearly explain what you are guiding the user through at the start
- Assume higher intent than a tutorial and show what the user needs to do to get that task done
- Assume familiarity of concepts, but explain why suggested actions are helpful
- Crosslink heavily to conceptual/reference pages
- Discuss alternatives and responses to real-world tradeoffs that may arise when solving a problem
- Use lots of example code, ideally within complete code blocks that the reader can copy and run.
- End with a recap/next steps section summarizing what the tutorial covered and future reading, such as other related how-to guides
### Conceptual guides
LangGraph's conceptual guides fall under the **Explanation** quadrant of Diataxis. They should cover LangChain terms and concepts
in a more abstract way than how-to guides or tutorials, and should be geared towards curious users interested in
gaining a deeper understanding of the framework. Try to avoid excessively large code examples. The goal here is to
impart perspective to the user rather than to finish a practical project. These guides should cover **why** things work they way they do.
To quote the Diataxis website:
> The perspective of explanation is higher and wider than that of the other types. It does not take the user’s eye-level view, as in a how-to guide, or a close-up view of the machinery, like reference material. Its scope in each case is a topic - “an area of knowledge”, that somehow has to be bounded in a reasonable, meaningful way.
Some examples include:
- [What does it mean to be agentic?](https://langchain-ai.github.io/langgraph/concepts/high_level/)
- [Tool calling](https://langchain-ai.github.io/langgraph/concepts/agentic_concepts/#tool-calling)
Here are some high-level tips on writing a good conceptual guide:
- Explain design decisions. Why does concept X exist and why was it designed this way?
- Use analogies and reference other concepts and alternatives
- Avoid blending in too much reference content
- You can and should reference content covered in other guides, but make sure to link to them
### References
References contain detailed, low-level information that describes exactly what functionality exists and how to use it.
In LangGraph, this is mainly our API reference pages, which are populated from docstrings within code.
References pages are generally not read end-to-end, but are consulted as necessary when a user needs to know
how to use something specific.
To quote the Diataxis website:
> The only purpose of a reference guide is to describe, as succinctly as possible, and in an orderly way. Whereas the content of tutorials and how-to guides are led by needs of the user, reference material is led by the product it describes.
Many of the reference pages in LangChain are automatically generated from code,
but here are some high-level tips on writing a good docstring:
- Be concise
- Discuss special cases and deviations from a user's expectations
- Go into detail on required inputs and outputs
- Light details on when one might use the feature are fine, but in-depth details belong in other sections.
Each category serves a distinct purpose and requires a specific approach to writing and structuring the content.
## General guidelines
Here are some other guidelines you should think about when writing and organizing documentation.
We generally do not merge new tutorials from outside contributors without an actue need.
We welcome updates as well as new integration docs, how-tos, and references.
### Avoid duplication
Multiple pages that cover the same material in depth are difficult to maintain and cause confusion. There should
be only one (very rarely two), canonical pages for a given concept or feature. Instead, you should link to other guides.
### Link to other sections
Because sections of the docs do not exist in a vacuum, it is important to link to other sections as often as possible
to allow a developer to learn more about an unfamiliar topic inline.
This includes linking to the API references as well as conceptual sections!
### Be concise
In general, take a less-is-more approach. If a section with a good explanation of a concept already exists, you should link to it rather than
re-explain it, unless the concept you are documenting presents some new wrinkle.
Be concise, including in code samples.
### General style
- Use active voice and present tense whenever possible
- Use examples and code snippets to illustrate concepts and usage
- Use appropriate header levels (`#`, `##`, `###`, etc.) to organize the content hierarchically
- Use fewer cells with more code to make copy/paste easier
- Use bullet points and numbered lists to break down information into easily digestible chunks
- Use tables (especially for **Reference** sections) and diagrams often to present information visually
- Include the table of contents for longer documentation pages to help readers navigate the content, but hide it for shorter pages
## Setup
LangChain documentation consists of two components:
1. Main Documentation: Hosted at [https://langchain-ai.github.io](https://langchain-ai.github.io/langgraph/),
this comprehensive resource serves as the primary user-facing documentation.
It covers a wide array of topics, including tutorials, use cases, integrations,
and more, offering extensive guidance on building with LangGraph.
The content for this documentation lives in the `/docs` directory of the monorepo.
2. In-code Documentation: This is documentation of the codebase itself, which is also
used to generate the externally facing [API Reference](https://langchain-ai.github.io/langgraph/reference/graphs/).
The content for the API reference is autogenerated by scanning the docstrings in the codebase. For this reason we ask that developers document their code well.
We appreciate all contributions to the documentation, whether it be fixing a typo,
adding a new tutorial or example and whether it be in the main documentation or the API Reference.
### 📜 Main Documentation
The content for the main documentation is located in the `/docs` directory of the monorepo.
The documentation is written using a combination of ipython notebooks (`.ipynb` files)
and markdown (`.md` files). The notebooks are converted to markdown
and then built using [MkDocs](https://www.mkdocs.org/).
Feel free to make contributions to the main documentation! 🥰
After modifying the documentation:
1. Run the linting and formatting commands (see below) to ensure that the documentation is well-formatted and free of errors.
2. Optionally build the documentation locally to verify that the changes look good.
3. Make a pull request with the changes.
### ⚒️ Linting and Building Documentation Locally
After writing up the documentation, you may want to lint and build the documentation
locally to ensure that it looks good and is free of errors.
If you're unable to build it locally that's okay as well, as you will be able to
see a preview of the documentation on the pull request page.
From the **monorepo root**, run the following command to install the dependencies:
```bash
poetry install --with docs --no-root
```
#### Building
The code that builds the documentation is located in the `/docs` directory of the monorepo.
Before building the documentation, it is always a good idea to clean the build directory:
```bash
make clean-docs
```
You can build and preview the documentation as outlined below:
```bash
make serve-docs
```
#### Linting
The documentation is linted from the **monorepo root**. To lint it, run the following from there:
```bash
make spellcheck
```
### ️In-code Documentation
The in-code documentation is autogenerated from docstrings.
For the API reference to be useful, the codebase must be well-documented. This means that all functions, classes, and methods should have a docstring that explains what they do, what the arguments are, and what the return value is. This is a good practice in general, but it is especially important for LangChain because the API reference is the primary resource for developers to understand how to use the codebase.
We generally follow the [Google Python Style Guide](https://google.github.io/styleguide/pyguide.html#38-comments-and-docstrings) for docstrings.
Here is an example of a well-documented function:
```python
def my_function(arg1: int, arg2: str) -> float:
"""This is a short description of the function. (It should be a single sentence.)
This is a longer description of the function. It should explain what
the function does, what the arguments are, and what the return value is.
It should wrap at 88 characters.
Examples:
This is a section for examples of how to use the function.
.. code-block:: python
my_function(1, "hello")
Args:
arg1: This is a description of arg1. We do not need to specify the type since
it is already specified in the function signature.
arg2: This is a description of arg2.
Returns:
This is a description of the return value.
"""
return 3.14
```
-68
View File
@@ -1,68 +0,0 @@
# Define the directories containing projects
LIBS_DIRS := $(wildcard libs/*)
# Default target
.PHONY: all
all: lint format lock test
# Install dependencies for all projects
.PHONY: install
install:
@echo "Creating virtual environment..."
@uv venv
@for dir in $(LIBS_DIRS); do \
if [ -f $$dir/pyproject.toml ]; then \
echo "Installing dependencies for $$dir"; \
uv pip install -e $$dir; \
fi; \
done
# Lint all projects
.PHONY: lint
lint:
@for dir in $(LIBS_DIRS); do \
if [ -f $$dir/Makefile ]; then \
echo "Running lint in $$dir"; \
$(MAKE) -C $$dir lint; \
fi; \
done
# Format all projects
.PHONY: format
format:
@for dir in $(LIBS_DIRS); do \
if [ -f $$dir/Makefile ]; then \
echo "Running format in $$dir"; \
$(MAKE) -C $$dir format; \
fi; \
done
# Lock all projects
.PHONY: lock
lock:
@for dir in $(LIBS_DIRS); do \
if [ -f $$dir/Makefile ]; then \
echo "Running lock in $$dir"; \
(cd $$dir && uv lock); \
fi; \
done
# Lock all projects and upgrade dependencies
.PHONY: lock-upgrade
lock-upgrade:
@for dir in $(LIBS_DIRS); do \
if [ -f $$dir/Makefile ]; then \
echo "Running lock-upgrade in $$dir"; \
(cd $$dir && uv lock --upgrade); \
fi; \
done
# Test all projects
.PHONY: test
test:
@for dir in $(LIBS_DIRS); do \
if [ -f $$dir/Makefile ]; then \
echo "Running test in $$dir"; \
$(MAKE) -C $$dir test; \
fi; \
done
+313 -65
View File
@@ -1,91 +1,339 @@
<picture class="github-only">
<source media="(prefers-color-scheme: light)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg">
<source media="(prefers-color-scheme: dark)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_light.svg">
<img alt="LangGraph Logo" src="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg" width="80%">
</picture>
# 🦜🕸️LangGraph
<div>
<br>
</div>
[![Version](https://img.shields.io/pypi/v/langgraph.svg)](https://pypi.org/project/langgraph/)
![Version](https://img.shields.io/pypi/v/langgraph)
[![Downloads](https://static.pepy.tech/badge/langgraph/month)](https://pepy.tech/project/langgraph)
[![Open Issues](https://img.shields.io/github/issues-raw/langchain-ai/langgraph)](https://github.com/langchain-ai/langgraph/issues)
[![Docs](https://img.shields.io/badge/docs-latest-blue)](https://docs.langchain.com/oss/python/langgraph/overview)
[![Docs](https://img.shields.io/badge/docs-latest-blue)](https://langchain-ai.github.io/langgraph/)
Trusted by companies shaping the future of agents – including Klarna, Replit, Elastic, and more – LangGraph is a low-level orchestration framework for building, managing, and deploying long-running, stateful agents.
⚡ Building language agents as graphs ⚡
## Get started
> [!NOTE]
> Looking for the JS version? See the [JS repo](https://github.com/langchain-ai/langgraphjs) and the [JS docs](https://langchain-ai.github.io/langgraphjs/).
Install LangGraph:
## Overview
```
[LangGraph](https://langchain-ai.github.io/langgraph/) is a library for building
stateful, multi-actor applications with LLMs, used to create agent and multi-agent
workflows. Check out an introductory tutorial [here](https://langchain-ai.github.io/langgraph/tutorials/introduction/).
LangGraph is inspired by [Pregel](https://research.google/pubs/pub37252/) and [Apache Beam](https://beam.apache.org/). The public interface draws inspiration from [NetworkX](https://networkx.org/documentation/latest/). LangGraph is built by LangChain Inc, the creators of LangChain, but can be used without LangChain.
### Why use LangGraph?
LangGraph powers [production-grade agents](https://www.langchain.com/built-with-langgraph), trusted by Linkedin, Uber, Klarna, GitLab, and many more. LangGraph provides fine-grained control over both the flow and state of your agent applications. It implements a central [persistence layer](https://langchain-ai.github.io/langgraph/concepts/persistence/), enabling features that are common to most agent architectures:
- **Memory**: LangGraph persists arbitrary aspects of your application's state,
supporting memory of conversations and other updates within and across user
interactions;
- **Human-in-the-loop**: Because state is checkpointed, execution can be interrupted
and resumed, allowing for decisions, validation, and corrections at key stages via
human input.
Standardizing these components allows individuals and teams to focus on the behavior
of their agent, instead of its supporting infrastructure.
Through [LangGraph Platform](#langgraph-platform), LangGraph also provides tooling for
the development, deployment, debugging, and monitoring of your applications.
LangGraph integrates seamlessly with
[LangChain](https://python.langchain.com/docs/introduction/) and
[LangSmith](https://docs.smith.langchain.com/) (but does not require them).
To learn more about LangGraph, check out our first LangChain Academy
course, *Introduction to LangGraph*, available for free
[here](https://academy.langchain.com/courses/intro-to-langgraph).
### LangGraph Platform
[LangGraph Platform](https://langchain-ai.github.io/langgraph/concepts/langgraph_platform) is infrastructure for deploying LangGraph agents. It is a commercial solution for deploying agentic applications to production, built on the open-source LangGraph framework. The LangGraph Platform consists of several components that work together to support the development, deployment, debugging, and monitoring of LangGraph applications: [LangGraph Server](https://langchain-ai.github.io/langgraph/concepts/langgraph_server) (APIs), [LangGraph SDKs](https://langchain-ai.github.io/langgraph/concepts/sdk) (clients for the APIs), [LangGraph CLI](https://langchain-ai.github.io/langgraph/concepts/langgraph_cli) (command line tool for building the server), and [LangGraph Studio](https://langchain-ai.github.io/langgraph/concepts/langgraph_studio) (UI/debugger).
See deployment options [here](https://langchain-ai.github.io/langgraph/concepts/deployment_options/)
(includes a free tier).
Here are some common issues that arise in complex deployments, which LangGraph Platform addresses:
- **Streaming support**: LangGraph Server provides [multiple streaming modes](https://langchain-ai.github.io/langgraph/concepts/streaming) optimized for various application needs
- **Background runs**: Runs agents asynchronously in the background
- **Support for long running agents**: Infrastructure that can handle long running processes
- **[Double texting](https://langchain-ai.github.io/langgraph/concepts/double_texting)**: Handle the case where you get two messages from the user before the agent can respond
- **Handle burstiness**: Task queue for ensuring requests are handled consistently without loss, even under heavy loads
## Installation
```shell
pip install -U langgraph
```
Create a simple workflow:
## Example
```python
from langgraph.graph import START, StateGraph
from typing_extensions import TypedDict
Let's build a tool-calling [ReAct-style](https://langchain-ai.github.io/langgraph/concepts/agentic_concepts/#react-implementation) agent that uses a search tool!
class State(TypedDict):
text: str
def node_a(state: State) -> dict:
return {"text": state["text"] + "a"}
def node_b(state: State) -> dict:
return {"text": state["text"] + "b"}
graph = StateGraph(State)
graph.add_node("node_a", node_a)
graph.add_node("node_b", node_b)
graph.add_edge(START, "node_a")
graph.add_edge("node_a", "node_b")
print(graph.compile().invoke({"text": ""}))
# {'text': 'ab'}
```shell
pip install langchain-anthropic
```
Get started with the [LangGraph Quickstart](https://docs.langchain.com/oss/python/langgraph/quickstart).
```shell
export ANTHROPIC_API_KEY=sk-...
```
To quickly build agents with LangChain's `create_agent` (built on LangGraph), see the [LangChain Agents documentation](https://docs.langchain.com/oss/python/langchain/agents).
Optionally, we can set up [LangSmith](https://docs.smith.langchain.com/) for best-in-class observability.
## Core benefits
```shell
export LANGSMITH_TRACING=true
export LANGSMITH_API_KEY=lsv2_sk_...
```
LangGraph provides low-level supporting infrastructure for *any* long-running, stateful workflow or agent. LangGraph does not abstract prompts or architecture, and provides the following central benefits:
The simplest way to create a tool-calling agent in LangGraph is to use `create_react_agent`:
- [Durable execution](https://docs.langchain.com/oss/python/langgraph/durable-execution): Build agents that persist through failures and can run for extended periods, automatically resuming from exactly where they left off.
- [Human-in-the-loop](https://docs.langchain.com/oss/python/langgraph/interrupts): Seamlessly incorporate human oversight by inspecting and modifying agent state at any point during execution.
- [Comprehensive memory](https://docs.langchain.com/oss/python/langgraph/memory): Create truly stateful agents with both short-term working memory for ongoing reasoning and long-term persistent memory across sessions.
- [Debugging with LangSmith](http://www.langchain.com/langsmith): Gain deep visibility into complex agent behavior with visualization tools that trace execution paths, capture state transitions, and provide detailed runtime metrics.
- [Production-ready deployment](https://docs.langchain.com/langsmith/app-development): Deploy sophisticated agent systems confidently with scalable infrastructure designed to handle the unique challenges of stateful, long-running workflows.
<details open>
<summary>High-level implementation</summary>
## LangGraph’s ecosystem
```python
from langgraph.prebuilt import create_react_agent
from langgraph.checkpoint.memory import MemorySaver
from langchain_anthropic import ChatAnthropic
from langchain_core.tools import tool
While LangGraph can be used standalone, it also integrates seamlessly with any LangChain product, giving developers a full suite of tools for building agents. To improve your LLM application development, pair LangGraph with:
# Define the tools for the agent to use
@tool
def search(query: str):
"""Call to surf the web."""
# This is a placeholder, but don't tell the LLM that...
if "sf" in query.lower() or "san francisco" in query.lower():
return "It's 60 degrees and foggy."
return "It's 90 degrees and sunny."
- [LangSmith](http://www.langchain.com/langsmith) — Helpful for agent evals and observability. Debug poor-performing LLM app runs, evaluate agent trajectories, gain visibility in production, and improve performance over time.
- [LangSmith Deployment](https://docs.langchain.com/langsmith/deployments) — Deploy and scale agents effortlessly with a purpose-built deployment platform for long running, stateful workflows. Discover, reuse, configure, and share agents across teams — and iterate quickly with visual prototyping in [LangGraph Studio](https://docs.langchain.com/oss/python/langgraph/studio).
- [LangChain](https://docs.langchain.com/oss/python/langchain/overview) – Provides integrations and composable components to streamline LLM application development.
> [!NOTE]
> Looking for the JS version of LangGraph? See the [JS repo](https://github.com/langchain-ai/langgraphjs) and the [JS docs](https://docs.langchain.com/oss/javascript/langgraph/overview).
tools = [search]
model = ChatAnthropic(model="claude-3-5-sonnet-latest", temperature=0)
## Additional resources
# Initialize memory to persist state between graph runs
checkpointer = MemorySaver()
- [Guides](https://docs.langchain.com/oss/python/langgraph/guides): Quick, actionable code snippets for topics such as streaming, adding memory & persistence, and design patterns (e.g. branching, subgraphs, etc.).
- [Reference](https://reference.langchain.com/python/langgraph/): Detailed reference on core classes, methods, how to use the graph and checkpointing APIs, and higher-level prebuilt components.
- [Examples](https://docs.langchain.com/oss/python/langgraph/agentic-rag): Guided examples on getting started with LangGraph.
- [LangChain Forum](https://forum.langchain.com/): Connect with the community and share all of your technical questions, ideas, and feedback.
- [LangChain Academy](https://academy.langchain.com/courses/intro-to-langgraph): Learn the basics of LangGraph in our free, structured course.
- [Case studies](https://www.langchain.com/built-with-langgraph): Hear how industry leaders use LangGraph to ship AI applications at scale.
app = create_react_agent(model, tools, checkpointer=checkpointer)
## Acknowledgements
# Use the agent
final_state = app.invoke(
{"messages": [{"role": "user", "content": "what is the weather in sf"}]},
config={"configurable": {"thread_id": 42}}
)
final_state["messages"][-1].content
```
```
"Based on the search results, I can tell you that the current weather in San Francisco is:\n\nTemperature: 60 degrees Fahrenheit\nConditions: Foggy\n\nSan Francisco is known for its microclimates and frequent fog, especially during the summer months. The temperature of 60°F (about 15.5°C) is quite typical for the city, which tends to have mild temperatures year-round. The fog, often referred to as "Karl the Fog" by locals, is a characteristic feature of San Francisco\'s weather, particularly in the mornings and evenings.\n\nIs there anything else you\'d like to know about the weather in San Francisco or any other location?"
```
LangGraph is inspired by [Pregel](https://research.google/pubs/pub37252/) and [Apache Beam](https://beam.apache.org/). The public interface draws inspiration from [NetworkX](https://networkx.org/documentation/latest/). LangGraph is built by LangChain Inc, the creators of LangChain, but can be used without LangChain.
Now when we pass the same <code>"thread_id"</code>, the conversation context is retained via the saved state (i.e. stored list of messages)
```python
final_state = app.invoke(
{"messages": [{"role": "user", "content": "what about ny"}]},
config={"configurable": {"thread_id": 42}}
)
final_state["messages"][-1].content
```
```
"Based on the search results, I can tell you that the current weather in New York City is:\n\nTemperature: 90 degrees Fahrenheit (approximately 32.2 degrees Celsius)\nConditions: Sunny\n\nThis weather is quite different from what we just saw in San Francisco. New York is experiencing much warmer temperatures right now. Here are a few points to note:\n\n1. The temperature of 90°F is quite hot, typical of summer weather in New York City.\n2. The sunny conditions suggest clear skies, which is great for outdoor activities but also means it might feel even hotter due to direct sunlight.\n3. This kind of weather in New York often comes with high humidity, which can make it feel even warmer than the actual temperature suggests.\n\nIt's interesting to see the stark contrast between San Francisco's mild, foggy weather and New York's hot, sunny conditions. This difference illustrates how varied weather can be across different parts of the United States, even on the same day.\n\nIs there anything else you'd like to know about the weather in New York or any other location?"
```
</details>
> [!TIP]
> LangGraph is a **low-level** framework that allows you to implement any custom agent
architectures. Click on the low-level implementation below to see how to implement a
tool-calling agent from scratch.
<details>
<summary>Low-level implementation</summary>
```python
from typing import Literal
from langchain_anthropic import ChatAnthropic
from langchain_core.tools import tool
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import END, START, StateGraph, MessagesState
from langgraph.prebuilt import ToolNode
# Define the tools for the agent to use
@tool
def search(query: str):
"""Call to surf the web."""
# This is a placeholder, but don't tell the LLM that...
if "sf" in query.lower() or "san francisco" in query.lower():
return "It's 60 degrees and foggy."
return "It's 90 degrees and sunny."
tools = [search]
tool_node = ToolNode(tools)
model = ChatAnthropic(model="claude-3-5-sonnet-latest", temperature=0).bind_tools(tools)
# Define the function that determines whether to continue or not
def should_continue(state: MessagesState) -> Literal["tools", END]:
messages = state['messages']
last_message = messages[-1]
# If the LLM makes a tool call, then we route to the "tools" node
if last_message.tool_calls:
return "tools"
# Otherwise, we stop (reply to the user)
return END
# Define the function that calls the model
def call_model(state: MessagesState):
messages = state['messages']
response = model.invoke(messages)
# We return a list, because this will get added to the existing list
return {"messages": [response]}
# Define a new graph
workflow = StateGraph(MessagesState)
# Define the two nodes we will cycle between
workflow.add_node("agent", call_model)
workflow.add_node("tools", tool_node)
# Set the entrypoint as `agent`
# This means that this node is the first one called
workflow.add_edge(START, "agent")
# We now add a conditional edge
workflow.add_conditional_edges(
# First, we define the start node. We use `agent`.
# This means these are the edges taken after the `agent` node is called.
"agent",
# Next, we pass in the function that will determine which node is called next.
should_continue,
)
# We now add a normal edge from `tools` to `agent`.
# This means that after `tools` is called, `agent` node is called next.
workflow.add_edge("tools", 'agent')
# Initialize memory to persist state between graph runs
checkpointer = MemorySaver()
# Finally, we compile it!
# This compiles it into a LangChain Runnable,
# meaning you can use it as you would any other runnable.
# Note that we're (optionally) passing the memory when compiling the graph
app = workflow.compile(checkpointer=checkpointer)
# Use the agent
final_state = app.invoke(
{"messages": [{"role": "user", "content": "what is the weather in sf"}]},
config={"configurable": {"thread_id": 42}}
)
final_state["messages"][-1].content
```
<b>Step-by-step Breakdown</b>:
<details>
<summary>Initialize the model and tools.</summary>
<ul>
<li>
We use <code>ChatAnthropic</code> as our LLM. <strong>NOTE:</strong> we need to make sure the model knows that it has these tools available to call. We can do this by converting the LangChain tools into the format for OpenAI tool calling using the <code>.bind_tools()</code> method.
</li>
<li>
We define the tools we want to use - a search tool in our case. It is really easy to create your own tools - see documentation here on how to do that <a href="https://python.langchain.com/docs/how_to/custom_tools/">here</a>.
</li>
</ul>
</details>
<details>
<summary>Initialize graph with state.</summary>
<ul>
<li>We initialize graph (<code>StateGraph</code>) by passing state schema (in our case <code>MessagesState</code>)</li>
<li><code>MessagesState</code> is a prebuilt state schema that has one attribute -- a list of LangChain <code>Message</code> objects, as well as logic for merging the updates from each node into the state.</li>
</ul>
</details>
<details>
<summary>Define graph nodes.</summary>
There are two main nodes we need:
<ul>
<li>The <code>agent</code> node: responsible for deciding what (if any) actions to take.</li>
<li>The <code>tools</code> node that invokes tools: if the agent decides to take an action, this node will then execute that action.</li>
</ul>
</details>
<details>
<summary>Define entry point and graph edges.</summary>
First, we need to set the entry point for graph execution - <code>agent</code> node.
Then we define one normal and one conditional edge. Conditional edge means that the destination depends on the contents of the graph's state (<code>MessagesState</code>). In our case, the destination is not known until the agent (LLM) decides.
<ul>
<li>Conditional edge: after the agent is called, we should either:
<ul>
<li>a. Run tools if the agent said to take an action, OR</li>
<li>b. Finish (respond to the user) if the agent did not ask to run tools</li>
</ul>
</li>
<li>Normal edge: after the tools are invoked, the graph should always return to the agent to decide what to do next</li>
</ul>
</details>
<details>
<summary>Compile the graph.</summary>
<ul>
<li>
When we compile the graph, we turn it into a LangChain
<a href="https://python.langchain.com/docs/concepts/runnables/">Runnable</a>,
which automatically enables calling <code>.invoke()</code>, <code>.stream()</code> and <code>.batch()</code>
with your inputs
</li>
<li>
We can also optionally pass checkpointer object for persisting state between graph runs, and enabling memory,
human-in-the-loop workflows, time travel and more. In our case we use <code>MemorySaver</code> -
a simple in-memory checkpointer
</li>
</ul>
</details>
<details>
<summary>Execute the graph.</summary>
<ol>
<li>LangGraph adds the input message to the internal state, then passes the state to the entrypoint node, <code>"agent"</code>.</li>
<li>The <code>"agent"</code> node executes, invoking the chat model.</li>
<li>The chat model returns an <code>AIMessage</code>. LangGraph adds this to the state.</li>
<li>Graph cycles the following steps until there are no more <code>tool_calls</code> on <code>AIMessage</code>:
<ul>
<li>If <code>AIMessage</code> has <code>tool_calls</code>, <code>"tools"</code> node executes</li>
<li>The <code>"agent"</code> node executes again and returns <code>AIMessage</code></li>
</ul>
</li>
<li>Execution progresses to the special <code>END</code> value and outputs the final state. And as a result, we get a list of all our chat messages as output.</li>
</ol>
</details>
</details>
## Documentation
* [Tutorials](https://langchain-ai.github.io/langgraph/tutorials/): Learn to build with LangGraph through guided examples.
* [How-to Guides](https://langchain-ai.github.io/langgraph/how-tos/): Accomplish specific things within LangGraph, from streaming, to adding memory & persistence, to common design patterns (branching, subgraphs, etc.), these are the place to go if you want to copy and run a specific code snippet.
* [Conceptual Guides](https://langchain-ai.github.io/langgraph/concepts/high_level/): In-depth explanations of the key concepts and principles behind LangGraph, such as nodes, edges, state and more.
* [API Reference](https://langchain-ai.github.io/langgraph/reference/graphs/): Review important classes and methods, simple examples of how to use the graph and checkpointing APIs, higher-level prebuilt components and more.
* [LangGraph Platform](https://langchain-ai.github.io/langgraph/concepts/#langgraph-platform): LangGraph Platform is a commercial solution for deploying agentic applications in production, built on the open-source LangGraph framework.
## Resources
* [Built with LangGraph](https://www.langchain.com/built-with-langgraph): Hear how industry leaders use LangGraph to ship powerful, production-ready AI applications.
## Contributing
For more information on how to contribute, see [here](https://github.com/langchain-ai/langgraph/blob/main/CONTRIBUTING.md).
+90
View File
@@ -0,0 +1,90 @@
##############################
## Java
##############################
.mtj.tmp/
*.class
*.jar
*.war
*.ear
*.nar
hs_err_pid*
replay_pid*
##############################
## Maven
##############################
target/
pom.xml.tag
pom.xml.releaseBackup
pom.xml.versionsBackup
pom.xml.next
pom.xml.bak
release.properties
dependency-reduced-pom.xml
buildNumber.properties
.mvn/timing.properties
.mvn/wrapper/maven-wrapper.jar
##############################
## Gradle
##############################
bin/
build/
.gradle
.gradletasknamecache
gradle-app.setting
!gradle-wrapper.jar
##############################
## IntelliJ
##############################
out/
.idea/
.idea_modules/
*.iml
*.ipr
*.iws
##############################
## Eclipse
##############################
.settings/
bin/
tmp/
.metadata
.classpath
.project
*.tmp
*.bak
*.swp
*~.nib
local.properties
.loadpath
.factorypath
##############################
## NetBeans
##############################
nbproject/private/
build/
nbbuild/
dist/
nbdist/
nbactions.xml
nb-configuration.xml
##############################
## Visual Studio Code
##############################
.vscode/
.code-workspace
##############################
## OS X
##############################
.DS_Store
##############################
## Miscellaneous
##############################
*.log
+107
View File
@@ -0,0 +1,107 @@
# Channel Initialization in Java LangGraph
This document explains how channel initialization is handled in the Java implementation of LangGraph, matching Python's behavior.
## Current Implementation
### Java Implementation (Python-Compatible)
In the Java implementation:
1. **First Superstep Behavior**:
- Only nodes that have the input channel as one of their triggers run in the first superstep
- This matches Python's behavior for graph execution
2. **Channel Reading**:
- Channels that haven't been initialized return `null` values (instead of throwing exceptions)
- Nodes are expected to handle potentially `null` values from uninitialized channels
3. **Graph Execution**:
- Subsequent supersteps only execute nodes that:
- Subscribe to a channel that was updated
- OR have a trigger matching a channel that was updated
## Key Distinctions
There's an important distinction between:
1. **Input channels** - Channels from which the node reads values when it executes
2. **Trigger channels** - Channels that determine when this node should execute
In our Java implementation:
- `channels` property defines which channels the node reads from
- `triggerChannels` property defines which channels can cause the node to execute
This naming is more intuitive and aligns better with the conceptual distinction between reading from a channel and being triggered by a channel.
## Implementation Details
The Python-compatible implementation in Java LangGraph makes the following changes:
1. Modified `TaskPlanner.plan()` to only execute nodes with input channel triggers in the first superstep
2. Updated tests to:
- Add triggers for nodes that should execute on first superstep (e.g., `trigger("input")`)
- Remove unnecessary manual channel initialization that was previously used to avoid EmptyChannelException
- Use input maps to provide initial values instead of `channel.update()`
- Only keep manual initialization in specific test cases that need it (like the mix of initialized/uninitialized channels test)
3. Clarified the distinction between "subscribe to read" and "trigger to execute" semantics
## Recommended Practices
When building graphs with the Java implementation, follow these practices to ensure Python compatibility:
### 1. Add trigger channels to nodes
Always add appropriate trigger channels to nodes that should execute in the first superstep:
```java
PregelNode node = new PregelNode.Builder("node", executable)
.channel("input") // Channel to read from
.triggerChannel("input") // Channel that triggers execution
.writer("output")
.build();
```
You can also add multiple trigger channels if needed:
```java
PregelNode node = new PregelNode.Builder("node", executable)
.channel("input1")
.channel("input2")
.triggerChannel("input1") // Will trigger on this channel
.triggerChannel("input2") // And also on this channel
.writer("output")
.build();
```
### 2. Handle uninitialized channels gracefully
Inside node execution logic, handle potentially uninitialized channels using default values:
```java
// Handle uninitialized channels with a default value
Integer input = 0; // Default value for uninitialized channel
if (inputs.containsKey("inputChannel") && inputs.get("inputChannel") != null) {
input = (Integer) inputs.get("inputChannel");
}
```
### 3. Provide initial values through input map
Instead of manually initializing channels, provide initial values through the input map:
```java
// DO NOT do this:
// channel.update(Collections.singletonList(initialValue));
// Instead, provide values in the input map:
Map<String, Object> input = new HashMap<>();
input.put("inputChannel", initialValue);
Object result = pregel.invoke(input, null);
```
### 4. Remember execution rules
- Only nodes with input channel as a trigger run in the first superstep
- In subsequent supersteps, nodes run if they subscribe to or have a trigger matching an updated channel
- Uninitialized channels return `null` or empty collections rather than throwing exceptions
+88
View File
@@ -0,0 +1,88 @@
# Python-Java Implementation Mapping
This document records the mapping between Python and Java implementations of LangGraph, highlighting any deliberate differences and their rationale.
## Core Components
### Channels
| Component | Python Path | Java Path | Deviations |
|-----------|-------------|-----------|------------|
| BaseChannel | langgraph/channels/base.py | com.langgraph.channels.BaseChannel | Java uses interface with default methods instead of Python's abstract base class. Channel returns null or empty values when uninitialized, rather than throwing exceptions. |
| AbstractChannel | langgraph/channels/base.py | com.langgraph.channels.AbstractChannel | Java implementation provides default functionality shared by channel implementations. Added Python compatibility for uninitialized channels. |
| TopicChannel | langgraph/channels/topic_channel.py | com.langgraph.channels.TopicChannel | Java implementation preserves Python's multi-value behavior while using Java collections. Returns empty list for uninitialized channels. |
| LastValue | langgraph/channels/last_value.py | com.langgraph.channels.LastValue | Returns null for uninitialized channels to match Python behavior. |
| EphemeralValue | langgraph/channels/ephemeral_value.py | com.langgraph.channels.EphemeralValue | Returns null for uninitialized channels to match Python behavior. |
| Channels (utility) | langgraph/channels/__init__.py | com.langgraph.channels.Channels | Java uses utility class with static methods instead of module-level functions. |
### Pregel Algorithm
| Component | Python Path | Java Path | Deviations |
|-----------|-------------|-----------|------------|
| PregelNode | langgraph/pregel/algorithm.py | com.langgraph.pregel.PregelNode | Java exposes these concepts with clearer naming: 'channels' (input channels to read from) and 'triggerChannels' (channels that trigger execution). Java now supports multiple trigger channels like Python. |
| Pregel | langgraph/pregel/pregel.py | com.langgraph.pregel.Pregel | Java uses Builder pattern instead of Python's initialization parameters. Functionally equivalent. |
| PregelLoop | langgraph/pregel/pregel_loop.py | com.langgraph.pregel.execute.PregelLoop | Implementation follows Java conventions with robust cycle detection. Ensures runs complete when possible by executing a final validation step before throwing recursion errors. |
| Runner Functions | langgraph/pregel/runner.py | com.langgraph.pregel.execute.SuperstepManager | Python's functional approach mapped to Java's object-oriented design. |
| Algorithm Functions | langgraph/pregel/algo.py | Various Java classes | Python's functional approach distributed across several Java classes according to responsibility. |
| TaskPlanner | langgraph/pregel/algo.py | com.langgraph.pregel.task.TaskPlanner | Java implementation now matches Python: only nodes with the input channel as a trigger execute on first run. See CHANNEL_INITIALIZATION.md for details. |
### Checkpoint
| Component | Python Path | Java Path | Deviations |
|-----------|-------------|-----------|------------|
| BaseCheckpointSaver | langgraph/checkpoint/base.py | com.langgraph.checkpoint.base.BaseCheckpointSaver | Java uses interfaces rather than abstract classes where appropriate. |
| MemoryCheckpointSaver | langgraph/checkpoint/memory.py | com.langgraph.checkpoint.base.memory.MemoryCheckpointSaver | Java implementation uses more type safety but maintains same functionality. |
| Serializer | langgraph/checkpoint/serde.py | com.langgraph.checkpoint.serde.Serializer | Java uses interface with specific implementations for different serialization approaches. |
## Method-Level Mappings
### PregelLoop (Python: langgraph/pregel/loop.py, Java: com.langgraph.pregel.execute.PregelLoop)
| Python Method | Java Method | Deviations |
|---------------|-------------|------------|
| `__init__` | Constructor + Builder pattern | Java uses Builder pattern for more flexible initialization. |
| `tick` | `execute` | Same core functionality, but with improved recursion detection that matches Python behavior while being more resilient. Java executes a final validation step before throwing recursion errors to ensure runs complete when possible. |
| `_first` | `initializeWithInput` | Similar initialization logic but with Java-specific patterns. |
| `stream` | `stream` | Both handle streaming with similar semantics but with improved robustness in Java. Stream mode includes more validation to prevent false recursion errors. |
| `_put_checkpoint` | `createCheckpoint` | Similar checkpoint creation but with Java-specific implementation. |
### Runner Functions (Python: langgraph/pregel/runner.py)
| Python Function | Java Method | Deviations |
|-----------------|-------------|------------|
| `commit` | `SuperstepManager.commit` | Java implementation encapsulates in object instead of standalone function. |
| `tick` | `SuperstepManager.tick` | Same core functionality but adapted to Java's object-oriented paradigm. |
### Algorithm Functions (Python: langgraph/pregel/algo.py)
| Python Function | Java Method | Deviations |
|-----------------|-------------|------------|
| `prepare_next_tasks` | `TaskPlanner.planTasks` | Java implementation encapsulates in object instead of standalone function. |
| `prepare_single_task` | `TaskPlanner.planSingleTask` | Same approach but with stronger typing in Java. |
| `apply_writes` | Multiple methods in ChannelRegistry | Java distributes responsibility across specialized classes. |
## Implementation Notes
### General Patterns
- Java uses more explicit type information compared to Python
- Builder pattern is used in Java where Python uses parameter initialization
- Java collections (List, Map) replace Python collections (list, dict)
- Java follows standard exception hierarchy rather than Python's exception model
- Python's functional approach is often translated to Java's object-oriented design using objects with state
- Uninitialized channels in Java return null or empty collections rather than throwing exceptions
- Nodes in Java follow Python's behavior: only nodes with input channel as a trigger run in the first superstep
- Both implementations handle uninitialized channels gracefully without requiring manual initialization
### Missing Features (To Be Implemented)
- Some stream modes are not yet fully implemented in Java
- Advanced graph features are still under development in Java
- Some error handling cases need refinement to match Python semantics fully
## When Adding New Components
When adding new Java classes that correspond to Python implementations:
1. Add an entry to this document
2. Document any deviations and justify according to allowed reasons:
- Different public interfaces to match Java developer expectations
- Different implementation details to match Java stdlib/patterns
- Not yet fully implemented Python behavior
3. Never introduce deviations just to take shortcuts or change behavior
+266
View File
@@ -0,0 +1,266 @@
# LangGraph Java
A Java implementation of the [LangGraph](https://github.com/langchain-ai/langgraph) framework for building stateful, streaming LLM applications.
## Overview
LangGraph Java is designed for building directed, stateful computational graphs suitable for orchestrating LLM-based applications. The framework is particularly useful for:
- Building agents with tools, memory, and planning abilities
- Creating multi-agent systems with communication channels
- Implementing retrieval augmented generation (RAG) pipelines
- Supporting streaming output for responsive UI experiences
Key features:
- **Type-safe execution** with Java generics
- **Stateful graph execution** with checkpoint persistence
- **Streaming output** for real-time feedback
- **Directed computation graphs** with deterministic execution
## Project Structure
- `langgraph-checkpoint`: Base persistence interfaces
- `langgraph-core`: Main library with channels, Pregel implementation
- `langgraph-examples`: Example applications
## Requirements
- Java 17 or higher
- Gradle 7.0 or higher
## Building
```bash
./gradlew build
```
## Getting Started
### Basic Example
Here's a simple example that creates a graph with a single node that adds 1 to its input:
```java
import com.langgraph.channels.LastValue;
import com.langgraph.pregel.Pregel;
import com.langgraph.pregel.PregelExecutable;
import com.langgraph.pregel.PregelNode;
import java.util.HashMap;
import java.util.Map;
public class SimpleExample {
public static void main(String[] args) {
// Create a node that adds 1 to the input
PregelNode<Integer, Integer> node = new PregelNode.Builder<>("adder",
new PregelExecutable<Integer, Integer>() {
@Override
public Map<String, Integer> execute(Map<String, Integer> inputs, Map<String, Object> context) {
// Get input value, default to 0 if not present
int inputValue = inputs.getOrDefault("input", 0);
// Return output with value increased by 1
Map<String, Integer> output = new HashMap<>();
output.put("output", inputValue + 1);
return output;
}
})
.channels("input") // Read from "input" channel
.triggerChannels("input") // Triggered by "input" updates
.writers("output") // Write to "output" channel
.build();
// Create channels
Map<String, BaseChannel<?, ?, ?>> channels = new HashMap<>();
channels.put("input", LastValue.<Integer>create("input"));
channels.put("output", LastValue.<Integer>create("output"));
// Create Pregel instance
Pregel<Integer, Integer> pregel = new Pregel.Builder<Integer, Integer>()
.addNode(node)
.addChannels(channels)
.build();
// Run with input 5
Map<String, Integer> input = new HashMap<>();
input.put("input", 5);
Map<String, Integer> result = pregel.invoke(input, null);
// Print result (should be 6)
System.out.println("Result: " + result.get("output"));
}
}
```
### Multi-Step Graph Example
Here's an example of a two-node graph that performs sequential processing:
```java
import com.langgraph.channels.BaseChannel;
import com.langgraph.channels.LastValue;
import com.langgraph.pregel.Pregel;
import com.langgraph.pregel.PregelExecutable;
import com.langgraph.pregel.PregelNode;
import java.util.*;
public class SequentialExample {
public static void main(String[] args) {
// First node: Add 1 to the input and write to intermediate channel
PregelNode<Integer, Integer> adder = new PregelNode.Builder<>("adder",
new PregelExecutable<Integer, Integer>() {
@Override
public Map<String, Integer> execute(Map<String, Integer> inputs, Map<String, Object> context) {
int inputValue = inputs.getOrDefault("input", 0);
System.out.println("Adder received input: " + inputValue);
// Add 1 to the input value
int result = inputValue + 1;
// Write to the intermediate channel "state"
Map<String, Integer> output = new HashMap<>();
output.put("state", result);
return output;
}
})
.channels("input")
.triggerChannels("input")
.writers("state")
.build();
// Second node: Multiply intermediate value by 2 and write to output
PregelNode<Integer, Integer> multiplier = new PregelNode.Builder<>("multiplier",
new PregelExecutable<Integer, Integer>() {
@Override
public Map<String, Integer> execute(Map<String, Integer> inputs, Map<String, Object> context) {
// Get state value, default to 1 if not present
int stateValue = inputs.getOrDefault("state", 1);
// Multiply by 2
int result = stateValue * 2;
// Write to the output channel
Map<String, Integer> output = new HashMap<>();
output.put("output", result);
return output;
}
})
.channels("state")
.triggerChannels("state")
.writers("output")
.build();
// Create and configure channels
Map<String, BaseChannel<?, ?, ?>> channels = new HashMap<>();
channels.put("input", LastValue.<Integer>create("input"));
channels.put("state", LastValue.<Integer>create("state"));
channels.put("output", LastValue.<Integer>create("output"));
// Create Pregel instance with both nodes
Pregel<Integer, Integer> pregel = new Pregel.Builder<Integer, Integer>()
.addNode(adder)
.addNode(multiplier)
.addChannels(channels)
.build();
// Run with input 5
Map<String, Integer> input = Collections.singletonMap("input", 5);
Map<String, Integer> result = pregel.invoke(input, null);
// Print result: (5 + 1) * 2 = 12
System.out.println("Result: " + result.get("output"));
}
}
```
## Advanced Usage
### Working with String Data
```java
// Create a node that processes string data
PregelNode<String, String> processor = new PregelNode.Builder<>("processor",
new PregelExecutable<String, String>() {
@Override
public Map<String, String> execute(Map<String, String> inputs, Map<String, Object> context) {
String input = inputs.getOrDefault("input", "");
Map<String, String> output = new HashMap<>();
output.put("output", input.toUpperCase());
return output;
}
})
.channels("input")
.triggerChannels("input")
.writers("output")
.build();
// Create channels
Map<String, BaseChannel<?, ?, ?>> channels = new HashMap<>();
channels.put("input", LastValue.<String>create("input"));
channels.put("output", LastValue.<String>create("output"));
// Create Pregel instance
Pregel<String, String> pregel = new Pregel.Builder<String, String>()
.addNode(processor)
.addChannels(channels)
.build();
```
### Working with JSON-like Data
```java
// Create a node that processes Map<String, Object> data (JSON-like)
PregelNode<Map<String, Object>, Map<String, Object>> processor =
new PregelNode.Builder<>("processor",
new PregelExecutable<Map<String, Object>, Map<String, Object>>() {
@Override
public Map<String, Map<String, Object>> execute(
Map<String, Map<String, Object>> inputs,
Map<String, Object> context) {
Map<String, Object> input = inputs.getOrDefault("input", Collections.emptyMap());
// Process input
Map<String, Object> result = new HashMap<>(input);
result.put("processed", true);
Map<String, Map<String, Object>> output = new HashMap<>();
output.put("output", result);
return output;
}
})
.channels("input")
.triggerChannels("input")
.writers("output")
.build();
// Create channels
Map<String, BaseChannel<?, ?, ?>> channels = new HashMap<>();
channels.put("input", LastValue.<Map<String, Object>>create("input"));
channels.put("output", LastValue.<Map<String, Object>>create("output"));
// Create Pregel instance
Pregel<Map<String, Object>, Map<String, Object>> pregel =
new Pregel.Builder<Map<String, Object>, Map<String, Object>>()
.addNode(processor)
.addChannels(channels)
.build();
```
## Channel Types
LangGraph Java provides different channel types for different use cases:
- **LastValue**: Stores the last value written to the channel
- **TopicChannel**: Collects multiple values into a list
- **EphemeralValue**: Only available for the current execution step
## Contributing
Contributions are welcome! Please feel free to submit a Pull Request.
## License
This project is licensed under the MIT License - see the LICENSE file for details.
+54
View File
@@ -0,0 +1,54 @@
# Type-Safe LangGraph Java Implementation Summary
## Changes Made
1. **PregelExecutable<I, O> Interface**
- Added generic type parameters for input and output
- Provides strict typing for node actions
- Added Legacy adapter for backward compatibility
2. **PregelNode<I, O> Class**
- Made generic to enforce type safety
- Added input and output type tracking
- Enhanced with type validation during execution
- Legacy factory methods for compatibility
3. **PregelProtocol<I, O> Interface**
- Added type parameters for input and output
- Typed API for graph I/O
- Legacy subinterface for backward compatibility
4. **Pregel<I, O> Class**
- Type-safe implementation
- Type validation for channels and nodes
- Enhanced builder pattern with types
- Legacy factory methods
## Type Safety Benefits
1. **Compile-time Type Checking**
- Input/output types checked at compile time
- Prevents type errors at runtime
- Clearer API for developers
2. **Enhanced Runtime Validation**
- Validates type compatibility at graph construction
- Checks node/channel compatibility
- Provides clear error messages for mismatches
3. **Reduced Need for Type Casting**
- Explicit type parameters eliminate need for casts
- Prevents ClassCastExceptions
- Better developer experience
4. **Documentation & API Clarity**
- Type parameters document expected types
- Self-documenting builder pattern
- Clearer type relationships
5. **Backward Compatibility**
- Legacy methods for existing code
- Gradual migration possible
- No breaking changes to existing APIs
+39
View File
@@ -0,0 +1,39 @@
plugins {
id 'java-library'
}
allprojects {
group = 'com.langgraph'
version = '0.1.0-SNAPSHOT'
repositories {
mavenCentral()
}
}
subprojects {
apply plugin: 'java-library'
java {
sourceCompatibility = JavaVersion.VERSION_17
targetCompatibility = JavaVersion.VERSION_17
}
tasks.withType(JavaCompile) {
options.encoding = 'UTF-8'
options.compilerArgs << '-parameters'
}
dependencies {
// Testing dependencies
testImplementation 'org.junit.jupiter:junit-jupiter-api:5.9.2'
testImplementation 'org.junit.jupiter:junit-jupiter-params:5.9.2'
testRuntimeOnly 'org.junit.jupiter:junit-jupiter-engine:5.9.2'
testImplementation 'org.mockito:mockito-core:5.2.0'
testImplementation 'org.assertj:assertj-core:3.24.2'
}
test {
useJUnitPlatform()
}
}
Binary file not shown.
@@ -0,0 +1,7 @@
distributionBase=GRADLE_USER_HOME
distributionPath=wrapper/dists
distributionUrl=https\://services.gradle.org/distributions/gradle-8.13-bin.zip
networkTimeout=10000
validateDistributionUrl=true
zipStoreBase=GRADLE_USER_HOME
zipStorePath=wrapper/dists
Vendored Executable
+251
View File
@@ -0,0 +1,251 @@
#!/bin/sh
#
# Copyright © 2015-2021 the original authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# SPDX-License-Identifier: Apache-2.0
#
##############################################################################
#
# Gradle start up script for POSIX generated by Gradle.
#
# Important for running:
#
# (1) You need a POSIX-compliant shell to run this script. If your /bin/sh is
# noncompliant, but you have some other compliant shell such as ksh or
# bash, then to run this script, type that shell name before the whole
# command line, like:
#
# ksh Gradle
#
# Busybox and similar reduced shells will NOT work, because this script
# requires all of these POSIX shell features:
# * functions;
# * expansions «$var», «${var}», «${var:-default}», «${var+SET}»,
# «${var#prefix}», «${var%suffix}», and «$( cmd )»;
# * compound commands having a testable exit status, especially «case»;
# * various built-in commands including «command», «set», and «ulimit».
#
# Important for patching:
#
# (2) This script targets any POSIX shell, so it avoids extensions provided
# by Bash, Ksh, etc; in particular arrays are avoided.
#
# The "traditional" practice of packing multiple parameters into a
# space-separated string is a well documented source of bugs and security
# problems, so this is (mostly) avoided, by progressively accumulating
# options in "$@", and eventually passing that to Java.
#
# Where the inherited environment variables (DEFAULT_JVM_OPTS, JAVA_OPTS,
# and GRADLE_OPTS) rely on word-splitting, this is performed explicitly;
# see the in-line comments for details.
#
# There are tweaks for specific operating systems such as AIX, CygWin,
# Darwin, MinGW, and NonStop.
#
# (3) This script is generated from the Groovy template
# https://github.com/gradle/gradle/blob/HEAD/platforms/jvm/plugins-application/src/main/resources/org/gradle/api/internal/plugins/unixStartScript.txt
# within the Gradle project.
#
# You can find Gradle at https://github.com/gradle/gradle/.
#
##############################################################################
# Attempt to set APP_HOME
# Resolve links: $0 may be a link
app_path=$0
# Need this for daisy-chained symlinks.
while
APP_HOME=${app_path%"${app_path##*/}"} # leaves a trailing /; empty if no leading path
[ -h "$app_path" ]
do
ls=$( ls -ld "$app_path" )
link=${ls#*' -> '}
case $link in #(
/*) app_path=$link ;; #(
*) app_path=$APP_HOME$link ;;
esac
done
# This is normally unused
# shellcheck disable=SC2034
APP_BASE_NAME=${0##*/}
# Discard cd standard output in case $CDPATH is set (https://github.com/gradle/gradle/issues/25036)
APP_HOME=$( cd -P "${APP_HOME:-./}" > /dev/null && printf '%s\n' "$PWD" ) || exit
# Use the maximum available, or set MAX_FD != -1 to use that value.
MAX_FD=maximum
warn () {
echo "$*"
} >&2
die () {
echo
echo "$*"
echo
exit 1
} >&2
# OS specific support (must be 'true' or 'false').
cygwin=false
msys=false
darwin=false
nonstop=false
case "$( uname )" in #(
CYGWIN* ) cygwin=true ;; #(
Darwin* ) darwin=true ;; #(
MSYS* | MINGW* ) msys=true ;; #(
NONSTOP* ) nonstop=true ;;
esac
CLASSPATH=$APP_HOME/gradle/wrapper/gradle-wrapper.jar
# Determine the Java command to use to start the JVM.
if [ -n "$JAVA_HOME" ] ; then
if [ -x "$JAVA_HOME/jre/sh/java" ] ; then
# IBM's JDK on AIX uses strange locations for the executables
JAVACMD=$JAVA_HOME/jre/sh/java
else
JAVACMD=$JAVA_HOME/bin/java
fi
if [ ! -x "$JAVACMD" ] ; then
die "ERROR: JAVA_HOME is set to an invalid directory: $JAVA_HOME
Please set the JAVA_HOME variable in your environment to match the
location of your Java installation."
fi
else
JAVACMD=java
if ! command -v java >/dev/null 2>&1
then
die "ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH.
Please set the JAVA_HOME variable in your environment to match the
location of your Java installation."
fi
fi
# Increase the maximum file descriptors if we can.
if ! "$cygwin" && ! "$darwin" && ! "$nonstop" ; then
case $MAX_FD in #(
max*)
# In POSIX sh, ulimit -H is undefined. That's why the result is checked to see if it worked.
# shellcheck disable=SC2039,SC3045
MAX_FD=$( ulimit -H -n ) ||
warn "Could not query maximum file descriptor limit"
esac
case $MAX_FD in #(
'' | soft) :;; #(
*)
# In POSIX sh, ulimit -n is undefined. That's why the result is checked to see if it worked.
# shellcheck disable=SC2039,SC3045
ulimit -n "$MAX_FD" ||
warn "Could not set maximum file descriptor limit to $MAX_FD"
esac
fi
# Collect all arguments for the java command, stacking in reverse order:
# * args from the command line
# * the main class name
# * -classpath
# * -D...appname settings
# * --module-path (only if needed)
# * DEFAULT_JVM_OPTS, JAVA_OPTS, and GRADLE_OPTS environment variables.
# For Cygwin or MSYS, switch paths to Windows format before running java
if "$cygwin" || "$msys" ; then
APP_HOME=$( cygpath --path --mixed "$APP_HOME" )
CLASSPATH=$( cygpath --path --mixed "$CLASSPATH" )
JAVACMD=$( cygpath --unix "$JAVACMD" )
# Now convert the arguments - kludge to limit ourselves to /bin/sh
for arg do
if
case $arg in #(
-*) false ;; # don't mess with options #(
/?*) t=${arg#/} t=/${t%%/*} # looks like a POSIX filepath
[ -e "$t" ] ;; #(
*) false ;;
esac
then
arg=$( cygpath --path --ignore --mixed "$arg" )
fi
# Roll the args list around exactly as many times as the number of
# args, so each arg winds up back in the position where it started, but
# possibly modified.
#
# NB: a `for` loop captures its iteration list before it begins, so
# changing the positional parameters here affects neither the number of
# iterations, nor the values presented in `arg`.
shift # remove old arg
set -- "$@" "$arg" # push replacement arg
done
fi
# Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script.
DEFAULT_JVM_OPTS='"-Xmx64m" "-Xms64m"'
# Collect all arguments for the java command:
# * DEFAULT_JVM_OPTS, JAVA_OPTS, and optsEnvironmentVar are not allowed to contain shell fragments,
# and any embedded shellness will be escaped.
# * For example: A user cannot expect ${Hostname} to be expanded, as it is an environment variable and will be
# treated as '${Hostname}' itself on the command line.
set -- \
"-Dorg.gradle.appname=$APP_BASE_NAME" \
-classpath "$CLASSPATH" \
org.gradle.wrapper.GradleWrapperMain \
"$@"
# Stop when "xargs" is not available.
if ! command -v xargs >/dev/null 2>&1
then
die "xargs is not available"
fi
# Use "xargs" to parse quoted args.
#
# With -n1 it outputs one arg per line, with the quotes and backslashes removed.
#
# In Bash we could simply go:
#
# readarray ARGS < <( xargs -n1 <<<"$var" ) &&
# set -- "${ARGS[@]}" "$@"
#
# but POSIX shell has neither arrays nor command substitution, so instead we
# post-process each arg (as a line of input to sed) to backslash-escape any
# character that might be a shell metacharacter, then use eval to reverse
# that process (while maintaining the separation between arguments), and wrap
# the whole thing up as a single "set" statement.
#
# This will of course break if any of these variables contains a newline or
# an unmatched quote.
#
eval "set -- $(
printf '%s\n' "$DEFAULT_JVM_OPTS $JAVA_OPTS $GRADLE_OPTS" |
xargs -n1 |
sed ' s~[^-[:alnum:]+,./:=@_]~\\&~g; ' |
tr '\n' ' '
)" '"$@"'
exec "$JAVACMD" "$@"
+94
View File
@@ -0,0 +1,94 @@
@rem
@rem Copyright 2015 the original author or authors.
@rem
@rem Licensed under the Apache License, Version 2.0 (the "License");
@rem you may not use this file except in compliance with the License.
@rem You may obtain a copy of the License at
@rem
@rem https://www.apache.org/licenses/LICENSE-2.0
@rem
@rem Unless required by applicable law or agreed to in writing, software
@rem distributed under the License is distributed on an "AS IS" BASIS,
@rem WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
@rem See the License for the specific language governing permissions and
@rem limitations under the License.
@rem
@rem SPDX-License-Identifier: Apache-2.0
@rem
@if "%DEBUG%"=="" @echo off
@rem ##########################################################################
@rem
@rem Gradle startup script for Windows
@rem
@rem ##########################################################################
@rem Set local scope for the variables with windows NT shell
if "%OS%"=="Windows_NT" setlocal
set DIRNAME=%~dp0
if "%DIRNAME%"=="" set DIRNAME=.
@rem This is normally unused
set APP_BASE_NAME=%~n0
set APP_HOME=%DIRNAME%
@rem Resolve any "." and ".." in APP_HOME to make it shorter.
for %%i in ("%APP_HOME%") do set APP_HOME=%%~fi
@rem Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script.
set DEFAULT_JVM_OPTS="-Xmx64m" "-Xms64m"
@rem Find java.exe
if defined JAVA_HOME goto findJavaFromJavaHome
set JAVA_EXE=java.exe
%JAVA_EXE% -version >NUL 2>&1
if %ERRORLEVEL% equ 0 goto execute
echo. 1>&2
echo ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. 1>&2
echo. 1>&2
echo Please set the JAVA_HOME variable in your environment to match the 1>&2
echo location of your Java installation. 1>&2
goto fail
:findJavaFromJavaHome
set JAVA_HOME=%JAVA_HOME:"=%
set JAVA_EXE=%JAVA_HOME%/bin/java.exe
if exist "%JAVA_EXE%" goto execute
echo. 1>&2
echo ERROR: JAVA_HOME is set to an invalid directory: %JAVA_HOME% 1>&2
echo. 1>&2
echo Please set the JAVA_HOME variable in your environment to match the 1>&2
echo location of your Java installation. 1>&2
goto fail
:execute
@rem Setup the command line
set CLASSPATH=%APP_HOME%\gradle\wrapper\gradle-wrapper.jar
@rem Execute Gradle
"%JAVA_EXE%" %DEFAULT_JVM_OPTS% %JAVA_OPTS% %GRADLE_OPTS% "-Dorg.gradle.appname=%APP_BASE_NAME%" -classpath "%CLASSPATH%" org.gradle.wrapper.GradleWrapperMain %*
:end
@rem End local scope for the variables with windows NT shell
if %ERRORLEVEL% equ 0 goto mainEnd
:fail
rem Set variable GRADLE_EXIT_CONSOLE if you need the _script_ return code instead of
rem the _cmd.exe /c_ return code!
set EXIT_CODE=%ERRORLEVEL%
if %EXIT_CODE% equ 0 set EXIT_CODE=1
if not ""=="%GRADLE_EXIT_CONSOLE%" exit %EXIT_CODE%
exit /b %EXIT_CODE%
:mainEnd
if "%OS%"=="Windows_NT" endlocal
:omega
@@ -0,0 +1,5 @@
dependencies {
// MessagePack for serialization
implementation 'org.msgpack:msgpack-core:0.9.3'
implementation 'org.msgpack:jackson-dataformat-msgpack:0.9.3'
}
@@ -0,0 +1,60 @@
package com.langgraph.checkpoint.base;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.concurrent.CompletableFuture;
/**
* Asynchronous interface for saving and loading checkpoints.
*/
public interface AsyncBaseCheckpointSaver {
/**
* Create a new checkpoint asynchronously.
*
* @param threadId The ID of the thread to checkpoint
* @param channelValues The values of the channels to checkpoint
* @return CompletableFuture with the ID of the new checkpoint
*/
CompletableFuture<String> checkpointAsync(String threadId, Map<String, Object> channelValues);
/**
* Get values from a checkpoint asynchronously.
*
* @param checkpointId The ID of the checkpoint to load
* @return CompletableFuture with the channel values from the checkpoint, or empty if not found
*/
CompletableFuture<Optional<Map<String, Object>>> getValuesAsync(String checkpointId);
/**
* List all checkpoints for a thread asynchronously.
*
* @param threadId The ID of the thread
* @return CompletableFuture with list of checkpoint IDs
*/
CompletableFuture<List<String>> listAsync(String threadId);
/**
* Get the latest checkpoint for a thread asynchronously.
*
* @param threadId The ID of the thread
* @return CompletableFuture with the ID of the latest checkpoint, or empty if none exists
*/
CompletableFuture<Optional<String>> latestAsync(String threadId);
/**
* Delete a checkpoint asynchronously.
*
* @param checkpointId The ID of the checkpoint to delete
* @return CompletableFuture completed when deletion is done
*/
CompletableFuture<Void> deleteAsync(String checkpointId);
/**
* Clear all checkpoints for a thread asynchronously.
*
* @param threadId The ID of the thread
* @return CompletableFuture completed when clearing is done
*/
CompletableFuture<Void> clearAsync(String threadId);
}
@@ -0,0 +1,57 @@
package com.langgraph.checkpoint.base;
import java.util.List;
import java.util.Map;
import java.util.Optional;
/**
* Interface for saving and loading checkpoints.
*/
public interface BaseCheckpointSaver {
/**
* Create a new checkpoint.
*
* @param threadId The ID of the thread to checkpoint
* @param channelValues The values of the channels to checkpoint
* @return The ID of the new checkpoint
*/
String checkpoint(String threadId, Map<String, Object> channelValues);
/**
* Get values from a checkpoint.
*
* @param checkpointId The ID of the checkpoint to load
* @return The channel values from the checkpoint, or empty if not found
*/
Optional<Map<String, Object>> getValues(String checkpointId);
/**
* List all checkpoints for a thread.
*
* @param threadId The ID of the thread
* @return List of checkpoint IDs
*/
List<String> list(String threadId);
/**
* Get the latest checkpoint for a thread.
*
* @param threadId The ID of the thread
* @return The ID of the latest checkpoint, or empty if none exists
*/
Optional<String> latest(String threadId);
/**
* Delete a checkpoint.
*
* @param checkpointId The ID of the checkpoint to delete
*/
void delete(String checkpointId);
/**
* Clear all checkpoints for a thread.
*
* @param threadId The ID of the thread
*/
void clear(String threadId);
}
@@ -0,0 +1,81 @@
package com.langgraph.checkpoint.base;
import java.nio.charset.StandardCharsets;
import java.security.MessageDigest;
import java.security.NoSuchAlgorithmException;
import java.util.Base64;
import java.util.UUID;
/**
* Utility class for generating deterministic IDs.
*/
public final class ID {
private ID() {
// Prevent instantiation
}
/**
* Generate a deterministic UUID based on a namespace and name.
*
* @param namespace The namespace for the ID
* @param name The name within the namespace
* @return A UUID
*/
public static UUID uuid(String namespace, String name) {
try {
MessageDigest md = MessageDigest.getInstance("SHA-1");
md.update(namespace.getBytes(StandardCharsets.UTF_8));
md.update(name.getBytes(StandardCharsets.UTF_8));
byte[] digest = md.digest();
// Set the version (4) and variant (RFC4122) bits
digest[6] = (byte) ((digest[6] & 0x0F) | 0x40);
digest[8] = (byte) ((digest[8] & 0x3F) | 0x80);
long msb = 0;
long lsb = 0;
for (int i = 0; i < 8; i++) {
msb = (msb << 8) | (digest[i] & 0xff);
}
for (int i = 8; i < 16; i++) {
lsb = (lsb << 8) | (digest[i] & 0xff);
}
return new UUID(msb, lsb);
} catch (NoSuchAlgorithmException e) {
throw new RuntimeException("SHA-1 algorithm not available", e);
}
}
/**
* Generate a checkpoint ID.
*
* @param threadId The thread ID
* @return A checkpoint ID
*/
public static String checkpointId(String threadId) {
return uuid("checkpoint", threadId + "/" + System.currentTimeMillis()).toString();
}
/**
* Generate a URL-safe base64 encoded ID.
*
* @param namespace The namespace for the ID
* @param name The name within the namespace
* @return A URL-safe base64-encoded ID
*/
public static String urlSafeId(String namespace, String name) {
try {
MessageDigest md = MessageDigest.getInstance("SHA-256");
md.update(namespace.getBytes(StandardCharsets.UTF_8));
md.update(name.getBytes(StandardCharsets.UTF_8));
byte[] digest = md.digest();
return Base64.getUrlEncoder().withoutPadding().encodeToString(digest);
} catch (NoSuchAlgorithmException e) {
throw new RuntimeException("SHA-256 algorithm not available", e);
}
}
}
@@ -0,0 +1,79 @@
package com.langgraph.checkpoint.base.memory;
import com.langgraph.checkpoint.base.AsyncBaseCheckpointSaver;
import com.langgraph.checkpoint.base.BaseCheckpointSaver;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.concurrent.CompletableFuture;
/**
* Asynchronous in-memory implementation of a checkpoint saver.
* This is a thin wrapper around the synchronous implementation that
* executes operations asynchronously.
*/
public class AsyncMemoryCheckpointSaver implements AsyncBaseCheckpointSaver {
private final BaseCheckpointSaver synchronousSaver;
/**
* Create an async memory checkpoint saver.
*/
public AsyncMemoryCheckpointSaver() {
this.synchronousSaver = new MemoryCheckpointSaver();
}
/**
* Create an async memory checkpoint saver with an existing synchronous saver.
*
* @param synchronousSaver The synchronous checkpoint saver to wrap
*/
public AsyncMemoryCheckpointSaver(BaseCheckpointSaver synchronousSaver) {
this.synchronousSaver = synchronousSaver;
}
@Override
public CompletableFuture<String> checkpointAsync(String threadId, Map<String, Object> channelValues) {
return CompletableFuture.supplyAsync(() ->
synchronousSaver.checkpoint(threadId, channelValues));
}
@Override
public CompletableFuture<Optional<Map<String, Object>>> getValuesAsync(String checkpointId) {
return CompletableFuture.supplyAsync(() ->
synchronousSaver.getValues(checkpointId));
}
@Override
public CompletableFuture<List<String>> listAsync(String threadId) {
return CompletableFuture.supplyAsync(() ->
synchronousSaver.list(threadId));
}
@Override
public CompletableFuture<Optional<String>> latestAsync(String threadId) {
return CompletableFuture.supplyAsync(() ->
synchronousSaver.latest(threadId));
}
@Override
public CompletableFuture<Void> deleteAsync(String checkpointId) {
return CompletableFuture.runAsync(() ->
synchronousSaver.delete(checkpointId));
}
@Override
public CompletableFuture<Void> clearAsync(String threadId) {
return CompletableFuture.runAsync(() ->
synchronousSaver.clear(threadId));
}
/**
* Get the underlying synchronous checkpoint saver.
*
* @return The synchronous checkpoint saver
*/
public BaseCheckpointSaver getSynchronousSaver() {
return synchronousSaver;
}
}
@@ -0,0 +1,77 @@
package com.langgraph.checkpoint.base.memory;
import com.langgraph.checkpoint.base.BaseCheckpointSaver;
import com.langgraph.checkpoint.base.ID;
import java.util.*;
import java.util.concurrent.ConcurrentHashMap;
/**
* In-memory implementation of a checkpoint saver.
*/
public class MemoryCheckpointSaver implements BaseCheckpointSaver {
private final Map<String, Map<String, Object>> checkpoints = new ConcurrentHashMap<>();
private final Map<String, List<String>> threadCheckpoints = new ConcurrentHashMap<>();
@Override
public String checkpoint(String threadId, Map<String, Object> channelValues) {
String checkpointId = ID.checkpointId(threadId);
// Store the checkpoint
checkpoints.put(checkpointId, new HashMap<>(channelValues));
// Add to thread's checkpoints
threadCheckpoints.computeIfAbsent(threadId, k ->
Collections.synchronizedList(new ArrayList<>())).add(checkpointId);
return checkpointId;
}
@Override
public Optional<Map<String, Object>> getValues(String checkpointId) {
Map<String, Object> values = checkpoints.get(checkpointId);
return Optional.ofNullable(values).map(HashMap::new);
}
@Override
public List<String> list(String threadId) {
List<String> result = threadCheckpoints.get(threadId);
return result != null ? new ArrayList<>(result) : Collections.emptyList();
}
@Override
public Optional<String> latest(String threadId) {
List<String> checkpoints = threadCheckpoints.get(threadId);
if (checkpoints == null || checkpoints.isEmpty()) {
return Optional.empty();
}
return Optional.of(checkpoints.get(checkpoints.size() - 1));
}
@Override
public void delete(String checkpointId) {
// Remove the checkpoint
Map<String, Object> removed = checkpoints.remove(checkpointId);
if (removed != null) {
// Find and remove from thread's checkpoints
for (List<String> checkpointsList : threadCheckpoints.values()) {
checkpointsList.remove(checkpointId);
}
}
}
@Override
public void clear(String threadId) {
List<String> checkpointIds = threadCheckpoints.remove(threadId);
if (checkpointIds != null) {
// Remove all checkpoints for this thread
for (String checkpointId : checkpointIds) {
checkpoints.remove(checkpointId);
}
}
}
}
@@ -0,0 +1,556 @@
package com.langgraph.checkpoint.serde;
import org.msgpack.core.MessageBufferPacker;
import org.msgpack.core.MessagePack;
import org.msgpack.core.MessageUnpacker;
import org.msgpack.core.MessageFormat;
import java.io.IOException;
import java.lang.reflect.*;
import java.time.Instant;
import java.time.LocalDate;
import java.time.LocalDateTime;
import java.time.LocalTime;
import java.util.*;
import java.util.concurrent.ConcurrentHashMap;
/**
* MessagePack-based serializer that uses reflection to handle arbitrary Java objects.
* Supports primitive types, collections, maps, records, and custom objects.
*/
public class MsgPackSerializer implements ReflectionSerializer {
private final Map<Class<?>, TypeSerializer<?>> serializers = new ConcurrentHashMap<>();
private final Map<Class<?>, TypeDeserializer<?>> deserializers = new ConcurrentHashMap<>();
private final Map<Class<?>, RecordInfo> recordInfoCache = new ConcurrentHashMap<>();
/**
* Record component information cache to avoid repeated reflection.
*/
private static class RecordInfo {
final RecordComponent[] components;
final Constructor<?> constructor;
RecordInfo(RecordComponent[] components, Constructor<?> constructor) {
this.components = components;
this.constructor = constructor;
}
}
/**
* Register built-in serializers for common types.
*/
public MsgPackSerializer() {
registerBuiltinTypes();
}
/**
* Register built-in serializers for common types.
*/
private void registerBuiltinTypes() {
// UUID serializer
registerSerializer(UUID.class, (uuid) -> uuid.toString());
registerDeserializer(UUID.class, (str) -> UUID.fromString((String) str));
// Date serializer
registerSerializer(java.util.Date.class, (date) -> date.getTime());
registerDeserializer(java.util.Date.class, (millis) -> new Date((Long) millis));
// Java 8 Date/Time API
registerSerializer(Instant.class, (instant) -> instant.toString());
registerDeserializer(Instant.class, (str) -> Instant.parse((String) str));
registerSerializer(LocalDate.class, (date) -> date.toString());
registerDeserializer(LocalDate.class, (str) -> LocalDate.parse((String) str));
registerSerializer(LocalTime.class, (time) -> time.toString());
registerDeserializer(LocalTime.class, (str) -> LocalTime.parse((String) str));
registerSerializer(LocalDateTime.class, (dateTime) -> dateTime.toString());
registerDeserializer(LocalDateTime.class, (str) -> LocalDateTime.parse((String) str));
// Add more built-in serializers as needed
}
@Override
public <T> void registerSerializer(Class<T> type, TypeSerializer<T> serializer) {
serializers.put(type, serializer);
}
@Override
public <T> void registerDeserializer(Class<T> type, TypeDeserializer<T> deserializer) {
deserializers.put(type, deserializer);
}
@Override
public byte[] serialize(Object obj) {
try {
MessageBufferPacker packer = MessagePack.newDefaultBufferPacker();
serializeObject(obj, packer);
return packer.toByteArray();
} catch (IOException e) {
throw new SerializationException("Failed to serialize object", e);
}
}
@Override
public Object deserialize(byte[] data) {
try {
MessageUnpacker unpacker = MessagePack.newDefaultUnpacker(data);
return deserializeObject(unpacker);
} catch (IOException e) {
throw new SerializationException("Failed to deserialize object", e);
}
}
/**
* Serialize an object to the MessagePack packer.
*
* @param obj Object to serialize
* @param packer MessagePack packer
* @throws IOException If packing fails
*/
private void serializeObject(Object obj, MessageBufferPacker packer) throws IOException {
if (obj == null) {
packer.packNil();
return;
}
Class<?> type = obj.getClass();
// Check for registered serializer
if (serializers.containsKey(type)) {
// This cast is safe because we only put serializers for a specific type in the map
TypeSerializer<?> untypedSerializer = serializers.get(type);
// We need this cast but it's type-safe because we only store TypeSerializer<T> for Class<T>
@SuppressWarnings("unchecked")
TypeSerializer<Object> serializer = (TypeSerializer<Object>) untypedSerializer;
Object serialized = serializer.toSerializable(obj);
// Pack as a special type
packer.packMapHeader(2);
packer.packString("__type__");
packer.packString(type.getName());
packer.packString("value");
serializeObject(serialized, packer);
return;
}
// Handle primitive types and common objects directly
if (obj instanceof String) {
packer.packString((String) obj);
} else if (obj instanceof Integer) {
packer.packInt((Integer) obj);
} else if (obj instanceof Long) {
packer.packLong((Long) obj);
} else if (obj instanceof Double) {
packer.packDouble((Double) obj);
} else if (obj instanceof Float) {
packer.packFloat((Float) obj);
} else if (obj instanceof Boolean) {
packer.packBoolean((Boolean) obj);
} else if (obj instanceof byte[]) {
packer.packBinaryHeader(((byte[]) obj).length);
packer.writePayload((byte[]) obj);
} else if (obj instanceof List) {
List<?> list = (List<?>) obj;
packer.packArrayHeader(list.size());
for (Object item : list) {
serializeObject(item, packer);
}
} else if (obj instanceof Map) {
Map<?, ?> map = (Map<?, ?>) obj;
packer.packMapHeader(map.size());
for (Map.Entry<?, ?> entry : map.entrySet()) {
serializeObject(entry.getKey(), packer);
serializeObject(entry.getValue(), packer);
}
} else if (obj instanceof Enum<?>) {
// Handle enums by name
packer.packMapHeader(2);
packer.packString("__type__");
packer.packString(type.getName());
packer.packString("value");
packer.packString(((Enum<?>) obj).name());
} else if (type.isRecord()) {
// Handle Record types
serializeRecord(obj, packer);
} else {
// Handle custom objects with reflection
serializeCustomObject(obj, packer);
}
}
/**
* Serialize a Record object.
*
* @param record The record to serialize
* @param packer The MessagePack packer
* @throws IOException If packing fails
*/
private void serializeRecord(Object record, MessageBufferPacker packer) throws IOException {
Class<?> recordClass = record.getClass();
// Pack as a special type with fields
packer.packMapHeader(2);
packer.packString("__type__");
packer.packString(recordClass.getName());
packer.packString("fields");
RecordComponent[] components = recordClass.getRecordComponents();
packer.packMapHeader(components.length);
for (RecordComponent component : components) {
packer.packString(component.getName());
try {
Method accessor = component.getAccessor();
Object value = accessor.invoke(record);
serializeObject(value, packer);
} catch (ReflectiveOperationException e) {
throw new SerializationException("Failed to access record component: " + component.getName(), e);
}
}
}
/**
* Serialize a custom object using reflection.
*
* @param obj The object to serialize
* @param packer The MessagePack packer
* @throws IOException If packing fails
*/
private void serializeCustomObject(Object obj, MessageBufferPacker packer) throws IOException {
Class<?> objClass = obj.getClass();
// Pack as a special type with fields
packer.packMapHeader(2);
packer.packString("__type__");
packer.packString(objClass.getName());
packer.packString("fields");
// Get all fields including inherited ones
List<Field> fields = getAllFields(objClass);
// Filter out transient fields
List<Field> serializableFields = fields.stream()
.filter(field -> !Modifier.isTransient(field.getModifiers()) &&
!Modifier.isStatic(field.getModifiers()))
.toList();
packer.packMapHeader(serializableFields.size());
for (Field field : serializableFields) {
packer.packString(field.getName());
try {
field.setAccessible(true);
Object value = field.get(obj);
serializeObject(value, packer);
} catch (IllegalAccessException e) {
throw new SerializationException("Failed to access field: " + field.getName(), e);
}
}
}
/**
* Get all fields for a class including inherited fields.
*
* @param clazz The class to get fields for
* @return List of all fields
*/
private List<Field> getAllFields(Class<?> clazz) {
List<Field> fields = new ArrayList<>();
Class<?> currentClass = clazz;
while (currentClass != null && currentClass != Object.class) {
fields.addAll(Arrays.asList(currentClass.getDeclaredFields()));
currentClass = currentClass.getSuperclass();
}
return fields;
}
/**
* Deserialize an object from the MessagePack unpacker.
*
* @param unpacker MessagePack unpacker
* @return Deserialized object
* @throws IOException If unpacking fails
*/
private Object deserializeObject(MessageUnpacker unpacker) throws IOException {
if (!unpacker.hasNext()) {
throw new SerializationException("Unexpected end of data");
}
if (unpacker.tryUnpackNil()) {
return null;
}
MessageFormat format = unpacker.getNextFormat();
if (format == MessageFormat.STR8 ||
format == MessageFormat.STR16 ||
format == MessageFormat.STR32 ||
format == MessageFormat.FIXSTR) {
return unpacker.unpackString();
} else if (format == MessageFormat.INT8 ||
format == MessageFormat.INT16 ||
format == MessageFormat.INT32 ||
format == MessageFormat.INT64 ||
format == MessageFormat.UINT8 ||
format == MessageFormat.UINT16 ||
format == MessageFormat.UINT32 ||
format == MessageFormat.UINT64 ||
format == MessageFormat.POSFIXINT ||
format == MessageFormat.NEGFIXINT) {
if (format == MessageFormat.INT64 || format == MessageFormat.UINT64) {
return unpacker.unpackLong();
} else {
try {
return unpacker.unpackInt();
} catch (Exception e) {
// Fallback to long if int unpacking fails
return unpacker.unpackLong();
}
}
} else if (format == MessageFormat.FLOAT32 ||
format == MessageFormat.FLOAT64) {
return unpacker.unpackDouble();
} else if (format == MessageFormat.BOOLEAN) {
return unpacker.unpackBoolean();
} else if (format == MessageFormat.BIN8 ||
format == MessageFormat.BIN16 ||
format == MessageFormat.BIN32) {
int binaryLength = unpacker.unpackBinaryHeader();
byte[] binary = new byte[binaryLength];
unpacker.readPayload(binary);
return binary;
} else if (format == MessageFormat.ARRAY16 ||
format == MessageFormat.ARRAY32 ||
format == MessageFormat.FIXARRAY) {
int arraySize = unpacker.unpackArrayHeader();
List<Object> list = new ArrayList<>(arraySize);
for (int i = 0; i < arraySize; i++) {
list.add(deserializeObject(unpacker));
}
return list;
} else if (format == MessageFormat.MAP16 ||
format == MessageFormat.MAP32 ||
format == MessageFormat.FIXMAP) {
int mapSize = unpacker.unpackMapHeader();
// Handle empty map
if (mapSize == 0) {
return new HashMap<>();
}
// Check for special type marker
Object firstKey = deserializeObject(unpacker);
if (mapSize == 2 && firstKey instanceof String && "__type__".equals(firstKey)) {
String typeName = (String) deserializeObject(unpacker);
// Get the second key
Object secondKey = deserializeObject(unpacker);
if (secondKey instanceof String) {
String secondKeyStr = (String) secondKey;
try {
Class<?> type = Class.forName(typeName);
// Check for registered deserializer
if ("value".equals(secondKeyStr) && deserializers.containsKey(type)) {
Object serialized = deserializeObject(unpacker);
TypeDeserializer<?> untypedDeserializer = deserializers.get(type);
// We need this cast but it's type-safe because we only store TypeDeserializer<T> for Class<T>
@SuppressWarnings("unchecked")
TypeDeserializer<Object> deserializer = (TypeDeserializer<Object>) untypedDeserializer;
return deserializer.fromSerialized(serialized);
}
// Handle enums
if ("value".equals(secondKeyStr) && type.isEnum()) {
String enumValue = (String) deserializeObject(unpacker);
// This cast is required for enum handling and is type-safe
@SuppressWarnings("unchecked")
Class<Enum> enumClass = (Class<Enum>) type;
return Enum.valueOf(enumClass, enumValue);
}
// Handle records
if ("fields".equals(secondKeyStr) && type.isRecord()) {
return deserializeRecord(type, unpacker);
}
// Handle custom objects
if ("fields".equals(secondKeyStr)) {
return deserializeCustomObject(type, unpacker);
}
} catch (ClassNotFoundException e) {
// If class not found, fall back to regular map deserialization
} catch (ReflectiveOperationException e) {
throw new SerializationException("Failed to deserialize object of type " + typeName, e);
}
// If special type handling failed, read the value to keep unpacker consistent
Object secondValue = deserializeObject(unpacker);
// Create a fallback map with the special type info
Map<Object, Object> fallbackMap = new HashMap<>();
fallbackMap.put(firstKey, typeName);
fallbackMap.put(secondKey, secondValue);
return fallbackMap;
}
// If the second key wasn't a string as expected, we need to handle it as a regular map
Object firstValue = deserializeObject(unpacker);
// Create a map with the first key-value pair
Map<Object, Object> map = new HashMap<>(mapSize);
map.put(firstKey, firstValue);
// Read the remaining entries
for (int i = 1; i < mapSize; i++) {
Object key = deserializeObject(unpacker);
Object value = deserializeObject(unpacker);
map.put(key, value);
}
return map;
} else {
// Regular map - we already read the first key
Map<Object, Object> map = new HashMap<>(mapSize);
// Read the first value
Object firstValue = deserializeObject(unpacker);
map.put(firstKey, firstValue);
// Read the remaining entries
for (int i = 1; i < mapSize; i++) {
Object key = deserializeObject(unpacker);
Object value = deserializeObject(unpacker);
map.put(key, value);
}
return map;
}
}
// Default case
throw new SerializationException("Unsupported MessagePack format: " + format);
}
/**
* Deserialize a record.
*
* @param recordClass The record class
* @param unpacker The unpacker containing the fields map
* @return The deserialized record
* @throws IOException If unpacking fails
* @throws ReflectiveOperationException If reflection operations fail
*/
private Object deserializeRecord(Class<?> recordClass, MessageUnpacker unpacker)
throws IOException, ReflectiveOperationException {
// Get record info from cache or create it
RecordInfo recordInfo = recordInfoCache.computeIfAbsent(recordClass, cls -> {
try {
RecordComponent[] components = cls.getRecordComponents();
Class<?>[] paramTypes = Arrays.stream(components)
.map(RecordComponent::getType)
.toArray(Class<?>[]::new);
Constructor<?> constructor = cls.getDeclaredConstructor(paramTypes);
constructor.setAccessible(true);
return new RecordInfo(components, constructor);
} catch (NoSuchMethodException e) {
throw new SerializationException("Failed to get constructor for record: " + cls.getName(), e);
}
});
// Read the fields map
int fieldCount = unpacker.unpackMapHeader();
Map<String, Object> fieldValues = new HashMap<>(fieldCount);
for (int i = 0; i < fieldCount; i++) {
String fieldName = (String) deserializeObject(unpacker);
Object fieldValue = deserializeObject(unpacker);
fieldValues.put(fieldName, fieldValue);
}
// Prepare constructor arguments in the correct order
Object[] constructorArgs = new Object[recordInfo.components.length];
for (int i = 0; i < recordInfo.components.length; i++) {
RecordComponent component = recordInfo.components[i];
Object value = fieldValues.get(component.getName());
constructorArgs[i] = value;
}
// Create the record instance
return recordInfo.constructor.newInstance(constructorArgs);
}
/**
* Deserialize a custom object.
*
* @param objectClass The object class
* @param unpacker The unpacker containing the fields map
* @return The deserialized object
* @throws IOException If unpacking fails
* @throws ReflectiveOperationException If reflection operations fail
*/
private Object deserializeCustomObject(Class<?> objectClass, MessageUnpacker unpacker)
throws IOException, ReflectiveOperationException {
// Create instance using default constructor
Constructor<?> constructor;
try {
constructor = objectClass.getDeclaredConstructor();
constructor.setAccessible(true);
} catch (NoSuchMethodException e) {
throw new SerializationException(
"Class " + objectClass.getName() + " must have a no-arg constructor for deserialization", e);
}
Object instance = constructor.newInstance();
// Read the fields map
int fieldCount = unpacker.unpackMapHeader();
for (int i = 0; i < fieldCount; i++) {
String fieldName = (String) deserializeObject(unpacker);
Object fieldValue = deserializeObject(unpacker);
try {
// Find the field (including in superclasses)
Field field = findField(objectClass, fieldName);
if (field != null) {
field.setAccessible(true);
field.set(instance, fieldValue);
}
} catch (NoSuchFieldException e) {
// Skip fields that don't exist in the current class version
}
}
return instance;
}
/**
* Find a field in a class or its superclasses.
*
* @param clazz The class to search
* @param fieldName The field name to find
* @return The found field
* @throws NoSuchFieldException If the field is not found
*/
private Field findField(Class<?> clazz, String fieldName) throws NoSuchFieldException {
Class<?> currentClass = clazz;
while (currentClass != null) {
try {
return currentClass.getDeclaredField(fieldName);
} catch (NoSuchFieldException e) {
currentClass = currentClass.getSuperclass();
}
}
throw new NoSuchFieldException("Field not found: " + fieldName);
}
}
@@ -0,0 +1,24 @@
package com.langgraph.checkpoint.serde;
/**
* Interface for a serializer that uses reflection to handle arbitrary Java objects.
*/
public interface ReflectionSerializer extends Serializer<Object> {
/**
* Register a custom serializer for a specific type.
*
* @param type Type to register
* @param serializer Custom serializer for the type
* @param <T> Type to register
*/
<T> void registerSerializer(Class<T> type, TypeSerializer<T> serializer);
/**
* Register a custom deserializer for a specific type.
*
* @param type Type to register
* @param deserializer Custom deserializer for the type
* @param <T> Type to register
*/
<T> void registerDeserializer(Class<T> type, TypeDeserializer<T> deserializer);
}
@@ -0,0 +1,25 @@
package com.langgraph.checkpoint.serde;
/**
* Exception thrown during serialization/deserialization.
*/
public class SerializationException extends RuntimeException {
/**
* Create a new serialization exception with a message.
*
* @param message Error message
*/
public SerializationException(String message) {
super(message);
}
/**
* Create a new serialization exception with a message and cause.
*
* @param message Error message
* @param cause Underlying cause
*/
public SerializationException(String message, Throwable cause) {
super(message, cause);
}
}
@@ -0,0 +1,24 @@
package com.langgraph.checkpoint.serde;
/**
* Interface for serializing and deserializing objects.
*
* @param <T> Type of object to serialize/deserialize
*/
public interface Serializer<T> {
/**
* Serialize an object to bytes.
*
* @param obj The object to serialize
* @return Serialized bytes
*/
byte[] serialize(T obj);
/**
* Deserialize bytes to an object.
*
* @param data The bytes to deserialize
* @return Deserialized object
*/
T deserialize(byte[] data);
}
@@ -0,0 +1,17 @@
package com.langgraph.checkpoint.serde;
/**
* Interface for deserializing a specific type from MessagePack.
*
* @param <T> Type to deserialize
*/
@FunctionalInterface
public interface TypeDeserializer<T> {
/**
* Convert from serialized representation to object.
*
* @param serialized Serialized representation
* @return Deserialized object
*/
T fromSerialized(Object serialized);
}
@@ -0,0 +1,17 @@
package com.langgraph.checkpoint.serde;
/**
* Interface for serializing a specific type to a format that can be included in MessagePack.
*
* @param <T> Type to serialize
*/
@FunctionalInterface
public interface TypeSerializer<T> {
/**
* Convert object to a serializable representation.
*
* @param obj Object to convert
* @return Serializable representation (must be compatible with MessagePack)
*/
Object toSerializable(T obj);
}
@@ -0,0 +1,59 @@
package com.langgraph.checkpoint.base;
import org.junit.jupiter.api.Test;
import java.util.UUID;
import static org.assertj.core.api.Assertions.assertThat;
public class IDTest {
@Test
public void testUuidDeterministic() {
// Same inputs should produce same UUIDs
UUID uuid1 = ID.uuid("test", "value");
UUID uuid2 = ID.uuid("test", "value");
assertThat(uuid1).isEqualTo(uuid2);
}
@Test
public void testUuidDifferentNamespace() {
// Different namespaces should produce different UUIDs
UUID uuid1 = ID.uuid("namespace1", "value");
UUID uuid2 = ID.uuid("namespace2", "value");
assertThat(uuid1).isNotEqualTo(uuid2);
}
@Test
public void testUuidDifferentName() {
// Different names should produce different UUIDs
UUID uuid1 = ID.uuid("test", "value1");
UUID uuid2 = ID.uuid("test", "value2");
assertThat(uuid1).isNotEqualTo(uuid2);
}
@Test
public void testCheckpointId() {
// Checkpoint IDs should be valid UUIDs
String id = ID.checkpointId("thread-123");
// Should be a valid UUID string
UUID uuid = UUID.fromString(id);
assertThat(uuid).isNotNull();
}
@Test
public void testUrlSafeId() {
// URL-safe IDs should be deterministic
String id1 = ID.urlSafeId("test", "value");
String id2 = ID.urlSafeId("test", "value");
assertThat(id1).isEqualTo(id2);
// Should not contain padding characters or unsafe URL characters
assertThat(id1).doesNotContain("=", "+", "/");
}
}
@@ -0,0 +1,155 @@
package com.langgraph.checkpoint.base.memory;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ExecutionException;
import static org.assertj.core.api.Assertions.assertThat;
public class AsyncMemoryCheckpointSaverTest {
private AsyncMemoryCheckpointSaver saver;
@BeforeEach
public void setUp() {
saver = new AsyncMemoryCheckpointSaver();
}
@Test
public void testCheckpointAsync() throws ExecutionException, InterruptedException {
// Create test data
String threadId = "test-thread";
Map<String, Object> values = new HashMap<>();
values.put("key1", "value1");
values.put("key2", 42);
// Create checkpoint asynchronously
CompletableFuture<String> future = saver.checkpointAsync(threadId, values);
// Wait for completion
String checkpointId = future.get();
// Verify checkpoint ID format (should be a UUID)
assertThat(checkpointId).matches("^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$");
// Verify thread has a checkpoint
CompletableFuture<List<String>> listFuture = saver.listAsync(threadId);
List<String> checkpoints = listFuture.get();
assertThat(checkpoints).hasSize(1);
assertThat(checkpoints).contains(checkpointId);
}
@Test
public void testGetValuesAsync() throws ExecutionException, InterruptedException {
// Create test data
String threadId = "test-thread";
Map<String, Object> values = new HashMap<>();
values.put("key1", "value1");
values.put("key2", 42);
// Create checkpoint
String checkpointId = saver.checkpointAsync(threadId, values).get();
// Get values asynchronously
CompletableFuture<Optional<Map<String, Object>>> future = saver.getValuesAsync(checkpointId);
Optional<Map<String, Object>> retrievedValues = future.get();
// Verify values
assertThat(retrievedValues).isPresent();
assertThat(retrievedValues.get()).containsEntry("key1", "value1");
assertThat(retrievedValues.get()).containsEntry("key2", 42);
}
@Test
public void testListAsync() throws ExecutionException, InterruptedException {
// Create test data
String threadId = "test-thread";
// Initially empty
CompletableFuture<List<String>> initialFuture = saver.listAsync(threadId);
List<String> initial = initialFuture.get();
assertThat(initial).isEmpty();
// Create multiple checkpoints
String id1 = saver.checkpointAsync(threadId, Map.of("key", "value1")).get();
String id2 = saver.checkpointAsync(threadId, Map.of("key", "value2")).get();
String id3 = saver.checkpointAsync(threadId, Map.of("key", "value3")).get();
// List checkpoints asynchronously
CompletableFuture<List<String>> future = saver.listAsync(threadId);
List<String> checkpoints = future.get();
// Verify order and content
assertThat(checkpoints).hasSize(3);
assertThat(checkpoints).containsExactly(id1, id2, id3);
}
@Test
public void testLatestAsync() throws ExecutionException, InterruptedException {
// Create test data
String threadId = "test-thread";
// Initially empty
CompletableFuture<Optional<String>> initialFuture = saver.latestAsync(threadId);
Optional<String> initial = initialFuture.get();
assertThat(initial).isEmpty();
// Create multiple checkpoints
saver.checkpointAsync(threadId, Map.of("key", "value1")).get();
saver.checkpointAsync(threadId, Map.of("key", "value2")).get();
String id3 = saver.checkpointAsync(threadId, Map.of("key", "value3")).get();
// Get latest asynchronously
CompletableFuture<Optional<String>> future = saver.latestAsync(threadId);
Optional<String> latest = future.get();
// Verify latest
assertThat(latest).isPresent();
assertThat(latest.get()).isEqualTo(id3);
}
@Test
public void testDeleteAsync() throws ExecutionException, InterruptedException {
// Create test data
String threadId = "test-thread";
// Create checkpoint
String checkpointId = saver.checkpointAsync(threadId, Map.of("key", "value")).get();
// Verify checkpoint exists
assertThat(saver.getValuesAsync(checkpointId).get()).isPresent();
// Delete checkpoint asynchronously
CompletableFuture<Void> future = saver.deleteAsync(checkpointId);
future.get(); // Wait for completion
// Verify checkpoint is deleted
assertThat(saver.getValuesAsync(checkpointId).get()).isEmpty();
}
@Test
public void testClearAsync() throws ExecutionException, InterruptedException {
// Create test data
String threadId = "test-thread";
// Create multiple checkpoints
String id1 = saver.checkpointAsync(threadId, Map.of("key", "value1")).get();
String id2 = saver.checkpointAsync(threadId, Map.of("key", "value2")).get();
// Verify checkpoints exist
assertThat(saver.listAsync(threadId).get()).hasSize(2);
// Clear thread asynchronously
CompletableFuture<Void> future = saver.clearAsync(threadId);
future.get(); // Wait for completion
// Verify checkpoints are deleted
assertThat(saver.listAsync(threadId).get()).isEmpty();
}
}
@@ -0,0 +1,161 @@
package com.langgraph.checkpoint.base.memory;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import static org.assertj.core.api.Assertions.assertThat;
public class MemoryCheckpointSaverTest {
private MemoryCheckpointSaver saver;
@BeforeEach
public void setUp() {
saver = new MemoryCheckpointSaver();
}
@Test
public void testCheckpoint() {
// Create test data
String threadId = "test-thread";
Map<String, Object> values = new HashMap<>();
values.put("key1", "value1");
values.put("key2", 42);
// Create checkpoint
String checkpointId = saver.checkpoint(threadId, values);
// Verify checkpoint ID format (should be a UUID)
assertThat(checkpointId).matches("^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$");
// Verify thread has a checkpoint
List<String> checkpoints = saver.list(threadId);
assertThat(checkpoints).hasSize(1);
assertThat(checkpoints).contains(checkpointId);
// Verify latest checkpoint
Optional<String> latest = saver.latest(threadId);
assertThat(latest).isPresent();
assertThat(latest.get()).isEqualTo(checkpointId);
}
@Test
public void testGetValues() {
// Create test data
String threadId = "test-thread";
Map<String, Object> values = new HashMap<>();
values.put("key1", "value1");
values.put("key2", 42);
// Create checkpoint
String checkpointId = saver.checkpoint(threadId, values);
// Get values
Optional<Map<String, Object>> retrievedValues = saver.getValues(checkpointId);
// Verify values
assertThat(retrievedValues).isPresent();
assertThat(retrievedValues.get()).containsEntry("key1", "value1");
assertThat(retrievedValues.get()).containsEntry("key2", 42);
// Verify non-existent checkpoint
Optional<Map<String, Object>> nonExistent = saver.getValues("non-existent");
assertThat(nonExistent).isEmpty();
}
@Test
public void testList() {
// Create test data
String threadId = "test-thread";
// Initially empty
List<String> initial = saver.list(threadId);
assertThat(initial).isEmpty();
// Create multiple checkpoints
String id1 = saver.checkpoint(threadId, Map.of("key", "value1"));
String id2 = saver.checkpoint(threadId, Map.of("key", "value2"));
String id3 = saver.checkpoint(threadId, Map.of("key", "value3"));
// List checkpoints
List<String> checkpoints = saver.list(threadId);
// Verify order and content
assertThat(checkpoints).hasSize(3);
assertThat(checkpoints).containsExactly(id1, id2, id3);
// Different thread should have no checkpoints
List<String> otherThread = saver.list("other-thread");
assertThat(otherThread).isEmpty();
}
@Test
public void testLatest() {
// Create test data
String threadId = "test-thread";
// Initially empty
Optional<String> initial = saver.latest(threadId);
assertThat(initial).isEmpty();
// Create multiple checkpoints
saver.checkpoint(threadId, Map.of("key", "value1"));
saver.checkpoint(threadId, Map.of("key", "value2"));
String id3 = saver.checkpoint(threadId, Map.of("key", "value3"));
// Get latest
Optional<String> latest = saver.latest(threadId);
// Verify latest
assertThat(latest).isPresent();
assertThat(latest.get()).isEqualTo(id3);
}
@Test
public void testDelete() {
// Create test data
String threadId = "test-thread";
// Create checkpoint
String checkpointId = saver.checkpoint(threadId, Map.of("key", "value"));
// Verify checkpoint exists
assertThat(saver.getValues(checkpointId)).isPresent();
assertThat(saver.list(threadId)).contains(checkpointId);
// Delete checkpoint
saver.delete(checkpointId);
// Verify checkpoint is deleted
assertThat(saver.getValues(checkpointId)).isEmpty();
assertThat(saver.list(threadId)).doesNotContain(checkpointId);
}
@Test
public void testClear() {
// Create test data
String threadId = "test-thread";
// Create multiple checkpoints
String id1 = saver.checkpoint(threadId, Map.of("key", "value1"));
String id2 = saver.checkpoint(threadId, Map.of("key", "value2"));
// Verify checkpoints exist
assertThat(saver.list(threadId)).hasSize(2);
assertThat(saver.getValues(id1)).isPresent();
assertThat(saver.getValues(id2)).isPresent();
// Clear thread
saver.clear(threadId);
// Verify checkpoints are deleted
assertThat(saver.list(threadId)).isEmpty();
assertThat(saver.getValues(id1)).isEmpty();
assertThat(saver.getValues(id2)).isEmpty();
}
}
@@ -0,0 +1,383 @@
package com.langgraph.checkpoint.serde;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import java.time.Instant;
import java.time.LocalDate;
import java.time.LocalDateTime;
import java.time.LocalTime;
import java.util.*;
import java.util.Objects;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.within;
public class MsgPackSerializerTest {
private MsgPackSerializer serializer;
@BeforeEach
public void setUp() {
serializer = new MsgPackSerializer();
}
@Test
public void testSerializeDeserializePrimitives() {
// Test with various primitive types
assertRoundTrip("Test string");
assertRoundTrip(123);
assertRoundTrip(123456789L);
assertRoundTrip(123.45);
assertRoundTrip(123.45f);
assertRoundTrip(true);
assertRoundTrip(false);
assertRoundTrip(null);
}
@Test
public void testSerializeDeserializeArrays() {
// Test with arrays and collections
assertRoundTrip(new byte[] {1, 2, 3, 4, 5});
assertRoundTrip(Arrays.asList("one", "two", "three"));
assertRoundTrip(Arrays.asList(1, 2, 3, 4, 5));
}
@Test
public void testSerializeDeserializeMaps() {
// Test with maps
Map<String, Object> map = new HashMap<>();
map.put("string", "value");
map.put("int", 123);
map.put("boolean", true);
assertRoundTrip(map);
}
@Test
public void testSerializeDeserializeNestedStructures() {
// Test with nested structures
Map<String, Object> nested = new HashMap<>();
nested.put("list", Arrays.asList(1, 2, 3));
nested.put("map", Map.of("key", "value"));
assertRoundTrip(nested);
}
@Test
public void testSerializeDeserializeEnums() {
// Test with enums
assertRoundTrip(TestEnum.VALUE1);
assertRoundTrip(TestEnum.VALUE2);
assertRoundTrip(TestEnum.VALUE3);
}
@Test
public void testSerializeDeserializeRecord() {
// Test with a record
TestRecord record = new TestRecord("test", 123, Arrays.asList("a", "b", "c"));
// Serialize and deserialize
byte[] serialized = serializer.serialize(record);
Object deserialized = serializer.deserialize(serialized);
// Verify
assertThat(deserialized).isInstanceOf(TestRecord.class);
TestRecord deserializedRecord = (TestRecord) deserialized;
assertThat(deserializedRecord.name()).isEqualTo("test");
assertThat(deserializedRecord.value()).isEqualTo(123);
assertThat(deserializedRecord.tags()).containsExactly("a", "b", "c");
}
@Test
public void testSerializeDeserializeNestedRecord() {
// Test with a nested record
NestedTestRecord record = new NestedTestRecord(
"parent",
new TestRecord("child", 456, Arrays.asList("x", "y", "z"))
);
// Serialize and deserialize
byte[] serialized = serializer.serialize(record);
Object deserialized = serializer.deserialize(serialized);
// Verify
assertThat(deserialized).isInstanceOf(NestedTestRecord.class);
NestedTestRecord deserializedRecord = (NestedTestRecord) deserialized;
assertThat(deserializedRecord.name()).isEqualTo("parent");
assertThat(deserializedRecord.child()).isInstanceOf(TestRecord.class);
assertThat(deserializedRecord.child().name()).isEqualTo("child");
assertThat(deserializedRecord.child().value()).isEqualTo(456);
assertThat(deserializedRecord.child().tags()).containsExactly("x", "y", "z");
}
@Test
public void testSerializeDeserializeCustomObject() {
// Test with a custom object
TestObject obj = new TestObject();
obj.setName("test");
obj.setValue(123);
obj.setActive(true);
// Serialize and deserialize
byte[] serialized = serializer.serialize(obj);
Object deserialized = serializer.deserialize(serialized);
// Verify
assertThat(deserialized).isInstanceOf(TestObject.class);
TestObject deserializedObj = (TestObject) deserialized;
assertThat(deserializedObj.getName()).isEqualTo("test");
assertThat(deserializedObj.getValue()).isEqualTo(123);
assertThat(deserializedObj.isActive()).isTrue();
}
@Test
public void testSerializeDeserializeInheritance() {
// Test with inheritance
ChildTestObject obj = new ChildTestObject();
obj.setName("parent");
obj.setValue(123);
obj.setActive(true);
obj.setChildProperty("child");
obj.setChildValue(456);
// Serialize and deserialize
byte[] serialized = serializer.serialize(obj);
Object deserialized = serializer.deserialize(serialized);
// Verify
assertThat(deserialized).isInstanceOf(ChildTestObject.class);
ChildTestObject deserializedObj = (ChildTestObject) deserialized;
assertThat(deserializedObj.getName()).isEqualTo("parent");
assertThat(deserializedObj.getValue()).isEqualTo(123);
assertThat(deserializedObj.isActive()).isTrue();
assertThat(deserializedObj.getChildProperty()).isEqualTo("child");
assertThat(deserializedObj.getChildValue()).isEqualTo(456);
}
@Test
public void testSerializeDeserializeWithCustomSerializer() {
// Register custom UUID serializer (although built-in one exists)
serializer.registerSerializer(UUID.class, (uuid) -> uuid.toString().replace("-", ""));
serializer.registerDeserializer(UUID.class, (str) -> {
String uuidStr = (String) str;
// Insert hyphens for standard UUID format
uuidStr = uuidStr.replaceFirst(
"(\\p{XDigit}{8})(\\p{XDigit}{4})(\\p{XDigit}{4})(\\p{XDigit}{4})(\\p{XDigit}+)",
"$1-$2-$3-$4-$5");
return UUID.fromString(uuidStr);
});
// Test with UUID
UUID uuid = UUID.randomUUID();
// Serialize and deserialize
byte[] serialized = serializer.serialize(uuid);
Object deserialized = serializer.deserialize(serialized);
// Verify
assertThat(deserialized).isInstanceOf(UUID.class);
assertThat(deserialized).isEqualTo(uuid);
}
@Test
public void testSerializeDeserializeDateTypes() {
// Test with Date
Date date = new Date();
// Serialize and deserialize
byte[] serialized = serializer.serialize(date);
Object deserialized = serializer.deserialize(serialized);
// Verify
assertThat(deserialized).isInstanceOf(Date.class);
assertThat(deserialized).isEqualTo(date);
// Test with Java 8 Date/Time types
Instant instant = Instant.now();
LocalDate localDate = LocalDate.now();
LocalTime localTime = LocalTime.now();
LocalDateTime localDateTime = LocalDateTime.now();
assertRoundTrip(instant);
assertRoundTrip(localDate);
assertRoundTrip(localTime);
assertRoundTrip(localDateTime);
}
@Test
public void testTransientFields() {
// Test with transient fields
ObjectWithTransient obj = new ObjectWithTransient();
obj.setPersistent("saved");
obj.setTransientField("not-saved");
// Serialize and deserialize
byte[] serialized = serializer.serialize(obj);
Object deserialized = serializer.deserialize(serialized);
// Verify
assertThat(deserialized).isInstanceOf(ObjectWithTransient.class);
ObjectWithTransient deserializedObj = (ObjectWithTransient) deserialized;
assertThat(deserializedObj.getPersistent()).isEqualTo("saved");
assertThat(deserializedObj.getTransientField()).isNull(); // Should be null after deserialization
}
/**
* Helper method to assert that an object survives a round trip through serialization.
*
* @param obj Object to test
*/
private void assertRoundTrip(Object obj) {
try {
// Serialize
byte[] serialized = serializer.serialize(obj);
// Deserialize
Object deserialized = serializer.deserialize(serialized);
// Verify
if (obj instanceof byte[]) {
// Arrays need special comparison
assertThat(deserialized).isInstanceOf(byte[].class);
assertThat((byte[]) deserialized).isEqualTo((byte[]) obj);
} else if (obj instanceof Number) {
// For any number type, compare by value instead of exact type
if (deserialized instanceof Number) {
double expected = ((Number) obj).doubleValue();
double actual = ((Number) deserialized).doubleValue();
assertThat(actual).isCloseTo(expected, within(0.0001));
} else {
throw new AssertionError("Expected Number, got " +
(deserialized != null ? deserialized.getClass().getName() : "null"));
}
} else {
// Special handling for lists
if (obj instanceof List && deserialized instanceof List) {
List<?> originalList = (List<?>) obj;
List<?> deserializedList = (List<?>) deserialized;
assertThat(deserializedList).hasSameSizeAs(originalList);
// Check each element
for (int i = 0; i < originalList.size(); i++) {
Object originalItem = originalList.get(i);
Object deserializedItem = deserializedList.get(i);
if (originalItem instanceof Number && deserializedItem instanceof Number) {
// Compare numbers by value instead of exact type
assertThat(((Number) deserializedItem).doubleValue())
.isCloseTo(((Number) originalItem).doubleValue(), within(0.0001));
} else {
assertThat(deserializedItem).isEqualTo(originalItem);
}
}
} else {
// Regular equality for other types
assertThat(deserialized).isEqualTo(obj);
}
}
} catch (Exception e) {
throw new AssertionError("Error in roundtrip for " + obj + ": " + e.getMessage(), e);
}
}
/**
* Test enum.
*/
public enum TestEnum {
VALUE1, VALUE2, VALUE3
}
/**
* Test record class.
*/
public record TestRecord(String name, int value, List<String> tags) {
}
/**
* Nested test record class.
*/
public record NestedTestRecord(String name, TestRecord child) {
}
/**
* Test class for custom object serialization.
*/
public static class TestObject {
private String name;
private int value;
private boolean active;
public String getName() {
return name;
}
public void setName(String name) {
this.name = name;
}
public int getValue() {
return value;
}
public void setValue(int value) {
this.value = value;
}
public boolean isActive() {
return active;
}
public void setActive(boolean active) {
this.active = active;
}
}
/**
* Child test class for inheritance testing.
*/
public static class ChildTestObject extends TestObject {
private String childProperty;
private int childValue;
public String getChildProperty() {
return childProperty;
}
public void setChildProperty(String childProperty) {
this.childProperty = childProperty;
}
public int getChildValue() {
return childValue;
}
public void setChildValue(int childValue) {
this.childValue = childValue;
}
}
/**
* Test class with transient fields.
*/
public static class ObjectWithTransient {
private String persistent;
private transient String transientField;
public String getPersistent() {
return persistent;
}
public void setPersistent(String persistent) {
this.persistent = persistent;
}
public String getTransientField() {
return transientField;
}
public void setTransientField(String transientField) {
this.transientField = transientField;
}
}
}
@@ -0,0 +1,16 @@
dependencies {
// Internal dependencies
implementation project(':langgraph-checkpoint')
// Test dependencies
testImplementation 'org.junit.jupiter:junit-jupiter-api:5.9.1'
testImplementation 'org.junit.jupiter:junit-jupiter-params:5.9.1'
testRuntimeOnly 'org.junit.jupiter:junit-jupiter-engine:5.9.1'
testImplementation 'org.mockito:mockito-core:5.2.0'
testImplementation 'org.mockito:mockito-junit-jupiter:5.2.0'
testImplementation 'org.assertj:assertj-core:3.24.2'
}
test {
useJUnitPlatform()
}
@@ -0,0 +1,144 @@
package com.langgraph.channels;
import java.lang.reflect.Type;
/**
* Abstract base implementation of BaseChannel that provides common functionality.
*
* @param <V> Type of the value stored in the channel
* @param <U> Type of the update received by the channel
* @param <C> Type of the checkpoint representation
*/
public abstract class AbstractChannel<V, U, C> implements BaseChannel<V, U, C> {
/**
* The full generic type information for value type.
*/
protected final TypeReference<V> valueTypeRef;
/**
* The full generic type information for update type.
*/
protected final TypeReference<U> updateTypeRef;
/**
* The full generic type information for checkpoint type.
*/
protected final TypeReference<C> checkpointTypeRef;
/**
* The channel key (name).
*/
protected String key = "";
/**
* Creates a new channel with full generic type information.
*
* @param valueTypeRef TypeReference for the value type
* @param updateTypeRef TypeReference for the update type
* @param checkpointTypeRef TypeReference for the checkpoint type
*/
protected AbstractChannel(TypeReference<V> valueTypeRef, TypeReference<U> updateTypeRef, TypeReference<C> checkpointTypeRef) {
this.valueTypeRef = valueTypeRef;
this.updateTypeRef = updateTypeRef;
this.checkpointTypeRef = checkpointTypeRef;
}
/**
* Creates a new channel with full generic type information and key.
*
* @param valueTypeRef TypeReference for the value type
* @param updateTypeRef TypeReference for the update type
* @param checkpointTypeRef TypeReference for the checkpoint type
* @param key The key (name) of this channel
*/
protected AbstractChannel(TypeReference<V> valueTypeRef, TypeReference<U> updateTypeRef,
TypeReference<C> checkpointTypeRef, String key) {
this.valueTypeRef = valueTypeRef;
this.updateTypeRef = updateTypeRef;
this.checkpointTypeRef = checkpointTypeRef;
this.key = key;
}
@Override
public String getKey() {
return key;
}
@Override
public void setKey(String key) {
this.key = key;
}
/**
* By default, checkpoint returns the current value.
* Note: This implementation assumes C and V are the same type for most channels.
* Subclasses where C and V differ MUST override this method.
*/
@Override
public C checkpoint() throws EmptyChannelException {
try {
// This cast is unavoidable due to Java generics limitations
// We can't enforce that C = V at compile time, so runtime cast is needed
// Each subclass properly implements fromCheckpoint to handle this correctly
@SuppressWarnings("unchecked")
C value = (C) get();
return value;
} catch (EmptyChannelException e) {
// For Python compatibility, allow checkpointing uninitialized channels
return null;
}
}
@Override
public Class<V> getValueType() {
return valueTypeRef.getRawClass();
}
@Override
public Class<U> getUpdateType() {
return updateTypeRef.getRawClass();
}
@Override
public Class<C> getCheckpointType() {
return checkpointTypeRef.getRawClass();
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
AbstractChannel<?, ?, ?> that = (AbstractChannel<?, ?, ?>) o;
// Compare key, valueTypeRef, updateTypeRef, and checkpointTypeRef
// The exact comparison of stored values is responsibility of subclasses
if (!key.equals(that.key)) return false;
if (!valueTypeRef.equals(that.valueTypeRef)) return false;
if (!updateTypeRef.equals(that.updateTypeRef)) return false;
return checkpointTypeRef.equals(that.checkpointTypeRef);
}
@Override
public int hashCode() {
int result = valueTypeRef.hashCode();
result = 31 * result + updateTypeRef.hashCode();
result = 31 * result + checkpointTypeRef.hashCode();
result = 31 * result + key.hashCode();
return result;
}
/**
* Handle a single-value update when channels support it.
* This method can be overridden by channels that want to support
* single-value updates (like TopicChannel). The default implementation
* returns false, indicating the update was not handled.
*
* @param singleValue The single value to update with
* @return true if the channel was updated, false otherwise
*/
public boolean updateSingleValue(U singleValue) {
// Default implementation doesn't support single value updates
return false;
}
}
@@ -0,0 +1,133 @@
package com.langgraph.channels;
import java.util.List;
/**
* Base interface for all channels in LangGraph.
* Channels are the primary mechanism for passing data between nodes in a LangGraph
* computational graph. They implement different semantics for handling updates
* (e.g., storing just the last value, aggregating values, etc.)
*
* @param <V> Type of the value stored in the channel
* @param <U> Type of the update received by the channel
* @param <C> Type of the checkpoint representation
*/
public interface BaseChannel<V, U, C> {
/**
* Returns the current value of the channel without type safety checks.
* This is mainly used internally by the framework.
*
* @return Current value as Object, or null if the channel has not been updated yet
*/
default Object getValue() {
try {
return get();
} catch (EmptyChannelException e) {
// Return null for Python compatibility when channel is not initialized
return null;
}
}
/**
* Marks the channel as updated. This is mainly used internally.
*/
default void resetUpdated() {
// Default implementation does nothing
}
/**
* Updates the channel with a sequence of values.
* The order of the updates in the list is arbitrary.
*
* @param values List of update values
* @return true if the channel was updated, false otherwise
* @throws InvalidUpdateException if the update is invalid for this channel type
*/
boolean update(List<U> values) throws InvalidUpdateException;
/**
* Returns the current value of the channel.
*
* @return Current value
* @throws EmptyChannelException if the channel has not been updated yet
*/
V get() throws EmptyChannelException;
/**
* Creates a checkpoint of the channel's current state.
*
* @return A serializable representation of the channel's state
* @throws EmptyChannelException if the channel has not been updated yet
*/
C checkpoint() throws EmptyChannelException;
/**
* Creates a new channel instance from a checkpoint.
*
* @param checkpoint Checkpoint data, or null if no prior state
* @return A new channel instance with the state from the checkpoint
*/
BaseChannel<V, U, C> fromCheckpoint(C checkpoint);
/**
* Marks the current value as consumed.
* By default, this is a no-op.
*
* @return true if the channel was updated, false otherwise
*/
default boolean consume() {
return false;
}
/**
* Returns the key/name of this channel.
*
* @return Channel key/name
*/
String getKey();
/**
* Sets the key/name for this channel.
*
* @param key Channel key/name
*/
void setKey(String key);
/**
* Returns the Class representing the type of values stored in this channel.
* This is useful for runtime type checking.
*
* @return The Class object for the value type
*/
Class<V> getValueType();
/**
* Returns the Class representing the type of updates this channel accepts.
* This enables runtime type checking of inputs.
*
* @return The Class object for the update type
*/
Class<U> getUpdateType();
/**
* Returns the Class representing the type of checkpoint data for this channel.
* Useful for serialization and deserialization.
*
* @return The Class object for the checkpoint type
*/
Class<C> getCheckpointType();
/**
* Updates the channel with a single value.
* This is a convenience method that some channel implementations may support
* for single-value updates. The default implementation returns false, indicating
* the single-value update was not handled.
*
* @param singleValue A single update value
* @return true if the channel was updated, false otherwise
* @throws InvalidUpdateException if the update is invalid for this channel type
*/
default boolean updateSingleValue(U singleValue) throws InvalidUpdateException {
return false;
}
}
@@ -0,0 +1,159 @@
package com.langgraph.channels;
import java.util.List;
import java.util.function.BinaryOperator;
/**
* A channel that aggregates values using a binary operator.
* This is useful for operations like sum, max, min, etc.
*
* @param <V> Type of the value stored in the channel
*/
public class BinaryOperatorChannel<V> extends AbstractChannel<V, V, V> {
/**
* The binary operator to apply for aggregation.
*/
private final BinaryOperator<V> operator;
/**
* The current value, null if the channel has not been updated yet.
*/
private V value;
/**
* The initial value to use if none has been set yet.
*/
private final V initialValue;
/**
* Flag to track if this channel has been initialized.
*/
private boolean initialized = false;
/**
* Creates a new BinaryOperatorChannel with the specified value type and operator.
*
* @param typeRef The TypeReference for the value type
* @param operator The binary operator to use for aggregation
* @param initialValue The initial value to use if none has been set yet
*/
protected BinaryOperatorChannel(TypeReference<V> typeRef, BinaryOperator<V> operator, V initialValue) {
super(typeRef, typeRef, typeRef); // For BinaryOperatorChannel, V=U=C
this.operator = operator;
this.initialValue = initialValue;
}
/**
* Creates a new BinaryOperatorChannel with the specified value type, key, and operator.
*
* @param typeRef The TypeReference for the value type
* @param key The key (name) of this channel
* @param operator The binary operator to use for aggregation
* @param initialValue The initial value to use if none has been set yet
*/
protected BinaryOperatorChannel(TypeReference<V> typeRef, String key, BinaryOperator<V> operator, V initialValue) {
super(typeRef, typeRef, typeRef, key); // For BinaryOperatorChannel, V=U=C
this.operator = operator;
this.initialValue = initialValue;
}
/**
* Factory method to create a BinaryOperatorChannel with proper generic type capture.
*
* <p>Example usage:
* <pre>
* BinaryOperatorChannel&lt;Integer&gt; channel = BinaryOperatorChannel.&lt;Integer&gt;create(Integer::sum, 0);
* </pre>
*
* @param <T> The type parameter for the channel
* @param operator The binary operator to use for aggregation
* @param initialValue The initial value to use if none has been set yet
* @return A new BinaryOperatorChannel with the captured type parameter
*/
public static <T> BinaryOperatorChannel<T> create(BinaryOperator<T> operator, T initialValue) {
return new BinaryOperatorChannel<>(new TypeReference<T>() {}, operator, initialValue);
}
/**
* Factory method to create a BinaryOperatorChannel with proper generic type capture
* and a specified key.
*
* <p>Example usage:
* <pre>
* BinaryOperatorChannel&lt;Integer&gt; channel = BinaryOperatorChannel.&lt;Integer&gt;create("counter", Integer::sum, 0);
* </pre>
*
* @param <T> The type parameter for the channel
* @param key The key (name) for the channel
* @param operator The binary operator to use for aggregation
* @param initialValue The initial value to use if none has been set yet
* @return A new BinaryOperatorChannel with the captured type parameter and specified key
*/
public static <T> BinaryOperatorChannel<T> create(String key, BinaryOperator<T> operator, T initialValue) {
return new BinaryOperatorChannel<>(new TypeReference<T>() {}, key, operator, initialValue);
}
@Override
public boolean update(List<V> values) {
if (values.isEmpty()) {
return false;
}
V current = initialized ? this.value : initialValue;
for (V val : values) {
current = operator.apply(current, val);
}
this.value = current;
initialized = true;
return true;
}
@Override
public V get() throws EmptyChannelException {
if (!initialized) {
throw new EmptyChannelException(
"BinaryOperatorChannel at key '" + key + "' is empty (never updated)");
}
return value;
}
@Override
public BaseChannel<V, V, V> fromCheckpoint(V checkpoint) {
BinaryOperatorChannel<V> newChannel = new BinaryOperatorChannel<>(
valueTypeRef, key, operator, initialValue);
// Even null is a valid checkpoint value - it means the channel was initialized with null
newChannel.value = checkpoint;
newChannel.initialized = true;
return newChannel;
}
/**
* Returns the string representation of this channel.
*
* @return String representation
*/
@Override
public String toString() {
return "BinaryOperator(" + (initialized ? value : "empty") + ")";
}
/**
* Returns the binary operator.
*
* @return The binary operator
*/
public BinaryOperator<V> getOperator() {
return operator;
}
/**
* Returns the initial value.
*
* @return The initial value
*/
public V getInitialValue() {
return initialValue;
}
}
@@ -0,0 +1,176 @@
package com.langgraph.channels;
import java.util.function.BinaryOperator;
/**
* Utility class for creating channels easily.
*/
public final class Channels {
private Channels() {
// Private constructor to prevent instantiation
}
/**
* Creates a LastValue channel.
*
* @param <V> The type of values
* @return A new LastValue channel
*/
public static <V> LastValue<V> lastValue() {
return LastValue.<V>create();
}
/**
* Creates a LastValue channel with the specified key.
*
* @param key The key (name) of the channel
* @param <V> The type of values
* @return A new LastValue channel
*/
public static <V> LastValue<V> lastValue(String key) {
return LastValue.<V>create(key);
}
/**
* Creates a Topic channel.
*
* @param <V> The type of values
* @return A new Topic channel
*/
public static <V> TopicChannel<V> topic() {
return TopicChannel.<V>create();
}
/**
* Creates a Topic channel with reset-on-consume behavior.
*
* @param resetOnConsume Whether to reset the channel when consumed
* @param <V> The type of values
* @return A new Topic channel
*/
public static <V> TopicChannel<V> topic(boolean resetOnConsume) {
return TopicChannel.<V>create(resetOnConsume);
}
/**
* Creates a Topic channel with the specified key.
*
* @param key The key (name) of the channel
* @param resetOnConsume Whether to reset the channel when consumed
* @param <V> The type of values
* @return A new Topic channel
*/
public static <V> TopicChannel<V> topic(String key, boolean resetOnConsume) {
return TopicChannel.<V>create(key, resetOnConsume);
}
/**
* Creates a BinaryOperator channel.
*
* @param operator The binary operator to use for aggregation
* @param initialValue The initial value
* @param <V> The type of values
* @return A new BinaryOperator channel
*/
public static <V> BinaryOperatorChannel<V> binaryOperator(
BinaryOperator<V> operator, V initialValue) {
return BinaryOperatorChannel.<V>create(operator, initialValue);
}
/**
* Creates a BinaryOperator channel with the specified key.
*
* @param key The key (name) of the channel
* @param operator The binary operator to use for aggregation
* @param initialValue The initial value
* @param <V> The type of values
* @return A new BinaryOperator channel
*/
public static <V> BinaryOperatorChannel<V> binaryOperator(
String key, BinaryOperator<V> operator, V initialValue) {
return BinaryOperatorChannel.<V>create(key, operator, initialValue);
}
/**
* Creates an EphemeralValue channel.
*
* @param <V> The type of values
* @return A new EphemeralValue channel
*/
public static <V> EphemeralValue<V> ephemeral() {
return EphemeralValue.<V>create();
}
/**
* Creates an EphemeralValue channel with the specified key.
*
* @param key The key (name) of the channel
* @param <V> The type of values
* @return A new EphemeralValue channel
*/
public static <V> EphemeralValue<V> ephemeral(String key) {
return EphemeralValue.<V>create(key);
}
// Common binary operators for numeric types
/**
* Creates an Integer adder binary operator channel.
*
* @param key The key (name) of the channel
* @return A new BinaryOperator channel for adding integers
*/
public static BinaryOperatorChannel<Integer> integerAdder(String key) {
return BinaryOperatorChannel.create(key, Integer::sum, 0);
}
/**
* Creates a Long adder binary operator channel.
*
* @param key The key (name) of the channel
* @return A new BinaryOperator channel for adding longs
*/
public static BinaryOperatorChannel<Long> longAdder(String key) {
return BinaryOperatorChannel.create(key, Long::sum, 0L);
}
/**
* Creates a Double adder binary operator channel.
*
* @param key The key (name) of the channel
* @return A new BinaryOperator channel for adding doubles
*/
public static BinaryOperatorChannel<Double> doubleAdder(String key) {
return BinaryOperatorChannel.create(key, Double::sum, 0.0);
}
/**
* Creates an Integer max binary operator channel.
*
* @param key The key (name) of the channel
* @return A new BinaryOperator channel for finding the maximum integer
*/
public static BinaryOperatorChannel<Integer> integerMax(String key) {
return BinaryOperatorChannel.create(key, Integer::max, Integer.MIN_VALUE);
}
/**
* Creates a Long max binary operator channel.
*
* @param key The key (name) of the channel
* @return A new BinaryOperator channel for finding the maximum long
*/
public static BinaryOperatorChannel<Long> longMax(String key) {
return BinaryOperatorChannel.create(key, Long::max, Long.MIN_VALUE);
}
/**
* Creates a Double max binary operator channel.
*
* @param key The key (name) of the channel
* @return A new BinaryOperator channel for finding the maximum double
*/
public static BinaryOperatorChannel<Double> doubleMax(String key) {
return BinaryOperatorChannel.create(key, Double::max, Double.MIN_VALUE);
}
}
@@ -0,0 +1,23 @@
package com.langgraph.channels;
/**
* Exception thrown when trying to access a value from a channel that hasn't been
* updated yet.
*/
public class EmptyChannelException extends RuntimeException {
/**
* Creates a new EmptyChannelException.
*/
public EmptyChannelException() {
super("Channel is empty (never updated)");
}
/**
* Creates a new EmptyChannelException with a custom message.
*
* @param message The error message
*/
public EmptyChannelException(String message) {
super(message);
}
}
@@ -0,0 +1,159 @@
package com.langgraph.channels;
import java.util.List;
/**
* A channel that stores the last value received but doesn't persist it across checkpoints.
* This is useful for values that should not be saved in the persistent state.
*
* @param <V> Type of the value stored in the channel
*/
public class EphemeralValue<V> extends AbstractChannel<V, V, Void> {
/**
* The current value, null if the channel has not been updated yet.
*/
private V value;
/**
* Flag to track if this channel has been initialized.
*/
private boolean initialized = false;
/**
* Creates a new EphemeralValue channel using TypeReference with the specified value type.
*
* @param valueTypeRef TypeReference capturing the value type
*/
@SuppressWarnings("unchecked")
protected EphemeralValue(TypeReference<V> valueTypeRef) {
// For EphemeralValue, V=U but C is Void (always null in checkpoint)
super(valueTypeRef, valueTypeRef, new TypeReference<Void>() {});
}
/**
* Creates a new EphemeralValue channel using TypeReference with specified key.
*
* @param valueTypeRef TypeReference capturing the value type
* @param key The key (name) of this channel
*/
@SuppressWarnings("unchecked")
protected EphemeralValue(TypeReference<V> valueTypeRef, String key) {
// For EphemeralValue, V=U but C is Void (always null in checkpoint)
super(valueTypeRef, valueTypeRef, new TypeReference<Void>() {}, key);
}
/**
* Factory method to create an EphemeralValue channel with proper generic type capture.
*
* <p>Example usage:
* <pre>
* EphemeralValue&lt;String&gt; channel = EphemeralValue.&lt;String&gt;create();
* </pre>
*
* @param <T> The type parameter for the channel
* @return A new EphemeralValue channel with the captured type parameter
*/
public static <T> EphemeralValue<T> create() {
return new EphemeralValue<>(new TypeReference<T>() {});
}
/**
* Factory method to create an EphemeralValue channel with proper generic type capture
* and a specified key.
*
* <p>Example usage:
* <pre>
* EphemeralValue&lt;String&gt; channel = EphemeralValue.&lt;String&gt;create("myChannel");
* </pre>
*
* @param <T> The type parameter for the channel
* @param key The key (name) for the channel
* @return A new EphemeralValue channel with the captured type parameter and specified key
*/
public static <T> EphemeralValue<T> create(String key) {
return new EphemeralValue<>(new TypeReference<T>() {}, key);
}
@Override
public boolean update(List<V> values) throws InvalidUpdateException {
if (values.isEmpty()) {
return false;
}
if (values.size() > 1) {
throw new InvalidUpdateException(
"At key '" + key + "': EphemeralValue channel can receive only one value per update. " +
"Use a different channel type to handle multiple values.");
}
value = values.get(0);
initialized = true;
return true;
}
@Override
public V get() throws EmptyChannelException {
if (!initialized) {
throw new EmptyChannelException("EphemeralValue channel at key '" + key + "' is empty (never updated)");
}
return value;
}
@Override
public Void checkpoint() {
// Ephemeral values don't persist in checkpoints
return null;
}
@Override
public BaseChannel<V, V, Void> fromCheckpoint(Void checkpoint) {
// Always start from an empty state, regardless of checkpoint
return new EphemeralValue<>(valueTypeRef, key);
}
/**
* Returns the string representation of this channel.
*
* @return String representation
*/
@Override
public String toString() {
return "EphemeralValue(" + (initialized ? value : "empty") + ")";
}
/**
* Checks if this channel is equal to another object.
*
* @param obj The object to compare with
* @return true if the objects are equal, false otherwise
*/
@Override
public boolean equals(Object obj) {
if (this == obj) {
return true;
}
if (!(obj instanceof EphemeralValue)) {
return false;
}
if (!super.equals(obj)) {
return false;
}
EphemeralValue<?> other = (EphemeralValue<?>) obj;
return initialized == other.initialized &&
(value == null ? other.value == null : value.equals(other.value));
}
/**
* Returns the hash code of this channel.
*
* @return The hash code
*/
@Override
public int hashCode() {
int result = super.hashCode();
result = 31 * result + (initialized ? 1 : 0);
result = 31 * result + (value != null ? value.hashCode() : 0);
return result;
}
}
@@ -0,0 +1,22 @@
package com.langgraph.channels;
/**
* Exception thrown when an invalid update is attempted on a channel.
*/
public class InvalidUpdateException extends RuntimeException {
/**
* Creates a new InvalidUpdateException.
*/
public InvalidUpdateException() {
super("Invalid update for channel");
}
/**
* Creates a new InvalidUpdateException with a custom message.
*
* @param message The error message
*/
public InvalidUpdateException(String message) {
super(message);
}
}
@@ -0,0 +1,157 @@
package com.langgraph.channels;
import java.util.List;
/**
* A channel that stores the last value received.
* Can receive at most one value per update.
*
* @param <V> Type of the value stored in the channel
*/
public class LastValue<V> extends AbstractChannel<V, V, V> {
/**
* The current value, null if the channel has not been updated yet.
*/
private V value;
/**
* Flag to track if this channel has been initialized.
*/
private boolean initialized = false;
/**
* Creates a new LastValue channel using TypeReference to preserve generic type information.
* This is especially useful for generic types like List&lt;Integer&gt;.
*
* @param typeRef The TypeReference that captures the full generic type
*/
protected LastValue(TypeReference<V> typeRef) {
// For LastValue, V=U=C (they are all the same type)
super(typeRef, typeRef, typeRef);
}
/**
* Creates a new LastValue channel using TypeReference to preserve generic type information,
* with the specified key.
*
* @param typeRef The TypeReference that captures the full generic type
* @param key The key (name) of this channel
*/
protected LastValue(TypeReference<V> typeRef, String key) {
// For LastValue, V=U=C (they are all the same type)
super(typeRef, typeRef, typeRef, key);
}
/**
* Factory method to create a LastValue channel with proper generic type capture.
* Use this instead of constructor when dealing with generic types like List&lt;Integer&gt;.
*
* <p>Example usage:
* <pre>
* LastValue&lt;List&lt;Integer&gt;&gt; channel = LastValue.&lt;List&lt;Integer&gt;&gt;create();
* </pre>
*
* @param <T> The type parameter for the channel
* @return A new LastValue channel with the captured type parameter
*/
public static <T> LastValue<T> create() {
return new LastValue<>(new TypeReference<T>() {});
}
/**
* Factory method to create a LastValue channel with proper generic type capture
* and a specified key.
*
* <p>Example usage:
* <pre>
* LastValue&lt;List&lt;Integer&gt;&gt; channel = LastValue.&lt;List&lt;Integer&gt;&gt;create("myChannel");
* </pre>
*
* @param <T> The type parameter for the channel
* @param key The key (name) for the channel
* @return A new LastValue channel with the captured type parameter and specified key
*/
public static <T> LastValue<T> create(String key) {
return new LastValue<>(new TypeReference<T>() {}, key);
}
@Override
public boolean update(List<V> values) throws InvalidUpdateException {
if (values.isEmpty()) {
return false;
}
if (values.size() > 1) {
throw new InvalidUpdateException(
"At key '" + key + "': LastValue channel can receive only one value per update. " +
"Use a different channel type to handle multiple values.");
}
value = values.get(0);
initialized = true;
return true;
}
@Override
public V get() throws EmptyChannelException {
// Return null if not initialized, for Python compatibility
// This prevents EmptyChannelException when accessing uninitialized channels
return value;
}
@Override
public BaseChannel<V, V, V> fromCheckpoint(V checkpoint) {
LastValue<V> newChannel = new LastValue<>(valueTypeRef, key);
// Even null is a valid checkpoint value - it means the channel was initialized with null
newChannel.value = checkpoint;
newChannel.initialized = true;
return newChannel;
}
/**
* Returns the string representation of this channel.
*
* @return String representation
*/
@Override
public String toString() {
return "LastValue(" + (initialized ? value : "empty") + ")";
}
/**
* Checks if this channel is equal to another object.
*
* @param obj The object to compare with
* @return true if the objects are equal, false otherwise
*/
@Override
public boolean equals(Object obj) {
if (this == obj) {
return true;
}
if (!(obj instanceof LastValue)) {
return false;
}
if (!super.equals(obj)) {
return false;
}
LastValue<?> other = (LastValue<?>) obj;
return initialized == other.initialized &&
(value == null ? other.value == null : value.equals(other.value));
}
/**
* Returns the hash code of this channel.
*
* @return The hash code
*/
@Override
public int hashCode() {
int result = super.hashCode();
result = 31 * result + (initialized ? 1 : 0);
result = 31 * result + (value != null ? value.hashCode() : 0);
return result;
}
}
@@ -0,0 +1,293 @@
package com.langgraph.channels;
import java.lang.reflect.ParameterizedType;
import java.lang.reflect.Type;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
/**
* A channel that collects values into a list.
* Unlike LastValue, it can receive multiple values per update.
*
* @param <V> Type of the value stored in the channel
*/
public class TopicChannel<V> extends AbstractChannel<List<V>, V, List<V>> {
/**
* The current list of values, empty if the channel has not been updated yet.
*/
private List<V> values = new ArrayList<>();
/**
* Flag to track if this channel has been initialized.
*/
private boolean initialized = false;
/**
* Flag to determine if the channel should reset after consumption.
*/
private final boolean resetOnConsume;
/**
* Creates a new Topic channel using TypeReference with specified reset behavior.
*
* @param elementTypeRef TypeReference capturing the element type
* @param resetOnConsume Whether to reset the channel after consume() is called
*/
protected TopicChannel(TypeReference<V> elementTypeRef, boolean resetOnConsume) {
// Create TypeReferences for the other types (List<V> in this case)
super(
createListTypeReference(elementTypeRef), // Value type (List<V>)
elementTypeRef, // Update type (V)
createListTypeReference(elementTypeRef) // Checkpoint type (List<V>)
);
this.resetOnConsume = resetOnConsume;
}
/**
* Creates a new Topic channel using TypeReference with specified key and reset behavior.
*
* @param elementTypeRef TypeReference capturing the element type
* @param key The key (name) of this channel
* @param resetOnConsume Whether to reset the channel after consume() is called
*/
protected TopicChannel(TypeReference<V> elementTypeRef, String key, boolean resetOnConsume) {
// Create TypeReferences for the other types (List<V> in this case)
super(
createListTypeReference(elementTypeRef), // Value type (List<V>)
elementTypeRef, // Update type (V)
createListTypeReference(elementTypeRef), // Checkpoint type (List<V>)
key
);
this.resetOnConsume = resetOnConsume;
}
/**
* Factory method to create a TypeReference for List<V> given a TypeReference for V.
*
* @param <V> The element type
* @param elementTypeRef The TypeReference for the element type
* @return A TypeReference for List<V>
*/
@SuppressWarnings("unchecked")
private static <V> TypeReference<List<V>> createListTypeReference(final TypeReference<V> elementTypeRef) {
final Type elementType = elementTypeRef.getType();
return new TypeReference<List<V>>() {
@Override
public Type getType() {
// Create a ParameterizedType for List<V>
return new ParameterizedType() {
@Override
public Type[] getActualTypeArguments() {
return new Type[] { elementType };
}
@Override
public Type getRawType() {
return List.class;
}
@Override
public Type getOwnerType() {
return null;
}
@Override
public String toString() {
return "java.util.List<" + elementType + ">";
}
};
}
@Override
public Class<List<V>> getRawClass() {
return (Class<List<V>>) (Class<?>) List.class;
}
};
}
/**
* Factory method to create a TopicChannel with proper generic type inference.
*
* <p>Example usage:
* <pre>
* TopicChannel&lt;Integer&gt; channel = TopicChannel.&lt;Integer&gt;create();
* </pre>
*
* @param <T> The element type parameter for the channel
* @return A new TopicChannel with the captured type parameter
*/
public static <T> TopicChannel<T> create() {
return new TopicChannel<>(new TypeReference<T>() {}, false);
}
/**
* Factory method to create a TopicChannel with proper generic type inference
* and a specified key.
*
* <p>Example usage:
* <pre>
* TopicChannel&lt;Integer&gt; channel = TopicChannel.&lt;Integer&gt;create("myChannel");
* </pre>
*
* @param <T> The element type parameter for the channel
* @param key The key (name) for the channel
* @return A new TopicChannel with the captured type parameter and specified key
*/
public static <T> TopicChannel<T> create(String key) {
return new TopicChannel<>(new TypeReference<T>() {}, key, false);
}
/**
* Factory method to create a TopicChannel with proper generic type inference,
* specified key, and reset behavior.
*
* <p>Example usage:
* <pre>
* TopicChannel&lt;Integer&gt; channel = TopicChannel.&lt;Integer&gt;create(true);
* TopicChannel&lt;String&gt; channel = TopicChannel.&lt;String&gt;create("myChannel", true);
* </pre>
*
* @param <T> The element type parameter for the channel
* @param resetOnConsume Whether to reset the channel after consume() is called
* @return A new TopicChannel with the captured type parameter
*/
public static <T> TopicChannel<T> create(boolean resetOnConsume) {
return new TopicChannel<>(new TypeReference<T>() {}, resetOnConsume);
}
/**
* Factory method to create a TopicChannel with proper generic type inference,
* specified key, and reset behavior.
*
* @param <T> The element type parameter for the channel
* @param key The key (name) for the channel
* @param resetOnConsume Whether to reset the channel after consume() is called
* @return A new TopicChannel with the captured type parameter and specified settings
*/
public static <T> TopicChannel<T> create(String key, boolean resetOnConsume) {
return new TopicChannel<>(new TypeReference<T>() {}, key, resetOnConsume);
}
@Override
public boolean update(List<V> newValues) {
if (newValues.isEmpty()) {
return false;
}
values.addAll(newValues);
initialized = true;
return true;
}
/**
* Updates the channel with a single new value.
* This is a convenience method for handling cases where the update comes as a single value
* instead of a list.
*
* @param newValue The new value to add to the topic
* @return true if the channel was updated, false otherwise
*/
@Override
public boolean updateSingleValue(V newValue) {
if (newValue == null) {
return false;
}
values.add(newValue);
initialized = true;
return true;
}
@Override
public List<V> get() throws EmptyChannelException {
// Always return the current list (empty or not) for Python compatibility
// This prevents EmptyChannelException when accessing uninitialized channels
return Collections.unmodifiableList(values);
}
@Override
public BaseChannel<List<V>, V, List<V>> fromCheckpoint(List<V> checkpoint) {
// Get the element type reference from the updateTypeRef
TypeReference<V> elementTypeRef = updateTypeRef;
// Create a new channel with the same type information
TopicChannel<V> newChannel = new TopicChannel<>(elementTypeRef, key, resetOnConsume);
// Restore the values from checkpoint
if (checkpoint != null) {
newChannel.values = new ArrayList<>(checkpoint);
newChannel.initialized = true;
}
return newChannel;
}
@Override
public boolean consume() {
if (resetOnConsume && initialized) {
values.clear();
initialized = false;
return true;
}
return false;
}
/**
* Returns the element type class.
*
* @return The element type class
*/
public Class<V> getElementType() {
return updateTypeRef.getRawClass();
}
/**
* Returns the string representation of this channel.
*
* @return String representation
*/
@Override
public String toString() {
return "Topic(" + (initialized ? values : "empty") + ")";
}
/**
* Checks if this channel is equal to another object.
*
* @param obj The object to compare with
* @return true if the objects are equal, false otherwise
*/
@Override
public boolean equals(Object obj) {
if (this == obj) {
return true;
}
if (!(obj instanceof TopicChannel)) {
return false;
}
if (!super.equals(obj)) {
return false;
}
TopicChannel<?> other = (TopicChannel<?>) obj;
return initialized == other.initialized &&
resetOnConsume == other.resetOnConsume &&
values.equals(other.values);
}
/**
* Returns the hash code of this channel.
*
* @return The hash code
*/
@Override
public int hashCode() {
int result = super.hashCode();
result = 31 * result + (initialized ? 1 : 0);
result = 31 * result + (resetOnConsume ? 1 : 0);
result = 31 * result + values.hashCode();
return result;
}
}
@@ -0,0 +1,84 @@
package com.langgraph.channels;
import java.lang.reflect.ParameterizedType;
import java.lang.reflect.Type;
/**
* A runtime type token for preserving generic type information.
* Used to capture generic type parameters that would otherwise be erased.
*
* <p>Usage example:
* <pre>
* TypeReference&lt;List&lt;String&gt;&gt; listStringType = new TypeReference&lt;List&lt;String&gt;&gt;() {};
* </pre>
*
* @param <T> The type to capture
*/
public abstract class TypeReference<T> {
private final Type type;
/**
* Creates a new type reference, capturing the generic type parameter T.
* Due to Java's type erasure, this constructor must be called from an
* anonymous subclass to capture the type information.
*/
protected TypeReference() {
Type superclass = getClass().getGenericSuperclass();
if (superclass instanceof ParameterizedType) {
// Extract the actual type argument from the anonymous subclass
type = ((ParameterizedType) superclass).getActualTypeArguments()[0];
} else {
throw new IllegalArgumentException("TypeReference must be created with type parameters");
}
}
/**
* Gets the captured type.
*
* @return The captured Type
*/
public Type getType() {
return type;
}
/**
* Returns the raw Class for this type reference.
*
* @return The raw Class
*/
@SuppressWarnings("unchecked")
public Class<T> getRawClass() {
try {
if (type instanceof Class<?>) {
return (Class<T>) type;
} else if (type instanceof ParameterizedType) {
return (Class<T>) ((ParameterizedType) type).getRawType();
} else {
// Handle type variables (like T) by returning Object.class as a fallback
return (Class<T>) Object.class;
}
} catch (Exception e) {
// If we encounter any other issue, fallback to Object.class
return (Class<T>) Object.class;
}
}
@Override
public String toString() {
return "TypeReference<" + type + ">";
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass().getSuperclass() != o.getClass().getSuperclass()) return false;
TypeReference<?> that = (TypeReference<?>) o;
return type.equals(that.type);
}
@Override
public int hashCode() {
return type.hashCode();
}
}
@@ -0,0 +1,273 @@
package com.langgraph.graph;
import com.langgraph.channels.BaseChannel;
import com.langgraph.channels.LastValue;
import com.langgraph.channels.TopicChannel;
import com.langgraph.checkpoint.base.BaseCheckpointSaver;
import com.langgraph.pregel.Pregel;
import com.langgraph.pregel.PregelExecutable;
import com.langgraph.pregel.PregelNode;
import com.langgraph.pregel.retry.RetryPolicy;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.function.Function;
/**
* A fluent builder for creating type-safe computational graphs.
* Provides a convenient interface for constructing complex graphs with
* properly typed nodes and channels.
*
* @param <I> The input type for the graph
* @param <O> The output type for the graph
*/
public class GraphBuilder<I, O> {
private final Map<String, BaseChannel<?, ?, ?>> channels = new HashMap<>();
private final Map<String, PregelNode<I, O>> nodes = new HashMap<>();
private BaseCheckpointSaver checkpointer;
private int maxSteps = 100;
/**
* Creates a new GraphBuilder with the specified input and output types.
*
* @param <I> Input type for the graph
* @param <O> Output type for the graph
* @return A new GraphBuilder instance
*/
public static <I, O> GraphBuilder<I, O> create() {
return new GraphBuilder<>();
}
/**
* Creates a new GraphBuilder with String input and output types.
* This is a convenience method for creating graphs that work with String data.
*
* @return A new GraphBuilder instance with String input and output types
*/
public static GraphBuilder<String, String> createStringGraph() {
return new GraphBuilder<>();
}
/**
* Creates a new GraphBuilder with Map input and output types.
* This is a convenience method for creating graphs that work with JSON-like data.
*
* @return A new GraphBuilder instance with Map input and output types
*/
public static GraphBuilder<Map<String, Object>, Map<String, Object>> createJsonGraph() {
return new GraphBuilder<>();
}
/**
* Adds a node to the graph with the specified name and executable.
* By default, the node is configured to read from "input" and write to "output",
* with "input" as its trigger channel.
*
* @param name The name of the node
* @param executable The executable for the node
* @return This builder for method chaining
*/
public GraphBuilder<I, O> addNode(String name, PregelExecutable<I, O> executable) {
PregelNode<I, O> node = new PregelNode.Builder<I, O>(name, executable)
.channels("input")
.triggerChannels("input")
.writers("output")
.build();
nodes.put(name, node);
return this;
}
/**
* Adds a node to the graph with the specified name, executable, and configuration.
* The configurator is a function that can be used to customize the node builder.
*
* @param name The name of the node
* @param executable The executable for the node
* @param configurator A function that configures the node builder
* @return This builder for method chaining
*/
public GraphBuilder<I, O> addNode(String name, PregelExecutable<I, O> executable,
Function<PregelNode.Builder<I, O>, PregelNode.Builder<I, O>> configurator) {
PregelNode.Builder<I, O> builder = new PregelNode.Builder<>(name, executable);
builder = configurator.apply(builder);
nodes.put(name, builder.build());
return this;
}
/**
* Adds a pre-built node to the graph.
*
* @param node The node to add
* @return This builder for method chaining
*/
public GraphBuilder<I, O> addNode(PregelNode<I, O> node) {
nodes.put(node.getName(), node);
return this;
}
/**
* Adds a LastValue channel to the graph.
*
* @param name The name of the channel
* @param <T> The type of the value stored in the channel
* @return This builder for method chaining
*/
public <T> GraphBuilder<I, O> addLastValueChannel(String name) {
LastValue<T> channel = LastValue.<T>create(name);
channels.put(name, channel);
return this;
}
/**
* Adds a TopicChannel to the graph.
*
* @param name The name of the channel
* @param <T> The type of the value stored in the channel
* @return This builder for method chaining
*/
public <T> GraphBuilder<I, O> addTopicChannel(String name) {
TopicChannel<T> channel = TopicChannel.<T>create(name);
channels.put(name, channel);
return this;
}
/**
* Adds a custom channel to the graph.
*
* @param name The name of the channel
* @param channel The channel to add
* @return This builder for method chaining
*/
public GraphBuilder<I, O> addChannel(String name, BaseChannel<?, ?, ?> channel) {
channels.put(name, channel);
return this;
}
/**
* Sets the checkpoint saver for the graph.
*
* @param checkpointer The checkpoint saver to use
* @return This builder for method chaining
*/
public GraphBuilder<I, O> setCheckpointer(BaseCheckpointSaver checkpointer) {
this.checkpointer = checkpointer;
return this;
}
/**
* Sets the maximum number of steps for the graph.
*
* @param maxSteps The maximum number of steps
* @return This builder for method chaining
*/
public GraphBuilder<I, O> setMaxSteps(int maxSteps) {
this.maxSteps = maxSteps;
return this;
}
/**
* Sets the retry policy for all nodes in the graph.
* Note: This creates new nodes with the specified retry policy.
*
* @param retryPolicy The retry policy to use
* @return This builder for method chaining
*/
public GraphBuilder<I, O> setRetryPolicy(RetryPolicy retryPolicy) {
// We need to recreate the nodes with the new retry policy
Map<String, PregelNode<I, O>> updatedNodes = new HashMap<>();
for (Map.Entry<String, PregelNode<I, O>> entry : nodes.entrySet()) {
String nodeName = entry.getKey();
PregelNode<I, O> node = entry.getValue();
// Create a new node with the same configuration but different retry policy
PregelNode<I, O> updatedNode = new PregelNode<>(
node.getName(),
node.getAction(),
node.getChannels(),
node.getTriggerChannels(),
node.getWriteEntries(),
retryPolicy
);
updatedNodes.put(nodeName, updatedNode);
}
// Replace all nodes with updated ones
nodes.clear();
nodes.putAll(updatedNodes);
return this;
}
/**
* Configures the node channels to form an implied sequence.
* This creates a chain of nodes where the output of one node is the input of the next.
*
* @param nodeNames The names of the nodes in the sequence
* @param inputChannel The name of the input channel
* @param outputChannel The name of the output channel
* @param intermediateChannel The name of the intermediate channel
* @return This builder for method chaining
*/
public GraphBuilder<I, O> configureSequence(List<String> nodeNames, String inputChannel,
String outputChannel, String intermediateChannel) {
if (nodeNames.size() < 2) {
throw new IllegalArgumentException("Sequence must have at least 2 nodes");
}
// First node reads from input, writes to intermediate
PregelNode<I, O> firstNode = nodes.get(nodeNames.get(0));
PregelNode.Builder<I, O> builder = new PregelNode.Builder<>(firstNode.getName(), firstNode.getAction())
.channels(inputChannel)
.triggerChannels(inputChannel)
.writers(intermediateChannel);
nodes.put(firstNode.getName(), builder.build());
// Middle nodes read from intermediate, write to intermediate
for (int i = 1; i < nodeNames.size() - 1; i++) {
PregelNode<I, O> node = nodes.get(nodeNames.get(i));
builder = new PregelNode.Builder<>(node.getName(), node.getAction())
.channels(intermediateChannel)
.triggerChannels(intermediateChannel)
.writers(intermediateChannel);
nodes.put(node.getName(), builder.build());
}
// Last node reads from intermediate, writes to output
PregelNode<I, O> lastNode = nodes.get(nodeNames.get(nodeNames.size() - 1));
builder = new PregelNode.Builder<>(lastNode.getName(), lastNode.getAction())
.channels(intermediateChannel)
.triggerChannels(intermediateChannel)
.writers(outputChannel);
nodes.put(lastNode.getName(), builder.build());
return this;
}
/**
* Builds a Pregel graph with the configured nodes and channels.
*
* @return A new Pregel instance
*/
public Pregel<I, O> build() {
// If we didn't add the input and output channels explicitly, add them
if (!channels.containsKey("input")) {
addLastValueChannel("input");
}
if (!channels.containsKey("output")) {
addLastValueChannel("output");
}
return new Pregel.Builder<I, O>()
.addNodes(new ArrayList<>(nodes.values()))
.addChannels(channels)
.setCheckpointer(checkpointer)
.setMaxSteps(maxSteps)
.build();
}
}
@@ -0,0 +1,27 @@
package com.langgraph.pregel;
/**
* Represents an error that occurs when a graph exceeds its recursion limit
* during execution.
*/
public class GraphRecursionError extends RuntimeException {
/**
* Creates a new GraphRecursionError with the specified message.
*
* @param message The error message
*/
public GraphRecursionError(String message) {
super(message);
}
/**
* Creates a new GraphRecursionError with the specified message and cause.
*
* @param message The error message
* @param cause The cause of the error
*/
public GraphRecursionError(String message, Throwable cause) {
super(message, cause);
}
}
@@ -0,0 +1,572 @@
package com.langgraph.pregel;
import com.langgraph.channels.BaseChannel;
import com.langgraph.checkpoint.base.BaseCheckpointSaver;
import com.langgraph.pregel.execute.PregelLoop;
import com.langgraph.pregel.execute.SuperstepManager;
import com.langgraph.pregel.registry.ChannelRegistry;
import com.langgraph.pregel.registry.NodeRegistry;
import java.util.*;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
/**
* The type-safe Pregel implementation.
* Orchestrates execution of a computational graph using the Bulk Synchronous Parallel model
* with strict type checking throughout the execution flow.
*
* @param <I> The input type for the overall graph
* @param <O> The output type for the overall graph
*/
public class Pregel<I, O> implements PregelProtocol<I, O> {
private final NodeRegistry nodeRegistry;
private final ChannelRegistry channelRegistry;
private final BaseCheckpointSaver checkpointer;
private final ExecutorService executor;
private final int maxSteps;
private final Set<String> inputChannels;
private final Set<String> outputChannels;
/**
* Get a channel by name (for debugging)
*
* @param name Channel name
* @return Channel with the given name
*/
public BaseChannel<?, ?, ?> getChannel(String name) {
return channelRegistry.get(name);
}
/**
* Create a type-safe Pregel instance with all parameters.
*
* @param nodes Map of node names to nodes
* @param channels Map of channel names to channels
* @param inputChannels Set of input channel names
* @param outputChannels Set of output channel names
* @param checkpointer Optional checkpointer for persisting state
* @param maxSteps Maximum number of steps to execute
*/
public Pregel(
Map<String, PregelNode<?, ?>> nodes,
Map<String, BaseChannel<?, ?, ?>> channels,
Set<String> inputChannels,
Set<String> outputChannels,
BaseCheckpointSaver checkpointer,
int maxSteps) {
// Initialize registries
this.nodeRegistry = new NodeRegistry(nodes);
this.channelRegistry = new ChannelRegistry(channels);
this.inputChannels = inputChannels != null ? inputChannels : new HashSet<>();
this.outputChannels = outputChannels != null ? outputChannels : new HashSet<>();
this.checkpointer = checkpointer;
this.executor = Executors.newWorkStealingPool();
this.maxSteps = maxSteps;
// Validate configuration
validate();
}
/**
* Validate the Pregel configuration.
* Checks that nodes and channels are properly configured and type-compatible.
*
* @throws IllegalStateException If the configuration is invalid
*/
private void validate() {
// Basic validation
nodeRegistry.validate();
// Validate channel references
Set<String> channelNames = channelRegistry.getNames();
nodeRegistry.validateSubscriptions(channelNames);
nodeRegistry.validateWriters(channelNames);
nodeRegistry.validateTriggers(channelNames);
// Type compatibility is ensured by generic type parameters
}
@Override
@SuppressWarnings("unchecked")
public Map<String, O> invoke(Map<String, I> input, Map<String, Object> config) {
// Extract configuration
String threadId = getThreadId(config);
Map<String, Object> context = createContext(threadId, config);
// Create input map with proper type safety
Map<String, Object> inputMap = new HashMap<>();
if (input != null) {
for (Map.Entry<String, I> entry : input.entrySet()) {
inputMap.put(entry.getKey(), entry.getValue());
}
}
// Initialize channels with input
initializeChannels(inputMap);
// Create execution components
SuperstepManager superstepManager = new SuperstepManager(nodeRegistry, channelRegistry);
PregelLoop pregelLoop = new PregelLoop(superstepManager, checkpointer, maxSteps);
// Execute to completion
Map<String, Object> result = pregelLoop.execute(inputMap, context, threadId);
// Filter the result
return filterOutput(result);
}
@Override
public Iterator<Map<String, O>> stream(Map<String, I> input, Map<String, Object> config, StreamMode streamMode) {
// Extract configuration
String threadId = getThreadId(config);
Map<String, Object> context = createContext(threadId, config);
// Create input map with proper type safety
Map<String, Object> inputMap = new HashMap<>();
if (input != null) {
for (Map.Entry<String, I> entry : input.entrySet()) {
inputMap.put(entry.getKey(), entry.getValue());
}
}
// Initialize channels with input
initializeChannels(inputMap);
// Create execution components
SuperstepManager superstepManager = new SuperstepManager(nodeRegistry, channelRegistry);
PregelLoop pregelLoop = new PregelLoop(superstepManager, checkpointer, maxSteps);
// Create iterator for streaming results
return new Iterator<Map<String, O>>() {
private final Queue<Map<String, O>> buffer = new LinkedList<>();
private boolean isDone = false;
@Override
public boolean hasNext() {
if (!buffer.isEmpty()) {
return true;
}
if (isDone) {
return false;
}
// Stream execution and collect results
pregelLoop.stream(
inputMap,
context,
threadId,
streamMode,
result -> {
// Filter the result to match output type
if (result instanceof Map) {
@SuppressWarnings("unchecked")
Map<String, Object> resultMap = (Map<String, Object>) result;
buffer.add(filterOutput(resultMap));
}
return true;
});
isDone = true;
return !buffer.isEmpty();
}
@Override
public Map<String, O> next() {
if (!hasNext()) {
throw new NoSuchElementException();
}
return buffer.poll();
}
};
}
@Override
public Map<String, O> getState(String threadId) {
if (threadId == null) {
throw new IllegalArgumentException("Thread ID is required");
}
if (checkpointer == null) {
return null;
}
// Get latest checkpoint
Optional<String> latestCheckpoint = checkpointer.latest(threadId);
if (!latestCheckpoint.isPresent()) {
return null;
}
// Get checkpoint values
Optional<Map<String, Object>> values = checkpointer.getValues(latestCheckpoint.get());
if (!values.isPresent()) {
return null;
}
// Filter to match output type
return filterOutput(values.get());
}
@Override
public void updateState(String threadId, Map<String, O> state) {
if (threadId == null) {
throw new IllegalArgumentException("Thread ID is required");
}
if (state == null) {
throw new IllegalArgumentException("State cannot be null");
}
// Convert typed state to Object map for backward compatibility
Map<String, Object> stateMap = new HashMap<>();
for (Map.Entry<String, O> entry : state.entrySet()) {
stateMap.put(entry.getKey(), entry.getValue());
}
// Validate state map
validateStateMap(stateMap);
// Update channels with the state
initializeChannels(stateMap);
// Create a checkpoint
if (checkpointer != null) {
checkpointer.checkpoint(threadId, stateMap);
}
}
@Override
public List<Map<String, O>> getStateHistory(String threadId) {
if (threadId == null) {
throw new IllegalArgumentException("Thread ID is required");
}
if (checkpointer == null) {
return Collections.emptyList();
}
List<String> checkpoints = checkpointer.list(threadId);
List<Map<String, O>> history = new ArrayList<>();
for (String checkpointId : checkpoints) {
Optional<Map<String, Object>> values = checkpointer.getValues(checkpointId);
if (values.isPresent()) {
// Filter to match output type
history.add(filterOutput(values.get()));
}
}
return history;
}
/**
* Filter result to include only designated output channels and validate type safety.
*
* @param result Result map to filter
* @return Map with typed output values including only designated channels
*/
@SuppressWarnings("unchecked")
private Map<String, O> filterOutput(Map<String, Object> result) {
if (result == null || result.isEmpty()) {
return Collections.emptyMap();
}
Map<String, O> typedResult = new HashMap<>();
// Filter the result to only include designated output channels
for (Map.Entry<String, Object> entry : result.entrySet()) {
String channelName = entry.getKey();
Object value = entry.getValue();
if (outputChannels.isEmpty() || outputChannels.contains(channelName)) {
// Type safety ensured by generic parameters
typedResult.put(channelName, (O) value);
}
}
return typedResult;
}
/**
* Validates a state map for compatibility with channels.
*
* @param stateMap State map to validate
* @throws IllegalArgumentException if state is invalid
*/
private void validateStateMap(Map<String, Object> stateMap) {
// Validate that the values are compatible with their corresponding channels
for (Map.Entry<String, Object> entry : stateMap.entrySet()) {
String channelName = entry.getKey();
Object value = entry.getValue();
// Type safety ensured by generic parameters
}
}
/**
* Get the thread ID from the configuration.
*
* @param config Configuration
* @return Thread ID
*/
private String getThreadId(Map<String, Object> config) {
if (config == null || !config.containsKey("thread_id")) {
return UUID.randomUUID().toString();
}
return config.get("thread_id").toString();
}
/**
* Create the execution context.
*
* @param threadId Thread ID
* @param config Configuration
* @return Context map
*/
private Map<String, Object> createContext(String threadId, Map<String, Object> config) {
Map<String, Object> context = new HashMap<>();
context.put("thread_id", threadId);
if (config != null) {
context.putAll(config);
}
return context;
}
/**
* Initialize channels with input.
*
* @param input Input map
* @throws IllegalArgumentException if any input value is incompatible with its channel
*/
private void initializeChannels(Map<String, Object> input) {
if (input == null || input.isEmpty()) {
return;
}
// Filter the input to only include designated input channels
if (!inputChannels.isEmpty()) {
Map<String, Object> filteredInput = new HashMap<>();
for (Map.Entry<String, Object> entry : input.entrySet()) {
String channelName = entry.getKey();
Object value = entry.getValue();
if (inputChannels.contains(channelName)) {
// Type safety ensured by generic parameters
filteredInput.put(channelName, value);
}
}
// Update channels with filtered input values
channelRegistry.updateAll(filteredInput);
} else {
// Check all input values for type compatibility
for (Map.Entry<String, Object> entry : input.entrySet()) {
String channelName = entry.getKey();
Object value = entry.getValue();
// Type safety ensured by generic parameters
}
// If no input channels are designated, use all input
channelRegistry.updateAll(input);
}
}
/**
* Get the NodeRegistry.
*
* @return NodeRegistry
*/
public NodeRegistry getNodeRegistry() {
return nodeRegistry;
}
/**
* Get the ChannelRegistry.
*
* @return ChannelRegistry
*/
public ChannelRegistry getChannelRegistry() {
return channelRegistry;
}
/**
* Get the checkpointer.
*
* @return BaseCheckpointSaver
*/
public BaseCheckpointSaver getCheckpointer() {
return checkpointer;
}
/**
* Shutdown the executor service.
*/
public void shutdown() {
executor.shutdown();
}
/**
* Builder for creating type-safe Pregel instances.
*
* @param <I> The input type for the graph
* @param <O> The output type for the graph
*/
public static class Builder<I, O> {
private final Map<String, PregelNode<?, ?>> nodes = new HashMap<>();
private final Map<String, BaseChannel<?, ?, ?>> channels = new HashMap<>();
private Set<String> inputChannels = new HashSet<>();
private Set<String> outputChannels = new HashSet<>();
private BaseCheckpointSaver checkpointer;
private int maxSteps = 100;
/**
* Create a Builder for a type-safe Pregel graph.
*/
public Builder() {
// No parameters needed - type parameters are inferred from usage
}
/**
* Add a node to the graph.
*
* @param node Node to add
* @return This builder
*/
public Builder<I, O> addNode(PregelNode<?, ?> node) {
if (node == null) {
throw new IllegalArgumentException("Node cannot be null");
}
nodes.put(node.getName(), node);
return this;
}
/**
* Add multiple nodes to the graph.
*
* @param nodes Collection of nodes to add
* @return This builder
*/
public Builder<I, O> addNodes(Collection<PregelNode<?, ?>> nodes) {
if (nodes != null) {
for (PregelNode<?, ?> node : nodes) {
addNode(node);
}
}
return this;
}
/**
* Add a channel to the graph.
*
* @param name Channel name
* @param channel Channel to add
* @return This builder
*/
public Builder<I, O> addChannel(String name, BaseChannel<?, ?, ?> channel) {
if (name == null || name.isEmpty()) {
throw new IllegalArgumentException("Channel name cannot be null or empty");
}
if (channel == null) {
throw new IllegalArgumentException("Channel cannot be null");
}
channels.put(name, channel);
return this;
}
/**
* Add multiple channels to the graph.
*
* @param channels Map of channel names to channels
* @return This builder
*/
public Builder<I, O> addChannels(Map<String, BaseChannel<?, ?, ?>> channels) {
if (channels != null) {
this.channels.putAll(channels);
}
return this;
}
/**
* Set input channels for this Pregel graph.
* Input channels will be populated from the input at invocation time.
*
* @param inputChannels Collection of input channel names
* @return This builder
*/
public Builder<I, O> setInputChannels(Collection<String> inputChannels) {
if (inputChannels != null) {
this.inputChannels = new HashSet<>(inputChannels);
}
return this;
}
/**
* Set output channels for this Pregel graph.
* Output channels will be included in the result.
*
* @param outputChannels Collection of output channel names
* @return This builder
*/
public Builder<I, O> setOutputChannels(Collection<String> outputChannels) {
if (outputChannels != null) {
this.outputChannels = new HashSet<>(outputChannels);
}
return this;
}
/**
* Set the checkpointer for persisting state.
*
* @param checkpointer Checkpointer to use
* @return This builder
*/
public Builder<I, O> setCheckpointer(BaseCheckpointSaver checkpointer) {
this.checkpointer = checkpointer;
return this;
}
/**
* Set the maximum number of steps to execute.
*
* @param maxSteps Maximum number of steps
* @return This builder
*/
public Builder<I, O> setMaxSteps(int maxSteps) {
if (maxSteps <= 0) {
throw new IllegalArgumentException("Max steps must be positive");
}
this.maxSteps = maxSteps;
return this;
}
/**
* Build the type-safe Pregel instance.
*
* @return Pregel instance with specified type parameters
*/
public Pregel<I, O> build() {
// If no input/output channels are explicitly set, auto-detect them
if (inputChannels.isEmpty()) {
// Use all channels as input channels by default
inputChannels.addAll(channels.keySet());
}
if (outputChannels.isEmpty()) {
// Use all channels as output channels by default
outputChannels.addAll(channels.keySet());
}
return new Pregel<>(nodes, channels, inputChannels, outputChannels, checkpointer, maxSteps);
}
}
}
@@ -0,0 +1,22 @@
package com.langgraph.pregel;
import java.util.Map;
/**
* Functional interface for actions that can be executed within Pregel.
* This represents the computations performed by nodes in the graph with type-safe input and output.
*
* @param <I> The input type that the node expects
* @param <O> The output type that the node produces
*/
@FunctionalInterface
public interface PregelExecutable<I, O> {
/**
* Execute the action with typed inputs from channels and context information.
*
* @param inputs Map of channel names to their current values with specified input type
* @param context Execution context containing thread ID and other configuration
* @return Map of channel names to values of the specified output type
*/
Map<String, O> execute(Map<String, I> inputs, Map<String, Object> context);
}
@@ -0,0 +1,479 @@
package com.langgraph.pregel;
import com.langgraph.pregel.channel.ChannelWriteEntry;
import com.langgraph.pregel.retry.RetryPolicy;
import java.util.*;
import java.util.stream.Collectors;
/**
* Represents a type-safe actor (node) in the Pregel system.
* A node is a computational unit that reads from input channels,
* executes an action, and writes results to output channels with type safety.
*
* <p>There are two key concepts for how nodes interact with channels:
* <ul>
* <li>Input Channels ({@link #channels}): Channels from which the node reads values.
* When a node executes, it receives values from all its input channels.
* </li>
* <li>Trigger Channels ({@link #triggerChannels}): Special channel(s) that determine when this node
* should execute. A node will execute when any of its trigger channels are updated.
* </li>
* </ul>
* </p>
*
* <p>In Python LangGraph, nodes only run on the first superstep if they have the input channel
* as one of their triggers. In Java LangGraph, we now match this behavior - nodes only run
* in the first superstep if they have appropriate trigger channels defined. For proper
* Python compatibility, it's important to explicitly define input channel as a trigger on
* nodes that should execute first.
* </p>
*
* @param <I> The input type that the node expects
* @param <O> The output type that the node produces
*/
public class PregelNode<I, O> {
private final String name;
private final PregelExecutable<I, O> action;
private final Set<String> channels; // Input channels
private final Set<String> triggerChannels; // Trigger channels
private final List<ChannelWriteEntry> writers;
private final RetryPolicy retryPolicy;
/**
* Create a typed PregelNode with write entries for outputs.
*
* @param name Unique identifier for the node
* @param action Function to execute when the node is triggered
* @param channels Channel names this node reads values from
* @param triggerChannels Channel(s) that determine when this node executes
* @param writeEntries Channel write entries that specify how to write outputs
* @param retryPolicy Strategy for handling execution failures
*/
public PregelNode(
String name,
PregelExecutable<I, O> action,
Collection<String> channels,
Collection<String> triggerChannels,
Collection<ChannelWriteEntry> writeEntries,
RetryPolicy retryPolicy) {
if (name == null || name.isEmpty()) {
throw new IllegalArgumentException("Node name cannot be null or empty");
}
if (action == null) {
throw new IllegalArgumentException("Action cannot be null");
}
this.name = name;
this.action = action;
this.channels = channels != null ? new HashSet<>(channels) : Collections.emptySet();
this.triggerChannels = triggerChannels != null ? new HashSet<>(triggerChannels) : Collections.emptySet();
this.writers = writeEntries != null ? new ArrayList<>(writeEntries) : Collections.emptyList();
this.retryPolicy = retryPolicy;
}
/**
* Get the name of the node.
*
* @return Node name
*/
public String getName() {
return name;
}
/**
* Get the action to execute.
*
* @return Node action
*/
public PregelExecutable<I, O> getAction() {
return action;
}
/**
* Get the input channels this node reads from.
*
* @return Set of channel names (immutable)
*/
public Set<String> getChannels() {
return Collections.unmodifiableSet(channels);
}
/**
* Get the trigger channels for this node.
*
* @return Set of trigger channels (immutable)
*/
public Set<String> getTriggerChannels() {
return Collections.unmodifiableSet(triggerChannels);
}
/**
* Get the write entries for this node.
*
* @return List of channel write entries (immutable)
*/
public List<ChannelWriteEntry> getWriteEntries() {
return Collections.unmodifiableList(writers);
}
/**
* Get the channels this node can write to.
*
* @return Set of channel names (immutable)
*/
public Set<String> getWriters() {
return writers.stream()
.map(ChannelWriteEntry::getChannel)
.collect(Collectors.toSet());
}
/**
* Get the retry policy for this node.
*
* @return Retry policy or null if using default policy
*/
public RetryPolicy getRetryPolicy() {
return retryPolicy;
}
/**
* Check if this node reads from a specific channel.
*
* @param channelName Channel name to check
* @return True if the node reads from the channel
*/
public boolean readsFrom(String channelName) {
return channels.contains(channelName);
}
/**
* Check if this node is triggered by a specific channel.
*
* @param channelName Channel name to check
* @return True if the node is triggered by the channel
*/
public boolean isTriggeredBy(String channelName) {
return triggerChannels.contains(channelName);
}
/**
* Check if this node can write to a specific channel.
*
* @param channelName Channel name to check
* @return True if the node can write to the channel
*/
public boolean canWriteTo(String channelName) {
return writers.stream()
.anyMatch(entry -> entry.getChannel().equals(channelName));
}
/**
* Find a write entry for a specific channel.
*
* @param channelName Channel name to look for
* @return Optional write entry for the channel
*/
public Optional<ChannelWriteEntry> getWriteEntry(String channelName) {
return writers.stream()
.filter(entry -> entry.getChannel().equals(channelName))
.findFirst();
}
/**
* Process node output according to write entries.
* This method preserves type safety by ensuring the output is of the expected type.
*
* @param nodeOutput Output from node execution
* @return Processed output with values transformed as specified by write entries
*/
@SuppressWarnings("unchecked")
public Map<String, O> processOutput(Map<String, O> nodeOutput) {
if (nodeOutput == null || nodeOutput.isEmpty()) {
return Collections.emptyMap();
}
Map<String, O> result = new HashMap<>();
// Process specific channel outputs
for (ChannelWriteEntry entry : writers) {
String channelName = entry.getChannel();
Object value = entry.isPassthrough() ? nodeOutput.get(channelName) : entry.getValue();
// Skip if explicit value is not found and this is a passthrough entry
if (entry.isPassthrough() && !nodeOutput.containsKey(channelName)) {
continue;
}
// Apply mapper if present
if (entry.hasMapper()) {
value = entry.getMapper().apply(value);
}
// Skip null values if configured to do so
if (value == null && entry.isSkipNone()) {
continue;
}
// Type safety is ensured by generic parameters
result.put(channelName, (O) value);
}
// If no write entries are specified, pass through all outputs
if (writers.isEmpty()) {
result.putAll(nodeOutput);
}
return result;
}
/**
* Execute the node's action with type safety for input and output.
* This method ensures type safety throughout the execution flow.
*
* @param inputs Map of input values
* @param context Execution context
* @return Map of typed output values
*/
@SuppressWarnings("unchecked")
public Map<String, O> executeTyped(Map<String, Object> inputs, Map<String, Object> context) {
// Convert inputs to the expected type using compile-time type safety
Map<String, I> typedInputs = new HashMap<>();
for (Map.Entry<String, Object> entry : inputs.entrySet()) {
String channelName = entry.getKey();
Object value = entry.getValue();
// Only include inputs for channels this node reads from
if (!channels.contains(channelName)) {
continue;
}
// Cast value to expected input type
typedInputs.put(channelName, (I) value);
}
// Execute the action with typed inputs
return action.execute(typedInputs, context);
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
PregelNode<?, ?> that = (PregelNode<?, ?>) o;
return Objects.equals(name, that.name);
}
@Override
public int hashCode() {
return Objects.hash(name);
}
@Override
public String toString() {
return "PregelNode{" +
"name='" + name + '\'' +
", channels=" + channels +
", triggerChannels=" + triggerChannels +
", writers=" + writers +
'}';
}
/**
* Builder for creating type-safe PregelNode instances.
*
* @param <I> The input type that the node expects
* @param <O> The output type that the node produces
*/
public static class Builder<I, O> {
private final String name;
private final PregelExecutable<I, O> action;
private Set<String> channels = new HashSet<>();
private Set<String> triggerChannels = new HashSet<>();
private List<ChannelWriteEntry> writers = new ArrayList<>();
private RetryPolicy retryPolicy;
/**
* Create a Builder with the required name and action.
*
* @param name Unique identifier for the node
* @param action Function to execute when the node is triggered
*/
public Builder(String name, PregelExecutable<I, O> action) {
if (name == null || name.isEmpty()) {
throw new IllegalArgumentException("Node name cannot be null or empty");
}
if (action == null) {
throw new IllegalArgumentException("Action cannot be null");
}
this.name = name;
this.action = action;
}
/**
* Add input channels that this node will read from.
*
* @param channelNames Channel names to read from
* @return This builder
*/
public Builder<I, O> channels(Collection<String> channelNames) {
if (channelNames != null) {
for (String channelName : channelNames) {
if (channelName != null && !channelName.isEmpty()) {
channels.add(channelName);
}
}
}
return this;
}
/**
* Add a single input channel that this node will read from.
*
* @param channelName Channel name to read from
* @return This builder
*/
public Builder<I, O> channels(String channelName) {
if (channelName != null && !channelName.isEmpty()) {
channels.add(channelName);
}
return this;
}
/**
* Add trigger channels that determine when this node executes.
*
* @param channelNames Channel names that trigger execution
* @return This builder
*/
public Builder<I, O> triggerChannels(Collection<String> channelNames) {
if (channelNames != null) {
for (String channelName : channelNames) {
if (channelName != null && !channelName.isEmpty()) {
triggerChannels.add(channelName);
}
}
}
return this;
}
/**
* Add a single trigger channel that determines when this node executes.
*
* @param channelName Channel name that triggers execution
* @return This builder
*/
public Builder<I, O> triggerChannels(String channelName) {
if (channelName != null && !channelName.isEmpty()) {
triggerChannels.add(channelName);
}
return this;
}
/**
* Add writers that specify where this node will write its output.
*
* @param entries Collection of ChannelWriteEntry objects
* @return This builder
*/
public Builder<I, O> writers(Collection<ChannelWriteEntry> entries) {
if (entries != null) {
for (ChannelWriteEntry entry : entries) {
if (entry != null) {
writers.add(entry);
}
}
}
return this;
}
/**
* Add a single writer that specifies where this node will write its output.
*
* @param entry ChannelWriteEntry object
* @return This builder
*/
public Builder<I, O> writers(ChannelWriteEntry entry) {
if (entry != null) {
writers.add(entry);
}
return this;
}
/**
* Add a simple writer to the specified channel.
* The node's output value for this channel will be passed through.
*
* @param channelName Channel name this node can write to
* @return This builder
*/
public Builder<I, O> writers(String channelName) {
if (channelName != null && !channelName.isEmpty()) {
writers.add(new ChannelWriteEntry(channelName));
}
return this;
}
/**
* Add multiple simple writers to the specified channels.
* The node's output values for these channels will be passed through.
*
* @param channelNames Channel names this node can write to
* @return This builder
*/
public Builder<I, O> writers(String... channelNames) {
if (channelNames != null) {
for (String name : channelNames) {
writers(name);
}
}
return this;
}
/**
* Add multiple simple writers from a collection of channel names.
* The node's output values for these channels will be passed through.
*
* @param channelNames Collection of channel names this node can write to
* @return This builder
*/
public Builder<I, O> writersFromCollection(Collection<String> channelNames) {
if (channelNames != null) {
for (String name : channelNames) {
writers(name);
}
}
return this;
}
/**
* Set the retry policy.
*
* @param retryPolicy Retry policy for handling failures
* @return This builder
*/
public Builder<I, O> retryPolicy(RetryPolicy retryPolicy) {
this.retryPolicy = retryPolicy;
return this;
}
/**
* Build the type-safe PregelNode.
*
* @return PregelNode instance with specified type parameters
*/
public PregelNode<I, O> build() {
return new PregelNode<>(name, action, channels, triggerChannels, writers, retryPolicy);
}
}
}
@@ -0,0 +1,58 @@
package com.langgraph.pregel;
import java.util.Iterator;
import java.util.List;
import java.util.Map;
/**
* Core interface defining the contract for all type-safe Pregel implementations.
* This protocol provides methods for execution, streaming results, and state management
* with proper generic type parameters.
*
* @param <I> The input type for the graph
* @param <O> The output type for the graph
*/
public interface PregelProtocol<I, O> {
/**
* Invoke the graph with typed input and run to completion.
*
* @param input Input to the graph, as a map of channel names to typed values
* @param config Optional configuration parameters
* @return Output from the graph after execution completes, as a map of channel names to typed values
*/
Map<String, O> invoke(Map<String, I> input, Map<String, Object> config);
/**
* Stream execution results as they are produced, with proper type safety.
*
* @param input Input to the graph, as a map of channel names to typed values
* @param config Optional configuration parameters
* @param streamMode Mode of streaming (VALUES, UPDATES, or DEBUG)
* @return Iterator of execution updates with proper types
*/
Iterator<Map<String, O>> stream(Map<String, I> input, Map<String, Object> config, StreamMode streamMode);
/**
* Get the current state for a thread with proper type safety.
*
* @param threadId Thread ID to get state for
* @return Current state as a map of channel names to typed values
*/
Map<String, O> getState(String threadId);
/**
* Update the state for a thread with type-safe values.
*
* @param threadId Thread ID to update
* @param state New state to set, as a map of channel names to typed values
*/
void updateState(String threadId, Map<String, O> state);
/**
* Get the state history for a thread with proper type safety.
*
* @param threadId Thread ID to get history for
* @return List of state snapshots in chronological order, each as a map of channel names to typed values
*/
List<Map<String, O>> getStateHistory(String threadId);
}
@@ -0,0 +1,21 @@
package com.langgraph.pregel;
/**
* Enum defining the different streaming options for Pregel execution.
*/
public enum StreamMode {
/**
* Stream the complete state after each superstep.
*/
VALUES,
/**
* Stream state deltas after each node execution.
*/
UPDATES,
/**
* Stream comprehensive execution information for debugging.
*/
DEBUG
}
@@ -0,0 +1,250 @@
package com.langgraph.pregel.channel;
import java.util.Objects;
import java.util.function.Function;
/**
* Represents a specification for writing to a channel.
* This defines both the channel to write to and how values should be processed before writing.
*/
public class ChannelWriteEntry {
/**
* Special marker value indicating that the node's output value should be passed through.
*/
public static final Object PASSTHROUGH = new Object() {
@Override
public String toString() {
return "PASSTHROUGH";
}
};
private final String channel;
private final Object value;
private final boolean skipNone;
private final Function<Object, Object> mapper;
/**
* Create a ChannelWriteEntry with all parameters.
*
* @param channel Channel name to write to
* @param value Value to write, or PASSTHROUGH to use the input
* @param skipNone Whether to skip writing if the value is null
* @param mapper Function to transform the value before writing
*/
public ChannelWriteEntry(
String channel,
Object value,
boolean skipNone,
Function<Object, Object> mapper) {
if (channel == null || channel.isEmpty()) {
throw new IllegalArgumentException("Channel name cannot be null or empty");
}
this.channel = channel;
this.value = value;
this.skipNone = skipNone;
this.mapper = mapper;
}
/**
* Create a ChannelWriteEntry with default parameters.
*
* @param channel Channel name to write to
*/
public ChannelWriteEntry(String channel) {
this(channel, PASSTHROUGH, false, null);
}
/**
* Create a ChannelWriteEntry with a specific value.
*
* @param channel Channel name to write to
* @param value Value to write
*/
public ChannelWriteEntry(String channel, Object value) {
this(channel, value, false, null);
}
/**
* Get the channel name.
*
* @return Channel name
*/
public String getChannel() {
return channel;
}
/**
* Get the value to write.
*
* @return Value or PASSTHROUGH
*/
public Object getValue() {
return value;
}
/**
* Check if writing should be skipped for null values.
*
* @return True if null values should be skipped
*/
public boolean isSkipNone() {
return skipNone;
}
/**
* Get the mapper function.
*
* @return Mapper or null if no mapping is required
*/
public Function<Object, Object> getMapper() {
return mapper;
}
/**
* Check if this entry uses a passthrough value.
*
* @return True if the value is PASSTHROUGH
*/
public boolean isPassthrough() {
return PASSTHROUGH.equals(value);
}
/**
* Check if this entry has a mapper.
*
* @return True if a mapper is present
*/
public boolean hasMapper() {
return mapper != null;
}
/**
* Process a value according to this entry's configuration.
*
* @param inputValue Input value (used if this entry is passthrough)
* @return Processed value to write, or null if writing should be skipped
*/
public Object processValue(Object inputValue) {
// Determine the base value (either fixed or passthrough)
Object baseValue = isPassthrough() ? inputValue : value;
// Apply mapper if present
Object processedValue = hasMapper() ? mapper.apply(baseValue) : baseValue;
// Skip null values if configured to do so
if (skipNone && processedValue == null) {
return null;
}
return processedValue;
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
ChannelWriteEntry that = (ChannelWriteEntry) o;
return skipNone == that.skipNone &&
Objects.equals(channel, that.channel) &&
Objects.equals(value, that.value);
}
@Override
public int hashCode() {
return Objects.hash(channel, value, skipNone);
}
@Override
public String toString() {
return "ChannelWriteEntry{" +
"channel='" + channel + '\'' +
", value=" + (isPassthrough() ? "PASSTHROUGH" : value) +
", skipNone=" + skipNone +
", hasMapper=" + (mapper != null) +
'}';
}
/**
* Create a builder for ChannelWriteEntry.
*
* @param channel Channel name
* @return Builder
*/
public static Builder builder(String channel) {
return new Builder(channel);
}
/**
* Builder for ChannelWriteEntry.
*/
public static class Builder {
private final String channel;
private Object value = PASSTHROUGH;
private boolean skipNone = false;
private Function<Object, Object> mapper = null;
/**
* Create a Builder.
*
* @param channel Channel name
*/
public Builder(String channel) {
if (channel == null || channel.isEmpty()) {
throw new IllegalArgumentException("Channel name cannot be null or empty");
}
this.channel = channel;
}
/**
* Set the value.
*
* @param value Value
* @return This builder
*/
public Builder value(Object value) {
this.value = value;
return this;
}
/**
* Set to passthrough mode.
*
* @return This builder
*/
public Builder passthrough() {
this.value = PASSTHROUGH;
return this;
}
/**
* Set whether to skip null values.
*
* @param skipNone Whether to skip null values
* @return This builder
*/
public Builder skipNone(boolean skipNone) {
this.skipNone = skipNone;
return this;
}
/**
* Set the mapper function.
*
* @param mapper Mapper function
* @return This builder
*/
public Builder mapper(Function<Object, Object> mapper) {
this.mapper = mapper;
return this;
}
/**
* Build the ChannelWriteEntry.
*
* @return ChannelWriteEntry
*/
public ChannelWriteEntry build() {
return new ChannelWriteEntry(channel, value, skipNone, mapper);
}
}
}
@@ -0,0 +1,144 @@
package com.langgraph.pregel.channel;
import java.util.Objects;
import java.util.function.Predicate;
/**
* Represents a permission to write to a channel with optional validation.
* This defines both which channel a node can write to and rules for validating the writes.
*/
public class ChannelWritePermission {
private final String channelName;
private final Predicate<Object> validator;
/**
* Create a ChannelWritePermission with a validator.
*
* @param channelName Name of the channel
* @param validator Optional validator for checking values written to the channel
*/
public ChannelWritePermission(String channelName, Predicate<Object> validator) {
if (channelName == null || channelName.isEmpty()) {
throw new IllegalArgumentException("Channel name cannot be null or empty");
}
this.channelName = channelName;
this.validator = validator;
}
/**
* Create a ChannelWritePermission without a validator.
*
* @param channelName Name of the channel
*/
public ChannelWritePermission(String channelName) {
this(channelName, null);
}
/**
* Get the channel name.
*
* @return Channel name
*/
public String getChannelName() {
return channelName;
}
/**
* Get the validator.
*
* @return Validator or null if no validation is required
*/
public Predicate<Object> getValidator() {
return validator;
}
/**
* Check if this permission has a validator.
*
* @return True if a validator is present
*/
public boolean hasValidator() {
return validator != null;
}
/**
* Validate a value.
*
* @param value Value to validate
* @return True if the value is valid or no validator is present
*/
public boolean validate(Object value) {
return validator == null || validator.test(value);
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
ChannelWritePermission that = (ChannelWritePermission) o;
return Objects.equals(channelName, that.channelName);
}
@Override
public int hashCode() {
return Objects.hash(channelName);
}
@Override
public String toString() {
return "ChannelWritePermission{" +
"channelName='" + channelName + '\'' +
", hasValidator=" + (validator != null) +
'}';
}
/**
* Create a builder for ChannelWritePermission.
*
* @param channelName Channel name
* @return Builder
*/
public static Builder builder(String channelName) {
return new Builder(channelName);
}
/**
* Builder for ChannelWritePermission.
*/
public static class Builder {
private final String channelName;
private Predicate<Object> validator;
/**
* Create a Builder.
*
* @param channelName Channel name
*/
public Builder(String channelName) {
if (channelName == null || channelName.isEmpty()) {
throw new IllegalArgumentException("Channel name cannot be null or empty");
}
this.channelName = channelName;
}
/**
* Set the validator.
*
* @param validator Validator
* @return This builder
*/
public Builder validator(Predicate<Object> validator) {
this.validator = validator;
return this;
}
/**
* Build the ChannelWritePermission.
*
* @return ChannelWritePermission
*/
public ChannelWritePermission build() {
return new ChannelWritePermission(channelName, validator);
}
}
}
@@ -0,0 +1,359 @@
package com.langgraph.pregel.execute;
import com.langgraph.checkpoint.base.BaseCheckpointSaver;
import com.langgraph.pregel.GraphRecursionError;
import com.langgraph.pregel.StreamMode;
import com.langgraph.pregel.registry.ChannelRegistry;
import com.langgraph.pregel.state.Checkpoint;
import java.util.*;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.function.Function;
/**
* Main execution loop for Pregel.
* Manages the execution of multiple supersteps until completion or interruption.
*/
public class PregelLoop {
private static final int DEFAULT_MAX_STEPS = 100;
private final SuperstepManager superstepManager;
private final BaseCheckpointSaver checkpointer;
private final int maxSteps;
private final AtomicInteger stepCount;
/**
* Create a PregelLoop.
*
* @param superstepManager Manager for executing supersteps
* @param checkpointer Optional checkpointer for persisting state
* @param maxSteps Maximum number of steps to execute before terminating
*/
public PregelLoop(
SuperstepManager superstepManager,
BaseCheckpointSaver checkpointer,
int maxSteps) {
this.superstepManager = superstepManager;
this.checkpointer = checkpointer;
this.maxSteps = maxSteps > 0 ? maxSteps : DEFAULT_MAX_STEPS;
this.stepCount = new AtomicInteger(0);
}
/**
* Create a PregelLoop with default max steps.
*
* @param superstepManager Manager for executing supersteps
* @param checkpointer Optional checkpointer for persisting state
*/
public PregelLoop(SuperstepManager superstepManager, BaseCheckpointSaver checkpointer) {
this(superstepManager, checkpointer, DEFAULT_MAX_STEPS);
}
/**
* Create a PregelLoop without checkpointing.
*
* @param superstepManager Manager for executing supersteps
* @param maxSteps Maximum number of steps
*/
public PregelLoop(SuperstepManager superstepManager, int maxSteps) {
this(superstepManager, null, maxSteps);
}
/**
* Create a PregelLoop with default configuration.
*
* @param superstepManager Manager for executing supersteps
*/
public PregelLoop(SuperstepManager superstepManager) {
this(superstepManager, null, DEFAULT_MAX_STEPS);
}
/**
* Execute the Pregel loop to completion and return the final state.
*
* @param input Initial input to the loop
* @param context Execution context
* @param threadId Thread ID for checkpointing
* @return Final state after execution
*/
public Map<String, Object> execute(
Map<String, Object> input,
Map<String, Object> context,
String threadId) {
if (input != null && !input.isEmpty()) {
// Initialize with input
initializeWithInput(input);
} else if (threadId != null && checkpointer != null) {
// Try to restore from checkpoint
restoreFromCheckpoint(threadId);
}
// Execute supersteps until completion
Map<String, Object> result = null;
stepCount.set(0);
while (stepCount.get() < maxSteps) {
stepCount.incrementAndGet();
// Execute a single superstep
SuperstepResult stepResult = superstepManager.executeStep(context);
// Capture result
result = stepResult.getState();
// Create checkpoint if configured
if (threadId != null && checkpointer != null) {
createCheckpoint(threadId, result);
}
// Check if we're done
if (!stepResult.hasMoreWork()) {
break;
}
}
// Since we've now exited the main loop, only throw an error if:
// 1. We've reached the max steps limit, AND
// 2. We still have more work to do (which means we didn't finish naturally)
if (stepCount.get() >= maxSteps) {
// Execute a "check" step to see if we still have more work
// This is also a form of final step which may complete the execution
SuperstepResult finalResult = superstepManager.executeStep(context);
result = finalResult.getState(); // update the result with this final step
// Only if this final step shows there's STILL more work after reaching limits,
// we have a genuine recursion issue - otherwise we just completed normally
if (finalResult.hasMoreWork()) {
throw new GraphRecursionError("Maximum iteration steps reached: " + maxSteps);
}
}
return result;
}
/**
* Execute the Pregel loop with streaming of intermediate states.
*
* @param input Initial input to the loop
* @param context Execution context
* @param threadId Thread ID for checkpointing
* @param streamMode Streaming mode
* @param callback Callback for each state update
*/
public void stream(
Map<String, Object> input,
Map<String, Object> context,
String threadId,
StreamMode streamMode,
Function<Map<String, Object>, Boolean> callback) {
if (input != null && !input.isEmpty()) {
// Initialize with input
initializeWithInput(input);
} else if (threadId != null && checkpointer != null) {
// Try to restore from checkpoint
restoreFromCheckpoint(threadId);
}
// Execute supersteps until completion
stepCount.set(0);
boolean continueExecution = true;
while (continueExecution && stepCount.get() < maxSteps) {
stepCount.incrementAndGet();
// Execute a single superstep
SuperstepResult stepResult = superstepManager.executeStep(context);
// Stream result based on mode
Map<String, Object> streamData = formatStreamOutput(stepResult, streamMode);
// Call the callback with the result
if (callback != null) {
continueExecution = callback.apply(streamData);
}
// Create checkpoint if configured
if (threadId != null && checkpointer != null) {
createCheckpoint(threadId, stepResult.getState());
}
// Check if we're done
if (!stepResult.hasMoreWork()) {
break;
}
}
// Only throw if we've both:
// 1. Reached max steps limit AND
// 2. The caller wants to continue (they returned true) AND
// 3. We actually still have more work in the execution engine
if (continueExecution && stepCount.get() >= maxSteps) {
// Execute one final step to see if it completes the execution
SuperstepResult finalResult = superstepManager.executeStep(context);
// If the final step shows we still have work to do after reaching limits
// AND the callback wanted to continue, then we have a genuine recursion issue
if (finalResult.hasMoreWork()) {
throw new GraphRecursionError("Maximum iteration steps reached in streaming: " + maxSteps);
}
// Otherwise we just completed normally on this final step
if (callback != null) {
// Call the callback with the final state
Map<String, Object> streamData = formatStreamOutput(finalResult, streamMode);
callback.apply(streamData); // Ignore the return value as we're done anyway
}
}
}
/**
* Initialize the Pregel loop with input.
*
* @param input Initial input
*/
private void initializeWithInput(Map<String, Object> input) {
if (input == null || input.isEmpty()) {
return;
}
// Update channel values with input values
ChannelRegistry channelRegistry = getChannelRegistry();
boolean anyChannelUpdated = false;
for (Map.Entry<String, Object> entry : input.entrySet()) {
String channelName = entry.getKey();
Object value = entry.getValue();
if (channelRegistry.contains(channelName) && value != null) {
// Update the channel with the input value
boolean updated = channelRegistry.update(channelName, value);
if (updated) {
anyChannelUpdated = true;
}
}
}
// Mark all input channels as updated for the initial superstep
superstepManager.addUpdatedChannels(input.keySet());
}
/**
* Restore state from checkpoint.
*
* @param threadId Thread ID
* @return True if state was restored, false otherwise
*/
private boolean restoreFromCheckpoint(String threadId) {
if (checkpointer == null || threadId == null) {
return false;
}
Optional<String> latestCheckpoint = checkpointer.latest(threadId);
if (!latestCheckpoint.isPresent()) {
return false;
}
Optional<Map<String, Object>> checkpoint = checkpointer.getValues(latestCheckpoint.get());
if (!checkpoint.isPresent()) {
return false;
}
// Restore channel values from checkpoint
ChannelRegistry channelRegistry = getChannelRegistry();
channelRegistry.restoreFromCheckpoint(checkpoint.get());
// Mark all channels as updated for the first superstep
superstepManager.addUpdatedChannels(checkpoint.get().keySet());
return true;
}
/**
* Create a checkpoint.
*
* @param threadId Thread ID
* @param state Current state
*/
private void createCheckpoint(String threadId, Map<String, Object> state) {
if (checkpointer == null || threadId == null) {
return;
}
// Create checkpoint with the current state
// We use threadId as the thread ID and construct a unique checkpoint ID
checkpointer.checkpoint(threadId, new HashMap<>(state));
}
/**
* Format the output for streaming based on the stream mode.
*
* @param result Superstep result
* @param streamMode Stream mode
* @return Formatted output
*/
private Map<String, Object> formatStreamOutput(SuperstepResult result, StreamMode streamMode) {
if (streamMode == null) {
streamMode = StreamMode.VALUES;
}
switch (streamMode) {
case VALUES:
// Return the full state
return result.getState();
case UPDATES:
// Return only the updated channels
Map<String, Object> updates = new HashMap<>();
for (String channelName : result.getUpdatedChannels()) {
if (result.getState().containsKey(channelName)) {
updates.put(channelName, result.getState().get(channelName));
}
}
return updates;
case DEBUG:
// Return detailed debug information
Map<String, Object> debug = new HashMap<>();
debug.put("state", result.getState());
debug.put("updated_channels", result.getUpdatedChannels());
debug.put("step", stepCount.get());
debug.put("has_more_work", result.hasMoreWork());
return debug;
default:
return result.getState();
}
}
/**
* Get the channel registry from the superstep manager.
*
* @return Channel registry
*/
private ChannelRegistry getChannelRegistry() {
// Access the private channelRegistry field from SuperstepManager for now
// In an ideal world, SuperstepManager would expose a getter for this
try {
java.lang.reflect.Field field = SuperstepManager.class.getDeclaredField("channelRegistry");
field.setAccessible(true);
return (ChannelRegistry) field.get(superstepManager);
} catch (Exception e) {
throw new RuntimeException("Failed to access channel registry", e);
}
}
/**
* Get the current step count.
*
* @return Current step count
*/
public int getStepCount() {
return stepCount.get();
}
/**
* Reset the step count.
*/
public void resetStepCount() {
stepCount.set(0);
}
}
@@ -0,0 +1,26 @@
package com.langgraph.pregel.execute;
/**
* Exception thrown when a superstep execution fails.
*/
public class SuperstepExecutionException extends RuntimeException {
/**
* Create a SuperstepExecutionException with a message.
*
* @param message Error message
*/
public SuperstepExecutionException(String message) {
super(message);
}
/**
* Create a SuperstepExecutionException with a message and cause.
*
* @param message Error message
* @param cause Cause of the error
*/
public SuperstepExecutionException(String message, Throwable cause) {
super(message, cause);
}
}
@@ -0,0 +1,234 @@
package com.langgraph.pregel.execute;
import com.langgraph.pregel.PregelNode;
import com.langgraph.pregel.registry.ChannelRegistry;
import com.langgraph.pregel.registry.NodeRegistry;
import com.langgraph.pregel.task.PregelExecutableTask;
import com.langgraph.pregel.task.PregelTask;
import com.langgraph.pregel.task.TaskExecutor;
import com.langgraph.pregel.task.TaskPlanner;
import java.util.*;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ExecutionException;
import java.util.stream.Collectors;
/**
* Manages the execution of a single superstep in the Pregel system.
* A superstep consists of planning, execution, and update phases.
*/
public class SuperstepManager {
private final NodeRegistry nodeRegistry;
private final ChannelRegistry channelRegistry;
private final TaskPlanner taskPlanner;
private final TaskExecutor taskExecutor;
private final Set<String> updatedChannels;
/**
* Create a SuperstepManager.
*
* @param nodeRegistry Node registry
* @param channelRegistry Channel registry
* @param taskPlanner Task planner
* @param taskExecutor Task executor
*/
public SuperstepManager(
NodeRegistry nodeRegistry,
ChannelRegistry channelRegistry,
TaskPlanner taskPlanner,
TaskExecutor taskExecutor) {
this.nodeRegistry = nodeRegistry;
this.channelRegistry = channelRegistry;
this.taskPlanner = taskPlanner;
this.taskExecutor = taskExecutor;
this.updatedChannels = new HashSet<>();
}
/**
* Create a SuperstepManager with default planner and executor.
*
* @param nodeRegistry Node registry
* @param channelRegistry Channel registry
*/
public SuperstepManager(NodeRegistry nodeRegistry, ChannelRegistry channelRegistry) {
this(
nodeRegistry,
channelRegistry,
new TaskPlanner(nodeRegistry.getAll()),
new TaskExecutor()
);
}
/**
* Execute a single superstep.
*
* @param context Execution context
* @return SuperstepResult containing the result of the superstep
*/
public SuperstepResult executeStep(Map<String, Object> context) {
// Plan phase: Determine nodes to execute based on channel updates
// For Python compatibility, this will now return tasks even if no channels
// have been updated, ensuring nodes run with uninitialized channels
List<PregelTask> tasks = taskPlanner.planAndPrioritize(updatedChannels);
if (tasks.isEmpty()) {
// No tasks to execute, superstep is complete
return new SuperstepResult(false, Collections.emptySet(), channelRegistry.collectValues());
}
// Clear updated channels for this superstep
updatedChannels.clear();
// Execute phase: Run all tasks and collect results
List<CompletableFuture<Map<String, Object>>> futures = new ArrayList<>();
Map<PregelTask, CompletableFuture<Map<String, Object>>> taskFutures = new HashMap<>();
for (PregelTask task : tasks) {
PregelNode node = nodeRegistry.get(task.getNode());
// Prepare inputs for this task
Map<String, Object> inputs = new HashMap<>();
Set<String> nodeChannels = node.getChannels();
for (String channelName : nodeChannels) {
if (channelRegistry.contains(channelName)) {
inputs.put(channelName, channelRegistry.get(channelName).getValue());
}
}
// Add trigger channel values if present
Set<String> triggerChannels = node.getTriggerChannels();
for (String triggerChannel : triggerChannels) {
if (channelRegistry.contains(triggerChannel)) {
inputs.put(triggerChannel, channelRegistry.get(triggerChannel).getValue());
}
}
// Create executable task
PregelExecutableTask executableTask = new PregelExecutableTask(task, inputs, context);
// Execute task asynchronously
CompletableFuture<Map<String, Object>> future = taskExecutor.executeAsync(node, executableTask);
futures.add(future);
taskFutures.put(task, future);
}
// Wait for all tasks to complete
CompletableFuture<Void> allFutures = CompletableFuture.allOf(
futures.toArray(new CompletableFuture[0]));
try {
// Block until all tasks complete
allFutures.join();
// Collect results
Map<String, Set<Object>> allUpdates = new ConcurrentHashMap<>();
for (Map.Entry<PregelTask, CompletableFuture<Map<String, Object>>> entry : taskFutures.entrySet()) {
PregelTask task = entry.getKey();
PregelNode node = nodeRegistry.get(task.getNode());
Map<String, Object> rawResult = entry.getValue().get();
if (rawResult != null) {
// Process the output according to write entries
Map<String, Object> result = node.processOutput(rawResult);
// Record updates for each channel
for (Map.Entry<String, Object> update : result.entrySet()) {
String channelName = update.getKey();
Object value = update.getValue();
// Skip null values
if (value == null) {
continue;
}
// Group updates by channel name
allUpdates.computeIfAbsent(channelName, k -> ConcurrentHashMap.newKeySet())
.add(value);
}
}
}
// Update phase: Apply updates to channels
Set<String> updated = new HashSet<>();
for (Map.Entry<String, Set<Object>> entry : allUpdates.entrySet()) {
String channelName = entry.getKey();
Set<Object> values = entry.getValue();
if (values.size() == 1) {
// Single update for this channel
Object value = values.iterator().next();
if (channelRegistry.update(channelName, value)) {
updated.add(channelName);
}
} else if (values.size() > 1) {
// Multiple updates for this channel
// For a TopicChannel, we should add all values
if (channelRegistry.get(channelName) instanceof com.langgraph.channels.TopicChannel) {
// Update each value in the TopicChannel
boolean anyUpdated = false;
for (Object value : values) {
if (channelRegistry.update(channelName, value)) {
anyUpdated = true;
}
}
if (anyUpdated) {
updated.add(channelName);
}
} else {
// For other channels, resolve conflicts by using the last value
Object lastValue = values.stream().reduce((a, b) -> b).orElse(null);
if (lastValue != null && channelRegistry.update(channelName, lastValue)) {
updated.add(channelName);
}
}
}
}
// Update our tracking of updated channels for the next superstep
updatedChannels.addAll(updated);
// Return superstep result
return new SuperstepResult(
!updated.isEmpty(),
updated,
channelRegistry.collectValues()
);
} catch (ExecutionException e) {
// Task execution failed
Throwable cause = e.getCause();
throw new SuperstepExecutionException("Superstep execution failed", cause);
} catch (Exception e) {
throw new SuperstepExecutionException("Superstep execution failed", e);
}
}
/**
* Get the updated channels from the previous superstep.
*
* @return Set of updated channel names
*/
public Set<String> getUpdatedChannels() {
return Collections.unmodifiableSet(updatedChannels);
}
/**
* Add channels to the set of updated channels.
*
* @param channelNames Channel names to add
*/
public void addUpdatedChannels(Collection<String> channelNames) {
if (channelNames != null) {
updatedChannels.addAll(channelNames);
}
}
/**
* Clear the set of updated channels.
*/
public void clearUpdatedChannels() {
updatedChannels.clear();
}
}
@@ -0,0 +1,87 @@
package com.langgraph.pregel.execute;
import java.util.Collections;
import java.util.Map;
import java.util.Set;
/**
* Represents the result of a single superstep execution.
* Contains information about channel updates and the current state.
*/
public class SuperstepResult {
private final boolean hasMoreWork;
private final Set<String> updatedChannels;
private final Map<String, Object> state;
/**
* Create a SuperstepResult.
*
* @param hasMoreWork True if there is more work to do (channels were updated)
* @param updatedChannels Set of channel names that were updated
* @param state Current state after the superstep
*/
public SuperstepResult(boolean hasMoreWork, Set<String> updatedChannels, Map<String, Object> state) {
this.hasMoreWork = hasMoreWork;
this.updatedChannels = updatedChannels != null
? Collections.unmodifiableSet(updatedChannels)
: Collections.emptySet();
this.state = state != null
? Collections.unmodifiableMap(state)
: Collections.emptyMap();
}
/**
* Check if there is more work to do.
*
* @return True if there is more work to do
*/
public boolean hasMoreWork() {
return hasMoreWork;
}
/**
* Get the channels that were updated in this superstep.
*
* @return Unmodifiable set of updated channel names
*/
public Set<String> getUpdatedChannels() {
return updatedChannels;
}
/**
* Get the current state after the superstep.
*
* @return Unmodifiable map of the current state
*/
public Map<String, Object> getState() {
return state;
}
/**
* Check if a specific channel was updated.
*
* @param channelName Channel name to check
* @return True if the channel was updated
*/
public boolean wasChannelUpdated(String channelName) {
return updatedChannels.contains(channelName);
}
/**
* Get the number of updated channels.
*
* @return Number of updated channels
*/
public int getUpdateCount() {
return updatedChannels.size();
}
@Override
public String toString() {
return "SuperstepResult{" +
"hasMoreWork=" + hasMoreWork +
", updatedChannelCount=" + updatedChannels.size() +
", stateSize=" + state.size() +
'}';
}
}
@@ -0,0 +1,291 @@
package com.langgraph.pregel.registry;
import com.langgraph.channels.BaseChannel;
import com.langgraph.channels.EmptyChannelException;
import java.util.*;
import java.util.stream.Collectors;
/**
* Registry for managing a collection of channels.
* Provides methods for registration, validation, and channel lookup.
*/
public class ChannelRegistry {
private final Map<String, BaseChannel<?, ?, ?>> channels;
/**
* Create an empty ChannelRegistry.
*/
public ChannelRegistry() {
this.channels = new HashMap<>();
}
/**
* Create a ChannelRegistry with initial channels.
*
* @param channels Map of channel names to channels
*/
public ChannelRegistry(Map<String, BaseChannel<?, ?, ?>> channels) {
this.channels = new HashMap<>();
if (channels != null) {
channels.forEach(this::register);
}
}
/**
* Register a channel.
*
* @param name Channel name
* @param channel Channel to register
* @return This registry
* @throws IllegalArgumentException If a channel with the same name is already registered
*/
public ChannelRegistry register(String name, BaseChannel<?, ?, ?> channel) {
if (name == null || name.isEmpty()) {
throw new IllegalArgumentException("Channel name cannot be null or empty");
}
if (channel == null) {
throw new IllegalArgumentException("Channel cannot be null");
}
if (channels.containsKey(name)) {
throw new IllegalArgumentException("Channel with name '" + name + "' is already registered");
}
channels.put(name, channel);
return this;
}
/**
* Register multiple channels.
*
* @param channelsToRegister Map of channel names to channels
* @return This registry
* @throws IllegalArgumentException If a channel with the same name is already registered
*/
public ChannelRegistry registerAll(Map<String, BaseChannel<?, ?, ?>> channelsToRegister) {
if (channelsToRegister != null) {
channelsToRegister.forEach(this::register);
}
return this;
}
/**
* Get a channel by name.
*
* @param name Name of the channel to get
* @return Channel with the given name
* @throws NoSuchElementException If no channel with the given name is registered
*/
public BaseChannel<?, ?, ?> get(String name) {
BaseChannel<?, ?, ?> channel = channels.get(name);
if (channel == null) {
throw new NoSuchElementException("No channel registered with name '" + name + "'");
}
return channel;
}
/**
* Check if a channel with the given name is registered.
*
* @param name Name to check
* @return True if a channel with the given name is registered
*/
public boolean contains(String name) {
return channels.containsKey(name);
}
/**
* Remove a channel by name.
*
* @param name Name of the channel to remove
* @return This registry
*/
public ChannelRegistry remove(String name) {
channels.remove(name);
return this;
}
/**
* Get all registered channels.
*
* @return Unmodifiable map of channel names to channels
*/
public Map<String, BaseChannel<?, ?, ?>> getAll() {
return Collections.unmodifiableMap(channels);
}
/**
* Get the number of registered channels.
*
* @return Number of registered channels
*/
public int size() {
return channels.size();
}
/**
* Get the names of all registered channels.
*
* @return Set of channel names
*/
public Set<String> getNames() {
return Collections.unmodifiableSet(channels.keySet());
}
/**
* Update a channel with a value.
* First tries to use the updateSingleValue method if supported by the channel,
* otherwise falls back to wrapping the value in a singleton list.
*
* @param name Channel name
* @param value Value to update the channel with
* @return True if the channel was updated, false otherwise
* @throws NoSuchElementException If no channel with the given name is registered
*/
public boolean update(String name, Object value) {
BaseChannel<?, ?, ?> channel = get(name);
// Since we don't know the exact type at compile time, we have to use an unchecked cast
// This is safe because the channel will validate the type at runtime
@SuppressWarnings("unchecked")
BaseChannel<Object, Object, Object> typedChannel = (BaseChannel<Object, Object, Object>) channel;
// First try to use the updateSingleValue method
boolean updated = typedChannel.updateSingleValue(value);
// Fall back to using the update method with a singleton list if updateSingleValue didn't work
if (!updated) {
updated = typedChannel.update(Collections.singletonList(value));
}
return updated;
}
/**
* Update multiple channels.
*
* @param updates Map of channel names to values
* @return Set of channel names that were updated
*/
public Set<String> updateAll(Map<String, Object> updates) {
if (updates == null || updates.isEmpty()) {
return Collections.emptySet();
}
Set<String> updatedChannels = new HashSet<>();
for (Map.Entry<String, Object> entry : updates.entrySet()) {
String name = entry.getKey();
Object value = entry.getValue();
if (contains(name) && update(name, value)) {
updatedChannels.add(name);
}
}
return updatedChannels;
}
/**
* Collect values from all channels.
*
* @return Map of channel names to their current values
*/
public Map<String, Object> collectValues() {
Map<String, Object> values = new HashMap<>();
for (Map.Entry<String, BaseChannel<?, ?, ?>> entry : channels.entrySet()) {
String name = entry.getKey();
BaseChannel<?, ?, ?> channel = entry.getValue();
// Get value, will return null for uninitialized channels (Python compatibility)
Object value = channel.getValue();
// Always include the channel in the output, even if value is null
// This ensures Python compatibility where channels are always present
values.put(name, value);
}
return values;
}
/**
* Capture checkpoint data from all channels.
*
* @return Map of channel names to their checkpoint data
*/
public Map<String, Object> checkpoint() {
Map<String, Object> checkpointData = new HashMap<>();
for (Map.Entry<String, BaseChannel<?, ?, ?>> entry : channels.entrySet()) {
String name = entry.getKey();
BaseChannel<?, ?, ?> channel = entry.getValue();
try {
Object data = channel.checkpoint();
// Always include the channel, even if data is null
checkpointData.put(name, data);
} catch (EmptyChannelException e) {
// Include null value for uninitialized channels for Python compatibility
checkpointData.put(name, null);
}
}
return checkpointData;
}
/**
* Restore channels from checkpoint data.
*
* @param checkpointData Map of channel names to checkpoint data
*/
public void restoreFromCheckpoint(Map<String, Object> checkpointData) {
if (checkpointData == null || checkpointData.isEmpty()) {
return;
}
for (Map.Entry<String, Object> entry : checkpointData.entrySet()) {
String name = entry.getKey();
Object data = entry.getValue();
if (contains(name)) {
BaseChannel<?, ?, ?> channel = channels.get(name);
// Since we don't know the exact type at compile time, we have to use an unchecked cast
// This is safe because the channel will validate the type at runtime
@SuppressWarnings("unchecked")
BaseChannel<Object, Object, Object> typedChannel =
(BaseChannel<Object, Object, Object>) channel;
typedChannel.fromCheckpoint(data);
}
}
}
/**
* Reset all channels, clearing any update flags.
*/
public void resetUpdated() {
for (BaseChannel<?, ?, ?> channel : channels.values()) {
channel.resetUpdated();
}
}
/**
* Get a subset of this registry with channels that match the given names.
*
* @param channelNames Names of channels to include
* @return New registry with only the specified channels
*/
public ChannelRegistry subset(Collection<String> channelNames) {
ChannelRegistry subset = new ChannelRegistry();
if (channelNames != null) {
for (String name : channelNames) {
if (contains(name)) {
subset.register(name, get(name));
}
}
}
return subset;
}
}
@@ -0,0 +1,258 @@
package com.langgraph.pregel.registry;
import com.langgraph.pregel.PregelNode;
import java.util.*;
import java.util.function.Function;
import java.util.stream.Collectors;
/**
* Registry for managing a collection of nodes.
* Provides methods for registration, validation, and node lookup.
*/
public class NodeRegistry {
private final Map<String, PregelNode<?, ?>> nodes;
/**
* Create an empty NodeRegistry.
*/
public NodeRegistry() {
this.nodes = new HashMap<>();
}
/**
* Create a NodeRegistry with initial nodes.
*
* @param nodes Collection of nodes to register
*/
public NodeRegistry(Collection<PregelNode<?, ?>> nodes) {
this.nodes = new HashMap<>();
if (nodes != null) {
nodes.forEach(this::register);
}
}
/**
* Create a NodeRegistry with initial nodes.
*
* @param nodes Map of node names to nodes
*/
public NodeRegistry(Map<String, PregelNode<?, ?>> nodes) {
this.nodes = new HashMap<>();
if (nodes != null) {
nodes.forEach((name, node) -> {
if (!name.equals(node.getName())) {
throw new IllegalArgumentException(
"Node name mismatch: key '" + name + "' != node name '" + node.getName() + "'");
}
register(node);
});
}
}
/**
* Register a node.
*
* @param node Node to register
* @return This registry
* @throws IllegalArgumentException If a node with the same name is already registered
*/
public NodeRegistry register(PregelNode<?, ?> node) {
if (node == null) {
throw new IllegalArgumentException("Node cannot be null");
}
String name = node.getName();
if (nodes.containsKey(name)) {
throw new IllegalArgumentException("Node with name '" + name + "' is already registered");
}
nodes.put(name, node);
return this;
}
/**
* Register multiple nodes.
*
* @param nodesToRegister Collection of nodes to register
* @return This registry
* @throws IllegalArgumentException If a node with the same name is already registered
*/
public NodeRegistry registerAll(Collection<PregelNode<?, ?>> nodesToRegister) {
if (nodesToRegister != null) {
nodesToRegister.forEach(this::register);
}
return this;
}
/**
* Get a node by name.
*
* @param name Name of the node to get
* @return Node with the given name
* @throws NoSuchElementException If no node with the given name is registered
*/
public PregelNode<?, ?> get(String name) {
PregelNode<?, ?> node = nodes.get(name);
if (node == null) {
throw new NoSuchElementException("No node registered with name '" + name + "'");
}
return node;
}
/**
* Check if a node with the given name is registered.
*
* @param name Name to check
* @return True if a node with the given name is registered
*/
public boolean contains(String name) {
return nodes.containsKey(name);
}
/**
* Remove a node by name.
*
* @param name Name of the node to remove
* @return This registry
*/
public NodeRegistry remove(String name) {
nodes.remove(name);
return this;
}
/**
* Get all registered nodes.
*
* @return Unmodifiable map of node names to nodes
*/
public Map<String, PregelNode<?, ?>> getAll() {
return Collections.unmodifiableMap(nodes);
}
/**
* Get all registered nodes as a collection.
*
* @return Collection of all registered nodes
*/
public Collection<PregelNode<?, ?>> getNodes() {
return Collections.unmodifiableCollection(nodes.values());
}
/**
* Get the number of registered nodes.
*
* @return Number of registered nodes
*/
public int size() {
return nodes.size();
}
/**
* Get all nodes that read from the given channel.
*
* @param channelName Channel name
* @return Set of nodes that read from the channel
*/
public Set<PregelNode<?, ?>> getSubscribers(String channelName) {
return nodes.values().stream()
.filter(node -> node.readsFrom(channelName))
.collect(Collectors.toSet());
}
/**
* Get all nodes that are triggered by the given channel.
*
* @param triggerName Trigger name
* @return Set of nodes that are triggered by the channel
*/
public Set<PregelNode<?, ?>> getTriggered(String triggerName) {
return nodes.values().stream()
.filter(node -> node.isTriggeredBy(triggerName))
.collect(Collectors.toSet());
}
/**
* Get all nodes that can write to the given channel.
*
* @param channelName Channel name
* @return Set of nodes that can write to the channel
*/
public Set<PregelNode<?, ?>> getWriters(String channelName) {
return nodes.values().stream()
.filter(node -> node.canWriteTo(channelName))
.collect(Collectors.toSet());
}
/**
* Validate the registry.
* Checks that all nodes have valid configurations.
*
* @throws IllegalStateException If the registry is invalid
*/
public void validate() {
// Validate that each node has a unique name
Set<String> nodeNames = new HashSet<>();
for (PregelNode<?, ?> node : nodes.values()) {
String name = node.getName();
if (nodeNames.contains(name)) {
throw new IllegalStateException("Duplicate node name: " + name);
}
nodeNames.add(name);
}
}
/**
* Validate that nodes only read from existing channels.
*
* @param channelNames Set of valid channel names
* @throws IllegalStateException If a node reads from a non-existent channel
*/
public void validateSubscriptions(Set<String> channelNames) {
for (PregelNode<?, ?> node : nodes.values()) {
Set<String> channels = node.getChannels();
for (String channelName : channels) {
if (!channelNames.contains(channelName)) {
throw new IllegalStateException(
"Node '" + node.getName() + "' reads from non-existent channel '" + channelName + "'");
}
}
}
}
/**
* Validate that nodes only write to existing channels.
*
* @param channelNames Set of valid channel names
* @throws IllegalStateException If a node writes to a non-existent channel
*/
public void validateWriters(Set<String> channelNames) {
for (PregelNode<?, ?> node : nodes.values()) {
Set<String> writers = node.getWriters();
for (String channelName : writers) {
if (!channelNames.contains(channelName)) {
throw new IllegalStateException(
"Node '" + node.getName() + "' writes to non-existent channel '" + channelName + "'");
}
}
}
}
/**
* Validate that nodes only use existing trigger channels.
*
* @param channelNames Set of valid channel names
* @throws IllegalStateException If a node uses a non-existent trigger channel
*/
public void validateTriggers(Set<String> channelNames) {
for (PregelNode<?, ?> node : nodes.values()) {
Set<String> triggers = node.getTriggerChannels();
for (String triggerChannel : triggers) {
if (!channelNames.contains(triggerChannel)) {
throw new IllegalStateException(
"Node '" + node.getName() + "' has non-existent trigger channel '" + triggerChannel + "'");
}
}
}
}
}
@@ -0,0 +1,196 @@
package com.langgraph.pregel.retry;
import java.time.Duration;
import java.util.function.Predicate;
/**
* Factory class for creating common retry policies.
*/
public final class RetryPolicies {
private RetryPolicies() {
// Prevent instantiation
}
/**
* Create a retry policy that never retries.
*
* @return Retry policy
*/
public static RetryPolicy noRetry() {
return RetryPolicy.noRetry();
}
/**
* Create a simple retry policy with a maximum number of attempts.
*
* @param maxAttempts Maximum number of attempts
* @return Retry policy
*/
public static RetryPolicy maxAttempts(int maxAttempts) {
return RetryPolicy.maxAttempts(maxAttempts);
}
/**
* Create a retry policy that always retries with a constant backoff.
*
* @param backoff Backoff duration between retries
* @return Retry policy
*/
public static RetryPolicy constantBackoff(Duration backoff) {
return RetryPolicy.constantBackoff(backoff);
}
/**
* Create a retry policy with exponential backoff.
*
* @param initialBackoff Initial backoff duration
* @param maxAttempts Maximum number of attempts
* @param maxBackoff Maximum backoff duration
* @return Retry policy
*/
public static RetryPolicy exponentialBackoff(Duration initialBackoff, int maxAttempts, Duration maxBackoff) {
return RetryPolicy.exponentialBackoff(initialBackoff, maxAttempts, maxBackoff);
}
/**
* Create a retry policy with exponential backoff and jitter.
*
* @param initialBackoff Initial backoff duration
* @param maxAttempts Maximum number of attempts
* @param maxBackoff Maximum backoff duration
* @param jitterFactor Jitter factor (0.0 to 1.0, where 0.0 means no jitter)
* @return Retry policy
*/
public static RetryPolicy exponentialBackoffWithJitter(Duration initialBackoff, int maxAttempts,
Duration maxBackoff, double jitterFactor) {
return RetryPolicy.exponentialBackoffWithJitter(initialBackoff, maxAttempts, maxBackoff, jitterFactor);
}
/**
* Create a retry policy that filters exceptions.
*
* @param basePolicy Base retry policy to delegate to
* @param filter Predicate to determine which exceptions should be retried
* @return Retry policy
*/
public static RetryPolicy withExceptionFilter(RetryPolicy basePolicy, Predicate<Throwable> filter) {
return RetryPolicy.withExceptionFilter(basePolicy, filter);
}
/**
* Create a retry policy that handles specific exception types.
*
* @param basePolicy Base retry policy to delegate to
* @param exceptionClass Exception class to retry
* @return Retry policy
*/
public static <T extends Throwable> RetryPolicy onException(RetryPolicy basePolicy, Class<T> exceptionClass) {
return withExceptionFilter(basePolicy, exceptionClass::isInstance);
}
/**
* Create a builder for RetryPolicy.
*
* @return Builder
*/
public static Builder builder() {
return new Builder();
}
/**
* Builder for RetryPolicy.
*/
public static class Builder {
private int maxAttempts = 3; // Default value
private Duration initialBackoff = Duration.ZERO;
private Duration maxBackoff = Duration.ofSeconds(1);
private double jitterFactor = 0.0;
private Predicate<Throwable> exceptionFilter = throwable -> true;
/**
* Set the maximum number of attempts.
*
* @param maxAttempts Maximum number of attempts
* @return This builder
*/
public Builder maxAttempts(int maxAttempts) {
this.maxAttempts = maxAttempts;
return this;
}
/**
* Set the initial backoff duration.
*
* @param initialBackoff Initial backoff duration
* @return This builder
*/
public Builder initialBackoff(Duration initialBackoff) {
this.initialBackoff = initialBackoff;
return this;
}
/**
* Set the maximum backoff duration.
*
* @param maxBackoff Maximum backoff duration
* @return This builder
*/
public Builder maxBackoff(Duration maxBackoff) {
this.maxBackoff = maxBackoff;
return this;
}
/**
* Set the jitter factor.
*
* @param jitterFactor Jitter factor (0.0 to 1.0, where 0.0 means no jitter)
* @return This builder
*/
public Builder jitterFactor(double jitterFactor) {
this.jitterFactor = jitterFactor;
return this;
}
/**
* Set the exception filter.
*
* @param exceptionFilter Predicate to determine which exceptions should be retried
* @return This builder
*/
public Builder exceptionFilter(Predicate<Throwable> exceptionFilter) {
this.exceptionFilter = exceptionFilter;
return this;
}
/**
* Build the RetryPolicy.
*
* @return RetryPolicy
*/
public RetryPolicy build() {
RetryPolicy basePolicy;
if (initialBackoff.equals(Duration.ZERO)) {
basePolicy = RetryPolicy.maxAttempts(maxAttempts);
} else if (jitterFactor > 0) {
basePolicy = RetryPolicy.exponentialBackoffWithJitter(
initialBackoff, maxAttempts, maxBackoff, jitterFactor);
} else {
basePolicy = RetryPolicy.exponentialBackoff(
initialBackoff, maxAttempts, maxBackoff);
}
if (exceptionFilter != null) {
// Create a predicate that always returns true
Predicate<Throwable> alwaysTrue = t -> true;
// If the filter is different from the always-true predicate, apply it
if (!exceptionFilter.equals(alwaysTrue)) {
return RetryPolicy.withExceptionFilter(basePolicy, exceptionFilter);
}
}
return basePolicy;
}
}
}
@@ -0,0 +1,187 @@
package com.langgraph.pregel.retry;
import java.time.Duration;
import java.util.concurrent.ThreadLocalRandom;
import java.util.function.Predicate;
/**
* Interface for handling execution failures by determining if and how to retry failed tasks.
*/
public interface RetryPolicy {
/**
* Decide how to handle a failed execution.
*
* @param attempt Current attempt number (1-based)
* @param error Error that occurred
* @return Retry decision with backoff information
*/
RetryDecision shouldRetry(int attempt, Throwable error);
/**
* Create a builder for RetryPolicy.
*
* @return Builder instance for creating RetryPolicy
*/
static RetryPolicies.Builder builder() {
return RetryPolicies.builder();
}
/**
* Class representing a retry decision.
*/
class RetryDecision {
private final boolean shouldRetry;
private final Duration backoff;
private RetryDecision(boolean shouldRetry, Duration backoff) {
this.shouldRetry = shouldRetry;
this.backoff = backoff;
}
/**
* Check if the task should be retried.
*
* @return true if the task should be retried, false otherwise
*/
public boolean shouldRetry() {
return shouldRetry;
}
/**
* Get the backoff duration before the next retry.
*
* @return Duration to wait before the next retry
*/
public Duration getBackoff() {
return backoff;
}
/**
* Create a decision to retry after the specified backoff.
*
* @param backoff Duration to wait before the next retry
* @return Retry decision
*/
public static RetryDecision retry(Duration backoff) {
return new RetryDecision(true, backoff);
}
/**
* Create a decision to not retry.
*
* @return Retry decision
*/
public static RetryDecision fail() {
return new RetryDecision(false, Duration.ZERO);
}
}
/**
* Create a simple retry policy with a maximum number of attempts.
*
* @param maxAttempts Maximum number of attempts
* @return Retry policy
*/
static RetryPolicy maxAttempts(int maxAttempts) {
return (attempt, error) ->
attempt < maxAttempts ? RetryDecision.retry(Duration.ZERO) : RetryDecision.fail();
}
/**
* Create a retry policy that never retries.
*
* @return Retry policy
*/
static RetryPolicy noRetry() {
return (attempt, error) -> RetryDecision.fail();
}
/**
* Create a retry policy that always retries with a constant backoff.
*
* @param backoff Backoff duration between retries
* @return Retry policy
*/
static RetryPolicy constantBackoff(Duration backoff) {
return (attempt, error) -> RetryDecision.retry(backoff);
}
/**
* Create a retry policy with exponential backoff.
*
* @param initialBackoff Initial backoff duration
* @param maxAttempts Maximum number of attempts
* @param maxBackoff Maximum backoff duration
* @return Retry policy
*/
static RetryPolicy exponentialBackoff(Duration initialBackoff, int maxAttempts, Duration maxBackoff) {
return (attempt, error) -> {
if (attempt >= maxAttempts) {
return RetryDecision.fail();
}
long initialBackoffMillis = initialBackoff.toMillis();
long maxBackoffMillis = maxBackoff.toMillis();
// Calculate exponential backoff: initialBackoff * 2^(attempt-1)
long backoffMillis = initialBackoffMillis * (1L << (attempt - 1));
// Ensure backoff doesn't exceed maxBackoff
backoffMillis = Math.min(backoffMillis, maxBackoffMillis);
return RetryDecision.retry(Duration.ofMillis(backoffMillis));
};
}
/**
* Create a retry policy with exponential backoff and jitter.
*
* @param initialBackoff Initial backoff duration
* @param maxAttempts Maximum number of attempts
* @param maxBackoff Maximum backoff duration
* @param jitterFactor Jitter factor (0.0 to 1.0, where 0.0 means no jitter)
* @return Retry policy
*/
static RetryPolicy exponentialBackoffWithJitter(Duration initialBackoff, int maxAttempts,
Duration maxBackoff, double jitterFactor) {
return (attempt, error) -> {
if (attempt >= maxAttempts) {
return RetryDecision.fail();
}
long initialBackoffMillis = initialBackoff.toMillis();
long maxBackoffMillis = maxBackoff.toMillis();
// Calculate exponential backoff: initialBackoff * 2^(attempt-1)
long backoffMillis = initialBackoffMillis * (1L << (attempt - 1));
// Ensure backoff doesn't exceed maxBackoff
backoffMillis = Math.min(backoffMillis, maxBackoffMillis);
if (jitterFactor > 0) {
// Apply jitter: backoff * (1 - jitterFactor + random * 2 * jitterFactor)
double jitter = 1.0 - jitterFactor + ThreadLocalRandom.current().nextDouble() * 2 * jitterFactor;
backoffMillis = (long) (backoffMillis * jitter);
}
return RetryDecision.retry(Duration.ofMillis(backoffMillis));
};
}
/**
* Create a retry policy that filters exceptions.
*
* @param basePolicy Base retry policy to delegate to
* @param filter Predicate to determine which exceptions should be retried
* @return Retry policy
*/
static RetryPolicy withExceptionFilter(RetryPolicy basePolicy, Predicate<Throwable> filter) {
return (attempt, error) -> {
if (filter.test(error)) {
return basePolicy.shouldRetry(attempt, error);
} else {
return RetryDecision.fail();
}
};
}
}
@@ -0,0 +1,160 @@
package com.langgraph.pregel.state;
import java.util.Collections;
import java.util.HashMap;
import java.util.Map;
import java.util.Objects;
/**
* Represents a snapshot of execution state at a superstep boundary.
* Checkpoints contain the serialized state of all channels at a point in time.
*/
public class Checkpoint {
private Map<String, Object> channelValues;
/**
* Create a Checkpoint with channel values.
*
* @param channelValues Map of channel names to their checkpoint values
*/
public Checkpoint(Map<String, Object> channelValues) {
this.channelValues = channelValues != null ? new HashMap<>(channelValues) : new HashMap<>();
}
/**
* Create an empty Checkpoint.
*/
public Checkpoint() {
this(Collections.emptyMap());
}
/**
* Get the channel values.
*
* @return Unmodifiable map of channel values
*/
public Map<String, Object> getValues() {
return Collections.unmodifiableMap(channelValues);
}
/**
* Get a channel value by name.
*
* @param channelName Channel name
* @return Channel value, or null if not present
*/
public Object getValue(String channelName) {
return channelValues.get(channelName);
}
/**
* Update the checkpoint with new channel values.
*
* @param channelValues New channel values
*/
public void update(Map<String, Object> channelValues) {
if (channelValues == null) {
throw new IllegalArgumentException("Channel values cannot be null");
}
this.channelValues = new HashMap<>(channelValues);
}
/**
* Update a single channel value.
*
* @param channelName Channel name
* @param value Channel value
*/
public void updateChannel(String channelName, Object value) {
if (channelName == null || channelName.isEmpty()) {
throw new IllegalArgumentException("Channel name cannot be null or empty");
}
if (value == null) {
channelValues.remove(channelName);
} else {
channelValues.put(channelName, value);
}
}
/**
* Check if this checkpoint contains a value for the given channel.
*
* @param channelName Channel name
* @return True if the checkpoint contains a value for the channel
*/
public boolean containsChannel(String channelName) {
return channelValues.containsKey(channelName);
}
/**
* Create a new Checkpoint with updated values.
*
* @param updates Map of channel names to values to update
* @return New Checkpoint with updated values
*/
public Checkpoint withUpdates(Map<String, Object> updates) {
if (updates == null || updates.isEmpty()) {
return this;
}
Map<String, Object> newValues = new HashMap<>(this.channelValues);
newValues.putAll(updates);
return new Checkpoint(newValues);
}
/**
* Create a new Checkpoint with only the specified channels.
*
* @param channelNames Channel names to include
* @return New Checkpoint with only the specified channels
*/
public Checkpoint subset(Iterable<String> channelNames) {
if (channelNames == null) {
return new Checkpoint();
}
Map<String, Object> subsetValues = new HashMap<>();
for (String name : channelNames) {
if (channelValues.containsKey(name)) {
subsetValues.put(name, channelValues.get(name));
}
}
return new Checkpoint(subsetValues);
}
/**
* Get the number of channels in this checkpoint.
*
* @return Number of channels
*/
public int size() {
return channelValues.size();
}
/**
* Check if this checkpoint is empty.
*
* @return True if the checkpoint contains no values
*/
public boolean isEmpty() {
return channelValues.isEmpty();
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
Checkpoint that = (Checkpoint) o;
return Objects.equals(channelValues, that.channelValues);
}
@Override
public int hashCode() {
return Objects.hash(channelValues);
}
@Override
public String toString() {
return "Checkpoint{channelCount=" + channelValues.size() + "}";
}
}
@@ -0,0 +1,155 @@
package com.langgraph.pregel.stream;
import com.langgraph.pregel.StreamMode;
import java.util.Map;
import java.util.Queue;
import java.util.Set;
import java.util.concurrent.ConcurrentLinkedQueue;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.function.Consumer;
/**
* Controls the streaming of results during Pregel execution.
* Manages backpressure, cancellation, and output formatting.
*/
public class StreamController {
private final Queue<Map<String, Object>> buffer;
private final AtomicBoolean isCancelled;
private final AtomicBoolean isPaused;
private final Consumer<Map<String, Object>> outputConsumer;
private final StreamMode streamMode;
private int stepCount;
/**
* Create a StreamController.
*
* @param outputConsumer Consumer to receive output
* @param streamMode Stream mode
*/
public StreamController(Consumer<Map<String, Object>> outputConsumer, StreamMode streamMode) {
this.buffer = new ConcurrentLinkedQueue<>();
this.isCancelled = new AtomicBoolean(false);
this.isPaused = new AtomicBoolean(false);
this.outputConsumer = outputConsumer;
this.streamMode = streamMode != null ? streamMode : StreamMode.VALUES;
this.stepCount = 0;
}
/**
* Process a new state update.
*
* @param state Current state
* @param updatedChannels Set of channel names that were updated in this step
* @param hasMoreWork Whether there is more work to do
* @return True if execution should continue, false if it should stop
*/
public boolean processUpdate(Map<String, Object> state, Set<String> updatedChannels, boolean hasMoreWork) {
if (isCancelled.get()) {
return false;
}
stepCount++;
// Format output based on stream mode
Map<String, Object> output = StreamOutput.format(
state,
updatedChannels,
stepCount,
hasMoreWork,
streamMode
);
// Add to buffer and consume if not paused
buffer.add(output);
consumeOutput();
return !isCancelled.get();
}
/**
* Consume output from the buffer.
*/
private void consumeOutput() {
if (isPaused.get() || outputConsumer == null) {
return;
}
Map<String, Object> output;
while ((output = buffer.poll()) != null) {
outputConsumer.accept(output);
// Check if we should stop consuming
if (isPaused.get() || isCancelled.get()) {
break;
}
}
}
/**
* Cancel streaming.
*/
public void cancel() {
isCancelled.set(true);
}
/**
* Pause streaming.
*/
public void pause() {
isPaused.set(true);
}
/**
* Resume streaming.
*/
public void resume() {
isPaused.set(false);
consumeOutput();
}
/**
* Check if streaming is cancelled.
*
* @return True if cancelled
*/
public boolean isCancelled() {
return isCancelled.get();
}
/**
* Check if streaming is paused.
*
* @return True if paused
*/
public boolean isPaused() {
return isPaused.get();
}
/**
* Get the current step count.
*
* @return Current step count
*/
public int getStepCount() {
return stepCount;
}
/**
* Get the buffer size.
*
* @return Buffer size
*/
public int getBufferSize() {
return buffer.size();
}
/**
* Get the stream mode.
*
* @return Stream mode
*/
public StreamMode getStreamMode() {
return streamMode;
}
}
@@ -0,0 +1,105 @@
package com.langgraph.pregel.stream;
import com.langgraph.pregel.StreamMode;
import java.util.HashMap;
import java.util.Map;
import java.util.Set;
/**
* Utility class for formatting output for streaming based on the stream mode.
*/
public class StreamOutput {
private StreamOutput() {
// Prevent instantiation
}
/**
* Format output for streaming based on the stream mode.
*
* @param state Current state
* @param updatedChannels Set of channel names that were updated in this step
* @param step Current step number
* @param hasMoreWork Whether there is more work to do
* @param streamMode Stream mode
* @return Formatted output
*/
public static Map<String, Object> format(
Map<String, Object> state,
Set<String> updatedChannels,
int step,
boolean hasMoreWork,
StreamMode streamMode) {
if (streamMode == null) {
streamMode = StreamMode.VALUES;
}
switch (streamMode) {
case VALUES:
// Return the full state
return state != null ? new HashMap<>(state) : new HashMap<>();
case UPDATES:
// Return only the updated channels
Map<String, Object> updates = new HashMap<>();
if (updatedChannels != null && state != null) {
for (String channelName : updatedChannels) {
if (state.containsKey(channelName)) {
updates.put(channelName, state.get(channelName));
}
}
}
return updates;
case DEBUG:
// Return detailed debug information
Map<String, Object> debug = new HashMap<>();
debug.put("state", state != null ? new HashMap<>(state) : new HashMap<>());
debug.put("updated_channels", updatedChannels);
debug.put("step", step);
debug.put("has_more_work", hasMoreWork);
return debug;
default:
return state != null ? new HashMap<>(state) : new HashMap<>();
}
}
/**
* Format values mode output.
*
* @param state Current state
* @return Formatted output
*/
public static Map<String, Object> formatValues(Map<String, Object> state) {
return format(state, null, 0, false, StreamMode.VALUES);
}
/**
* Format updates mode output.
*
* @param state Current state
* @param updatedChannels Set of channel names that were updated in this step
* @return Formatted output
*/
public static Map<String, Object> formatUpdates(Map<String, Object> state, Set<String> updatedChannels) {
return format(state, updatedChannels, 0, false, StreamMode.UPDATES);
}
/**
* Format debug mode output.
*
* @param state Current state
* @param updatedChannels Set of channel names that were updated in this step
* @param step Current step number
* @param hasMoreWork Whether there is more work to do
* @return Formatted output
*/
public static Map<String, Object> formatDebug(
Map<String, Object> state,
Set<String> updatedChannels,
int step,
boolean hasMoreWork) {
return format(state, updatedChannels, step, hasMoreWork, StreamMode.DEBUG);
}
}
@@ -0,0 +1,86 @@
package com.langgraph.pregel.task;
import java.util.Collections;
import java.util.Map;
import java.util.Objects;
import java.util.HashMap;
/**
* Represents an executable task with inputs and context.
* This is a concrete task ready for execution with all required data.
*/
public class PregelExecutableTask {
private final PregelTask task;
private final Map<String, Object> inputs;
private final Map<String, Object> context;
/**
* Create a PregelExecutableTask with all parameters.
*
* @param task Task to execute
* @param inputs Channel inputs for the task
* @param context Execution context
*/
public PregelExecutableTask(
PregelTask task,
Map<String, Object> inputs,
Map<String, Object> context) {
if (task == null) {
throw new IllegalArgumentException("Task cannot be null");
}
this.task = task;
this.inputs = inputs != null ? new HashMap<>(inputs) : Collections.emptyMap();
this.context = context != null ? new HashMap<>(context) : Collections.emptyMap();
}
/**
* Get the task.
*
* @return Task
*/
public PregelTask getTask() {
return task;
}
/**
* Get the inputs.
*
* @return Map of channel inputs (immutable)
*/
public Map<String, Object> getInputs() {
return Collections.unmodifiableMap(inputs);
}
/**
* Get the context.
*
* @return Map of context values (immutable)
*/
public Map<String, Object> getContext() {
return Collections.unmodifiableMap(context);
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
PregelExecutableTask that = (PregelExecutableTask) o;
return Objects.equals(task, that.task) &&
Objects.equals(inputs, that.inputs) &&
Objects.equals(context, that.context);
}
@Override
public int hashCode() {
return Objects.hash(task, inputs, context);
}
@Override
public String toString() {
return "PregelExecutableTask{" +
"task=" + task +
", inputs=" + inputs.keySet() +
", context=" + context.keySet() +
'}';
}
}
@@ -0,0 +1,99 @@
package com.langgraph.pregel.task;
import com.langgraph.pregel.retry.RetryPolicy;
import java.util.Objects;
/**
* Represents a task to be executed within the Pregel system.
* A task identifies a node to execute, an optional trigger, and a retry policy.
*/
public class PregelTask {
private final String node;
private final String trigger;
private final RetryPolicy retryPolicy;
/**
* Create a PregelTask with all parameters.
*
* @param node Node name to execute
* @param trigger Optional trigger that caused this task
* @param retryPolicy Optional retry policy for execution failures
*/
public PregelTask(String node, String trigger, RetryPolicy retryPolicy) {
if (node == null || node.isEmpty()) {
throw new IllegalArgumentException("Node name cannot be null or empty");
}
this.node = node;
this.trigger = trigger;
this.retryPolicy = retryPolicy;
}
/**
* Create a PregelTask with just a node name.
*
* @param node Node name to execute
*/
public PregelTask(String node) {
this(node, null, null);
}
/**
* Create a PregelTask with node name and trigger.
*
* @param node Node name to execute
* @param trigger Trigger that caused this task
*/
public PregelTask(String node, String trigger) {
this(node, trigger, null);
}
/**
* Get the node name.
*
* @return Node name
*/
public String getNode() {
return node;
}
/**
* Get the trigger.
*
* @return Trigger or null if not triggered
*/
public String getTrigger() {
return trigger;
}
/**
* Get the retry policy.
*
* @return RetryPolicy or null if using default policy
*/
public RetryPolicy getRetryPolicy() {
return retryPolicy;
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
PregelTask that = (PregelTask) o;
return Objects.equals(node, that.node) &&
Objects.equals(trigger, that.trigger);
}
@Override
public int hashCode() {
return Objects.hash(node, trigger);
}
@Override
public String toString() {
return "PregelTask{" +
"node='" + node + '\'' +
(trigger != null ? ", trigger='" + trigger + '\'' : "") +
'}';
}
}
@@ -0,0 +1,50 @@
package com.langgraph.pregel.task;
/**
* Exception thrown when task execution fails after all retry attempts.
*/
public class TaskExecutionException extends RuntimeException {
private final int attempt;
/**
* Create a TaskExecutionException with a message.
*
* @param message Error message
*/
public TaskExecutionException(String message) {
super(message);
this.attempt = 0;
}
/**
* Create a TaskExecutionException with a message and cause.
*
* @param message Error message
* @param cause Cause of the error
*/
public TaskExecutionException(String message, Throwable cause) {
super(message, cause);
this.attempt = 0;
}
/**
* Create a TaskExecutionException with a message, cause, and attempt number.
*
* @param message Error message
* @param cause Cause of the error
* @param attempt Attempt number that failed
*/
public TaskExecutionException(String message, Throwable cause, int attempt) {
super(message, cause);
this.attempt = attempt;
}
/**
* Get the attempt number that failed.
*
* @return Attempt number
*/
public int getAttempt() {
return attempt;
}
}
@@ -0,0 +1,145 @@
package com.langgraph.pregel.task;
import com.langgraph.pregel.PregelNode;
import com.langgraph.pregel.retry.RetryPolicy;
import java.time.Duration;
import java.util.Map;
import java.util.concurrent.Callable;
import java.util.concurrent.CancellationException;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.CompletionException;
/**
* Executes PregelExecutableTask instances with retry logic.
*/
public class TaskExecutor {
private static final int DEFAULT_MAX_ATTEMPTS = 3;
private final RetryPolicy defaultRetryPolicy;
/**
* Create a TaskExecutor with a default retry policy.
*
* @param defaultRetryPolicy Default retry policy for tasks
*/
public TaskExecutor(RetryPolicy defaultRetryPolicy) {
this.defaultRetryPolicy = defaultRetryPolicy;
}
/**
* Create a TaskExecutor with a default retry policy allowing up to 3 attempts.
*/
public TaskExecutor() {
this(RetryPolicy.maxAttempts(DEFAULT_MAX_ATTEMPTS));
}
/**
* Execute a task.
*
* @param node Node to execute
* @param task Task to execute
* @return Result of the execution
* @throws TaskExecutionException If execution fails after all retry attempts
*/
public Map<String, Object> execute(PregelNode node, PregelExecutableTask task) throws TaskExecutionException {
if (node == null) {
throw new IllegalArgumentException("Node cannot be null");
}
if (task == null) {
throw new IllegalArgumentException("Task cannot be null");
}
// Get the appropriate retry policy
RetryPolicy retryPolicy = task.getTask().getRetryPolicy();
if (retryPolicy == null) {
retryPolicy = node.getRetryPolicy();
if (retryPolicy == null) {
retryPolicy = defaultRetryPolicy;
}
}
// Execute with retry
return executeWithRetry(() -> {
Map<String, Object> inputs = task.getInputs();
Map<String, Object> context = task.getContext();
try {
// Execute the action
return node.getAction().execute(inputs, context);
} catch (Exception e) {
// Wrap and rethrow
throw new TaskExecutionException("Error executing node " + node.getName(), e);
}
}, retryPolicy);
}
/**
* Execute a task asynchronously.
*
* @param node Node to execute
* @param task Task to execute
* @return CompletableFuture with the result of the execution
*/
public CompletableFuture<Map<String, Object>> executeAsync(PregelNode node, PregelExecutableTask task) {
return CompletableFuture.supplyAsync(() -> execute(node, task));
}
/**
* Execute a callable with retry logic.
*
* @param <T> Type of the result
* @param callable Callable to execute
* @param retryPolicy Retry policy to use
* @return Result of the callable
* @throws TaskExecutionException If execution fails after all retry attempts
*/
private <T> T executeWithRetry(Callable<T> callable, RetryPolicy retryPolicy) throws TaskExecutionException {
int attempt = 1;
Throwable lastError = null;
while (true) {
try {
return callable.call();
} catch (CancellationException | InterruptedException e) {
// Do not retry cancellation or interruption
Thread.currentThread().interrupt();
throw new TaskExecutionException("Task execution was cancelled or interrupted", e);
} catch (CompletionException e) {
// Unwrap CompletionException
lastError = e.getCause() != null ? e.getCause() : e;
} catch (Exception e) {
lastError = e;
}
// If we get here, execution failed
if (retryPolicy != null) {
RetryPolicy.RetryDecision decision = retryPolicy.shouldRetry(attempt, lastError);
if (decision.shouldRetry()) {
// Sleep if backoff is specified
Duration backoff = decision.getBackoff();
if (!backoff.isZero() && !backoff.isNegative()) {
try {
Thread.sleep(backoff.toMillis());
} catch (InterruptedException ie) {
Thread.currentThread().interrupt();
throw new TaskExecutionException("Retry was interrupted", ie);
}
}
// Increment attempt counter
attempt++;
} else {
// Do not retry
break;
}
} else {
// No retry policy, fail immediately
break;
}
}
// All retries failed or no retry policy
throw new TaskExecutionException("Task execution failed after " + attempt + " attempts", lastError);
}
}
@@ -0,0 +1,167 @@
package com.langgraph.pregel.task;
import com.langgraph.pregel.PregelNode;
import java.util.*;
import java.util.stream.Collectors;
/**
* Plans which nodes to execute based on channel updates.
*
* <p>The TaskPlanner determines which nodes to execute in each superstep based on which
* channels have been updated. There are two important cases:</p>
*
* <ol>
* <li>First Superstep (no channels updated yet):
* <ul>
* <li>Current Behavior: All nodes are executed, regardless of subscriptions or triggers</li>
* <li>Python LangGraph Behavior: Only nodes with the input channel as their trigger would execute</li>
* </ul>
* </li>
* <li>Subsequent Supersteps:
* <ul>
* <li>Nodes execute if either:
* <ol>
* <li>They subscribe to a channel that was updated</li>
* <li>They have a trigger matching a channel that was updated</li>
* </ol>
* </li>
* </ul>
* </li>
* </ol>
*
* <p>Note: For Python compatibility, a future version of this implementation will likely change
* to only execute nodes with the appropriate input channel trigger in the first superstep.</p>
*/
public class TaskPlanner {
private final Map<String, PregelNode<?, ?>> nodes;
// The input channel name, used to determine which nodes should run in first superstep
private final String inputChannelName;
/**
* Create a TaskPlanner with default input channel name "input".
*
* @param nodes Map of node names to nodes
*/
public TaskPlanner(Map<String, PregelNode<?, ?>> nodes) {
this(nodes, "input");
}
/**
* Create a TaskPlanner with a specific input channel name.
*
* @param nodes Map of node names to nodes
* @param inputChannelName The name of the input channel
*/
public TaskPlanner(Map<String, PregelNode<?, ?>> nodes, String inputChannelName) {
if (nodes == null) {
throw new IllegalArgumentException("Nodes cannot be null");
}
this.nodes = new HashMap<>(nodes);
this.inputChannelName = inputChannelName;
}
/**
* Plan which nodes to execute based on updated channels.
* With full Python compatibility for uninitialized channels.
*
* @param updatedChannels Set of channel names that were updated
* @return List of tasks to execute
*/
public List<PregelTask> plan(Collection<String> updatedChannels) {
// For first superstep when no channels have been updated yet
if (updatedChannels == null || updatedChannels.isEmpty()) {
// Proper Python compatibility: only run nodes with input channel trigger
List<PregelTask> tasks = new ArrayList<>();
for (PregelNode<?, ?> node : nodes.values()) {
// Use the newer method for checking trigger channels
if (node.isTriggeredBy(inputChannelName)) {
// Use the first trigger channel for Task creation
Set<String> triggerChannels = node.getTriggerChannels();
String trigger = triggerChannels.isEmpty() ?
null : triggerChannels.iterator().next();
tasks.add(new PregelTask(node.getName(), trigger, node.getRetryPolicy()));
}
}
return tasks;
}
// Convert to set for O(1) lookups
Set<String> updatedChannelSet = new HashSet<>(updatedChannels);
// Collect tasks to execute
List<PregelTask> tasks = new ArrayList<>();
for (PregelNode<?, ?> node : nodes.values()) {
// Check if the node reads from any updated channels
boolean shouldExecute = false;
Set<String> nodeChannels = node.getChannels();
for (String channelName : nodeChannels) {
if (updatedChannelSet.contains(channelName)) {
shouldExecute = true;
break;
}
}
// Check if the node is triggered by any updated channels
if (!shouldExecute) {
Set<String> nodeTriggers = node.getTriggerChannels();
for (String channelName : nodeTriggers) {
if (updatedChannelSet.contains(channelName)) {
shouldExecute = true;
break;
}
}
}
if (shouldExecute) {
// Use the first trigger channel for Task creation
Set<String> triggerChannels = node.getTriggerChannels();
String trigger = triggerChannels.isEmpty() ?
null : triggerChannels.iterator().next();
tasks.add(new PregelTask(node.getName(), trigger, node.getRetryPolicy()));
}
}
return tasks;
}
/**
* Prioritize tasks for execution.
* This method can be overridden to implement custom prioritization logic.
*
* @param tasks List of tasks to prioritize
* @return Prioritized list of tasks
*/
public List<PregelTask> prioritize(List<PregelTask> tasks) {
// Default implementation does not change the order
return new ArrayList<>(tasks);
}
/**
* Filter tasks based on dependencies.
* This method can be overridden to implement custom filtering logic.
*
* @param tasks List of tasks to filter
* @return Filtered list of tasks
*/
protected List<PregelTask> filter(List<PregelTask> tasks) {
return tasks.stream()
.filter(task -> nodes.containsKey(task.getNode()))
.collect(Collectors.toList());
}
/**
* Plan, filter, and prioritize tasks for execution.
*
* @param updatedChannels Set of channel names that were updated
* @return Prioritized list of tasks to execute
*/
public List<PregelTask> planAndPrioritize(Collection<String> updatedChannels) {
List<PregelTask> tasks = plan(updatedChannels);
tasks = filter(tasks);
return prioritize(tasks);
}
}
@@ -0,0 +1,147 @@
package com.langgraph.channels;
import org.junit.jupiter.api.Test;
import java.util.Arrays;
import java.util.Collections;
import java.util.function.BinaryOperator;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
public class BinaryOperatorChannelTest {
@Test
void testEmptyChannel() {
BinaryOperatorChannel<Integer> channel = BinaryOperatorChannel.create(Integer::sum, 0);
assertThatThrownBy(channel::get)
.isInstanceOf(EmptyChannelException.class)
.hasMessageContaining("empty");
}
@Test
void testSumOperator() {
BinaryOperatorChannel<Integer> channel = BinaryOperatorChannel.create(Integer::sum, 0);
// Initial update
boolean updated = channel.update(Collections.singletonList(5));
assertThat(updated).isTrue();
assertThat(channel.get()).isEqualTo(5);
// Add more values
updated = channel.update(Arrays.asList(10, 7, 3));
assertThat(updated).isTrue();
assertThat(channel.get()).isEqualTo(25); // 5 + 10 + 7 + 3 = 25
}
@Test
void testMaxOperator() {
BinaryOperatorChannel<Integer> channel = BinaryOperatorChannel.create(Integer::max, Integer.MIN_VALUE);
// Initial update
channel.update(Collections.singletonList(5));
assertThat(channel.get()).isEqualTo(5);
// Add higher values
channel.update(Arrays.asList(10, 7));
assertThat(channel.get()).isEqualTo(10);
// Add lower values
channel.update(Collections.singletonList(3));
assertThat(channel.get()).isEqualTo(10); // Max is still 10
}
@Test
void testStringConcatenation() {
BinaryOperator<String> concat = (a, b) -> a + b;
BinaryOperatorChannel<String> channel = BinaryOperatorChannel.create(concat, "");
// Initial update
channel.update(Collections.singletonList("Hello"));
assertThat(channel.get()).isEqualTo("Hello");
// Add more values
channel.update(Arrays.asList(", ", "World", "!"));
assertThat(channel.get()).isEqualTo("Hello, World!");
}
@Test
void testEmptyUpdate() {
BinaryOperatorChannel<Integer> channel = BinaryOperatorChannel.create(Integer::sum, 0);
// Empty update should return false
boolean updated = channel.update(Collections.emptyList());
assertThat(updated).isFalse();
// Channel should still be empty
assertThatThrownBy(channel::get)
.isInstanceOf(EmptyChannelException.class);
}
@Test
void testUpdateOrder() {
// Using subtraction to check order (not commutative)
BinaryOperator<Integer> subtract = (a, b) -> a - b;
BinaryOperatorChannel<Integer> channel = BinaryOperatorChannel.create(subtract, 100);
// Subtract values from 100
channel.update(Arrays.asList(20, 30));
// Result should be 100 - 20 - 30 = 50
assertThat(channel.get()).isEqualTo(50);
}
@Test
void testCheckpoint() {
BinaryOperatorChannel<Integer> channel = BinaryOperatorChannel.create(Integer::sum, 0);
// Update the channel
channel.update(Arrays.asList(5, 10, 15));
// Create a checkpoint
Integer checkpoint = channel.checkpoint();
assertThat(checkpoint).isEqualTo(30);
// Create a new channel from the checkpoint
BinaryOperatorChannel<Integer> newChannel =
(BinaryOperatorChannel<Integer>) channel.fromCheckpoint(checkpoint);
// Verify the new channel has the same accumulated value
assertThat(newChannel.get()).isEqualTo(30);
// Add more to the original
channel.update(Collections.singletonList(20));
assertThat(channel.get()).isEqualTo(50);
// New channel should be unchanged
assertThat(newChannel.get()).isEqualTo(30);
// Add to the new channel
newChannel.update(Collections.singletonList(5));
assertThat(newChannel.get()).isEqualTo(35);
}
@Test
void testUtilityMethods() {
// Test the utility method for Integer adder
BinaryOperatorChannel<Integer> intAdder = Channels.integerAdder("counter");
intAdder.update(Arrays.asList(1, 2, 3));
assertThat(intAdder.get()).isEqualTo(6);
// Test the utility method for Long adder
BinaryOperatorChannel<Long> longAdder = Channels.longAdder("longCounter");
longAdder.update(Arrays.asList(1L, 2L, 3L));
assertThat(longAdder.get()).isEqualTo(6L);
// Test the utility method for Double adder
BinaryOperatorChannel<Double> doubleAdder = Channels.doubleAdder("doubleCounter");
doubleAdder.update(Arrays.asList(1.5, 2.5));
assertThat(doubleAdder.get()).isEqualTo(4.0);
// Test the utility method for Integer max
BinaryOperatorChannel<Integer> intMax = Channels.integerMax("maxValue");
intMax.update(Arrays.asList(5, 10, 3));
assertThat(intMax.get()).isEqualTo(10);
}
}
@@ -0,0 +1,125 @@
package com.langgraph.channels;
import org.junit.jupiter.api.Test;
import java.util.Arrays;
import java.util.Collections;
import java.util.List;
import java.util.function.BinaryOperator;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
public class ChannelsTest {
@Test
void testLastValueFactory() {
// Create channel using factory
LastValue<String> channel = LastValue.<String>create();
// Update the channel
channel.update(Collections.singletonList("test"));
// Verify it works
assertThat(channel.get()).isEqualTo("test");
// Create channel with key
LastValue<String> namedChannel = LastValue.<String>create("input");
assertThat(namedChannel.getKey()).isEqualTo("input");
}
@Test
void testTopicFactory() {
// Create topic channel using factory
TopicChannel<String> channel = TopicChannel.<String>create();
// Update with values
channel.update(Arrays.asList("one", "two"));
// Verify it works
assertThat(channel.get()).containsExactly("one", "two");
// Create reset-on-consume topic
TopicChannel<String> resetChannel = TopicChannel.<String>create(true);
resetChannel.update(Collections.singletonList("test"));
// Consume should reset the channel
boolean consumed = resetChannel.consume();
assertThat(consumed).isTrue();
// Channel should be empty but not throw with Python compatibility
assertThat(resetChannel.get()).isEmpty();
// Create with key
TopicChannel<String> namedChannel = TopicChannel.<String>create("messages", false);
assertThat(namedChannel.getKey()).isEqualTo("messages");
}
@Test
void testBinaryOperatorFactory() {
// Create a binary operator channel using factory
BinaryOperator<Integer> sum = Integer::sum;
BinaryOperatorChannel<Integer> channel = Channels.<Integer>binaryOperator(sum, 0);
// Update with values
channel.update(Arrays.asList(1, 2, 3));
// Verify it works
assertThat(channel.get()).isEqualTo(6);
// Create with key
BinaryOperatorChannel<Integer> namedChannel =
Channels.<Integer>binaryOperator("counter", sum, 0);
assertThat(namedChannel.getKey()).isEqualTo("counter");
}
@Test
void testEphemeralFactory() {
// Create ephemeral channel using factory
EphemeralValue<String> channel = Channels.<String>ephemeral();
// Update with value
channel.update(Collections.singletonList("test"));
// Verify it works
assertThat(channel.get()).isEqualTo("test");
// Create with key
EphemeralValue<String> namedChannel = Channels.<String>ephemeral("temporary");
assertThat(namedChannel.getKey()).isEqualTo("temporary");
}
@Test
void testNumericOperatorFactories() {
// Test integer adder
BinaryOperatorChannel<Integer> intAdder = Channels.integerAdder("int-sum");
intAdder.update(Arrays.asList(1, 2, 3));
assertThat(intAdder.get()).isEqualTo(6);
assertThat(intAdder.getKey()).isEqualTo("int-sum");
// Test long adder
BinaryOperatorChannel<Long> longAdder = Channels.longAdder("long-sum");
longAdder.update(Arrays.asList(100L, 200L, 300L));
assertThat(longAdder.get()).isEqualTo(600L);
// Test double adder
BinaryOperatorChannel<Double> doubleAdder = Channels.doubleAdder("double-sum");
doubleAdder.update(Arrays.asList(1.5, 2.5, 3.0));
assertThat(doubleAdder.get()).isEqualTo(7.0);
// Test integer max
BinaryOperatorChannel<Integer> intMax = Channels.integerMax("max-value");
intMax.update(Arrays.asList(5, 10, 3));
assertThat(intMax.get()).isEqualTo(10);
// Test long max
BinaryOperatorChannel<Long> longMax = Channels.longMax("long-max");
longMax.update(Arrays.asList(100L, 500L, 200L));
assertThat(longMax.get()).isEqualTo(500L);
// Test double max
BinaryOperatorChannel<Double> doubleMax = Channels.doubleMax("double-max");
doubleMax.update(Arrays.asList(1.5, 3.5, 2.0));
assertThat(doubleMax.get()).isEqualTo(3.5);
}
}

Some files were not shown because too many files have changed in this diff Show More