Update nullability, clean up formatting, and improve naming

This commit is contained in:
Gunpal Jain 2025-04-04 06:48:14 +05:30
parent 98d706415c
commit f46737e118
14 changed files with 75 additions and 72 deletions

View file

@ -23,7 +23,7 @@ public class FunctionDef
public string? Impact { get; set; }
[JsonPropertyName("parameters")]
public FunctionParametersDef Parameters { get; set; } = new FunctionParametersDef();
public FunctionParametersDef? Parameters { get; set; } = new FunctionParametersDef();
[JsonPropertyName("output")]
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]

View file

@ -89,17 +89,20 @@ public class ChatCompletionProvider : IChatCompletion
return responseMessage;
}
public Task<bool> GetChatCompletionsAsync(Agent agent, List<RoleDialogModel> conversations, Func<RoleDialogModel, Task> onMessageReceived, Func<RoleDialogModel, Task> onFunctionExecuting)
public Task<bool> GetChatCompletionsAsync(Agent agent, List<RoleDialogModel> conversations,
Func<RoleDialogModel, Task> onMessageReceived, Func<RoleDialogModel, Task> onFunctionExecuting)
{
throw new NotImplementedException();
}
public Task<bool> GetChatCompletionsStreamingAsync(Agent agent, List<RoleDialogModel> conversations, Func<RoleDialogModel, Task> onMessageReceived)
public Task<bool> GetChatCompletionsStreamingAsync(Agent agent, List<RoleDialogModel> conversations,
Func<RoleDialogModel, Task> onMessageReceived)
{
throw new NotImplementedException();
}
private (string, MessageParameters) PrepareOptions(Agent agent, List<RoleDialogModel> conversations, LlmModelSetting settings)
private (string, MessageParameters) PrepareOptions(Agent agent, List<RoleDialogModel> conversations,
LlmModelSetting settings)
{
var instruction = "";
renderedInstructions = [];
@ -178,8 +181,8 @@ public class ChatCompletionProvider : IChatCompletion
var state = _services.GetRequiredService<IConversationStateService>();
var temperature = decimal.Parse(state.GetState("temperature", "0.0"));
var maxTokens = int.TryParse(state.GetState("max_tokens"), out var tokens)
? tokens
: agent.LlmConfig?.MaxOutputTokens ?? LlmConstant.DEFAULT_MAX_OUTPUT_TOKEN;
? tokens
: agent.LlmConfig?.MaxOutputTokens ?? LlmConstant.DEFAULT_MAX_OUTPUT_TOKEN;
var parameters = new MessageParameters()
{
@ -197,9 +200,11 @@ public class ChatCompletionProvider : IChatCompletion
{
new SystemMessage(instruction)
};
};
}
JsonSerializerOptions jsonSerializationOptions = new()
;
JsonSerializerOptions? jsonSerializationOptions = new()
{
DefaultIgnoreCondition = JsonIgnoreCondition.WhenWritingNull,
Converters = { new JsonStringEnumConverter() },
@ -238,7 +243,7 @@ public class ChatCompletionProvider : IChatCompletion
private string GetPrompt(MessageParameters parameters)
{
var prompt = $"{string.Join("\r\n", (parameters.System??new List<SystemMessage>()).Select(x => x.Text))}\r\n";
var prompt = $"{string.Join("\r\n", (parameters.System ?? new List<SystemMessage>()).Select(x => x.Text))}\r\n";
prompt += "\r\n[CONVERSATION]";
var verbose = string.Join("\r\n", parameters.Messages
@ -272,6 +277,7 @@ public class ChatCompletionProvider : IChatCompletion
}));
return $"{role}: {content}";
}
return string.Empty;
}));
@ -279,10 +285,12 @@ public class ChatCompletionProvider : IChatCompletion
if (parameters.Tools != null && parameters.Tools.Count > 0)
{
var functions = string.Join("\r\n", parameters.Tools.Select(x =>
{
return $"\r\n{x.Function.Name}: {x.Function.Description}\r\n{JsonSerializer.Serialize(x.Function.Parameters)}";
}));
var functions = string.Join("\r\n",
parameters.Tools.Select(x =>
{
return
$"\r\n{x.Function.Name}: {x.Function.Description}\r\n{JsonSerializer.Serialize(x.Function.Parameters)}";
}));
prompt += $"\r\n[FUNCTIONS]\r\n{functions}\r\n";
}

View file

@ -21,13 +21,13 @@ public class GeminiChatCompletionProvider : IChatCompletion
public string Provider => "google-ai";
public string Model => _model;
private GoogleAiSettings _googleSettings;
private GoogleAiSettings _settings;
public GeminiChatCompletionProvider(
IServiceProvider services,
GoogleAiSettings googleSettings,
ILogger<GeminiChatCompletionProvider> logger)
{
_googleSettings = googleSettings;
_settings = googleSettings;
_services = services;
_logger = logger;
}
@ -42,7 +42,7 @@ public class GeminiChatCompletionProvider : IChatCompletion
await hook.BeforeGenerating(agent, conversations);
}
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _googleSettings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var aiModel = client.CreateGenerativeModel(_model);
var (prompt, request) = PrepareOptions(aiModel, agent, conversations);
@ -101,7 +101,7 @@ public class GeminiChatCompletionProvider : IChatCompletion
await hook.BeforeGenerating(agent, conversations);
}
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _googleSettings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var chatClient = client.CreateGenerativeModel(_model);
var (prompt, messages) = PrepareOptions(chatClient, agent, conversations);
@ -166,7 +166,7 @@ public class GeminiChatCompletionProvider : IChatCompletion
public async Task<bool> GetChatCompletionsStreamingAsync(Agent agent, List<RoleDialogModel> conversations, Func<RoleDialogModel, Task> onMessageReceived)
{
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _googleSettings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var chatClient = client.CreateGenerativeModel(_model);
var (prompt, messages) = PrepareOptions(chatClient,agent, conversations);

View file

@ -30,7 +30,7 @@ public class TextEmbeddingProvider : ITextEmbedding
public async Task<float[]> GetVectorAsync(string text)
{
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _settings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var embeddingClient = client.CreateEmbeddingModel(_model);
var response = await embeddingClient.EmbedContentAsync(text);
@ -40,7 +40,7 @@ public class TextEmbeddingProvider : ITextEmbedding
public async Task<List<float[]>> GetVectorsAsync(List<string> texts)
{
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _settings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var embeddingClient = client.CreateEmbeddingModel(_model);
var response = await embeddingClient.BatchEmbedContentAsync(texts.Select(s=>new Content(s, Roles.User)));

View file

@ -5,20 +5,20 @@ namespace BotSharp.Plugin.GoogleAi.Providers;
public static class ProviderHelper
{
public static GenerativeAI.GoogleAi GetGeminiClient(string provider, string model, IServiceProvider services, GoogleAiSettings? aiSettings, ILogger? _logger)
public static GenerativeAI.GoogleAi GetGeminiClient(string provider, string model, IServiceProvider services)
{
var aiSettings = services.GetRequiredService<GoogleAiSettings>();
if (aiSettings == null || aiSettings.Gemini ==null || string.IsNullOrEmpty(aiSettings.Gemini.ApiKey))
{
var settingsService = services.GetRequiredService<ILlmProviderService>();
var settings = settingsService.GetSetting(provider, model);
var client = new GenerativeAI.GoogleAi(settings.ApiKey, logger:_logger);
var client = new GenerativeAI.GoogleAi(settings.ApiKey);
return client;
}
else
{
return new GenerativeAI.GoogleAi(aiSettings.Gemini.ApiKey, logger:_logger);
return new GenerativeAI.GoogleAi(aiSettings.Gemini.ApiKey);
}
}
public static GooglePalmClient GetPalmClient(string provider, string model, IServiceProvider services)

View file

@ -33,13 +33,14 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
private readonly ILogger<GeminiChatCompletionProvider> _logger;
private List<string> renderedInstructions = [];
private readonly GoogleAiSettings _googleSettings;
private readonly GoogleAiSettings _settings;
public GoogleRealTimeProvider(
IServiceProvider services,
GoogleAiSettings googleSettings,
ILogger<GeminiChatCompletionProvider> logger)
{
_googleSettings = googleSettings;
_settings = googleSettings;
_services = services;
_logger = logger;
}
@ -88,7 +89,7 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
public async Task AppenAudioBuffer(string message)
{
await _client.SendAudioAsync(Convert.FromBase64String(message));
await _client.SendAudioAsync(Convert.FromBase64String(message));
}
public async Task TriggerModelInference(string? instructions = null)
@ -101,13 +102,10 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
public async Task CancelModelResponse()
{
}
public async Task RemoveConversationItem(string itemId)
{
}
private async Task AttachEvents()
@ -124,8 +122,8 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
{
if (e.Payload.ServerContent.TurnComplete == true)
{
var responseDone = await ResponseDone(conn,e.Payload.ServerContent);
onModelResponseDone(responseDone);
var responseDone = await ResponseDone(conn, e.Payload.ServerContent);
onModelResponseDone(responseDone);
}
}
};
@ -139,13 +137,11 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
onInputAudioTranscriptionCompleted(new RoleDialogModel(AgentRole.Assistant, e.Text));
};
_client.GenerationInterrupted += async (sender, e) => { onUserInterrupted(); };
_client.AudioReceiveCompleted += async (sender, e) =>
{
onModelAudioResponseDone();
};
_client.AudioReceiveCompleted += async (sender, e) => { onModelAudioResponseDone(); };
}
private async Task<List<RoleDialogModel>> ResponseDone(RealtimeHubConnection conn, BidiGenerateContentServerContent serverContent)
private async Task<List<RoleDialogModel>> ResponseDone(RealtimeHubConnection conn,
BidiGenerateContentServerContent serverContent)
{
var outputs = new List<RoleDialogModel>();
@ -194,6 +190,7 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
Model = _model,
});
}
return outputs;
}
@ -206,29 +203,25 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
{
var contentHooks = _services.GetServices<IContentGeneratingHook>().ToList();
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _googleSettings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var chatClient = client.CreateGenerativeModel(_model);
var (prompt, request) = PrepareOptions(chatClient, agent, conversations);
var config = request.GenerationConfig;
//Output Modality can either be text or audio
config.ResponseModalities = new List<Modality>([Modality.AUDIO]);
var settingsService = _services.GetRequiredService<ILlmProviderService>();
var settings = settingsService.GetSetting(Provider, _model);
_client = chatClient.CreateMultiModalLiveClient(config,
systemInstruction: request.SystemInstruction?.Parts.FirstOrDefault()?.Text);
_client.UseGoogleSearch = _googleSettings.Gemini.UseGoogleSearch;
_client.UseGoogleSearch = _settings.Gemini.UseGoogleSearch;
if (request.Tools != null && request.Tools.Count > 0)
{
var lst = (request.Tools.Select(s => (IFunctionTool)new FakeFunctionTool(s)).ToList());
var lst = (request.Tools.Select(s => (IFunctionTool)new TemporaryFunctionTool(s)).ToList());
_client.AddFunctionTools(lst, new ToolConfig()
{
@ -261,7 +254,7 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
var agentService = _services.GetRequiredService<IAgentService>();
var agent = await agentService.LoadAgent(conn.CurrentAgentId);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _googleSettings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var chatClient = client.CreateGenerativeModel(_model);
var (prompt, request) = PrepareOptions(chatClient, agent, new List<RoleDialogModel>());
@ -330,7 +323,7 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
await _client.SendClientContentAsync(new BidiGenerateContentClientContent()
{
TurnComplete = true,
Turns = new []{new Content(message.Content, AgentRole.User)}
Turns = new[] { new Content(message.Content, AgentRole.User) }
});
}
else
@ -354,7 +347,7 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
List<RoleDialogModel> conversations)
{
var agentService = _services.GetRequiredService<IAgentService>();
var googleSettings = _googleSettings;
var googleSettings = _settings;
renderedInstructions = [];
// Add settings
@ -472,7 +465,7 @@ namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
GenerationConfig = new()
{
Temperature = temperature,
MaxOutputTokens = maxTokens
MaxOutputTokens = maxTokens
}
};

View file

@ -4,11 +4,11 @@ using GenerativeAI.Types;
namespace BotSharp.Plugin.GoogleAi.Providers.Realtime
{
public class FakeFunctionTool:IFunctionTool
public class TemporaryFunctionTool:IFunctionTool
{
public Tool Tool { get; set; }
public FakeFunctionTool(Tool tool)
public TemporaryFunctionTool(Tool tool)
{
this.Tool = tool;
}

View file

@ -17,14 +17,14 @@ public class GeminiTextCompletionProvider : ITextCompletion
public string Provider => "google-ai";
public string Model => _model;
private GoogleAiSettings _googleSettings;
private GoogleAiSettings _settings;
public GeminiTextCompletionProvider(
IServiceProvider services,
GoogleAiSettings googleSettings,
ILogger<GeminiTextCompletionProvider> logger,
ITokenStatistics tokenStatistics)
{
_googleSettings = googleSettings;
_settings = googleSettings;
_services = services;
_logger = logger;
_tokenStatistics = tokenStatistics;
@ -50,7 +50,7 @@ public class GeminiTextCompletionProvider : ITextCompletion
await hook.BeforeGenerating(agent, new List<RoleDialogModel> { userMessage });
}
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services, _googleSettings, _logger);
var client = ProviderHelper.GetGeminiClient(Provider, _model, _services);
var aiModel = client.CreateGenerativeModel(_model);
PrepareOptions(aiModel);

View file

@ -169,11 +169,18 @@
"Provider": "google-ai",
"Models": [
{
"Name": "gemini-2.0-flash-exp",
"Name": "gemini-2.0-flash",
"ApiKey": "",
"Type": "chat",
"MultiModal": true,
"RealTime": true,
"PromptCost": 0.0015,
"CompletionCost": 0.002
},
{
"Name": "gemini-2.0-flash-exp",
"ApiKey": "",
"Type": "realtime",
"MultiModal": true,
"PromptCost": 0.0015,
"CompletionCost": 0.002
}

View file

@ -11,7 +11,7 @@ using Shouldly;
namespace BotSharp.Plugin.Google.Core
{
public class ChatCompletion_Tests:TestBase
public class ChatCompletionTests:TestBase
{
protected static Agent CreateTestAgent()
{
@ -56,7 +56,7 @@ namespace BotSharp.Plugin.Google.Core
yield return new object[] { services.BuildServiceProvider().GetService<IChatCompletion>() ?? throw new Exception("Error while initializing"), agent, modelName };
}
}
public ChatCompletion_Tests()
public ChatCompletionTests()
{
}

View file

@ -9,7 +9,7 @@ using Shouldly;
namespace BotSharp.Plugin.Google.Core
{
public class Embedding_Tests:TestBase
public class EmbeddingTests:TestBase
{
protected static Agent CreateTestAgent()
{

View file

@ -10,7 +10,7 @@ using Shouldly;
namespace BotSharp.Plugin.Google.Core
{
public class FunctionCalling_Tests : TestBase
public class FunctionCallingTests : TestBase
{
public static IEnumerable<object[]> CreateTestLLMProviders()
{

View file

@ -11,7 +11,7 @@ using Shouldly;
namespace BotSharp.Plugin.Google.Core
{
public class GoogleRealTime_Tests : TestBase
public class GoogleRealTimeTests : TestBase
{
protected static Agent CreateTestAgent()
{

View file

@ -100,9 +100,4 @@
<CopyToOutputDirectory>PreserveNewest</CopyToOutputDirectory>
</Content>
</ItemGroup>
<ItemGroup>
<ProjectReference Include="..\..\src\Infrastructure\BotSharp.Abstraction\BotSharp.Abstraction.csproj" />
</ItemGroup>
</Project>