diff --git a/.github/workflows/python-package.yml b/.github/workflows/python-package.yml index 6e11a4882..cb4e202c9 100644 --- a/.github/workflows/python-package.yml +++ b/.github/workflows/python-package.yml @@ -153,6 +153,17 @@ jobs: image: redis:7 ports: - 6379:6379 + aerospike: + image: aerospike/aerospike-server:8.1.2.4 + ports: + - 3000:3000 + - 3001:3001 + - 3002:3002 + options: >- + --health-cmd "asinfo -v build" + --health-interval 5s + --health-timeout 3s + --health-retries 20 strategy: fail-fast: false matrix: diff --git a/burr/integrations/persisters/b_aerospike.py b/burr/integrations/persisters/b_aerospike.py new file mode 100644 index 000000000..0eabd5f7c --- /dev/null +++ b/burr/integrations/persisters/b_aerospike.py @@ -0,0 +1,1014 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import hashlib +import json +import random +import sys +import threading +import time +from collections import OrderedDict +from datetime import datetime, timezone +from typing import Any, Optional + +from burr.core import persistence, state +from burr.integrations import base + +if sys.version_info < (3, 10): + raise ImportError( + "The Aerospike persister requires Python 3.10 or newer. " + "Install Burr with the 'aerospike' extra on a supported Python version." + ) + +try: + import aerospike + import aerospike_helpers.expressions as expr + from aerospike_helpers.operations import map_operations as map_ops + from aerospike_helpers.operations import operations as aero_ops +except ImportError as e: + base.require_plugin(e, "aerospike") + +_CODE_VERSION = 1 +_SYSTEM = "burr-as" + +# Aerospike bin names are capped at 15 characters. +_PART_BIN = "partition" +_PREFIX_BIN = "key_prefix" +_APP_BIN = "app_id" +_APPS_BIN = "app_ids" +_TAIL_BIN = "tail_page" +_MEMBERSHIP_BIN = "member_confirm" +_SEQ_BIN = "sequence_id" +_POS_BIN = "position" +_STATE_BIN = "state" +_STATUS_BIN = "status" +_CREATED_BIN = "created_at" + +_MAX_WRITE_ATTEMPTS = 3 +_MAX_TAIL_HINTS = 128 +_MEMBERSHIP_BATCH_SIZE = 1000 +_RECORD_NOT_FOUND_CODE = 2 +_RETRYABLE_CODES = {9, -10, 7, 14} # Timeout, Connection, ClusterChange, KEY_BUSY + + +class AerospikePersistenceError(Exception): + """Base class for Aerospike persister errors.""" + + +class AerospikePersistenceConsistencyError(AerospikePersistenceError): + """Raised when stored data violates the persister's invariants.""" + + +class AerospikePersistenceSerializationError(AerospikePersistenceError, ValueError): + """Raised when state cannot be serialized to strict JSON.""" + + +class AerospikePersistenceUncertainOutcomeError(AerospikePersistenceError): + """Raised when a write retry budget is exhausted without a definitive result.""" + + +class AerospikeBasePersister(persistence.BaseStatePersister): + """Synchronous Aerospike-backed implementation of Burr's ``BaseStatePersister``. + + The persister stores one immutable history record per ``(partition_key, + app_id, sequence_id)`` and a small mutable head record per ``(partition_key, + app_id)`` that points to the latest sequence. Application IDs within a + partition are listed through one materialized membership record per + logical partition. + + This optional integration requires Python 3.10+ and the official Aerospike + Python client. Burr core remains compatible with Python 3.9+. + """ + + @classmethod + def from_config(cls, config: dict) -> "AerospikeBasePersister": + """Create a persister from a configuration dictionary.""" + return cls.from_values(**config) + + @classmethod + def from_values( + cls, + hosts: Optional[list] = None, + client_config: Optional[dict] = None, + namespace: str = "test", + history_set: str = "burr_state", + head_set: str = "burr_head", + membership_set: str = "burr_apps", + key_prefix: Optional[str] = "", + serde_kwargs: Optional[dict] = None, + ) -> "AerospikeBasePersister": + """Create a persister from seed hosts and client configuration. + + :param hosts: Aerospike seed hosts as ``[(host, port), ...]``. + :param client_config: Additional official client configuration options. + :param namespace: Aerospike namespace. + :param history_set: Set for immutable checkpoint records. + :param head_set: Set for mutable latest-sequence heads. + :param membership_set: Set for per-partition application membership. + :param key_prefix: Optional logical prefix for key isolation. + :param serde_kwargs: Kwargs for Burr ``State`` serialization. + """ + if hosts is None: + hosts = [("127.0.0.1", 3000)] + aerospike_config = {"hosts": hosts} + if client_config: + aerospike_config.update(client_config) + + client = aerospike.client(aerospike_config) + + return cls( + client, + namespace=namespace, + history_set=history_set, + head_set=head_set, + membership_set=membership_set, + key_prefix=key_prefix if key_prefix is not None else "", + serde_kwargs=serde_kwargs, + _client_config=aerospike_config, + _owned=True, + ) + + def __init__( + self, + client, + *, + namespace: str = "test", + history_set: str = "burr_state", + head_set: str = "burr_head", + membership_set: str = "burr_apps", + key_prefix: Optional[str] = "", + serde_kwargs: Optional[dict] = None, + _client_config: Optional[dict] = None, + _owned: bool = False, + ): + """Initialize the persister with an Aerospike client. + + Direct construction accepts an already-connected caller-owned client. + Use :meth:`from_values` or :meth:`from_config` for persister-owned + clients. + """ + self._client = client + self._owned = _owned + self.namespace = namespace + self.history_set = history_set + self.head_set = head_set + self.membership_set = membership_set + self.key_prefix = key_prefix if key_prefix is not None else "" + self.serde_kwargs = serde_kwargs or {} + self._client_config = _client_config + self._initialized = False + self._closed = False + self._tail_hints = OrderedDict() + self._tail_hints_lock = threading.Lock() + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.cleanup() + return False + + def is_async(self) -> bool: + return False + + def cleanup(self): + """Close the persister-owned client, if any.""" + if self._owned and not self._closed and self._client is not None: + self._client.close() + self._closed = True + + def __getstate__(self) -> dict: + if not self._owned: + raise TypeError( + "An AerospikeBasePersister constructed with an injected client cannot be pickled" + ) + if self._client_config is None: + raise TypeError( + "Cannot pickle an AerospikeBasePersister without reconnectable client configuration" + ) + state = self.__dict__.copy() + del state["_client"] + del state["_tail_hints"] + del state["_tail_hints_lock"] + state["_initialized"] = False + state["_closed"] = False + return state + + def __setstate__(self, state: dict): + client_config = state.get("_client_config") + if client_config is None: + raise TypeError( + "Cannot unpickle an AerospikeBasePersister without client configuration" + ) + self.__dict__.update(state) + try: + self._client = aerospike.client(client_config) + except Exception as e: + raise AerospikePersistenceError( + f"Failed to reconnect Aerospike client from configuration: {e}" + ) from e + self._owned = True + self._initialized = False + self._closed = False + self._tail_hints = OrderedDict() + self._tail_hints_lock = threading.Lock() + + def set_serde_kwargs(self, serde_kwargs: dict): + self.serde_kwargs = serde_kwargs + + def initialize(self): + """Mark the dynamically provisioned persister ready for use.""" + self._initialized = True + + def is_initialized(self) -> bool: + return self._initialized + + def save( + self, + partition_key: Optional[str], + app_id: str, + sequence_id: int, + position: str, + state: state.State, + status: str, + **kwargs, + ): + self._validate_sequence_id(sequence_id) + + try: + state_json = json.dumps(state.serialize(**self.serde_kwargs), allow_nan=False) + except (TypeError, ValueError) as e: + raise AerospikePersistenceSerializationError( + f"State is not JSON-serializable: {e}" + ) from e + created_at = datetime.now(timezone.utc).isoformat(timespec="microseconds") + + key_prefix = self.key_prefix + partition_canonical = self._canonical_partition(partition_key) + key_prefix_canonical = self._canonical_key_prefix() + membership_key = self._key( + self.membership_set, + self._derive_membership_key(key_prefix, partition_key), + ) + + history_key = self._key( + self.history_set, + self._derive_history_key(key_prefix, partition_key, app_id, sequence_id), + ) + head_key = self._key( + self.head_set, + self._derive_head_key(key_prefix, partition_key, app_id), + ) + + history_bins = { + _PART_BIN: partition_canonical, + _PREFIX_BIN: key_prefix_canonical, + _APP_BIN: app_id, + _SEQ_BIN: sequence_id, + _POS_BIN: position, + _STATE_BIN: state_json, + _STATUS_BIN: status, + _CREATED_BIN: created_at, + } + + # History must exist before the head can be advanced. + self._write_history(history_key, history_bins) + + head_bins = { + _PART_BIN: partition_canonical, + _PREFIX_BIN: key_prefix_canonical, + _APP_BIN: app_id, + _SEQ_BIN: sequence_id, + } + _, membership_registered = self._advance_head(head_key, head_bins) + if membership_registered is not True: + self._register_membership(membership_key, partition_key, app_id) + self._confirm_membership( + head_key, + partition_canonical, + key_prefix_canonical, + app_id, + ) + + def load( + self, + partition_key: Optional[str], + app_id: Optional[str], + sequence_id: Optional[int] = None, + **kwargs, + ) -> Optional[persistence.PersistedStateData]: + if app_id is None: + raise ValueError("app_id is required for load") + if sequence_id is not None: + self._validate_sequence_id(sequence_id) + + if sequence_id is None: + return self._load_latest(partition_key, app_id) + + # Exact-sequence load. + key_prefix = self.key_prefix + partition_canonical = self._canonical_partition(partition_key) + key_prefix_canonical = self._canonical_key_prefix() + history_key = self._key( + self.history_set, + self._derive_history_key(key_prefix, partition_key, app_id, sequence_id), + ) + history_record = self._read_history(history_key) + if history_record is None: + return None + self._validate_history_identity( + history_record, partition_canonical, key_prefix_canonical, app_id, sequence_id + ) + return self._build_persisted_data(history_record, partition_key, app_id, sequence_id) + + def _load_latest( + self, + partition_key: Optional[str], + app_id: str, + ) -> Optional[persistence.PersistedStateData]: + """Load the latest state for an application by reading the head, then the referenced history.""" + key_prefix = self.key_prefix + partition_canonical = self._canonical_partition(partition_key) + key_prefix_canonical = self._canonical_key_prefix() + head_key = self._key( + self.head_set, + self._derive_head_key(key_prefix, partition_key, app_id), + ) + head_record = self._read_head(head_key) + if head_record is None: + return None + stored_seq, _ = self._validate_head_identity( + head_record, + partition_canonical, + key_prefix_canonical, + app_id, + ) + history_key = self._key( + self.history_set, + self._derive_history_key(key_prefix, partition_key, app_id, stored_seq), + ) + history_record = self._read_history(history_key) + if history_record is None: + raise AerospikePersistenceConsistencyError( + f"Head for app_id={app_id} references missing history sequence {stored_seq}" + ) + self._validate_history_identity( + history_record, partition_canonical, key_prefix_canonical, app_id, stored_seq + ) + return self._build_persisted_data(history_record, partition_key, app_id, stored_seq) + + def list_app_ids(self, partition_key: Optional[str], **kwargs) -> list[str]: + membership_key = self._membership_page_key(partition_key, 0) + root = self._read_membership_root(membership_key) + if root is None: + return [] + app_ids, tail_page = root + if tail_page == 0: + return list(app_ids) + combined = set(app_ids) + for first_page in range(1, tail_page + 1, _MEMBERSHIP_BATCH_SIZE): + keys = [ + self._membership_page_key(partition_key, page) + for page in range( + first_page, + min(first_page + _MEMBERSHIP_BATCH_SIZE, tail_page + 1), + ) + ] + try: + results = self._client.batch_read( + keys, + bins=[_APPS_BIN], + policy={"replica": aerospike.POLICY_REPLICA_MASTER}, + ) + except aerospike.exception.AerospikeError as e: + raise AerospikePersistenceError(f"Membership batch read failed: {e}") from e + membership_pages = results.batch_records + if len(membership_pages) != len(keys): + raise AerospikePersistenceConsistencyError("Membership batch read was incomplete") + for membership_page in membership_pages: + if membership_page.result == _RECORD_NOT_FOUND_CODE: + continue + if membership_page.result != 0: + raise AerospikePersistenceError( + f"Membership batch entry failed with result code {membership_page.result}" + ) + if membership_page.record is None: + raise AerospikePersistenceConsistencyError( + "Successful membership batch entry lacks a record" + ) + combined.update(self._validate_membership_map(membership_page.record[2])) + return list(combined) + + def _canonical_json(self, value: Any) -> str: + """Deterministic JSON encoding used for canonical identities and state snapshots.""" + return json.dumps( + value, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + allow_nan=False, + ) + + def _identity_history( + self, key_prefix: Optional[str], partition_key: Optional[str], app_id: str, sequence_id: int + ) -> list: + return [ + _SYSTEM, + _CODE_VERSION, + "history", + key_prefix, + partition_key, + app_id, + sequence_id, + ] + + def _identity_head( + self, key_prefix: Optional[str], partition_key: Optional[str], app_id: str + ) -> list: + return [ + _SYSTEM, + _CODE_VERSION, + "head", + key_prefix, + partition_key, + app_id, + ] + + def _identity_membership(self, key_prefix: Optional[str], partition_key: Optional[str]) -> list: + return [ + _SYSTEM, + _CODE_VERSION, + "membership", + key_prefix, + partition_key, + ] + + def _identity_membership_page( + self, key_prefix: Optional[str], partition_key: Optional[str], page: int + ) -> list: + return [ + _SYSTEM, + _CODE_VERSION, + "membership-page", + key_prefix, + partition_key, + page, + ] + + def _sha256_hex(self, data: str) -> str: + return hashlib.sha256(data.encode("utf-8")).hexdigest() + + def _derive_history_key( + self, key_prefix: Optional[str], partition_key: Optional[str], app_id: str, sequence_id: int + ) -> str: + return self._sha256_hex( + self._canonical_json( + self._identity_history(key_prefix, partition_key, app_id, sequence_id) + ) + ) + + def _derive_head_key( + self, key_prefix: Optional[str], partition_key: Optional[str], app_id: str + ) -> str: + return self._sha256_hex( + self._canonical_json(self._identity_head(key_prefix, partition_key, app_id)) + ) + + def _derive_membership_key( + self, key_prefix: Optional[str], partition_key: Optional[str] + ) -> str: + return self._sha256_hex( + self._canonical_json(self._identity_membership(key_prefix, partition_key)) + ) + + def _derive_membership_page_key( + self, key_prefix: Optional[str], partition_key: Optional[str], page: int + ) -> str: + if type(page) is not int or isinstance(page, bool) or page <= 0: + raise ValueError("membership overflow page must be a positive integer") + return self._sha256_hex( + self._canonical_json(self._identity_membership_page(key_prefix, partition_key, page)) + ) + + def _membership_page_key(self, partition_key: Optional[str], page: int): + user_key = ( + self._derive_membership_key(self.key_prefix, partition_key) + if page == 0 + else self._derive_membership_page_key(self.key_prefix, partition_key, page) + ) + return self._key(self.membership_set, user_key) + + def _validate_sequence_id(self, sequence_id: Any) -> None: + if type(sequence_id) is not int or isinstance(sequence_id, bool): + raise ValueError( + f"sequence_id must be a 64-bit signed integer, got {sequence_id!r} (type {type(sequence_id).__name__})" + ) + if sequence_id < -(2**63) or sequence_id > 2**63 - 1: + raise ValueError( + f"sequence_id {sequence_id} is outside the signed 64-bit integer range" + ) + + def _key(self, set_name: str, user_key: str): + return (self.namespace, set_name, user_key) + + def _canonical_partition(self, partition_key: Optional[str]) -> str: + return self._canonical_json(partition_key) + + def _canonical_key_prefix(self) -> str: + return self._canonical_json(self.key_prefix) + + def _read_record(self, key, policy=None): + read_policy = policy or {"replica": aerospike.POLICY_REPLICA_MASTER} + try: + return self._client.get(key, policy=read_policy) + except aerospike.exception.RecordNotFound: + return None + + def _is_retryable(self, exc: aerospike.exception.AerospikeError) -> bool: + return getattr(exc, "code", None) in _RETRYABLE_CODES + + def _backoff(self, attempt: int) -> None: + # Bounded exponential backoff with jitter (seconds). + delay = min(0.05 * (2**attempt), 1.0) + time.sleep(delay * (0.5 + random.random())) + + def _read_history(self, key): + record = self._read_record(key) + if record is None: + return None + _, _, bins = record + return bins + + def _read_head(self, key): + return self._read_history(key) + + def _validate_membership_map(self, bins: dict) -> dict: + app_ids = bins.get(_APPS_BIN) + if not isinstance(app_ids, dict): + raise AerospikePersistenceConsistencyError( + "Membership page lacks a usable application-ID map" + ) + return app_ids + + def _validate_tail_page(self, value: Any) -> int: + if type(value) is not int or isinstance(value, bool) or value < 0 or value > 2**63 - 1: + raise AerospikePersistenceConsistencyError( + "Membership root tail_page must be a non-negative signed integer" + ) + return value + + def _read_membership_page(self, key) -> Optional[dict]: + try: + _, _, bins = self._client.select( + key, + [_APPS_BIN], + policy={"replica": aerospike.POLICY_REPLICA_MASTER}, + ) + return self._validate_membership_map(bins) + except aerospike.exception.RecordNotFound: + return None + except aerospike.exception.AerospikeError as e: + raise AerospikePersistenceError(f"Membership read failed: {e}") from e + + def _read_membership_root(self, key) -> Optional[tuple[dict, int]]: + try: + _, _, bins = self._client.select( + key, + [_APPS_BIN, _TAIL_BIN], + policy={"replica": aerospike.POLICY_REPLICA_MASTER}, + ) + return self._validate_membership_map(bins), self._validate_tail_page( + bins.get(_TAIL_BIN) + ) + except aerospike.exception.RecordNotFound: + return None + except aerospike.exception.AerospikeError as e: + raise AerospikePersistenceError(f"Membership read failed: {e}") from e + + def _advance_membership_tail(self, root_key, expected: int) -> int: + policy = { + "commit_level": aerospike.POLICY_COMMIT_LEVEL_ALL, + "expressions": expr.Eq(expr.IntBin(_TAIL_BIN), expected).compile(), + "ttl": aerospike.TTL_NEVER_EXPIRE, + "max_retries": 0, + } + operations = [aero_ops.write(_TAIL_BIN, expected + 1), aero_ops.read(_TAIL_BIN)] + for attempt in range(1, _MAX_WRITE_ATTEMPTS + 1): + try: + _, _, bins = self._client.operate(root_key, operations, policy=policy) + return self._validate_tail_page(bins[_TAIL_BIN]) + except (aerospike.exception.FilteredOut, aerospike.exception.AerospikeError) as e: + if not isinstance(e, aerospike.exception.FilteredOut) and not self._is_retryable(e): + raise AerospikePersistenceError( + f"Membership tail allocation failed: {e}" + ) from e + root = self._read_membership_root(root_key) + if root is not None and root[1] > expected: + return root[1] + if attempt == _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "Membership tail allocation failed with an uncertain outcome" + ) from e + self._backoff(attempt) + + def _put_membership_page(self, key, app_id: str, root: bool = False) -> None: + operations = [map_ops.map_put(_APPS_BIN, app_id, 1)] + policy = { + "commit_level": aerospike.POLICY_COMMIT_LEVEL_ALL, + "ttl": aerospike.TTL_NEVER_EXPIRE, + "max_retries": 0, + } + if root: + operations.insert(0, aero_ops.write(_TAIL_BIN, 0)) + policy["expressions"] = expr.Or( + expr.Not(expr.BinExists(_TAIL_BIN)), expr.Eq(expr.IntBin(_TAIL_BIN), 0) + ).compile() + self._client.operate(key, operations, policy=policy) + + def _get_tail_hint(self, membership_key) -> Optional[int]: + with self._tail_hints_lock: + tail_page = self._tail_hints.get(membership_key) + if tail_page is not None: + self._tail_hints.move_to_end(membership_key) + return tail_page + + def _set_tail_hint(self, membership_key, tail_page: int) -> None: + if tail_page <= 0: + return + with self._tail_hints_lock: + existing = self._tail_hints.get(membership_key, 0) + self._tail_hints[membership_key] = max(existing, tail_page) + self._tail_hints.move_to_end(membership_key) + while len(self._tail_hints) > _MAX_TAIL_HINTS: + self._tail_hints.popitem(last=False) + + def _register_membership( + self, membership_key, partition_key: Optional[str], app_id: str + ) -> None: + hinted_page = self._get_tail_hint(membership_key) + current_page = hinted_page if hinted_page is not None else 0 + current_key = self._membership_page_key(partition_key, current_page) + uncertain_attempt = 1 + rollovers = 0 + while True: + try: + self._put_membership_page(current_key, app_id, root=current_page == 0) + self._set_tail_hint(membership_key, current_page) + return + except aerospike.exception.FilteredOut: + root = self._read_membership_root(membership_key) + if root is None: + raise AerospikePersistenceConsistencyError( + "Membership root disappeared after filtering" + ) + current_page = root[1] + self._set_tail_hint(membership_key, current_page) + current_key = self._membership_page_key(partition_key, current_page) + uncertain_attempt = 1 + except aerospike.exception.RecordTooBig as e: + root = self._read_membership_root(membership_key) + if root is None: + raise AerospikePersistenceConsistencyError( + "Membership root disappeared during rollover" + ) + authoritative_tail = root[1] + self._set_tail_hint(membership_key, authoritative_tail) + if current_page < authoritative_tail: + current_page = authoritative_tail + current_key = self._membership_page_key(partition_key, current_page) + uncertain_attempt = 1 + continue + rollovers += 1 + if rollovers > _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "Membership rollover retry budget exhausted" + ) from e + current_page = self._advance_membership_tail(membership_key, authoritative_tail) + self._set_tail_hint(membership_key, current_page) + current_key = self._membership_page_key(partition_key, current_page) + uncertain_attempt = 1 + except aerospike.exception.AerospikeError as e: + if not self._is_retryable(e): + raise AerospikePersistenceError(f"Membership registration failed: {e}") from e + app_ids = self._read_membership_page(current_key) + if app_ids is not None and app_id in app_ids: + self._set_tail_hint(membership_key, current_page) + return + if uncertain_attempt == _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "Membership registration failed with an uncertain outcome" + ) from e + self._backoff(uncertain_attempt) + uncertain_attempt += 1 + + def _write_history_once(self, history_key, bins): + """Attempt one create-only history write.""" + try: + self._client.put( + history_key, + bins, + policy={ + "commit_level": aerospike.POLICY_COMMIT_LEVEL_ALL, + "exists": aerospike.POLICY_EXISTS_CREATE, + "ttl": aerospike.TTL_NEVER_EXPIRE, + }, + ) + except aerospike.exception.RecordTooBig as e: + raise AerospikePersistenceSerializationError( + "Checkpoint record exceeds the namespace max-record-size" + ) from e + + def _write_history(self, history_key, bins): + """Create-only history write with bounded retries.""" + for attempt in range(1, _MAX_WRITE_ATTEMPTS + 1): + try: + self._write_history_once(history_key, bins) + return None + except aerospike.exception.RecordExistsError: + # The checkpoint already exists; history is immutable, so this is idempotent. + return None + except aerospike.exception.AerospikeError as e: + if self._is_retryable(e): + # Reconcile: if the record is now present, the write succeeded. + existing = self._read_history(history_key) + if existing is not None: + return existing + if attempt == _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "History write timed out or failed with an uncertain outcome" + ) from e + self._backoff(attempt) + continue + raise AerospikePersistenceError(f"History write failed: {e}") from e + + raise AerospikePersistenceUncertainOutcomeError("History write retry budget exhausted") + + def _head_filter_expression( + self, + partition_canonical: str, + key_prefix_canonical: str, + app_id: str, + sequence_id: int, + ): + """Build a conditional expression for the head operate(). + + The operation is allowed when the record does not exist, or when the + existing record has the same identity and a lower stored sequence. + """ + return expr.Or( + expr.Not(expr.BinExists(_SEQ_BIN)), + expr.And( + expr.Eq(expr.StrBin(_PART_BIN), partition_canonical), + expr.Eq(expr.StrBin(_PREFIX_BIN), key_prefix_canonical), + expr.Eq(expr.StrBin(_APP_BIN), app_id), + expr.LT(expr.IntBin(_SEQ_BIN), sequence_id), + ), + ).compile() + + def _head_ops(self, bins: dict): + """Return operate() ops that write all head bins and read the sequence.""" + return [ + aero_ops.write(_PART_BIN, bins[_PART_BIN]), + aero_ops.write(_PREFIX_BIN, bins[_PREFIX_BIN]), + aero_ops.write(_APP_BIN, bins[_APP_BIN]), + aero_ops.write(_SEQ_BIN, bins[_SEQ_BIN]), + aero_ops.read(_SEQ_BIN), + aero_ops.read(_MEMBERSHIP_BIN), + ] + + def _validate_head_identity( + self, + bins: dict, + partition_canonical: str, + key_prefix_canonical: str, + app_id: str, + ) -> tuple[int, Optional[bool]]: + """Validate head identity and return its sequence and optional membership confirmation.""" + if bins.get(_PART_BIN) != partition_canonical: + raise AerospikePersistenceConsistencyError( + "Head record partition identity does not match the requested identity" + ) + if bins.get(_PREFIX_BIN) != key_prefix_canonical: + raise AerospikePersistenceConsistencyError( + "Head record key_prefix does not match the requested key_prefix" + ) + if bins.get(_APP_BIN) != app_id: + raise AerospikePersistenceConsistencyError( + "Head record app_id does not match the requested app_id" + ) + stored_seq = bins[_SEQ_BIN] + self._validate_sequence_id(stored_seq) + membership_registered = bins.get(_MEMBERSHIP_BIN) + if membership_registered is not None and membership_registered is not True: + raise AerospikePersistenceConsistencyError( + "Head membership confirmation must be true when present" + ) + return stored_seq, membership_registered + + def _advance_head( + self, + head_key, + head_bins: dict, + ): + """Atomically create or advance the application head. + + Uses a conditional operate() with server-side filtering. A filtered or + ambiguous result is classified through a correctness-sensitive master + read before any retry or final decision. + """ + sequence_id = head_bins[_SEQ_BIN] + partition_canonical = head_bins[_PART_BIN] + key_prefix_canonical = head_bins[_PREFIX_BIN] + app_id = head_bins[_APP_BIN] + + ops = self._head_ops(head_bins) + filter_expr = self._head_filter_expression( + partition_canonical, key_prefix_canonical, app_id, sequence_id + ) + + for attempt in range(1, _MAX_WRITE_ATTEMPTS + 1): + try: + _, _, result_bins = self._client.operate( + head_key, + ops, + policy={ + "commit_level": aerospike.POLICY_COMMIT_LEVEL_ALL, + "expressions": filter_expr, + "ttl": aerospike.TTL_NEVER_EXPIRE, + "max_retries": 0, + }, + ) + except aerospike.exception.FilteredOut: + # The record exists and the condition was false. Read and classify. + existing = self._read_head(head_key) + if existing is None: + if attempt == _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "Head update was filtered out but the record could not be reconciled" + ) + self._backoff(attempt) + continue + stored_seq, membership_registered = self._validate_head_identity( + existing, partition_canonical, key_prefix_canonical, app_id + ) + if stored_seq >= sequence_id: + return stored_seq, membership_registered + # Stored sequence is lower than requested but filter was false: + # a concurrent writer may have changed the record, retry. + if attempt == _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "Head update could not be reconciled within the retry budget" + ) + self._backoff(attempt) + continue + except aerospike.exception.AerospikeError as e: + if self._is_retryable(e): + existing = self._read_head(head_key) + if existing is not None: + stored_seq, membership_registered = self._validate_head_identity( + existing, + partition_canonical, + key_prefix_canonical, + app_id, + ) + if stored_seq >= sequence_id: + return stored_seq, membership_registered + if attempt == _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "Head update failed with an uncertain outcome" + ) from e + self._backoff(attempt) + continue + raise AerospikePersistenceError(f"Head update failed: {e}") from e + + stored_seq = result_bins[_SEQ_BIN] + self._validate_sequence_id(stored_seq) + if stored_seq < sequence_id: + raise AerospikePersistenceConsistencyError( + "Head operate returned a sequence lower than the requested sequence" + ) + membership_registered = result_bins.get(_MEMBERSHIP_BIN) + if membership_registered is not None and membership_registered is not True: + raise AerospikePersistenceConsistencyError( + "Head membership confirmation must be true when present" + ) + return stored_seq, membership_registered + + raise AerospikePersistenceUncertainOutcomeError("Head update retry budget exhausted") + + def _confirm_membership( + self, + head_key, + partition_canonical: str, + key_prefix_canonical: str, + app_id: str, + ) -> None: + """Monotonically mark an application head's membership as confirmed.""" + filter_expr = expr.And( + expr.Eq(expr.StrBin(_PART_BIN), partition_canonical), + expr.Eq(expr.StrBin(_PREFIX_BIN), key_prefix_canonical), + expr.Eq(expr.StrBin(_APP_BIN), app_id), + expr.Not(expr.BinExists(_MEMBERSHIP_BIN)), + ).compile() + operations = [ + aero_ops.write(_MEMBERSHIP_BIN, True), + aero_ops.read(_SEQ_BIN), + aero_ops.read(_MEMBERSHIP_BIN), + ] + policy = { + "commit_level": aerospike.POLICY_COMMIT_LEVEL_ALL, + "expressions": filter_expr, + "ttl": aerospike.TTL_NEVER_EXPIRE, + "max_retries": 0, + } + + for attempt in range(1, _MAX_WRITE_ATTEMPTS + 1): + try: + _, _, result_bins = self._client.operate(head_key, operations, policy=policy) + self._validate_sequence_id(result_bins[_SEQ_BIN]) + if result_bins.get(_MEMBERSHIP_BIN) is not True: + raise AerospikePersistenceConsistencyError( + "Head confirmation operation did not return true" + ) + return + except ( + aerospike.exception.FilteredOut, + aerospike.exception.AerospikeError, + ) as e: + if not isinstance(e, aerospike.exception.FilteredOut) and not self._is_retryable(e): + raise AerospikePersistenceError( + f"Head membership confirmation failed: {e}" + ) from e + existing = self._read_head(head_key) + if existing is not None: + _, membership_registered = self._validate_head_identity( + existing, partition_canonical, key_prefix_canonical, app_id + ) + if membership_registered is True: + return + if attempt == _MAX_WRITE_ATTEMPTS: + raise AerospikePersistenceUncertainOutcomeError( + "Head membership confirmation failed with an uncertain outcome" + ) from e + self._backoff(attempt) + + def _validate_history_identity( + self, + bins: dict, + partition_canonical: str, + key_prefix_canonical: str, + app_id: str, + sequence_id: int, + ): + stored_seq = bins[_SEQ_BIN] + self._validate_sequence_id(stored_seq) + + if bins[_PREFIX_BIN] != key_prefix_canonical: + raise AerospikePersistenceConsistencyError( + "History record key_prefix does not match the requested key_prefix" + ) + if bins[_PART_BIN] != partition_canonical: + raise AerospikePersistenceConsistencyError( + "History record partition identity does not match the requested identity" + ) + if bins[_APP_BIN] != app_id: + raise AerospikePersistenceConsistencyError( + "History record app_id does not match the requested app_id" + ) + if bins[_SEQ_BIN] != sequence_id: + raise AerospikePersistenceConsistencyError( + "History record sequence_id does not match the requested sequence_id" + ) + + def _build_persisted_data( + self, + history_bins: dict, + partition_key: Optional[str], + app_id: str, + sequence_id: int, + ) -> persistence.PersistedStateData: + return { + "partition_key": partition_key, + "app_id": app_id, + "sequence_id": sequence_id, + "position": history_bins[_POS_BIN], + "state": state.State.deserialize( + json.loads(history_bins[_STATE_BIN]), **self.serde_kwargs + ), + "created_at": history_bins[_CREATED_BIN], + "status": history_bins[_STATUS_BIN], + } diff --git a/docs/getting_started/install.rst b/docs/getting_started/install.rst index 956ce13fb..57f65bc8b 100644 --- a/docs/getting_started/install.rst +++ b/docs/getting_started/install.rst @@ -133,6 +133,12 @@ This installs the dependencies for Pydantic. This installs the dependencies for Redis. +.. code-block:: bash + + pip install "apache-burr[aerospike]" + +This installs the dependencies for Aerospike. It requires Python 3.10 or newer; Burr core remains compatible with Python 3.9. + .. code-block:: bash pip install "apache-burr[start]" diff --git a/docs/reference/persister.rst b/docs/reference/persister.rst index 9564bc81e..9f55008c4 100644 --- a/docs/reference/persister.rst +++ b/docs/reference/persister.rst @@ -42,6 +42,8 @@ We currently support the following database integrations: +-------------+-----------+-----------------------------------------------------+---------------+-----------------------------------------------------+ | MongoDB | pymongo | :ref:`MongoDBBasePersister ` | ❌ | ❌ | +-------------+-----------+-----------------------------------------------------+---------------+-----------------------------------------------------+ + | Aerospike | aerospike | :ref:`AerospikeBasePersister ` | ❌ | ❌ | + +-------------+-----------+-----------------------------------------------------+---------------+-----------------------------------------------------+ We follow the naming convention ``b_dependency-library``, where the ``b_`` is used to avoid name clashing with the underlying library. We chose the library name in case we implement the same database @@ -129,6 +131,14 @@ Currently we support the following, although we highly recommend you contribute .. automethod:: __init__ +.. _syncaerospikeref: + +.. autoclass:: burr.integrations.persisters.b_aerospike.AerospikeBasePersister + :members: + + .. automethod:: __init__ + + Note that the :py:class:`LocalTrackingClient ` leverages the :py:class:`BaseStateLoader ` to allow loading state, although it uses different mechanisms to save state (as it tracks more than just state). diff --git a/examples/integrations/aerospike/README.md b/examples/integrations/aerospike/README.md new file mode 100644 index 000000000..d761a0fca --- /dev/null +++ b/examples/integrations/aerospike/README.md @@ -0,0 +1,118 @@ + + +# Aerospike + Burr + +This example extends the Burr simple chatbot with durable state in Aerospike. The persister saves one immutable checkpoint after every action and reloads the latest checkpoint when an application with the same identifiers starts again. + +## Install dependencies + +Python 3.10 or newer is required for the Aerospike integration. + +```bash +pip install "apache-burr[aerospike]" openai +``` + +When running from this repository, install the checkout into its virtual environment instead: + +```bash +uv pip install --python .venv/bin/python -e ".[aerospike]" openai +``` + +## Start Aerospike + +Start a local Aerospike Community Edition database with Docker: + +```bash +docker run --name burr-aerospike -d \ + -p 3000:3000 \ + aerospike/aerospike-server:8.1.2.4 +``` + +The example uses the default `test` namespace. Aerospike creates the `burr_state` and `burr_head` sets on the first write. + +## Run the chatbot + +Set `OPENAI_API_KEY`, then run: + +```bash +export OPENAI_API_KEY="your-key" +python aerospike_local.py +``` + +From the repository root with the existing virtual environment: + +```bash +export OPENAI_API_KEY="your-key" +.venv/bin/python examples/integrations/aerospike/aerospike_local.py +``` + +Enter `exit` or `quit` to stop. Restarting with the same `app_id` and `partition_key` restores the conversation from Aerospike. + +Use command-line options to select a conversation or send one prompt: + +```bash +python aerospike_local.py \ + --app-id another-conversation \ + --partition-key another-user \ + --prompt "What did we discuss last time?" +``` + +The example keeps at most 20 chat messages in the current state. The Aerospike persister still retains an immutable checkpoint for every completed Burr action. + +## Inspect persisted checkpoints + +Run AQL from the Aerospike tools image on the server container's network: + +```bash +docker run --rm -it \ + --network container:burr-aerospike \ + aerospike/aerospike-tools \ + aql -h 127.0.0.1 +``` + +At the AQL prompt, list the sets and inspect the saved checkpoints: + +```sql +SHOW SETS +SELECT * FROM test.burr_state +``` + +Each `burr_state` record is an immutable checkpoint. The relevant bins include: + +- `app_id`: conversation identifier, `aerospike-chatbot` by default +- `sequence_id`: action sequence number +- `position`: action that produced the checkpoint +- `state`: serialized Burr state, including the chat history +- `status`: action completion status + +The mutable record that points to the latest checkpoint for each conversation is stored separately: + +```sql +SELECT * FROM test.burr_head +``` + +You can also execute a query without opening an interactive AQL prompt: + +```bash +docker run --rm \ + --network container:burr-aerospike \ + aerospike/aerospike-tools \ + aql -h 127.0.0.1 -c "SELECT * FROM test.burr_state" +``` diff --git a/examples/integrations/aerospike/aerospike_local.py b/examples/integrations/aerospike/aerospike_local.py new file mode 100644 index 000000000..68dc32858 --- /dev/null +++ b/examples/integrations/aerospike/aerospike_local.py @@ -0,0 +1,100 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import argparse + +import openai + +from burr.core import ApplicationBuilder, State, action +from burr.integrations.persisters.b_aerospike import AerospikeBasePersister + +MODEL = "gpt-4o-mini" +MAX_HISTORY_ITEMS = 20 +client = openai.Client() + + +@action(reads=["chat_history"], writes=["prompt", "chat_history"]) +def human_input(state: State, prompt: str) -> State: + chat_item = {"content": prompt, "role": "user"} + chat_history = [*state["chat_history"], chat_item][-MAX_HISTORY_ITEMS:] + return state.update(prompt=prompt, chat_history=chat_history) + + +@action(reads=["chat_history"], writes=["response", "chat_history"]) +def ai_response(state: State) -> State: + content = ( + client.chat.completions.create( + model=MODEL, + messages=state["chat_history"], + ) + .choices[0] + .message.content + ) + chat_item = {"content": content, "role": "assistant"} + chat_history = [*state["chat_history"], chat_item][-MAX_HISTORY_ITEMS:] + return state.update(response=content, chat_history=chat_history) + + +def build_application(persister: AerospikeBasePersister, app_id: str, partition_key: str): + return ( + ApplicationBuilder() + .with_actions(human_input, ai_response) + .with_transitions( + ("human_input", "ai_response"), + ("ai_response", "human_input"), + ) + .initialize_from( + persister, + resume_at_next_action=True, + default_state={"chat_history": []}, + default_entrypoint="human_input", + ) + .with_state_persister(persister) + .with_identifiers(app_id=app_id, partition_key=partition_key) + .build() + ) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--app-id", default="aerospike-chatbot") + parser.add_argument("--partition-key", default="example-user") + parser.add_argument("--prompt") + args = parser.parse_args() + + with AerospikeBasePersister.from_values( + hosts=[("127.0.0.1", 3000)], + namespace="test", + key_prefix="aerospike-example", + ) as persister: + persister.initialize() + app = build_application(persister, args.app_id, args.partition_key) + if args.prompt: + *_, state = app.run(halt_after=["ai_response"], inputs={"prompt": args.prompt}) + print(state["response"]) + return + + while True: + prompt = input("you: ").strip() + if prompt.lower() in {"exit", "quit"}: + break + *_, state = app.run(halt_after=["ai_response"], inputs={"prompt": prompt}) + print("assistant:", state["response"]) + + +if __name__ == "__main__": + main() diff --git a/pyproject.toml b/pyproject.toml index 6c3d597f4..01829d529 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -83,6 +83,10 @@ redis = [ "redis" ] +aerospike = [ + "aerospike>=19.0.0; python_version >= \"3.10\"" +] + release = [ "jinja2", ] @@ -101,6 +105,7 @@ tests = [ "apache-burr[psycopg2]", "apache-burr[pymongo]", "apache-burr[redis]", + "apache-burr[aerospike]", "apache-burr[opentelemetry]", "apache-burr[langfuse]", "apache-burr[haystack]", @@ -121,6 +126,7 @@ documentation = [ "apache-burr[psycopg2]", "apache-burr[pymongo]", "apache-burr[redis]", + "apache-burr[aerospike]", "apache-burr[ray]", "apache-burr[streamlit]", "sphinxcontrib-googleanalytics" diff --git a/tests/integrations/persisters/test_b_aerospike.py b/tests/integrations/persisters/test_b_aerospike.py new file mode 100644 index 000000000..727a9b192 --- /dev/null +++ b/tests/integrations/persisters/test_b_aerospike.py @@ -0,0 +1,699 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import concurrent.futures +import os +import pickle +import sys +import uuid + +import pytest + +if os.environ.get("BURR_CI_INTEGRATION_TESTS") != "true" or sys.version_info < (3, 10): + pytest.skip("Skipping Aerospike integration tests", allow_module_level=True) + +from unittest.mock import Mock, patch + +import aerospike + +from burr.core import state +from burr.core.persistence import BaseStatePersister +from burr.integrations.persisters.b_aerospike import ( + AerospikeBasePersister, + AerospikePersistenceConsistencyError, + AerospikePersistenceUncertainOutcomeError, +) + + +@pytest.fixture +def aerospike_persister(): + persister = AerospikeBasePersister.from_values(key_prefix=f"test-{uuid.uuid4().hex}") + persister.initialize() + yield persister + persister.cleanup() + + +def test_persister_exposes_the_synchronous_burr_interface(aerospike_persister): + assert isinstance(aerospike_persister, BaseStatePersister) + assert aerospike_persister.is_async() is False + assert aerospike_persister.is_initialized() is True + + +def test_exact_checkpoint_round_trip_returns_all_persisted_fields(aerospike_persister): + aerospike_persister.save( + "tenant-a", + "application-a", + 7, + "answer", + state.State({"answer": 42}), + "completed", + ) + + loaded = aerospike_persister.load("tenant-a", "application-a", 7) + + assert loaded["partition_key"] == "tenant-a" + assert loaded["app_id"] == "application-a" + assert loaded["sequence_id"] == 7 + assert loaded["position"] == "answer" + assert loaded["state"].get_all() == {"answer": 42} + assert loaded["status"] == "completed" + assert loaded["created_at"].endswith("+00:00") + + +def test_owned_persister_pickle_round_trip_reconnects_and_loads_existing_state( + aerospike_persister, +): + aerospike_persister.save( + "tenant-a", + "pickle-app", + 3, + "checkpoint", + state.State({"restored": True}), + "completed", + ) + + reconstructed = pickle.loads(pickle.dumps(aerospike_persister)) + try: + assert reconstructed.is_initialized() is False + loaded = reconstructed.load("tenant-a", "pickle-app", 3) + assert loaded["state"].get_all() == {"restored": True} + finally: + reconstructed.cleanup() + + +def test_owned_persister_pickle_preserves_membership_set(): + persister = AerospikeBasePersister.from_values( + membership_set=f"membership_{uuid.uuid4().hex[:12]}" + ) + reconstructed = pickle.loads(pickle.dumps(persister)) + try: + assert reconstructed.membership_set == persister.membership_set + finally: + reconstructed.cleanup() + persister.cleanup() + + +def test_latest_load_returns_the_checkpoint_with_the_greatest_sequence( + aerospike_persister, +): + aerospike_persister.save("pk", "app", 2, "second", state.State({"value": 2}), "completed") + aerospike_persister.save("pk", "app", 1, "first", state.State({"value": 1}), "completed") + + loaded = aerospike_persister.load("pk", "app") + + assert loaded["sequence_id"] == 2 + assert loaded["state"].get_all() == {"value": 2} + + +def test_absent_exact_checkpoint_and_absent_head_return_none(aerospike_persister): + assert aerospike_persister.load("pk", "missing", 1) is None + assert aerospike_persister.load("pk", "missing") is None + + +def test_load_requires_an_application_id(aerospike_persister): + with pytest.raises(ValueError, match="app_id"): + aerospike_persister.load("pk", None) + + +def test_none_empty_and_literal_none_partitions_remain_distinct(aerospike_persister): + for partition_key, value in [(None, "null"), ("", "empty"), ("None", "literal")]: + aerospike_persister.save( + partition_key, + "app", + 1, + "position", + state.State({"value": value}), + "completed", + ) + + assert aerospike_persister.load(None, "app", 1)["state"].get_all() == {"value": "null"} + assert aerospike_persister.load("", "app", 1)["state"].get_all() == {"value": "empty"} + assert aerospike_persister.load("None", "app", 1)["state"].get_all() == {"value": "literal"} + assert set(aerospike_persister.list_app_ids(None)) == {"app"} + assert set(aerospike_persister.list_app_ids("")) == {"app"} + assert set(aerospike_persister.list_app_ids("None")) == {"app"} + + +def test_list_app_ids_returns_each_application_once_without_an_order_contract( + aerospike_persister, +): + aerospike_persister.save("pk", "app-a", 1, "one", state.State({"v": 1}), "completed") + aerospike_persister.save("pk", "app-a", 2, "two", state.State({"v": 2}), "completed") + aerospike_persister.save("pk", "app-b", 1, "one", state.State({"v": 3}), "completed") + + assert set(aerospike_persister.list_app_ids("pk")) == {"app-a", "app-b"} + assert aerospike_persister.list_app_ids("another-partition") == [] + + +def test_identical_save_is_idempotent_and_preserves_the_first_creation_time( + aerospike_persister, +): + checkpoint = state.State({"message": "hello", "nested": {"b": 2, "a": 1}}) + aerospike_persister.save("pk", "app", 1, "position", checkpoint, "completed") + first = aerospike_persister.load("pk", "app", 1) + + aerospike_persister.save( + "pk", + "app", + 1, + "position", + state.State({"nested": {"a": 1, "b": 2}, "message": "hello"}), + "completed", + ) + + assert aerospike_persister.load("pk", "app", 1)["created_at"] == first["created_at"] + + +@pytest.mark.parametrize( + ("position", "saved_state", "status"), + [ + ("different", state.State({"value": 1}), "completed"), + ("position", state.State({"value": 2}), "completed"), + ("position", state.State({"value": 1}), "failed"), + ], +) +def test_duplicate_checkpoint_save_is_idempotent( + aerospike_persister, position, saved_state, status +): + """A duplicate save by primary key is a no-op; the first checkpoint is preserved.""" + aerospike_persister.save("pk", "app", 1, "position", state.State({"value": 1}), "completed") + + # A second write with the same key but different content must not raise + # and must not overwrite the immutable history record. + aerospike_persister.save("pk", "app", 1, position, saved_state, status) + + loaded = aerospike_persister.load("pk", "app", 1) + assert loaded["position"] == "position" + assert loaded["state"].get_all() == {"value": 1} + assert loaded["status"] == "completed" + + +@pytest.mark.parametrize("sequence_id", [True, -(2**63) - 1, 2**63]) +def test_invalid_sequence_is_rejected_before_persistence(aerospike_persister, sequence_id): + with pytest.raises(ValueError, match="sequence"): + aerospike_persister.save( + "pk", + "invalid-sequence", + sequence_id, + "position", + state.State({}), + "completed", + ) + + assert aerospike_persister.list_app_ids("pk") == [] + + +@pytest.mark.parametrize("value", [float("nan"), float("inf")]) +def test_non_finite_state_is_rejected_without_creating_a_head(aerospike_persister, value): + with pytest.raises((TypeError, ValueError), match="serializ|JSON|finite"): + aerospike_persister.save( + "pk", + "invalid-state", + 1, + "position", + state.State({"value": value}), + "completed", + ) + + assert aerospike_persister.load("pk", "invalid-state") is None + + +def test_oversized_membership_preserves_existing_membership_and_loadable_state( + aerospike_persister, +): + partition = f"oversized-membership-{uuid.uuid4().hex}" + first_app = "a" * 600_000 + second_app = "b" * 600_000 + aerospike_persister.save(partition, first_app, 1, "first", state.State({"v": 1}), "completed") + + aerospike_persister.save(partition, second_app, 1, "second", state.State({"v": 2}), "completed") + + assert set(aerospike_persister.list_app_ids(partition)) == {first_app, second_app} + assert aerospike_persister.load(partition, first_app)["sequence_id"] == 1 + assert aerospike_persister.load(partition, second_app)["sequence_id"] == 1 + + +def test_repeated_initialization_is_idempotent(aerospike_persister): + aerospike_persister.initialize() + assert aerospike_persister.is_initialized() is True + + +def test_head_never_regresses_when_an_older_checkpoint_is_retried(aerospike_persister): + old = state.State({"value": "old"}) + aerospike_persister.save("pk", "app", 1, "old", old, "completed") + aerospike_persister.save("pk", "app", 2, "new", state.State({"value": "new"}), "completed") + aerospike_persister.save("pk", "app", 1, "old", old, "completed") + + assert aerospike_persister.load("pk", "app")["sequence_id"] == 2 + + +class CallerOwnedClient: + def __init__(self): + self.close_calls = 0 + self.database_calls = 0 + + def close(self): + self.close_calls += 1 + + def get(self, *args, **kwargs): + self.database_calls += 1 + raise AssertionError("the database must not be accessed") + + def put(self, *args, **kwargs): + self.database_calls += 1 + raise AssertionError("the database must not be accessed") + + def operate(self, *args, **kwargs): + self.database_calls += 1 + raise AssertionError("the database must not be accessed") + + +def test_injected_client_constructs_a_synchronous_base_persister(): + persister = AerospikeBasePersister(client=CallerOwnedClient()) + + assert isinstance(persister, BaseStatePersister) + assert persister.is_async() is False + + +def test_cleanup_and_context_exit_never_close_an_injected_client(): + client = CallerOwnedClient() + persister = AerospikeBasePersister(client=client) + + persister.cleanup() + persister.cleanup() + with persister: + pass + + assert client.close_calls == 0 + + +def test_cleanup_closes_an_internally_constructed_client_once(): + client = Mock() + with patch( + "burr.integrations.persisters.b_aerospike.aerospike.client", return_value=client + ) as factory: + persister = AerospikeBasePersister.from_values() + + persister.cleanup() + persister.cleanup() + + factory.assert_called_once() + client.close.assert_called_once_with() + + +def test_from_config_accepts_official_client_configuration(): + client = Mock() + config = { + "hosts": [("aerospike.internal", 3000)], + "client_config": {"user": "service-user", "password": "secret"}, + "namespace": "test", + "history_set": "history", + "head_set": "heads", + "membership_set": "applications", + "key_prefix": "service-a", + } + with patch( + "burr.integrations.persisters.b_aerospike.aerospike.client", return_value=client + ) as factory: + persister = AerospikeBasePersister.from_config(config) + + try: + supplied_config = factory.call_args.args[0] + assert supplied_config["hosts"] == [("aerospike.internal", 3000)] + assert supplied_config["user"] == "service-user" + assert supplied_config["password"] == "secret" + assert persister.membership_set == "applications" + finally: + persister.cleanup() + + +@pytest.mark.parametrize("sequence_id", [True, -(2**63) - 1, 2**63, 1.0, "1"]) +def test_invalid_sequence_is_rejected_before_client_access(sequence_id): + client = CallerOwnedClient() + persister = AerospikeBasePersister(client=client) + + with pytest.raises(ValueError, match="sequence"): + persister.save("pk", "app", sequence_id, "position", state.State({}), "completed") + + assert client.database_calls == 0 + + +def test_load_without_an_app_id_is_rejected_before_client_access(): + client = CallerOwnedClient() + persister = AerospikeBasePersister(client=client) + + with pytest.raises(ValueError, match="app_id"): + persister.load("pk", None) + + assert client.database_calls == 0 + + +def test_concurrent_saves_to_same_application_advance_monotonically(aerospike_persister): + """Concurrent saves to one application should create every checkpoint and leave the head at the max sequence.""" + sequences = list(range(1, 11)) + + def save(seq): + aerospike_persister.save( + "pk", + "concurrent-app", + seq, + f"position-{seq}", + state.State({"seq": seq}), + "completed", + ) + + with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor: + list(executor.map(save, sequences)) + + latest = aerospike_persister.load("pk", "concurrent-app") + assert latest is not None + assert latest["sequence_id"] == max(sequences) + assert aerospike_persister.list_app_ids("pk") == ["concurrent-app"] + + for seq in sequences: + loaded = aerospike_persister.load("pk", "concurrent-app", seq) + assert loaded is not None + assert loaded["sequence_id"] == seq + + +def test_owned_persister_uses_the_factory_connected_client(): + client = Mock() + with patch("burr.integrations.persisters.b_aerospike.aerospike.client", return_value=client): + persister = AerospikeBasePersister.from_values() + + try: + client.connect.assert_not_called() + finally: + persister.cleanup() + + +class RecordingClient: + def __init__(self): + self.put_policy = None + self.operate_policies = [] + self.get_calls = 0 + self.head_operates = 0 + self.membership_operations = None + + def put(self, key, bins, policy): + self.put_policy = policy + + def operate(self, key, operations, policy): + self.operate_policies.append(policy) + if key[1] == "burr_head": + self.head_operates += 1 + if self.head_operates == 1: + return key, {}, {"sequence_id": 1} + return key, {}, {"sequence_id": 1, "member_confirm": True} + self.membership_operations = operations + return key, {}, {} + + def get(self, key, policy): + self.get_calls += 1 + raise AssertionError("successful operate must not require a reconciliation read") + + +def test_save_uses_write_policies_and_the_successful_operate_result(): + client = RecordingClient() + persister = AerospikeBasePersister(client=client) + + persister.save( + "partition", + "application", + 1, + "position", + state.State({"value": 1}), + "completed", + ) + + assert "replica" not in client.put_policy + assert all("replica" not in policy for policy in client.operate_policies) + assert all( + policy["commit_level"] == aerospike.POLICY_COMMIT_LEVEL_ALL + for policy in client.operate_policies + ) + assert all(policy["ttl"] == aerospike.TTL_NEVER_EXPIRE for policy in client.operate_policies) + assert client.operate_policies[-1]["max_retries"] == 0 + assert len(client.membership_operations) == 2 + assert client.membership_operations[0]["bin"] == "tail_page" + assert client.membership_operations[0]["val"] == 0 + assert client.membership_operations[1]["bin"] == "app_ids" + assert client.get_calls == 0 + + +class ConditionalMembershipClient: + def __init__(self, confirmation_error=None): + self.confirmed = False + self.sequence_id = None + self.membership = {} + self.membership_accesses = 0 + self.confirmation_error = confirmation_error + + def put(self, key, bins, policy): + return None + + def operate(self, key, operations, policy): + if key[1] == "burr_apps": + self.membership_accesses += 1 + self.membership["app"] = 1 + return key, {}, {} + if self.sequence_id is None or not self.confirmed: + if self.sequence_id is None: + self.sequence_id = 1 + return key, {}, {"sequence_id": self.sequence_id} + if self.confirmation_error is not None: + error, self.confirmation_error = self.confirmation_error, None + self.confirmed = True + raise error + self.confirmed = True + return key, {}, {"sequence_id": self.sequence_id, "member_confirm": True} + self.sequence_id = max(self.sequence_id, 2) + return key, {}, {"sequence_id": self.sequence_id, "member_confirm": True} + + def get(self, key, policy): + return ( + key, + {}, + { + "partition": '"pk"', + "key_prefix": '""', + "app_id": "app", + "sequence_id": self.sequence_id, + **({"member_confirm": True} if self.confirmed else {}), + }, + ) + + def select(self, key, bins, policy): + self.membership_accesses += 1 + return ( + key, + {}, + { + "partition": '"pk"', + "key_prefix": '""', + "app_ids": self.membership, + "tail_page": 0, + }, + ) + + +def test_first_save_confirms_membership_and_subsequent_save_skips_membership_access(): + client = ConditionalMembershipClient() + persister = AerospikeBasePersister(client=client) + + persister.save("pk", "app", 1, "one", state.State({"value": 1}), "completed") + assert persister.list_app_ids("pk") == ["app"] + accesses_after_first_save_and_list = client.membership_accesses + + persister.save("pk", "app", 2, "two", state.State({"value": 2}), "completed") + + assert client.confirmed is True + assert client.membership_accesses == accesses_after_first_save_and_list + + +def test_ambiguous_confirmation_is_reconciled_from_the_head(): + client = ConditionalMembershipClient(aerospike.exception.TimeoutError()) + persister = AerospikeBasePersister(client=client) + + persister.save("pk", "app", 1, "one", state.State({"value": 1}), "completed") + + assert client.confirmed is True + assert persister.list_app_ids("pk") == ["app"] + + +def test_save_rejects_a_stored_false_membership_confirmation_without_membership_access(): + client = ConditionalMembershipClient() + client.sequence_id = 1 + client.operate = Mock(side_effect=aerospike.exception.FilteredOut()) + client.get = Mock( + return_value=( + ("test", "burr_head", "key"), + {}, + { + "partition": '"pk"', + "key_prefix": '""', + "app_id": "app", + "sequence_id": 1, + "member_confirm": False, + }, + ) + ) + persister = AerospikeBasePersister(client=client) + + with pytest.raises(AerospikePersistenceConsistencyError, match="must be true"): + persister.save("pk", "app", 1, "one", state.State({"value": 1}), "completed") + + assert client.membership_accesses == 0 + + +class MembershipRetryClient: + def __init__(self, membership_results, select_results): + self.membership_results = iter(membership_results) + self.select_results = iter(select_results) + self.operate_calls = 0 + + def put(self, key, bins, policy): + return None + + def operate(self, key, operations, policy): + self.operate_calls += 1 + if self.operate_calls == 1: + return key, {}, {"sequence_id": 1} + if key[1] == "burr_head": + return key, {}, {"sequence_id": 1, "member_confirm": True} + result = next(self.membership_results) + if isinstance(result, Exception): + raise result + return result + + def select(self, key, bins, policy): + result = next(self.select_results) + if isinstance(result, Exception): + raise result + return result + + +def test_membership_registration_retries_after_unobserved_timeout(): + client = MembershipRetryClient( + [aerospike.exception.TimeoutError(), (("test", "burr_apps", "key"), {}, {})], + [aerospike.exception.RecordNotFound()], + ) + persister = AerospikeBasePersister(client=client) + + with patch.object(persister, "_backoff") as backoff: + persister.save("pk", "app", 1, "position", state.State({"value": 1}), "completed") + + assert client.operate_calls == 4 + backoff.assert_called_once_with(1) + + +def test_membership_registration_reconciles_observed_timeout(): + client = MembershipRetryClient( + [aerospike.exception.TimeoutError()], + [ + ( + ("test", "burr_apps", "key"), + {}, + {"partition": '"pk"', "key_prefix": '""', "app_ids": {"app": 1}}, + ) + ], + ) + persister = AerospikeBasePersister(client=client) + + persister.save("pk", "app", 1, "position", state.State({"value": 1}), "completed") + + assert client.operate_calls == 3 + + +def test_membership_registration_exhausts_uncertain_retry_budget(): + client = MembershipRetryClient( + [aerospike.exception.TimeoutError()] * 3, + [aerospike.exception.RecordNotFound()] * 3, + ) + persister = AerospikeBasePersister(client=client) + + with patch.object(persister, "_backoff"), pytest.raises( + AerospikePersistenceUncertainOutcomeError, match="Membership" + ): + persister.save("pk", "app", 1, "position", state.State({"value": 1}), "completed") + + assert client.operate_calls == 4 + + +def test_list_app_ids_uses_one_membership_primary_key_read(): + client = Mock() + client.select.return_value = ( + ("test", "burr_apps", "digest"), + {}, + {"app_ids": {"a": 1, "b": 1}, "tail_page": 0, "unknown": "ignored"}, + ) + persister = AerospikeBasePersister(client=client) + + assert set(persister.list_app_ids("pk")) == {"a", "b"} + + client.select.assert_called_once() + key, bins = client.select.call_args.args + assert key == ( + "test", + "burr_apps", + "aeb3748692d1637fc73dba70cd2887713db3fa6c26ca6ce8a01240d426228c5b", + ) + assert bins == ["app_ids", "tail_page"] + client.query.assert_not_called() + + +def test_list_app_ids_treats_an_absent_membership_record_as_empty(): + client = Mock() + client.select.side_effect = aerospike.exception.RecordNotFound() + persister = AerospikeBasePersister(client=client) + + assert persister.list_app_ids("pk") == [] + client.query.assert_not_called() + + +def test_list_app_ids_propagates_missing_membership_map(): + client = Mock() + client.select.return_value = (("test", "burr_apps", "digest"), {}, {}) + persister = AerospikeBasePersister(client=client) + + with pytest.raises(AerospikePersistenceConsistencyError, match="application-ID map"): + persister.list_app_ids("pk") + + +def test_list_app_ids_propagates_non_iterable_membership_map(): + client = Mock() + client.select.return_value = ( + ("test", "burr_apps", "digest"), + {}, + {"app_ids": None}, + ) + persister = AerospikeBasePersister(client=client) + + with pytest.raises(AerospikePersistenceConsistencyError, match="application-ID map"): + persister.list_app_ids("pk") + + +def test_initialize_is_idempotent_without_remote_calls(): + client = Mock() + persister = AerospikeBasePersister(client=client) + + persister.initialize() + persister.initialize() + + assert persister.is_initialized() is True + client.assert_not_called() + assert client.method_calls == []