using AIStudio.Assistants.ERI; using AIStudio.Provider; using AIStudio.Settings; using AIStudio.Settings.DataModel; using AIStudio.Tools.ERIClient; using AIStudio.Tools.ERIClient.DataModel; namespace AIStudio.Tools.Services; public sealed class DataSourceService { private readonly record struct ParticipatingProvider(string Role, bool IsTrusted, ConfidenceLevel ConfidenceLevel); private readonly RustService rustService; private readonly SettingsManager settingsManager; private readonly ILogger logger; public DataSourceService(SettingsManager settingsManager, ILogger logger, RustService rustService) { this.logger = logger; this.rustService = rustService; this.settingsManager = settingsManager; this.logger.LogInformation("The data source service has been initialized."); } /// /// Returns a list of data sources that are allowed for the selected LLM provider. /// It also returns the data sources selected before when they are still allowed. /// /// The selected LLM provider. /// The active data source options, which determine which agent providers participate. /// The data sources selected before. /// The allowed data sources and the data sources selected before -- when they are still allowed. public async Task GetDataSources(AIStudio.Settings.Provider selectedLLMProvider, DataSourceOptions dataSourceOptions, IReadOnlyCollection? previousSelectedDataSources = null) { // // Case: Somehow the selected LLM provider was not set. The default provider // does not mean anything. We cannot filter the data sources by any means. // We return an empty list. Better safe than sorry. // if (selectedLLMProvider == Settings.Provider.NONE) { this.logger.LogWarning("The selected LLM provider is not set. We cannot filter the data sources by any means."); return new([], []); } 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); } /// /// Returns a list of data sources that are allowed for the selected LLM provider. /// It also returns the data sources selected before when they are still allowed. /// /// The selected LLM provider. /// The active data source options, which determine which agent providers participate. /// The data sources selected before. /// The allowed data sources and the data sources selected before -- when they are still allowed. public async Task GetDataSources(IProvider selectedLLMProvider, DataSourceOptions dataSourceOptions, IReadOnlyCollection? previousSelectedDataSources = null) { // // Case: Somehow the selected LLM provider was not set. The default provider // does not mean anything. We cannot filter the data sources by any means. // We return an empty list. Better safe than sorry. // if (selectedLLMProvider is NoProvider) { this.logger.LogWarning("The selected LLM provider is the default provider. We cannot filter the data sources by any means."); return new([], []); } 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 IReadOnlyList GetParticipatingProviders(string currentProviderId, DataSourceOptions dataSourceOptions, ParticipatingProvider currentProvider) { var providers = new List { 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 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 GetDataSources(bool usingTrustedProvider, IReadOnlyList participatingProviders, IReadOnlyCollection? previousSelectedDataSources = null) { var allDataSources = this.settingsManager.ConfigurationData.DataSources.ToList(); var previousSelectedDataSourceIds = previousSelectedDataSources?.Select(source => source.Id).ToHashSet(StringComparer.Ordinal) ?? []; var filteredDataSources = new List(allDataSources.Count); var filteredSelectedDataSources = new List(previousSelectedDataSourceIds.Count); var tasks = new List>(allDataSources.Count); // Start all checks in parallel: foreach (var source in allDataSources) tasks.Add(this.CheckOneDataSource(source, usingTrustedProvider, participatingProviders)); // Wait for all checks and collect the results: foreach (var task in tasks) { var source = await task; if (source is not null) { filteredDataSources.Add(source); if (previousSelectedDataSourceIds.Contains(source.Id)) filteredSelectedDataSources.Add(source); } } return new(filteredDataSources, filteredSelectedDataSources); } private async Task CheckOneDataSource(IDataSource source, bool usingTrustedProvider, IReadOnlyList 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. // SecurityRequirements? eriSourceRequirements = null; if (source is DataSourceERI_V1 eriSource) { using var cancellationTokenSource = new CancellationTokenSource(TimeSpan.FromSeconds(6)); using var client = ERIClientFactory.Get(ERIVersion.V1, eriSource); if(client is null) { this.logger.LogError($"Could not create ERI client for source '{source.Name}' (id={source.Id}). We skip this source."); return null; } this.logger.LogInformation($"Authenticating with ERI source '{source.Name}' (id={source.Id})..."); var loginResult = await client.AuthenticateAsync(this.rustService, cancellationToken: cancellationTokenSource.Token); if (!loginResult.Successful) { this.logger.LogWarning($"Authentication with ERI source '{source.Name}' (id={source.Id}) failed. We skip this source. Reason: {loginResult.Message}"); return null; } this.logger.LogInformation($"Checking security requirements for ERI source '{source.Name}' (id={source.Id})..."); var securityRequest = await client.GetSecurityRequirementsAsync(cancellationTokenSource.Token); if (!securityRequest.Successful) { this.logger.LogWarning($"Could not retrieve security requirements for ERI source '{source.Name}' (id={source.Id}). We skip this source. Reason: {loginResult.Message}"); return null; } eriSourceRequirements = securityRequest.Data; this.logger.LogInformation($"Security requirements for ERI source '{source.Name}' (id={source.Id}) retrieved successfully."); } if (source is not IExternalDataSource externalSource) return source; switch (externalSource.SecurityPolicy) { case DataSourceSecurity.ALLOW_ANY: // // Case: The data source allows any provider type. We want to use a trusted provider. // There is no issue with this source. Accept it. // if(usingTrustedProvider) return source; // // Case: This is a local data source. When the source allows any provider type, we can use it. // Accept it. // if(eriSourceRequirements is null) return source; // // Case: The ERI source requires a self-hosted or organization-trusted provider. This misconfiguration happens // when the ERI server operator changes the security requirements. The ERI server // operator owns the data -- we have to respect their rules. We skip this source. // if (eriSourceRequirements is { AllowedProviderType: ProviderType.SELF_HOSTED }) { this.logger.LogWarning($"The ERI source '{source.Name}' (id={source.Id}) requires a self-hosted or organization-trusted provider. We skip this source."); return null; } // // Case: The ERI source allows any provider type. The data source configuration is correct. // Accept it. // if(eriSourceRequirements is { AllowedProviderType: ProviderType.ANY }) return source; // // Case: Missing rules. We skip this source. Better safe than sorry. // this.logger.LogWarning($"The ERI source '{source.Name}' (id={source.Id}) was filtered out due to missing rules."); return null; // // Case: The data source requires a trusted provider. We want to use a trusted provider. // There is no issue with this source. Accept it. // case DataSourceSecurity.SELF_HOSTED when usingTrustedProvider: return source; // // Case: The data source requires a trusted provider. We want to use an untrusted provider. // We skip this source. // case DataSourceSecurity.SELF_HOSTED when !usingTrustedProvider: this.logger.LogWarning($"The data source '{source.Name}' (id={source.Id}) requires a self-hosted or organization-trusted provider. We skip this source."); return null; // // Case: The data source did not specify a security policy. We skip this source. // Better safe than sorry. // case DataSourceSecurity.NOT_SPECIFIED: this.logger.LogWarning($"The data source '{source.Name}' (id={source.Id}) has no security policy. We skip this source."); return null; // // Case: Some developer forgot to implement a security policy. We skip this source. // Better safe than sorry. // default: this.logger.LogWarning($"The data source '{source.Name}' (id={source.Id}) was filtered out due unknown security policy."); return null; } } }