mirror of
https://github.com/MindWorkAI/AI-Studio.git
synced 2026-10-08 19:29:40 +00:00
Merge branch 'main' into chunk-data
Resolved 29 conflicting files. The notable decisions: Confidence: main's tool-calling gate (RequiredProviderConfidence) and this branch's local-RAG gate (DataConfidenceLevel) turned out to be the same rule on the same axis, so they are now one field. Both tool results and data sources raise it through RequireProviderConfidence(). The gate checks the level strictly and no longer exempts providers trusted by configuration: TrustedProviderIds is documented as applying to data-source security checks only, and organizations set confidence through DataConfidence .CustomConfidenceScheme instead. The security axis (DataSecurity, ERI, IsTrustedForDataSourceSecurityChecks) is unchanged. Provider creation: main's CreateProvider signature won (hfEndpointKind, capabilityOverrides, no model parameter); tokenizerPath was added to it and is set for every provider, including the new Hetzner, IONOS and LiteLLM. Provider and EmbeddingProvider combine the record parameters, Lua parsing and Lua serialization of both sides. File types: main's hierarchy (ODT leaf, WORD parent, PowerPoint without the legacy .ppt, TABULAR instead of DELIMITED_TABLE) plus this branch's SPREADSHEET parent with ODS and the xlsm/xlsb/xla/xlam extensions, which the runtime already reads. Both sides had added a conflicting HTML filter; the reading family keeps the name, and the export path uses a narrow HTML_DOCUMENT, following the existing LATEX/TEX split. Runtime: main's file_data.rs is the base, including the prompt-injection sanitizer and the extraction routes. Token counting and chunk segmentation moved into take_released, so they act on the text the filter has released rather than on text it is still holding. A failed count is logged and left out instead of ending the extraction, because the app counts such a segment itself. Data sources: the participating-provider checks of this branch are kept, and main's GetAllowedDataSources overload now builds on them. DirectChatService resolves the launched chat's data source options before the check, so filter and chat see the same options. .NET and Rust both build clean; I18N regenerated to 4060 keys.
This commit is contained in:
640 files changed
+45059
-4952
No files matched your search
@@ -29,7 +29,7 @@ public sealed class ProviderAlibabaCloud() : BaseProvider(LLMProviders.ALIBABA_C
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -44,6 +44,7 @@ public sealed class ProviderAlibabaCloud() : BaseProvider(LLMProviders.ALIBABA_C
|
||||
Messages = [systemPrompt, ..messages],
|
||||
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -102,10 +103,12 @@ public sealed class ProviderAlibabaCloud() : BaseProvider(LLMProviders.ALIBABA_C
|
||||
new Model("qwen2.5-vl-3b-instruct", "Qwen2.5-VL 3b"),
|
||||
};
|
||||
|
||||
var result = await this.LoadModels(["q"], SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(["q"], SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Concat(additionalModels).OrderBy(x => x.Id)]
|
||||
// The API is the authority: when it reports a model we also keep as a fallback above,
|
||||
// its entry comes first and the fallback is dropped.
|
||||
Models = [..result.Models.Concat(additionalModels).DistinctBy(x => x.Id).OrderBy(x => x.Id)]
|
||||
};
|
||||
}
|
||||
|
||||
@@ -124,10 +127,12 @@ public sealed class ProviderAlibabaCloud() : BaseProvider(LLMProviders.ALIBABA_C
|
||||
new Model("text-embedding-v3", "text-embedding-v3"),
|
||||
};
|
||||
|
||||
var result = await this.LoadModels(["text-embedding-"], SecretStoreType.EMBEDDING_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(["text-embedding-"], SecretStoreType.EMBEDDING_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Concat(additionalModels).OrderBy(x => x.Id)]
|
||||
// The API is the authority: when it reports a model we also keep as a fallback above,
|
||||
// its entry comes first and the fallback is dropped.
|
||||
Models = [..result.Models.Concat(additionalModels).DistinctBy(x => x.Id).OrderBy(x => x.Id)]
|
||||
};
|
||||
}
|
||||
|
||||
@@ -143,14 +148,12 @@ public sealed class ProviderAlibabaCloud() : BaseProvider(LLMProviders.ALIBABA_C
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(string[] prefixes, SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadModels(string[] prefixes, SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data.Where(model => prefixes.Any(prefix => model.Id.StartsWith(prefix, StringComparison.InvariantCulture))),
|
||||
token,
|
||||
apiKeyProvisional);
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
using System.Text.Json;
|
||||
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
/// <summary>
|
||||
/// One turn of the model, handed back unchanged so the conversation can continue.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The content blocks are kept as raw JSON on purpose. A turn can carry text, tool uses, and
|
||||
/// thinking blocks, and the thinking blocks have to return exactly as they arrived — reading and
|
||||
/// rebuilding them would risk changing them.
|
||||
/// </remarks>
|
||||
public sealed record AnthropicMessage(IList<JsonElement> Content, string Role = "assistant") : IMessage<IList<JsonElement>>;
|
||||
@@ -0,0 +1,48 @@
|
||||
using System.Text.Json;
|
||||
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
/// <summary>
|
||||
/// One non-streamed answer of the Anthropic messages API.
|
||||
/// </summary>
|
||||
public sealed record AnthropicResponse
|
||||
{
|
||||
public string StopReason { get; init; } = string.Empty;
|
||||
|
||||
public IList<JsonElement> Content { get; init; } = [];
|
||||
|
||||
/// <summary>
|
||||
/// The tool calls the model asked for.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// A block without an ID or a name cannot be answered and is dropped here, so the harness
|
||||
/// sees a well-formed list. Anthropic supplies both for every real tool use.
|
||||
/// </remarks>
|
||||
public IReadOnlyList<AnthropicToolUse> GetToolUses() => this.Content
|
||||
.Where(x => ReadString(x, "type").Equals("tool_use", StringComparison.Ordinal))
|
||||
.Select(x => new AnthropicToolUse
|
||||
{
|
||||
Id = ReadString(x, "id"),
|
||||
Name = ReadString(x, "name"),
|
||||
Input = x.TryGetProperty("input", out var input) ? input : default,
|
||||
})
|
||||
.Where(x => !string.IsNullOrWhiteSpace(x.Id) && !string.IsNullOrWhiteSpace(x.Name))
|
||||
.ToList();
|
||||
|
||||
/// <summary>
|
||||
/// The text the model wrote, with its blocks joined.
|
||||
/// </summary>
|
||||
public string GetTextOutput() => string.Concat(this.Content
|
||||
.Where(x => ReadString(x, "type").Equals("text", StringComparison.Ordinal))
|
||||
.Select(x => ReadString(x, "text")));
|
||||
|
||||
private static string ReadString(JsonElement item, string propertyName)
|
||||
{
|
||||
if (item.ValueKind is not JsonValueKind.Object ||
|
||||
!item.TryGetProperty(propertyName, out var property) ||
|
||||
property.ValueKind is not JsonValueKind.String)
|
||||
return string.Empty;
|
||||
|
||||
return property.GetString() ?? string.Empty;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
using System.Text.Json;
|
||||
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
/// <summary>
|
||||
/// A tool as the Anthropic messages API expects it.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Anthropic names the schema field input schema, where the chat completions and responses APIs
|
||||
/// call it parameters. The description is the plain one, without their nesting.
|
||||
/// </remarks>
|
||||
public sealed record AnthropicTool
|
||||
{
|
||||
public string Name { get; init; } = string.Empty;
|
||||
|
||||
public string Description { get; init; } = string.Empty;
|
||||
|
||||
public bool Strict { get; init; }
|
||||
|
||||
public JsonElement InputSchema { get; init; }
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
using AIStudio.Tools.ToolCallingSystem;
|
||||
using AIStudio.Tools.ToolCallingSystem.Harness;
|
||||
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
/// <summary>
|
||||
/// Speaks the Anthropic messages wire format for the tool calling loop.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Anthropic works in content blocks rather than in separate message kinds: the model's turn is
|
||||
/// one assistant message whose blocks may mix text, thinking, and tool uses, and the results go
|
||||
/// back as tool result blocks inside a single user message. That difference is what made this
|
||||
/// provider hard to support before the loop and the wire format were separated — it is now the
|
||||
/// only thing this class is about.
|
||||
/// </remarks>
|
||||
public sealed class AnthropicToolCallingAdapter(Model chatModel, IList<IMessageBase> baseMessages, string systemPrompt, int maxTokens,
|
||||
IDictionary<string, object> apiParameters, IReadOnlyList<(ToolDefinition Definition, IToolImplementation Implementation)> runnableTools,
|
||||
Func<ChatRequest, CancellationToken, Task<AnthropicResponse?>> executeRequestAsync) : IToolCallingProviderAdapter
|
||||
{
|
||||
private readonly List<IMessageBase> internalMessages = [];
|
||||
private readonly List<AnthropicToolResultContent> pendingToolResults = [];
|
||||
private readonly List<AnthropicTool> tools = runnableTools.Select(x => ProviderToolAdapters.ToAnthropicTool(x.Definition)).ToList();
|
||||
private AnthropicResponse? lastResponse;
|
||||
|
||||
/// <inheritdoc />
|
||||
public async Task<ToolCallingRound?> ExecuteRoundAsync(string? finalResponseInstruction, bool includeTools, CancellationToken token = default)
|
||||
{
|
||||
//
|
||||
// The results of the previous round are flushed here rather than when they were recorded:
|
||||
// they all belong in one user message, and only now is it certain that no more are coming.
|
||||
//
|
||||
if (this.pendingToolResults.Count > 0)
|
||||
{
|
||||
this.internalMessages.Add(new AnthropicToolResultMessage([..this.pendingToolResults]));
|
||||
this.pendingToolResults.Clear();
|
||||
}
|
||||
|
||||
var response = await executeRequestAsync(new ChatRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
Messages = [..baseMessages, ..this.internalMessages],
|
||||
System = finalResponseInstruction is null
|
||||
? systemPrompt
|
||||
: $"{systemPrompt}{Environment.NewLine}{Environment.NewLine}{finalResponseInstruction}",
|
||||
|
||||
MaxTokens = maxTokens,
|
||||
Stream = false,
|
||||
Tools = includeTools && this.tools.Count > 0 ? this.tools : null,
|
||||
AdditionalApiParameters = apiParameters,
|
||||
}, token);
|
||||
|
||||
if (response is null)
|
||||
return null;
|
||||
|
||||
this.lastResponse = response;
|
||||
return new ToolCallingRound(
|
||||
response.GetTextOutput(),
|
||||
response.GetToolUses()
|
||||
.Select(toolUse => new ToolCallingRequestedCall(
|
||||
toolUse.Id,
|
||||
toolUse.Name,
|
||||
toolUse.Arguments,
|
||||
ToolExecutor.IsValidArgumentsJson(toolUse.Arguments)))
|
||||
.ToList(),
|
||||
[]);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public void RecordAssistantTurn()
|
||||
{
|
||||
if (this.lastResponse is null)
|
||||
return;
|
||||
|
||||
//
|
||||
// The blocks go back exactly as they arrived. Thinking blocks in particular have to be
|
||||
// returned unchanged for the model to continue from them.
|
||||
//
|
||||
this.internalMessages.Add(new AnthropicMessage([..this.lastResponse.Content]));
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public void RecordToolResult(string callId, string content, bool isError = false) => this.pendingToolResults.Add(new AnthropicToolResultContent
|
||||
{
|
||||
ToolUseId = callId,
|
||||
Content = content,
|
||||
IsError = isError,
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
using System.Text.Json.Serialization;
|
||||
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
public sealed record AnthropicToolResultContent
|
||||
{
|
||||
public string Type { get; init; } = "tool_result";
|
||||
|
||||
public string ToolUseId { get; init; } = string.Empty;
|
||||
|
||||
public string Content { get; init; } = string.Empty;
|
||||
|
||||
/// <summary>
|
||||
/// Whether the tool failed rather than returning a result.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Only sent when true: Anthropic reads its absence as success, and this way a successful
|
||||
/// result stays byte-identical to what earlier versions sent.
|
||||
/// </remarks>
|
||||
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingDefault)]
|
||||
public bool IsError { get; init; }
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
/// <summary>
|
||||
/// The results of the tools the model asked for, as one user turn.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// All results of one turn belong in a single message. Splitting them across several messages
|
||||
/// teaches the model to stop asking for more than one tool at a time.
|
||||
/// </remarks>
|
||||
public sealed record AnthropicToolResultMessage(IList<AnthropicToolResultContent> Content, string Role = "user") : IMessage<IList<AnthropicToolResultContent>>;
|
||||
@@ -0,0 +1,17 @@
|
||||
using System.Text.Json;
|
||||
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
public sealed record AnthropicToolUse
|
||||
{
|
||||
public string Id { get; init; } = string.Empty;
|
||||
|
||||
public string Name { get; init; } = string.Empty;
|
||||
|
||||
public JsonElement Input { get; init; }
|
||||
|
||||
/// <summary>
|
||||
/// The arguments as JSON text, which is what the tool executor works with.
|
||||
/// </summary>
|
||||
public string Arguments => this.Input.ValueKind is JsonValueKind.Undefined ? "{}" : this.Input.GetRawText();
|
||||
}
|
||||
@@ -18,7 +18,17 @@ public readonly record struct ChatRequest(
|
||||
string System
|
||||
)
|
||||
{
|
||||
/// <summary>
|
||||
/// The tools the model may call, or null when it should answer without them.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Omitted from the request when null: sending an empty list is not the same as sending no
|
||||
/// tools at all, and the final round of a tool conversation has to offer none.
|
||||
/// </remarks>
|
||||
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
|
||||
public IList<AnthropicTool>? Tools { get; init; }
|
||||
|
||||
// Attention: The "required" modifier is not supported for [JsonExtensionData].
|
||||
[JsonExtensionData]
|
||||
public IDictionary<string, object> AdditionalApiParameters { get; init; } = new Dictionary<string, object>();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
// ReSharper disable NotAccessedPositionalProperty.Global
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
/// <summary>
|
||||
/// The delta object of a response line.
|
||||
/// </summary>
|
||||
/// <param name="Type">The type of the delta.</param>
|
||||
/// <param name="Text">The text of the delta.</param>
|
||||
public readonly record struct Delta(string Type, string Text);
|
||||
@@ -5,13 +5,15 @@ using System.Text.Json;
|
||||
using AIStudio.Chat;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Settings;
|
||||
using AIStudio.Tools.Rust;
|
||||
using AIStudio.Tools.ToolCallingSystem;
|
||||
using AIStudio.Tools.ToolCallingSystem.Harness;
|
||||
|
||||
namespace AIStudio.Provider.Anthropic;
|
||||
|
||||
public sealed class ProviderAnthropic() : BaseProvider(LLMProviders.ANTHROPIC, new Uri("https://api.anthropic.com/v1/"), ExternalHttpTrustPolicy.SYSTEM_TRUST_ONLY, LOGGER)
|
||||
{
|
||||
private static readonly ILogger<ProviderAnthropic> LOGGER = Program.LOGGER_FACTORY.CreateLogger<ProviderAnthropic>();
|
||||
|
||||
#region Implementation of IProvider
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -70,18 +72,59 @@ public sealed class ProviderAnthropic() : BaseProvider(LLMProviders.ANTHROPIC, n
|
||||
}
|
||||
}
|
||||
);
|
||||
|
||||
|
||||
//
|
||||
// Prepare the tools we want to use. When the model may call one, the conversation runs
|
||||
// through the harness instead of being streamed straight away: tool rounds are not
|
||||
// streamed, only the final answer is.
|
||||
//
|
||||
var toolRegistry = Program.SERVICE_PROVIDER.GetService<ToolRegistry>();
|
||||
var toolExecutor = Program.SERVICE_PROVIDER.GetService<ToolExecutor>();
|
||||
var currentAssistantContent = chatThread.Blocks.LastOrDefault(x => x.Role is ChatRole.AI)?.Content as ContentText;
|
||||
currentAssistantContent?.ToolInvocations.Clear();
|
||||
|
||||
var providerSettings = this.CreateSettingsProvider(chatModel);
|
||||
var runnableTools = toolRegistry is null
|
||||
? []
|
||||
: await toolRegistry.GetRunnableToolsAsync(providerSettings, chatThread.RuntimeComponent, chatThread.RuntimeSelectedToolIds,
|
||||
this.Provider.GetConfidence(settingsManager).Level, chatThread.MayRunTools(settingsManager));
|
||||
|
||||
var systemPrompt = chatThread.PrepareSystemPrompt(settingsManager, runnableTools.Select(x => x.Definition));
|
||||
if (toolExecutor is not null && runnableTools.Count > 0)
|
||||
{
|
||||
var adapter = new AnthropicToolCallingAdapter(chatModel, [..messages], systemPrompt, maxTokens, apiParameters, runnableTools,
|
||||
(requestDto, requestToken) => this.ExecuteMessagesRequest(requestDto, requestedSecret, requestToken));
|
||||
|
||||
var loop = Program.SERVICE_PROVIDER.GetRequiredService<IToolCallingLoop>();
|
||||
var loopContext = new ToolCallingLoopContext
|
||||
{
|
||||
ChatThread = chatThread,
|
||||
RunnableTools = runnableTools,
|
||||
ToolExecutor = toolExecutor,
|
||||
Provider = this,
|
||||
CurrentAssistantContent = currentAssistantContent,
|
||||
ProviderInstanceName = this.InstanceName,
|
||||
ProviderType = this.Provider,
|
||||
ModelId = chatModel.Id,
|
||||
};
|
||||
|
||||
await foreach (var content in loop.RunAsync(adapter, loopContext, token))
|
||||
yield return content;
|
||||
|
||||
yield break;
|
||||
}
|
||||
|
||||
// Prepare the Anthropic HTTP chat request:
|
||||
var chatRequest = JsonSerializer.Serialize(new ChatRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
|
||||
|
||||
// Build the messages:
|
||||
Messages = [..messages],
|
||||
|
||||
System = chatThread.PrepareSystemPrompt(settingsManager),
|
||||
|
||||
System = systemPrompt,
|
||||
MaxTokens = maxTokens,
|
||||
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
AdditionalApiParameters = apiParameters
|
||||
@@ -107,6 +150,33 @@ public sealed class ProviderAnthropic() : BaseProvider(LLMProviders.ANTHROPIC, n
|
||||
yield return content;
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Runs one non-streamed messages request, as the tool rounds need it.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Tool rounds are not streamed: the whole answer has to be there before its tool calls can
|
||||
/// be executed. Only the final answer reaches the user through the streaming path.
|
||||
/// </remarks>
|
||||
/// <returns>The answer, or null when the request failed and the user was already told.</returns>
|
||||
private async Task<AnthropicResponse?> ExecuteMessagesRequest(ChatRequest requestDto, RequestedSecret requestedSecret, CancellationToken token)
|
||||
{
|
||||
using var request = new HttpRequestMessage(HttpMethod.Post, "messages");
|
||||
request.Headers.Add("x-api-key", await requestedSecret.Secret.Decrypt(Program.ENCRYPTION));
|
||||
request.Headers.Add("anthropic-version", "2023-06-01");
|
||||
request.Content = new StringContent(JsonSerializer.Serialize(requestDto, JSON_SERIALIZER_OPTIONS), Encoding.UTF8, "application/json");
|
||||
|
||||
using var response = await this.HttpClient.SendAsync(request, token);
|
||||
if (!response.IsSuccessStatusCode)
|
||||
{
|
||||
var responseBody = await response.Content.ReadAsStringAsync(token);
|
||||
LOGGER.LogError("Tool calling messages request failed with status code {ResponseStatusCode} and body: '{ResponseBody}'.", response.StatusCode, responseBody);
|
||||
await ToolCallingMessages.SendToolCallingRequestFailedAsync((int)response.StatusCode);
|
||||
return null;
|
||||
}
|
||||
|
||||
return await response.Content.ReadFromJsonAsync<AnthropicResponse>(JSON_SERIALIZER_OPTIONS, token);
|
||||
}
|
||||
|
||||
#pragma warning disable CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ImageURL> StreamImageCompletion(Model imageModel, string promptPositive, string promptNegative = FilterOperator.String.Empty, ImageURL referenceImageURL = default, [EnumeratorCancellation] CancellationToken token = default)
|
||||
@@ -140,10 +210,12 @@ public sealed class ProviderAnthropic() : BaseProvider(LLMProviders.ANTHROPIC, n
|
||||
new Model("claude-3-opus-latest", "Claude 3 Opus (Latest)"),
|
||||
};
|
||||
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Concat(additionalModels).OrderBy(x => x.Id)]
|
||||
// The API is the authority: when it reports a model we also keep as a fallback above,
|
||||
// its entry comes first and the fallback is dropped.
|
||||
Models = [..result.Models.Concat(additionalModels).DistinctBy(x => x.Id).OrderBy(x => x.Id)]
|
||||
};
|
||||
}
|
||||
|
||||
@@ -164,16 +236,14 @@ public sealed class ProviderAnthropic() : BaseProvider(LLMProviders.ANTHROPIC, n
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
}
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models?limit=100",
|
||||
modelResponse => modelResponse.Data,
|
||||
token,
|
||||
apiKeyProvisional,
|
||||
failureReasonSelector: (response, _) => response.StatusCode switch
|
||||
{
|
||||
@@ -187,6 +257,6 @@ public sealed class ProviderAnthropic() : BaseProvider(LLMProviders.ANTHROPIC, n
|
||||
request.Headers.Add("x-api-key", secretKey);
|
||||
request.Headers.Add("anthropic-version", "2023-06-01");
|
||||
},
|
||||
jsonSerializerOptions: JSON_SERIALIZER_OPTIONS);
|
||||
jsonSerializerOptions: JSON_SERIALIZER_OPTIONS, token: token);
|
||||
}
|
||||
}
|
||||
@@ -29,11 +29,4 @@ public readonly record struct ResponseStreamLine(string Type, int Index, Delta D
|
||||
public IList<ISource> GetSources() => [];
|
||||
|
||||
#endregion
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// The delta object of a response line.
|
||||
/// </summary>
|
||||
/// <param name="Type">The type of the delta.</param>
|
||||
/// <param name="Text">The text of the delta.</param>
|
||||
public readonly record struct Delta(string Type, string Text);
|
||||
}
|
||||
@@ -10,6 +10,8 @@ using AIStudio.Provider.Anthropic;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Provider.SelfHosted;
|
||||
using AIStudio.Settings;
|
||||
using AIStudio.Tools.ToolCallingSystem;
|
||||
using AIStudio.Tools.ToolCallingSystem.Harness;
|
||||
using AIStudio.Tools.MIME;
|
||||
using AIStudio.Tools.PluginSystem;
|
||||
using AIStudio.Tools.Rust;
|
||||
@@ -87,6 +89,8 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
/// <inheritdoc />
|
||||
public string AdditionalJsonApiParameters { get; init; } = string.Empty;
|
||||
|
||||
internal ProviderCapabilityOverrides? CapabilityOverrides { get; set; }
|
||||
|
||||
/// <inheritdoc />
|
||||
public string TokenizerPath { get; init; } = string.Empty;
|
||||
|
||||
@@ -179,16 +183,15 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
_ => GetDefaultModelLoadFailureReason(response),
|
||||
};
|
||||
|
||||
protected async Task<ModelLoadResult> LoadModelsResponse<TResponse>(
|
||||
SecretStoreType storeType,
|
||||
protected async Task<ModelLoadResult> LoadModelsResponse<TResponse>(SecretStoreType storeType,
|
||||
string requestPath,
|
||||
Func<TResponse, IEnumerable<Model>> modelFactory,
|
||||
CancellationToken token,
|
||||
string? apiKeyProvisional = null,
|
||||
Func<HttpResponseMessage, string, ModelLoadFailureReason>? failureReasonSelector = null,
|
||||
Action<HttpRequestMessage, string>? requestConfigurator = null,
|
||||
JsonSerializerOptions? jsonSerializerOptions = null,
|
||||
bool isTryingSecret = false)
|
||||
bool isTryingSecret = false,
|
||||
CancellationToken token = default)
|
||||
{
|
||||
var secretKey = await this.GetModelLoadingSecretKey(storeType, apiKeyProvisional, isTryingSecret);
|
||||
if (string.IsNullOrWhiteSpace(secretKey) && !isTryingSecret)
|
||||
@@ -272,10 +275,10 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
{
|
||||
exception = new();
|
||||
|
||||
if (!line.StartsWith("data: ", StringComparison.InvariantCulture))
|
||||
if (!TryGetServerSentEventData(line, out var jsonData))
|
||||
return false;
|
||||
|
||||
var jsonData = line[6..].Trim();
|
||||
jsonData = jsonData.Trim();
|
||||
if (string.IsNullOrWhiteSpace(jsonData) || jsonData is "[DONE]")
|
||||
return false;
|
||||
|
||||
@@ -307,6 +310,21 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
}
|
||||
}
|
||||
|
||||
private static bool TryGetServerSentEventData(string line, out string data)
|
||||
{
|
||||
const string DATA_PREFIX = "data:";
|
||||
data = string.Empty;
|
||||
|
||||
if (!line.StartsWith(DATA_PREFIX, StringComparison.InvariantCulture))
|
||||
return false;
|
||||
|
||||
data = line[DATA_PREFIX.Length..];
|
||||
if (data.StartsWith(' '))
|
||||
data = data[1..];
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
private static bool IsProviderStreamFailure(JsonElement root)
|
||||
{
|
||||
var eventType = TryGetString(root, "type");
|
||||
@@ -363,7 +381,41 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
|
||||
errorCode = TryGetString(root, "code");
|
||||
errorType = TryGetString(root, "type");
|
||||
errorMessage = TryGetString(root, "message");
|
||||
|
||||
//
|
||||
// Services built on FastAPI, such as Helmholtz Blablador, word their errors as "detail".
|
||||
// And some providers put the sentence straight into "error" instead of an object, e.g.
|
||||
// {"error": "Model not supported by provider novita"}. The object form was handled above,
|
||||
// so reading "error" here can only meet the plain sentence:
|
||||
//
|
||||
errorMessage = TryGetString(root, "message") ?? TryGetString(root, "detail") ?? TryGetString(root, "error");
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Reads the error message a provider sent in the body of a failed response.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Providers word their errors differently, but they all put a sentence somewhere into the
|
||||
/// body. Passing that sentence on is what lets a user act on the problem instead of only
|
||||
/// learning that something went wrong.
|
||||
/// </remarks>
|
||||
/// <param name="responseBody">The body of the failed response.</param>
|
||||
/// <returns>The message, or an empty string when the body carries none.</returns>
|
||||
private static string ReadProviderErrorMessage(string responseBody)
|
||||
{
|
||||
if (string.IsNullOrWhiteSpace(responseBody))
|
||||
return string.Empty;
|
||||
|
||||
try
|
||||
{
|
||||
using var document = JsonDocument.Parse(responseBody);
|
||||
TryGetProviderStreamError(document.RootElement, out _, out _, out var errorMessage);
|
||||
return errorMessage ?? string.Empty;
|
||||
}
|
||||
catch (JsonException)
|
||||
{
|
||||
return string.Empty;
|
||||
}
|
||||
}
|
||||
|
||||
private static bool TryGetErrorElement(JsonElement root, out JsonElement errorElement)
|
||||
@@ -397,6 +449,10 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
/// <summary>
|
||||
/// Sends a request and handles rate limiting by exponential backoff.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Two cancellation tokens, so one of them cannot be the last parameter: the user token
|
||||
/// survives a retry, while the request token belongs to the single attempt being made.
|
||||
/// </remarks>
|
||||
/// <param name="requestBuilder">A function that builds the request.</param>
|
||||
/// <param name="userCancellationToken">The user cancellation token.</param>
|
||||
/// <param name="requestCancellationToken">The token to use for the HTTP request.</param>
|
||||
@@ -469,16 +525,35 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
|
||||
if(nextResponse.StatusCode is HttpStatusCode.BadRequest)
|
||||
{
|
||||
//
|
||||
// The provider explains the problem in the body, while the reason phrase says no
|
||||
// more than "Bad Request". We show that explanation and fall back to the phrase
|
||||
// only when the body carries none:
|
||||
//
|
||||
var badRequestMessage = ReadProviderErrorMessage(errorBody);
|
||||
if (string.IsNullOrWhiteSpace(badRequestMessage))
|
||||
badRequestMessage = nextResponse.ReasonPhrase;
|
||||
|
||||
//
|
||||
// When we recognize what went wrong, we say what it means for the user instead of
|
||||
// guessing at the message format. The classification happened above already:
|
||||
//
|
||||
var classifiedMessage = this.GetProviderRequestFailureUserMessage(providerRequestFailure);
|
||||
if(!string.IsNullOrWhiteSpace(classifiedMessage))
|
||||
{
|
||||
await MessageBus.INSTANCE.SendError(new(Icons.Material.Filled.CloudOff, classifiedMessage));
|
||||
}
|
||||
|
||||
// Check if the error body contains "context" and "token" (case-insensitive),
|
||||
// which indicates that the context window is likely exceeded:
|
||||
if(errorBody.Contains("context", StringComparison.InvariantCultureIgnoreCase) &&
|
||||
else if(errorBody.Contains("context", StringComparison.InvariantCultureIgnoreCase) &&
|
||||
errorBody.Contains("token", StringComparison.InvariantCultureIgnoreCase))
|
||||
{
|
||||
await MessageBus.INSTANCE.SendError(new(Icons.Material.Filled.CloudOff, string.Format(TB("We tried to communicate with the LLM provider '{0}' (type={1}). The data of the chat, including all file attachments, is probably too large for the selected model and provider. The provider message is: '{2}'"), this.InstanceName, this.Provider, nextResponse.ReasonPhrase)));
|
||||
await MessageBus.INSTANCE.SendError(new(Icons.Material.Filled.CloudOff, string.Format(TB("We tried to communicate with the LLM provider '{0}' (type={1}). The data of the chat, including all file attachments, is probably too large for the selected model and provider. The provider message is: '{2}'"), this.InstanceName, this.Provider, badRequestMessage)));
|
||||
}
|
||||
else
|
||||
{
|
||||
await MessageBus.INSTANCE.SendError(new(Icons.Material.Filled.CloudOff, string.Format(TB("We tried to communicate with the LLM provider '{0}' (type={1}). The required message format might be changed. The provider message is: '{2}'"), this.InstanceName, this.Provider, nextResponse.ReasonPhrase)));
|
||||
await MessageBus.INSTANCE.SendError(new(Icons.Material.Filled.CloudOff, string.Format(TB("We tried to communicate with the LLM provider '{0}' (type={1}). The required message format might be changed. The provider message is: '{2}'"), this.InstanceName, this.Provider, badRequestMessage)));
|
||||
}
|
||||
|
||||
this.logger.LogError("Failed request with status code {ResponseStatusCode} (message = '{ResponseReasonPhrase}', error body = '{ErrorBody}').", nextResponse.StatusCode, nextResponse.ReasonPhrase, errorBody);
|
||||
@@ -664,13 +739,13 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
if (this.TryCreateProviderRequestExceptionFromStreamLine(providerName, line, out var providerRequestException))
|
||||
throw providerRequestException;
|
||||
|
||||
// Skip lines that do not start with "data: ". Regard
|
||||
// Skip lines that do not start with "data:". According
|
||||
// to the specification, we only want to read the data lines:
|
||||
if (!line.StartsWith("data: ", StringComparison.InvariantCulture))
|
||||
if (!TryGetServerSentEventData(line, out var jsonData))
|
||||
continue;
|
||||
|
||||
// Check if the line is the end of the stream:
|
||||
if (line.StartsWith("data: [DONE]", StringComparison.InvariantCulture))
|
||||
if (jsonData is "[DONE]")
|
||||
yield break;
|
||||
|
||||
//
|
||||
@@ -684,10 +759,6 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
|
||||
try
|
||||
{
|
||||
// We know that the line starts with "data: ". Hence, we can
|
||||
// skip the first 6 characters to get the JSON data after that.
|
||||
var jsonData = line[6..];
|
||||
|
||||
// Deserialize the JSON data:
|
||||
providerResponse = JsonSerializer.Deserialize<TAnnotation>(jsonData, JSON_SERIALIZER_OPTIONS);
|
||||
|
||||
@@ -716,10 +787,6 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
TDelta? providerResponse;
|
||||
try
|
||||
{
|
||||
// We know that the line starts with "data: ". Hence, we can
|
||||
// skip the first 6 characters to get the JSON data after that.
|
||||
var jsonData = line[6..];
|
||||
|
||||
// Deserialize the JSON data:
|
||||
providerResponse = JsonSerializer.Deserialize<TDelta>(jsonData, JSON_SERIALIZER_OPTIONS);
|
||||
|
||||
@@ -869,20 +936,19 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
if (line.StartsWith("event: response.completed", StringComparison.InvariantCulture))
|
||||
yield break;
|
||||
|
||||
if (!TryGetServerSentEventData(line, out var jsonData))
|
||||
continue;
|
||||
|
||||
//
|
||||
// Find delta lines:
|
||||
//
|
||||
if (line.StartsWith("""
|
||||
data: {"type":"response.output_text.delta"
|
||||
""", StringComparison.InvariantCulture))
|
||||
if (jsonData.StartsWith("""
|
||||
{"type":"response.output_text.delta"
|
||||
""", StringComparison.InvariantCulture))
|
||||
{
|
||||
TDelta? providerResponse;
|
||||
try
|
||||
{
|
||||
// We know that the line starts with "data: ". Hence, we can
|
||||
// skip the first 6 characters to get the JSON data after that.
|
||||
var jsonData = line[6..];
|
||||
|
||||
// Deserialize the JSON data:
|
||||
providerResponse = JsonSerializer.Deserialize<TDelta>(jsonData, JSON_SERIALIZER_OPTIONS);
|
||||
|
||||
@@ -906,18 +972,14 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
//
|
||||
// Find annotation added lines:
|
||||
//
|
||||
else if (annotationSupported && line.StartsWith(
|
||||
else if (annotationSupported && jsonData.StartsWith(
|
||||
"""
|
||||
data: {"type":"response.output_text.annotation.added"
|
||||
{"type":"response.output_text.annotation.added"
|
||||
""", StringComparison.InvariantCulture))
|
||||
{
|
||||
TAnnotation? providerResponse;
|
||||
try
|
||||
{
|
||||
// We know that the line starts with "data: ". Hence, we can
|
||||
// skip the first 6 characters to get the JSON data after that.
|
||||
var jsonData = line[6..];
|
||||
|
||||
// Deserialize the JSON data:
|
||||
providerResponse = JsonSerializer.Deserialize<TAnnotation>(jsonData, JSON_SERIALIZER_OPTIONS);
|
||||
|
||||
@@ -965,13 +1027,14 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
Model chatModel,
|
||||
ChatThread chatThread,
|
||||
SettingsManager settingsManager,
|
||||
Func<TextMessage, IDictionary<string, object>, Task<TRequest>> requestFactory,
|
||||
Func<TextMessage, IDictionary<string, object>, IList<object>?, Task<TRequest>> requestFactory,
|
||||
SecretStoreType storeType = SecretStoreType.LLM_PROVIDER,
|
||||
bool isTryingSecret = false,
|
||||
string systemPromptRole = "system",
|
||||
string requestPath = "chat/completions",
|
||||
Action<HttpRequestHeaders>? headersAction = null,
|
||||
[EnumeratorCancellation] CancellationToken token = default)
|
||||
where TRequest : ChatCompletionAPIRequest
|
||||
where TDelta : IResponseStreamLine
|
||||
where TAnnotation : IAnnotationStreamLine
|
||||
{
|
||||
@@ -980,18 +1043,69 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
if(!requestedSecret.Success && !isTryingSecret)
|
||||
yield break;
|
||||
|
||||
// Prepare the system prompt:
|
||||
var systemPrompt = new TextMessage
|
||||
{
|
||||
Role = systemPromptRole,
|
||||
Content = chatThread.PrepareSystemPrompt(settingsManager),
|
||||
};
|
||||
|
||||
// Parse the API parameters:
|
||||
var apiParameters = this.ParseAdditionalApiParameters();
|
||||
var apiParameters = this.ParseAdditionalApiParameters("parallel_tool_calls");
|
||||
|
||||
var toolRegistry = Program.SERVICE_PROVIDER.GetService<ToolRegistry>();
|
||||
var toolExecutor = Program.SERVICE_PROVIDER.GetService<ToolExecutor>();
|
||||
var currentAssistantContent = chatThread.Blocks.LastOrDefault(x => x.Role is ChatRole.AI)?.Content as ContentText;
|
||||
currentAssistantContent?.ToolInvocations.Clear();
|
||||
|
||||
TextMessage systemPrompt;
|
||||
if (toolRegistry is not null && toolExecutor is not null)
|
||||
{
|
||||
var providerSettings = this.CreateSettingsProvider(chatModel);
|
||||
var runnableTools = await toolRegistry.GetRunnableToolsAsync(
|
||||
providerSettings,
|
||||
chatThread.RuntimeComponent,
|
||||
chatThread.RuntimeSelectedToolIds,
|
||||
this.Provider.GetConfidence(settingsManager).Level,
|
||||
chatThread.MayRunTools(settingsManager));
|
||||
|
||||
systemPrompt = new TextMessage
|
||||
{
|
||||
Role = systemPromptRole,
|
||||
Content = chatThread.PrepareSystemPrompt(settingsManager, runnableTools.Select(x => x.Definition)),
|
||||
};
|
||||
|
||||
if (runnableTools.Count > 0)
|
||||
{
|
||||
var adapter = new ChatCompletionToolCallingAdapter<TRequest>(requestFactory, systemPrompt, apiParameters,
|
||||
runnableTools.Select(x => ProviderToolAdapters.ToChatCompletionTool(x.Definition)).ToList(), runnableTools,
|
||||
(requestDto, requestToken) => this.ExecuteChatCompletionRequest(requestDto, requestPath, requestedSecret, headersAction, requestToken),
|
||||
this.InstanceName, this.logger);
|
||||
|
||||
var loop = Program.SERVICE_PROVIDER.GetRequiredService<IToolCallingLoop>();
|
||||
var loopContext = new ToolCallingLoopContext
|
||||
{
|
||||
ChatThread = chatThread,
|
||||
RunnableTools = runnableTools,
|
||||
ToolExecutor = toolExecutor,
|
||||
Provider = this,
|
||||
CurrentAssistantContent = currentAssistantContent,
|
||||
ProviderInstanceName = this.InstanceName,
|
||||
ProviderType = this.Provider,
|
||||
ModelId = chatModel.Id,
|
||||
};
|
||||
|
||||
await foreach (var content in loop.RunAsync(adapter, loopContext, token))
|
||||
yield return content;
|
||||
|
||||
yield break;
|
||||
}
|
||||
|
||||
}
|
||||
else
|
||||
{
|
||||
systemPrompt = new TextMessage
|
||||
{
|
||||
Role = systemPromptRole,
|
||||
Content = chatThread.PrepareSystemPrompt(settingsManager),
|
||||
};
|
||||
}
|
||||
|
||||
// Prepare the provider HTTP chat request:
|
||||
var providerChatRequest = JsonSerializer.Serialize(await requestFactory(systemPrompt, apiParameters), JSON_SERIALIZER_OPTIONS);
|
||||
var providerChatRequest = JsonSerializer.Serialize(await requestFactory(systemPrompt, apiParameters, null), JSON_SERIALIZER_OPTIONS);
|
||||
|
||||
async Task<HttpRequestMessage> RequestBuilder()
|
||||
{
|
||||
@@ -1014,6 +1128,73 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
yield return content;
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Describes this provider instance with the given model as configured provider settings.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Anything asking about model capabilities must go through this, because the expert
|
||||
/// capability overrides live on the settings object: a provider that builds its own settings
|
||||
/// instance without them silently ignores what the user configured.
|
||||
/// </remarks>
|
||||
protected AIStudio.Settings.Provider CreateSettingsProvider(Model chatModel) => new()
|
||||
{
|
||||
UsedLLMProvider = this.Provider,
|
||||
Model = chatModel,
|
||||
InstanceName = this.InstanceName,
|
||||
CapabilityOverrides = this.CapabilityOverrides,
|
||||
};
|
||||
|
||||
private async Task<ChatCompletionResponse?> ExecuteChatCompletionRequest(ChatCompletionAPIRequest requestDto, string requestPath, RequestedSecret requestedSecret,
|
||||
Action<HttpRequestHeaders>? headersAction, CancellationToken token)
|
||||
{
|
||||
var responseData = await this.SendRequest(RequestBuilder, token);
|
||||
if (responseData.IsFailedAfterAllRetries)
|
||||
return null;
|
||||
|
||||
using var response = responseData.Response!;
|
||||
return await response.Content.ReadFromJsonAsync<ChatCompletionResponse>(JSON_SERIALIZER_OPTIONS, token);
|
||||
|
||||
async Task<HttpRequestMessage> RequestBuilder()
|
||||
{
|
||||
var request = new HttpRequestMessage(HttpMethod.Post, requestPath);
|
||||
if (requestedSecret.Success)
|
||||
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", await requestedSecret.Secret.Decrypt(Program.ENCRYPTION));
|
||||
|
||||
headersAction?.Invoke(request.Headers);
|
||||
request.Content = new StringContent(JsonSerializer.Serialize(requestDto, JSON_SERIALIZER_OPTIONS), Encoding.UTF8, "application/json");
|
||||
return request;
|
||||
}
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Builds the message a user gets to see when a transcription request failed.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// AI Studio always sends WebM/Opus. Some providers run their speech recognition behind a
|
||||
/// decoder which reads WAV only, and they answer with a bad request whose body says that it
|
||||
/// could not decode the file. Nobody can do anything about that inside AI Studio, so we name
|
||||
/// the likely cause and point at the provider instead of showing the raw message.
|
||||
/// </remarks>
|
||||
/// <param name="statusCode">The status code the provider answered with.</param>
|
||||
/// <param name="responseBody">The body the provider answered with.</param>
|
||||
/// <returns>The message to show, or an empty string when we have nothing to say.</returns>
|
||||
private string GetTranscriptionFailureUserMessage(HttpStatusCode statusCode, string responseBody)
|
||||
{
|
||||
var failureReason = this.ClassifyProviderRequestFailure(statusCode, responseBody);
|
||||
var classifiedMessage = this.GetProviderRequestFailureUserMessage(failureReason);
|
||||
if (!string.IsNullOrWhiteSpace(classifiedMessage))
|
||||
return classifiedMessage;
|
||||
|
||||
if (statusCode is HttpStatusCode.BadRequest && responseBody.Contains("not decode", StringComparison.OrdinalIgnoreCase))
|
||||
return string.Format(TB("The provider '{0}' was not able to read the audio file. It probably does not support the WebM/Opus format which AI Studio sends. Please contact the provider about it."), this.InstanceName);
|
||||
|
||||
var providerMessage = ReadProviderErrorMessage(responseBody);
|
||||
if (!string.IsNullOrWhiteSpace(providerMessage))
|
||||
return string.Format(TB("The provider '{0}' reported an error: {1}"), this.InstanceName, providerMessage);
|
||||
|
||||
return string.Empty;
|
||||
}
|
||||
|
||||
protected async Task<TranscriptionResult> PerformStandardTranscriptionRequest(RequestedSecret requestedSecret, Model transcriptionModel, string audioFilePath, Host host = Host.NONE, CancellationToken token = default)
|
||||
{
|
||||
try
|
||||
@@ -1040,6 +1221,15 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
|
||||
form.Add(new StringContent(modelName), "model");
|
||||
|
||||
//
|
||||
// Ask for the plain JSON format explicitly. We only ever read the 'text' field, so the
|
||||
// additional data of 'verbose_json' would be wasted anyway. More importantly, gateways
|
||||
// fill in a format of their own when the client names none: LiteLLM asks for
|
||||
// 'verbose_json' to get the duration it needs for its cost tracking, and the newer
|
||||
// transcription models of OpenAI reject that format.
|
||||
//
|
||||
form.Add(new StringContent("json"), "response_format");
|
||||
|
||||
using var request = new HttpRequestMessage(HttpMethod.Post, host.TranscriptionURL());
|
||||
request.Content = form;
|
||||
|
||||
@@ -1083,8 +1273,7 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
if (!response.IsSuccessStatusCode)
|
||||
{
|
||||
this.logger.LogError("Transcription request failed with status code {ResponseStatusCode} and body: '{ResponseBody}'.", response.StatusCode, responseBody);
|
||||
var providerRequestFailure = this.ClassifyProviderRequestFailure(response.StatusCode, responseBody);
|
||||
return TranscriptionResult.Failure(this.GetProviderRequestFailureUserMessage(providerRequestFailure));
|
||||
return TranscriptionResult.Failure(this.GetTranscriptionFailureUserMessage(response.StatusCode, responseBody));
|
||||
}
|
||||
|
||||
var transcriptionResponse = JsonSerializer.Deserialize<TranscriptionResponse>(responseBody, JSON_SERIALIZER_OPTIONS);
|
||||
@@ -1110,6 +1299,10 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
}
|
||||
}
|
||||
|
||||
/// <remarks>
|
||||
/// The cancellation token is not the last parameter, unlike everywhere else in this codebase:
|
||||
/// C# demands that a params parameter comes last.
|
||||
/// </remarks>
|
||||
protected async Task<IReadOnlyList<IReadOnlyList<float>>> PerformStandardTextEmbeddingRequest(RequestedSecret requestedSecret, Model embeddingModel, Host host = Host.NONE, CancellationToken token = default, params List<string> texts)
|
||||
{
|
||||
try
|
||||
@@ -1163,6 +1356,15 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
this.logger.LogError("Embedding request failed with status code {ResponseStatusCode} and body: '{ResponseBody}'.", response.StatusCode, responseBody);
|
||||
var providerRequestFailure = this.ClassifyProviderRequestFailure(response.StatusCode, responseBody);
|
||||
var userMessage = this.GetProviderRequestFailureUserMessage(providerRequestFailure);
|
||||
|
||||
// We know nothing about this failure, so we pass on what the provider said about it:
|
||||
if (string.IsNullOrWhiteSpace(userMessage))
|
||||
{
|
||||
var providerMessage = ReadProviderErrorMessage(responseBody);
|
||||
if (!string.IsNullOrWhiteSpace(providerMessage))
|
||||
userMessage = string.Format(TB("The provider '{0}' reported an error: {1}"), this.InstanceName, providerMessage);
|
||||
}
|
||||
|
||||
if (!string.IsNullOrWhiteSpace(userMessage))
|
||||
await MessageBus.INSTANCE.SendError(new(Icons.Material.Filled.CloudOff, userMessage));
|
||||
|
||||
@@ -1215,7 +1417,7 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
|
||||
protected static bool TryPopIntParameter(IDictionary<string, object> parameters, string key, out int value)
|
||||
{
|
||||
value = default;
|
||||
value = 0;
|
||||
if (!TryPopParameter(parameters, key, out var raw) || raw is null)
|
||||
return false;
|
||||
|
||||
@@ -1225,15 +1427,15 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
value = i;
|
||||
return true;
|
||||
|
||||
case long l when l is >= int.MinValue and <= int.MaxValue:
|
||||
case long l and >= int.MinValue and <= int.MaxValue:
|
||||
value = (int)l;
|
||||
return true;
|
||||
|
||||
case double d when d is >= int.MinValue and <= int.MaxValue:
|
||||
case double d and >= int.MinValue and <= int.MaxValue:
|
||||
value = (int)d;
|
||||
return true;
|
||||
|
||||
case decimal m when m is >= int.MinValue and <= int.MaxValue:
|
||||
case decimal m and >= int.MinValue and <= int.MaxValue:
|
||||
value = (int)m;
|
||||
return true;
|
||||
}
|
||||
@@ -1243,7 +1445,7 @@ public abstract class BaseProvider : IProvider, ISecretId
|
||||
|
||||
protected static bool TryPopBoolParameter(IDictionary<string, object> parameters, string key, out bool value)
|
||||
{
|
||||
value = default;
|
||||
value = false;
|
||||
if (!TryPopParameter(parameters, key, out var raw) || raw is null)
|
||||
return false;
|
||||
|
||||
|
||||
@@ -47,6 +47,12 @@ public sealed record Confidence
|
||||
Description = TB("The trust level of this provider **has not yet** been thoroughly **investigated and evaluated**. We do not know if your data is safe."),
|
||||
};
|
||||
|
||||
public static readonly Confidence USER_OPERATED_GATEWAY = new()
|
||||
{
|
||||
Level = ConfidenceLevel.UNKNOWN,
|
||||
Description = TB("You or your organization operate this gateway. However, it forwards your data to **whichever providers you configured behind it**, which may be cloud services in any jurisdiction. We cannot know where your data ends up, so **please assign the trust level yourself**."),
|
||||
};
|
||||
|
||||
public static readonly Confidence USA_NO_TRAINING = new()
|
||||
{
|
||||
Level = ConfidenceLevel.MODERATE,
|
||||
@@ -64,6 +70,12 @@ public sealed record Confidence
|
||||
Level = ConfidenceLevel.MEDIUM,
|
||||
Description = TB("The provider is located in the EU and is subject to the **GDPR** (General Data Protection Regulation). Additionally, the provider states that **your data is not used for training**."),
|
||||
};
|
||||
|
||||
public static readonly Confidence GDPR_EXPERIMENTAL_OPEN_SOURCE = new()
|
||||
{
|
||||
Level = ConfidenceLevel.MEDIUM,
|
||||
Description = TB("The provider operates its service in the EU and is subject to the **GDPR** (General Data Protection Regulation). It provides access to **open source models**. However, the service is currently **experimental**, and performance and availability are not guaranteed. We have no provider-specific information about whether submitted data is used for training."),
|
||||
};
|
||||
|
||||
public static readonly Confidence SELF_HOSTED = new()
|
||||
{
|
||||
|
||||
@@ -29,7 +29,7 @@ public sealed class ProviderDeepSeek() : BaseProvider(LLMProviders.DEEP_SEEK, ne
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingDirectImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -44,6 +44,7 @@ public sealed class ProviderDeepSeek() : BaseProvider(LLMProviders.DEEP_SEEK, ne
|
||||
Messages = [systemPrompt, ..messages],
|
||||
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -74,7 +75,7 @@ public sealed class ProviderDeepSeek() : BaseProvider(LLMProviders.DEEP_SEEK, ne
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -97,13 +98,12 @@ public sealed class ProviderDeepSeek() : BaseProvider(LLMProviders.DEEP_SEEK, ne
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data,
|
||||
token,
|
||||
apiKeyProvisional);
|
||||
modelResponse => modelResponse.Data.Where(model => model.IsChatModel()),
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
namespace AIStudio.Provider.Fireworks;
|
||||
|
||||
/// <summary>
|
||||
/// Data model for a choice made by the AI.
|
||||
/// </summary>
|
||||
/// <param name="Index">The index of the choice.</param>
|
||||
/// <param name="Delta">The delta text of the choice.</param>
|
||||
public readonly record struct Choice(int Index, Delta Delta);
|
||||
@@ -0,0 +1,7 @@
|
||||
namespace AIStudio.Provider.Fireworks;
|
||||
|
||||
/// <summary>
|
||||
/// The delta text of a choice.
|
||||
/// </summary>
|
||||
/// <param name="Content">The content of the delta text.</param>
|
||||
public readonly record struct Delta(string Content);
|
||||
@@ -29,7 +29,7 @@ public class ProviderFireworks() : BaseProvider(LLMProviders.FIREWORKS, new Uri(
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -45,6 +45,7 @@ public class ProviderFireworks() : BaseProvider(LLMProviders.FIREWORKS, new Uri(
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
|
||||
@@ -29,17 +29,4 @@ public readonly record struct ResponseStreamLine(string Id, string Object, uint
|
||||
public IList<ISource> GetSources() => [];
|
||||
|
||||
#endregion
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Data model for a choice made by the AI.
|
||||
/// </summary>
|
||||
/// <param name="Index">The index of the choice.</param>
|
||||
/// <param name="Delta">The delta text of the choice.</param>
|
||||
public readonly record struct Choice(int Index, Delta Delta);
|
||||
|
||||
/// <summary>
|
||||
/// The delta text of a choice.
|
||||
/// </summary>
|
||||
/// <param name="Content">The content of the delta text.</param>
|
||||
public readonly record struct Delta(string Content);
|
||||
}
|
||||
@@ -10,6 +10,14 @@ public sealed class ProviderGWDG() : BaseProvider(LLMProviders.GWDG, new Uri("ht
|
||||
{
|
||||
private static readonly ILogger<ProviderGWDG> LOGGER = Program.LOGGER_FACTORY.CreateLogger<ProviderGWDG>();
|
||||
|
||||
// Source: https://docs.hpc.gwdg.de/services/saia/index.html#embeddings
|
||||
private static readonly Model[] KNOWN_EMBEDDING_MODELS =
|
||||
[
|
||||
new("e5-mistral-7b-instruct", "E5 Mistral 7B Instruct"),
|
||||
new("multilingual-e5-large-instruct", "Multilingual E5 Large Instruct"),
|
||||
new("qwen3-embedding-4b", "Qwen3 Embedding 4B"),
|
||||
];
|
||||
|
||||
#region Implementation of IProvider
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -29,7 +37,7 @@ public sealed class ProviderGWDG() : BaseProvider(LLMProviders.GWDG, new Uri("ht
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -44,6 +52,7 @@ public sealed class ProviderGWDG() : BaseProvider(LLMProviders.GWDG, new Uri("ht
|
||||
Messages = [systemPrompt, ..messages],
|
||||
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -67,18 +76,19 @@ public sealed class ProviderGWDG() : BaseProvider(LLMProviders.GWDG, new Uri("ht
|
||||
}
|
||||
|
||||
/// <inhertidoc />
|
||||
public override Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts)
|
||||
public override async Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts)
|
||||
{
|
||||
return Task.FromResult<IReadOnlyList<IReadOnlyList<float>>>([]);
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.EMBEDDING_PROVIDER);
|
||||
return await this.PerformStandardTextEmbeddingRequest(requestedSecret, embeddingModel, token: token, texts: texts);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Where(model => !model.Id.StartsWith("e5-mistral-7b-instruct", StringComparison.InvariantCultureIgnoreCase))]
|
||||
Models = [..result.Models.Where(model => model.IsChatModel())]
|
||||
};
|
||||
}
|
||||
|
||||
@@ -89,12 +99,26 @@ public sealed class ProviderGWDG() : BaseProvider(LLMProviders.GWDG, new Uri("ht
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
/// <remarks>
|
||||
/// SAIA answers the models endpoint with its chat models only, so asking it for the embedding
|
||||
/// models comes back empty. We therefore fall back to the models the documentation names. The
|
||||
/// endpoint is still asked first: should SAIA start reporting them one day, its answer wins
|
||||
/// over our list. A failed request is passed on unchanged, so a wrong API key stays visible
|
||||
/// as such instead of being covered up by the fallback.
|
||||
/// </remarks>
|
||||
public override async Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, apiKeyProvisional, token);
|
||||
if (!result.Success)
|
||||
return result;
|
||||
|
||||
var embeddingModels = result.Models.Where(model => model.IsEmbeddingModel()).ToList();
|
||||
if (embeddingModels.Count is 0)
|
||||
return ModelLoadResult.FromModels(KNOWN_EMBEDDING_MODELS);
|
||||
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Where(model => model.Id.StartsWith("e5-", StringComparison.InvariantCultureIgnoreCase))]
|
||||
Models = [..embeddingModels]
|
||||
};
|
||||
}
|
||||
|
||||
@@ -110,18 +134,17 @@ public sealed class ProviderGWDG() : BaseProvider(LLMProviders.GWDG, new Uri("ht
|
||||
|
||||
#endregion
|
||||
|
||||
private async Task<ModelLoadResult> LoadModels(SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private async Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
var result = await this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data,
|
||||
token,
|
||||
apiKeyProvisional);
|
||||
apiKeyProvisional, token: token);
|
||||
|
||||
if (!result.Success)
|
||||
LOGGER.LogWarning("Failed to load models for provider {ProviderId}. FailureReason: {FailureReason}. TechnicalDetails: {TechnicalDetails}", this.Id, result.FailureReason, result.TechnicalDetails);
|
||||
|
||||
return result;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -31,7 +31,7 @@ public class ProviderGoogle() : BaseProvider(LLMProviders.GOOGLE, new Uri("https
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -47,6 +47,7 @@ public class ProviderGoogle() : BaseProvider(LLMProviders.GOOGLE, new Uri("https
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -146,7 +147,7 @@ public class ProviderGoogle() : BaseProvider(LLMProviders.GOOGLE, new Uri("https
|
||||
/// <inheritdoc />
|
||||
public override async Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models =
|
||||
@@ -167,7 +168,7 @@ public class ProviderGoogle() : BaseProvider(LLMProviders.GOOGLE, new Uri("https
|
||||
|
||||
public override async Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models =
|
||||
@@ -186,7 +187,7 @@ public class ProviderGoogle() : BaseProvider(LLMProviders.GOOGLE, new Uri("https
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
@@ -194,7 +195,6 @@ public class ProviderGoogle() : BaseProvider(LLMProviders.GOOGLE, new Uri("https
|
||||
modelResponse => modelResponse.Data
|
||||
.Where(model => !string.IsNullOrWhiteSpace(model.Id))
|
||||
.Select(model => new Model(this.NormalizeModelId(model.Id), model.DisplayName)),
|
||||
token,
|
||||
apiKeyProvisional,
|
||||
failureReasonSelector: (response, _) => response.StatusCode switch
|
||||
{
|
||||
@@ -202,7 +202,8 @@ public class ProviderGoogle() : BaseProvider(LLMProviders.GOOGLE, new Uri("https
|
||||
System.Net.HttpStatusCode.Unauthorized => ModelLoadFailureReason.INVALID_OR_MISSING_API_KEY,
|
||||
System.Net.HttpStatusCode.TooManyRequests => ModelLoadFailureReason.TOO_MANY_REQUESTS,
|
||||
_ => ModelLoadFailureReason.PROVIDER_UNAVAILABLE,
|
||||
});
|
||||
},
|
||||
token: token);
|
||||
}
|
||||
|
||||
private bool IsEmbeddingModel(string modelId)
|
||||
|
||||
@@ -29,7 +29,7 @@ public class ProviderGroq() : BaseProvider(LLMProviders.GROQ, new Uri("https://a
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
if (TryPopIntParameter(apiParameters, "seed", out var parsedSeed))
|
||||
apiParameters["seed"] = parsedSeed;
|
||||
@@ -48,6 +48,7 @@ public class ProviderGroq() : BaseProvider(LLMProviders.GROQ, new Uri("https://a
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -64,9 +65,10 @@ public class ProviderGroq() : BaseProvider(LLMProviders.GROQ, new Uri("https://a
|
||||
#pragma warning restore CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
public override async Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(TranscriptionResult.Failure());
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.TRANSCRIPTION_PROVIDER);
|
||||
return await this.PerformStandardTranscriptionRequest(requestedSecret, transcriptionModel, audioFilePath, token: token);
|
||||
}
|
||||
|
||||
/// <inhertidoc />
|
||||
@@ -76,9 +78,13 @@ public class ProviderGroq() : BaseProvider(LLMProviders.GROQ, new Uri("https://a
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
public override async Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Where(model => model.IsChatModel())]
|
||||
};
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -94,23 +100,23 @@ public class ProviderGroq() : BaseProvider(LLMProviders.GROQ, new Uri("https://a
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
public override async Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
var result = await this.LoadModels(SecretStoreType.TRANSCRIPTION_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Where(model => model.IsTranscriptionModel())]
|
||||
};
|
||||
}
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data.Where(n =>
|
||||
!n.Id.StartsWith("whisper-", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.StartsWith("distil-", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("-tts", StringComparison.OrdinalIgnoreCase)),
|
||||
token,
|
||||
apiKeyProvisional);
|
||||
modelResponse => modelResponse.Data,
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -31,7 +31,7 @@ public sealed class ProviderHelmholtz() : BaseProvider(LLMProviders.HELMHOLTZ, n
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -46,6 +46,7 @@ public sealed class ProviderHelmholtz() : BaseProvider(LLMProviders.HELMHOLTZ, n
|
||||
Messages = [systemPrompt, ..messages],
|
||||
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -62,9 +63,10 @@ public sealed class ProviderHelmholtz() : BaseProvider(LLMProviders.HELMHOLTZ, n
|
||||
#pragma warning restore CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
public override async Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(TranscriptionResult.Failure());
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.TRANSCRIPTION_PROVIDER);
|
||||
return await this.PerformStandardTranscriptionRequest(requestedSecret, transcriptionModel, audioFilePath, token: token);
|
||||
}
|
||||
|
||||
/// <inhertidoc />
|
||||
@@ -77,14 +79,12 @@ public sealed class ProviderHelmholtz() : BaseProvider(LLMProviders.HELMHOLTZ, n
|
||||
/// <inheritdoc />
|
||||
public override async Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models =
|
||||
[
|
||||
..result.Models.Where(model => !model.Id.StartsWith("text-", StringComparison.InvariantCultureIgnoreCase) &&
|
||||
!model.Id.Contains("-embedding", StringComparison.InvariantCultureIgnoreCase)
|
||||
)
|
||||
..result.Models.Where(model => model.IsChatModel())
|
||||
]
|
||||
};
|
||||
}
|
||||
@@ -98,28 +98,32 @@ public sealed class ProviderHelmholtz() : BaseProvider(LLMProviders.HELMHOLTZ, n
|
||||
/// <inheritdoc />
|
||||
public override async Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models =
|
||||
[
|
||||
..result.Models.Where(model =>
|
||||
model.Id.Contains("-embedding", StringComparison.InvariantCultureIgnoreCase) ||
|
||||
model.Id.StartsWith("text-", StringComparison.InvariantCultureIgnoreCase) ||
|
||||
model.Id.Contains("gritlm", StringComparison.InvariantCultureIgnoreCase))
|
||||
..result.Models.Where(model => model.IsEmbeddingModel())
|
||||
]
|
||||
};
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
public override async Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
var result = await this.LoadModels(SecretStoreType.TRANSCRIPTION_PROVIDER, apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models =
|
||||
[
|
||||
..result.Models.Where(model => model.IsTranscriptionModel())
|
||||
]
|
||||
};
|
||||
}
|
||||
|
||||
#endregion
|
||||
|
||||
private async Task<ModelLoadResult> LoadModels(SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private async Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
var secretKey = await this.GetModelLoadingSecretKey(storeType, apiKeyProvisional);
|
||||
if (string.IsNullOrWhiteSpace(secretKey))
|
||||
@@ -161,4 +165,4 @@ public sealed class ProviderHelmholtz() : BaseProvider(LLMProviders.HELMHOLTZ, n
|
||||
return FailedModelLoadResult(ModelLoadFailureReason.PROVIDER_UNAVAILABLE, e.Message);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
using System.Runtime.CompilerServices;
|
||||
|
||||
using AIStudio.Chat;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Settings;
|
||||
|
||||
namespace AIStudio.Provider.Hetzner;
|
||||
|
||||
public sealed class ProviderHetzner() : BaseProvider(LLMProviders.HETZNER, new Uri("https://inference.hetzner.com/api/v1/"), ExternalHttpTrustPolicy.SYSTEM_TRUST_ONLY, LOGGER)
|
||||
{
|
||||
private static readonly ILogger<ProviderHetzner> LOGGER = Program.LOGGER_FACTORY.CreateLogger<ProviderHetzner>();
|
||||
|
||||
#region Implementation of IProvider
|
||||
|
||||
/// <inheritdoc />
|
||||
public override string Id => LLMProviders.HETZNER.ToSecretId();
|
||||
|
||||
/// <inheritdoc />
|
||||
public override string InstanceName { get; set; } = "Hetzner (Experimental)";
|
||||
|
||||
/// <inheritdoc />
|
||||
public override bool HasModelLoadingCapability => true;
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ContentStreamChunk> StreamChatCompletion(Model chatModel, ChatThread chatThread, SettingsManager settingsManager, [EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
await foreach (var content in this.StreamOpenAICompatibleChatCompletion<ChatCompletionAPIRequest, ChatCompletionDeltaStreamLine, NoChatCompletionAnnotationStreamLine>(
|
||||
"Hetzner",
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
|
||||
return new ChatCompletionAPIRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
Messages = [systemPrompt, ..messages],
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
token: token))
|
||||
yield return content;
|
||||
}
|
||||
|
||||
#pragma warning disable CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ImageURL> StreamImageCompletion(Model imageModel, string promptPositive, string promptNegative = FilterOperator.String.Empty, ImageURL referenceImageURL = default, [EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
yield break;
|
||||
}
|
||||
#pragma warning restore CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(TranscriptionResult.Failure());
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts)
|
||||
{
|
||||
return Task.FromResult<IReadOnlyList<IReadOnlyList<float>>>([]);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(SecretStoreType.LLM_PROVIDER, "models", modelResponse => modelResponse.Data.Where(model => model.IsChatModel()), apiKeyProvisional, token: token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetImageModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
}
|
||||
|
||||
#endregion
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
|
||||
/// <summary>
|
||||
/// Which of the Hugging Face endpoints a provider instance talks to.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Hugging Face serves chatting and everything else from different places. Chat completions go to
|
||||
/// the router's own OpenAI-compatible endpoint, which accepts the model IDs as the hub writes them
|
||||
/// and picks an inference provider from a suffix. Embeddings and transcription do not exist there
|
||||
/// at all and have to be asked of one provider's own route. Because the base URL is fixed when a
|
||||
/// provider instance is built, the instance has to know from the start which one it is for.
|
||||
/// </remarks>
|
||||
public enum HFEndpointKind
|
||||
{
|
||||
/// <summary>
|
||||
/// The router's own endpoint, which serves chat completions.
|
||||
/// </summary>
|
||||
CHAT,
|
||||
|
||||
/// <summary>
|
||||
/// The OpenAI-compatible route of one inference provider, which serves embeddings.
|
||||
/// </summary>
|
||||
EMBEDDING,
|
||||
|
||||
/// <summary>
|
||||
/// The OpenAI-compatible route of one inference provider, which transcribes audio.
|
||||
/// </summary>
|
||||
TRANSCRIPTION,
|
||||
}
|
||||
@@ -3,16 +3,46 @@
|
||||
/// <summary>
|
||||
/// Enum for inference providers that Hugging Face supports.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Besides the providers themselves, this enum carries the routing strategies Hugging Face offers.
|
||||
/// They are no providers, but they take the same place: the router picks a provider for us instead
|
||||
/// of us naming one.
|
||||
///
|
||||
/// NONE must stay the first value: settings are read through the tolerant enum converter, which
|
||||
/// falls back to the first value whenever it meets a name we no longer know. That is what happens
|
||||
/// to a configuration naming one of the providers Hugging Face stopped routing in July 2026
|
||||
/// (Hyperbolic, SambaNova, Nebius, NVIDIA, Clarifai, Black Forest Labs), and to one naming the
|
||||
/// Hugging Face Inference API, which serves no model we can reach: it has no chat models at all,
|
||||
/// and its OpenAI-compatible routes for embeddings and transcription do not exist. Such a provider
|
||||
/// has to end up on NONE, where the validation asks the user to choose again. Were a routing
|
||||
/// strategy first, those configurations would silently switch to automatic routing instead.
|
||||
/// </remarks>
|
||||
public enum HFInferenceProvider
|
||||
{
|
||||
NONE,
|
||||
|
||||
|
||||
//
|
||||
// Routing strategies. Hugging Face writes them where a provider name would go:
|
||||
//
|
||||
AUTOMATIC,
|
||||
CHEAPEST,
|
||||
PREFERRED,
|
||||
|
||||
//
|
||||
// The providers Hugging Face routes:
|
||||
//
|
||||
BASETEN,
|
||||
CEREBRAS,
|
||||
NEBIUS_AI_STUDIO,
|
||||
SAMBANOVA,
|
||||
NOVITA,
|
||||
HYPERBOLIC,
|
||||
TOGETHER_AI,
|
||||
COHERE,
|
||||
DEEPINFRA,
|
||||
FEATHERLESS_AI,
|
||||
FIREWORKS,
|
||||
HF_INFERENCE_API,
|
||||
GROQ,
|
||||
NOVITA,
|
||||
NSCALE,
|
||||
OVHCLOUD,
|
||||
PUBLIC_AI,
|
||||
SCALEWAY,
|
||||
TOGETHER_AI,
|
||||
ZAI,
|
||||
}
|
||||
@@ -1,43 +1,158 @@
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
using AIStudio.Tools.PluginSystem;
|
||||
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
|
||||
public static class HFInferenceProviderExtensions
|
||||
{
|
||||
public static string Endpoints(this HFInferenceProvider provider, Model model) => provider switch
|
||||
{
|
||||
HFInferenceProvider.CEREBRAS => "cerebras/v1/",
|
||||
HFInferenceProvider.NEBIUS_AI_STUDIO => "nebius/v1/",
|
||||
HFInferenceProvider.SAMBANOVA => "sambanova/v1/",
|
||||
HFInferenceProvider.NOVITA => "novita/v3/openai/",
|
||||
HFInferenceProvider.HYPERBOLIC => "hyperbolic/v1/",
|
||||
HFInferenceProvider.TOGETHER_AI => "together/v1/",
|
||||
HFInferenceProvider.FIREWORKS => "fireworks-ai/inference/v1/",
|
||||
HFInferenceProvider.HF_INFERENCE_API => $"hf-inference/models/{model.ToString()}/v1/",
|
||||
_ => string.Empty,
|
||||
};
|
||||
|
||||
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(HFInferenceProviderExtensions).Namespace, nameof(HFInferenceProviderExtensions));
|
||||
|
||||
/// <summary>
|
||||
/// The slug Hugging Face uses for this inference provider.
|
||||
/// </summary>
|
||||
/// <param name="provider">The inference provider.</param>
|
||||
/// <returns>The slug, or an empty string for the routing strategies, which name no provider.</returns>
|
||||
public static string EndpointsId(this HFInferenceProvider provider) => provider switch
|
||||
{
|
||||
HFInferenceProvider.BASETEN => "baseten",
|
||||
HFInferenceProvider.CEREBRAS => "cerebras",
|
||||
HFInferenceProvider.NEBIUS_AI_STUDIO => "nebius",
|
||||
HFInferenceProvider.SAMBANOVA => "sambanova",
|
||||
HFInferenceProvider.COHERE => "cohere",
|
||||
HFInferenceProvider.DEEPINFRA => "deepinfra",
|
||||
HFInferenceProvider.FEATHERLESS_AI => "featherless-ai",
|
||||
HFInferenceProvider.FIREWORKS => "fireworks-ai",
|
||||
HFInferenceProvider.GROQ => "groq",
|
||||
HFInferenceProvider.NOVITA => "novita",
|
||||
HFInferenceProvider.HYPERBOLIC => "hyperbolic",
|
||||
HFInferenceProvider.NSCALE => "nscale",
|
||||
HFInferenceProvider.OVHCLOUD => "ovhcloud",
|
||||
HFInferenceProvider.PUBLIC_AI => "publicai",
|
||||
HFInferenceProvider.SCALEWAY => "scaleway",
|
||||
HFInferenceProvider.TOGETHER_AI => "together",
|
||||
HFInferenceProvider.FIREWORKS => "fireworks",
|
||||
HFInferenceProvider.HF_INFERENCE_API => "hf-inference",
|
||||
HFInferenceProvider.ZAI => "zai-org",
|
||||
|
||||
_ => string.Empty,
|
||||
};
|
||||
|
||||
|
||||
/// <summary>
|
||||
/// The suffix which tells the router where to send the request.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The router serves every provider through one endpoint. Which provider answers is decided by
|
||||
/// a suffix on the model name, e.g. "google/gemma-4-31B-it:novita". Without a suffix, the router
|
||||
/// picks the fastest provider itself.
|
||||
/// </remarks>
|
||||
/// <param name="provider">The inference provider.</param>
|
||||
/// <returns>The suffix including its colon, or an empty string when the router should choose.</returns>
|
||||
public static string ModelSuffix(this HFInferenceProvider provider) => provider switch
|
||||
{
|
||||
HFInferenceProvider.NONE or HFInferenceProvider.AUTOMATIC => string.Empty,
|
||||
|
||||
HFInferenceProvider.CHEAPEST => ":cheapest",
|
||||
HFInferenceProvider.PREFERRED => ":preferred",
|
||||
|
||||
_ => $":{provider.EndpointsId()}",
|
||||
};
|
||||
|
||||
/// <summary>
|
||||
/// Whether this inference provider serves models to chat with.
|
||||
/// </summary>
|
||||
/// <param name="provider">The inference provider.</param>
|
||||
/// <returns>True, when the provider serves chat models.</returns>
|
||||
public static bool SupportsChat(this HFInferenceProvider provider) => provider is not HFInferenceProvider.NONE;
|
||||
|
||||
/// <summary>
|
||||
/// Whether this inference provider creates embeddings for us.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Embeddings are a much shorter story than chatting. The router serves them nowhere near its
|
||||
/// own endpoint, only through the route of a provider, and only two of those answer the
|
||||
/// OpenAI-compatible form we send. The routing strategies are out by their nature: without a
|
||||
/// named provider there is no route to address.
|
||||
/// </remarks>
|
||||
/// <param name="provider">The inference provider.</param>
|
||||
/// <returns>True, when we can create embeddings through this provider.</returns>
|
||||
public static bool SupportsEmbeddings(this HFInferenceProvider provider) => provider is HFInferenceProvider.TOGETHER_AI or HFInferenceProvider.DEEPINFRA;
|
||||
|
||||
/// <summary>
|
||||
/// Whether this inference provider transcribes audio for us.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The same two providers as for embeddings, and for the same reason: transcription lives on a
|
||||
/// provider's own route, and only these two answer the OpenAI-compatible form there. Others do
|
||||
/// transcribe for Hugging Face, but not in a shape we could send an audio file to: fal-ai and
|
||||
/// Replicate both turn the request down with "Model not supported by provider".
|
||||
/// </remarks>
|
||||
/// <param name="provider">The inference provider.</param>
|
||||
/// <returns>True, when we can transcribe audio through this provider.</returns>
|
||||
public static bool SupportsTranscription(this HFInferenceProvider provider) => provider is HFInferenceProvider.TOGETHER_AI or HFInferenceProvider.DEEPINFRA;
|
||||
|
||||
/// <summary>
|
||||
/// The base URL of the provider's own OpenAI-compatible route.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Only chatting goes through the router's own endpoint. Everything else has to address the
|
||||
/// provider directly, and they do not agree on where their OpenAI-compatible API sits: DeepInfra
|
||||
/// keeps it below an additional "openai" segment, and answers the path without it with
|
||||
/// "Not allowed to POST /v1/embeddings for provider deepinfra".
|
||||
/// </remarks>
|
||||
/// <param name="provider">The inference provider.</param>
|
||||
/// <returns>The base URL, or an empty string when the provider has no such route.</returns>
|
||||
public static string ProviderBaseURL(this HFInferenceProvider provider) => provider switch
|
||||
{
|
||||
HFInferenceProvider.TOGETHER_AI => "https://router.huggingface.co/together/v1/",
|
||||
HFInferenceProvider.DEEPINFRA => "https://router.huggingface.co/deepinfra/v1/openai/",
|
||||
|
||||
_ => string.Empty,
|
||||
};
|
||||
|
||||
/// <summary>
|
||||
/// Removes the routing suffix from a model, if it carries one.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The suffix says where a request goes, not what the model is. Everything asking what a model
|
||||
/// can do has to look at the bare name: "google/gemma-4-31B-it:novita" is the same model as
|
||||
/// "google/gemma-4-31B-it", and a name detection which never heard of the suffix would miss it.
|
||||
/// Model IDs on the hub are written as "org/model" and carry no colon of their own, so the last
|
||||
/// colon always starts the suffix.
|
||||
/// </remarks>
|
||||
/// <param name="model">The model as it is configured.</param>
|
||||
/// <returns>The model without its routing suffix.</returns>
|
||||
public static Model WithoutRoutingSuffix(this Model model)
|
||||
{
|
||||
var separatorIndex = model.Id.LastIndexOf(':');
|
||||
return separatorIndex is -1 ? model : model with { Id = model.Id[..separatorIndex] };
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// The value to filter the Hugging Face model catalog by.
|
||||
/// </summary>
|
||||
/// <param name="provider">The inference provider.</param>
|
||||
/// <returns>The provider slug, or "all" when no particular provider was chosen.</returns>
|
||||
public static string CatalogFilter(this HFInferenceProvider provider)
|
||||
{
|
||||
var slug = provider.EndpointsId();
|
||||
return string.IsNullOrEmpty(slug) ? "all" : slug;
|
||||
}
|
||||
|
||||
public static string ToName(this HFInferenceProvider provider) => provider switch
|
||||
{
|
||||
HFInferenceProvider.AUTOMATIC => TB("Automatic: the fastest provider"),
|
||||
HFInferenceProvider.CHEAPEST => TB("Automatic: the cheapest provider"),
|
||||
HFInferenceProvider.PREFERRED => TB("Automatic: your preferred order"),
|
||||
|
||||
HFInferenceProvider.BASETEN => "Baseten",
|
||||
HFInferenceProvider.CEREBRAS => "Cerebras",
|
||||
HFInferenceProvider.NEBIUS_AI_STUDIO => "Nebius AI Studio",
|
||||
HFInferenceProvider.SAMBANOVA => "Sambanova",
|
||||
HFInferenceProvider.NOVITA => "Novita",
|
||||
HFInferenceProvider.HYPERBOLIC => "Hyperbolic",
|
||||
HFInferenceProvider.TOGETHER_AI => "Together AI",
|
||||
HFInferenceProvider.COHERE => "Cohere",
|
||||
HFInferenceProvider.DEEPINFRA => "DeepInfra",
|
||||
HFInferenceProvider.FEATHERLESS_AI => "Featherless AI",
|
||||
HFInferenceProvider.FIREWORKS => "Fireworks AI",
|
||||
HFInferenceProvider.HF_INFERENCE_API => "Hugging Face Inference API",
|
||||
HFInferenceProvider.GROQ => "Groq",
|
||||
HFInferenceProvider.NOVITA => "Novita",
|
||||
HFInferenceProvider.NSCALE => "Nscale",
|
||||
HFInferenceProvider.OVHCLOUD => "OVHcloud",
|
||||
HFInferenceProvider.PUBLIC_AI => "Public AI",
|
||||
HFInferenceProvider.SCALEWAY => "Scaleway",
|
||||
HFInferenceProvider.TOGETHER_AI => "Together AI",
|
||||
HFInferenceProvider.ZAI => "Z.ai",
|
||||
|
||||
_ => string.Empty,
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
|
||||
/// <summary>
|
||||
/// One model as the Hugging Face router describes it.
|
||||
/// </summary>
|
||||
/// <param name="Id">The ID of the model, written as "org/model".</param>
|
||||
/// <param name="Providers">The inference providers serving this model.</param>
|
||||
public readonly record struct HFModel(string Id, IList<HFModelProvider>? Providers);
|
||||
@@ -0,0 +1,8 @@
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
|
||||
/// <summary>
|
||||
/// One inference provider serving a model.
|
||||
/// </summary>
|
||||
/// <param name="Provider">The slug of the inference provider, e.g. "novita".</param>
|
||||
/// <param name="Status">Whether the provider currently serves the model. Known value: "live".</param>
|
||||
public readonly record struct HFModelProvider(string Provider, string Status);
|
||||
@@ -0,0 +1,12 @@
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
|
||||
/// <summary>
|
||||
/// One model as the Hugging Face hub lists it.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The hub answers with a plain array of models and describes each of them in far more detail than
|
||||
/// we need here, from tags to download counts. We only ever ask for the models of one provider and
|
||||
/// one task, so the ID is all that is left to read.
|
||||
/// </remarks>
|
||||
/// <param name="Id">The ID of the model, written as "org/model".</param>
|
||||
public readonly record struct HubModel(string Id);
|
||||
@@ -0,0 +1,12 @@
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
|
||||
/// <summary>
|
||||
/// A data model for the response from the model endpoint of the Hugging Face router.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The router says more about a model than the OpenAI model list does: which inference providers
|
||||
/// serve it, and which kinds of input it takes. That is why this provider brings its own data model
|
||||
/// instead of using the shared one.
|
||||
/// </remarks>
|
||||
/// <param name="Data">The models the router knows.</param>
|
||||
public readonly record struct ModelsResponse(IList<HFModel> Data);
|
||||
@@ -1,8 +1,10 @@
|
||||
using System.Runtime.CompilerServices;
|
||||
using System.Net;
|
||||
using System.Runtime.CompilerServices;
|
||||
|
||||
using AIStudio.Chat;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Settings;
|
||||
using AIStudio.Tools.PluginSystem;
|
||||
|
||||
namespace AIStudio.Provider.HuggingFace;
|
||||
|
||||
@@ -10,11 +12,138 @@ public sealed class ProviderHuggingFace : BaseProvider
|
||||
{
|
||||
private static readonly ILogger<ProviderHuggingFace> LOGGER = Program.LOGGER_FACTORY.CreateLogger<ProviderHuggingFace>();
|
||||
|
||||
public ProviderHuggingFace(HFInferenceProvider hfProvider, Model model) : base(LLMProviders.HUGGINGFACE, new Uri($"https://router.huggingface.co/{hfProvider.Endpoints(model)}"), ExternalHttpTrustPolicy.SYSTEM_TRUST_ONLY, LOGGER)
|
||||
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(ProviderHuggingFace).Namespace, nameof(ProviderHuggingFace));
|
||||
|
||||
/// <summary>
|
||||
/// The OpenAI-compatible endpoint which serves every inference provider.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Hugging Face also keeps a route per provider, such as "/novita/v3/openai/". Those expect the
|
||||
/// model ID as that provider spells it, which differs from the ID on the hub: Novita knows
|
||||
/// "google/gemma-4-31B-it" as "google/gemma-4-31b-it", and the router is case-sensitive. Asking
|
||||
/// for the hub spelling there is answered with "Model not supported by provider novita". This
|
||||
/// endpoint takes the hub spelling and translates it for us, so it is the one we use.
|
||||
/// </remarks>
|
||||
private const string ROUTER_BASE_URL = "https://router.huggingface.co/v1/";
|
||||
|
||||
/// <summary>
|
||||
/// Where the models of an inference provider are listed.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The router lists the chat models it routes, but nothing else. Which embedding models a
|
||||
/// provider offers is known to the hub alone, which answers this without a token. The URL is
|
||||
/// absolute on purpose: it addresses the hub, not the router this provider is built on.
|
||||
/// </remarks>
|
||||
private const string HUB_MODELS_URL = "https://huggingface.co/api/models?limit=100&sort=downloads&direction=-1&inference_provider=";
|
||||
|
||||
private readonly HFInferenceProvider hfProvider;
|
||||
|
||||
public ProviderHuggingFace(HFInferenceProvider hfProvider, HFEndpointKind endpointKind = HFEndpointKind.CHAT) : base(LLMProviders.HUGGINGFACE, new Uri(BuildBaseURL(hfProvider, endpointKind)), ExternalHttpTrustPolicy.SYSTEM_TRUST_ONLY, LOGGER)
|
||||
{
|
||||
LOGGER.LogInformation($"We use the inference provider '{hfProvider}'. Thus we use the base URL 'https://router.huggingface.co/{hfProvider.Endpoints(model)}'.");
|
||||
this.hfProvider = hfProvider;
|
||||
LOGGER.LogInformation($"We use the inference provider '{hfProvider}' for {endpointKind}. Thus, we use the base URL '{BuildBaseURL(hfProvider, endpointKind)}'.");
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Determines the base URL for the endpoint this provider instance talks to.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// A provider which serves no embeddings has no route of its own to offer, and neither have the
|
||||
/// routing strategies. We still have to hand a URL to the base class, so we fall back to the
|
||||
/// router. A request sent there is answered with a plain "Not Found", which is the honest
|
||||
/// outcome: the user selected something we told them we cannot do, and the validation of the
|
||||
/// dialog says so before it ever comes to a request.
|
||||
/// </remarks>
|
||||
/// <param name="hfProvider">The chosen inference provider.</param>
|
||||
/// <param name="endpointKind">Which endpoint this instance is built for.</param>
|
||||
/// <returns>The base URL to use.</returns>
|
||||
private static string BuildBaseURL(HFInferenceProvider hfProvider, HFEndpointKind endpointKind)
|
||||
{
|
||||
if (endpointKind is HFEndpointKind.CHAT)
|
||||
return ROUTER_BASE_URL;
|
||||
|
||||
var providerBaseURL = hfProvider.ProviderBaseURL();
|
||||
return string.IsNullOrEmpty(providerBaseURL) ? ROUTER_BASE_URL : providerBaseURL;
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Builds the model name to send to the router.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The router picks the inference provider from a suffix on the model name. When the user wrote
|
||||
/// a suffix themselves, we keep theirs: appending a second one would name a model nobody knows.
|
||||
/// </remarks>
|
||||
/// <param name="model">The model the user chose.</param>
|
||||
/// <returns>The model name including the provider suffix, when one applies.</returns>
|
||||
private string BuildModelIdentifier(Model model)
|
||||
{
|
||||
var modelId = model.Id;
|
||||
if (string.IsNullOrWhiteSpace(modelId) || modelId.Contains(':'))
|
||||
return modelId;
|
||||
|
||||
return $"{modelId}{this.hfProvider.ModelSuffix()}";
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Recognizes the router's answer for a model the chosen inference provider does not serve.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Not every model is available at every inference provider, and the router says so with a bad
|
||||
/// request. Without this, the user would be told that the message format might have changed,
|
||||
/// which points them at something they cannot fix and away from the one thing they can: picking
|
||||
/// another provider. The router words this failure as the error code "model_not_supported",
|
||||
/// while the providers behind it word it as a sentence of their own.
|
||||
/// </remarks>
|
||||
/// <param name="value">A piece of the failed response: an error code, a message, or the body.</param>
|
||||
/// <returns>True, when this text names an unsupported model.</returns>
|
||||
private static bool IsModelNotSupportedError(string? value)
|
||||
{
|
||||
if (string.IsNullOrWhiteSpace(value))
|
||||
return false;
|
||||
|
||||
return value.Contains("model_not_supported", StringComparison.OrdinalIgnoreCase) ||
|
||||
value.Contains("not supported by provider", StringComparison.OrdinalIgnoreCase) ||
|
||||
value.Contains("not supported by any provider", StringComparison.OrdinalIgnoreCase);
|
||||
}
|
||||
|
||||
#region Overrides of BaseProvider
|
||||
|
||||
/// <inheritdoc />
|
||||
protected override ProviderRequestFailureReason ClassifyProviderRequestFailure(HttpStatusCode statusCode, string responseBody)
|
||||
{
|
||||
if (statusCode is HttpStatusCode.BadRequest && IsModelNotSupportedError(responseBody))
|
||||
return ProviderRequestFailureReason.MODEL_NOT_SUPPORTED_BY_PROVIDER;
|
||||
|
||||
return base.ClassifyProviderRequestFailure(statusCode, responseBody);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
protected override ProviderRequestFailureReason ClassifyProviderRequestFailure(string? errorCode, string? errorType, string? errorMessage, string responseBody)
|
||||
{
|
||||
if (IsModelNotSupportedError(errorCode) || IsModelNotSupportedError(errorType) || IsModelNotSupportedError(errorMessage))
|
||||
return ProviderRequestFailureReason.MODEL_NOT_SUPPORTED_BY_PROVIDER;
|
||||
|
||||
return base.ClassifyProviderRequestFailure(errorCode, errorType, errorMessage, responseBody);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
protected override string GetProviderRequestFailureUserMessage(ProviderRequestFailureReason failureReason)
|
||||
{
|
||||
if (failureReason is not ProviderRequestFailureReason.MODEL_NOT_SUPPORTED_BY_PROVIDER)
|
||||
return base.GetProviderRequestFailureUserMessage(failureReason);
|
||||
|
||||
//
|
||||
// When Hugging Face chose the provider itself, naming it back to the user would help
|
||||
// nobody: they never picked it, and no other choice of provider is left to try:
|
||||
//
|
||||
if (this.hfProvider is HFInferenceProvider.NONE or HFInferenceProvider.AUTOMATIC or HFInferenceProvider.CHEAPEST or HFInferenceProvider.PREFERRED)
|
||||
return TB("No Hugging Face inference provider offers the selected model. Please check the model name and whether it is still available on Hugging Face.");
|
||||
|
||||
return string.Format(TB("The Hugging Face inference provider '{0}' does not offer the selected model. Please select another inference provider, or let Hugging Face choose one for you."), this.hfProvider.ToName());
|
||||
}
|
||||
|
||||
#endregion
|
||||
|
||||
#region Implementation of IProvider
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -24,7 +153,7 @@ public sealed class ProviderHuggingFace : BaseProvider
|
||||
public override string InstanceName { get; set; } = "HuggingFace";
|
||||
|
||||
/// <inheritdoc />
|
||||
public override bool HasModelLoadingCapability => false;
|
||||
public override bool HasModelLoadingCapability => true;
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ContentStreamChunk> StreamChatCompletion(Model chatModel, ChatThread chatThread, SettingsManager settingsManager, [EnumeratorCancellation] CancellationToken token = default)
|
||||
@@ -34,14 +163,14 @@ public sealed class ProviderHuggingFace : BaseProvider
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
|
||||
return new ChatCompletionAPIRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
Model = this.BuildModelIdentifier(chatModel),
|
||||
|
||||
// Build the messages:
|
||||
// - First of all the system prompt
|
||||
@@ -49,6 +178,7 @@ public sealed class ProviderHuggingFace : BaseProvider
|
||||
Messages = [systemPrompt, ..messages],
|
||||
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -65,21 +195,66 @@ public sealed class ProviderHuggingFace : BaseProvider
|
||||
#pragma warning restore CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
public override async Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(TranscriptionResult.Failure());
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.TRANSCRIPTION_PROVIDER);
|
||||
|
||||
//
|
||||
// Note that we send the model as it is: this request goes to the provider's own route,
|
||||
// where a routing suffix would be part of the name and name nothing:
|
||||
//
|
||||
return await this.PerformStandardTranscriptionRequest(requestedSecret, transcriptionModel, audioFilePath, token: token);
|
||||
}
|
||||
|
||||
/// <inhertidoc />
|
||||
public override Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts)
|
||||
public override async Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts)
|
||||
{
|
||||
return Task.FromResult<IReadOnlyList<IReadOnlyList<float>>>([]);
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.EMBEDDING_PROVIDER);
|
||||
|
||||
//
|
||||
// Note that we send the model as it is: this request goes to the provider's own route,
|
||||
// where a routing suffix would be part of the name and name nothing:
|
||||
//
|
||||
return await this.PerformStandardTextEmbeddingRequest(requestedSecret, embeddingModel, token: token, texts: texts);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
return this.LoadModelsResponse<ModelsResponse>(SecretStoreType.LLM_PROVIDER, "models", this.SelectChatModels, apiKeyProvisional, token: token);
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Picks the models the user may chat with through the chosen inference provider.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The router reports every model it knows, together with the providers serving it. When the
|
||||
/// user named a provider, we show what that provider offers and nothing else. Showing more
|
||||
/// would be a disservice: every model outside that list is answered with a bad request, and the
|
||||
/// user would only learn about it once they try to chat.
|
||||
/// </remarks>
|
||||
/// <param name="response">The response of the model endpoint.</param>
|
||||
/// <returns>The models to offer.</returns>
|
||||
private IEnumerable<Model> SelectChatModels(ModelsResponse response)
|
||||
{
|
||||
var chatModels = response.Data.Where(hfModel => new Model(hfModel.Id, null).IsChatModel());
|
||||
var providerSlug = this.hfProvider.EndpointsId();
|
||||
if (string.IsNullOrEmpty(providerSlug))
|
||||
return ToModels(chatModels);
|
||||
|
||||
return ToModels(chatModels.Where(hfModel => IsServedBy(hfModel, providerSlug)));
|
||||
}
|
||||
|
||||
private static IEnumerable<Model> ToModels(IEnumerable<HFModel> hfModels) => hfModels.Select(hfModel => new Model(hfModel.Id, null));
|
||||
|
||||
private static bool IsServedBy(HFModel hfModel, string providerSlug)
|
||||
{
|
||||
if (hfModel.Providers is null)
|
||||
return false;
|
||||
|
||||
return hfModel.Providers.Any(provider =>
|
||||
string.Equals(provider.Provider, providerSlug, StringComparison.OrdinalIgnoreCase) &&
|
||||
string.Equals(provider.Status, "live", StringComparison.OrdinalIgnoreCase));
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -91,14 +266,34 @@ public sealed class ProviderHuggingFace : BaseProvider
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
if (!this.hfProvider.SupportsEmbeddings())
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
|
||||
return this.LoadHubModels(SecretStoreType.EMBEDDING_PROVIDER, "feature-extraction", apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Loads the models one inference provider offers for a task, as the hub lists them.
|
||||
/// </summary>
|
||||
/// <param name="storeType">Which stored API key to use.</param>
|
||||
/// <param name="pipelineTag">The task to ask for, as the hub names it.</param>
|
||||
/// <param name="apiKeyProvisional">An API key which is not stored yet.</param>
|
||||
/// <param name="token">The cancellation token to use.</param>
|
||||
/// <returns>The models of that provider for that task.</returns>
|
||||
private Task<ModelLoadResult> LoadHubModels(SecretStoreType storeType, string pipelineTag, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
var requestURL = $"{HUB_MODELS_URL}{this.hfProvider.EndpointsId()}&pipeline_tag={pipelineTag}";
|
||||
return this.LoadModelsResponse<IList<HubModel>>(storeType, requestURL, hubModels => hubModels.Select(hubModel => new Model(hubModel.Id, null)), apiKeyProvisional, token: token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
if (!this.hfProvider.SupportsTranscription())
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
|
||||
return this.LoadHubModels(SecretStoreType.TRANSCRIPTION_PROVIDER, "automatic-speech-recognition", apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
#endregion
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
using System.Runtime.CompilerServices;
|
||||
|
||||
using AIStudio.Chat;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Settings;
|
||||
|
||||
namespace AIStudio.Provider.IONOS;
|
||||
|
||||
public sealed class ProviderIONOS() : BaseProvider(LLMProviders.IONOS, new Uri("https://openai.inference.de-txl.ionos.com/v1/"), ExternalHttpTrustPolicy.SYSTEM_TRUST_ONLY, LOGGER)
|
||||
{
|
||||
/// <summary>
|
||||
/// IONOS keeps an alias of some embedding models around, so that customers can migrate away from
|
||||
/// the previous naming. Those aliases point to the very same models we already offer, which is
|
||||
/// why we hide them instead of listing every embedding model twice.
|
||||
/// </summary>
|
||||
private const string MIGRATION_ALIAS_SUFFIX = "-migration";
|
||||
|
||||
private static readonly ILogger<ProviderIONOS> LOGGER = Program.LOGGER_FACTORY.CreateLogger<ProviderIONOS>();
|
||||
|
||||
#region Implementation of IProvider
|
||||
|
||||
/// <inheritdoc />
|
||||
public override string Id => LLMProviders.IONOS.ToSecretId();
|
||||
|
||||
/// <inheritdoc />
|
||||
public override string InstanceName { get; set; } = "IONOS";
|
||||
|
||||
/// <inheritdoc />
|
||||
public override bool HasModelLoadingCapability => true;
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ContentStreamChunk> StreamChatCompletion(Model chatModel, ChatThread chatThread, SettingsManager settingsManager, [EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
await foreach (var content in this.StreamOpenAICompatibleChatCompletion<ChatCompletionAPIRequest, ChatCompletionDeltaStreamLine, NoChatCompletionAnnotationStreamLine>(
|
||||
"IONOS",
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
|
||||
return new ChatCompletionAPIRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
|
||||
// Build the messages:
|
||||
// - First of all the system prompt
|
||||
// - Then none-empty user and AI messages
|
||||
Messages = [systemPrompt, ..messages],
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
token: token))
|
||||
yield return content;
|
||||
}
|
||||
|
||||
#pragma warning disable CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ImageURL> StreamImageCompletion(Model imageModel, string promptPositive, string promptNegative = FilterOperator.String.Empty, ImageURL referenceImageURL = default, [EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
yield break;
|
||||
}
|
||||
#pragma warning restore CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(TranscriptionResult.Failure());
|
||||
}
|
||||
|
||||
/// <inhertidoc />
|
||||
public override async Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts)
|
||||
{
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.EMBEDDING_PROVIDER);
|
||||
return await this.PerformStandardTextEmbeddingRequest(requestedSecret, embeddingModel, token: token, texts: texts);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, model => model.IsChatModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetImageModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, model => model.IsEmbeddingModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
}
|
||||
|
||||
#endregion
|
||||
|
||||
/// <summary>
|
||||
/// Loads the models of one kind from IONOS.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// IONOS serves chat, embedding, reranking, OCR, and image models through one endpoint, and its
|
||||
/// response tells us nothing but the model's name. We therefore let the shared model kind
|
||||
/// detection sort them apart.
|
||||
/// </remarks>
|
||||
/// <param name="storeType">The secret store to read the API key from.</param>
|
||||
/// <param name="isWantedKind">Decides whether a model belongs to the requested kind.</param>
|
||||
/// <param name="apiKeyProvisional">An API key which was not stored yet.</param>
|
||||
/// <param name="token">The cancellation token.</param>
|
||||
/// <returns>The models of the requested kind.</returns>
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, Func<Model, bool> isWantedKind, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data
|
||||
.Where(model => !model.Id.EndsWith(MIGRATION_ALIAS_SUFFIX, StringComparison.OrdinalIgnoreCase))
|
||||
.Where(isWantedKind),
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
}
|
||||
@@ -78,10 +78,14 @@ public interface IProvider
|
||||
/// <summary>
|
||||
/// Embed a text file.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The cancellation token is not the last parameter, unlike everywhere else in this codebase:
|
||||
/// C# demands that a params parameter comes last, and every implementation inherits that order.
|
||||
/// </remarks>
|
||||
/// <param name="embeddingModel">The model to use for embedding.</param>
|
||||
/// <param name="settingsManager">The settings manager instance to use.</param>
|
||||
/// <param name="token">The cancellation token.</param>
|
||||
/// /// <param name="texts">A single string or a list of strings to embed.</param>
|
||||
/// <param name="texts">A single string or a list of strings to embed.</param>
|
||||
/// <returns>>The embedded text as a single vector or as a list of vectors.</returns>
|
||||
public Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts);
|
||||
|
||||
|
||||
@@ -16,6 +16,9 @@ public enum LLMProviders
|
||||
ALIBABA_CLOUD = 12,
|
||||
PERPLEXITY = 14,
|
||||
OPEN_ROUTER = 15,
|
||||
HETZNER = 16,
|
||||
IONOS = 17,
|
||||
LITE_LLM = 18,
|
||||
|
||||
FIREWORKS = 5,
|
||||
GROQ = 6,
|
||||
|
||||
@@ -6,7 +6,10 @@ using AIStudio.Provider.Google;
|
||||
using AIStudio.Provider.Groq;
|
||||
using AIStudio.Provider.GWDG;
|
||||
using AIStudio.Provider.Helmholtz;
|
||||
using AIStudio.Provider.Hetzner;
|
||||
using AIStudio.Provider.HuggingFace;
|
||||
using AIStudio.Provider.IONOS;
|
||||
using AIStudio.Provider.LiteLLM;
|
||||
using AIStudio.Provider.Mistral;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Provider.OpenRouter;
|
||||
@@ -56,11 +59,14 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.ALIBABA_CLOUD => "Alibaba Cloud",
|
||||
LLMProviders.PERPLEXITY => "Perplexity",
|
||||
LLMProviders.OPEN_ROUTER => "OpenRouter",
|
||||
LLMProviders.HETZNER => "Hetzner (Experimental)",
|
||||
LLMProviders.IONOS => "IONOS",
|
||||
LLMProviders.LITE_LLM => "LiteLLM",
|
||||
|
||||
LLMProviders.GROQ => "Groq",
|
||||
LLMProviders.FIREWORKS => "Fireworks.ai",
|
||||
LLMProviders.HUGGINGFACE => "Hugging Face",
|
||||
|
||||
|
||||
LLMProviders.SELF_HOSTED => translate ? TB("Self-hosted") : "Self-hosted",
|
||||
|
||||
LLMProviders.HELMHOLTZ => "Helmholtz Blablador",
|
||||
@@ -91,6 +97,9 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.ALIBABA_CLOUD => "Alibaba Cloud",
|
||||
LLMProviders.PERPLEXITY => "Perplexity",
|
||||
LLMProviders.OPEN_ROUTER => "OpenRouter",
|
||||
LLMProviders.HETZNER => "Hetzner",
|
||||
LLMProviders.IONOS => "IONOS",
|
||||
LLMProviders.LITE_LLM => "LiteLLM",
|
||||
|
||||
LLMProviders.GROQ => "Groq",
|
||||
LLMProviders.FIREWORKS => "Fireworks.ai",
|
||||
@@ -144,6 +153,23 @@ public static class LLMProvidersExtensions
|
||||
|
||||
LLMProviders.OPEN_ROUTER => Confidence.USA_HUB.WithRegion("America, U.S.").WithSources("https://openrouter.ai/privacy", "https://openrouter.ai/terms").WithLevel(settingsManager.GetConfiguredConfidenceLevel(llmProvider)),
|
||||
|
||||
LLMProviders.HETZNER => Confidence.GDPR_EXPERIMENTAL_OPEN_SOURCE.WithRegion("Europe, Germany").WithSources(
|
||||
"https://experiments.hetzner.com/docs/inference",
|
||||
"https://www.hetzner.com/legal/privacy-policy/",
|
||||
"https://www.hetzner.com/legal/terms-and-conditions/"
|
||||
).WithLevel(settingsManager.GetConfiguredConfidenceLevel(llmProvider)),
|
||||
|
||||
LLMProviders.IONOS => Confidence.GDPR_NO_TRAINING.WithRegion("Europe, Germany").WithSources(
|
||||
"https://docs.ionos.com/cloud/ai/ai-model-hub/governance-and-compliance/data-handling",
|
||||
"https://www.ionos.com/terms-gtc/privacy-policy/"
|
||||
).WithLevel(settingsManager.GetConfiguredConfidenceLevel(llmProvider)),
|
||||
|
||||
// LiteLLM is a gateway the user runs, but it is not a self-hosted LLM: the proxy owner decides
|
||||
// which downstream providers it routes to, and those are usually cloud services. Self-hosting
|
||||
// the proxy therefore says nothing about where the data ends up, so we do not claim the trust
|
||||
// of a self-hosted model here and let the user assign the level themselves.
|
||||
LLMProviders.LITE_LLM => Confidence.USER_OPERATED_GATEWAY.WithSources("https://docs.litellm.ai/docs/data_security").WithLevel(settingsManager.GetConfiguredConfidenceLevel(llmProvider)),
|
||||
|
||||
LLMProviders.SELF_HOSTED => Confidence.SELF_HOSTED.WithLevel(settingsManager.GetConfiguredConfidenceLevel(llmProvider)),
|
||||
|
||||
LLMProviders.HELMHOLTZ => Confidence.GDPR_NO_TRAINING.WithRegion("Europe, Germany").WithSources("https://helmholtz.cloud/services/?serviceID=d7d5c597-a2f6-4bd1-b71e-4d6499d98570").WithLevel(settingsManager.GetConfiguredConfidenceLevel(llmProvider)),
|
||||
@@ -167,7 +193,11 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.GOOGLE => true,
|
||||
LLMProviders.HELMHOLTZ => true,
|
||||
LLMProviders.ALIBABA_CLOUD => true,
|
||||
|
||||
LLMProviders.IONOS => true,
|
||||
LLMProviders.GWDG => true,
|
||||
LLMProviders.OPEN_ROUTER => true,
|
||||
LLMProviders.LITE_LLM => true,
|
||||
|
||||
//
|
||||
// Providers that do not support embeddings:
|
||||
//
|
||||
@@ -175,20 +205,26 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.ANTHROPIC => false,
|
||||
LLMProviders.FIREWORKS => false,
|
||||
LLMProviders.X => false,
|
||||
LLMProviders.GWDG => false,
|
||||
LLMProviders.DEEP_SEEK => false,
|
||||
LLMProviders.HUGGINGFACE => false,
|
||||
LLMProviders.PERPLEXITY => false,
|
||||
LLMProviders.OPEN_ROUTER => true,
|
||||
LLMProviders.HETZNER => false,
|
||||
|
||||
//
|
||||
// Hugging Face serves embeddings, but not through the router endpoint we chat with: that
|
||||
// one answers "/v1/embeddings" with a plain "Not Found". They have to be asked of one
|
||||
// inference provider directly, and only some of them answer the OpenAI-compatible form.
|
||||
// Which ones is decided by HFInferenceProviderExtensions.SupportsEmbeddings.
|
||||
//
|
||||
LLMProviders.HUGGINGFACE => true,
|
||||
|
||||
//
|
||||
// Self-hosted providers are treated as a special case anyway.
|
||||
//
|
||||
LLMProviders.SELF_HOSTED => true,
|
||||
|
||||
|
||||
_ => false,
|
||||
};
|
||||
|
||||
|
||||
public static bool ProvideTranscriptionAPI(this LLMProviders llmProvider) => llmProvider switch
|
||||
{
|
||||
//
|
||||
@@ -198,7 +234,10 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.MISTRAL => true,
|
||||
LLMProviders.FIREWORKS => true,
|
||||
LLMProviders.GWDG => true,
|
||||
|
||||
LLMProviders.HELMHOLTZ => true,
|
||||
LLMProviders.GROQ => true,
|
||||
LLMProviders.LITE_LLM => true,
|
||||
|
||||
//
|
||||
// Providers that support transcription but provide no OpenAI-compatible API yet:
|
||||
//
|
||||
@@ -209,20 +248,26 @@ public static class LLMProvidersExtensions
|
||||
// Providers that do not support transcription:
|
||||
//
|
||||
LLMProviders.OPEN_ROUTER => false,
|
||||
LLMProviders.GROQ => false,
|
||||
LLMProviders.HETZNER => false,
|
||||
LLMProviders.IONOS => false,
|
||||
LLMProviders.ANTHROPIC => false,
|
||||
LLMProviders.X => false,
|
||||
LLMProviders.DEEP_SEEK => false,
|
||||
LLMProviders.HUGGINGFACE => false,
|
||||
LLMProviders.PERPLEXITY => false,
|
||||
|
||||
LLMProviders.HELMHOLTZ => false,
|
||||
|
||||
//
|
||||
// Hugging Face transcribes audio, but like embeddings, not through the router endpoint we
|
||||
// chat with: that one answers "/v1/audio/transcriptions" with a plain "Not Found". Only
|
||||
// some of the inference providers answer the OpenAI-compatible form, which
|
||||
// HFInferenceProviderExtensions.SupportsTranscription decides.
|
||||
//
|
||||
LLMProviders.HUGGINGFACE => true,
|
||||
|
||||
//
|
||||
// Self-hosted providers are treated as a special case anyway.
|
||||
//
|
||||
LLMProviders.SELF_HOSTED => true,
|
||||
|
||||
|
||||
_ => false,
|
||||
};
|
||||
|
||||
@@ -233,7 +278,7 @@ public static class LLMProvidersExtensions
|
||||
/// <returns>The provider instance.</returns>
|
||||
public static IProvider CreateProvider(this AIStudio.Settings.Provider providerSettings)
|
||||
{
|
||||
return providerSettings.UsedLLMProvider.CreateProvider(providerSettings.InstanceName, providerSettings.Host, providerSettings.Hostname, providerSettings.Model, providerSettings.HFInferenceProvider, providerSettings.Id, providerSettings.AdditionalJsonApiParameters, tokenizerPath: providerSettings.TokenizerPath, providerSettings.IsEnterpriseConfiguration);
|
||||
return providerSettings.UsedLLMProvider.CreateProvider(providerSettings.InstanceName, providerSettings.Host, providerSettings.Hostname, providerSettings.HFInferenceProvider, providerSettings.Id, providerSettings.AdditionalJsonApiParameters, providerSettings.IsEnterpriseConfiguration, capabilityOverrides: providerSettings.CapabilityOverrides, tokenizerPath: providerSettings.TokenizerPath);
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
@@ -243,7 +288,7 @@ public static class LLMProvidersExtensions
|
||||
/// <returns>The provider instance.</returns>
|
||||
public static IProvider CreateProvider(this EmbeddingProvider embeddingProviderSettings)
|
||||
{
|
||||
return embeddingProviderSettings.UsedLLMProvider.CreateProvider(embeddingProviderSettings.Name, embeddingProviderSettings.Host, embeddingProviderSettings.Hostname, embeddingProviderSettings.Model, HFInferenceProvider.NONE, configuredProviderId: embeddingProviderSettings.Id, tokenizerPath: embeddingProviderSettings.TokenizerPath, isEnterpriseConfiguration: embeddingProviderSettings.IsEnterpriseConfiguration);
|
||||
return embeddingProviderSettings.UsedLLMProvider.CreateProvider(embeddingProviderSettings.Name, embeddingProviderSettings.Host, embeddingProviderSettings.Hostname, embeddingProviderSettings.HFInferenceProvider, configuredProviderId: embeddingProviderSettings.Id, isEnterpriseConfiguration: embeddingProviderSettings.IsEnterpriseConfiguration, hfEndpointKind: HFEndpointKind.EMBEDDING, tokenizerPath: embeddingProviderSettings.TokenizerPath);
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
@@ -253,14 +298,14 @@ public static class LLMProvidersExtensions
|
||||
/// <returns>The provider instance.</returns>
|
||||
public static IProvider CreateProvider(this TranscriptionProvider transcriptionProviderSettings)
|
||||
{
|
||||
return transcriptionProviderSettings.UsedLLMProvider.CreateProvider(transcriptionProviderSettings.Name, transcriptionProviderSettings.Host, transcriptionProviderSettings.Hostname, transcriptionProviderSettings.Model, HFInferenceProvider.NONE, configuredProviderId: transcriptionProviderSettings.Id, tokenizerPath: string.Empty, isEnterpriseConfiguration: transcriptionProviderSettings.IsEnterpriseConfiguration);
|
||||
return transcriptionProviderSettings.UsedLLMProvider.CreateProvider(transcriptionProviderSettings.Name, transcriptionProviderSettings.Host, transcriptionProviderSettings.Hostname, transcriptionProviderSettings.HFInferenceProvider, configuredProviderId: transcriptionProviderSettings.Id, isEnterpriseConfiguration: transcriptionProviderSettings.IsEnterpriseConfiguration, hfEndpointKind: HFEndpointKind.TRANSCRIPTION);
|
||||
}
|
||||
|
||||
private static IProvider CreateProvider(this LLMProviders provider, string instanceName, Host host, string hostname, Model model, HFInferenceProvider inferenceProvider, string configuredProviderId = "", string expertProviderApiParameter = "", string tokenizerPath = "", bool isEnterpriseConfiguration = false)
|
||||
|
||||
private static IProvider CreateProvider(this LLMProviders provider, string instanceName, Host host, string hostname, HFInferenceProvider inferenceProvider, string configuredProviderId = "", string expertProviderApiParameter = "", bool isEnterpriseConfiguration = false, HFEndpointKind hfEndpointKind = HFEndpointKind.CHAT, ProviderCapabilityOverrides? capabilityOverrides = null, string tokenizerPath = "")
|
||||
{
|
||||
try
|
||||
{
|
||||
return provider switch
|
||||
IProvider providerInstance = provider switch
|
||||
{
|
||||
LLMProviders.OPEN_AI => new ProviderOpenAI { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.ANTHROPIC => new ProviderAnthropic { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
@@ -271,10 +316,13 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.ALIBABA_CLOUD => new ProviderAlibabaCloud { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.PERPLEXITY => new ProviderPerplexity { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.OPEN_ROUTER => new ProviderOpenRouter { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.HETZNER => new ProviderHetzner { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.IONOS => new ProviderIONOS { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.LITE_LLM => new ProviderLiteLLM(hostname) { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
|
||||
LLMProviders.GROQ => new ProviderGroq { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.FIREWORKS => new ProviderFireworks { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.HUGGINGFACE => new ProviderHuggingFace(inferenceProvider, model) { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
LLMProviders.HUGGINGFACE => new ProviderHuggingFace(inferenceProvider, hfEndpointKind) { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
|
||||
LLMProviders.SELF_HOSTED => new ProviderSelfHosted(host, hostname) { InstanceName = instanceName, ConfiguredProviderId = configuredProviderId, AdditionalJsonApiParameters = expertProviderApiParameter, TokenizerPath = tokenizerPath, IsEnterpriseConfiguration = isEnterpriseConfiguration },
|
||||
|
||||
@@ -283,6 +331,11 @@ public static class LLMProvidersExtensions
|
||||
|
||||
_ => new NoProvider(),
|
||||
};
|
||||
|
||||
if (providerInstance is BaseProvider baseProvider)
|
||||
baseProvider.CapabilityOverrides = capabilityOverrides;
|
||||
|
||||
return providerInstance;
|
||||
}
|
||||
catch (Exception e)
|
||||
{
|
||||
@@ -302,6 +355,8 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.ALIBABA_CLOUD => "https://account.alibabacloud.com/register/intl_register.htm",
|
||||
LLMProviders.PERPLEXITY => "https://www.perplexity.ai/account/api",
|
||||
LLMProviders.OPEN_ROUTER => "https://openrouter.ai/keys",
|
||||
LLMProviders.HETZNER => "https://experiments.hetzner.com",
|
||||
LLMProviders.IONOS => "https://cloud.ionos.com/compute/sign-up",
|
||||
|
||||
LLMProviders.GROQ => "https://console.groq.com/",
|
||||
LLMProviders.FIREWORKS => "https://fireworks.ai/login",
|
||||
@@ -327,6 +382,8 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.PERPLEXITY => "https://www.perplexity.ai/account/api/",
|
||||
LLMProviders.OPEN_ROUTER => "https://openrouter.ai/activity",
|
||||
LLMProviders.HUGGINGFACE => "https://huggingface.co/settings/billing",
|
||||
LLMProviders.HETZNER => "https://experiments.hetzner.com",
|
||||
LLMProviders.IONOS => "https://dcd.ionos.com/latest/?page=dcd-ai-model-hub",
|
||||
|
||||
_ => string.Empty,
|
||||
};
|
||||
@@ -345,6 +402,8 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.PERPLEXITY => true,
|
||||
LLMProviders.OPEN_ROUTER => true,
|
||||
LLMProviders.HUGGINGFACE => true,
|
||||
LLMProviders.HETZNER => true,
|
||||
LLMProviders.IONOS => true,
|
||||
|
||||
_ => false,
|
||||
};
|
||||
@@ -352,14 +411,13 @@ public static class LLMProvidersExtensions
|
||||
public static string GetModelsOverviewURL(this LLMProviders provider, HFInferenceProvider inferenceProvider) => provider switch
|
||||
{
|
||||
LLMProviders.FIREWORKS => "https://fireworks.ai/models?show=Serverless",
|
||||
LLMProviders.HUGGINGFACE => $"https://huggingface.co/models?inference_provider={inferenceProvider.EndpointsId()}",
|
||||
LLMProviders.HUGGINGFACE => $"https://huggingface.co/models?inference_provider={inferenceProvider.CatalogFilter()}",
|
||||
_ => string.Empty,
|
||||
};
|
||||
|
||||
public static bool IsLLMModelProvidedManually(this LLMProviders provider) => provider switch
|
||||
{
|
||||
LLMProviders.FIREWORKS => true,
|
||||
LLMProviders.HUGGINGFACE => true,
|
||||
_ => false,
|
||||
};
|
||||
|
||||
@@ -408,6 +466,7 @@ public static class LLMProvidersExtensions
|
||||
public static bool IsHostnameNeeded(this LLMProviders provider) => provider switch
|
||||
{
|
||||
LLMProviders.SELF_HOSTED => true,
|
||||
LLMProviders.LITE_LLM => true,
|
||||
_ => false,
|
||||
};
|
||||
|
||||
@@ -422,15 +481,18 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.ALIBABA_CLOUD => true,
|
||||
LLMProviders.PERPLEXITY => true,
|
||||
LLMProviders.OPEN_ROUTER => true,
|
||||
LLMProviders.HETZNER => true,
|
||||
LLMProviders.IONOS => true,
|
||||
LLMProviders.LITE_LLM => true,
|
||||
|
||||
LLMProviders.GROQ => true,
|
||||
LLMProviders.FIREWORKS => true,
|
||||
LLMProviders.HELMHOLTZ => true,
|
||||
LLMProviders.GWDG => true,
|
||||
LLMProviders.HUGGINGFACE => true,
|
||||
|
||||
|
||||
LLMProviders.SELF_HOSTED => host is (Host.OLLAMA or Host.VLLM),
|
||||
|
||||
|
||||
_ => false,
|
||||
};
|
||||
|
||||
@@ -445,6 +507,8 @@ public static class LLMProvidersExtensions
|
||||
LLMProviders.ALIBABA_CLOUD => true,
|
||||
LLMProviders.PERPLEXITY => true,
|
||||
LLMProviders.OPEN_ROUTER => true,
|
||||
LLMProviders.HETZNER => true,
|
||||
LLMProviders.IONOS => true,
|
||||
|
||||
LLMProviders.GROQ => true,
|
||||
LLMProviders.FIREWORKS => true,
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
namespace AIStudio.Provider;
|
||||
|
||||
public static class LLMProvidersIconExtensions
|
||||
{
|
||||
private const string ICON_ROOT = "/images/provider-icons";
|
||||
private const string SVG_DATA_URL_PREFIX = "data:image/svg+xml;base64,";
|
||||
|
||||
public static string GetIconUrl(this AIStudio.Settings.Provider provider, bool isDarkMode)
|
||||
=> provider.UsedLLMProvider.GetIconUrl(isDarkMode, provider.CustomIconDataUrl);
|
||||
|
||||
public static string GetIconUrl(this LLMProviders provider, bool isDarkMode, string? customIconDataUrl)
|
||||
{
|
||||
if (customIconDataUrl?.StartsWith(SVG_DATA_URL_PREFIX, StringComparison.Ordinal) == true)
|
||||
return customIconDataUrl;
|
||||
|
||||
return provider.GetIconUrl(isDarkMode);
|
||||
}
|
||||
|
||||
public static string GetIconUrl(this LLMProviders provider, bool isDarkMode) => provider switch
|
||||
{
|
||||
LLMProviders.NONE => $"{ICON_ROOT}/provider{DarkVariant(isDarkMode)}.svg",
|
||||
LLMProviders.OPEN_AI => $"{ICON_ROOT}/openai{DarkVariant(isDarkMode)}.svg",
|
||||
LLMProviders.ANTHROPIC => $"{ICON_ROOT}/anthropic{DarkVariant(isDarkMode)}.svg",
|
||||
LLMProviders.MISTRAL => $"{ICON_ROOT}/mistral.svg",
|
||||
LLMProviders.GOOGLE => $"{ICON_ROOT}/google.svg",
|
||||
LLMProviders.X => $"{ICON_ROOT}/x{DarkVariant(isDarkMode)}.svg",
|
||||
LLMProviders.DEEP_SEEK => $"{ICON_ROOT}/deepseek.svg",
|
||||
LLMProviders.ALIBABA_CLOUD => $"{ICON_ROOT}/alibaba-cloud.svg",
|
||||
LLMProviders.PERPLEXITY => $"{ICON_ROOT}/perplexity.svg",
|
||||
LLMProviders.OPEN_ROUTER => $"{ICON_ROOT}/openrouter.svg",
|
||||
LLMProviders.HETZNER => $"{ICON_ROOT}/hetzner.svg",
|
||||
LLMProviders.IONOS => $"{ICON_ROOT}/ionos.svg",
|
||||
LLMProviders.LITE_LLM => $"{ICON_ROOT}/litellm.svg",
|
||||
LLMProviders.GROQ => $"{ICON_ROOT}/groq.svg",
|
||||
LLMProviders.FIREWORKS => $"{ICON_ROOT}/fireworks.svg",
|
||||
LLMProviders.HUGGINGFACE => $"{ICON_ROOT}/hugging-face.svg",
|
||||
LLMProviders.SELF_HOSTED => $"{ICON_ROOT}/self-hosted{DarkVariant(isDarkMode)}.svg",
|
||||
LLMProviders.HELMHOLTZ => $"{ICON_ROOT}/helmholtz.svg",
|
||||
LLMProviders.GWDG => $"{ICON_ROOT}/gwdg.svg",
|
||||
_ => $"{ICON_ROOT}/provider{DarkVariant(isDarkMode)}.svg",
|
||||
};
|
||||
|
||||
private static string DarkVariant(bool isDarkMode) => isDarkMode ? "-dark" : string.Empty;
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
using System.Runtime.CompilerServices;
|
||||
|
||||
using AIStudio.Chat;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Settings;
|
||||
|
||||
namespace AIStudio.Provider.LiteLLM;
|
||||
|
||||
public sealed class ProviderLiteLLM(string hostname) : BaseProvider(LLMProviders.LITE_LLM, BuildBaseUri(hostname), ExternalHttpTrustPolicy.ALLOW_CUSTOM_ROOTS_WHEN_HOST_WHITELISTED, LOGGER)
|
||||
{
|
||||
private static readonly ILogger<ProviderLiteLLM> LOGGER = Program.LOGGER_FACTORY.CreateLogger<ProviderLiteLLM>();
|
||||
|
||||
#region Implementation of IProvider
|
||||
|
||||
/// <inheritdoc />
|
||||
public override string Id => LLMProviders.LITE_LLM.ToSecretId();
|
||||
|
||||
/// <inheritdoc />
|
||||
public override string InstanceName { get; set; } = "LiteLLM";
|
||||
|
||||
/// <inheritdoc />
|
||||
public override bool HasModelLoadingCapability => true;
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ContentStreamChunk> StreamChatCompletion(Model chatModel, ChatThread chatThread, SettingsManager settingsManager, [EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
await foreach (var content in this.StreamOpenAICompatibleChatCompletion<ChatCompletionAPIRequest, ChatCompletionDeltaStreamLine, NoChatCompletionAnnotationStreamLine>(
|
||||
"LiteLLM",
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingDirectImageUrlAsync(this.Provider, chatModel);
|
||||
|
||||
return new ChatCompletionAPIRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
|
||||
// Build the messages:
|
||||
// - First of all the system prompt
|
||||
// - Then none-empty user and AI messages
|
||||
Messages = [systemPrompt, ..messages],
|
||||
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
token: token))
|
||||
yield return content;
|
||||
}
|
||||
|
||||
#pragma warning disable CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
/// <inheritdoc />
|
||||
public override async IAsyncEnumerable<ImageURL> StreamImageCompletion(Model imageModel, string promptPositive, string promptNegative = FilterOperator.String.Empty, ImageURL referenceImageURL = default, [EnumeratorCancellation] CancellationToken token = default)
|
||||
{
|
||||
yield break;
|
||||
}
|
||||
#pragma warning restore CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async Task<TranscriptionResult> TranscribeAudioAsync(Model transcriptionModel, string audioFilePath, SettingsManager settingsManager, CancellationToken token = default)
|
||||
{
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.TRANSCRIPTION_PROVIDER);
|
||||
return await this.PerformStandardTranscriptionRequest(requestedSecret, transcriptionModel, audioFilePath, token: token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async Task<IReadOnlyList<IReadOnlyList<float>>> EmbedTextAsync(Model embeddingModel, SettingsManager settingsManager, CancellationToken token = default, params List<string> texts)
|
||||
{
|
||||
var requestedSecret = await Program.RUST_SERVICE.GetAPIKey(this, SecretStoreType.EMBEDDING_PROVIDER);
|
||||
return await this.PerformStandardTextEmbeddingRequest(requestedSecret, embeddingModel, token: token, texts: texts);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, static model => model.IsChatModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetImageModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return Task.FromResult(ModelLoadResult.FromModels([]));
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, static model => model.IsEmbeddingModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.TRANSCRIPTION_PROVIDER, static model => model.IsTranscriptionModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
#endregion
|
||||
|
||||
private static Uri BuildBaseUri(string hostname)
|
||||
{
|
||||
// LiteLLM exposes an OpenAI-compatible API under the "/v1/" path. Users configure the
|
||||
// base URL of their LiteLLM proxy (e.g. http://localhost:4000); we normalize any trailing
|
||||
// slash and append the OpenAI-compatible path.
|
||||
var normalizedHostname = hostname.TrimEnd('/');
|
||||
return new Uri($"{normalizedHostname}/v1/");
|
||||
}
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, Func<Model, bool> isWantedKind, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
//
|
||||
// The gateway serves every kind of model through one endpoint, so we have to sort
|
||||
// them apart ourselves. We use the shared model kind detection for that, which every
|
||||
// other provider uses as well:
|
||||
//
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data.Where(IsRealModel).Where(isWantedKind),
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Checks whether this entry is a model at all, or one of LiteLLM's wildcards.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// A LiteLLM configuration may pass a whole provider through at once, written as "openai/*" or
|
||||
/// just "*". Those patterns show up among the models, but they are no models: asking the gateway
|
||||
/// for one of them fails. No model carries an asterisk in its name, which makes it a safe mark.
|
||||
/// </remarks>
|
||||
/// <param name="model">The entry to check.</param>
|
||||
/// <returns>True, when the entry is a model rather than a wildcard.</returns>
|
||||
private static bool IsRealModel(Model model) => !model.Id.Contains('*');
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
namespace AIStudio.Provider.Mistral;
|
||||
|
||||
public readonly record struct Model(string Id, string Object, int Created, string OwnedBy);
|
||||
@@ -1,5 +1,3 @@
|
||||
namespace AIStudio.Provider.Mistral;
|
||||
|
||||
public readonly record struct ModelsResponse(string Object, Model[] Data);
|
||||
|
||||
public readonly record struct Model(string Id, string Object, int Created, string OwnedBy);
|
||||
public readonly record struct ModelsResponse(string Object, Model[] Data);
|
||||
@@ -29,7 +29,7 @@ public sealed class ProviderMistral() : BaseProvider(LLMProviders.MISTRAL, new U
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
if (TryPopBoolParameter(apiParameters, "safe_prompt", out var parsedSafePrompt))
|
||||
apiParameters["safe_prompt"] = parsedSafePrompt;
|
||||
@@ -51,6 +51,7 @@ public sealed class ProviderMistral() : BaseProvider(LLMProviders.MISTRAL, new U
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -91,10 +92,12 @@ public sealed class ProviderMistral() : BaseProvider(LLMProviders.MISTRAL, new U
|
||||
{
|
||||
Models =
|
||||
[
|
||||
// Codestral is a fill-in-the-middle model, which we cannot use for chats. That is
|
||||
// specific to Mistral's catalog, which is why it is not part of the shared model
|
||||
// kind detection:
|
||||
..modelResponse.Models.Where(n =>
|
||||
!n.Id.StartsWith("code", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("embed", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("moderation", StringComparison.OrdinalIgnoreCase))
|
||||
n.IsChatModel())
|
||||
]
|
||||
};
|
||||
}
|
||||
@@ -108,7 +111,7 @@ public sealed class ProviderMistral() : BaseProvider(LLMProviders.MISTRAL, new U
|
||||
|
||||
return modelResponse with
|
||||
{
|
||||
Models = [..modelResponse.Models.Where(n => n.Id.Contains("embed", StringComparison.InvariantCulture))]
|
||||
Models = [..modelResponse.Models.Where(n => n.IsEmbeddingModel())]
|
||||
};
|
||||
}
|
||||
|
||||
@@ -136,7 +139,6 @@ public sealed class ProviderMistral() : BaseProvider(LLMProviders.MISTRAL, new U
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data.Select(n => new Provider.Model(n.Id, null)),
|
||||
token,
|
||||
apiKeyProvisional);
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -42,15 +42,33 @@ public readonly record struct Model(string Id, string? DisplayName)
|
||||
|
||||
#endregion
|
||||
|
||||
#region Implementation of IEquatable<Model?>
|
||||
#region Implementation of IEquatable<Model>
|
||||
|
||||
public bool Equals(Model? other)
|
||||
{
|
||||
if(other is null)
|
||||
return false;
|
||||
|
||||
return this.Id == other.Value.Id;
|
||||
}
|
||||
/// <summary>
|
||||
/// Two models are the same model when they carry the same ID.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The display name is decoration. A provider may report a model under a display name of its own,
|
||||
/// while we know the very same model as a hardcoded fallback under a different one. Comparing the
|
||||
/// ID alone keeps those two the same model, so that removing duplicates works.
|
||||
///
|
||||
/// Note that this overload is the one the runtime uses, for example for Distinct(). The overload
|
||||
/// taking a nullable model below is a separate one and never gets called on its behalf, which is
|
||||
/// why the hash code has to follow this one.
|
||||
/// </remarks>
|
||||
/// <param name="other">The model to compare with.</param>
|
||||
/// <returns>True, when both models carry the same ID.</returns>
|
||||
public bool Equals(Model other) => string.Equals(this.Id, other.Id, StringComparison.Ordinal);
|
||||
|
||||
/// <summary>
|
||||
/// Two models are the same model when they carry the same ID.
|
||||
/// </summary>
|
||||
/// <param name="other">The model to compare with, which may be null.</param>
|
||||
/// <returns>True, when the other model exists and carries the same ID.</returns>
|
||||
public bool Equals(Model? other) => other is not null && this.Equals(other.Value);
|
||||
|
||||
/// <inheritdoc />
|
||||
public override int GetHashCode() => this.Id?.GetHashCode(StringComparison.Ordinal) ?? 0;
|
||||
|
||||
#endregion
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
namespace AIStudio.Provider;
|
||||
|
||||
/// <summary>
|
||||
/// The kind of an AI model, i.e. what the model is made for.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// This describes what kind of model we are dealing with. It answers a different question than the
|
||||
/// Capability enum: capabilities describe what a chat model is able to do, for example whether it
|
||||
/// accepts images or performs reasoning. Note that Capability.EMBEDDING marks a chat model which is
|
||||
/// able to create embeddings as well, whereas ModelKind.EMBEDDING marks a model whose only purpose
|
||||
/// is creating embeddings.
|
||||
/// </remarks>
|
||||
public enum ModelKind
|
||||
{
|
||||
/// <summary>
|
||||
/// The model is used for chat completions.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// This is the fallback: we report a model as a chat model whenever we do not recognize any
|
||||
/// other kind. Providers keep adding models we have never heard of, and a model we fail to
|
||||
/// recognize must stay visible to the user instead of silently disappearing from their list.
|
||||
/// </remarks>
|
||||
CHAT,
|
||||
|
||||
/// <summary>
|
||||
/// The model continues a text instead of answering in a conversation.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// These are the models from the era before chat completions, such as OpenAI's text-davinci-003.
|
||||
/// Some providers still offer them, but they only work through the completions endpoint. Asking
|
||||
/// them for a chat completion fails, so they must not show up as chat models.
|
||||
/// </remarks>
|
||||
TEXT_COMPLETION,
|
||||
|
||||
/// <summary>
|
||||
/// The model maps text or images into a vector space.
|
||||
/// </summary>
|
||||
EMBEDDING,
|
||||
|
||||
/// <summary>
|
||||
/// The model scores documents against a query to reorder search results.
|
||||
/// </summary>
|
||||
RERANKING,
|
||||
|
||||
/// <summary>
|
||||
/// The model generates or edits images.
|
||||
/// </summary>
|
||||
IMAGE_GENERATION,
|
||||
|
||||
/// <summary>
|
||||
/// The model generates or edits videos.
|
||||
/// </summary>
|
||||
VIDEO_GENERATION,
|
||||
|
||||
/// <summary>
|
||||
/// The model transcribes audio into text.
|
||||
/// </summary>
|
||||
TRANSCRIPTION,
|
||||
|
||||
/// <summary>
|
||||
/// The model speaks: it synthesizes speech from text, or answers in audio itself.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// This covers the pure text-to-speech models as well as those which hold a conversation in
|
||||
/// audio, such as the audio models of OpenAI. The latter do accept text, but they are made for
|
||||
/// spoken input and output, so they do not belong among the chat models.
|
||||
/// </remarks>
|
||||
SPEECH_SYNTHESIS,
|
||||
|
||||
/// <summary>
|
||||
/// The model holds a spoken conversation over a live connection.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// These models expect a streaming connection of their own, usually a WebSocket, instead of the
|
||||
/// chat completion API. They cannot be used for a normal chat.
|
||||
/// </remarks>
|
||||
REALTIME,
|
||||
|
||||
/// <summary>
|
||||
/// The model extracts text from images or scanned documents.
|
||||
/// </summary>
|
||||
OCR,
|
||||
|
||||
/// <summary>
|
||||
/// The model classifies content for policy violations.
|
||||
/// </summary>
|
||||
MODERATION,
|
||||
|
||||
/// <summary>
|
||||
/// Not a model at all.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Some providers list entries in their models endpoint which are no models, such as OpenAI's
|
||||
/// 'container' resource for its code interpreter. A provider talking to such an entry gets an
|
||||
/// error, so they must not appear in any of the model lists we show.
|
||||
/// </remarks>
|
||||
OTHER,
|
||||
}
|
||||
@@ -0,0 +1,219 @@
|
||||
namespace AIStudio.Provider;
|
||||
|
||||
/// <summary>
|
||||
/// Determines what kind of model we are dealing with, based on its name.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Many providers serve every kind of model through one models endpoint, without telling us what
|
||||
/// kind each model is. Before this class existed, every provider carried its own list of name
|
||||
/// fragments to sort those models apart. Those lists disagreed with each other: a model like
|
||||
/// nomic-embed-text was recognized as an embedding model by some providers, while others offered it
|
||||
/// as a chat model. The knowledge about model families is the same for all providers, so it lives
|
||||
/// here now.
|
||||
///
|
||||
/// This class recognizes what a model is NOT made for. Everything we do not recognize is reported as
|
||||
/// a chat model. That direction matters: when a provider adds a model family we have never seen, the
|
||||
/// user still gets to use it. Getting it wrong the other way around would hide a model the user is
|
||||
/// paying for.
|
||||
///
|
||||
/// What this class must not become is a place for provider-specific knowledge. That a model called
|
||||
/// "codestral" is a fill-in-the-middle model at Mistral, or that Alibaba's chat models all start
|
||||
/// with a "q", is true for that one provider only. Such rules stay in the provider.
|
||||
/// </remarks>
|
||||
public static class ModelKindExtensions
|
||||
{
|
||||
//
|
||||
// Checked first, because these entries are no models at all: whatever else their name might
|
||||
// suggest, none of the other kinds applies to them.
|
||||
//
|
||||
private static readonly string[] OTHER_MARKERS = ["container"];
|
||||
|
||||
//
|
||||
// Reranking is checked before embedding: rerankers are commonly named after the embedding model
|
||||
// they belong to, e.g. Qwen3-VL-Reranker-8B next to Qwen3-VL-Embedding-8B.
|
||||
//
|
||||
private static readonly string[] RERANKING_MARKERS = ["rerank"];
|
||||
|
||||
private static readonly string[] EMBEDDING_MARKERS = ["embed", "bge", "mpnet", "paraphrase", "sentence-transformers", "gte-", "e5-", "gritlm"];
|
||||
|
||||
//
|
||||
// The models from before chat completions existed. Providers keep offering some of them, and
|
||||
// Helmholtz Blablador still reports 'text-davinci-003', but asking any of them for a chat
|
||||
// completion fails. We deliberately do not look for 'ada' here: three letters appear in far too
|
||||
// many unrelated model names, and losing a chat model weighs heavier than keeping a dead one.
|
||||
//
|
||||
private static readonly string[] TEXT_COMPLETION_MARKERS = ["davinci", "babbage", "curie", "gpt-3.5-turbo-instruct"];
|
||||
|
||||
private static readonly string[] IMAGE_GENERATION_MARKERS = ["flux", "stable-diffusion", "sdxl", "dall-e", "midjourney", "gpt-image"];
|
||||
|
||||
private static readonly string[] VIDEO_GENERATION_MARKERS = ["sora", "veo-", "runway"];
|
||||
|
||||
//
|
||||
// Markers which have to stand as a word of their own. "kling" is such a case: taken as a plain
|
||||
// substring, it also matches the organization "Klingspor", the model "Inkling", and the
|
||||
// fine-tune "Llama-2-7b-chat-klingon" -- all of them models to chat with, which would vanish
|
||||
// from the user's list. The video models themselves are named "kling-v1" or "kling-video",
|
||||
// where the name ends at a separator.
|
||||
//
|
||||
private static readonly string[] VIDEO_GENERATION_WORD_MARKERS = ["kling"];
|
||||
|
||||
//
|
||||
// Voxtral is marketed as an audio model which understands speech, so one could expect it to work
|
||||
// in a chat as well. It does not: asking Mistral for a chat completion with 'voxtral-mini-latest'
|
||||
// is answered with 'Invalid model'. Voxtral therefore belongs here, next to the models which do
|
||||
// nothing but transcribe.
|
||||
//
|
||||
private static readonly string[] TRANSCRIPTION_MARKERS = ["whisper", "-transcribe", "wav2vec", "parakeet", "voxtral"];
|
||||
|
||||
//
|
||||
// Besides the pure text-to-speech models, this covers the models which answer in audio, such as
|
||||
// 'gpt-audio' and 'gpt-4o-audio-preview'. Those do accept a text-only request, but they are made
|
||||
// for spoken conversations, and the providers offering them directly keep them out of their chat
|
||||
// model lists as well.
|
||||
//
|
||||
private static readonly string[] SPEECH_SYNTHESIS_MARKERS = ["-tts", "tts-", "-speech", "speech-", "-audio", "audio-"];
|
||||
|
||||
//
|
||||
// The models for spoken conversations over a live connection. They speak their own protocol,
|
||||
// usually a WebSocket, and answer a chat completion request with an error. Checked before
|
||||
// transcription, because some of them carry the name of a transcription model, such as
|
||||
// OpenAI's 'gpt-realtime-whisper'. Those still need the live connection.
|
||||
//
|
||||
private static readonly string[] REALTIME_MARKERS = ["realtime"];
|
||||
|
||||
private static readonly string[] OCR_MARKERS = ["ocr"];
|
||||
|
||||
private static readonly string[] MODERATION_MARKERS = ["moderation", "guard"];
|
||||
|
||||
/// <summary>
|
||||
/// Determines what kind of model this is, based on its name.
|
||||
/// </summary>
|
||||
/// <param name="model">The model to inspect.</param>
|
||||
/// <returns>The recognized kind, or ModelKind.CHAT when we recognize no other kind.</returns>
|
||||
public static ModelKind DetermineKind(this Model model)
|
||||
{
|
||||
if (string.IsNullOrWhiteSpace(model.Id) || model.IsSystemModel)
|
||||
return ModelKind.CHAT;
|
||||
|
||||
if (HasAnyMarker(model.Id, OTHER_MARKERS))
|
||||
return ModelKind.OTHER;
|
||||
|
||||
if (HasAnyMarker(model.Id, RERANKING_MARKERS))
|
||||
return ModelKind.RERANKING;
|
||||
|
||||
if (HasAnyMarker(model.Id, EMBEDDING_MARKERS))
|
||||
return ModelKind.EMBEDDING;
|
||||
|
||||
if (HasAnyMarker(model.Id, TEXT_COMPLETION_MARKERS))
|
||||
return ModelKind.TEXT_COMPLETION;
|
||||
|
||||
if (HasAnyMarker(model.Id, IMAGE_GENERATION_MARKERS))
|
||||
return ModelKind.IMAGE_GENERATION;
|
||||
|
||||
if (HasAnyMarker(model.Id, VIDEO_GENERATION_MARKERS) || HasAnyWordMarker(model.Id, VIDEO_GENERATION_WORD_MARKERS))
|
||||
return ModelKind.VIDEO_GENERATION;
|
||||
|
||||
if (HasAnyMarker(model.Id, REALTIME_MARKERS))
|
||||
return ModelKind.REALTIME;
|
||||
|
||||
if (HasAnyMarker(model.Id, TRANSCRIPTION_MARKERS))
|
||||
return ModelKind.TRANSCRIPTION;
|
||||
|
||||
if (HasAnyMarker(model.Id, SPEECH_SYNTHESIS_MARKERS))
|
||||
return ModelKind.SPEECH_SYNTHESIS;
|
||||
|
||||
if (HasAnyMarker(model.Id, OCR_MARKERS))
|
||||
return ModelKind.OCR;
|
||||
|
||||
if (HasAnyMarker(model.Id, MODERATION_MARKERS))
|
||||
return ModelKind.MODERATION;
|
||||
|
||||
return ModelKind.CHAT;
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Checks whether this model can be used for chatting.
|
||||
/// </summary>
|
||||
/// <param name="model">The model to check.</param>
|
||||
/// <returns>True, when the model is a chat model or when we recognize no other kind.</returns>
|
||||
public static bool IsChatModel(this Model model) => model.DetermineKind() is ModelKind.CHAT;
|
||||
|
||||
/// <summary>
|
||||
/// Checks whether this model creates embeddings.
|
||||
/// </summary>
|
||||
/// <param name="model">The model to check.</param>
|
||||
/// <returns>True, when the model is an embedding model.</returns>
|
||||
public static bool IsEmbeddingModel(this Model model) => model.DetermineKind() is ModelKind.EMBEDDING;
|
||||
|
||||
/// <summary>
|
||||
/// Checks whether this model transcribes audio.
|
||||
/// </summary>
|
||||
/// <param name="model">The model to check.</param>
|
||||
/// <returns>True, when the model is a transcription model.</returns>
|
||||
public static bool IsTranscriptionModel(this Model model) => model.DetermineKind() is ModelKind.TRANSCRIPTION;
|
||||
|
||||
/// <summary>
|
||||
/// Checks whether this model generates images.
|
||||
/// </summary>
|
||||
/// <param name="model">The model to check.</param>
|
||||
/// <returns>True, when the model is an image generation model.</returns>
|
||||
public static bool IsImageModel(this Model model) => model.DetermineKind() is ModelKind.IMAGE_GENERATION;
|
||||
|
||||
private static bool HasAnyMarker(string modelId, string[] markers)
|
||||
{
|
||||
foreach (var marker in markers)
|
||||
if (modelId.Contains(marker, StringComparison.OrdinalIgnoreCase))
|
||||
return true;
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Checks whether the model name contains one of the markers as a word of its own.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// A short marker which is also a common syllable cannot be looked for as a plain substring:
|
||||
/// it would match names which have nothing to do with it, and the model would be sorted into
|
||||
/// the wrong kind. Such a marker counts only where a name segment begins and ends with it.
|
||||
/// </remarks>
|
||||
/// <param name="modelId">The ID of the model.</param>
|
||||
/// <param name="markers">The markers to look for.</param>
|
||||
/// <returns>True, when one of the markers stands as a word of its own.</returns>
|
||||
private static bool HasAnyWordMarker(string modelId, string[] markers)
|
||||
{
|
||||
foreach (var marker in markers)
|
||||
{
|
||||
var searchIndex = 0;
|
||||
while (searchIndex <= modelId.Length - marker.Length)
|
||||
{
|
||||
var markerIndex = modelId.IndexOf(marker, searchIndex, StringComparison.OrdinalIgnoreCase);
|
||||
if (markerIndex is -1)
|
||||
break;
|
||||
|
||||
if (IsWholeWord(modelId, marker, markerIndex))
|
||||
return true;
|
||||
|
||||
// The same marker may appear again later in the name, so we keep looking:
|
||||
searchIndex = markerIndex + 1;
|
||||
}
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
private static bool IsWholeWord(string modelId, string marker, int markerIndex)
|
||||
{
|
||||
if (markerIndex > 0 && !IsSeparator(modelId[markerIndex - 1]))
|
||||
return false;
|
||||
|
||||
var endIndex = markerIndex + marker.Length;
|
||||
return endIndex >= modelId.Length || IsSeparator(modelId[endIndex]);
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// The characters which separate the parts of a model name, such as in "fal-ai/kling-video".
|
||||
/// </summary>
|
||||
/// <param name="character">The character to check.</param>
|
||||
/// <returns>True, when the character separates two parts of a name.</returns>
|
||||
private static bool IsSeparator(char character) => character is '/' or '-' or '_' or '.' or ' ' or ':';
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
using System.Text.Json;
|
||||
using System.Text.Json.Serialization;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
public sealed record AssistantToolCallMessage : IMessageBase
|
||||
{
|
||||
public string Role { get; init; } = "assistant";
|
||||
|
||||
public JsonElement? Content { get; init; }
|
||||
|
||||
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
|
||||
public string? ReasoningContent { get; init; }
|
||||
|
||||
public IList<ChatCompletionToolCall> ToolCalls { get; init; } = [];
|
||||
}
|
||||
@@ -17,8 +17,14 @@ public record ChatCompletionAPIRequest(
|
||||
public ChatCompletionAPIRequest() : this(string.Empty, [], true)
|
||||
{
|
||||
}
|
||||
|
||||
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
|
||||
public IList<object>? Tools { get; init; }
|
||||
|
||||
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
|
||||
public bool? ParallelToolCalls { get; init; }
|
||||
|
||||
// Attention: The "required" modifier is not supported for [JsonExtensionData].
|
||||
[JsonExtensionData]
|
||||
public IDictionary<string, object> AdditionalApiParameters { get; init; } = new Dictionary<string, object>();
|
||||
}
|
||||
}
|
||||
@@ -7,7 +7,7 @@ namespace AIStudio.Provider.OpenAI;
|
||||
/// <param name="Delta">The delta text of the choice.</param>
|
||||
public record ChatCompletionChoice(int Index, ChatCompletionDelta Delta)
|
||||
{
|
||||
public ChatCompletionChoice() : this(0, new (string.Empty))
|
||||
public ChatCompletionChoice() : this(0, new())
|
||||
{
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
using System.Text;
|
||||
using System.Text.Json;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
internal static class ChatCompletionContent
|
||||
{
|
||||
public static string? GetText(JsonElement? content)
|
||||
{
|
||||
if (content is not { } value)
|
||||
return null;
|
||||
|
||||
if (value.ValueKind is JsonValueKind.String)
|
||||
return value.GetString();
|
||||
|
||||
if (value.ValueKind is not JsonValueKind.Array)
|
||||
return null;
|
||||
|
||||
var text = new StringBuilder();
|
||||
foreach (var chunk in value.EnumerateArray())
|
||||
{
|
||||
if (chunk.ValueKind is JsonValueKind.String)
|
||||
{
|
||||
text.Append(chunk.GetString());
|
||||
continue;
|
||||
}
|
||||
|
||||
if (chunk.ValueKind is not JsonValueKind.Object ||
|
||||
!chunk.TryGetProperty("type", out var type) ||
|
||||
type.ValueKind is not JsonValueKind.String ||
|
||||
!string.Equals(type.GetString(), "text", StringComparison.Ordinal) ||
|
||||
!chunk.TryGetProperty("text", out var textElement) ||
|
||||
textElement.ValueKind is not JsonValueKind.String)
|
||||
continue;
|
||||
|
||||
text.Append(textElement.GetString());
|
||||
}
|
||||
|
||||
return text.Length == 0 ? null : text.ToString();
|
||||
}
|
||||
}
|
||||
@@ -1,12 +1,16 @@
|
||||
using System.Text.Json;
|
||||
using System.Text.Json.Serialization;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// The delta text of a choice.
|
||||
/// </summary>
|
||||
/// <param name="Content">The content of the delta text.</param>
|
||||
public record ChatCompletionDelta(string Content)
|
||||
public sealed record ChatCompletionDelta
|
||||
{
|
||||
public ChatCompletionDelta() : this(string.Empty)
|
||||
{
|
||||
}
|
||||
}
|
||||
[JsonPropertyName("content")]
|
||||
public JsonElement? RawContent { get; init; }
|
||||
|
||||
[JsonIgnore]
|
||||
public string Content => ChatCompletionContent.GetText(this.RawContent) ?? string.Empty;
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
public sealed record ChatCompletionResponse
|
||||
{
|
||||
public string Id { get; init; } = string.Empty;
|
||||
|
||||
public string Model { get; init; } = string.Empty;
|
||||
|
||||
public IList<ChatCompletionResponseChoice> Choices { get; init; } = [];
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
public sealed record ChatCompletionResponseChoice
|
||||
{
|
||||
public int Index { get; init; }
|
||||
|
||||
public string FinishReason { get; init; } = string.Empty;
|
||||
|
||||
public ChatCompletionResponseMessage Message { get; init; } = new();
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
using System.Text.Json;
|
||||
using System.Text.Json.Serialization;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
public sealed record ChatCompletionResponseMessage
|
||||
{
|
||||
public string Role { get; init; } = string.Empty;
|
||||
|
||||
[JsonPropertyName("content")]
|
||||
public JsonElement? RawContent { get; init; }
|
||||
|
||||
[JsonIgnore]
|
||||
public string? Content => ChatCompletionContent.GetText(this.RawContent);
|
||||
|
||||
public string? ReasoningContent { get; init; }
|
||||
|
||||
public IList<ChatCompletionToolCall?>? ToolCalls { get; init; }
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
using System.Text.Json;
|
||||
using System.Text.Json.Serialization;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
public sealed record ChatCompletionToolCall
|
||||
{
|
||||
public string? Id { get; init; }
|
||||
|
||||
public string? Type { get; init; } = "function";
|
||||
|
||||
public ChatCompletionToolFunction? Function { get; init; }
|
||||
|
||||
[JsonExtensionData]
|
||||
public IDictionary<string, JsonElement> AdditionalMetadata { get; init; } = new Dictionary<string, JsonElement>();
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
using System.Text.Json;
|
||||
|
||||
using AIStudio.Tools.ToolCallingSystem;
|
||||
using AIStudio.Tools.ToolCallingSystem.Harness;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// Speaks the Chat Completions wire format for the tool calling loop.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Tool calls arrive in the assistant message, and results go back as tool messages correlated by
|
||||
/// tool call ID. Every OpenAI-compatible provider uses this shape.
|
||||
/// </remarks>
|
||||
public sealed class ChatCompletionToolCallingAdapter<TRequest>(
|
||||
Func<TextMessage, IDictionary<string, object>, IList<object>?, Task<TRequest>> requestFactory,
|
||||
TextMessage systemPrompt, IDictionary<string, object> apiParameters,
|
||||
IList<object> providerTools,
|
||||
IReadOnlyList<(ToolDefinition Definition, IToolImplementation Implementation)> runnableTools,
|
||||
Func<ChatCompletionAPIRequest, CancellationToken, Task<ChatCompletionResponse?>> executeRequestAsync,
|
||||
string providerInstanceName, ILogger logger)
|
||||
: IToolCallingProviderAdapter where TRequest : ChatCompletionAPIRequest
|
||||
{
|
||||
private readonly List<IMessageBase> internalMessages = [];
|
||||
private ChatCompletionResponseMessage? lastResponseMessage;
|
||||
private List<ChatCompletionToolCall> lastToolCalls = [];
|
||||
|
||||
/// <inheritdoc />
|
||||
public async Task<ToolCallingRound?> ExecuteRoundAsync(string? finalResponseInstruction, bool includeTools, CancellationToken token = default)
|
||||
{
|
||||
var requestSystemPrompt = finalResponseInstruction is null
|
||||
? systemPrompt : systemPrompt with
|
||||
{
|
||||
Content = $"{systemPrompt.Content}{Environment.NewLine}{Environment.NewLine}{finalResponseInstruction}",
|
||||
};
|
||||
|
||||
ChatCompletionAPIRequest requestDtoBase = await requestFactory(requestSystemPrompt, apiParameters, includeTools ? providerTools : null);
|
||||
var requestDto = requestDtoBase with
|
||||
{
|
||||
Messages = [..requestDtoBase.Messages, ..this.internalMessages],
|
||||
Stream = false,
|
||||
|
||||
//
|
||||
// AI Studio runs tool calls one after another, so asking for parallel calls would
|
||||
// only produce work it then has to serialize anyway. Requests without tools omit the
|
||||
// parameter because some providers reject it then.
|
||||
//
|
||||
ParallelToolCalls = requestDtoBase.Tools is null ? null : false,
|
||||
};
|
||||
|
||||
var response = await executeRequestAsync(requestDto, token);
|
||||
if (response is null)
|
||||
return null;
|
||||
|
||||
// The response comes from a provider, so its shape is a promise rather than a guarantee:
|
||||
// a JSON null for the choices field overwrites the initialized property with null.
|
||||
// ReSharper disable once ConditionalAccessQualifierIsNonNullableAccordingToAPIContract
|
||||
var responseChoice = response.Choices?.FirstOrDefault();
|
||||
if (responseChoice?.Message is null)
|
||||
{
|
||||
logger.LogError(
|
||||
"The tool calling response did not contain a usable choice. ProviderInstanceName={ProviderInstanceName}, ChoiceCount={ChoiceCount}",
|
||||
providerInstanceName,
|
||||
response.Choices?.Count ?? 0);
|
||||
|
||||
throw ToolCallingMessages.InvalidToolCallingResponse(providerInstanceName);
|
||||
}
|
||||
|
||||
this.lastResponseMessage = responseChoice.Message;
|
||||
var preparedCalls = this.PrepareToolCalls(responseChoice.Message.ToolCalls ?? []);
|
||||
this.lastToolCalls = preparedCalls.Select(x => x.ToolCall).ToList();
|
||||
|
||||
return new ToolCallingRound(
|
||||
responseChoice.Message.Content ?? string.Empty,
|
||||
preparedCalls
|
||||
.Select(x => new ToolCallingRequestedCall(x.ToolCall.Id!, x.ToolCall.Function!.Name!, x.ToolCall.Function!.Arguments!, x.IsValid))
|
||||
.ToList(),
|
||||
[]);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public void RecordAssistantTurn() => this.internalMessages.Add(new AssistantToolCallMessage
|
||||
{
|
||||
Content = this.lastResponseMessage?.RawContent,
|
||||
ReasoningContent = this.lastResponseMessage?.ReasoningContent,
|
||||
ToolCalls = this.lastToolCalls,
|
||||
});
|
||||
|
||||
/// <inheritdoc />
|
||||
/// <remarks>
|
||||
/// Chat Completions has no error flag on a tool message, so a failure travels in the content
|
||||
/// like any other result.
|
||||
/// </remarks>
|
||||
public void RecordToolResult(string callId, string content, bool isError = false) => this.internalMessages.Add(new ToolResultMessage
|
||||
{
|
||||
Content = content,
|
||||
ToolCallId = callId,
|
||||
});
|
||||
|
||||
/// <summary>
|
||||
/// Normalizes the tool calls of one response.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Models get this wrong in several ways: a missing call ID, a missing function name, or
|
||||
/// arguments that are not valid JSON. None of that may reach a tool, but none of it may be
|
||||
/// dropped either — a call the model never hears about again leaves it waiting. So each call
|
||||
/// is either marked invalid and answered with an error, or corrected where that is safe.
|
||||
/// </remarks>
|
||||
private List<PreparedChatCompletionToolCall> PrepareToolCalls(IEnumerable<ChatCompletionToolCall?> toolCalls)
|
||||
{
|
||||
var preparedToolCalls = new List<PreparedChatCompletionToolCall>();
|
||||
foreach (var returnedToolCall in toolCalls)
|
||||
{
|
||||
//
|
||||
// Unlike the Responses API, Chat Completions does not need the ID to come from the
|
||||
// provider: it only has to match between our request and our answer. So a missing one
|
||||
// can be supplied instead of failing the call.
|
||||
//
|
||||
var toolCallId = string.IsNullOrWhiteSpace(returnedToolCall?.Id)
|
||||
? $"call_{Guid.NewGuid():N}"
|
||||
: returnedToolCall.Id;
|
||||
|
||||
var returnedFunctionName = returnedToolCall?.Function?.Name;
|
||||
var returnedArguments = returnedToolCall?.Function?.Arguments;
|
||||
var isValid = returnedToolCall?.Function is not null &&
|
||||
!string.IsNullOrWhiteSpace(returnedFunctionName) &&
|
||||
ToolExecutor.IsValidArgumentsJson(returnedArguments);
|
||||
|
||||
var normalizedToolCall = new ChatCompletionToolCall
|
||||
{
|
||||
Id = toolCallId,
|
||||
Type = string.IsNullOrWhiteSpace(returnedToolCall?.Type) ? "function" : returnedToolCall.Type,
|
||||
AdditionalMetadata = returnedToolCall?.AdditionalMetadata ?? new Dictionary<string, JsonElement>(),
|
||||
Function = new ChatCompletionToolFunction
|
||||
{
|
||||
Name = string.IsNullOrWhiteSpace(returnedFunctionName) ? "invalid_tool_call" : returnedFunctionName,
|
||||
Arguments = returnedArguments ?? "{}",
|
||||
},
|
||||
};
|
||||
|
||||
if (!isValid)
|
||||
{
|
||||
logger.LogWarning("Received an invalid Chat Completions tool call. ToolCallId={ToolCallId}", toolCallId);
|
||||
preparedToolCalls.Add(new PreparedChatCompletionToolCall(normalizedToolCall, false));
|
||||
continue;
|
||||
}
|
||||
|
||||
var canonicalName = runnableTools
|
||||
.Select(x => x.Definition.Function.Name)
|
||||
.FirstOrDefault(x => x.Equals(returnedFunctionName!.Trim(), StringComparison.Ordinal));
|
||||
|
||||
if (canonicalName is not null && !canonicalName.Equals(returnedFunctionName, StringComparison.Ordinal))
|
||||
{
|
||||
logger.LogWarning("Canonicalized tool call function name '{ReturnedFunctionName}' to '{CanonicalFunctionName}'.", returnedFunctionName, canonicalName);
|
||||
normalizedToolCall = normalizedToolCall with
|
||||
{
|
||||
Function = normalizedToolCall.Function! with
|
||||
{
|
||||
Name = canonicalName,
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
preparedToolCalls.Add(new PreparedChatCompletionToolCall(normalizedToolCall, true));
|
||||
}
|
||||
|
||||
return preparedToolCalls;
|
||||
}
|
||||
|
||||
private readonly record struct PreparedChatCompletionToolCall(ChatCompletionToolCall ToolCall, bool IsValid);
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
public sealed record ChatCompletionToolFunction
|
||||
{
|
||||
public string? Name { get; init; }
|
||||
|
||||
public string? Arguments { get; init; }
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
using System.Text.Json;
|
||||
using System.Text.Json.Nodes;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// Translates a tool's parameter schema into the form OpenAI's strict mode requires.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Strict mode does not accept an optional argument the ordinary JSON Schema way. It insists that
|
||||
/// every property appears in <c>required</c>, and an argument that may be left out has to say so
|
||||
/// by allowing null instead — <c>"type": ["string", "null"]</c>, and <c>null</c> among its enum
|
||||
/// values where it has any.<br/><br/>
|
||||
/// Tool definitions are written the ordinary way, so this converts on the way out. Both forms mean
|
||||
/// the same to a tool: a null argument and an absent one are treated alike.
|
||||
/// </remarks>
|
||||
public static class OpenAIStrictToolSchema
|
||||
{
|
||||
private const string NULL_TYPE = "null";
|
||||
|
||||
/// <summary>
|
||||
/// Converts one parameter schema, leaving it untouched when every argument is required
|
||||
/// anyway.
|
||||
/// </summary>
|
||||
public static JsonElement FromToolParameters(JsonElement parameters)
|
||||
{
|
||||
if (parameters.ValueKind is not JsonValueKind.Object)
|
||||
return parameters;
|
||||
|
||||
if (JsonNode.Parse(parameters.GetRawText()) is not JsonObject schema)
|
||||
return parameters;
|
||||
|
||||
if (schema["properties"] is not JsonObject properties)
|
||||
return parameters;
|
||||
|
||||
var requiredNames = schema["required"] is JsonArray required
|
||||
? required.Select(entry => entry?.GetValue<string>()).Where(entry => entry is not null).ToHashSet(StringComparer.Ordinal)
|
||||
: [];
|
||||
|
||||
var optionalPropertyNames = properties
|
||||
.Select(property => property.Key)
|
||||
.Where(propertyName => !requiredNames.Contains(propertyName))
|
||||
.ToList();
|
||||
|
||||
if (optionalPropertyNames.Count is 0)
|
||||
return parameters;
|
||||
|
||||
foreach (var propertyName in optionalPropertyNames)
|
||||
{
|
||||
if (properties[propertyName] is not JsonObject property)
|
||||
continue;
|
||||
|
||||
AllowNullType(property);
|
||||
AllowNullEnumValue(property);
|
||||
}
|
||||
|
||||
//
|
||||
// Every property is required in strict mode. The order follows the properties, so the
|
||||
// schema stays stable across requests, which prompt caching depends on.
|
||||
//
|
||||
schema["required"] = new JsonArray([..properties.Select(property => JsonValue.Create(property.Key))]);
|
||||
return JsonSerializer.Deserialize<JsonElement>(schema.ToJsonString());
|
||||
}
|
||||
|
||||
private static void AllowNullType(JsonObject property)
|
||||
{
|
||||
switch (property["type"])
|
||||
{
|
||||
case JsonValue singleType when singleType.TryGetValue<string>(out var typeName) && !typeName.Equals(NULL_TYPE, StringComparison.Ordinal):
|
||||
property["type"] = new JsonArray(JsonValue.Create(typeName), JsonValue.Create(NULL_TYPE));
|
||||
break;
|
||||
|
||||
case JsonArray types when types.All(entry => entry?.GetValue<string>() != NULL_TYPE):
|
||||
types.Add(JsonValue.Create(NULL_TYPE));
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
private static void AllowNullEnumValue(JsonObject property)
|
||||
{
|
||||
// Only where the property restricts its values at all: adding null to an absent enum
|
||||
// would turn an unrestricted argument into one that may only be null.
|
||||
if (property["enum"] is not JsonArray enumValues || enumValues.Any(entry => entry is null))
|
||||
return;
|
||||
|
||||
enumValues.Insert(0, null);
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,9 @@ using System.Text.Json;
|
||||
using AIStudio.Chat;
|
||||
using AIStudio.Settings;
|
||||
using AIStudio.Tools.PluginSystem;
|
||||
using AIStudio.Tools.Rust;
|
||||
using AIStudio.Tools.ToolCallingSystem;
|
||||
using AIStudio.Tools.ToolCallingSystem.Harness;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
@@ -16,7 +19,6 @@ namespace AIStudio.Provider.OpenAI;
|
||||
public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Uri("https://api.openai.com/v1/"), ExternalHttpTrustPolicy.SYSTEM_TRUST_ONLY, LOGGER)
|
||||
{
|
||||
private static readonly ILogger<ProviderOpenAI> LOGGER = Program.LOGGER_FACTORY.CreateLogger<ProviderOpenAI>();
|
||||
|
||||
private static string TB(string fallbackEN) => I18N.I.T(fallbackEN, typeof(ProviderOpenAI).Namespace, nameof(ProviderOpenAI));
|
||||
|
||||
#region Implementation of IProvider
|
||||
@@ -87,8 +89,10 @@ public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Ur
|
||||
_ => systemPromptRole,
|
||||
};
|
||||
|
||||
// Read the model capabilities:
|
||||
var modelCapabilities = this.Provider.GetModelCapabilities(chatModel);
|
||||
// Read the model capabilities. Through the settings provider, so that the user's expert
|
||||
// capability overrides apply:
|
||||
var providerSettings = this.CreateSettingsProvider(chatModel);
|
||||
var modelCapabilities = providerSettings.GetModelCapabilities();
|
||||
|
||||
// Check if we are using the Responses API or the Chat Completion API:
|
||||
var usingResponsesAPI = modelCapabilities.Contains(Capability.RESPONSES_API);
|
||||
@@ -98,81 +102,156 @@ public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Ur
|
||||
|
||||
LOGGER.LogInformation("Using the system prompt role '{SystemPromptRole}' and the '{RequestPath}' API for model '{ChatModelId}'.", systemPromptRole, requestPath, chatModel.Id);
|
||||
|
||||
// Prepare the system prompt:
|
||||
var systemPrompt = new TextMessage
|
||||
{
|
||||
Role = systemPromptRole,
|
||||
Content = chatThread.PrepareSystemPrompt(settingsManager),
|
||||
};
|
||||
|
||||
//
|
||||
// Prepare the tools we want to use:
|
||||
//
|
||||
IList<ProviderTool> providerTools = modelCapabilities.Contains(Capability.WEB_SEARCH) switch
|
||||
{
|
||||
true => [ ProviderTools.WEB_SEARCH ],
|
||||
_ => []
|
||||
};
|
||||
var toolRegistry = Program.SERVICE_PROVIDER.GetService<ToolRegistry>();
|
||||
var providerConfidence = this.Provider.GetConfidence(settingsManager).Level;
|
||||
|
||||
//
|
||||
// The provider-native web search is held to the same confidence the local web search tool
|
||||
// asks for: to the user it is the same act, whoever performs the search.
|
||||
//
|
||||
var minimumWebSearchConfidence = toolRegistry?.GetMinimumProviderConfidence(ToolSelectionRules.WEB_SEARCH_TOOL_ID) ?? ConfidenceLevel.NONE;
|
||||
var isWebSearchAllowed = settingsManager.IsToolActive(ToolSelectionRules.WEB_SEARCH_TOOL_ID) &&
|
||||
ToolSelectionRules.IsProviderConfidenceAllowed(providerConfidence, minimumWebSearchConfidence);
|
||||
IList<object> providerTools = modelCapabilities.Contains(Capability.WEB_SEARCH) && isWebSearchAllowed
|
||||
? [ ProviderTools.WEB_SEARCH ]
|
||||
: [];
|
||||
|
||||
|
||||
// Parse the API parameters:
|
||||
var apiParameters = this.ParseAdditionalApiParameters("input", "store", "tools");
|
||||
var additionalApiParameters = this.ParseAdditionalApiParameters("input", "store", "tools");
|
||||
|
||||
if (!usingResponsesAPI)
|
||||
{
|
||||
await foreach (var content in this.StreamOpenAICompatibleChatCompletion<ChatCompletionAPIRequest, ChatCompletionDeltaStreamLine, ChatCompletionAnnotationStreamLine>(
|
||||
"OpenAI",
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
var messages = await chatThread.Blocks.BuildMessagesAsync(
|
||||
this.Provider,
|
||||
chatModel,
|
||||
role => role switch
|
||||
{
|
||||
ChatRole.USER => "user",
|
||||
ChatRole.AI => "assistant",
|
||||
ChatRole.AGENT => "assistant",
|
||||
ChatRole.SYSTEM => systemPromptRole,
|
||||
_ => "user",
|
||||
},
|
||||
text => new SubContentText
|
||||
{
|
||||
Text = text,
|
||||
},
|
||||
async attachment => new SubContentImageUrlNested
|
||||
{
|
||||
ImageUrl = new SubContentImageUrlData
|
||||
{
|
||||
Url = await attachment.TryAsBase64(token: token) is (true, var base64Content)
|
||||
? $"data:{attachment.DetermineMimeType()};base64,{base64Content}"
|
||||
: string.Empty,
|
||||
},
|
||||
});
|
||||
|
||||
return new ChatCompletionAPIRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
Messages = [systemPrompt, ..messages],
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters,
|
||||
};
|
||||
},
|
||||
systemPromptRole: systemPromptRole,
|
||||
requestPath: "chat/completions",
|
||||
token: token))
|
||||
yield return content;
|
||||
|
||||
yield break;
|
||||
}
|
||||
|
||||
var toolExecutor = Program.SERVICE_PROVIDER.GetService<ToolExecutor>();
|
||||
var currentAssistantContent = chatThread.Blocks.LastOrDefault(x => x.Role is ChatRole.AI)?.Content as ContentText;
|
||||
currentAssistantContent?.ToolInvocations.Clear();
|
||||
|
||||
IReadOnlyList<(ToolDefinition Definition, IToolImplementation Implementation)> runnableTools = toolRegistry is null
|
||||
? []
|
||||
: await toolRegistry.GetRunnableToolsAsync(
|
||||
providerSettings,
|
||||
chatThread.RuntimeComponent,
|
||||
chatThread.RuntimeSelectedToolIds,
|
||||
providerConfidence,
|
||||
chatThread.MayRunTools(settingsManager));
|
||||
|
||||
var toolAwareDefinitions = toolExecutor is null
|
||||
? Enumerable.Empty<ToolDefinition>()
|
||||
: runnableTools.Select(x => x.Definition);
|
||||
var systemPrompt = new TextMessage
|
||||
{
|
||||
Role = systemPromptRole,
|
||||
Content = chatThread.PrepareSystemPrompt(settingsManager, toolAwareDefinitions),
|
||||
};
|
||||
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesAsync(
|
||||
this.Provider, chatModel,
|
||||
|
||||
// OpenAI-specific role mapping:
|
||||
role => role switch
|
||||
{
|
||||
ChatRole.USER => "user",
|
||||
ChatRole.AI => "assistant",
|
||||
ChatRole.AGENT => "assistant",
|
||||
ChatRole.SYSTEM => systemPromptRole,
|
||||
|
||||
_ => "user",
|
||||
},
|
||||
|
||||
// OpenAI's text sub-content depends on the model, whether we are using
|
||||
// the Responses API or the Chat Completion API:
|
||||
text => usingResponsesAPI switch
|
||||
text => new SubContentInputText
|
||||
{
|
||||
// Responses API uses INPUT_TEXT:
|
||||
true => new SubContentInputText
|
||||
{
|
||||
Text = text,
|
||||
},
|
||||
|
||||
// Chat Completion API uses TEXT:
|
||||
false => new SubContentText
|
||||
{
|
||||
Text = text,
|
||||
},
|
||||
Text = text,
|
||||
},
|
||||
|
||||
// OpenAI's image sub-content depends on the model as well,
|
||||
// whether we are using the Responses API or the Chat Completion API:
|
||||
async attachment => usingResponsesAPI switch
|
||||
async attachment => new SubContentInputImage
|
||||
{
|
||||
// Responses API uses INPUT_IMAGE:
|
||||
true => new SubContentInputImage
|
||||
{
|
||||
ImageUrl = await attachment.TryAsBase64(token: token) is (true, var base64Content)
|
||||
? $"data:{attachment.DetermineMimeType()};base64,{base64Content}"
|
||||
: string.Empty,
|
||||
},
|
||||
|
||||
// Chat Completion API uses IMAGE_URL:
|
||||
false => new SubContentImageUrlNested
|
||||
{
|
||||
ImageUrl = new SubContentImageUrlData
|
||||
{
|
||||
Url = await attachment.TryAsBase64(token: token) is (true, var base64Content)
|
||||
? $"data:{attachment.DetermineMimeType()};base64,{base64Content}"
|
||||
: string.Empty,
|
||||
},
|
||||
}
|
||||
ImageUrl = await attachment.TryAsBase64(token: token) is (true, var base64Content)
|
||||
? $"data:{attachment.DetermineMimeType()};base64,{base64Content}"
|
||||
: string.Empty,
|
||||
});
|
||||
|
||||
var baseInput = new List<object> { systemPrompt };
|
||||
baseInput.AddRange(messages);
|
||||
|
||||
if (usingResponsesAPI && toolExecutor is not null && runnableTools.Count > 0)
|
||||
{
|
||||
var adapter = new ResponsesToolCallingAdapter(
|
||||
chatModel,
|
||||
baseInput,
|
||||
additionalApiParameters,
|
||||
providerTools,
|
||||
runnableTools,
|
||||
(requestDto, requestToken) => this.ExecuteResponsesRequest(requestDto, requestedSecret, requestToken));
|
||||
|
||||
var loop = Program.SERVICE_PROVIDER.GetRequiredService<IToolCallingLoop>();
|
||||
var loopContext = new ToolCallingLoopContext
|
||||
{
|
||||
ChatThread = chatThread,
|
||||
RunnableTools = runnableTools,
|
||||
ToolExecutor = toolExecutor,
|
||||
Provider = this,
|
||||
CurrentAssistantContent = currentAssistantContent,
|
||||
ProviderInstanceName = this.InstanceName,
|
||||
ProviderType = this.Provider,
|
||||
ModelId = chatModel.Id,
|
||||
};
|
||||
|
||||
await foreach (var content in loop.RunAsync(adapter, loopContext, token))
|
||||
yield return content;
|
||||
|
||||
yield break;
|
||||
}
|
||||
|
||||
if (runnableTools.Count > 0)
|
||||
providerTools = [];
|
||||
|
||||
//
|
||||
// Create the request: either for the Responses API or the Chat Completion API
|
||||
@@ -189,16 +268,16 @@ public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Ur
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
AdditionalApiParameters = apiParameters
|
||||
AdditionalApiParameters = additionalApiParameters
|
||||
}, JSON_SERIALIZER_OPTIONS),
|
||||
|
||||
|
||||
// Responses API request:
|
||||
true => JsonSerializer.Serialize(new ResponsesAPIRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
|
||||
// All messages go into the input field:
|
||||
Input = [systemPrompt, ..messages],
|
||||
Input = baseInput,
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
@@ -207,10 +286,10 @@ public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Ur
|
||||
Store = false,
|
||||
|
||||
// Tools we want to use:
|
||||
ProviderTools = providerTools,
|
||||
Tools = providerTools,
|
||||
|
||||
// Additional API parameters:
|
||||
AdditionalApiParameters = apiParameters
|
||||
AdditionalApiParameters = additionalApiParameters
|
||||
|
||||
}, JSON_SERIALIZER_OPTIONS),
|
||||
};
|
||||
@@ -237,6 +316,24 @@ public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Ur
|
||||
yield return content;
|
||||
}
|
||||
|
||||
private async Task<ResponsesResponse?> ExecuteResponsesRequest(ResponsesAPIRequest requestDto, RequestedSecret requestedSecret, CancellationToken token)
|
||||
{
|
||||
using var request = new HttpRequestMessage(HttpMethod.Post, "responses");
|
||||
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", await requestedSecret.Secret.Decrypt(Program.ENCRYPTION));
|
||||
request.Content = new StringContent(JsonSerializer.Serialize(requestDto, JSON_SERIALIZER_OPTIONS), Encoding.UTF8, "application/json");
|
||||
|
||||
using var response = await this.HttpClient.SendAsync(request, token);
|
||||
if (!response.IsSuccessStatusCode)
|
||||
{
|
||||
var responseBody = await response.Content.ReadAsStringAsync(token);
|
||||
LOGGER.LogError("Tool calling Responses API request failed with status code {ResponseStatusCode} and body: '{ResponseBody}'.", response.StatusCode, responseBody);
|
||||
await ToolCallingMessages.SendToolCallingRequestFailedAsync((int)response.StatusCode);
|
||||
return null;
|
||||
}
|
||||
|
||||
return await response.Content.ReadFromJsonAsync<ResponsesResponse>(JSON_SERIALIZER_OPTIONS, token);
|
||||
}
|
||||
|
||||
#pragma warning disable CS1998 // Async method lacks 'await' operators and will run synchronously
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -261,59 +358,46 @@ public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Ur
|
||||
return await this.PerformStandardTextEmbeddingRequest(requestedSecret, embeddingModel, token: token, texts: texts);
|
||||
}
|
||||
|
||||
//
|
||||
// OpenAI offers every kind of model through one models endpoint, so we have to sort them apart
|
||||
// ourselves. We used to do that with lists of name prefixes kept here. The shared model kind
|
||||
// detection knows those families as well, and it knows them for every provider, so we ask it
|
||||
// instead of maintaining a second set of rules which only ever lagged behind.
|
||||
//
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, ["chatgpt-", "gpt-", "o1-", "o3-", "o4-"], token, apiKeyProvisional);
|
||||
return result with
|
||||
{
|
||||
Models =
|
||||
[
|
||||
..result.Models.Where(model => !model.Id.Contains("image", StringComparison.OrdinalIgnoreCase) &&
|
||||
!model.Id.Contains("realtime", StringComparison.OrdinalIgnoreCase) &&
|
||||
!model.Id.Contains("audio", StringComparison.OrdinalIgnoreCase) &&
|
||||
!model.Id.Contains("tts", StringComparison.OrdinalIgnoreCase) &&
|
||||
!model.Id.Contains("transcribe", StringComparison.OrdinalIgnoreCase))
|
||||
]
|
||||
};
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, static model => model.IsChatModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetImageModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.IMAGE_PROVIDER, ["dall-e-", "gpt-image"], token, apiKeyProvisional);
|
||||
return this.LoadModels(SecretStoreType.IMAGE_PROVIDER, static model => model.IsImageModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, ["text-embedding-"], token, apiKeyProvisional);
|
||||
return this.LoadModels(SecretStoreType.EMBEDDING_PROVIDER, static model => model.IsEmbeddingModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
|
||||
/// <inheritdoc />
|
||||
public override async Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
public override Task<ModelLoadResult> GetTranscriptionModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.TRANSCRIPTION_PROVIDER, ["whisper-", "gpt-"], token, apiKeyProvisional);
|
||||
return result with
|
||||
{
|
||||
Models =
|
||||
[
|
||||
..result.Models.Where(model => model.Id.StartsWith("whisper-", StringComparison.InvariantCultureIgnoreCase) ||
|
||||
model.Id.Contains("-transcribe", StringComparison.InvariantCultureIgnoreCase))
|
||||
]
|
||||
};
|
||||
return this.LoadModels(SecretStoreType.TRANSCRIPTION_PROVIDER, static model => model.IsTranscriptionModel(), apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string[] prefixes, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, Func<Model, bool> isWantedKind, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data.Where(model => prefixes.Any(prefix => model.Id.StartsWith(prefix, StringComparison.InvariantCulture))),
|
||||
token,
|
||||
apiKeyProvisional);
|
||||
modelResponse => modelResponse.Data.Where(isWantedKind),
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
|
||||
private static bool HasInsufficientQuotaError(string responseBody)
|
||||
@@ -370,4 +454,4 @@ public sealed class ProviderOpenAI() : BaseProvider(LLMProviders.OPEN_AI, new Ur
|
||||
propertyElement.ValueKind is JsonValueKind.String &&
|
||||
string.Equals(propertyElement.GetString(), expectedValue, StringComparison.OrdinalIgnoreCase);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -6,16 +6,16 @@ namespace AIStudio.Provider.OpenAI;
|
||||
/// The request body for the Responses API.
|
||||
/// </summary>
|
||||
/// <param name="Model">Which model to use.</param>
|
||||
/// <param name="Input">The chat messages.</param>
|
||||
/// <param name="Input">The chat messages and Responses API input items.</param>
|
||||
/// <param name="Stream">Whether to stream the response.</param>
|
||||
/// <param name="Store">Whether to store the response on the server (usually OpenAI's infrastructure).</param>
|
||||
/// <param name="ProviderTools">The provider-side tools to use for the request.</param>
|
||||
/// <param name="Tools">The provider-side tools and local function tools to use for the request.</param>
|
||||
public record ResponsesAPIRequest(
|
||||
string Model,
|
||||
IList<IMessageBase> Input,
|
||||
IList<object> Input,
|
||||
bool Stream,
|
||||
bool Store,
|
||||
[property: JsonPropertyName("tools")] IList<ProviderTool> ProviderTools)
|
||||
IList<object> Tools)
|
||||
{
|
||||
public ResponsesAPIRequest() : this(string.Empty, [], true, false, [])
|
||||
{
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// A function call item returned by the OpenAI Responses API.
|
||||
/// </summary>
|
||||
public sealed record ResponsesFunctionCallItem
|
||||
{
|
||||
public string? Type { get; init; }
|
||||
|
||||
public string? CallId { get; init; }
|
||||
|
||||
public string? Name { get; init; }
|
||||
|
||||
public string? Arguments { get; init; }
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// A local function result item sent back to the OpenAI Responses API.
|
||||
/// </summary>
|
||||
public sealed record ResponsesFunctionCallOutputItem
|
||||
{
|
||||
public string Type { get; init; } = "function_call_output";
|
||||
|
||||
public string CallId { get; init; } = string.Empty;
|
||||
|
||||
public string Output { get; init; } = string.Empty;
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
using System.Text.Json;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// The flat function tool definition shape expected by the OpenAI Responses API.
|
||||
/// </summary>
|
||||
public sealed record ResponsesFunctionTool
|
||||
{
|
||||
public string Type { get; init; } = "function";
|
||||
|
||||
public string Name { get; init; } = string.Empty;
|
||||
|
||||
public string Description { get; init; } = string.Empty;
|
||||
|
||||
public JsonElement Parameters { get; init; }
|
||||
|
||||
public bool Strict { get; init; }
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
using System.Text.Json;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// Non-streaming OpenAI Responses API result used during local tool execution.
|
||||
/// </summary>
|
||||
public sealed record ResponsesResponse
|
||||
{
|
||||
public string Id { get; init; } = string.Empty;
|
||||
|
||||
public string Model { get; init; } = string.Empty;
|
||||
|
||||
public string? OutputText { get; init; }
|
||||
|
||||
public IList<JsonElement> Output { get; init; } = [];
|
||||
|
||||
public IReadOnlyList<ResponsesFunctionCallItem> GetFunctionCalls() => this.Output
|
||||
.Where(x => ReadString(x, "type").Equals("function_call", StringComparison.Ordinal))
|
||||
.Select(x => new ResponsesFunctionCallItem
|
||||
{
|
||||
Type = ReadString(x, "type"),
|
||||
CallId = ReadString(x, "call_id"),
|
||||
Name = ReadString(x, "name"),
|
||||
Arguments = ReadString(x, "arguments"),
|
||||
})
|
||||
.ToList();
|
||||
|
||||
public string GetTextOutput()
|
||||
{
|
||||
if (!string.IsNullOrWhiteSpace(this.OutputText))
|
||||
return this.OutputText;
|
||||
|
||||
return string.Concat(this.Output
|
||||
.Where(x => ReadString(x, "type").Equals("message", StringComparison.Ordinal))
|
||||
.SelectMany(ReadContentItems)
|
||||
.Select(x => ReadString(x, "type") switch
|
||||
{
|
||||
"output_text" => ReadString(x, "text"),
|
||||
"refusal" => ReadString(x, "refusal"),
|
||||
_ => string.Empty,
|
||||
}));
|
||||
}
|
||||
|
||||
public IReadOnlyList<Source> GetSources() => this.Output
|
||||
.Where(x => ReadString(x, "type").Equals("message", StringComparison.Ordinal))
|
||||
.SelectMany(ReadContentItems)
|
||||
.SelectMany(x => ReadArrayItems(x, "annotations"))
|
||||
.Where(x => ReadString(x, "type").Equals("url_citation", StringComparison.Ordinal))
|
||||
.Select(x => new Source(ReadString(x, "title"), ReadString(x, "url"), SourceOrigin.LLM))
|
||||
.Where(x => !string.IsNullOrWhiteSpace(x.Title) && !string.IsNullOrWhiteSpace(x.URL))
|
||||
.ToList();
|
||||
|
||||
private static IEnumerable<JsonElement> ReadContentItems(JsonElement outputItem)
|
||||
=> ReadArrayItems(outputItem, "content");
|
||||
|
||||
private static IEnumerable<JsonElement> ReadArrayItems(JsonElement item, string propertyName)
|
||||
{
|
||||
if (item.ValueKind is not JsonValueKind.Object ||
|
||||
!item.TryGetProperty(propertyName, out var array) ||
|
||||
array.ValueKind is not JsonValueKind.Array)
|
||||
yield break;
|
||||
|
||||
foreach (var arrayItem in array.EnumerateArray())
|
||||
yield return arrayItem;
|
||||
}
|
||||
|
||||
private static string ReadString(JsonElement item, string propertyName)
|
||||
{
|
||||
if (item.ValueKind is not JsonValueKind.Object ||
|
||||
!item.TryGetProperty(propertyName, out var property) ||
|
||||
property.ValueKind is not JsonValueKind.String)
|
||||
return string.Empty;
|
||||
|
||||
return property.GetString() ?? string.Empty;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
using AIStudio.Tools.ToolCallingSystem;
|
||||
using AIStudio.Tools.ToolCallingSystem.Harness;
|
||||
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
/// <summary>
|
||||
/// Speaks the OpenAI Responses wire format for the tool calling loop.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Function calls arrive as output items and results go back as function call output items,
|
||||
/// correlated by call ID. Unlike Chat Completions, the whole output of a round has to be sent
|
||||
/// back for the next one, reasoning items included, or the API refuses to continue.
|
||||
/// </remarks>
|
||||
public sealed class ResponsesToolCallingAdapter(Model chatModel, IList<object> baseInput, IDictionary<string, object> apiParameters, IList<object> providerTools,
|
||||
IReadOnlyList<(ToolDefinition Definition, IToolImplementation Implementation)> runnableTools,
|
||||
Func<ResponsesAPIRequest, CancellationToken, Task<ResponsesResponse?>> executeRequestAsync) : IToolCallingProviderAdapter
|
||||
{
|
||||
private readonly List<object> internalItems = [];
|
||||
private ResponsesResponse? lastResponse;
|
||||
|
||||
/// <summary>
|
||||
/// The tools offered to the model: the provider-native ones plus our local functions.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// A provider-native tool whose type collides with one of our function names is dropped
|
||||
/// because the model could not tell the two apart.
|
||||
/// </remarks>
|
||||
private readonly IList<object> effectiveProviderTools = BuildEffectiveProviderTools(providerTools, runnableTools);
|
||||
|
||||
/// <inheritdoc />
|
||||
public async Task<ToolCallingRound?> ExecuteRoundAsync(string? finalResponseInstruction, bool includeTools, CancellationToken token = default)
|
||||
{
|
||||
var requestInput = new List<object>(baseInput);
|
||||
if (finalResponseInstruction is not null && requestInput.FirstOrDefault() is TextMessage systemPrompt)
|
||||
{
|
||||
requestInput[0] = systemPrompt with
|
||||
{
|
||||
Content = $"{systemPrompt.Content}{Environment.NewLine}{Environment.NewLine}{finalResponseInstruction}",
|
||||
};
|
||||
}
|
||||
|
||||
requestInput.AddRange(this.internalItems);
|
||||
|
||||
var response = await executeRequestAsync(new ResponsesAPIRequest
|
||||
{
|
||||
Model = chatModel.Id,
|
||||
Input = requestInput,
|
||||
Stream = false,
|
||||
Store = false,
|
||||
Tools = includeTools ? this.effectiveProviderTools : [],
|
||||
AdditionalApiParameters = apiParameters,
|
||||
}, token);
|
||||
|
||||
if (response is null)
|
||||
return null;
|
||||
|
||||
this.lastResponse = response;
|
||||
return new ToolCallingRound(
|
||||
response.GetTextOutput(),
|
||||
response.GetFunctionCalls()
|
||||
.Select(call => new ToolCallingRequestedCall(
|
||||
call.CallId ?? string.Empty,
|
||||
call.Name ?? string.Empty,
|
||||
call.Arguments ?? string.Empty,
|
||||
!string.IsNullOrWhiteSpace(call.Name) && ToolExecutor.IsValidArgumentsJson(call.Arguments)))
|
||||
.ToList(),
|
||||
|
||||
response.GetSources());
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
public void RecordAssistantTurn()
|
||||
{
|
||||
if (this.lastResponse is null)
|
||||
return;
|
||||
|
||||
// Every output item, not just the function calls: the API rejects a continuation whose
|
||||
// reasoning items are missing.
|
||||
foreach (var outputItem in this.lastResponse.Output)
|
||||
this.internalItems.Add(outputItem);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
/// <remarks>
|
||||
/// The Responses API has no error flag on a function call output, so a failure travels in the
|
||||
/// output like any other result.
|
||||
/// </remarks>
|
||||
public void RecordToolResult(string callId, string content, bool isError = false) => this.internalItems.Add(new ResponsesFunctionCallOutputItem
|
||||
{
|
||||
CallId = callId,
|
||||
Output = content,
|
||||
});
|
||||
|
||||
private static IList<object> BuildEffectiveProviderTools(IList<object> providerTools, IReadOnlyList<(ToolDefinition Definition, IToolImplementation Implementation)> runnableTools)
|
||||
{
|
||||
var localFunctionNames = runnableTools
|
||||
.Select(x => x.Definition.Function.Name)
|
||||
.ToHashSet(StringComparer.Ordinal);
|
||||
|
||||
return providerTools
|
||||
.Where(x => x is not ProviderTool providerTool || !localFunctionNames.Contains(providerTool.Type))
|
||||
.Concat(runnableTools.Select(x => (object)ProviderToolAdapters.ToResponsesTool(x.Definition)))
|
||||
.ToList();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
namespace AIStudio.Provider.OpenAI;
|
||||
|
||||
public sealed record ToolResultMessage : IMessage<string>
|
||||
{
|
||||
public string Role { get; init; } = "tool";
|
||||
|
||||
public string Content { get; init; } = string.Empty;
|
||||
|
||||
public string ToolCallId { get; init; } = string.Empty;
|
||||
}
|
||||
@@ -33,7 +33,7 @@ public sealed class ProviderOpenRouter() : BaseProvider(LLMProviders.OPEN_ROUTER
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -49,6 +49,7 @@ public sealed class ProviderOpenRouter() : BaseProvider(LLMProviders.OPEN_ROUTER
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -86,7 +87,7 @@ public sealed class ProviderOpenRouter() : BaseProvider(LLMProviders.OPEN_ROUTER
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, token, apiKeyProvisional);
|
||||
return this.LoadModels(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -98,7 +99,7 @@ public sealed class ProviderOpenRouter() : BaseProvider(LLMProviders.OPEN_ROUTER
|
||||
/// <inheritdoc />
|
||||
public override Task<ModelLoadResult> GetEmbeddingModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
return this.LoadEmbeddingModels(token, apiKeyProvisional);
|
||||
return this.LoadEmbeddingModels(apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
/// <inheritdoc />
|
||||
@@ -109,45 +110,37 @@ public sealed class ProviderOpenRouter() : BaseProvider(LLMProviders.OPEN_ROUTER
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<OpenRouterModelsResponse>(
|
||||
storeType,
|
||||
"models",
|
||||
modelResponse => modelResponse.Data
|
||||
.Where(n =>
|
||||
!n.Id.Contains("whisper", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("dall-e", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("tts", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("embedding", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("moderation", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("stable-diffusion", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("flux", StringComparison.OrdinalIgnoreCase) &&
|
||||
!n.Id.Contains("midjourney", StringComparison.OrdinalIgnoreCase))
|
||||
.Select(n => new Model(n.Id, n.Name)),
|
||||
token,
|
||||
.Select(n => new Model(n.Id, n.Name))
|
||||
.Where(model => model.IsChatModel()),
|
||||
apiKeyProvisional,
|
||||
requestConfigurator: (request, secretKey) =>
|
||||
{
|
||||
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", secretKey);
|
||||
request.Headers.Add("HTTP-Referer", PROJECT_WEBSITE);
|
||||
request.Headers.Add("X-Title", PROJECT_NAME);
|
||||
});
|
||||
},
|
||||
token: token);
|
||||
}
|
||||
|
||||
private Task<ModelLoadResult> LoadEmbeddingModels(CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadEmbeddingModels(string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<OpenRouterModelsResponse>(
|
||||
SecretStoreType.EMBEDDING_PROVIDER,
|
||||
"embeddings/models",
|
||||
modelResponse => modelResponse.Data.Select(n => new Model(n.Id, n.Name)),
|
||||
token,
|
||||
apiKeyProvisional,
|
||||
requestConfigurator: (request, secretKey) =>
|
||||
{
|
||||
request.Headers.Authorization = new AuthenticationHeaderValue("Bearer", secretKey);
|
||||
request.Headers.Add("HTTP-Referer", PROJECT_WEBSITE);
|
||||
request.Headers.Add("X-Title", PROJECT_NAME);
|
||||
});
|
||||
},
|
||||
token: token);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -38,7 +38,7 @@ public sealed class ProviderPerplexity() : BaseProvider(LLMProviders.PERPLEXITY,
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -52,6 +52,7 @@ public sealed class ProviderPerplexity() : BaseProvider(LLMProviders.PERPLEXITY,
|
||||
// - Then none-empty user and AI messages
|
||||
Messages = [systemPrompt, ..messages],
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -106,4 +107,4 @@ public sealed class ProviderPerplexity() : BaseProvider(LLMProviders.PERPLEXITY,
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels() => Task.FromResult(ModelLoadResult.FromModels(KNOWN_MODELS));
|
||||
}
|
||||
}
|
||||
@@ -5,4 +5,13 @@ public enum ProviderRequestFailureReason
|
||||
NONE,
|
||||
INSUFFICIENT_QUOTA,
|
||||
TOO_MANY_REQUESTS,
|
||||
|
||||
/// <summary>
|
||||
/// The provider does not serve the requested model.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// This applies to gateways which route to other providers: the model exists, but the one
|
||||
/// meant to answer for it does not offer it.
|
||||
/// </remarks>
|
||||
MODEL_NOT_SUPPORTED_BY_PROVIDER,
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
using AIStudio.Provider.Anthropic;
|
||||
using AIStudio.Provider.OpenAI;
|
||||
using AIStudio.Tools.ToolCallingSystem;
|
||||
|
||||
namespace AIStudio.Provider;
|
||||
|
||||
/// <summary>
|
||||
/// Converts a tool definition into the wire shape one provider API expects.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// The definitions state a tool once, in plain JSON Schema. What differs per API is not only the
|
||||
/// field names but how an optional argument is expressed, which is why the OpenAI shapes convert
|
||||
/// the schema while Anthropic takes it as written.
|
||||
/// </remarks>
|
||||
public static class ProviderToolAdapters
|
||||
{
|
||||
/// <summary>
|
||||
/// Builds the nested function tool shape used by Chat Completions compatible APIs.
|
||||
/// </summary>
|
||||
public static object ToChatCompletionTool(ToolDefinition definition) => new
|
||||
{
|
||||
type = "function",
|
||||
function = new
|
||||
{
|
||||
name = definition.Function.Name,
|
||||
description = definition.Function.DescriptionForLLM,
|
||||
parameters = ToOpenAIParameters(definition),
|
||||
strict = definition.Function.Strict,
|
||||
}
|
||||
};
|
||||
|
||||
/// <summary>
|
||||
/// Builds the flat function tool shape used by the OpenAI Responses API.
|
||||
/// </summary>
|
||||
public static ResponsesFunctionTool ToResponsesTool(ToolDefinition definition) => new()
|
||||
{
|
||||
Name = definition.Function.Name,
|
||||
Description = definition.Function.DescriptionForLLM,
|
||||
Parameters = ToOpenAIParameters(definition),
|
||||
Strict = definition.Function.Strict,
|
||||
};
|
||||
|
||||
/// <summary>
|
||||
/// Builds the tool shape used by the Anthropic messages API.
|
||||
/// </summary>
|
||||
/// <remarks>
|
||||
/// Different field names — Anthropic calls the parameters an input schema and takes the
|
||||
/// description without nesting it under a function object — but the schema itself needs no
|
||||
/// conversion: Anthropic reads optionality the same way the definitions write it.
|
||||
/// </remarks>
|
||||
public static AnthropicTool ToAnthropicTool(ToolDefinition definition) => new()
|
||||
{
|
||||
Name = definition.Function.Name,
|
||||
Description = definition.Function.DescriptionForLLM,
|
||||
InputSchema = definition.Function.Parameters,
|
||||
Strict = definition.Function.Strict,
|
||||
};
|
||||
|
||||
/// <summary>
|
||||
/// The parameter schema for the OpenAI APIs, converted only when strict mode asks for it.
|
||||
/// </summary>
|
||||
private static System.Text.Json.JsonElement ToOpenAIParameters(ToolDefinition definition) => definition.Function.Strict
|
||||
? OpenAIStrictToolSchema.FromToolParameters(definition.Function.Parameters)
|
||||
: definition.Function.Parameters;
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
namespace AIStudio.Provider.SelfHosted;
|
||||
|
||||
public readonly record struct Model(string Id, string? Object, string? OwnedBy, ModelArchitecture? Architecture);
|
||||
@@ -0,0 +1,3 @@
|
||||
namespace AIStudio.Provider.SelfHosted;
|
||||
|
||||
public readonly record struct ModelArchitecture(string[]? InputModalities, string[]? OutputModalities);
|
||||
@@ -1,7 +1,3 @@
|
||||
namespace AIStudio.Provider.SelfHosted;
|
||||
|
||||
public readonly record struct ModelsResponse(string? Object, Model[]? Data);
|
||||
|
||||
public readonly record struct Model(string Id, string? Object, string? OwnedBy, ModelArchitecture? Architecture);
|
||||
|
||||
public readonly record struct ModelArchitecture(string[]? InputModalities, string[]? OutputModalities);
|
||||
public readonly record struct ModelsResponse(string? Object, Model[]? Data);
|
||||
@@ -35,7 +35,7 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
effectiveChatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages. The image format depends on the host:
|
||||
// - Ollama uses the direct image URL format: { "type": "image_url", "image_url": "data:..." }
|
||||
@@ -57,6 +57,7 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -95,12 +96,12 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
switch (host)
|
||||
{
|
||||
case Host.LLAMA_CPP:
|
||||
return await this.LoadLlamaCppTextModels(["embed"], [], token, apiKeyProvisional);
|
||||
return await this.LoadLlamaCppTextModels(["embed"], [], apiKeyProvisional, token);
|
||||
|
||||
case Host.LM_STUDIO:
|
||||
case Host.OLLAMA:
|
||||
case Host.VLLM:
|
||||
return await this.LoadModels( SecretStoreType.LLM_PROVIDER, ["embed"], [], token, apiKeyProvisional);
|
||||
return await this.LoadModels( SecretStoreType.LLM_PROVIDER, ["embed"], [], apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
return ModelLoadResult.FromModels([]);
|
||||
@@ -127,7 +128,7 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
case Host.LM_STUDIO:
|
||||
case Host.OLLAMA:
|
||||
case Host.VLLM:
|
||||
return await this.LoadModels( SecretStoreType.EMBEDDING_PROVIDER, [], ["embed"], token, apiKeyProvisional);
|
||||
return await this.LoadModels( SecretStoreType.EMBEDDING_PROVIDER, [], ["embed"], apiKeyProvisional, token);
|
||||
}
|
||||
|
||||
return ModelLoadResult.FromModels([]);
|
||||
@@ -154,7 +155,7 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
|
||||
case Host.OLLAMA:
|
||||
case Host.VLLM:
|
||||
return await this.LoadModels(SecretStoreType.TRANSCRIPTION_PROVIDER, [], [], token, apiKeyProvisional);
|
||||
return await this.LoadModels(SecretStoreType.TRANSCRIPTION_PROVIDER, [], [], apiKeyProvisional, token);
|
||||
|
||||
default:
|
||||
return ModelLoadResult.FromModels([]);
|
||||
@@ -169,7 +170,7 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
|
||||
#endregion
|
||||
|
||||
private async Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string[] ignorePhrases, string[] filterPhrases, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private async Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string[] ignorePhrases, string[] filterPhrases, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
var secretKey = await this.GetModelLoadingSecretKey(storeType, apiKeyProvisional, true);
|
||||
|
||||
@@ -202,13 +203,12 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
return FailedModelLoadResult(ModelLoadFailureReason.PROVIDER_UNAVAILABLE, e.Message);
|
||||
}
|
||||
}
|
||||
|
||||
private async Task<Provider.Model> ResolveChatModelForRequest(Provider.Model chatModel, CancellationToken token)
|
||||
{
|
||||
if (host is not Host.LLAMA_CPP || !chatModel.IsSystemModel)
|
||||
return chatModel;
|
||||
|
||||
var modelLoadResult = await this.LoadLlamaCppTextModels(["embed"], [], token);
|
||||
var modelLoadResult = await this.LoadLlamaCppTextModels(["embed"], [], null, token);
|
||||
if (!modelLoadResult.Success)
|
||||
return chatModel;
|
||||
|
||||
@@ -245,7 +245,7 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
return chatModel;
|
||||
}
|
||||
|
||||
private async Task<ModelLoadResult> LoadLlamaCppTextModels(string[] ignorePhrases, string[] filterPhrases, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private async Task<ModelLoadResult> LoadLlamaCppTextModels(string[] ignorePhrases, string[] filterPhrases, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
var secretKey = await this.GetModelLoadingSecretKey(SecretStoreType.LLM_PROVIDER, apiKeyProvisional, true);
|
||||
|
||||
@@ -325,4 +325,4 @@ public sealed class ProviderSelfHosted(Host host, string hostname) : BaseProvide
|
||||
{
|
||||
return ModelLoadResult.FromModels([ AIStudio.Provider.Model.SYSTEM_MODEL ]);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -29,7 +29,7 @@ public sealed class ProviderX() : BaseProvider(LLMProviders.X, new Uri("https://
|
||||
chatModel,
|
||||
chatThread,
|
||||
settingsManager,
|
||||
async (systemPrompt, apiParameters) =>
|
||||
async (systemPrompt, apiParameters, tools) =>
|
||||
{
|
||||
// Build the list of messages:
|
||||
var messages = await chatThread.Blocks.BuildMessagesUsingNestedImageUrlAsync(this.Provider, chatModel);
|
||||
@@ -45,6 +45,7 @@ public sealed class ProviderX() : BaseProvider(LLMProviders.X, new Uri("https://
|
||||
|
||||
// Right now, we only support streaming completions:
|
||||
Stream = true,
|
||||
Tools = tools,
|
||||
AdditionalApiParameters = apiParameters
|
||||
};
|
||||
},
|
||||
@@ -75,7 +76,7 @@ public sealed class ProviderX() : BaseProvider(LLMProviders.X, new Uri("https://
|
||||
/// <inheritdoc />
|
||||
public override async Task<ModelLoadResult> GetTextModels(string? apiKeyProvisional = null, CancellationToken token = default)
|
||||
{
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, ["grok-"], token, apiKeyProvisional);
|
||||
var result = await this.LoadModels(SecretStoreType.LLM_PROVIDER, ["grok-"], apiKeyProvisional, token);
|
||||
return result with
|
||||
{
|
||||
Models = [..result.Models.Where(n => !n.Id.Contains("-image", StringComparison.OrdinalIgnoreCase))]
|
||||
@@ -102,7 +103,7 @@ public sealed class ProviderX() : BaseProvider(LLMProviders.X, new Uri("https://
|
||||
|
||||
#endregion
|
||||
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string[] prefixes, CancellationToken token, string? apiKeyProvisional = null)
|
||||
private Task<ModelLoadResult> LoadModels(SecretStoreType storeType, string[] prefixes, string? apiKeyProvisional, CancellationToken token)
|
||||
{
|
||||
return this.LoadModelsResponse<ModelsResponse>(
|
||||
storeType,
|
||||
@@ -115,7 +116,6 @@ public sealed class ProviderX() : BaseProvider(LLMProviders.X, new Uri("https://
|
||||
DisplayName = "Grok 2.0 (latest)",
|
||||
}
|
||||
]),
|
||||
token,
|
||||
apiKeyProvisional);
|
||||
apiKeyProvisional, token: token);
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user