Let data sources be searched page by page with a query of their own

This commit is contained in:
Thorsten Sommer 2026-09-24 15:25:14 +02:00
parent 4617aa5230
commit 92339479d0
Signed by untrusted user who does not match committer: tsommer
GPG Key ID: 371BBA77A02C0108
9 changed files with 449 additions and 68 deletions

View File

@ -75,6 +75,39 @@ public readonly record struct DataSourceERI_V1 : IERIDataSource
/// <inheritdoc /> /// <inheritdoc />
public async Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) public async Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default)
{
var latestUserPrompt = lastUserPrompt switch
{
ContentText text => text.Text,
ContentImage image => await image.TryAsBase64(token) is (success: true, { } base64Image)
? base64Image
: string.Empty,
_ => string.Empty
};
return await this.RetrieveDataAsync(latestUserPrompt, lastUserPrompt.ToERIContentType, thread, this.MaxMatches, token);
}
/// <inheritdoc />
public async Task<RetrievalPage> RetrieveDataAsync(string query, int page, ChatThread thread, CancellationToken token = default)
{
var window = RetrievalPaging.GetWindowSize(page, this.MaxMatches);
if (this.MaxMatches == 0)
return RetrievalPage.EMPTY;
//
// ERI v1 knows no query apart from the latest user prompt, so the query takes its place; the
// thread still tells the server what the conversation is about. Nor does it know an offset:
// the server returns the whole window, and the page is cut from it here. Hence, the pages
// are only as stable as the order in which the server returns its matches. A server which
// returns fewer matches than asked for ends the paging early, which errs on the safe side.
//
var contexts = await this.RetrieveDataAsync(query, ContentType.TEXT, thread, window, token);
var (pageContexts, hasMore) = RetrievalPaging.Cut(contexts, page, this.MaxMatches);
return new RetrievalPage(pageContexts, hasMore);
}
private async Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(string latestUserPrompt, ContentType latestUserPromptType, ChatThread thread, int maxMatches, CancellationToken token)
{ {
// Important: Do not dispose the RustService here, as it is a singleton. // Important: Do not dispose the RustService here, as it is a singleton.
var rustService = Program.SERVICE_PROVIDER.GetRequiredService<RustService>(); var rustService = Program.SERVICE_PROVIDER.GetRequiredService<RustService>();
@ -86,18 +119,11 @@ public readonly record struct DataSourceERI_V1 : IERIDataSource
{ {
var retrievalRequest = new RetrievalRequest var retrievalRequest = new RetrievalRequest
{ {
LatestUserPromptType = lastUserPrompt.ToERIContentType, LatestUserPromptType = latestUserPromptType,
LatestUserPrompt = lastUserPrompt switch LatestUserPrompt = latestUserPrompt,
{
ContentText text => text.Text,
ContentImage image => await image.TryAsBase64(token) is (success: true, { } base64Image)
? base64Image
: string.Empty,
_ => string.Empty
},
Thread = await thread.ToERIChatThread(token), Thread = await thread.ToERIChatThread(token),
MaxMatches = this.MaxMatches, MaxMatches = maxMatches,
RetrievalProcessId = this.SelectedRetrievalId, RetrievalProcessId = this.SelectedRetrievalId,
Parameters = null, // The ERI server selects useful default parameters Parameters = null, // The ERI server selects useful default parameters
}; };

View File

@ -57,6 +57,10 @@ public readonly record struct DataSourceLocalDirectory : IInternalDataSource
public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) => public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) =>
Program.SERVICE_PROVIDER.GetRequiredService<DataSourceLocalRetrievalService>().RetrieveDataAsync(this, lastUserPrompt, thread, token); Program.SERVICE_PROVIDER.GetRequiredService<DataSourceLocalRetrievalService>().RetrieveDataAsync(this, lastUserPrompt, thread, token);
/// <inheritdoc />
public Task<RetrievalPage> RetrieveDataAsync(string query, int page, ChatThread thread, CancellationToken token = default) =>
Program.SERVICE_PROVIDER.GetRequiredService<DataSourceLocalRetrievalService>().RetrieveDataAsync(this, query, page, thread, token);
/// <summary> /// <summary>
/// The path to the directory. /// The path to the directory.
/// </summary> /// </summary>

View File

@ -57,6 +57,10 @@ public readonly record struct DataSourceLocalFile : IInternalDataSource
public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) => public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) =>
Program.SERVICE_PROVIDER.GetRequiredService<DataSourceLocalRetrievalService>().RetrieveDataAsync(this, lastUserPrompt, thread, token); Program.SERVICE_PROVIDER.GetRequiredService<DataSourceLocalRetrievalService>().RetrieveDataAsync(this, lastUserPrompt, thread, token);
/// <inheritdoc />
public Task<RetrievalPage> RetrieveDataAsync(string query, int page, ChatThread thread, CancellationToken token = default) =>
Program.SERVICE_PROVIDER.GetRequiredService<DataSourceLocalRetrievalService>().RetrieveDataAsync(this, query, page, thread, token);
/// <summary> /// <summary>
/// The path to the file. /// The path to the file.
/// </summary> /// </summary>

View File

@ -22,7 +22,7 @@ public interface IDataSource : IConfigurationObject
public DataSourceType Type { get; init; } public DataSourceType Type { get; init; }
/// <summary> /// <summary>
/// The maximum number of matches to return when retrieving data from the ERI server. /// The maximum number of matches one retrieval returns. Searched page by page, it is the size of a page.
/// </summary> /// </summary>
public ushort MaxMatches { get; init; } public ushort MaxMatches { get; init; }
@ -34,4 +34,21 @@ public interface IDataSource : IConfigurationObject
/// <param name="token">The cancellation token.</param> /// <param name="token">The cancellation token.</param>
/// <returns>The retrieved data context.</returns> /// <returns>The retrieved data context.</returns>
public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default); public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default);
/// <summary>
/// Search the data source for a query of its own, one page at a time.
/// </summary>
/// <remarks>
/// Unlike the retrieval above, the query need not be what the user wrote last: Semantic Search
/// lets the model work it out from the conversation, and search as often as it takes. The first
/// page holds what the retrieval above finds for the same text. How the pages are cut is
/// described in RetrievalPaging.
/// </remarks>
/// <param name="query">What to search for.</param>
/// <param name="page">The page to retrieve, from 1 up to RetrievalPaging.GetLastPage for MaxMatches.</param>
/// <param name="thread">The chat thread.</param>
/// <param name="token">The cancellation token.</param>
/// <returns>The retrieved data contexts of this page, and whether the next page is worth asking for.</returns>
/// <exception cref="ArgumentOutOfRangeException">The page is below 1 or beyond the last page.</exception>
public Task<RetrievalPage> RetrieveDataAsync(string query, int page, ChatThread thread, CancellationToken token = default);
} }

View File

@ -310,6 +310,11 @@ public sealed class SqliteIndexStoreClientImplementation(string name, string dat
if (string.IsNullOrWhiteSpace(ftsQuery)) if (string.IsNullOrWhiteSpace(ftsQuery))
return []; return [];
//
// Chunks of the same score keep the order of their rows. The results are cut into pages by
// asking for more of them each time, cf. RetrievalPaging. If ties could fall differently
// with every limit, a page might show a chunk again or skip one.
//
await using var context = this.CreateContext(); await using var context = this.CreateContext();
var results = await context.SearchResults var results = await context.SearchResults
.FromSqlInterpolated($""" .FromSqlInterpolated($"""
@ -338,7 +343,7 @@ public sealed class SqliteIndexStoreClientImplementation(string name, string dat
JOIN data_sources ds ON ds.data_source_id = f.data_source_id JOIN data_sources ds ON ds.data_source_id = f.data_source_id
WHERE ds.data_source_id = {dataSourceId} WHERE ds.data_source_id = {dataSourceId}
AND embedding_chunks_fts MATCH {ftsQuery} AND embedding_chunks_fts MATCH {ftsQuery}
ORDER BY Score ORDER BY Score, c.id
LIMIT {maxMatches} LIMIT {maxMatches}
""") """)
.AsNoTracking() .AsNoTracking()

View File

@ -0,0 +1,22 @@
namespace AIStudio.Tools.RAG;
/// <summary>
/// One page of what a search in a data source found.
/// </summary>
/// <remarks>
/// A page does not say how many matches there are in total, and it could not: a vector search has
/// no total, since every chunk matches, only less similar ones match less. What a page does say is
/// whether asking for the next one is worth it.
/// </remarks>
/// <param name="Contexts">What this page found, the most relevant first.</param>
/// <param name="HasMore">True when the next page can be retrieved and may hold further matches. That
/// page can still turn out empty, when everything on it was already shown on an earlier page. False
/// when the search is exhausted, or when this page is the last one which can be retrieved at all, cf.
/// RetrievalPaging.GetLastPage.</param>
public sealed record RetrievalPage(IReadOnlyList<IRetrievalContext> Contexts, bool HasMore)
{
/// <summary>
/// A page without any matches and nothing after it.
/// </summary>
public static readonly RetrievalPage EMPTY = new([], false);
}

View File

@ -0,0 +1,170 @@
namespace AIStudio.Tools.RAG;
/// <summary>
/// Cuts what a search found into pages, without keeping anything between two of them.
/// </summary>
/// <remarks>
/// <para>
/// Neither the vector store nor the keyword index knows an offset, and neither needs one: page p
/// of size k is cut from the first p·k + 1 matches of every channel. The one match beyond the
/// page tells whether a next page is worth asking for. The first page is therefore exactly what a
/// search for k matches always returned.
/// </para>
/// <para>
/// Staying without state is not a shortcut but a requirement: tool results do not travel into
/// later turns, so a page has to come out of the query and its number alone.
/// </para>
/// </remarks>
public static class RetrievalPaging
{
/// <summary>
/// How many matches a page beyond the first may fetch at most, per channel.
/// </summary>
/// <remarks>
/// Every page fetches its whole window again, from the vector store and the keyword index, or
/// from the ERI server. The first page is exempt: its size is what the user or the organization
/// configured, and fetching it is what the retrieval always did.
/// </remarks>
public const int MAX_RESULT_WINDOW = 100;
/// <summary>
/// The last page which can be retrieved for the given page size.
/// </summary>
/// <param name="pageSize">The number of matches per page.</param>
/// <returns>The number of the last page, which is at least 1.</returns>
public static int GetLastPage(int pageSize) => pageSize < 1 ? 1 : Math.Max(1, (MAX_RESULT_WINDOW - 1) / pageSize);
/// <summary>
/// How many matches every channel has to deliver for the given page.
/// </summary>
/// <param name="page">The page, starting at 1.</param>
/// <param name="pageSize">The number of matches per page.</param>
/// <returns>The size of the window, i.e., the page, all pages before it, and one match more.</returns>
/// <exception cref="ArgumentOutOfRangeException">The page is below 1 or beyond the last page.</exception>
public static int GetWindowSize(int page, int pageSize) => GetPageEnd(page, pageSize) + 1;
/// <summary>
/// Cuts one page out of what a single channel found.
/// </summary>
/// <param name="matches">What the channel found, the most relevant first, fetched with the window of this page.</param>
/// <param name="page">The page, starting at 1.</param>
/// <param name="pageSize">The number of matches per page.</param>
/// <returns>The matches of this page, and whether the next page is worth asking for.</returns>
/// <exception cref="ArgumentOutOfRangeException">The page is below 1 or beyond the last page.</exception>
public static (IReadOnlyList<T> Matches, bool HasMore) Cut<T>(IReadOnlyList<T> matches, int page, int pageSize)
{
var end = GetPageEnd(page, pageSize);
var start = end - pageSize;
var pageMatches = matches.Skip(start).Take(pageSize).ToList();
return (pageMatches, HasMore(page, pageSize, matches.Count));
}
/// <summary>
/// Cuts one page out of what two channels found.
/// </summary>
/// <remarks>
/// <para>
/// A page holds the page of the first channel, followed by the page of the second one. This
/// order is deterministic on purpose; reranking would replace it, and change the first page
/// with it.
/// </para>
/// <para>
/// A match both channels found is shown once, on the earlier of its two pages; on the same
/// page, in the part of the first channel. Hence, no match turns up on two pages. Matches
/// without a key are never taken for one another.
/// </para>
/// </remarks>
/// <param name="first">What the first channel found, the most relevant first, fetched with the window of this page.</param>
/// <param name="second">What the second channel found, likewise.</param>
/// <param name="getKey">What identifies a match across both channels. Letter case does not matter.</param>
/// <param name="page">The page, starting at 1.</param>
/// <param name="pageSize">The number of matches per page and channel.</param>
/// <returns>The matches of this page, and whether the next page is worth asking for.</returns>
/// <exception cref="ArgumentOutOfRangeException">The page is below 1 or beyond the last page.</exception>
public static (IReadOnlyList<T> Matches, bool HasMore) Merge<T>(IReadOnlyList<T> first, IReadOnlyList<T> second, Func<T, string> getKey, int page, int pageSize)
{
var end = GetPageEnd(page, pageSize);
var start = end - pageSize;
var firstRanks = GetFirstRanks(first, end + 1, getKey);
var secondRanks = GetFirstRanks(second, end + 1, getKey);
var pageMatches = new List<T>(2 * pageSize);
for (var rank = start; rank < Math.Min(end, first.Count); rank++)
{
var match = first[rank];
var key = getKey(match);
if (!string.IsNullOrWhiteSpace(key))
{
// The first channel found it further up already:
if (firstRanks[key] != rank)
continue;
// The second channel showed it on an earlier page:
if (secondRanks.TryGetValue(key, out var secondRank) && secondRank < start)
continue;
}
pageMatches.Add(match);
}
for (var rank = start; rank < Math.Min(end, second.Count); rank++)
{
var match = second[rank];
var key = getKey(match);
if (!string.IsNullOrWhiteSpace(key))
{
// The second channel found it further up already:
if (secondRanks[key] != rank)
continue;
// The first channel shows it on this page or showed it on an earlier one:
if (firstRanks.TryGetValue(key, out var firstRank) && firstRank < end)
continue;
}
pageMatches.Add(match);
}
return (pageMatches, HasMore(page, pageSize, first.Count, second.Count));
}
/// <summary>
/// Where the given page ends, i.e., the number of matches on it and on all pages before it.
/// </summary>
private static int GetPageEnd(int page, int pageSize)
{
var lastPage = GetLastPage(pageSize);
if (page < 1 || page > lastPage)
throw new ArgumentOutOfRangeException(nameof(page), page, $"With {pageSize} matches per page, the page has to be between 1 and {lastPage}.");
return page * Math.Max(0, pageSize);
}
/// <remarks>
/// Whatever a channel found beyond this page is enough to ask for the next one. That page can
/// still turn out empty, when the other channel showed all of it before. Saying there is more
/// when there is not costs one empty page; saying the opposite would hide matches.
/// </remarks>
private static bool HasMore(int page, int pageSize, params int[] channelCounts)
{
if (page >= GetLastPage(pageSize))
return false;
var end = page * pageSize;
return channelCounts.Any(count => count > end);
}
private static Dictionary<string, int> GetFirstRanks<T>(IReadOnlyList<T> matches, int window, Func<T, string> getKey)
{
var ranks = new Dictionary<string, int>(StringComparer.OrdinalIgnoreCase);
for (var rank = 0; rank < Math.Min(window, matches.Count); rank++)
{
var key = getKey(matches[rank]);
if (!string.IsNullOrWhiteSpace(key))
ranks.TryAdd(key, rank);
}
return ranks;
}
}

View File

@ -60,18 +60,31 @@ public sealed class DataSourceLocalRetrievalService(
public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(DataSourceLocalDirectory dataSource, IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) => public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(DataSourceLocalDirectory dataSource, IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) =>
this.RetrieveDataAsync(dataSource, lastUserPrompt, token); this.RetrieveDataAsync(dataSource, lastUserPrompt, token);
public Task<RetrievalPage> RetrieveDataAsync(DataSourceLocalFile dataSource, string query, int page, ChatThread thread, CancellationToken token = default) =>
this.RetrievePageAsync(dataSource, query, page, token);
public Task<RetrievalPage> RetrieveDataAsync(DataSourceLocalDirectory dataSource, string query, int page, ChatThread thread, CancellationToken token = default) =>
this.RetrievePageAsync(dataSource, query, page, token);
private async Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IInternalDataSource dataSource, IContent lastUserPrompt, CancellationToken token) private async Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IInternalDataSource dataSource, IContent lastUserPrompt, CancellationToken token)
{ {
var query = GetQueryText(lastUserPrompt); // The first page is what this retrieval has always returned:
var firstPage = await this.RetrievePageAsync(dataSource, GetQueryText(lastUserPrompt), 1, token);
return firstPage.Contexts;
}
private async Task<RetrievalPage> RetrievePageAsync(IInternalDataSource dataSource, string query, int page, CancellationToken token)
{
var pageSize = (int)dataSource.MaxMatches;
var window = RetrievalPaging.GetWindowSize(page, pageSize);
if (string.IsNullOrWhiteSpace(query)) if (string.IsNullOrWhiteSpace(query))
{ {
logger.LogDebug("Skipping local retrieval for data source '{DataSourceName}' ({DataSourceId}) because the latest prompt does not contain text.", dataSource.Name, dataSource.Id); logger.LogDebug("Skipping local retrieval for data source '{DataSourceName}' ({DataSourceId}) because there is no text to search for.", dataSource.Name, dataSource.Id);
return []; return RetrievalPage.EMPTY;
} }
var maxMatches = (int)dataSource.MaxMatches; if (pageSize == 0)
if (maxMatches == 0) return RetrievalPage.EMPTY;
return [];
// //
// A data source waiting for its index is kept out of the selection before the RAG process // A data source waiting for its index is kept out of the selection before the RAG process
@ -87,30 +100,40 @@ public sealed class DataSourceLocalRetrievalService(
{ {
logger.LogWarning("Skipping local retrieval for data source '{DataSourceName}' ({DataSourceId}) because its index has to be built anew.", dataSource.Name, dataSource.Id); logger.LogWarning("Skipping local retrieval for data source '{DataSourceName}' ({DataSourceId}) because its index has to be built anew.", dataSource.Name, dataSource.Id);
await this.ReportRetrievalGapAsync(dataSource, "index-rebuilding", string.Format(TB("The data source '{0}' was left out of the answer: it is being indexed again and cannot be searched until that is finished."), dataSource.Name)); await this.ReportRetrievalGapAsync(dataSource, "index-rebuilding", string.Format(TB("The data source '{0}' was left out of the answer: it is being indexed again and cannot be searched until that is finished."), dataSource.Name));
return []; return RetrievalPage.EMPTY;
} }
var collectionName = DataSourceEmbeddingNames.GetCollectionName(dataSource.Id); var collectionName = DataSourceEmbeddingNames.GetCollectionName(dataSource.Id);
var vectorTask = this.SearchVectorAsync(dataSource, query, maxMatches, collectionName, token); var vectorTask = this.SearchVectorAsync(dataSource, query, window, collectionName, token);
var bm25Task = this.SearchBm25Async(dataSource, query, maxMatches, token); var bm25Task = this.SearchBm25Async(dataSource, query, window, token);
await Task.WhenAll(vectorTask, bm25Task); await Task.WhenAll(vectorTask, bm25Task);
token.ThrowIfCancellationRequested(); token.ThrowIfCancellationRequested();
var hits = MergeResults(vectorTask.Result, bm25Task.Result, maxMatches); var (hits, hasMore) = RetrievalPaging.Merge(
vectorTask.Result.Select((result, index) => FromVectorResult(result, index + 1)).ToList(),
bm25Task.Result.Select((result, index) => FromBm25Result(result, index + 1)).ToList(),
hit => hit.ChunkId,
page,
pageSize);
logger.LogInformation( logger.LogInformation(
"Retrieved {MergedHits} local RAG hits for data source '{DataSourceName}' ({DataSourceId}). VectorCandidates={VectorHits}, BM25Candidates={BM25Hits}, RequestedPerChannel={RequestedPerChannel}.", "Retrieved {MergedHits} local RAG hits on page {Page} for data source '{DataSourceName}' ({DataSourceId}). VectorCandidates={VectorHits}, BM25Candidates={BM25Hits}, RequestedPerChannel={RequestedPerChannel}, HasMore={HasMore}.",
hits.Count, hits.Count,
page,
dataSource.Name, dataSource.Name,
dataSource.Id, dataSource.Id,
vectorTask.Result.Count, vectorTask.Result.Count,
bm25Task.Result.Count, bm25Task.Result.Count,
maxMatches); window,
hasMore);
return hits var contexts = hits
.Where(hit => !string.IsNullOrWhiteSpace(hit.Text)) .Where(hit => !string.IsNullOrWhiteSpace(hit.Text))
.Select(hit => ToRetrievalContext(hit, dataSource)) .Select(hit => ToRetrievalContext(hit, dataSource))
.ToList(); .ToList();
return new RetrievalPage(contexts, hasMore);
} }
private async Task<IReadOnlyList<VectorSearchResult>> SearchVectorAsync( private async Task<IReadOnlyList<VectorSearchResult>> SearchVectorAsync(
@ -312,7 +335,7 @@ public sealed class DataSourceLocalRetrievalService(
return results; return results;
logger.LogWarning( logger.LogWarning(
"Local RAG {SearchName} search returned {ReturnedHits} chunks for data source '{DataSourceName}' ({DataSourceId}), which exceeds the configured maximum {MaxMatches}. Truncating to the datasource limit.", "Local RAG {SearchName} search returned {ReturnedHits} chunks for data source '{DataSourceName}' ({DataSourceId}), which exceeds the requested maximum {MaxMatches}. Truncating to it.",
searchName, searchName,
results.Count, results.Count,
dataSource.Name, dataSource.Name,
@ -322,47 +345,6 @@ public sealed class DataSourceLocalRetrievalService(
return results.Take(maxMatches).ToList(); return results.Take(maxMatches).ToList();
} }
private static IReadOnlyList<LocalRetrievalHit> MergeResults(
IReadOnlyList<VectorSearchResult> vectorResults,
IReadOnlyList<IndexStoreSearchResult> bm25Results,
int maxMatches)
{
// Future reranking should replace this deterministic channel merge.
var merged = new List<LocalRetrievalHit>(maxMatches * 2);
var seenChunkIds = new HashSet<string>(StringComparer.OrdinalIgnoreCase);
AppendHits(
merged,
seenChunkIds,
vectorResults
.Select((result, index) => FromVectorResult(result, index + 1)),
maxMatches);
AppendHits(
merged,
seenChunkIds,
bm25Results
.Select((result, index) => FromBm25Result(result, index + 1)),
maxMatches);
return merged;
}
private static void AppendHits(List<LocalRetrievalHit> merged, HashSet<string> seenChunkIds, IEnumerable<LocalRetrievalHit> hits, int maxNewHits)
{
var added = 0;
foreach (var hit in hits)
{
if (!string.IsNullOrWhiteSpace(hit.ChunkId) && !seenChunkIds.Add(hit.ChunkId))
continue;
merged.Add(hit);
added++;
if (added >= maxNewHits)
return;
}
}
private static LocalRetrievalHit FromVectorResult(VectorSearchResult result, int rank) => private static LocalRetrievalHit FromVectorResult(VectorSearchResult result, int rank) =>
new( new(
RetrievalChannel.VECTOR, RetrievalChannel.VECTOR,

View File

@ -0,0 +1,151 @@
using AIStudio.Tools.RAG;
namespace AIStudio.Tests.Tools;
/// <summary>
/// Checks how what a search found is cut into pages.
/// </summary>
/// <remarks>
/// Semantic Search lets the model page through a data source, yet nothing is kept between two
/// pages: every page is cut anew from a larger window. Three things have to hold for that. The
/// first page is what the classic RAG process always received, so the way of searching changes
/// nothing about what is found. No match turns up on two pages, although both channels of a local
/// data source often find the same chunk. And the model is told there is more whenever there might
/// be, and never that there is nothing when there is.
/// </remarks>
[TestFixture]
public sealed class RetrievalPagingTests
{
private const int PAGE_SIZE = 2;
[Test]
public void TheFirstPageShowsTheFirstChannelThenWhatOnlyTheSecondFound()
{
// The vector search found a, b, and c; the keyword search b, d, and e:
var (matches, _) = Merge(["a", "b", "c"], ["b", "d", "e"], page: 1);
Assert.That(matches, Is.EqualTo(new[] { "a", "b", "d" }), "This is what the RAG process always sent: the vector matches first, then the keyword matches it did not have yet.");
}
[Test]
public void EveryMatchTurnsUpOnExactlyOnePage()
{
string[] first = ["a", "b", "c", "d", "e", "f", "g"];
string[] second = ["c", "h", "a", "i", "e", "j", "k"];
var shown = new List<string>();
for (var page = 1; page <= 3; page++)
shown.AddRange(Merge(first, second, page).Matches);
Assert.Multiple(() =>
{
Assert.That(shown, Is.Unique, "A chunk both channels found is shown on the earlier of its two pages only.");
Assert.That(shown, Is.EquivalentTo(first.Take(6).Union(second.Take(6))), "Leaving out the duplicates must not leave out anything else.");
});
}
[Test]
public void ThereIsMoreWhenAChannelFoundMoreThanThePageHolds()
{
var (_, hasMore) = Merge(["a", "b", "c"], [], page: 1);
Assert.That(hasMore, Is.True);
}
[Test]
public void ThereIsNothingMoreWhenEveryChannelEndsOnThisPage()
{
var (_, hasMore) = Merge(["a", "b"], ["c", "d"], page: 1);
Assert.That(hasMore, Is.False, "Neither channel found anything beyond this page, so the next one would be empty.");
}
[Test]
public void TheLastPageHasNothingAfterIt()
{
var lastPage = RetrievalPaging.GetLastPage(PAGE_SIZE);
var everything = Enumerable.Range(0, RetrievalPaging.MAX_RESULT_WINDOW).Select(number => $"chunk-{number}").ToArray();
Assert.Multiple(() =>
{
Assert.That(Merge(everything, [], lastPage - 1).HasMore, Is.True);
Assert.That(Merge(everything, [], lastPage).HasMore, Is.False, "No page beyond this one can be retrieved, however much the search found.");
});
}
[TestCase(0)]
[TestCase(-1)]
public void APageBelowTheFirstCannotBeRetrieved(int page)
{
Assert.Throws<ArgumentOutOfRangeException>(() => RetrievalPaging.GetWindowSize(page, PAGE_SIZE));
}
[Test]
public void APageBeyondTheLastCannotBeRetrieved()
{
var lastPage = RetrievalPaging.GetLastPage(PAGE_SIZE);
Assert.Throws<ArgumentOutOfRangeException>(() => RetrievalPaging.GetWindowSize(lastPage + 1, PAGE_SIZE));
}
[TestCase(1)]
[TestCase(7)]
[TestCase(10)]
[TestCase(33)]
[TestCase(50)]
public void NoPageBeyondTheFirstFetchesMoreThanTheLimit(int pageSize)
{
var lastPage = RetrievalPaging.GetLastPage(pageSize);
Assert.That(RetrievalPaging.GetWindowSize(lastPage, pageSize), Is.LessThanOrEqualTo(RetrievalPaging.MAX_RESULT_WINDOW), "Every page fetches its whole window again.");
}
[Test]
public void TheFirstPageAlwaysHoldsTheConfiguredNumberOfMatches()
{
Assert.Multiple(() =>
{
Assert.That(RetrievalPaging.GetLastPage(500), Is.EqualTo(1));
Assert.That(RetrievalPaging.GetWindowSize(1, 500), Is.EqualTo(501), "The limit is for paging deeper. It must not shorten what the user asked for per search.");
});
}
[Test]
public void MatchesWithoutAKeyAreNeverTakenForOneAnother()
{
var (matches, _) = Merge(["", "a"], ["", "b"], page: 1);
Assert.That(matches, Is.EqualTo(new[] { "", "a", "", "b" }));
}
[Test]
public void LetterCaseDoesNotTellMatchesApart()
{
var (matches, _) = Merge(["CHUNK-1"], ["chunk-1", "chunk-2"], page: 1);
Assert.That(matches, Is.EqualTo(new[] { "CHUNK-1", "chunk-2" }));
}
[Test]
public void ASingleChannelIsCutInOrder()
{
string[] matches = ["a", "b", "c", "d", "e"];
var secondPage = RetrievalPaging.Cut(matches, 2, PAGE_SIZE);
var thirdPage = RetrievalPaging.Cut(matches, 3, PAGE_SIZE);
Assert.Multiple(() =>
{
Assert.That(secondPage.Matches, Is.EqualTo(new[] { "c", "d" }));
Assert.That(secondPage.HasMore, Is.True);
Assert.That(thirdPage.Matches, Is.EqualTo(new[] { "e" }));
Assert.That(thirdPage.HasMore, Is.False, "An ERI server which found fewer than asked for ends the paging.");
});
}
private static (IReadOnlyList<string> Matches, bool HasMore) Merge(string[] first, string[] second, int page)
{
// A channel never returns more than it is asked for:
var window = RetrievalPaging.GetWindowSize(page, PAGE_SIZE);
return RetrievalPaging.Merge(first.Take(window).ToList(), second.Take(window).ToList(), match => match, page, PAGE_SIZE);
}
}