From 9c2c01cb246c060245f1374eaa3b472969fa15be Mon Sep 17 00:00:00 2001 From: Empiree Date: Sat, 26 Sep 2026 14:37:22 +0200 Subject: [PATCH] Fix AzureDocumentDB score threshold for EuclideanDistance --- .../AzureDocumentDB/AzureDocumentDB.csproj | 2 +- .../AzureDocumentDB/DocumentDBCollection.cs | 2 +- .../DocumentDBCollectionSearchMapping.cs | 20 ++++++++++++++----- .../DocumentDBDistanceFunctionTests.cs | 6 ------ 4 files changed, 17 insertions(+), 13 deletions(-) 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;