diff --git a/MEVD/src/SqliteVec/SqliteCollection.cs b/MEVD/src/SqliteVec/SqliteCollection.cs index 893dfc1..475368e 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; @@ -51,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; } @@ -103,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); @@ -559,53 +569,89 @@ 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; - 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) == default) + || (keyProperty.Type == typeof(long) && keyProperty.GetValue(record) == default)); - foreach (var record in recordsList) + DbCommand insertCommand; + if (isRecordKeyDatabaseGenerated) + { + insertCommandReturning ??= SqliteCommandBuilder.BuildInsertCommand( + connection, + _dataTableName, + _model, + _dataInsertProperties, + isRecordKeyDatabaseGenerated: true, + replaceIfExists: true); + insertCommand = insertCommandReturning; + } + else { - switch (keyProperty.Type) + insertCommandNonReturning ??= SqliteCommandBuilder.BuildInsertCommand( + connection, + _dataTableName, + _model, + _dataInsertProperties, + isRecordKeyDatabaseGenerated: false, + replaceIfExists: true); + insertCommand = insertCommandNonReturning; + } + + insertCommand.Transaction = transaction; + + SqliteCommandBuilder.SetInsertParameterValues( + insertCommand, _dataInsertProperties, 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 +666,22 @@ private async Task DoUpsertAsync(IEnumerable records, CancellationToken connection, _vectorTableName, _model, - recordsList, - generatedEmbeddings, - data: false); + _vectorInsertProperties, + isRecordKeyDatabaseGenerated: false); vectorInsertCommand.Transaction = transaction; - await connection.ExecuteWithErrorHandlingAsync( - _collectionMetadata, - "VectorInsert", - () => vectorInsertCommand.ExecuteNonQueryAsync(cancellationToken), - cancellationToken).ConfigureAwait(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..33e8045 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, + IReadOnlyList properties, + 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 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(" INTO ").AppendIdentifier(tableName).Append(" ("); - sql.Append("INSERT"); + var propertyIndex = 0; + foreach (var property in properties) + { + if (property is KeyPropertyModel && isRecordKeyDatabaseGenerated) + { + continue; + } - if (replaceIfExists && !isRecordKeyDatabaseGenerated) + if (propertyIndex++ > 0) { - sql.Append(" OR REPLACE"); + sql.Append(", "); } - sql.Append(" INTO ").AppendIdentifier(tableName).Append(" ("); + sql.AppendIdentifier(property.StorageName); + } + + sql.AppendLine(")"); + + sql.Append("VALUES ("); - var propertyIndex = 0; - foreach (var property in properties) + 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); + + command.Parameters.Add(new SqliteParameter { ParameterName = parameterName }); + } - sql.Append("VALUES ("); + sql.AppendLine(")"); - propertyIndex = 0; - foreach (var property in properties) - { - var value = property.GetValueAsObject(record); + if (isRecordKeyDatabaseGenerated) + { + sql.Append("RETURNING ").AppendIdentifier(keyProperty.StorageName); + } - switch (property) - { - case KeyPropertyModel { IsAutoGenerated: true }: + sql.AppendLine(";"); + + command.CommandText = sql.ToString(); + + return command; + } + + 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, IReadOnlyList properties, bool isRecordKeyDatabaseGenerated, + object record, int recordIndex = 0, + Dictionary>>? generatedEmbeddings = null) + { + for (var i = 0; i < properties.Count; i++) + { + var property = properties[i]; + 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..54ba637 100644 --- a/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs +++ b/MEVD/test/SqliteVec.UnitTests/SqliteCommandBuilderTests.cs @@ -3,7 +3,10 @@ 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; using Microsoft.Extensions.VectorData.ProviderServices; using CommunityToolkit.VectorData.SqliteVec; @@ -129,53 +132,33 @@ 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, TableName, model, - records, - generatedEmbeddings: null, - data: true, + SqliteCommandBuilder.GetDataInsertProperties(model), + 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] @@ -194,35 +177,149 @@ 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, TableName, model, - records, - generatedEmbeddings: null, - data: true, + SqliteCommandBuilder.GetDataInsertProperties(model), + 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" }; - Assert.Equal("@Name1", command.Parameters[3].ParameterName); - Assert.Equal("NameValue2", command.Parameters[3].Value); + var properties = SqliteCommandBuilder.GetDataInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, properties, isRecordKeyDatabaseGenerated: false); + + // 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 properties = SqliteCommandBuilder.GetVectorInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "VectorTable", model, properties, isRecordKeyDatabaseGenerated: 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 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); + + // 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 properties = SqliteCommandBuilder.GetDataInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, properties, isRecordKeyDatabaseGenerated: false); + + // Act + SqliteCommandBuilder.SetInsertParameterValues(command, properties, isRecordKeyDatabaseGenerated: false, record); + + // 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 properties = SqliteCommandBuilder.GetDataInsertProperties(model); + var command = SqliteCommandBuilder.BuildInsertCommand(this._connection, "TestTable", model, properties, isRecordKeyDatabaseGenerated: 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 +499,7 @@ public void Dispose() private static CollectionModel BuildModel(List properties) => new SqliteModelBuilder() .BuildDynamic(new() { Properties = properties }, defaultEmbeddingGenerator: null); + + private static byte[] FloatToBytes(params float[] values) + => MemoryMarshal.AsBytes(values.AsSpan()).ToArray(); }