From 06ca07432d3a1fb3591bb08095e6bed6fe854c5d Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Thu, 13 Mar 2025 18:55:16 -0700 Subject: [PATCH 1/4] Add sweeper --- .../langgraph/store/postgres/aio.py | 45 +++ .../langgraph/store/postgres/base.py | 304 +++++++++++++----- 2 files changed, 264 insertions(+), 85 deletions(-) diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py index 0615f0ec7..5ffcaaf5a 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py @@ -115,7 +115,9 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con "supports_pipeline", "index_config", "embeddings", + "supports_ttl", ) + supports_ttl: bool = True def __init__( self, @@ -256,6 +258,49 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con "INSERT INTO vector_migrations (v) VALUES (%s)", (v,) ) + async def sweep_ttl(self) -> int: + """Delete expired store items based on TTL. + + Returns: + int: The number of deleted items. + """ + async with self._cursor() as cur: + await cur.execute( + """ + DELETE FROM store + WHERE expires_at IS NOT NULL AND expires_at < NOW() + """ + ) + deleted_count = cur.rowcount + return deleted_count + + async def start_ttl_sweeper( + self, sweep_interval_minutes: Optional[int] = None + ) -> None: + """Periodically delete expired store items based on TTL.""" + if not self.ttl_config: + return + sweep_interval_minutes_ = float( + cast( + float, + sweep_interval_minutes + or self.ttl_config.get("sweep_interval_minutes") + or 5, + ) + ) + logger.info( + f"Starting store TTL sweeper with interval {sweep_interval_minutes_} minutes", + ) + + while True: + await asyncio.sleep(sweep_interval_minutes_ * 60) + try: + expired_items = await self.sweep_ttl() + if expired_items > 0: + logger.info(f"Store swept {expired_items} expired items") + except Exception as exc: + logger.exception("Store TTL sweep iteration failed", exc_info=exc) + async def _execute_batch( self, grouped_ops: dict, diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/base.py b/libs/checkpoint-postgres/langgraph/store/postgres/base.py index c8445dd42..954b8e88d 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/base.py @@ -2,6 +2,7 @@ import asyncio import json import logging import threading +import time from collections import defaultdict from collections.abc import Iterable, Iterator, Sequence from contextlib import contextmanager @@ -74,6 +75,16 @@ CREATE TABLE IF NOT EXISTS store ( """ -- For faster lookups by prefix CREATE INDEX CONCURRENTLY IF NOT EXISTS store_prefix_idx ON store USING btree (prefix text_pattern_ops); +""", + """ +-- Add expires_at column to store table +ALTER TABLE store +ADD COLUMN expires_at TIMESTAMP WITH TIME ZONE, +ADD COLUMN ttl_minutes INT; + +-- Add indexes for efficient TTL sweeping +CREATE INDEX idx_store_expires_at ON store (expires_at) +WHERE expires_at IS NOT NULL; """, ] @@ -225,20 +236,55 @@ class BasePostgresStore(Generic[C]): self, get_ops: Sequence[tuple[int, GetOp]], ) -> list[tuple[str, tuple, tuple[str, ...], list]]: + """ + Build queries to fetch (and optionally refresh the TTL of) multiple keys per namespace. + + Each returned element is a tuple of: + (sql_query_string, sql_params, namespace, items_for_this_namespace) + + where items_for_this_namespace is the original list of (idx, key, refresh_ttl). + """ + namespace_groups = defaultdict(list) + refresh_ttls = defaultdict(list) for idx, op in get_ops: namespace_groups[op.namespace].append((idx, op.key)) + refresh_ttls[op.namespace].append(op.refresh_ttl) + results = [] for namespace, items in namespace_groups.items(): - _, keys = zip(*items) - keys_to_query = ",".join(["%s"] * len(keys)) - query = f""" - SELECT key, value, created_at, updated_at - FROM store - WHERE prefix = %s AND key IN ({keys_to_query}) + _, keys = zip(*items, strict=True) + this_refresh_ttls = refresh_ttls[namespace] + + query = """ + WITH passed_in AS ( + SELECT unnest(%s::text[]) AS key, + unnest(%s::bool[]) AS do_refresh + ), + updated AS ( + UPDATE store s + SET expires_at = NOW() + (s.ttl_minutes || ' minutes')::interval + FROM passed_in p + WHERE s.prefix = %s + AND s.key = p.key + AND p.do_refresh = TRUE + AND s.ttl_minutes IS NOT NULL + RETURNING s.key + ) + SELECT s.key, s.value, s.created_at, s.updated_at + FROM store s + JOIN passed_in p ON s.key = p.key + WHERE s.prefix = %s """ - params = (_namespace_to_text(namespace), *keys) + ns_text = _namespace_to_text(namespace) + params = ( + list(keys), # -> unnest(%s::text[]) + list(this_refresh_ttls), # -> unnest(%s::bool[]) + ns_text, # -> prefix = %s (for UPDATE) + ns_text, # -> prefix = %s (for final SELECT) + ) results.append((query, params, namespace, items)) + return results def _prepare_batch_PUT_queries( @@ -246,9 +292,8 @@ class BasePostgresStore(Generic[C]): put_ops: Sequence[tuple[int, PutOp]], ) -> tuple[ list[tuple[str, Sequence]], - Optional[tuple[str, Sequence[tuple[str, str, str, str]]]], + tuple[str, Sequence[tuple[str, str, str, str]]] | None, ]: - # Last-write wins dedupped_ops: dict[tuple[tuple[str, ...], str], PutOp] = {} for _, op in put_ops: dedupped_ops[(op.namespace, op.key)] = op @@ -274,23 +319,32 @@ class BasePostgresStore(Generic[C]): ) params = (_namespace_to_text(namespace), *keys) queries.append((query, params)) - embedding_request: Optional[tuple[str, Sequence[tuple[str, str, str, str]]]] = ( - None - ) + embedding_request: tuple[str, Sequence[tuple[str, str, str, str]]] | None = None if inserts: values = [] insertion_params = [] vector_values = [] embedding_request_params = [] + # Handle TTL expiration # First handle main store insertions for op in inserts: - values.append("(%s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)") + if op.ttl is not None: + expires_at_str = f"NOW() + INTERVAL '{op.ttl*60} seconds'" + ttl_minutes = op.ttl + else: + expires_at_str = "NULL" + ttl_minutes = None + + values.append( + f"(%s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, {expires_at_str}, %s)" + ) insertion_params.extend( [ _namespace_to_text(op.namespace), op.key, Jsonb(cast(dict, op.value)), + ttl_minutes, ] ) @@ -304,7 +358,7 @@ class BasePostgresStore(Generic[C]): k = op.key if op.index is None: - paths = self.index_config["__tokenized_fields"] + paths = cast(dict, self.index_config)["__tokenized_fields"] else: paths = [(ix, tokenize_path(ix)) for ix in op.index] @@ -319,11 +373,13 @@ class BasePostgresStore(Generic[C]): values_str = ",".join(values) query = f""" - INSERT INTO store (prefix, key, value, created_at, updated_at) + INSERT INTO store (prefix, key, value, created_at, updated_at, expires_at, ttl_minutes) VALUES {values_str} ON CONFLICT (prefix, key) DO UPDATE SET value = EXCLUDED.value, - updated_at = CURRENT_TIMESTAMP + updated_at = CURRENT_TIMESTAMP, + expires_at = EXCLUDED.expires_at, + ttl_minutes = EXCLUDED.ttl_minutes """ queries.append((query, insertion_params)) @@ -344,95 +400,109 @@ class BasePostgresStore(Generic[C]): self, search_ops: Sequence[tuple[int, SearchOp]], ) -> tuple[ - list[tuple[str, list[Union[None, str, list[float]]]]], # queries, params + list[tuple[str, list[None | str | list[float]]]], # queries, params list[tuple[int, str]], # idx, query_text pairs to embed ]: + """ + Build per-SearchOp SQL queries (with optional TTL refresh) plus embedding requests. + Returns: + - queries: list of (SQL, param_list) + - embedding_requests: list of (original_index_in_search_ops, text_query) + """ + queries = [] embedding_requests = [] for idx, (_, op) in enumerate(search_ops): - # Build filter conditions first filter_params = [] - filter_conditions = [] + filter_clauses = [] if op.filter: for key, value in op.filter.items(): if isinstance(value, dict): for op_name, val in value.items(): - condition, filter_params_ = self._get_filter_condition( + condition, params_ = self._get_filter_condition( key, op_name, val ) - filter_conditions.append(condition) - filter_params.extend(filter_params_) + filter_clauses.append(condition) + filter_params.extend(params_) else: - filter_conditions.append("value->%s = %s::jsonb") - filter_params.extend([key, json.dumps(value)]) + filter_clauses.append("value->%s = %s::jsonb") + filter_params.extend([key, orjson.dumps(value).decode("utf-8")]) + + ns_condition = "TRUE" + ns_param: Sequence[str] | None = None + if op.namespace_prefix: + ns_condition = "store.prefix LIKE %s" + ns_param = (f"{_namespace_to_text(op.namespace_prefix)}%",) + else: + ns_param = () + + extra_filters = ( + " AND " + " AND ".join(filter_clauses) if filter_clauses else "" + ) - # Vector search branch if op.query and self.index_config: + # We'll embed the text later, so record the request. embedding_requests.append((idx, op.query)) score_operator, post_operator = get_distance_operator(self) + post_operator = post_operator.replace("scored", "uniq") vector_type = ( cast(PostgresIndexConfig, self.index_config) .get("ann_index_config", {}) .get("vector_type", "vector") ) + # For hamming bit vectors, or “regular” vectors if ( vector_type == "bit" - and self.index_config.get("distance_type") == "hamming" + and cast(dict, self.index_config).get("distance_type") == "hamming" ): score_operator = score_operator % ( "%s", - self.index_config["dims"], + cast(dict, self.index_config)["dims"], ) else: - score_operator = score_operator % ( - "%s", - vector_type, - ) + score_operator = score_operator % ("%s", vector_type) - vectors_per_doc_estimate = self.index_config["__estimated_num_vectors"] + vectors_per_doc_estimate = cast(dict, self.index_config)[ + "__estimated_num_vectors" + ] expanded_limit = (op.limit * vectors_per_doc_estimate * 2) + 1 - # Vector search with CTE for proper score handling - filter_str = ( - "" - if not filter_conditions - else " AND " + " AND ".join(filter_conditions) - ) - if op.namespace_prefix: - prefix_filter_str = f"WHERE s.prefix LIKE %s {filter_str} " - ns_args: Sequence = (f"{_namespace_to_text(op.namespace_prefix)}%",) - else: - ns_args = () - if filter_str: - prefix_filter_str = f"WHERE {filter_str} " - else: - prefix_filter_str = "" - - base_query = f""" - WITH scored AS ( - SELECT s.prefix, s.key, s.value, s.created_at, s.updated_at, {score_operator} AS neg_score - FROM store s - JOIN store_vectors sv ON s.prefix = sv.prefix AND s.key = sv.key - {prefix_filter_str} - ORDER BY {score_operator} ASC + # “sub_scored” does the main vector search + # Then we do DISTINCT ON to drop duplicates if your store can have them + # Finally we limit & offset + vector_search_cte = f""" + SELECT store.prefix, store.key, store.value, store.created_at, store.updated_at, + {score_operator} AS neg_score + FROM store + JOIN store_vectors sv ON store.prefix = sv.prefix AND store.key = sv.key + WHERE {ns_condition} {extra_filters} + ORDER BY {score_operator} ASC LIMIT %s - ) - SELECT * FROM ( - SELECT DISTINCT ON (prefix, key) - prefix, key, value, created_at, updated_at, {post_operator} as score - FROM scored - ORDER BY prefix, key, score DESC - ) AS unique_docs - ORDER BY score DESC - LIMIT %s - OFFSET %s - """ - params = [ - PLACEHOLDER, # Vector placeholder - *ns_args, + """ + + search_results_sql = f""" + WITH scored AS ( + {vector_search_cte} + ) + SELECT uniq.prefix, uniq.key, uniq.value, uniq.created_at, uniq.updated_at, + {post_operator} AS score + FROM ( + SELECT DISTINCT ON (scored.prefix, scored.key) + scored.prefix, scored.key, scored.value, scored.created_at, scored.updated_at, scored.neg_score + FROM scored + ORDER BY scored.prefix, scored.key, scored.neg_score ASC + ) uniq + ORDER BY score DESC + LIMIT %s + OFFSET %s + """ + + search_results_params = [ + PLACEHOLDER, + *ns_param, *filter_params, PLACEHOLDER, expanded_limit, @@ -440,24 +510,45 @@ class BasePostgresStore(Generic[C]): op.offset, ] - # Regular search branch else: - base_query = """ - SELECT prefix, key, value, created_at, updated_at - FROM store - WHERE prefix LIKE %s - """ - params = [f"{_namespace_to_text(op.namespace_prefix)}%"] + base_query = f""" + SELECT store.prefix, store.key, store.value, store.created_at, store.updated_at, 0 AS score + FROM store + WHERE {ns_condition} {extra_filters} + ORDER BY store.updated_at DESC + LIMIT %s + OFFSET %s + """ + search_results_sql = base_query + search_results_params = [ + *ns_param, + *filter_params, + op.limit, + op.offset, + ] - if filter_conditions: - params.extend(filter_params) - base_query += " AND " + " AND ".join(filter_conditions) - - base_query += " ORDER BY updated_at DESC" - base_query += " LIMIT %s OFFSET %s" - params.extend([op.limit, op.offset]) - - queries.append((base_query, params)) + if op.refresh_ttl: + # Wrap entire primary query in a CTE, then perform "update_at" + final_sql = f""" + WITH search_results AS ( + {search_results_sql} + ), + updated AS ( + UPDATE store s + SET expires_at = NOW() + (s.ttl_minutes || ' minutes')::interval + FROM search_results sr + WHERE s.prefix = sr.prefix + AND s.key = sr.key + AND s.ttl_minutes IS NOT NULL + ) + SELECT sr.prefix, sr.key, sr.value, sr.created_at, sr.updated_at, sr.score + FROM search_results sr + """ + final_params = search_results_params[:] # copy + else: + final_sql = search_results_sql + final_params = search_results_params + queries.append((final_sql, final_params)) return queries, embedding_requests @@ -612,7 +703,9 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): "supports_pipeline", "index_config", "embeddings", + "supports_ttl", ) + supports_ttl: bool = True def __init__( self, @@ -689,6 +782,47 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): else: yield cls(conn, index=index) + def sweep_ttl(self) -> int: + """Delete expired store items based on TTL. + + Returns: + int: The number of deleted items. + """ + with self._cursor() as cur: + cur.execute( + """ + DELETE FROM store + WHERE expires_at IS NOT NULL AND expires_at < NOW() + """ + ) + deleted_count = cur.rowcount + return deleted_count + + def start_ttl_sweeper(self, sweep_interval_minutes: Optional[int] = None) -> None: + """Periodically delete expired store items based on TTL.""" + if not self.ttl_config: + return + sweep_interval_minutes_ = float( + cast( + float, + sweep_interval_minutes + or self.ttl_config.get("sweep_interval_minutes") + or 5, + ) + ) + logger.info( + f"Starting store TTL sweeper with interval {sweep_interval_minutes_} minutes", + ) + + while True: + time.sleep(sweep_interval_minutes_ * 60) + try: + expired_items = self.sweep_ttl() + if expired_items > 0: + logger.info(f"Store swept {expired_items} expired items") + except Exception as exc: + logger.exception("Store TTL sweep iteration failed", exc_info=exc) + @contextmanager def _cursor(self, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]: """Create a database cursor as a context manager. From 9741d9bdf001117e041caf7cc845fa99390a94e6 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Fri, 14 Mar 2025 09:56:21 -0700 Subject: [PATCH 2/4] Add tests for sweeper (sync) --- .../langgraph/store/postgres/aio.py | 138 ++++++-- .../langgraph/store/postgres/base.py | 157 +++++++-- .../tests/test_async_store.py | 62 +++- libs/checkpoint-postgres/tests/test_store.py | 318 +++++++++++------- .../langgraph/store/base/__init__.py | 6 + libs/cli/langgraph_cli/config.py | 6 + 6 files changed, 493 insertions(+), 194 deletions(-) diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py index 5ffcaaf5a..a5e7604cd 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py @@ -2,6 +2,7 @@ import asyncio import logging from collections.abc import AsyncIterator, Iterable, Sequence from contextlib import asynccontextmanager +from types import TracebackType from typing import Any, Callable, Optional, Union, cast import orjson @@ -25,6 +26,7 @@ from langgraph.store.postgres.base import ( PoolConfig, PostgresIndexConfig, Row, + TTLConfig, _decode_ns_bytes, _ensure_index_config, _group_ops, @@ -106,6 +108,11 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con Semantic search is disabled by default. You can enable it by providing an `index` configuration when creating the store. Without this configuration, all `index` arguments passed to `put` or `aput` will have no effect. + + Note: + If you provide a TTL configuration, you must explicitly call `start_ttl_sweeper()` to begin + the background task that removes expired items. Call `stop_ttl_sweeper()` to properly + clean up resources when you're done with the store. """ __slots__ = ( @@ -115,7 +122,9 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con "supports_pipeline", "index_config", "embeddings", - "supports_ttl", + "ttl_config", + "_ttl_sweeper_task", + "_ttl_stop_event", ) supports_ttl: bool = True @@ -128,6 +137,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]] ] = None, index: Optional[PostgresIndexConfig] = None, + ttl: Optional[TTLConfig] = None, ) -> None: if isinstance(conn, AsyncConnectionPool) and pipe is not None: raise ValueError( @@ -143,10 +153,13 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con self.index_config = index if self.index_config: self.embeddings, self.index_config = _ensure_index_config(self.index_config) - else: self.embeddings = None + self.ttl_config = ttl + self._ttl_sweeper_task: Optional[asyncio.Task[None]] = None + self._ttl_stop_event = asyncio.Event() + async def abatch(self, ops: Iterable[Op]) -> list[Result]: grouped_ops, num_ops = _group_ops(ops) results: list[Result] = [None] * num_ops @@ -169,6 +182,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con pipeline: bool = False, pool_config: Optional[PoolConfig] = None, index: Optional[PostgresIndexConfig] = None, + ttl: Optional[TTLConfig] = None, ) -> AsyncIterator["AsyncPostgresStore"]: """Create a new AsyncPostgresStore instance from a connection string. @@ -200,16 +214,16 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con **cast(dict, pc), ), ) as pool: - yield cls(conn=pool, index=index) + yield cls(conn=pool, index=index, ttl=ttl) else: async with await AsyncConnection.connect( conn_string, autocommit=True, prepare_threshold=0, row_factory=dict_row ) as conn: if pipeline: async with conn.pipeline() as pipe: - yield cls(conn=conn, pipe=pipe, index=index) + yield cls(conn=conn, pipe=pipe, index=index, ttl=ttl) else: - yield cls(conn=conn, index=index) + yield cls(conn=conn, index=index, ttl=ttl) async def setup(self) -> None: """Set up the store database asynchronously. @@ -276,30 +290,100 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con async def start_ttl_sweeper( self, sweep_interval_minutes: Optional[int] = None - ) -> None: - """Periodically delete expired store items based on TTL.""" - if not self.ttl_config: - return - sweep_interval_minutes_ = float( - cast( - float, - sweep_interval_minutes - or self.ttl_config.get("sweep_interval_minutes") - or 5, - ) - ) - logger.info( - f"Starting store TTL sweeper with interval {sweep_interval_minutes_} minutes", - ) + ) -> asyncio.Task[None]: + """Periodically delete expired store items based on TTL. - while True: - await asyncio.sleep(sweep_interval_minutes_ * 60) + Returns: + Task that can be awaited or cancelled. + """ + if not self.ttl_config: + return asyncio.create_task(asyncio.sleep(0)) + + if self._ttl_sweeper_task is not None and not self._ttl_sweeper_task.done(): + return self._ttl_sweeper_task + + self._ttl_stop_event.clear() + + interval = float( + sweep_interval_minutes or self.ttl_config.get("sweep_interval_minutes") or 5 + ) + logger.info(f"Starting store TTL sweeper with interval {interval} minutes") + + async def _sweep_loop() -> None: + while not self._ttl_stop_event.is_set(): + try: + try: + await asyncio.wait_for( + self._ttl_stop_event.wait(), + timeout=interval * 60, + ) + break + except asyncio.TimeoutError: + pass + + expired_items = await self.sweep_ttl() + if expired_items > 0: + logger.info(f"Store swept {expired_items} expired items") + except asyncio.CancelledError: + break + except Exception as exc: + logger.exception("Store TTL sweep iteration failed", exc_info=exc) + + task = asyncio.create_task(_sweep_loop()) + task.set_name("ttl_sweeper") + self._ttl_sweeper_task = task + return task + + async def stop_ttl_sweeper(self, timeout: Optional[float] = None) -> bool: + """Stop the TTL sweeper task if it's running. + + Args: + timeout: Maximum time to wait for the task to stop, in seconds. + If None, wait indefinitely. + + Returns: + bool: True if the task was successfully stopped or wasn't running, + False if the timeout was reached before the task stopped. + """ + if self._ttl_sweeper_task is None or self._ttl_sweeper_task.done(): + return True + + logger.info("Stopping TTL sweeper task") + self._ttl_stop_event.set() + + if timeout is not None: try: - expired_items = await self.sweep_ttl() - if expired_items > 0: - logger.info(f"Store swept {expired_items} expired items") - except Exception as exc: - logger.exception("Store TTL sweep iteration failed", exc_info=exc) + await asyncio.wait_for(self._ttl_sweeper_task, timeout=timeout) + success = True + except asyncio.TimeoutError: + success = False + else: + await self._ttl_sweeper_task + success = True + + if success: + self._ttl_sweeper_task = None + logger.info("TTL sweeper task stopped") + else: + logger.warning("Timed out waiting for TTL sweeper task to stop") + + return success + + async def __aenter__(self) -> "AsyncPostgresStore": + return self + + async def __aexit__( + self, + exc_type: Optional[type[BaseException]], + exc_val: Optional[BaseException], + exc_tb: Optional["TracebackType"], + ) -> None: + # Ensure the TTL sweeper task is stopped when exiting the context + if hasattr(self, "_ttl_sweeper_task") and self._ttl_sweeper_task is not None: + # Set the event to signal the task to stop + self._ttl_stop_event.set() + # We don't wait for the task to complete here to avoid blocking + # The task will clean up itself gracefully async def _execute_batch( self, diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/base.py b/libs/checkpoint-postgres/langgraph/store/postgres/base.py index 954b8e88d..d88ae19ed 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/base.py @@ -1,8 +1,8 @@ import asyncio +import concurrent.futures import json import logging import threading -import time from collections import defaultdict from collections.abc import Iterable, Iterator, Sequence from contextlib import contextmanager @@ -81,7 +81,8 @@ CREATE INDEX CONCURRENTLY IF NOT EXISTS store_prefix_idx ON store USING btree (p ALTER TABLE store ADD COLUMN expires_at TIMESTAMP WITH TIME ZONE, ADD COLUMN ttl_minutes INT; - +""", + """ -- Add indexes for efficient TTL sweeping CREATE INDEX idx_store_expires_at ON store (expires_at) WHERE expires_at IS NOT NULL; @@ -253,7 +254,7 @@ class BasePostgresStore(Generic[C]): results = [] for namespace, items in namespace_groups.items(): - _, keys = zip(*items, strict=True) + _, keys = zip(*items) this_refresh_ttls = refresh_ttls[namespace] query = """ @@ -292,7 +293,7 @@ class BasePostgresStore(Generic[C]): put_ops: Sequence[tuple[int, PutOp]], ) -> tuple[ list[tuple[str, Sequence]], - tuple[str, Sequence[tuple[str, str, str, str]]] | None, + Optional[tuple[str, Sequence[tuple[str, str, str, str]]]], ]: dedupped_ops: dict[tuple[tuple[str, ...], str], PutOp] = {} for _, op in put_ops: @@ -319,7 +320,9 @@ class BasePostgresStore(Generic[C]): ) params = (_namespace_to_text(namespace), *keys) queries.append((query, params)) - embedding_request: tuple[str, Sequence[tuple[str, str, str, str]]] | None = None + embedding_request: Optional[tuple[str, Sequence[tuple[str, str, str, str]]]] = ( + None + ) if inserts: values = [] insertion_params = [] @@ -400,7 +403,7 @@ class BasePostgresStore(Generic[C]): self, search_ops: Sequence[tuple[int, SearchOp]], ) -> tuple[ - list[tuple[str, list[None | str | list[float]]]], # queries, params + list[tuple[str, list[Union[None, str, list[float]]]]], # queries, params list[tuple[int, str]], # idx, query_text pairs to embed ]: """ @@ -412,7 +415,6 @@ class BasePostgresStore(Generic[C]): queries = [] embedding_requests = [] - for idx, (_, op) in enumerate(search_ops): filter_params = [] filter_clauses = [] @@ -430,7 +432,7 @@ class BasePostgresStore(Generic[C]): filter_params.extend([key, orjson.dumps(value).decode("utf-8")]) ns_condition = "TRUE" - ns_param: Sequence[str] | None = None + ns_param: Optional[Sequence[Union[str]]] = None if op.namespace_prefix: ns_condition = "store.prefix LIKE %s" ns_param = (f"{_namespace_to_text(op.namespace_prefix)}%",) @@ -512,7 +514,7 @@ class BasePostgresStore(Generic[C]): else: base_query = f""" - SELECT store.prefix, store.key, store.value, store.created_at, store.updated_at, 0 AS score + SELECT store.prefix, store.key, store.value, store.created_at, store.updated_at, NULL AS score FROM store WHERE {ns_condition} {extra_filters} ORDER BY store.updated_at DESC @@ -625,7 +627,7 @@ class BasePostgresStore(Generic[C]): class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): - """Postgres-backed store with optional vector search using pgvector. + """Postgres-backed store with oktional vector search using pgvector. !!! example "Examples" Basic setup and usage: @@ -694,6 +696,11 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): Make sure to call `setup()` before first use to create necessary tables and indexes. The pgvector extension must be available to use vector search. + Note: + If you provide a TTL configuration, you must explicitly call `start_ttl_sweeper()` to begin + the background thread that removes expired items. Call `stop_ttl_sweeper()` to properly + clean up resources when you're done with the store. + """ __slots__ = ( @@ -703,7 +710,8 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): "supports_pipeline", "index_config", "embeddings", - "supports_ttl", + "_ttl_sweeper_thread", + "_ttl_stop_event", ) supports_ttl: bool = True @@ -730,6 +738,8 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): else: self.embeddings = None self.ttl_config = ttl + self._ttl_sweeper_thread: Optional[threading.Thread] = None + self._ttl_stop_event = threading.Event() @classmethod @contextmanager @@ -740,6 +750,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): pipeline: bool = False, pool_config: Optional[PoolConfig] = None, index: Optional[PostgresIndexConfig] = None, + ttl: Optional[TTLConfig] = None, ) -> Iterator["PostgresStore"]: """Create a new PostgresStore instance from a connection string. @@ -771,16 +782,16 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): **cast(dict, pc), ), ) as pool: - yield cls(conn=pool, index=index) + yield cls(conn=pool, index=index, ttl=ttl) else: with Connection.connect( conn_string, autocommit=True, prepare_threshold=0, row_factory=dict_row ) as conn: if pipeline: with conn.pipeline() as pipe: - yield cls(conn, pipe=pipe, index=index) + yield cls(conn, pipe=pipe, index=index, ttl=ttl) else: - yield cls(conn, index=index) + yield cls(conn, index=index, ttl=ttl) def sweep_ttl(self) -> int: """Delete expired store items based on TTL. @@ -798,30 +809,96 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): deleted_count = cur.rowcount return deleted_count - def start_ttl_sweeper(self, sweep_interval_minutes: Optional[int] = None) -> None: - """Periodically delete expired store items based on TTL.""" - if not self.ttl_config: - return - sweep_interval_minutes_ = float( - cast( - float, - sweep_interval_minutes - or self.ttl_config.get("sweep_interval_minutes") - or 5, - ) - ) - logger.info( - f"Starting store TTL sweeper with interval {sweep_interval_minutes_} minutes", - ) + def start_ttl_sweeper( + self, sweep_interval_minutes: Optional[int] = None + ) -> concurrent.futures.Future[None]: + """Periodically delete expired store items based on TTL. - while True: - time.sleep(sweep_interval_minutes_ * 60) + Returns: + Future that can be waited on or cancelled. + """ + if not self.ttl_config: + future: concurrent.futures.Future[None] = concurrent.futures.Future() + future.set_result(None) + return future + + if self._ttl_sweeper_thread and self._ttl_sweeper_thread.is_alive(): + logger.info("TTL sweeper thread is already running") + # Return a future that can be used to cancel the existing thread + future = concurrent.futures.Future() + future.add_done_callback( + lambda f: self._ttl_stop_event.set() if f.cancelled() else None + ) + return future + + self._ttl_stop_event.clear() + + interval = float( + sweep_interval_minutes or self.ttl_config.get("sweep_interval_minutes") or 5 + ) + logger.info(f"Starting store TTL sweeper with interval {interval} minutes") + + future = concurrent.futures.Future() + + def _sweep_loop() -> None: try: - expired_items = self.sweep_ttl() - if expired_items > 0: - logger.info(f"Store swept {expired_items} expired items") + while not self._ttl_stop_event.is_set(): + if self._ttl_stop_event.wait(interval * 60): + break + + try: + expired_items = self.sweep_ttl() + if expired_items > 0: + logger.info(f"Store swept {expired_items} expired items") + except Exception as exc: + logger.exception( + "Store TTL sweep iteration failed", exc_info=exc + ) + future.set_result(None) except Exception as exc: - logger.exception("Store TTL sweep iteration failed", exc_info=exc) + future.set_exception(exc) + + thread = threading.Thread(target=_sweep_loop, daemon=True, name="ttl-sweeper") + self._ttl_sweeper_thread = thread + thread.start() + + future.add_done_callback( + lambda f: self._ttl_stop_event.set() if f.cancelled() else None + ) + return future + + def stop_ttl_sweeper(self, timeout: Optional[float] = None) -> bool: + """Stop the TTL sweeper thread if it's running. + + Args: + timeout: Maximum time to wait for the thread to stop, in seconds. + If None, wait indefinitely. + + Returns: + bool: True if the thread was successfully stopped or wasn't running, + False if the timeout was reached before the thread stopped. + """ + if not self._ttl_sweeper_thread or not self._ttl_sweeper_thread.is_alive(): + return True + + logger.info("Stopping TTL sweeper thread") + self._ttl_stop_event.set() + + self._ttl_sweeper_thread.join(timeout) + success = not self._ttl_sweeper_thread.is_alive() + + if success: + self._ttl_sweeper_thread = None + logger.info("TTL sweeper thread stopped") + else: + logger.warning("Timed out waiting for TTL sweeper thread to stop") + + return success + + def __del__(self) -> None: + """Ensure the TTL sweeper thread is stopped when the object is garbage collected.""" + if hasattr(self, "_ttl_stop_event") and hasattr(self, "_ttl_sweeper_thread"): + self.stop_ttl_sweeper(timeout=0.1) @contextmanager def _cursor(self, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]: @@ -1020,8 +1097,14 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): with self._cursor() as cur: version = _get_version(cur, table="store_migrations") for v, sql in enumerate(self.MIGRATIONS[version + 1 :], start=version + 1): - cur.execute(sql) - cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,)) + try: + cur.execute(sql) + cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,)) + except Exception as e: + logger.error( + f"Failed to apply migration {v}.\nSql={sql}\nError={e}" + ) + raise if self.index_config: version = _get_version(cur, table="vector_migrations") diff --git a/libs/checkpoint-postgres/tests/test_async_store.py b/libs/checkpoint-postgres/tests/test_async_store.py index 068ec1502..09502403d 100644 --- a/libs/checkpoint-postgres/tests/test_async_store.py +++ b/libs/checkpoint-postgres/tests/test_async_store.py @@ -26,6 +26,9 @@ from tests.conftest import ( CharacterEmbeddings, ) +TTL_SECONDS = 6 +TTL_MINUTES = TTL_SECONDS / 60 + @pytest.fixture(scope="function", params=["default", "pipe", "pool"]) async def store(request) -> AsyncIterator[AsyncPostgresStore]: @@ -42,28 +45,52 @@ async def store(request) -> AsyncIterator[AsyncPostgresStore]: conn_string = f"{uri_base}/{database}{query_params}" admin_conn_string = DEFAULT_URI - + ttl_config = { + "default_ttl": TTL_MINUTES, + "refresh_on_read": True, + "sweep_interval_minutes": TTL_MINUTES / 2, + } async with await AsyncConnection.connect( admin_conn_string, autocommit=True ) as conn: await conn.execute(f"CREATE DATABASE {database}") try: - async with AsyncPostgresStore.from_conn_string(conn_string) as store: + async with AsyncPostgresStore.from_conn_string( + conn_string, ttl=ttl_config + ) as store: + store.MIGRATIONS = [ + ( + mig.replace( + "ADD COLUMN ttl_minutes INT;", "ADD COLUMN ttl_minutes FLOAT;" + ) + if isinstance(mig, str) + else mig + ) + for mig in store.MIGRATIONS + ] await store.setup() if request.param == "pipe": async with AsyncPostgresStore.from_conn_string( - conn_string, pipeline=True + conn_string, pipeline=True, ttl=ttl_config ) as store: + await store.start_ttl_sweeper() yield store + await store.stop_ttl_sweeper() elif request.param == "pool": async with AsyncPostgresStore.from_conn_string( - conn_string, pool_config={"min_size": 1, "max_size": 10} + conn_string, pool_config={"min_size": 1, "max_size": 10}, ttl=ttl_config ) as store: + await store.start_ttl_sweeper() yield store + await store.stop_ttl_sweeper() else: # default - async with AsyncPostgresStore.from_conn_string(conn_string) as store: + async with AsyncPostgresStore.from_conn_string( + conn_string, ttl=ttl_config + ) as store: + await store.start_ttl_sweeper() yield store + await store.stop_ttl_sweeper() finally: async with await AsyncConnection.connect( admin_conn_string, autocommit=True @@ -635,3 +662,28 @@ async def test_search_sorting( assert len(set(r.key for r in results)) == 10 assert results[0].key == "M" assert results[0].score > results[1].score + + +async def test_store_ttl(store): + # Assumes a TTL of 1 minute = 60 seconds + ns = ("foo",) + await store.start_ttl_sweeper() + await store.aput( + ns, + key="item1", + value={"foo": "bar"}, + ttl=TTL_MINUTES, # type: ignore + ) + await asyncio.sleep(TTL_SECONDS - 2) + res = await store.aget(ns, key="item1", refresh_ttl=True) + assert res is not None + await asyncio.sleep(TTL_SECONDS - 2) + results = await store.asearch(ns, query="foo", refresh_ttl=True) + assert len(results) == 1 + await asyncio.sleep(TTL_SECONDS - 2) + res = await store.aget(ns, key="item1", refresh_ttl=False) + assert res is not None + await asyncio.sleep(TTL_SECONDS - 1) + # Now has been (TTL_SECONDS-2)*2 > TTL_SECONDS + TTL_SECONDS/2 + results = await store.asearch(ns, query="bar", refresh_ttl=False) + assert len(results) == 0 diff --git a/libs/checkpoint-postgres/tests/test_store.py b/libs/checkpoint-postgres/tests/test_store.py index 50d697962..4ee37484e 100644 --- a/libs/checkpoint-postgres/tests/test_store.py +++ b/libs/checkpoint-postgres/tests/test_store.py @@ -1,6 +1,7 @@ # type: ignore import re +import time from contextlib import contextmanager from typing import Any, Optional from uuid import uuid4 @@ -24,6 +25,9 @@ from tests.conftest import ( CharacterEmbeddings, ) +TTL_SECONDS = 6 +TTL_MINUTES = TTL_SECONDS / 60 + @pytest.fixture(scope="function", params=["default", "pipe", "pool"]) def store(request) -> PostgresStore: @@ -32,29 +36,58 @@ def store(request) -> PostgresStore: uri_base = "/".join(uri_parts[:-1]) query_params = "" if "?" in uri_parts[-1]: - db_name, query_params = uri_parts[-1].split("?", 1) + _, query_params = uri_parts[-1].split("?", 1) query_params = "?" + query_params conn_string = f"{uri_base}/{database}{query_params}" admin_conn_string = DEFAULT_URI - + ttl_config = { + "default_ttl": TTL_MINUTES, + "refresh_on_read": True, + "sweep_interval_minutes": TTL_MINUTES / 2, + } with Connection.connect(admin_conn_string, autocommit=True) as conn: conn.execute(f"CREATE DATABASE {database}") try: - with PostgresStore.from_conn_string(conn_string) as store: + with PostgresStore.from_conn_string(conn_string, ttl=ttl_config) as store: + store.MIGRATIONS = [ + ( + mig.replace( + "ADD COLUMN ttl_minutes INT;", "ADD COLUMN ttl_minutes FLOAT;" + ) + if isinstance(mig, str) + else mig + ) + for mig in store.MIGRATIONS + ] store.setup() if request.param == "pipe": - with PostgresStore.from_conn_string(conn_string, pipeline=True) as store: + with PostgresStore.from_conn_string( + conn_string, + pipeline=True, + ttl=ttl_config, + ) as store: + store.start_ttl_sweeper() yield store + + store.stop_ttl_sweeper() elif request.param == "pool": with PostgresStore.from_conn_string( - conn_string, pool_config={"min_size": 1, "max_size": 10} + conn_string, + pool_config={"min_size": 1, "max_size": 10}, + ttl=ttl_config, ) as store: + store.start_ttl_sweeper() yield store + + store.stop_ttl_sweeper() else: # default - with PostgresStore.from_conn_string(conn_string) as store: + with PostgresStore.from_conn_string(conn_string, ttl=ttl_config) as store: + store.start_ttl_sweeper() yield store + + store.stop_ttl_sweeper() finally: with Connection.connect(admin_conn_string, autocommit=True) as conn: conn.execute(f"DROP DATABASE {database}") @@ -220,134 +253,127 @@ def test_batch_list_namespaces_ops(store: PostgresStore) -> None: assert all(ns[-1] == "public" for ns in results[2]) -class TestPostgresStore: - @pytest.fixture(autouse=True) - def setup(self) -> None: - with PostgresStore.from_conn_string(DEFAULT_URI) as store: - store.setup() +def test_basic_store_ops(store) -> None: + namespace = ("test", "documents") + item_id = "doc1" + item_value = {"title": "Test Document", "content": "Hello, World!"} - def test_basic_store_ops(self) -> None: - with PostgresStore.from_conn_string(DEFAULT_URI) as store: - namespace = ("test", "documents") - item_id = "doc1" - item_value = {"title": "Test Document", "content": "Hello, World!"} + store.put(namespace, item_id, item_value) + item = store.get(namespace, item_id) - store.put(namespace, item_id, item_value) - item = store.get(namespace, item_id) + assert item + assert item.namespace == namespace + assert item.key == item_id + assert item.value == item_value - assert item - assert item.namespace == namespace - assert item.key == item_id - assert item.value == item_value + # Test update + updated_value = {"title": "Updated Document", "content": "Hello, Updated!"} + store.put(namespace, item_id, updated_value) + updated_item = store.get(namespace, item_id) - # Test update - updated_value = {"title": "Updated Document", "content": "Hello, Updated!"} - store.put(namespace, item_id, updated_value) - updated_item = store.get(namespace, item_id) + assert updated_item.value == updated_value + assert updated_item.updated_at > item.updated_at - assert updated_item.value == updated_value - assert updated_item.updated_at > item.updated_at + # Test get from non-existent namespace + different_namespace = ("test", "other_documents") + item_in_different_namespace = store.get(different_namespace, item_id) + assert item_in_different_namespace is None - # Test get from non-existent namespace - different_namespace = ("test", "other_documents") - item_in_different_namespace = store.get(different_namespace, item_id) - assert item_in_different_namespace is None + # Test delete + store.delete(namespace, item_id) + deleted_item = store.get(namespace, item_id) + assert deleted_item is None - # Test delete - store.delete(namespace, item_id) - deleted_item = store.get(namespace, item_id) - assert deleted_item is None - def test_list_namespaces(self) -> None: - with PostgresStore.from_conn_string(DEFAULT_URI) as store: - # Create test data with various namespaces - test_namespaces = [ - ("test", "documents", "public"), - ("test", "documents", "private"), - ("test", "images", "public"), - ("test", "images", "private"), - ("prod", "documents", "public"), - ("prod", "documents", "private"), - ] +def test_list_namespaces(store) -> None: + # Create test data with various namespaces + test_namespaces = [ + ("test", "documents", "public"), + ("test", "documents", "private"), + ("test", "images", "public"), + ("test", "images", "private"), + ("prod", "documents", "public"), + ("prod", "documents", "private"), + ] - # Insert test data - for namespace in test_namespaces: - store.put(namespace, "dummy", {"content": "dummy"}) + # Insert test data + for namespace in test_namespaces: + store.put(namespace, "dummy", {"content": "dummy"}) - # Test listing with various filters - all_namespaces = store.list_namespaces() - assert len(all_namespaces) == len(test_namespaces) + # Test listing with various filters + all_namespaces = store.list_namespaces() + assert len(all_namespaces) == len(test_namespaces) - # Test prefix filtering - test_prefix_namespaces = store.list_namespaces(prefix=["test"]) - assert len(test_prefix_namespaces) == 4 - assert all(ns[0] == "test" for ns in test_prefix_namespaces) + # Test prefix filtering + test_prefix_namespaces = store.list_namespaces(prefix=["test"]) + assert len(test_prefix_namespaces) == 4 + assert all(ns[0] == "test" for ns in test_prefix_namespaces) - # Test suffix filtering - public_namespaces = store.list_namespaces(suffix=["public"]) - assert len(public_namespaces) == 3 - assert all(ns[-1] == "public" for ns in public_namespaces) + # Test suffix filtering + public_namespaces = store.list_namespaces(suffix=["public"]) + assert len(public_namespaces) == 3 + assert all(ns[-1] == "public" for ns in public_namespaces) - # Test max depth - depth_2_namespaces = store.list_namespaces(max_depth=2) - assert all(len(ns) <= 2 for ns in depth_2_namespaces) + # Test max depth + depth_2_namespaces = store.list_namespaces(max_depth=2) + assert all(len(ns) <= 2 for ns in depth_2_namespaces) - # Test pagination - paginated_namespaces = store.list_namespaces(limit=3) - assert len(paginated_namespaces) == 3 + # Test pagination + paginated_namespaces = store.list_namespaces(limit=3) + assert len(paginated_namespaces) == 3 - # Cleanup - for namespace in test_namespaces: - store.delete(namespace, "dummy") + # Cleanup + for namespace in test_namespaces: + store.delete(namespace, "dummy") - def test_search(self) -> None: - with PostgresStore.from_conn_string(DEFAULT_URI) as store: - # Create test data - test_data = [ - ( - ("test", "docs"), - "doc1", - {"title": "First Doc", "author": "Alice", "tags": ["important"]}, - ), - ( - ("test", "docs"), - "doc2", - {"title": "Second Doc", "author": "Bob", "tags": ["draft"]}, - ), - ( - ("test", "images"), - "img1", - {"title": "Image 1", "author": "Alice", "tags": ["final"]}, - ), - ] - for namespace, key, value in test_data: - store.put(namespace, key, value) +def test_search(store) -> None: + # Create test data + test_data = [ + ( + ("test", "docs"), + "doc1", + {"title": "First Doc", "author": "Alice", "tags": ["important"]}, + ), + ( + ("test", "docs"), + "doc2", + {"title": "Second Doc", "author": "Bob", "tags": ["draft"]}, + ), + ( + ("test", "images"), + "img1", + {"title": "Image 1", "author": "Alice", "tags": ["final"]}, + ), + ] - # Test basic search - all_items = store.search(["test"]) - assert len(all_items) == 3 + for namespace, key, value in test_data: + store.put(namespace, key, value) - # Test namespace filtering - docs_items = store.search(["test", "docs"]) - assert len(docs_items) == 2 - assert all(item.namespace == ("test", "docs") for item in docs_items) + # Test basic search + all_items = store.search(["test"]) + assert len(all_items) == 3 - # Test value filtering - alice_items = store.search(["test"], filter={"author": "Alice"}) - assert len(alice_items) == 2 - assert all(item.value["author"] == "Alice" for item in alice_items) + # Test namespace filtering + docs_items = store.search(["test", "docs"]) + assert len(docs_items) == 2 + assert all(item.namespace == ("test", "docs") for item in docs_items) - # Test pagination - paginated_items = store.search(["test"], limit=2) - assert len(paginated_items) == 2 + # Test value filtering + alice_items = store.search(["test"], filter={"author": "Alice"}) + assert len(alice_items) == 2 + assert all(item.value["author"] == "Alice" for item in alice_items) - offset_items = store.search(["test"], offset=2) - assert len(offset_items) == 1 + # Test pagination + paginated_items = store.search(["test"], limit=2) + assert len(paginated_items) == 2 - # Cleanup - for namespace, key, _ in test_data: - store.delete(namespace, key) + offset_items = store.search(["test"], offset=2) + assert len(offset_items) == 1 + + # Cleanup + for namespace, key, _ in test_data: + store.delete(namespace, key) @contextmanager @@ -356,6 +382,7 @@ def _create_vector_store( distance_type: str, fake_embeddings: Embeddings, text_fields: Optional[list[str]] = None, + enable_ttl: bool = True, ) -> PostgresStore: """Create a store with vector search enabled.""" database = f"test_{uuid4().hex[:16]}" @@ -385,6 +412,7 @@ def _create_vector_store( with PostgresStore.from_conn_string( conn_string, index=index_config, + ttl={"default_ttl": 2, "refresh_on_read": True} if enable_ttl else None, ) as store: store.setup() yield store @@ -393,15 +421,19 @@ def _create_vector_store( conn.execute(f"DROP DATABASE {database}") +_vector_params = [ + (vector_type, distance_type, True) + for vector_type in VECTOR_TYPES + for distance_type in ( + ["hamming"] if vector_type == "bit" else ["l2", "inner_product", "cosine"] + ) +] +_vector_params += [(*_vector_params[-1][:2], False)] + + @pytest.fixture( scope="function", - params=[ - (vector_type, distance_type) - for vector_type in VECTOR_TYPES - for distance_type in ( - ["hamming"] if vector_type == "bit" else ["l2", "inner_product", "cosine"] - ) - ], + params=_vector_params, ids=lambda p: f"{p[0]}_{p[1]}", ) def vector_store( @@ -409,8 +441,10 @@ def vector_store( fake_embeddings: Embeddings, ) -> PostgresStore: """Create a store with vector search enabled.""" - vector_type, distance_type = request.param - with _create_vector_store(vector_type, distance_type, fake_embeddings) as store: + vector_type, distance_type, enable_ttl = request.param + with _create_vector_store( + vector_type, distance_type, fake_embeddings, enable_ttl=enable_ttl + ) as store: yield store @@ -474,7 +508,10 @@ def test_vector_update_with_embedding(vector_store: PostgresStore) -> None: assert not any(r.key == "doc4" for r in results_new) -def test_vector_search_with_filters(vector_store: PostgresStore) -> None: +@pytest.mark.parametrize("refresh_ttl", [True, False]) +def test_vector_search_with_filters( + vector_store: PostgresStore, refresh_ttl: bool +) -> None: """Test combining vector search with filters.""" # Insert test documents docs = [ @@ -487,16 +524,23 @@ def test_vector_search_with_filters(vector_store: PostgresStore) -> None: for key, value in docs: vector_store.put(("test",), key, value) - results = vector_store.search(("test",), query="apple", filter={"color": "red"}) + results = vector_store.search( + ("test",), query="apple", filter={"color": "red"}, refresh_ttl=refresh_ttl + ) assert len(results) == 2 assert results[0].key == "doc1" - results = vector_store.search(("test",), query="car", filter={"color": "red"}) + results = vector_store.search( + ("test",), query="car", filter={"color": "red"}, refresh_ttl=refresh_ttl + ) assert len(results) == 2 assert results[0].key == "doc2" results = vector_store.search( - ("test",), query="bbbbluuu", filter={"score": {"$gt": 3.2}} + ("test",), + query="bbbbluuu", + filter={"score": {"$gt": 3.2}}, + refresh_ttl=refresh_ttl, ) assert len(results) == 3 assert results[0].key == "doc4" @@ -688,7 +732,7 @@ def test_embed_with_path_operation_config( store.put(("test",), "doc5", doc5, index=False) results = store.search(("test",)) assert len(results) == 3 - assert all(r.score is None for r in results) + assert all(r.score is None for r in results), f"{results}" assert any(r.key == "doc5" for r in results) results = store.search(("test",), query="hhh") @@ -790,3 +834,27 @@ def test_nonnull_migrations() -> None: for migration in PostgresStore.MIGRATIONS: statement = _leading_comment_remover.sub("", migration).split()[0] assert statement.strip() + + +def test_store_ttl(store): + # Assumes a TTL of 1 minute = 60 seconds + ns = ("foo",) + store.put( + ns, + key="item1", + value={"foo": "bar"}, + ttl=TTL_MINUTES, # type: ignore + ) + time.sleep(TTL_SECONDS - 2) + res = store.get(ns, key="item1", refresh_ttl=True) + assert res is not None + time.sleep(TTL_SECONDS - 2) + results = store.search(ns, query="foo", refresh_ttl=True) + assert len(results) == 1 + time.sleep(TTL_SECONDS - 2) + res = store.get(ns, key="item1", refresh_ttl=False) + assert res is not None + time.sleep(TTL_SECONDS - 1) + # Now has been (TTL_SECONDS-2)*2 > TTL_SECONDS + TTL_SECONDS/2 + res = store.search(ns, query="bar", refresh_ttl=False) + assert len(res) == 0 diff --git a/libs/checkpoint/langgraph/store/base/__init__.py b/libs/checkpoint/langgraph/store/base/__init__.py index fff90cc56..de914a177 100644 --- a/libs/checkpoint/langgraph/store/base/__init__.py +++ b/libs/checkpoint/langgraph/store/base/__init__.py @@ -537,6 +537,12 @@ class TTLConfig(TypedDict, total=False): The expiration timer refreshes on both read and write operations. Defaults to None (no expiration). """ + sweep_interval_minutes: Optional[int] + """Interval in minutes between TTL sweep operations. + + If provided, the store will periodically delete expired items based on TTL. + Defaults to None (no sweeping). + """ class IndexConfig(TypedDict, total=False): diff --git a/libs/cli/langgraph_cli/config.py b/libs/cli/langgraph_cli/config.py index 6200b1ad7..011432147 100644 --- a/libs/cli/langgraph_cli/config.py +++ b/libs/cli/langgraph_cli/config.py @@ -27,6 +27,12 @@ class TTLConfig(TypedDict, total=False): If provided, all new items will have this TTL unless explicitly overridden. If omitted, items will have no TTL by default. """ + sweep_interval_minutes: Optional[int] + """Optional. Interval in minutes between TTL sweep iterations. + + If provided, the store will periodically delete expired items based on the TTL. + If omitted, no automatic sweeping will occur. + """ class IndexConfig(TypedDict, total=False): From 394a9fa85fcfe7bc9f236f490627576031bf7be0 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Fri, 14 Mar 2025 13:48:23 -0700 Subject: [PATCH 3/4] Update schema --- libs/cli/Makefile | 5 ++++- libs/cli/schemas/schema.json | 10 ++++++++++ libs/cli/schemas/schema.v0.json | 10 ++++++++++ 3 files changed, 24 insertions(+), 1 deletion(-) diff --git a/libs/cli/Makefile b/libs/cli/Makefile index 22506684d..eecf8fcdd 100644 --- a/libs/cli/Makefile +++ b/libs/cli/Makefile @@ -1,4 +1,4 @@ -.PHONY: test lint format test-integration +.PHONY: test lint format test-integration update-schema ###################### # TESTING AND COVERAGE @@ -31,3 +31,6 @@ lint lint_diff lint_package lint_tests: format format_diff: poetry run ruff format $(PYTHON_FILES) poetry run ruff check --select I --fix $(PYTHON_FILES) + +update-schema: + poetry run python generate_schema.py diff --git a/libs/cli/schemas/schema.json b/libs/cli/schemas/schema.json index 72eeec475..3cca859b2 100644 --- a/libs/cli/schemas/schema.json +++ b/libs/cli/schemas/schema.json @@ -459,6 +459,16 @@ }, "refresh_on_read": { "type": "boolean" + }, + "sweep_interval_minutes": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ] } }, "required": [] diff --git a/libs/cli/schemas/schema.v0.json b/libs/cli/schemas/schema.v0.json index 72eeec475..3cca859b2 100644 --- a/libs/cli/schemas/schema.v0.json +++ b/libs/cli/schemas/schema.v0.json @@ -459,6 +459,16 @@ }, "refresh_on_read": { "type": "boolean" + }, + "sweep_interval_minutes": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ] } }, "required": [] From 18ed044c27165910ec23071ea51900d01e3c69f0 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Fri, 14 Mar 2025 14:03:30 -0700 Subject: [PATCH 4/4] Bump patch version --- libs/checkpoint-postgres/langgraph/store/postgres/base.py | 2 +- libs/checkpoint-postgres/pyproject.toml | 2 +- libs/checkpoint/pyproject.toml | 2 +- libs/cli/pyproject.toml | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/base.py b/libs/checkpoint-postgres/langgraph/store/postgres/base.py index d88ae19ed..d6035bb70 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/base.py @@ -627,7 +627,7 @@ class BasePostgresStore(Generic[C]): class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): - """Postgres-backed store with oktional vector search using pgvector. + """Postgres-backed store with optional vector search using pgvector. !!! example "Examples" Basic setup and usage: diff --git a/libs/checkpoint-postgres/pyproject.toml b/libs/checkpoint-postgres/pyproject.toml index 307181d40..0d9ac1c12 100644 --- a/libs/checkpoint-postgres/pyproject.toml +++ b/libs/checkpoint-postgres/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "langgraph-checkpoint-postgres" -version = "2.0.16" +version = "2.0.17" description = "Library with a Postgres implementation of LangGraph checkpoint saver." authors = [] license = "MIT" diff --git a/libs/checkpoint/pyproject.toml b/libs/checkpoint/pyproject.toml index 354205b2a..a019e55d5 100644 --- a/libs/checkpoint/pyproject.toml +++ b/libs/checkpoint/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "langgraph-checkpoint" -version = "2.0.19" +version = "2.0.20" description = "Library with base interfaces for LangGraph checkpoint savers." authors = [] license = "MIT" diff --git a/libs/cli/pyproject.toml b/libs/cli/pyproject.toml index 1f7c38cb8..aaeddd4a6 100644 --- a/libs/cli/pyproject.toml +++ b/libs/cli/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "langgraph-cli" -version = "0.1.76" +version = "0.1.77" description = "CLI for interacting with LangGraph API" authors = [] license = "MIT"