diff --git a/app/MindWork AI Studio/Program.cs b/app/MindWork AI Studio/Program.cs index 21073095..39b9627e 100644 --- a/app/MindWork AI Studio/Program.cs +++ b/app/MindWork AI Studio/Program.cs @@ -133,6 +133,7 @@ internal sealed class Program builder.Services.AddSingleton(); builder.Services.AddSingleton(); builder.Services.AddSingleton(); + builder.Services.AddSingleton(); builder.Services.AddScoped(); builder.Services.AddTransient(); builder.Services.AddTransient(); diff --git a/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalDirectory.cs b/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalDirectory.cs index 3c5c0c47..efed00aa 100644 --- a/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalDirectory.cs +++ b/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalDirectory.cs @@ -1,6 +1,7 @@ using AIStudio.Chat; using AIStudio.Provider; using AIStudio.Tools.RAG; +using AIStudio.Tools.Services; namespace AIStudio.Settings.DataModel; @@ -56,11 +57,8 @@ public readonly record struct DataSourceLocalDirectory : IInternalDataSource public ushort MaxMatches { get; init; } = 10; /// - public Task> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) - { - IReadOnlyList retrievalContext = new List(); - return Task.FromResult(retrievalContext); - } + public Task> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) => + Program.SERVICE_PROVIDER.GetRequiredService().RetrieveDataAsync(this, lastUserPrompt, thread, token); /// /// The path to the directory. diff --git a/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalFile.cs b/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalFile.cs index 641c5355..130c161a 100644 --- a/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalFile.cs +++ b/app/MindWork AI Studio/Settings/DataModel/DataSourceLocalFile.cs @@ -1,6 +1,7 @@ using AIStudio.Chat; using AIStudio.Provider; using AIStudio.Tools.RAG; +using AIStudio.Tools.Services; namespace AIStudio.Settings.DataModel; @@ -56,11 +57,8 @@ public readonly record struct DataSourceLocalFile : IInternalDataSource public ushort MaxMatches { get; init; } = 10; /// - public Task> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) - { - IReadOnlyList retrievalContext = new List(); - return Task.FromResult(retrievalContext); - } + public Task> RetrieveDataAsync(IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) => + Program.SERVICE_PROVIDER.GetRequiredService().RetrieveDataAsync(this, lastUserPrompt, thread, token); /// /// The path to the file. diff --git a/app/MindWork AI Studio/Tools/Databases/EmbeddingState/EmbeddingStateClient.cs b/app/MindWork AI Studio/Tools/Databases/EmbeddingState/EmbeddingStateClient.cs index 5ebf4488..4cf8e513 100644 --- a/app/MindWork AI Studio/Tools/Databases/EmbeddingState/EmbeddingStateClient.cs +++ b/app/MindWork AI Studio/Tools/Databases/EmbeddingState/EmbeddingStateClient.cs @@ -26,6 +26,8 @@ public abstract class EmbeddingStateClient(string name, string path) : DatabaseC public abstract Task UpsertChunksAsync(string dataSourceId, IReadOnlyList chunks, CancellationToken token); + public abstract Task> SearchChunksAsync(string dataSourceId, string query, int maxMatches, CancellationToken token); + public abstract Task DeleteDataSourceAsync(string dataSourceId, CancellationToken token); } @@ -51,3 +53,26 @@ public sealed record EmbeddingStateChunk( int ChunkIndex, string ChunkText, DateTime EmbeddedAtUtc); + +public sealed record EmbeddingStateSearchResult( + 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, + DateTime CreationUtc, + DateTime LastWriteUtc, + DateTime EmbeddedAtUtc, + int ChunkCount, + string ComplianceLevel, + int ComplianceLevelRank); diff --git a/app/MindWork AI Studio/Tools/Databases/EmbeddingState/NoEmbeddingStateClient.cs b/app/MindWork AI Studio/Tools/Databases/EmbeddingState/NoEmbeddingStateClient.cs index e043b38a..e7696cfa 100644 --- a/app/MindWork AI Studio/Tools/Databases/EmbeddingState/NoEmbeddingStateClient.cs +++ b/app/MindWork AI Studio/Tools/Databases/EmbeddingState/NoEmbeddingStateClient.cs @@ -46,6 +46,9 @@ public sealed class NoEmbeddingStateClient(string name, string? unavailableReaso public override Task UpsertChunksAsync(string dataSourceId, IReadOnlyList chunks, CancellationToken token) => Task.CompletedTask; + public override Task> SearchChunksAsync(string dataSourceId, string query, int maxMatches, CancellationToken token) => + Task.FromResult>([]); + public override Task DeleteDataSourceAsync(string dataSourceId, CancellationToken token) => Task.CompletedTask; public override void Dispose() diff --git a/app/MindWork AI Studio/Tools/Databases/EmbeddingState/SqliteEmbeddingStateClientImplementation.cs b/app/MindWork AI Studio/Tools/Databases/EmbeddingState/SqliteEmbeddingStateClientImplementation.cs index a6fe24c6..905d38e6 100644 --- a/app/MindWork AI Studio/Tools/Databases/EmbeddingState/SqliteEmbeddingStateClientImplementation.cs +++ b/app/MindWork AI Studio/Tools/Databases/EmbeddingState/SqliteEmbeddingStateClientImplementation.cs @@ -1,4 +1,5 @@ using System.Globalization; +using System.Text.RegularExpressions; using AIStudio.Settings; using AIStudio.Tools.PluginSystem; @@ -16,6 +17,9 @@ public sealed class SqliteEmbeddingStateClientImplementation( { 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 static readonly Regex FTS_TOKEN_REGEX = new(@"[\p{L}\p{Nd}_]+", RegexOptions.Compiled | RegexOptions.CultureInvariant); private readonly string databasePath = databasePath; private readonly string connectionString = new SqliteConnectionStringBuilder @@ -316,6 +320,84 @@ public sealed class SqliteEmbeddingStateClientImplementation( await transaction.CommitAsync(token); } + public override async Task> SearchChunksAsync(string dataSourceId, string query, int maxMatches, CancellationToken token) + { + if (maxMatches <= 0) + return []; + + var ftsQuery = BuildFtsQuery(query); + if (string.IsNullOrWhiteSpace(ftsQuery)) + return []; + + var results = new List(maxMatches); + await using var connection = await this.OpenConnectionAsync(token); + await using var command = CreateCommand(connection, """ + SELECT + c.chunk_id, + c.parent_file_id, + ds.data_source_id, + ds.data_source_name, + ds.data_source_type, + f.absolute_path, + f.file_name, + f.relative_path, + f.file_type, + c.page_number, + c.chunk_index, + c.chunk_text, + bm25(embedding_chunks_fts) AS score, + f.fingerprint, + f.file_size, + f.creation_utc, + f.last_write_utc, + c.embedded_at_utc, + f.chunk_count, + f.compliance_level, + f.compliance_level_rank + 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 $query + ORDER BY score + LIMIT $maxMatches + """); + + command.Parameters.AddWithValue("$dataSourceId", dataSourceId); + command.Parameters.AddWithValue("$query", ftsQuery); + command.Parameters.AddWithValue("$maxMatches", maxMatches); + + await using var reader = await command.ExecuteReaderAsync(token); + while (await reader.ReadAsync(token)) + { + results.Add(new EmbeddingStateSearchResult( + reader.GetString(0), + reader.GetString(1), + reader.GetString(2), + reader.GetString(3), + reader.GetString(4), + reader.GetString(5), + reader.GetString(6), + reader.GetString(7), + reader.GetString(8), + reader.IsDBNull(9) ? null : reader.GetInt32(9), + reader.GetInt32(10), + reader.GetString(11), + reader.GetDouble(12), + reader.GetString(13), + reader.GetInt64(14), + ParseUtc(reader.GetString(15)), + ParseUtc(reader.GetString(16)), + ParseUtc(reader.GetString(17)), + reader.GetInt32(18), + reader.GetString(19), + reader.GetInt32(20))); + } + + return results; + } + public override async Task DeleteDataSourceAsync(string dataSourceId, CancellationToken token) { await using var connection = await this.OpenConnectionAsync(token); @@ -525,6 +607,20 @@ public sealed class SqliteEmbeddingStateClientImplementation( : DateTime.UnixEpoch; } + 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 NoEmbeddingStateClient CreateNoEmbeddingStateClient(string name, string? unavailableReason, DatabaseClientStatus status, ILogger databaseClientLogger) { var client = new NoEmbeddingStateClient(name, unavailableReason, status); diff --git a/app/MindWork AI Studio/Tools/Databases/VectorStore/NoVectorStoreClient.cs b/app/MindWork AI Studio/Tools/Databases/VectorStore/NoVectorStoreClient.cs index b7d77bfb..c681a4c2 100644 --- a/app/MindWork AI Studio/Tools/Databases/VectorStore/NoVectorStoreClient.cs +++ b/app/MindWork AI Studio/Tools/Databases/VectorStore/NoVectorStoreClient.cs @@ -28,6 +28,9 @@ public sealed class NoVectorStoreClient(string name, string? unavailableReason, public override Task InsertEmbedding(string storeName, IReadOnlyList points, CancellationToken token) => Task.FromException(this.CreateUnavailableException()); + public override Task> SearchEmbeddingAsync(string storeName, IReadOnlyList vector, int maxMatches, CancellationToken token) => + Task.FromException>(this.CreateUnavailableException()); + public override Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token) => Task.FromException(this.CreateUnavailableException()); diff --git a/app/MindWork AI Studio/Tools/Databases/VectorStore/QdrantEdgeClientImplementation.cs b/app/MindWork AI Studio/Tools/Databases/VectorStore/QdrantEdgeClientImplementation.cs index da519703..7a8d3b2f 100644 --- a/app/MindWork AI Studio/Tools/Databases/VectorStore/QdrantEdgeClientImplementation.cs +++ b/app/MindWork AI Studio/Tools/Databases/VectorStore/QdrantEdgeClientImplementation.cs @@ -15,6 +15,7 @@ public sealed class QdrantEdgeClientImplementation( 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 DELETE_STORE_PATH = "/system/qdrant-edge/delete-store"; @@ -86,6 +87,18 @@ public sealed class QdrantEdgeClientImplementation( public override Task InsertEmbedding(string storeName, IReadOnlyList points, CancellationToken token) => rustService.ExecuteDatabaseOperation(DATABASE_NAME, INSERT_PATH, new InsertEmbeddingRequest(storeName, points), token); + public override async Task> SearchEmbeddingAsync(string storeName, IReadOnlyList vector, int maxMatches, CancellationToken token) + { + if (maxMatches <= 0) + return []; + + return await rustService.ExecuteDatabaseQuery>( + 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); @@ -108,6 +121,8 @@ public sealed class QdrantEdgeClientImplementation( private sealed record InsertEmbeddingRequest(string StoreName, IReadOnlyList Points); + private sealed record SearchEmbeddingRequest(string StoreName, IReadOnlyList Vector, int MaxMatches); + private sealed record DeleteEmbeddingByFileRequest(string StoreName, string FilePath); private sealed record DeleteVectorStoreRequest(string StoreName); diff --git a/app/MindWork AI Studio/Tools/Databases/VectorStore/VectorSearchResult.cs b/app/MindWork AI Studio/Tools/Databases/VectorStore/VectorSearchResult.cs new file mode 100644 index 00000000..94319c3c --- /dev/null +++ b/app/MindWork AI Studio/Tools/Databases/VectorStore/VectorSearchResult.cs @@ -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 ComplianceLevel, + int ComplianceLevelRank); diff --git a/app/MindWork AI Studio/Tools/Databases/VectorStore/VectorStoreClient.cs b/app/MindWork AI Studio/Tools/Databases/VectorStore/VectorStoreClient.cs index c5b3a50e..6831cdc0 100644 --- a/app/MindWork AI Studio/Tools/Databases/VectorStore/VectorStoreClient.cs +++ b/app/MindWork AI Studio/Tools/Databases/VectorStore/VectorStoreClient.cs @@ -6,6 +6,8 @@ public abstract class VectorStoreClient(string name, string path): DatabaseClien public abstract Task InsertEmbedding(string storeName, IReadOnlyList points, CancellationToken token); + public abstract Task> SearchEmbeddingAsync(string storeName, IReadOnlyList vector, int maxMatches, CancellationToken token); + public abstract Task DeleteEmbeddingByFile(string storeName, string filePath, CancellationToken token); public abstract Task DeleteVectorStore(string storeName, CancellationToken token); diff --git a/app/MindWork AI Studio/Tools/RAG/RAGProcesses/AISrcSelWithRetCtxVal.cs b/app/MindWork AI Studio/Tools/RAG/RAGProcesses/AISrcSelWithRetCtxVal.cs index 7c867619..d24df2ea 100644 --- a/app/MindWork AI Studio/Tools/RAG/RAGProcesses/AISrcSelWithRetCtxVal.cs +++ b/app/MindWork AI Studio/Tools/RAG/RAGProcesses/AISrcSelWithRetCtxVal.cs @@ -205,17 +205,7 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess var ragSources = new List(); 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 +215,63 @@ public sealed class AISrcSelWithRetCtxVal : IRagProcess } #endregion -} \ No newline at end of file + + private static IReadOnlyList CreateSources(IRetrievalContext retrievalContext) + { + var sources = new List(); + AddSource(sources, GetReferenceTitle(retrievalContext), GetReferenceLink(retrievalContext)); + foreach (var link in retrievalContext.Links) + AddSource(sources, retrievalContext.DataSourceName, link); + + return sources; + } + + private static void AddSource(ICollection 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); +} diff --git a/app/MindWork AI Studio/Tools/RAG/RetrievalTextContext.cs b/app/MindWork AI Studio/Tools/RAG/RetrievalTextContext.cs index 9997e571..362acf99 100644 --- a/app/MindWork AI Studio/Tools/RAG/RetrievalTextContext.cs +++ b/app/MindWork AI Studio/Tools/RAG/RetrievalTextContext.cs @@ -40,4 +40,14 @@ public sealed class RetrievalTextContext : IRetrievalContext /// For example, one sentence or paragraph before and after the matched text. /// public IReadOnlyList SurroundingContent { get; set; } = []; + + /// + /// Optional title used when this context is displayed as a source reference. + /// + public string ReferenceTitle { get; init; } = string.Empty; + + /// + /// Optional link used when this context is displayed as a source reference. + /// + public string ReferenceLink { get; init; } = string.Empty; } \ No newline at end of file diff --git a/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingNames.cs b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingNames.cs new file mode 100644 index 00000000..3e95ae4c --- /dev/null +++ b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingNames.cs @@ -0,0 +1,21 @@ +namespace AIStudio.Tools.Services; + +internal static class DataSourceEmbeddingNames +{ + public static string GetCollectionName(string dataSourceName, string dataSourceId) + { + var safeId = dataSourceId + .ToLowerInvariant() + .Replace("-", string.Empty, StringComparison.Ordinal); + + var safeName = new string(dataSourceName + .ToLowerInvariant() + .Where(c => c is >= 'a' and <= 'z' or >= '0' and <= '9') + .Take(32) + .ToArray()); + + safeName = string.IsNullOrWhiteSpace(safeName) ? "datasource" : safeName; + + return $"rag_{safeName}_{safeId}"; + } +} diff --git a/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingProviders.cs b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingProviders.cs new file mode 100644 index 00000000..c544f05f --- /dev/null +++ b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingProviders.cs @@ -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; + } +} diff --git a/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.Files.cs b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.Files.cs index 0682ba37..56160ec6 100644 --- a/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.Files.cs +++ b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.Files.cs @@ -905,22 +905,8 @@ public sealed partial class DataSourceEmbeddingService : null; } - private string GetCollectionName(string dataSourceName, string dataSourceId) - { - var safeId = dataSourceId - .ToLowerInvariant() - .Replace("-", string.Empty, StringComparison.Ordinal); - - var safeName = new string(dataSourceName - .ToLowerInvariant() - .Where(c => c is >= 'a' and <= 'z' or >= '0' and <= '9') - .Take(32) - .ToArray()); - - safeName = string.IsNullOrWhiteSpace(safeName) ? "datasource" : safeName; - - return $"rag_{safeName}_{safeId}"; - } + private string GetCollectionName(string dataSourceName, string dataSourceId) => + DataSourceEmbeddingNames.GetCollectionName(dataSourceName, dataSourceId); private string CreatePointId(string dataSourceId, string fingerprint, int chunkIndex) => CreateStableGuid($"{dataSourceId}:chunk:{fingerprint}:{chunkIndex}"); diff --git a/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.cs b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.cs index bbd5fdb7..b51369b2 100644 --- a/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.cs +++ b/app/MindWork AI Studio/Tools/Services/DataSourceEmbeddingService.cs @@ -866,13 +866,7 @@ public sealed partial class DataSourceEmbeddingService(SettingsManager settingsM } private bool TryResolveEmbeddingProvider(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; - } + => DataSourceEmbeddingProviders.TryResolve(settingsManager, dataSource, out embeddingProvider); private async Task EnsureCompatibleManifestAsync( IDataSource dataSource, diff --git a/app/MindWork AI Studio/Tools/Services/DataSourceLocalRetrievalService.cs b/app/MindWork AI Studio/Tools/Services/DataSourceLocalRetrievalService.cs new file mode 100644 index 00000000..839ae493 --- /dev/null +++ b/app/MindWork AI Studio/Tools/Services/DataSourceLocalRetrievalService.cs @@ -0,0 +1,307 @@ +using AIStudio.Chat; +using AIStudio.Provider; +using AIStudio.Settings; +using AIStudio.Settings.DataModel; +using AIStudio.Tools.Databases; +using AIStudio.Tools.Databases.EmbeddingState; +using AIStudio.Tools.Databases.VectorStore; +using AIStudio.Tools.RAG; + +namespace AIStudio.Tools.Services; + +public sealed class DataSourceLocalRetrievalService( + SettingsManager settingsManager, + DatabaseClientProvider databaseClientProvider, + ILogger logger) +{ + private enum RetrievalChannel + { + VECTOR, + BM25, + } + + 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 ComplianceLevel, + int ComplianceLevelRank); + + public Task> RetrieveDataAsync(DataSourceLocalFile dataSource, IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) => + this.RetrieveDataAsync((IInternalDataSource)dataSource, lastUserPrompt, token); + + public Task> RetrieveDataAsync(DataSourceLocalDirectory dataSource, IContent lastUserPrompt, ChatThread thread, CancellationToken token = default) => + this.RetrieveDataAsync((IInternalDataSource)dataSource, lastUserPrompt, token); + + private async Task> 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 candidateLimit = maxMatches * 2; + var collectionName = DataSourceEmbeddingNames.GetCollectionName(dataSource.Name, dataSource.Id); + var vectorTask = this.SearchVectorAsync(dataSource, query, candidateLimit, collectionName, token); + var bm25Task = this.SearchBm25Async(dataSource, query, candidateLimit, 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> 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); + 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); + 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); + return []; + } + + return await vectorStore.SearchEmbeddingAsync(collectionName, vector, maxMatches, token); + } + catch (OperationCanceledException) when (token.IsCancellationRequested) + { + throw; + } + catch (Exception exception) + { + logger.LogWarning(exception, "Vector retrieval failed for data source '{DataSourceName}' ({DataSourceId}).", dataSource.Name, dataSource.Id); + return []; + } + } + + private async Task> SearchBm25Async(IInternalDataSource dataSource, string query, int maxMatches, CancellationToken token) + { + try + { + var embeddingState = await databaseClientProvider.GetEmbeddingStateAsync(token); + if (!embeddingState.IsAvailable) + { + logger.LogWarning( + "Skipping BM25 retrieval for data source '{DataSourceName}' ({DataSourceId}) because local RAG index '{DatabaseName}' is unavailable.", + dataSource.Name, + dataSource.Id, + embeddingState.Name); + return []; + } + + return await embeddingState.SearchChunksAsync(dataSource.Id, query, maxMatches, token); + } + 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 static IReadOnlyList MergeResults( + IReadOnlyList vectorResults, + IReadOnlyList bm25Results, + int maxMatches) + { + // Future reranking should replace this deterministic channel merge. + var merged = new List(maxMatches * 2); + var seenChunkIds = new HashSet(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 merged, HashSet seenChunkIds, IEnumerable 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.ComplianceLevel, + result.ComplianceLevelRank); + + private static LocalRetrievalHit FromBm25Result(EmbeddingStateSearchResult 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.ComplianceLevel, + result.ComplianceLevelRank); + + 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); + var page = hit.PageNumber is > 0 ? $", page {hit.PageNumber}" : string.Empty; + return $"{sourceName} (chunk {hit.ChunkIndex + 1}{page})"; + } + + 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) => fileType.TrimStart('.').ToLowerInvariant() switch + { + "csv" or "tsv" or "ods" or "xls" or "xlsx" or "xlsm" or "xlsb" => RetrievalContentType.TEXT_SPREADSHEET, + "odp" or "ppt" or "pptx" => RetrievalContentType.TEXT_PRESENTATION, + "htm" or "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; +} diff --git a/app/MindWork AI Studio/Tools/Services/RustService.Databases.cs b/app/MindWork AI Studio/Tools/Services/RustService.Databases.cs index 3f101d70..d4e8bf0f 100644 --- a/app/MindWork AI Studio/Tools/Services/RustService.Databases.cs +++ b/app/MindWork AI Studio/Tools/Services/RustService.Databases.cs @@ -49,5 +49,22 @@ public sealed partial class RustService throw new InvalidOperationException(operation?.Issue ?? $"The {databaseName} operation failed."); } + public async Task ExecuteDatabaseQuery(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>(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(bool Success, string Issue, TResult? Data); } diff --git a/app/MindWork AI Studio/Tools/SourceExtensions.cs b/app/MindWork AI Studio/Tools/SourceExtensions.cs index c7ecf314..e71ffa7b 100644 --- a/app/MindWork AI Studio/Tools/SourceExtensions.cs +++ b/app/MindWork AI Studio/Tools/SourceExtensions.cs @@ -73,7 +73,7 @@ public static class SourceExtensions public static void MergeSources(this IList sources, IList addedSources) { foreach (var addedSource in addedSources) - if (sources.All(s => s.URL != addedSource.URL && s.Title != addedSource.Title)) + if (sources.All(s => s.URL != addedSource.URL || s.Title != addedSource.Title)) sources.Add((Source)addedSource); } } \ No newline at end of file diff --git a/runtime/src/qdrant_edge_database.rs b/runtime/src/qdrant_edge_database.rs index a77b3213..151c3ff9 100644 --- a/runtime/src/qdrant_edge_database.rs +++ b/runtime/src/qdrant_edge_database.rs @@ -6,12 +6,13 @@ use std::sync::Mutex; use axum::Json; use log::{error, info, warn}; use once_cell::sync::Lazy; -use qdrant_edge::external::serde_json::json; +use qdrant_edge::external::serde_json::{json, Value}; use qdrant_edge::external::uuid::Uuid; use qdrant_edge::{ Condition, Distance, EdgeConfig, EdgeOptimizersConfig, EdgeShard, EdgeVectorParams, - FieldCondition, Filter, HnswIndexConfig, Match, MatchValue, PointId, PointInsertOperations, - PointOperations, PointStruct, UpdateOperation, ValueVariants, Vectors, + FieldCondition, Filter, HnswIndexConfig, Match, MatchValue, NamedQuery, Payload, PointId, + PointInsertOperations, PointOperations, PointStruct, QueryEnum, ScoredPoint, SearchRequest, + UpdateOperation, ValueVariants, VectorInternal, Vectors, WithPayloadInterface, WithVector, }; use serde::{Deserialize, Serialize}; use tauri::Manager; @@ -97,6 +98,13 @@ pub struct InsertQdrantEdgeEmbeddingRequest { pub points: Vec, } +#[derive(Deserialize)] +pub struct SearchQdrantEdgeEmbeddingRequest { + pub store_name: String, + pub vector: Vec, + pub max_matches: usize, +} + #[derive(Deserialize)] pub struct DeleteQdrantEdgeEmbeddingByFileRequest { pub store_name: String, @@ -114,6 +122,38 @@ pub struct QdrantEdgeOperationResponse { pub issue: String, } +#[derive(Serialize)] +pub struct QdrantEdgeSearchResponse { + pub success: bool, + pub issue: String, + pub data: Vec, +} + +#[derive(Serialize)] +pub struct QdrantEdgeSearchResult { + pub point_id: String, + pub score: f32, + pub data_source_id: String, + pub data_source_name: String, + pub data_source_type: String, + pub chunk_id: String, + pub parent_file_id: String, + pub file_path: String, + pub absolute_path: String, + pub file_name: String, + pub relative_path: String, + pub file_type: String, + pub page_number: Option, + pub chunk_index: i32, + pub text: String, + pub fingerprint: String, + pub creation_utc: String, + pub last_write_utc: String, + pub embedded_at_utc: String, + pub compliance_level: String, + pub compliance_level_rank: i32, +} + #[derive(Clone, Serialize)] pub struct QdrantEdgeInfo { pub name: String, @@ -223,6 +263,36 @@ impl QdrantEdgeDatabase { Ok(()) } + fn search_embedding(&mut self, store_name: &str, vector: Vec, max_matches: usize) -> QdrantEdgeResult> { + if max_matches == 0 { + return Ok(vec![]); + } + + validate_vector_size(vector.len())?; + let Some(shard) = self.get_existing_store(store_name)? else { + return Ok(vec![]); + }; + + let search_results = shard.search(SearchRequest { + query: QueryEnum::Nearest(NamedQuery::new( + VectorInternal::Dense(vector), + VECTOR_NAME, + )), + filter: None, + params: None, + limit: max_matches, + offset: 0, + with_payload: Some(WithPayloadInterface::Bool(true)), + with_vector: Some(WithVector::Bool(false)), + score_threshold: None, + })?; + + Ok(search_results + .into_iter() + .map(to_qdrant_edge_search_result) + .collect()) + } + fn delete_embedding_by_file(&mut self, store_name: &str, file_path: &str) -> QdrantEdgeResult<()> { let Some(shard) = self.get_existing_store(store_name)? else { return Ok(()); @@ -296,6 +366,12 @@ pub async fn insert_qdrant_edge_embedding(_token: APIToken, Json(request): Json< }) } +pub async fn search_qdrant_edge_embeddings(_token: APIToken, Json(request): Json) -> Json { + execute_qdrant_edge_query(|database| { + database.search_embedding(&request.store_name, request.vector, request.max_matches) + }) +} + pub async fn delete_qdrant_edge_embedding_by_file(_token: APIToken, Json(request): Json) -> Json { execute_qdrant_edge_operation(|database| { database.delete_embedding_by_file(&request.store_name, &request.file_path) @@ -374,6 +450,37 @@ where } } +fn execute_qdrant_edge_query(operation: F) -> Json +where + F: FnOnce(&mut QdrantEdgeDatabase) -> QdrantEdgeResult>, +{ + let mut database_guard = QDRANT_EDGE_DATABASE.lock().unwrap(); + let Some(database) = database_guard.as_mut() else { + return Json(QdrantEdgeSearchResponse { + success: false, + issue: "Qdrant Edge is not available.".to_string(), + data: vec![], + }); + }; + + match operation(database) { + Ok(data) => Json(QdrantEdgeSearchResponse { + success: true, + issue: String::new(), + data, + }), + Err(e) => { + let issue = e.to_string(); + error!(Source = "Qdrant Edge"; "Qdrant Edge query failed: {issue}"); + Json(QdrantEdgeSearchResponse { + success: false, + issue, + data: vec![], + }) + }, + } +} + fn set_qdrant_edge_available() { let mut status = QDRANT_EDGE_STATUS.lock().unwrap(); status.status = QdrantEdgeStatus::Available; @@ -538,12 +645,63 @@ fn to_qdrant_edge_point(point: QdrantEdgeStoragePoint) -> qdrant_edge::PointStru .into() } +fn to_qdrant_edge_search_result(point: ScoredPoint) -> QdrantEdgeSearchResult { + let payload = point.payload.unwrap_or_default(); + QdrantEdgeSearchResult { + point_id: point_id_to_string(point.id), + score: point.score, + data_source_id: payload_string(&payload, "data_source_id"), + data_source_name: payload_string(&payload, "data_source_name"), + data_source_type: payload_string(&payload, "data_source_type"), + chunk_id: payload_string(&payload, "chunk_id"), + parent_file_id: payload_string(&payload, "parent_file_id"), + file_path: payload_string(&payload, "file_path"), + absolute_path: payload_string(&payload, "absolute_path"), + file_name: payload_string(&payload, "file_name"), + relative_path: payload_string(&payload, "relative_path"), + file_type: payload_string(&payload, "file_type"), + page_number: payload_i32(&payload, "page_number"), + chunk_index: payload_i32(&payload, "chunk_index").unwrap_or_default(), + text: payload_string(&payload, "text"), + fingerprint: payload_string(&payload, "fingerprint"), + creation_utc: payload_string(&payload, "creation_utc"), + last_write_utc: payload_string(&payload, "last_write_utc"), + embedded_at_utc: payload_string(&payload, "embedded_at_utc"), + compliance_level: payload_string(&payload, "compliance_level"), + compliance_level_rank: payload_i32(&payload, "compliance_level_rank").unwrap_or_default(), + } +} + fn to_point_id(point_id: &str) -> PointId { Uuid::parse_str(point_id) .map(PointId::Uuid) .unwrap_or_else(|_| PointId::NumId(stable_u64(point_id))) } +fn point_id_to_string(point_id: PointId) -> String { + match point_id { + PointId::NumId(id) => id.to_string(), + PointId::Uuid(uuid) => uuid.to_string(), + } +} + +fn payload_string(payload: &Payload, key: &str) -> String { + payload + .0 + .get(key) + .and_then(Value::as_str) + .unwrap_or_default() + .to_string() +} + +fn payload_i32(payload: &Payload, key: &str) -> Option { + payload + .0 + .get(key) + .and_then(Value::as_i64) + .and_then(|value| i32::try_from(value).ok()) +} + fn stable_u64(value: &str) -> u64 { let mut hash = 0xcbf29ce484222325_u64; for byte in value.as_bytes() { diff --git a/runtime/src/runtime_api.rs b/runtime/src/runtime_api.rs index f50913f2..c52affea 100644 --- a/runtime/src/runtime_api.rs +++ b/runtime/src/runtime_api.rs @@ -36,6 +36,7 @@ pub fn start_runtime_api() { .route("/system/qdrant-edge/info", get(crate::qdrant_edge_database::qdrant_edge_info)) .route("/system/qdrant-edge/ensure", post(crate::qdrant_edge_database::ensure_qdrant_edge_store)) .route("/system/qdrant-edge/insert", post(crate::qdrant_edge_database::insert_qdrant_edge_embedding)) + .route("/system/qdrant-edge/search", post(crate::qdrant_edge_database::search_qdrant_edge_embeddings)) .route("/system/qdrant-edge/delete-file", post(crate::qdrant_edge_database::delete_qdrant_edge_embedding_by_file)) .route("/system/qdrant-edge/delete-store", post(crate::qdrant_edge_database::delete_qdrant_edge_store)) .route("/clipboard/set", post(crate::clipboard::set_clipboard))