From 72ff5f59f78952a3b75bfa5b9280b19f8f956ddc Mon Sep 17 00:00:00 2001 From: Ross Donald Date: Fri, 25 Sep 2026 16:44:06 +1000 Subject: [PATCH 1/3] SqliteVec: re-use a single command for upserts --- MEVD/src/SqliteVec/SqliteCollection.cs | 126 ++++++++---- MEVD/src/SqliteVec/SqliteCommandBuilder.cs | 136 +++++++------ .../SqliteCommandBuilderTests.cs | 183 +++++++++++++++--- 3 files changed, 310 insertions(+), 135 deletions(-) diff --git a/MEVD/src/SqliteVec/SqliteCollection.cs b/MEVD/src/SqliteVec/SqliteCollection.cs index 893dfc1..6a9f985 100644 --- a/MEVD/src/SqliteVec/SqliteCollection.cs +++ b/MEVD/src/SqliteVec/SqliteCollection.cs @@ -1,6 +1,7 @@ // Licensed to the .NET Foundation under one or more agreements. // The .NET Foundation licenses this file to you under the MIT license. +using System.Data.Common; using System.Diagnostics; using System.Diagnostics.CodeAnalysis; using System.Linq.Expressions; @@ -559,53 +560,90 @@ private async Task DoUpsertAsync(IEnumerable records, CancellationToken } var keyProperty = _model.KeyProperty; + var isKeyPossiblyDatabaseGenerated = keyProperty.IsAutoGenerated + && (keyProperty.Type == typeof(int) || keyProperty.Type == typeof(long)); using var connection = await GetConnectionAsync(cancellationToken).ConfigureAwait(false); using var transaction = connection.BeginTransaction(); - using var dataCommand = SqliteCommandBuilder.BuildInsertCommand( - connection, - _dataTableName, - _model, - recordsList, - generatedEmbeddings, - data: true, - replaceIfExists: true); - dataCommand.Transaction = transaction; + DbCommand? insertCommandReturning = null; + DbCommand? insertCommandNonReturning = null; + var insertProperties = SqliteCommandBuilder.GetInsertProperties(_model, data: true); - using (var reader = await connection.ExecuteWithErrorHandlingAsync( - _collectionMetadata, - "updateData", - () => dataCommand.ExecuteReaderAsync(cancellationToken), - cancellationToken).ConfigureAwait(false)) + try { - // If the key property is auto-generated, we need to read the generated keys from the database and inject them into the records - // (except for GUIDs which are generated client-side and have already been injected). - if (keyProperty is KeyPropertyModel { IsAutoGenerated: true } && keyProperty.Type != typeof(Guid)) + foreach (var record in records) { - int? keyOrdinal = null; + var isRecordKeyDatabaseGenerated = isKeyPossiblyDatabaseGenerated && ( + (keyProperty.Type == typeof(int) && keyProperty.GetValue(record) is var i && i == 0) + || (keyProperty.Type == typeof(long) && keyProperty.GetValue(record) is var l && l == 0L)); - foreach (var record in recordsList) + DbCommand insertCommand; + if (isRecordKeyDatabaseGenerated) { - switch (keyProperty.Type) + insertCommandReturning ??= SqliteCommandBuilder.BuildInsertCommand( + connection, + _dataTableName, + _model, + data: true, + isRecordKeyDatabaseGenerated: true, + replaceIfExists: true); + insertCommand = insertCommandReturning; + } + else + { + insertCommandNonReturning ??= SqliteCommandBuilder.BuildInsertCommand( + connection, + _dataTableName, + _model, + data: true, + isRecordKeyDatabaseGenerated: false, + replaceIfExists: true); + insertCommand = insertCommandNonReturning; + } + + insertCommand.Transaction = transaction; + + SqliteCommandBuilder.SetInsertParameterValues( + insertCommand, insertProperties, isRecordKeyDatabaseGenerated, record); + + if (isRecordKeyDatabaseGenerated) + { + var generatedKey = await connection.ExecuteWithErrorHandlingAsync( + _collectionMetadata, + "updateData", + () => insertCommand.ExecuteScalarAsync(cancellationToken), + cancellationToken).ConfigureAwait(false); + + // Write the database-generated key back into the record so the vector table insert below references the right row. + if (generatedKey is not null) { - case var t when t == typeof(int) && keyProperty.GetValue(record) == 0: - keyOrdinal ??= reader.GetOrdinal(keyProperty.StorageName); - await reader.ReadAsync(cancellationToken).ConfigureAwait(false); - keyProperty.SetValue(record, reader.GetFieldValue(keyOrdinal.Value)); - await reader.NextResultAsync(cancellationToken).ConfigureAwait(false); - continue; - case var t when t == typeof(long) && keyProperty.GetValue(record) == 0L: - keyOrdinal ??= reader.GetOrdinal(keyProperty.StorageName); - await reader.ReadAsync(cancellationToken).ConfigureAwait(false); - keyProperty.SetValue(record, reader.GetFieldValue(keyOrdinal.Value)); - await reader.NextResultAsync(cancellationToken).ConfigureAwait(false); - continue; + if (keyProperty.Type == typeof(int)) + { + keyProperty.SetValue(record, Convert.ToInt32(generatedKey)); + } + else + { + keyProperty.SetValue(record, Convert.ToInt64(generatedKey)); + } } } + else + { + await connection.ExecuteWithErrorHandlingAsync( + _collectionMetadata, + "updateData", + () => insertCommand.ExecuteNonQueryAsync(cancellationToken), + cancellationToken).ConfigureAwait(false); + } } } + finally + { + insertCommandReturning?.Dispose(); + insertCommandNonReturning?.Dispose(); + } // We've inserted the main data records, now insert the records into the vector virtual table as well. if (_vectorPropertiesExist) @@ -620,16 +658,24 @@ private async Task DoUpsertAsync(IEnumerable records, CancellationToken connection, _vectorTableName, _model, - recordsList, - generatedEmbeddings, - data: false); + data: false, + isRecordKeyDatabaseGenerated: false); vectorInsertCommand.Transaction = transaction; - await connection.ExecuteWithErrorHandlingAsync( - _collectionMetadata, - "VectorInsert", - () => vectorInsertCommand.ExecuteNonQueryAsync(cancellationToken), - cancellationToken).ConfigureAwait(false); + var vectorInsertProperties = SqliteCommandBuilder.GetInsertProperties(_model, data: false); + + for (int i = 0; i < recordsList.Count; i++) + { + SqliteCommandBuilder.SetInsertParameterValues( + vectorInsertCommand, vectorInsertProperties, isRecordKeyDatabaseGenerated: false, + recordsList[i], i, generatedEmbeddings); + + await connection.ExecuteWithErrorHandlingAsync( + _collectionMetadata, + "VectorInsert", + () => vectorInsertCommand.ExecuteNonQueryAsync(cancellationToken), + cancellationToken).ConfigureAwait(false); + } } transaction.Commit(); diff --git a/MEVD/src/SqliteVec/SqliteCommandBuilder.cs b/MEVD/src/SqliteVec/SqliteCommandBuilder.cs index 5e0817e..1316577 100644 --- a/MEVD/src/SqliteVec/SqliteCommandBuilder.cs +++ b/MEVD/src/SqliteVec/SqliteCommandBuilder.cs @@ -115,64 +115,97 @@ public static DbCommand BuildInsertCommand( SqliteConnection connection, string tableName, CollectionModel model, - IReadOnlyList records, - Dictionary>>? generatedEmbeddings, bool data, + bool isRecordKeyDatabaseGenerated, bool replaceIfExists = false) { var sql = new StringBuilder(); var command = connection.CreateCommand(); - var recordIndex = 0; - - var properties = model.KeyProperties.Concat(data ? model.DataProperties : (IEnumerable)model.VectorProperties).ToList(); + var properties = GetInsertProperties(model, data); var keyProperty = model.KeyProperty; - var isKeyPossiblyDatabaseGenerated = keyProperty.IsAutoGenerated && (keyProperty.Type == typeof(int) || keyProperty.Type == typeof(long)); - foreach (var record in records) + sql.Append("INSERT"); + + if (replaceIfExists && !isRecordKeyDatabaseGenerated) { - var isRecordKeyDatabaseGenerated = isKeyPossiblyDatabaseGenerated - ? (keyProperty.Type == typeof(int) && keyProperty.GetValue(record) is var i && i == 0) - || (keyProperty.Type == typeof(long) && keyProperty.GetValue(record) is var l && l == 0L) - : false; + sql.Append(" OR REPLACE"); + } - sql.Append("INSERT"); + sql.Append(" INTO ").AppendIdentifier(tableName).Append(" ("); - if (replaceIfExists && !isRecordKeyDatabaseGenerated) + var propertyIndex = 0; + foreach (var property in properties) + { + if (property is KeyPropertyModel && isRecordKeyDatabaseGenerated) { - sql.Append(" OR REPLACE"); + continue; } - sql.Append(" INTO ").AppendIdentifier(tableName).Append(" ("); + if (propertyIndex++ > 0) + { + sql.Append(", "); + } - var propertyIndex = 0; - foreach (var property in properties) + sql.AppendIdentifier(property.StorageName); + } + + sql.AppendLine(")"); + + sql.Append("VALUES ("); + + propertyIndex = 0; + foreach (var property in properties) + { + if (property is KeyPropertyModel && isRecordKeyDatabaseGenerated) { - if (property is KeyPropertyModel && isRecordKeyDatabaseGenerated) - { - continue; - } + continue; + } - if (propertyIndex++ > 0) - { - sql.Append(", "); - } + var parameterName = GetParameterName(property.StorageName); - sql.AppendIdentifier(property.StorageName); + if (propertyIndex++ > 0) + { + sql.Append(", "); } - sql.AppendLine(")"); + sql.Append(parameterName); - sql.Append("VALUES ("); + command.Parameters.Add(new SqliteParameter { ParameterName = parameterName }); + } - propertyIndex = 0; - foreach (var property in properties) - { - var value = property.GetValueAsObject(record); + sql.AppendLine(")"); - switch (property) - { - case KeyPropertyModel { IsAutoGenerated: true }: + if (isRecordKeyDatabaseGenerated) + { + sql.Append("RETURNING ").AppendIdentifier(keyProperty.StorageName); + } + + sql.AppendLine(";"); + + command.CommandText = sql.ToString(); + + return command; + } + + public static List GetInsertProperties(CollectionModel model, bool data) + { + return model.KeyProperties + .Concat(data ? model.DataProperties : (IEnumerable)model.VectorProperties).ToList(); + } + + public static void SetInsertParameterValues( + DbCommand command, List properties, bool isRecordKeyDatabaseGenerated, + object record, int recordIndex = 0, + Dictionary>>? generatedEmbeddings = null) + { + foreach (var property in properties) + { + var value = property.GetValueAsObject(record); + + switch (property) + { + case KeyPropertyModel { IsAutoGenerated: true } keyProperty: { switch (value) { @@ -206,7 +239,7 @@ public static DbCommand BuildInsertCommand( break; } - case VectorPropertyModel vectorProperty: + case VectorPropertyModel vectorProperty: { if (generatedEmbeddings?[vectorProperty] is IReadOnlyList ge) { @@ -224,35 +257,10 @@ public static DbCommand BuildInsertCommand( }; break; } - } - - var parameterName = GetParameterName(property.StorageName, recordIndex); - - if (propertyIndex++ > 0) - { - sql.Append(", "); - } - - sql.Append(parameterName); - - command.Parameters.Add(new SqliteParameter(parameterName, value ?? DBNull.Value)); - } - - sql.AppendLine(")"); - - if (isRecordKeyDatabaseGenerated) - { - sql.Append("RETURNING \"").Append(keyProperty.StorageName).Append('"'); } - sql.AppendLine(";"); - - recordIndex++; + command.Parameters[GetParameterName(property.StorageName)].Value = value ?? DBNull.Value; } - - command.CommandText = sql.ToString(); - - return command; } public static DbCommand BuildSelectDataCommand( @@ -574,7 +582,9 @@ private static (DbCommand Command, string WhereClause) GetCommandWithWhereClause } private static string GetParameterName(string propertyName, int index) - => $"@{propertyName}{index}"; + => GetParameterName(propertyName) + index; + + private static string GetParameterName(string propertyName) => $"@{propertyName}"; #endregion } diff --git a/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs b/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs index 9f209ca..298e325 100644 --- a/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs +++ b/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs @@ -3,7 +3,9 @@ using System; using System.Collections.Generic; +using System.Linq; using Microsoft.Data.Sqlite; +using Microsoft.Extensions.AI; using Microsoft.Extensions.VectorData; using Microsoft.Extensions.VectorData.ProviderServices; using CommunityToolkit.VectorData.SqliteVec; @@ -140,42 +142,28 @@ public void ItBuildsInsertCommand_without_autogenerated_key(bool replaceIfExists this._connection, TableName, model, - records, - generatedEmbeddings: null, data: true, + isRecordKeyDatabaseGenerated: false, replaceIfExists: replaceIfExists); // Assert Assert.Equal(replaceIfExists, command.CommandText.Contains("OR REPLACE")); Assert.Contains($"INTO \"{TableName}\" (\"Id\", \"Name\", \"Age\", \"Address\")", command.CommandText); - Assert.Contains("VALUES (@Id0, @Name0, @Age0, @Address0)", command.CommandText); - Assert.Contains("VALUES (@Id1, @Name1, @Age1, @Address1)", command.CommandText); + Assert.Contains("VALUES (@Id, @Name, @Age, @Address)", command.CommandText); Assert.DoesNotContain("RETURNING", command.CommandText); - Assert.Equal("@Id0", command.Parameters[0].ParameterName); - Assert.Equal(1, command.Parameters[0].Value); - - Assert.Equal("@Name0", command.Parameters[1].ParameterName); - Assert.Equal("NameValue1", command.Parameters[1].Value); - - Assert.Equal("@Age0", command.Parameters[2].ParameterName); - Assert.Equal("AgeValue1", command.Parameters[2].Value); - - Assert.Equal("@Address0", command.Parameters[3].ParameterName); - Assert.Equal("AddressValue1", command.Parameters[3].Value); - - Assert.Equal("@Id1", command.Parameters[4].ParameterName); - Assert.Equal(2, command.Parameters[4].Value); + Assert.Equal("@Id", command.Parameters[0].ParameterName); + Assert.Null(command.Parameters[0].Value); - Assert.Equal("@Name1", command.Parameters[5].ParameterName); - Assert.Equal("NameValue2", command.Parameters[5].Value); + Assert.Equal("@Name", command.Parameters[1].ParameterName); + Assert.Null(command.Parameters[1].Value); - Assert.Equal("@Age1", command.Parameters[6].ParameterName); - Assert.Equal("AgeValue2", command.Parameters[6].Value); + Assert.Equal("@Age", command.Parameters[2].ParameterName); + Assert.Null(command.Parameters[2].Value); - Assert.Equal("@Address1", command.Parameters[7].ParameterName); - Assert.Equal("AddressValue2", command.Parameters[7].Value); + Assert.Equal("@Address", command.Parameters[3].ParameterName); + Assert.Null(command.Parameters[3].Value); } [Theory] @@ -205,24 +193,144 @@ public void ItBuildsInsertCommand_with_autogenerated_key(bool replaceIfExists) this._connection, TableName, model, - records, - generatedEmbeddings: null, data: true, + isRecordKeyDatabaseGenerated: true, replaceIfExists: replaceIfExists); // Assert Assert.DoesNotContain("OR REPLACE", command.CommandText); Assert.Contains($"INTO \"{TableName}\" (\"Name\", \"Age\", \"Address\")", command.CommandText); - Assert.Contains("VALUES (@Name0, @Age0, @Address0)", command.CommandText); - Assert.Contains("VALUES (@Name1, @Age1, @Address1)", command.CommandText); + Assert.Contains("VALUES (@Name, @Age, @Address)", command.CommandText); Assert.Contains("RETURNING \"Id\"", command.CommandText); - Assert.Equal("@Name0", command.Parameters[0].ParameterName); - Assert.Equal("NameValue1", command.Parameters[0].Value); + Assert.Equal("@Name", command.Parameters[0].ParameterName); + Assert.Null(command.Parameters[0].Value); + } + + [Fact] + public void ItSetsInsertParameterValues_for_data_properties() + { + // Arrange + var model = BuildModel( + [ + new VectorStoreKeyProperty("Id", typeof(string)), + new VectorStoreDataProperty("Name", typeof(string)), + ]); + + var record = new Dictionary { ["Id"] = "KeyValue", ["Name"] = "NameValue" }; + + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, data: true, isRecordKeyDatabaseGenerated: false); + var properties = SqliteCommandBuilder.GetInsertProperties(model, data: true); + + // Act + SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record); + + // Assert + Assert.Equal("KeyValue", command.Parameters["@Id"].Value); + Assert.Equal("NameValue", command.Parameters["@Name"].Value); + } + + [Fact] + public void ItSetsInsertParameterValues_converts_vectors_to_blobs() + { + // Arrange + var model = BuildModel( + [ + new VectorStoreKeyProperty("Id", typeof(string)), + new VectorStoreVectorProperty("Embedding", typeof(ReadOnlyMemory), 2), + ]); + + float[] embeddings = [1f, 2f]; + var record = new Dictionary { ["Id"] = "KeyValue", ["Embedding"] = new ReadOnlyMemory(embeddings) }; + + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "VectorTable", model, data: false, isRecordKeyDatabaseGenerated: false); + var properties = SqliteCommandBuilder.GetInsertProperties(model, data: false); + + // Act + SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record); + + // Assert + Assert.Equal(FloatToBytes(embeddings), (byte[])command.Parameters["@Embedding"].Value!); + } + + [Fact] + public void ItSetsInsertParameterValues_uses_generated_embeddings_by_record_index() + { + // Arrange + var model = BuildModel( + [ + new VectorStoreKeyProperty("Id", typeof(string)), + new VectorStoreVectorProperty("Embedding", typeof(ReadOnlyMemory), 2), + ]); + + var record = new Dictionary { ["Id"] = "Key2", ["Embedding"] = new ReadOnlyMemory([1f, 1f]) }; + + var vectorProperty = model.VectorProperties.Single(v => v.StorageName == "Embedding"); + + var embeddings = new ReadOnlyMemory([7f, 7f]); + var generatedEmbeddings = new Dictionary>> + { + [vectorProperty] = [new(new float[] { 9f, 9f }), new(embeddings)], + }; + + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "VectorTable", model, data: false, isRecordKeyDatabaseGenerated: false); + var properties = SqliteCommandBuilder.GetInsertProperties(model, data: false); + + // Act + SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record, recordIndex: 1, generatedEmbeddings); + + // Assert - The generated embedding at the record's batch index should be used instead of the value in the record's Embedding property. + Assert.Equal(SqlitePropertyMapping.MapVectorForStorageModel(embeddings), (byte[])command.Parameters["@Embedding"].Value!); + } + + [Fact] + public void ItSetsInsertParameterValues_generates_missing_guid_keys() + { + // Arrange + var model = BuildModel( + [ + new VectorStoreKeyProperty("Id", typeof(Guid)), + new VectorStoreDataProperty("Name", typeof(string)), + ]); + + var record = new Dictionary { ["Id"] = Guid.Empty, ["Name"] = "NameValue" }; + + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, data: true, isRecordKeyDatabaseGenerated: false); + var properties = SqliteCommandBuilder.GetInsertProperties(model, data: true); + + // Act + SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record); - Assert.Equal("@Name1", command.Parameters[3].ParameterName); - Assert.Equal("NameValue2", command.Parameters[3].Value); + // Assert + var generatedKey = Assert.IsType(command.Parameters["@Id"].Value); + Assert.NotEqual(Guid.Empty, generatedKey); + Assert.Equal(generatedKey, record["Id"]); + } + + [Fact] + public void ItSetsInsertParameterValues_skips_database_generated_keys() + { + // Arrange + var model = BuildModel( + [ + new VectorStoreKeyProperty("Id", typeof(int)), + new VectorStoreDataProperty("Name", typeof(string)), + ]); + + var record = new Dictionary { ["Id"] = 0, ["Name"] = "NameValue" }; + + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, data: true, isRecordKeyDatabaseGenerated: true); + var properties = SqliteCommandBuilder.GetInsertProperties(model, data: true); + + // Act + SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: true, record); + + // Assert + // The key column is omitted from the statement, so only the data parameters are bound. + Assert.Single(command.Parameters); + Assert.Equal("NameValue", command.Parameters["@Name"].Value); + Assert.Equal(0, record["Id"]); } [Theory] @@ -402,4 +510,15 @@ public void Dispose() private static CollectionModel BuildModel(List properties) => new SqliteModelBuilder() .BuildDynamic(new() { Properties = properties }, defaultEmbeddingGenerator: null); + + private static byte[] FloatToBytes(params float[] values) + { + var bytes = new byte[values.Length * sizeof(float)]; + for (var i = 0; i < values.Length; i++) + { + BitConverter.GetBytes(values[i]).CopyTo(bytes, i * sizeof(float)); + } + + return bytes; + } } From dcea24783f3c2b360acd0bae68a0016ac7257fd6 Mon Sep 17 00:00:00 2001 From: Ross Donald Date: Fri, 25 Sep 2026 19:05:54 +1000 Subject: [PATCH 2/3] Remove redundant record definitions from tests --- .../SqliteVec.UnitTests/SqliteCommandBuilderTests.cs | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs b/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs index 298e325..01dba8f 100644 --- a/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs +++ b/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs @@ -131,12 +131,6 @@ public void ItBuildsInsertCommand_without_autogenerated_key(bool replaceIfExists new VectorStoreDataProperty("Address", typeof(string)), ]); - var records = new List> - { - new() { ["Id"] = 1, ["Name"] = "NameValue1", ["Age"] = "AgeValue1", ["Address"] = "AddressValue1" }, - new() { ["Id"] = 2, ["Name"] = "NameValue2", ["Age"] = "AgeValue2", ["Address"] = "AddressValue2" }, - }; - // Act var command = SqliteCommandBuilder.BuildInsertCommand( this._connection, @@ -182,12 +176,6 @@ public void ItBuildsInsertCommand_with_autogenerated_key(bool replaceIfExists) new VectorStoreDataProperty("Address", typeof(string)), ]); - var records = new List> - { - new() { ["Id"] = default(int), ["Name"] = "NameValue1", ["Age"] = "AgeValue1", ["Address"] = "AddressValue1" }, - new() { ["Id"] = default(int), ["Name"] = "NameValue2", ["Age"] = "AgeValue2", ["Address"] = "AddressValue2" }, - }; - // Act var command = SqliteCommandBuilder.BuildInsertCommand( this._connection, From 253a756099a8060e7f8191a84c41a860dcb55dd5 Mon Sep 17 00:00:00 2001 From: Ross Donald Date: Fri, 2 Oct 2026 17:05:49 +1000 Subject: [PATCH 3/3] Cache insert property lists, use BCL AsBytes, simplify key value check. --- MEVD/src/SqliteVec/SqliteCollection.cs | 26 ++++++++------ MEVD/src/SqliteVec/SqliteCommandBuilder.cs | 18 +++++----- .../SqliteCommandBuilderTests.cs | 35 ++++++++----------- 3 files changed, 39 insertions(+), 40 deletions(-) diff --git a/MEVD/src/SqliteVec/SqliteCollection.cs b/MEVD/src/SqliteVec/SqliteCollection.cs index 6a9f985..475368e 100644 --- a/MEVD/src/SqliteVec/SqliteCollection.cs +++ b/MEVD/src/SqliteVec/SqliteCollection.cs @@ -52,6 +52,12 @@ public class SqliteCollection : VectorStoreCollectionTable name in SQLite for vector properties. private readonly string _vectorTableName; + /// Properties to bind on inserts into the data table. + private readonly IReadOnlyList _dataInsertProperties; + + /// Properties to bind on inserts into the vector table. + private readonly IReadOnlyList _vectorInsertProperties; + /// public override string Name { get; } @@ -104,6 +110,9 @@ internal SqliteCollection(string connectionString, string name, Func 0; + _dataInsertProperties = SqliteCommandBuilder.GetDataInsertProperties(_model); + _vectorInsertProperties = SqliteCommandBuilder.GetVectorInsertProperties(_model); + // Populate some collections of properties _keyStorageName = _model.KeyProperty.StorageName; _mapper = new SqliteMapper(_model); @@ -569,15 +578,14 @@ private async Task DoUpsertAsync(IEnumerable records, CancellationToken DbCommand? insertCommandReturning = null; DbCommand? insertCommandNonReturning = null; - var insertProperties = SqliteCommandBuilder.GetInsertProperties(_model, data: true); try { foreach (var record in records) { var isRecordKeyDatabaseGenerated = isKeyPossiblyDatabaseGenerated && ( - (keyProperty.Type == typeof(int) && keyProperty.GetValue(record) is var i && i == 0) - || (keyProperty.Type == typeof(long) && keyProperty.GetValue(record) is var l && l == 0L)); + (keyProperty.Type == typeof(int) && keyProperty.GetValue(record) == default) + || (keyProperty.Type == typeof(long) && keyProperty.GetValue(record) == default)); DbCommand insertCommand; if (isRecordKeyDatabaseGenerated) @@ -586,7 +594,7 @@ private async Task DoUpsertAsync(IEnumerable records, CancellationToken connection, _dataTableName, _model, - data: true, + _dataInsertProperties, isRecordKeyDatabaseGenerated: true, replaceIfExists: true); insertCommand = insertCommandReturning; @@ -597,7 +605,7 @@ private async Task DoUpsertAsync(IEnumerable records, CancellationToken connection, _dataTableName, _model, - data: true, + _dataInsertProperties, isRecordKeyDatabaseGenerated: false, replaceIfExists: true); insertCommand = insertCommandNonReturning; @@ -606,7 +614,7 @@ private async Task DoUpsertAsync(IEnumerable records, CancellationToken insertCommand.Transaction = transaction; SqliteCommandBuilder.SetInsertParameterValues( - insertCommand, insertProperties, isRecordKeyDatabaseGenerated, record); + insertCommand, _dataInsertProperties, isRecordKeyDatabaseGenerated, record); if (isRecordKeyDatabaseGenerated) { @@ -658,16 +666,14 @@ await connection.ExecuteWithErrorHandlingAsync( connection, _vectorTableName, _model, - data: false, + _vectorInsertProperties, isRecordKeyDatabaseGenerated: false); vectorInsertCommand.Transaction = transaction; - var vectorInsertProperties = SqliteCommandBuilder.GetInsertProperties(_model, data: false); - for (int i = 0; i < recordsList.Count; i++) { SqliteCommandBuilder.SetInsertParameterValues( - vectorInsertCommand, vectorInsertProperties, isRecordKeyDatabaseGenerated: false, + vectorInsertCommand, _vectorInsertProperties, isRecordKeyDatabaseGenerated: false, recordsList[i], i, generatedEmbeddings); await connection.ExecuteWithErrorHandlingAsync( diff --git a/MEVD/src/SqliteVec/SqliteCommandBuilder.cs b/MEVD/src/SqliteVec/SqliteCommandBuilder.cs index 1316577..33e8045 100644 --- a/MEVD/src/SqliteVec/SqliteCommandBuilder.cs +++ b/MEVD/src/SqliteVec/SqliteCommandBuilder.cs @@ -115,14 +115,13 @@ public static DbCommand BuildInsertCommand( SqliteConnection connection, string tableName, CollectionModel model, - bool data, + IReadOnlyList properties, bool isRecordKeyDatabaseGenerated, bool replaceIfExists = false) { var sql = new StringBuilder(); var command = connection.CreateCommand(); - var properties = GetInsertProperties(model, data); var keyProperty = model.KeyProperty; sql.Append("INSERT"); @@ -188,19 +187,20 @@ public static DbCommand BuildInsertCommand( return command; } - public static List GetInsertProperties(CollectionModel model, bool data) - { - return model.KeyProperties - .Concat(data ? model.DataProperties : (IEnumerable)model.VectorProperties).ToList(); - } + public static IReadOnlyList GetDataInsertProperties(CollectionModel model) + => [.. model.KeyProperties, .. model.DataProperties]; + + public static IReadOnlyList GetVectorInsertProperties(CollectionModel model) + => [.. model.KeyProperties, .. model.VectorProperties]; public static void SetInsertParameterValues( - DbCommand command, List properties, bool isRecordKeyDatabaseGenerated, + DbCommand command, IReadOnlyList properties, bool isRecordKeyDatabaseGenerated, object record, int recordIndex = 0, Dictionary>>? generatedEmbeddings = null) { - foreach (var property in properties) + for (var i = 0; i < properties.Count; i++) { + var property = properties[i]; var value = property.GetValueAsObject(record); switch (property) diff --git a/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs b/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs index 01dba8f..54ba637 100644 --- a/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs +++ b/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs @@ -4,6 +4,7 @@ using System; using System.Collections.Generic; using System.Linq; +using System.Runtime.InteropServices; using Microsoft.Data.Sqlite; using Microsoft.Extensions.AI; using Microsoft.Extensions.VectorData; @@ -136,7 +137,7 @@ public void ItBuildsInsertCommand_without_autogenerated_key(bool replaceIfExists this._connection, TableName, model, - data: true, + SqliteCommandBuilder.GetDataInsertProperties(model), isRecordKeyDatabaseGenerated: false, replaceIfExists: replaceIfExists); @@ -181,7 +182,7 @@ public void ItBuildsInsertCommand_with_autogenerated_key(bool replaceIfExists) this._connection, TableName, model, - data: true, + SqliteCommandBuilder.GetDataInsertProperties(model), isRecordKeyDatabaseGenerated: true, replaceIfExists: replaceIfExists); @@ -208,8 +209,8 @@ public void ItSetsInsertParameterValues_for_data_properties() var record = new Dictionary { ["Id"] = "KeyValue", ["Name"] = "NameValue" }; - var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, data: true, isRecordKeyDatabaseGenerated: false); - var properties = SqliteCommandBuilder.GetInsertProperties(model, data: true); + var properties = SqliteCommandBuilder.GetDataInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, properties, isRecordKeyDatabaseGenerated: false); // Act SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record); @@ -232,8 +233,8 @@ public void ItSetsInsertParameterValues_converts_vectors_to_blobs() float[] embeddings = [1f, 2f]; var record = new Dictionary { ["Id"] = "KeyValue", ["Embedding"] = new ReadOnlyMemory(embeddings) }; - var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "VectorTable", model, data: false, isRecordKeyDatabaseGenerated: false); - var properties = SqliteCommandBuilder.GetInsertProperties(model, data: false); + var properties = SqliteCommandBuilder.GetVectorInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "VectorTable", model, properties, isRecordKeyDatabaseGenerated: false); // Act SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record); @@ -262,8 +263,8 @@ public void ItSetsInsertParameterValues_uses_generated_embeddings_by_record_inde [vectorProperty] = [new(new float[] { 9f, 9f }), new(embeddings)], }; - var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "VectorTable", model, data: false, isRecordKeyDatabaseGenerated: false); - var properties = SqliteCommandBuilder.GetInsertProperties(model, data: false); + var properties = SqliteCommandBuilder.GetVectorInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "VectorTable", model, properties, isRecordKeyDatabaseGenerated: false); // Act SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record, recordIndex: 1, generatedEmbeddings); @@ -284,8 +285,8 @@ public void ItSetsInsertParameterValues_generates_missing_guid_keys() var record = new Dictionary { ["Id"] = Guid.Empty, ["Name"] = "NameValue" }; - var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, data: true, isRecordKeyDatabaseGenerated: false); - var properties = SqliteCommandBuilder.GetInsertProperties(model, data: true); + var properties = SqliteCommandBuilder.GetDataInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, properties, isRecordKeyDatabaseGenerated: false); // Act SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record); @@ -308,8 +309,8 @@ public void ItSetsInsertParameterValues_skips_database_generated_keys() var record = new Dictionary { ["Id"] = 0, ["Name"] = "NameValue" }; - var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, data: true, isRecordKeyDatabaseGenerated: true); - var properties = SqliteCommandBuilder.GetInsertProperties(model, data: true); + var properties = SqliteCommandBuilder.GetDataInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, properties, isRecordKeyDatabaseGenerated: true); // Act SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: true, record); @@ -500,13 +501,5 @@ private static CollectionModel BuildModel(List properties) .BuildDynamic(new() { Properties = properties }, defaultEmbeddingGenerator: null); private static byte[] FloatToBytes(params float[] values) - { - var bytes = new byte[values.Length * sizeof(float)]; - for (var i = 0; i < values.Length; i++) - { - BitConverter.GetBytes(values[i]).CopyTo(bytes, i * sizeof(float)); - } - - return bytes; - } + => MemoryMarshal.AsBytes(values.AsSpan()).ToArray(); }