add search filter

This commit is contained in:
Jicheng Lu 2024-08-28 17:45:33 -05:00
parent 8674ac9830
commit f4e20ac64c
4 changed files with 59 additions and 39 deletions

View file

@ -0,0 +1,10 @@
namespace BotSharp.Abstraction.Models;
public class KeyValue
{
[JsonPropertyName("key")]
public string Key { get; set; }
[JsonPropertyName("value")]
public string Value { get; set; }
}

View file

@ -4,4 +4,7 @@ public class VectorFilter : StringIdPagination
{
[JsonPropertyName("with_vector")]
public bool WithVector { get; set; }
}
[JsonPropertyName("search_pairs")]
public IEnumerable<KeyValue>? SearchPairs { get; set; }
}

View file

@ -51,10 +51,34 @@ public class QdrantDb : IVectorDb
return new StringIdPagedItems<VectorCollectionData>();
}
var totalPointCount = await client.CountAsync(collectionName);
// Build query filter
Filter? queryFilter = null;
if (!filter.SearchPairs.IsNullOrEmpty())
{
var conditions = filter.SearchPairs.Select(x => new Condition
{
Field = new FieldCondition
{
Key = x.Key,
Match = new Match { Text = x.Value },
}
});
queryFilter = new Filter
{
Should =
{
conditions
}
};
}
var totalPointCount = await client.CountAsync(collectionName, filter: queryFilter);
var response = await client.ScrollAsync(collectionName, limit: (uint)filter.Size,
offset: !string.IsNullOrWhiteSpace(filter.StartId) ? new PointId { Uuid = filter.StartId } : 0,
offset: !string.IsNullOrWhiteSpace(filter.StartId) ? new PointId { Uuid = filter.StartId } : null,
filter: queryFilter,
vectorsSelector: filter.WithVector);
var points = response?.Result?.Select(x => new VectorCollectionData
{
Id = x.Id?.Uuid ?? string.Empty,
@ -160,43 +184,26 @@ public class QdrantDb : IVectorDb
return results;
}
var points = await client.SearchAsync(collectionName,
vector,
limit: (ulong)limit,
scoreThreshold: confidence,
vectorsSelector: new WithVectorsSelector { Enable = withVector });
var pickFields = fields != null;
foreach (var point in points)
var payloadSelector = new WithPayloadSelector { Enable = true };
if (fields != null)
{
var data = new Dictionary<string, string>();
if (pickFields)
{
foreach (var field in fields)
{
if (point.Payload.ContainsKey(field))
{
data[field] = point.Payload[field].StringValue;
}
else
{
data[field] = "";
}
}
}
else
{
data = point.Payload.ToDictionary(k => k.Key, v => v.Value.StringValue);
}
results.Add(new VectorCollectionData
{
Id = point.Id.Uuid,
Data = data,
Score = point.Score,
Vector = withVector ? point.Vectors?.Vector?.Data?.ToArray() : null
});
payloadSelector.Include = new PayloadIncludeSelector { Fields = { fields.ToArray() } };
}
var points = await client.SearchAsync(collectionName,
vector,
limit: (ulong)limit,
scoreThreshold: confidence,
payloadSelector: payloadSelector,
vectorsSelector: withVector);
results = points.Select(x => new VectorCollectionData
{
Id = x.Id.Uuid,
Data = x.Payload.ToDictionary(x => x.Key, x => x.Value.StringValue),
Score = x.Score,
Vector = x.Vectors?.Vector?.Data?.ToArray()
}).ToList();
return results;
}

View file

@ -31,7 +31,7 @@ namespace BotSharp.Plugin.SemanticKernel
public Task<StringIdPagedItems<VectorCollectionData>> GetPagedCollectionData(string collectionName, VectorFilter filter)
{
throw new System.NotImplementedException();
throw new NotImplementedException();
}
public Task<IEnumerable<VectorCollectionData>> GetCollectionData(string collectionName, IEnumerable<Guid> ids,