package com.suno.android.common_data.repos.orpheus import com.suno.android.common_core_utils.DispatcherIO import com.suno.android.common_core_utils.Id import com.suno.android.common_core_utils.SunoLogger import com.suno.android.common_data.repos.orpheus.models.OrpheusSession import com.suno.android.common_networking.orpheus.OrpheusRealtimeClient import com.suno.android.common_networking.orpheus.OrpheusRealtimeEvent import com.suno.android.common_networking.sse.RealtimeEvent import kotlinx.coroutines.CoroutineDispatcher import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.Job import kotlinx.coroutines.SupervisorJob import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.MutableSharedFlow import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.flow.StateFlow import kotlinx.coroutines.flow.asSharedFlow import kotlinx.coroutines.flow.asStateFlow import kotlinx.coroutines.flow.catch import kotlinx.coroutines.flow.launchIn import kotlinx.coroutines.flow.onEach import javax.inject.Inject import javax.inject.Singleton @Singleton class OrpheusConnectionManager @Inject constructor( loggerFactory: SunoLogger.Factory, private val realtimeClient: OrpheusRealtimeClient, @DispatcherIO private val dispatcher: CoroutineDispatcher, ) { private val logger = loggerFactory.create(this@OrpheusConnectionManager) private val scope = CoroutineScope(SupervisorJob() + dispatcher) private var subscriptionJob: Job? = null private val _realtimeEvents = MutableSharedFlow(extraBufferCapacity = 64) val realtimeEvents: Flow = _realtimeEvents.asSharedFlow() private val _connectionState = MutableStateFlow(OrpheusConnectionState.Disconnected) val connectionState: StateFlow = _connectionState.asStateFlow() fun connect( sessionId: Id, ) { disconnect() _connectionState.value = OrpheusConnectionState.Connecting(sessionId) subscriptionJob = realtimeClient.subscribe(sessionId.value) .onEach { event -> _realtimeEvents.emit(event) updateConnectionState(event) } .catch { error -> logger.e(error) { "Error in realtime subscription" } } .launchIn(scope) logger.d { "Started realtime connection for session: ${sessionId.value}" } } fun disconnect() { subscriptionJob?.cancel() subscriptionJob = null _connectionState.value = OrpheusConnectionState.Disconnected logger.d { "Stopped realtime connection" } } private fun updateConnectionState( event: OrpheusRealtimeEvent, ) { when (event) { is RealtimeEvent.Connected -> { val connectingState = _connectionState.value as? OrpheusConnectionState.Connecting if (connectingState != null) { _connectionState.value = OrpheusConnectionState.Connected(connectingState.id) } } is RealtimeEvent.Disconnected -> { _connectionState.value = OrpheusConnectionState.Disconnected } is RealtimeEvent.Error -> { _connectionState.value = OrpheusConnectionState.Error } is RealtimeEvent.Data -> { // Data events don't affect connection state } } } } sealed class OrpheusConnectionState { data object Disconnected : OrpheusConnectionState() data class Connecting( val id: Id, ) : OrpheusConnectionState() data class Connected( val id: Id, ) : OrpheusConnectionState() data object Error : OrpheusConnectionState() val isTerminal: Boolean get() = this is Connected || this is Error || this is Disconnected val sessionId: Id? get() = when (this) { is Connecting -> id is Connected -> id else -> null } }