using System; using System.ClientModel; using System.ClientModel.Primitives; using System.Collections.Generic; using System.Net.ServerSentEvents; using System.Runtime.CompilerServices; using System.Text.Json; using System.Threading; using System.Threading.Tasks; #nullable enable namespace OpenAI; /// /// Implementation of collection abstraction over streaming chat updates. /// internal class AsyncSseUpdateCollection : AsyncCollectionResult { private readonly Func> _sendRequestAsync; private readonly Func, IEnumerable> _eventDeserializerFunc; private readonly CancellationToken _cancellationToken; public List AdditionalDisposalActions { get; } = []; public AsyncSseUpdateCollection( Func> sendRequestAsync, Func> jsonMultiDeserializerFunc, CancellationToken cancellationToken) : this( sendRequestAsync, DeserializeSseToMultipleViaJson(jsonMultiDeserializerFunc), cancellationToken) { Argument.AssertNotNull(jsonMultiDeserializerFunc, nameof(jsonMultiDeserializerFunc)); } public AsyncSseUpdateCollection( Func> sendRequestAsync, Func jsonSingleDeserializerFunc, CancellationToken cancellationToken) : this( sendRequestAsync, DeserializeSseToSingleViaJson(jsonSingleDeserializerFunc), cancellationToken) { Argument.AssertNotNull(jsonSingleDeserializerFunc, nameof(jsonSingleDeserializerFunc)); } public AsyncSseUpdateCollection( Func> sendRequestAsync, Func> jsonMultiDeserializerFunc, CancellationToken cancellationToken) : this( sendRequestAsync, DeserializeSseToMultipleViaJson(jsonMultiDeserializerFunc), cancellationToken) { Argument.AssertNotNull(jsonMultiDeserializerFunc, nameof(jsonMultiDeserializerFunc)); } public AsyncSseUpdateCollection( Func> sendRequestAsync, Func jsonSingleDeserializerFunc, CancellationToken cancellationToken) : this( sendRequestAsync, DeserializeSseToSingleViaJson(jsonSingleDeserializerFunc), cancellationToken) { Argument.AssertNotNull(jsonSingleDeserializerFunc, nameof(jsonSingleDeserializerFunc)); } public AsyncSseUpdateCollection( Func> sendRequestAsync, Func, IEnumerable> eventDeserializerFunc, CancellationToken cancellationToken) { Argument.AssertNotNull(sendRequestAsync, nameof(sendRequestAsync)); Argument.AssertNotNull(eventDeserializerFunc, nameof(eventDeserializerFunc)); _sendRequestAsync = sendRequestAsync; _eventDeserializerFunc = eventDeserializerFunc; _cancellationToken = cancellationToken; } public override ContinuationToken? GetContinuationToken(ClientResult page) // Continuation is not supported for SSE streams. => null; public async override IAsyncEnumerable GetRawPagesAsync() { // We don't currently support resuming a dropped connection from the // last received event, so the response collection has a single element. yield return await _sendRequestAsync(); } protected async override IAsyncEnumerable GetValuesFromPageAsync(ClientResult page) { await using IAsyncEnumerator enumerator = new AsyncSseUpdateEnumerator(_eventDeserializerFunc, page, _cancellationToken, AdditionalDisposalActions); while (await enumerator.MoveNextAsync().ConfigureAwait(false)) { yield return enumerator.Current; } } [MethodImpl(MethodImplOptions.AggressiveInlining)] internal static Func, IEnumerable> DeserializeSseToMultipleViaJson( Func> jsonDeserializationFunc) { return (item) => { using JsonDocument document = JsonDocument.Parse(item.Data); return jsonDeserializationFunc.Invoke(document.RootElement, ModelSerializationExtensions.WireOptions); }; } [MethodImpl(MethodImplOptions.AggressiveInlining)] internal static Func, IEnumerable> DeserializeSseToSingleViaJson( Func jsonSingleDeserializationFunc) => DeserializeSseToMultipleViaJson((e, o) => [jsonSingleDeserializationFunc.Invoke(e, o)]); [MethodImpl(MethodImplOptions.AggressiveInlining)] internal static Func, IEnumerable> DeserializeSseToMultipleViaJson( Func> jsonDeserializationFunc) { return (item) => { using JsonDocument document = JsonDocument.Parse(item.Data); return jsonDeserializationFunc.Invoke(document.RootElement, BinaryData.FromBytes(item.Data), ModelSerializationExtensions.WireOptions); }; } [MethodImpl(MethodImplOptions.AggressiveInlining)] internal static Func, IEnumerable> DeserializeSseToSingleViaJson( Func jsonSingleDeserializationFunc) => DeserializeSseToMultipleViaJson((e, d, o) => [jsonSingleDeserializationFunc.Invoke(e, d, o)]); private sealed class AsyncSseUpdateEnumerator : IAsyncEnumerator { private static ReadOnlySpan TerminalData => "[DONE]"u8; private List _additionalDisposalActions; private readonly CancellationToken _cancellationToken; private readonly PipelineResponse _response; // These enumerators represent what is effectively a doubly-nested // loop over the outer event collection and the inner update collection, // i.e.: // foreach (var sse in _events) { // // get _updates from sse event // foreach (var update in _updates) { ... } // } private IAsyncEnumerator>? _events; private IEnumerator? _updates; private readonly Func, IEnumerable> _deserializerFunc; private U? _current; private bool _started; public AsyncSseUpdateEnumerator( Func, IEnumerable> deserializerFunc, ClientResult page, CancellationToken cancellationToken, List additionalDisposalActions) { Argument.AssertNotNull(page, nameof(page)); _deserializerFunc = deserializerFunc; _response = page.GetRawResponse(); _cancellationToken = cancellationToken; _additionalDisposalActions = additionalDisposalActions; } U IAsyncEnumerator.Current => _current!; async ValueTask IAsyncEnumerator.MoveNextAsync() { if (_events is null && _started) { throw new ObjectDisposedException(nameof(AsyncSseUpdateEnumerator)); } _cancellationToken.ThrowIfCancellationRequested(); _events ??= CreateEventEnumeratorAsync(); _started = true; if (_updates is not null && _updates.MoveNext()) { _current = _updates.Current; return true; } if (await _events.MoveNextAsync().ConfigureAwait(false)) { if (_events.Current.Data.AsSpan().SequenceEqual(TerminalData)) { _current = default; return false; } _updates = _deserializerFunc .Invoke(_events.Current) .GetEnumerator(); if (_updates.MoveNext()) { _current = _updates.Current; return true; } } _current = default; return false; } private IAsyncEnumerator> CreateEventEnumeratorAsync() { if (_response.ContentStream is null) { throw new InvalidOperationException("Unable to create result from response with null ContentStream"); } IAsyncEnumerable> enumerable = SseParser.Create(_response.ContentStream, (_, bytes) => bytes.ToArray()).EnumerateAsync(); return enumerable.GetAsyncEnumerator(_cancellationToken); } public async ValueTask DisposeAsync() { await DisposeAsyncCore().ConfigureAwait(false); GC.SuppressFinalize(this); } private async ValueTask DisposeAsyncCore() { if (_events is not null) { await _events.DisposeAsync().ConfigureAwait(false); _events = null; // Dispose the response so we don't leave the network connection open. _response?.Dispose(); } foreach (Action additionalDisposalAction in _additionalDisposalActions ?? []) { additionalDisposalAction.Invoke(); } _additionalDisposalActions?.Clear(); } } }