chore(langgraph): Add passthrough params/headers to invoke/stream/etc. (#5940)

This commit is contained in:
William FH
2025-08-18 19:38:08 +00:00
committed by GitHub
parent a86eb4c5d0
commit c0b29a6df5
3 changed files with 890 additions and 682 deletions
+67 -5
View File
@@ -25,9 +25,17 @@ from langgraph_sdk.client import (
get_client,
get_sync_client,
)
from langgraph_sdk.schema import Checkpoint, ThreadState
from langgraph_sdk.schema import Command as CommandSDK
from langgraph_sdk.schema import StreamMode as StreamModeSDK
from langgraph_sdk.schema import (
Checkpoint,
QueryParamTypes,
ThreadState,
)
from langgraph_sdk.schema import (
Command as CommandSDK,
)
from langgraph_sdk.schema import (
StreamMode as StreamModeSDK,
)
from typing_extensions import Self
from langgraph._internal._config import merge_configs
@@ -208,6 +216,8 @@ class RemoteGraph(PregelProtocol):
config: RunnableConfig | None = None,
*,
xray: int | bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> DrawableGraph:
"""Get graph by graph name.
@@ -226,6 +236,8 @@ class RemoteGraph(PregelProtocol):
graph = sync_client.assistants.get_graph(
assistant_id=self.assistant_id,
xray=xray,
headers=headers,
params=params,
)
return DrawableGraph(
nodes=self._get_drawable_nodes(graph),
@@ -237,6 +249,8 @@ class RemoteGraph(PregelProtocol):
config: RunnableConfig | None = None,
*,
xray: int | bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> DrawableGraph:
"""Get graph by graph name.
@@ -255,6 +269,8 @@ class RemoteGraph(PregelProtocol):
graph = await client.assistants.get_graph(
assistant_id=self.assistant_id,
xray=xray,
headers=headers,
params=params,
)
return DrawableGraph(
nodes=self._get_drawable_nodes(graph),
@@ -376,7 +392,12 @@ class RemoteGraph(PregelProtocol):
return sanitized
def get_state(
self, config: RunnableConfig, *, subgraphs: bool = False
self,
config: RunnableConfig,
*,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> StateSnapshot:
"""Get the state of a thread.
@@ -388,6 +409,8 @@ class RemoteGraph(PregelProtocol):
config: A `RunnableConfig` that includes `thread_id` in the
`configurable` field.
subgraphs: Include subgraphs in the state.
headers: Optional custom headers to include with the request.
params: Optional query parameters to include with the request.
Returns:
The latest state of the thread.
@@ -399,11 +422,18 @@ class RemoteGraph(PregelProtocol):
thread_id=merged_config["configurable"]["thread_id"],
checkpoint=self._get_checkpoint(merged_config),
subgraphs=subgraphs,
headers=headers,
params=params,
)
return self._create_state_snapshot(state)
async def aget_state(
self, config: RunnableConfig, *, subgraphs: bool = False
self,
config: RunnableConfig,
*,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> StateSnapshot:
"""Get the state of a thread.
@@ -415,6 +445,8 @@ class RemoteGraph(PregelProtocol):
config: A `RunnableConfig` that includes `thread_id` in the
`configurable` field.
subgraphs: Include subgraphs in the state.
headers: Optional custom headers to include with the request.
params: Optional query parameters to include with the request.
Returns:
The latest state of the thread.
@@ -426,6 +458,8 @@ class RemoteGraph(PregelProtocol):
thread_id=merged_config["configurable"]["thread_id"],
checkpoint=self._get_checkpoint(merged_config),
subgraphs=subgraphs,
headers=headers,
params=params,
)
return self._create_state_snapshot(state)
@@ -436,6 +470,8 @@ class RemoteGraph(PregelProtocol):
filter: dict[str, Any] | None = None,
before: RunnableConfig | None = None,
limit: int | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> Iterator[StateSnapshot]:
"""Get the state history of a thread.
@@ -460,6 +496,8 @@ class RemoteGraph(PregelProtocol):
before=self._get_checkpoint(before),
metadata=filter,
checkpoint=self._get_checkpoint(merged_config),
headers=headers,
params=params,
)
for state in states:
yield self._create_state_snapshot(state)
@@ -471,6 +509,8 @@ class RemoteGraph(PregelProtocol):
filter: dict[str, Any] | None = None,
before: RunnableConfig | None = None,
limit: int | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> AsyncIterator[StateSnapshot]:
"""Get the state history of a thread.
@@ -482,6 +522,8 @@ class RemoteGraph(PregelProtocol):
filter: Metadata to filter on.
before: A `RunnableConfig` that includes checkpoint metadata.
limit: Max number of states to return.
headers: Optional custom headers to include with the request.
params: Optional query parameters to include with the request.
Returns:
States of the thread.
@@ -495,6 +537,8 @@ class RemoteGraph(PregelProtocol):
before=self._get_checkpoint(before),
metadata=filter,
checkpoint=self._get_checkpoint(merged_config),
headers=headers,
params=params,
)
for state in states:
yield self._create_state_snapshot(state)
@@ -518,6 +562,9 @@ class RemoteGraph(PregelProtocol):
config: RunnableConfig,
values: dict[str, Any] | Any | None,
as_node: str | None = None,
*,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> RunnableConfig:
"""Update the state of a thread.
@@ -540,6 +587,8 @@ class RemoteGraph(PregelProtocol):
values=values,
as_node=as_node,
checkpoint=self._get_checkpoint(merged_config),
headers=headers,
params=params,
)
return self._get_config(response["checkpoint"])
@@ -548,6 +597,9 @@ class RemoteGraph(PregelProtocol):
config: RunnableConfig,
values: dict[str, Any] | Any | None,
as_node: str | None = None,
*,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> RunnableConfig:
"""Update the state of a thread.
@@ -570,6 +622,8 @@ class RemoteGraph(PregelProtocol):
values=values,
as_node=as_node,
checkpoint=self._get_checkpoint(merged_config),
headers=headers,
params=params,
)
return self._get_config(response["checkpoint"])
@@ -634,6 +688,7 @@ class RemoteGraph(PregelProtocol):
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
**kwargs: Any,
) -> Iterator[dict[str, Any] | Any]:
"""Create a run and stream the results.
@@ -681,6 +736,7 @@ class RemoteGraph(PregelProtocol):
headers=_merge_tracing_headers(headers)
if self.distributed_tracing
else headers,
params=params,
**kwargs,
):
# split mode and ns
@@ -741,6 +797,7 @@ class RemoteGraph(PregelProtocol):
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
**kwargs: Any,
) -> AsyncIterator[dict[str, Any] | Any]:
"""Create a run and stream the results.
@@ -788,6 +845,7 @@ class RemoteGraph(PregelProtocol):
headers=_merge_tracing_headers(headers)
if self.distributed_tracing
else headers,
params=params,
**kwargs,
):
# split mode and ns
@@ -862,6 +920,7 @@ class RemoteGraph(PregelProtocol):
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
**kwargs: Any,
) -> dict[str, Any] | Any:
"""Create a run, wait until it finishes and return the final state.
@@ -884,6 +943,7 @@ class RemoteGraph(PregelProtocol):
interrupt_after=interrupt_after,
headers=headers,
stream_mode="values",
params=params,
**kwargs,
):
pass
@@ -900,6 +960,7 @@ class RemoteGraph(PregelProtocol):
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
**kwargs: Any,
) -> dict[str, Any] | Any:
"""Create a run, wait until it finishes and return the final state.
@@ -922,6 +983,7 @@ class RemoteGraph(PregelProtocol):
interrupt_after=interrupt_after,
headers=headers,
stream_mode="values",
params=params,
**kwargs,
):
pass
+1 -1
View File
@@ -14,7 +14,7 @@ license-files = ['LICENSE']
dependencies = [
"langchain-core>=0.1",
"langgraph-checkpoint>=2.1.0,<3.0.0",
"langgraph-sdk>=0.2.0,<0.3.0",
"langgraph-sdk>=0.2.2,<0.3.0",
"langgraph-prebuilt>=0.6.0,<0.7.0",
"xxhash>=3.5.0",
"pydantic>=2.7.4",
+822 -676
View File
File diff suppressed because it is too large Load Diff