Fix filtering on ns

This commit is contained in:
Nuno Campos
2024-08-27 16:09:11 -07:00
parent 499fe10ef2
commit e338a4ec98
3 changed files with 8 additions and 3 deletions
@@ -260,7 +260,8 @@ class BasePostgresSaver(BaseCheckpointSaver):
if config:
wheres.append("thread_id = %s ")
param_values.append(config["configurable"]["thread_id"])
if checkpoint_ns := config["configurable"].get("checkpoint_ns"):
checkpoint_ns = config["configurable"].get("checkpoint_ns")
if checkpoint_ns is not None:
wheres.append("checkpoint_ns = %s")
param_values.append(checkpoint_ns)
@@ -70,7 +70,8 @@ def search_where(
if config is not None:
wheres.append("thread_id = ?")
param_values.append(config["configurable"]["thread_id"])
if checkpoint_ns := config["configurable"].get("checkpoint_ns"):
checkpoint_ns = config["configurable"].get("checkpoint_ns")
if checkpoint_ns is not None:
wheres.append("checkpoint_ns = ?")
param_values.append(checkpoint_ns)
@@ -210,7 +210,10 @@ class MemorySaver(
config_checkpoint_id = get_checkpoint_id(config) if config else None
for thread_id in thread_ids:
for checkpoint_ns in self.storage[thread_id].keys():
if config_checkpoint_ns and checkpoint_ns != config_checkpoint_ns:
if (
config_checkpoint_ns is not None
and checkpoint_ns != config_checkpoint_ns
):
continue
for checkpoint_id, (