temp save

This commit is contained in:
Jicheng Lu 2025-05-13 17:33:23 -05:00
parent da9c3c6ab0
commit d9fe6da099
5 changed files with 161 additions and 22 deletions

View file

@ -0,0 +1,18 @@
using GenerativeAI.Types;
namespace BotSharp.Plugin.GoogleAI.Models.Realtime;
internal class RealtimeClientPayload
{
[JsonPropertyName("setup")]
public RealtimeGenerateContentSetup? Setup { get; set; }
[JsonPropertyName("clientContent")]
public BidiGenerateContentClientContent? ClientContent { get; set; }
[JsonPropertyName("realtimeInput")]
public BidiGenerateContentRealtimeInput? RealtimeInput { get; set; }
[JsonPropertyName("toolResponse")]
public BidiGenerateContentToolResponse? ToolResponse { get; set; }
}

View file

@ -0,0 +1,26 @@
using GenerativeAI.Types;
namespace BotSharp.Plugin.GoogleAI.Models.Realtime;
internal class RealtimeGenerateContentSetup
{
[JsonPropertyName("model")]
public string? Model { get; set; }
[JsonPropertyName("generationConfig")]
public GenerationConfig? GenerationConfig { get; set; }
[JsonPropertyName("systemInstruction")]
public Content? SystemInstruction { get; set; }
[JsonPropertyName("tools")]
public Tool[]? Tools { get; set; }
[JsonPropertyName("inputAudioTranscription")]
public AudioTranscriptionConfig? InputAudioTranscription { get; set; }
[JsonPropertyName("outputAudioTranscription")]
public AudioTranscriptionConfig? OutputAudioTranscription { get; set; }
}
internal class AudioTranscriptionConfig { }

View file

@ -30,6 +30,12 @@ internal class RealtimeGenerateContentServerContent
[JsonPropertyName("modelTurn")]
public Content? ModelTurn { get; set; }
[JsonPropertyName("inputTranscription")]
public RealtimeGenerateContentTranscription? InputTranscription { get; set; }
[JsonPropertyName("outputTranscription")]
public RealtimeGenerateContentTranscription? OutputTranscription { get; set; }
}
internal class RealtimeUsageMetaData
@ -58,4 +64,10 @@ internal class RealtimeTokenDetail
[JsonPropertyName("tokenCount")]
public int? TokenCount { get; set; }
}
internal class RealtimeGenerateContentTranscription
{
[JsonPropertyName("text")]
public string? Text { get; set; }
}

View file

@ -107,8 +107,9 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
JsonOptions = _jsonOptions
});
var uri = BuildWebsocketUri(modelSettings.ApiKey, "v1beta");
await _session.ConnectAsync(
uri: new Uri($"wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.v1beta.GenerativeService.BidiGenerateContent?key={modelSettings.ApiKey}"),
uri: uri,
cancellationToken: CancellationToken.None);
await onModelReady();
@ -148,9 +149,12 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
Func<string, Task> onModelAudioTranscriptDone,
Func<List<RoleDialogModel>, Task> onModelResponseDone,
Func<string, Task> onConversationItemCreated,
Func<RoleDialogModel, Task> onInputAudioTranscriptionCompleted,
Func<RoleDialogModel, Task> onInputAudioTranscriptionDone,
Func<Task> onInterruptionDetected)
{
var inputTranscription = string.Empty;
var outputTranscription = string.Empty;
await foreach (ChatSessionUpdate update in _session.ReceiveUpdatesAsync(CancellationToken.None))
{
var receivedText = update?.RawResponse;
@ -163,7 +167,6 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
try
{
var response = JsonSerializer.Deserialize<RealtimeServerResponse>(receivedText, _jsonOptions);
if (response == null)
{
continue;
@ -175,10 +178,29 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
}
else if (response.ServerContent != null)
{
if (response.ServerContent.InputTranscription?.Text != null)
{
outputTranscription = string.Empty;
inputTranscription += response.ServerContent.InputTranscription.Text;
}
if (response.ServerContent.OutputTranscription?.Text != null)
{
outputTranscription += response.ServerContent.OutputTranscription.Text;
}
if (response.ServerContent.ModelTurn != null)
{
_logger.LogInformation($"Model audio delta received.");
var parts = response.ServerContent.ModelTurn.Parts;
if (!string.IsNullOrEmpty(inputTranscription))
{
var message = await OnUserAudioTranscriptionCompleted(conn, inputTranscription);
await onInputAudioTranscriptionDone(message);
inputTranscription = string.Empty;
}
if (!parts.IsNullOrEmpty())
{
foreach (var part in parts)
@ -197,13 +219,23 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
else if (response.ServerContent.TurnComplete == true)
{
_logger.LogInformation($"Model turn completed.");
if (!string.IsNullOrEmpty(outputTranscription))
{
var messages = await OnResponseDone(conn, outputTranscription, response.UsageMetaData);
await onModelResponseDone(messages);
// Reset input/output transcription
inputTranscription = string.Empty;
outputTranscription = string.Empty;
}
}
}
}
catch (Exception ex)
{
_logger.LogError(ex, $"Error when deserializing server response.");
continue;
_logger.LogError(ex, $"Error when deserializing server response. {ex.Message}");
break;
}
}
@ -288,7 +320,7 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
client.Connected += (sender, e) =>
{
_logger.LogInformation("Google Realtime Client connected.");
_onModelReady();
_onModelReady().ConfigureAwait(false).GetAwaiter().GetResult();
};
client.Disconnected += (sender, e) =>
@ -301,7 +333,7 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
_logger.LogInformation("User message received.");
if (e.Payload.SetupComplete != null)
{
_onConversationItemCreated(_client.ConnectionId.ToString());
_onConversationItemCreated(_client.ConnectionId.ToString()).ConfigureAwait(false).GetAwaiter().GetResult();
}
if (e.Payload.ServerContent != null)
@ -309,31 +341,31 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
if (e.Payload.ServerContent.TurnComplete == true)
{
var responseDone = await ResponseDone(_conn, e.Payload.ServerContent);
_onModelResponseDone(responseDone);
_onModelResponseDone(responseDone).ConfigureAwait(false).GetAwaiter().GetResult();
}
}
};
client.AudioChunkReceived += (sender, e) =>
{
_onModelAudioDeltaReceived(Convert.ToBase64String(e.Buffer), Guid.NewGuid().ToString());
_onModelAudioDeltaReceived(Convert.ToBase64String(e.Buffer), Guid.NewGuid().ToString()).ConfigureAwait(false).GetAwaiter().GetResult();
};
client.TextChunkReceived += (sender, e) =>
{
_onInputAudioTranscriptionDone(new RoleDialogModel(AgentRole.Assistant, e.Text));
_onInputAudioTranscriptionDone(new RoleDialogModel(AgentRole.Assistant, e.Text)).ConfigureAwait(false).GetAwaiter().GetResult();
};
client.GenerationInterrupted += (sender, e) =>
{
_logger.LogInformation("Audio generation interrupted.");
_onUserInterrupted();
_onUserInterrupted().ConfigureAwait(false).GetAwaiter().GetResult();
};
client.AudioReceiveCompleted += (sender, e) =>
{
_logger.LogInformation("Audio receive completed.");
_onModelAudioResponseDone();
_onModelAudioResponseDone().ConfigureAwait(false).GetAwaiter().GetResult();
};
client.ErrorOccurred += (sender, e) =>
@ -345,6 +377,43 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
return Task.CompletedTask;
}
private async Task<List<RoleDialogModel>> OnResponseDone(RealtimeHubConnection conn, string text, RealtimeUsageMetaData? useage)
{
var outputs = new List<RoleDialogModel>
{
new(AgentRole.Assistant, text)
{
CurrentAgentId = conn.CurrentAgentId,
MessageId = Guid.NewGuid().ToString(),
MessageType = MessageTypeName.Plain
}
};
if (useage != null)
{
var contentHooks = _services.GetServices<IContentGeneratingHook>();
foreach (var hook in contentHooks)
{
await hook.AfterGenerated(new RoleDialogModel(AgentRole.Assistant, text)
{
CurrentAgentId = conn.CurrentAgentId
},
new TokenStatsModel
{
Provider = Provider,
Model = _model,
Prompt = text,
TextInputTokens = useage.PromptTokensDetails?.FirstOrDefault(x => x.Modality == Modality.TEXT.ToString())?.TokenCount ?? 0,
AudioInputTokens = useage.PromptTokensDetails?.FirstOrDefault(x => x.Modality == Modality.AUDIO.ToString())?.TokenCount ?? 0,
TextOutputTokens = useage.ResponseTokensDetails?.FirstOrDefault(x => x.Modality == Modality.TEXT.ToString())?.TokenCount ?? 0,
AudioOutputTokens = useage.ResponseTokensDetails?.FirstOrDefault(x => x.Modality == Modality.AUDIO.ToString())?.TokenCount ?? 0
});
}
}
return outputs;
}
private async Task<List<RoleDialogModel>> ResponseDone(RealtimeHubConnection conn,
BidiGenerateContentServerContent serverContent)
{
@ -401,8 +470,6 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
public async Task SendEventToModel(object message)
{
//todo Send Audio Chunks to Model, Botsharp RealTime Implementation seems to be incomplete
if (_session == null) return;
await _session.SendEventToModel(message);
@ -419,9 +486,9 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
var (prompt, request) = PrepareOptions(agent, []);
var config = request.GenerationConfig;
//Output Modality can either be text or audio
if (config != null)
{
//Output Modality can either be text or audio
config.ResponseModalities = [Modality.AUDIO];
var words = new List<string>();
@ -467,14 +534,16 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
// //Tools = request.Tools?.ToArray(),
//});
await SendEventToModel(new BidiClientPayload
await SendEventToModel(new RealtimeClientPayload
{
Setup = new BidiGenerateContentSetup()
Setup = new RealtimeGenerateContentSetup()
{
GenerationConfig = config,
Model = Model.ToModelId(),
SystemInstruction = request.SystemInstruction,
Tools = []
Tools = [],
InputAudioTranscription = new(),
OutputAudioTranscription = new()
}
});
@ -532,7 +601,7 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
}
else
{
throw new NotImplementedException("");
throw new NotImplementedException($"Unrecognized role {message.Role}.");
}
}
@ -542,9 +611,9 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
}
public async Task<RoleDialogModel> OnConversationItemCreated(RealtimeHubConnection conn, string response)
public async Task<RoleDialogModel> OnConversationItemCreated(RealtimeHubConnection conn, string text)
{
return await Task.FromResult(new RoleDialogModel(AgentRole.User, response));
return await Task.FromResult(new RoleDialogModel(AgentRole.User, text));
}
private (string, GenerateContentRequest) PrepareOptions(Agent agent,
@ -688,4 +757,18 @@ public class GoogleRealTimeProvider : IRealTimeCompletion
return prompt;
}
private async Task<RoleDialogModel> OnUserAudioTranscriptionCompleted(RealtimeHubConnection conn, string text)
{
return new RoleDialogModel(AgentRole.User, text)
{
CurrentAgentId = conn.CurrentAgentId
};
}
private Uri BuildWebsocketUri(string apiKey, string version = "v1alpha")
{
return new Uri($"wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.{version}.GenerativeService.BidiGenerateContent?key={apiKey}");
}
}

View file

@ -402,7 +402,7 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
}
else
{
throw new NotImplementedException("");
throw new NotImplementedException($"Unrecognized role {message.Role}.");
}
}