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
138 changes: 138 additions & 0 deletions src/Weaviate.Client.Tests/Integration/TestSearchHybrid.cs
Original file line number Diff line number Diff line change
Expand Up @@ -880,6 +880,144 @@ await collection.Query.Hybrid(
Assert.Equal(splitAcross, andCrossObjs[0].UUID);
}

/// <summary>
/// Creates a collection with 3 tight clusters (a, b, c) of vectors in 3D.
/// </summary>
private async Task<CollectionClient> CreateClusteredCollection()
{
var collection = await CollectionFactory(
properties: new[] { Property.Text("name") },
vectorConfig: Configure.Vector(t => t.SelfProvided())
);

var data = new (string Name, float[] Vector)[]
{
("a1", [1.0f, 0.0f, 0.0f]),
("a2", [0.95f, 0.05f, 0.0f]),
("a3", [0.9f, 0.1f, 0.0f]),
("b1", [0.0f, 1.0f, 0.0f]),
("b2", [0.05f, 0.95f, 0.0f]),
("c1", [0.0f, 0.0f, 1.0f]),
};
foreach (var (name, vector) in data)
{
await collection.Data.Insert(
new { name },
vectors: vector,
cancellationToken: TestContext.Current.CancellationToken
);
}

return collection;
}

/// <summary>
/// Tests that test hybrid diversity balance zero differs from balance one
/// </summary>
[Fact]
public async Task Test_Hybrid_Diversity_Balance_Reorders()
{
RequireVersion("1.38.6");

var collection = await CreateClusteredCollection();

var baseline = (
await collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
limit: 3,
cancellationToken: TestContext.Current.CancellationToken
)
)
.Select(o => o.UUID)
.ToList();

var balanceZero = (
await collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 3, Balance: 0.0f),
limit: 3,
cancellationToken: TestContext.Current.CancellationToken
)
)
.Select(o => o.UUID)
.ToList();

var balanceOne = (
await collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 3, Balance: 1.0f),
limit: 3,
cancellationToken: TestContext.Current.CancellationToken
)
)
.Select(o => o.UUID)
.ToList();

// Pure diversity picks across clusters, so it must differ from pure relevance,
// while pure relevance matches the plain hybrid baseline.
Assert.NotEqual(balanceZero, balanceOne);
Assert.Equal(baseline, balanceOne);
}

/// <summary>
/// Tests that test hybrid diversity mmr limit caps results
/// </summary>
[Fact]
public async Task Test_Hybrid_Diversity_MMR_Limit_Caps_Results()
{
RequireVersion("1.38.6");

var collection = await CollectionFactory(
properties: new[] { Property.Text("name") },
vectorConfig: Configure.Vector(t => t.SelfProvided())
);

// Enough items (>25) that a small mmr limit is distinguishable from the server's default limit.
for (var i = 0; i < 50; i++)
{
await collection.Data.Insert(
new { name = $"t{i}" },
vectors: new float[] { 1.0f - 0.001f * i, 0f, 0f },
cancellationToken: TestContext.Current.CancellationToken
);
}

var objs = await collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 5, Balance: 0.5f),
cancellationToken: TestContext.Current.CancellationToken
);

Assert.Equal(5, objs.Count());
}

/// <summary>
/// Tests that test hybrid diversity missing limit errors
/// </summary>
[Fact]
public async Task Test_Hybrid_Diversity_MissingLimit_Errors()
{
RequireVersion("1.38.6");

var collection = await CreateClusteredCollection();

// The server requires the MMR limit; the client forwards the request unvalidated.
var exception = await Assert.ThrowsAnyAsync<WeaviateException>(async () =>
await collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Balance: 0.5f),
cancellationToken: TestContext.Current.CancellationToken
)
);
Assert.NotNull(exception.InnerException);
Assert.Contains("MMR limit", exception.InnerException.Message);
}

/// <summary>
/// Tests that test aggregate max vector distance
/// </summary>
Expand Down
191 changes: 191 additions & 0 deletions src/Weaviate.Client.Tests/Unit/TestHybridDiversitySyntax.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,191 @@
using Weaviate.Client.Models;
using Weaviate.Client.Tests.Unit.Mocks;
using Weaviate.Client.Typed;
using V1 = Weaviate.Client.Grpc.Protobuf.V1;

namespace Weaviate.Client.Tests.Unit;

/// <summary>
/// Unit tests verifying the hybrid diversity selection maps to the expected
/// Selection message in the gRPC request across the query, generate and typed paths.
/// </summary>
[Collection("Unit Tests")]
public class TestHybridDiversitySyntax : IAsyncLifetime
{
private const string CollectionName = "TestCollection";

private Func<V1.SearchRequest?> _getRequest = null!;
private CollectionClient _collection = null!;

/// <summary>
/// The article test type
/// </summary>
private class Article
{
/// <summary>
/// The title
/// </summary>
public string? Title { get; set; }
}

/// <summary>
/// Initializes this instance
/// </summary>
/// <returns>The value task</returns>
public ValueTask InitializeAsync()
{
var (client, getRequest) = MockGrpcClient.CreateWithSearchCapture(new Version(1, 39, 0));
_getRequest = getRequest;
_collection = client.Collections.Use(CollectionName);
return ValueTask.CompletedTask;
}

/// <summary>
/// Disposes this instance
/// </summary>
/// <returns>The value task</returns>
public ValueTask DisposeAsync()
{
GC.SuppressFinalize(this);
return ValueTask.CompletedTask;
}

/// <summary>
/// Tests that hybrid diversity selection sets selection on the request
/// </summary>
[Fact]
public async Task Hybrid_DiversitySelection_MMR_SetsSelection()
{
// Act
await _collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 7, Balance: 0.5f),
limit: 7,
cancellationToken: TestContext.Current.CancellationToken
);

// Assert
var request = _getRequest();
Assert.NotNull(request);
Assert.NotNull(request.HybridSearch.Selection);
Assert.Equal(7u, request.HybridSearch.Selection.Mmr.Limit);
Assert.Equal(0.5f, request.HybridSearch.Selection.Mmr.Balance);
}

/// <summary>
/// Tests that hybrid diversity selection with omitted balance leaves balance unset
/// </summary>
[Fact]
public async Task Hybrid_DiversitySelection_BalanceOmitted_LeavesBalanceUnset()
{
// Act
await _collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 3),
cancellationToken: TestContext.Current.CancellationToken
);

// Assert
var request = _getRequest();
Assert.NotNull(request);
Assert.NotNull(request.HybridSearch.Selection);
Assert.True(request.HybridSearch.Selection.Mmr.HasLimit);
Assert.False(request.HybridSearch.Selection.Mmr.HasBalance);
}

/// <summary>
/// Tests that hybrid without diversity selection leaves selection unset
/// </summary>
[Fact]
public async Task Hybrid_NoDiversitySelection_LeavesSelectionUnset()
{
// Act
await _collection.Query.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
limit: 5,
cancellationToken: TestContext.Current.CancellationToken
);

// Assert
var request = _getRequest();
Assert.NotNull(request);
Assert.Null(request.HybridSearch.Selection);
}

/// <summary>
/// Tests that generate hybrid diversity selection sets selection on the request
/// </summary>
[Fact]
public async Task Generate_Hybrid_DiversitySelection_SetsSelection()
{
// Act
await _collection.Generate.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 4, Balance: 0.25f),
singlePrompt: "Describe {title}",
cancellationToken: TestContext.Current.CancellationToken
);

// Assert
var request = _getRequest();
Assert.NotNull(request);
Assert.NotNull(request.HybridSearch.Selection);
Assert.Equal(4u, request.HybridSearch.Selection.Mmr.Limit);
Assert.Equal(0.25f, request.HybridSearch.Selection.Mmr.Balance);
}

/// <summary>
/// Tests that typed hybrid diversity selection sets selection on the request
/// </summary>
[Fact]
public async Task Typed_Hybrid_DiversitySelection_SetsSelection()
{
// Arrange
var typedQueryClient = new TypedQueryClient<Article>(_collection.Query);

// Act
await typedQueryClient.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 2, Balance: 1.0f),
cancellationToken: TestContext.Current.CancellationToken
);

// Assert
var request = _getRequest();
Assert.NotNull(request);
Assert.NotNull(request.HybridSearch.Selection);
Assert.Equal(2u, request.HybridSearch.Selection.Mmr.Limit);
Assert.Equal(1.0f, request.HybridSearch.Selection.Mmr.Balance);
}

/// <summary>
/// Tests that typed generate hybrid diversity selection sets selection on the request
/// </summary>
[Fact]
public async Task Typed_Generate_Hybrid_DiversitySelection_SetsSelection()
{
// Arrange
var typedGenerateClient = new TypedGenerateClient<Article>(_collection.Generate);

// Act
await typedGenerateClient.Hybrid(
query: null,
vectors: new float[] { 1f, 0f, 0f },
diversitySelection: new Diversity.MMR(Limit: 6),
singlePrompt: "Describe {title}",
cancellationToken: TestContext.Current.CancellationToken
);

// Assert
var request = _getRequest();
Assert.NotNull(request);
Assert.NotNull(request.HybridSearch.Selection);
Assert.Equal(6u, request.HybridSearch.Selection.Mmr.Limit);
Assert.False(request.HybridSearch.Selection.Mmr.HasBalance);
}
}
Loading
Loading