package com.suno.android.common_networking.sse import com.suno.android.common_core_utils.DispatcherIO import com.suno.android.common_core_utils.SunoLogger import com.suno.android.common_core_utils.clock.Clock import com.suno.android.common_networking.di.Sse import kotlinx.coroutines.CancellationException import kotlinx.coroutines.CoroutineDispatcher import kotlinx.coroutines.currentCoroutineContext import kotlinx.coroutines.delay import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.flow import kotlinx.coroutines.flow.flowOn import kotlinx.coroutines.isActive import kotlinx.coroutines.withTimeoutOrNull import okhttp3.HttpUrl.Companion.toHttpUrl import okhttp3.OkHttpClient import okhttp3.Request import java.io.BufferedReader import java.io.IOException import javax.inject.Inject import javax.inject.Singleton import kotlin.math.min import kotlin.math.pow import kotlin.time.DurationUnit import kotlin.time.toDuration private const val TOKEN_REFRESH_BUFFER_MS = 30000L // 30 seconds /** * Generic SSE client implementing the WHATWG Server-Sent Events specification. * https://html.spec.whatwg.org/multipage/server-sent-events.html */ @Singleton class SseClient @Inject constructor( loggerFactory: SunoLogger.Factory, @Sse private val sseHttpClient: OkHttpClient, private val lineParser: SseLineParser, private val clock: Clock, @DispatcherIO private val dispatcherIO: CoroutineDispatcher, ) { private val logger = loggerFactory.create(this@SseClient) /** * Subscribes to an SSE stream with automatic reconnection. */ fun subscribe( config: SseConnectionConfig, authTokenProvider: suspend () -> Result, ): Flow = flow { var retryAttempt = 0 var lastEventId: String? = null while (currentCoroutineContext().isActive) { val connectionResult = connect( config = config, lastEventId = lastEventId, authTokenProvider = authTokenProvider, onEvent = { event -> if (event is SseEvent.Message) { event.id?.let { lastEventId = it } } emit(event) }, ) when (connectionResult) { is ConnectionResult.StreamEnded -> { emit(SseEvent.Disconnected) retryAttempt = 0 } is ConnectionResult.Failure -> { emit(SseEvent.Error(connectionResult.throwable)) if (retryAttempt >= config.maxRetries) { logger.w { "Max retry attempts (${config.maxRetries}) exceeded, stopping reconnection" } break } val delayMs = calculateRetryDelay( attempt = retryAttempt, config = config, ) delay(delayMs) retryAttempt++ } } } }.flowOn(dispatcherIO) private suspend fun connect( config: SseConnectionConfig, lastEventId: String?, authTokenProvider: suspend () -> Result, onEvent: suspend (SseEvent) -> Unit, ): ConnectionResult = runCatching { val authToken = authTokenProvider().getOrElse { error -> throw IOException("Failed to fetch auth token: $error") } val request = createRequest( config = config, lastEventId = lastEventId, authToken = authToken, ) sseHttpClient.newCall(request).execute().use { response -> if (!response.isSuccessful) { throw IOException("SSE connection failed: ${response.code} ${response.message}") } onEvent(SseEvent.Connected) // Long-lived connection with optional timeout val connectionTimeoutMs = calculateConnectionTimeout(authToken) if (connectionTimeoutMs != null && connectionTimeoutMs <= 0L) { throw IOException("Token expired or expiring too soon") } else { logger.d { val timeoutMinutes = connectionTimeoutMs?.toDuration(DurationUnit.MILLISECONDS)?.inWholeMinutes "Connection will timeout in $timeoutMinutes minutes to refresh token" } } val streamResult = if (connectionTimeoutMs != null) { withTimeoutOrNull(timeMillis = connectionTimeoutMs) { response.body?.byteStream()?.bufferedReader()?.use { reader -> processSseStream( reader = reader, config = config, onEvent = onEvent, ) } } } else { response.body?.byteStream()?.bufferedReader()?.use { reader -> processSseStream( reader = reader, config = config, onEvent = onEvent, ) } } if (streamResult == null) { logger.d { "Connection timed out, reconnecting" } } ConnectionResult.StreamEnded } }.getOrElse { e -> when (e) { is CancellationException -> { logger.d { "Connection cancelled" } throw e // Re-throw to properly propagate cancellation } else -> { logger.e(e) { "Connection attempt failed" } ConnectionResult.Failure(e) } } } private fun createRequest( config: SseConnectionConfig, lastEventId: String?, authToken: SseAuthToken, ): Request { val urlBuilder = config.url.toHttpUrl().newBuilder() config.queryParameters.forEach { (key, value) -> urlBuilder.addQueryParameter(name = key, value = value) } // Add lastEventId as query parameter for reconnection lastEventId?.let { eventId -> urlBuilder.addQueryParameter(name = "lastEventId", value = eventId) logger.d { "Reconnecting with lastEventId: $eventId" } } return Request.Builder() .url(urlBuilder.build()) .header("Accept", "text/event-stream") .header("Cache-Control", "no-cache") .header("Authorization", authToken.authorizationHeader) .build() } private fun calculateConnectionTimeout( authToken: SseAuthToken, ): Long? { val tokenBasedTimeout = authToken.expiresAtMs?.let { expiresAtMs -> val timeUntilExpiry = expiresAtMs - clock.currentTime.inWholeMilliseconds timeUntilExpiry - TOKEN_REFRESH_BUFFER_MS } return tokenBasedTimeout } private suspend fun processSseStream( reader: BufferedReader, config: SseConnectionConfig, onEvent: suspend (SseEvent) -> Unit, ) { var currentEventId: String? = null val dataAccumulator = StringBuilder() while (currentCoroutineContext().isActive) { val line = reader.readLine() ?: break when (val parsedLine = lineParser.parseLine(line)) { is SseLineParser.ParsedLine.Id -> { logger.d { "Received and parsed line: $parsedLine" } currentEventId = parsedLine.id } is SseLineParser.ParsedLine.Data -> { logger.d { "Received and parsed line: $parsedLine" } // Prevent memory leak from malformed SSE streams if (dataAccumulator.length > config.maxEventSizeBytes) { logger.w { "SSE event exceeded max size (${dataAccumulator.length} bytes), discarding incomplete event" } currentEventId = null dataAccumulator.clear() continue } // Per SSE spec: multiple data lines are joined with newlines if (dataAccumulator.isNotEmpty()) { dataAccumulator.append('\n') } dataAccumulator.append(parsedLine.data) } is SseLineParser.ParsedLine.EndOfEvent -> { if (dataAccumulator.isNotEmpty()) { logger.d { "Received end of event: $dataAccumulator" } val event = SseEvent.Message( data = dataAccumulator.toString(), id = currentEventId, ) onEvent(event) dataAccumulator.clear() } // Reset for next event currentEventId = null } else -> { logger.d { "Received and parsed $parsedLine, ignoring" } } } } } private fun calculateRetryDelay( attempt: Int, config: SseConnectionConfig, ): Long { val cappedAttempt = min(attempt, config.maxRetryAttemptForBackoff) val exponentialDelay = config.initialRetryDelayMs * config.retryBackoffMultiplier.pow(cappedAttempt).toLong() return min(exponentialDelay, config.maxRetryDelayMs) } } private sealed class ConnectionResult { data object StreamEnded : ConnectionResult() data class Failure( val throwable: Throwable, ) : ConnectionResult() }