Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 10 additions & 4 deletions python/packages/core/agent_framework/_workflows/_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,7 +127,13 @@ def from_dict(cls, data: Mapping[str, Any]) -> WorkflowCheckpoint:


class CheckpointStorage(Protocol):
"""Protocol for checkpoint storage backends."""
"""Protocol for checkpoint storage backends.

Reads return objects owned by the caller: mutating a checkpoint returned by ``load``,
``list_checkpoints`` or ``get_latest`` must not change what a later read returns.
Backends that serialize their storage satisfy this by construction; backends that hold
checkpoints in memory must copy on read.
"""

async def save(self, checkpoint: WorkflowCheckpoint) -> CheckpointID:
"""Save a checkpoint and return its ID.
Expand Down Expand Up @@ -217,12 +223,12 @@ async def load(self, checkpoint_id: CheckpointID) -> WorkflowCheckpoint:
checkpoint = self._checkpoints.get(checkpoint_id)
if checkpoint:
logger.debug(f"Loaded checkpoint {checkpoint_id} from memory")
return checkpoint
return copy.deepcopy(checkpoint)
raise WorkflowCheckpointException(f"No checkpoint found with ID {checkpoint_id}")

async def list_checkpoints(self, *, workflow_name: str) -> list[WorkflowCheckpoint]:
"""List checkpoint objects for a given workflow name."""
return [cp for cp in self._checkpoints.values() if cp.workflow_name == workflow_name]
return [copy.deepcopy(cp) for cp in self._checkpoints.values() if cp.workflow_name == workflow_name]

async def delete(self, checkpoint_id: CheckpointID) -> bool:
"""Delete a checkpoint by ID."""
Expand All @@ -239,7 +245,7 @@ async def get_latest(self, *, workflow_name: str) -> WorkflowCheckpoint | None:
return None
latest_checkpoint = max(checkpoints, key=lambda cp: datetime.fromisoformat(cp.timestamp))
logger.debug(f"Latest checkpoint for workflow {workflow_name} is {latest_checkpoint.checkpoint_id}")
return latest_checkpoint
return copy.deepcopy(latest_checkpoint)

async def list_checkpoint_ids(self, *, workflow_name: str) -> list[CheckpointID]:
"""List checkpoint IDs. If workflow_id is provided, filter by that workflow."""
Expand Down
40 changes: 40 additions & 0 deletions python/packages/core/tests/workflow/test_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -1765,4 +1765,44 @@ async def test_file_checkpoint_storage_roundtrip_empty_collections():
assert loaded.pending_request_info_events == {}


async def test_memory_checkpoint_storage_load_returns_caller_owned_copy():
"""Mutating a loaded checkpoint must not change what a later read returns."""
storage = InMemoryCheckpointStorage()
checkpoint = WorkflowCheckpoint(
workflow_name="test-workflow",
graph_signature_hash="test-hash",
state={"answer": "original"},
)
await storage.save(checkpoint)

loaded = await storage.load(checkpoint.checkpoint_id)
loaded.state["answer"] = "mutated"

reloaded = await storage.load(checkpoint.checkpoint_id)
assert reloaded.state["answer"] == "original"


async def test_memory_checkpoint_storage_list_and_get_latest_return_caller_owned_copies():
"""list_checkpoints and get_latest follow the same read-isolation contract as load."""
storage = InMemoryCheckpointStorage()
checkpoint = WorkflowCheckpoint(
workflow_name="test-workflow",
graph_signature_hash="test-hash",
state={"answer": "original"},
)
await storage.save(checkpoint)

listed = await storage.list_checkpoints(workflow_name="test-workflow")
listed[0].state["answer"] = "mutated-via-list"

latest = await storage.get_latest(workflow_name="test-workflow")
assert latest is not None
assert latest.state["answer"] == "original"

latest.state["answer"] = "mutated-via-get-latest"

reloaded = await storage.load(checkpoint.checkpoint_id)
assert reloaded.state["answer"] == "original"


# endregion
Loading