Skip to content

Commit f193423

Browse files
committed
Add support to v3 websocket in C# SDK
1 parent 1cf8147 commit f193423

10 files changed

Lines changed: 592 additions & 50 deletions

sdks/csharp/src/CompressionHelpers.cs

Lines changed: 26 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -41,21 +41,20 @@ internal static GZipStream GzipReader(Stream stream)
4141
}
4242

4343
/// <summary>
44-
/// Decompresses and decodes a serialized <see cref="ServerMessage"/> from a byte array,
44+
/// Decompresses a serialized <see cref="ServerMessage"/> from a byte array,
4545
/// automatically handling the specified compression algorithm (None, Brotli, or Gzip).
4646
/// Ensures efficient decompression by reading the entire stream at once to avoid
4747
/// performance issues with certain stream implementations.
4848
/// Throws <see cref="InvalidOperationException"/> if an unknown compression type is encountered.
4949
/// </summary>
5050
/// <param name="bytes">The compressed and encoded server message as a byte array.</param>
51-
/// <returns>The deserialized <see cref="ServerMessage"/> object.</returns>
52-
internal static ServerMessage DecompressDecodeMessage(byte[] bytes)
51+
/// <returns>The decompressed encoded <see cref="ServerMessage"/> object.</returns>
52+
internal static byte[] DecompressMessagePayload(byte[] bytes)
5353
{
5454
using var stream = new MemoryStream(bytes);
5555

5656
// The stream will never be empty. It will at least contain the compression algo.
5757
var compression = (CompressionAlgos)stream.ReadByte();
58-
// Conditionally decompress and decode.
5958
Stream decompressedStream = compression switch
6059
{
6160
CompressionAlgos.None => stream,
@@ -67,11 +66,31 @@ internal static ServerMessage DecompressDecodeMessage(byte[] bytes)
6766
// TODO: consider pooling these.
6867
// DO NOT TRY TO TAKE THIS OUT. The BrotliStream ReadByte() implementation allocates an array
6968
// PER BYTE READ. You have to do it all at once to avoid that problem.
70-
MemoryStream memoryStream = new MemoryStream();
69+
using var memoryStream = new MemoryStream();
7170
decompressedStream.CopyTo(memoryStream);
72-
memoryStream.Seek(0, SeekOrigin.Begin);
73-
return new ServerMessage.BSATN().Read(new BinaryReader(memoryStream));
71+
return memoryStream.ToArray();
7472
}
73+
/// <summary>
74+
/// Decodes a serialized <see cref="ServerMessage"/> from a byte array.
75+
/// </summary>
76+
/// <param name="bytes">The encoded server message as a byte array.</param>
77+
/// <returns>The deserialized <see cref="ServerMessage"/> object.</returns>
78+
internal static ServerMessage DecodeServerMessage(byte[] bytes)
79+
{
80+
using var stream = new MemoryStream(bytes);
81+
using var reader = new BinaryReader(stream);
82+
return new ServerMessage.BSATN().Read(reader);
83+
}
84+
/// <summary>
85+
/// Decompresses and decodes a serialized <see cref="ServerMessage"/> from a byte array,
86+
/// automatically handling the specified compression algorithm (None, Brotli, or Gzip).
87+
/// Ensures efficient decompression by reading the entire stream at once to avoid
88+
/// performance issues with certain stream implementations.
89+
/// Throws <see cref="InvalidOperationException"/> if an unknown compression type is encountered.
90+
/// </summary>
91+
/// <param name="bytes">The compressed and encoded server message as a byte array.</param>
92+
/// <returns>The deserialized <see cref="ServerMessage"/> object.</returns>
93+
internal static ServerMessage DecompressDecodeMessage(byte[] bytes) => DecodeServerMessage(DecompressMessagePayload(bytes));
7594

7695
/// <summary>
7796
/// Prepare to read a BsatnRowList.

sdks/csharp/src/Plugins/WebSocket.jslib

Lines changed: 21 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,9 @@ mergeInto(LibraryManager.library, {
3333
var host = UTF8ToString(baseUriPtr);
3434
var uri = UTF8ToString(uriPtr);
3535
var protocol = UTF8ToString(protocolPtr);
36+
// The C# WebGL bridge can only pass one string argument here, so
37+
// multiple offered subprotocols are marshalled as a comma-separated string.
38+
var offeredProtocols = protocol.indexOf(',') === -1 ? protocol : protocol.split(',');
3639
var authToken = UTF8ToString(authTokenPtr);
3740
if (authToken)
3841
{
@@ -55,15 +58,31 @@ mergeInto(LibraryManager.library, {
5558
}
5659
}
5760

58-
var socket = new window.WebSocket(uri, protocol);
61+
var socket = new window.WebSocket(uri, offeredProtocols);
5962
socket.binaryType = "arraybuffer";
6063

6164
var socketId = manager.nextId++;
6265
manager.instances[socketId] = socket;
6366

6467
socket.onopen = function() {
6568
if (manager.callbacks.open) {
66-
WebSocketDynCall('vi', manager.callbacks.open, [socketId]);
69+
var protocolStr = socket.protocol || "";
70+
// Marshal the negotiated subprotocol to C# just for the duration of
71+
// this callback. We use stack allocation because the pointer only
72+
// needs to remain valid while dynCall is executing synchronously.
73+
var protocolLength = lengthBytesUTF8(protocolStr) + 1;
74+
var stack = stackSave();
75+
try {
76+
var protocolPtr = stackAlloc(protocolLength);
77+
// Write a temporary null-terminated UTF-8 string into the
78+
// Emscripten stack frame so the C# callback can copy it.
79+
stringToUTF8(protocolStr, protocolPtr, protocolLength);
80+
WebSocketDynCall('vii', manager.callbacks.open, [socketId, protocolPtr]);
81+
} finally {
82+
// Release the temporary stack allocation immediately after
83+
// the callback returns; C# must not retain the pointer.
84+
stackRestore(stack);
85+
}
6786
}
6887
};
6988

sdks/csharp/src/SpacetimeDBClient.cs

Lines changed: 32 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -169,6 +169,8 @@ public abstract class DbConnectionBase<DbConnection, Tables, Reducer> : IDbConne
169169
protected abstract IErrorContext ToErrorContext(Exception errorContext);
170170
protected abstract IProcedureEventContext ToProcedureEventContext(ProcedureEvent procedureEvent);
171171

172+
private Func<byte[], byte[][]> decodeTransportMessages = DecodeV2TransportMessages;
173+
172174
private readonly ConcurrentDictionary<uint, TaskCompletionSource<OneOffQueryResult>> waitingOneOffQueries = new();
173175

174176
private readonly ConcurrentDictionary<uint, PendingReducerCall> pendingReducerCalls = new();
@@ -220,10 +222,16 @@ protected DbConnectionBase()
220222
{
221223
var options = new WebSocket.ConnectOptions
222224
{
223-
Protocol = "v2.bsatn.spacetimedb"
225+
Protocols = WebSocketProtocols.Preferred
224226
};
225227
webSocket = new WebSocket(options);
226228
webSocket.OnMessage += OnMessageReceived;
229+
webSocket.OnProtocolNegotiated += protocolVersion =>
230+
{
231+
decodeTransportMessages = protocolVersion == WebSocketProtocolVersion.V3
232+
? WebSocketV3Payload.DecodeServerMessages
233+
: DecodeV2TransportMessages;
234+
};
227235
webSocket.OnSendError += a => onSendError?.Invoke(a);
228236
#if UNITY_5_3_OR_NEWER
229237
webSocket.OnClose += (e) =>
@@ -289,6 +297,8 @@ internal struct ParsedMessage
289297

290298
private static readonly Status Committed = new Status.Committed(default);
291299

300+
private static byte[][] DecodeV2TransportMessages(byte[] payload) => new[] { payload };
301+
292302
/// <summary>
293303
/// Get a description of a message suitable for storing in the tracker metadata.
294304
/// </summary>
@@ -427,9 +437,18 @@ void ParseOneOffQuery(OneOffQueryResult resp)
427437
#endif
428438
try
429439
{
430-
var message = _parseQueue.Take(_parseCancellationToken);
431-
var parsedMessage = ParseMessage(message);
432-
_applyQueue.Add(parsedMessage, _parseCancellationToken);
440+
var unparsed = _parseQueue.Take(_parseCancellationToken);
441+
var payload = CompressionHelpers.DecompressMessagePayload(unparsed.bytes);
442+
var decodedMessages = decodeTransportMessages(payload);
443+
stats.ParseMessageQueueTracker.FinishTrackingRequest(
444+
unparsed.parseQueueTrackerId,
445+
$"type=ws_payload,count={decodedMessages.Length}"
446+
);
447+
foreach (var messageBytes in decodedMessages)
448+
{
449+
var parsedMessage = ParseMessage(messageBytes, unparsed.timestamp);
450+
_applyQueue.Add(parsedMessage, _parseCancellationToken);
451+
}
433452
}
434453
catch (OperationCanceledException)
435454
{
@@ -452,13 +471,12 @@ void ParseOneOffQuery(OneOffQueryResult resp)
452471
}
453472
}
454473

455-
ParsedMessage ParseMessage(UnparsedMessage unparsed)
474+
ParsedMessage ParseMessage(byte[] messageBytes, DateTime timestamp)
456475
{
457476
var dbOps = ParsedDatabaseUpdate.New();
458-
var message = CompressionHelpers.DecompressDecodeMessage(unparsed.bytes);
477+
var message = CompressionHelpers.DecodeServerMessage(messageBytes);
459478
var trackerMetadata = TrackerMetadataForMessage(message);
460479

461-
stats.ParseMessageQueueTracker.FinishTrackingRequest(unparsed.parseQueueTrackerId, trackerMetadata);
462480
var parseStart = DateTime.UtcNow;
463481

464482
ReducerEvent<Reducer>? reducerEvent = default;
@@ -469,11 +487,11 @@ ParsedMessage ParseMessage(UnparsedMessage unparsed)
469487
case ServerMessage.InitialConnection:
470488
break;
471489
case ServerMessage.SubscribeApplied(var subscribeApplied):
472-
stats.SubscriptionRequestTracker.FinishTrackingRequest(subscribeApplied.RequestId, unparsed.timestamp);
490+
stats.SubscriptionRequestTracker.FinishTrackingRequest(subscribeApplied.RequestId, timestamp);
473491
dbOps = ParseSubscribeRows(subscribeApplied.Rows);
474492
break;
475493
case ServerMessage.UnsubscribeApplied(var unsubscribeApplied):
476-
stats.SubscriptionRequestTracker.FinishTrackingRequest(unsubscribeApplied.RequestId, unparsed.timestamp);
494+
stats.SubscriptionRequestTracker.FinishTrackingRequest(unsubscribeApplied.RequestId, timestamp);
477495
if (unsubscribeApplied.Rows != null)
478496
{
479497
dbOps = ParseUnsubscribeRows(unsubscribeApplied.Rows);
@@ -482,7 +500,7 @@ ParsedMessage ParseMessage(UnparsedMessage unparsed)
482500
case ServerMessage.SubscriptionError(var subscriptionError):
483501
if (subscriptionError.RequestId.HasValue)
484502
{
485-
stats.SubscriptionRequestTracker.FinishTrackingRequest(subscriptionError.RequestId.Value, unparsed.timestamp);
503+
stats.SubscriptionRequestTracker.FinishTrackingRequest(subscriptionError.RequestId.Value, timestamp);
486504
}
487505
break;
488506
case ServerMessage.TransactionUpdate(var transactionUpdate):
@@ -492,7 +510,7 @@ ParsedMessage ParseMessage(UnparsedMessage unparsed)
492510
ParseOneOffQuery(resp);
493511
break;
494512
case ServerMessage.ReducerResult(var reducerResult):
495-
if (!stats.ReducerRequestTracker.FinishTrackingRequest(reducerResult.RequestId, unparsed.timestamp))
513+
if (!stats.ReducerRequestTracker.FinishTrackingRequest(reducerResult.RequestId, timestamp))
496514
{
497515
Log.Warn($"Failed to finish tracking reducer request: {reducerResult.RequestId}");
498516
}
@@ -545,7 +563,7 @@ ParsedMessage ParseMessage(UnparsedMessage unparsed)
545563
procedureResult.RequestId
546564
);
547565

548-
if (!stats.ProcedureRequestTracker.FinishTrackingRequest(procedureResult.RequestId, unparsed.timestamp))
566+
if (!stats.ProcedureRequestTracker.FinishTrackingRequest(procedureResult.RequestId, timestamp))
549567
{
550568
Log.Warn($"Failed to finish tracking procedure request: {procedureResult.RequestId}");
551569
}
@@ -558,7 +576,7 @@ ParsedMessage ParseMessage(UnparsedMessage unparsed)
558576
stats.ParseMessageTracker.InsertRequest(parseStart, trackerMetadata);
559577
var applyTracker = stats.ApplyMessageQueueTracker.StartTrackingRequest(trackerMetadata);
560578

561-
return new ParsedMessage { message = message, dbOps = dbOps, receiveTimestamp = unparsed.timestamp, applyQueueTrackerId = applyTracker, reducerEvent = reducerEvent, procedureEvent = procedureEvent };
579+
return new ParsedMessage { message = message, dbOps = dbOps, receiveTimestamp = timestamp, applyQueueTrackerId = applyTracker, reducerEvent = reducerEvent, procedureEvent = procedureEvent };
562580
}
563581
}
564582

@@ -609,6 +627,7 @@ void IDbConnection.Connect(string? token, string uri, string addressOrName, Comp
609627
{
610628
isClosing = false;
611629
connectionClosed = false;
630+
decodeTransportMessages = DecodeV2TransportMessages;
612631
Identity = null;
613632
initialConnectionId = null;
614633
onConnectInvoked = false;

0 commit comments

Comments
 (0)