Skip to content
Draft
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
6 changes: 2 additions & 4 deletions mssql-tds/src/connection/client_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -511,8 +511,7 @@ impl ClientContext {
port: 1433,
instance_name: None,
},
// TODO: make V2 as default when full V2 support is added
vector_version: VectorVersion::V1,
vector_version: VectorVersion::V2,
column_encryption_setting: ColumnEncryptionSetting::Disabled,
column_encryption_key_store_providers: std::sync::Arc::new(
crate::security::keystore::ColumnEncryptionKeyStoreProviderRegistry::new(),
Expand Down Expand Up @@ -574,8 +573,7 @@ impl ClientContext {
port: 1433,
instance_name: None,
},
// TODO: make V2 as default when full V2 support is added
vector_version: VectorVersion::V1,
vector_version: VectorVersion::V2,
column_encryption_setting: ColumnEncryptionSetting::Disabled,
column_encryption_key_store_providers: std::sync::Arc::new(
crate::security::keystore::ColumnEncryptionKeyStoreProviderRegistry::new(),
Expand Down
36 changes: 35 additions & 1 deletion mssql-tds/src/datatypes/bulk_copy_metadata.rs
Original file line number Diff line number Diff line change
Expand Up @@ -710,8 +710,14 @@ impl BulkCopyColumnMetadata {
SqlDbType::Variant => "sql_variant".to_string(),
SqlDbType::Json => "nvarchar(max)".to_string(),
SqlDbType::Vector => {
use crate::datatypes::sqldatatypes::VectorBaseType;

let dims = self.vector_dimensions()?;
format!("vector({})", dims)
// `vector(N)` implies the float32 base type; float16 must be spelled out.
match VectorBaseType::try_from(self.scale)? {
VectorBaseType::Float32 => format!("vector({})", dims),
VectorBaseType::Float16 => format!("vector({}, float16)", dims),
}
}
})
}
Expand Down Expand Up @@ -1522,6 +1528,34 @@ mod tests {
assert_eq!(meta.get_sql_type_definition().unwrap(), "vector(3)");
}

#[test]
fn vector_dimensions_valid_float16() {
use crate::datatypes::sqldatatypes::VECTOR_HEADER_SIZE;
// Float16 (scale=1), element_size=2, 3 dimensions: header(8) + 3*2 = 14
let meta = BulkCopyColumnMetadata::new("v", SqlDbType::Vector, 0xF5)
.with_length(
(VECTOR_HEADER_SIZE + 3 * 2) as i32,
TypeLength::Variable((VECTOR_HEADER_SIZE + 3 * 2) as i32),
)
.with_scale(1);
assert_eq!(meta.vector_dimensions().unwrap(), 3);
}

#[test]
fn vector_sql_type_definition_float16() {
use crate::datatypes::sqldatatypes::VECTOR_HEADER_SIZE;
let meta = BulkCopyColumnMetadata::new("v", SqlDbType::Vector, 0xF5)
.with_length(
(VECTOR_HEADER_SIZE + 3 * 2) as i32,
TypeLength::Variable((VECTOR_HEADER_SIZE + 3 * 2) as i32),
)
.with_scale(1);
assert_eq!(
meta.get_sql_type_definition().unwrap(),
"vector(3, float16)"
);
}

#[test]
fn encoding_type_latin1_non_latin() {
let enc = EncodingType::Latin1;
Expand Down
16 changes: 16 additions & 0 deletions mssql-tds/src/datatypes/sql_vector.rs
Original file line number Diff line number Diff line change
Expand Up @@ -306,4 +306,20 @@ mod tests {
let vector = SqlVector::try_from_f32(values).unwrap();
assert_eq!(vector.base_type(), VectorBaseType::Float32);
}

#[test]
fn test_from_f16_valid() {
let vector = SqlVector::try_from_f16(vec![1.0, 2.0, 3.0]).unwrap();
assert_eq!(vector.base_type(), VectorBaseType::Float16);
assert_eq!(vector.dimension_count(), 3);
assert_eq!(vector.as_f32(), Some(&[1.0, 2.0, 3.0][..]));
assert_eq!(vector.total_size(), VECTOR_HEADER_SIZE + 3 * 2);
}

#[test]
fn test_f16_allows_more_dimensions_than_f32() {
let values = vec![0.0f32; VectorBaseType::Float16.max_dimensions() as usize];
assert!(SqlVector::try_from_f16(values.clone()).is_ok());
assert!(SqlVector::try_from_f32(values).is_err());
}
}
75 changes: 73 additions & 2 deletions mssql-tds/src/datatypes/tds_value_serializer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1717,8 +1717,13 @@ impl TdsValueSerializer {
writer.write_i32_async((*f).to_bits() as i32).await?;
}
}
VectorData::Float16(_) => {
todo!("Phase 2: Float16 serialization");
VectorData::Float16(vs) => {
// Narrow to IEEE 754 half-precision (round-to-nearest-even)
for f in vs {
writer
.write_u16_async(half::f16::from_f32(*f).to_bits())
.await?;
}
}
}

Expand Down Expand Up @@ -3742,6 +3747,72 @@ mod serializer_tests {
assert_eq!(p[4], 0xE7); // NVARCHAR
assert_eq!(p[5], 0x07); // prop_len = 7
}

// ── serialize_vector ──

fn vector_ctx(exact_size: usize) -> TdsTypeContext {
let mut ctx = nullable_ctx(crate::datatypes::sqldatatypes::TdsDataType::Vector as u8);
ctx.max_size = exact_size;
ctx
}

#[test]
fn vector_float32_elements() {
let mut mock = MockNetworkWriter::new(64);
let mut w = PacketWriter::new(PacketType::TabularResult, &mut mock, None, None);
let vector =
crate::datatypes::sql_vector::SqlVector::try_from_f32(vec![1.0, 2.0, 3.0]).unwrap();
let ctx = vector_ctx(8 + 3 * 4);
block_on(TdsValueSerializer::serialize_value(
&mut w,
&ColumnValues::Vector(vector),
&ctx,
))
.unwrap();
let p = payload(&w);
assert_eq!(&p[0..2], &20u16.to_le_bytes()); // 8 header + 3 * 4
assert_eq!(&p[2..10], &[0xA9, 0x01, 0x03, 0x00, 0x00, 0x00, 0x00, 0x00]);
assert_eq!(&p[10..14], &1.0f32.to_le_bytes());
assert_eq!(&p[14..18], &2.0f32.to_le_bytes());
assert_eq!(&p[18..22], &3.0f32.to_le_bytes());
}

#[test]
fn vector_float16_elements() {
let mut mock = MockNetworkWriter::new(64);
let mut w = PacketWriter::new(PacketType::TabularResult, &mut mock, None, None);
let vector =
crate::datatypes::sql_vector::SqlVector::try_from_f16(vec![1.0, 2.0, 3.0]).unwrap();
let ctx = vector_ctx(8 + 3 * 2);
block_on(TdsValueSerializer::serialize_value(
&mut w,
&ColumnValues::Vector(vector),
&ctx,
))
.unwrap();
let p = payload(&w);
assert_eq!(&p[0..2], &14u16.to_le_bytes()); // 8 header + 3 * 2
// base type byte is 0x01 (float16)
assert_eq!(&p[2..10], &[0xA9, 0x01, 0x03, 0x00, 0x01, 0x00, 0x00, 0x00]);
assert_eq!(&p[10..16], &[0x00, 0x3C, 0x00, 0x40, 0x00, 0x42]);
}

#[test]
fn vector_float16_rounds_to_nearest_even() {
let mut mock = MockNetworkWriter::new(64);
let mut w = PacketWriter::new(PacketType::TabularResult, &mut mock, None, None);
// 0.1 is not representable in f16; nearest half is 0x2E66.
let vector = crate::datatypes::sql_vector::SqlVector::try_from_f16(vec![0.1]).unwrap();
let ctx = vector_ctx(8 + 2);
block_on(TdsValueSerializer::serialize_value(
&mut w,
&ColumnValues::Vector(vector),
&ctx,
))
.unwrap();
let p = payload(&w);
assert_eq!(&p[10..12], &0x2E66u16.to_le_bytes());
Comment on lines +3804 to +3814
}
}

#[cfg(test)]
Expand Down
9 changes: 9 additions & 0 deletions mssql-tds/src/message/features/vectorfeature.rs
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,15 @@ mod tests {
);
}

#[test]
fn test_deserialize_v2_acknowledged() {
let mut feature = VectorFeature::default();
feature.set_acknowledged(true);
feature.deserialize(&[2u8]).unwrap();
assert!(feature.is_acknowledged());
assert_eq!(feature.negotiated_version(), 2);
}

#[test]
fn test_deserialize_invalid_length() {
let mut feature = VectorFeature::default();
Expand Down
2 changes: 1 addition & 1 deletion mssql-tds/tests/common/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@ pub fn create_context() -> ClientContext {
.map_err(|_| std::env::VarError::NotPresent)
})
.expect("SQL_PASSWORD environment variable not set and /tmp/password could not be read");
context.database = "master".to_string();
context.database = env::var("DB_DATABASE").unwrap_or_else(|_| "master".to_string());
context.encryption_options = EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: trust_server_certificate(),
Expand Down
Loading
Loading