diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index a4d90de..c3dd36a 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -36,6 +36,7 @@ jobs: embedding) PROJECT="VectorSharp.Embedding" ;; nomic-embed) PROJECT="VectorSharp.Embedding.NomicEmbed" ;; chunking) PROJECT="VectorSharp.Chunking" ;; + reranking) PROJECT="VectorSharp.Reranking" ;; *) echo "Unknown package: $PACKAGE" && exit 1 ;; esac echo "PROJECT=${PROJECT}" >> $GITHUB_OUTPUT diff --git a/Publishing.md b/Publishing.md index 286e566..3df245c 100644 --- a/Publishing.md +++ b/Publishing.md @@ -18,6 +18,7 @@ Where `{package}` is the lowercase package name and `{version}` is a valid semve | `embedding` | VectorSharp.Embedding | | `nomic-embed` | VectorSharp.Embedding.NomicEmbed | | `chunking` | VectorSharp.Chunking | +| `reranking` | VectorSharp.Reranking | ## Publishing a Package @@ -28,29 +29,19 @@ git push origin storage-v1.0.0 This triggers the publish workflow which will build, run tests, pack the specified project with the given version, and push to nuget.org. -## Examples +Packages are versioned independently: a tag names one package and one version, and publishing one +leaves the others alone. Take the prefix from the table above. -```bash -# Publish VectorSharp.Storage 1.0.0 -git tag storage-v1.0.0 -git push origin storage-v1.0.0 +### The tag is what sets the version -# Publish VectorSharp.Storage 1.1.0 (independent of other packages) -git tag storage-v1.1.0 -git push origin storage-v1.1.0 +The workflow builds and packs with `-p:Version=`, which overrides the +project file's `` for every project in the build. **The tag decides what is published; +the project file does not.** `embedding-v1.0.3` shipped from a project file that read +`1.0.0`, and nothing noticed. -# Publish VectorSharp.Embedding 1.0.0 (independent of other packages) -git tag embedding-v1.0.0 -git push origin embedding-v1.0.0 - -# Publish VectorSharp.Embedding.NomicEmbed 1.0.0 -git tag nomic-embed-v1.0.0 -git push origin nomic-embed-v1.0.0 - -# Publish VectorSharp.Chunking 1.0.0 -git tag chunking-v1.0.0 -git push origin chunking-v1.0.0 -``` +Keep `` in step with the last tag anyway, and bump it in the same change that adds the +behaviour it names. It is what a local `dotnet pack` produces, what a developer reads to see what +is in the working tree, and how a reviewer tells a patch from a minor. ## Prerequisites @@ -58,6 +49,19 @@ A `NUGET_API_KEY` secret must be configured in the repository settings (Settings ## Adding a New Package -1. Add the project to the solution +1. Add both the package project and its test project to `VectorSharp.slnx` 2. Add a new case in `.github/workflows/publish.yml` mapping the tag prefix to the project name -3. Update the table above +3. Update the tag prefix table above +4. Add a row to the Packages table in the root `README.md` — a package missing from it is invisible + to anyone arriving at the repository +5. Write the package's own `README.md`, and wire it up in the project file with both + `README.md` and + ``. Without the second, the build still + succeeds and the nuget.org listing renders blank +6. Add a `` to it in `VectorSharp.Packaging.Tests.csproj` + +`VectorSharp.Packaging.Tests` takes the set of packages from `VectorSharp.slnx` rather than a list +of its own, so step 1 is what puts a new package under its checks — there is no list to forget to +update. Step 6 is what lets those checks find the built `.nupkg`; without it they fail naming the +missing reference. A package that legitimately has dependencies must be named in that project's +`PackagesWithDependencies`, with the reason. diff --git a/README.md b/README.md index 84332c2..b718221 100644 --- a/README.md +++ b/README.md @@ -3,16 +3,17 @@ [![Tests](https://github.com/AdamTovatt/vector-sharp/actions/workflows/dotnet.yml/badge.svg)](https://github.com/AdamTovatt/vector-sharp/actions/workflows/dotnet.yml) [![License: MIT](https://img.shields.io/badge/License-MIT-green.svg)](https://opensource.org/licenses/MIT) -A high-performance .NET library for in-process vector similarity search and text embedding. No external services, no Docker, no infrastructure — just NuGet packages that work. +A high-performance .NET library for in-process vector similarity search and text embedding. No external services required, no Docker, no infrastructure — just NuGet packages that work. Search, storage and chunking run entirely in your process; the embedding and reranking packages let you reach a hosted API when you want one, but never oblige you to. ## Packages | Package | Description | |---------|-------------| | [VectorSharp.Storage](VectorSharp.Storage/README.md) | In-memory and disk-backed vector storage with SIMD-optimized cosine similarity | -| [VectorSharp.Embedding](VectorSharp.Embedding/README.md) | Channel-based embedding service with configurable parallelism | +| [VectorSharp.Embedding](VectorSharp.Embedding/README.md) | Channel-based embedding service with configurable parallelism, request batching and token usage reporting | | [VectorSharp.Embedding.NomicEmbed](VectorSharp.Embedding.NomicEmbed/README.md) | Nomic Embed Text v1.5 model for local inference (768-dim, 8192 token context) | -| [VectorSharp.Chunking](VectorSharp.Chunking/README.md) | Streaming text chunker with predefined formats for Markdown and C# | +| [VectorSharp.Chunking](VectorSharp.Chunking/README.md) | Streaming text chunker with a hard token bound, position-carrying chunks and predefined formats for Markdown, C#, JavaScript/TypeScript, HTML, CSS, Python and plain text | +| [VectorSharp.Reranking](VectorSharp.Reranking/README.md) | Reranking abstractions for the second stage after a vector or hybrid search | ## Quick Start diff --git a/VectorSharp.Chunking.Tests/ChunkReaderFormatTests.cs b/VectorSharp.Chunking.Tests/ChunkReaderFormatTests.cs index 7d96198..d85c1ab 100644 --- a/VectorSharp.Chunking.Tests/ChunkReaderFormatTests.cs +++ b/VectorSharp.Chunking.Tests/ChunkReaderFormatTests.cs @@ -2,33 +2,6 @@ namespace VectorSharp.Chunking.Tests { public class ChunkReaderFormatTests { - private static StreamReader ReaderFrom(string text) - { - MemoryStream stream = new MemoryStream(System.Text.Encoding.UTF8.GetBytes(text)); - return new StreamReader(stream); - } - - private static async Task> ReadAllChunks(ChunkReader reader) - { - List chunks = new List(); - await foreach (string chunk in reader.ReadAllAsync()) - { - chunks.Add(chunk); - } - return chunks; - } - - private static ChunkReader CreateReader(string input, IReadOnlyList breakStrings, IReadOnlyList stopSignals, int maxTokens = 10) - { - StreamReader streamReader = ReaderFrom(input); - return ChunkReader.Create(streamReader, TokenCounter.CountWords, new ChunkReaderOptions - { - MaxTokensPerChunk = maxTokens, - BreakStrings = breakStrings, - StopSignals = stopSignals - }); - } - #region PlainText [Fact] diff --git a/VectorSharp.Chunking.Tests/ChunkReaderPositionTests.cs b/VectorSharp.Chunking.Tests/ChunkReaderPositionTests.cs new file mode 100644 index 0000000..4f9dea3 --- /dev/null +++ b/VectorSharp.Chunking.Tests/ChunkReaderPositionTests.cs @@ -0,0 +1,275 @@ +using System.Text; + +namespace VectorSharp.Chunking.Tests +{ + /// + /// What adds over : + /// each chunk's position in the input and its token count. Everything the two enumerations + /// share — which readers Create accepts, the single-pass rule, cancellation — is in + /// , since it is not about position. + /// + public class ChunkReaderPositionTests + { + /// + /// Asserts the two properties that make offsets trustworthy: the chunks concatenate back + /// into the input exactly, and every chunk's offsets point at that chunk's own text in + /// the original. The checks are stated as the contract + /// the type documents — it meets the next chunk's start, and the last one meets the end + /// of the input — rather than by repeating the property's own definition. + /// + private static void AssertLosslessWithOffsets(string original, List chunks) + { + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + + long expectedOffset = 0; + for (int i = 0; i < chunks.Count; i++) + { + TextChunk chunk = chunks[i]; + + Assert.Equal(expectedOffset, chunk.StartOffset); + Assert.Equal(chunk.Text, original.Substring((int)chunk.StartOffset, chunk.Text.Length)); + + if (i + 1 < chunks.Count) + Assert.Equal(chunks[i + 1].StartOffset, chunk.EndOffset); + + expectedOffset += chunk.Text.Length; + } + + Assert.Equal(original.Length, expectedOffset); + + if (chunks.Count > 0) + Assert.Equal(original.Length, chunks[^1].EndOffset); + } + + #region ReadChunksAsync (basic behavior) + + [Fact] + public async Task ReadChunksAsync_EmptyInput_YieldsNoChunks() + { + using StringReader stringReader = new StringReader(""); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Empty(chunks); + } + + [Fact] + public async Task ReadChunksAsync_NoChunks_YieldsNullFromFirstOrDefaultRatherThanAnEmptyChunk() + { + // Why TextChunk is a class: as a struct this would hand back a chunk whose Text is + // null despite being non-nullable, and whose EndOffset throws. Asserted through + // behaviour so that changing the type back reddens a test rather than passing quietly. + using StringReader stringReader = new StringReader(""); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Null(chunks.FirstOrDefault()); + } + + [Fact] + public async Task ReadChunksAsync_FirstChunk_StartsAtZero() + { + using StringReader stringReader = new StringReader("one. two. three. four. five. six."); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 2, BreakStrings = [". "], StopSignals = [] }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.NotEmpty(chunks); + Assert.Equal(0, chunks[0].StartOffset); + } + + [Fact] + public async Task ReadChunksAsync_YieldsSameTextAsReadAllAsync() + { + string original = "# Main Title\n\nIntro paragraph with multiple sentences. This is the second one! Is this a question?\n\n## Sub Section\n\n1. First item\n- Second item\n+ Third item\n\n```\ncode block\n```\n\n### Deeper\n\nFinal text.\n"; + + using StringReader chunkSource = new StringReader(original); + List chunks = await ReadAllTextChunks( + ChunkReader.Create(chunkSource, TokenCounter.CountWords, MarkdownOptions(5))); + + using StringReader stringSource = new StringReader(original); + List strings = await ReadAllChunks( + ChunkReader.Create(stringSource, TokenCounter.CountWords, MarkdownOptions(5))); + + Assert.Equal(strings, chunks.Select(chunk => chunk.Text)); + } + + #endregion + + #region ReadChunksAsync (lossless reconstruction and offsets) + + [Fact] + public async Task ReadChunksAsync_MarkdownContent_OffsetsAreRunningSumOfLengths() + { + using StringReader stringReader = new StringReader(MarkdownSample); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, MarkdownOptions(5)); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Equal(5, chunks.Count); + AssertLosslessWithOffsets(MarkdownSample, chunks); + } + + [Fact] + public async Task ReadChunksAsync_CSharpContent_OffsetsAreRunningSumOfLengths() + { + string original = "namespace Foo\n{\n public class Bar\n {\n /// \n /// Does something.\n /// \n public void Method()\n {\n Console.WriteLine(\"hello\");\n }\n }\n}\n"; + using StringReader stringReader = new StringReader(original); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions + { + MaxTokensPerChunk = 5, + BreakStrings = BreakStrings.CSharp, + StopSignals = StopSignals.CSharp + }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Equal(5, chunks.Count); + AssertLosslessWithOffsets(original, chunks); + } + + [Fact] + public async Task ReadChunksAsync_SegmentCutAtTheLimit_OffsetsStayCorrect() + { + // Neither line contains a break point, so each arrives over the limit and is cut. The + // offsets have to describe the pieces, not the segments they were cut from — a cut + // that left every piece carrying its segment's offset would still reconstruct the + // input and still point at the wrong place in it. + string original = "one two three four five six seven eight nine ten\neleven twelve thirteen fourteen\n"; + using StringReader stringReader = new StringReader(original); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 3, BreakStrings = ["\n"], StopSignals = [] }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.True(chunks.Count > 2); + AssertLosslessWithOffsets(original, chunks); + Assert.All(chunks, chunk => Assert.True(chunk.TokenCount <= 3)); + } + + [Fact] + public async Task ReadChunksAsync_StopSignals_OffsetsStayCorrect() + { + string original = "intro text\n# Heading\nmore text\n# Second Heading\ntail text"; + using StringReader stringReader = new StringReader(original); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 100, BreakStrings = ["\n"], StopSignals = ["# "] }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Equal(3, chunks.Count); + AssertLosslessWithOffsets(original, chunks); + } + + [Fact] + public async Task ReadChunksAsync_LongInput_OffsetsStayCorrectAcrossManyChunks() + { + StringBuilder builder = new StringBuilder(); + for (int i = 0; i < 200; i++) + { + builder.Append($"Sentence number {i} of the document. "); + } + string original = builder.ToString(); + + using StringReader stringReader = new StringReader(original); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 12, BreakStrings = [". "], StopSignals = [] }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Equal(100, chunks.Count); + AssertLosslessWithOffsets(original, chunks); + } + + [Fact] + public async Task ReadChunksAsync_NonAsciiInput_OffsetsAreUtf16CodeUnits() + { + // Accented letters, an astral-plane emoji (one surrogate pair, two code units) and a + // ZWJ sequence: the offsets are documented as UTF-16 code units, so they have to + // track string.Length rather than characters a reader would count by eye. + string original = "café ☕ naïve\n👨‍👩‍👧 family emoji 🎉\nlast line here\n"; + using StringReader stringReader = new StringReader(original); + // Four, so that the longest line fits and no line is cut. What offsets do across a cut + // is ChunkReaderTokenBoundTests' subject; this one is about the unit they are in. + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 4, BreakStrings = ["\n"], StopSignals = [] }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Equal(3, chunks.Count); + AssertLosslessWithOffsets(original, chunks); + + // The second chunk starts after "café ☕ naïve\n", which is 13 characters but the + // emoji in it costs nothing extra — the pair in the second line is what would show + // up as a discrepancy if offsets counted anything other than code units. + Assert.Equal(13, chunks[1].StartOffset); + Assert.Equal(original.IndexOf("last line here", StringComparison.Ordinal), chunks[2].StartOffset); + } + + #endregion + + #region ReadChunksAsync (token counts) + + [Fact] + public async Task ReadChunksAsync_TokenCount_IsTheCountOfTheChunkText() + { + using StringReader stringReader = new StringReader("one two three. four five six. seven eight nine."); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 3, BreakStrings = [". "], StopSignals = [] }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Equal(3, chunks.Count); + Assert.Equal(3, chunks[0].TokenCount); + Assert.Equal(3, chunks[1].TokenCount); + Assert.Equal(3, chunks[2].TokenCount); + } + + [Fact] + public async Task ReadChunksAsync_TokenCount_MatchesTheSuppliedCounter() + { + using StringReader stringReader = new StringReader(MarkdownSample); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, MarkdownOptions(10)); + + List chunks = await ReadAllTextChunks(reader); + + Assert.NotEmpty(chunks); + foreach (TextChunk chunk in chunks) + { + Assert.Equal(TokenCounter.CountWords(chunk.Text), chunk.TokenCount); + } + } + + [Fact] + public async Task ReadChunksAsync_MergedChunks_CountTheTokenCounterOnceEachTime() + { + // The counter can be an actual tokenizer, so a chunk that grows by one segment must + // carry the count it already computed forward rather than recounting the same text. + int calls = 0; + int CountingCounter(string text) + { + calls++; + return TokenCounter.CountWords(text); + } + + using StringReader stringReader = new StringReader("a. b. c. d. e. f. g. h. i. j."); + ChunkReader reader = ChunkReader.Create(stringReader, CountingCounter, + new ChunkReaderOptions { MaxTokensPerChunk = 3, BreakStrings = [". "], StopSignals = [] }); + + List chunks = await ReadAllTextChunks(reader); + + // The input holds 10 segments, and each one is counted exactly once: either when it + // starts a chunk or when it is merged into the chunk being built. Nothing is counted + // twice, so the total lands on the segment count rather than above it. + Assert.Equal(4, chunks.Count); + Assert.Equal(10, calls); + } + + #endregion + + } +} diff --git a/VectorSharp.Chunking.Tests/ChunkReaderTests.cs b/VectorSharp.Chunking.Tests/ChunkReaderTests.cs index 286652b..69879c8 100644 --- a/VectorSharp.Chunking.Tests/ChunkReaderTests.cs +++ b/VectorSharp.Chunking.Tests/ChunkReaderTests.cs @@ -2,29 +2,53 @@ namespace VectorSharp.Chunking.Tests { public class ChunkReaderTests { - private static StreamReader ReaderFrom(string text) + #region Create (validation) + + // Both overloads are named explicitly: a bare null literal would be ambiguous between + // them, and each has to reject null on its own. + [Fact] + public void Create_NullTextReader_Throws() { - MemoryStream stream = new MemoryStream(System.Text.Encoding.UTF8.GetBytes(text)); - return new StreamReader(stream); + Assert.Throws(() => + ChunkReader.Create((TextReader)null!, TokenCounter.CountWords)); } - private static async Task> ReadAllChunks(ChunkReader reader) + [Fact] + public void Create_NullStreamReader_Throws() { - List chunks = new List(); - await foreach (string chunk in reader.ReadAllAsync()) - { - chunks.Add(chunk); - } - return chunks; + Assert.Throws(() => + ChunkReader.Create((StreamReader)null!, TokenCounter.CountWords)); } - #region Create (validation) + /// + /// Looks for a Create overload whose first parameter is exactly . + /// The parameter types are compared by identity rather than through + /// , which resolves like the compiler does and + /// would happily answer with the TextReader overload for any type derived from it. + /// + private static bool HasCreateOverloadTaking(Type readerType) + { + Type[] expected = [readerType, typeof(Func), typeof(ChunkReaderOptions)]; + + return typeof(ChunkReader) + .GetMethods(System.Reflection.BindingFlags.Public | System.Reflection.BindingFlags.Static) + .Where(method => method.Name == nameof(ChunkReader.Create)) + .Any(method => method.GetParameters().Select(parameter => parameter.ParameterType).SequenceEqual(expected)); + } + // Asserted through reflection because no call site can detect this: source that passes a + // StreamReader compiles whether or not the StreamReader overload exists, silently binding + // to the TextReader one instead. The already published signature only survives in IL. [Fact] - public void Create_NullReader_Throws() + public void Create_KeepsTheStreamReaderOverloadEarlierAssembliesBindTo() { - Assert.Throws(() => - ChunkReader.Create(null!, TokenCounter.CountWords)); + Assert.True(HasCreateOverloadTaking(typeof(StreamReader))); + } + + [Fact] + public void Create_HasTheWidenedTextReaderOverload() + { + Assert.True(HasCreateOverloadTaking(typeof(TextReader))); } [Fact] @@ -121,7 +145,7 @@ public async Task ReadAllAsync_SegmentsExceedingLimit_SplitsIntoMultipleChunks() List chunks = await ReadAllChunks(reader); Assert.True(chunks.Count > 1); - Assert.All(chunks, chunk => Assert.True(TokenCounter.CountWords(chunk) <= 4)); // small overshoot possible for single segments + Assert.All(chunks, chunk => Assert.True(TokenCounter.CountWords(chunk) <= 4)); } [Fact] @@ -309,28 +333,31 @@ public async Task ReadAllAsync_AllChunksWithinTokenLimit() List chunks = await ReadAllChunks(reader); - // Each chunk should be within the limit (except single oversized segments) Assert.True(chunks.Count > 1); + Assert.All(chunks, chunk => Assert.True(TokenCounter.CountWords(chunk) <= 3)); } [Fact] - public async Task ReadAllAsync_SingleOversizedSegment_ReturnedAsIs() + public async Task ReadAllAsync_SingleOversizedSegment_IsCutAtTheLimit() { - // A segment with no break points that exceeds the token limit - using StreamReader streamReader = ReaderFrom("one two three four five six seven eight nine ten"); + // A segment with no break point in it is cut at the limit rather than coming back + // whole, and the pieces still reproduce it exactly. + string input = "one two three four five six seven eight nine ten"; + using StreamReader streamReader = ReaderFrom(input); ChunkReader reader = ChunkReader.Create(streamReader, TokenCounter.CountWords, new ChunkReaderOptions { MaxTokensPerChunk = 3, BreakStrings = ["\n"], StopSignals = [] }); List chunks = await ReadAllChunks(reader); - // Should return the whole thing as one chunk since there are no break points - Assert.Single(chunks); - Assert.Equal("one two three four five six seven eight nine ten", chunks[0]); + Assert.True(chunks.Count > 1); + Assert.All(chunks, chunk => Assert.True(TokenCounter.CountWords(chunk) <= 3, + $"'{chunk}' counts {TokenCounter.CountWords(chunk)} tokens against a limit of 3.")); + Assert.Equal(input, string.Join("", chunks)); } #endregion - #region ReadAllAsync (cancellation) + #region Cancellation [Fact] public async Task ReadAllAsync_CancelledToken_ThrowsOperationCancelled() @@ -349,6 +376,133 @@ await Assert.ThrowsAnyAsync(async () => }); } + [Fact] + public async Task ReadChunksAsync_CancelledToken_ThrowsOperationCancelled() + { + using StringReader stringReader = new StringReader("hello world this is some text\nmore text here\n"); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 2, BreakStrings = ["\n"], StopSignals = [] }); + CancellationToken cancelled = new CancellationToken(true); + + await Assert.ThrowsAnyAsync(async () => + { + await foreach (TextChunk chunk in reader.ReadChunksAsync(cancelled)) + { + // Should not reach here + } + }); + } + + [Fact] + public async Task ReadChunksAsync_CancelledMidEnumeration_StopsAtTheNextChunk() + { + // The case a caller actually hits, and the one the loop's own + // ThrowIfCancellationRequested exists for. A token cancelled before the call throws + // out of the segment reader's guard before the loop is entered at all, so the two + // tests reach different checks despite reading almost the same. + using StringReader stringReader = new StringReader("aa bb\ncc dd\nee ff\ngg hh\nii jj\nkk ll\n"); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 2, BreakStrings = ["\n"], StopSignals = [] }); + + using CancellationTokenSource cancellationSource = new CancellationTokenSource(); + List received = new List(); + + await Assert.ThrowsAnyAsync(async () => + { + await foreach (TextChunk chunk in reader.ReadChunksAsync(cancellationSource.Token)) + { + received.Add(chunk); + + if (received.Count == 2) + await cancellationSource.CancelAsync(); + } + }); + + // Everything up to the cancellation is still delivered: the chunks already yielded + // are valid, and a caller that stops mid-stream keeps what it read. + Assert.Equal(2, received.Count); + Assert.Equal(0, received[0].StartOffset); + Assert.Equal(received[0].EndOffset, received[1].StartOffset); + } + + #endregion + + #region Create (reader types) + + [Fact] + public async Task Create_StringReader_ReadsChunks() + { + using StringReader stringReader = new StringReader("hello world"); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 10, BreakStrings = ["\n"], StopSignals = [] }); + + List chunks = await ReadAllTextChunks(reader); + + Assert.Single(chunks); + Assert.Equal("hello world", chunks[0].Text); + } + + [Fact] + public async Task Create_StringReader_MatchesStreamReaderOutput() + { + using StringReader stringReader = new StringReader(MarkdownSample); + List fromString = await ReadAllTextChunks( + ChunkReader.Create(stringReader, TokenCounter.CountWords, MarkdownOptions(10))); + + using StreamReader streamReader = ReaderFrom(MarkdownSample); + List fromStream = await ReadAllTextChunks( + ChunkReader.Create(streamReader, TokenCounter.CountWords, MarkdownOptions(10))); + + Assert.Equal(fromStream.Select(chunk => chunk.Text), fromString.Select(chunk => chunk.Text)); + Assert.Equal(fromStream.Select(chunk => chunk.StartOffset), fromString.Select(chunk => chunk.StartOffset)); + } + + #endregion + + #region Single pass + + [Fact] + public async Task ReadChunksAsync_SecondEnumeration_Throws() + { + using StringReader stringReader = new StringReader("one. two. three. four. five. six."); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 2, BreakStrings = [". "], StopSignals = [] }); + + await ReadAllTextChunks(reader); + + await Assert.ThrowsAsync(() => ReadAllTextChunks(reader)); + } + + [Fact] + public async Task ReadChunksAsync_AbandonedThenReEnumerated_Throws() + { + // The abandoned pass has already consumed a segment it never yielded. Continuing + // would drop that text and mislabel everything after it, so the second pass is + // refused rather than left to produce chunks that reconstruct nothing. + using StringReader stringReader = new StringReader("aa bb cc\ndd ee ff\ngg hh ii\njj kk ll\n"); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 3, BreakStrings = ["\n"], StopSignals = [] }); + + await foreach (TextChunk chunk in reader.ReadChunksAsync()) + { + break; + } + + await Assert.ThrowsAsync(() => ReadAllTextChunks(reader)); + } + + [Fact] + public async Task ReadAllAsync_AfterReadChunksAsync_Throws() + { + using StringReader stringReader = new StringReader("one. two. three. four."); + ChunkReader reader = ChunkReader.Create(stringReader, TokenCounter.CountWords, + new ChunkReaderOptions { MaxTokensPerChunk = 2, BreakStrings = [". "], StopSignals = [] }); + + await ReadAllTextChunks(reader); + + await Assert.ThrowsAsync(() => ReadAllChunks(reader)); + } + #endregion #region ReadAllAsync (defaults) diff --git a/VectorSharp.Chunking.Tests/ChunkReaderTokenBoundTests.cs b/VectorSharp.Chunking.Tests/ChunkReaderTokenBoundTests.cs new file mode 100644 index 0000000..dfd0bfe --- /dev/null +++ b/VectorSharp.Chunking.Tests/ChunkReaderTokenBoundTests.cs @@ -0,0 +1,475 @@ +namespace VectorSharp.Chunking.Tests +{ + /// + /// as a bound rather than a target: text + /// running past the limit with no break point in it is cut, and the two guarantees that make + /// offsets usable have to survive the cut. + /// + public class ChunkReaderTokenBoundTests + { + /// + /// Text with no break point anywhere in it, long enough to need cutting many times over. + /// Joined with spaces: the break strings below are newlines only, so this is one segment + /// however long it gets, while still being many tokens to a word-counting tokenizer. + /// + private static string UnbreakableText(int words) + { + return string.Join(" ", Enumerable.Range(0, words).Select(i => $"word{i}")); + } + + private static ChunkReader ReaderOver(string input, int maxTokens, Func? countTokens = null) + { + return CreateReader(input, breakStrings: ["\n"], stopSignals: [], maxTokens, countTokens); + } + + #region The bound holds + + [Fact] + public async Task ReadChunksAsync_TextWithNoBreakPoint_EveryChunkIsWithinTheLimit() + { + string original = UnbreakableText(500); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 10)); + + // Pinned rather than "more than one": one chunk per character also satisfies "more + // than one", and that is exactly the degenerate output a broken search produces. + Assert.Equal(50, chunks.Count); + Assert.All(chunks, chunk => Assert.True( + TokenCounter.CountWords(chunk.Text) <= 10, + $"chunk at {chunk.StartOffset} counts {TokenCounter.CountWords(chunk.Text)} tokens against a limit of 10.")); + } + + [Fact] + public async Task ReadChunksAsync_TextWithNoBreakPoint_CutsAtTheLimitNotBelowIt() + { + // A bound that held by cutting after every character would satisfy the assertion above + // and be useless. With one token per character and a limit of 10, every chunk but the + // last has to be exactly 10 characters. + string original = new string('x', 95); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 10, TokenCounter.CountCharacters)); + + Assert.Equal(10, chunks.Count); + Assert.All(chunks.Take(9), chunk => Assert.Equal(10, chunk.Text.Length)); + Assert.Equal(5, chunks[^1].Text.Length); + } + + [Fact] + public async Task ReadChunksAsync_TokenCountIsCountedForThePieceNotDividedOutOfTheWhole() + { + string original = new string('x', 95); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 10, TokenCounter.CountCharacters)); + + Assert.All(chunks, chunk => Assert.Equal(TokenCounter.CountCharacters(chunk.Text), chunk.TokenCount)); + } + + [Fact] + public async Task ReadChunksAsync_SegmentExactlyAtTheLimit_IsNotCut() + { + // The boundary between "full" and "over". A segment exactly at the limit is already + // correct, and cutting it would change output for text that never needed it. + string original = new string('x', 10); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 10, TokenCounter.CountCharacters)); + + Assert.Single(chunks); + Assert.Equal(original, chunks[0].Text); + } + + [Fact] + public async Task ReadChunksAsync_SegmentsWithinTheLimit_AreNotTouchedByTheCuttingPath() + { + // Text whose segments already fit chunks by accumulation alone. Cutting must not reach + // it: the search stops at the token limit, so text put through it needlessly comes + // back in pieces the size of the limit instead of the size its break strings implied. + string original = "alpha beta\ngamma delta\nepsilon zeta\n"; + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 100)); + + Assert.Single(chunks); + Assert.Equal(original, chunks[0].Text); + } + + #endregion + + #region What the cut must not break + + [Fact] + public async Task ReadChunksAsync_CutText_StillReconstructsExactly() + { + string original = UnbreakableText(500); + + // Seven is simply well under the word count of the text, so the segment is cut many + // times rather than once — a guarantee that held only for a single cut would not be + // one. + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 7)); + + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + } + + [Fact] + public async Task ReadChunksAsync_CutText_OffsetsPointAtEachPiecesOwnText() + { + string original = UnbreakableText(500); + + // Seven is simply well under the word count of the text, so the segment is cut many + // times rather than once — a guarantee that held only for a single cut would not be + // one. + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 7)); + + long expectedOffset = 0; + foreach (TextChunk chunk in chunks) + { + Assert.Equal(expectedOffset, chunk.StartOffset); + Assert.Equal(chunk.Text, original.Substring((int)chunk.StartOffset, chunk.Text.Length)); + expectedOffset = chunk.EndOffset; + } + + Assert.Equal(original.Length, expectedOffset); + } + + [Fact] + public async Task ReadChunksAsync_CutTextFollowedByMore_OffsetsStayCorrectAcrossTheJoin() + { + // A cut segment, then ordinary segments after it. The reader's own cursor has to stay + // in step with text that was consumed as one segment and emitted as several. + // + // Nine, so that the two tails together are under the limit and would merge into one + // chunk if the reader carried on normally after the cut. A cursor left pointing at the + // end of the segment rather than the end of the last piece shows up as an offset gap + // right at that join. + string original = UnbreakableText(200) + "\nshort tail\nanother tail\n"; + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 9)); + + long expectedOffset = 0; + foreach (TextChunk chunk in chunks) + { + Assert.Equal(expectedOffset, chunk.StartOffset); + expectedOffset = chunk.EndOffset; + } + + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + Assert.Equal(original.Length, expectedOffset); + } + + [Fact] + public async Task ReadChunksAsync_CutText_NeverYieldsAnEmptyChunk() + { + string original = UnbreakableText(300); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 5)); + + Assert.All(chunks, chunk => Assert.NotEmpty(chunk.Text)); + } + + [Theory] + [InlineData(5, 100)] + [InlineData(6, 67)] + [InlineData(7, 67)] + public async Task ReadChunksAsync_TextWithSurrogatePairs_NeverCutsAPairInHalf(int maxTokens, int expectedChunks) + { + // Each emoji is two UTF-16 code units that are one character. A cut landing between + // them produces two chunks that are not text, and a lone surrogate is what reaches the + // caller's tokenizer and their embedding API. + // + // The odd limits are the ones that matter, and are why this is a theory. With one + // token per code unit, an odd limit puts the largest prefix that fits exactly halfway + // through a pair, so the cut has to be pulled back deliberately. An even limit lands + // between pairs on its own and would hold however the search was written. + string original = string.Concat(Enumerable.Repeat("\U0001F600", 200)); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens, TokenCounter.CountCharacters)); + + // Pinned per case: at an odd limit every chunk is one code unit shorter than the limit + // allows, because the cut has to step back off the pair, so the counts differ between + // 5 and 6 by more than the limits do. + Assert.Equal(expectedChunks, chunks.Count); + Assert.All(chunks, chunk => + { + Assert.False(char.IsLowSurrogate(chunk.Text[0]), + $"a chunk of {chunk.Text.Length} code units starts with the second half of a surrogate pair."); + Assert.False(char.IsHighSurrogate(chunk.Text[^1]), + $"a chunk of {chunk.Text.Length} code units ends with the first half of a surrogate pair."); + }); + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + } + + [Fact] + public async Task ReadChunksAsync_CancelledWhileCuttingOneSegment_StopsWithoutFinishingIt() + { + // Cutting one long segment can be hundreds of pieces and hundreds of calls into the + // caller's tokenizer, all inside a single pass of the reader's own loop. A token + // checked only by that loop is not checked again until the whole segment is done, so a + // caller who has stopped waiting waits for all of it. + string original = new string('x', 200000); + using CancellationTokenSource cancellationSource = new CancellationTokenSource(); + int received = 0; + + await Assert.ThrowsAnyAsync(async () => + { + await foreach (TextChunk chunk in ReaderOver(original, maxTokens: 5, TokenCounter.CountCharacters).ReadChunksAsync(cancellationSource.Token)) + { + received++; + + if (received == 3) + await cancellationSource.CancelAsync(); + } + }); + + // Three delivered, then it stops. Without a check inside the cut it would run to the + // end of the segment, which is 40000 pieces. + Assert.Equal(3, received); + } + + #endregion + + #region Alongside the reader's other rules + + [Fact] + public async Task ReadChunksAsync_SegmentThatIsBothAStopSignalAndOverTheLimit_IsCutAndStillStartsItsOwnChunk() + { + // A stop signal forces a new chunk before the limit is ever consulted, so an oversized + // one reaches the cut by a different route than an ordinary segment does. + string oversizedHeading = "# " + UnbreakableText(30); + string original = $"intro\n{oversizedHeading}\ntail\n"; + + List chunks = await ReadAllTextChunks( + CreateReader(original, breakStrings: ["\n"], stopSignals: ["#"], maxTokens: 5)); + + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + Assert.All(chunks, chunk => Assert.True(chunk.TokenCount <= 5)); + + long expectedOffset = 0; + foreach (TextChunk chunk in chunks) + { + Assert.Equal(expectedOffset, chunk.StartOffset); + expectedOffset = chunk.EndOffset; + } + + // The heading still begins a chunk of its own rather than being appended to "intro". + Assert.Contains(chunks, chunk => chunk.Text.StartsWith("# ", StringComparison.Ordinal)); + } + + [Fact] + public async Task ReadChunksAsync_DefaultMarkdownOptions_CutsAnOversizedLineInAFencedBlock() + { + // The configuration a real caller uses. Every other test here reduces the break + // strings to a newline, which is not what the package ships with, and a fenced code + // block is where unbroken text actually turns up in Markdown. + string original = $"# Title\n\nSome intro text.\n\n```\n{UnbreakableText(60)}\n```\n\nClosing text.\n"; + + List chunks = await ReadAllTextChunks( + ChunkReader.Create(new StringReader(original), TokenCounter.CountWords, MarkdownOptions(8))); + + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + Assert.All(chunks, chunk => Assert.True( + chunk.TokenCount <= 8, + $"chunk at {chunk.StartOffset} counts {chunk.TokenCount} tokens against a limit of 8.")); + + long expectedOffset = 0; + foreach (TextChunk chunk in chunks) + { + Assert.Equal(expectedOffset, chunk.StartOffset); + expectedOffset = chunk.EndOffset; + } + } + + [Fact] + public async Task ReadChunksAsync_SegmentsAccumulatingToExactlyTheLimit_AreNotCut() + { + // The same "already fits" guard the single-segment case reaches, arrived at by + // accumulation instead: five segments of two characters each land on the limit + // together, and the chunk they form must come back whole. + string original = "ab\ncd\ncf\ngh\nij\n"; + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 15, TokenCounter.CountCharacters)); + + Assert.Single(chunks); + Assert.Equal(original, chunks[0].Text); + Assert.Equal(15, chunks[0].TokenCount); + } + + [Fact] + public async Task ReadChunksAsync_SubwordCounter_CutsMidWordExactlyAsTheReadmeSays() + { + // The README warns that a cut lands at the token limit rather than at a boundary, and + // prints this as the example. Pinned here so the warning cannot quietly stop being + // true — a reader who tries the sample and gets something else trusts neither. + // + // Four characters to a token approximates a subword tokenizer closely enough to show + // the behaviour: a word counter cannot, because adding a character to a word never + // raises its count. + Func subwordCounter = text => (text.Length + 3) / 4; + string original = "the quick brown fox jumps"; + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 3, subwordCounter)); + + Assert.Equal(["the quick br", "own fox jump", "s"], chunks.Select(chunk => chunk.Text)); + } + + #endregion + + #region Degenerate counters + + [Fact] + public async Task ReadChunksAsync_SingleCharacterAlreadyOverTheLimit_EmitsItAloneRatherThanLooping() + { + // The one case the docs admit can come back over the limit. What matters is that it + // terminates, reconstructs, and cuts as small as it possibly can rather than giving up + // and emitting the whole segment. + string original = new string('x', 20); + + List chunks = await ReadAllTextChunks( + ReaderOver(original, maxTokens: 1, text => text.Length * 5)); + + Assert.Equal(20, chunks.Count); + Assert.All(chunks, chunk => Assert.Equal(1, chunk.Text.Length)); + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + } + + [Fact] + public async Task ReadChunksAsync_SurrogatePairAlreadyOverTheLimit_EmitsThePairWholeNotItsHalves() + { + string original = string.Concat(Enumerable.Repeat("\U0001F600", 8)); + + List chunks = await ReadAllTextChunks( + ReaderOver(original, maxTokens: 1, text => text.Length * 5)); + + Assert.Equal(8, chunks.Count); + Assert.All(chunks, chunk => Assert.Equal("\U0001F600", chunk.Text)); + } + + [Fact] + public async Task ReadChunksAsync_CounterThatSometimesShrinksAsTextGrows_StillNeverExceedsTheLimit() + { + // Real tokenizers are not perfectly monotonic in the length of their input: a + // byte-pair encoder can charge fewer tokens for "the" than for "th", because the + // longer string merges into one token. The search must not infer a count for a length + // it never measured, or a dip like that becomes an over-limit chunk. + // + // This counter charges one token per character and then refunds four whenever the + // length is a multiple of ten, so counts dip as the text grows, exactly where a + // doubling search lands. + Func dippingCounter = text => text.Length % 10 == 0 ? text.Length - 4 : text.Length; + string original = new string('x', 500); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 25, dippingCounter)); + + Assert.True(chunks.Count > 1); + Assert.All(chunks, chunk => Assert.True( + dippingCounter(chunk.Text) <= 25, + $"a chunk of {chunk.Text.Length} characters counts {dippingCounter(chunk.Text)} tokens against a limit of 25.")); + Assert.All(chunks, chunk => Assert.Equal(dippingCounter(chunk.Text), chunk.TokenCount)); + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + } + + [Fact] + public async Task ReadChunksAsync_ChunkThatFitsWholeButNotInPieces_IsNotCutUp() + { + // The splitter answers "does this fit?" from the count already measured for the whole + // chunk, not by searching upward from one character. It matters for a counter that + // charges more for a short piece than for the whole: here the text counts exactly the + // limit, while any piece of it counts far over. Searching would take the piece at its + // word and hand back one chunk per character, all of them over the limit, for text + // that was correct as it stood. + Func wholeIsCheaperCounter = text => text.Length == 20 ? 10 : 99; + string original = new string('x', 20); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 10, wholeIsCheaperCounter)); + + Assert.Single(chunks); + Assert.Equal(original, chunks[0].Text); + Assert.Equal(10, chunks[0].TokenCount); + } + + [Fact] + public async Task ReadChunksAsync_CounterThatChargesMoreForOneCharacterThanTheWholeLimit_TerminatesAndReconstructs() + { + // The documented exception, stated as what it is rather than as a bound that holds. + // The search stops probing upward at the first piece that does not fit, so a counter + // that reports over the limit for one character and under it for forty is not + // something it will discover — it emits single characters, over the limit, because the + // only smaller cut would divide a character in half. What has to hold even here is + // that it terminates and loses nothing. + Func perverseCounter = text => text.Length % 2 == 0 ? 1 : 99; + string original = new string('x', 41); + + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 50, perverseCounter)); + + Assert.Equal(41, chunks.Count); + Assert.All(chunks, chunk => Assert.Equal(1, chunk.Text.Length)); + Assert.All(chunks, chunk => Assert.Equal(perverseCounter(chunk.Text), chunk.TokenCount)); + Assert.Equal(original, string.Join("", chunks.Select(chunk => chunk.Text))); + + long expectedOffset = 0; + foreach (TextChunk chunk in chunks) + { + Assert.Equal(expectedOffset, chunk.StartOffset); + expectedOffset = chunk.EndOffset; + } + } + + #endregion + + #region Cost + + [Fact] + public async Task ReadChunksAsync_LargeUnbreakableInput_DoesNotProbeQuadratically() + { + // The cut is found by measuring candidate prefixes, and every measurement copies the + // prefix it measures. Searching the whole remaining text for each piece would make one + // large input quadratic in its length, which trades an over-sized chunk for a stall. + // Bounding the characters handed to the counter is what rules that out. + string original = new string('x', 20000); + long charactersMeasured = 0; + + List chunks = await ReadAllTextChunks( + ReaderOver(original, maxTokens: 100, text => + { + charactersMeasured += text.Length; + return text.Length; + })); + + Assert.Equal(200, chunks.Count); + + // Stated as a multiple of the input rather than a number, because what is being ruled + // out is a shape: bounded probing stays proportional to the input, unbounded probing + // grows with its square. Searching the whole remainder for each piece measures around + // 113 times the input here and more as it grows; bounding it measures around 10 times, + // whatever the length. The multiple below sits well clear of both, so it fails on the + // shape rather than on a change of a few probes either way. + Assert.True(charactersMeasured < original.Length * 30L, + $"the token counter was handed {charactersMeasured} characters for a {original.Length} character input."); + } + + [Fact] + public async Task ReadChunksAsync_LargeUnbreakableInput_DoesNotCopyTheRemainderPerPiece() + { + // The other half of the same cost, and invisible to the probe count above: cutting a + // piece off the front by re-slicing what is left copies everything behind it, so the + // work is quadratic in the chunk even though every probe is small. Measured as + // allocation because that is what re-slicing spends. + // + // The enumeration below never actually suspends — a StringReader completes + // synchronously — so it stays on this thread and the counter sees all of it. + string original = new string('x', 20000); + + long before = GC.GetAllocatedBytesForCurrentThread(); + List chunks = await ReadAllTextChunks(ReaderOver(original, maxTokens: 100, TokenCounter.CountCharacters)); + long allocated = GC.GetAllocatedBytesForCurrentThread() - before; + + Assert.Equal(200, chunks.Count); + + // Two bytes per character. Measured, the pieces and probes together come to about 14 + // times the input; re-slicing takes it to about 114 times, and further as the input + // grows. The multiple below sits between them with room on both sides, so it fails on + // the shape rather than on a change of a few allocations either way. + Assert.True(allocated < original.Length * 2L * 40, + $"cutting a {original.Length} character chunk allocated {allocated} bytes."); + } + + #endregion + } +} diff --git a/VectorSharp.Chunking.Tests/SegmentReaderTests.cs b/VectorSharp.Chunking.Tests/SegmentReaderTests.cs index dfcb473..66b6051 100644 --- a/VectorSharp.Chunking.Tests/SegmentReaderTests.cs +++ b/VectorSharp.Chunking.Tests/SegmentReaderTests.cs @@ -2,12 +2,6 @@ namespace VectorSharp.Chunking.Tests { public class SegmentReaderTests { - private static StreamReader ReaderFrom(string text) - { - MemoryStream stream = new MemoryStream(System.Text.Encoding.UTF8.GetBytes(text)); - return new StreamReader(stream); - } - private static async Task> ReadAllSegments(SegmentReader reader) { List segments = new List(); @@ -95,6 +89,42 @@ public async Task ReadNextAsync_DoubleNewlineOverSingle() Assert.Equal("still second", segments[2]); } + [Fact] + public async Task LastSegmentStartOffset_TracksEachSegmentAsItIsReturned() + { + using StreamReader streamReader = ReaderFrom("one. two. three"); + SegmentReader reader = new SegmentReader(streamReader, [". "]); + List offsets = new List(); + + while (await reader.ReadNextAsync() != null) + { + offsets.Add(reader.LastSegmentStartOffset); + } + + Assert.Equal([0L, 5L, 10L], offsets); + } + + [Fact] + public async Task LastSegmentStartOffset_LookaheadPushedBack_IsNotCountedUntilItsSegmentIsReturned() + { + // The whole of what the property promises. Reading "hello\n\n" requires looking one + // character past the first "\n" to find out a longer break string follows, and the + // 'w' after "second\n" is read and pushed back the same way. A cursor advanced by what + // was read rather than by what was returned would put the second segment a character + // early and every one after it out by the same amount. + using StreamReader streamReader = ReaderFrom("hello\n\nsecond\nworld"); + SegmentReader reader = new SegmentReader(streamReader, ["\n", "\n\n"]); + List<(string Segment, long Offset)> read = new List<(string, long)>(); + + string? segment; + while ((segment = await reader.ReadNextAsync()) != null) + { + read.Add((segment, reader.LastSegmentStartOffset)); + } + + Assert.Equal([("hello\n\n", 0L), ("second\n", 7L), ("world", 14L)], read); + } + [Fact] public async Task ReadNextAsync_CustomBreakStrings_Work() { diff --git a/VectorSharp.Chunking.Tests/TestReaders.cs b/VectorSharp.Chunking.Tests/TestReaders.cs new file mode 100644 index 0000000..cfde077 --- /dev/null +++ b/VectorSharp.Chunking.Tests/TestReaders.cs @@ -0,0 +1,71 @@ +using System.Text; + +namespace VectorSharp.Chunking.Tests +{ + /// + /// Shared fixtures for the chunking tests: readers over in-memory text, a configured + /// , and a drain for each of its two enumerations. + /// Imported as a global static in the test project, so the members are callable unqualified + /// the way the per-file copies they replaced were. + /// + internal static class TestReaders + { + /// + /// A short document exercising every Markdown break string: headings, paragraphs and a + /// list. Shared because both chunk-reader test files read the same sample. + /// + internal const string MarkdownSample = "# Title\n\nFirst paragraph with text. More sentences here.\n\n## Section\n\n- item one\n- item two\n\nFinal paragraph.\n"; + + internal static ChunkReaderOptions MarkdownOptions(int maxTokens) => new ChunkReaderOptions + { + MaxTokensPerChunk = maxTokens, + BreakStrings = BreakStrings.Markdown, + StopSignals = StopSignals.Markdown + }; + + /// + /// Wraps text in a , which keeps a real stream and its decoding + /// in the picture. Tests specifically about which reader type is accepted use the + /// directly instead. + /// + internal static StreamReader ReaderFrom(string text) + { + MemoryStream stream = new MemoryStream(Encoding.UTF8.GetBytes(text)); + return new StreamReader(stream); + } + + /// The token counter to chunk by. Defaults to counting words, + /// which is what all but the token-bound tests want; those pass a counter that reaches the + /// limit at a position they can state exactly. + internal static ChunkReader CreateReader(string input, IReadOnlyList breakStrings, IReadOnlyList stopSignals, int maxTokens = 10, Func? countTokens = null) + { + StreamReader streamReader = ReaderFrom(input); + return ChunkReader.Create(streamReader, countTokens ?? TokenCounter.CountWords, new ChunkReaderOptions + { + MaxTokensPerChunk = maxTokens, + BreakStrings = breakStrings, + StopSignals = stopSignals + }); + } + + internal static async Task> ReadAllChunks(ChunkReader reader) + { + List chunks = new List(); + await foreach (string chunk in reader.ReadAllAsync()) + { + chunks.Add(chunk); + } + return chunks; + } + + internal static async Task> ReadAllTextChunks(ChunkReader reader) + { + List chunks = new List(); + await foreach (TextChunk chunk in reader.ReadChunksAsync()) + { + chunks.Add(chunk); + } + return chunks; + } + } +} diff --git a/VectorSharp.Chunking.Tests/TokenCounter.cs b/VectorSharp.Chunking.Tests/TokenCounter.cs index 159c366..02028c6 100644 --- a/VectorSharp.Chunking.Tests/TokenCounter.cs +++ b/VectorSharp.Chunking.Tests/TokenCounter.cs @@ -2,6 +2,14 @@ namespace VectorSharp.Chunking.Tests { internal static class TokenCounter { + /// + /// One token per character. A word counter can never cut mid-word — adding a character to + /// a word does not raise its count — so it cannot express where a limit falls inside a + /// run of text. This one reaches the limit at a position a test can state exactly, and is + /// the closer model of a real subword tokenizer. + /// + internal static int CountCharacters(string text) => text.Length; + internal static int CountWords(string text) { if (string.IsNullOrWhiteSpace(text)) diff --git a/VectorSharp.Chunking.Tests/VectorSharp.Chunking.Tests.csproj b/VectorSharp.Chunking.Tests/VectorSharp.Chunking.Tests.csproj index 447b5d9..691bb13 100644 --- a/VectorSharp.Chunking.Tests/VectorSharp.Chunking.Tests.csproj +++ b/VectorSharp.Chunking.Tests/VectorSharp.Chunking.Tests.csproj @@ -17,6 +17,7 @@ + diff --git a/VectorSharp.Chunking/ChunkReader.cs b/VectorSharp.Chunking/ChunkReader.cs index b659f19..21481b9 100644 --- a/VectorSharp.Chunking/ChunkReader.cs +++ b/VectorSharp.Chunking/ChunkReader.cs @@ -4,17 +4,25 @@ namespace VectorSharp.Chunking { /// /// A streaming text chunker that splits text into token-bounded chunks suitable for embedding. - /// Reads from a and yields chunks as . + /// Reads from a and yields chunks as . /// + /// + /// An instance is a single forward pass over its reader and is consumed by enumerating it. + /// Enumerating the same instance a second time throws, because the reader has already moved + /// past the text the first pass consumed. Create a new instance to chunk the same input again. + /// public sealed class ChunkReader { private readonly SegmentReader _segmentReader; + private readonly ChunkSplitter _chunkSplitter; private readonly Func _countTokens; private readonly ChunkReaderOptions _options; + private int _enumerationStarted; - private ChunkReader(SegmentReader segmentReader, Func countTokens, ChunkReaderOptions options) + private ChunkReader(SegmentReader segmentReader, ChunkSplitter chunkSplitter, Func countTokens, ChunkReaderOptions options) { _segmentReader = segmentReader; + _chunkSplitter = chunkSplitter; _countTokens = countTokens; _options = options; } @@ -22,7 +30,9 @@ private ChunkReader(SegmentReader segmentReader, Func countTokens, /// /// Creates a new instance. /// - /// The stream reader to read text from. + /// The reader to read text from. Any works, + /// including for files and streams and + /// for text already held in memory. /// /// A function that counts the number of tokens in a string. /// Used to enforce . @@ -33,7 +43,7 @@ private ChunkReader(SegmentReader segmentReader, Func countTokens, /// A configured instance. /// Thrown when or is null. /// Thrown when options contain invalid values. - public static ChunkReader Create(StreamReader reader, Func countTokens, ChunkReaderOptions? options = null) + public static ChunkReader Create(TextReader reader, Func countTokens, ChunkReaderOptions? options = null) { if (reader == null) throw new ArgumentNullException(nameof(reader)); @@ -48,73 +58,148 @@ public static ChunkReader Create(StreamReader reader, Func countTok throw new ArgumentException("BreakStrings must contain at least one entry.", nameof(options)); SegmentReader segmentReader = new SegmentReader(reader, effectiveOptions.BreakStrings); - return new ChunkReader(segmentReader, countTokens, effectiveOptions); + ChunkSplitter chunkSplitter = new ChunkSplitter(countTokens, effectiveOptions.MaxTokensPerChunk); + return new ChunkReader(segmentReader, chunkSplitter, countTokens, effectiveOptions); + } + + /// + /// Creates a new instance reading from a . + /// Behaves exactly like the overload. + /// + /// + /// Not redundant with the overload, and not to be deleted as such. + /// It exists for binary compatibility with the published chunking-v1.0.0, where Create took + /// a : assemblies compiled against that version reference this + /// exact signature in IL, so removing it leaves them throwing MissingMethodException at run + /// time even though their source still compiles unchanged against the wider overload. + /// + /// The stream reader to read text from. + /// + /// A function that counts the number of tokens in a string. + /// Used to enforce . + /// + /// + /// Optional configuration. Defaults to Markdown format with 300 tokens per chunk. + /// + /// A configured instance. + /// Thrown when or is null. + /// Thrown when options contain invalid values. + public static ChunkReader Create(StreamReader reader, Func countTokens, ChunkReaderOptions? options = null) + { + return Create((TextReader)reader, countTokens, options); } /// - /// Reads all chunks from the stream asynchronously. + /// Reads all chunks from the reader asynchronously. /// Concatenating all yielded chunks reproduces the original input exactly. + /// Consumes the reader: see the remarks on . /// /// Cancellation token for the operation. /// An async enumerable of text chunks. + /// Thrown when this instance has already been enumerated. public async IAsyncEnumerable ReadAllAsync( [EnumeratorCancellation] CancellationToken cancellationToken = default) { + await foreach (TextChunk chunk in ReadChunksAsync(cancellationToken)) + { + yield return chunk.Text; + } + } + + /// + /// Reads all chunks from the reader asynchronously, each carrying its position in the + /// input and its token count. + /// Concatenating of all yielded chunks reproduces the + /// original input exactly, which is what makes + /// usable for linking a chunk back to its place in the source document. + /// Consumes the reader: see the remarks on . + /// + /// Cancellation token for the operation. + /// An async enumerable of text chunks with position and token information. + /// Thrown when this instance has already been enumerated. + public async IAsyncEnumerable ReadChunksAsync( + [EnumeratorCancellation] CancellationToken cancellationToken = default) + { + // Guards both public enumerations, since ReadAllAsync runs through this one. + // A second pass would silently skip the segment the abandoned pass had already + // consumed but not yet yielded, and produce chunks that reconstruct nothing. + if (Interlocked.Exchange(ref _enumerationStarted, 1) == 1) + throw new InvalidOperationException("This ChunkReader has already been enumerated. Create a new one to read the input again."); + string? firstSegment = await _segmentReader.ReadNextAsync(cancellationToken); if (firstSegment == null) yield break; - string currentChunk = firstSegment; + PendingChunk current = StartChunk(firstSegment); while (true) { cancellationToken.ThrowIfCancellationRequested(); - int currentTokens = _countTokens(currentChunk); - if (currentTokens >= _options.MaxTokensPerChunk) + if (current.TokenCount >= _options.MaxTokensPerChunk) { - yield return currentChunk; - currentChunk = string.Empty; + // Full, or past full. A segment with no break point inside it arrives here + // whole however long it is, so this is the only place a chunk over the limit + // could ever be produced, and the splitter is what stops one being produced. + foreach (Cut cut in _chunkSplitter.Split(current.Text, current.TokenCount, cancellationToken)) + { + yield return new TextChunk + { + Text = current.Text.Substring(cut.Start, cut.Length), + StartOffset = current.StartOffset + cut.Start, + TokenCount = cut.TokenCount + }; + } - string? nextAfterOversize = await _segmentReader.ReadNextAsync(cancellationToken); - if (nextAfterOversize == null) + string? nextAfterFullChunk = await _segmentReader.ReadNextAsync(cancellationToken); + if (nextAfterFullChunk == null) yield break; - currentChunk = nextAfterOversize; + current = StartChunk(nextAfterFullChunk); continue; } string? nextSegment = await _segmentReader.ReadNextAsync(cancellationToken); if (nextSegment == null) { - if (!string.IsNullOrEmpty(currentChunk)) - yield return currentChunk; + if (current.Text.Length > 0) + yield return current.ToTextChunk(); yield break; } if (StartsWithStopSignal(nextSegment)) { - if (!string.IsNullOrEmpty(currentChunk)) - yield return currentChunk; - currentChunk = nextSegment; + if (current.Text.Length > 0) + yield return current.ToTextChunk(); + current = StartChunk(nextSegment); continue; } - string potentialChunk = currentChunk + nextSegment; - int potentialTokens = _countTokens(potentialChunk); + string combinedText = current.Text + nextSegment; + int combinedTokens = _countTokens(combinedText); - if (potentialTokens <= _options.MaxTokensPerChunk) + if (combinedTokens <= _options.MaxTokensPerChunk) { - currentChunk = potentialChunk; + current = current.Extend(combinedText, combinedTokens); } else { - yield return currentChunk; - currentChunk = nextSegment; + yield return current.ToTextChunk(); + current = StartChunk(nextSegment); } } } + /// + /// Starts a chunk at the segment the segment reader just returned, taking the offset + /// from the reader's own cursor instead of a running total kept here, so the offset + /// describes where the text actually came from. + /// + private PendingChunk StartChunk(string segment) + { + return new PendingChunk(segment, _countTokens(segment), _segmentReader.LastSegmentStartOffset); + } + private bool StartsWithStopSignal(string segment) { if (string.IsNullOrEmpty(segment) || _options.StopSignals == null || _options.StopSignals.Count == 0) @@ -132,5 +217,46 @@ private bool StartsWithStopSignal(string segment) return false; } + + /// + /// A chunk under construction: its text, the token count of exactly that text, and the + /// offset it starts at. The three travel as one value so they cannot drift apart, and so + /// the token count is carried forward when a chunk grows instead of being recounted. + /// + private readonly struct PendingChunk + { + internal PendingChunk(string text, int tokenCount, long startOffset) + { + Text = text; + TokenCount = tokenCount; + StartOffset = startOffset; + } + + internal string Text { get; } + + internal int TokenCount { get; } + + internal long StartOffset { get; } + + /// + /// Returns this chunk with more text appended. The start offset does not move, + /// and the caller supplies the count it already had to compute for the decision + /// to extend. + /// + internal PendingChunk Extend(string extendedText, int extendedTokenCount) + { + return new PendingChunk(extendedText, extendedTokenCount, StartOffset); + } + + internal TextChunk ToTextChunk() + { + return new TextChunk + { + Text = Text, + StartOffset = StartOffset, + TokenCount = TokenCount + }; + } + } } } diff --git a/VectorSharp.Chunking/ChunkReaderOptions.cs b/VectorSharp.Chunking/ChunkReaderOptions.cs index 71a9231..531aa8a 100644 --- a/VectorSharp.Chunking/ChunkReaderOptions.cs +++ b/VectorSharp.Chunking/ChunkReaderOptions.cs @@ -8,6 +8,22 @@ public sealed class ChunkReaderOptions /// /// Gets the maximum number of tokens allowed per chunk. Default is 300. /// + /// + /// A hard bound, not a target: text with no break point in it for long enough to pass this + /// is cut at the limit rather than emitted whole. Set it to what the model behind your + /// embedding provider actually accepts — a chunk over a hosted model's context window is + /// rejected or silently truncated, and a truncated chunk is a piece of the document that + /// is no longer searchable with nothing in the result to say so. + /// + /// A cut lands where the token counter says the limit falls, which for a subword or + /// character counter is in the middle of a word. Break strings, not this, are what put a + /// boundary somewhere meaningful. + /// + /// The bound rests on a counter that does not report over the limit for a short piece of + /// text and under it for a longer one. Every real tokenizer satisfies that; one that does + /// not can produce a chunk over the limit, because the alternative is cutting a single + /// character in half. + /// public int MaxTokensPerChunk { get; init; } = 300; /// diff --git a/VectorSharp.Chunking/ChunkSplitter.cs b/VectorSharp.Chunking/ChunkSplitter.cs new file mode 100644 index 0000000..4f04d46 --- /dev/null +++ b/VectorSharp.Chunking/ChunkSplitter.cs @@ -0,0 +1,241 @@ +namespace VectorSharp.Chunking +{ + /// + /// One piece of a chunk: where it starts and ends in the text it was cut from, and the number + /// of tokens counted for exactly that piece. + /// + /// + /// The three travel together for the same reason a chunk under construction does — a position + /// and a count kept apart drift apart, and a count that belongs to a different span of text is + /// worse than no count at all. + /// + internal readonly struct Cut + { + internal Cut(int start, int end, int tokenCount) + { + Start = start; + End = end; + TokenCount = tokenCount; + } + + internal int Start { get; } + + internal int End { get; } + + internal int TokenCount { get; } + + internal int Length => End - Start; + } + + /// + /// Cuts text into consecutive pieces that stay within a token limit. + /// + /// + /// Separate from , which assembles segments into chunks: this decides + /// where text may be divided, which is a different question answered with a different tool — + /// a search against the caller's token counter rather than a scan for break strings. Keeping + /// it apart also means the search can be exercised directly instead of only through a reader. + /// + /// The pieces cover the text exactly, in order and with no gaps, which is what carries the + /// round-trip guarantee and the offsets through a cut. + /// + internal sealed class ChunkSplitter + { + private readonly Func _countTokens; + private readonly int _maxTokensPerChunk; + + internal ChunkSplitter(Func countTokens, int maxTokensPerChunk) + { + _countTokens = countTokens; + _maxTokensPerChunk = maxTokensPerChunk; + } + + /// + /// Cuts into pieces within the token limit. + /// + /// The text to cut. + /// The token count already measured for the whole of + /// , so that text needing no cutting costs no further counting. + /// Observed once per piece. Cutting a long stretch of text + /// with no break point in it can take many pieces and many calls into the caller's + /// tokenizer, and a caller who has stopped waiting should not have to wait for all of them. + /// + /// Text already within the limit comes back as one piece without being searched. For an + /// ordinary counter that is only a saving — the search would arrive at the same answer, + /// after a handful of needless calls into the caller's tokenizer. + /// + /// What it also does is answer from the measurement of the whole text rather than from the + /// smallest piece upward. A counter that reports over the limit for one character and + /// under it for the whole chunk would otherwise be taken at its word about the character: + /// the search works upward and stops at the first piece that does not fit, so it would cut + /// a chunk that fits perfectly into single characters. + /// + internal IEnumerable Split(string text, int knownTokenCount, CancellationToken cancellationToken = default) + { + if (knownTokenCount <= _maxTokensPerChunk) + { + yield return new Cut(0, text.Length, knownTokenCount); + yield break; + } + + // Walks an index rather than re-slicing what is left. Trimming the front off a string + // copies everything behind it, which for a long stretch of text is the same quadratic + // cost the search is written to avoid, just moved somewhere less visible. + int start = 0; + while (start < text.Length) + { + cancellationToken.ThrowIfCancellationRequested(); + + Cut cut = NextCut(text, start); + yield return cut; + + start = cut.End; + } + } + + /// + /// Finds how much of can be taken from + /// while staying within the token limit, and the token count of exactly that much. + /// + /// + /// The token counter belongs to the caller, so nothing here can predict what it will say. + /// Every length this returns is one that was measured, never one interpolated between two + /// probes — which is what makes the result safe against a tokenizer that is not perfectly + /// monotonic in the length of its input. Byte-pair encoders are not: "the" can cost fewer + /// tokens than "th". Such a counter only makes a cut shorter than it needed to be. + /// + /// What the search does assume is that a counter which reports over the limit for a short + /// piece of text will not report under it for a longer one, since it stops probing upward + /// at the first piece that does not fit. A counter that violates that badly enough — one + /// that charges more for one character than the whole limit allows — is the one case where + /// a piece comes back over the limit, because the alternative is cutting a character in + /// half. + /// + /// The search runs outward from the smallest possible piece, doubling until one does not + /// fit, and then narrows inside that range. Searching the whole remaining text for every + /// piece would be quadratic in the length of the input, because each probe has to copy the + /// text it measures; this keeps every probe within twice the piece it finds. + /// + /// A piece of at least one character, never ending inside a surrogate pair. + internal Cut NextCut(string text, int start) + { + int smallestEnd = SmallestPieceEnd(text, start); + int smallestTokens = CountFrom(text, start, smallestEnd); + + if (smallestTokens > _maxTokensPerChunk) + { + // The smallest piece that can be taken is already over the limit. It goes out + // anyway: cutting further would divide one character into halves that are not + // characters, which hands the caller's tokenizer and their embedding API a lone + // surrogate. See the remarks above for when this is reachable. + return new Cut(start, smallestEnd, smallestTokens); + } + + int bestEnd = smallestEnd; + int bestTokens = smallestTokens; + + if (bestEnd == text.Length) + return new Cut(start, bestEnd, bestTokens); + + // Outward: double the piece until one does not fit, so the range to narrow is bounded + // by the answer rather than by the length of the input. + int probeEnd = bestEnd; + while (true) + { + probeEnd = NextLegalEndAtOrAbove(text, start + (int)Math.Min((long)(probeEnd - start) * 2, text.Length - start)); + + // Upward, never downward, so that stepping off a surrogate pair cannot leave the + // probe where it already was and stall the doubling. + if (probeEnd <= bestEnd) + probeEnd = NextLegalEndAtOrAbove(text, Math.Min(bestEnd + 1, text.Length)); + + int count = CountFrom(text, start, probeEnd); + if (count > _maxTokensPerChunk) + break; + + bestEnd = probeEnd; + bestTokens = count; + + if (bestEnd == text.Length) + return new Cut(start, bestEnd, bestTokens); + } + + // Inward: bestEnd fits, probeEnd does not, so the answer is between them. + int low = bestEnd + 1; + int high = probeEnd - 1; + + while (low <= high) + { + int middle = low + (high - low) / 2; + int candidate = NextLegalEndAtOrBelow(text, middle); + + // Stepping down off a surrogate pair can land below the range. Step up instead — + // clamping back to the bottom of the range would put the cut back inside the pair, + // which is the one place it must never be. + if (candidate < low) + candidate = NextLegalEndAtOrAbove(text, middle); + + if (candidate > high) + { + // Every position left in the range is inside a surrogate pair, so there is no + // longer piece to be found. + break; + } + + int count = CountFrom(text, start, candidate); + + if (count <= _maxTokensPerChunk) + { + bestEnd = candidate; + bestTokens = count; + low = candidate + 1; + } + else + { + high = candidate - 1; + } + } + + return new Cut(start, bestEnd, bestTokens); + } + + private int CountFrom(string text, int start, int end) + { + return _countTokens(text.Substring(start, end - start)); + } + + /// + /// Where the smallest piece starting at ends: one character + /// later, or two when the first is the leading half of a surrogate pair. + /// + internal static int SmallestPieceEnd(string text, int start) + { + return start + 2 <= text.Length && char.IsHighSurrogate(text[start]) && char.IsLowSurrogate(text[start + 1]) ? start + 2 : start + 1; + } + + /// + /// Whether ending a piece at this position would fall between the two halves of one + /// character. + /// + /// + /// Cutting there hands out two pieces that are not text: a lone surrogate is not a + /// character, and it is what would reach the caller's tokenizer and their embedding API. + /// This protects code points and nothing wider — a cut can still separate a base character + /// from a combining mark that follows it, and can land in the middle of a word. + /// + internal static bool IsInsideSurrogatePair(string text, int position) + { + return position > 0 && position < text.Length && char.IsHighSurrogate(text[position - 1]) && char.IsLowSurrogate(text[position]); + } + + internal static int NextLegalEndAtOrBelow(string text, int position) + { + return IsInsideSurrogatePair(text, position) ? position - 1 : position; + } + + internal static int NextLegalEndAtOrAbove(string text, int position) + { + return IsInsideSurrogatePair(text, position) ? position + 1 : position; + } + } +} diff --git a/VectorSharp.Chunking/README.md b/VectorSharp.Chunking/README.md index 6a7433f..3fed532 100644 --- a/VectorSharp.Chunking/README.md +++ b/VectorSharp.Chunking/README.md @@ -14,13 +14,16 @@ dotnet add package VectorSharp.Chunking ## Features -- **Stream-based** — reads character-by-character, never loads the full file into memory -- **Token-bounded** — chunks respect a configurable token limit via your own token counter +- **Stream-based** — reads character-by-character from any `TextReader`, holding one segment at a time rather than the whole file +- **Token-bounded** — every chunk respects the configured token limit, counted by your own token counter. Text that runs past the limit with no break point in it is cut at the limit rather than emitted whole +- **Position-aware** — `ReadChunksAsync` yields chunks that carry their offset in the source and their token count - **Format-aware** — ships with predefined break strings and stop signals for Markdown, C#, JavaScript/TypeScript/JSX/TSX, HTML, CSS, Python, and generic plain text - **Round-trip safe** — concatenating all chunks reproduces the original text exactly - **Stop signals** — headings, code blocks, and other structural elements always start a new chunk - **Zero dependencies** — pure text processing, no embedding or tokenizer dependency +> **Memory, as distinct from chunk size.** Input containing none of the configured break strings — minified JavaScript, a base64 blob, a single-line CSV — is one segment, and a segment is accumulated whole before it is cut into chunks. The chunks that come out are within the limit, so nothing over-sized reaches an embedding API, but peak memory is still the size of the longest run of text with no break point in it. Give such input break strings that do occur in it (`","`, `";"`, a fixed run of characters) if that matters. + ## Quick Start ```csharp @@ -31,11 +34,66 @@ ChunkReader chunker = ChunkReader.Create(reader, text => myTokenizer.CountTokens await foreach (string chunk in chunker.ReadAllAsync()) { - // Each chunk is within the token limit and splits at natural boundaries + // Each chunk is within the token limit, splitting at break strings where it can float[] embedding = await embedder.EmbedAsync(chunk); } ``` +Any `TextReader` works, so text already in memory needs no stream wrapper: + +```csharp +using StringReader reader = new StringReader(documentText); +ChunkReader chunker = ChunkReader.Create(reader, myTokenCounter); +``` + +## Chunks With Position + +`ReadChunksAsync` yields `TextChunk` values instead of bare strings. Each one knows where it came +from, which is what lets a search hit link back to its place in the source document: + +```csharp +await foreach (TextChunk chunk in chunker.ReadChunksAsync()) +{ + // chunk.StartOffset — character offset from the start of the input + // chunk.EndOffset — one past the last character, the next chunk's StartOffset + // chunk.TokenCount — as counted by the token counter you supplied + float[] embedding = await embedder.EmbedAsync(chunk.Text); +} +``` + +An offset is the position the chunk's text was read from, not an estimate reconstructed afterwards. +Because the chunks reproduce the input exactly, it is also the sum of the lengths of all preceding +chunks. Offsets are measured in UTF-16 code units, the same unit as `string.Length`, and typed as +`long` because they accumulate over the whole input, which may run past `int.MaxValue` characters. + +`ReadAllAsync` yields the same text as `ReadChunksAsync` — it is the string-only view of the same +pass. + +A `ChunkReader` is a single forward pass over its reader: it is consumed by enumerating it, and +enumerating the same instance again throws `InvalidOperationException`. Create a new one to chunk +the same input a second time. + +## Upgrading from 1.x + +The only published 1.x release is 1.0.0, so this is the whole of what changes for every upgrader. +Two behaviours differ; no signature does. + +**`MaxTokensPerChunk` is a bound rather than a target.** In 1.0.0, a segment with no break point +inside it was emitted whole however far over the limit it was; it is now cut at the limit. Text +whose segments already fit chunks identically — the cut only reaches text that used to come back +over the limit, which is text a model with a fixed context window would have rejected or silently +truncated. Read the note under [How It Works](#how-it-works) on where a cut lands before relying +on this, because it is not at a word boundary. + +**Enumerating a `ChunkReader` twice throws.** In 1.0.0, a second enumeration returned an empty +sequence — which looked like "this document has no chunks" rather than "this reader is spent" — +and after an abandoned pass it silently dropped the segment that pass had consumed but not yet +yielded. Code that enumerated once, which is the only use that ever produced correct output, is +unaffected. + +The public surface is otherwise source- and binary-compatible: nothing was removed or changed, and +the `StreamReader` overload of `Create` is still there. + ## Configuration ```csharp @@ -128,15 +186,22 @@ ChunkReader chunker = ChunkReader.Create(reader, myTokenCounter, new ChunkReader ## How It Works ``` -StreamReader ──▶ SegmentReader ──▶ ChunkReader ──▶ IAsyncEnumerable - (break strings) (token limits, - stop signals) +TextReader ──▶ SegmentReader ──▶ ChunkReader ──▶ IAsyncEnumerable + (break strings) (token limits, or IAsyncEnumerable + stop signals, + cutting) ``` 1. **Segment reading** — text is read character-by-character and split at break string boundaries. Longer break strings are matched first (e.g., `\n\n` is preferred over `\n`). 2. **Chunk assembly** — segments are concatenated into chunks until adding the next segment would exceed the token limit. If a segment starts with a stop signal, it forces a new chunk to begin. +3. **Cutting** — a segment that is over the limit on its own is cut into consecutive pieces that are not. The cut is found by measuring candidate prefixes with your token counter, so it lands where your tokenizer says the limit falls rather than at a guessed character count. The pieces still concatenate back into the original text, and each carries its own offset, so round-tripping and `StartOffset` hold across a cut. + +> **A cut lands at the token limit, not at anything you would call a boundary.** It is reached only by text that had no break point in it for long enough to run past the limit, and at that point there is nothing structural left to cut at. With a character- or subword-based counter — which is what real embedding models use — that means cutting mid-word: at a limit of 3 tokens, `"the quick brown fox jumps"` comes back as `"the quick br"`, `"own fox jump"`, `"s"`. A word-counting tokenizer never does this, because adding a character to a word does not raise its count, so a word counter is a poor model of what your provider will actually do. +> +> The one thing a cut will not do is divide a single character: it never falls between the two halves of a surrogate pair. It can still separate a base character from a combining mark that follows it. If you need cuts on grapheme or word boundaries, give the chunker break strings that occur in your text so the limit is never the thing that decides. + ## End-to-End with VectorSharp ```csharp @@ -146,19 +211,36 @@ using VectorSharp.Embedding; using VectorSharp.Embedding.NomicEmbed; await using EmbeddingService embedder = new EmbeddingService(NomicEmbedProvider.Create); -using CosineVectorStore store = VectorStore.Create("docs", embedder.Dimension); +using CosineVectorStore store = VectorStore.Create("docs", embedder.Dimension); using StreamReader reader = new StreamReader("document.md"); ChunkReader chunker = ChunkReader.Create(reader, text => myTokenizer.CountTokens(text)); -int id = 0; -await foreach (string chunk in chunker.ReadAllAsync()) +// Collect first, so the whole document goes to the embedder in one call rather than one call +// per chunk. Against a remote embedding API that is the difference between one HTTP round trip +// and one per chunk; the service splits the texts into provider-sized batches itself. +List chunks = new List(); +await foreach (TextChunk chunk in chunker.ReadChunksAsync()) { - float[] embedding = await embedder.EmbedAsync(chunk, EmbeddingPurpose.Document); - await store.AddAsync(id++, embedding); + chunks.Add(chunk); +} + +float[][] embeddings = await embedder.EmbedBatchAsync( + chunks.Select(chunk => chunk.Text).ToList(), + EmbeddingPurpose.Document); + +for (int i = 0; i < chunks.Count; i++) +{ + // Keying on the offset rather than a running counter is what StartOffset is for: it points + // back into the source document, so a hit can be quoted or linked in context later. + await store.AddAsync(chunks[i].StartOffset, embeddings[i]); } ``` +`AddAsync` is keyed by the store's key type — `VectorStore.Create` here, to match `StartOffset`. Use whatever key your own records are addressed by; the point is that the chunk carries enough to build one. + +If the document is large enough that holding every chunk in memory matters, batch as you go: fill a `List` up to a few hundred entries, embed and store that, then clear it. What to avoid is a single `EmbedAsync` per chunk. + ## API Reference ### ChunkReader @@ -166,9 +248,28 @@ await foreach (string chunk in chunker.ReadAllAsync()) ```csharp public sealed class ChunkReader { + public static ChunkReader Create(TextReader reader, Func countTokens, + ChunkReaderOptions? options = null); public static ChunkReader Create(StreamReader reader, Func countTokens, ChunkReaderOptions? options = null); public IAsyncEnumerable ReadAllAsync(CancellationToken cancellationToken = default); + public IAsyncEnumerable ReadChunksAsync(CancellationToken cancellationToken = default); +} +``` + +The `StreamReader` overload does nothing the `TextReader` one does not. It stays because +chunking-v1.0.0 shipped `Create` taking a `StreamReader`, and assemblies compiled against that +version bind to that exact signature at run time. + +### TextChunk + +```csharp +public sealed class TextChunk +{ + public required string Text { get; init; } + public required long StartOffset { get; init; } // characters from the start of the input + public required int TokenCount { get; init; } + public long EndOffset { get; } // StartOffset + Text.Length } ``` diff --git a/VectorSharp.Chunking/SegmentReader.cs b/VectorSharp.Chunking/SegmentReader.cs index 04ed21c..cc375a2 100644 --- a/VectorSharp.Chunking/SegmentReader.cs +++ b/VectorSharp.Chunking/SegmentReader.cs @@ -3,22 +3,31 @@ namespace VectorSharp.Chunking { /// - /// Reads text from a stream and splits it into segments at break string boundaries. + /// Reads text from a reader and splits it into segments at break string boundaries. /// Uses character-by-character reading with a lookahead buffer for memory efficiency. /// internal sealed class SegmentReader { - private readonly StreamReader _contentReader; + private readonly TextReader _contentReader; private readonly string[] _breakStrings; private readonly Queue _unreadBuffer; private readonly char[] _readBuffer = new char[1]; + private long _returnedCharacters; + + /// + /// Gets the offset in the input of the segment returned by the most recent + /// call that returned one. + /// Characters pushed back into the lookahead buffer are not counted until the segment + /// that contains them is returned, so this is the true position of the segment text. + /// + internal long LastSegmentStartOffset { get; private set; } /// /// Initializes a new instance of the class. /// - /// The stream reader to read content from. + /// The reader to read content from. /// The strings that indicate break points. Will be sorted longest-first internally. - internal SegmentReader(StreamReader contentReader, IReadOnlyList breakStrings) + internal SegmentReader(TextReader contentReader, IReadOnlyList breakStrings) { _contentReader = contentReader; _breakStrings = breakStrings.ToArray(); @@ -66,32 +75,50 @@ internal SegmentReader(StreamReader contentReader, IReadOnlyList breakSt } else if (currentMatch != null) { - string result = segmentBuilder.ToString(0, matchEndPosition); - - string extraChars = segmentBuilder.ToString(matchEndPosition, segmentBuilder.Length - matchEndPosition); - foreach (char c in extraChars) - { - _unreadBuffer.Enqueue(c); - } - - return result; + return TakeSegmentAndPushBackTail(segmentBuilder, matchEndPosition); } } if (currentMatch != null) { - string result = segmentBuilder.ToString(0, matchEndPosition); + return TakeSegmentAndPushBackTail(segmentBuilder, matchEndPosition); + } - string extraChars = segmentBuilder.ToString(matchEndPosition, segmentBuilder.Length - matchEndPosition); - foreach (char c in extraChars) - { - _unreadBuffer.Enqueue(c); - } + return segmentBuilder.Length > 0 ? TrackPosition(segmentBuilder.ToString()) : null; + } - return result; + /// + /// Returns everything up to the end of the break string match, and puts the lookahead read + /// past it back for the next call. + /// + /// + /// One helper for both of the paths that end a segment on a match: the loop's, when a + /// longer break string turned out not to follow, and the end-of-input one. They were + /// identical, and identical is how they have to stay — a change made to one and not the + /// other would move the reader's position on some inputs and not others. + /// + private string TakeSegmentAndPushBackTail(StringBuilder segmentBuilder, int matchEndPosition) + { + string segment = segmentBuilder.ToString(0, matchEndPosition); + + for (int i = matchEndPosition; i < segmentBuilder.Length; i++) + { + _unreadBuffer.Enqueue(segmentBuilder[i]); } - return segmentBuilder.Length > 0 ? segmentBuilder.ToString() : null; + return TrackPosition(segment); + } + + /// + /// Records where the segment about to be returned starts and advances the cursor past it. + /// Every return path of that produces a segment goes through + /// here, which is what keeps in step with the input. + /// + private string TrackPosition(string segment) + { + LastSegmentStartOffset = _returnedCharacters; + _returnedCharacters += segment.Length; + return segment; } private string? FindLongestBreakString(StringBuilder content) diff --git a/VectorSharp.Chunking/TextChunk.cs b/VectorSharp.Chunking/TextChunk.cs new file mode 100644 index 0000000..40b09bc --- /dev/null +++ b/VectorSharp.Chunking/TextChunk.cs @@ -0,0 +1,44 @@ +namespace VectorSharp.Chunking +{ + // A class rather than a struct: `required` is enforced only on object-creation expressions, + // so a struct here would let `default(TextChunk)` and `List.FirstOrDefault()` + // hand out an instance whose non-nullable Text is null and whose EndOffset throws. + // As a class those same expressions produce a plain null, which the caller can see. + + /// + /// A single chunk produced by , together with + /// its position in the source text and its token count. + /// + public sealed class TextChunk + { + /// + /// Gets the chunk text. Concatenating this value for every chunk in the order they were + /// yielded reproduces the original input exactly. + /// + public required string Text { get; init; } + + /// + /// Gets the character offset of this chunk from the start of the input, which is the + /// position the chunk's first segment was read from. Because the chunks reproduce the + /// input exactly, it is also the sum of the lengths of all preceding chunks. + /// Measured in UTF-16 code units, the same unit as , and + /// typed as because offsets accumulate over the whole input, which + /// may run past characters. + /// + public required long StartOffset { get; init; } + + /// + /// Gets the number of tokens in , as counted by the token counter + /// passed to . + /// Bounded by , whose documentation + /// states the bound and the one counter that can defeat it. + /// + public required int TokenCount { get; init; } + + /// + /// Gets the character offset one past the last character of this chunk, which is the + /// of the next chunk. + /// + public long EndOffset => StartOffset + Text.Length; + } +} diff --git a/VectorSharp.Chunking/VectorSharp.Chunking.csproj b/VectorSharp.Chunking/VectorSharp.Chunking.csproj index d1b37fb..d98c489 100644 --- a/VectorSharp.Chunking/VectorSharp.Chunking.csproj +++ b/VectorSharp.Chunking/VectorSharp.Chunking.csproj @@ -7,10 +7,10 @@ True VectorSharp.Chunking - 1.0.0 + 2.0.0 VectorSharp.Chunking Adam Tovatt - Streaming text chunker that splits text into token-bounded chunks suitable for embedding. Ships with predefined formats for Markdown and C#. Zero dependencies. + Streaming text chunker that splits text into token-bounded chunks suitable for embedding. Chunks carry their offset in the source text and their token count. Ships with predefined formats for Markdown, C#, JavaScript/TypeScript, HTML, CSS, Python and plain text. Zero dependencies. MIT vector;chunking;text-chunking;text-splitting;embedding;rag README.md diff --git a/VectorSharp.Embedding.Tests/BatchRecordingProvider.cs b/VectorSharp.Embedding.Tests/BatchRecordingProvider.cs new file mode 100644 index 0000000..239c05d --- /dev/null +++ b/VectorSharp.Embedding.Tests/BatchRecordingProvider.cs @@ -0,0 +1,161 @@ +namespace VectorSharp.Embedding.Tests +{ + /// + /// What the providers were handed, shared across every instance the factory creates. + /// Held outside the provider because the service creates one provider per worker, so state a + /// test needs to read across workers cannot live on the provider itself. + /// + internal sealed class BatchRecorder + { + private readonly object _lock = new object(); + private readonly List _calls = new List(); + private int _singleCallCount; + + /// + /// The calls the providers received, in arrival order. A batch of N texts appears as one + /// entry with N texts; N single calls appear as N entries of one. + /// + public IReadOnlyList Calls + { + get + { + lock (_lock) + { + return _calls.ToList(); + } + } + } + + public IReadOnlyList> Batches + { + get + { + return Calls.Select(call => call.Texts).ToList(); + } + } + + /// + /// How many times the single-text method was called. Stays at zero while the service uses + /// the batch path, which is what distinguishes real batching from a loop. + /// + public int SingleCallCount => Volatile.Read(ref _singleCallCount); + + public void RecordBatch(IReadOnlyList texts, EmbeddingPurpose purpose) + { + lock (_lock) + { + _calls.Add(new RecordedCall { Texts = texts.ToList(), Purpose = purpose }); + } + } + + public void RecordSingleCall() + { + Interlocked.Increment(ref _singleCallCount); + } + + internal sealed class RecordedCall + { + public required IReadOnlyList Texts { get; init; } + + public required EmbeddingPurpose Purpose { get; init; } + } + } + + /// + /// A provider that embeds a whole batch in one call, the way a remote API provider does, and + /// reports what it received to a shared . Tests assert against what + /// arrived here rather than against what the service set out to send. + /// + internal sealed class BatchRecordingProvider : IEmbeddingProvider + { + private readonly BatchRecorder _recorder; + private readonly Func, EmbeddingUsage?>? _usageForBatch; + + /// + /// Fixed rather than configurable, so it cannot disagree with the length + /// actually produces. + /// + public int Dimension => VectorLength; + + /// + /// True by default, the way a provider backed by an API that takes an array of inputs + /// answers. Set false to stand in for one that embeds a text at a time. + /// + public bool SupportsBatching { get; } + + private const int VectorLength = 8; + + /// Collects what every instance of this provider received. Omit it + /// in a test that does not read what arrived: the provider then records into one of its + /// own, so the test's Arrange block holds only what the test is actually about. Pass one + /// whenever the assertions look at the calls, since the service builds a provider per + /// worker and only a shared recorder sees all of them. + /// Produces the usage this provider reports for a batch, or + /// null to report none. Null by default, which is the local-model case. + /// Whether this provider tells the service it batches. + public BatchRecordingProvider(BatchRecorder? recorder = null, Func, EmbeddingUsage?>? usageForBatch = null, bool supportsBatching = true) + { + _recorder = recorder ?? new BatchRecorder(); + _usageForBatch = usageForBatch; + SupportsBatching = supportsBatching; + } + + public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + _recorder.RecordSingleCall(); + return Task.FromResult(VectorFor(text)); + } + + public Task EmbedBatchWithUsageAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + _recorder.RecordBatch(texts, purpose); + + float[][] vectors = new float[texts.Count][]; + for (int i = 0; i < texts.Count; i++) + { + vectors[i] = VectorFor(texts[i]); + } + + return Task.FromResult(new EmbeddingResult + { + Vectors = vectors, + Usage = _usageForBatch?.Invoke(texts) + }); + } + + /// + /// A vector that identifies the text it came from, so a test can tell whether vectors came + /// back matched to the texts that produced them. + /// + /// + /// Every element is derived from the whole string, in order. An earlier version stored the + /// length and a plain sum of the characters, which is equal for any two anagrams and for + /// any two same-length strings whose characters happen to add up the same — so a test + /// asserting that vectors came back in input order held only while its literals happened + /// to avoid the collision, and adding one more text could silently stop it detecting + /// misordering. VectorFor_DistinguishesTextsAPlainLengthAndSumWouldNot pins this. + /// + public static float[] VectorFor(string text) + { + float[] vector = new float[VectorLength]; + + // FNV-1a over the text, then one further round per element, so position matters and + // every element depends on every character. + uint hash = 2166136261; + foreach (char character in text) + { + hash = (hash ^ character) * 16777619; + } + + for (int i = 0; i < VectorLength; i++) + { + hash = (hash ^ (uint)i) * 16777619; + vector[i] = hash % 65536; + } + + return vector; + } + + public void Dispose() { } + } +} diff --git a/VectorSharp.Embedding.Tests/BatchingPolicyTests.cs b/VectorSharp.Embedding.Tests/BatchingPolicyTests.cs new file mode 100644 index 0000000..4a36a24 --- /dev/null +++ b/VectorSharp.Embedding.Tests/BatchingPolicyTests.cs @@ -0,0 +1,134 @@ +namespace VectorSharp.Embedding.Tests +{ + /// + /// The splitting rule on its own. The service-level tests in + /// EmbeddingBatchingTests still cover that these partitions reach a provider as calls; + /// these cover the partitions themselves, which do not need a worker pool to be wrong. + /// + public class BatchingPolicyTests + { + private static string[] Texts(int count, int lengthEach = 4) + { + string[] texts = new string[count]; + for (int i = 0; i < count; i++) + { + texts[i] = new string((char)('a' + (i % 26)), lengthEach); + } + return texts; + } + + private static int[] Sizes(IReadOnlyList> batches) + { + return batches.Select(batch => batch.Count).ToArray(); + } + + [Fact] + public void Split_WithinBothLimits_IsOneBatch() + { + BatchingPolicy policy = new BatchingPolicy(maxTextsPerBatch: 64, maxCharactersPerBatch: 1000); + + Assert.Equal([5], Sizes(policy.Split(Texts(5)))); + } + + [Fact] + public void Split_OnTextCount_FillsEachBatchBeforeStartingTheNext() + { + BatchingPolicy policy = new BatchingPolicy(maxTextsPerBatch: 3, maxCharactersPerBatch: 1000); + + Assert.Equal([3, 3, 1], Sizes(policy.Split(Texts(7)))); + } + + [Fact] + public void Split_OnCharacterCount_ClosesTheBatchBeforeItWouldExceed() + { + BatchingPolicy policy = new BatchingPolicy(maxTextsPerBatch: 100, maxCharactersPerBatch: 25); + + Assert.Equal([2, 2, 2, 2, 2], Sizes(policy.Split(Texts(10, lengthEach: 10)))); + } + + [Fact] + public void Split_BothLimitsInPlay_AppliesWhicheverBindsFirst() + { + BatchingPolicy policy = new BatchingPolicy(maxTextsPerBatch: 4, maxCharactersPerBatch: 30); + + // Four would fit the count limit, but three of ten characters reach the character + // limit first. + Assert.Equal([3, 3, 2], Sizes(policy.Split(Texts(8, lengthEach: 10)))); + } + + [Fact] + public void Split_TextLongerThanTheCharacterLimit_GoesInABatchOfItsOwn() + { + BatchingPolicy policy = new BatchingPolicy(maxTextsPerBatch: 64, maxCharactersPerBatch: 10); + string oversized = new string('x', 50); + + IReadOnlyList> batches = policy.Split(["short", oversized, "tiny"]); + + Assert.Contains(batches, batch => batch.Count == 1 && batch[0] == oversized); + Assert.Equal(3, batches.Sum(batch => batch.Count)); + } + + [Fact] + public void Split_PreservesOrderAcrossBatches() + { + // The service reassembles vectors by position, so the batches concatenated in order + // have to reproduce the input. A split that reordered would misalign every vector. + BatchingPolicy policy = new BatchingPolicy(maxTextsPerBatch: 2, maxCharactersPerBatch: 1000); + string[] texts = ["a", "b", "c", "d", "e"]; + + IReadOnlyList> batches = policy.Split(texts); + + Assert.Equal(texts, batches.SelectMany(batch => batch)); + } + + [Fact] + public void Split_NoTexts_ProducesNoBatches() + { + BatchingPolicy policy = new BatchingPolicy(maxTextsPerBatch: 64, maxCharactersPerBatch: 1000); + + Assert.Empty(policy.Split([])); + } + + [Fact] + public void For_ProviderThatDoesNotBatch_CapsAtOneTextWhateverTheOptionsSay() + { + BatchingPolicy policy = BatchingPolicy.For( + new EmbeddingServiceOptions { MaxTextsPerBatch = 64 }, + providerSupportsBatching: false); + + Assert.Equal(1, policy.MaxTextsPerBatch); + } + + [Fact] + public void For_ProviderThatBatches_TakesTheConfiguredLimit() + { + BatchingPolicy policy = BatchingPolicy.For( + new EmbeddingServiceOptions { MaxTextsPerBatch = 64 }, + providerSupportsBatching: true); + + Assert.Equal(64, policy.MaxTextsPerBatch); + } + + [Fact] + public void For_CharacterLimit_IsTakenWhetherOrNotTheProviderBatches() + { + // Only the text count depends on the provider's answer. A single text is a whole + // request either way, so the character limit is carried through unchanged. + EmbeddingServiceOptions options = new EmbeddingServiceOptions { MaxCharactersPerBatch = 4096 }; + + Assert.Equal(4096, BatchingPolicy.For(options, providerSupportsBatching: false).MaxCharactersPerBatch); + Assert.Equal(4096, BatchingPolicy.For(options, providerSupportsBatching: true).MaxCharactersPerBatch); + } + + [Theory] + [InlineData(0, 100)] + [InlineData(-1, 100)] + [InlineData(64, 0)] + [InlineData(64, -1)] + public void Constructor_LimitBelowOne_Throws(int maxTextsPerBatch, int maxCharactersPerBatch) + { + Assert.Throws(() => + new BatchingPolicy(maxTextsPerBatch, maxCharactersPerBatch)); + } + } +} diff --git a/VectorSharp.Embedding.Tests/EmbeddingBatchingTests.cs b/VectorSharp.Embedding.Tests/EmbeddingBatchingTests.cs new file mode 100644 index 0000000..ff8ad8e --- /dev/null +++ b/VectorSharp.Embedding.Tests/EmbeddingBatchingTests.cs @@ -0,0 +1,483 @@ +namespace VectorSharp.Embedding.Tests +{ + public class EmbeddingBatchingTests + { + private static string[] Texts(int count, int lengthEach = 4) + { + string[] texts = new string[count]; + for (int i = 0; i < count; i++) + { + texts[i] = new string((char)('a' + (i % 26)), lengthEach); + } + return texts; + } + + #region Batching reaches the provider as one call + + [Fact] + public async Task EmbedBatchAsync_WithinLimits_ReachesTheProviderAsOneCall() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder)); + + await service.EmbedBatchAsync(Texts(5)); + + Assert.Single(recorder.Batches); + Assert.Equal(5, recorder.Batches[0].Count); + Assert.Equal(0, recorder.SingleCallCount); + } + + [Fact] + public async Task EmbedBatchAsync_SplitsOnTextCountLimit() + { + BatchRecorder recorder = new BatchRecorder(); + // Concurrency 1 so the batches arrive in the order they were queued; this asserts the + // sizes in sequence, which several workers would interleave. + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder), + new EmbeddingServiceOptions { MaxTextsPerBatch = 3, Concurrency = 1 }); + + await service.EmbedBatchAsync(Texts(7)); + + Assert.Equal(3, recorder.Batches.Count); + Assert.Equal([3, 3, 1], recorder.Batches.Select(batch => batch.Count)); + } + + [Fact] + public async Task EmbedBatchAsync_SplitsOnCharacterLimit() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder), + new EmbeddingServiceOptions { MaxTextsPerBatch = 100, MaxCharactersPerBatch = 25 }); + + // Ten texts of ten characters: the character limit closes a batch after two. + await service.EmbedBatchAsync(Texts(10, lengthEach: 10)); + + Assert.Equal(5, recorder.Batches.Count); + Assert.All(recorder.Batches, batch => Assert.Equal(2, batch.Count)); + } + + [Fact] + public async Task EmbedBatchAsync_AppliesWhicheverLimitBindsFirst() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder), + new EmbeddingServiceOptions { MaxTextsPerBatch = 4, MaxCharactersPerBatch = 30, Concurrency = 1 }); + + // Four texts would fit the count limit, but three of ten characters reach the + // character limit first, so both limits have to hold at once. + await service.EmbedBatchAsync(Texts(8, lengthEach: 10)); + + Assert.Equal([3, 3, 2], recorder.Batches.Select(batch => batch.Count)); + } + + [Fact] + public async Task EmbedBatchAsync_TextLongerThanTheCharacterLimit_IsSentAloneRatherThanDropped() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder), + new EmbeddingServiceOptions { MaxCharactersPerBatch = 10 }); + + string oversized = new string('x', 50); + float[][] vectors = await service.EmbedBatchAsync(["short", oversized, "tiny"]); + + Assert.Equal(3, vectors.Length); + Assert.Contains(recorder.Batches, batch => batch.Count == 1 && batch[0] == oversized); + } + + [Fact] + public async Task EmbedBatchAsync_AcrossBatches_ReturnsVectorsInInputOrder() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder), + new EmbeddingServiceOptions { MaxTextsPerBatch = 2, Concurrency = 4 }); + + // Deliberately including same-length texts and an anagram pair. Literals of seven + // distinct lengths would let this pass against a provider double whose vector only + // encoded the length, so the misordering it exists to catch has to be catchable + // without relying on the inputs happening to differ in an easy way. + string[] texts = ["alpha", "bravo", "ab", "ba", "cccccc", "e", "ggg"]; + + float[][] vectors = await service.EmbedBatchAsync(texts); + + Assert.Equal(texts.Length, vectors.Length); + for (int i = 0; i < texts.Length; i++) + { + Assert.Equal(BatchRecordingProvider.VectorFor(texts[i]), vectors[i]); + } + } + + [Fact] + public async Task EmbedAsync_SingleText_ReachesTheProviderAsABatchOfOne() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder)); + + float[] vector = await service.EmbedAsync("only this"); + + Assert.Single(recorder.Batches); + Assert.Equal(["only this"], recorder.Batches[0]); + Assert.Equal(BatchRecordingProvider.VectorFor("only this"), vector); + } + + #endregion + + #region Purpose + + [Theory] + [InlineData(EmbeddingPurpose.Query)] + [InlineData(EmbeddingPurpose.Document)] + public async Task EmbedBatchAsync_DeliversThePurposeToTheProvider(EmbeddingPurpose purpose) + { + // Asserted at the provider rather than at the call: the purpose travels through the + // request and the worker, and a service that dropped it would return correct-looking + // vectors from the wrong side of the model. + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder)); + + await service.EmbedBatchAsync(Texts(3), purpose); + + Assert.Single(recorder.Calls); + Assert.Equal(purpose, recorder.Calls[0].Purpose); + } + + [Theory] + [InlineData(EmbeddingPurpose.Query)] + [InlineData(EmbeddingPurpose.Document)] + public async Task EmbedAsync_DeliversThePurposeToTheProvider(EmbeddingPurpose purpose) + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder)); + + await service.EmbedAsync("a query", purpose); + + Assert.Single(recorder.Calls); + Assert.Equal(purpose, recorder.Calls[0].Purpose); + } + + [Fact] + public async Task EmbedBatchAsync_DefaultPurpose_IsDocument() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder)); + + await service.EmbedBatchAsync(Texts(2)); + + Assert.Single(recorder.Calls); + Assert.Equal(EmbeddingPurpose.Document, recorder.Calls[0].Purpose); + } + + [Fact] + public void VectorFor_DistinguishesTextsAPlainLengthAndSumWouldNot() + { + // The property the order assertions above rest on. Both of these pairs are equal under + // "length plus sum of characters", so a double encoding only that would hand back + // identical vectors for different texts and every ordering assertion would hold + // whatever order the vectors came back in. + Assert.NotEqual(BatchRecordingProvider.VectorFor("ab"), BatchRecordingProvider.VectorFor("ba")); + Assert.NotEqual(BatchRecordingProvider.VectorFor("ac"), BatchRecordingProvider.VectorFor("bb")); + } + + #endregion + + #region Batching only for providers that gain from it + + [Fact] + public async Task EmbedBatchAsync_ProviderThatDoesNotBatch_StillGetsOneTextPerCall() + { + // The batching policy says 64, but grouping texts for a provider that embeds one at a + // time would only move them into a single worker. It is the provider's answer, not the + // policy, that decides whether grouping happens at all. + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(recorder, supportsBatching: false), + new EmbeddingServiceOptions { MaxTextsPerBatch = 64 }); + + await service.EmbedBatchAsync(Texts(6)); + + Assert.Equal(6, recorder.Batches.Count); + Assert.All(recorder.Batches, batch => Assert.Single(batch)); + } + + [Fact] + public async Task EmbedBatchAsync_ProviderThatDoesNotBatch_KeepsWorkingAcrossWorkers() + { + // The regression this guards against is invisible to a count assertion: a batch looped + // inside one worker still returns the right vectors, just serially. The provider here + // only completes a call once two have arrived, so a serialised path never finishes and + // the timeout fails the test rather than hanging it. + ArrivalLatch latch = new ArrivalLatch(arrivalsRequired: 2); + await using EmbeddingService service = new EmbeddingService( + () => new LatchedSingleTextProvider(latch), + new EmbeddingServiceOptions { Concurrency = 2, MaxTextsPerBatch = 64 }); + + float[][] vectors = await service.EmbedBatchAsync(Texts(4)).WaitAsync(TimeSpan.FromSeconds(20)); + + Assert.Equal(4, vectors.Length); + } + + [Fact] + public async Task EmbedBatchAsync_ProviderThatBatches_GetsWholeBatches() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(recorder, supportsBatching: true), + new EmbeddingServiceOptions { MaxTextsPerBatch = 64 }); + + await service.EmbedBatchAsync(Texts(6)); + + Assert.Single(recorder.Batches); + Assert.Equal(6, recorder.Batches[0].Count); + } + + [Fact] + public void SupportsBatching_DefaultsToFalse() + { + // A provider written before the batch methods existed answers this without knowing it, + // and has to answer no. + IEmbeddingProvider provider = new TestEmbeddingProvider(); + + Assert.False(provider.SupportsBatching); + } + + #endregion + + #region The default interface implementation + + [Fact] + public async Task EmbedBatchAsync_ProviderImplementingOnlyTheSingleTextMethod_StillWorks() + { + // TestEmbeddingProvider implements EmbedAsync and nothing else, so everything here + // runs through the default implementations on IEmbeddingProvider. + TestEmbeddingProvider provider = new TestEmbeddingProvider(dimension: 16); + await using EmbeddingService service = new EmbeddingService(() => provider); + + float[][] vectors = await service.EmbedBatchAsync(["one", "two", "three"]); + + Assert.Equal(3, vectors.Length); + Assert.All(vectors, vector => Assert.Equal(16, vector.Length)); + Assert.Equal(3, provider.CallCount); + } + + [Fact] + public async Task EmbedBatchAsync_ProviderThatBatchesButDoesNotMeter_StillGetsOneCallPerBatch() + { + // The combination the interface documents: EmbedBatchAsync overridden, + // EmbedBatchWithUsageAsync inherited. The service calls the usage method, whose default + // routes into the overridden one, so the batch stays a single call and reports nothing. + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchingOnlyProvider(recorder), + new EmbeddingServiceOptions { MaxTextsPerBatch = 64 }); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(Texts(5)); + + Assert.Equal(5, result.Vectors.Count); + Assert.Null(result.Usage); + Assert.Equal([5], recorder.Batches.Select(batch => batch.Count)); + Assert.Equal(0, recorder.SingleCallCount); + } + + [Fact] + public async Task DefaultEmbedBatchAsync_CalledDirectlyOnTheInterface_LoopsTheSingleTextMethod() + { + TestEmbeddingProvider provider = new TestEmbeddingProvider(dimension: 16); + + float[][] vectors = await ((IEmbeddingProvider)provider).EmbedBatchAsync(["a", "b"]); + + // The call count is what says "loops the single-text method". Comparing the batch's + // vectors against a fresh EmbedAsync would pass for any implementation that is merely + // deterministic, including one that never called EmbedAsync at all. + Assert.Equal(2, provider.CallCount); + Assert.Equal(2, vectors.Length); + } + + [Fact] + public async Task DefaultEmbedBatchWithUsageAsync_ReportsNoUsage() + { + IEmbeddingProvider provider = new TestEmbeddingProvider(); + + EmbeddingResult result = await provider.EmbedBatchWithUsageAsync(["a", "b"]); + + Assert.Equal(2, result.Vectors.Count); + Assert.Null(result.Usage); + } + + [Fact] + public async Task DefaultEmbedBatchAsync_NullList_Throws() + { + IEmbeddingProvider provider = new TestEmbeddingProvider(); + + await Assert.ThrowsAsync(() => provider.EmbedBatchAsync(null!)); + } + + [Fact] + public async Task DefaultEmbedBatchAsync_EmptyList_ReturnsEmptyAndCallsNothing() + { + TestEmbeddingProvider provider = new TestEmbeddingProvider(); + + float[][] vectors = await ((IEmbeddingProvider)provider).EmbedBatchAsync([]); + + Assert.Empty(vectors); + Assert.Equal(0, provider.CallCount); + } + + #endregion + + #region Validation + + [Fact] + public async Task EmbedBatchAsync_NullElement_Throws() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder)); + + await Assert.ThrowsAsync(() => + service.EmbedBatchAsync(["fine", null!])); + + Assert.Empty(recorder.Batches); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_NullList_Throws() + { + // Its own public entry point with its own documented ArgumentNullException, not just + // the path EmbedBatchAsync happens to reach it through. + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider()); + + await Assert.ThrowsAsync(() => service.EmbedBatchWithUsageAsync(null!)); + } + + [Fact] + public async Task MaxCharactersPerBatch_ProviderThatDoesNotBatch_ChangesNothing() + { + // The README says the character limit is only reached by a provider that batches, + // because one text is already a whole request for a provider that does not. A limit + // low enough to split these texts several ways still leaves one text per call. + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(recorder, supportsBatching: false), + new EmbeddingServiceOptions { MaxTextsPerBatch = 64, MaxCharactersPerBatch = 1 }); + + await service.EmbedBatchAsync(Texts(4, lengthEach: 10)); + + Assert.Equal(4, recorder.Batches.Count); + Assert.All(recorder.Batches, batch => Assert.Single(batch)); + } + + [Fact] + public async Task EmbedBatchAsync_ProviderReturnsWrongNumberOfVectors_Fails() + { + // Vectors are matched to texts by position, so a short response is a lost mapping + // rather than a partial success. + await using EmbeddingService service = new EmbeddingService(() => new ShortResponseProvider()); + + await Assert.ThrowsAsync(() => + service.EmbedBatchAsync(["a", "b", "c"])); + } + + [Fact] + public async Task Constructor_FactoryThrowsPartway_DisposesTheProvidersItAlreadyBuilt() + { + // The constructor never returns, so nothing else can dispose these. + List built = new List(); + + Assert.Throws(() => + new EmbeddingService(() => + { + if (built.Count == 2) + throw new InvalidOperationException("factory failed"); + + TestEmbeddingProvider provider = new TestEmbeddingProvider(); + built.Add(provider); + return provider; + }, + new EmbeddingServiceOptions { Concurrency = 4 })); + + Assert.Equal(2, built.Count); + foreach (TestEmbeddingProvider provider in built) + { + await Assert.ThrowsAsync(() => provider.EmbedAsync("anything")); + } + } + + #endregion + + #region Cancellation + + [Fact] + public async Task EmbedBatchAsync_CancelledWhileInFlight_ReachesTheProvider() + { + // Completing the caller's task is not enough: the batch itself is what a remote API + // bills, so the token has to reach the call that is running. + CancellationTokenSource cts = new CancellationTokenSource(); + ObservedCancellationProvider provider = new ObservedCancellationProvider(); + + await using EmbeddingService service = new EmbeddingService(() => provider); + + Task embedTask = service.EmbedBatchAsync(Texts(3), EmbeddingPurpose.Document, cts.Token); + + await provider.CallStarted.Task.WaitAsync(TimeSpan.FromSeconds(20)); + await cts.CancelAsync(); + + await Assert.ThrowsAnyAsync(() => embedTask); + Assert.True(await provider.SawCancellation.Task.WaitAsync(TimeSpan.FromSeconds(20))); + } + + [Fact] + public async Task EmbedBatchAsync_ProviderCancelsAgainstItsOwnToken_SurfacesTheProvidersFailure() + { + // The third of the worker's three OperationCanceledException catches, and the only one + // nothing else covers. Neither the caller's token nor the service's is cancelled here, + // so this is a failed call wearing a cancellation's exception type. The worker has to + // pass the provider's own exception through rather than replacing it with a + // cancellation of its own — the message is the only thing that tells a caller their + // provider timed out rather than something in here deciding to stop. + await using EmbeddingService service = new EmbeddingService(() => new SelfCancellingProvider()); + + OperationCanceledException failure = await Assert.ThrowsAnyAsync(() => + service.EmbedBatchAsync(Texts(2))); + + Assert.Equal(SelfCancellingProvider.FailureMessage, failure.Message); + } + + [Fact] + public async Task EmbedBatchAsync_TokenAlreadyCancelled_Throws() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider(recorder)); + + CancellationToken cancelled = new CancellationToken(true); + + await Assert.ThrowsAnyAsync(() => + service.EmbedBatchAsync(Texts(2), EmbeddingPurpose.Document, cancelled)); + } + + [Fact] + public async Task EmbedBatchAsync_QueueFullAndServiceDisposed_FailsTheCallerRatherThanStranding() + { + // The blocked-writer path: capacity 1 with a worker held open, so later batches are + // still queued or waiting to be written when the service is disposed under them. + // Every one of them has to end up somewhere — a batch left in the channel after the + // workers stop would leave this caller awaiting a result nothing will ever produce. + ArrivalLatch neverReached = new ArrivalLatch(arrivalsRequired: 99); + EmbeddingService service = new EmbeddingService( + () => new LatchedSingleTextProvider(neverReached), + new EmbeddingServiceOptions { Concurrency = 1, ChannelCapacity = 1, MaxTextsPerBatch = 1 }); + + Task blocked = service.EmbedBatchAsync(Texts(8)); + + await service.DisposeAsync().AsTask().WaitAsync(TimeSpan.FromSeconds(30)); + + // Bounded so a regression fails this test instead of hanging the run — and the failure + // has to be a shutdown failure. A TimeoutException from the bound would mean the + // caller was left waiting, which is the very thing under test, so it must not count + // as the exception this expects. + Exception failure = await Assert.ThrowsAnyAsync(() => blocked.WaitAsync(TimeSpan.FromSeconds(30))); + + Assert.True( + failure is ObjectDisposedException or OperationCanceledException, + $"Expected the call to fail because the service was disposed, but it failed with {failure.GetType().Name}: {failure.Message}"); + } + + #endregion + } +} diff --git a/VectorSharp.Embedding.Tests/EmbeddingProviderDoubles.cs b/VectorSharp.Embedding.Tests/EmbeddingProviderDoubles.cs new file mode 100644 index 0000000..b80fcaa --- /dev/null +++ b/VectorSharp.Embedding.Tests/EmbeddingProviderDoubles.cs @@ -0,0 +1,252 @@ +namespace VectorSharp.Embedding.Tests +{ + // The convention for this project: an IEmbeddingProvider double is a top-level type in a file + // of its own, never nested inside the test class that happens to use it first. Doubles get + // reused across files as soon as a second test needs the same behaviour, and one nested in a + // test class is invisible from the next file — which is how this project ended up with two + // hand-rolled recorders doing the same job. TestEmbeddingProvider and BatchRecordingProvider + // are the two big enough to warrant their own files; the small ones live here together. + + /// + /// Releases every waiter once enough calls have arrived, counting arrivals cumulatively + /// rather than tracking how many are inside at once — a call that has already returned + /// still counts. That is all the tests using it need, because a serialised path never + /// produces a second arrival while the first is held. + /// + internal sealed class ArrivalLatch + { + private readonly int _arrivalsRequired; + private readonly TaskCompletionSource _reached = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + private int _arrivals; + + public ArrivalLatch(int arrivalsRequired) + { + _arrivalsRequired = arrivalsRequired; + } + + public Task ArriveAndWaitAsync() + { + if (Interlocked.Increment(ref _arrivals) >= _arrivalsRequired) + _reached.TrySetResult(); + + return _reached.Task; + } + } + + /// + /// A provider that only implements the single-text method, standing in for a local model. + /// + internal sealed class LatchedSingleTextProvider : IEmbeddingProvider + { + private readonly ArrivalLatch _latch; + + public int Dimension => 8; + + public LatchedSingleTextProvider(ArrivalLatch latch) + { + _latch = latch; + } + + public async Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + // Bounded and cancellable on purpose. If the calls are ever serialised the latch + // is never reached, and a wait that could not be broken out of would hang the + // service's disposal — and with it the whole test run — instead of failing. + await _latch.ArriveAndWaitAsync().WaitAsync(TimeSpan.FromSeconds(10), cancellationToken); + return new float[8]; + } + + public void Dispose() { } + } + + /// + /// Reports whether the token the service handed it was the one the caller cancelled. + /// + internal sealed class ObservedCancellationProvider : IEmbeddingProvider + { + public int Dimension => 8; + + public bool SupportsBatching => true; + + public TaskCompletionSource CallStarted { get; } = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + public TaskCompletionSource SawCancellation { get; } = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + return Task.FromResult(new float[8]); + } + + public async Task EmbedBatchWithUsageAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + CallStarted.TrySetResult(); + + try + { + await Task.Delay(Timeout.Infinite, cancellationToken); + } + catch (OperationCanceledException) + { + SawCancellation.TrySetResult(true); + throw; + } + + throw new InvalidOperationException("Unreachable: the delay never completes on its own."); + } + + public void Dispose() { } + } + + /// + /// Fails a call by cancelling a token of its own, which no one outside it holds. + /// + /// + /// The case a library provider produces without meaning to: an HTTP client with its own + /// per-request timeout raises from that timeout, not + /// from the token it was handed. It is a failed call, and the service must not report it to + /// the caller as a cancellation the caller asked for. + /// + internal sealed class SelfCancellingProvider : IEmbeddingProvider + { + internal const string FailureMessage = "the provider's own token timed out"; + + public int Dimension => 8; + + public bool SupportsBatching => true; + + public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + return Task.FromResult(new float[8]); + } + + public Task EmbedBatchWithUsageAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + using CancellationTokenSource ownSource = new CancellationTokenSource(); + ownSource.Cancel(); + + throw new OperationCanceledException(FailureMessage, ownSource.Token); + } + + public void Dispose() { } + } + + /// + /// Batches but meters nothing: overrides only the plain batch method and inherits the usage + /// one. Records through the shared rather than keeping its own + /// tally, so there is one implementation of "what did the providers receive". + /// + internal sealed class BatchingOnlyProvider : IEmbeddingProvider + { + private readonly BatchRecorder _recorder; + + public int Dimension => 8; + + public bool SupportsBatching => true; + + public BatchingOnlyProvider(BatchRecorder recorder) + { + _recorder = recorder; + } + + public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + _recorder.RecordSingleCall(); + return Task.FromResult(new float[8]); + } + + public Task EmbedBatchAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + _recorder.RecordBatch(texts, purpose); + + float[][] vectors = new float[texts.Count][]; + for (int i = 0; i < texts.Count; i++) + { + vectors[i] = new float[8]; + } + + return Task.FromResult(vectors); + } + + public void Dispose() { } + } + + /// + /// A provider that drops a vector, the way a remote API would if it silently skipped an input. + /// + internal sealed class ShortResponseProvider : IEmbeddingProvider + { + public int Dimension => 8; + + public bool SupportsBatching => true; + + public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + return Task.FromResult(new float[8]); + } + + public Task EmbedBatchWithUsageAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + float[][] vectors = new float[Math.Max(0, texts.Count - 1)][]; + for (int i = 0; i < vectors.Length; i++) + { + vectors[i] = new float[8]; + } + + return Task.FromResult(new EmbeddingResult { Vectors = vectors }); + } + + public void Dispose() { } + } + + /// + /// Throws from every call, for the paths that check a provider's failure reaches the caller. + /// + internal sealed class FailingProvider : IEmbeddingProvider + { + public int Dimension => 768; + + public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + throw new InvalidOperationException("Intentional test failure"); + } + + public void Dispose() { } + } + + /// + /// A count shared across the provider instances the service builds, one per worker. + /// + internal sealed class SharedCounter + { + private int _value; + + public int Value => Volatile.Read(ref _value); + + public void Increment() => Interlocked.Increment(ref _value); + } + + /// + /// Counts every call into , so a test can see the total across all + /// of the service's workers. + /// + internal sealed class CountingProvider : IEmbeddingProvider + { + private readonly SharedCounter _counter; + + public int Dimension { get; } + + public CountingProvider(int dimension, SharedCounter counter) + { + Dimension = dimension; + _counter = counter; + } + + public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + _counter.Increment(); + return Task.FromResult(new float[Dimension]); + } + + public void Dispose() { } + } +} diff --git a/VectorSharp.Embedding.Tests/EmbeddingServiceOptionsTests.cs b/VectorSharp.Embedding.Tests/EmbeddingServiceOptionsTests.cs index 5fa4603..9c1dfd8 100644 --- a/VectorSharp.Embedding.Tests/EmbeddingServiceOptionsTests.cs +++ b/VectorSharp.Embedding.Tests/EmbeddingServiceOptionsTests.cs @@ -33,5 +33,43 @@ public void ChannelCapacity_CustomValue_IsRetained() Assert.Equal(500, options.ChannelCapacity); } + + // Both defaults are quoted as fact in the XML docs, the README and the package + // description, so they are pinned here as literals rather than left to drift. + [Fact] + public void DefaultMaxTextsPerBatch_IsSixtyFour() + { + EmbeddingServiceOptions options = new EmbeddingServiceOptions(); + + Assert.Equal(64, options.MaxTextsPerBatch); + } + + [Fact] + public void DefaultMaxCharactersPerBatch_IsOneHundredThousand() + { + EmbeddingServiceOptions options = new EmbeddingServiceOptions(); + + Assert.Equal(100000, options.MaxCharactersPerBatch); + } + + [Fact] + public void MaxTextsPerBatch_CustomValue_IsRetained() + { + EmbeddingServiceOptions options = new EmbeddingServiceOptions { MaxTextsPerBatch = 8 }; + + Assert.Equal(8, options.MaxTextsPerBatch); + } + + [Fact] + public void MaxCharactersPerBatch_CustomValue_IsRetained() + { + EmbeddingServiceOptions options = new EmbeddingServiceOptions { MaxCharactersPerBatch = 2048 }; + + Assert.Equal(2048, options.MaxCharactersPerBatch); + } + + // The service constructor's rejection of out-of-range options is tested in + // EmbeddingServiceTests, where the rest of its validation lives. This file is about the + // options type itself: its defaults, and that it keeps what it is given. } } diff --git a/VectorSharp.Embedding.Tests/EmbeddingServiceTests.cs b/VectorSharp.Embedding.Tests/EmbeddingServiceTests.cs index 08db5c1..8d4e507 100644 --- a/VectorSharp.Embedding.Tests/EmbeddingServiceTests.cs +++ b/VectorSharp.Embedding.Tests/EmbeddingServiceTests.cs @@ -19,28 +19,33 @@ public void Constructor_NullFactory_Throws() new EmbeddingService(null!)); } - [Fact] - public void Constructor_ZeroConcurrency_Throws() - { - Assert.Throws(() => - new EmbeddingService(() => new TestEmbeddingProvider(), - new EmbeddingServiceOptions { Concurrency = 0 })); - } - - [Fact] - public void Constructor_NegativeConcurrency_Throws() - { - Assert.Throws(() => - new EmbeddingService(() => new TestEmbeddingProvider(), - new EmbeddingServiceOptions { Concurrency = -1 })); - } - - [Fact] - public void Constructor_ZeroChannelCapacity_Throws() + // Every option the constructor range-checks, checked the same way. One theory rather than + // a test per option, so that an option added later without its guard is a missing row + // here rather than a case nobody noticed was absent — which is how MaxTextsPerBatch and + // MaxCharactersPerBatch ended up with half the coverage Concurrency had. + public static TheoryData OutOfRangeOptions => + new TheoryData + { + { "Concurrency zero", new EmbeddingServiceOptions { Concurrency = 0 } }, + { "Concurrency negative", new EmbeddingServiceOptions { Concurrency = -1 } }, + { "ChannelCapacity zero", new EmbeddingServiceOptions { ChannelCapacity = 0 } }, + { "ChannelCapacity negative", new EmbeddingServiceOptions { ChannelCapacity = -1 } }, + { "MaxTextsPerBatch zero", new EmbeddingServiceOptions { MaxTextsPerBatch = 0 } }, + { "MaxTextsPerBatch negative", new EmbeddingServiceOptions { MaxTextsPerBatch = -1 } }, + { "MaxCharactersPerBatch zero", new EmbeddingServiceOptions { MaxCharactersPerBatch = 0 } }, + { "MaxCharactersPerBatch negative", new EmbeddingServiceOptions { MaxCharactersPerBatch = -1 } }, + }; + + /// Names the case in the test output. xUnit cannot render the + /// options object itself, so without this every row reports under the same name and a + /// failure does not say which option let a bad value through. + /// The options the constructor has to reject. + [Theory] + [MemberData(nameof(OutOfRangeOptions))] + public void Constructor_OptionBelowOne_Throws(string description, EmbeddingServiceOptions options) { Assert.Throws(() => - new EmbeddingService(() => new TestEmbeddingProvider(), - new EmbeddingServiceOptions { ChannelCapacity = 0 })); + new EmbeddingService(() => new TestEmbeddingProvider(), options)); } #endregion @@ -234,50 +239,5 @@ public async Task DisposeAsync_Idempotent_SecondCallNoOp() } #endregion - - #region Helpers - - private sealed class FailingProvider : IEmbeddingProvider - { - public int Dimension => 768; - - public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) - { - throw new InvalidOperationException("Intentional test failure"); - } - - public void Dispose() { } - } - - private sealed class SharedCounter - { - private int _value; - public int Value => _value; - public void Increment() => Interlocked.Increment(ref _value); - } - - private sealed class CountingProvider : IEmbeddingProvider - { - private readonly SharedCounter _counter; - - public int Dimension { get; } - - public CountingProvider(int dimension, SharedCounter counter) - { - Dimension = dimension; - _counter = counter; - } - - public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) - { - _counter.Increment(); - float[] result = new float[Dimension]; - return Task.FromResult(result); - } - - public void Dispose() { } - } - - #endregion } } diff --git a/VectorSharp.Embedding.Tests/EmbeddingUsageTests.cs b/VectorSharp.Embedding.Tests/EmbeddingUsageTests.cs new file mode 100644 index 0000000..dee2707 --- /dev/null +++ b/VectorSharp.Embedding.Tests/EmbeddingUsageTests.cs @@ -0,0 +1,334 @@ +namespace VectorSharp.Embedding.Tests +{ + public class EmbeddingUsageTests + { + private static EmbeddingUsage TokensPerText(IReadOnlyList batch, string? model = "test-model") + { + return new EmbeddingUsage { TokenCount = batch.Count * 10, Model = model }; + } + + #region Nothing reported + + [Fact] + public async Task EmbedBatchWithUsageAsync_ProviderReportsNothing_UsageIsNullNotZero() + { + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider()); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["a", "b"]); + + Assert.Null(result.Usage); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_ProviderMetersNothingAtAll_UsageIsNullNotZero() + { + // The local-model case end to end: a provider that implements only the single-text + // method has nothing to report, and the service must not turn that into a zero a + // caller could mistake for a free call. + await using EmbeddingService service = new EmbeddingService(() => new TestEmbeddingProvider()); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["a", "b"]); + + Assert.Equal(2, result.Vectors.Count); + Assert.Null(result.Usage); + } + + [Fact] + public async Task EmbedWithUsageAsync_ProviderReportsNothing_UsageIsNullNotZero() + { + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider()); + + EmbeddingResult result = await service.EmbedWithUsageAsync("a"); + + Assert.Null(result.Usage); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_EmptyList_ReturnsNoVectorsAndNoUsage() + { + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(usageForBatch: batch => TokensPerText(batch))); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync([]); + + Assert.Empty(result.Vectors); + Assert.Null(result.Usage); + } + + #endregion + + #region Reported usage + + [Fact] + public async Task EmbedBatchWithUsageAsync_ProviderReportsZeroTokens_UsageSurvivesAsAReportedZero() + { + // The other half of "null means unknown, never free": a reported zero is a real + // measurement and has to arrive as one, or the distinction the null carries is empty. + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(usageForBatch: batch => new EmbeddingUsage { TokenCount = 0, Model = "free-tier" })); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["a", "b"]); + + Assert.NotNull(result.Usage); + Assert.Equal(0, result.Usage.TokenCount); + Assert.Equal("free-tier", result.Usage.Model); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_ProviderReportsUsage_PassesItThrough() + { + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(usageForBatch: batch => TokensPerText(batch))); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["a", "b", "c"]); + + Assert.NotNull(result.Usage); + Assert.Equal(30, result.Usage.TokenCount); + Assert.Equal("test-model", result.Usage.Model); + } + + [Fact] + public async Task EmbedWithUsageAsync_ProviderReportsUsage_PassesItThrough() + { + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(usageForBatch: batch => TokensPerText(batch))); + + EmbeddingResult result = await service.EmbedWithUsageAsync("a"); + + Assert.Single(result.Vectors); + Assert.NotNull(result.Usage); + Assert.Equal(10, result.Usage.TokenCount); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_SplitIntoBatches_SumsTheReportedTokens() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(recorder, batch => TokensPerText(batch)), + new EmbeddingServiceOptions { MaxTextsPerBatch = 2 }); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["a", "b", "c", "d", "e"]); + + Assert.Equal(3, recorder.Batches.Count); + Assert.NotNull(result.Usage); + Assert.Equal(50, result.Usage.TokenCount); + Assert.Equal("test-model", result.Usage.Model); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_EveryBatchNamesNoModel_KeepsTheTotalAndReportsNoModel() + { + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(recorder, batch => TokensPerText(batch, model: null)), + new EmbeddingServiceOptions { MaxTextsPerBatch = 2 }); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["a", "b", "c", "d"]); + + Assert.Equal(2, recorder.Batches.Count); + Assert.NotNull(result.Usage); + Assert.Equal(40, result.Usage.TokenCount); + Assert.Null(result.Usage.Model); + } + + #endregion + + #region Combine, directly + + [Fact] + public void Combine_NothingToCombine_IsNullRatherThanAReportedZero() + { + // Reachable through the service only as the empty-input case, which returns before it + // gets here — but it is the case the whole type guards: totalling nothing up would + // produce a usage record saying the call cost zero, which means "this was free" + // rather than "nothing was reported". + Assert.Null(EmbeddingUsage.Combine([])); + } + + [Fact] + public void Combine_OneBatch_PassesItThrough() + { + EmbeddingUsage? combined = EmbeddingUsage.Combine([new EmbeddingUsage { TokenCount = 7, Model = "m" }]); + + Assert.NotNull(combined); + Assert.Equal(7, combined.TokenCount); + Assert.Equal("m", combined.Model); + } + + [Fact] + public void Combine_EveryBatchReported_SumsTheTokens() + { + EmbeddingUsage? combined = EmbeddingUsage.Combine( + [ + new EmbeddingUsage { TokenCount = 10, Model = "m" }, + new EmbeddingUsage { TokenCount = 5, Model = "m" } + ]); + + Assert.NotNull(combined); + Assert.Equal(15, combined.TokenCount); + Assert.Equal("m", combined.Model); + } + + [Fact] + public void Combine_AnyBatchReportedNothing_IsNull() + { + EmbeddingUsage? combined = EmbeddingUsage.Combine( + [ + new EmbeddingUsage { TokenCount = 10, Model = "m" }, + null + ]); + + Assert.Null(combined); + } + + [Fact] + public void Combine_LastBatchReportedNothing_IsNull() + { + // The other order, because a check that returned early on the first null would pass + // the case above and still total a partial spend when the gap comes last. + EmbeddingUsage? combined = EmbeddingUsage.Combine( + [ + null, + new EmbeddingUsage { TokenCount = 10, Model = "m" } + ]); + + Assert.Null(combined); + } + + [Fact] + public void Combine_BatchesDisagreeOnTheModel_KeepsTheTotalAndDropsTheName() + { + EmbeddingUsage? combined = EmbeddingUsage.Combine( + [ + new EmbeddingUsage { TokenCount = 10, Model = "m1" }, + new EmbeddingUsage { TokenCount = 10, Model = "m2" } + ]); + + Assert.NotNull(combined); + Assert.Equal(20, combined.TokenCount); + Assert.Null(combined.Model); + } + + [Fact] + public void Combine_FirstBatchNamedNoModel_DoesNotBorrowALaterOnesName() + { + EmbeddingUsage? combined = EmbeddingUsage.Combine( + [ + new EmbeddingUsage { TokenCount = 10, Model = null }, + new EmbeddingUsage { TokenCount = 10, Model = "m" } + ]); + + Assert.NotNull(combined); + Assert.Equal(20, combined.TokenCount); + Assert.Null(combined.Model); + } + + [Fact] + public void Combine_EveryBatchNamedNoModel_KeepsTheTotalAndReportsNoModel() + { + EmbeddingUsage? combined = EmbeddingUsage.Combine( + [ + new EmbeddingUsage { TokenCount = 10, Model = null }, + new EmbeddingUsage { TokenCount = 10, Model = null } + ]); + + Assert.NotNull(combined); + Assert.Equal(20, combined.TokenCount); + Assert.Null(combined.Model); + } + + #endregion + + #region Combining across batches + + [Fact] + public async Task EmbedBatchWithUsageAsync_OnlySomeBatchesReport_ReturnsNullRatherThanAPartialTotal() + { + // A total assembled from some of the batches would understate the spend, and + // under-reporting is the failure mode that matters when the caller is billed. + // The discriminator is a literal from this test's own input, so the mixed split it + // depends on cannot quietly stop happening. + BatchRecorder recorder = new BatchRecorder(); + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(recorder, batch => batch.Contains("metered") ? TokensPerText(batch) : null), + new EmbeddingServiceOptions { MaxTextsPerBatch = 1 }); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["metered", "silent", "also silent"]); + + // The premise of the test: some batches reported and some did not. + Assert.Equal(3, recorder.Batches.Count); + Assert.Contains(recorder.Batches, batch => batch.Contains("metered")); + Assert.Contains(recorder.Batches, batch => !batch.Contains("metered")); + + Assert.Null(result.Usage); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_FirstBatchNamesNoModel_DoesNotBorrowAnotherBatchsModel() + { + // The two are not the same model just because one of them declined to say. Attributing + // the whole call to the only name that turned up would misreport spend per model. + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(usageForBatch: batch => TokensPerText(batch, model: batch.Contains("unnamed") ? null : "model-x")), + new EmbeddingServiceOptions { MaxTextsPerBatch = 1, Concurrency = 1 }); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["unnamed", "named", "also named"]); + + Assert.NotNull(result.Usage); + Assert.Equal(30, result.Usage.TokenCount); + Assert.Null(result.Usage.Model); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_BatchesReportDifferentModels_DropsTheModelButKeepsTheTotal() + { + int batchNumber = 0; + await using EmbeddingService service = new EmbeddingService( + () => new BatchRecordingProvider(usageForBatch: batch => new EmbeddingUsage + { + TokenCount = 10, + Model = $"model-{Interlocked.Increment(ref batchNumber)}" + }), + new EmbeddingServiceOptions { MaxTextsPerBatch = 1 }); + + EmbeddingResult result = await service.EmbedBatchWithUsageAsync(["a", "b", "c"]); + + Assert.NotNull(result.Usage); + Assert.Equal(30, result.Usage.TokenCount); + Assert.Null(result.Usage.Model); + } + + #endregion + + #region EmbedWithUsageAsync guards + + [Fact] + public async Task EmbedWithUsageAsync_NullText_Throws() + { + await using EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider()); + + await Assert.ThrowsAsync(() => service.EmbedWithUsageAsync(null!)); + } + + [Fact] + public async Task EmbedWithUsageAsync_AfterDispose_Throws() + { + EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider()); + await service.DisposeAsync(); + + await Assert.ThrowsAsync(() => service.EmbedWithUsageAsync("a")); + } + + [Fact] + public async Task EmbedBatchWithUsageAsync_AfterDispose_Throws() + { + EmbeddingService service = new EmbeddingService(() => new BatchRecordingProvider()); + await service.DisposeAsync(); + + await Assert.ThrowsAsync(() => service.EmbedBatchWithUsageAsync(["a"])); + } + + #endregion + } +} diff --git a/VectorSharp.Embedding/BatchingPolicy.cs b/VectorSharp.Embedding/BatchingPolicy.cs new file mode 100644 index 0000000..da411c3 --- /dev/null +++ b/VectorSharp.Embedding/BatchingPolicy.cs @@ -0,0 +1,93 @@ +namespace VectorSharp.Embedding +{ + /// + /// Decides how a caller's texts are divided into the batches the service sends to a provider. + /// + /// + /// Separate from so that the splitting rule can be exercised on + /// its own. Reached through the service it needs a service, a factory, a provider and a + /// recording double to observe a partition, which buries the rule being checked. + /// + /// Internal rather than public: this is how the service behaves, not a knob a caller sets. + /// Callers configure it through . + /// + internal sealed class BatchingPolicy + { + /// + /// The largest number of texts one batch may hold. This is the effective limit rather than + /// the configured one: see . + /// + internal int MaxTextsPerBatch { get; } + + /// + /// The largest total number of characters one batch may hold, except where a single text + /// exceeds it on its own. + /// + internal int MaxCharactersPerBatch { get; } + + internal BatchingPolicy(int maxTextsPerBatch, int maxCharactersPerBatch) + { + if (maxTextsPerBatch < 1) + throw new ArgumentOutOfRangeException(nameof(maxTextsPerBatch), "MaxTextsPerBatch must be at least 1."); + + if (maxCharactersPerBatch < 1) + throw new ArgumentOutOfRangeException(nameof(maxCharactersPerBatch), "MaxCharactersPerBatch must be at least 1."); + + MaxTextsPerBatch = maxTextsPerBatch; + MaxCharactersPerBatch = maxCharactersPerBatch; + } + + /// + /// Builds the policy the service will actually apply to a given provider. + /// + /// + /// A provider that answers false to gets + /// one text per batch whatever the options say, because grouping for it would run the + /// texts serially inside one worker instead of spreading them across the pool. See + /// for the whole of that reasoning. + /// + internal static BatchingPolicy For(EmbeddingServiceOptions options, bool providerSupportsBatching) + { + return new BatchingPolicy( + providerSupportsBatching ? options.MaxTextsPerBatch : 1, + options.MaxCharactersPerBatch); + } + + /// + /// Splits the texts into batches that respect both limits. + /// + /// + /// A single text over the character limit becomes a batch of its own rather than being + /// dropped or split, since only the caller can decide how to shorten it. Order is + /// preserved: the batches concatenated in order reproduce the input, which is what lets + /// the service reassemble the vectors by position. + /// + internal IReadOnlyList> Split(IReadOnlyList texts) + { + List> batches = new List>(); + List currentBatch = new List(); + long currentCharacters = 0; + + foreach (string text in texts) + { + bool wouldExceedCount = currentBatch.Count >= MaxTextsPerBatch; + bool wouldExceedCharacters = currentBatch.Count > 0 && currentCharacters + text.Length > MaxCharactersPerBatch; + + if (wouldExceedCount || wouldExceedCharacters) + { + batches.Add(currentBatch); + currentBatch = new List(); + currentCharacters = 0; + } + + currentBatch.Add(text); + currentCharacters += text.Length; + } + + if (currentBatch.Count > 0) + batches.Add(currentBatch); + + return batches; + } + } +} diff --git a/VectorSharp.Embedding/EmbeddingRequest.cs b/VectorSharp.Embedding/EmbeddingRequest.cs index 6ac70a3..b35565f 100644 --- a/VectorSharp.Embedding/EmbeddingRequest.cs +++ b/VectorSharp.Embedding/EmbeddingRequest.cs @@ -2,23 +2,33 @@ namespace VectorSharp.Embedding { /// /// Internal message type representing a pending embedding request in the channel. + /// Carries a whole batch rather than a single text, so one batch is one unit of work for + /// one worker and reaches the provider as one call. /// internal sealed class EmbeddingRequest { /// - /// The text to produce an embedding for. + /// The texts to produce embeddings for. Already sized to the batching policy by the + /// service, so a worker hands this list to the provider as it stands. /// - public required string Text { get; init; } + public required IReadOnlyList Texts { get; init; } /// - /// The intended purpose of the embedding. + /// The intended purpose of the embeddings. /// public required EmbeddingPurpose Purpose { get; init; } + /// + /// The token the caller passed. Carried through so the worker can stop the provider call + /// itself: completing the caller's task is not enough when the call in flight is a whole + /// batch being billed by a remote API. + /// + public required CancellationToken CallerToken { get; init; } + /// /// The completion source that the caller awaits. The worker sets the result - /// after producing the embedding, or sets an exception if inference fails. + /// after producing the embeddings, or sets an exception if inference fails. /// - public required TaskCompletionSource CompletionSource { get; init; } + public required TaskCompletionSource CompletionSource { get; init; } } } diff --git a/VectorSharp.Embedding/EmbeddingResult.cs b/VectorSharp.Embedding/EmbeddingResult.cs new file mode 100644 index 0000000..1204271 --- /dev/null +++ b/VectorSharp.Embedding/EmbeddingResult.cs @@ -0,0 +1,23 @@ +namespace VectorSharp.Embedding +{ + /// + /// The vectors produced for a set of texts, together with what the provider reported spending + /// on them if it reported anything. + /// + public sealed class EmbeddingResult + { + /// + /// Gets the embeddings, one per input text, in the order the texts were given. + /// + public required IReadOnlyList Vectors { get; init; } + + /// + /// Gets what the call cost, or null when the provider reported nothing. + /// + /// + /// Null means unknown, never free. A caller metering spend has to treat it as a gap in its + /// accounting rather than as a zero, which is why an unreported count is not defaulted. + /// + public EmbeddingUsage? Usage { get; init; } + } +} diff --git a/VectorSharp.Embedding/EmbeddingService.cs b/VectorSharp.Embedding/EmbeddingService.cs index 954a20d..efedbef 100644 --- a/VectorSharp.Embedding/EmbeddingService.cs +++ b/VectorSharp.Embedding/EmbeddingService.cs @@ -13,6 +13,7 @@ public sealed class EmbeddingService : IAsyncDisposable private readonly Task[] _workerTasks; private readonly IEmbeddingProvider[] _providers; private readonly CancellationTokenSource _cts; + private readonly BatchingPolicy _batchingPolicy; private int _disposed; /// @@ -28,7 +29,7 @@ public sealed class EmbeddingService : IAsyncDisposable /// Called once per worker. Each worker owns its own provider instance. /// Configuration options. If null, defaults are used. /// Thrown when providerFactory is null. - /// Thrown when concurrency or channel capacity is less than 1. + /// Thrown when concurrency, channel capacity or either batch limit is less than 1. public EmbeddingService(Func providerFactory, EmbeddingServiceOptions? options = null) { ArgumentNullException.ThrowIfNull(providerFactory); @@ -41,6 +42,12 @@ public EmbeddingService(Func providerFactory, EmbeddingServi if (effectiveOptions.ChannelCapacity < 1) throw new ArgumentOutOfRangeException(nameof(options), "ChannelCapacity must be at least 1."); + if (effectiveOptions.MaxTextsPerBatch < 1) + throw new ArgumentOutOfRangeException(nameof(options), "MaxTextsPerBatch must be at least 1."); + + if (effectiveOptions.MaxCharactersPerBatch < 1) + throw new ArgumentOutOfRangeException(nameof(options), "MaxCharactersPerBatch must be at least 1."); + BoundedChannelOptions channelOptions = new BoundedChannelOptions(effectiveOptions.ChannelCapacity) { FullMode = BoundedChannelFullMode.Wait, @@ -53,14 +60,42 @@ public EmbeddingService(Func providerFactory, EmbeddingServi _providers = new IEmbeddingProvider[effectiveOptions.Concurrency]; _workerTasks = new Task[effectiveOptions.Concurrency]; - // Create provider instances - for (int i = 0; i < effectiveOptions.Concurrency; i++) + // Create provider instances. If the factory fails partway, the ones already built are + // disposed here: the constructor never returns, so the caller has no service to + // dispose and an ONNX session would keep its native memory with nothing referencing it. + int created = 0; + try + { + for (; created < effectiveOptions.Concurrency; created++) + { + _providers[created] = providerFactory(); + } + } + catch { - _providers[i] = providerFactory(); + for (int i = 0; i < created; i++) + { + try + { + _providers[i].Dispose(); + } + catch + { + // The factory's failure is what the caller needs to see, so a provider + // failing to clean up must not replace it. + } + } + + _cts.Dispose(); + throw; } Dimension = _providers[0].Dimension; + // The capability is read from the first provider because the factory produces one kind + // of provider, which Dimension above already assumes. + _batchingPolicy = BatchingPolicy.For(effectiveOptions, _providers[0].SupportsBatching); + // Start worker tasks for (int i = 0; i < effectiveOptions.Concurrency; i++) { @@ -80,69 +115,161 @@ public EmbeddingService(Func providerFactory, EmbeddingServi /// Thrown when text is null. /// Thrown when the service has been disposed. public async Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + EmbeddingResult result = await EmbedWithUsageAsync(text, purpose, cancellationToken).ConfigureAwait(false); + return result.Vectors[0]; + } + + /// + /// Produces a vector embedding for the given text, reporting what it cost if the provider + /// meters that. + /// + /// The text to embed. + /// The intended purpose of the embedding. + /// A token to cancel the operation. + /// A result holding the single embedding, and usage when the provider reports it. + /// Thrown when text is null. + /// Thrown when the service has been disposed. + public async Task EmbedWithUsageAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) { ObjectDisposedException.ThrowIf(_disposed == 1, this); ArgumentNullException.ThrowIfNull(text); - TaskCompletionSource tcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - - CancellationTokenRegistration registration = default; - if (cancellationToken.CanBeCanceled) - { - registration = cancellationToken.Register(() => tcs.TrySetCanceled(cancellationToken)); - } - - EmbeddingRequest request = new EmbeddingRequest { Text = text, Purpose = purpose, CompletionSource = tcs }; - - try - { - await _channel.Writer.WriteAsync(request, cancellationToken); - } - catch (ChannelClosedException) - { - throw new ObjectDisposedException(GetType().FullName); - } - - try - { - return await tcs.Task; - } - finally - { - await registration.DisposeAsync(); - } + return await SubmitBatchAsync([text], purpose, cancellationToken).ConfigureAwait(false); } /// - /// Produces vector embeddings for multiple texts. Requests are queued and distributed - /// across available workers. + /// Produces vector embeddings for multiple texts. The texts are split into batches + /// according to the batching policy and each batch is queued as one request. /// /// The texts to embed. /// The intended purpose of the embeddings. /// A token to cancel the operation. /// An array of float arrays, one embedding per input text, in the same order. - /// Thrown when texts is null. + /// Thrown when texts is null, or when any element of it is null. /// Thrown when the service has been disposed. + /// Thrown when the provider returns a number of + /// vectors that does not match the number of texts it was given, which leaves them + /// unmatchable to their inputs. public async Task EmbedBatchAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + EmbeddingResult result = await EmbedBatchWithUsageAsync(texts, purpose, cancellationToken).ConfigureAwait(false); + return result.Vectors.ToArray(); + } + + /// + /// Produces vector embeddings for multiple texts, reporting what they cost if the provider + /// meters that. The texts are split into batches according to the batching policy and each + /// batch is queued as one request. + /// + /// The texts to embed. + /// The intended purpose of the embeddings. + /// A token to cancel the operation. + /// A result holding one embedding per input text in the same order, and the + /// combined usage when the provider reports it. The usage is null unless every batch the + /// call was split into reported one. + /// Thrown when texts is null, or when any element of it is null. + /// Thrown when the service has been disposed. + /// Thrown when the provider returns a number of + /// vectors that does not match the number of texts it was given, which leaves them + /// unmatchable to their inputs. + public async Task EmbedBatchWithUsageAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) { ObjectDisposedException.ThrowIf(_disposed == 1, this); ArgumentNullException.ThrowIfNull(texts); if (texts.Count == 0) - return Array.Empty(); + return new EmbeddingResult { Vectors = Array.Empty(), Usage = null }; - Task[] tasks = new Task[texts.Count]; + // Checked up front rather than left to the provider: a null in the middle of a batch + // would otherwise reach a provider that has no way to report which text was bad. for (int i = 0; i < texts.Count; i++) { - tasks[i] = EmbedAsync(texts[i], purpose, cancellationToken); + if (texts[i] == null) + throw new ArgumentNullException($"{nameof(texts)}[{i}]"); } - return await Task.WhenAll(tasks); + IReadOnlyList> batches = _batchingPolicy.Split(texts); + + Task[] batchTasks = new Task[batches.Count]; + for (int i = 0; i < batches.Count; i++) + { + batchTasks[i] = SubmitBatchAsync(batches[i], purpose, cancellationToken); + } + + EmbeddingResult[] batchResults = await Task.WhenAll(batchTasks).ConfigureAwait(false); + + float[][] vectors = new float[texts.Count][]; + int nextVector = 0; + foreach (EmbeddingResult batchResult in batchResults) + { + foreach (float[] vector in batchResult.Vectors) + { + vectors[nextVector++] = vector; + } + } + + return new EmbeddingResult + { + Vectors = vectors, + Usage = EmbeddingUsage.Combine(batchResults.Select(batchResult => batchResult.Usage)) + }; + } + + /// + /// Queues one batch and awaits the worker that takes it. + /// + private async Task SubmitBatchAsync(IReadOnlyList texts, EmbeddingPurpose purpose, CancellationToken cancellationToken) + { + TaskCompletionSource tcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + CancellationTokenRegistration registration = default; + if (cancellationToken.CanBeCanceled) + { + registration = cancellationToken.Register(() => tcs.TrySetCanceled(cancellationToken)); + } + + EmbeddingRequest request = new EmbeddingRequest + { + Texts = texts, + Purpose = purpose, + CallerToken = cancellationToken, + CompletionSource = tcs + }; + + // One finally over both the write and the wait. Disposing only around the wait leaks + // the registration whenever the write throws, which is the disposed-mid-queue path + // below, and the registration outlives the call attached to the caller's token. + try + { + try + { + await _channel.Writer.WriteAsync(request, cancellationToken).ConfigureAwait(false); + } + catch (ChannelClosedException) + { + throw new ObjectDisposedException(GetType().FullName); + } + + return await tcs.Task.ConfigureAwait(false); + } + finally + { + await registration.DisposeAsync().ConfigureAwait(false); + } } /// /// Stops all workers and disposes all provider instances. /// + /// + /// Cancels the shutdown token and then waits for every worker to return, so a provider + /// call already in flight decides how long this takes. A provider that observes the + /// cancellation token it is handed makes disposal prompt; one that ignores it holds + /// disposal open for a full round trip, which against a remote API can be seconds. There + /// is deliberately no timeout: abandoning the wait would dispose a provider while a call + /// was still running inside it, which is worse than waiting. + /// public async ValueTask DisposeAsync() { if (Interlocked.Exchange(ref _disposed, 1) == 1) @@ -152,25 +279,49 @@ public async ValueTask DisposeAsync() _channel.Writer.TryComplete(); // Cancel workers - await _cts.CancelAsync(); + await _cts.CancelAsync().ConfigureAwait(false); // Wait for workers to drain try { - await Task.WhenAll(_workerTasks); + await Task.WhenAll(_workerTasks).ConfigureAwait(false); } catch (OperationCanceledException) { // Expected during shutdown } - - // Dispose providers - foreach (IEmbeddingProvider provider in _providers) + finally { - provider.Dispose(); + // In a finally rather than after the catch, so that a worker failing in a way that + // is not a cancellation still leaves the providers disposed and the queued callers + // released. The constructor already goes to this trouble when the factory fails + // partway; letting a native ONNX session leak on a different failure would give + // that guarantee away. The worker's exception still propagates from here. + ReleaseAbandonedRequests(); + + foreach (IEmbeddingProvider provider in _providers) + { + provider.Dispose(); + } + + _cts.Dispose(); } + } - _cts.Dispose(); + /// + /// Fails every request still sitting in the channel once the workers have stopped. + /// + /// + /// Nothing will pick these up any more, and each one has a caller awaiting a completion + /// source that nothing else will ever complete. Failing them here is what keeps disposal + /// from leaving a caller waiting forever. + /// + private void ReleaseAbandonedRequests() + { + while (_channel.Reader.TryRead(out EmbeddingRequest? abandoned)) + { + abandoned.CompletionSource.TrySetException(new ObjectDisposedException(GetType().FullName)); + } } private async Task RunWorkerAsync(IEmbeddingProvider provider, CancellationToken cancellationToken) @@ -179,17 +330,47 @@ private async Task RunWorkerAsync(IEmbeddingProvider provider, CancellationToken { await foreach (EmbeddingRequest request in _channel.Reader.ReadAllAsync(cancellationToken)) { + // The caller's token has to reach the provider, not just complete the caller's + // task: an abandoned batch is up to MaxTextsPerBatch texts still being billed + // by a remote API, with this worker held for the whole round trip. Linked with + // the shutdown token so either can stop the call. Skipped entirely when the + // caller passed a token that can never fire, which is the common case. + using CancellationTokenSource? linkedSource = request.CallerToken.CanBeCanceled + ? CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, request.CallerToken) + : null; + + CancellationToken callToken = linkedSource?.Token ?? cancellationToken; + try { - float[] result = await provider.EmbedAsync(request.Text, request.Purpose, cancellationToken); + // One channel item is one provider call, which is what makes a batch cost + // one request against a remote API rather than one per text. + EmbeddingResult result = await provider.EmbedBatchWithUsageAsync(request.Texts, request.Purpose, callToken).ConfigureAwait(false); + + // Vectors are matched to texts by position all the way back to the caller, + // so a provider that returns a different number of them has already lost + // the mapping. Failing the request beats handing back misaligned vectors. + if (result.Vectors.Count != request.Texts.Count) + { + throw new InvalidOperationException( + $"{provider.GetType().Name} returned {result.Vectors.Count} vectors for {request.Texts.Count} texts."); + } + request.CompletionSource.TrySetResult(result); } catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) { request.CompletionSource.TrySetCanceled(cancellationToken); } + catch (OperationCanceledException) when (request.CallerToken.IsCancellationRequested) + { + request.CompletionSource.TrySetCanceled(request.CallerToken); + } catch (Exception ex) { + // Including an OperationCanceledException raised from a token of the + // provider's own, which is a failure of that call rather than a + // cancellation of this one. request.CompletionSource.TrySetException(ex); } } diff --git a/VectorSharp.Embedding/EmbeddingServiceOptions.cs b/VectorSharp.Embedding/EmbeddingServiceOptions.cs index 2325631..5f8b220 100644 --- a/VectorSharp.Embedding/EmbeddingServiceOptions.cs +++ b/VectorSharp.Embedding/EmbeddingServiceOptions.cs @@ -14,9 +14,44 @@ public sealed class EmbeddingServiceOptions /// /// Gets the maximum number of pending embedding requests in the channel. - /// When the channel is full, will - /// wait until space becomes available. Default is 1000. + /// When the channel is full, any of the service's embed methods will wait until space + /// becomes available. Default is 1000. /// + /// + /// Counted in batches, not texts: one queued item is one batch of up to + /// texts. Against a provider that does not batch a batch is + /// one text, so this bounds the queue at the same number of texts; against one that does, + /// the bound on texts is this multiplied by . Lower this + /// when the queued text itself is large enough to matter, since the queue holds the + /// strings until a worker takes them. + /// public int ChannelCapacity { get; init; } = 1000; + + /// + /// Gets the largest number of texts the service will put in one request to the provider. + /// Default is 64. + /// + /// + /// Hosted embedding APIs cap the number of inputs per request, so a batch larger than a + /// provider accepts has to be split before it is sent. Set this to what the API accepts. + /// + /// Ignored by a provider whose is false, + /// which is sent one text per request whatever this says. + /// + public int MaxTextsPerBatch { get; init; } = 64; + + /// + /// Gets the largest total number of characters the service will put in one request to the + /// provider. Default is 100000. + /// + /// + /// The other cap hosted APIs apply, independent of the text count: a batch is closed when + /// either limit would be exceeded. A single text longer than this limit is still sent, on + /// its own — splitting a text is the chunker's job, not this one's. + /// + /// Also only reached by a provider that batches, since one text is already a whole request + /// for a provider that does not. + /// + public int MaxCharactersPerBatch { get; init; } = 100000; } } diff --git a/VectorSharp.Embedding/EmbeddingUsage.cs b/VectorSharp.Embedding/EmbeddingUsage.cs new file mode 100644 index 0000000..a47ec71 --- /dev/null +++ b/VectorSharp.Embedding/EmbeddingUsage.cs @@ -0,0 +1,78 @@ +namespace VectorSharp.Embedding +{ + /// + /// What a provider reported spending on a call. Only ever present when the provider actually + /// reported it: a local model that meters nothing produces no usage at all rather than a + /// usage record full of zeroes. + /// + public sealed class EmbeddingUsage + { + /// + /// Gets the number of tokens the provider reported for the call. + /// + /// + /// Required rather than nullable, so that an cannot be built + /// without a real count. "Nothing was reported" is expressed by the absence of the whole + /// record, which keeps it distinguishable from a call that genuinely cost zero tokens. + /// + public required int TokenCount { get; init; } + + /// + /// Gets the model identifier the provider attributed the call to, or null when the + /// provider does not name one. + /// + public string? Model { get; init; } + + /// + /// Combines the usage of the batches one call was split into. + /// + /// + /// Reports a total only when every batch reported one. A partial total would understate + /// what the call actually cost, and under-reporting spend is worse than reporting nothing: + /// null already means "unknown", which is what a partially reported call is. An empty + /// sequence is the same case — there is nothing to report, and totalling nothing up would + /// produce a reported zero, which means "this was free". + /// + /// Lives beside the type rather than inside because what it + /// encodes is this type's own "null means unknown, never free" rule, and because a rule + /// reachable only through a worker pool is a rule that can only be tested through one. + /// + internal static EmbeddingUsage? Combine(IEnumerable reported) + { + int totalTokens = 0; + string? model = null; + bool haveSeenABatch = false; + bool modelsAgree = true; + + foreach (EmbeddingUsage? usage in reported) + { + if (usage == null) + return null; + + totalTokens += usage.TokenCount; + + // Tracked with its own flag rather than by testing model for null, which cannot + // tell "no batch seen yet" from "this batch named no model" — and would then + // attribute the whole call to a model only one of the batches reported. + if (!haveSeenABatch) + { + model = usage.Model; + haveSeenABatch = true; + } + else if (usage.Model != model) + { + modelsAgree = false; + } + } + + if (!haveSeenABatch) + return null; + + return new EmbeddingUsage + { + TokenCount = totalTokens, + Model = modelsAgree ? model : null + }; + } + } +} diff --git a/VectorSharp.Embedding/IEmbeddingProvider.cs b/VectorSharp.Embedding/IEmbeddingProvider.cs index 99f7e32..37148c7 100644 --- a/VectorSharp.Embedding/IEmbeddingProvider.cs +++ b/VectorSharp.Embedding/IEmbeddingProvider.cs @@ -7,6 +7,12 @@ namespace VectorSharp.Embedding /// Individual instances are NOT required to be thread-safe — the /// creates one instance per worker to avoid shared state. /// + /// + /// There is deliberately no single-text counterpart to : + /// reaches a provider only through the batch method, a single + /// text being a batch of one, so a single-text usage member here would never be called. + /// The service exposes the single-text form to callers instead. + /// public interface IEmbeddingProvider : IDisposable { /// @@ -14,6 +20,24 @@ public interface IEmbeddingProvider : IDisposable /// int Dimension { get; } + /// + /// Gets whether this provider embeds several texts in one call more cheaply than one at a + /// time. False unless a provider says otherwise. + /// + /// + /// This is what uses to decide whether to group a caller's + /// texts into batches at all. A provider that has not overridden the batch methods gains + /// nothing from grouping — the default implementations loop the single-text method, so a + /// group would run serially inside one worker instead of spreading across them. Answering + /// false keeps such a provider, a local ONNX model in particular, on one text per unit of + /// work and therefore parallel across workers. + /// + /// Override to true together with , and ideally + /// as well: claiming batching without implementing it is what + /// turns the parallel path into a serial one. + /// + bool SupportsBatching => false; + /// /// Produces a vector embedding for the given text. /// @@ -24,5 +48,68 @@ public interface IEmbeddingProvider : IDisposable /// A token to cancel the operation. /// A float array of length representing the embedding. Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default); + + // Deliberately a virtual interface member, not an extension method over + // EmbedBatchWithUsageAsync. Folding it in would tidy the shape — there would be no way + // left to override the wrong one of the two — at the cost of the case this member exists + // for: a provider that batches but meters nothing can override this alone, and would + // otherwise have to implement a usage-shaped method to say it has no usage to report. + // That trade is the whole of the decision, and it does not improve with age: an interface + // member is a signature callers bind to, so removing one is a break for anything compiled + // against a version that had it. + + /// + /// Produces vector embeddings for several texts in one call. + /// + /// + /// The default implementation loops , so a provider that only + /// implements the single-text method keeps working unchanged. Override it in providers + /// that can embed an array in one request — against a remote API this is the difference + /// between one HTTP round trip and one per text — and set + /// to true so the service groups texts for it. + /// + /// The texts to embed. + /// The intended purpose of the embeddings. + /// A token to cancel the operation. + /// One embedding per input text, in the same order. + async Task EmbedBatchAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + ArgumentNullException.ThrowIfNull(texts); + + float[][] vectors = new float[texts.Count][]; + for (int i = 0; i < texts.Count; i++) + { + vectors[i] = await EmbedAsync(texts[i], purpose, cancellationToken).ConfigureAwait(false); + } + + return vectors; + } + + /// + /// Produces vector embeddings for several texts in one call, reporting what they cost if + /// this provider meters that. + /// + /// + /// This is the method calls, and the one to override in a + /// provider that meters tokens. The default implementation calls + /// and reports no usage, because a provider that has not + /// been written to report usage has nothing truthful to say about it. + /// + /// The two batch methods are separately overridable, so a provider that batches and meters + /// should override both: overriding only this one leaves a direct caller of + /// looping single calls. That direction is deliberate — the + /// service reaches a provider through this method, so a provider that overrides only + /// (batches, does not meter) still gets one request per + /// batch through the service. + /// + /// The texts to embed. + /// The intended purpose of the embeddings. + /// A token to cancel the operation. + /// The embeddings, and usage when this provider reports it. + async Task EmbedBatchWithUsageAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default) + { + float[][] vectors = await EmbedBatchAsync(texts, purpose, cancellationToken).ConfigureAwait(false); + return new EmbeddingResult { Vectors = vectors, Usage = null }; + } } } diff --git a/VectorSharp.Embedding/README.md b/VectorSharp.Embedding/README.md index 59b5a82..db52633 100644 --- a/VectorSharp.Embedding/README.md +++ b/VectorSharp.Embedding/README.md @@ -4,7 +4,7 @@ [![NuGet](https://img.shields.io/badge/nuget-VectorSharp.Embedding-blue.svg)](https://www.nuget.org/packages/VectorSharp.Embedding) -Channel-based embedding service with configurable parallelism. Provider-agnostic — works with local ONNX models, remote HTTP endpoints, or any custom embedding source. +Channel-based embedding service with configurable parallelism, request batching and token usage reporting. Provider-agnostic — works with local ONNX models, remote HTTP endpoints, or any custom embedding source. ## Install @@ -18,7 +18,9 @@ This package contains the core abstractions and service. For a ready-to-use mode - **Channel-based architecture** — any code can request embeddings, workers process them in the background - **Configurable parallelism** — N concurrent workers, each with its own provider instance -- **Backpressure** — bounded channel prevents unbounded memory growth +- **Real batching** — providers that embed an array in one call get whole batches under configurable size limits, so a remote API sees one request rather than one per text; providers that do not stay one text per call, spread across workers +- **Usage reporting** — a provider that meters tokens reports them back; one that does not reports nothing at all, rather than a zero you could mistake for a free call +- **Backpressure** — a bounded channel caps how many batches wait to be picked up, so a queue cannot outgrow it faster than the workers drain it - **Provider-agnostic** — implement `IEmbeddingProvider` for any embedding source - **Purpose-aware** — `EmbeddingPurpose.Document` vs `EmbeddingPurpose.Query` for models that distinguish between them - **Zero dependencies** — only uses in-box `System.Threading.Channels` @@ -41,23 +43,88 @@ float[] embedding = await embedder.EmbedAsync("some text"); float[] docEmbedding = await embedder.EmbedAsync("document text", EmbeddingPurpose.Document); float[] queryEmbedding = await embedder.EmbedAsync("search query", EmbeddingPurpose.Query); -// Batch embedding — requests are distributed across workers +// Batch embedding — split into batches, each batch one call to the provider float[][] embeddings = await embedder.EmbedBatchAsync(new[] { "text1", "text2", "text3" }); + +// The same call, with what the provider reported spending +EmbeddingResult result = await embedder.EmbedBatchWithUsageAsync(new[] { "text1", "text2" }); +int? tokens = result.Usage?.TokenCount; // null when the provider meters nothing ``` ## How It Works ``` -Caller ──EmbedAsync──▶ Channel Writer ──▶ [bounded channel] ──▶ Channel Reader ──▶ Worker N - ▲ │ - │ ▼ - └──── await TCS.Task ◀──── TCS.SetResult(float[]) ◀── provider.EmbedAsync(text) ───┘ +Caller ──EmbedAsync────────▶ one batch of one ──┐ + EmbedBatchAsync ──▶ split into batches ─┴─▶ [bounded channel] ──▶ Worker N + ▲ │ + │ ▼ + └── await TCS.Task ◀── TCS.SetResult(EmbeddingResult) ◀── provider.EmbedBatchWithUsageAsync(texts) ``` -- Requests are queued in a bounded channel with natural backpressure +- Batches are queued in a bounded channel with natural backpressure - N workers consume from the channel, each owning its own `IEmbeddingProvider` instance -- Errors in one request don't affect other requests or workers -- Disposal completes the channel, drains workers, and disposes all providers +- One queued item is one batch, so a batch is one provider call rather than one per text +- A failing batch fails only that batch; other batches and the workers carry on +- Disposal completes the channel, drains workers, and disposes all providers. Any batch still + queued at that point is failed with `ObjectDisposedException` rather than left awaiting a worker + that has stopped +- **Disposal waits for whatever call is already inside a provider, with no timeout.** Observe the + `CancellationToken` you are handed and disposal is prompt; ignore it and disposal takes as long + as your slowest round trip. There is deliberately no bound — abandoning the wait would dispose a + provider with a call still running inside it + +## Batching + +Batching is per provider, not per caller. A provider answers `SupportsBatching`, which is `false` +unless it says otherwise: + +- **`false`** — the texts are sent one per request, spread across `Concurrency` workers. This is + what a local ONNX model wants: it embeds one text at a time either way, so grouping would only + move the work into a single worker and run it serially. +- **`true`** — the texts are grouped by the limits below and each group is queued as one unit of + work, so a hosted API sees one request instead of one per text. This applies to whichever method + the caller used: `EmbedBatchAsync` and `EmbedBatchWithUsageAsync` take the same path. + +Every hosted embedding API accepts an array of inputs and caps both how many it takes and how much +text, so both limits are applied and whichever binds first closes a batch: + +```csharp +EmbeddingServiceOptions options = new EmbeddingServiceOptions +{ + MaxTextsPerBatch = 64, // default + MaxCharactersPerBatch = 100000 // default +}; +``` + +A single text longer than `MaxCharactersPerBatch` is sent on its own rather than split — deciding +how to shorten a text belongs to the caller, or to `VectorSharp.Chunking` before this point. + +Unrelated single `EmbedAsync` calls are never coalesced into a shared batch. A batch is a batch +because the caller asked for one. + +## Usage Reporting + +`EmbedWithUsageAsync` and `EmbedBatchWithUsageAsync` return an `EmbeddingResult`: the vectors, plus +what the provider reported spending. + +```csharp +EmbeddingResult result = await embedder.EmbedBatchWithUsageAsync(chunks); + +if (result.Usage != null) +{ + meter.Record(result.Usage.TokenCount, result.Usage.Model); +} +else +{ + // This provider reports nothing. Not the same as a call that cost nothing. + meter.RecordUnknownSpend(); +} +``` + +`Usage` is null unless the provider actually reported something, and an unreported count is never +defaulted to `0` — a caller billed per token has to be able to tell "not reported" from "free". +When a call is split into several batches, the totals are summed only if every batch reported one; +a total assembled from some of them would understate the spend, so it comes back null instead. ## EmbeddingPurpose @@ -78,8 +145,10 @@ Providers that don't distinguish between purposes simply ignore the parameter. ```csharp EmbeddingServiceOptions options = new EmbeddingServiceOptions { - Concurrency = 4, // Number of concurrent workers (default: 1) - ChannelCapacity = 500 // Max pending requests before backpressure (default: 1000) + Concurrency = 4, // Number of concurrent workers (default: 1) + ChannelCapacity = 500, // Max pending batches before backpressure (default: 1000) + MaxTextsPerBatch = 64, // Max texts in one provider call (default: 64) + MaxCharactersPerBatch = 100000 // Max characters in one provider call (default: 100000) }; ``` @@ -87,6 +156,11 @@ Note: each worker creates its own provider instance via the factory. For ONNX mo ## Implementing a Custom Provider +Only `Dimension`, `EmbedAsync` and `Dispose` have to be implemented. `SupportsBatching` and the +batch methods are default interface implementations — the batch ones loop the single-text method +and the capability answers false — so an existing provider keeps working, and keeps behaving, +untouched: + ```csharp public class HttpEmbeddingProvider : IEmbeddingProvider { @@ -121,6 +195,42 @@ await using EmbeddingService embedder = new EmbeddingService( ); ``` +A provider talking to an API that embeds an array in one request should say so and override the +batch methods, which is where the round-trip saving comes from. The service only groups texts for a +provider that advertises the capability, so both parts are needed: + +```csharp +public bool SupportsBatching => true; + +public async Task EmbedBatchWithUsageAsync(IReadOnlyList texts, + EmbeddingPurpose purpose = EmbeddingPurpose.Document, + CancellationToken cancellationToken = default) +{ + ApiResponse response = await PostAsync(texts, purpose, cancellationToken); + + return new EmbeddingResult + { + Vectors = response.Embeddings, // one per input text, in the same order + Usage = response.TokenCount is int used // only when the API actually reported it + ? new EmbeddingUsage { TokenCount = used, Model = response.Model } + : null + }; +} + +// Override the plain batch method too, so callers that reach for it get one request as well +public async Task EmbedBatchAsync(IReadOnlyList texts, + EmbeddingPurpose purpose = EmbeddingPurpose.Document, + CancellationToken cancellationToken = default) +{ + EmbeddingResult result = await EmbedBatchWithUsageAsync(texts, purpose, cancellationToken); + return result.Vectors.ToArray(); +} +``` + +The service matches vectors to texts by position, so a batch must come back with exactly one vector +per input text, in order. A response of a different length fails the request rather than returning +vectors attributed to the wrong text. + ## API Reference ### IEmbeddingProvider @@ -131,6 +241,32 @@ public interface IEmbeddingProvider : IDisposable int Dimension { get; } Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default); + + bool SupportsBatching { get; } // default false — decides whether the service groups texts + + // Default implementations — override to embed an array in one request + Task EmbedBatchAsync(IReadOnlyList texts, + EmbeddingPurpose purpose = EmbeddingPurpose.Document, + CancellationToken cancellationToken = default); + Task EmbedBatchWithUsageAsync(IReadOnlyList texts, + EmbeddingPurpose purpose = EmbeddingPurpose.Document, + CancellationToken cancellationToken = default); +} +``` + +### EmbeddingResult and EmbeddingUsage + +```csharp +public sealed class EmbeddingResult +{ + public required IReadOnlyList Vectors { get; init; } + public EmbeddingUsage? Usage { get; init; } // null means unknown, never free +} + +public sealed class EmbeddingUsage +{ + public required int TokenCount { get; init; } + public string? Model { get; init; } } ``` @@ -143,9 +279,15 @@ public sealed class EmbeddingService : IAsyncDisposable public int Dimension { get; } public Task EmbedAsync(string text, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default); + public Task EmbedWithUsageAsync(string text, + EmbeddingPurpose purpose = EmbeddingPurpose.Document, + CancellationToken cancellationToken = default); public Task EmbedBatchAsync(IReadOnlyList texts, EmbeddingPurpose purpose = EmbeddingPurpose.Document, CancellationToken cancellationToken = default); + public Task EmbedBatchWithUsageAsync(IReadOnlyList texts, + EmbeddingPurpose purpose = EmbeddingPurpose.Document, + CancellationToken cancellationToken = default); public ValueTask DisposeAsync(); } ``` diff --git a/VectorSharp.Embedding/VectorSharp.Embedding.csproj b/VectorSharp.Embedding/VectorSharp.Embedding.csproj index 2a14268..bd00d8b 100644 --- a/VectorSharp.Embedding/VectorSharp.Embedding.csproj +++ b/VectorSharp.Embedding/VectorSharp.Embedding.csproj @@ -7,10 +7,10 @@ True VectorSharp.Embedding - 1.0.0 + 1.1.0 VectorSharp.Embedding Adam Tovatt - Channel-based embedding service with configurable parallelism for VectorSharp. Provider-agnostic — supports local ONNX models, remote HTTP endpoints, or any custom embedding source. + Channel-based embedding service with configurable parallelism for VectorSharp. Batches texts into single provider requests for providers that support it, under configurable size limits, and reports provider token usage when the provider meters it. Provider-agnostic — supports local ONNX models, remote HTTP endpoints, or any custom embedding source. MIT vector;embeddings;embedding-service;channels;parallelism README.md diff --git a/VectorSharp.Packaging.Tests/PackagingTests.cs b/VectorSharp.Packaging.Tests/PackagingTests.cs new file mode 100644 index 0000000..60a08ee --- /dev/null +++ b/VectorSharp.Packaging.Tests/PackagingTests.cs @@ -0,0 +1,257 @@ +using System.Diagnostics; +using System.IO.Compression; +using System.Xml.Linq; + +namespace VectorSharp.Packaging.Tests +{ + /// + /// Checks what the repository actually publishes: for every project the solution packs, the + /// .nupkg a build produces is opened and inspected. These read files off disk rather than the + /// loaded assembly, which is why they live in a project of their own rather than beside any + /// one package's API surface tests. + /// + /// + /// The set of packages is taken from the solution rather than listed here, so a package added + /// later is covered without anyone remembering to add it. A package that needs an exception + /// has to be named in , which makes the exception + /// visible instead of making the coverage invisible. + /// + public class PackagingTests + { + /// + /// Packages that legitimately declare dependencies, and are therefore exempt from + /// . + /// + /// + /// An exclusion list rather than an inclusion list on purpose: a package added to the + /// solution is guarded by default and its author has to come here and say why not, rather + /// than a new package being silently unchecked because nobody updated a list. + /// + private static readonly IReadOnlySet PackagesWithDependencies = new HashSet(StringComparer.Ordinal) + { + // Depends on the ONNX runtime and the tokenizer, and claims no otherwise. + "VectorSharp.Embedding.NomicEmbed" + }; + + public static TheoryData PackablePackages + { + get + { + TheoryData data = []; + foreach (PackableProject project in PackableProjects()) + { + data.Add(project.PackageId); + } + return data; + } + } + + public static TheoryData ZeroDependencyPackages + { + get + { + TheoryData data = []; + foreach (PackableProject project in PackableProjects().Where(project => !PackagesWithDependencies.Contains(project.PackageId))) + { + data.Add(project.PackageId); + } + return data; + } + } + + [Theory] + [MemberData(nameof(ZeroDependencyPackages))] + public void PackedPackage_DeclaresNoDependencies(string packageId) + { + // Asserted against the nuspec inside the built .nupkg, which is the artifact that + // reaches nuget.org. The project file alone would miss a dependency injected by a + // Directory.Build.props, and the compiled assembly would miss one that is declared + // but unused, since the compiler omits references nothing touches. + XDocument nuspec = ReadNuspec(packageId); + XNamespace ns = nuspec.Root!.GetDefaultNamespace(); + + string[] dependencies = nuspec.Descendants(ns + "dependency") + .Select(dependency => dependency.Attribute("id")?.Value ?? "") + .ToArray(); + + // Asserted on the joined names rather than on the sequence, because a failure has to + // name the dependency that appeared. Assert.Empty prints no contents. + Assert.Equal(string.Empty, string.Join(", ", dependencies)); + } + + [Theory] + [MemberData(nameof(PackablePackages))] + public void PackedPackage_ShipsItsReadme(string packageId) + { + // Without the readme file packed alongside PackageReadmeFile, the nuget.org listing + // renders blank — a failure invisible from inside the repository. + using ZipArchive package = ZipFile.OpenRead(PackagePath(packageId)); + + Assert.Contains(package.Entries, entry => entry.FullName.Equals("README.md", StringComparison.Ordinal)); + } + + [Fact] + public void PackableProjects_AreAllReferencedByThisProject() + { + // The theories above resolve a package by finding its built .nupkg, so a packable + // project this test project does not reference would fail on a missing file with no + // hint as to why. Naming the omission directly is the difference between a puzzling + // failure and an instruction. + string[] missing = PackableProjects() + .Where(project => !File.Exists(Path.Combine(AppContext.BaseDirectory, $"{project.AssemblyName}.dll"))) + .Select(project => project.PackageId) + .ToArray(); + + Assert.Equal( + string.Empty, + string.Join(", ", missing)); + } + + private static XDocument ReadNuspec(string packageId) + { + using ZipArchive package = ZipFile.OpenRead(PackagePath(packageId)); + + ZipArchiveEntry nuspec = package.Entries.FirstOrDefault(entry => entry.FullName.EndsWith(".nuspec", StringComparison.Ordinal)) + ?? throw new InvalidOperationException($"No nuspec inside the {packageId} package."); + + using Stream content = nuspec.Open(); + return XDocument.Load(content); + } + + /// + /// Finds the package this build produced for the given package id. + /// + /// + /// The version comes from the package's own built assembly, not from the project file's + /// <Version>. The publish workflow builds and packs with + /// -p:Version=<tag version>, an MSBuild global property that overrides + /// <Version> for every project — so the project file names the version of a + /// local build and nothing else, and reading it here would look for a file the release + /// build never wrote. The assembly is stamped from the same property, so it agrees with + /// the .nupkg however the build was invoked. + /// + /// The configuration is pinned to the one these tests were built in, rather than searched + /// for. Searching the whole bin tree lets a Release run pass by inspecting a Debug + /// artifact left over from a developer's last build, and configuration is exactly what CI + /// varies. + /// + private static string PackagePath(string packageId) + { + PackableProject project = PackableProjects().FirstOrDefault(candidate => candidate.PackageId == packageId) + ?? throw new InvalidOperationException($"No packable project in the solution produces {packageId}."); + + string version = BuiltVersion(project); + string expectedFileName = $"{packageId}.{version}.nupkg"; + string packageDirectory = Path.Combine(project.Directory, "bin", Configuration); + string packagePath = Path.Combine(packageDirectory, expectedFileName); + + if (!File.Exists(packagePath)) + { + throw new InvalidOperationException( + $"No {expectedFileName} in {packageDirectory}. The packages are produced by a build, so build the solution before running these."); + } + + return packagePath; + } + + /// + /// The version the current build stamped into the package's assembly, which is the version + /// it also named the .nupkg after. Build metadata after a '+' is dropped, since the package + /// file name carries the version alone. + /// + private static string BuiltVersion(PackableProject project) + { + string assemblyPath = Path.Combine(AppContext.BaseDirectory, $"{project.AssemblyName}.dll"); + + if (!File.Exists(assemblyPath)) + { + throw new InvalidOperationException( + $"No {project.AssemblyName}.dll beside the tests. Add a ProjectReference to {project.PackageId} in VectorSharp.Packaging.Tests.csproj."); + } + + string informationalVersion = FileVersionInfo.GetVersionInfo(assemblyPath).ProductVersion + ?? throw new InvalidOperationException($"{assemblyPath} carries no product version."); + + int metadataStart = informationalVersion.IndexOf('+'); + return metadataStart < 0 ? informationalVersion : informationalVersion[..metadataStart]; + } + + /// + /// The build configuration these tests were built in, read from their own output path: + /// bin/<Configuration>/<TargetFramework>. + /// + private static string Configuration + { + get + { + DirectoryInfo? directory = new DirectoryInfo(AppContext.BaseDirectory.TrimEnd(Path.DirectorySeparatorChar)); + + while (directory?.Parent != null && !directory.Parent.Name.Equals("bin", StringComparison.OrdinalIgnoreCase)) + { + directory = directory.Parent; + } + + return directory?.Name + ?? throw new InvalidOperationException($"No bin directory above {AppContext.BaseDirectory}, so the build configuration cannot be determined."); + } + } + + /// + /// Every project the solution packs, taken from the solution file so that this set cannot + /// drift from what the repository actually publishes. + /// + private static IEnumerable PackableProjects() + { + string root = RepositoryRoot(); + XDocument solution = XDocument.Load(Path.Combine(root, "VectorSharp.slnx")); + + foreach (XElement projectElement in solution.Descendants("Project")) + { + string? relativePath = projectElement.Attribute("Path")?.Value; + if (relativePath == null) + continue; + + string projectFile = Path.Combine(root, relativePath.Replace('\\', Path.DirectorySeparatorChar).Replace('/', Path.DirectorySeparatorChar)); + if (!File.Exists(projectFile)) + throw new InvalidOperationException($"The solution lists {relativePath}, which does not exist."); + + XDocument project = XDocument.Load(projectFile); + + bool packable = project.Descendants("GeneratePackageOnBuild") + .Any(element => bool.TryParse(element.Value, out bool value) && value); + + if (!packable) + continue; + + string packageId = project.Descendants("PackageId").FirstOrDefault()?.Value + ?? throw new InvalidOperationException($"{relativePath} is packed but declares no PackageId."); + + string assemblyName = project.Descendants("AssemblyName").FirstOrDefault()?.Value + ?? Path.GetFileNameWithoutExtension(projectFile); + + yield return new PackableProject(packageId, assemblyName, Path.GetDirectoryName(projectFile)!); + } + } + + /// + /// Walks up from the test binary to the directory holding the solution file. Throws rather + /// than returning something empty if it never finds one. + /// + private static string RepositoryRoot() + { + DirectoryInfo? directory = new DirectoryInfo(AppContext.BaseDirectory); + + while (directory != null && !File.Exists(Path.Combine(directory.FullName, "VectorSharp.slnx"))) + { + directory = directory.Parent; + } + + if (directory == null) + throw new InvalidOperationException($"No VectorSharp.slnx found above {AppContext.BaseDirectory}."); + + return directory.FullName; + } + + private sealed record PackableProject(string PackageId, string AssemblyName, string Directory); + } +} diff --git a/VectorSharp.Packaging.Tests/VectorSharp.Packaging.Tests.csproj b/VectorSharp.Packaging.Tests/VectorSharp.Packaging.Tests.csproj new file mode 100644 index 0000000..9663217 --- /dev/null +++ b/VectorSharp.Packaging.Tests/VectorSharp.Packaging.Tests.csproj @@ -0,0 +1,38 @@ + + + + net10.0 + latest + enable + enable + false + true + + + + + + + + + + + + + + + + + + + + + + diff --git a/VectorSharp.Reranking.Tests/KeywordRerankProvider.cs b/VectorSharp.Reranking.Tests/KeywordRerankProvider.cs new file mode 100644 index 0000000..364510f --- /dev/null +++ b/VectorSharp.Reranking.Tests/KeywordRerankProvider.cs @@ -0,0 +1,68 @@ +namespace VectorSharp.Reranking.Tests +{ + /// + /// A deterministic reranker for testing, scoring each document by how many of the query's + /// words it contains. Stands in for a real provider closely enough to exercise the contract: + /// it scores every candidate, sorts, truncates to topN, and reports the positions it was + /// given rather than the text it was given. + /// + internal sealed class KeywordRerankProvider : IRerankProvider + { + private readonly RerankUsage? _usage; + private bool _disposed; + + public KeywordRerankProvider(RerankUsage? usage = null) + { + _usage = usage; + } + + public Task RerankAsync(string query, IReadOnlyList documents, int topN, CancellationToken cancellationToken = default) + { + ObjectDisposedException.ThrowIf(_disposed, this); + ArgumentNullException.ThrowIfNull(query); + ArgumentNullException.ThrowIfNull(documents); + ArgumentOutOfRangeException.ThrowIfLessThan(topN, 1); + + for (int i = 0; i < documents.Count; i++) + { + if (documents[i] == null) + throw new ArgumentNullException($"{nameof(documents)}[{i}]"); + } + + cancellationToken.ThrowIfCancellationRequested(); + + string[] queryWords = query.Split(' ', StringSplitOptions.RemoveEmptyEntries); + + List scored = new List(documents.Count); + for (int i = 0; i < documents.Count; i++) + { + scored.Add(new RerankMatch { Index = i, Score = ScoreOf(documents[i], queryWords) }); + } + + List best = scored + .OrderByDescending(match => match.Score) + .ThenBy(match => match.Index) + .Take(topN) + .ToList(); + + return Task.FromResult(new RerankResult { Matches = best, Usage = _usage }); + } + + private static float ScoreOf(string document, string[] queryWords) + { + int hits = 0; + foreach (string word in queryWords) + { + if (document.Contains(word, StringComparison.OrdinalIgnoreCase)) + hits++; + } + + return hits; + } + + public void Dispose() + { + _disposed = true; + } + } +} diff --git a/VectorSharp.Reranking.Tests/KeywordRerankProviderTests.cs b/VectorSharp.Reranking.Tests/KeywordRerankProviderTests.cs new file mode 100644 index 0000000..d091bb6 --- /dev/null +++ b/VectorSharp.Reranking.Tests/KeywordRerankProviderTests.cs @@ -0,0 +1,14 @@ +namespace VectorSharp.Reranking.Tests +{ + /// + /// Runs the contract against the in-repo double. A provider + /// added later inherits the same base and gets the same checks. + /// + public class KeywordRerankProviderTests : RerankProviderContractTests + { + protected override IRerankProvider CreateProvider(RerankUsage? usage = null) + { + return new KeywordRerankProvider(usage); + } + } +} diff --git a/VectorSharp.Reranking.Tests/PackageSurfaceTests.cs b/VectorSharp.Reranking.Tests/PackageSurfaceTests.cs new file mode 100644 index 0000000..a3047ce --- /dev/null +++ b/VectorSharp.Reranking.Tests/PackageSurfaceTests.cs @@ -0,0 +1,66 @@ +using System.Reflection; + +namespace VectorSharp.Reranking.Tests +{ + /// + /// Pins what the package exposes. The promise is that it stays small — one interface and its + /// result types — and that is the kind of claim that erodes one convenience type at a time. + /// + public class PackageSurfaceTests + { + private static readonly Assembly PackageAssembly = typeof(IRerankProvider).Assembly; + + [Fact] + public void PublicSurface_IsTheInterfaceAndItsResultTypes() + { + // Enumerated rather than counted, so a new public type has to be named here to pass + // and a reviewer sees exactly what the package exposes. + string[] expected = + [ + nameof(IRerankProvider), + nameof(RerankMatch), + nameof(RerankResult), + nameof(RerankUsage) + ]; + + string[] actual = PackageAssembly + .GetExportedTypes() + .Select(type => type.Name) + .OrderBy(name => name, StringComparer.Ordinal) + .ToArray(); + + Assert.Equal(expected.OrderBy(name => name, StringComparer.Ordinal), actual); + } + + [Fact] + public void PublicSurface_HasNoServiceOrWorkerMachinery() + { + // A second gate behind the enumeration above, for the likely path: someone adds a + // service type, sees that test go red, and makes it green by adding the name to the + // list. This one is what stops them. + // + // It matches on names, so it catches the shapes anyone would reach for first and not + // a name that avoids them. The enumeration is the real fence; this is the reason to + // stop and think before widening it. + string[] serviceLikeNames = PackageAssembly + .GetExportedTypes() + .Select(type => type.Name) + .Where(name => name.EndsWith("Service", StringComparison.Ordinal) + || name.EndsWith("Options", StringComparison.Ordinal) + || name.Contains("Worker", StringComparison.Ordinal) + || name.Contains("Queue", StringComparison.Ordinal)) + .ToArray(); + + Assert.Empty(serviceLikeNames); + } + + [Fact] + public void RerankMatch_IsAReferenceType() + { + // Backs up the behavioural test in the contract suite: as a struct, an absent match + // reads as a real one at index 0. Stated here too because the type decision is what + // the contract test is protecting, and this names it directly. + Assert.True(typeof(RerankMatch).IsClass); + } + } +} diff --git a/VectorSharp.Reranking.Tests/RerankProviderContractTests.cs b/VectorSharp.Reranking.Tests/RerankProviderContractTests.cs new file mode 100644 index 0000000..9a5f904 --- /dev/null +++ b/VectorSharp.Reranking.Tests/RerankProviderContractTests.cs @@ -0,0 +1,270 @@ +namespace VectorSharp.Reranking.Tests +{ + /// + /// The behaviour every is expected to have, written against the + /// interface rather than any one implementation. A new provider inherits this class and + /// supplies itself, and gets the contract checked instead of hand-rolling its own version of + /// these tests. + /// + /// + /// The package ships abstractions only, so these tests necessarily run against a double. What + /// keeps them from being a test of the double is that the double is not where they live: the + /// assertions are the contract, and the first real provider runs the same ones. + /// + public abstract class RerankProviderContractTests + { + /// + /// The candidate list every test below ranks. A fresh array each time rather than one + /// shared instance: 's own remarks say that mutating the list + /// passed to a call invalidates every index it returned, and a static array a subclass or + /// a provider double could write into would carry that between tests, where the failure + /// would land in whichever test happened to run second. + /// + protected static string[] Documents => + [ + "cats sleep most of the day", + "sorting a list in python", + "dogs and cats living together", + "how to sort a python list quickly" + ]; + + /// + /// Ranks the documents above without a tie: document 3 matches every word, document 1 + /// matches all but "quickly", and the other two match none. A query that tied the top two + /// would leave the ordering assertions resting on the provider's tiebreak instead of on + /// the scores. + /// + protected const string UnambiguousQuery = "python list sort quickly"; + + /// + /// Creates the provider under test. It must score against + /// by relevance, and report . + /// + protected abstract IRerankProvider CreateProvider(RerankUsage? usage = null); + + #region Results point back at the caller's own records + + [Fact] + public async Task RerankAsync_MatchesCarryTheIndexOfTheDocumentTheyScored() + { + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync(UnambiguousQuery, Documents, topN: 2); + + Assert.Equal(2, result.Matches.Count); + Assert.Equal(3, result.Matches[0].Index); + Assert.Equal(1, result.Matches[1].Index); + } + + [Fact] + public async Task RerankAsync_IndexesMapBackToTheCallersMetadata() + { + // A worked illustration of why matches are indexes rather than text, kept for what it + // shows a reader: the same two positions the test above pins, resolved through the + // caller's own array. + int[] documentIds = [101, 102, 103, 104]; + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync(UnambiguousQuery, Documents, topN: 2); + + int[] rerankedIds = result.Matches.Select(match => documentIds[match.Index]).ToArray(); + + Assert.Equal([104, 102], rerankedIds); + } + + [Fact] + public async Task RerankAsync_IndexesAreWithinTheSuppliedList() + { + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync("cats", Documents, topN: Documents.Length); + + // Pinned first: without it, a provider returning nothing satisfies both assertions + // below, since Assert.All passes over an empty sequence and 0 distinct equals 0. + Assert.Equal(Documents.Length, result.Matches.Count); + + Assert.All(result.Matches, match => Assert.InRange(match.Index, 0, Documents.Length - 1)); + Assert.Equal(result.Matches.Count, result.Matches.Select(match => match.Index).Distinct().Count()); + } + + #endregion + + #region Ordering and topN + + [Fact] + public async Task RerankAsync_ReturnsMatchesInDescendingScoreOrder() + { + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync(UnambiguousQuery, Documents, topN: 4); + + Assert.Equal(4, result.Matches.Count); + + // The top two are strictly ordered, so this cannot pass on a run of equal scores. + Assert.True( + result.Matches[0].Score > result.Matches[1].Score, + $"Expected a strict ordering at the top, got {result.Matches[0].Score} then {result.Matches[1].Score}."); + + for (int i = 1; i < result.Matches.Count; i++) + { + Assert.True( + result.Matches[i - 1].Score >= result.Matches[i].Score, + $"Match {i - 1} scored {result.Matches[i - 1].Score} and match {i} scored {result.Matches[i].Score}."); + } + } + + [Fact] + public async Task RerankAsync_ReturnsAtMostTopN() + { + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync("cats", Documents, topN: 1); + + Assert.Single(result.Matches); + } + + [Fact] + public async Task RerankAsync_FewerDocumentsThanTopN_ReturnsWhatThereIs() + { + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync("cats", ["cats sleep", "dogs bark"], topN: 10); + + Assert.Equal(2, result.Matches.Count); + } + + [Fact] + public async Task RerankAsync_NoDocuments_ReturnsNoMatchesRatherThanFailing() + { + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync("cats", [], topN: 5); + + Assert.Empty(result.Matches); + } + + [Fact] + public async Task RerankAsync_NoMatches_YieldsNullFromFirstOrDefaultRatherThanAnIndexZeroMatch() + { + // Why RerankMatch is a class: as a struct this would hand back Index 0 with score 0, + // a match pointing at the caller's first document and indistinguishable from a real + // one. Asserted through behaviour so that changing the type back reddens a test. + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync("cats", [], topN: 5); + + Assert.Null(result.Matches.FirstOrDefault()); + } + + #endregion + + #region Usage + + [Fact] + public async Task RerankAsync_ProviderMetersNothing_UsageIsNullNotZero() + { + // A provider that reports nothing must be distinguishable from one reporting a call + // that cost nothing, the same rule the embedding package follows. + using IRerankProvider provider = CreateProvider(); + + RerankResult result = await provider.RerankAsync("cats", Documents, topN: 2); + + Assert.Null(result.Usage); + } + + [Fact] + public async Task RerankAsync_ProviderReportsZeroTokens_UsageSurvivesAsAReportedZero() + { + // The other half of the rule: a reported zero is a real measurement and has to come + // back as one, or "null means unknown, never free" says nothing. + using IRerankProvider provider = CreateProvider(new RerankUsage { TokenCount = 0, Model = "free-tier" }); + + RerankResult result = await provider.RerankAsync("cats", Documents, topN: 2); + + Assert.NotNull(result.Usage); + Assert.Equal(0, result.Usage.TokenCount); + Assert.Equal("free-tier", result.Usage.Model); + } + + [Fact] + public async Task RerankAsync_ProviderReportsUsage_PassesItThrough() + { + using IRerankProvider provider = CreateProvider(new RerankUsage { TokenCount = 42, Model = "test-reranker" }); + + RerankResult result = await provider.RerankAsync("cats", Documents, topN: 2); + + Assert.NotNull(result.Usage); + Assert.Equal(42, result.Usage.TokenCount); + Assert.Equal("test-reranker", result.Usage.Model); + } + + [Fact] + public async Task RerankAsync_ProviderReportsUsageWithoutAModel_KeepsTheCountAndLeavesTheModelNull() + { + using IRerankProvider provider = CreateProvider(new RerankUsage { TokenCount = 7 }); + + RerankResult result = await provider.RerankAsync("cats", Documents, topN: 2); + + Assert.NotNull(result.Usage); + Assert.Equal(7, result.Usage.TokenCount); + Assert.Null(result.Usage.Model); + } + + #endregion + + #region Argument handling + + [Fact] + public async Task RerankAsync_NullQuery_Throws() + { + using IRerankProvider provider = CreateProvider(); + + await Assert.ThrowsAsync(() => provider.RerankAsync(null!, Documents, topN: 2)); + } + + [Fact] + public async Task RerankAsync_NullDocuments_Throws() + { + using IRerankProvider provider = CreateProvider(); + + await Assert.ThrowsAsync(() => provider.RerankAsync("cats", null!, topN: 2)); + } + + [Fact] + public async Task RerankAsync_NullDocumentElement_Throws() + { + using IRerankProvider provider = CreateProvider(); + + await Assert.ThrowsAsync(() => provider.RerankAsync("cats", ["fine", null!], topN: 2)); + } + + [Fact] + public async Task RerankAsync_TopNBelowOne_Throws() + { + using IRerankProvider provider = CreateProvider(); + + await Assert.ThrowsAsync(() => provider.RerankAsync("cats", Documents, topN: 0)); + } + + [Fact] + public async Task RerankAsync_CancelledToken_Throws() + { + using IRerankProvider provider = CreateProvider(); + CancellationToken cancelled = new CancellationToken(true); + + await Assert.ThrowsAnyAsync(() => + provider.RerankAsync("cats", Documents, topN: 2, cancelled)); + } + + [Fact] + public async Task RerankAsync_AfterDispose_Throws() + { + IRerankProvider provider = CreateProvider(); + provider.Dispose(); + + await Assert.ThrowsAsync(() => provider.RerankAsync("cats", Documents, topN: 2)); + } + + #endregion + } +} diff --git a/VectorSharp.Reranking.Tests/VectorSharp.Reranking.Tests.csproj b/VectorSharp.Reranking.Tests/VectorSharp.Reranking.Tests.csproj new file mode 100644 index 0000000..dd616da --- /dev/null +++ b/VectorSharp.Reranking.Tests/VectorSharp.Reranking.Tests.csproj @@ -0,0 +1,26 @@ + + + + net10.0 + latest + enable + enable + false + true + + + + + + + + + + + + + + + + + diff --git a/VectorSharp.Reranking/IRerankProvider.cs b/VectorSharp.Reranking/IRerankProvider.cs new file mode 100644 index 0000000..ff61348 --- /dev/null +++ b/VectorSharp.Reranking/IRerankProvider.cs @@ -0,0 +1,56 @@ +namespace VectorSharp.Reranking +{ + /// + /// Defines the contract for a component that reorders a candidate set against a query. + /// This is the second stage of retrieval: a vector or hybrid search produces candidates + /// cheaply, and a reranker scores them properly. + /// Implementations may use a local model, remote API calls, or any other source. + /// Each instance owns its own resources and must be disposed when no longer needed. + /// Individual instances are NOT required to be thread-safe. + /// + /// + /// There is no service around this interface on purpose. Reranking happens once per user + /// query against a candidate set that is already in hand, not as a bulk pipeline, so the + /// worker pool and queue that earn their place in embedding would only add machinery here. + /// + public interface IRerankProvider : IDisposable + { + /// + /// Scores the given documents against the query and returns the best ones. + /// + /// The query the documents are judged against. + /// The candidates, typically the results of a vector search. + /// Positions in this list are what the returned matches refer to, so the list must not be + /// modified until the call has returned and its indexes have been resolved. Passing a + /// collection the caller goes on mutating invalidates every index silently: the count + /// still matches, nothing throws, and the matches simply point at different documents. + /// The maximum number of matches to return. A provider returns fewer + /// when fewer documents were supplied. Named for the term the rerank APIs themselves use, + /// rather than the count of VectorSharp.Storage searches. + /// A token to cancel the operation. + /// The matches ordered by score descending, and usage when the provider reports it. + /// + /// Implementations should reject a null query, a null document list, and a null element + /// within it, with ; and a below + /// 1 with . Rejecting is the point: an empty + /// result is a real answer here, so it must not double as an error signal. + /// An empty document list is not an error either — the answer is no matches. + /// + /// A provider whose backing API caps how many candidates it accepts should refuse a longer + /// list rather than truncate it. Truncation leaves every returned index valid and every + /// assertion satisfied while candidates disappear, which is the failure a caller cannot see. + /// + /// Argument rejection may surface either synchronously or as a faulted task, since a + /// provider that validates inside an async method necessarily does the latter. + /// Callers should await the call to observe it rather than expecting a throw at the call + /// site. + /// + /// Implementations should also observe before doing + /// any work, so that a token cancelled before the call throws + /// rather than spending a metered request whose + /// result the caller has already stopped waiting for. A provider that only reaches its + /// token at an HTTP call has already paid for the request by then. + /// + Task RerankAsync(string query, IReadOnlyList documents, int topN, CancellationToken cancellationToken = default); + } +} diff --git a/VectorSharp.Reranking/README.md b/VectorSharp.Reranking/README.md new file mode 100644 index 0000000..acad662 --- /dev/null +++ b/VectorSharp.Reranking/README.md @@ -0,0 +1,209 @@ +# VectorSharp.Reranking + +[← Back to VectorSharp](../README.md) + +[![NuGet](https://img.shields.io/badge/nuget-VectorSharp.Reranking-blue.svg)](https://www.nuget.org/packages/VectorSharp.Reranking) + +Reranking abstractions for the second stage of retrieval. A vector or hybrid search produces +candidates cheaply; a reranker scores them properly. Zero dependencies. + +## Install + +``` +dotnet add package VectorSharp.Reranking +``` + +This package contains the abstraction and its result types. Implement `IRerankProvider` against +whichever reranker you use — a hosted API or a local model. + +## Features + +- **Index-based results** — matches say *which* candidate they scored, so you keep your own metadata +- **Usage reporting** — a provider that meters tokens reports them back; one that does not reports nothing at all, rather than a zero you could mistake for a free call +- **No machinery** — one interface and three result types, with no service, queue or worker pool +- **Zero dependencies** — not even on the sibling VectorSharp packages + +## Usage + +The candidate side of this example comes from `VectorSharp.Storage`; this package only takes the +list of strings and gives back positions into it. + +```csharp +using VectorSharp.Reranking; +using VectorSharp.Storage; // only for the search that produces the candidates + +using IRerankProvider reranker = new MyRerankProvider(); + +// Candidates from a vector search, and whatever you know about each of them +IReadOnlyList> candidates = await store.FindMostSimilarAsync(queryVector, count: 50); +string[] documents = candidates.Select(candidate => textById[candidate.Id]).ToArray(); + +RerankResult result = await reranker.RerankAsync("how do I sort a list", documents, topN: 5); + +foreach (RerankMatch match in result.Matches) +{ + // match.Index is a position in `documents`, so it maps straight back to your own records + int id = candidates[match.Index].Id; + Console.WriteLine($"{id} scored {match.Score}"); +} +``` + +## Why Indexes, Not Text + +A reranker is given text and returns a judgement about it. What the caller actually needs back is +the *record* behind that text — a row id, a chunk offset, a file path — and only the caller has +that. Returning the document text would force a lookup by content, which is slower and ambiguous +the moment two candidates read the same. An index into the list you passed in is unambiguous and +free to resolve. + +Matches come back ordered by score descending, and there are at most `topN` of them — fewer when +you supplied fewer documents. + +Because the results are positions, the list you pass in has to hold still: don't modify it until +the call has returned and you have resolved the indexes. A collection mutated mid-call invalidates +every index without anything failing — the count still matches and the matches simply point at +different documents. + +Scores are comparable only within a single call. Providers differ in range and scale, so a score +means nothing next to one from another query or another model. + +## Usage Reporting + +`RerankResult.Usage` carries what the provider reported spending, and is null when it reported +nothing: + +```csharp +RerankResult result = await reranker.RerankAsync(query, documents, topN: 5); + +if (result.Usage != null) +{ + meter.Record(result.Usage.TokenCount, result.Usage.Model); +} +else +{ + // This provider reports nothing. Not the same as a call that cost nothing. + meter.RecordUnknownSpend(); +} +``` + +An unreported count is never defaulted to `0`, because a caller billed per token has to be able to +tell "not reported" from "free". + +## Implementing a Custom Provider + +An instance owns its own resources and is not required to be thread-safe — this package has no +service serializing access to it, so a provider registered as a singleton has to be safe for +concurrent calls on its own, or be registered per scope. + +```csharp +public sealed class HttpRerankProvider : IRerankProvider +{ + private readonly HttpClient _client; + private readonly string _endpoint; + + public HttpRerankProvider(string endpoint) + { + _endpoint = endpoint; + _client = new HttpClient(); + } + + public async Task RerankAsync(string query, IReadOnlyList documents, + int topN, CancellationToken cancellationToken = default) + { + ArgumentNullException.ThrowIfNull(query); + ArgumentNullException.ThrowIfNull(documents); + ArgumentOutOfRangeException.ThrowIfLessThan(topN, 1); + + if (documents.Count == 0) + return new RerankResult { Matches = [] }; + + ApiResponse response = await PostAsync(query, documents, topN, cancellationToken); + + return new RerankResult + { + Matches = response.Results + .Select(result => new RerankMatch { Index = result.Index, Score = result.Score }) + .ToArray(), + Usage = response.TokenCount is int used // only when the API actually reported it + ? new RerankUsage { TokenCount = used, Model = response.Model } + : null + }; + } + + public void Dispose() => _client.Dispose(); +} +``` + +The provider owns the `HttpClient` it created here, which is why it disposes it. A provider handed +a client from `IHttpClientFactory` or a typed-client registration must not — that handler belongs +to the container. + +An empty candidate list is not an error — the answer is no matches. Reserve exceptions for a null +query, a null list, a null element inside it, or a `topN` below 1, so that an empty result never +doubles as a failure signal. + +If the API behind a provider caps how many candidates it accepts, refuse a longer list rather than +truncating it. Truncation leaves every returned index valid while candidates vanish, which is the +one failure the caller has no way to notice. + +Observe the `CancellationToken` before doing any work, so that a token already cancelled when the +call arrives throws instead of spending a metered request nobody is waiting for. + +**As a caller:** `await` the call to see an argument rejection. A provider that validates inside an +`async` method necessarily returns a faulted task rather than throwing at the call site, so a +synchronous `try`/`catch` wrapped around the call itself catches nothing. + +## Where It Fits + +``` +query ──▶ embed ──▶ vector search ──▶ candidates ──▶ rerank ──▶ top N +``` + +- **embed** — [VectorSharp.Embedding](../VectorSharp.Embedding/README.md) +- **vector search** — [VectorSharp.Storage](../VectorSharp.Storage/README.md) +- **rerank** — this package + +Reranking is one call per user query against a candidate set already in hand, so this package has +no service, queue or worker pool. The concurrency machinery that earns its place in +[VectorSharp.Embedding](../VectorSharp.Embedding/README.md), where a document can produce hundreds +of chunks, has nothing to do here. + +## API Reference + +### IRerankProvider + +```csharp +public interface IRerankProvider : IDisposable +{ + Task RerankAsync(string query, IReadOnlyList documents, int topN, + CancellationToken cancellationToken = default); +} +``` + +Instances own their own resources and are not required to be thread-safe. + +### RerankResult, RerankMatch and RerankUsage + +```csharp +public sealed class RerankResult +{ + public required IReadOnlyList Matches { get; init; } // best first, at most topN + public RerankUsage? Usage { get; init; } // null means unknown, never free +} + +public sealed class RerankMatch +{ + public required int Index { get; init; } // position in the documents you passed in + public required float Score { get; init; } // higher is more relevant +} + +public sealed class RerankUsage +{ + public required int TokenCount { get; init; } + public string? Model { get; init; } +} +``` + +## License + +MIT diff --git a/VectorSharp.Reranking/RerankMatch.cs b/VectorSharp.Reranking/RerankMatch.cs new file mode 100644 index 0000000..100dcda --- /dev/null +++ b/VectorSharp.Reranking/RerankMatch.cs @@ -0,0 +1,36 @@ +namespace VectorSharp.Reranking +{ + // A class rather than a struct, for the same reason TextChunk in VectorSharp.Chunking is one: + // `required` is enforced only on object-creation expressions, so default(RerankMatch) or + // FirstOrDefault() over an empty list would hand back Index 0 with score 0 — a match pointing + // at the caller's first document, indistinguishable from a real one. As a class those + // expressions produce a null the caller can see. + + /// + /// One reranked candidate: which document it was, and how well the reranker judged it to + /// answer the query. + /// + public sealed class RerankMatch + { + /// + /// Gets the position of this document in the list passed to + /// . + /// + /// + /// An index rather than the document text, because the caller is the one holding whatever + /// sits behind each candidate — a row id, a chunk offset, a file path — and needs to get + /// back to it. Returning the text would force a lookup by content, which is both slower + /// and ambiguous when two candidates read the same. + /// + public required int Index { get; init; } + + /// + /// Gets the relevance score the provider gave this document. Higher is more relevant. + /// + /// + /// Comparable only within one call: providers differ in range and scale, and a score + /// carries no meaning against a score from a different query or model. + /// + public required float Score { get; init; } + } +} diff --git a/VectorSharp.Reranking/RerankResult.cs b/VectorSharp.Reranking/RerankResult.cs new file mode 100644 index 0000000..11de5b7 --- /dev/null +++ b/VectorSharp.Reranking/RerankResult.cs @@ -0,0 +1,25 @@ +namespace VectorSharp.Reranking +{ + /// + /// The outcome of a rerank call: the candidates that survived, best first, and what the + /// provider reported spending if it reported anything. + /// + public sealed class RerankResult + { + /// + /// Gets the reranked candidates, ordered by descending. + /// Holds at most the topN the caller asked for, and may hold fewer when fewer documents + /// were supplied. + /// + public required IReadOnlyList Matches { get; init; } + + /// + /// Gets what the call cost, or null when the provider reported nothing. + /// + /// + /// Null means unknown, never free. A caller metering spend has to treat it as a gap in its + /// accounting rather than as a zero, which is why an unreported count is not defaulted. + /// + public RerankUsage? Usage { get; init; } + } +} diff --git a/VectorSharp.Reranking/RerankUsage.cs b/VectorSharp.Reranking/RerankUsage.cs new file mode 100644 index 0000000..1e89a83 --- /dev/null +++ b/VectorSharp.Reranking/RerankUsage.cs @@ -0,0 +1,34 @@ +namespace VectorSharp.Reranking +{ + /// + /// What a provider reported spending on a rerank call. Only ever present when the provider + /// actually reported it: a provider that meters nothing produces no usage at all rather than a + /// usage record full of zeroes. + /// + /// + /// Reads the same as VectorSharp.Embedding.EmbeddingUsage today, and is deliberately a + /// separate type rather than a shared one: this package takes no dependencies, including on + /// its sibling packages, so a caller can rerank without pulling in an embedding service it is + /// not using. The resemblance is a convenience for anyone reading both, not a contract — + /// what a rerank API meters is its own business, and a field added to one of the two does not + /// oblige the other. + /// + public sealed class RerankUsage + { + /// + /// Gets the number of tokens the provider reported for the call. + /// + /// + /// Required rather than nullable, so that a cannot be built + /// without a real count. "Nothing was reported" is expressed by the absence of the whole + /// record, which keeps it distinguishable from a call that genuinely cost zero tokens. + /// + public required int TokenCount { get; init; } + + /// + /// Gets the model identifier the provider attributed the call to, or null when the + /// provider does not name one. + /// + public string? Model { get; init; } + } +} diff --git a/VectorSharp.Reranking/VectorSharp.Reranking.csproj b/VectorSharp.Reranking/VectorSharp.Reranking.csproj new file mode 100644 index 0000000..bf06e8d --- /dev/null +++ b/VectorSharp.Reranking/VectorSharp.Reranking.csproj @@ -0,0 +1,29 @@ + + + + net10.0 + enable + enable + True + + VectorSharp.Reranking + 1.0.0 + VectorSharp.Reranking + Adam Tovatt + Reranking abstractions for VectorSharp — the second stage after a vector or hybrid search. Results carry the index of the candidate they scored, so callers keep their own metadata, and providers report token usage when they meter it. Zero dependencies. + MIT + vector;reranking;rerank;search;retrieval;rag + README.md + https://github.com/AdamTovatt/vector-sharp + git + + True + True + snupkg + + + + + + + diff --git a/VectorSharp.slnx b/VectorSharp.slnx index 14fddd7..67612af 100644 --- a/VectorSharp.slnx +++ b/VectorSharp.slnx @@ -1,10 +1,13 @@ + + +