diff --git a/lib/cassandra/cluster/schema/cql_type_parser.rb b/lib/cassandra/cluster/schema/cql_type_parser.rb index a10f08af4..03864b76e 100644 --- a/lib/cassandra/cluster/schema/cql_type_parser.rb +++ b/lib/cassandra/cluster/schema/cql_type_parser.rb @@ -74,6 +74,8 @@ def lookup_type(node, types) Cassandra::Types.frozen(lookup_type(node.children.first, types)) when 'empty' then Cassandra::Types.custom('org.apache.cassandra.db.marshal.EmptyType') + when 'vector' then + Cassandra::Types.vector(lookup_type(node.children[0], types), node.children[1]) when /\A'/ then # Custom type. Cassandra::Types.custom(node.name[1..-2]) diff --git a/lib/cassandra/driver.rb b/lib/cassandra/driver.rb index 096d8bacc..f99df3e34 100644 --- a/lib/cassandra/driver.rb +++ b/lib/cassandra/driver.rb @@ -253,6 +253,9 @@ def create_schema_fetcher_picker picker.when('4.') do Cluster::Schema::Fetchers::V3_0_x.new(schema_cql_type_parser, cluster_schema) end + picker.when('5.') do + Cluster::Schema::Fetchers::V3_0_x.new(schema_cql_type_parser, cluster_schema) + end picker end diff --git a/lib/cassandra/protocol/coder.rb b/lib/cassandra/protocol/coder.rb index f83c174b6..e3e6cbf8c 100644 --- a/lib/cassandra/protocol/coder.rb +++ b/lib/cassandra/protocol/coder.rb @@ -49,6 +49,16 @@ def write_list_v4(buffer, list, type) buffer.append_bytes(raw) end + def write_vector_v4(buffer, vector, type) + raw = CqlByteBuffer.new + + vector.each do |element| + write_value_v4(raw, element, type, false) + end + + buffer.append_bytes(raw) + end + def write_map_v4(buffer, map, key_type, value_type) raw = CqlByteBuffer.new @@ -81,7 +91,7 @@ def write_tuple_v4(buffer, value, members) buffer.append_bytes(raw) end - def write_value_v4(buffer, value, type) + def write_value_v4(buffer, value, type, with_size = true) if value.nil? buffer.append_int(-1) return @@ -94,24 +104,25 @@ def write_value_v4(buffer, value, type) case type.kind when :ascii then write_ascii(buffer, value) - when :bigint, :counter then write_bigint(buffer, value) + when :bigint, :counter then write_bigint(buffer, value, with_size) when :blob then write_blob(buffer, value) when :boolean then write_boolean(buffer, value) when :custom then write_custom(buffer, value, type) when :decimal then write_decimal(buffer, value) - when :double then write_double(buffer, value) - when :float then write_float(buffer, value) - when :int then write_int(buffer, value) + when :double then write_double(buffer, value, with_size) + when :float then write_float(buffer, value, with_size) + when :int then write_int(buffer, value, with_size) when :inet then write_inet(buffer, value) when :timestamp then write_timestamp(buffer, value) when :uuid, :timeuuid then write_uuid(buffer, value) when :text then write_text(buffer, value) when :varint then write_varint(buffer, value) - when :tinyint then write_tinyint(buffer, value) - when :smallint then write_smallint(buffer, value) + when :tinyint then write_tinyint(buffer, value, with_size) + when :smallint then write_smallint(buffer, value, with_size) when :time then write_time(buffer, value) when :date then write_date(buffer, value) when :list, :set then write_list_v4(buffer, value, type.value_type) + when :vector then write_vector_v4(buffer, value, type.value_type) when :map then write_map_v4(buffer, value, type.key_type, type.value_type) @@ -234,24 +245,24 @@ def read_values_v4(buffer, column_metadata, custom_type_handlers) end end - def read_value_v4(buffer, type, custom_type_handlers) + def read_value_v4(buffer, type, custom_type_handlers, with_size = true) case type.kind when :ascii then read_ascii(buffer) - when :bigint, :counter then read_bigint(buffer) + when :bigint, :counter then read_bigint(buffer, with_size) when :blob then buffer.read_bytes when :boolean then read_boolean(buffer) when :decimal then read_decimal(buffer) - when :double then read_double(buffer) - when :float then read_float(buffer) - when :int then read_int(buffer) + when :double then read_double(buffer, with_size) + when :float then read_float(buffer, with_size) + when :int then read_int(buffer, with_size) when :timestamp then read_timestamp(buffer) when :uuid then read_uuid(buffer) when :timeuuid then read_uuid(buffer, TimeUuid) when :text then read_text(buffer) when :varint then read_varint(buffer) when :inet then read_inet(buffer) - when :tinyint then read_tinyint(buffer) - when :smallint then read_smallint(buffer) + when :tinyint then read_tinyint(buffer, with_size) + when :smallint then read_smallint(buffer, with_size) when :time then read_time(buffer) when :date then read_date(buffer) when :custom then read_custom(buffer, type, custom_type_handlers) @@ -317,6 +328,15 @@ def read_value_v4(buffer, type, custom_type_handlers) values.fill(nil, values.length, (members.length - values.length)) Cassandra::Tuple::Strict.new(members, values) + when :vector + return nil unless read_size(buffer) + + value_type = type.value_type + value = ::Array.new + type.dimension.to_i.times do + value << read_value_v4(buffer, value_type, custom_type_handlers, false) + end + value else raise Errors::DecodingError, %(Unsupported value type: #{type}) end @@ -767,8 +787,9 @@ def read_ascii(buffer) value && value.force_encoding(::Encoding::ASCII) end - def read_bigint(buffer) - read_size(buffer) && buffer.read_long + def read_bigint(buffer, with_size = true) + read_size(buffer) || return if with_size + buffer.read_long end alias read_counter read_bigint @@ -791,16 +812,19 @@ def read_decimal(buffer) size && buffer.read_decimal(size) end - def read_double(buffer) - read_size(buffer) && buffer.read_double + def read_double(buffer, with_size = true) + read_size(buffer) || return if with_size + buffer.read_double end - def read_float(buffer) - read_size(buffer) && buffer.read_float + def read_float(buffer, with_size = true) + read_size(buffer) || return if with_size + buffer.read_float end - def read_int(buffer) - read_size(buffer) && buffer.read_signed_int + def read_int(buffer, with_size = true) + read_size(buffer) || return if with_size + buffer.read_signed_int end def read_timestamp(buffer) @@ -832,12 +856,14 @@ def read_inet(buffer) size && ::IPAddr.new_ntoh(buffer.read(size)) end - def read_tinyint(buffer) - read_size(buffer) && buffer.read_tinyint + def read_tinyint(buffer, with_size = true) + read_size(buffer) || return if with_size + buffer.read_tinyint end - def read_smallint(buffer) - read_size(buffer) && buffer.read_smallint + def read_smallint(buffer, with_size = true) + read_size(buffer) || return if with_size + buffer.read_smallint end def read_time(buffer) @@ -856,8 +882,8 @@ def write_ascii(buffer, value) buffer.append_bytes(value.encode(::Encoding::ASCII)) end - def write_bigint(buffer, value) - buffer.append_int(8) + def write_bigint(buffer, value, with_size = true) + buffer.append_int(8) if with_size buffer.append_long(value) end @@ -885,18 +911,18 @@ def write_decimal(buffer, value) buffer.append_bytes(CqlByteBuffer.new.append_decimal(value)) end - def write_double(buffer, value) - buffer.append_int(8) + def write_double(buffer, value, with_size = true) + buffer.append_int(8) if with_size buffer.append_double(value) end - def write_float(buffer, value) - buffer.append_int(4) + def write_float(buffer, value, with_size = true) + buffer.append_int(4) if with_size buffer.append_float(value) end - def write_int(buffer, value) - buffer.append_int(4) + def write_int(buffer, value, with_size = true) + buffer.append_int(4) if with_size buffer.append_int(value) end @@ -924,13 +950,13 @@ def write_varint(buffer, value) buffer.append_bytes(CqlByteBuffer.new.append_varint(value)) end - def write_tinyint(buffer, value) - buffer.append_int(1) + def write_tinyint(buffer, value, with_size = true) + buffer.append_int(1) if with_size buffer.append_tinyint(value) end - def write_smallint(buffer, value) - buffer.append_int(2) + def write_smallint(buffer, value, with_size = true) + buffer.append_int(2) if with_size buffer.append_smallint(value) end diff --git a/lib/cassandra/types.rb b/lib/cassandra/types.rb index accea6331..2d53ef362 100644 --- a/lib/cassandra/types.rb +++ b/lib/cassandra/types.rb @@ -844,6 +844,70 @@ def eql?(other) alias == eql? end + class Vector < Type + # @private + attr_reader :value_type, :dimension + + # @private + def initialize(value_type, dimension) + super(:vector) + @value_type = value_type + @dimension = dimension + end + + # Coerces the value to Array + # @param value [Object] original value + # @return [Array] value + # @see Cassandra::Type#new + def new(*value) + value = Array(value.first) if value.one? + + value.each do |v| + Util.assert_type(@value_type, v) + end + value + end + + # Asserts that a given value is an Array with matching length + # @param value [Object] value to be validated + # @param message [String] error message to use when assertion fails + # @yieldreturn [String] error message to use when assertion fails + # @raise [ArgumentError] if the value is not an Array + # @raise [ArgumentError] if the Array is not the correct size + # @return [void] + # @see Cassandra::Type#assert + def assert(value, message = nil, &block) + Util.assert_instance_of(::Array, value, message, &block) + Util.assert_size(@dimension.to_i, value, message, &block) + value.each do |v| + Util.assert_type(@value_type, v, message, &block) + end + nil + end + + # @return [String] `"vector"` + # @see Cassandra::Type#to_s + def to_s + "vector<#{@value_type},#{@dimension}>" + end + + def hash + @hash ||= begin + h = 17 + h = 31 * h + @kind.hash + h = 31 * h + @value_type.hash + h = 31 * h + @dimension.hash + h + end + end + + def eql?(other) + other.is_a?(Vector) && @value_type == other.value_type && @dimension == other.dimension + end + + alias == eql? + end + # @!parse # class Smallint < Type # # @return [Symbol] `:smallint` @@ -1611,6 +1675,16 @@ def set(value_type) Set.new(value_type) end + # @param value_type [Cassandra::Type] the type of elements in this list + # @param dimension [Integer] the dimension of the vector + # @return [Cassandra::Types::Vector] vector type + def vector(value_type, dimension) + Util.assert_instance_of(Cassandra::Type, value_type, + "list type must be a Cassandra::Type, #{value_type.inspect} given") + + Vector.new(value_type, dimension) + end + # @param members [*Cassandra::Type] types of members of this tuple # @return [Cassandra::Types::Tuple] tuple type def tuple(*members) @@ -1698,7 +1772,14 @@ def udt(keyspace, name, *fields) # @param name [String] name of the custom type # @return [Cassandra::Types::Custom] custom type def custom(name) - Custom.new(name) + case name + when /VectorType\(([\w.]+),\s?(\d+)\)/ + type = Cluster::Schema::FQCNTypeParser.new.parse($1).results[0][0] + dimension = $2 + Vector.new(type, dimension) + else + Custom.new(name) + end end def duration