diff --git a/test/unit/org/apache/cassandra/io/sstable/CQLSSTableWriterTest.java b/test/unit/org/apache/cassandra/io/sstable/CQLSSTableWriterTest.java index 45de365a6c..b7b67ba9bd 100644 --- a/test/unit/org/apache/cassandra/io/sstable/CQLSSTableWriterTest.java +++ b/test/unit/org/apache/cassandra/io/sstable/CQLSSTableWriterTest.java @@ -39,6 +39,9 @@ import java.util.stream.StreamSupport; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; + +import org.apache.cassandra.cql3.constraints.ConstraintViolationException; + import org.junit.Before; import org.junit.Ignore; import org.junit.Rule; @@ -54,6 +57,7 @@ import org.apache.cassandra.cql3.functions.types.LocalDate; import org.apache.cassandra.cql3.functions.types.TypeCodec; import org.apache.cassandra.cql3.functions.types.UDTValue; import org.apache.cassandra.cql3.functions.types.UserType; +import org.apache.cassandra.db.marshal.FloatType; import org.apache.cassandra.db.marshal.UTF8Type; import org.apache.cassandra.dht.ByteOrderedPartitioner; import org.apache.cassandra.dht.Murmur3Partitioner; @@ -77,6 +81,7 @@ import org.apache.cassandra.utils.ByteBufferUtil; import org.apache.cassandra.utils.FBUtilities; import org.apache.cassandra.utils.JavaDriverUtils; import org.apache.cassandra.utils.OutputHandler; +import org.assertj.core.api.Assertions; import static org.apache.cassandra.utils.Clock.Global.currentTimeMillis; import static org.junit.Assert.assertEquals; @@ -1578,6 +1583,83 @@ public abstract class CQLSSTableWriterTest assertFalse(indexDescriptor.isPerColumnIndexBuildComplete(new IndexIdentifier(keyspace, table, "idx2"))); } + @Test + public void testWritingVectorData() throws Exception + { + final String schema = "CREATE TABLE " + qualifiedTable + " (" + + " k int," + + " v1 VECTOR," + + " PRIMARY KEY (k)" + + ")"; + + CQLSSTableWriter writer = CQLSSTableWriter.builder() + .inDirectory(dataDir) + .forTable(schema) + .using("INSERT INTO " + keyspace + "." + table + " (k, v1) " + + "VALUES (?, ?)").build(); + + for (int i = 0; i < 100; i++) + { + writer.addRow(i, List.of( (float)i, (float)i, (float)i, (float)i, (float)i)); + } + + writer.close(); + loadSSTables(dataDir, keyspace, table); + + if (verifyDataAfterLoading) + { + UntypedResultSet resultSet = QueryProcessor.executeInternal("SELECT * FROM " + keyspace + "." + table); + + assertEquals(resultSet.size(), 100); + int cnt = 0; + for (UntypedResultSet.Row row : resultSet) + { + assertEquals(cnt, row.getInt("k")); + List vector = row.getVector("v1", FloatType.instance, 5); + Assertions.assertThat(vector).hasSize(5); + final float floatCount = (float)cnt; + Assertions.assertThat(vector).allMatch(val -> val == floatCount); + cnt++; + } + } + } + + @Test + public void testConstraintViolation() throws Exception + { + final String schema = "CREATE TABLE " + qualifiedTable + " (" + + " k int," + + " v1 int CHECK v1 < 5 ," + + " PRIMARY KEY (k)" + + ")"; + + CQLSSTableWriter writer = CQLSSTableWriter.builder() + .inDirectory(dataDir) + .forTable(schema) + .using("INSERT INTO " + keyspace + "." + table + " (k, v1) " + + "VALUES (?, ?)").build(); + + writer.addRow(1, 4); + + Assertions.assertThatThrownBy(() -> writer.addRow(2, 11)) + .describedAs("Should throw when adding a row that violates constraints") + .isInstanceOf(ConstraintViolationException.class) + .hasMessageContaining("Column value does not satisfy value constraint for column 'v1'. It should be v1 < 5"); + + writer.close(); + loadSSTables(dataDir, keyspace, table); + + if (verifyDataAfterLoading) + { + UntypedResultSet resultSet = QueryProcessor.executeInternal("SELECT * FROM " + keyspace + "." + table); + + assertEquals(resultSet.size(), 1); + UntypedResultSet.Row row = resultSet.one(); + assertEquals(1, row.getInt("k")); + assertEquals(4, row.getInt("v1")); + } + } + protected static void loadSSTables(File dataDir, final String ks, final String tb) throws ExecutionException, InterruptedException { SSTableLoader loader = new SSTableLoader(dataDir, new SSTableLoader.Client()