Skip to content
Merged
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
132 changes: 92 additions & 40 deletions MEVD/src/SqliteVec/SqliteCollection.cs
Original file line number Diff line number Diff line change
@@ -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;
Expand Down Expand Up @@ -51,6 +52,12 @@ public class SqliteCollection<TKey, TRecord> : VectorStoreCollection<TKey, TReco
/// <summary>Table name in SQLite for vector properties.</summary>
private readonly string _vectorTableName;

/// <summary>Properties to bind on inserts into the data table.</summary>
private readonly IReadOnlyList<PropertyModel> _dataInsertProperties;

/// <summary>Properties to bind on inserts into the vector table.</summary>
private readonly IReadOnlyList<PropertyModel> _vectorInsertProperties;

/// <inheritdoc />
public override string Name { get; }

Expand Down Expand Up @@ -103,6 +110,9 @@ internal SqliteCollection(string connectionString, string name, Func<SqliteColle

_vectorPropertiesExist = _model.VectorProperties.Count > 0;

_dataInsertProperties = SqliteCommandBuilder.GetDataInsertProperties(_model);
_vectorInsertProperties = SqliteCommandBuilder.GetVectorInsertProperties(_model);

// Populate some collections of properties
_keyStorageName = _model.KeyProperty.StorageName;
_mapper = new SqliteMapper<TRecord>(_model);
Expand Down Expand Up @@ -559,53 +569,89 @@ private async Task DoUpsertAsync(IEnumerable<TRecord> 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<int>(record) == default)
|| (keyProperty.Type == typeof(long) && keyProperty.GetValue<long>(record) == default));

foreach (var record in recordsList)
DbCommand insertCommand;
if (isRecordKeyDatabaseGenerated)
{
insertCommandReturning ??= SqliteCommandBuilder.BuildInsertCommand(
connection,
_dataTableName,
_model,
_dataInsertProperties,
isRecordKeyDatabaseGenerated: true,
replaceIfExists: true);
Comment thread
adamsitnik marked this conversation as resolved.
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<int>(record) == 0:
keyOrdinal ??= reader.GetOrdinal(keyProperty.StorageName);
await reader.ReadAsync(cancellationToken).ConfigureAwait(false);
keyProperty.SetValue<int>(record, reader.GetFieldValue<int>(keyOrdinal.Value));
await reader.NextResultAsync(cancellationToken).ConfigureAwait(false);
continue;
case var t when t == typeof(long) && keyProperty.GetValue<long>(record) == 0L:
keyOrdinal ??= reader.GetOrdinal(keyProperty.StorageName);
await reader.ReadAsync(cancellationToken).ConfigureAwait(false);
keyProperty.SetValue<long>(record, reader.GetFieldValue<long>(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)
Expand All @@ -620,16 +666,22 @@ private async Task DoUpsertAsync(IEnumerable<TRecord> 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();
Expand Down
138 changes: 74 additions & 64 deletions MEVD/src/SqliteVec/SqliteCommandBuilder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -115,64 +115,97 @@ public static DbCommand BuildInsertCommand(
SqliteConnection connection,
string tableName,
CollectionModel model,
IReadOnlyList<object> records,
Dictionary<VectorPropertyModel, IReadOnlyList<Embedding<float>>>? generatedEmbeddings,
bool data,
IReadOnlyList<PropertyModel> 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<PropertyModel>)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<int>(record) is var i && i == 0)
|| (keyProperty.Type == typeof(long) && keyProperty.GetValue<long>(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<PropertyModel> GetDataInsertProperties(CollectionModel model)
=> [.. model.KeyProperties, .. model.DataProperties];

public static IReadOnlyList<PropertyModel> GetVectorInsertProperties(CollectionModel model)
=> [.. model.KeyProperties, .. model.VectorProperties];

public static void SetInsertParameterValues(
DbCommand command, IReadOnlyList<PropertyModel> properties, bool isRecordKeyDatabaseGenerated,
object record, int recordIndex = 0,
Dictionary<VectorPropertyModel, IReadOnlyList<Embedding<float>>>? 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)
{
Expand Down Expand Up @@ -206,7 +239,7 @@ public static DbCommand BuildInsertCommand(
break;
}

case VectorPropertyModel vectorProperty:
case VectorPropertyModel vectorProperty:
{
if (generatedEmbeddings?[vectorProperty] is IReadOnlyList<Embedding> ge)
{
Expand All @@ -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<TRecord>(
Expand Down Expand Up @@ -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
}
Loading
Loading