diff --git a/MEVD/src/AzureDocumentDB/AzureDocumentDB.csproj b/MEVD/src/AzureDocumentDB/AzureDocumentDB.csproj index cbad69e..0fc87c9 100644 --- a/MEVD/src/AzureDocumentDB/AzureDocumentDB.csproj +++ b/MEVD/src/AzureDocumentDB/AzureDocumentDB.csproj @@ -1,7 +1,7 @@  - 1.0.0 + 1.0.1 CommunityToolkit.VectorData.AzureDocumentDB $(AssemblyName) net10.0;net8.0;netstandard2.1;net472 diff --git a/MEVD/src/AzureDocumentDB/DocumentDBCollection.cs b/MEVD/src/AzureDocumentDB/DocumentDBCollection.cs index 797a8da..88902da 100644 --- a/MEVD/src/AzureDocumentDB/DocumentDBCollection.cs +++ b/MEVD/src/AzureDocumentDB/DocumentDBCollection.cs @@ -406,7 +406,7 @@ _ when vectorProperty.EmbeddingGenerationDispatcher is not null // Add score threshold filter as a $match stage if specified if (options.ScoreThreshold.HasValue) { - pipeline.Add(DocumentDBCollectionSearchMapping.GetScoreThresholdMatchQuery(ScorePropertyName, options.ScoreThreshold.Value)); + pipeline.Add(DocumentDBCollectionSearchMapping.GetScoreThresholdMatchQuery(ScorePropertyName, options.ScoreThreshold.Value, vectorProperty.DistanceFunction)); } const string OperationName = "Aggregate"; diff --git a/MEVD/src/AzureDocumentDB/DocumentDBCollectionSearchMapping.cs b/MEVD/src/AzureDocumentDB/DocumentDBCollectionSearchMapping.cs index 0e76f7a..b343d4f 100644 --- a/MEVD/src/AzureDocumentDB/DocumentDBCollectionSearchMapping.cs +++ b/MEVD/src/AzureDocumentDB/DocumentDBCollectionSearchMapping.cs @@ -98,17 +98,27 @@ public static BsonDocument GetProjectionQuery(string scorePropertyName, string d /// Returns a $match stage to filter results by score threshold. /// - /// Azure DocumentDB returns a similarity score where higher values mean more similar, - /// so we filter with $gte to keep results at or above the threshold. + /// For COS and IP Azure DocumentDB returns a similarity score where higher values mean more similar, + /// so we filter with $gte. For L2 it returns a distance where lower values mean more similar, + /// so we filter with $lte. /// - public static BsonDocument GetScoreThresholdMatchQuery(string scorePropertyName, double scoreThreshold) - => new() + public static BsonDocument GetScoreThresholdMatchQuery(string scorePropertyName, double scoreThreshold, string? distanceFunction) + { + var comparisonOperator = GetVectorPropertyDistanceFunction(distanceFunction) switch + { + DistanceFunction.CosineDistance or DistanceFunction.DotProductSimilarity => "$gte", + DistanceFunction.EuclideanDistance => "$lte", + _ => throw new NotSupportedException($"Score threshold is not supported for distance function '{distanceFunction}'.") + }; + + return new() { { "$match", new BsonDocument { - { scorePropertyName, new BsonDocument { { "$gte", scoreThreshold } } } + { scorePropertyName, new BsonDocument { { comparisonOperator, scoreThreshold } } } } } }; + } } diff --git a/MEVD/test/AzureDocumentDB.ConformanceTests/DocumentDBDistanceFunctionTests.cs b/MEVD/test/AzureDocumentDB.ConformanceTests/DocumentDBDistanceFunctionTests.cs index ac8d25f..8cd94cd 100644 --- a/MEVD/test/AzureDocumentDB.ConformanceTests/DocumentDBDistanceFunctionTests.cs +++ b/MEVD/test/AzureDocumentDB.ConformanceTests/DocumentDBDistanceFunctionTests.cs @@ -2,7 +2,6 @@ // The .NET Foundation licenses this file to you under the MIT license. using AzureDocumentDB.ConformanceTests.Support; -using Microsoft.Extensions.VectorData; using VectorData.ConformanceTests; using VectorData.ConformanceTests.Support; using Xunit; @@ -18,11 +17,6 @@ public class DocumentDBDistanceFunctionTests(DocumentDBDistanceFunctionTests.Fix public override Task HammingDistance() => Assert.ThrowsAsync(base.HammingDistance); public override Task ManhattanDistance() => Assert.ThrowsAsync(base.ManhattanDistance); - // AzureDocumentDB EuclideanDistance doesn't correctly filter by score threshold (returns all results). - // See https://github.com/CommunityToolkit/AI/issues/6. - protected override Task TestScoreThreshold(VectorStoreCollection.SearchRecord> collection) - => Task.CompletedTask; - public new class Fixture() : DistanceFunctionTests.Fixture { public override TestStore TestStore => DocumentDBTestStore.Instance;