Merge pull request #915 from hchen2020/master

refactor realtime code.
This commit is contained in:
Haiping 2025-03-05 16:59:32 -06:00 committed by GitHub
commit 4a5ec1cd08
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 263 additions and 204 deletions

View file

@ -24,7 +24,7 @@ public interface IRealTimeCompletion
Task Disconnect();
Task<RealtimeSession> CreateSession(Agent agent, List<RoleDialogModel> conversations);
Task UpdateSession(RealtimeHubConnection conn);
Task UpdateSession(RealtimeHubConnection conn, bool turnDetection = true);
Task InsertConversationItem(RoleDialogModel message);
Task RemoveConversationItem(string itemId);
Task TriggerModelInference(string? instructions = null);

View file

@ -1,9 +1,16 @@
using System.Collections.Concurrent;
namespace BotSharp.Abstraction.Realtime.Models;
public class RealtimeHubConnection
{
public string Event { get; set; } = null!;
public string StreamId { get; set; } = null!;
public string? LastAssistantItem { get; set; } = null!;
public long LatestMediaTimestamp { get; set; }
public long? ResponseStartTimestamp { get; set; }
public string KeypadInputBuffer { get; set; } = string.Empty;
public ConcurrentQueue<string> MarkQueue { get; set; } = new();
public string CurrentAgentId { get; set; } = null!;
public string ConversationId { get; set; } = null!;
public string Data { get; set; } = string.Empty;

View file

@ -4,6 +4,10 @@ using BotSharp.Abstraction.Realtime.Models;
using BotSharp.Abstraction.MLTasks;
using BotSharp.Abstraction.Conversations.Enums;
using BotSharp.Abstraction.Routing.Models;
using NetTopologySuite.Index.HPRtree;
using BotSharp.Abstraction.Agents.Models;
using Microsoft.Identity.Client.Extensions.Msal;
using Microsoft.AspNetCore.Cors.Infrastructure;
namespace BotSharp.Core.Realtime;
@ -46,9 +50,14 @@ public class RealtimeHub : IRealtimeHub
{
await completer.AppenAudioBuffer(conn.Data);
}
else if (conn.Event == "user_dtmf_received")
{
await HandleUserDtmfReceived(completer, conn);
}
else if (conn.Event == "user_disconnected")
{
await completer.Disconnect();
await HandleUserDisconnected(conn);
}
} while (!result.CloseStatus.HasValue);
@ -58,8 +67,6 @@ public class RealtimeHub : IRealtimeHub
private async Task ConnectToModel(IRealTimeCompletion completer, WebSocket userWebSocket, RealtimeHubConnection conn)
{
var hookProvider = _services.GetRequiredService<ConversationHookProvider>();
var storage = _services.GetRequiredService<IConversationStorage>();
var convService = _services.GetRequiredService<IConversationService>();
convService.SetConversationId(conn.ConversationId, []);
var conversation = await convService.GetConversation(conn.ConversationId);
@ -92,8 +99,8 @@ public class RealtimeHub : IRealtimeHub
await completer.Connect(conn,
onModelReady: async () =>
{
// Control initial session
await completer.UpdateSession(conn);
// Control initial session, prevent initial response interruption
await completer.UpdateSession(conn, turnDetection: false);
// Add dialog history
foreach (var item in dialogs)
@ -103,17 +110,41 @@ public class RealtimeHub : IRealtimeHub
if (dialogs.LastOrDefault()?.Role == AgentRole.Assistant)
{
// await completer.TriggerModelInference($"Rephase your last response:\r\n{dialogs.LastOrDefault()?.Content}");
await completer.TriggerModelInference($"Rephase your last response:\r\n{dialogs.LastOrDefault()?.Content}");
}
else
{
await completer.TriggerModelInference("Reply based on the conversation context.");
}
// Start turn detection
await Task.Delay(1000 * 8);
await completer.UpdateSession(conn, turnDetection: true);
},
onModelAudioDeltaReceived: async audioDeltaData =>
{
// If this is the first delta of a new response, set the start timestamp
if (!conn.ResponseStartTimestamp.HasValue)
{
conn.ResponseStartTimestamp = conn.LatestMediaTimestamp;
_logger.LogDebug($"Setting start timestamp for new response: {conn.ResponseStartTimestamp}ms");
}
var data = conn.OnModelMessageReceived(audioDeltaData);
await SendEventToUser(userWebSocket, data);
// Send mark messages to Media Streams so we know if and when AI response playback is finished
if (!string.IsNullOrEmpty(conn.StreamId))
{
var markEvent = new
{
@event = "mark",
streamSid = conn.StreamId,
mark = new { name = "responsePart" }
};
await SendEventToUser(userWebSocket, markEvent);
conn.MarkQueue.Enqueue("responsePart");
}
},
onModelAudioResponseDone: async () =>
{
@ -160,16 +191,18 @@ public class RealtimeHub : IRealtimeHub
await completer.TriggerModelInference("Reply based on the function's output.");
}
}
// append output audio transcript to conversation
storage.Append(conn.ConversationId, message);
dialogs.Add(message);
foreach (var hook in hookProvider.HooksOrderByPriority)
else
{
hook.SetAgent(agent)
.SetConversation(conversation);
// append output audio transcript to conversation
dialogs.Add(message);
await hook.OnResponseGenerated(message);
foreach (var hook in hookProvider.HooksOrderByPriority)
{
hook.SetAgent(agent)
.SetConversation(conversation);
await hook.OnResponseGenerated(message);
}
}
}
},
@ -180,7 +213,6 @@ public class RealtimeHub : IRealtimeHub
onInputAudioTranscriptionCompleted: async message =>
{
// append input audio transcript to conversation
storage.Append(conn.ConversationId, message);
dialogs.Add(message);
foreach (var hook in hookProvider.HooksOrderByPriority)
@ -193,11 +225,56 @@ public class RealtimeHub : IRealtimeHub
},
onUserInterrupted: async () =>
{
// Reset states
conn.MarkQueue.Clear();
conn.LastAssistantItem = null;
conn.ResponseStartTimestamp = null;
var data = conn.OnModelUserInterrupted();
await SendEventToUser(userWebSocket, data);
});
}
private async Task HandleUserDtmfReceived(IRealTimeCompletion completer, RealtimeHubConnection conn)
{
var routing = _services.GetRequiredService<IRoutingService>();
var hookProvider = _services.GetRequiredService<ConversationHookProvider>();
var agentService = _services.GetRequiredService<IAgentService>();
var agent = await agentService.LoadAgent(conn.CurrentAgentId);
var dialogs = routing.Context.GetDialogs();
var convService = _services.GetRequiredService<IConversationService>();
var conversation = await convService.GetConversation(conn.ConversationId);
var message = new RoleDialogModel(AgentRole.User, conn.Data)
{
CurrentAgentId = routing.Context.GetCurrentAgentId()
};
dialogs.Add(message);
foreach (var hook in hookProvider.HooksOrderByPriority)
{
hook.SetAgent(agent)
.SetConversation(conversation);
await hook.OnMessageReceived(message);
}
await completer.InsertConversationItem(message);
await completer.TriggerModelInference("Reply based on the user input");
}
private async Task HandleUserDisconnected(RealtimeHubConnection conn)
{
// Save dialog history
var routing = _services.GetRequiredService<IRoutingService>();
var storage = _services.GetRequiredService<IConversationStorage>();
var dialogs = routing.Context.GetDialogs();
foreach (var item in dialogs)
{
storage.Append(conn.ConversationId, item);
}
}
private async Task SendEventToUser(WebSocket webSocket, object message)
{
var data = JsonSerializer.Serialize(message);

View file

@ -1,4 +1,5 @@
using BotSharp.Abstraction.Agents;
using BotSharp.Abstraction.Conversations.Enums;
using BotSharp.Abstraction.Infrastructures.Enums;
using BotSharp.Abstraction.Translation;
using System;
@ -26,6 +27,11 @@ namespace BotSharp.Logger.Hooks
{
return;
}
if (_states.GetState("channel") == ConversationChannel.Phone)
{
return;
}
// Handle multi-language for output
var agentService = _services.GetRequiredService<IAgentService>();

View file

@ -35,7 +35,7 @@ public class VerboseLogHook : IContentGeneratingHook
public async Task AfterGenerated(RoleDialogModel message, TokenStatsModel tokenStats)
{
if (!_convSettings.ShowVerboseLog) return;
if (!_convSettings.ShowVerboseLog || string.IsNullOrEmpty(tokenStats.Prompt)) return;
var agentService = _services.GetRequiredService<IAgentService>();
var agent = await agentService.LoadAgent(message.CurrentAgentId);

View file

@ -47,7 +47,7 @@ public class RealtimeSessionBody
public FunctionDef[] Tools { get; set; } = [];
[JsonPropertyName("turn_detection")]
public RealtimeSessionTurnDetection TurnDetection { get; set; } = new();
public RealtimeSessionTurnDetection? TurnDetection { get; set; } = new();
}
public class RealtimeSessionTurnDetection

View file

@ -1,8 +1,27 @@
using BotSharp.Abstraction.Functions.Models;
namespace BotSharp.Plugin.OpenAI.Models.Realtime;
public class RealtimeSessionCreationRequest : RealtimeSessionBody
public class RealtimeSessionCreationRequest
{
[JsonPropertyName("model")]
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
public string Model { get; set; } = null!;
[JsonPropertyName("modalities")]
public string[] Modalities { get; set; } = ["audio", "text"];
[JsonPropertyName("instructions")]
public string Instructions { get; set; } = null!;
[JsonPropertyName("tool_choice")]
public string ToolChoice { get; set; } = "auto";
[JsonPropertyName("tools")]
public FunctionDef[] Tools { get; set; } = [];
[JsonPropertyName("turn_detection")]
public RealtimeSessionTurnDetection TurnDetection { get; set; } = new();
}
/// <summary>

View file

@ -59,10 +59,9 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
if (_webSocket.State == WebSocketState.Open)
{
onModelReady();
// Receive a message
_ = ReceiveMessage(conn,
onModelReady,
onModelAudioDeltaReceived,
onModelAudioResponseDone,
onAudioTranscriptDone,
@ -75,7 +74,10 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
public async Task Disconnect()
{
await _webSocket.CloseAsync(WebSocketCloseStatus.Empty, null, CancellationToken.None);
if (_webSocket.State == WebSocketState.Open)
{
await _webSocket.CloseAsync(WebSocketCloseStatus.Empty, null, CancellationToken.None);
}
}
public async Task AppenAudioBuffer(string message)
@ -119,7 +121,8 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
});
}
private async Task ReceiveMessage(RealtimeHubConnection conn,
private async Task ReceiveMessage(RealtimeHubConnection conn,
Action onModelReady,
Action<string> onModelAudioDeltaReceived,
Action onModelAudioResponseDone,
Action<string> onAudioTranscriptDone,
@ -128,9 +131,8 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
Action<RoleDialogModel> onInputAudioTranscriptionCompleted,
Action onUserInterrupted)
{
var buffer = new byte[1024 * 256];
var buffer = new byte[1024 * 16];
WebSocketReceiveResult result;
string? lastAssistantItem = null;
do
{
@ -154,6 +156,7 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
else if (response.Type == "session.created")
{
_logger.LogInformation($"{response.Type}: {receivedText}");
onModelReady();
}
else if (response.Type == "session.updated")
{
@ -173,7 +176,16 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
else if (response.Type == "response.audio.delta")
{
var audio = JsonSerializer.Deserialize<ResponseAudioDelta>(receivedText);
lastAssistantItem = audio?.ItemId;
// Record last assistant item ID for interruption handling
if (conn.ResponseStartTimestamp.HasValue)
{
conn.ResponseStartTimestamp = conn.LatestMediaTimestamp;
}
if (!string.IsNullOrEmpty(conn.StreamId))
{
conn.LastAssistantItem = audio?.ItemId;
}
if (audio != null && audio.Delta != null)
{
@ -205,19 +217,24 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
}
else if (response.Type == "input_audio_buffer.speech_started")
{
// var elapsedTime = latestMediaTimestamp - responseStartTimestampTwilio;
// handle use interuption
if (!string.IsNullOrEmpty(lastAssistantItem))
// Handle user interuption
if (conn.MarkQueue.Count > 0 && conn.ResponseStartTimestamp != null)
{
var truncateEvent = new
{
type = "conversation.item.truncate",
item_id = lastAssistantItem,
content_index = 0,
audio_end_ms = 300
};
var elapsedTime = conn.LatestMediaTimestamp - conn.ResponseStartTimestamp;
if (!string.IsNullOrEmpty(conn.LastAssistantItem))
{
var truncateEvent = new
{
type = "conversation.item.truncate",
item_id = conn.LastAssistantItem,
content_index = 0,
audio_end_ms = elapsedTime
};
await SendEventToModel(truncateEvent);
}
await SendEventToModel(truncateEvent);
onUserInterrupted();
}
}
@ -256,6 +273,7 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
var args = new RealtimeSessionCreationRequest
{
Model = _model,
Instructions = instruction,
ToolChoice = "auto",
Tools = options.Tools.Select(x =>
@ -271,14 +289,14 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
};
var settingsService = _services.GetRequiredService<ILlmProviderService>();
var settings = settingsService.GetSetting(Provider, args.Model);
var settings = settingsService.GetSetting(Provider, args.Model ?? _model);
var api = _services.GetRequiredService<IOpenAiRealtimeApi>();
var session = await api.GetSessionAsync(args, settings.ApiKey);
return session;
}
public async Task UpdateSession(RealtimeHubConnection conn)
public async Task UpdateSession(RealtimeHubConnection conn, bool turnDetection = true)
{
var convService = _services.GetRequiredService<IConversationService>();
var conv = await convService.GetConversation(conn.ConversationId);
@ -318,16 +336,22 @@ public class RealTimeCompletionProvider : IRealTimeCompletion
ToolChoice = "auto",
Tools = functions,
Modalities = [ "text", "audio" ],
Temperature = Math.Max(options.Temperature ?? 0f, 0.6f),
Temperature = Math.Max(options.Temperature ?? 0f, 0.8f),
MaxResponseOutputTokens = 512,
TurnDetection = new RealtimeSessionTurnDetection
{
Threshold = 0.8f,
SilenceDuration = 800
Threshold = 0.5f,
PrefixPadding = 300,
SilenceDuration = 500
}
}
};
if (!turnDetection)
{
sessionUpdate.session.TurnDetection = null;
}
await HookEmitter.Emit<IContentGeneratingHook>(_services, async hook =>
{
await hook.OnSessionUpdated(agent, instruction, functions);

View file

@ -0,0 +1,17 @@
using BotSharp.Plugin.Twilio.Models.Stream;
using System.Text.Json.Serialization;
public class StreamEventDtmfResponse : StreamEventResponse
{
[JsonPropertyName("dtmf")]
public StreamEventDtmfBody Body { get; set; }
}
public class StreamEventDtmfBody
{
[JsonPropertyName("track")]
public string Track { get; set; }
[JsonPropertyName("digit")]
public string Digit { get; set; }
}

View file

@ -1,13 +1,9 @@
using BotSharp.Abstraction.Realtime;
using BotSharp.Abstraction.Realtime.Models;
using BotSharp.Core.Infrastructures;
using BotSharp.Plugin.Twilio.Interfaces;
using BotSharp.Plugin.Twilio.Models.Stream;
using Microsoft.AspNetCore.Http;
using System.Net.WebSockets;
using System.Text.Json;
using System.Collections.Concurrent;
using System.Text;
using Task = System.Threading.Tasks.Task;
namespace BotSharp.Plugin.Twilio.Services.Stream;
@ -37,7 +33,14 @@ public class TwilioStreamMiddleware
var services = httpContext.RequestServices;
var conversationId = request.Path.Value.Split("/").Last();
using WebSocket webSocket = await httpContext.WebSockets.AcceptWebSocketAsync();
await HandleWebSocket(services, conversationId, webSocket);
try
{
await HandleWebSocket(services, conversationId, webSocket);
}
catch (Exception ex)
{
_logger.LogError(ex, $"Error in WebSocket communication: {ex.Message} for conversation {conversationId}");
}
return;
}
}
@ -48,23 +51,15 @@ public class TwilioStreamMiddleware
private async Task HandleWebSocket(IServiceProvider services, string conversationId, WebSocket webSocket)
{
var hub = services.GetRequiredService<IRealtimeHub>();
var convService = services.GetRequiredService<IConversationService>();
// Session state
var conn = new RealtimeHubConnection
{
ConversationId = conversationId
};
// Variables for timestamp and interruption handling
string streamSid = null;
long latestMediaTimestamp = 0;
string lastAssistantItem = null;
var markQueue = new ConcurrentQueue<string>();
long? responseStartTimestampTwilio = null;
// Load session and state
convService.SetConversationId(conversationId, new List<MessageState>());
// load conversation and state
var convService = services.GetRequiredService<IConversationService>();
convService.SetConversationId(conversationId, []);
var hooks = services.GetServices<ITwilioSessionHook>();
foreach (var hook in hooks)
{
@ -72,159 +67,73 @@ public class TwilioStreamMiddleware
}
convService.States.Save();
// Set up event handlers
conn.OnModelMessageReceived = message =>
await hub.Listen(webSocket, (receivedText) =>
{
// Record last assistant item ID for interruption handling
if (!string.IsNullOrEmpty(conn.StreamId))
var response = JsonSerializer.Deserialize<StreamEventResponse>(receivedText);
conn.StreamId = response.StreamSid;
conn.Event = response.Event switch
{
lastAssistantItem = conn.StreamId;
}
// If this is the first delta of a new response, set the start timestamp
if (!responseStartTimestampTwilio.HasValue)
{
responseStartTimestampTwilio = latestMediaTimestamp;
_logger.LogDebug($"Setting start timestamp for new response: {responseStartTimestampTwilio}ms");
}
// Add mark to queue
markQueue.Enqueue("responsePart");
return new
{
@event = "media",
streamSid = conn.StreamId,
media = new { payload = message }
"start" => "user_connected",
"media" => "user_data_received",
"stop" => "user_disconnected",
_ => response.Event
};
};
conn.OnModelAudioResponseDone = () =>
{
return new
if (string.IsNullOrEmpty(conn.Event))
{
@event = "mark",
streamSid = conn.StreamId,
mark = new { name = "responsePart" }
};
};
conn.OnModelUserInterrupted = () =>
{
// Reset states
markQueue.Clear();
lastAssistantItem = null;
responseStartTimestampTwilio = null;
return new
{
@event = "clear",
streamSid = conn.StreamId
};
};
try
{
await hub.Listen(webSocket, receivedText =>
{
var response = JsonSerializer.Deserialize<StreamEventResponse>(receivedText);
if (response == null)
{
_logger.LogWarning("Failed to parse received WebSocket message");
return conn;
}
conn.StreamId = response.StreamSid;
switch (response.Event)
{
case "start":
conn.Event = "user_connected";
streamSid = response.StreamSid;
_logger.LogInformation($"Incoming stream started: {streamSid}");
// Reset start and media timestamps
responseStartTimestampTwilio = null;
latestMediaTimestamp = 0;
var startResponse = JsonSerializer.Deserialize<StreamEventStartResponse>(receivedText);
if (startResponse?.Body?.CustomParameters != null)
{
conn.Data = JsonSerializer.Serialize(startResponse.Body.CustomParameters);
}
break;
case "media":
conn.Event = "user_data_received";
var mediaResponse = JsonSerializer.Deserialize<StreamEventMediaResponse>(receivedText);
if (mediaResponse?.Body != null)
{
conn.Data = mediaResponse.Body.Payload;
// Update latest media timestamp
if (long.TryParse(mediaResponse.Body.Timestamp, out latestMediaTimestamp))
{
_logger.LogDebug($"Received media message with timestamp: {latestMediaTimestamp}ms");
}
// Check if user started speaking (interruption handling)
if (markQueue.Count > 0 && responseStartTimestampTwilio.HasValue &&
!string.IsNullOrEmpty(lastAssistantItem))
{
// Detect voice activity - more complex logic can be added here
// e.g., check audio energy levels or use VAD (Voice Activity Detection)
// If voice activity detected, handle interruption
if (ShouldHandleInterruption(mediaResponse.Body.Payload))
{
conn.Event = "user_interrupted";
long elapsedTime = latestMediaTimestamp - responseStartTimestampTwilio.Value;
_logger.LogDebug($"Calculating elapsed time for truncation: {latestMediaTimestamp} - {responseStartTimestampTwilio} = {elapsedTime}ms");
}
}
}
break;
case "mark":
// Handle mark event
if (markQueue.TryDequeue(out _))
{
_logger.LogDebug("Processing mark event, removing one mark from queue");
}
break;
case "stop":
conn.Event = "user_disconnected";
break;
default:
_logger.LogInformation($"Received non-media event: {response.Event}");
break;
}
return conn;
});
}
catch (Exception ex)
{
_logger.LogError(ex, "Error in WebSocket communication");
}
}
}
// Simple interruption detection logic - can be extended as needed
private bool ShouldHandleInterruption(string audioPayload)
{
// Here should implement actual voice activity detection logic
// e.g., analyze audio energy levels or use VAD algorithm
// Simple example - should be replaced with real detection logic in production
if (!string.IsNullOrEmpty(audioPayload))
{
// Check if audio payload contains sufficient energy
// This is just a placeholder - needs actual VAD implementation
return false; // Default to false to avoid false interruptions
}
return false;
conn.OnModelMessageReceived = message =>
new
{
@event = "media",
streamSid = response.StreamSid,
media = new { payload = message }
};
conn.OnModelAudioResponseDone = () =>
new
{
@event = "mark",
streamSid = response.StreamSid,
mark = new { name = "responsePart" }
};
conn.OnModelUserInterrupted = () =>
new
{
@event = "clear",
streamSid = response.StreamSid
};
if (response.Event == "start")
{
var startResponse = JsonSerializer.Deserialize<StreamEventStartResponse>(receivedText);
conn.LatestMediaTimestamp = 0;
conn.ResponseStartTimestamp = null;
conn.Data = JsonSerializer.Serialize(startResponse.Body.CustomParameters);
}
else if (response.Event == "media")
{
var mediaResponse = JsonSerializer.Deserialize<StreamEventMediaResponse>(receivedText);
conn.LatestMediaTimestamp = long.Parse(mediaResponse.Body.Timestamp);
conn.Data = mediaResponse.Body.Payload;
}
else if (response.Event == "dtmf")
{
var dtmfResponse = JsonSerializer.Deserialize<StreamEventDtmfResponse>(receivedText);
if (dtmfResponse.Body.Digit == "#")
{
conn.Event = "user_dtmf_received";
conn.Data = conn.KeypadInputBuffer;
conn.KeypadInputBuffer = string.Empty;
}
else
{
conn.KeypadInputBuffer += dtmfResponse.Body.Digit;
}
}
return conn;
});
}
}