mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 20:45:05 +02:00
Add spec
This commit is contained in:
@@ -0,0 +1,50 @@
|
||||
# LangGraph Coding Guide
|
||||
|
||||
## Build/Test/Lint Commands
|
||||
|
||||
- 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`
|
||||
|
||||
## Code Style Guidelines
|
||||
|
||||
- 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
|
||||
|
||||
## Feature Overview
|
||||
|
||||
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:
|
||||
|
||||
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
|
||||
|
||||
## Repository Structure
|
||||
|
||||
LangGraph follows a monorepo organization, with the following structure:
|
||||
|
||||
- `docs/` contains the source code (markdown and jupyter notebooks) for our documentation (hosted at https://langchain-ai.github.io/langgraph/)
|
||||
- `libs/langgraph` is the main 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.
|
||||
@@ -0,0 +1,379 @@
|
||||
# LangGraph Architecture Specification
|
||||
|
||||
## Overview
|
||||
|
||||
LangGraph is a framework for building stateful, observable applications with large language models (LLMs). It uses a graph-based architecture with explicit state management to provide features like streaming output, cyclical workflows, human-in-the-loop capabilities, and persistence.
|
||||
|
||||
This document provides a comprehensive overview of LangGraph's architecture, how the components interact, and the design principles that guide its implementation.
|
||||
|
||||
## Architectural Layers
|
||||
|
||||
LangGraph follows a layered architecture that provides different levels of abstraction:
|
||||
|
||||
```
|
||||
┌───────────────────────────────────────────────────────┐
|
||||
│ Application Layer │
|
||||
│ (User-defined agents and cognitive architectures) │
|
||||
└────────────────────────────┬──────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────────────────────────────────────────┐
|
||||
│ High-Level API Layer │
|
||||
│ (StateGraph, Functional API, etc.) │
|
||||
└────────────────────────────┬──────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────────────────────────────────────────┐
|
||||
│ Execution Layer │
|
||||
│ (Pregel) │
|
||||
└────────────────────────────┬──────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────────────────────────────────────────┐
|
||||
│ State Management Layer │
|
||||
│ (Channels, Schemas, Checkpoints) │
|
||||
└────────────────────────────┬──────────────────────────┘
|
||||
│
|
||||
▼
|
||||
┌───────────────────────────────────────────────────────┐
|
||||
│ Persistence Layer │
|
||||
│ (Memory, Disk, Database implementations) │
|
||||
└───────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
### Application Layer
|
||||
|
||||
Where users define their specific LLM applications, agents, and workflows using the LangGraph API.
|
||||
|
||||
### High-Level API Layer
|
||||
|
||||
Provides intuitive interfaces like `StateGraph` for defining computation graphs with minimal boilerplate.
|
||||
|
||||
### Execution Layer
|
||||
|
||||
Implements the Pregel computation model for executing the graph in a deterministic, observable way.
|
||||
|
||||
### State Management Layer
|
||||
|
||||
Handles state definition, validation, transformation, and propagation through the graph.
|
||||
|
||||
### Persistence Layer
|
||||
|
||||
Provides storage implementations for checkpoints and long-term memory.
|
||||
|
||||
## Core Components
|
||||
|
||||
### StateGraph
|
||||
|
||||
The primary user-facing API for defining computation graphs:
|
||||
|
||||
```python
|
||||
from langgraph.graph import StateGraph
|
||||
from typing import TypedDict, Annotated
|
||||
|
||||
# Define state schema
|
||||
class State(TypedDict):
|
||||
messages: list[str]
|
||||
counter: int
|
||||
|
||||
# Create graph with schema
|
||||
graph = StateGraph(State)
|
||||
|
||||
# Add nodes (functions)
|
||||
graph.add_node("process", process_func)
|
||||
graph.add_node("decide", decide_func)
|
||||
|
||||
# Add edges
|
||||
graph.add_edge("process", "decide")
|
||||
graph.add_conditional_edges(
|
||||
"decide",
|
||||
lambda state: "continue" if state["counter"] < 5 else "end"
|
||||
)
|
||||
|
||||
# Compile graph into a runnable
|
||||
workflow = graph.compile()
|
||||
```
|
||||
|
||||
Key features:
|
||||
|
||||
- Type-safe state schema
|
||||
- Conditional routing
|
||||
- Cyclical execution patterns
|
||||
- Checkpoint integration
|
||||
- Streaming support
|
||||
|
||||
### Pregel Execution Engine
|
||||
|
||||
The computational backbone that executes the graph:
|
||||
|
||||
- Implements the Bulk Synchronous Parallel computation model
|
||||
- Manages the lifecycle of node execution and state updates
|
||||
- Ensures deterministic execution despite parallel processing
|
||||
- Integrates with the checkpoint system
|
||||
- Provides streaming capabilities
|
||||
|
||||
### Channel System
|
||||
|
||||
Provides communication between nodes with specialized behaviors:
|
||||
|
||||
- LastValue: Stores a single value, ensuring type safety
|
||||
- Topic: Pub/sub pattern for multi-consumer updates
|
||||
- BinaryOperatorAggregate: Combines values using operators
|
||||
- Barrier channels: Synchronization mechanisms
|
||||
- And more specialized channel types
|
||||
|
||||
### State Schema
|
||||
|
||||
Defines the structure and behavior of application state:
|
||||
|
||||
```python
|
||||
class ConversationState(TypedDict):
|
||||
messages: list[dict] # Regular list
|
||||
context: Annotated[dict, untracked()] # Excluded from checkpoints
|
||||
history: Annotated[list[str], append()] # Append-only list
|
||||
```
|
||||
|
||||
Features:
|
||||
|
||||
- Type validation
|
||||
- Custom reducers via annotations
|
||||
- Integration with channels
|
||||
- Multiple definition formats (TypedDict, Pydantic, dataclass)
|
||||
|
||||
### Checkpoint System
|
||||
|
||||
Enables persistence and human-in-the-loop capabilities:
|
||||
|
||||
- Thread-based execution isolation
|
||||
- Checkpoint creation and restoration
|
||||
- State history tracking
|
||||
- Time travel debugging
|
||||
- Hierarchical checkpoint namespaces
|
||||
|
||||
### Human-in-the-Loop
|
||||
|
||||
Support for interactive workflows:
|
||||
|
||||
- Interruption at specific points
|
||||
- State inspection during interruption
|
||||
- State modification
|
||||
- Resumption from interrupted state
|
||||
|
||||
## Key Interfaces
|
||||
|
||||
### StateGraph API
|
||||
|
||||
```python
|
||||
class StateGraph:
|
||||
def __init__(self, state_schema: Type) -> None: ...
|
||||
|
||||
def add_node(self, name: str, action: Callable) -> None: ...
|
||||
|
||||
def add_edge(self, start: str, end: str) -> None: ...
|
||||
|
||||
def add_conditional_edges(
|
||||
self,
|
||||
start: str,
|
||||
condition: Callable[[Any], str]
|
||||
) -> None: ...
|
||||
|
||||
def compile(self, **kwargs) -> PregelRunnable: ...
|
||||
```
|
||||
|
||||
### PregelRunnable API
|
||||
|
||||
```python
|
||||
class PregelRunnable:
|
||||
def invoke(self, input: Any, config: dict = None) -> Any: ...
|
||||
|
||||
def stream(
|
||||
self,
|
||||
input: Any,
|
||||
config: dict = None,
|
||||
stream_mode: StreamMode = None
|
||||
) -> Iterator[Any]: ...
|
||||
|
||||
def get_state(self, thread_id: str = None) -> Any: ...
|
||||
|
||||
def update_state(self, thread_id: str, state: Any) -> None: ...
|
||||
|
||||
def get_state_history(self, thread_id: str) -> list[Any]: ...
|
||||
```
|
||||
|
||||
## Design Principles
|
||||
|
||||
LangGraph's architecture is guided by the following principles:
|
||||
|
||||
### 1. Explicit State
|
||||
|
||||
State is always explicitly defined and validated, providing type safety and preventing many classes of bugs.
|
||||
|
||||
### 2. Composability
|
||||
|
||||
Components are designed to be combined in various ways:
|
||||
|
||||
- Nodes can be nested graphs
|
||||
- Channels can be composed for complex behaviors
|
||||
- States can be nested for hierarchical organization
|
||||
|
||||
### 3. Observability
|
||||
|
||||
Execution is transparent and observable:
|
||||
|
||||
- Streaming support for real-time visibility
|
||||
- State history tracking
|
||||
- Detailed tracing
|
||||
- Checkpoint inspection
|
||||
|
||||
### 4. Determinism
|
||||
|
||||
Given the same input and thread ID, execution produces identical results:
|
||||
|
||||
- Consistent ordering of parallel operations
|
||||
- Atomic state updates
|
||||
- Reliable checkpoint restoration
|
||||
|
||||
### 5. Extensibility
|
||||
|
||||
The framework is designed for extension:
|
||||
|
||||
- Custom channel types
|
||||
- Pluggable storage backends
|
||||
- Custom state schema formats
|
||||
- Integrations with other frameworks
|
||||
|
||||
## Implementation Invariants
|
||||
|
||||
These invariants are maintained and tested throughout the codebase:
|
||||
|
||||
### State Management Invariants
|
||||
|
||||
1. **Type Safety**: All state updates must conform to the schema
|
||||
2. **Atomic Updates**: State updates are all-or-nothing
|
||||
3. **State Isolation**: Updates are not visible until the end of a superstep
|
||||
4. **Schema Compatibility**: State schemas must be compatible with serialization
|
||||
|
||||
### Execution Invariants
|
||||
|
||||
1. **Deterministic Ordering**: Node execution order is consistent
|
||||
2. **Termination**: Execution always completes for valid graphs
|
||||
3. **Error Handling**: Node failures are handled gracefully
|
||||
4. **Checkpoint Fidelity**: Execution resumes correctly from checkpoints
|
||||
|
||||
### Channel Invariants
|
||||
|
||||
1. **Type Enforcement**: Channel values must match declared types
|
||||
2. **Update Validation**: Updates are validated before application
|
||||
3. **Serialization**: Channels must serialize/deserialize correctly
|
||||
4. **Behavior Consistency**: Each channel type must maintain its contract
|
||||
|
||||
## Reimplementation Guide
|
||||
|
||||
If reimplementing LangGraph from scratch, follow these steps:
|
||||
|
||||
1. **Start with State Schemas**: Implement the state validation system
|
||||
2. **Build Channel Types**: Create the basic channel implementations
|
||||
3. **Implement Pregel Core**: Build the execution engine
|
||||
4. **Add Checkpoint Support**: Implement persistence
|
||||
5. **Create StateGraph API**: Build the high-level interface
|
||||
6. **Add HITL Features**: Implement interruption/resumption
|
||||
|
||||
Challenging aspects:
|
||||
|
||||
- Maintaining determinism with parallel execution
|
||||
- Ensuring type safety across the system
|
||||
- Implementing efficient checkpointing
|
||||
- Managing complex state transitions
|
||||
|
||||
## Testing Strategy
|
||||
|
||||
LangGraph's test suite focuses on:
|
||||
|
||||
1. **Unit Tests**: For individual components
|
||||
2. **Integration Tests**: For component interactions
|
||||
3. **Property Tests**: For invariant verification
|
||||
4. **Snapshot Tests**: For regression prevention
|
||||
5. **Performance Tests**: For optimization
|
||||
|
||||
## Optimization and Performance
|
||||
|
||||
LangGraph includes several optimizations:
|
||||
|
||||
1. **Parallel Execution**: Nodes execute in parallel when possible
|
||||
2. **Lazy Checkpointing**: Only changed state is serialized
|
||||
3. **Channel-specific Optimizations**: Each channel type optimizes its pattern
|
||||
4. **Batched Operations**: Tasks are batched for efficiency
|
||||
5. **Memory Management**: Large states use specialized handling
|
||||
|
||||
## Security Considerations
|
||||
|
||||
When implementing or extending LangGraph, consider:
|
||||
|
||||
1. **Input Validation**: All external inputs must be validated
|
||||
2. **Serialization Safety**: Avoid security issues in serialization
|
||||
3. **Access Control**: Proper thread isolation to prevent data leakage
|
||||
4. **Resource Limits**: Prevent unbounded resource consumption
|
||||
5. **Secrets Management**: Avoid storing secrets in checkpoints
|
||||
|
||||
## Advanced Patterns
|
||||
|
||||
### Nested Graphs
|
||||
|
||||
```python
|
||||
# Main graph
|
||||
main_graph = StateGraph(MainState)
|
||||
|
||||
# Subgraph
|
||||
subgraph = StateGraph(SubState)
|
||||
subgraph.add_node("sub_process", sub_process)
|
||||
compiled_subgraph = subgraph.compile()
|
||||
|
||||
# Include subgraph in main graph
|
||||
main_graph.add_node("subprocess", compiled_subgraph)
|
||||
```
|
||||
|
||||
### Complex Routing
|
||||
|
||||
```python
|
||||
# Define routing logic
|
||||
def router(state: State) -> str:
|
||||
if state["error"]:
|
||||
return "error_handler"
|
||||
elif state["counter"] > 10:
|
||||
return "summarize"
|
||||
else:
|
||||
return "continue"
|
||||
|
||||
# Add conditional edges
|
||||
graph.add_conditional_edges("process", router)
|
||||
```
|
||||
|
||||
## Related Systems
|
||||
|
||||
LangGraph draws inspiration from and can be compared to:
|
||||
|
||||
1. **Airflow/Temporal**: Workflow orchestration systems
|
||||
2. **Actor Frameworks**: Like Akka and Ray
|
||||
3. **Stream Processing**: Systems like Apache Flink
|
||||
4. **State Machines**: Like XState and statecharts
|
||||
|
||||
Key differentiators:
|
||||
|
||||
- Not for DAGs, full support for cycles
|
||||
- Focus on LLM-specific workflows
|
||||
- Type-safe state management
|
||||
- Human-in-the-loop capabilities
|
||||
- Checkpoint-based persistence
|
||||
|
||||
## Conclusion
|
||||
|
||||
LangGraph's layered architecture provides a flexible, type-safe, and observable framework for building complex LLM applications. By understanding the components and their interactions, developers can leverage the full power of the framework while maintaining robust, maintainable code.
|
||||
|
||||
For implementation details on specific components, refer to the other specification documents:
|
||||
|
||||
- [Pregel.md](Pregel.md): Execution engine details
|
||||
- [Channels.md](Channels.md): Communication mechanism
|
||||
- [StateSchema.md](StateSchema.md): State definition
|
||||
- [StateGraph.md](StateGraph.md): High-level API
|
||||
- [CheckpointSystem.md](CheckpointSystem.md): Persistence
|
||||
- [HumanInTheLoop.md](HumanInTheLoop.md): Interruption/resumption
|
||||
@@ -0,0 +1,807 @@
|
||||
# Java Channel Interfaces
|
||||
|
||||
This document defines the Java interfaces for the channel system of LangGraph, closely aligned with the Python implementation.
|
||||
|
||||
## `Channel` Interface
|
||||
|
||||
The base interface for all channels, providing methods for getting values, applying updates, and checkpoint management.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
/**
|
||||
* Interface for communication channels between nodes in a graph.
|
||||
*/
|
||||
public interface Channel {
|
||||
/**
|
||||
* Get the current value of the channel.
|
||||
*
|
||||
* @return Current value
|
||||
*/
|
||||
Object getValue();
|
||||
|
||||
/**
|
||||
* Update the channel with a new value.
|
||||
*
|
||||
* @param value New value
|
||||
* @return True if the update was applied, false otherwise
|
||||
*/
|
||||
boolean update(Object value);
|
||||
|
||||
/**
|
||||
* Get the value to save in a checkpoint.
|
||||
*
|
||||
* @return Checkpointed value
|
||||
*/
|
||||
Object checkpoint();
|
||||
|
||||
/**
|
||||
* Restore the channel from a checkpoint.
|
||||
*
|
||||
* @param value Checkpointed value
|
||||
*/
|
||||
void fromCheckpoint(Object value);
|
||||
}
|
||||
```
|
||||
|
||||
## Channel Implementations
|
||||
|
||||
### `LastValue` Channel
|
||||
|
||||
A channel that stores a single value, replacing it with each update.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.Objects;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
/**
|
||||
* Channel that stores the last value set, rejecting multiple updates within a single step.
|
||||
* This is the default channel type for most use cases.
|
||||
*/
|
||||
public class LastValue implements Channel {
|
||||
private final AtomicReference<Object> value = new AtomicReference<>();
|
||||
private boolean updated = false;
|
||||
|
||||
/**
|
||||
* Create a LastValue channel with an optional initial value.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
*/
|
||||
public LastValue(Object initialValue) {
|
||||
value.set(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty LastValue channel.
|
||||
*/
|
||||
public LastValue() {
|
||||
this(null);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object getValue() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean update(Object newValue) {
|
||||
if (updated) {
|
||||
throw new IllegalStateException("LastValue channel cannot be updated multiple times in one step");
|
||||
}
|
||||
|
||||
// Skip update if value hasn't changed
|
||||
if (Objects.equals(value.get(), newValue)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
value.set(newValue);
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object checkpoint() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void fromCheckpoint(Object checkpointValue) {
|
||||
value.set(checkpointValue);
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the update flag at the end of a superstep.
|
||||
*/
|
||||
public void resetUpdated() {
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if the channel was updated in the current step.
|
||||
*
|
||||
* @return True if updated, false otherwise
|
||||
*/
|
||||
public boolean wasUpdated() {
|
||||
return updated;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### `AnyValue` Channel
|
||||
|
||||
A channel that accepts multiple updates within a step, storing only the last one.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.Objects;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
/**
|
||||
* Channel that accepts multiple updates within a step, storing only the last one.
|
||||
*/
|
||||
public class AnyValue implements Channel {
|
||||
private final AtomicReference<Object> value = new AtomicReference<>();
|
||||
private boolean updated = false;
|
||||
|
||||
/**
|
||||
* Create an AnyValue channel with an optional initial value.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
*/
|
||||
public AnyValue(Object initialValue) {
|
||||
value.set(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty AnyValue channel.
|
||||
*/
|
||||
public AnyValue() {
|
||||
this(null);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object getValue() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean update(Object newValue) {
|
||||
// Skip update if value hasn't changed
|
||||
if (Objects.equals(value.get(), newValue)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
value.set(newValue);
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object checkpoint() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void fromCheckpoint(Object checkpointValue) {
|
||||
value.set(checkpointValue);
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the update flag at the end of a superstep.
|
||||
*/
|
||||
public void resetUpdated() {
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if the channel was updated in the current step.
|
||||
*
|
||||
* @return True if updated, false otherwise
|
||||
*/
|
||||
public boolean wasUpdated() {
|
||||
return updated;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### `EphemeralValue` Channel
|
||||
|
||||
A channel that clears its value after being read.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.Objects;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
/**
|
||||
* Channel that clears its value after being read.
|
||||
* Useful for temporary values that should only be processed once.
|
||||
*/
|
||||
public class EphemeralValue implements Channel {
|
||||
private final AtomicReference<Object> value = new AtomicReference<>();
|
||||
private boolean updated = false;
|
||||
|
||||
/**
|
||||
* Create an EphemeralValue channel with an optional initial value.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
*/
|
||||
public EphemeralValue(Object initialValue) {
|
||||
value.set(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty EphemeralValue channel.
|
||||
*/
|
||||
public EphemeralValue() {
|
||||
this(null);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object getValue() {
|
||||
Object currentValue = value.getAndSet(null);
|
||||
return currentValue;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean update(Object newValue) {
|
||||
// Skip update if value hasn't changed
|
||||
if (Objects.equals(value.get(), newValue)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
value.set(newValue);
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object checkpoint() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void fromCheckpoint(Object checkpointValue) {
|
||||
value.set(checkpointValue);
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the update flag at the end of a superstep.
|
||||
*/
|
||||
public void resetUpdated() {
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if the channel was updated in the current step.
|
||||
*
|
||||
* @return True if updated, false otherwise
|
||||
*/
|
||||
public boolean wasUpdated() {
|
||||
return updated;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### `UntrackedValue` Channel
|
||||
|
||||
A channel that's excluded from checkpoints.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.Objects;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
/**
|
||||
* Channel that's excluded from checkpoints.
|
||||
* Useful for storing large temporary data that shouldn't be persisted.
|
||||
*/
|
||||
public class UntrackedValue implements Channel {
|
||||
private final AtomicReference<Object> value = new AtomicReference<>();
|
||||
private boolean updated = false;
|
||||
|
||||
/**
|
||||
* Create an UntrackedValue channel with an optional initial value.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
*/
|
||||
public UntrackedValue(Object initialValue) {
|
||||
value.set(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty UntrackedValue channel.
|
||||
*/
|
||||
public UntrackedValue() {
|
||||
this(null);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object getValue() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean update(Object newValue) {
|
||||
// Skip update if value hasn't changed
|
||||
if (Objects.equals(value.get(), newValue)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
value.set(newValue);
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object checkpoint() {
|
||||
// Return null for checkpoint as this value is not tracked
|
||||
return null;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void fromCheckpoint(Object checkpointValue) {
|
||||
// No-op as this channel isn't tracked in checkpoints
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the update flag at the end of a superstep.
|
||||
*/
|
||||
public void resetUpdated() {
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if the channel was updated in the current step.
|
||||
*
|
||||
* @return True if updated, false otherwise
|
||||
*/
|
||||
public boolean wasUpdated() {
|
||||
return updated;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### `Topic` Channel
|
||||
|
||||
A publish-subscribe channel supporting multiple values and subscribers.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
|
||||
/**
|
||||
* Publish-subscribe channel supporting multiple values and subscribers.
|
||||
* Each subscriber receives all published values.
|
||||
*/
|
||||
public class Topic implements Channel {
|
||||
private final List<Object> values = new CopyOnWriteArrayList<>();
|
||||
private boolean updated = false;
|
||||
|
||||
@Override
|
||||
public Object getValue() {
|
||||
List<Object> result = new ArrayList<>(values);
|
||||
values.clear();
|
||||
return result;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean update(Object value) {
|
||||
values.add(value);
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object checkpoint() {
|
||||
return new ArrayList<>(values);
|
||||
}
|
||||
|
||||
@Override
|
||||
@SuppressWarnings("unchecked")
|
||||
public void fromCheckpoint(Object checkpointValue) {
|
||||
values.clear();
|
||||
if (checkpointValue != null) {
|
||||
values.addAll((List<Object>) checkpointValue);
|
||||
}
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the update flag at the end of a superstep.
|
||||
*/
|
||||
public void resetUpdated() {
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if the channel was updated in the current step.
|
||||
*
|
||||
* @return True if updated, false otherwise
|
||||
*/
|
||||
public boolean wasUpdated() {
|
||||
return updated;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### `BinaryOperatorAggregate` Channel
|
||||
|
||||
A channel that aggregates values using a binary operator.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.Objects;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
import java.util.function.BinaryOperator;
|
||||
|
||||
/**
|
||||
* Channel that aggregates values using a binary operator.
|
||||
*
|
||||
* @param <T> Type of values to aggregate
|
||||
*/
|
||||
public class BinaryOperatorAggregate<T> implements Channel {
|
||||
private final AtomicReference<T> value = new AtomicReference<>();
|
||||
private final BinaryOperator<T> operator;
|
||||
private boolean updated = false;
|
||||
|
||||
/**
|
||||
* Create a BinaryOperatorAggregate channel with an operator and optional initial value.
|
||||
*
|
||||
* @param operator Binary operator for combining values
|
||||
* @param initialValue Optional initial value
|
||||
*/
|
||||
public BinaryOperatorAggregate(BinaryOperator<T> operator, T initialValue) {
|
||||
this.operator = operator;
|
||||
value.set(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create a BinaryOperatorAggregate channel with an operator.
|
||||
*
|
||||
* @param operator Binary operator for combining values
|
||||
*/
|
||||
public BinaryOperatorAggregate(BinaryOperator<T> operator) {
|
||||
this(operator, null);
|
||||
}
|
||||
|
||||
@Override
|
||||
@SuppressWarnings("unchecked")
|
||||
public T getValue() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
@SuppressWarnings("unchecked")
|
||||
public boolean update(Object newValue) {
|
||||
T typedValue = (T) newValue;
|
||||
T currentValue = value.get();
|
||||
|
||||
if (currentValue == null) {
|
||||
value.set(typedValue);
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
// Apply the binary operator to combine values
|
||||
T combinedValue = operator.apply(currentValue, typedValue);
|
||||
|
||||
// Skip update if value hasn't changed
|
||||
if (Objects.equals(currentValue, combinedValue)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
value.set(combinedValue);
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object checkpoint() {
|
||||
return value.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
@SuppressWarnings("unchecked")
|
||||
public void fromCheckpoint(Object checkpointValue) {
|
||||
value.set((T) checkpointValue);
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the update flag at the end of a superstep.
|
||||
*/
|
||||
public void resetUpdated() {
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if the channel was updated in the current step.
|
||||
*
|
||||
* @return True if updated, false otherwise
|
||||
*/
|
||||
public boolean wasUpdated() {
|
||||
return updated;
|
||||
}
|
||||
|
||||
/**
|
||||
* Factory for creating sum aggregates.
|
||||
*
|
||||
* @param <T> Type of values to sum
|
||||
* @return A channel that sums values
|
||||
*/
|
||||
public static <T extends Number> BinaryOperatorAggregate<T> sum() {
|
||||
return new BinaryOperatorAggregate<>((a, b) -> {
|
||||
if (a instanceof Integer) {
|
||||
return (T) Integer.valueOf(((Integer) a) + ((Integer) b));
|
||||
} else if (a instanceof Long) {
|
||||
return (T) Long.valueOf(((Long) a) + ((Long) b));
|
||||
} else if (a instanceof Double) {
|
||||
return (T) Double.valueOf(((Double) a) + ((Double) b));
|
||||
} else if (a instanceof Float) {
|
||||
return (T) Float.valueOf(((Float) a) + ((Float) b));
|
||||
} else {
|
||||
throw new IllegalArgumentException("Unsupported number type: " + a.getClass());
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### `NamedBarrierValue` Channel
|
||||
|
||||
A synchronization mechanism requiring all named values to be received.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.HashMap;
|
||||
import java.util.HashSet;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
/**
|
||||
* Synchronization channel requiring all named values to be received.
|
||||
* Triggers when all expected names have provided values.
|
||||
*/
|
||||
public class NamedBarrierValue implements Channel {
|
||||
private final Set<String> expectedNames;
|
||||
private final Map<String, Object> values = new ConcurrentHashMap<>();
|
||||
private boolean updated = false;
|
||||
|
||||
/**
|
||||
* Create a NamedBarrierValue channel with expected names.
|
||||
*
|
||||
* @param expectedNames Set of names expected to provide values
|
||||
*/
|
||||
public NamedBarrierValue(Set<String> expectedNames) {
|
||||
this.expectedNames = new HashSet<>(expectedNames);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object getValue() {
|
||||
// Return the map of collected values if all expected names have provided values
|
||||
if (values.keySet().containsAll(expectedNames)) {
|
||||
Map<String, Object> result = new HashMap<>(values);
|
||||
values.clear();
|
||||
return result;
|
||||
}
|
||||
|
||||
// Return null if the barrier is not satisfied
|
||||
return null;
|
||||
}
|
||||
|
||||
@Override
|
||||
@SuppressWarnings("unchecked")
|
||||
public boolean update(Object value) {
|
||||
if (!(value instanceof Map)) {
|
||||
throw new IllegalArgumentException("NamedBarrierValue requires a Map<String, Object> update");
|
||||
}
|
||||
|
||||
Map<String, Object> update = (Map<String, Object>) value;
|
||||
if (update.size() != 1) {
|
||||
throw new IllegalArgumentException("NamedBarrierValue update must contain exactly one entry");
|
||||
}
|
||||
|
||||
String name = update.keySet().iterator().next();
|
||||
if (!expectedNames.contains(name)) {
|
||||
throw new IllegalArgumentException("Unexpected name in NamedBarrierValue update: " + name);
|
||||
}
|
||||
|
||||
values.put(name, update.get(name));
|
||||
updated = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Object checkpoint() {
|
||||
return new HashMap<>(values);
|
||||
}
|
||||
|
||||
@Override
|
||||
@SuppressWarnings("unchecked")
|
||||
public void fromCheckpoint(Object checkpointValue) {
|
||||
values.clear();
|
||||
if (checkpointValue != null) {
|
||||
values.putAll((Map<String, Object>) checkpointValue);
|
||||
}
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the update flag at the end of a superstep.
|
||||
*/
|
||||
public void resetUpdated() {
|
||||
updated = false;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if the channel was updated in the current step.
|
||||
*
|
||||
* @return True if updated, false otherwise
|
||||
*/
|
||||
public boolean wasUpdated() {
|
||||
return updated;
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if all expected names have provided values.
|
||||
*
|
||||
* @return True if the barrier is satisfied, false otherwise
|
||||
*/
|
||||
public boolean isBarrierSatisfied() {
|
||||
return values.keySet().containsAll(expectedNames);
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Channel Factory
|
||||
|
||||
A factory class for creating channels.
|
||||
|
||||
```java
|
||||
package com.langgraph.channels;
|
||||
|
||||
import java.util.Set;
|
||||
import java.util.function.BinaryOperator;
|
||||
|
||||
/**
|
||||
* Factory for creating channels.
|
||||
*/
|
||||
public final class Channels {
|
||||
private Channels() {}
|
||||
|
||||
/**
|
||||
* Create a LastValue channel.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
* @return LastValue channel
|
||||
*/
|
||||
public static LastValue lastValue(Object initialValue) {
|
||||
return new LastValue(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty LastValue channel.
|
||||
*
|
||||
* @return LastValue channel
|
||||
*/
|
||||
public static LastValue lastValue() {
|
||||
return new LastValue();
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an AnyValue channel.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
* @return AnyValue channel
|
||||
*/
|
||||
public static AnyValue anyValue(Object initialValue) {
|
||||
return new AnyValue(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty AnyValue channel.
|
||||
*
|
||||
* @return AnyValue channel
|
||||
*/
|
||||
public static AnyValue anyValue() {
|
||||
return new AnyValue();
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an EphemeralValue channel.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
* @return EphemeralValue channel
|
||||
*/
|
||||
public static EphemeralValue ephemeralValue(Object initialValue) {
|
||||
return new EphemeralValue(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty EphemeralValue channel.
|
||||
*
|
||||
* @return EphemeralValue channel
|
||||
*/
|
||||
public static EphemeralValue ephemeralValue() {
|
||||
return new EphemeralValue();
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an UntrackedValue channel.
|
||||
*
|
||||
* @param initialValue Optional initial value
|
||||
* @return UntrackedValue channel
|
||||
*/
|
||||
public static UntrackedValue untrackedValue(Object initialValue) {
|
||||
return new UntrackedValue(initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create an empty UntrackedValue channel.
|
||||
*
|
||||
* @return UntrackedValue channel
|
||||
*/
|
||||
public static UntrackedValue untrackedValue() {
|
||||
return new UntrackedValue();
|
||||
}
|
||||
|
||||
/**
|
||||
* Create a Topic channel.
|
||||
*
|
||||
* @return Topic channel
|
||||
*/
|
||||
public static Topic topic() {
|
||||
return new Topic();
|
||||
}
|
||||
|
||||
/**
|
||||
* Create a BinaryOperatorAggregate channel.
|
||||
*
|
||||
* @param operator Binary operator for combining values
|
||||
* @param initialValue Optional initial value
|
||||
* @param <T> Type of values to aggregate
|
||||
* @return BinaryOperatorAggregate channel
|
||||
*/
|
||||
public static <T> BinaryOperatorAggregate<T> binaryOperatorAggregate(
|
||||
BinaryOperator<T> operator, T initialValue) {
|
||||
return new BinaryOperatorAggregate<>(operator, initialValue);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create a BinaryOperatorAggregate channel.
|
||||
*
|
||||
* @param operator Binary operator for combining values
|
||||
* @param <T> Type of values to aggregate
|
||||
* @return BinaryOperatorAggregate channel
|
||||
*/
|
||||
public static <T> BinaryOperatorAggregate<T> binaryOperatorAggregate(BinaryOperator<T> operator) {
|
||||
return new BinaryOperatorAggregate<>(operator);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create a NamedBarrierValue channel.
|
||||
*
|
||||
* @param expectedNames Set of names expected to provide values
|
||||
* @return NamedBarrierValue channel
|
||||
*/
|
||||
public static NamedBarrierValue namedBarrierValue(Set<String> expectedNames) {
|
||||
return new NamedBarrierValue(expectedNames);
|
||||
}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,153 @@
|
||||
# Java Checkpoint Interfaces
|
||||
|
||||
This document defines the Java interfaces for the checkpoint layer of LangGraph, aligned with the Python implementation.
|
||||
|
||||
## `BaseCheckpointSaver` Interface
|
||||
|
||||
The `BaseCheckpointSaver` interface provides methods for creating, loading, and managing checkpoints.
|
||||
|
||||
```java
|
||||
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);
|
||||
}
|
||||
```
|
||||
|
||||
## `ID` Utility
|
||||
|
||||
A utility class for generating deterministic IDs, similar to the Python implementation.
|
||||
|
||||
```java
|
||||
package com.langgraph.checkpoint.base;
|
||||
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.UUID;
|
||||
import java.security.MessageDigest;
|
||||
import java.security.NoSuchAlgorithmException;
|
||||
import java.util.Base64;
|
||||
|
||||
/**
|
||||
* Utility class for generating 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,682 @@
|
||||
# Java Port Implementation Plan for LangGraph
|
||||
|
||||
This document outlines the plan for implementing a Java port of LangGraph, following the specifications defined in the other documents.
|
||||
|
||||
## Project Structure
|
||||
|
||||
```
|
||||
langgraph-java/
|
||||
├── build.gradle
|
||||
├── settings.gradle
|
||||
├── README.md
|
||||
├── langgraph-checkpoint/ # Base persistence interfaces (equiv. to libs/checkpoint)
|
||||
│ └── src/
|
||||
│ ├── main/java/
|
||||
│ │ └── com/langgraph/checkpoint/
|
||||
│ │ ├── base/ # Base interfaces and utilities
|
||||
│ │ └── serde/ # Serialization
|
||||
│ └── test/java/
|
||||
├── langgraph-core/ # Main library (equiv. to libs/langgraph)
|
||||
│ └── src/
|
||||
│ ├── main/java/
|
||||
│ │ └── com/langgraph/
|
||||
│ │ ├── channels/ # Channel implementations
|
||||
│ │ ├── graph/ # StateGraph API
|
||||
│ │ └── pregel/ # Pregel execution engine
|
||||
│ └── test/java/
|
||||
├── langgraph-examples/ # Example applications
|
||||
│ └── src/
|
||||
│ └── main/java/
|
||||
│ └── com/langgraph/examples/
|
||||
└── langgraph-checkpoint-postgres/ # Postgres implementation (future)
|
||||
```
|
||||
|
||||
## Development Phases
|
||||
|
||||
The implementation will follow a bottom-up approach, starting with the lowest-level components and building up to the high-level APIs.
|
||||
|
||||
### Phase 1: Foundation (langgraph-checkpoint)
|
||||
|
||||
1. **Set up project structure**
|
||||
|
||||
- Initialize Gradle project
|
||||
- Configure dependencies
|
||||
- Set up testing framework (JUnit 5)
|
||||
|
||||
2. **Implement ID utilities**
|
||||
|
||||
- `ID.java` - For deterministic ID generation
|
||||
- Test with different inputs
|
||||
|
||||
3. **Implement serialization framework**
|
||||
|
||||
- `Serializer` interface
|
||||
- `ReflectionSerializer` interface
|
||||
- `TypeSerializer` and `TypeDeserializer` interfaces
|
||||
- `MsgPackSerializer` implementation
|
||||
- Comprehensive tests for serialization/deserialization
|
||||
|
||||
4. **Implement checkpoint interfaces**
|
||||
- `BaseCheckpointSaver` interface
|
||||
- `AsyncBaseCheckpointSaver` interface
|
||||
- `MemoryCheckpointSaver` implementation
|
||||
- Tests for checkpoint operations
|
||||
|
||||
### Phase 2: State Management (langgraph-core)
|
||||
|
||||
1. **Implement channel system**
|
||||
|
||||
- `Channel` interface
|
||||
- Basic channel implementations:
|
||||
- `LastValue`
|
||||
- `AnyValue`
|
||||
- `EphemeralValue`
|
||||
- `UntrackedValue`
|
||||
- Advanced channel implementations:
|
||||
- `Topic`
|
||||
- `BinaryOperatorAggregate`
|
||||
- `NamedBarrierValue`
|
||||
- Tests for each channel type
|
||||
|
||||
2. **Implement schema validation**
|
||||
- `SchemaValidator` interface
|
||||
- `RecordSchemaValidator` implementation
|
||||
- Tests for schema validation
|
||||
|
||||
### Phase 3: Execution Engine (langgraph-core)
|
||||
|
||||
1. **Implement core Pregel interfaces**
|
||||
|
||||
- `PregelProtocol` interface
|
||||
- `StreamMode` enum
|
||||
- `PregelExecutable` interface
|
||||
- `PregelNode` class
|
||||
- `PregelTask` classes
|
||||
- Tests for interfaces
|
||||
|
||||
2. **Implement execution system**
|
||||
- `RetryPolicy` interface
|
||||
- `Checkpoint` class
|
||||
- `Pregel` implementation
|
||||
- Tests for execution
|
||||
|
||||
### Phase 4: High-Level API (langgraph-core)
|
||||
|
||||
1. **Implement StateGraph**
|
||||
|
||||
- `GraphConstants` constants
|
||||
- `NodeAction` and `EdgeCondition` interfaces
|
||||
- `StateGraph` class
|
||||
- `CompiledStateGraph` class
|
||||
- Tests for graph construction and execution
|
||||
|
||||
2. **Build examples**
|
||||
- Counter example
|
||||
- Conversation example
|
||||
- Multi-step workflow example
|
||||
|
||||
### Phase 5: Extensions (future)
|
||||
|
||||
1. **Implement Postgres adaptors**
|
||||
- PostgresCheckpointSaver
|
||||
- Tests for Postgres integration
|
||||
|
||||
## Testing Strategy
|
||||
|
||||
Each component will be implemented following Test-Driven Development (TDD):
|
||||
|
||||
1. Write test cases based on Python implementation
|
||||
2. Implement interface to satisfy tests
|
||||
3. Implement concrete class(es)
|
||||
4. Run tests and refine implementation
|
||||
5. Document completed component
|
||||
|
||||
### Key Test Areas
|
||||
|
||||
- **Serialization**: Test with various object types, nested structures, and cycles
|
||||
- **Channels**: Test all channel behaviors and edge cases
|
||||
- **Pregel**: Test deterministic execution, error handling, and checkpoint integration
|
||||
- **StateGraph**: Test graph construction, validation, and execution
|
||||
|
||||
## Implementation Notes
|
||||
|
||||
### Java-Specific Adaptations
|
||||
|
||||
1. **Records for state schemas**
|
||||
|
||||
- Java Records (Java 14+) as TypedDict/Pydantic alternative
|
||||
- Reflection-based validation and conversion
|
||||
|
||||
2. **CompletableFuture for async**
|
||||
|
||||
- Async interfaces use CompletableFuture
|
||||
- Parallel execution with work-stealing pools
|
||||
|
||||
3. **Generics for type safety**
|
||||
|
||||
- Extensive use of generics for type safety
|
||||
- Runtime type checking for dynamic aspects
|
||||
|
||||
4. **Builder pattern**
|
||||
- Builder pattern for complex object construction
|
||||
- Fluent interfaces for API usability
|
||||
|
||||
### Performance Considerations
|
||||
|
||||
- Thread-safe implementations using concurrent collections
|
||||
- Atomic operations for parallel safety
|
||||
- Minimizing object creation and copying
|
||||
- Efficient serialization with MessagePack
|
||||
|
||||
## Initial Implementation Tasks
|
||||
|
||||
1. Set up project structure and build system
|
||||
2. Implement ID utilities in checkpoint module
|
||||
3. Implement MsgPackSerializer
|
||||
4. Implement MemoryCheckpointSaver
|
||||
5. Write comprehensive tests for foundation layer
|
||||
|
||||
## Timeline
|
||||
|
||||
1. **Phase 1: Foundation** - 2 weeks
|
||||
2. **Phase 2: State Management** - 2 weeks
|
||||
3. **Phase 3: Execution Engine** - 3 weeks
|
||||
4. **Phase 4: High-Level API** - 2 weeks
|
||||
5. **Integration and Examples** - 1 week
|
||||
|
||||
Total estimated time: ~10 weeks for core functionality
|
||||
|
||||
## Dependencies
|
||||
|
||||
- **MessagePack**: `org.msgpack:msgpack-core:0.9.3`
|
||||
- **JUnit 5**: `org.junit.jupiter:junit-jupiter:5.8.2`
|
||||
- **Mockito**: `org.mockito:mockito-core:4.5.1`
|
||||
- **AssertJ**: `org.assertj:assertj-core:3.22.0`
|
||||
|
||||
## Phase 3 in detail:
|
||||
|
||||
This implementation plan is organized to match the structure of the Python implementation while prioritizing testability at each step. Don't forget to use Java patterns and idioms where appropriate.
|
||||
|
||||
1. Core Data Structures & Interfaces (Week 1)
|
||||
|
||||
Day 1-2: Base Interfaces and Types
|
||||
|
||||
1. StreamMode Enum
|
||||
- Implementation: Define values (VALUES, UPDATES, DEBUG)
|
||||
- Test: Verify serialization and string representation
|
||||
2. PregelExecutable Interface
|
||||
- Implementation: Define execute method
|
||||
- Test: Create mock implementations for testing
|
||||
3. RetryPolicy Interface & Implementations
|
||||
- Implementation: RetryPolicy interface with factory methods
|
||||
- Test: Verify retry decision logic
|
||||
|
||||
Day 3-4: Task Management
|
||||
|
||||
1. PregelTask Class
|
||||
- Implementation: Node name, trigger, retry policy
|
||||
- Test: Constructor, getters, equals, hashCode
|
||||
2. PregelExecutableTask Class
|
||||
- Implementation: Task with inputs and context
|
||||
- Test: Construction, input/context immutability
|
||||
|
||||
Day 5: Core Protocol
|
||||
|
||||
1. PregelProtocol Interface
|
||||
|
||||
- Implementation: Define methods from spec
|
||||
- Test: Mock implementation for testing
|
||||
|
||||
2. Node & Channel Integration (Week 2)
|
||||
|
||||
Day 1-2: Node Implementation
|
||||
|
||||
1. PregelNode Class
|
||||
- Implementation: Full implementation with accessors
|
||||
- Test: Subscription, trigger, write permissions
|
||||
2. NodeRegistry
|
||||
- Implementation: Managing collections of nodes
|
||||
- Test: Registration, lookup, validation
|
||||
|
||||
Day 3-5: Message Handling
|
||||
|
||||
1. Checkpoint Class
|
||||
|
||||
- Implementation: State snapshot storage
|
||||
- Test: Capture and restore state
|
||||
|
||||
2. ChannelRegistry
|
||||
|
||||
- Implementation: Managing channel collections
|
||||
- Test: Registration, lookup, validation
|
||||
|
||||
3. Core Execution Components (Week 3)
|
||||
|
||||
Day 1-2: Task Planning
|
||||
|
||||
1. PregelTaskPlanner
|
||||
- Implementation: Determine tasks to execute based on updates
|
||||
- Test: Task selection with different update patterns
|
||||
2. TaskPrioritizer
|
||||
- Implementation: Order tasks for execution
|
||||
- Test: Priority ordering with different dependency patterns
|
||||
|
||||
Day 3-5: Task Execution
|
||||
|
||||
1. TaskExecutor
|
||||
- Implementation: Execute tasks with retry logic
|
||||
- Test: Successful execution, error handling, retry behavior
|
||||
2. ExecutionContext
|
||||
|
||||
- Implementation: Thread-local context for execution
|
||||
- Test: Context propagation, thread safety
|
||||
|
||||
3. Superstep Management (Week 4)
|
||||
|
||||
Day 1-2: Superstep Core
|
||||
|
||||
1. SuperstepManager
|
||||
- Implementation: Manage a single superstep execution
|
||||
- Test: Plan, execute, update phases with mock nodes
|
||||
2. UpdateCollector
|
||||
- Implementation: Collect and apply channel updates
|
||||
- Test: Update ordering, conflict resolution
|
||||
|
||||
Day 3-5: Execution Loop
|
||||
|
||||
1. PregelLoop
|
||||
- Implementation: Core execution logic, step iteration
|
||||
- Test: Loop termination, state tracking
|
||||
2. CheckpointManager
|
||||
|
||||
- Implementation: Integration with checkpointing
|
||||
- Test: Checkpoint captures, restore behavior
|
||||
|
||||
3. Full Engine & Streaming (Week 5)
|
||||
|
||||
Day 1-3: Pregel Engine
|
||||
|
||||
1. Pregel Class
|
||||
- Implementation: Core engine with all components
|
||||
- Test: End-to-end execution, configuration
|
||||
2. PregelBuilder
|
||||
- Implementation: Fluent builder interface
|
||||
- Test: Configuration options, validation
|
||||
|
||||
Day 4-5: Streaming Support
|
||||
|
||||
1. StreamOutput
|
||||
- Implementation: Format output for streaming
|
||||
- Test: Different stream modes
|
||||
2. StreamController
|
||||
- Implementation: Manage streaming state
|
||||
- Test: Backpressure, cancellation
|
||||
|
||||
Implementation Strategy
|
||||
|
||||
1. Incremental Testing
|
||||
|
||||
Create test classes for each component that can be used in isolation:
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testTaskPlanning() {
|
||||
// Create mock channels with updates
|
||||
Map<String, Channel> channels = createMockChannels(
|
||||
Map.of("input", true, "other", false));
|
||||
|
||||
// Create nodes with subscriptions
|
||||
Set<PregelNode> nodes = createTestNodes();
|
||||
|
||||
// Create planner
|
||||
PregelTaskPlanner planner = new PregelTaskPlanner(nodes);
|
||||
|
||||
// Test planning logic
|
||||
List<PregelTask> tasks = planner.plan(channels);
|
||||
|
||||
// Verify correct tasks selected
|
||||
assertThat(tasks).hasSize(1);
|
||||
assertThat(tasks.get(0).getNode()).isEqualTo("processor");
|
||||
}
|
||||
```
|
||||
|
||||
2. Test Each Component in Isolation
|
||||
|
||||
For each component, test:
|
||||
|
||||
- Normal operation
|
||||
- Edge cases
|
||||
- Error conditions
|
||||
- Integration with dependencies
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testTaskExecution() {
|
||||
// Create mock executor
|
||||
TaskExecutor executor = new TaskExecutor();
|
||||
|
||||
// Create task with expected inputs
|
||||
PregelExecutableTask task = createTestTask();
|
||||
|
||||
// Execute and capture results
|
||||
Map<String, Object> result = executor.execute(task);
|
||||
|
||||
// Verify results
|
||||
assertThat(result)
|
||||
.containsKey("output")
|
||||
.containsEntry("output", "processed");
|
||||
}
|
||||
```
|
||||
|
||||
3. Incremental Integration
|
||||
|
||||
1. Start with simplest components: PregelTask, StreamMode, etc.
|
||||
1. Build TaskPlanner with mocked nodes
|
||||
1. Create TaskExecutor with mocked actions
|
||||
1. Integrate into SuperstepManager
|
||||
1. Combine in PregelLoop
|
||||
1. Build complete Pregel engine
|
||||
|
||||
```java
|
||||
// First, test task planner alone
|
||||
@Test
|
||||
void testTaskPlannerInIsolation() {
|
||||
PregelTaskPlanner planner = new PregelTaskPlanner(mockNodes);
|
||||
List<PregelTask> tasks = planner.plan(updatedChannels);
|
||||
// Verify tasks
|
||||
}
|
||||
|
||||
// Then, test executor alone
|
||||
@Test
|
||||
void testTaskExecutorInIsolation() {
|
||||
TaskExecutor executor = new TaskExecutor();
|
||||
Map<String, Object> result = executor.execute(mockTask);
|
||||
// Verify result
|
||||
}
|
||||
|
||||
// Finally, test them together in SuperstepManager
|
||||
@Test
|
||||
void testSuperstepIntegration() {
|
||||
SuperstepManager manager = new SuperstepManager(
|
||||
planner,
|
||||
executor,
|
||||
channels
|
||||
);
|
||||
SuperstepResult result = manager.executeStep();
|
||||
// Verify complete superstep behavior
|
||||
}
|
||||
```
|
||||
|
||||
4. Use Real Components When Possible
|
||||
|
||||
1. Use real Channel implementations from previous phase
|
||||
1. Create simple test PregelExecutables
|
||||
1. Build test workflows of increasing complexity
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testSimpleWorkflow() {
|
||||
// Create real channels
|
||||
Map<String, Channel> channels = new HashMap<>();
|
||||
channels.put("input", new LastValue<>(String.class));
|
||||
channels.put("output", new LastValue<>(String.class));
|
||||
|
||||
// Create real nodes
|
||||
Map<String, PregelNode> nodes = new HashMap<>();
|
||||
nodes.put("processor", new PregelNode(
|
||||
"processor",
|
||||
(inputs, context) -> {
|
||||
String input = (String) inputs.get("input");
|
||||
return Map.of("output", input.toUpperCase());
|
||||
},
|
||||
Set.of("input"),
|
||||
null,
|
||||
Set.of("output"),
|
||||
null
|
||||
));
|
||||
|
||||
// Create real Pregel instance
|
||||
Pregel pregel = new Pregel(nodes, channels, null);
|
||||
|
||||
// Run and verify
|
||||
Object result = pregel.invoke(Map.of("input", "hello"), null);
|
||||
|
||||
// Verify complete execution
|
||||
@SuppressWarnings("unchecked")
|
||||
Map<String, Object> resultMap = (Map<String, Object>) result;
|
||||
assertThat(resultMap).containsEntry("output", "HELLO");
|
||||
}
|
||||
```
|
||||
|
||||
File Organization
|
||||
|
||||
Based on the Python structure, here's how the Java implementation will be organized:
|
||||
|
||||
com.langgraph.pregel/
|
||||
├── PregelProtocol.java # Core interface
|
||||
├── StreamMode.java # Enum for streaming options
|
||||
├── PregelExecutable.java # Interface for node functions
|
||||
├── PregelNode.java # Node definition
|
||||
├── task/
|
||||
│ ├── PregelTask.java # Task representation
|
||||
│ ├── PregelExecutableTask.java # Task with inputs
|
||||
│ ├── TaskPlanner.java # Task planning logic
|
||||
│ └── TaskExecutor.java # Task execution
|
||||
├── state/
|
||||
│ ├── Checkpoint.java # State checkpoint
|
||||
│ ├── ChannelRegistry.java # Channel management
|
||||
│ └── NodeRegistry.java # Node management
|
||||
├── execute/
|
||||
│ ├── SuperstepManager.java # Single superstep execution
|
||||
│ ├── PregelLoop.java # Main execution loop
|
||||
│ ├── ExecutionContext.java # Context for execution
|
||||
│ └── UpdateCollector.java # Collect updates
|
||||
├── stream/
|
||||
│ ├── StreamController.java # Manage streaming
|
||||
│ └── StreamOutput.java # Format output
|
||||
├── retry/
|
||||
│ ├── RetryPolicy.java # Retry interface
|
||||
│ └── RetryPolicies.java # Standard policies
|
||||
└── Pregel.java # Main implementation
|
||||
|
||||
Testing Step-by-Step
|
||||
|
||||
The testing strategy follows a specific progression:
|
||||
|
||||
1. Unit Testing: Test each component in isolation
|
||||
2. Component Testing: Test related components together
|
||||
3. Integration Testing: Test main subsystems
|
||||
4. System Testing: Test complete workflows
|
||||
|
||||
Example Test Progression for TaskPlanner:
|
||||
|
||||
1. Unit Test: Mock everything
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testPlannerWithMocks() {
|
||||
Set<String> updatedChannels = Set.of("input");
|
||||
Map<String, PregelNode> mockNodes = createMockNodes();
|
||||
|
||||
TaskPlanner planner = new TaskPlanner(mockNodes);
|
||||
List<PregelTask> tasks = planner.plan(updatedChannels);
|
||||
|
||||
// Test with various update patterns
|
||||
}
|
||||
```
|
||||
|
||||
2. Component Test: Use real nodes, mock channels
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testPlannerWithRealNodes() {
|
||||
Set<String> updatedChannels = Set.of("input");
|
||||
Map<String, PregelNode> realNodes = createRealNodes();
|
||||
|
||||
TaskPlanner planner = new TaskPlanner(realNodes);
|
||||
List<PregelTask> tasks = planner.plan(updatedChannels);
|
||||
|
||||
// Verify with real node behavior
|
||||
}
|
||||
```
|
||||
|
||||
3. Integration Test: Use real nodes and channels
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testPlannerIntegration() {
|
||||
Map<String, Channel> channels = createRealChannels();
|
||||
// Update channels
|
||||
channels.get("input").update("test");
|
||||
|
||||
Map<String, PregelNode> nodes = createRealNodes();
|
||||
|
||||
// Get updated channel names
|
||||
Set<String> updatedChannels = getUpdatedChannelNames(channels);
|
||||
|
||||
TaskPlanner planner = new TaskPlanner(nodes);
|
||||
List<PregelTask> tasks = planner.plan(updatedChannels);
|
||||
|
||||
// Verify end-to-end planning
|
||||
}
|
||||
```
|
||||
|
||||
4. System Test: Use in a full Pregel execution
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testPlannerInFullSystem() {
|
||||
// Set up complete Pregel system
|
||||
Pregel pregel = createTestPregelSystem();
|
||||
|
||||
// Execute a workflow that will trigger planning
|
||||
pregel.invoke(Map.of("input", "test"), null);
|
||||
|
||||
// Verify entire execution via output
|
||||
}
|
||||
```
|
||||
|
||||
Sample Test Case Implementations
|
||||
|
||||
To illustrate the TDD approach, here are key test cases for early components:
|
||||
|
||||
1. PregelTask
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testPregelTask() {
|
||||
// Basic construction
|
||||
PregelTask task = new PregelTask("node1", "trigger1", RetryPolicy.noRetry());
|
||||
|
||||
assertThat(task.getNode()).isEqualTo("node1");
|
||||
assertThat(task.getTrigger()).isEqualTo("trigger1");
|
||||
assertThat(task.getRetryPolicy()).isNotNull();
|
||||
|
||||
// Equality
|
||||
PregelTask sameTask = new PregelTask("node1", "trigger1", RetryPolicy.maxAttempts(3));
|
||||
PregelTask differentNode = new PregelTask("node2", "trigger1", RetryPolicy.noRetry());
|
||||
PregelTask differentTrigger = new PregelTask("node1", "trigger2", RetryPolicy.noRetry());
|
||||
|
||||
assertThat(task).isEqualTo(sameTask);
|
||||
assertThat(task).isNotEqualTo(differentNode);
|
||||
assertThat(task).isNotEqualTo(differentTrigger);
|
||||
}
|
||||
```
|
||||
|
||||
2. PregelNode
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testPregelNode() {
|
||||
// Create a simple action
|
||||
PregelExecutable action = (inputs, context) -> Map.of("output", "result");
|
||||
|
||||
// Basic construction
|
||||
PregelNode node = new PregelNode(
|
||||
"processor",
|
||||
action,
|
||||
Set.of("input1", "input2"),
|
||||
"trigger1",
|
||||
Set.of("output1", "output2"),
|
||||
RetryPolicy.maxAttempts(3)
|
||||
);
|
||||
|
||||
// Test properties
|
||||
assertThat(node.getName()).isEqualTo("processor");
|
||||
assertThat(node.getAction()).isSameAs(action);
|
||||
assertThat(node.getSubscribe()).containsExactlyInAnyOrder("input1", "input2");
|
||||
assertThat(node.getTrigger()).isEqualTo("trigger1");
|
||||
assertThat(node.getWriters()).containsExactlyInAnyOrder("output1", "output2");
|
||||
assertThat(node.getRetryPolicy()).isNotNull();
|
||||
|
||||
// Test helper methods
|
||||
assertThat(node.subscribesTo("input1")).isTrue();
|
||||
assertThat(node.subscribesTo("input3")).isFalse();
|
||||
assertThat(node.hasTrigger("trigger1")).isTrue();
|
||||
assertThat(node.hasTrigger("trigger2")).isFalse();
|
||||
assertThat(node.canWriteTo("output1")).isTrue();
|
||||
assertThat(node.canWriteTo("output3")).isFalse();
|
||||
}
|
||||
```
|
||||
|
||||
3. TaskPlanner
|
||||
|
||||
```java
|
||||
@Test
|
||||
void testTaskPlanner() {
|
||||
// Create nodes
|
||||
PregelNode node1 = new PregelNode(
|
||||
"node1",
|
||||
(inputs, context) -> Map.of(),
|
||||
Set.of("channel1"),
|
||||
null,
|
||||
Set.of("output1"),
|
||||
null
|
||||
);
|
||||
|
||||
PregelNode node2 = new PregelNode(
|
||||
"node2",
|
||||
(inputs, context) -> Map.of(),
|
||||
Set.of("channel2"),
|
||||
null,
|
||||
Set.of("output2"),
|
||||
null
|
||||
);
|
||||
|
||||
PregelNode node3 = new PregelNode(
|
||||
"node3",
|
||||
(inputs, context) -> Map.of(),
|
||||
null,
|
||||
"trigger1",
|
||||
Set.of("output3"),
|
||||
null
|
||||
);
|
||||
|
||||
Map<String, PregelNode> nodes = Map.of(
|
||||
"node1", node1,
|
||||
"node2", node2,
|
||||
"node3", node3
|
||||
);
|
||||
|
||||
TaskPlanner planner = new TaskPlanner(nodes);
|
||||
|
||||
// Test with different updated channels
|
||||
Set<String> update1 = Set.of("channel1");
|
||||
List<PregelTask> tasks1 = planner.plan(update1);
|
||||
assertThat(tasks1).hasSize(1);
|
||||
assertThat(tasks1.get(0).getNode()).isEqualTo("node1");
|
||||
|
||||
Set<String> update2 = Set.of("channel1", "channel2");
|
||||
List<PregelTask> tasks2 = planner.plan(update2);
|
||||
assertThat(tasks2).hasSize(2);
|
||||
|
||||
Set<String> update3 = Set.of("trigger1");
|
||||
List<PregelTask> tasks3 = planner.plan(update3);
|
||||
assertThat(tasks3).hasSize(1);
|
||||
assertThat(tasks3.get(0).getNode()).isEqualTo("node3");
|
||||
|
||||
Set<String> update4 = Set.of("channel3");
|
||||
List<PregelTask> tasks4 = planner.plan(update4);
|
||||
assertThat(tasks4).isEmpty();
|
||||
}
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,520 @@
|
||||
# Java Serialization Interfaces
|
||||
|
||||
This document defines the Java interfaces for the serialization layer of LangGraph, closely aligned with the Python implementation.
|
||||
|
||||
## `Serializer` Interface
|
||||
|
||||
The base serializer interface providing methods for serializing and deserializing objects.
|
||||
|
||||
```java
|
||||
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);
|
||||
}
|
||||
```
|
||||
|
||||
## `ReflectionSerializer` Interface
|
||||
|
||||
A specialized serializer that can handle arbitrary Java objects by using reflection.
|
||||
|
||||
```java
|
||||
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);
|
||||
}
|
||||
```
|
||||
|
||||
## `TypeSerializer` and `TypeDeserializer` Interfaces
|
||||
|
||||
Interfaces for custom type serialization and deserialization.
|
||||
|
||||
```java
|
||||
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
|
||||
*/
|
||||
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);
|
||||
}
|
||||
|
||||
/**
|
||||
* Interface for deserializing a specific type from MessagePack.
|
||||
*
|
||||
* @param <T> Type to deserialize
|
||||
*/
|
||||
public interface TypeDeserializer<T> {
|
||||
/**
|
||||
* Convert from serialized representation to object.
|
||||
*
|
||||
* @param serialized Serialized representation
|
||||
* @return Deserialized object
|
||||
*/
|
||||
T fromSerialized(Object serialized);
|
||||
}
|
||||
```
|
||||
|
||||
## `MsgPackSerializer` Implementation
|
||||
|
||||
A concrete implementation of `ReflectionSerializer` using MessagePack.
|
||||
|
||||
```java
|
||||
package com.langgraph.checkpoint.serde;
|
||||
|
||||
import org.msgpack.core.MessageBufferPacker;
|
||||
import org.msgpack.core.MessagePack;
|
||||
import org.msgpack.core.MessageUnpacker;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.lang.reflect.Constructor;
|
||||
import java.lang.reflect.Field;
|
||||
import java.lang.reflect.Method;
|
||||
import java.util.*;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
/**
|
||||
* MessagePack-based serializer that uses reflection to handle arbitrary Java objects.
|
||||
* This implementation mirrors the Python serialization approach by saving constructor
|
||||
* and import information to reconstruct 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<String, Class<?>> classCache = new ConcurrentHashMap<>();
|
||||
|
||||
/**
|
||||
* Register built-in serializers for common types.
|
||||
*/
|
||||
public MsgPackSerializer() {
|
||||
// Register common built-in types
|
||||
registerBuiltinTypes();
|
||||
}
|
||||
|
||||
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));
|
||||
|
||||
// ... other built-in types 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
|
||||
*/
|
||||
@SuppressWarnings("unchecked")
|
||||
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)) {
|
||||
TypeSerializer<Object> serializer = (TypeSerializer<Object>) serializers.get(type);
|
||||
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 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 {
|
||||
// Custom object - serialize using reflection
|
||||
serializeCustomObject(obj, packer);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Serialize a custom object using reflection.
|
||||
*
|
||||
* @param obj Object to serialize
|
||||
* @param packer MessagePack packer
|
||||
* @throws IOException If packing fails
|
||||
*/
|
||||
private void serializeCustomObject(Object obj, MessageBufferPacker packer) throws IOException {
|
||||
Class<?> type = obj.getClass();
|
||||
|
||||
// Pack object with type information
|
||||
packer.packMapHeader(3);
|
||||
packer.packString("__type__");
|
||||
packer.packString(type.getName());
|
||||
|
||||
// Save constructor info
|
||||
packer.packString("__constructor__");
|
||||
packer.packString(type.getName());
|
||||
|
||||
// Save fields using reflection
|
||||
Map<String, Object> fields = getObjectFields(obj);
|
||||
packer.packString("__fields__");
|
||||
packer.packMapHeader(fields.size());
|
||||
|
||||
for (Map.Entry<String, Object> entry : fields.entrySet()) {
|
||||
packer.packString(entry.getKey());
|
||||
serializeObject(entry.getValue(), packer);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Get all fields from an object using reflection.
|
||||
*
|
||||
* @param obj Object to extract fields from
|
||||
* @return Map of field name to field value
|
||||
*/
|
||||
private Map<String, Object> getObjectFields(Object obj) {
|
||||
Map<String, Object> result = new HashMap<>();
|
||||
Class<?> type = obj.getClass();
|
||||
|
||||
// Get all declared fields, including private ones
|
||||
for (Field field : type.getDeclaredFields()) {
|
||||
try {
|
||||
field.setAccessible(true);
|
||||
result.put(field.getName(), field.get(obj));
|
||||
} catch (IllegalAccessException e) {
|
||||
throw new SerializationException("Failed to access field: " + field.getName(), e);
|
||||
}
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
/**
|
||||
* Deserialize an object from the MessagePack unpacker.
|
||||
*
|
||||
* @param unpacker MessagePack unpacker
|
||||
* @return Deserialized object
|
||||
* @throws IOException If unpacking fails
|
||||
*/
|
||||
@SuppressWarnings("unchecked")
|
||||
private Object deserializeObject(MessageUnpacker unpacker) throws IOException {
|
||||
if (unpacker.tryUnpackNil()) {
|
||||
return null;
|
||||
}
|
||||
|
||||
// Handle different types based on MessagePack format
|
||||
switch (unpacker.getNextFormat()) {
|
||||
case STRING:
|
||||
return unpacker.unpackString();
|
||||
|
||||
case INTEGER:
|
||||
return unpacker.unpackInt();
|
||||
|
||||
case FLOAT:
|
||||
return unpacker.unpackDouble();
|
||||
|
||||
case BOOLEAN:
|
||||
return unpacker.unpackBoolean();
|
||||
|
||||
case BINARY:
|
||||
int binaryLength = unpacker.unpackBinaryHeader();
|
||||
byte[] binary = new byte[binaryLength];
|
||||
unpacker.readPayload(binary);
|
||||
return binary;
|
||||
|
||||
case ARRAY:
|
||||
int arraySize = unpacker.unpackArrayHeader();
|
||||
List<Object> list = new ArrayList<>(arraySize);
|
||||
for (int i = 0; i < arraySize; i++) {
|
||||
list.add(deserializeObject(unpacker));
|
||||
}
|
||||
return list;
|
||||
|
||||
case MAP:
|
||||
int mapSize = unpacker.unpackMapHeader();
|
||||
|
||||
// Check if this is a typed object
|
||||
if (mapSize == 2 || mapSize == 3) {
|
||||
String firstKey = unpacker.unpackString();
|
||||
if ("__type__".equals(firstKey)) {
|
||||
String typeName = unpacker.unpackString();
|
||||
String secondKey = unpacker.unpackString();
|
||||
|
||||
if ("value".equals(secondKey)) {
|
||||
// This is a simple typed value
|
||||
Object value = deserializeObject(unpacker);
|
||||
Class<?> type = loadClass(typeName);
|
||||
|
||||
if (deserializers.containsKey(type)) {
|
||||
TypeDeserializer<Object> deserializer =
|
||||
(TypeDeserializer<Object>) deserializers.get(type);
|
||||
return deserializer.fromSerialized(value);
|
||||
}
|
||||
|
||||
return value;
|
||||
} else if ("__constructor__".equals(secondKey)) {
|
||||
// This is a complex object with fields
|
||||
String constructorName = unpacker.unpackString();
|
||||
String fieldsKey = unpacker.unpackString();
|
||||
|
||||
if ("__fields__".equals(fieldsKey)) {
|
||||
int fieldsCount = unpacker.unpackMapHeader();
|
||||
Map<String, Object> fields = new HashMap<>(fieldsCount);
|
||||
|
||||
for (int i = 0; i < fieldsCount; i++) {
|
||||
String fieldName = unpacker.unpackString();
|
||||
Object fieldValue = deserializeObject(unpacker);
|
||||
fields.put(fieldName, fieldValue);
|
||||
}
|
||||
|
||||
return reconstructObject(typeName, constructorName, fields);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Regular map
|
||||
Map<Object, Object> map = new HashMap<>(mapSize);
|
||||
for (int i = 0; i < mapSize; i++) {
|
||||
Object key = deserializeObject(unpacker);
|
||||
Object value = deserializeObject(unpacker);
|
||||
map.put(key, value);
|
||||
}
|
||||
return map;
|
||||
|
||||
default:
|
||||
throw new SerializationException("Unsupported MessagePack format: " + unpacker.getNextFormat());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Reconstruct an object using its class name, constructor, and field values.
|
||||
*
|
||||
* @param typeName Full class name
|
||||
* @param constructorName Constructor class name
|
||||
* @param fields Map of field names to values
|
||||
* @return Reconstructed object
|
||||
*/
|
||||
private Object reconstructObject(String typeName, String constructorName, Map<String, Object> fields) {
|
||||
try {
|
||||
Class<?> type = loadClass(typeName);
|
||||
|
||||
// Try to create instance using no-arg constructor
|
||||
Object instance = type.getDeclaredConstructor().newInstance();
|
||||
|
||||
// Set all fields using reflection
|
||||
for (Map.Entry<String, Object> entry : fields.entrySet()) {
|
||||
setField(instance, entry.getKey(), entry.getValue());
|
||||
}
|
||||
|
||||
return instance;
|
||||
} catch (Exception e) {
|
||||
throw new SerializationException("Failed to reconstruct object of type: " + typeName, e);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Set a field value using reflection.
|
||||
*
|
||||
* @param obj Object to set field on
|
||||
* @param fieldName Field name
|
||||
* @param value Field value
|
||||
*/
|
||||
private void setField(Object obj, String fieldName, Object value) {
|
||||
try {
|
||||
Field field = obj.getClass().getDeclaredField(fieldName);
|
||||
field.setAccessible(true);
|
||||
field.set(obj, value);
|
||||
} catch (Exception e) {
|
||||
throw new SerializationException("Failed to set field: " + fieldName, e);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Load a class by name, with caching.
|
||||
*
|
||||
* @param className Class name to load
|
||||
* @return Class object
|
||||
*/
|
||||
private Class<?> loadClass(String className) {
|
||||
return classCache.computeIfAbsent(className, name -> {
|
||||
try {
|
||||
return Class.forName(name);
|
||||
} catch (ClassNotFoundException e) {
|
||||
throw new SerializationException("Failed to load class: " + name, e);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Exception thrown during serialization/deserialization.
|
||||
*/
|
||||
public static class SerializationException extends RuntimeException {
|
||||
public SerializationException(String message) {
|
||||
super(message);
|
||||
}
|
||||
|
||||
public SerializationException(String message, Throwable cause) {
|
||||
super(message, cause);
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Record Class Support
|
||||
|
||||
Java's Record classes (Java 14+) can be used as a close alternative to Python's TypedDict/Pydantic models for state schemas. The serializer can handle them via reflection.
|
||||
|
||||
```java
|
||||
package com.langgraph.checkpoint.serde;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* Example of a state schema using Java Record (Java 14+).
|
||||
* Records provide immutable data classes with automatic getters,
|
||||
* equals/hashCode, and toString implementations.
|
||||
*/
|
||||
public record ConversationState(
|
||||
List<Map<String, Object>> messages,
|
||||
Map<String, Object> context,
|
||||
List<String> history
|
||||
) {
|
||||
// Can include custom methods if needed
|
||||
}
|
||||
|
||||
/**
|
||||
* Example of how to use Records for state schemas
|
||||
*/
|
||||
public class StateExample {
|
||||
public static void main(String[] args) {
|
||||
// Create a state instance
|
||||
ConversationState state = new ConversationState(
|
||||
List.of(Map.of("role", "user", "content", "Hello")),
|
||||
Map.of("session_id", "12345"),
|
||||
List.of("Started conversation")
|
||||
);
|
||||
|
||||
// Serialize the state
|
||||
MsgPackSerializer serializer = new MsgPackSerializer();
|
||||
byte[] serialized = serializer.serialize(state);
|
||||
|
||||
// Deserialize the state
|
||||
ConversationState deserialized = (ConversationState) serializer.deserialize(serialized);
|
||||
|
||||
// Access fields using generated getters
|
||||
System.out.println(deserialized.messages());
|
||||
System.out.println(deserialized.context());
|
||||
System.out.println(deserialized.history());
|
||||
}
|
||||
}
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
+369
@@ -0,0 +1,369 @@
|
||||
# Pregel Execution Engine Specification
|
||||
|
||||
## Overview
|
||||
|
||||
Pregel is the foundational execution engine for LangGraph, implementing a Bulk Synchronous Parallel (BSP) computation model. Inspired by Google's original Pregel system, it provides a framework for executing computational graphs with stateful communication between nodes while maintaining strict invariants around execution order, state updates, and error handling.
|
||||
|
||||
## Conceptual Model
|
||||
|
||||
Pregel follows the Bulk Synchronous Parallel model, where computation proceeds in a series of synchronized steps called "supersteps." Each superstep consists of three distinct phases:
|
||||
|
||||
1. **Plan**: Determine which actors to execute based on pending channel updates
|
||||
2. **Execute**: Run selected actors in parallel, collecting their outputs
|
||||
3. **Update**: Apply all updates to channels atomically at the end of the step
|
||||
|
||||
This model ensures several critical properties:
|
||||
|
||||
- **Determinism**: Given the same input, execution produces the same output
|
||||
- **Isolation**: Node executions within a superstep cannot observe each other's updates until the next superstep
|
||||
- **Atomicity**: All channel updates from a superstep are applied at once
|
||||
- **Checkpoint-ability**: The system state can be captured at superstep boundaries
|
||||
|
||||
## Core Components
|
||||
|
||||
### PregelProtocol
|
||||
|
||||
Abstract interface defining the contract for all Pregel implementations:
|
||||
|
||||
```python
|
||||
class PregelProtocol(Protocol):
|
||||
def invoke(self, input: Any, config: Optional[dict] = None, **kwargs: Any) -> Any: ...
|
||||
|
||||
def stream(
|
||||
self,
|
||||
input: Any,
|
||||
config: Optional[dict] = None,
|
||||
stream_mode: Optional[StreamMode] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterator[Any]: ...
|
||||
|
||||
def get_state(self, thread_id: Optional[str] = None) -> Any: ...
|
||||
|
||||
def update_state(self, thread_id: str, state: Any) -> None: ...
|
||||
|
||||
def get_state_history(self, thread_id: str) -> list[Any]: ...
|
||||
```
|
||||
|
||||
### PregelNode
|
||||
|
||||
Represents an actor in the system with the following properties:
|
||||
|
||||
```python
|
||||
class PregelNode:
|
||||
def __init__(
|
||||
self,
|
||||
name: str,
|
||||
action: PregelExecutable,
|
||||
*,
|
||||
subscribe: Optional[Collection[str]] = None,
|
||||
trigger: Optional[str] = None,
|
||||
writers: Optional[Collection[str]] = None,
|
||||
retry_policy: Optional[RetryPolicy] = None,
|
||||
) -> None: ...
|
||||
```
|
||||
|
||||
Key properties:
|
||||
|
||||
- **name**: Unique identifier for the node
|
||||
- **action**: Function to execute when the node is triggered
|
||||
- **subscribe**: Channel names this node listens to for updates
|
||||
- **trigger**: Special condition for node execution
|
||||
- **writers**: Channels this node can write to (for validation)
|
||||
- **retry_policy**: Strategy for handling execution failures
|
||||
|
||||
### Channels
|
||||
|
||||
Communication paths that store values and propagate them between nodes:
|
||||
|
||||
```python
|
||||
class Channel(Protocol):
|
||||
def get_value(self) -> Any: ...
|
||||
|
||||
def update(self, value: Any) -> bool: ...
|
||||
|
||||
def checkpoint(self) -> Any: ...
|
||||
|
||||
def from_checkpoint(self, value: Any) -> None: ...
|
||||
```
|
||||
|
||||
Each channel type implements this interface with specific behaviors:
|
||||
|
||||
1. **LastValue**: Stores only the most recent value; rejects multiple updates in a single step
|
||||
2. **AnyValue**: Permits multiple updates within a step, storing the last one
|
||||
3. **EphemeralValue**: Temporary storage that clears after being read
|
||||
4. **UntrackedValue**: Like LastValue but excluded from checkpoints
|
||||
5. **NamedBarrierValue**: Synchronization mechanism requiring all named values to be received
|
||||
6. **BinaryOperatorAggregate**: Applies operations to combine values (sum, join, etc.)
|
||||
7. **Topic**: PubSub topic supporting multiple subscribers and values
|
||||
|
||||
### Tasks
|
||||
|
||||
Units of work representing computations to execute:
|
||||
|
||||
```python
|
||||
class PregelTask:
|
||||
def __init__(
|
||||
self,
|
||||
node: str,
|
||||
trigger: Optional[str] = None,
|
||||
retry_policy: Optional[RetryPolicy] = None,
|
||||
) -> None: ...
|
||||
|
||||
class PregelExecutableTask:
|
||||
def __init__(
|
||||
self,
|
||||
task: PregelTask,
|
||||
inputs: dict[str, Any],
|
||||
context: dict[str, Any],
|
||||
) -> None: ...
|
||||
```
|
||||
|
||||
### Checkpoints
|
||||
|
||||
Snapshots of execution state at superstep boundaries:
|
||||
|
||||
```python
|
||||
class Checkpoint:
|
||||
def __init__(self, channel_values: dict[str, Any]) -> None: ...
|
||||
|
||||
def get_values(self) -> dict[str, Any]: ...
|
||||
|
||||
def update(self, channel_values: dict[str, Any]) -> None: ...
|
||||
```
|
||||
|
||||
## Execution Flow
|
||||
|
||||
When `invoke()` or `stream()` is called, Pregel performs the following steps:
|
||||
|
||||
1. **Initialization**:
|
||||
|
||||
- Create or retrieve a checkpoint for the specified thread
|
||||
- Load initial channel values from the checkpoint
|
||||
- Transform input into channel updates
|
||||
|
||||
2. **Superstep Loop**:
|
||||
|
||||
- **Plan**: Identify nodes to execute based on channel updates
|
||||
- **Execute**: Run selected nodes in parallel, collecting updates
|
||||
- **Update**: Apply all updates to channels atomically
|
||||
- **Checkpoint**: Save the current state if checkpointing is enabled
|
||||
- Repeat until no more nodes are active or an END token is received
|
||||
|
||||
3. **Termination**:
|
||||
- Extract output from designated output channels
|
||||
- Return result or stream updates based on the stream mode
|
||||
|
||||
This process is depicted in the diagram below:
|
||||
|
||||
```
|
||||
┌──────────────────────────────────┐
|
||||
│ INPUT │
|
||||
└──────────────────┬───────────────┘
|
||||
│
|
||||
▼
|
||||
┌──────────────────────────────────┐
|
||||
│ INITIALIZE STATE │
|
||||
└──────────────────┬───────────────┘
|
||||
│
|
||||
▼
|
||||
┌──────────────────────────────────┐
|
||||
│ SUPERSTEP LOOP │
|
||||
│ ┌────────────────────────────┐ │
|
||||
│ │ PLAN │ │
|
||||
│ │ (Identify active nodes) │ │
|
||||
│ └────────────┬───────────────┘ │
|
||||
│ │ │
|
||||
│ ▼ │
|
||||
│ ┌────────────────────────────┐ │
|
||||
│ │ EXECUTE │ │
|
||||
│ │ (Run nodes in parallel) │ │
|
||||
│ └────────────┬───────────────┘ │
|
||||
│ │ │
|
||||
│ ▼ │
|
||||
│ ┌────────────────────────────┐ │
|
||||
│ │ UPDATE │ │
|
||||
│ │ (Apply channel updates) │ │
|
||||
│ └────────────┬───────────────┘ │
|
||||
│ │ │
|
||||
│ ▼ │
|
||||
│ ┌────────────────────────────┐ │
|
||||
│ │ CHECKPOINT │ │
|
||||
│ │ (Save current state) │ │
|
||||
│ └────────────┬───────────────┘ │
|
||||
│ │ │
|
||||
│ ▼ │
|
||||
│ [More nodes active?] │
|
||||
│ / \ │
|
||||
│ Yes No │
|
||||
│ │ │ │
|
||||
│ └──────────────────┘ │
|
||||
└──────────────────┬───────────────┘
|
||||
│
|
||||
▼
|
||||
┌──────────────────────────────────┐
|
||||
│ OUTPUT │
|
||||
└──────────────────────────────────┘
|
||||
```
|
||||
|
||||
## Stream Modes
|
||||
|
||||
Pregel supports different streaming options to provide visibility into execution:
|
||||
|
||||
- **values**: Stream the complete state after each superstep
|
||||
- **updates**: Stream state deltas after each node execution
|
||||
- **debug**: Stream comprehensive execution information for debugging
|
||||
|
||||
Each mode has different performance characteristics and use cases.
|
||||
|
||||
## Critical Invariants
|
||||
|
||||
Based on LangGraph's test suite, Pregel maintains the following invariants:
|
||||
|
||||
1. **Execution Determinism**:
|
||||
|
||||
- Given the same input and thread ID, execution produces identical results
|
||||
- Channel update order within a superstep does not affect the final state
|
||||
|
||||
2. **State Isolation**:
|
||||
|
||||
- Updates from one node are not visible to other nodes within the same superstep
|
||||
- Nodes cannot observe partial updates from incomplete executions
|
||||
|
||||
3. **Checkpoint Consistency**:
|
||||
|
||||
- Checkpoints capture the complete system state at superstep boundaries
|
||||
- Restored checkpoints resume execution with the exact same state
|
||||
|
||||
4. **Task Management**:
|
||||
|
||||
- Tasks execute in a deterministic order based on node dependencies
|
||||
- Parallel execution optimizes performance but maintains consistency
|
||||
|
||||
5. **Error Handling**:
|
||||
|
||||
- Node execution failures can be handled with retry policies
|
||||
- Unrecoverable errors propagate without corrupting the execution state
|
||||
- Failures during checkpointing do not corrupt previous checkpoints
|
||||
|
||||
6. **Termination Guarantees**:
|
||||
- Execution always terminates for acyclic graphs
|
||||
- Cyclic graphs require explicit exit conditions to ensure termination
|
||||
- Execution timeouts prevent infinite loops
|
||||
|
||||
## Implementation Notes
|
||||
|
||||
### Type Safety
|
||||
|
||||
Pregel enforces type safety through:
|
||||
|
||||
- Input/output schema validation
|
||||
- Channel type checking for updates
|
||||
- Runtime validation of node return values
|
||||
|
||||
### Concurrency Model
|
||||
|
||||
Pregel balances parallelism with determinism:
|
||||
|
||||
- Nodes within a superstep can execute in parallel
|
||||
- Channel updates are collected and applied sequentially
|
||||
- Execution order is deterministic despite parallel processing
|
||||
|
||||
### Performance Optimizations
|
||||
|
||||
- **Eager Planning**: Pregel identifies all active nodes at the start of a superstep
|
||||
- **Task Batching**: Similar tasks can be batched for efficient execution
|
||||
- **Lazy Checkpointing**: Only modified channels are included in checkpoints
|
||||
- **Channel-specific optimizations**: Different channel types use specialized storage strategies
|
||||
|
||||
### Testing Approach
|
||||
|
||||
The test suite for Pregel focuses on:
|
||||
|
||||
1. **Functional correctness**: Ensuring proper node execution and state updates
|
||||
2. **Concurrency safety**: Verifying parallel execution does not affect determinism
|
||||
3. **Error handling**: Confirming failures are properly managed
|
||||
4. **Checkpoint fidelity**: Validating checkpoint creation and restoration
|
||||
5. **Edge cases**: Testing unusual graph topologies and execution patterns
|
||||
|
||||
## Reimplementation Guidance
|
||||
|
||||
When reimplementing Pregel from scratch, consider the following approach:
|
||||
|
||||
1. Start with a simplified sequential execution model that maintains basic invariants
|
||||
2. Add channel implementations one at a time, focusing on correctness
|
||||
3. Implement the checkpoint system with proper serialization
|
||||
4. Add parallel execution with careful attention to update ordering
|
||||
5. Implement error handling and retry policies
|
||||
6. Optimize for performance and resource usage
|
||||
|
||||
The most challenging aspects are:
|
||||
|
||||
- Maintaining determinism with parallel execution
|
||||
- Ensuring checkpoint consistency
|
||||
- Properly handling error cases
|
||||
- Managing memory usage for large state objects
|
||||
|
||||
## Example Usage
|
||||
|
||||
```python
|
||||
# Define node functions
|
||||
def process_input(state):
|
||||
# Process input data
|
||||
return {"output_channel": processed_data}
|
||||
|
||||
def make_decision(state):
|
||||
# Make a decision based on processed data
|
||||
return {"decision_channel": decision}
|
||||
|
||||
# Create PregelNodes
|
||||
input_node = PregelNode(
|
||||
name="input_processor",
|
||||
action=process_input,
|
||||
subscribe=["input_channel"],
|
||||
writers=["output_channel"]
|
||||
)
|
||||
|
||||
decision_node = PregelNode(
|
||||
name="decision_maker",
|
||||
action=make_decision,
|
||||
subscribe=["output_channel"],
|
||||
writers=["decision_channel"]
|
||||
)
|
||||
|
||||
# Create channels
|
||||
channels = {
|
||||
"input_channel": LastValue(),
|
||||
"output_channel": LastValue(),
|
||||
"decision_channel": LastValue()
|
||||
}
|
||||
|
||||
# Create Pregel instance
|
||||
pregel = Pregel(
|
||||
nodes={"input_processor": input_node, "decision_maker": decision_node},
|
||||
channels=channels,
|
||||
checkpoint_factories={"memory": memory_checkpointer_factory}
|
||||
)
|
||||
|
||||
# Invoke the graph
|
||||
result = pregel.invoke(
|
||||
{"input_channel": input_data},
|
||||
config={"thread_id": "conversation_123"}
|
||||
)
|
||||
|
||||
# Stream execution with updates
|
||||
for update in pregel.stream(
|
||||
{"input_channel": input_data},
|
||||
config={"thread_id": "conversation_456"},
|
||||
stream_mode="updates"
|
||||
):
|
||||
print(update)
|
||||
```
|
||||
|
||||
## Related Components
|
||||
|
||||
Pregel interacts closely with:
|
||||
|
||||
- **StateGraph**: High-level API that compiles to Pregel
|
||||
- **Channels**: Communication primitives used by Pregel
|
||||
- **Checkpoint System**: Provides persistence for Pregel execution
|
||||
- **Human-in-the-Loop**: Uses Pregel's checkpointing for interruption/resumption
|
||||
@@ -0,0 +1,182 @@
|
||||
# StateGraph API Specification
|
||||
|
||||
## Overview
|
||||
|
||||
`StateGraph` is the primary high-level API in LangGraph for building stateful computation graphs. It represents a graph structure where nodes communicate by reading and writing to a shared state, enabling complex multi-step workflows with LLMs, tools, and other components.
|
||||
|
||||
## Constructor
|
||||
|
||||
```python
|
||||
def __init__(
|
||||
self,
|
||||
state_schema: Optional[Type[Any]] = None,
|
||||
config_schema: Optional[Type[Any]] = None,
|
||||
*,
|
||||
input: Optional[Type[Any]] = None,
|
||||
output: Optional[Type[Any]] = None,
|
||||
) -> None
|
||||
```
|
||||
|
||||
### Parameters
|
||||
|
||||
- **state_schema**: The schema defining the state structure, typically a TypedDict or Pydantic model
|
||||
- **config_schema**: Optional schema defining configuration parameters
|
||||
- **input**: Optional schema for graph inputs (defaults to state_schema)
|
||||
- **output**: Optional schema for graph outputs (defaults to state_schema)
|
||||
|
||||
## Core Methods
|
||||
|
||||
### Node Management
|
||||
|
||||
```python
|
||||
def add_node(
|
||||
self,
|
||||
node: Union[str, RunnableLike],
|
||||
action: Optional[RunnableLike] = None,
|
||||
*,
|
||||
metadata: Optional[dict[str, Any]] = None,
|
||||
input: Optional[Type[Any]] = None,
|
||||
retry: Optional[RetryPolicy] = None,
|
||||
destinations: Optional[Union[dict[str, str], tuple[str]]] = None,
|
||||
subgraphs: list[PregelProtocol] = EMPTY_SEQ,
|
||||
) -> Self
|
||||
```
|
||||
|
||||
Adds a new node to the graph. The node can be specified as a string ID with an action callable, or as a Runnable.
|
||||
|
||||
```python
|
||||
def add_sequence(
|
||||
self,
|
||||
nodes: Sequence[Union[RunnableLike, tuple[str, RunnableLike]]],
|
||||
) -> Self
|
||||
```
|
||||
|
||||
Adds a sequence of nodes to be executed in order, automatically creating edges between them.
|
||||
|
||||
### Edge Management
|
||||
|
||||
```python
|
||||
def add_edge(self, start_key: Union[str, list[str]], end_key: str) -> Self
|
||||
```
|
||||
|
||||
Adds a directed edge from start node to end node. The start node can be a single node or a list of nodes.
|
||||
|
||||
```python
|
||||
def add_conditional_edges(
|
||||
self,
|
||||
source: str,
|
||||
path: Union[Callable[..., Union[Hashable, list[Hashable]]], Runnable[Any, Union[Hashable, list[Hashable]]]],
|
||||
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
|
||||
then: Optional[str] = None,
|
||||
) -> Self
|
||||
```
|
||||
|
||||
Adds conditional routing logic between nodes. The path callable examines the state and returns a value that determines the next node to execute.
|
||||
|
||||
### Graph Entry/Exit
|
||||
|
||||
```python
|
||||
def set_entry_point(self, key: str) -> Self
|
||||
```
|
||||
|
||||
Defines the starting node for graph execution. Only needed if not using START node.
|
||||
|
||||
```python
|
||||
def set_conditional_entry_point(
|
||||
self,
|
||||
path: Union[Callable[..., Union[Hashable, list[Hashable]]], Runnable[Any, Union[Hashable, list[Hashable]]]],
|
||||
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
|
||||
then: Optional[str] = None,
|
||||
) -> Self
|
||||
```
|
||||
|
||||
Sets a conditional starting point based on the input state.
|
||||
|
||||
```python
|
||||
def set_finish_point(self, key: str) -> Self
|
||||
```
|
||||
|
||||
Marks a node as an exit point for the graph.
|
||||
|
||||
### Validation and Compilation
|
||||
|
||||
```python
|
||||
def validate(self, interrupt: Optional[Sequence[str]] = None) -> Self
|
||||
```
|
||||
|
||||
Checks the graph for correctness, ensuring there are no disconnected nodes or unreachable states. Called by compile.
|
||||
|
||||
```python
|
||||
def compile(
|
||||
self,
|
||||
checkpointer: Checkpointer = None,
|
||||
*,
|
||||
store: Optional[BaseStore] = None,
|
||||
interrupt_before: Optional[Union[All, list[str]]] = None,
|
||||
interrupt_after: Optional[Union[All, list[str]]] = None,
|
||||
name: Optional[str] = None,
|
||||
) -> "CompiledStateGraph"
|
||||
```
|
||||
|
||||
Transforms the graph into an executable CompiledStateGraph. The checkpointer enables state persistence.
|
||||
|
||||
## Constants
|
||||
|
||||
Two special constants are provided for graph construction:
|
||||
|
||||
- **START**: Special value representing the entry point to the graph
|
||||
- **END**: Special value representing an exit point from the graph
|
||||
|
||||
## Implementation Details
|
||||
|
||||
When a `StateGraph` is compiled, it is transformed into a Pregel instance with:
|
||||
|
||||
1. **Node Translation**: Each graph node becomes a `PregelNode` with associated actions
|
||||
2. **Channel Creation**: State fields are represented as channels with appropriate behaviors
|
||||
3. **Edge Mapping**: Graph edges determine message routing between nodes
|
||||
4. **Branch Handling**: Conditional edges are implemented as special routing logic
|
||||
5. **Checkpoint Configuration**: If provided, enables state persistence and resumption
|
||||
|
||||
The StateGraph API handles the complexities of the underlying Pregel execution model, providing a more intuitive interface for building stateful workflows.
|
||||
|
||||
## Example Usage
|
||||
|
||||
```python
|
||||
from typing import TypedDict
|
||||
from langgraph.graph import StateGraph, START, END
|
||||
|
||||
# Define the state schema
|
||||
class State(TypedDict):
|
||||
count: int
|
||||
message: str
|
||||
|
||||
# Create a StateGraph with our schema
|
||||
graph = StateGraph(State)
|
||||
|
||||
# Add nodes
|
||||
def increment(state: State):
|
||||
return {"count": state["count"] + 1}
|
||||
|
||||
def check(state: State):
|
||||
if state["count"] >= 3:
|
||||
return "finish"
|
||||
return "increment"
|
||||
|
||||
def finish(state: State):
|
||||
return {"message": f"Finished with count {state['count']}"}
|
||||
|
||||
graph.add_node("increment", increment)
|
||||
graph.add_node("check", check)
|
||||
graph.add_node("finish", finish)
|
||||
|
||||
# Add edges
|
||||
graph.add_edge(START, "increment")
|
||||
graph.add_edge("increment", "check")
|
||||
graph.add_conditional_edges("check", check, {"finish": "finish"})
|
||||
graph.add_edge("finish", END)
|
||||
|
||||
# Compile and run
|
||||
compiled_graph = graph.compile()
|
||||
result = compiled_graph.invoke({"count": 0, "message": ""})
|
||||
# result will be {"count": 3, "message": "Finished with count 3"}
|
||||
```
|
||||
Reference in New Issue
Block a user