remove refactors

This commit is contained in:
vbarda
2024-08-21 18:06:18 -04:00
parent 0a87b9fa1c
commit e7bc74e918
+70 -93
View File
@@ -258,73 +258,6 @@ def _assemble_state_snapshot_hierarchy(
return state_snapshot
def _prepare_state_snapshot(
saved: CheckpointTuple,
graph: Pregel,
) -> StateSnapshot:
with ChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in graph.channels.items()
},
saved.checkpoint,
saved.config,
) as channels, ManagedValuesManager(
graph.managed_values_dict, ensure_config(saved.config)
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
return StateSnapshot(
values=read_channels(channels, graph.stream_channels_asis),
next=tuple(t.name for t in next_tasks),
config=saved.config,
metadata=saved.metadata,
created_at=saved.checkpoint["ts"],
parent_config=saved.parent_config,
tasks=tasks_w_writes(next_tasks, saved.pending_writes),
)
async def _prepare_state_snapshot_async(
saved: CheckpointTuple, graph: Pregel
) -> StateSnapshot:
async with AsyncChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in graph.channels.items()
},
saved.checkpoint,
saved.config,
) as channels, AsyncManagedValuesManager(
graph.managed_values_dict, ensure_config(saved.config)
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
return StateSnapshot(
values=read_channels(channels, graph.stream_channels_asis),
next=tuple(t.name for t in next_tasks),
config=saved.config,
metadata=saved.metadata,
created_at=saved.checkpoint["ts"],
parent_config=saved.parent_config,
tasks=tasks_w_writes(next_tasks, saved.pending_writes),
)
def _has_nested_interrupts(
graph: Pregel,
) -> bool:
@@ -517,20 +450,16 @@ class Pregel(
if not self.checkpointer:
raise ValueError("No checkpointer set")
checkpoint_tuple = self.checkpointer.get_tuple(config)
checkpoint_config = checkpoint_tuple.config if checkpoint_tuple else config
saved = self.checkpointer.get_tuple(config)
checkpoint_config = saved.config if saved else config
checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "")
checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id")
checkpoint_ns_to_checkpoint_id: dict[str, str] = {}
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {}
checkpoint_ns_to_graph: dict[str, Pregel] = {}
for checkpoint_tuple in self.checkpointer.list(config):
saved_checkpoint_ns = checkpoint_tuple.config["configurable"][
"checkpoint_ns"
]
saved_checkpoint_id = checkpoint_tuple.config["configurable"][
"checkpoint_id"
]
for saved in self.checkpointer.list(config):
saved_checkpoint_ns = saved.config["configurable"]["checkpoint_ns"]
saved_checkpoint_id = saved.config["configurable"]["checkpoint_id"]
if checkpoint_id != saved_checkpoint_id:
continue
@@ -547,10 +476,36 @@ class Pregel(
self, saved_checkpoint_ns
)
state_snapshot = _prepare_state_snapshot(
checkpoint_tuple,
checkpoint_ns_to_graph[saved_checkpoint_ns],
)
graph = checkpoint_ns_to_graph[saved_checkpoint_ns]
with ChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in graph.channels.items()
},
saved.checkpoint,
saved.config,
) as channels, ManagedValuesManager(
graph.managed_values_dict, ensure_config(saved.config)
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
state_snapshot = StateSnapshot(
read_channels(channels, graph.stream_channels_asis),
tuple(t.name for t in next_tasks),
saved.config,
saved.metadata,
saved.checkpoint["ts"],
saved.parent_config,
tasks_w_writes(next_tasks, saved.pending_writes),
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
saved_checkpoint_ns
@@ -576,20 +531,16 @@ class Pregel(
if not self.checkpointer:
raise ValueError("No checkpointer set")
checkpoint_tuple = await self.checkpointer.aget_tuple(config)
checkpoint_config = checkpoint_tuple.config if checkpoint_tuple else config
saved = await self.checkpointer.aget_tuple(config)
checkpoint_config = saved.config if saved else config
checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "")
checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id")
checkpoint_ns_to_checkpoint_id: dict[str, str] = {}
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {}
checkpoint_ns_to_graph: dict[str, Pregel] = {}
async for checkpoint_tuple in self.checkpointer.alist(config):
saved_checkpoint_ns = checkpoint_tuple.config["configurable"][
"checkpoint_ns"
]
saved_checkpoint_id = checkpoint_tuple.config["configurable"][
"checkpoint_id"
]
async for saved in self.checkpointer.alist(config):
saved_checkpoint_ns = saved.config["configurable"]["checkpoint_ns"]
saved_checkpoint_id = saved.config["configurable"]["checkpoint_id"]
if checkpoint_id != saved_checkpoint_id:
continue
@@ -606,10 +557,36 @@ class Pregel(
self, saved_checkpoint_ns
)
state_snapshot = await _prepare_state_snapshot_async(
checkpoint_tuple,
checkpoint_ns_to_graph[saved_checkpoint_ns],
)
graph = checkpoint_ns_to_graph[saved_checkpoint_ns]
async with AsyncChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in graph.channels.items()
},
saved.checkpoint,
saved.config,
) as channels, AsyncManagedValuesManager(
graph.managed_values_dict, ensure_config(saved.config)
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
graph.nodes,
channels,
managed,
saved.config,
saved.metadata.get("step", -1) + 1,
for_execution=False,
)
state_snapshot = StateSnapshot(
read_channels(channels, graph.stream_channels_asis),
tuple(t.name for t in next_tasks),
saved.config,
saved.metadata,
saved.checkpoint["ts"],
saved.parent_config,
tasks_w_writes(next_tasks, saved.pending_writes),
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
saved_checkpoint_ns