Added local RAG (#756)

Co-authored-by: Thorsten Sommer <SommerEngineering@users.noreply.github.com>
This commit is contained in:
Paul KoudelkaandThorsten Sommer authored and GitHub committed 2026-09-09 18:43:37 +02:00
1 parent d043fbc8f0
commit c7b42bee96
179 files changed
+276977 -1421

No files matched your search

+20
View File
@@ -0,0 +1,20 @@
namespace AIStudio.Tools;
/// <summary>
/// Icons we draw ourselves, because the Material icon set MudBlazor ships does not contain them.
/// </summary>
/// <remarks>
/// The strings follow the same convention as the MudBlazor icons: they contain the SVG child
/// elements only, drawn on a 24 by 24 canvas. MudIcon and every component taking an icon wrap them
/// into the svg element themselves, which is why there must be no svg root element here.
/// </remarks>
public static class AppIcons
{
/// <summary>
/// The classic database symbol: a cylinder made of three stacked discs.
/// </summary>
public const string DATABASE =
"""
<path d="M5 4.6A7 2.6 0 0 1 19 4.6L19 8.9A7 2.6 0 0 1 5 8.9Z"/><path d="M5 9.9A7 2.6 0 0 0 19 9.9L19 14.1A7 2.6 0 0 1 5 14.1Z"/><path d="M5 15.1A7 2.6 0 0 0 19 15.1L19 19.4A7 2.6 0 0 1 5 19.4Z"/>
""";
}
@@ -0,0 +1,28 @@
namespace AIStudio.Tools;
/// <summary>
/// Content which a reader held back, together with the token count of exactly that content.
/// </summary>
/// <remarks>
/// Readers which assemble a page or a slide from several stream events cannot pass their content
/// on right away. Its token count has to travel with it: the count describes the content, not the
/// event which happened to arrive at the moment the content was released. Keeping the two together
/// is what stops a page from being sized by the text of the page after it.
/// </remarks>
/// <param name="Content">The assembled content.</param>
/// <param name="TokenCount">The number of tokens of that content, or null when it is unknown.</param>
public readonly record struct ContentStreamPendingContent(string Content, int? TokenCount)
{
/// <summary>
/// Adds up two token counts, where an unknown count makes the sum unknown as well.
/// </summary>
/// <remarks>
/// A partial sum would understate the whole and would let the chunking size a chunk by a part
/// of what it holds. Reporting the count as unknown is the honest answer, because the caller
/// can still count the content itself.
/// </remarks>
/// <param name="left">The first count, or null when it is unknown.</param>
/// <param name="right">The second count, or null when it is unknown.</param>
/// <returns>The sum, or null when either count is unknown.</returns>
public static int? AddTokenCounts(int? left, int? right) => left is null || right is null ? null : left + right;
}
@@ -11,14 +11,25 @@ namespace AIStudio.Tools;
/// <param name="Content">The content to append, or null when this event carries none.</param>
/// <param name="Error">The reported failure, or null when the event was processed successfully.</param>
/// <param name="PromptInjection">What the runtime filtered out of the content, or null when it filtered nothing.</param>
public readonly record struct ContentStreamProcessedEvent(string? Content, ContentStreamErrorDetails? Error, ContentStreamPromptInjectionDetails? PromptInjection = null)
/// <param name="TokenCount">The number of tokens of the content, or null when it is unknown.</param>
public readonly record struct ContentStreamProcessedEvent(string? Content, ContentStreamErrorDetails? Error, ContentStreamPromptInjectionDetails? PromptInjection = null, int? TokenCount = null)
{
/// <summary>
/// An event which neither produced content nor reported a failure.
/// </summary>
public static readonly ContentStreamProcessedEvent NOTHING = new(null, null);
public static ContentStreamProcessedEvent FromContent(string? content) => new(content, null);
/// <summary>
/// An event which produced content, with the token count of that very content.
/// </summary>
/// <remarks>
/// The count travels with the content because a reader may hold content back across several
/// events: pairing it with the count of the event which released it would size it by the
/// wrong text.
/// </remarks>
/// <param name="content">The content to append.</param>
/// <param name="tokenCount">The number of tokens of that content, or null when it is unknown.</param>
public static ContentStreamProcessedEvent FromContent(string? content, int? tokenCount = null) => new(content, null, TokenCount: tokenCount);
public static ContentStreamProcessedEvent FromError(ContentStreamErrorDetails? error) => new(null, error);
@@ -12,4 +12,7 @@ public sealed class ContentStreamSseEvent
[JsonPropertyName("metadata")]
public ContentStreamSseMetadata? Metadata { get; init; }
[JsonPropertyName("token_count")]
public int? TokenCount { get; init; }
}
@@ -17,7 +17,7 @@ public static class ContentStreamSseHandler
switch (sseEvent.Metadata)
{
case ContentStreamTextMetadata:
return ContentStreamProcessedEvent.FromContent(sseEvent.Content);
return ContentStreamProcessedEvent.FromContent(sseEvent.Content, sseEvent.TokenCount);
case ContentStreamPdfMetadata pdfMetadata:
var pageNumber = pdfMetadata.Pdf?.PageNumber ?? 0;
@@ -25,7 +25,7 @@ public static class ContentStreamSseHandler
# Page {pageNumber}
{sseEvent.Content}
""");
""", sseEvent.TokenCount);
case ContentStreamSpreadsheetMetadata spreadsheetMetadata:
var sheetName = spreadsheetMetadata.Spreadsheet?.SheetName;
@@ -38,31 +38,43 @@ public static class ContentStreamSseHandler
}
spreadSheetResult.Append(sseEvent.Content);
return ContentStreamProcessedEvent.FromContent(spreadSheetResult.ToString());
return ContentStreamProcessedEvent.FromContent(spreadSheetResult.ToString(), sseEvent.TokenCount);
//
// Documents which the runtime reads page by page are buffered, so the images of
// a page can follow its Markdown. Documents converted as a whole, e.g. by Pandoc,
// carry no page number and are passed on unchanged.
//
// The buffering is why the count comes back from the reader rather than from
// this event: the page which is released here arrived one event ago, and this
// event's count belongs to the page which is now being buffered.
//
case ContentStreamDocumentMetadata documentMetadata:
if (documentMetadata.Document?.PageNumber is not > 0)
return ContentStreamProcessedEvent.FromContent(sseEvent.Content);
return ContentStreamProcessedEvent.FromContent(sseEvent.Content, sseEvent.TokenCount);
var documentManager = DOCUMENT_MANAGERS.GetOrAdd(sseEvent.StreamId!, _ => new());
var documentContent = documentManager.AddPage(documentMetadata, sseEvent.Content, extractImages);
return documentContent is null ? ContentStreamProcessedEvent.NOTHING : ContentStreamProcessedEvent.FromContent(documentContent);
var documentContent = documentManager.AddPage(documentMetadata, sseEvent.Content, sseEvent.TokenCount, extractImages);
return documentContent is null ? ContentStreamProcessedEvent.NOTHING : ContentStreamProcessedEvent.FromContent(documentContent.Value.Content, documentContent.Value.TokenCount);
case ContentStreamImageMetadata:
return ContentStreamProcessedEvent.FromContent(sseEvent.Content);
return ContentStreamProcessedEvent.FromContent(sseEvent.Content, sseEvent.TokenCount);
case ContentStreamPresentationMetadata presentationMetadata:
if (!extractImages)
{
var slideNumber = presentationMetadata.Presentation?.SlideNumber ?? 0;
return ContentStreamProcessedEvent.FromContent(slideNumber > 0
? $"# Slide {slideNumber}\n{sseEvent.Content}"
: sseEvent.Content, sseEvent.TokenCount);
}
var slideManager = SLIDE_MANAGERS.GetOrAdd(
sseEvent.StreamId!,
_ => new()
);
slideManager.AddSlide(presentationMetadata, sseEvent.Content, extractImages);
slideManager.AddSlide(presentationMetadata, sseEvent.Content, sseEvent.TokenCount, extractImages);
return ContentStreamProcessedEvent.NOTHING;
//
@@ -82,11 +94,11 @@ public static class ContentStreamSseHandler
return ContentStreamProcessedEvent.FromPromptInjection(promptInjectionMetadata.PromptInjection);
default:
return ContentStreamProcessedEvent.FromContent(sseEvent.Content);
return ContentStreamProcessedEvent.FromContent(sseEvent.Content, sseEvent.TokenCount);
}
case { Content: not null, Metadata: null }:
return ContentStreamProcessedEvent.FromContent(sseEvent.Content);
return ContentStreamProcessedEvent.FromContent(sseEvent.Content, sseEvent.TokenCount);
default:
return ContentStreamProcessedEvent.NOTHING;
@@ -166,32 +178,45 @@ public static class ContentStreamSseHandler
return $"![Image](data:{imageMediaType};base64,{base64Image})";
}
public static string? Clear(string streamId)
/// <summary>
/// Releases what the readers of a stream still hold back and forgets the stream.
/// </summary>
/// <remarks>
/// The readers which assemble pages or slides always keep the last one of them: nothing tells
/// them that no further image is coming. It is released here, and it carries its own token
/// count, because a chunk without one cannot be sized by the caller.
/// </remarks>
/// <param name="streamId">The stream to release and forget.</param>
/// <returns>The content which was held back, or null when there was none.</returns>
public static ContentStreamPendingContent? Clear(string streamId)
{
if (string.IsNullOrWhiteSpace(streamId))
return null;
var finalContentChunk = new StringBuilder();
if(SLIDE_MANAGERS.TryGetValue(streamId, out var slideManager))
int? tokenCount = 0;
if(SLIDE_MANAGERS.TryGetValue(streamId, out var slideManager)
&& slideManager.GetAllSlidesInOrder() is { } slides
&& !string.IsNullOrWhiteSpace(slides.Content))
{
var result = slideManager.GetAllSlidesInOrder();
if (!string.IsNullOrWhiteSpace(result))
finalContentChunk.Append(result);
finalContentChunk.Append(slides.Content);
tokenCount = ContentStreamPendingContent.AddTokenCounts(tokenCount, slides.TokenCount);
}
if (DOCUMENT_MANAGERS.TryGetValue(streamId, out var documentManager))
if (DOCUMENT_MANAGERS.TryGetValue(streamId, out var documentManager)
&& documentManager.Flush() is { } page
&& !string.IsNullOrWhiteSpace(page.Content))
{
var result = documentManager.Flush();
if (!string.IsNullOrWhiteSpace(result))
finalContentChunk.Append(result);
finalContentChunk.Append(page.Content);
tokenCount = ContentStreamPendingContent.AddTokenCounts(tokenCount, page.TokenCount);
}
SLIDE_MANAGERS.TryRemove(streamId, out _);
DOCUMENT_MANAGERS.TryRemove(streamId, out _);
var imageIdPrefix = $"{streamId}-";
foreach (var key in CHUNKED_IMAGES.Keys.Where(k => k.StartsWith(imageIdPrefix, StringComparison.InvariantCultureIgnoreCase)))
CHUNKED_IMAGES.TryRemove(key, out _);
return finalContentChunk.Length > 0 ? finalContentChunk.ToString() : null;
return finalContentChunk.Length > 0 ? new ContentStreamPendingContent(finalContentChunk.ToString(), tokenCount) : null;
}
}
}
@@ -39,14 +39,15 @@ public abstract class DatabaseClient(string name, string path)
{
string[] suffixes = { "B", "KB", "MB", "GB", "TB", "PB" };
int suffixIndex = 0;
double convertedSize = size;
while (size >= 1024 && suffixIndex < suffixes.Length - 1)
while (convertedSize >= 1024 && suffixIndex < suffixes.Length - 1)
{
size /= 1024;
convertedSize /= 1024;
suffixIndex++;
}
return $"{size:0##} {suffixes[suffixIndex]}";
return $"{convertedSize:0.##} {suffixes[suffixIndex]}";
}
public void SetLogger(ILogger<DatabaseClient> logService)
@@ -1,5 +1,6 @@
using AIStudio.Tools.Services;
using AIStudio.Tools.Databases.IndexStore;
using AIStudio.Tools.Databases.VectorStore;
using AIStudio.Tools.Services;
namespace AIStudio.Tools.Databases;
@@ -44,10 +45,10 @@ public sealed class DatabaseClientProvider(RustService rustService, ILoggerFacto
}
}
public async Task<IVectorStoreClient> GetVectorStoreAsync(CancellationToken cancellationToken = default)
public async Task<VectorStoreClient> GetVectorStoreAsync(CancellationToken cancellationToken = default)
{
var client = await this.GetClientAsync(DatabaseRole.VECTOR_STORE, cancellationToken);
if (client is IVectorStoreClient vectorStore)
if (client is VectorStoreClient vectorStore)
return vectorStore;
return new NoVectorStoreClient(
@@ -56,6 +57,18 @@ public sealed class DatabaseClientProvider(RustService rustService, ILoggerFacto
client.Status);
}
public async Task<IndexStoreClient> GetIndexStoreAsync(CancellationToken cancellationToken = default)
{
var client = await this.GetClientAsync(DatabaseRole.INDEX_STORE, cancellationToken);
if (client is IndexStoreClient indexStore)
return indexStore;
return new NoIndexStoreClient(
client.Name,
"The configured database client does not support local RAG index operations.",
client.Status);
}
private DatabaseClient CacheIfAvailable(DatabaseRole databaseRole, DatabaseClient client)
{
if (!client.IsAvailable)
@@ -92,6 +105,7 @@ public sealed class DatabaseClientProvider(RustService rustService, ILoggerFacto
private async Task<DatabaseClient> CreateClientAsync(DatabaseRole databaseRole, CancellationToken cancellationToken) => databaseRole switch
{
DatabaseRole.VECTOR_STORE => await QdrantEdgeClientImplementation.CreateAsync(rustService, this.logger, this.databaseClientLogger, cancellationToken),
DatabaseRole.INDEX_STORE => await SqliteIndexStoreClientImplementation.CreateAsync(this.logger, this.databaseClientLogger, cancellationToken),
_ => new NoDatabaseClient(databaseRole.ToString(), "The requested database role is not supported.")
};
@@ -3,4 +3,5 @@ namespace AIStudio.Tools.Databases;
public enum DatabaseRole
{
VECTOR_STORE,
INDEX_STORE,
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Databases.IndexStore;
public sealed record EmbeddingStateChunk(string ChunkId, string ParentFileId, int? PageNumber, int ChunkIndex, string ChunkText, DateTimeOffset EmbeddedAtUtc);
@@ -0,0 +1,20 @@
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class EmbeddingStateChunkEntity
{
public int Id { get; set; }
public string ChunkId { get; set; } = string.Empty;
public string ParentFileId { get; set; } = string.Empty;
public int? PageNumber { get; set; }
public int ChunkIndex { get; set; }
public string ChunkText { get; set; } = string.Empty;
public DateTimeOffset EmbeddedAtUtc { get; set; }
public EmbeddingStateFileEntity? File { get; set; }
}
@@ -0,0 +1,24 @@
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class EmbeddingStateDataSourceEntity
{
public string DataSourceId { get; set; } = string.Empty;
public string DataSourceName { get; set; } = string.Empty;
public string DataSourceType { get; set; } = string.Empty;
public string EmbeddingProviderId { get; set; } = string.Empty;
public string EmbeddingSignature { get; set; } = string.Empty;
public string SourceHash { get; set; } = string.Empty;
public int VectorSize { get; set; }
public DateTimeOffset UpdatedAtUtc { get; set; }
public List<EmbeddingStateFileEntity> Files { get; set; } = [];
public List<IndexingFailureEntity> PermanentIndexingFailures { get; set; } = [];
}
@@ -0,0 +1,16 @@
namespace AIStudio.Tools.Databases.IndexStore;
public sealed record EmbeddingStateFile(
string ParentFileId,
string AbsolutePath,
string FileName,
string RelativePath,
string FileType,
string Fingerprint,
long FileSize,
DateTimeOffset CreationUtc,
DateTimeOffset LastWriteUtc,
DateTimeOffset EmbeddedAtUtc,
int ChunkCount,
string ConfidenceLevel,
int ConfidenceLevelRank);
@@ -0,0 +1,36 @@
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class EmbeddingStateFileEntity
{
public string ParentFileId { get; set; } = string.Empty;
public string DataSourceId { get; set; } = string.Empty;
public string AbsolutePath { get; set; } = string.Empty;
public string FileName { get; set; } = string.Empty;
public string RelativePath { get; set; } = string.Empty;
public string FileType { get; set; } = string.Empty;
public string Fingerprint { get; set; } = string.Empty;
public long FileSize { get; set; }
public DateTimeOffset CreationUtc { get; set; }
public DateTimeOffset LastWriteUtc { get; set; }
public DateTimeOffset EmbeddedAtUtc { get; set; }
public int ChunkCount { get; set; }
public string ConfidenceLevel { get; set; } = string.Empty;
public int ConfidenceLevelRank { get; set; }
public EmbeddingStateDataSourceEntity? DataSource { get; set; }
public List<EmbeddingStateChunkEntity> Chunks { get; set; } = [];
}
@@ -0,0 +1,36 @@
using AIStudio.Tools.Services;
namespace AIStudio.Tools.Databases.IndexStore;
public abstract class IndexStoreClient(string name, string path) : DatabaseClient(name, path)
{
public abstract Task<DataSourceEmbeddingManifest> GetManifestAsync(string dataSourceId, CancellationToken token);
public abstract Task UpsertDataSourceAsync(
string dataSourceId,
string dataSourceName,
string dataSourceType,
string embeddingProviderId,
string embeddingSignature,
string sourceHash,
int vectorSize,
CancellationToken token);
public abstract Task UpdateVectorSizeAsync(string dataSourceId, int vectorSize, CancellationToken token);
public abstract Task UpdateDataSourceHashAsync(string dataSourceId, string sourceHash, CancellationToken token);
public abstract Task UpsertFileAsync(string dataSourceId, EmbeddingStateFile file, CancellationToken token);
public abstract Task DeleteFileAsync(string dataSourceId, string filePath, CancellationToken token);
public abstract Task UpsertPermanentFailureAsync(string dataSourceId, PermanentIndexingFailure failure, CancellationToken token);
public abstract Task DeletePermanentFailureAsync(string dataSourceId, string filePath, CancellationToken token);
public abstract Task UpsertChunksAsync(string dataSourceId, IReadOnlyList<EmbeddingStateChunk> chunks, CancellationToken token);
public abstract Task<IReadOnlyList<IndexStoreSearchResult>> SearchChunksAsync(string dataSourceId, string query, int maxMatches, CancellationToken token);
public abstract Task DeleteDataSourceAsync(string dataSourceId, CancellationToken token);
}
@@ -0,0 +1,18 @@
using System.Globalization;
namespace AIStudio.Tools.Databases.IndexStore;
internal static class IndexStoreDateTimeOffset
{
public static string ToUtcText(DateTimeOffset dateTime)
{
return dateTime.ToUniversalTime().ToString("O", CultureInfo.InvariantCulture);
}
public static DateTimeOffset ParseUtc(string value)
{
return DateTimeOffset.TryParse(value, CultureInfo.InvariantCulture, DateTimeStyles.None, out var dateTime)
? dateTime.ToUniversalTime()
: DateTimeOffset.UnixEpoch;
}
}
@@ -0,0 +1,7 @@
using Microsoft.EntityFrameworkCore.Storage.ValueConversion;
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class IndexStoreDateTimeOffsetConverter() : ValueConverter<DateTimeOffset, string>(
value => IndexStoreDateTimeOffset.ToUtcText(value),
value => IndexStoreDateTimeOffset.ParseUtc(value));
@@ -0,0 +1,141 @@
using Microsoft.Data.Sqlite;
using Microsoft.EntityFrameworkCore;
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class IndexStoreDbContext(DbContextOptions<IndexStoreDbContext> options) : DbContext(options)
{
public static DbContextOptions<IndexStoreDbContext> CreateOptions(string databasePath) => new DbContextOptionsBuilder<IndexStoreDbContext>()
.UseSqlite(BuildConnectionString(databasePath))
.Options;
public DbSet<EmbeddingStateDataSourceEntity> DataSources => this.Set<EmbeddingStateDataSourceEntity>();
public DbSet<EmbeddingStateFileEntity> EmbeddedFiles => this.Set<EmbeddingStateFileEntity>();
public DbSet<EmbeddingStateChunkEntity> EmbeddingChunks => this.Set<EmbeddingStateChunkEntity>();
public DbSet<IndexingFailureEntity> PermanentIndexingFailures => this.Set<IndexingFailureEntity>();
public DbSet<IndexStoreSearchResultEntity> SearchResults => this.Set<IndexStoreSearchResultEntity>();
protected override void OnModelCreating(ModelBuilder modelBuilder)
{
var utcDateTimeOffsetConverter = new IndexStoreDateTimeOffsetConverter();
modelBuilder.Entity<EmbeddingStateDataSourceEntity>(entity =>
{
entity.ToTable("data_sources");
entity.HasKey(dataSource => dataSource.DataSourceId);
entity.Property(dataSource => dataSource.DataSourceId).HasColumnName("data_source_id");
entity.Property(dataSource => dataSource.DataSourceName).HasColumnName("data_source_name").IsRequired();
entity.Property(dataSource => dataSource.DataSourceType).HasColumnName("data_source_type").IsRequired();
entity.Property(dataSource => dataSource.EmbeddingProviderId).HasColumnName("embedding_provider_id").IsRequired();
entity.Property(dataSource => dataSource.EmbeddingSignature).HasColumnName("embedding_signature").IsRequired();
entity.Property(dataSource => dataSource.SourceHash).HasColumnName("source_hash").IsRequired().HasDefaultValue(string.Empty);
entity.Property(dataSource => dataSource.VectorSize).HasColumnName("vector_size").HasDefaultValue(0);
entity.Property(dataSource => dataSource.UpdatedAtUtc).HasColumnName("updated_at_utc").HasConversion(utcDateTimeOffsetConverter).IsRequired();
entity
.HasMany(dataSource => dataSource.Files)
.WithOne(file => file.DataSource)
.HasForeignKey(file => file.DataSourceId)
.OnDelete(DeleteBehavior.Cascade);
entity
.HasMany(dataSource => dataSource.PermanentIndexingFailures)
.WithOne(failure => failure.DataSource)
.HasForeignKey(failure => failure.DataSourceId)
.OnDelete(DeleteBehavior.Cascade);
});
modelBuilder.Entity<EmbeddingStateFileEntity>(entity =>
{
entity.ToTable("embedded_files");
entity.HasKey(file => file.ParentFileId);
entity.Property(file => file.ParentFileId).HasColumnName("parent_file_id");
entity.Property(file => file.DataSourceId).HasColumnName("data_source_id").IsRequired();
entity.Property(file => file.AbsolutePath).HasColumnName("absolute_path").UseCollation("NOCASE").IsRequired();
entity.Property(file => file.FileName).HasColumnName("file_name").IsRequired();
entity.Property(file => file.RelativePath).HasColumnName("relative_path").IsRequired();
entity.Property(file => file.FileType).HasColumnName("file_type").IsRequired();
entity.Property(file => file.Fingerprint).HasColumnName("fingerprint").IsRequired();
entity.Property(file => file.FileSize).HasColumnName("file_size");
entity.Property(file => file.CreationUtc).HasColumnName("creation_utc").HasConversion(utcDateTimeOffsetConverter).IsRequired();
entity.Property(file => file.LastWriteUtc).HasColumnName("last_write_utc").HasConversion(utcDateTimeOffsetConverter).IsRequired();
entity.Property(file => file.EmbeddedAtUtc).HasColumnName("embedded_at_utc").HasConversion(utcDateTimeOffsetConverter).IsRequired();
entity.Property(file => file.ChunkCount).HasColumnName("chunk_count");
entity.Property(file => file.ConfidenceLevel).HasColumnName("confidence_level").IsRequired();
entity.Property(file => file.ConfidenceLevelRank).HasColumnName("confidence_level_rank");
entity.HasIndex(file => file.DataSourceId).HasDatabaseName("idx_embedded_files_data_source");
entity.HasIndex(file => file.AbsolutePath).HasDatabaseName("idx_embedded_files_absolute_path");
entity.HasIndex(file => file.FileType).HasDatabaseName("idx_embedded_files_file_type");
entity.HasIndex(file => file.ConfidenceLevelRank).HasDatabaseName("idx_embedded_files_confidence");
entity.HasIndex(file => new { file.DataSourceId, file.AbsolutePath }).HasDatabaseName("idx_embedded_files_data_source_absolute_path").IsUnique();
entity
.HasMany(file => file.Chunks)
.WithOne(chunk => chunk.File)
.HasForeignKey(chunk => chunk.ParentFileId)
.OnDelete(DeleteBehavior.Cascade);
});
modelBuilder.Entity<EmbeddingStateChunkEntity>(entity =>
{
entity.ToTable("embedding_chunks");
entity.HasKey(chunk => chunk.Id);
entity.Property(chunk => chunk.Id).HasColumnName("id").ValueGeneratedOnAdd();
entity.Property(chunk => chunk.ChunkId).HasColumnName("chunk_id").IsRequired();
entity.Property(chunk => chunk.ParentFileId).HasColumnName("parent_file_id").IsRequired();
entity.Property(chunk => chunk.PageNumber).HasColumnName("page_number");
entity.Property(chunk => chunk.ChunkIndex).HasColumnName("chunk_index");
entity.Property(chunk => chunk.ChunkText).HasColumnName("chunk_text").IsRequired();
entity.Property(chunk => chunk.EmbeddedAtUtc).HasColumnName("embedded_at_utc").HasConversion(utcDateTimeOffsetConverter).IsRequired();
entity.HasIndex(chunk => chunk.ChunkId).HasDatabaseName("idx_embedding_chunks_chunk_id").IsUnique();
entity.HasIndex(chunk => chunk.ParentFileId).HasDatabaseName("idx_embedding_chunks_parent_file");
entity.HasIndex(chunk => chunk.PageNumber).HasDatabaseName("idx_embedding_chunks_page");
entity.HasIndex(chunk => new { chunk.ParentFileId, chunk.ChunkIndex }).HasDatabaseName("idx_embedding_chunks_parent_file_chunk_index").IsUnique();
});
modelBuilder.Entity<IndexingFailureEntity>(entity =>
{
entity.ToTable("permanent_indexing_failures");
entity.HasKey(failure => failure.ParentFileId);
entity.Property(failure => failure.ParentFileId).HasColumnName("parent_file_id");
entity.Property(failure => failure.DataSourceId).HasColumnName("data_source_id").IsRequired();
entity.Property(failure => failure.AbsolutePath).HasColumnName("absolute_path").UseCollation("NOCASE").IsRequired();
entity.Property(failure => failure.Fingerprint).HasColumnName("fingerprint").IsRequired();
entity.Property(failure => failure.FailureCode).HasColumnName("failure_code").IsRequired();
entity.Property(failure => failure.FailureMessage).HasColumnName("failure_message").IsRequired();
entity.Property(failure => failure.OccurredAtUtc).HasColumnName("occurred_at_utc").HasConversion(utcDateTimeOffsetConverter).IsRequired();
entity.HasIndex(failure => failure.DataSourceId).HasDatabaseName("idx_permanent_indexing_failures_data_source");
entity.HasIndex(failure => new { failure.DataSourceId, failure.AbsolutePath }).HasDatabaseName("idx_permanent_indexing_failures_data_source_absolute_path").IsUnique();
});
modelBuilder.Entity<IndexStoreSearchResultEntity>(entity =>
{
entity.HasNoKey();
entity.ToView("embedding_chunk_search_results");
entity.Property(result => result.CreationUtc).HasConversion(utcDateTimeOffsetConverter);
entity.Property(result => result.LastWriteUtc).HasConversion(utcDateTimeOffsetConverter);
entity.Property(result => result.EmbeddedAtUtc).HasConversion(utcDateTimeOffsetConverter);
});
}
private static string BuildConnectionString(string databasePath) => new SqliteConnectionStringBuilder
{
DataSource = databasePath,
Mode = SqliteOpenMode.ReadWriteCreate,
Cache = SqliteCacheMode.Shared,
ForeignKeys = true,
DefaultTimeout = 30,
}.ToString();
}
@@ -0,0 +1,15 @@
using Microsoft.EntityFrameworkCore.Design;
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class IndexStoreDesignTimeDbContextFactory : IDesignTimeDbContextFactory<IndexStoreDbContext>
{
public IndexStoreDbContext CreateDbContext(string[] args)
{
var databasePath = args.FirstOrDefault(argument => argument.EndsWith(".sqlite3", StringComparison.OrdinalIgnoreCase));
if (string.IsNullOrWhiteSpace(databasePath))
databasePath = Path.Combine(Path.GetTempPath(), "mindwork-ai-studio-rag-index-design.sqlite3");
return new IndexStoreDbContext(IndexStoreDbContext.CreateOptions(databasePath));
}
}
@@ -0,0 +1,16 @@
using System.Diagnostics.CodeAnalysis;
using Microsoft.EntityFrameworkCore;
namespace AIStudio.Tools.Databases.IndexStore;
internal static class IndexStoreSchemaMigrator
{
[DynamicDependency(DynamicallyAccessedMemberTypes.All, typeof(Migrations.InitialRagIndex))]
[DynamicDependency(DynamicallyAccessedMemberTypes.All, typeof(Migrations.PermanentIndexingFailures))]
public static async Task MigrateAsync(IndexStoreDbContext context, CancellationToken token)
{
await context.Database.MigrateAsync(token);
await context.Database.ExecuteSqlRawAsync("PRAGMA journal_mode=WAL;", token);
}
}
@@ -0,0 +1,24 @@
namespace AIStudio.Tools.Databases.IndexStore;
public sealed record IndexStoreSearchResult(
string ChunkId,
string ParentFileId,
string DataSourceId,
string DataSourceName,
string DataSourceType,
string AbsolutePath,
string FileName,
string RelativePath,
string FileType,
int? PageNumber,
int ChunkIndex,
string ChunkText,
double Score,
string Fingerprint,
long FileSize,
DateTimeOffset CreationUtc,
DateTimeOffset LastWriteUtc,
DateTimeOffset EmbeddedAtUtc,
int ChunkCount,
string ConfidenceLevel,
int ConfidenceLevelRank);
@@ -0,0 +1,46 @@
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class IndexStoreSearchResultEntity
{
public string ChunkId { get; set; } = string.Empty;
public string ParentFileId { get; set; } = string.Empty;
public string DataSourceId { get; set; } = string.Empty;
public string DataSourceName { get; set; } = string.Empty;
public string DataSourceType { get; set; } = string.Empty;
public string AbsolutePath { get; set; } = string.Empty;
public string FileName { get; set; } = string.Empty;
public string RelativePath { get; set; } = string.Empty;
public string FileType { get; set; } = string.Empty;
public int? PageNumber { get; set; }
public int ChunkIndex { get; set; }
public string ChunkText { get; set; } = string.Empty;
public double Score { get; set; }
public string Fingerprint { get; set; } = string.Empty;
public long FileSize { get; set; }
public DateTimeOffset CreationUtc { get; set; }
public DateTimeOffset LastWriteUtc { get; set; }
public DateTimeOffset EmbeddedAtUtc { get; set; }
public int ChunkCount { get; set; }
public string ConfidenceLevel { get; set; } = string.Empty;
public int ConfidenceLevelRank { get; set; }
}
@@ -0,0 +1,27 @@
namespace AIStudio.Tools.Databases.IndexStore;
internal sealed class IndexingFailureEntity
{
public string ParentFileId { get; set; } = string.Empty;
public string DataSourceId { get; set; } = string.Empty;
public string AbsolutePath { get; set; } = string.Empty;
public string Fingerprint { get; set; } = string.Empty;
/// <summary>
/// The failure code, stored by name.
/// </summary>
/// <remarks>
/// The enum has no explicit numbers, so storing the name keeps the rows readable across
/// versions which add or reorder codes.
/// </remarks>
public string FailureCode { get; set; } = string.Empty;
public string FailureMessage { get; set; } = string.Empty;
public DateTimeOffset OccurredAtUtc { get; set; }
public EmbeddingStateDataSourceEntity? DataSource { get; set; }
}
@@ -0,0 +1,202 @@
#nullable disable
using Microsoft.EntityFrameworkCore.Infrastructure;
using Microsoft.EntityFrameworkCore.Migrations;
namespace AIStudio.Tools.Databases.IndexStore.Migrations;
[DbContext(typeof(IndexStoreDbContext))]
[Migration("20260804000000_InitialRagIndex")]
public partial class InitialRagIndex : Migration
{
protected override void Up(MigrationBuilder migrationBuilder)
{
migrationBuilder.CreateTable(
name: "data_sources",
columns: table => new
{
data_source_id = table.Column<string>(type: "TEXT", nullable: false),
data_source_name = table.Column<string>(type: "TEXT", nullable: false),
data_source_type = table.Column<string>(type: "TEXT", nullable: false),
embedding_provider_id = table.Column<string>(type: "TEXT", nullable: false),
embedding_signature = table.Column<string>(type: "TEXT", nullable: false),
source_hash = table.Column<string>(type: "TEXT", nullable: false, defaultValue: string.Empty),
vector_size = table.Column<int>(type: "INTEGER", nullable: false, defaultValue: 0),
updated_at_utc = table.Column<string>(type: "TEXT", nullable: false),
},
constraints: table =>
{
table.PrimaryKey("PK_data_sources", source => source.data_source_id);
});
migrationBuilder.CreateTable(
name: "embedded_files",
columns: table => new
{
parent_file_id = table.Column<string>(type: "TEXT", nullable: false),
data_source_id = table.Column<string>(type: "TEXT", nullable: false),
absolute_path = table.Column<string>(type: "TEXT", nullable: false, collation: "NOCASE"),
file_name = table.Column<string>(type: "TEXT", nullable: false),
relative_path = table.Column<string>(type: "TEXT", nullable: false),
file_type = table.Column<string>(type: "TEXT", nullable: false),
fingerprint = table.Column<string>(type: "TEXT", nullable: false),
file_size = table.Column<long>(type: "INTEGER", nullable: false),
creation_utc = table.Column<string>(type: "TEXT", nullable: false),
last_write_utc = table.Column<string>(type: "TEXT", nullable: false),
embedded_at_utc = table.Column<string>(type: "TEXT", nullable: false),
chunk_count = table.Column<int>(type: "INTEGER", nullable: false),
confidence_level = table.Column<string>(type: "TEXT", nullable: false),
confidence_level_rank = table.Column<int>(type: "INTEGER", nullable: false),
},
constraints: table =>
{
table.PrimaryKey("PK_embedded_files", file => file.parent_file_id);
table.ForeignKey(
name: "FK_embedded_files_data_sources_data_source_id",
column: file => file.data_source_id,
principalTable: "data_sources",
principalColumn: "data_source_id",
onDelete: ReferentialAction.Cascade);
});
migrationBuilder.CreateTable(
name: "embedding_chunks",
columns: table => new
{
id = table.Column<int>(type: "INTEGER", nullable: false)
.Annotation("Sqlite:Autoincrement", true),
chunk_id = table.Column<string>(type: "TEXT", nullable: false),
parent_file_id = table.Column<string>(type: "TEXT", nullable: false),
page_number = table.Column<int>(type: "INTEGER", nullable: true),
chunk_index = table.Column<int>(type: "INTEGER", nullable: false),
chunk_text = table.Column<string>(type: "TEXT", nullable: false),
embedded_at_utc = table.Column<string>(type: "TEXT", nullable: false),
},
constraints: table =>
{
table.PrimaryKey("PK_embedding_chunks", chunk => chunk.id);
table.ForeignKey(
name: "FK_embedding_chunks_embedded_files_parent_file_id",
column: chunk => chunk.parent_file_id,
principalTable: "embedded_files",
principalColumn: "parent_file_id",
onDelete: ReferentialAction.Cascade);
});
migrationBuilder.CreateIndex(
name: "idx_embedded_files_absolute_path",
table: "embedded_files",
column: "absolute_path");
migrationBuilder.CreateIndex(
name: "idx_embedded_files_confidence",
table: "embedded_files",
column: "confidence_level_rank");
migrationBuilder.CreateIndex(
name: "idx_embedded_files_data_source",
table: "embedded_files",
column: "data_source_id");
migrationBuilder.CreateIndex(
name: "idx_embedded_files_data_source_absolute_path",
table: "embedded_files",
columns: ["data_source_id", "absolute_path"],
unique: true);
migrationBuilder.CreateIndex(
name: "idx_embedded_files_file_type",
table: "embedded_files",
column: "file_type");
migrationBuilder.CreateIndex(
name: "idx_embedding_chunks_chunk_id",
table: "embedding_chunks",
column: "chunk_id",
unique: true);
migrationBuilder.CreateIndex(
name: "idx_embedding_chunks_page",
table: "embedding_chunks",
column: "page_number");
migrationBuilder.CreateIndex(
name: "idx_embedding_chunks_parent_file",
table: "embedding_chunks",
column: "parent_file_id");
migrationBuilder.CreateIndex(
name: "idx_embedding_chunks_parent_file_chunk_index",
table: "embedding_chunks",
columns: ["parent_file_id", "chunk_index"],
unique: true);
migrationBuilder.Sql("""
CREATE VIRTUAL TABLE IF NOT EXISTS embedding_chunks_fts
USING fts5(chunk_id UNINDEXED, file_name, chunk_text);
CREATE TRIGGER IF NOT EXISTS embedding_chunks_ai
AFTER INSERT ON embedding_chunks
BEGIN
INSERT INTO embedding_chunks_fts(rowid, chunk_id, file_name, chunk_text)
VALUES (
new.id,
new.chunk_id,
(SELECT file_name FROM embedded_files WHERE parent_file_id = new.parent_file_id),
new.chunk_text);
END;
CREATE TRIGGER IF NOT EXISTS embedding_chunks_ad
AFTER DELETE ON embedding_chunks
BEGIN
DELETE FROM embedding_chunks_fts
WHERE rowid = old.id;
END;
CREATE TRIGGER IF NOT EXISTS embedding_chunks_au
AFTER UPDATE ON embedding_chunks
BEGIN
DELETE FROM embedding_chunks_fts
WHERE rowid = old.id;
INSERT INTO embedding_chunks_fts(rowid, chunk_id, file_name, chunk_text)
VALUES (
new.id,
new.chunk_id,
(SELECT file_name FROM embedded_files WHERE parent_file_id = new.parent_file_id),
new.chunk_text);
END;
CREATE TRIGGER IF NOT EXISTS embedded_files_file_name_au
AFTER UPDATE OF file_name ON embedded_files
BEGIN
DELETE FROM embedding_chunks_fts
WHERE rowid IN (
SELECT id
FROM embedding_chunks
WHERE parent_file_id = new.parent_file_id
);
INSERT INTO embedding_chunks_fts(rowid, chunk_id, file_name, chunk_text)
SELECT id, chunk_id, new.file_name, chunk_text
FROM embedding_chunks
WHERE parent_file_id = new.parent_file_id;
END;
""");
}
protected override void Down(MigrationBuilder migrationBuilder)
{
migrationBuilder.Sql("""
DROP TRIGGER IF EXISTS embedded_files_file_name_au;
DROP TRIGGER IF EXISTS embedding_chunks_au;
DROP TRIGGER IF EXISTS embedding_chunks_ad;
DROP TRIGGER IF EXISTS embedding_chunks_ai;
DROP TABLE IF EXISTS embedding_chunks_fts;
""");
migrationBuilder.DropTable(name: "embedding_chunks");
migrationBuilder.DropTable(name: "embedded_files");
migrationBuilder.DropTable(name: "data_sources");
}
}
@@ -0,0 +1,53 @@
#nullable disable
using Microsoft.EntityFrameworkCore.Infrastructure;
using Microsoft.EntityFrameworkCore.Migrations;
namespace AIStudio.Tools.Databases.IndexStore.Migrations;
[DbContext(typeof(IndexStoreDbContext))]
[Migration("20260909000000_PermanentIndexingFailures")]
public partial class PermanentIndexingFailures : Migration
{
protected override void Up(MigrationBuilder migrationBuilder)
{
migrationBuilder.CreateTable(
name: "permanent_indexing_failures",
columns: table => new
{
parent_file_id = table.Column<string>(type: "TEXT", nullable: false),
data_source_id = table.Column<string>(type: "TEXT", nullable: false),
absolute_path = table.Column<string>(type: "TEXT", nullable: false, collation: "NOCASE"),
fingerprint = table.Column<string>(type: "TEXT", nullable: false),
failure_code = table.Column<string>(type: "TEXT", nullable: false),
failure_message = table.Column<string>(type: "TEXT", nullable: false),
occurred_at_utc = table.Column<string>(type: "TEXT", nullable: false),
},
constraints: table =>
{
table.PrimaryKey("PK_permanent_indexing_failures", failure => failure.parent_file_id);
table.ForeignKey(
name: "FK_permanent_indexing_failures_data_sources_data_source_id",
column: failure => failure.data_source_id,
principalTable: "data_sources",
principalColumn: "data_source_id",
onDelete: ReferentialAction.Cascade);
});
migrationBuilder.CreateIndex(
name: "idx_permanent_indexing_failures_data_source",
table: "permanent_indexing_failures",
column: "data_source_id");
migrationBuilder.CreateIndex(
name: "idx_permanent_indexing_failures_data_source_absolute_path",
table: "permanent_indexing_failures",
columns: ["data_source_id", "absolute_path"],
unique: true);
}
protected override void Down(MigrationBuilder migrationBuilder)
{
migrationBuilder.DropTable(name: "permanent_indexing_failures");
}
}
@@ -0,0 +1,392 @@
#nullable disable
using Microsoft.EntityFrameworkCore;
using Microsoft.EntityFrameworkCore.Infrastructure;
namespace AIStudio.Tools.Databases.IndexStore.Migrations;
// EF Core writes this file and declares the type as partial. Dropping the keyword would only
// last until the next migration is added, so the inspection is silenced instead.
// ReSharper disable once PartialTypeWithSinglePart
[DbContext(typeof(IndexStoreDbContext))]
partial class IndexStoreDbContextModelSnapshot : ModelSnapshot
{
protected override void BuildModel(ModelBuilder modelBuilder)
{
#pragma warning disable 612, 618
modelBuilder.HasAnnotation("ProductVersion", "9.0.18");
var utcDateTimeOffsetConverter = new IndexStoreDateTimeOffsetConverter();
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.EmbeddingStateDataSourceEntity", entity =>
{
entity.Property<string>("DataSourceId")
.HasColumnType("TEXT")
.HasColumnName("data_source_id");
entity.Property<string>("DataSourceName")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("data_source_name");
entity.Property<string>("DataSourceType")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("data_source_type");
entity.Property<string>("EmbeddingProviderId")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("embedding_provider_id");
entity.Property<string>("EmbeddingSignature")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("embedding_signature");
entity.Property<string>("SourceHash")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("source_hash")
.HasDefaultValue(string.Empty);
entity.Property<DateTimeOffset>("UpdatedAtUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT")
.HasColumnName("updated_at_utc");
entity.Property<int>("VectorSize")
.HasColumnType("INTEGER")
.HasColumnName("vector_size")
.HasDefaultValue(0);
entity.HasKey("DataSourceId");
entity.ToTable("data_sources");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.EmbeddingStateFileEntity", entity =>
{
entity.Property<string>("ParentFileId")
.HasColumnType("TEXT")
.HasColumnName("parent_file_id");
entity.Property<string>("AbsolutePath")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("absolute_path")
.UseCollation("NOCASE");
entity.Property<int>("ChunkCount")
.HasColumnType("INTEGER")
.HasColumnName("chunk_count");
entity.Property<string>("ConfidenceLevel")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("confidence_level");
entity.Property<int>("ConfidenceLevelRank")
.HasColumnType("INTEGER")
.HasColumnName("confidence_level_rank");
entity.Property<DateTimeOffset>("CreationUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT")
.HasColumnName("creation_utc");
entity.Property<string>("DataSourceId")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("data_source_id");
entity.Property<DateTimeOffset>("EmbeddedAtUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT")
.HasColumnName("embedded_at_utc");
entity.Property<string>("FileName")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("file_name");
entity.Property<long>("FileSize")
.HasColumnType("INTEGER")
.HasColumnName("file_size");
entity.Property<string>("FileType")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("file_type");
entity.Property<string>("Fingerprint")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("fingerprint");
entity.Property<DateTimeOffset>("LastWriteUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT")
.HasColumnName("last_write_utc");
entity.Property<string>("RelativePath")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("relative_path");
entity.HasKey("ParentFileId");
entity.HasIndex("AbsolutePath")
.HasDatabaseName("idx_embedded_files_absolute_path");
entity.HasIndex("ConfidenceLevelRank")
.HasDatabaseName("idx_embedded_files_confidence");
entity.HasIndex("DataSourceId")
.HasDatabaseName("idx_embedded_files_data_source");
entity.HasIndex("DataSourceId", "AbsolutePath")
.IsUnique()
.HasDatabaseName("idx_embedded_files_data_source_absolute_path");
entity.HasIndex("FileType")
.HasDatabaseName("idx_embedded_files_file_type");
entity.ToTable("embedded_files");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.EmbeddingStateChunkEntity", entity =>
{
entity.Property<int>("Id")
.ValueGeneratedOnAdd()
.HasColumnType("INTEGER")
.HasColumnName("id")
.HasAnnotation("Sqlite:Autoincrement", true);
entity.Property<string>("ChunkId")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("chunk_id");
entity.Property<int>("ChunkIndex")
.HasColumnType("INTEGER")
.HasColumnName("chunk_index");
entity.Property<string>("ChunkText")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("chunk_text");
entity.Property<DateTimeOffset>("EmbeddedAtUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT")
.HasColumnName("embedded_at_utc");
entity.Property<int?>("PageNumber")
.HasColumnType("INTEGER")
.HasColumnName("page_number");
entity.Property<string>("ParentFileId")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("parent_file_id");
entity.HasKey("Id");
entity.HasIndex("ChunkId")
.IsUnique()
.HasDatabaseName("idx_embedding_chunks_chunk_id");
entity.HasIndex("PageNumber")
.HasDatabaseName("idx_embedding_chunks_page");
entity.HasIndex("ParentFileId")
.HasDatabaseName("idx_embedding_chunks_parent_file");
entity.HasIndex("ParentFileId", "ChunkIndex")
.IsUnique()
.HasDatabaseName("idx_embedding_chunks_parent_file_chunk_index");
entity.ToTable("embedding_chunks");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.IndexingFailureEntity", entity =>
{
entity.Property<string>("ParentFileId")
.HasColumnType("TEXT")
.HasColumnName("parent_file_id");
entity.Property<string>("AbsolutePath")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("absolute_path")
.UseCollation("NOCASE");
entity.Property<string>("DataSourceId")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("data_source_id");
entity.Property<string>("FailureCode")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("failure_code");
entity.Property<string>("FailureMessage")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("failure_message");
entity.Property<string>("Fingerprint")
.IsRequired()
.HasColumnType("TEXT")
.HasColumnName("fingerprint");
entity.Property<DateTimeOffset>("OccurredAtUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT")
.HasColumnName("occurred_at_utc");
entity.HasKey("ParentFileId");
entity.HasIndex("DataSourceId")
.HasDatabaseName("idx_permanent_indexing_failures_data_source");
entity.HasIndex("DataSourceId", "AbsolutePath")
.IsUnique()
.HasDatabaseName("idx_permanent_indexing_failures_data_source_absolute_path");
entity.ToTable("permanent_indexing_failures");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.IndexStoreSearchResultEntity", entity =>
{
entity.Property<string>("AbsolutePath")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<int>("ChunkCount")
.HasColumnType("INTEGER");
entity.Property<string>("ChunkId")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<int>("ChunkIndex")
.HasColumnType("INTEGER");
entity.Property<string>("ChunkText")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<string>("ConfidenceLevel")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<int>("ConfidenceLevelRank")
.HasColumnType("INTEGER");
entity.Property<DateTimeOffset>("CreationUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT");
entity.Property<string>("DataSourceId")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<string>("DataSourceName")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<string>("DataSourceType")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<DateTimeOffset>("EmbeddedAtUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT");
entity.Property<string>("FileName")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<long>("FileSize")
.HasColumnType("INTEGER");
entity.Property<string>("FileType")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<string>("Fingerprint")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<DateTimeOffset>("LastWriteUtc")
.HasConversion(utcDateTimeOffsetConverter)
.HasColumnType("TEXT");
entity.Property<int?>("PageNumber")
.HasColumnType("INTEGER");
entity.Property<string>("ParentFileId")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<string>("RelativePath")
.IsRequired()
.HasColumnType("TEXT");
entity.Property<double>("Score")
.HasColumnType("REAL");
entity.HasNoKey();
entity.ToView("embedding_chunk_search_results");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.EmbeddingStateFileEntity", entity =>
{
entity.HasOne("AIStudio.Tools.Databases.IndexStore.EmbeddingStateDataSourceEntity", "DataSource")
.WithMany("Files")
.HasForeignKey("DataSourceId")
.OnDelete(DeleteBehavior.Cascade)
.IsRequired();
entity.Navigation("DataSource");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.EmbeddingStateChunkEntity", entity =>
{
entity.HasOne("AIStudio.Tools.Databases.IndexStore.EmbeddingStateFileEntity", "File")
.WithMany("Chunks")
.HasForeignKey("ParentFileId")
.OnDelete(DeleteBehavior.Cascade)
.IsRequired();
entity.Navigation("File");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.IndexingFailureEntity", entity =>
{
entity.HasOne("AIStudio.Tools.Databases.IndexStore.EmbeddingStateDataSourceEntity", "DataSource")
.WithMany("PermanentIndexingFailures")
.HasForeignKey("DataSourceId")
.OnDelete(DeleteBehavior.Cascade)
.IsRequired();
entity.Navigation("DataSource");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.EmbeddingStateDataSourceEntity", entity =>
{
entity.Navigation("Files");
entity.Navigation("PermanentIndexingFailures");
});
modelBuilder.Entity("AIStudio.Tools.Databases.IndexStore.EmbeddingStateFileEntity", entity =>
{
entity.Navigation("Chunks");
});
#pragma warning restore 612, 618
}
}
@@ -0,0 +1,61 @@
using AIStudio.Tools.PluginSystem;
using AIStudio.Tools.Services;
namespace AIStudio.Tools.Databases.IndexStore;
public sealed class NoIndexStoreClient(string name, string? unavailableReason, DatabaseClientStatus status = DatabaseClientStatus.UNAVAILABLE) : IndexStoreClient(name, string.Empty)
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(NoIndexStoreClient).Namespace, nameof(NoIndexStoreClient));
public override DatabaseClientStatus Status => status;
public override async IAsyncEnumerable<(string Label, string Value)> GetDisplayInfo()
{
yield return (TB("Status"), status switch
{
DatabaseClientStatus.STARTING => TB("Starting"),
_ => TB("Unavailable")
});
if (!string.IsNullOrWhiteSpace(unavailableReason))
yield return (TB("Reason"), unavailableReason);
await Task.CompletedTask;
}
public override Task<DataSourceEmbeddingManifest> GetManifestAsync(string dataSourceId, CancellationToken token) =>
Task.FromResult(new DataSourceEmbeddingManifest());
public override Task UpsertDataSourceAsync(
string dataSourceId,
string dataSourceName,
string dataSourceType,
string embeddingProviderId,
string embeddingSignature,
string sourceHash,
int vectorSize,
CancellationToken token) => Task.CompletedTask;
public override Task UpdateVectorSizeAsync(string dataSourceId, int vectorSize, CancellationToken token) => Task.CompletedTask;
public override Task UpdateDataSourceHashAsync(string dataSourceId, string sourceHash, CancellationToken token) => Task.CompletedTask;
public override Task UpsertFileAsync(string dataSourceId, EmbeddingStateFile file, CancellationToken token) => Task.CompletedTask;
public override Task DeleteFileAsync(string dataSourceId, string filePath, CancellationToken token) => Task.CompletedTask;
public override Task UpsertPermanentFailureAsync(string dataSourceId, PermanentIndexingFailure failure, CancellationToken token) => Task.CompletedTask;
public override Task DeletePermanentFailureAsync(string dataSourceId, string filePath, CancellationToken token) => Task.CompletedTask;
public override Task UpsertChunksAsync(string dataSourceId, IReadOnlyList<EmbeddingStateChunk> chunks, CancellationToken token) => Task.CompletedTask;
public override Task<IReadOnlyList<IndexStoreSearchResult>> SearchChunksAsync(string dataSourceId, string query, int maxMatches, CancellationToken token) =>
Task.FromResult<IReadOnlyList<IndexStoreSearchResult>>([]);
public override Task DeleteDataSourceAsync(string dataSourceId, CancellationToken token) => Task.CompletedTask;
public override void Dispose()
{
}
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Databases.IndexStore;
public sealed record PermanentIndexingFailure(string ParentFileId, string AbsolutePath, string Fingerprint, FileExtractionErrorCode Code, string Message, DateTimeOffset OccurredAtUtc);
@@ -0,0 +1,472 @@
using System.Globalization;
using System.Text.RegularExpressions;
using AIStudio.Settings;
using AIStudio.Tools.PluginSystem;
using AIStudio.Tools.Services;
using Microsoft.EntityFrameworkCore;
namespace AIStudio.Tools.Databases.IndexStore;
public sealed class SqliteIndexStoreClientImplementation(string name, string databasePath, string basePath, string version) : IndexStoreClient(name, basePath)
{
private const string DATABASE_NAME = "Local RAG Index";
private const string DATABASE_FILENAME = "rag-index.sqlite3";
private const int MAX_FTS_QUERY_TERMS = 32;
private const int CHUNK_UPSERT_BATCH_SIZE = 500;
private static readonly Regex FTS_TOKEN_REGEX = new(@"[\p{L}\p{Nd}_]+", RegexOptions.Compiled | RegexOptions.CultureInvariant);
private readonly string databasePath = databasePath;
private readonly DbContextOptions<IndexStoreDbContext> dbContextOptions = IndexStoreDbContext.CreateOptions(databasePath);
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(SqliteIndexStoreClientImplementation).Namespace, nameof(SqliteIndexStoreClientImplementation));
public override string CacheKey => $"{this.Name}:{this.databasePath}:{version}";
public static async Task<DatabaseClient> CreateAsync(
ILogger logger,
ILogger<DatabaseClient> databaseClientLogger,
CancellationToken cancellationToken)
{
if (string.IsNullOrWhiteSpace(SettingsManager.DataDirectory))
return CreateNoIndexStoreClient(DATABASE_NAME, "The application data directory is not available yet.", DatabaseClientStatus.STARTING, databaseClientLogger);
try
{
SQLitePCL.Batteries_V2.Init();
var basePath = Path.Combine(SettingsManager.DataDirectory, "databases", "sqlite");
Directory.CreateDirectory(basePath);
var databasePath = Path.Combine(basePath, DATABASE_FILENAME);
var client = new SqliteIndexStoreClientImplementation(DATABASE_NAME, databasePath, basePath, string.Empty);
await client.InitializeAsync(cancellationToken);
var version = await client.GetSqliteVersionAsync(cancellationToken);
client = new SqliteIndexStoreClientImplementation(DATABASE_NAME, databasePath, basePath, version);
client.SetLogger(databaseClientLogger);
return client;
}
catch (Exception exception)
{
logger.LogWarning(exception, "{DatabaseName} is not available. Indexed file fingerprints and search chunks are disabled.", DATABASE_NAME);
return CreateNoIndexStoreClient(DATABASE_NAME, exception.Message, DatabaseClientStatus.UNAVAILABLE, databaseClientLogger);
}
}
public override async IAsyncEnumerable<(string Label, string Value)> GetDisplayInfo()
{
await using var context = this.CreateContext();
yield return (TB("Reported version"), version);
yield return (TB("Database path"), this.databasePath);
yield return (TB("Storage size"), this.GetStorageSize());
yield return (TB("Indexed data sources"), (await context.DataSources.CountAsync(CancellationToken.None)).ToString(CultureInfo.InvariantCulture));
yield return (TB("Indexed files"), (await context.EmbeddedFiles.CountAsync(CancellationToken.None)).ToString(CultureInfo.InvariantCulture));
yield return (TB("Search chunks"), (await context.EmbeddingChunks.CountAsync(CancellationToken.None)).ToString(CultureInfo.InvariantCulture));
yield return (TB("Permanently skipped files"), (await context.PermanentIndexingFailures.CountAsync(CancellationToken.None)).ToString(CultureInfo.InvariantCulture));
}
public override async Task<DataSourceEmbeddingManifest> GetManifestAsync(string dataSourceId, CancellationToken token)
{
await using var context = this.CreateContext();
var manifest = new DataSourceEmbeddingManifest();
var dataSource = await context.DataSources
.AsNoTracking()
.FirstOrDefaultAsync(source => source.DataSourceId == dataSourceId, token);
if (dataSource is null)
return manifest;
manifest.EmbeddingProviderId = dataSource.EmbeddingProviderId;
manifest.EmbeddingSignature = dataSource.EmbeddingSignature;
manifest.SourceHash = dataSource.SourceHash;
manifest.VectorSize = dataSource.VectorSize;
var files = await context.EmbeddedFiles
.AsNoTracking()
.Where(file => file.DataSourceId == dataSourceId && file.ChunkCount > 0)
.ToListAsync(token);
foreach (var file in files)
{
manifest.Files[file.AbsolutePath] = new EmbeddedFileRecord(
file.Fingerprint,
file.FileSize,
file.LastWriteUtc,
file.EmbeddedAtUtc,
file.ChunkCount);
}
var permanentFailures = await context.PermanentIndexingFailures
.AsNoTracking()
.Where(failure => failure.DataSourceId == dataSourceId)
.ToListAsync(token);
foreach (var failure in permanentFailures)
{
manifest.PermanentFailures[failure.AbsolutePath] = new PermanentIndexingFailureRecord(
failure.Fingerprint,
ParseFailureCode(failure.FailureCode),
failure.FailureMessage,
failure.OccurredAtUtc);
}
return manifest;
}
public override async Task UpsertDataSourceAsync(
string dataSourceId,
string dataSourceName,
string dataSourceType,
string embeddingProviderId,
string embeddingSignature,
string sourceHash,
int vectorSize,
CancellationToken token)
{
await using var context = this.CreateContext();
var dataSource = await context.DataSources.FirstOrDefaultAsync(source => source.DataSourceId == dataSourceId, token);
if (dataSource is null)
{
dataSource = new EmbeddingStateDataSourceEntity
{
DataSourceId = dataSourceId,
};
context.DataSources.Add(dataSource);
}
ApplyDataSource(dataSource, dataSourceName, dataSourceType, embeddingProviderId, embeddingSignature, sourceHash, vectorSize);
await context.SaveChangesAsync(token);
}
public override async Task UpdateVectorSizeAsync(string dataSourceId, int vectorSize, CancellationToken token)
{
await using var context = this.CreateContext();
var dataSource = await context.DataSources.FirstOrDefaultAsync(source => source.DataSourceId == dataSourceId, token);
if (dataSource is null)
return;
dataSource.VectorSize = vectorSize;
dataSource.UpdatedAtUtc = DateTimeOffset.UtcNow;
await context.SaveChangesAsync(token);
}
public override async Task UpdateDataSourceHashAsync(string dataSourceId, string sourceHash, CancellationToken token)
{
await using var context = this.CreateContext();
var dataSource = await context.DataSources.FirstOrDefaultAsync(source => source.DataSourceId == dataSourceId, token);
if (dataSource is null)
return;
dataSource.SourceHash = sourceHash;
dataSource.UpdatedAtUtc = DateTimeOffset.UtcNow;
await context.SaveChangesAsync(token);
}
public override async Task UpsertFileAsync(string dataSourceId, EmbeddingStateFile file, CancellationToken token)
{
await using var context = this.CreateContext();
var fileEntity = await context.EmbeddedFiles.FirstOrDefaultAsync(entity => entity.ParentFileId == file.ParentFileId, token);
if (fileEntity is null)
{
fileEntity = new EmbeddingStateFileEntity
{
ParentFileId = file.ParentFileId,
};
context.EmbeddedFiles.Add(fileEntity);
}
ApplyFile(fileEntity, dataSourceId, file);
await context.SaveChangesAsync(token);
}
public override async Task DeleteFileAsync(string dataSourceId, string filePath, CancellationToken token)
{
await using var context = this.CreateContext();
await using var transaction = await context.Database.BeginTransactionAsync(token);
var parentFileIds = await context.EmbeddedFiles
.Where(file => file.DataSourceId == dataSourceId && file.AbsolutePath == filePath)
.Select(file => file.ParentFileId)
.ToListAsync(token);
foreach (var parentFileIdBatch in parentFileIds.Chunk(CHUNK_UPSERT_BATCH_SIZE))
await context.EmbeddingChunks
.Where(chunk => parentFileIdBatch.Contains(chunk.ParentFileId))
.ExecuteDeleteAsync(token);
await context.EmbeddedFiles
.Where(file => file.DataSourceId == dataSourceId && file.AbsolutePath == filePath)
.ExecuteDeleteAsync(token);
await transaction.CommitAsync(token);
}
public override async Task UpsertPermanentFailureAsync(string dataSourceId, PermanentIndexingFailure failure, CancellationToken token)
{
await using var context = this.CreateContext();
var failureEntity = await context.PermanentIndexingFailures.FirstOrDefaultAsync(entity => entity.ParentFileId == failure.ParentFileId, token);
if (failureEntity is null)
{
failureEntity = new IndexingFailureEntity
{
ParentFileId = failure.ParentFileId,
};
context.PermanentIndexingFailures.Add(failureEntity);
}
ApplyPermanentFailure(failureEntity, dataSourceId, failure);
await context.SaveChangesAsync(token);
}
public override async Task DeletePermanentFailureAsync(string dataSourceId, string filePath, CancellationToken token)
{
await using var context = this.CreateContext();
await context.PermanentIndexingFailures
.Where(failure => failure.DataSourceId == dataSourceId && failure.AbsolutePath == filePath)
.ExecuteDeleteAsync(token);
}
public override async Task UpsertChunksAsync(string dataSourceId, IReadOnlyList<EmbeddingStateChunk> chunks, CancellationToken token)
{
if (chunks.Count == 0)
return;
await using var context = this.CreateContext();
await using var transaction = await context.Database.BeginTransactionAsync(token);
foreach (var chunkBatch in chunks.Chunk(CHUNK_UPSERT_BATCH_SIZE))
{
token.ThrowIfCancellationRequested();
var chunkIds = chunkBatch
.Select(chunk => chunk.ChunkId)
.Distinct(StringComparer.Ordinal)
.ToArray();
var existingChunks = await context.EmbeddingChunks
.Where(chunk => chunkIds.Contains(chunk.ChunkId))
.ToDictionaryAsync(chunk => chunk.ChunkId, StringComparer.Ordinal, token);
foreach (var chunk in chunkBatch)
{
if (!existingChunks.TryGetValue(chunk.ChunkId, out var chunkEntity))
{
chunkEntity = new EmbeddingStateChunkEntity
{
ChunkId = chunk.ChunkId,
};
context.EmbeddingChunks.Add(chunkEntity);
existingChunks[chunk.ChunkId] = chunkEntity;
}
ApplyChunk(chunkEntity, chunk);
}
await context.SaveChangesAsync(token);
context.ChangeTracker.Clear();
}
await transaction.CommitAsync(token);
}
public override async Task<IReadOnlyList<IndexStoreSearchResult>> SearchChunksAsync(string dataSourceId, string query, int maxMatches, CancellationToken token)
{
if (maxMatches <= 0)
return [];
var ftsQuery = BuildFtsQuery(query);
if (string.IsNullOrWhiteSpace(ftsQuery))
return [];
await using var context = this.CreateContext();
var results = await context.SearchResults
.FromSqlInterpolated($"""
SELECT
c.chunk_id AS ChunkId,
c.parent_file_id AS ParentFileId,
ds.data_source_id AS DataSourceId,
ds.data_source_name AS DataSourceName,
ds.data_source_type AS DataSourceType,
f.absolute_path AS AbsolutePath,
f.file_name AS FileName,
f.relative_path AS RelativePath,
f.file_type AS FileType,
c.page_number AS PageNumber,
c.chunk_index AS ChunkIndex,
c.chunk_text AS ChunkText,
bm25(embedding_chunks_fts) AS Score,
f.fingerprint AS Fingerprint,
f.file_size AS FileSize,
f.creation_utc AS CreationUtc,
f.last_write_utc AS LastWriteUtc,
c.embedded_at_utc AS EmbeddedAtUtc,
f.chunk_count AS ChunkCount,
f.confidence_level AS ConfidenceLevel,
f.confidence_level_rank AS ConfidenceLevelRank
FROM embedding_chunks_fts
JOIN embedding_chunks c ON c.id = embedding_chunks_fts.rowid
JOIN embedded_files f ON f.parent_file_id = c.parent_file_id
JOIN data_sources ds ON ds.data_source_id = f.data_source_id
WHERE ds.data_source_id = {dataSourceId}
AND embedding_chunks_fts MATCH {ftsQuery}
ORDER BY Score
LIMIT {maxMatches}
""")
.AsNoTracking()
.ToListAsync(token);
return results.Select(ToSearchResult).ToList();
}
public override async Task DeleteDataSourceAsync(string dataSourceId, CancellationToken token)
{
await using var context = this.CreateContext();
await using var transaction = await context.Database.BeginTransactionAsync(token);
var parentFileIds = await context.EmbeddedFiles
.Where(file => file.DataSourceId == dataSourceId)
.Select(file => file.ParentFileId)
.ToListAsync(token);
foreach (var parentFileIdBatch in parentFileIds.Chunk(CHUNK_UPSERT_BATCH_SIZE))
await context.EmbeddingChunks
.Where(chunk => parentFileIdBatch.Contains(chunk.ParentFileId))
.ExecuteDeleteAsync(token);
await context.DataSources
.Where(source => source.DataSourceId == dataSourceId)
.ExecuteDeleteAsync(token);
await transaction.CommitAsync(token);
}
public override void Dispose()
{
}
private async Task InitializeAsync(CancellationToken token)
{
await using var context = this.CreateContext();
await IndexStoreSchemaMigrator.MigrateAsync(context, token);
}
private async Task<string> GetSqliteVersionAsync(CancellationToken token)
{
await using var context = this.CreateContext();
var versions = await context.Database
.SqlQueryRaw<string>("SELECT sqlite_version() AS Value")
.ToListAsync(token);
return versions.FirstOrDefault() ?? string.Empty;
}
private IndexStoreDbContext CreateContext() => new(this.dbContextOptions);
private static void ApplyDataSource(
EmbeddingStateDataSourceEntity dataSource,
string dataSourceName,
string dataSourceType,
string embeddingProviderId,
string embeddingSignature,
string sourceHash,
int vectorSize)
{
dataSource.DataSourceName = dataSourceName;
dataSource.DataSourceType = dataSourceType;
dataSource.EmbeddingProviderId = embeddingProviderId;
dataSource.EmbeddingSignature = embeddingSignature;
dataSource.SourceHash = sourceHash;
dataSource.VectorSize = vectorSize;
dataSource.UpdatedAtUtc = DateTimeOffset.UtcNow;
}
private static void ApplyFile(EmbeddingStateFileEntity fileEntity, string dataSourceId, EmbeddingStateFile file)
{
fileEntity.DataSourceId = dataSourceId;
fileEntity.AbsolutePath = file.AbsolutePath;
fileEntity.FileName = file.FileName;
fileEntity.RelativePath = file.RelativePath;
fileEntity.FileType = file.FileType;
fileEntity.Fingerprint = file.Fingerprint;
fileEntity.FileSize = file.FileSize;
fileEntity.CreationUtc = file.CreationUtc;
fileEntity.LastWriteUtc = file.LastWriteUtc;
fileEntity.EmbeddedAtUtc = file.EmbeddedAtUtc;
fileEntity.ChunkCount = file.ChunkCount;
fileEntity.ConfidenceLevel = file.ConfidenceLevel;
fileEntity.ConfidenceLevelRank = file.ConfidenceLevelRank;
}
private static void ApplyPermanentFailure(IndexingFailureEntity failureEntity, string dataSourceId, PermanentIndexingFailure failure)
{
failureEntity.DataSourceId = dataSourceId;
failureEntity.AbsolutePath = failure.AbsolutePath;
failureEntity.Fingerprint = failure.Fingerprint;
failureEntity.FailureCode = failure.Code.ToString();
failureEntity.FailureMessage = failure.Message;
failureEntity.OccurredAtUtc = failure.OccurredAtUtc;
}
/// <remarks>
/// A row written by a newer version may name a code this one does not know. Such a row still
/// says that the file failed permanently, so it keeps its place in the manifest and only loses
/// the reason it names.
/// </remarks>
private static FileExtractionErrorCode ParseFailureCode(string failureCode) =>
Enum.TryParse<FileExtractionErrorCode>(failureCode, ignoreCase: true, out var parsedCode) ? parsedCode : FileExtractionErrorCode.UNKNOWN;
private static void ApplyChunk(EmbeddingStateChunkEntity chunkEntity, EmbeddingStateChunk chunk)
{
chunkEntity.ChunkId = chunk.ChunkId;
chunkEntity.ParentFileId = chunk.ParentFileId;
chunkEntity.PageNumber = chunk.PageNumber;
chunkEntity.ChunkIndex = chunk.ChunkIndex;
chunkEntity.ChunkText = chunk.ChunkText;
chunkEntity.EmbeddedAtUtc = chunk.EmbeddedAtUtc;
}
private static IndexStoreSearchResult ToSearchResult(IndexStoreSearchResultEntity result) => new(
result.ChunkId,
result.ParentFileId,
result.DataSourceId,
result.DataSourceName,
result.DataSourceType,
result.AbsolutePath,
result.FileName,
result.RelativePath,
result.FileType,
result.PageNumber,
result.ChunkIndex,
result.ChunkText,
result.Score,
result.Fingerprint,
result.FileSize,
result.CreationUtc,
result.LastWriteUtc,
result.EmbeddedAtUtc,
result.ChunkCount,
result.ConfidenceLevel,
result.ConfidenceLevelRank);
private static string BuildFtsQuery(string query)
{
var terms = FTS_TOKEN_REGEX
.Matches(query)
.Select(match => match.Value)
.Where(term => !string.IsNullOrWhiteSpace(term))
.Distinct(StringComparer.OrdinalIgnoreCase)
.Take(MAX_FTS_QUERY_TERMS)
.Select(term => $"\"{term.Replace("\"", "\"\"", StringComparison.Ordinal)}\"")
.ToList();
return terms.Count == 0 ? string.Empty : string.Join(" OR ", terms);
}
private static NoIndexStoreClient CreateNoIndexStoreClient(string name, string? unavailableReason, DatabaseClientStatus status, ILogger<DatabaseClient> databaseClientLogger)
{
var client = new NoIndexStoreClient(name, unavailableReason, status);
client.SetLogger(databaseClientLogger);
return client;
}
}
@@ -7,14 +7,10 @@ public sealed class NoDatabaseClient(string name, string? unavailableReason, Dat
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(NoDatabaseClient).Namespace, nameof(NoDatabaseClient));
public override DatabaseClientStatus Status => status;
public override async IAsyncEnumerable<(string Label, string Value)> GetDisplayInfo()
{
yield return (TB("Status"), status switch
{
DatabaseClientStatus.STARTING => TB("Starting"),
_ => TB("Unavailable")
});
yield return (TB("Status"), TB("Unavailable"));
if (!string.IsNullOrWhiteSpace(unavailableReason))
yield return (TB("Reason"), unavailableReason);
@@ -1,12 +0,0 @@
namespace AIStudio.Tools.Databases.VectorStore;
public interface IVectorStoreClient
{
Task EnsureVectorStoreExists(string storeName, int vectorSize, CancellationToken token);
Task InsertEmbedding(string storeName, IReadOnlyList<VectorStoragePoint> points, CancellationToken token);
Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token);
Task DeleteVectorStore(string storeName, CancellationToken token);
}
@@ -2,7 +2,7 @@ using AIStudio.Tools.PluginSystem;
namespace AIStudio.Tools.Databases.VectorStore;
public sealed class NoVectorStoreClient(string name, string? unavailableReason, DatabaseClientStatus status = DatabaseClientStatus.UNAVAILABLE) : DatabaseClient(name, string.Empty), IVectorStoreClient
public sealed class NoVectorStoreClient(string name, string? unavailableReason, DatabaseClientStatus status = DatabaseClientStatus.UNAVAILABLE) : VectorStoreClient(name, string.Empty)
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(NoVectorStoreClient).Namespace, nameof(NoVectorStoreClient));
@@ -22,16 +22,22 @@ public sealed class NoVectorStoreClient(string name, string? unavailableReason,
await Task.CompletedTask;
}
public Task EnsureVectorStoreExists(string storeName, int vectorSize, CancellationToken token) =>
public override Task<VectorStoreEnsureResult> EnsureVectorStoreExists(string storeName, string dataSourceName, int vectorSize, CancellationToken token) =>
Task.FromException<VectorStoreEnsureResult>(this.CreateUnavailableException());
public override Task InsertEmbedding(string storeName, IReadOnlyList<VectorStoragePoint> points, CancellationToken token) =>
Task.FromException(this.CreateUnavailableException());
public Task InsertEmbedding(string storeName, IReadOnlyList<VectorStoragePoint> points, CancellationToken token) =>
public override Task<IReadOnlyList<VectorSearchResult>> SearchEmbeddingAsync(string storeName, IReadOnlyList<float> vector, int maxMatches, CancellationToken token) =>
Task.FromException<IReadOnlyList<VectorSearchResult>>(this.CreateUnavailableException());
public override Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token) =>
Task.FromException(this.CreateUnavailableException());
public Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token) =>
public override Task OptimizeVectorStore(string storeName, CancellationToken token) =>
Task.FromException(this.CreateUnavailableException());
public Task DeleteVectorStore(string storeName, CancellationToken token) =>
public override Task DeleteVectorStore(string storeName, CancellationToken token) =>
Task.FromException(this.CreateUnavailableException());
private InvalidOperationException CreateUnavailableException() =>
@@ -9,13 +9,15 @@ public sealed class QdrantEdgeClientImplementation(
string path,
string version,
int storesCount,
RustService rustService) : DatabaseClient(name, path), IVectorStoreClient
RustService rustService) : VectorStoreClient(name, path)
{
private const string DATABASE_NAME = "Qdrant Edge";
private const string INFO_PATH = "/system/qdrant-edge/info";
private const string ENSURE_PATH = "/system/qdrant-edge/ensure";
private const string INSERT_PATH = "/system/qdrant-edge/insert";
private const string SEARCH_PATH = "/system/qdrant-edge/search";
private const string DELETE_FILE_PATH = "/system/qdrant-edge/delete-file";
private const string OPTIMIZE_PATH = "/system/qdrant-edge/optimize";
private const string DELETE_STORE_PATH = "/system/qdrant-edge/delete-store";
private readonly string path = path;
@@ -80,16 +82,32 @@ public sealed class QdrantEdgeClientImplementation(
yield return (TB("Number of vector stores"), displayStoresCount.ToString());
}
public Task EnsureVectorStoreExists(string storeName, int vectorSize, CancellationToken token) =>
rustService.ExecuteDatabaseOperation(DATABASE_NAME, ENSURE_PATH, new EnsureVectorStoreRequest(storeName, vectorSize), token);
public override async Task<VectorStoreEnsureResult> EnsureVectorStoreExists(string storeName, string dataSourceName, int vectorSize, CancellationToken token) =>
await rustService.ExecuteDatabaseQuery<EnsureVectorStoreRequest, VectorStoreEnsureResult>(DATABASE_NAME, ENSURE_PATH,
new EnsureVectorStoreRequest(storeName, dataSourceName, vectorSize), token) ?? throw new InvalidOperationException("The vector store ensure response was empty.");
public Task InsertEmbedding(string storeName, IReadOnlyList<VectorStoragePoint> points, CancellationToken token) =>
public override Task InsertEmbedding(string storeName, IReadOnlyList<VectorStoragePoint> points, CancellationToken token) =>
rustService.ExecuteDatabaseOperation(DATABASE_NAME, INSERT_PATH, new InsertEmbeddingRequest(storeName, points), token);
public Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token) =>
public override async Task<IReadOnlyList<VectorSearchResult>> SearchEmbeddingAsync(string storeName, IReadOnlyList<float> vector, int maxMatches, CancellationToken token)
{
if (maxMatches <= 0)
return [];
return await rustService.ExecuteDatabaseQuery<SearchEmbeddingRequest, List<VectorSearchResult>>(
DATABASE_NAME,
SEARCH_PATH,
new SearchEmbeddingRequest(storeName, vector, maxMatches),
token) ?? [];
}
public override Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token) =>
rustService.ExecuteDatabaseOperation(DATABASE_NAME, DELETE_FILE_PATH, new DeleteEmbeddingByFileRequest(storeName, filePath), token);
public Task DeleteVectorStore(string storeName, CancellationToken token) =>
public override Task OptimizeVectorStore(string storeName, CancellationToken token) =>
rustService.ExecuteDatabaseOperation(DATABASE_NAME, OPTIMIZE_PATH, new OptimizeVectorStoreRequest(storeName), token);
public override Task DeleteVectorStore(string storeName, CancellationToken token) =>
rustService.ExecuteDatabaseOperation(DATABASE_NAME, DELETE_STORE_PATH, new DeleteVectorStoreRequest(storeName), token);
public override void Dispose()
@@ -104,11 +122,15 @@ public sealed class QdrantEdgeClientImplementation(
}
// ReSharper disable NotAccessedPositionalProperty.Local
private sealed record EnsureVectorStoreRequest(string StoreName, int VectorSize);
private sealed record EnsureVectorStoreRequest(string StoreName, string DataSourceName, int VectorSize);
private sealed record InsertEmbeddingRequest(string StoreName, IReadOnlyList<VectorStoragePoint> Points);
private sealed record SearchEmbeddingRequest(string StoreName, IReadOnlyList<float> Vector, int MaxMatches);
private sealed record DeleteEmbeddingByFileRequest(string StoreName, string FilePath);
private sealed record OptimizeVectorStoreRequest(string StoreName);
private sealed record DeleteVectorStoreRequest(string StoreName);
// ReSharper restore NotAccessedPositionalProperty.Local
@@ -0,0 +1,24 @@
namespace AIStudio.Tools.Databases.VectorStore;
public sealed record VectorSearchResult(
string PointId,
double Score,
string DataSourceId,
string DataSourceName,
string DataSourceType,
string ChunkId,
string ParentFileId,
string FilePath,
string AbsolutePath,
string FileName,
string RelativePath,
string FileType,
int? PageNumber,
int ChunkIndex,
string Text,
string Fingerprint,
string CreationUtc,
string LastWriteUtc,
string EmbeddedAtUtc,
string ConfidenceLevel,
int ConfidenceLevelRank);
@@ -6,11 +6,19 @@ public sealed record VectorStoragePoint(
string DataSourceId,
string DataSourceName,
string DataSourceType,
string ChunkId,
string ParentFileId,
string FilePath,
string AbsolutePath,
string FileName,
string RelativePath,
string FileType,
int? PageNumber,
int ChunkIndex,
string Text,
string Fingerprint,
DateTime LastWriteUtc,
DateTime EmbeddedAtUtc);
DateTimeOffset CreationUtc,
DateTimeOffset LastWriteUtc,
DateTimeOffset EmbeddedAtUtc,
string ConfidenceLevel,
int ConfidenceLevelRank);
@@ -0,0 +1,16 @@
namespace AIStudio.Tools.Databases.VectorStore;
public abstract class VectorStoreClient(string name, string path): DatabaseClient(name, path)
{
public abstract Task<VectorStoreEnsureResult> EnsureVectorStoreExists(string storeName, string dataSourceName, int vectorSize, CancellationToken token);
public abstract Task InsertEmbedding(string storeName, IReadOnlyList<VectorStoragePoint> points, CancellationToken token);
public abstract Task<IReadOnlyList<VectorSearchResult>> SearchEmbeddingAsync(string storeName, IReadOnlyList<float> vector, int maxMatches, CancellationToken token);
public abstract Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token);
public abstract Task OptimizeVectorStore(string storeName, CancellationToken token);
public abstract Task DeleteVectorStore(string storeName, CancellationToken token);
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Databases.VectorStore;
public sealed record VectorStoreEnsureResult(bool Created);
@@ -9,12 +9,13 @@ namespace AIStudio.Tools;
public sealed class DocumentManager
{
private StringBuilder? currentPageContent;
private int? currentPageTokenCount;
public string? AddPage(ContentStreamDocumentMetadata metadata, string? content, bool extractImages)
public ContentStreamPendingContent? AddPage(ContentStreamDocumentMetadata metadata, string? content, int? tokenCount, bool extractImages)
{
var pageNumber = metadata.Document?.PageNumber ?? 0;
if (pageNumber == 0)
return content;
return content is null ? null : new ContentStreamPendingContent(content, tokenCount);
var image = metadata.Document?.Image;
if (image is null)
@@ -32,6 +33,12 @@ public sealed class DocumentManager
this.currentPageContent.AppendLine($"<!-- Estimated page {pageNumber} -->");
this.currentPageContent.AppendLine();
this.currentPageContent.Append(content);
//
// The count waits here together with the page it belongs to. Handing it out along with
// the page we just completed would size that page by the text of this one.
//
this.currentPageTokenCount = tokenCount;
return completedPage;
}
@@ -45,19 +52,28 @@ public sealed class DocumentManager
{
this.currentPageContent.AppendLine();
this.currentPageContent.AppendLine(markdownImage);
//
// The runtime counted the text of this page, not the image we just embedded into it.
// A data URI is orders of magnitude larger than that text, so the count no longer
// describes the page: we drop it, and whoever needs one counts the page itself.
//
this.currentPageTokenCount = null;
}
}
return null;
}
public string? Flush()
public ContentStreamPendingContent? Flush()
{
if (this.currentPageContent is null)
return null;
var result = this.currentPageContent.ToString();
var tokenCount = this.currentPageTokenCount;
this.currentPageContent = null;
return string.IsNullOrWhiteSpace(result) ? null : result;
this.currentPageTokenCount = null;
return string.IsNullOrWhiteSpace(result) ? null : new ContentStreamPendingContent(result, tokenCount);
}
}
+1
View File
@@ -203,6 +203,7 @@ public enum Event
/// Carries data sources that were automatically selected for retrieval-augmented generation.
/// </summary>
RAG_AUTO_DATA_SOURCES_SELECTED,
RAG_EMBEDDING_STATUS_CHANGED,
@@ -179,7 +179,7 @@ public static class FileExportFormatExtensions
FileExportFormat.OPEN_DOCUMENT_TEXT => FileTypes.ODT,
FileExportFormat.LATEX => FileTypes.TEX,
FileExportFormat.MARKDOWN => FileTypes.MARKDOWN,
FileExportFormat.HTML => FileTypes.HTML,
FileExportFormat.HTML => FileTypes.HTML_DOCUMENT,
FileExportFormat.CSV => FileTypes.CSV,
FileExportFormat.TSV => FileTypes.TSV,
@@ -32,6 +32,13 @@ public enum FileExtractionErrorCode
FORMAT_DETECTION_FAILED,
NOT_A_VALID_PDF,
NOT_A_VALID_SPREADSHEET,
/// <summary>
/// The package of a Word, OpenDocument, or presentation file is broken, e.g. a damaged
/// archive or a missing part inside it.
/// </summary>
NOT_A_VALID_DOCUMENT,
PDFIUM_UNAVAILABLE,
PDF_ENCRYPTED,
PAGE_EXTRACTION_FAILED,
@@ -0,0 +1,140 @@
using AIStudio.Tools.PluginSystem;
namespace AIStudio.Tools;
/// <summary>
/// Tells failures which lie in the file apart from failures which lie in its surroundings, and
/// puts both into words for the indexing user interface.
/// </summary>
/// <remarks>
/// A file without readable text fails the same way on every run, so the indexer remembers it and
/// waits for the file to change. An offline network drive or an overloaded provider says nothing
/// about the file itself, which is why those keep being retried.
/// </remarks>
internal static class FileExtractionErrorCodeExtensions
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(FileExtractionErrorCodeExtensions).Namespace, nameof(FileExtractionErrorCodeExtensions));
/// <summary>
/// Gets a value indicating whether reading the file again will fail again, as long as the file
/// itself does not change.
/// </summary>
/// <param name="code">The stable failure code.</param>
/// <returns>True, when the reason lies in the file itself.</returns>
internal static bool IsPermanentIndexingFailure(this FileExtractionErrorCode code) => code switch
{
//
// The reason lies in the file. Reading it again without changing it produces the same
// outcome, so the indexer waits for a new fingerprint:
//
FileExtractionErrorCode.NO_TEXT_EXTRACTED => true,
FileExtractionErrorCode.NO_CONTENT => true,
FileExtractionErrorCode.NOT_TEXT_CONTENT => true,
FileExtractionErrorCode.NOT_A_VALID_PDF => true,
FileExtractionErrorCode.NOT_A_VALID_SPREADSHEET => true,
FileExtractionErrorCode.NOT_A_VALID_DOCUMENT => true,
FileExtractionErrorCode.PDF_ENCRYPTED => true,
FileExtractionErrorCode.FORMAT_DETECTION_FAILED => true,
FileExtractionErrorCode.EXECUTABLE_REJECTED => true,
FileExtractionErrorCode.UNSUPPORTED => true,
// Pages holding nothing but images are one of the recurring cases here, and that is a
// property of the document, not of the environment:
FileExtractionErrorCode.PAGE_EXTRACTION_FAILED => true,
//
// Everything else depends on the surroundings: an unavailable drive, a file someone else
// has open, a missing engine, or a runtime which did not answer in time. All of them are
// worth another attempt during the next run:
//
_ => false,
};
/// <summary>
/// Names the cause in a few words.
/// </summary>
/// <remarks>
/// Used to group the files of an indexing run by what happened to them: nine hundred entries
/// which all say the same sentence are one cause, not nine hundred.
/// </remarks>
/// <param name="code">The stable failure code.</param>
/// <returns>The localized name of the cause, or an empty text when the code has none.</returns>
internal static string GetIndexingCauseName(this FileExtractionErrorCode code) => code switch
{
FileExtractionErrorCode.NO_TEXT_EXTRACTED => TB("No readable text"),
FileExtractionErrorCode.NO_CONTENT => TB("No content"),
FileExtractionErrorCode.NOT_TEXT_CONTENT => TB("Not a text file"),
FileExtractionErrorCode.NOT_A_VALID_PDF => TB("Not a readable PDF"),
FileExtractionErrorCode.NOT_A_VALID_SPREADSHEET => TB("Not a readable spreadsheet"),
FileExtractionErrorCode.NOT_A_VALID_DOCUMENT => TB("Not a readable document"),
FileExtractionErrorCode.PDF_ENCRYPTED => TB("Protected PDF"),
FileExtractionErrorCode.FORMAT_DETECTION_FAILED => TB("Unknown file type"),
FileExtractionErrorCode.EXECUTABLE_REJECTED => TB("Executable program"),
FileExtractionErrorCode.UNSUPPORTED => TB("Unsupported file type"),
FileExtractionErrorCode.PAGE_EXTRACTION_FAILED => TB("Pages without readable text"),
FileExtractionErrorCode.FILE_NOT_FOUND => TB("File does not exist anymore"),
FileExtractionErrorCode.FILE_NOT_READABLE => TB("File could not be read"),
FileExtractionErrorCode.FILE_LOCKED => TB("File is open elsewhere"),
FileExtractionErrorCode.TIMEOUT => TB("Reading took too long"),
FileExtractionErrorCode.PDFIUM_UNAVAILABLE => TB("PDF system unavailable"),
FileExtractionErrorCode.PANDOC_UNAVAILABLE => TB("Pandoc unavailable"),
// Nothing about these lies in the file: AI Studio asked its runtime for the content and
// got back something it cannot work with. One name for all of them, because that is the
// one thing the user can tell from them:
FileExtractionErrorCode.INVALID_RESPONSE => TB("Internal error"),
FileExtractionErrorCode.INVALID_REQUEST => TB("Internal error"),
FileExtractionErrorCode.REQUEST_FAILED => TB("Internal error"),
FileExtractionErrorCode.INTERNAL => TB("Internal error"),
// Codes which say nothing beyond the message of the single file. The caller names those
// files itself and shows their messages instead:
_ => string.Empty,
};
/// <summary>
/// Gets the localized message which explains why a file was not indexed.
/// </summary>
/// <remarks>
/// These texts are the counterpart of the ones used for chat attachments: there, a file which
/// cannot be read is simply not sent, while here it stays out of the index and the user needs
/// to know whether AI Studio will come back to it on its own.
/// </remarks>
/// <param name="code">The stable failure code.</param>
/// <param name="fileName">The name of the file, as shown to the user.</param>
/// <returns>The localized message.</returns>
internal static string ToIndexingUserMessage(this FileExtractionErrorCode code, string fileName) => string.Format(ToIndexingMessageFormat(code), fileName);
private static string ToIndexingMessageFormat(FileExtractionErrorCode code) => code switch
{
//
// Permanent failures. Each of them names what is wrong with the file and says that AI
// Studio comes back to it once the file changes:
//
FileExtractionErrorCode.NO_TEXT_EXTRACTED => TB("No text could be read from the file '{0}', so it was not indexed. It might contain images only, such as a scanned PDF without a text layer. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.NO_CONTENT => TB("The file '{0}' did not provide any content, so it was not indexed. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.NOT_TEXT_CONTENT => TB("The file '{0}' is not a text file, so it was not indexed. Its content could not be read as text, which means it might have a wrong file extension. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.NOT_A_VALID_PDF => TB("The file '{0}' is not a readable PDF, so it was not indexed. It might be damaged or transferred incompletely. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.NOT_A_VALID_SPREADSHEET => TB("The file '{0}' is not a readable spreadsheet, so it was not indexed. It might be damaged or transferred incompletely. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.NOT_A_VALID_DOCUMENT => TB("The file '{0}' is not a readable document, so it was not indexed. It might be damaged or transferred incompletely. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.PDF_ENCRYPTED => TB("The file '{0}' is protected and could not be opened, so it was not indexed. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.FORMAT_DETECTION_FAILED => TB("The file type of '{0}' could not be determined, so the file was not indexed. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.EXECUTABLE_REJECTED => TB("The file '{0}' is an executable program and was not indexed, regardless of its file extension."),
FileExtractionErrorCode.UNSUPPORTED => TB("The file type of '{0}' is not supported, so the file was not indexed. AI Studio reads it again as soon as the file changes."),
FileExtractionErrorCode.PAGE_EXTRACTION_FAILED => TB("Pages of the file '{0}' could not be read, so it was not indexed. They might contain images only. AI Studio reads it again as soon as the file changes."),
//
// Temporary failures. They name what the user can act on, and every one of them is tried
// again during the next run:
//
FileExtractionErrorCode.FILE_NOT_FOUND => TB("The file '{0}' does not exist anymore and was not indexed."),
FileExtractionErrorCode.FILE_NOT_READABLE => TB("The file '{0}' could not be read and was not indexed. When the file is stored on a network drive, the drive might be unavailable, or another program might be blocking the file. AI Studio tries again during the next run."),
FileExtractionErrorCode.FILE_LOCKED => TB("The file '{0}' is currently open in another program, which is why it was not indexed. When the file is stored on a shared network drive, a colleague might have it open. AI Studio tries again during the next run."),
FileExtractionErrorCode.TIMEOUT => TB("Reading the file '{0}' took too long and was stopped, so the file was not indexed. When the file is stored on a network drive, the connection might be slow or interrupted. AI Studio tries again during the next run."),
FileExtractionErrorCode.PDFIUM_UNAVAILABLE => TB("AI Studio was not able to start its PDF engine, so the file '{0}' was not indexed. AI Studio tries again during the next run."),
FileExtractionErrorCode.PANDOC_UNAVAILABLE => TB("Reading the file '{0}' needs Pandoc, which is not available, so the file was not indexed. AI Studio tries again during the next run."),
_ => TB("The file '{0}' could not be read and was not indexed. AI Studio tries again during the next run."),
};
}
@@ -0,0 +1,27 @@
namespace AIStudio.Tools;
/// <summary>
/// Thrown when a file could not be read, carrying the stable failure code along with the message.
/// </summary>
/// <remarks>
/// The plain message alone does not say whether another attempt is worth anything. The code does,
/// which is what the indexer needs to tell a file without readable text apart from a network drive
/// which happens to be offline.
/// </remarks>
public sealed class FileExtractionException(FileExtractionErrorCode code, string message, int? pageNumber = null, string? detectedFormat = null) : Exception(message)
{
/// <summary>
/// Gets the stable failure code.
/// </summary>
public FileExtractionErrorCode Code { get; } = code;
/// <summary>
/// Gets the page the failure belongs to, when the failure affects a single page only.
/// </summary>
public int? PageNumber { get; } = pageNumber;
/// <summary>
/// Gets the format the runtime identified by looking at the content.
/// </summary>
public string? DetectedFormat { get; } = detectedFormat;
}
@@ -78,6 +78,7 @@ internal static class FileExtractionResultExtensions
FileExtractionErrorCode.TIMEOUT => TB("Reading the file '{0}' took too long and was stopped, so the file was not sent. When the file is stored on a network drive, the connection might be slow or interrupted."),
FileExtractionErrorCode.NOT_A_VALID_PDF => TB("The file '{0}' is not a readable PDF and was not sent. It might be damaged or transferred incompletely."),
FileExtractionErrorCode.NOT_A_VALID_SPREADSHEET => TB("The file '{0}' is not a readable spreadsheet and was not sent. It might be damaged or transferred incompletely."),
FileExtractionErrorCode.NOT_A_VALID_DOCUMENT => TB("The file '{0}' is not a readable document and was not sent. It might be damaged or transferred incompletely."),
FileExtractionErrorCode.PDF_ENCRYPTED => TB("The file '{0}' is protected and could not be opened, so it was not sent."),
FileExtractionErrorCode.PDFIUM_UNAVAILABLE => TB("AI Studio was not able to start its PDF engine, so the file '{0}' was not sent."),
FileExtractionErrorCode.PANDOC_UNAVAILABLE => TB("Reading the file '{0}' needs Pandoc, which is not available, so the file was not sent."),
@@ -70,6 +70,8 @@ public sealed class PluginConfiguration(bool isInternal, LuaState state, PluginT
if (!dryRun)
{
await PluginConfigurationObject.SyncManagedTokenizersAsync(this.Id, this.PluginPath);
// Store any decrypted API keys from enterprise configuration in the OS keyring:
await StoreEnterpriseApiKeysAsync();
await StoreEnterpriseSecretsAsync();
@@ -1,3 +1,4 @@
using System.Diagnostics.CodeAnalysis;
using System.Linq.Expressions;
using AIStudio.Settings;
@@ -206,6 +207,43 @@ public sealed record PluginConfigurationObject
return true;
}
[SuppressMessage("Usage", "MWAIS0001:Direct access to `Providers` is not allowed", Justification = "Tokenizer synchronization needs indexed access to update enterprise-managed providers in place.")]
public static async Task<bool> SyncManagedTokenizersAsync(Guid configPluginId, string pluginPath)
{
var wasConfigurationChanged = false;
var localSettingsManager = SettingsManagerAccess;
for (var i = 0; i < localSettingsManager.ConfigurationData.Providers.Count; i++)
{
var provider = localSettingsManager.ConfigurationData.Providers[i];
if (!provider.IsEnterpriseConfiguration || provider.EnterpriseConfigurationPluginId != configPluginId)
continue;
var syncedProvider = await SyncProviderTokenizerAsync(provider, pluginPath);
if (syncedProvider == provider)
continue;
localSettingsManager.ConfigurationData.Providers[i] = syncedProvider;
wasConfigurationChanged = true;
}
for (var i = 0; i < localSettingsManager.ConfigurationData.EmbeddingProviders.Count; i++)
{
var provider = localSettingsManager.ConfigurationData.EmbeddingProviders[i];
if (!provider.IsEnterpriseConfiguration || provider.EnterpriseConfigurationPluginId != configPluginId)
continue;
var syncedProvider = await SyncEmbeddingTokenizerAsync(provider, pluginPath);
if (syncedProvider == provider)
continue;
localSettingsManager.ConfigurationData.EmbeddingProviders[i] = syncedProvider;
wasConfigurationChanged = true;
}
return wasConfigurationChanged;
}
/// <summary>
/// Parses configured data sources from a configuration plugin.
/// </summary>
@@ -396,6 +434,19 @@ public sealed record PluginConfigurationObject
var wasConfigurationChanged = leftOverObjects.Count > 0;
foreach (var item in leftOverObjects.Distinct())
{
if (item is Settings.Provider provider)
{
var deleteTokenizerResult = await RustService.DeleteTokenizer(TokenizerModelId.ForProvider(provider));
if (!deleteTokenizerResult.Success)
LOG.LogWarning("Failed to delete tokenizer for removed enterprise provider '{ProviderName}': {Issue}", provider.InstanceName, deleteTokenizerResult.Message);
}
else if (item is EmbeddingProvider embeddingProvider)
{
var deleteTokenizerResult = await RustService.DeleteTokenizer(TokenizerModelId.ForEmbeddingProvider(embeddingProvider));
if (!deleteTokenizerResult.Success)
LOG.LogWarning("Failed to delete tokenizer for removed enterprise embedding provider '{ProviderName}': {Issue}", embeddingProvider.Name, deleteTokenizerResult.Message);
}
configuredObjects.Remove(item);
// Delete the API key from the OS keyring if the removed object has one:
@@ -426,4 +477,89 @@ public sealed record PluginConfigurationObject
return wasConfigurationChanged;
}
private static async Task<Settings.Provider> SyncProviderTokenizerAsync(Settings.Provider provider, string pluginPath)
{
var syncedTokenizerPath = await SyncTokenizerAsync(
provider.TokenizerPath,
pluginPath,
TokenizerModelId.ForProvider(provider),
$"provider '{provider.InstanceName}'");
return provider with { TokenizerPath = syncedTokenizerPath };
}
private static async Task<EmbeddingProvider> SyncEmbeddingTokenizerAsync(EmbeddingProvider provider, string pluginPath)
{
var syncedTokenizerPath = await SyncTokenizerAsync(
provider.TokenizerPath,
pluginPath,
TokenizerModelId.ForEmbeddingProvider(provider),
$"embedding provider '{provider.Name}'");
return provider with { TokenizerPath = syncedTokenizerPath };
}
private static async Task<string> SyncTokenizerAsync(string configuredTokenizerPath, string pluginPath, string modelId, string logName)
{
if (string.IsNullOrWhiteSpace(configuredTokenizerPath))
{
var deleteResult = await RustService.DeleteTokenizer(modelId);
if (!deleteResult.Success)
LOG.LogWarning("Failed to delete tokenizer for {LogName}: {Issue}", logName, deleteResult.Message);
return string.Empty;
}
var resolvedPath = ResolvePluginTokenizerPath(configuredTokenizerPath, pluginPath);
if (resolvedPath is null)
{
var deleteResult = await RustService.DeleteTokenizer(modelId);
if (!deleteResult.Success)
LOG.LogWarning("Failed to delete tokenizer after invalid path for {LogName}: {Issue}", logName, deleteResult.Message);
LOG.LogWarning("The configured tokenizer path '{TokenizerPath}' for {LogName} is invalid. The tokenizer path must stay within the plugin directory '{PluginPath}'.", configuredTokenizerPath, logName, pluginPath);
return string.Empty;
}
var validateResult = await RustService.ValidateTokenizer(resolvedPath);
if (!validateResult.Success)
{
var deleteResult = await RustService.DeleteTokenizer(modelId);
if (!deleteResult.Success)
LOG.LogWarning("Failed to delete tokenizer after validation failure for {LogName}: {Issue}", logName, deleteResult.Message);
LOG.LogWarning("The configured tokenizer for {LogName} is invalid. Path='{TokenizerPath}', issue='{Issue}'", logName, resolvedPath, validateResult.Message);
return string.Empty;
}
var storeResult = await RustService.StoreTokenizer(modelId, resolvedPath);
if (!storeResult.Success)
{
LOG.LogWarning("Failed to store tokenizer for {LogName}. Path='{TokenizerPath}', issue='{Issue}'", logName, resolvedPath, storeResult.Message);
return string.Empty;
}
return storeResult.StoredPath;
}
private static string? ResolvePluginTokenizerPath(string configuredTokenizerPath, string pluginPath)
{
if (string.IsNullOrWhiteSpace(pluginPath))
return null;
var fullPluginPath = Path.GetFullPath(pluginPath);
var candidatePath = Path.GetFullPath(Path.Combine(fullPluginPath, configuredTokenizerPath));
if (candidatePath.Equals(fullPluginPath, StringComparison.OrdinalIgnoreCase))
return null;
var pluginPrefix = fullPluginPath.EndsWith(Path.DirectorySeparatorChar)
? fullPluginPath
: fullPluginPath + Path.DirectorySeparatorChar;
return candidatePath.StartsWith(pluginPrefix, StringComparison.OrdinalIgnoreCase)
? candidatePath
: null;
}
}
@@ -177,4 +177,4 @@ public static partial class PluginFactory
HOT_RELOAD_SEMAPHORE.Release();
}
}
}
}
@@ -43,22 +43,40 @@ public sealed class AugmentationOne : IAugmentationProcess
{
// Let's get the validation agent & set up its provider:
var validationAgent = Program.SERVICE_PROVIDER.GetService<AgentRetrievalContextValidation>()!;
validationAgent.SetLLMProvider(provider);
// Let's validate all retrieval contexts:
var validationResults = await validationAgent.ValidateRetrievalContextsAsync(lastUserPrompt, chatThread, retrievalContexts, token);
//
// Now, filter the retrieval contexts to the most relevant ones:
//
var targetWindow = validationResults.DetermineTargetWindow(TargetWindowStrategy.TOP10_BETTER_THAN_GUESSING);
var threshold = validationResults.GetConfidenceThreshold(targetWindow);
// Filter the retrieval contexts:
retrievalContexts = validationResults.Where(x => x.RetrievalContext is not null && x.Confidence >= threshold).Select(x => x.RetrievalContext!).ToList();
if (validationAgent.SetLLMProvider(provider, chatThread.DataSecurity, chatThread.RequiredProviderConfidence))
{
try
{
// Let's validate all retrieval contexts:
var validationResults = await validationAgent.ValidateRetrievalContextsAsync(lastUserPrompt, chatThread, retrievalContexts, token);
if (validationResults.Count == 0)
LOGGER.LogWarning("Retrieval context validation returned no results. Continuing augmentation with all retrieved contexts.");
else
{
//
// Now, filter the retrieval contexts to the most relevant ones:
//
var targetWindow = validationResults.DetermineTargetWindow(TargetWindowStrategy.TOP10_BETTER_THAN_GUESSING);
var threshold = validationResults.GetConfidenceThreshold(targetWindow);
// Filter the retrieval contexts:
retrievalContexts = validationResults.Where(x => x.RetrievalContext is not null && x.Confidence >= threshold).Select(x => x.RetrievalContext!).ToList();
}
}
catch (OperationCanceledException) when (token.IsCancellationRequested)
{
throw;
}
catch (Exception exception)
{
LOGGER.LogError(exception, "Retrieval context validation failed. Continuing augmentation with all retrieved contexts.");
}
}
else
LOGGER.LogWarning("Skipping retrieval context validation because no sufficiently trusted validation agent provider is available. Continuing augmentation with all retrieved contexts.");
}
LOGGER.LogInformation($"Starting the augmentation process over {numTotalRetrievalContexts:###,###,###,###} retrieval contexts.");
LOGGER.LogInformation($"Starting the augmentation process over {retrievalContexts.Count:###,###,###,###} of {numTotalRetrievalContexts:###,###,###,###} retrieved contexts.");
//
// We build a huge prompt from all retrieval contexts:
@@ -31,11 +31,9 @@ public class AgenticSrcSelWithDynHeur : IDataSourceSelectionProcess
IReadOnlyList<IDataSource> selectedDataSources = [];
IReadOnlyList<DataSourceAgentSelected> finalAISelection = [];
// Get the settings manager:
var settings = Program.SERVICE_PROVIDER.GetService<SettingsManager>()!;
// Get the agent for the data source selection:
var selectionAgent = Program.SERVICE_PROVIDER.GetService<AgentDataSourceSelection>()!;
var allowedDataSources = dataSources.AllowedDataSources.ToDictionary(ds => ds.Id, StringComparer.Ordinal);
try
{
@@ -61,14 +59,14 @@ public class AgenticSrcSelWithDynHeur : IDataSourceSelectionProcess
var totalAISelectedDataSources = aiSelectedDataSources.Count;
// Filter out the data sources that are not available:
aiSelectedDataSources = aiSelectedDataSources.Where(x => settings.ConfigurationData.DataSources.FirstOrDefault(ds => ds.Id == x.Id) is not null).ToList();
aiSelectedDataSources = aiSelectedDataSources.Where(x => allowedDataSources.ContainsKey(x.Id)).ToList();
// Store the real AI-selected data sources:
finalAISelection = aiSelectedDataSources.Select(x => new DataSourceAgentSelected { DataSource = settings.ConfigurationData.DataSources.First(ds => ds.Id == x.Id), AIDecision = x, Selected = false }).ToList();
finalAISelection = aiSelectedDataSources.Select(x => new DataSourceAgentSelected { DataSource = allowedDataSources[x.Id], AIDecision = x, Selected = false }).ToList();
var numHallucinatedSources = totalAISelectedDataSources - aiSelectedDataSources.Count;
if (numHallucinatedSources > 0)
LOGGER.LogWarning($"The AI hallucinated {numHallucinatedSources} data source(s). We ignore them.");
LOGGER.LogWarning($"The AI selected {numHallucinatedSources} unavailable data source(s). We ignore them.");
if (aiSelectedDataSources.Count > 3)
{
@@ -87,7 +85,7 @@ public class AgenticSrcSelWithDynHeur : IDataSourceSelectionProcess
LOGGER.LogInformation($"The AI selected {aiSelectedDataSources.Count} data source(s) with a confidence of at least {threshold}.");
// Transform the final data sources to the actual data sources:
selectedDataSources = aiSelectedDataSources.Select(x => settings.ConfigurationData.DataSources.FirstOrDefault(ds => ds.Id == x.Id)).Where(ds => ds is not null).ToList()!;
selectedDataSources = aiSelectedDataSources.Select(x => allowedDataSources[x.Id]).ToList();
return new(proceedWithRAG, selectedDataSources);
}
@@ -96,7 +94,7 @@ public class AgenticSrcSelWithDynHeur : IDataSourceSelectionProcess
//
// Transform the selected data sources to the actual data sources:
selectedDataSources = aiSelectedDataSources.Select(x => settings.ConfigurationData.DataSources.FirstOrDefault(ds => ds.Id == x.Id)).Where(ds => ds is not null).ToList()!;
selectedDataSources = aiSelectedDataSources.Select(x => allowedDataSources[x.Id]).ToList();
// Mark the data sources as selected:
foreach (var dataSource in finalAISelection)
@@ -35,7 +35,13 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess
//
// 1. Check if the user wants to bind any data sources to the chat:
//
if (chatThread.DataSourceOptions.IsEnabled())
//
// Data sources are a preview feature. The check belongs here rather than in the options
// themselves: a chat keeps its data source options while the feature is switched off, and
// organizations may preselect data sources through a configuration plugin. Without this,
// such a chat would still run the entire RAG process with the feature disabled.
//
if (PreviewFeatures.PRE_RAG_2024.IsEnabled(settings) && chatThread.DataSourceOptions.IsEnabled())
{
LOGGER.LogInformation("Data sources are enabled for this chat.");
@@ -74,7 +80,7 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess
// data sources changed its security requirements.
//
List<IDataSource> preselectedDataSources = chatThread.DataSourceOptions.PreselectedDataSourceIds.Select(id => settings.ConfigurationData.DataSources.FirstOrDefault(ds => ds.Id == id)).Where(ds => ds is not null).ToList()!;
var dataSources = await dataSourceService.GetDataSources(provider, preselectedDataSources);
var dataSources = await dataSourceService.GetDataSources(provider, chatThread.DataSourceOptions, preselectedDataSources);
var selectedDataSources = dataSources.SelectedDataSources;
//
@@ -92,7 +98,7 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess
//
// No, the user made the choice manually:
//
var selectedDataSourceInfo = selectedDataSources.Select(ds => ds.Name).Aggregate((a, b) => $"'{a}', '{b}'");
var selectedDataSourceInfo = string.Join(", ", selectedDataSources.Select(ds => $"'{ds.Name}'"));
LOGGER.LogInformation($"The user selected the data sources manually. {selectedDataSources.Count} data source(s) are selected: {selectedDataSourceInfo}.");
}
@@ -104,12 +110,15 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess
else
{
var previousDataSecurity = chatThread.DataSecurity;
var previousRequiredProviderConfidence = chatThread.RequiredProviderConfidence;
//
// Update the data security of the chat thread. We consider the current data security
// of the chat thread and the data security of the selected data sources:
//
var dataSecurityRestrictedToSelfHosted = selectedDataSources.Any(x => x.SecurityPolicy is DataSourceSecurity.SELF_HOSTED);
var dataSecurityRestrictedToSelfHosted = selectedDataSources
.OfType<IExternalDataSource>()
.Any(dataSource => dataSource.SecurityPolicy is DataSourceSecurity.SELF_HOSTED);
chatThread.DataSecurity = dataSecurityRestrictedToSelfHosted switch
{
//
@@ -150,6 +159,12 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess
if (previousDataSecurity != chatThread.DataSecurity)
LOGGER.LogInformation($"The data security of the chat thread was updated from '{previousDataSecurity}' to '{chatThread.DataSecurity}'.");
foreach (var dataSource in selectedDataSources.OfType<IInternalDataSource>())
chatThread.RequireProviderConfidence(dataSource.ConfidenceLevel);
if (previousRequiredProviderConfidence != chatThread.RequiredProviderConfidence)
LOGGER.LogInformation($"The required provider confidence of the chat thread was updated from '{previousRequiredProviderConfidence.GetName()}' to '{chatThread.RequiredProviderConfidence.GetName()}'.");
}
//
@@ -205,17 +220,7 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess
var ragSources = new List<ISource>();
foreach (var retrievalContext in dataContexts)
{
var title = retrievalContext.DataSourceName;
if(string.IsNullOrWhiteSpace(title))
continue;
var link = retrievalContext.Path;
if(!link.StartsWith("http", StringComparison.OrdinalIgnoreCase))
continue;
ragSources.Add(new Source(title, link, SourceOrigin.RAG));
}
ragSources.AddRange(CreateSources(retrievalContext));
// Merge the sources, avoiding duplicates:
aiAnswerSources.MergeSources(ragSources);
@@ -225,4 +230,63 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess
}
#endregion
}
private static IReadOnlyList<ISource> CreateSources(IRetrievalContext retrievalContext)
{
var sources = new List<ISource>();
AddSource(sources, GetReferenceTitle(retrievalContext), GetReferenceLink(retrievalContext));
foreach (var link in retrievalContext.Links)
AddSource(sources, retrievalContext.DataSourceName, link);
return sources;
}
private static void AddSource(ICollection<ISource> sources, string title, string link)
{
if (string.IsNullOrWhiteSpace(title) || !TryNormalizeSourceLink(link, out var normalizedLink))
return;
sources.Add(new Source(title, normalizedLink, SourceOrigin.RAG));
}
private static string GetReferenceTitle(IRetrievalContext retrievalContext) =>
retrievalContext is RetrievalTextContext { ReferenceTitle: { Length: > 0 } referenceTitle }
? referenceTitle
: retrievalContext.DataSourceName;
private static string GetReferenceLink(IRetrievalContext retrievalContext) =>
retrievalContext is RetrievalTextContext { ReferenceLink: { Length: > 0 } referenceLink }
? referenceLink
: retrievalContext.Path;
private static bool TryNormalizeSourceLink(string link, out string normalizedLink)
{
normalizedLink = string.Empty;
if (string.IsNullOrWhiteSpace(link))
return false;
if (Uri.TryCreate(link, UriKind.Absolute, out var absoluteUri) && IsSupportedSourceUri(absoluteUri))
{
normalizedLink = absoluteUri.AbsoluteUri;
return true;
}
try
{
if (!Path.IsPathRooted(link))
return false;
normalizedLink = new Uri(Path.GetFullPath(link)).AbsoluteUri;
return true;
}
catch
{
return false;
}
}
private static bool IsSupportedSourceUri(Uri uri) =>
string.Equals(uri.Scheme, Uri.UriSchemeHttp, StringComparison.OrdinalIgnoreCase)
|| string.Equals(uri.Scheme, Uri.UriSchemeHttps, StringComparison.OrdinalIgnoreCase)
|| string.Equals(uri.Scheme, Uri.UriSchemeFile, StringComparison.OrdinalIgnoreCase);
}
@@ -40,4 +40,14 @@ public sealed class RetrievalTextContext : IRetrievalContext
/// For example, one sentence or paragraph before and after the matched text.
/// </remarks>
public IReadOnlyList<string> SurroundingContent { get; set; } = [];
/// <summary>
/// Optional title used when this context is displayed as a source reference.
/// </summary>
public string ReferenceTitle { get; init; } = string.Empty;
/// <summary>
/// Optional link used when this context is displayed as a source reference.
/// </summary>
public string ReferenceLink { get; init; } = string.Empty;
}
@@ -0,0 +1,41 @@
namespace AIStudio.Tools.Rust;
/// <summary>
/// Represents a file type that can optionally contain child file types.
/// Use the static helpers <see cref="Leaf"/>, <see cref="Parent"/> and <see cref="Composite"/> to build readable trees.
/// </summary>
/// <param name="FilterName">Display name of the type (e.g., "Document").</param>
/// <param name="FilterExtensions">File extensions belonging to this type (without dot).</param>
/// <param name="Children">Nested file types that are included when this type is selected.</param>
public sealed record FileType(string FilterName, string[] FilterExtensions, IReadOnlyList<FileType> Children)
{
/// <summary>
/// Factory for a leaf node.
/// Example: <c>FileType.Leaf(".NET", "cs", "razor")</c>
/// </summary>
public static FileType Leaf(string name, params string[] extensions) =>
new(name, extensions, []);
/// <summary>
/// Factory for a parent node that only has children.
/// Example: <c>FileType.Parent("Source Code", dotnet, java)</c>
/// </summary>
public static FileType Parent(string name, params FileType[]? children) =>
new(name, [], children ?? []);
/// <summary>
/// Factory for a composite node that has its own extensions in addition to children.
/// </summary>
public static FileType Composite(string name, string[] extensions, params FileType[] children) =>
new(name, extensions, children);
/// <summary>
/// Collects all extensions for this type, including children.
/// </summary>
public IEnumerable<string> FlattenExtensions()
{
return this.FilterExtensions
.Concat(this.Children.SelectMany(child => child.FlattenExtensions()))
.Distinct(StringComparer.OrdinalIgnoreCase);
}
}
+24 -11
View File
@@ -31,13 +31,18 @@ public static class FileTypes
public static readonly FileTypeFilter RUST = FileTypeFilter.Leaf("Rust", "rs");
public static readonly FileTypeFilter LUA = FileTypeFilter.Leaf("Lua", "lua");
public static readonly FileTypeFilter PHP = FileTypeFilter.Leaf("PHP", "php");
public static readonly FileTypeFilter WEB = FileTypeFilter.Leaf("HTML/CSS", "html", "css");
public static readonly FileTypeFilter HTML = FileTypeFilter.Leaf("HTML", "html", "htm");
public static readonly FileTypeFilter CSS = FileTypeFilter.Leaf("CSS", "css");
public static readonly FileTypeFilter WEB = FileTypeFilter.Parent("HTML/CSS", HTML, CSS);
/// <summary>
/// Gets the standalone HTML filter used for visual briefing import and export.
/// </summary>
public static readonly FileTypeFilter VISUAL_BRIEFING_HTML = FileTypeFilter.Leaf(TB("Visual briefing"), "html");
public static readonly FileTypeFilter HTML = FileTypeFilter.Leaf("HTML", "html");
// Only the canonical extension, without the legacy ".htm": this is what we write when
// exporting, whereas the HTML family above is what we accept when reading.
public static readonly FileTypeFilter HTML_DOCUMENT = FileTypeFilter.Leaf("HTML", "html");
public static readonly FileTypeFilter APP = FileTypeFilter.Leaf("Swift/Kotlin", "swift", "kt");
public static readonly FileTypeFilter SHELL = FileTypeFilter.Leaf("Shell", "sh", "bash", "zsh");
public static readonly FileTypeFilter LOG = FileTypeFilter.Leaf("Log", "log");
@@ -59,8 +64,10 @@ public static class FileTypes
public static readonly FileTypeFilter MS_WORD = FileTypeFilter.Leaf("Microsoft Word", "docx");
public static readonly FileTypeFilter ODT = FileTypeFilter.Leaf("OpenDocument Text", "odt");
public static readonly FileTypeFilter WORD = FileTypeFilter.Parent("Word", ODT, MS_WORD);
public static readonly FileTypeFilter EXCEL = FileTypeFilter.Leaf("Excel", "xls", "xlsx");
public static readonly FileTypeFilter EXCEL = FileTypeFilter.Leaf("Excel", "xls", "xlsx", "xlsm", "xlsb", "xla", "xlam");
public static readonly FileTypeFilter ODS = FileTypeFilter.Leaf("OpenDocument Spreadsheet", "ods");
public static readonly FileTypeFilter SPREADSHEET = FileTypeFilter.Parent(TB("Spreadsheet"), EXCEL, ODS);
// The legacy binary ".ppt" is missing on purpose: AI Studio has no reader for it, so offering
// it would only let users attach a file which cannot be read.
public static readonly FileTypeFilter POWER_POINT = FileTypeFilter.Leaf("PowerPoint", "pptx", "odp");
@@ -72,13 +79,13 @@ public static class FileTypes
public static readonly FileTypeFilter TEX = FileTypeFilter.Leaf("LaTeX", "tex");
public static readonly FileTypeFilter OFFICE_FILES = FileTypeFilter.Parent(TB("Office Files"),
WORD, EXCEL, POWER_POINT, PDF);
WORD, SPREADSHEET, POWER_POINT, PDF);
public static readonly FileTypeFilter DOCUMENT = FileTypeFilter.Parent(TB("Document"),
TEXT, TABULAR, OFFICE_FILES, SOURCE_CODE, LATEX);
// Media hierarchy
public static readonly FileTypeFilter IMAGE = FileTypeFilter.Leaf(TB("Image"),
"jpg", "jpeg", "png", "gif", "bmp", "tiff", "svg", "webp", "heic");
"jpg", "jpeg", "png", "gif", "bmp", "tiff", "svg", "webp", "heic", "avif");
/// <summary>
/// Gets the prototype visual-asset image formats.
@@ -95,6 +102,7 @@ public static class FileTypes
// Other standalone types
public static readonly FileTypeFilter CERTIFICATE_BUNDLE = FileTypeFilter.Leaf(TB("Certificate bundle"), "pem", "crt", "cer");
public static readonly FileTypeFilter EXECUTABLES = FileTypeFilter.Leaf(TB("Executable"), "exe", "app", "bin", "appimage");
public static readonly FileTypeFilter SHORTCUT = FileTypeFilter.Leaf(TB("Shortcut"), "lnk");
public static readonly FileTypeFilter PLUGIN_ARCHIVE = FileTypeFilter.Leaf(TB("Plugin archive"), PluginArchive.PLUGIN_FILE_EXTENSION.TrimStart('.'), "zip");
/// <summary>
@@ -140,6 +148,14 @@ public static class FileTypes
.ToArray();
}
public static bool IsAllowedExtension(string extension, params FileTypeFilter[]? types)
{
if (types == null || types.Length == 0 || string.IsNullOrWhiteSpace(extension))
return false;
return OnlyAllowTypes(types).Contains(extension.TrimStart('.'), StringComparer.OrdinalIgnoreCase);
}
/// <summary>
/// Validates a file path against the provided filters.
/// Supports extension-based matching and source-like file names (e.g. Dockerfile).
@@ -150,11 +166,8 @@ public static class FileTypes
return false;
var extension = Path.GetExtension(filePath).TrimStart('.');
if (!string.IsNullOrWhiteSpace(extension))
{
if (OnlyAllowTypes(types).Contains(extension, StringComparer.OrdinalIgnoreCase))
return true;
}
if (IsAllowedExtension(extension, types))
return true;
var fileName = Path.GetFileName(filePath);
if (string.IsNullOrWhiteSpace(fileName))
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Rust;
public readonly record struct TokenizerHandlingResponse(int Success, string Response);
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Rust;
public readonly record struct TokenizerResponse(bool Success, int TokenCount, string Message, string StoredPath = "");
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Services;
public sealed record ArbitraryFileDataSegment(string Content, int TokenCount);
@@ -0,0 +1,24 @@
using System.Net;
using AIStudio.Provider;
namespace AIStudio.Tools.Services;
/// <summary>
/// One input which could not be embedded, together with everything known about why.
/// </summary>
/// <remarks>
/// The reason alone used to be all we kept, which made every failure look alike in the UI: a
/// rejected API key, an unreachable provider, and a file nobody may read were one and the same
/// list entry. The surrounding fields are what lets the UI offer the matching way out.
/// </remarks>
/// <param name="FilePath">The file that failed or the name of the data source when the failure was not about one file.</param>
/// <param name="Reason">What to tell the user about it, ready to show.</param>
/// <param name="OccurredAtUtc">When it happened, so the list still makes sense when the user looks at it later.</param>
/// <param name="FailureReason">What kind of failure it was. Everything that did not come from a provider stays at NONE.</param>
/// <param name="StatusCode">What the provider answered, where it answered at all.</param>
/// <param name="EmbeddingProviderName">The embedding provider that was asked.</param>
/// <param name="ExtractionCode">Why reading the file failed, where the failure was about reading it at all.</param>
/// <param name="IsPermanent">Whether the file stays out of the index until it changes.</param>
public sealed record DataSourceEmbeddingFailure(string FilePath, string Reason, DateTimeOffset OccurredAtUtc, ProviderRequestFailureReason FailureReason = ProviderRequestFailureReason.NONE,
HttpStatusCode? StatusCode = null, string EmbeddingProviderName = "", FileExtractionErrorCode ExtractionCode = FileExtractionErrorCode.NONE, bool IsPermanent = false);
@@ -0,0 +1,25 @@
namespace AIStudio.Tools.Services;
public sealed class DataSourceEmbeddingManifest
{
public string EmbeddingProviderId { get; set; } = string.Empty;
public string EmbeddingSignature { get; set; } = string.Empty;
public string SourceHash { get; set; } = string.Empty;
public int VectorSize { get; set; }
public Dictionary<string, EmbeddedFileRecord> Files { get; init; } = new(StringComparer.OrdinalIgnoreCase);
/// <summary>
/// The files whose indexing failed for a reason which lies in the file itself, keyed by their
/// absolute path.
/// </summary>
/// <remarks>
/// These files are not read again as long as their fingerprint stays the same. Without this,
/// a folder holding hundreds of scanned documents without a text layer would be read again on
/// every single run, with the outcome known in advance.
/// </remarks>
public Dictionary<string, PermanentIndexingFailureRecord> PermanentFailures { get; init; } = new(StringComparer.OrdinalIgnoreCase);
}
@@ -0,0 +1,12 @@
namespace AIStudio.Tools.Services;
internal static class DataSourceEmbeddingNames
{
public static string GetCollectionName(string dataSourceId)
{
if (!Guid.TryParse(dataSourceId, out var parsedDataSourceId))
throw new ArgumentException("Data source ID must be a valid GUID.", nameof(dataSourceId));
return $"rag_{parsedDataSourceId:N}";
}
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Services;
public sealed record DataSourceEmbeddingOverview(DataSourceEmbeddingState State, int IndexedFiles, int TotalFiles, int FailedFiles);
@@ -0,0 +1,18 @@
using System.Diagnostics.CodeAnalysis;
using AIStudio.Provider;
using AIStudio.Settings;
namespace AIStudio.Tools.Services;
internal static class DataSourceEmbeddingProviders
{
public static bool TryResolve(SettingsManager settingsManager, IDataSource dataSource, [NotNullWhen(true)] out EmbeddingProvider? embeddingProvider)
{
embeddingProvider = settingsManager.ConfigurationData.EmbeddingProviders.FirstOrDefault(provider =>
dataSource is IInternalDataSource internalDataSource &&
provider.Id.Equals(internalDataSource.EmbeddingId, StringComparison.OrdinalIgnoreCase));
return embeddingProvider != default && embeddingProvider.UsedLLMProvider is not LLMProviders.NONE;
}
}
File diff suppressed because it is too large. Load diff
@@ -0,0 +1,26 @@
using AIStudio.Tools.Databases.IndexStore;
using AIStudio.Tools.Databases.VectorStore;
namespace AIStudio.Tools.Services;
public sealed partial class DataSourceEmbeddingService
{
private async Task ResetPersistedStateAsync(
string dataSourceId,
VectorStoreClient? vectorStore,
IndexStoreClient? indexStore,
CancellationToken token)
{
await this.DeleteCollectionAsync(DataSourceEmbeddingNames.GetCollectionName(dataSourceId), vectorStore, token);
indexStore ??= await databaseClientProvider.GetIndexStoreAsync(token);
if (!indexStore.IsAvailable)
{
logger.LogWarning("Could not delete local RAG embedding state for data source '{DataSourceId}' because the database '{DatabaseName}' is unavailable.", dataSourceId, indexStore.Name);
return;
}
await indexStore.DeleteDataSourceAsync(dataSourceId, token);
logger.LogInformation("Reset persisted local RAG embedding state for data source '{DataSourceId}'.", dataSourceId);
}
}
@@ -0,0 +1,289 @@
using System.Collections.Concurrent;
using AIStudio.Settings;
using AIStudio.Settings.DataModel;
namespace AIStudio.Tools.Services;
public sealed partial class DataSourceEmbeddingService
{
private const int WATCHER_DEBOUNCE_SECONDS = 2;
private readonly ConcurrentDictionary<string, DataSourceWatcherRegistration> watchers = new(StringComparer.OrdinalIgnoreCase);
private readonly Dictionary<string, CancellationTokenSource> watcherDebounceTokens = new(StringComparer.OrdinalIgnoreCase);
private readonly object watcherDebounceLock = new();
private void RefreshWatchers()
{
if (!settingsManager.ConfigurationData.App.DataSourceIndexing.AutomaticRefresh)
{
this.RemoveAllWatchers();
return;
}
if (Volatile.Read(ref this.startupHashCheckCompleted) == 0)
{
logger.LogDebug("File watchers are not activated yet because the startup persisted hash check has not completed.");
this.RemoveAllWatchers();
return;
}
var supportedSources = settingsManager.ConfigurationData.DataSources
.Where(this.IsSupportedInternalDataSource)
.ToDictionary(source => source.Id, StringComparer.OrdinalIgnoreCase);
foreach (var existingWatcherId in this.watchers.Keys.Except(supportedSources.Keys, StringComparer.OrdinalIgnoreCase).ToList())
this.RemoveWatcher(existingWatcherId);
foreach (var dataSource in supportedSources.Values)
this.EnsureWatcher(dataSource);
}
private void EnsureWatcher(IDataSource dataSource)
{
if (!settingsManager.ConfigurationData.App.DataSourceIndexing.AutomaticRefresh)
return;
var configuration = GetWatchConfiguration(dataSource);
if (configuration is null)
return;
if (this.watchers.TryGetValue(dataSource.Id, out var existingRegistration))
{
if (IsSameWatchConfiguration(existingRegistration.Configuration, configuration))
return;
this.RemoveWatcher(dataSource.Id);
}
var watcher = this.CreateWatcher(dataSource.Id, configuration);
if (watcher is null)
return;
if (!this.watchers.TryAdd(dataSource.Id, new DataSourceWatcherRegistration(watcher, configuration)))
watcher.Dispose();
}
private FileSystemWatcher? CreateWatcher(string dataSourceId, DataSourceWatcherConfiguration configuration)
{
try
{
var watcher = new FileSystemWatcher(configuration.RootPath)
{
Filter = configuration.Filter,
IncludeSubdirectories = configuration.IncludeSubdirectories,
NotifyFilter = NotifyFilters.FileName | NotifyFilters.DirectoryName | NotifyFilters.LastWrite | NotifyFilters.CreationTime | NotifyFilters.Size,
};
watcher.Changed += (_, args) => this.OnWatchedDataSourceChanged(dataSourceId, configuration, args);
watcher.Deleted += (_, args) => this.OnWatchedDataSourceChanged(dataSourceId, configuration, args);
watcher.Created += (_, args) => this.OnWatchedDataSourceChanged(dataSourceId, configuration, args);
watcher.Renamed += (_, args) => this.OnWatchedDataSourceChanged(dataSourceId, configuration, args);
watcher.Error += (_, args) =>
{
logger.LogWarning(args.GetException(), "The file watcher for data source '{DataSourceId}' failed. Recreating it.", dataSourceId);
this.RemoveWatcher(dataSourceId);
this.EnsureWatcher(dataSourceId);
this.ScheduleWatchedDataSourceRefresh(dataSourceId);
};
watcher.EnableRaisingEvents = true;
return watcher;
}
catch (Exception exception)
{
logger.LogWarning(exception, "Failed to create file watcher for data source '{DataSourceId}' at '{RootPath}'.", dataSourceId, configuration.RootPath);
return null;
}
}
private void RemoveWatcher(string dataSourceId)
{
this.CancelPendingWatcherRefresh(dataSourceId);
if (this.watchers.TryRemove(dataSourceId, out var registration))
registration.Watcher.Dispose();
}
private void RemoveAllWatchers()
{
foreach (var watcherId in this.watchers.Keys.ToList())
this.RemoveWatcher(watcherId);
}
private void DisposeWatchers()
{
this.CancelAllPendingWatcherRefreshes();
foreach (var registration in this.watchers.Values)
registration.Watcher.Dispose();
this.watchers.Clear();
}
private void OnWatchedDataSourceChanged(string dataSourceId, DataSourceWatcherConfiguration configuration, FileSystemEventArgs args)
{
if (!this.IsRelevantWatcherEvent(configuration, args))
{
logger.LogDebug(
"Ignoring file system change for data source '{DataSourceId}' at '{Path}' (event={ChangeType}) because the path is not part of the RAG index.",
dataSourceId,
args.FullPath,
args.ChangeType);
return;
}
logger.LogDebug(
"Detected relevant file system change for data source '{DataSourceId}' at '{Path}' (event={ChangeType}). Scheduling a debounced embedding run.",
dataSourceId,
args.FullPath,
args.ChangeType);
this.ScheduleWatchedDataSourceRefresh(dataSourceId);
}
private void ScheduleWatchedDataSourceRefresh(string dataSourceId)
{
if (!settingsManager.ConfigurationData.App.DataSourceIndexing.AutomaticRefresh)
return;
var debounceToken = new CancellationTokenSource();
lock (this.watcherDebounceLock)
{
if (this.watcherDebounceTokens.Remove(dataSourceId, out var existingToken))
existingToken.Cancel();
this.watcherDebounceTokens[dataSourceId] = debounceToken;
}
_ = Task.Run(async () =>
{
try
{
await Task.Delay(TimeSpan.FromSeconds(WATCHER_DEBOUNCE_SECONDS), debounceToken.Token);
if (!this.TryCompletePendingWatcherRefresh(dataSourceId, debounceToken))
return;
var dataSource = settingsManager.ConfigurationData.DataSources
.FirstOrDefault(source => source.Id.Equals(dataSourceId, StringComparison.OrdinalIgnoreCase));
if (dataSource is not null)
{
logger.LogInformation("Queueing data source '{DataSourceName}' ({DataSourceId}) after file system changes settled. The hash pipeline will reindex only changed files.", dataSource.Name, dataSource.Id);
await this.QueueDataSourceAsync(dataSource, true, DataSourceEmbeddingRefreshMode.WATCHER_HASH_CHECK);
}
}
catch (OperationCanceledException)
{
}
catch (Exception exception)
{
logger.LogWarning(exception, "Failed to queue watched data source '{DataSourceId}' after a file system change.", dataSourceId);
}
finally
{
debounceToken.Dispose();
}
});
}
private void EnsureWatcher(string dataSourceId)
{
var dataSource = settingsManager.ConfigurationData.DataSources
.FirstOrDefault(source => source.Id.Equals(dataSourceId, StringComparison.OrdinalIgnoreCase));
if (dataSource is not null)
this.EnsureWatcher(dataSource);
}
private void CancelPendingWatcherRefresh(string dataSourceId)
{
lock (this.watcherDebounceLock)
{
if (this.watcherDebounceTokens.Remove(dataSourceId, out var token))
token.Cancel();
}
}
private void CancelAllPendingWatcherRefreshes()
{
lock (this.watcherDebounceLock)
{
foreach (var token in this.watcherDebounceTokens.Values)
token.Cancel();
this.watcherDebounceTokens.Clear();
}
}
private bool TryCompletePendingWatcherRefresh(string dataSourceId, CancellationTokenSource debounceToken)
{
lock (this.watcherDebounceLock)
{
if (!this.watcherDebounceTokens.TryGetValue(dataSourceId, out var currentToken) || !ReferenceEquals(currentToken, debounceToken))
return false;
this.watcherDebounceTokens.Remove(dataSourceId);
return true;
}
}
private bool IsRelevantWatcherEvent(DataSourceWatcherConfiguration configuration, FileSystemEventArgs args)
{
if (args is RenamedEventArgs renamedArgs)
{
return this.IsRelevantWatcherPath(configuration, renamedArgs.FullPath, args.ChangeType)
|| this.IsRelevantWatcherPath(configuration, renamedArgs.OldFullPath, args.ChangeType);
}
return this.IsRelevantWatcherPath(configuration, args.FullPath, args.ChangeType);
}
private bool IsRelevantWatcherPath(DataSourceWatcherConfiguration configuration, string path, WatcherChangeTypes changeType)
{
if (string.IsNullOrWhiteSpace(path))
return false;
var fileName = Path.GetFileName(path);
if (string.IsNullOrWhiteSpace(fileName))
return true;
if (!configuration.IncludeSubdirectories && !string.Equals(fileName, configuration.Filter, StringComparison.OrdinalIgnoreCase))
return false;
if (Directory.Exists(path))
return true;
if (IsSkippedRagFileName(fileName))
return false;
if (this.IsSupportedRagFilePath(path))
return true;
return changeType is WatcherChangeTypes.Deleted or WatcherChangeTypes.Renamed
&& string.IsNullOrWhiteSpace(Path.GetExtension(path));
}
private static DataSourceWatcherConfiguration? GetWatchConfiguration(IDataSource dataSource) => dataSource switch
{
DataSourceLocalDirectory localDirectory when Directory.Exists(localDirectory.Path) => new DataSourceWatcherConfiguration(
localDirectory.Path,
"*.*",
true),
DataSourceLocalFile localFile when File.Exists(localFile.FilePath) && !string.IsNullOrWhiteSpace(Path.GetDirectoryName(localFile.FilePath)) => new DataSourceWatcherConfiguration(
Path.GetDirectoryName(localFile.FilePath)!,
Path.GetFileName(localFile.FilePath),
false),
_ => null,
};
private static bool IsSameWatchConfiguration(DataSourceWatcherConfiguration left, DataSourceWatcherConfiguration right)
{
return left.IncludeSubdirectories == right.IncludeSubdirectories
&& string.Equals(left.RootPath, right.RootPath, StringComparison.OrdinalIgnoreCase)
&& string.Equals(left.Filter, right.Filter, StringComparison.OrdinalIgnoreCase);
}
private sealed record DataSourceWatcherConfiguration(string RootPath, string Filter, bool IncludeSubdirectories);
private sealed record DataSourceWatcherRegistration(FileSystemWatcher Watcher, DataSourceWatcherConfiguration Configuration);
}
File diff suppressed because it is too large. Load diff
@@ -0,0 +1,10 @@
namespace AIStudio.Tools.Services;
public enum DataSourceEmbeddingState
{
IDLE,
QUEUED,
RUNNING,
COMPLETED,
FAILED,
}
@@ -0,0 +1,45 @@
using AIStudio.Settings.DataModel;
using AIStudio.Tools.PluginSystem;
namespace AIStudio.Tools.Services;
public sealed record DataSourceEmbeddingStatus(
string DataSourceId,
string DataSourceName,
DataSourceType DataSourceType,
DataSourceEmbeddingState State,
int TotalFiles,
int IndexedFiles,
int FailedFiles,
string CurrentFile,
string LastError,
IReadOnlyList<DataSourceEmbeddingFailure> Failures,
int PermanentlySkippedFiles = 0)
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(DataSourceEmbeddingStatus).Namespace, nameof(DataSourceEmbeddingStatus));
/// <remarks>
/// Files which were skipped for good are done, even though nothing was indexed of them.
/// Leaving them out would keep the bar short of the end for a data source which has nothing
/// left to do.
/// </remarks>
public int ProgressPercent => this.TotalFiles <= 0 ? 0 : Math.Clamp((int)Math.Round((this.IndexedFiles + this.PermanentlySkippedFiles) * 100d / this.TotalFiles), 0, 100);
public string StateLabel => this.State switch
{
DataSourceEmbeddingState.QUEUED => TB("Queued"),
DataSourceEmbeddingState.RUNNING => TB("Running"),
DataSourceEmbeddingState.COMPLETED => TB("Completed"),
DataSourceEmbeddingState.FAILED => TB("Needs attention"),
_ => TB("Idle")
};
public int SortOrder => this.State switch
{
DataSourceEmbeddingState.RUNNING => 0,
DataSourceEmbeddingState.QUEUED => 1,
DataSourceEmbeddingState.FAILED => 2,
DataSourceEmbeddingState.COMPLETED => 3,
_ => 4,
};
}
@@ -0,0 +1,506 @@
using AIStudio.Chat;
using AIStudio.Provider;
using AIStudio.Settings;
using AIStudio.Settings.DataModel;
using AIStudio.Tools.Databases;
using AIStudio.Tools.Databases.IndexStore;
using AIStudio.Tools.Databases.VectorStore;
using AIStudio.Tools.PluginSystem;
using AIStudio.Tools.RAG;
using AIStudio.Tools.Rust;
namespace AIStudio.Tools.Services;
public sealed class DataSourceLocalRetrievalService(
SettingsManager settingsManager, RustService rustService, DatabaseClientProvider databaseClientProvider,
ILogger<DataSourceLocalRetrievalService> logger)
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(DataSourceLocalRetrievalService).Namespace, nameof(DataSourceLocalRetrievalService));
//
// Which gaps the user was already told about in this session. Retrieval runs for every single
// message, so without this one broken embedding provider would put a warning on every prompt.
//
private readonly HashSet<string> reportedRetrievalGaps = new(StringComparer.Ordinal);
private readonly Lock retrievalGapLock = new();
private enum RetrievalChannel
{
VECTOR,
BM25,
}
//
// A hit keeps the complete shape both retrieval channels deliver, even where nothing reads a
// value yet. Merging is deterministic on purpose for now, so Channel, Score and Rank have no
// consumer until reranking arrives. Naming them still beats handing an unlabelled tuple of
// strings and numbers through the service.
//
// ReSharper disable NotAccessedPositionalProperty.Local
private sealed record LocalRetrievalHit(
RetrievalChannel Channel,
string ChunkId,
string ParentFileId,
string DataSourceId,
string DataSourceName,
string DataSourceType,
string AbsolutePath,
string FileName,
string RelativePath,
string FileType,
int? PageNumber,
int ChunkIndex,
string Text,
double Score,
int Rank,
string ConfidenceLevel,
int ConfidenceLevelRank);
// ReSharper restore NotAccessedPositionalProperty.Local
public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(DataSourceLocalFile dataSource, IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) =>
this.RetrieveDataAsync(dataSource, lastUserPrompt, token);
public Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(DataSourceLocalDirectory dataSource, IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) =>
this.RetrieveDataAsync(dataSource, lastUserPrompt, token);
private async Task<IReadOnlyList<IRetrievalContext>> RetrieveDataAsync(IInternalDataSource dataSource, IContent lastUserPrompt, CancellationToken token)
{
var query = GetQueryText(lastUserPrompt);
if (string.IsNullOrWhiteSpace(query))
{
logger.LogDebug("Skipping local retrieval for data source '{DataSourceName}' ({DataSourceId}) because the latest prompt does not contain text.", dataSource.Name, dataSource.Id);
return [];
}
var maxMatches = (int)dataSource.MaxMatches;
if (maxMatches == 0)
return [];
var collectionName = DataSourceEmbeddingNames.GetCollectionName(dataSource.Id);
var vectorTask = this.SearchVectorAsync(dataSource, query, maxMatches, collectionName, token);
var bm25Task = this.SearchBm25Async(dataSource, query, maxMatches, token);
await Task.WhenAll(vectorTask, bm25Task);
token.ThrowIfCancellationRequested();
var hits = MergeResults(vectorTask.Result, bm25Task.Result, maxMatches);
logger.LogInformation(
"Retrieved {MergedHits} local RAG hits for data source '{DataSourceName}' ({DataSourceId}). VectorCandidates={VectorHits}, BM25Candidates={BM25Hits}, RequestedPerChannel={RequestedPerChannel}.",
hits.Count,
dataSource.Name,
dataSource.Id,
vectorTask.Result.Count,
bm25Task.Result.Count,
maxMatches);
return hits
.Where(hit => !string.IsNullOrWhiteSpace(hit.Text))
.Select(ToRetrievalContext)
.ToList();
}
private async Task<IReadOnlyList<VectorSearchResult>> SearchVectorAsync(
IInternalDataSource dataSource,
string query,
int maxMatches,
string collectionName,
CancellationToken token)
{
try
{
var vectorStore = await databaseClientProvider.GetVectorStoreAsync(token);
if (!vectorStore.IsAvailable)
{
logger.LogWarning(
"Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because vector store '{VectorStoreName}' is unavailable.",
dataSource.Name,
dataSource.Id,
vectorStore.Name);
await this.ReportRetrievalGapAsync(dataSource, "no-vector-store", string.Format(TB("The data source '{0}' was left out of the answer: its local index is not available."), dataSource.Name));
return [];
}
if (!DataSourceEmbeddingProviders.TryResolve(settingsManager, dataSource, out var embeddingProvider))
{
logger.LogWarning("Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because the selected embedding provider is not available.", dataSource.Name, dataSource.Id);
await this.ReportRetrievalGapAsync(dataSource, "no-embedding-provider", string.Format(TB("The data source '{0}' was left out of the answer: its embedding provider is not available. Please check it in the settings."), dataSource.Name));
return [];
}
if (!await this.QueryFitsEmbeddingProviderAsync(dataSource, embeddingProvider, query, token))
return [];
var provider = embeddingProvider.CreateProvider();
var vectors = await provider.EmbedTextAsync(embeddingProvider.Model, settingsManager, token, [query]);
token.ThrowIfCancellationRequested();
var vector = vectors.FirstOrDefault();
if (vector is null || vector.Count == 0)
{
logger.LogWarning("Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because query embedding returned no vector.", dataSource.Name, dataSource.Id);
await this.ReportRetrievalGapAsync(dataSource, "no-query-vector", string.Format(TB("The data source '{0}' was left out of the answer: its embedding provider '{1}' did not return a vector for your message."), dataSource.Name, embeddingProvider.Name));
return [];
}
var results = this.LimitSearchResults(
dataSource,
"vector",
await vectorStore.SearchEmbeddingAsync(collectionName, vector, maxMatches, token),
maxMatches);
this.LogVectorResults(dataSource, results);
return results;
}
catch (OperationCanceledException) when (token.IsCancellationRequested)
{
throw;
}
catch (ProviderRequestException exception)
{
//
// The embedding provider named the cause and what to do about it. That sentence is
// worth far more to the user than the fact that a search came back empty:
//
logger.LogWarning(
exception,
"Vector retrieval failed for data source '{DataSourceName}' ({DataSourceId}) because the embedding provider failed. FailureReason={FailureReason}, StatusCode={StatusCode}.",
dataSource.Name, dataSource.Id, exception.FailureReason, exception.StatusCode);
await this.ReportRetrievalGapAsync(dataSource, $"provider-{exception.FailureReason}", string.Format(TB("The data source '{0}' was left out of the answer. {1}"), dataSource.Name, exception.UserMessage));
return [];
}
catch (Exception exception)
{
logger.LogWarning(exception, "Vector retrieval failed for data source '{DataSourceName}' ({DataSourceId}).", dataSource.Name, dataSource.Id);
await this.ReportRetrievalGapAsync(dataSource, "vector-search-failed", string.Format(TB("The data source '{0}' was left out of the answer because searching it failed."), dataSource.Name));
return [];
}
}
/// <summary>
/// Tells the user once that a data source cannot take part in answering.
/// </summary>
/// <remarks>
/// A failed search is not an error of the chat: the model still answers, only without what
/// this data source knows. Saying so once is what keeps somebody from trusting an answer
/// which was put together without half of its sources. Saying it with every prompt would be
/// worse than saying nothing, which is why every gap is reported once per session.
/// </remarks>
/// <param name="dataSource">The data source which could not be searched.</param>
/// <param name="gapKey">What kind of gap this is, so a different problem is reported again.</param>
/// <param name="userMessage">What to tell the user.</param>
private async Task ReportRetrievalGapAsync(IInternalDataSource dataSource, string gapKey, string userMessage)
{
lock (this.retrievalGapLock)
{
if (!this.reportedRetrievalGaps.Add($"{dataSource.Id}::{gapKey}"))
return;
}
await MessageBus.INSTANCE.SendWarning(new(Icons.Material.Filled.SearchOff, userMessage));
}
private async Task<bool> QueryFitsEmbeddingProviderAsync(
IInternalDataSource dataSource,
EmbeddingProvider embeddingProvider,
string query,
CancellationToken token)
{
var providerTokenLimit = Math.Max(1, embeddingProvider.EffectiveTokenLimit);
if (query.Length > RustService.MAX_TOKEN_COUNT_REQUEST_TEXT_LENGTH)
{
logger.LogWarning(
"Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because the latest prompt has {CharacterCount} characters and exceeds the safe tokenizer request length of {MaxCharacterCount}. ProviderTokenLimit={ProviderTokenLimit}.",
dataSource.Name,
dataSource.Id,
query.Length,
RustService.MAX_TOKEN_COUNT_REQUEST_TEXT_LENGTH,
providerTokenLimit);
await this.ReportRetrievalGapAsync(dataSource, "query-too-long", string.Format(TB("The data source '{0}' was left out of the answer because your message is too long to search with."), dataSource.Name));
return false;
}
var tokenCountResponse = await rustService.GetTokenCount(embeddingProvider, query, token);
if (tokenCountResponse is not { Success: true })
{
logger.LogWarning(
"Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because the token count for embedding provider '{EmbeddingProviderName}' could not be determined. Reason='{Reason}'.",
dataSource.Name,
dataSource.Id,
embeddingProvider.Name,
tokenCountResponse?.Message ?? "No response was returned by the tokenizer service.");
await this.ReportRetrievalGapAsync(dataSource, "no-token-count", string.Format(TB("The data source '{0}' was left out of the answer: the tokenizer of its embedding provider '{1}' is not available."), dataSource.Name, embeddingProvider.Name));
return false;
}
var queryTokenCount = tokenCountResponse.Value.TokenCount;
if (queryTokenCount > providerTokenLimit)
{
logger.LogWarning(
"Skipping vector retrieval for data source '{DataSourceName}' ({DataSourceId}) because the latest prompt has {QueryTokenCount} tokens, exceeding embedding provider '{EmbeddingProviderName}' limit of {ProviderTokenLimit} tokens.",
dataSource.Name,
dataSource.Id,
queryTokenCount,
embeddingProvider.Name,
providerTokenLimit);
await this.ReportRetrievalGapAsync(dataSource, "query-over-token-limit", string.Format(TB("The data source '{0}' was left out of the answer because your message is longer than its embedding provider '{1}' accepts."), dataSource.Name, embeddingProvider.Name));
return false;
}
return true;
}
private async Task<IReadOnlyList<IndexStoreSearchResult>> SearchBm25Async(IInternalDataSource dataSource, string query, int maxMatches, CancellationToken token)
{
try
{
var indexStore = await databaseClientProvider.GetIndexStoreAsync(token);
if (!indexStore.IsAvailable)
{
logger.LogWarning(
"Skipping BM25 retrieval for data source '{DataSourceName}' ({DataSourceId}) because local RAG index '{DatabaseName}' is unavailable.",
dataSource.Name,
dataSource.Id,
indexStore.Name);
return [];
}
var results = this.LimitSearchResults(
dataSource,
"BM25",
await indexStore.SearchChunksAsync(dataSource.Id, query, maxMatches, token),
maxMatches);
this.LogBm25Results(dataSource, results);
return results;
}
catch (OperationCanceledException) when (token.IsCancellationRequested)
{
throw;
}
catch (Exception exception)
{
logger.LogWarning(exception, "BM25 retrieval failed for data source '{DataSourceName}' ({DataSourceId}).", dataSource.Name, dataSource.Id);
return [];
}
}
private IReadOnlyList<T> LimitSearchResults<T>(IInternalDataSource dataSource, string searchName, IReadOnlyList<T> results, int maxMatches)
{
if (results.Count <= maxMatches)
return results;
logger.LogWarning(
"Local RAG {SearchName} search returned {ReturnedHits} chunks for data source '{DataSourceName}' ({DataSourceId}), which exceeds the configured maximum {MaxMatches}. Truncating to the datasource limit.",
searchName,
results.Count,
dataSource.Name,
dataSource.Id,
maxMatches);
return results.Take(maxMatches).ToList();
}
private static IReadOnlyList<LocalRetrievalHit> MergeResults(
IReadOnlyList<VectorSearchResult> vectorResults,
IReadOnlyList<IndexStoreSearchResult> bm25Results,
int maxMatches)
{
// Future reranking should replace this deterministic channel merge.
var merged = new List<LocalRetrievalHit>(maxMatches * 2);
var seenChunkIds = new HashSet<string>(StringComparer.OrdinalIgnoreCase);
AppendHits(
merged,
seenChunkIds,
vectorResults
.Select((result, index) => FromVectorResult(result, index + 1)),
maxMatches);
AppendHits(
merged,
seenChunkIds,
bm25Results
.Select((result, index) => FromBm25Result(result, index + 1)),
maxMatches);
return merged;
}
private static void AppendHits(List<LocalRetrievalHit> merged, HashSet<string> seenChunkIds, IEnumerable<LocalRetrievalHit> hits, int maxNewHits)
{
var added = 0;
foreach (var hit in hits)
{
if (!string.IsNullOrWhiteSpace(hit.ChunkId) && !seenChunkIds.Add(hit.ChunkId))
continue;
merged.Add(hit);
added++;
if (added >= maxNewHits)
return;
}
}
private static LocalRetrievalHit FromVectorResult(VectorSearchResult result, int rank) =>
new(
RetrievalChannel.VECTOR,
result.ChunkId,
result.ParentFileId,
result.DataSourceId,
result.DataSourceName,
result.DataSourceType,
FirstNonEmpty(result.AbsolutePath, result.FilePath),
result.FileName,
result.RelativePath,
result.FileType,
result.PageNumber,
result.ChunkIndex,
result.Text,
result.Score,
rank,
result.ConfidenceLevel,
result.ConfidenceLevelRank);
private static LocalRetrievalHit FromBm25Result(IndexStoreSearchResult result, int rank) =>
new(
RetrievalChannel.BM25,
result.ChunkId,
result.ParentFileId,
result.DataSourceId,
result.DataSourceName,
result.DataSourceType,
result.AbsolutePath,
result.FileName,
result.RelativePath,
result.FileType,
result.PageNumber,
result.ChunkIndex,
result.ChunkText,
result.Score,
rank,
result.ConfidenceLevel,
result.ConfidenceLevelRank);
private static RetrievalTextContext ToRetrievalContext(LocalRetrievalHit hit)
{
var sourceName = FirstNonEmpty(hit.FileName, hit.DataSourceName);
var path = FirstNonEmpty(hit.AbsolutePath, hit.RelativePath);
var referenceLink = string.IsNullOrWhiteSpace(path) ? string.Empty : BuildReferenceLink(path, hit);
return new RetrievalTextContext
{
DataSourceName = sourceName,
Category = RetrievalContentCategory.TEXT,
Type = GetRetrievalContentType(hit.FileType),
Path = path,
Links = [],
MatchedText = hit.Text,
SurroundingContent = [],
ReferenceTitle = BuildReferenceTitle(hit),
ReferenceLink = referenceLink,
};
}
private static string BuildReferenceTitle(LocalRetrievalHit hit)
{
var sourceName = FirstNonEmpty(hit.FileName, hit.DataSourceName);
return BuildLocatedReferenceTitle(sourceName, hit.ChunkIndex, hit.PageNumber);
}
private static string BuildLocatedReferenceTitle(string sourceName, int chunkIndex, int? pageNumber)
{
var location = pageNumber is > 0
? string.Format(TB("Page {0}"), pageNumber)
: string.Format(TB("Chunk {0}"), chunkIndex + 1);
return $"{sourceName} ({location})";
}
private static string BuildReferenceLink(string path, LocalRetrievalHit hit)
{
var link = NormalizeLocalReferencePath(path);
var separator = link.Contains('#', StringComparison.Ordinal) ? "&" : "#";
return $"{link}{separator}chunk={hit.ChunkIndex}";
}
private static string NormalizeLocalReferencePath(string path)
{
try
{
return Path.IsPathRooted(path)
? new Uri(Path.GetFullPath(path)).AbsoluteUri
: path;
}
catch
{
return path;
}
}
private static RetrievalContentType GetRetrievalContentType(string fileType)
{
if (FileTypes.IsAllowedExtension(fileType, FileTypes.TABULAR, FileTypes.SPREADSHEET))
return RetrievalContentType.TEXT_SPREADSHEET;
if (FileTypes.IsAllowedExtension(fileType, FileTypes.POWER_POINT))
return RetrievalContentType.TEXT_PRESENTATION;
return FileTypes.IsAllowedExtension(fileType, FileTypes.HTML)
? RetrievalContentType.TEXT_WEBSITE
: RetrievalContentType.TEXT_DOCUMENT;
}
private static string GetQueryText(IContent lastUserPrompt) => lastUserPrompt switch
{
ContentText text => text.Text,
_ => string.Empty
};
private static string FirstNonEmpty(params string[] values) =>
values.FirstOrDefault(value => !string.IsNullOrWhiteSpace(value)) ?? string.Empty;
private void LogVectorResults(IInternalDataSource dataSource, IReadOnlyList<VectorSearchResult> results)
{
if (results.Count == 0)
{
logger.LogInformation("Local RAG vector search found no chunks for data source '{DataSourceName}' ({DataSourceId}).", dataSource.Name, dataSource.Id);
return;
}
foreach (var result in results.Select((result, index) => (Result: result, Rank: index + 1)))
{
logger.LogInformation(
"Local RAG vector search found chunk for data source '{DataSourceName}' ({DataSourceId}). Rank={Rank}, Score={Score}, ChunkId='{ChunkId}', ParentFileId='{ParentFileId}', File='{FileName}', Path='{Path}', Title='{Title}'.",
dataSource.Name,
dataSource.Id,
result.Rank,
result.Result.Score,
result.Result.ChunkId,
result.Result.ParentFileId,
result.Result.FileName,
FirstNonEmpty(result.Result.AbsolutePath, result.Result.FilePath),
BuildLocatedReferenceTitle(FirstNonEmpty(result.Result.FileName, dataSource.Name), result.Result.ChunkIndex, result.Result.PageNumber));
}
}
private void LogBm25Results(IInternalDataSource dataSource, IReadOnlyList<IndexStoreSearchResult> results)
{
if (results.Count == 0)
{
logger.LogInformation("Local RAG BM25 search found no chunks for data source '{DataSourceName}' ({DataSourceId}).", dataSource.Name, dataSource.Id);
return;
}
foreach (var result in results.Select((result, index) => (Result: result, Rank: index + 1)))
{
logger.LogInformation(
"Local RAG BM25 search found chunk for data source '{DataSourceName}' ({DataSourceId}). Rank={Rank}, Score={Score}, ChunkId='{ChunkId}', ParentFileId='{ParentFileId}', File='{FileName}', Path='{Path}', Title='{Title}'.",
dataSource.Name,
dataSource.Id,
result.Rank,
result.Result.Score,
result.Result.ChunkId,
result.Result.ParentFileId,
result.Result.FileName,
result.Result.AbsolutePath,
BuildLocatedReferenceTitle(FirstNonEmpty(result.Result.FileName, dataSource.Name), result.Result.ChunkIndex, result.Result.PageNumber));
}
}
}
@@ -9,6 +9,15 @@ namespace AIStudio.Tools.Services;
public sealed class DataSourceService
{
//
// Trust is recorded for every participating provider, while only the chat provider's trust
// decides about external data sources below. The agent which validates retrieval contexts
// checks its own provider before it runs, so nothing slips through today. Keeping the value
// named here is what makes that asymmetry visible.
//
// ReSharper disable once NotAccessedPositionalProperty.Local
private readonly record struct ParticipatingProvider(string Role, bool IsTrusted, ConfidenceLevel ConfidenceLevel);
private readonly RustService rustService;
private readonly SettingsManager settingsManager;
private readonly ILogger<DataSourceService> logger;
@@ -27,9 +36,10 @@ public sealed class DataSourceService
/// It also returns the data sources selected before when they are still allowed.
/// </summary>
/// <param name="selectedLLMProvider">The selected LLM provider.</param>
/// <param name="dataSourceOptions">The active data source options, which determine which agent providers participate.</param>
/// <param name="previousSelectedDataSources">The data sources selected before.</param>
/// <returns>The allowed data sources and the data sources selected before -- when they are still allowed.</returns>
public async Task<AllowedSelectedDataSources> GetDataSources(AIStudio.Settings.Provider selectedLLMProvider, IReadOnlyCollection<IDataSource>? previousSelectedDataSources = null)
public async Task<AllowedSelectedDataSources> GetDataSources(AIStudio.Settings.Provider selectedLLMProvider, DataSourceOptions dataSourceOptions, IReadOnlyCollection<IDataSource>? previousSelectedDataSources = null)
{
//
// Case: Somehow the selected LLM provider was not set. The default provider
@@ -42,7 +52,10 @@ public sealed class DataSourceService
return new([], []);
}
return await this.GetDataSources(selectedLLMProvider.IsTrustedForDataSourceSecurityChecks(this.settingsManager), previousSelectedDataSources);
var usingTrustedProvider = selectedLLMProvider.IsTrustedForDataSourceSecurityChecks(this.settingsManager);
var participatingProviders = this.GetParticipatingProviders(selectedLLMProvider.Id, dataSourceOptions,
new("chat provider", usingTrustedProvider, selectedLLMProvider.GetConfidenceLevel(this.settingsManager)));
return await this.GetDataSources(usingTrustedProvider, participatingProviders, previousSelectedDataSources);
}
/// <summary>
@@ -51,9 +64,10 @@ public sealed class DataSourceService
/// this method checks only the supplied data sources.
/// </summary>
/// <param name="selectedLLMProvider">The selected LLM provider.</param>
/// <param name="dataSourceOptions">The active data source options, which determine which agent providers participate.</param>
/// <param name="requestedDataSources">The data sources to check.</param>
/// <returns>The requested data sources that are allowed for the provider.</returns>
public async Task<IReadOnlyList<IDataSource>> GetAllowedDataSources(AIStudio.Settings.Provider selectedLLMProvider, IReadOnlyCollection<IDataSource> requestedDataSources)
public async Task<IReadOnlyList<IDataSource>> GetAllowedDataSources(AIStudio.Settings.Provider selectedLLMProvider, DataSourceOptions dataSourceOptions, IReadOnlyCollection<IDataSource> requestedDataSources)
{
if (selectedLLMProvider == Settings.Provider.NONE)
{
@@ -61,7 +75,10 @@ public sealed class DataSourceService
return [];
}
return await this.GetAllowedDataSources(selectedLLMProvider.IsTrustedForDataSourceSecurityChecks(this.settingsManager), requestedDataSources);
var usingTrustedProvider = selectedLLMProvider.IsTrustedForDataSourceSecurityChecks(this.settingsManager);
var participatingProviders = this.GetParticipatingProviders(selectedLLMProvider.Id, dataSourceOptions,
new("chat provider", usingTrustedProvider, selectedLLMProvider.GetConfidenceLevel(this.settingsManager)));
return await this.GetAllowedDataSources(usingTrustedProvider, participatingProviders, requestedDataSources);
}
/// <summary>
@@ -69,9 +86,10 @@ public sealed class DataSourceService
/// It also returns the data sources selected before when they are still allowed.
/// </summary>
/// <param name="selectedLLMProvider">The selected LLM provider.</param>
/// <param name="dataSourceOptions">The active data source options, which determine which agent providers participate.</param>
/// <param name="previousSelectedDataSources">The data sources selected before.</param>
/// <returns>The allowed data sources and the data sources selected before -- when they are still allowed.</returns>
public async Task<AllowedSelectedDataSources> GetDataSources(IProvider selectedLLMProvider, IReadOnlyCollection<IDataSource>? previousSelectedDataSources = null)
public async Task<AllowedSelectedDataSources> GetDataSources(IProvider selectedLLMProvider, DataSourceOptions dataSourceOptions, IReadOnlyCollection<IDataSource>? previousSelectedDataSources = null)
{
//
// Case: Somehow the selected LLM provider was not set. The default provider
@@ -84,27 +102,60 @@ public sealed class DataSourceService
return new([], []);
}
return await this.GetDataSources(selectedLLMProvider.IsTrustedForDataSourceSecurityChecks(this.settingsManager), previousSelectedDataSources);
var usingTrustedProvider = selectedLLMProvider.IsTrustedForDataSourceSecurityChecks(this.settingsManager);
var participatingProviders = this.GetParticipatingProviders(selectedLLMProvider.ConfiguredProviderId, dataSourceOptions,
new("chat provider", usingTrustedProvider, selectedLLMProvider.GetConfidenceLevel(this.settingsManager)));
return await this.GetDataSources(usingTrustedProvider, participatingProviders, previousSelectedDataSources);
}
private async Task<AllowedSelectedDataSources> GetDataSources(bool usingTrustedProvider, IReadOnlyCollection<IDataSource>? previousSelectedDataSources = null)
private IReadOnlyList<ParticipatingProvider> GetParticipatingProviders(string currentProviderId, DataSourceOptions dataSourceOptions, ParticipatingProvider currentProvider)
{
var providers = new List<ParticipatingProvider> { currentProvider };
if (dataSourceOptions.AutomaticDataSourceSelection)
this.AddAgentProvider(providers, Components.AGENT_DATA_SOURCE_SELECTION, currentProviderId, "data source selection agent");
if (dataSourceOptions.AutomaticValidation && this.settingsManager.ConfigurationData.AgentRetrievalContextValidation.EnableRetrievalContextValidation)
this.AddAgentProvider(providers, Components.AGENT_RETRIEVAL_CONTEXT_VALIDATION, currentProviderId, "retrieval context validation agent");
return providers;
}
private void AddAgentProvider(List<ParticipatingProvider> providers, Components component, string currentProviderId, string role)
{
var provider = this.settingsManager.GetPreselectedProvider(component, currentProviderId, true);
if (provider == Settings.Provider.NONE)
{
this.logger.LogWarning($"No provider is available for the {role}. Data sources cannot be made available while this agent is enabled.");
providers.Add(new(role, false, ConfidenceLevel.NONE));
return;
}
providers.Add(new(
role,
provider.IsTrustedForDataSourceSecurityChecks(this.settingsManager),
provider.GetConfidenceLevel(this.settingsManager)));
}
private async Task<AllowedSelectedDataSources> GetDataSources(bool usingTrustedProvider, IReadOnlyList<ParticipatingProvider> participatingProviders, IReadOnlyCollection<IDataSource>? previousSelectedDataSources = null)
{
var allDataSources = this.settingsManager.ConfigurationData.DataSources.ToList();
var previousSelectedDataSourceIds = previousSelectedDataSources?.Select(source => source.Id).ToHashSet(StringComparer.Ordinal) ?? [];
var filteredDataSources = await this.GetAllowedDataSources(usingTrustedProvider, allDataSources);
var filteredDataSources = await this.GetAllowedDataSources(usingTrustedProvider, participatingProviders, allDataSources);
var filteredSelectedDataSources = filteredDataSources.Where(source => previousSelectedDataSourceIds.Contains(source.Id)).ToList();
return new(filteredDataSources, filteredSelectedDataSources);
}
private async Task<IReadOnlyList<IDataSource>> GetAllowedDataSources(bool usingTrustedProvider, IReadOnlyCollection<IDataSource> requestedDataSources)
private async Task<IReadOnlyList<IDataSource>> GetAllowedDataSources(bool usingTrustedProvider, IReadOnlyList<ParticipatingProvider> participatingProviders, IReadOnlyCollection<IDataSource> requestedDataSources)
{
var filteredDataSources = new List<IDataSource>(requestedDataSources.Count);
var tasks = new List<Task<IDataSource?>>(requestedDataSources.Count);
// Start all checks in parallel:
foreach (var source in requestedDataSources)
tasks.Add(this.CheckOneDataSource(source, usingTrustedProvider));
tasks.Add(this.CheckOneDataSource(source, usingTrustedProvider, participatingProviders));
// Wait for all checks and collect the results:
foreach (var task in tasks)
@@ -117,8 +168,35 @@ public sealed class DataSourceService
return filteredDataSources;
}
private async Task<IDataSource?> CheckOneDataSource(IDataSource source, bool usingTrustedProvider)
private async Task<IDataSource?> CheckOneDataSource(IDataSource source, bool usingTrustedProvider, IReadOnlyList<ParticipatingProvider> participatingProviders)
{
if (source is IInternalDataSource internalSource)
{
foreach (var provider in participatingProviders)
{
if (!provider.ConfidenceLevel.AllowsDataSourceConfidenceLevel(internalSource.ConfidenceLevel))
{
this.logger.LogWarning($"The internal data source '{source.Name}' (id={source.Id}) requires provider confidence '{internalSource.ConfidenceLevel.GetName()}'. The {provider.Role} only has confidence '{provider.ConfidenceLevel.GetName()}'. We skip this source.");
return null;
}
}
if (!DataSourceEmbeddingProviders.TryResolve(this.settingsManager, source, out var embeddingProvider))
{
this.logger.LogWarning($"The internal data source '{source.Name}' (id={source.Id}) has no usable embedding provider. We skip this source.");
return null;
}
var embeddingProviderConfidence = embeddingProvider.GetConfidenceLevel(this.settingsManager);
if (!embeddingProviderConfidence.AllowsDataSourceConfidenceLevel(internalSource.ConfidenceLevel))
{
this.logger.LogWarning($"The internal data source '{source.Name}' (id={source.Id}) requires provider confidence '{internalSource.ConfidenceLevel.GetName()}'. Its embedding provider '{embeddingProvider.Name}' only has confidence '{embeddingProviderConfidence.GetName()}'. We skip this source.");
return null;
}
return source;
}
//
// Unfortunately, we have to live-check any ERI source for its security requirements.
// Because the ERI server operator might change the security requirements at any time.
@@ -153,8 +231,11 @@ public sealed class DataSourceService
eriSourceRequirements = securityRequest.Data;
this.logger.LogInformation($"Security requirements for ERI source '{source.Name}' (id={source.Id}) retrieved successfully.");
}
switch (source.SecurityPolicy)
if (source is not IExternalDataSource externalSource)
return source;
switch (externalSource.SecurityPolicy)
{
case DataSourceSecurity.ALLOW_ANY:
@@ -228,4 +309,4 @@ public sealed class DataSourceService
return null;
}
}
}
}
@@ -187,10 +187,24 @@ public sealed class DirectChatService(SettingsManager settingsManager, DataSourc
requestedDataSources.Add(dataSource);
}
//
// The options the launched chat will run under. We build them here already, because the
// data-source check depends on them: they decide which agent providers take part, and an
// agent with too little confidence makes a data source unavailable.
//
var standardOptions = settingsManager.ConfigurationData.Chat.PreselectedDataSourceOptions;
var launchedDataSourceOptions = new DataSourceOptions
{
DisableDataSources = false,
AutomaticDataSourceSelection = false,
AutomaticValidation = standardOptions.AutomaticValidation,
PreselectedDataSourceIds = requestedDataSources.Select(source => source.Id).ToList(),
};
IReadOnlyList<IDataSource> availableDataSources;
try
{
availableDataSources = await dataSourceService.GetAllowedDataSources(provider, requestedDataSources);
availableDataSources = await dataSourceService.GetAllowedDataSources(provider, launchedDataSourceOptions, requestedDataSources);
}
catch (Exception exception)
{
@@ -203,13 +217,6 @@ public sealed class DirectChatService(SettingsManager settingsManager, DataSourc
if (unavailableDataSources.Count > 0)
return new(null, string.Format(TB("The following data sources selected by the assistant chat launcher are currently unavailable or not permitted for the selected provider: {0}"), string.Join(", ", unavailableDataSources)));
var standardOptions = settingsManager.ConfigurationData.Chat.PreselectedDataSourceOptions;
return new(new()
{
DisableDataSources = false,
AutomaticDataSourceSelection = false,
AutomaticValidation = standardOptions.AutomaticValidation,
PreselectedDataSourceIds = requestedDataSources.Select(source => source.Id).ToList(),
}, string.Empty);
return new(launchedDataSourceOptions, string.Empty);
}
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Services;
public sealed record EmbeddedFileRecord(string Fingerprint, long FileSize, DateTimeOffset LastWriteUtc, DateTimeOffset EmbeddedAtUtc, int ChunkCount);
@@ -0,0 +1,19 @@
namespace AIStudio.Tools.Services;
public sealed class FileEnumerationResult
{
public List<FileInfo> Files { get; } = [];
public List<DataSourceEmbeddingFailure> Failures { get; } = [];
public int FailedFiles { get; set; }
public string LastError { get; set; } = string.Empty;
public void AddFailure(string filePath, string reason)
{
this.Failures.Add(new DataSourceEmbeddingFailure(filePath, reason, DateTimeOffset.UtcNow));
this.FailedFiles = this.Failures.Count;
this.LastError = reason;
}
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Services;
public sealed record PermanentIndexingFailureRecord(string Fingerprint, FileExtractionErrorCode Code, string Message, DateTimeOffset OccurredAtUtc);
@@ -49,5 +49,22 @@ public sealed partial class RustService
throw new InvalidOperationException(operation?.Issue ?? $"The {databaseName} operation failed.");
}
public async Task<TResult?> ExecuteDatabaseQuery<TRequest, TResult>(string databaseName, string path, TRequest request, CancellationToken cancellationToken = default)
{
using var cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
cts.CancelAfter(TimeSpan.FromMinutes(5));
using var response = await this.http.PostAsJsonAsync(path, request, this.jsonRustSerializerOptions, cts.Token);
response.EnsureSuccessStatusCode();
var operation = await response.Content.ReadFromJsonAsync<DatabaseQueryResponse<TResult>>(this.jsonRustSerializerOptions, cts.Token);
if (operation is not { Success: true })
throw new InvalidOperationException(operation?.Issue ?? $"The {databaseName} query failed.");
return operation.Data;
}
private sealed record DatabaseOperationResponse(bool Success, string Issue);
private sealed record DatabaseQueryResponse<TResult>(bool Success, string Issue, TResult? Data);
}
@@ -1,5 +1,8 @@
using System.Text;
using System.Text.Json;
using System.Runtime.CompilerServices;
using AIStudio.Settings;
using AIStudio.Tools.Security;
namespace AIStudio.Tools.Services;
@@ -38,7 +41,7 @@ public sealed partial class RustService
var guardService = Program.SERVICE_PROVIDER.GetRequiredService<PromptInjectionGuardService>();
var streamId = Guid.NewGuid().ToString();
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}";
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}&include_token_count=false";
//
// Both reasons to stop end the same read, so we combine them: our own timeout bounds the
@@ -169,7 +172,11 @@ public sealed partial class RustService
}
catch (JsonException e)
{
this.logger?.LogError(e, "Failed to deserialize SSE event while reading '{Path}': {JsonContent}", path, jsonContent);
// The runtime may report a failure as a bare JSON string instead of a chunk.
// That form still carries a readable reason, so we log it as such -- but it
// remains a failure and must reach the caller like any other:
if (!this.TryLogSseErrorMessage(jsonContent, path))
this.logger?.LogError(e, "Failed to deserialize SSE event while reading '{Path}': {JsonContent}", path, jsonContent);
if (failureCode is FileExtractionErrorCode.NONE)
{
@@ -201,9 +208,9 @@ public sealed partial class RustService
}
finally
{
var finalContentChunk = ContentStreamSseHandler.Clear(streamId);
if (!string.IsNullOrWhiteSpace(finalContentChunk))
resultBuilder.AppendLine(finalContentChunk);
// Reading the whole file at once needs no token counts, so only the content is used here:
if (ContentStreamSseHandler.Clear(streamId) is { } finalContentChunk && !string.IsNullOrWhiteSpace(finalContentChunk.Content))
resultBuilder.AppendLine(finalContentChunk.Content);
}
if (failureCode is not FileExtractionErrorCode.NONE)
@@ -246,4 +253,207 @@ public sealed partial class RustService
PromptInjectionRedactedCount = promptInjectionRedactedCount,
};
}
}
public async IAsyncEnumerable<string> StreamArbitraryFileData(string path, bool extractImages = false, [EnumeratorCancellation] CancellationToken token = default)
{
await foreach (var segment in this.StreamArbitraryFileDataCore(path, extractImages, false, string.Empty, token))
yield return segment.Content;
}
public async IAsyncEnumerable<ArbitraryFileDataSegment> StreamArbitraryFileDataWithTokenCounts(
string path,
EmbeddingProvider embeddingProvider,
[EnumeratorCancellation] CancellationToken token = default)
{
await foreach (var segment in this.StreamArbitraryFileDataCore(path, false, true, embeddingProvider.TokenizerPath, token))
{
if (segment.TokenCount is { } tokenCount)
{
yield return new(segment.Content, tokenCount);
continue;
}
//
// A segment the runtime did not count, e.g. a page which carries an embedded image on
// top of its text. The runtime leaves such a count out on purpose instead of failing
// the extraction, because we can count the segment ourselves. Without this, a document
// would be dropped over a number we are able to produce.
//
var countedSegment = await this.GetTokenCount(embeddingProvider, segment.Content, token);
if (countedSegment is { Success: true } counted)
{
yield return new(segment.Content, counted.TokenCount);
continue;
}
//
// Carries a code so callers can classify it: the file itself is fine, the answer of
// the runtime was not, which makes this worth another attempt.
//
throw new FileExtractionException(FileExtractionErrorCode.INVALID_RESPONSE, $"Rust did not return a token count for an extracted segment from '{path}' using provider '{embeddingProvider.Name}', and counting it afterwards failed as well: {countedSegment?.Message}");
}
}
private async IAsyncEnumerable<(string Content, int? TokenCount)> StreamArbitraryFileDataCore(
string path,
bool extractImages,
bool includeTokenCount,
string tokenizerPath,
[EnumeratorCancellation] CancellationToken token)
{
var streamId = Guid.NewGuid().ToString();
var requestUri = $"/retrieval/fs/extract?path={Uri.EscapeDataString(path)}&stream_id={streamId}&extract_images={extractImages}&include_token_count={includeTokenCount}&tokenizer_path={Uri.EscapeDataString(tokenizerPath)}";
using var request = new HttpRequestMessage(HttpMethod.Get, requestUri);
using var response = await this.http.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, token);
if (!response.IsSuccessStatusCode)
{
var responseBody = await response.Content.ReadAsStringAsync(token);
this.logger?.LogError(
"Failed to stream arbitrary file data from Rust runtime. Status: {StatusCode}, reason: '{ReasonPhrase}', path: '{Path}', body: '{Body}'",
response.StatusCode,
response.ReasonPhrase,
path,
responseBody);
if (includeTokenCount)
throw new InvalidOperationException($"Rust could not extract and count '{path}'. HTTP {(int)response.StatusCode} ({response.ReasonPhrase}): {responseBody}");
yield break;
}
var promptInjectionFindings = new List<PromptInjectionFinding>();
var promptInjectionRedactedCount = 0;
ContentStreamPendingContent? finalContentChunk;
try
{
await using var stream = await response.Content.ReadAsStreamAsync(token);
using var reader = new StreamReader(stream);
while (!reader.EndOfStream && !token.IsCancellationRequested)
{
var line = await reader.ReadLineAsync(token);
if (string.IsNullOrWhiteSpace(line))
continue;
if (!line.StartsWith("data:", StringComparison.InvariantCulture))
continue;
var jsonContent = line[5..];
ContentStreamSseEvent? sseEvent = null;
try
{
sseEvent = JsonSerializer.Deserialize<ContentStreamSseEvent>(jsonContent);
}
catch (JsonException)
{
if (this.TryLogSseErrorMessage(jsonContent, path))
{
if (includeTokenCount)
throw new InvalidOperationException($"Rust could not extract and count a segment from '{path}'. See the runtime log for details.");
continue;
}
this.logger?.LogError("Failed to deserialize SSE event: {JsonContent}", jsonContent);
}
if (sseEvent is null)
continue;
var processedEvent = ContentStreamSseHandler.ProcessEvent(sseEvent, extractImages);
if (processedEvent.Error is { } error)
{
// A notice says something about the file without failing the read, so the
// remaining content still belongs into the index:
if (error.IsNotice)
{
this.logger?.LogInformation(
"The runtime reported a notice while reading '{Path}' for embedding: code={ErrorCode}, detectedFormat='{DetectedFormat}', message='{Message}'",
path,
error.ParsedCode,
error.DetectedFormat,
error.Message);
continue;
}
//
// Everything else stops the read. Embedding a document which was only read in
// part would put a silently incomplete text into the index, and nothing after
// this point would reveal the gap:
//
this.logger?.LogError(
"The runtime reported a failure while reading '{Path}' for embedding: code={ErrorCode}, page={PageNumber}, detectedFormat='{DetectedFormat}', message='{Message}'",
path,
error.ParsedCode,
error.PageNumber,
error.DetectedFormat,
error.Message);
throw new FileExtractionException(error.ParsedCode, $"Rust could not extract '{path}': {error.Message}", error.PageNumber, error.DetectedFormat);
}
if (processedEvent.PromptInjection is { } promptInjection)
{
//
// Not a failure: the passages were removed and the document around them is
// intact, so what remains still belongs into the index. It only has to reach
// the user, because from here on the indexed document is no longer the one
// sitting on their disk.
//
promptInjectionRedactedCount += promptInjection.RedactedCount;
if (promptInjection.Findings is { } findings)
promptInjectionFindings.AddRange(findings);
continue;
}
//
// The count comes from the processed event, not from the event which was just read:
// a reader may hold content back across several events, and the count of the content
// it releases is the count of that content, not of the event that released it.
//
if (!string.IsNullOrWhiteSpace(processedEvent.Content))
yield return (processedEvent.Content, processedEvent.TokenCount);
}
}
finally
{
finalContentChunk = ContentStreamSseHandler.Clear(streamId);
}
if (finalContentChunk is { } pendingContent && !string.IsNullOrWhiteSpace(pendingContent.Content))
yield return (pendingContent.Content, pendingContent.TokenCount);
if (promptInjectionRedactedCount is 0)
yield break;
//
// Reported from here for the same reason as in ReadArbitraryFileData above: these two
// methods together are every way of reading a file, so they are the only two places
// where no caller can forget the report. Here it was missing, which is why a whole
// indexing run could filter documents without ever saying so.
//
var guardService = Program.SERVICE_PROVIDER.GetRequiredService<PromptInjectionGuardService>();
await guardService.ReportAsync(new(PromptInjectionSource.FileContent(path), promptInjectionFindings, promptInjectionRedactedCount));
}
private bool TryLogSseErrorMessage(string jsonContent, string path)
{
try
{
var errorMessage = JsonSerializer.Deserialize<string>(jsonContent);
if (string.IsNullOrWhiteSpace(errorMessage))
return false;
this.logger?.LogError("Rust retrieval stream error for '{Path}': {ErrorMessage}", path, errorMessage);
return true;
}
catch (JsonException)
{
return false;
}
}
}
@@ -0,0 +1,87 @@
using AIStudio.Settings;
using AIStudio.Tools.Rust;
namespace AIStudio.Tools.Services;
public sealed partial class RustService
{
internal const int MAX_TOKEN_COUNT_REQUEST_TEXT_LENGTH = 200_000;
private static TokenizerResponse CreateUnavailableTokenizerResponse(string message) => new(
false,
0,
message,
string.Empty);
public async Task<TokenizerResponse> ValidateTokenizer(string filePath)
{
var result = await this.http.PostAsJsonAsync("/tokenizer/validate", new {
file_path = filePath,
}, this.jsonRustSerializerOptions);
if (!result.IsSuccessStatusCode)
{
this.logger!.LogError($"Failed to validate the tokenizer '{result.StatusCode}'");
return CreateUnavailableTokenizerResponse("An error occured while sending the path to the Rust framework for validation: "+result.StatusCode);
}
var response = await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
return response;
}
public async Task<TokenizerResponse> StoreTokenizer(string modelId, string filePath)
{
this.logger!.LogInformation($"Storing tokenizer for model '{modelId}' from file '{filePath}'");
var result = await this.http.PostAsJsonAsync("/tokenizer/store", new {
model_id = modelId,
file_path = filePath,
}, this.jsonRustSerializerOptions);
if (!result.IsSuccessStatusCode)
{
this.logger!.LogError($"Failed to store the tokenizer '{result.StatusCode}'");
return CreateUnavailableTokenizerResponse("An error occured while sending the path to the Rust framework for storing: "+result.StatusCode);
}
return await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
}
public async Task<TokenizerResponse> DeleteTokenizer(string modelId)
{
this.logger!.LogInformation($"Deleting tokenizer for model '{modelId}'");
var result = await this.http.PostAsJsonAsync("/tokenizer/delete", new {
model_id = modelId,
}, this.jsonRustSerializerOptions);
if (!result.IsSuccessStatusCode)
{
this.logger!.LogError($"Failed to delete the tokenizer '{result.StatusCode}'");
return CreateUnavailableTokenizerResponse("An error occured while sending the tokenizer delete request to the Rust framework: "+result.StatusCode);
}
return await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions);
}
public Task<TokenizerResponse?> GetTokenCount(AIStudio.Settings.Provider provider, string text, CancellationToken cancellationToken = default) =>
this.GetTokenCount(provider.InstanceName, provider.TokenizerPath, text, cancellationToken);
public Task<TokenizerResponse?> GetTokenCount(EmbeddingProvider provider, string text, CancellationToken cancellationToken = default) =>
this.GetTokenCount(provider.Name, provider.TokenizerPath, text, cancellationToken);
private async Task<TokenizerResponse?> GetTokenCount(string providerName, string tokenizerPath, string text, CancellationToken cancellationToken)
{
var result = await this.http.PostAsJsonAsync("/tokenizer/count", new {
text,
tokenizer_path = tokenizerPath,
}, this.jsonRustSerializerOptions, cancellationToken);
if (!result.IsSuccessStatusCode)
{
this.logger!.LogError("Failed to get the token count for provider '{ProviderName}': {StatusCode}", providerName, result.StatusCode);
return CreateUnavailableTokenizerResponse("Error while getting token count from Rust service: "+result.StatusCode);
}
return await result.Content.ReadFromJsonAsync<TokenizerResponse>(this.jsonRustSerializerOptions, cancellationToken);
}
}
@@ -0,0 +1,20 @@
namespace AIStudio.Tools.Services;
public static class TokenizerModelId
{
public static string ForProvider(Settings.Provider provider) => ForProviderId(provider.Id);
public static string ForProviderId(string guid) => "chat_" + NormalizeGuid(guid);
public static string ForEmbeddingProvider(Settings.EmbeddingProvider provider) => ForEmbeddingProviderId(provider.Id);
public static string ForEmbeddingProviderId(string guid) => "embedding_" + NormalizeGuid(guid);
private static string NormalizeGuid(string guid)
{
if (Guid.TryParse(guid, out var parsedGuid))
return parsedGuid.ToString("D");
return guid.Trim();
}
}
+11 -1
View File
@@ -5,6 +5,16 @@ public sealed class Slide
public bool Delivered { get; set; }
public int Position { get; init; }
public List<ISlideContent> Content { get; } = new();
/// <summary>
/// The number of tokens of everything this slide holds, or null when it is unknown.
/// </summary>
/// <remarks>
/// A slide grows across several stream events, so its count grows with it. It becomes unknown
/// as soon as an image is embedded: the runtime counted the text of the slide, and a data URI
/// is orders of magnitude larger than that.
/// </remarks>
public int? TokenCount { get; set; }
}
+30 -8
View File
@@ -6,7 +6,7 @@ public sealed class SlideManager
{
private readonly Dictionary<int, Slide> slides = new();
public void AddSlide(ContentStreamPresentationMetadata metadata, string? content, bool extractImages = false)
public void AddSlide(ContentStreamPresentationMetadata metadata, string? content, int? tokenCount, bool extractImages = false)
{
var slideNumber = metadata.Presentation?.SlideNumber ?? 0;
if(slideNumber is 0)
@@ -42,11 +42,15 @@ public sealed class SlideManager
var createdSlide = new Slide
{
Delivered = false,
Position = slideNumber
Position = slideNumber,
// The count of the text we just added. It travels with the slide, because the slide
// is delivered long after this event:
TokenCount = tokenCount
};
createdSlide.Content.Add(slideText);
//
// Add image content to the slide?
//
@@ -54,7 +58,12 @@ public sealed class SlideManager
{
var markdownImage = ContentStreamSseHandler.BuildImageMarkdown(image!.Id!, image.MediaType);
if (markdownImage is not null)
{
createdSlide.Content.Add(new SlideImageContent(markdownImage));
// The runtime counted the text of the slide, not the data URI we just added:
createdSlide.TokenCount = null;
}
}
this.slides[slideNumber] = createdSlide;
@@ -70,24 +79,37 @@ public sealed class SlideManager
{
var textContent = slide.Content.OfType<SlideTextContent>().First();
textContent.Text.AppendLine(content);
slide.TokenCount = ContentStreamPendingContent.AddTokenCounts(slide.TokenCount, tokenCount);
}
// Add any image content?
if (addImage)
{
var markdownImage = ContentStreamSseHandler.BuildImageMarkdown(image!.Id!, image.MediaType);
if (markdownImage is not null)
{
slide.Content.Add(new SlideImageContent(markdownImage));
// The runtime counted the text of the slide, not the data URI we just added:
slide.TokenCount = null;
}
}
}
}
public string? GetAllSlidesInOrder()
public ContentStreamPendingContent? GetAllSlidesInOrder()
{
var content = new StringBuilder();
// Starts at zero and stays a number only as long as every slide contributes a count of its
// own. One slide without one makes the total unknown, which is what the caller has to know:
int? tokenCount = 0;
foreach (var slide in this.slides.Values.Where(s => !s.Delivered).OrderBy(s => s.Position))
{
slide.Delivered = true;
tokenCount = ContentStreamPendingContent.AddTokenCounts(tokenCount, slide.TokenCount);
foreach (var text in slide.Content.OfType<SlideTextContent>())
{
content.AppendLine(text.Text.ToString());
@@ -100,7 +122,7 @@ public sealed class SlideManager
content.AppendLine();
}
}
return content.Length > 0 ? content.ToString() : null;
return content.Length > 0 ? new ContentStreamPendingContent(content.ToString(), tokenCount) : null;
}
}
@@ -1,3 +1,5 @@
using AIStudio.Provider;
using AIStudio.Settings;
using AIStudio.Settings.DataModel;
using AIStudio.Tools.ERIClient.DataModel;
using AIStudio.Tools.PluginSystem;
@@ -6,7 +8,11 @@ namespace AIStudio.Tools.Validation;
public sealed class DataSourceValidation
{
public const int MAX_NAME_LENGTH = 40;
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(DataSourceValidation).Namespace, nameof(DataSourceValidation));
public static bool IsNameValid(string name) => !string.IsNullOrWhiteSpace(name) && name.Length <= MAX_NAME_LENGTH && !name.Any(char.IsControl);
public Func<string> GetSecretStorageIssue { get; init; } = () => string.Empty;
@@ -17,8 +23,14 @@ public sealed class DataSourceValidation
public Func<AuthMethod> GetAuthMethod { get; init; } = () => AuthMethod.NONE;
public Func<SecurityRequirements?> GetSecurityRequirements { get; init; } = () => null;
public Func<bool> GetSelectedCloudEmbedding { get; init; } = () => false;
public Func<EmbeddingProvider?> GetSelectedEmbeddingProvider { get; init; } = () => null;
public Func<ConfidenceLevel> GetConfidenceLevel { get; init; } = () => ConfidenceLevel.NONE;
public Func<SettingsManager?> GetSettingsManager { get; init; } = () => null;
public Func<bool> GetTestedConnection { get; init; } = () => false;
@@ -47,19 +59,19 @@ public sealed class DataSourceValidation
return null;
}
public string? ValidateSecurityPolicy(DataSourceSecurity securityPolicy)
{
if(securityPolicy is DataSourceSecurity.NOT_SPECIFIED)
return TB("Please select your security policy.");
var dataSourceSecurity = this.GetSecurityRequirements();
if (dataSourceSecurity is null)
return null;
if(dataSourceSecurity.Value.AllowedProviderType is ProviderType.SELF_HOSTED && securityPolicy is not DataSourceSecurity.SELF_HOSTED)
return TB("This data source can only be used with a self-hosted LLM provider. Please change the security policy.");
return null;
}
@@ -106,11 +118,14 @@ public sealed class DataSourceValidation
public string? ValidatingName(string dataSourceName)
{
if(string.IsNullOrWhiteSpace(dataSourceName))
if (string.IsNullOrWhiteSpace(dataSourceName))
return TB("The name must not be empty.");
if (dataSourceName.Length > 40)
if (dataSourceName.Length > MAX_NAME_LENGTH)
return TB("The name must not exceed 40 characters.");
if (dataSourceName.Any(char.IsControl))
return TB("The name must not contain control characters.");
var lowerName = dataSourceName.ToLowerInvariant();
if(lowerName != this.GetPreviousDataSourceName() && this.GetUsedDataSourceNames().Contains(lowerName))
@@ -149,6 +164,20 @@ public sealed class DataSourceValidation
return null;
}
public string? ValidateEmbeddingProviderAccess(string embeddingId)
{
var embeddingIssue = this.ValidateEmbeddingId(embeddingId);
return embeddingIssue ?? this.ValidateSelectedEmbeddingProviderAccess();
}
public string? ValidateDataSourceConfidenceLevel(ConfidenceLevel confidenceLevel)
{
if(confidenceLevel is ConfidenceLevel.NONE)
return TB("Please select a required provider confidence level.");
return this.ValidateSelectedEmbeddingProviderAccess();
}
public string? ValidateUserAcknowledgedCloudEmbedding(bool value)
{
if(this.GetSelectedCloudEmbedding() && !value)
@@ -175,4 +204,21 @@ public sealed class DataSourceValidation
return null;
}
}
private string? ValidateSelectedEmbeddingProviderAccess()
{
var selectedEmbedding = this.GetSelectedEmbeddingProvider();
var settingsManager = this.GetSettingsManager();
if(selectedEmbedding is null || settingsManager is null)
return null;
var confidenceLevel = this.GetConfidenceLevel();
if(selectedEmbedding.GetConfidenceLevel(settingsManager).AllowsDataSourceConfidenceLevel(confidenceLevel))
return null;
return string.Format(
TB("The selected embedding provider has confidence '{0}', but this data source requires provider confidence '{1}'. Select an embedding provider with equal or higher confidence or lower the required confidence level."),
selectedEmbedding.GetConfidenceLevel(settingsManager).GetName(),
confidenceLevel.GetName());
}
}
@@ -22,6 +22,7 @@ public sealed class ProviderValidation
public Func<bool> IsModelProvidedManually { get; init; } = () => false;
public Func<string> GetCustomTokenizerValidationIssue { get; init; } = () => string.Empty;
public Func<bool> IsModelSelectionHidden { get; init; } = () => false;
public string? ValidatingHostname(string hostname)
@@ -132,6 +133,15 @@ public sealed class ProviderValidation
return null;
}
public string? ValidatingCustomTokenizer(string _)
{
var issue = this.GetCustomTokenizerValidationIssue();
if (string.IsNullOrWhiteSpace(issue))
return null;
return issue;
}
/// <summary>
/// Validates the Hugging Face inference provider chosen for embeddings.
/// </summary>
@@ -179,4 +189,4 @@ public sealed class ProviderValidation
return null;
}
}
}