Merge branch 'main' into chunk-data

This commit is contained in:
Paul Koudelka committed 2026-06-10 14:11:44 +02:00
commit b7b18aa45c
109 files changed
+6469 -1691

No files matched your search

@@ -91,17 +91,10 @@ public sealed partial class DatabaseClientProvider(RustService rustService, ILog
private async Task<DatabaseClient> CreateClientAsync(DatabaseRole databaseRole, CancellationToken cancellationToken) => databaseRole switch
{
DatabaseRole.VECTOR_STORE => await this.CreateQdrantClientAsync(cancellationToken),
DatabaseRole.VECTOR_STORE => await QdrantEdgeClientImplementation.CreateAsync(rustService, this.logger, this.databaseClientLogger, cancellationToken),
_ => new NoDatabaseClient(databaseRole.ToString(), "The requested database role is not supported.")
};
private NoDatabaseClient CreateNoDatabaseClient(string name, string? unavailableReason, DatabaseClientStatus status)
{
var client = new NoDatabaseClient(name, unavailableReason, status);
client.SetLogger(this.databaseClientLogger);
return client;
}
private static bool IsSameClient(DatabaseClient left, DatabaseClient right) =>
left.IsAvailable
&& right.IsAvailable
@@ -2,14 +2,6 @@
public interface IVectorStoreClient
{
string Name { get; }
DatabaseClientStatus Status { get; }
bool IsAvailable { get; }
IAsyncEnumerable<(string Label, string Value)> GetDisplayInfo();
Task EnsureVectorStoreExists(string storeName, int vectorSize, CancellationToken token);
Task InsertEmbedding(string storeName, IReadOnlyList<VectorStoragePoint> points, CancellationToken token);
@@ -2,19 +2,19 @@ using AIStudio.Tools.PluginSystem;
namespace AIStudio.Tools.Databases.VectorStore;
public sealed class NoVectorStoreClient(string name, string? unavailableReason, DatabaseClientStatus status = DatabaseClientStatus.UNAVAILABLE) : IVectorStoreClient
public sealed class NoVectorStoreClient(string name, string? unavailableReason, DatabaseClientStatus status = DatabaseClientStatus.UNAVAILABLE) : DatabaseClient(name, string.Empty), IVectorStoreClient
{
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(NoVectorStoreClient).Namespace, nameof(NoVectorStoreClient));
public string Name => name;
public override DatabaseClientStatus Status => status;
public DatabaseClientStatus Status => status;
public bool IsAvailable => false;
public async IAsyncEnumerable<(string Label, string Value)> GetDisplayInfo()
public override async IAsyncEnumerable<(string Label, string Value)> GetDisplayInfo()
{
yield return (TB("Status"), TB("Unavailable"));
yield return (TB("Status"), status switch
{
DatabaseClientStatus.STARTING => TB("Starting"),
_ => TB("Unavailable")
});
if (!string.IsNullOrWhiteSpace(unavailableReason))
yield return (TB("Reason"), unavailableReason);
@@ -36,4 +36,8 @@ public sealed class NoVectorStoreClient(string name, string? unavailableReason,
private InvalidOperationException CreateUnavailableException() =>
new(unavailableReason ?? "The vector store is not available.");
public override void Dispose()
{
}
}
@@ -0,0 +1,115 @@
using AIStudio.Tools.PluginSystem;
using AIStudio.Tools.Rust;
using AIStudio.Tools.Services;
namespace AIStudio.Tools.Databases.VectorStore;
public sealed class QdrantEdgeClientImplementation(
string name,
string path,
string version,
int storesCount,
RustService rustService) : DatabaseClient(name, path), IVectorStoreClient
{
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 DELETE_FILE_PATH = "/system/qdrant-edge/delete-file";
private const string DELETE_STORE_PATH = "/system/qdrant-edge/delete-store";
private readonly string path = path;
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(QdrantEdgeClientImplementation).Namespace, nameof(QdrantEdgeClientImplementation));
public override string CacheKey => $"{this.Name}:{this.path}:{version}";
public static async Task<DatabaseClient> CreateAsync(
RustService rustService,
ILogger logger,
ILogger<DatabaseClient> databaseClientLogger,
CancellationToken cancellationToken)
{
var qdrantEdgeInfo = await rustService.GetDatabaseInfo(
DATABASE_NAME,
INFO_PATH,
QdrantEdgeInfo.Unavailable,
cancellationToken);
if (qdrantEdgeInfo.Status is QdrantEdgeStatus.STARTING)
{
return CreateNoVectorStoreClient(
DATABASE_NAME,
$"{DATABASE_NAME} is starting. Details will appear shortly.",
DatabaseClientStatus.STARTING,
databaseClientLogger);
}
if (!qdrantEdgeInfo.IsAvailable || qdrantEdgeInfo.Status is QdrantEdgeStatus.UNAVAILABLE)
{
var reason = qdrantEdgeInfo.UnavailableReason ?? "unknown";
// ReSharper disable DuplicateItemInLoggerTemplate
logger.LogWarning("{VectorStoreName} is not available. Starting without {VectorStoreName} vector store. Reason: '{Reason}'.", DATABASE_NAME, DATABASE_NAME, reason);
// ReSharper restore DuplicateItemInLoggerTemplate
return CreateNoVectorStoreClient(DATABASE_NAME, qdrantEdgeInfo.UnavailableReason, DatabaseClientStatus.UNAVAILABLE, databaseClientLogger);
}
if (qdrantEdgeInfo.Path == string.Empty)
return CreateNoVectorStoreClient(DATABASE_NAME, $"Failed to get the {DATABASE_NAME} path from Rust.", DatabaseClientStatus.UNAVAILABLE, databaseClientLogger);
var name = string.IsNullOrWhiteSpace(qdrantEdgeInfo.Name) ? DATABASE_NAME : qdrantEdgeInfo.Name;
var client = new QdrantEdgeClientImplementation(name, qdrantEdgeInfo.Path, qdrantEdgeInfo.Version, qdrantEdgeInfo.StoresCount, rustService);
client.SetLogger(databaseClientLogger);
return client;
}
public override async IAsyncEnumerable<(string Label, string Value)> GetDisplayInfo()
{
var currentInfo = await rustService.GetDatabaseInfo(
DATABASE_NAME,
INFO_PATH,
QdrantEdgeInfo.Unavailable);
var displayVersion = currentInfo.IsAvailable && !string.IsNullOrWhiteSpace(currentInfo.Version) ? currentInfo.Version : version;
var displayStoresCount = currentInfo.IsAvailable ? currentInfo.StoresCount : storesCount;
if (!currentInfo.IsAvailable)
yield return (TB("Status"), currentInfo.UnavailableReason ?? TB("Qdrant Edge is not available."));
yield return (TB("Reported version"), displayVersion);
yield return (TB("Storage size"), $"{this.GetStorageSize()}");
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 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) =>
rustService.ExecuteDatabaseOperation(DATABASE_NAME, DELETE_FILE_PATH, new DeleteEmbeddingByFileRequest(storeName, filePath), token);
public Task DeleteVectorStore(string storeName, CancellationToken token) =>
rustService.ExecuteDatabaseOperation(DATABASE_NAME, DELETE_STORE_PATH, new DeleteVectorStoreRequest(storeName), token);
public override void Dispose()
{
}
private static NoVectorStoreClient CreateNoVectorStoreClient(string name, string? unavailableReason, DatabaseClientStatus status, ILogger<DatabaseClient> databaseClientLogger)
{
var client = new NoVectorStoreClient(name, unavailableReason, status);
client.SetLogger(databaseClientLogger);
return client;
}
// ReSharper disable NotAccessedPositionalProperty.Local
private sealed record EnsureVectorStoreRequest(string StoreName, int VectorSize);
private sealed record InsertEmbeddingRequest(string StoreName, IReadOnlyList<VectorStoragePoint> Points);
private sealed record DeleteEmbeddingByFileRequest(string StoreName, string FilePath);
private sealed record DeleteVectorStoreRequest(string StoreName);
// ReSharper restore NotAccessedPositionalProperty.Local
}
@@ -0,0 +1,16 @@
namespace AIStudio.Tools.Databases.VectorStore;
public sealed record VectorStoragePoint(
string PointId,
IReadOnlyList<float> Vector,
string DataSourceId,
string DataSourceName,
string DataSourceType,
string FilePath,
string FileName,
string RelativePath,
int ChunkIndex,
string Text,
string Fingerprint,
DateTime LastWriteUtc,
DateTime EmbeddedAtUtc);
@@ -23,7 +23,7 @@ public abstract class ERIClientBase(IERIDataSource dataSource) : IDisposable
}
};
protected readonly HttpClient HttpClient = ExternalHttpClientTimeout.CreateHttpClient(new Uri($"{dataSource.Hostname}:{dataSource.Port}"));
protected readonly HttpClient HttpClient = ExternalHttpClientTimeout.CreateHttpClient(new Uri($"{dataSource.Hostname}:{dataSource.Port}"), ExternalHttpTrustPolicy.ALLOW_CUSTOM_ROOTS_WHEN_HOST_WHITELISTED);
protected string SecurityToken = string.Empty;
@@ -2,7 +2,7 @@ using System.Net.Http.Headers;
namespace AIStudio.Tools;
public readonly record struct EnterpriseEnvironment(string ConfigurationServerUrl, Guid ConfigurationId, EntityTagHeaderValue? ETag)
public readonly record struct EnterpriseEnvironment(string ConfigurationServerUrl, Guid ConfigurationId, string Source, string SourceDetail, string Slot, EntityTagHeaderValue? ETag)
{
public bool IsActive => !string.IsNullOrWhiteSpace(this.ConfigurationServerUrl) && this.ConfigurationId != Guid.Empty;
}
+10
View File
@@ -155,6 +155,16 @@ public enum Event
/// Requests the chat workspace overlay to be toggled.
/// </summary>
WORKSPACE_TOGGLE_OVERLAY,
/// <summary>
/// Notifies receivers that a workspace was renamed.
/// </summary>
WORKSPACE_RENAMED,
/// <summary>
/// Notifies receivers that a workspace was created.
/// </summary>
WORKSPACE_CREATED,
@@ -1,3 +1,7 @@
using System.Net.Security;
using System.Security.Cryptography;
using System.Security.Cryptography.X509Certificates;
using AIStudio.Settings;
namespace AIStudio.Tools;
@@ -12,15 +16,38 @@ public static class ExternalHttpClientTimeout
public const int MAX_HTTP_CLIENT_TIMEOUT_SECONDS = 3600;
public const int DEFAULT_HTTP_CLIENT_TIMEOUT_SECONDS = 3600;
private static readonly Lazy<SettingsManager> SETTINGS_MANAGER = new(() => Program.SERVICE_PROVIDER.GetRequiredService<SettingsManager>());
private const string ENV_CUSTOM_ROOT_CERTIFICATES_ENABLED = "MINDWORK_AI_STUDIO_EXTERNAL_HTTP_CUSTOM_ROOT_CERTIFICATES_ENABLED";
private const string ENV_CUSTOM_ROOT_CERTIFICATE_BUNDLE_PATH = "MINDWORK_AI_STUDIO_EXTERNAL_HTTP_CUSTOM_ROOT_CERTIFICATE_BUNDLE_PATH";
private const string ENV_CUSTOM_ROOT_CERTIFICATE_ALLOWED_HOSTS = "MINDWORK_AI_STUDIO_EXTERNAL_HTTP_CUSTOM_ROOT_CERTIFICATE_ALLOWED_HOSTS";
public static HttpClient CreateHttpClient(Uri? baseAddress = null)
// id-kp-serverAuth: Extended Key Usage for TLS server authentication.
// See RFC 5280, section 4.2.1.12: https://www.rfc-editor.org/rfc/rfc5280#section-4.2.1.12
private const string TLS_SERVER_AUTHENTICATION_EKU_OID = "1.3.6.1.5.5.7.3.1";
private static string TB(string fallbackEN) => PluginSystem.I18N.I.T(fallbackEN, typeof(ExternalHttpClientTimeout).Namespace, nameof(ExternalHttpClientTimeout));
private static readonly Lazy<ILogger> LOGGER = new(() => Program.LOGGER_FACTORY.CreateLogger(nameof(ExternalHttpClientTimeout)));
private static readonly Lazy<SettingsManager> SETTINGS_MANAGER = new(() => Program.SERVICE_PROVIDER.GetRequiredService<SettingsManager>());
private static readonly Lock CUSTOM_ROOT_CERTIFICATE_LOCK = new();
private static CustomRootCertificateCache? CUSTOM_ROOT_CERTIFICATE_CACHE;
public static HttpClient CreateHttpClient(ExternalHttpTrustPolicy trustPolicy) => CreateHttpClient(null, trustPolicy);
public static HttpClient CreateHttpClient(Uri? baseAddress, ExternalHttpTrustPolicy trustPolicy)
{
var httpClient = new HttpClient();
var customRootCertificateCache = GetCustomRootCertificateCache();
var httpClient = customRootCertificateCache.State.IsUsable
? new HttpClient(new HttpClientHandler
{
ServerCertificateCustomValidationCallback = (request, certificate, chain, sslPolicyErrors) =>
ValidateServerCertificateWithCustomRootCertificates(request, certificate, chain, sslPolicyErrors, customRootCertificateCache, trustPolicy)
})
: new HttpClient();
Configure(httpClient, baseAddress);
return httpClient;
}
public static ExternalHttpCustomRootCertificateState CustomRootCertificateState => GetCustomRootCertificateCache().State;
public static string GetTimeoutDescription()
{
var timeout = GetTimeout();
@@ -78,4 +105,364 @@ public static class ExternalHttpClientTimeout
if (baseAddress is not null)
httpClient.BaseAddress = baseAddress;
}
private static CustomRootCertificateCache GetCustomRootCertificateCache()
{
var configuration = ReadCustomRootCertificateConfiguration();
var cacheKey = $"{configuration.Enabled}|{configuration.BundlePath}|{string.Join(";", configuration.AllowedHostPatterns)}|{ReadCertificateBundleFileSignature(configuration.BundlePath)}";
lock (CUSTOM_ROOT_CERTIFICATE_LOCK)
{
if (CUSTOM_ROOT_CERTIFICATE_CACHE is not null && CUSTOM_ROOT_CERTIFICATE_CACHE.CacheKey == cacheKey)
return CUSTOM_ROOT_CERTIFICATE_CACHE;
CUSTOM_ROOT_CERTIFICATE_CACHE = LoadCustomRootCertificateCache(cacheKey, configuration);
LogCustomRootCertificateState(CUSTOM_ROOT_CERTIFICATE_CACHE.State);
return CUSTOM_ROOT_CERTIFICATE_CACHE;
}
}
private static CustomRootCertificateConfiguration ReadCustomRootCertificateConfiguration()
{
var envEnabled = Environment.GetEnvironmentVariable(ENV_CUSTOM_ROOT_CERTIFICATES_ENABLED);
var envBundlePath = Environment.GetEnvironmentVariable(ENV_CUSTOM_ROOT_CERTIFICATE_BUNDLE_PATH);
var envAllowedHosts = Environment.GetEnvironmentVariable(ENV_CUSTOM_ROOT_CERTIFICATE_ALLOWED_HOSTS);
var enabled = TryParseBooleanEnvironmentValue(envEnabled, out var parsedEnvEnabled)
? parsedEnvEnabled
: SETTINGS_MANAGER.Value.ConfigurationData.App.ExternalHttpCustomRootCertificatesEnabled;
var bundlePath = !string.IsNullOrWhiteSpace(envBundlePath)
? envBundlePath.Trim()
: SETTINGS_MANAGER.Value.ConfigurationData.App.ExternalHttpCustomRootCertificateBundlePath.Trim();
var allowedHostPatterns = ReadAllowedHostPatterns(envAllowedHosts);
var source = ReadCustomRootCertificateConfigurationSource(envEnabled, envBundlePath, envAllowedHosts);
return new(enabled, bundlePath, allowedHostPatterns, source);
}
private static string ReadCustomRootCertificateConfigurationSource(string? envEnabled, string? envBundlePath, string? envAllowedHosts)
{
if (!string.IsNullOrWhiteSpace(envEnabled) || !string.IsNullOrWhiteSpace(envBundlePath) || !string.IsNullOrWhiteSpace(envAllowedHosts))
return TB("environment variables");
var enabledIsManaged = ManagedConfiguration.TryGet(x => x.App, x => x.ExternalHttpCustomRootCertificatesEnabled, out var enabledMeta) && enabledMeta.IsLocked;
var bundlePathIsManaged = ManagedConfiguration.TryGet(x => x.App, x => x.ExternalHttpCustomRootCertificateBundlePath, out var bundlePathMeta) && bundlePathMeta.IsLocked;
var allowedHostsIsManaged = ManagedConfiguration.TryGet(x => x.App, x => x.ExternalHttpCustomRootCertificateAllowedHosts, out var allowedHostsMeta) && allowedHostsMeta.IsLocked;
return enabledIsManaged || bundlePathIsManaged || allowedHostsIsManaged
? TB("configuration plugin")
: TB("app settings");
}
private static IReadOnlyList<string> ReadAllowedHostPatterns(string? envAllowedHosts)
{
IEnumerable<string> rawPatterns = !string.IsNullOrWhiteSpace(envAllowedHosts)
? envAllowedHosts.Split([';', ','], StringSplitOptions.RemoveEmptyEntries | StringSplitOptions.TrimEntries)
: SETTINGS_MANAGER.Value.ConfigurationData.App.ExternalHttpCustomRootCertificateAllowedHosts;
var patterns = new HashSet<string>(StringComparer.OrdinalIgnoreCase);
foreach (var rawPattern in rawPatterns)
{
if (TryNormalizeAllowedHostPattern(rawPattern, out var pattern))
patterns.Add(pattern);
else
LOGGER.Value.LogWarning($"Ignoring invalid external HTTP custom root certificate host pattern: '{rawPattern}'.");
}
return patterns.Order(StringComparer.OrdinalIgnoreCase).ToList();
}
private static bool TryNormalizeAllowedHostPattern(string? rawPattern, out string pattern)
{
pattern = string.Empty;
if (string.IsNullOrWhiteSpace(rawPattern))
return false;
var normalized = rawPattern.Trim().TrimEnd('.').ToLowerInvariant();
if (normalized.Contains("://", StringComparison.Ordinal) || normalized.Contains('/', StringComparison.Ordinal) || normalized.Contains(':', StringComparison.Ordinal))
return false;
if (normalized.StartsWith("*.", StringComparison.Ordinal))
{
var suffix = normalized[2..];
if (!IsValidDnsHost(suffix))
return false;
pattern = $"*.{suffix}";
return true;
}
if (normalized.Contains('*', StringComparison.Ordinal))
return false;
if (!IsValidDnsHost(normalized))
return false;
pattern = normalized;
return true;
}
private static bool IsValidDnsHost(string host)
{
if (string.IsNullOrWhiteSpace(host))
return false;
if (Uri.CheckHostName(host) is not UriHostNameType.Dns)
return false;
return host.Split('.').All(label => !string.IsNullOrWhiteSpace(label) && !label.StartsWith('-') && !label.EndsWith('-'));
}
private static string ReadCertificateBundleFileSignature(string bundlePath)
{
if (string.IsNullOrWhiteSpace(bundlePath))
return string.Empty;
try
{
var fileInfo = new FileInfo(bundlePath);
return fileInfo.Exists
? $"{fileInfo.Length}|{fileInfo.LastWriteTimeUtc.Ticks}"
: "missing";
}
catch
{
return "unavailable";
}
}
private static bool TryParseBooleanEnvironmentValue(string? value, out bool parsedValue)
{
parsedValue = false;
if (string.IsNullOrWhiteSpace(value))
return false;
var normalized = value.Trim();
if (bool.TryParse(normalized, out parsedValue))
return true;
if (normalized is "1" || normalized.Equals("yes", StringComparison.OrdinalIgnoreCase) || normalized.Equals("on", StringComparison.OrdinalIgnoreCase))
{
parsedValue = true;
return true;
}
if (normalized is "0" || normalized.Equals("no", StringComparison.OrdinalIgnoreCase) || normalized.Equals("off", StringComparison.OrdinalIgnoreCase))
{
parsedValue = false;
return true;
}
return false;
}
private static CustomRootCertificateCache LoadCustomRootCertificateCache(string cacheKey, CustomRootCertificateConfiguration configuration)
{
var certificates = new X509Certificate2Collection();
if (!configuration.Enabled)
{
return new(
cacheKey,
certificates,
new ExternalHttpCustomRootCertificateState(false, configuration.Source, configuration.BundlePath, configuration.AllowedHostPatterns, false, 0, [], string.Empty));
}
if (string.IsNullOrWhiteSpace(configuration.BundlePath))
{
return new(
cacheKey,
certificates,
new ExternalHttpCustomRootCertificateState(true, configuration.Source, configuration.BundlePath, configuration.AllowedHostPatterns, false, 0, [], TB("No certificate bundle path is configured.")));
}
if (!File.Exists(configuration.BundlePath))
{
return new(
cacheKey,
certificates,
new ExternalHttpCustomRootCertificateState(true, configuration.Source, configuration.BundlePath, configuration.AllowedHostPatterns, false, 0, [], TB("The configured certificate bundle file does not exist.")));
}
try
{
var importedCertificates = new X509Certificate2Collection();
importedCertificates.ImportFromPemFile(configuration.BundlePath);
foreach (var certificate in importedCertificates)
{
if (!IsRootCertificateAuthority(certificate))
continue;
certificates.Add(certificate);
}
var fingerprints = certificates
.Select(certificate => certificate.GetCertHashString(HashAlgorithmName.SHA256))
.Order(StringComparer.OrdinalIgnoreCase)
.ToList();
var issue = certificates.Count == 0
? TB("The configured certificate bundle does not contain usable root CA certificates.")
: string.Empty;
return new(
cacheKey,
certificates,
new ExternalHttpCustomRootCertificateState(true, configuration.Source, configuration.BundlePath, configuration.AllowedHostPatterns, certificates.Count > 0, certificates.Count, fingerprints, issue));
}
catch (Exception e)
{
return new(
cacheKey,
certificates,
new ExternalHttpCustomRootCertificateState(true, configuration.Source, configuration.BundlePath, configuration.AllowedHostPatterns, false, 0, [], e.Message));
}
}
private static bool IsRootCertificateAuthority(X509Certificate2 certificate)
{
if (!certificate.SubjectName.RawData.SequenceEqual(certificate.IssuerName.RawData))
return false;
return certificate.Extensions
.OfType<X509BasicConstraintsExtension>()
.Any(extension => extension.CertificateAuthority);
}
private static bool ValidateServerCertificateWithCustomRootCertificates(
HttpRequestMessage request,
X509Certificate? certificate,
X509Chain? originalChain,
SslPolicyErrors sslPolicyErrors,
CustomRootCertificateCache customRootCertificateCache,
ExternalHttpTrustPolicy trustPolicy)
{
if (sslPolicyErrors is SslPolicyErrors.None)
return true;
if (sslPolicyErrors is not SslPolicyErrors.RemoteCertificateChainErrors || certificate is null)
return false;
var host = ReadRequestHost(request);
if (trustPolicy is ExternalHttpTrustPolicy.SYSTEM_TRUST_ONLY)
{
LOGGER.Value.LogError($"Rejected external HTTPS certificate for '{HostForLog(host)}' because this request requires system trust only. Configured custom root certificates are not allowed for this request.");
return false;
}
if (!IsAllowedCustomRootCertificateHost(host, customRootCertificateCache.State.AllowedHostPatterns))
{
LOGGER.Value.LogError($"Rejected external HTTPS certificate for '{HostForLog(host)}' because the host is not allowed to use configured custom root certificates.");
return false;
}
var ownsServerCertificate = certificate is not X509Certificate2;
var serverCertificate = certificate as X509Certificate2 ?? new X509Certificate2(certificate);
try
{
using var customChain = new X509Chain();
customChain.ChainPolicy.TrustMode = X509ChainTrustMode.CustomRootTrust;
customChain.ChainPolicy.CustomTrustStore.AddRange(customRootCertificateCache.Certificates);
customChain.ChainPolicy.ApplicationPolicy.Add(new Oid(TLS_SERVER_AUTHENTICATION_EKU_OID));
if (originalChain is not null)
{
foreach (var element in originalChain.ChainElements)
{
if (element.Certificate.Thumbprint == serverCertificate.Thumbprint)
continue;
customChain.ChainPolicy.ExtraStore.Add(element.Certificate);
}
}
var isValid = customChain.Build(serverCertificate);
if (isValid)
LogCustomRootCertificateAccepted(request);
return isValid;
}
finally
{
if (ownsServerCertificate)
serverCertificate.Dispose();
}
}
private static bool IsAllowedCustomRootCertificateHost(string host, IReadOnlyList<string> allowedHostPatterns)
{
if (string.IsNullOrWhiteSpace(host))
return false;
var normalizedHost = host.Trim().TrimEnd('.').ToLowerInvariant();
foreach (var pattern in allowedHostPatterns)
{
if (!pattern.StartsWith("*.", StringComparison.Ordinal))
{
if (normalizedHost.Equals(pattern, StringComparison.OrdinalIgnoreCase))
return true;
continue;
}
var suffix = pattern[2..];
if (!normalizedHost.EndsWith($".{suffix}", StringComparison.OrdinalIgnoreCase))
continue;
var prefix = normalizedHost[..^(suffix.Length + 1)];
if (!prefix.Contains('.', StringComparison.Ordinal))
return true;
}
return false;
}
private static void LogCustomRootCertificateState(ExternalHttpCustomRootCertificateState state)
{
if (!state.IsEnabled)
{
LOGGER.Value.LogInformation("External HTTP custom root certificates are disabled.");
return;
}
if (state.IsUsable)
{
LOGGER.Value.LogWarning($"External HTTP custom root certificates are enabled from {state.Source}. Loaded {state.CertificateCount} root certificate(s) from '{state.BundlePath}'. Allowed hosts: {FormatAllowedHostPatternsForLog(state.AllowedHostPatterns)}. Fingerprints: {string.Join(", ", state.CertificateFingerprints)}");
return;
}
LOGGER.Value.LogWarning($"External HTTP custom root certificates are enabled from {state.Source}, but no additional root certificates are usable. Bundle path: '{state.BundlePath}'. Issue: {state.Issue}");
}
private static void LogCustomRootCertificateAccepted(HttpRequestMessage request)
{
var host = ReadRequestHost(request);
LOGGER.Value.LogWarning($"Accepted an external HTTPS certificate for '{host}' using configured custom root certificates.");
}
private static string ReadRequestHost(HttpRequestMessage request)
{
var host = request.RequestUri?.IdnHost;
if (string.IsNullOrWhiteSpace(host))
host = request.RequestUri?.Host;
return host ?? string.Empty;
}
private static string HostForLog(string host) => string.IsNullOrWhiteSpace(host) ? "unknown host" : host;
private static string FormatAllowedHostPatternsForLog(IReadOnlyList<string> allowedHostPatterns)
{
if (allowedHostPatterns.Count == 0)
return "none";
return string.Join(", ", allowedHostPatterns);
}
private readonly record struct CustomRootCertificateConfiguration(bool Enabled, string BundlePath, IReadOnlyList<string> AllowedHostPatterns, string Source);
private sealed record CustomRootCertificateCache(
string CacheKey,
X509Certificate2Collection Certificates,
ExternalHttpCustomRootCertificateState State);
}
@@ -0,0 +1,11 @@
namespace AIStudio.Tools;
public sealed record ExternalHttpCustomRootCertificateState(
bool IsEnabled,
string Source,
string BundlePath,
IReadOnlyList<string> AllowedHostPatterns,
bool IsUsable,
int CertificateCount,
IReadOnlyList<string> CertificateFingerprints,
string Issue);
@@ -0,0 +1,7 @@
namespace AIStudio.Tools;
public enum ExternalHttpTrustPolicy
{
SYSTEM_TRUST_ONLY,
ALLOW_CUSTOM_ROOTS_WHEN_HOST_WHITELISTED
}
@@ -5,7 +5,7 @@ public class I18N : ILang
public static readonly I18N I = new();
private static readonly ILogger<I18N> LOG = Program.LOGGER_FACTORY.CreateLogger<I18N>();
private ILanguagePlugin? language = PluginFactory.BaseLanguage;
private ILanguagePlugin? language;
private I18N()
{
@@ -85,6 +85,14 @@ public abstract partial class PluginBase : IPluginMetadata
if(!string.IsNullOrWhiteSpace(parseError))
issues.Add(parseError);
if (this is NoPlugin or NoPluginLanguage)
{
this.IsInternal = isInternal;
this.IconSVG = string.Empty;
this.baseIssues = issues;
return;
}
// Notice: when no icon is specified, the default icon will be used.
this.TryInitIconSVG(out _, out var iconSVG);
this.IconSVG = iconSVG;
@@ -155,6 +155,9 @@ public sealed class PluginConfiguration(bool isInternal, LuaState state, PluginT
// Config: what should be the start page?
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.StartPage, this.Id, settingsTable, dryRun);
// Config: show quick start guide on the home page?
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.ShowQuickStartGuide, this.Id, settingsTable, dryRun);
// Config: allow the user to add providers?
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.AllowUserToAddProvider, this.Id, settingsTable, dryRun);
@@ -176,6 +179,11 @@ public sealed class PluginConfiguration(bool isInternal, LuaState state, PluginT
// Config: timeout for external HTTP requests
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.HttpClientTimeoutSeconds, this.Id, settingsTable, dryRun);
// Config: custom root certificates for external HTTP requests
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.ExternalHttpCustomRootCertificatesEnabled, this.Id, settingsTable, dryRun);
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.ExternalHttpCustomRootCertificateBundlePath, this.Id, settingsTable, dryRun);
ManagedConfiguration.TryProcessConfiguration(x => x.App, x => x.ExternalHttpCustomRootCertificateAllowedHosts, this.Id, settingsTable, dryRun);
// Handle configured LLM providers:
PluginConfigurationObject.TryParse(PluginConfigurationObjectType.LLM_PROVIDER, x => x.Providers, x => x.NextProviderNum, mainTable, this.Id, ref this.configObjects, dryRun);
@@ -15,7 +15,7 @@ public static partial class PluginFactory
var serverUrl = configServerUrl.EndsWith('/') ? configServerUrl[..^1] : configServerUrl;
var downloadUrl = $"{serverUrl}/{configPlugId}.zip";
using var http = ExternalHttpClientTimeout.CreateHttpClient();
using var http = ExternalHttpClientTimeout.CreateHttpClient(ExternalHttpTrustPolicy.ALLOW_CUSTOM_ROOTS_WHEN_HOST_WHITELISTED);
using var request = new HttpRequestMessage(HttpMethod.Get, downloadUrl);
var response = await http.SendAsync(request, HttpCompletionOption.ResponseHeadersRead, cancellationToken);
if (!response.IsSuccessStatusCode)
@@ -52,7 +52,7 @@ public static partial class PluginFactory
try
{
await LockHotReloadAsync();
using var httpClient = ExternalHttpClientTimeout.CreateHttpClient();
using var httpClient = ExternalHttpClientTimeout.CreateHttpClient(ExternalHttpTrustPolicy.ALLOW_CUSTOM_ROOTS_WHEN_HOST_WHITELISTED);
var response = await httpClient.GetAsync(downloadUrl, cancellationToken);
if (!response.IsSuccessStatusCode)
{
@@ -213,6 +213,10 @@ public static partial class PluginFactory
// Check for the start page:
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.App, x => x.StartPage, AVAILABLE_PLUGINS))
wasConfigurationChanged = true;
// Check for the quick start guide visibility:
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.App, x => x.ShowQuickStartGuide, AVAILABLE_PLUGINS))
wasConfigurationChanged = true;
// Check for users allowed to added providers:
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.App, x => x.AllowUserToAddProvider, AVAILABLE_PLUGINS))
@@ -249,6 +253,16 @@ public static partial class PluginFactory
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.App, x => x.HttpClientTimeoutSeconds, AVAILABLE_PLUGINS))
wasConfigurationChanged = true;
// Check for custom root certificates for external HTTP requests:
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.App, x => x.ExternalHttpCustomRootCertificatesEnabled, AVAILABLE_PLUGINS))
wasConfigurationChanged = true;
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.App, x => x.ExternalHttpCustomRootCertificateBundlePath, AVAILABLE_PLUGINS))
wasConfigurationChanged = true;
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.App, x => x.ExternalHttpCustomRootCertificateAllowedHosts, AVAILABLE_PLUGINS))
wasConfigurationChanged = true;
// Check if audit is required before it can be activated
if(ManagedConfiguration.IsConfigurationLeftOver(x => x.AssistantPluginAudit, x => x.RequireAuditBeforeActivation, AVAILABLE_PLUGINS))
wasConfigurationChanged = true;
@@ -1,3 +1,3 @@
namespace AIStudio.Tools.Rust;
public sealed record EnterpriseConfig(string Id, string ServerUrl);
public sealed record EnterpriseConfig(string Id, string ServerUrl, string Source, string SourceDetail, string Slot);
@@ -68,6 +68,7 @@ public static class FileTypes
public static readonly FileTypeFilter MEDIA = FileTypeFilter.Parent(TB("Media"), IMAGE, AUDIO, VIDEO);
// 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 FileTypeFilter? AsOneFileType(params FileTypeFilter[]? types)
@@ -0,0 +1,27 @@
namespace AIStudio.Tools.Rust;
/// <summary>
/// The response of the Qdrant Edge information request.
/// </summary>
public readonly record struct QdrantEdgeInfo
{
public QdrantEdgeStatus Status { get; init; }
public bool IsAvailable { get; init; }
public string? UnavailableReason { get; init; }
public string Name { get; init; }
public string Version { get; init; }
public string Path { get; init; }
public int StoresCount { get; init; }
public static QdrantEdgeInfo Unavailable(string reason) => new()
{
Status = QdrantEdgeStatus.UNAVAILABLE,
UnavailableReason = reason
};
}
@@ -1,8 +1,8 @@
namespace AIStudio.Tools.Rust;
public enum QdrantStatus
public enum QdrantEdgeStatus
{
STARTING,
AVAILABLE,
UNAVAILABLE,
}
}
@@ -1,23 +0,0 @@
namespace AIStudio.Tools.Rust;
/// <summary>
/// The response of the Qdrant information request.
/// </summary>
public readonly record struct QdrantInfo
{
public QdrantStatus Status { get; init; }
public bool IsAvailable { get; init; }
public string? UnavailableReason { get; init; }
public string Path { get; init; }
public int PortHttp { get; init; }
public int PortGrpc { get; init; }
public string Fingerprint { get; init; }
public string ApiToken { get; init; }
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools.Rust;
public readonly record struct RuntimeInfoResponse(string WorkingDirectory, string ExecutablePath, string LinuxPackageType);
@@ -14,4 +14,4 @@ public enum Shortcut
/// Toggles voice recording on/off.
/// </summary>
VOICE_RECORDING_TOGGLE,
}
}
@@ -14,7 +14,7 @@ public sealed class EnterpriseEnvironmentService(ILogger<EnterpriseEnvironmentSe
private static EnterpriseSecretSnapshot CURRENT_SECRET_SNAPSHOT;
private readonly record struct EnterpriseEnvironmentSnapshot(Guid ConfigurationId, string ConfigurationServerUrl, string? ETag);
private readonly record struct EnterpriseEnvironmentSnapshot(Guid ConfigurationId, string ConfigurationServerUrl, string Source, string SourceDetail, string Slot, string? ETag);
private readonly record struct EnterpriseSecretSnapshot(bool HasSecret, string Fingerprint);
@@ -224,6 +224,9 @@ public sealed class EnterpriseEnvironmentService(ILogger<EnterpriseEnvironmentSe
.Select(environment => new EnterpriseEnvironmentSnapshot(
environment.ConfigurationId,
NormalizeServerUrl(environment.ConfigurationServerUrl),
environment.Source,
environment.SourceDetail,
environment.Slot,
environment.ETag?.ToString()))
.OrderBy(environment => environment.ConfigurationId)
.ToList();
@@ -120,6 +120,12 @@ public sealed partial class RustService
return await response.Content.ReadAsStringAsync();
}
public async Task<RuntimeInfoResponse> GetRuntimeInfo()
{
var response = await this.http.GetFromJsonAsync<RuntimeInfoResponse>("/system/runtime/info", this.jsonRustSerializerOptions);
return response;
}
/// <summary>
/// Requests the Rust runtime to exit the entire desktop application.
/// </summary>
@@ -1,44 +1,53 @@
using AIStudio.Tools.Databases;
using AIStudio.Tools.Rust;
namespace AIStudio.Tools.Services;
public sealed partial class RustService
{
public async Task<QdrantInfo> GetQdrantInfo(CancellationToken cancellationToken = default)
public async Task<TDatabaseInfo> GetDatabaseInfo<TDatabaseInfo>(
string databaseName,
string infoPath,
Func<string, TDatabaseInfo> unavailableFactory,
CancellationToken cancellationToken = default)
{
try
{
using var cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
cts.CancelAfter(TimeSpan.FromSeconds(45));
return await this.http.GetFromJsonAsync<QdrantInfo>("/system/qdrant/info", this.jsonRustSerializerOptions, cts.Token);
var databaseInfo = await this.http.GetFromJsonAsync<TDatabaseInfo>(infoPath, this.jsonRustSerializerOptions, cts.Token);
return databaseInfo ?? unavailableFactory("The database information response was empty.");
}
catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested)
{
if(this.logger is not null)
this.logger.LogWarning("Fetching Qdrant info from Rust service was cancelled by caller.");
this.logger.LogWarning("Fetching {DatabaseName} info from Rust service was cancelled by caller.", databaseName);
else
Console.WriteLine("Fetching Qdrant info from Rust service was cancelled by caller.");
return new QdrantInfo
{
Status = QdrantStatus.UNAVAILABLE,
UnavailableReason = "Operation cancelled by caller."
};
Console.WriteLine($"Fetching {databaseName} info from Rust service was cancelled by caller.");
return unavailableFactory("Operation cancelled by caller.");
}
catch (Exception e)
{
if(this.logger is not null)
this.logger.LogError(e, "Error while fetching Qdrant info from Rust service.");
this.logger.LogError(e, "Error while fetching {DatabaseName} info from Rust service.", databaseName);
else
Console.WriteLine($"Error while fetching Qdrant info from Rust service: '{e}'.");
return new QdrantInfo
{
Status = QdrantStatus.UNAVAILABLE,
UnavailableReason = e.Message
};
Console.WriteLine($"Error while fetching {databaseName} info from Rust service: '{e}'.");
return unavailableFactory(e.Message);
}
}
}
public async Task ExecuteDatabaseOperation<TRequest>(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<DatabaseOperationResponse>(this.jsonRustSerializerOptions, cts.Token);
if (operation is not { Success: true })
throw new InvalidOperationException(operation?.Issue ?? $"The {databaseName} operation failed.");
}
private sealed record DatabaseOperationResponse(bool Success, string Issue);
}
@@ -47,7 +47,7 @@ public sealed partial class RustService
foreach (var config in configs)
{
if (Guid.TryParse(config.Id, out var id))
environments.Add(new EnterpriseEnvironment(config.ServerUrl, id, null));
environments.Add(new EnterpriseEnvironment(config.ServerUrl, id, config.Source, config.SourceDetail, config.Slot, null));
else
this.logger!.LogWarning($"Skipping enterprise config with invalid ID: '{config.Id}'.");
}
@@ -12,6 +12,8 @@ namespace AIStudio.Tools;
public static class WorkspaceBehaviour
{
public readonly record struct TryCreateWorkspaceResult(bool Success, WorkspaceTreeWorkspace Workspace);
private sealed class WorkspaceChatCacheEntry
{
public Guid WorkspaceId { get; init; }
@@ -76,9 +78,9 @@ public static class WorkspaceBehaviour
private static readonly TimeSpan PREFETCH_DELAY_DURATION = TimeSpan.FromMilliseconds(45);
private static string WorkspaceRootDirectory => Path.Join(SettingsManager.DataDirectory, "workspaces");
private static readonly string WORKSPACE_ROOT_DIRECTORY = Path.Join(SettingsManager.DataDirectory, "workspaces");
private static string TemporaryChatsRootDirectory => Path.Join(SettingsManager.DataDirectory, "tempChats");
private static readonly string TEMPORARY_CHATS_ROOT_DIRECTORY = Path.Join(SettingsManager.DataDirectory, "tempChats");
private static SemaphoreSlim GetChatSemaphore(Guid workspaceId, Guid chatId)
{
@@ -156,9 +158,9 @@ public static class WorkspaceBehaviour
private static async Task<List<WorkspaceChatCacheEntry>> ReadTemporaryChatsCoreAsync()
{
var chats = new List<WorkspaceChatCacheEntry>();
Directory.CreateDirectory(TemporaryChatsRootDirectory);
Directory.CreateDirectory(TEMPORARY_CHATS_ROOT_DIRECTORY);
foreach (var tempChatPath in Directory.EnumerateDirectories(TemporaryChatsRootDirectory))
foreach (var tempChatPath in Directory.EnumerateDirectories(TEMPORARY_CHATS_ROOT_DIRECTORY))
{
if (!Guid.TryParse(Path.GetFileName(tempChatPath), out var chatId))
continue;
@@ -188,8 +190,8 @@ public static class WorkspaceBehaviour
WORKSPACE_TREE_CACHE.Workspaces.Clear();
WORKSPACE_TREE_CACHE.WorkspaceOrder.Clear();
Directory.CreateDirectory(WorkspaceRootDirectory);
foreach (var workspacePath in Directory.EnumerateDirectories(WorkspaceRootDirectory))
Directory.CreateDirectory(WORKSPACE_ROOT_DIRECTORY);
foreach (var workspacePath in Directory.EnumerateDirectories(WORKSPACE_ROOT_DIRECTORY))
{
if (!Guid.TryParse(Path.GetFileName(workspacePath), out var workspaceId))
continue;
@@ -230,6 +232,99 @@ public static class WorkspaceBehaviour
chats.RemoveAt(existingIndex);
}
private static IReadOnlyList<string> ParseSearchTerms(string searchText) => searchText
.Split(' ', StringSplitOptions.RemoveEmptyEntries | StringSplitOptions.TrimEntries)
.Where(term => !string.IsNullOrWhiteSpace(term))
.Distinct(StringComparer.OrdinalIgnoreCase)
.ToList();
private static IReadOnlyList<string> GetMissingTerms(string text, IReadOnlyList<string> terms) => terms
.Where(term => text.IndexOf(term, StringComparison.OrdinalIgnoreCase) < 0)
.ToList();
private static bool ChatThreadContainsTerms(ChatThread thread, IReadOnlyList<string> terms)
{
var matchedTerms = new HashSet<string>(StringComparer.OrdinalIgnoreCase);
foreach (var block in thread.Blocks)
{
if (block.HideFromUser || block.Content is not ContentText textContent || string.IsNullOrWhiteSpace(textContent.Text))
continue;
foreach (var term in terms)
if (textContent.Text.Contains(term, StringComparison.OrdinalIgnoreCase))
matchedTerms.Add(term);
if (matchedTerms.Count == terms.Count)
return true;
}
return false;
}
private static bool WorkspaceNameExistsCore(string workspaceName, Guid excludedWorkspaceId = default)
{
return WORKSPACE_TREE_CACHE.Workspaces.Values.Any(workspace =>
workspace.WorkspaceId != excludedWorkspaceId &&
string.Equals(workspace.WorkspaceName.Trim(), workspaceName, StringComparison.OrdinalIgnoreCase));
}
private static async Task<bool> ThreadContainsTermsAsync(WorkspaceTreeChat chat, IReadOnlyList<string> terms, CancellationToken token)
{
var (acquired, semaphore) = await TryAcquireChatSemaphoreAsync(chat.WorkspaceId, chat.ChatId, nameof(ThreadContainsTermsAsync));
if (!acquired)
return false;
try
{
var threadPath = Path.Join(chat.ChatPath, "thread.json");
if (!File.Exists(threadPath))
return false;
var chatData = await File.ReadAllTextAsync(threadPath, Encoding.UTF8, token);
token.ThrowIfCancellationRequested();
var thread = JsonSerializer.Deserialize<ChatThread>(chatData, JSON_OPTIONS);
return thread is not null && ChatThreadContainsTerms(thread, terms);
}
catch (OperationCanceledException)
{
throw;
}
catch (Exception ex)
{
LOG.LogWarning(ex, "Failed to search chat thread for workspace '{WorkspaceId}', chat '{ChatId}'.", chat.WorkspaceId, chat.ChatId);
return false;
}
finally
{
semaphore.Release();
}
}
private static async Task<List<WorkspaceSearchResult>> SearchChatsAsync(IReadOnlyList<WorkspaceTreeChat> chats, IReadOnlyList<string> terms, bool includeThreadContents, CancellationToken token)
{
var results = new List<WorkspaceSearchResult>();
foreach (var chat in chats)
{
token.ThrowIfCancellationRequested();
var missingTerms = GetMissingTerms(chat.Name, terms);
if (missingTerms.Count == 0)
{
results.Add(new(chat, NameMatched: true, ThreadMatched: false));
continue;
}
if (!includeThreadContents)
continue;
var threadMatched = await ThreadContainsTermsAsync(chat, missingTerms, token);
if (threadMatched)
results.Add(new(chat, NameMatched: false, ThreadMatched: true));
}
return results;
}
private static async Task UpdateCacheAfterChatStored(Guid workspaceId, Guid chatId, string chatDirectory, string chatName, DateTimeOffset lastEditTime)
{
await WORKSPACE_TREE_CACHE_SEMAPHORE.WaitAsync();
@@ -348,6 +443,55 @@ public static class WorkspaceBehaviour
}
}
public static async Task<WorkspaceSearchSnapshot> SearchWorkspaceChatsAsync(string searchText, bool includeThreadContents, CancellationToken token = default)
{
var terms = ParseSearchTerms(searchText);
if (terms.Count == 0)
return new([], []);
List<WorkspaceTreeWorkspace> workspaces;
List<WorkspaceTreeChat> temporaryChats;
await WORKSPACE_TREE_CACHE_SEMAPHORE.WaitAsync(token);
try
{
await EnsureTreeShellLoadedCoreAsync();
workspaces = [];
foreach (var workspaceId in WORKSPACE_TREE_CACHE.WorkspaceOrder)
{
token.ThrowIfCancellationRequested();
if (!WORKSPACE_TREE_CACHE.Workspaces.TryGetValue(workspaceId, out var workspace))
continue;
if (!workspace.ChatsLoaded)
{
workspace.Chats = await ReadWorkspaceChatsCoreAsync(workspaceId, workspace.WorkspacePath);
workspace.ChatsLoaded = true;
}
workspaces.Add(ToPublicWorkspace(workspace));
}
temporaryChats = WORKSPACE_TREE_CACHE.TemporaryChats.Select(ToPublicChat).ToList();
}
finally
{
WORKSPACE_TREE_CACHE_SEMAPHORE.Release();
}
var matchingWorkspaces = new List<WorkspaceSearchWorkspace>();
foreach (var workspace in workspaces)
{
token.ThrowIfCancellationRequested();
var matchingChats = await SearchChatsAsync(workspace.Chats, terms, includeThreadContents, token);
if (matchingChats.Count > 0)
matchingWorkspaces.Add(new(workspace.WorkspaceId, workspace.WorkspacePath, workspace.Name, matchingChats));
}
var matchingTemporaryChats = await SearchChatsAsync(temporaryChats, terms, includeThreadContents, token);
return new(matchingWorkspaces, matchingTemporaryChats);
}
public static async Task TryPrefetchRemainingChatsAsync(Func<Guid, Task>? onWorkspaceUpdated = null, CancellationToken token = default)
{
while (true)
@@ -452,6 +596,100 @@ public static class WorkspaceBehaviour
WORKSPACE_TREE_CACHE_SEMAPHORE.Release();
}
}
public static string NormalizeWorkspaceName(string workspaceName) => workspaceName.Trim();
public static async Task<bool> IsWorkspaceNameExistingAsync(string workspaceName, Guid excludedWorkspaceId = default)
{
var normalizedWorkspaceName = NormalizeWorkspaceName(workspaceName);
if (string.IsNullOrWhiteSpace(normalizedWorkspaceName))
return false;
await WORKSPACE_TREE_CACHE_SEMAPHORE.WaitAsync();
try
{
await EnsureTreeShellLoadedCoreAsync();
return WorkspaceNameExistsCore(normalizedWorkspaceName, excludedWorkspaceId);
}
finally
{
WORKSPACE_TREE_CACHE_SEMAPHORE.Release();
}
}
public static async Task<TryCreateWorkspaceResult> TryCreateWorkspaceAsync(string workspaceName)
{
var normalizedWorkspaceName = NormalizeWorkspaceName(workspaceName);
if (string.IsNullOrWhiteSpace(normalizedWorkspaceName))
return new(false, default);
await WORKSPACE_TREE_CACHE_SEMAPHORE.WaitAsync();
try
{
await EnsureTreeShellLoadedCoreAsync();
if (WorkspaceNameExistsCore(normalizedWorkspaceName))
return new(false, default);
var workspaceId = Guid.NewGuid();
var workspacePath = Path.Join(WORKSPACE_ROOT_DIRECTORY, workspaceId.ToString());
Directory.CreateDirectory(workspacePath);
var workspaceNamePath = Path.Join(workspacePath, "name");
await File.WriteAllTextAsync(workspaceNamePath, normalizedWorkspaceName, Encoding.UTF8);
var workspace = new WorkspaceCacheEntry
{
WorkspaceId = workspaceId,
WorkspacePath = workspacePath,
WorkspaceName = normalizedWorkspaceName,
Chats = [],
ChatsLoaded = false,
};
WORKSPACE_TREE_CACHE.Workspaces[workspaceId] = workspace;
WORKSPACE_TREE_CACHE.WorkspaceOrder.Add(workspaceId);
return new(true, ToPublicWorkspace(workspace));
}
finally
{
WORKSPACE_TREE_CACHE_SEMAPHORE.Release();
}
}
public static async Task<bool> RenameWorkspaceAsync(Guid workspaceId, string workspaceName)
{
var normalizedWorkspaceName = NormalizeWorkspaceName(workspaceName);
if (string.IsNullOrWhiteSpace(normalizedWorkspaceName))
return false;
await WORKSPACE_TREE_CACHE_SEMAPHORE.WaitAsync();
try
{
await EnsureTreeShellLoadedCoreAsync();
if (!WORKSPACE_TREE_CACHE.Workspaces.TryGetValue(workspaceId, out var workspace))
return false;
var workspaceNamePath = Path.Join(workspace.WorkspacePath, "name");
if (string.Equals(workspace.WorkspaceName.Trim(), normalizedWorkspaceName, StringComparison.OrdinalIgnoreCase))
{
await File.WriteAllTextAsync(workspaceNamePath, normalizedWorkspaceName, Encoding.UTF8);
workspace.WorkspaceName = normalizedWorkspaceName;
return true;
}
if (WorkspaceNameExistsCore(normalizedWorkspaceName, workspaceId))
return false;
await File.WriteAllTextAsync(workspaceNamePath, normalizedWorkspaceName, Encoding.UTF8);
workspace.WorkspaceName = normalizedWorkspaceName;
return true;
}
finally
{
WORKSPACE_TREE_CACHE_SEMAPHORE.Release();
}
}
public static bool IsChatExisting(LoadChat loadChat)
{
@@ -533,7 +771,7 @@ public static class WorkspaceBehaviour
// Not in cache — read from disk and update cache in the same semaphore scope
// to avoid a second semaphore acquisition via UpdateWorkspaceNameInCacheAsync:
var workspacePath = Path.Join(WorkspaceRootDirectory, workspaceId.ToString());
var workspacePath = Path.Join(WORKSPACE_ROOT_DIRECTORY, workspaceId.ToString());
var workspaceNamePath = Path.Join(workspacePath, "name");
string workspaceName;
@@ -621,7 +859,7 @@ public static class WorkspaceBehaviour
private static async Task EnsureWorkspace(Guid workspaceId, string workspaceName)
{
var workspacePath = Path.Join(WorkspaceRootDirectory, workspaceId.ToString());
var workspacePath = Path.Join(WORKSPACE_ROOT_DIRECTORY, workspaceId.ToString());
var workspaceNamePath = Path.Join(workspacePath, "name");
if (!Path.Exists(workspacePath))
@@ -651,4 +889,4 @@ public static class WorkspaceBehaviour
public static async Task EnsureBiasWorkspace() => await EnsureWorkspace(KnownWorkspaces.BIAS_WORKSPACE_ID, "Bias of the Day");
public static async Task EnsureERIServerWorkspace() => await EnsureWorkspace(KnownWorkspaces.ERI_SERVER_WORKSPACE_ID, "ERI Servers");
}
}
@@ -0,0 +1,3 @@
namespace AIStudio.Tools;
public readonly record struct WorkspaceSearchResult(WorkspaceTreeChat Chat, bool NameMatched, bool ThreadMatched);
@@ -0,0 +1,3 @@
namespace AIStudio.Tools;
public readonly record struct WorkspaceSearchSnapshot(IReadOnlyList<WorkspaceSearchWorkspace> Workspaces, IReadOnlyList<WorkspaceSearchResult> TemporaryChats);
@@ -0,0 +1,3 @@
namespace AIStudio.Tools;
public readonly record struct WorkspaceSearchWorkspace(Guid WorkspaceId, string WorkspacePath, string Name, IReadOnlyList<WorkspaceSearchResult> Chats);