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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 38 additions & 8 deletions datasketch/storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -326,8 +326,14 @@ class CassandraClient:
key blob,
value blob,
ts bigint,
PRIMARY KEY (key, value)
) WITH CLUSTERING ORDER BY (value DESC)
PRIMARY KEY (key, value, ts)
) WITH CLUSTERING ORDER BY (value DESC, ts ASC)
"""

QUERY_GET_TABLE_SCHEMA = """
SELECT column_name, kind, position
FROM system_schema.columns
WHERE keyspace_name = ? AND table_name = ?
"""

QUERY_DROP_TABLE = "DROP TABLE IF EXISTS {}"
Expand Down Expand Up @@ -366,11 +372,7 @@ class CassandraClient:
WHERE key = ? AND value = ?
"""

QUERY_UPSERT = """
UPDATE {}
SET ts = ?
WHERE key = ? AND value = ?
"""
QUERY_UPSERT = "INSERT INTO {} (key, value, ts) VALUES (?, ?, 0)"

QUERY_INSERT = "INSERT INTO {} (key, value, ts) VALUES (?, ?, ?)"

Expand Down Expand Up @@ -406,6 +408,7 @@ def __init__(self, cassandra_params, name, buffer_size):
if cassandra_params.get("drop_tables", False):
self._session.execute(self.QUERY_DROP_TABLE.format(table_name))
self._session.execute(self.QUERY_CREATE_TABLE.format(table_name))
self._validate_table_schema(table_name)

# Prepare all the statements for this table
self._stmt_insert = self._session.prepare(self.QUERY_INSERT.format(table_name))
Expand All @@ -417,6 +420,30 @@ def __init__(self, cassandra_params, name, buffer_size):
self._stmt_delete_key = self._session.prepare(self.QUERY_DELETE_KEY.format(table_name))
self._stmt_delete_val = self._session.prepare(self.QUERY_DELETE_VAL.format(table_name))

def _validate_table_schema(self, table_name):
"""Reject tables created with the old duplicate-collapsing key."""
statement = self._session.prepare(self.QUERY_GET_TABLE_SCHEMA)
rows = self._session.execute(
statement,
(self._session.keyspace, table_name),
)
primary_key = {
(row.kind, row.position): row.column_name
for row in rows
if row.kind in ("partition_key", "clustering")
}
expected = {
("partition_key", 0): "key",
("clustering", 0): "value",
("clustering", 1): "ts",
}
if primary_key != expected:
raise RuntimeError(
"Cassandra table %r uses an incompatible primary key. "
"Rebuild the LSH table (or set drop_tables=True while rebuilding) "
"so duplicate ordered values can be preserved." % table_name
)

@property
def buffer_size(self):
"""Get the buffer size.
Expand Down Expand Up @@ -527,7 +554,10 @@ def upsert(self, key, vals, buffer=False):
:param iterable[byte|str] vals: the iterable of values
:param boolean buffer: whether the upsert statements should be buffered
"""
statements_and_parameters = [(self._stmt_upsert, (self._ts(), key, val)) for val in vals]
# The timestamp is part of the primary key so ordered storage can
# preserve duplicates. Sets use the fixed timestamp in
# QUERY_UPSERT, making repeated (key, value) inserts idempotent.
statements_and_parameters = [(self._stmt_upsert, (key, val)) for val in vals]
if buffer:
self._buffer(statements_and_parameters)
else:
Expand Down
72 changes: 72 additions & 0 deletions test/test_integration.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,14 @@
import os
from types import SimpleNamespace
from unittest.mock import Mock, patch

import numpy as np
import pytest

import datasketch.storage as storage_module
from datasketch.lsh import MinHashLSH
from datasketch.minhash import MinHash
from datasketch.storage import ordered_storage
from datasketch.weighted_minhash import WeightedMinHashGenerator

STORAGE_CONFIG_REDIS = {
Expand All @@ -29,6 +33,55 @@
DO_TEST_CASSANDRA = os.environ.get("DO_TEST_CASSANDRA") == "true"


@pytest.mark.skipif(
not hasattr(storage_module, "CassandraClient"),
reason="cassandra-driver is unavailable",
)
def test_cassandra_schema_validation_and_set_upsert():
client = object.__new__(storage_module.CassandraClient)
client._session = Mock(keyspace="lsh_test")
client._session.prepare.return_value = "prepared-schema-query"
client._session.execute.return_value = [
SimpleNamespace(kind="partition_key", position=0, column_name="key"),
SimpleNamespace(kind="clustering", position=0, column_name="value"),
SimpleNamespace(kind="clustering", position=1, column_name="ts"),
SimpleNamespace(kind="regular", position=-1, column_name="ignored"),
]

client._validate_table_schema("lsh_table")
client._session.prepare.assert_called_once_with(client.QUERY_GET_TABLE_SCHEMA)
client._session.execute.assert_called_once_with(
"prepared-schema-query",
("lsh_test", "lsh_table"),
)

client._session.execute.return_value = [
SimpleNamespace(kind="partition_key", position=0, column_name="key"),
SimpleNamespace(kind="clustering", position=0, column_name="value"),
]
with pytest.raises(RuntimeError, match="incompatible primary key"):
client._validate_table_schema("legacy_table")

client._stmt_upsert = "set-upsert"
client._execute = Mock()
client.upsert(b"key", [b"same", b"same"])
client._execute.assert_called_once_with(
[
("set-upsert", (b"key", b"same")),
("set-upsert", (b"key", b"same")),
]
)

session = Mock()
with (
patch.object(storage_module.CassandraSharedSession, "get_session", return_value=session),
patch.object(storage_module.c_cluster, "MonotonicTimestampGenerator", return_value=Mock()),
patch.object(storage_module.CassandraClient, "_validate_table_schema") as validate_schema,
):
storage_module.CassandraClient({}, b"table", 100)
validate_schema.assert_called_once_with("lsh_table")


def _clear_redis_keys(pattern="lsh_test*"):
if not DO_TEST_REDIS:
return
Expand Down Expand Up @@ -108,6 +161,25 @@ def test_insert(self, storage_config):
for i, H in enumerate(lsh.keys[b"a"]):
assert b"a" in lsh.hashtables[i][H]

def test_ordered_storage_preserves_duplicate_values(self, storage_config):
storage = ordered_storage(storage_config, name=b"lsh_test_duplicate_values")
storage.insert(b"key", b"same", b"same")
assert list(storage.get(b"key")) == [b"same", b"same"]

def test_remove_empty_minhash_clears_every_band(self, storage_config):
lsh = MinHashLSH(threshold=0.5, num_perm=16, storage_config=storage_config, prepickle=False)
empty = MinHash(16)
lsh.insert(b"empty", empty)

band_hashes = list(lsh.keys[b"empty"])
assert len(band_hashes) == lsh.b
assert len(set(band_hashes)) == 1

lsh.remove(b"empty")
assert b"empty" not in lsh.keys
for band_hash, hashtable in zip(band_hashes, lsh.hashtables):
assert b"empty" not in hashtable.get(band_hash)

def test_insert_non_bytes_key_raises_error(self, storage_config):
"""Test that inserting non-bytes keys with prepickle=False raises TypeError."""
lsh = MinHashLSH(threshold=0.5, num_perm=16, storage_config=storage_config, prepickle=False)
Expand Down
Loading