Kotlin Flow tips for ViewModel

Many months have passed since I last wrote about Kotlin Flow: Domain Model’s StateFlow Sharing. So, let’s take a look at something related to Flow handling in a ViewModel.

StateIn shortcut

Usually, you want a StateFlow‘s state to survive a configuration change. AndroidX’s default time limit is 5 seconds — but only if it’s subscribed. Since we don’t have context parameters yet, add this to your BaseViewModel.

class BaseViewModel : ViewModel() {
	fun <T> Flow<T>.stateIn(initialValue: T): StateFlow<T> =
		stateIn(viewModelScope, SharingStarted.WhileSubscribed(5.seconds), initialValue)
}

class MyViewModel(...) : BaseViewModel() {
	val data: StateFlow<...> = repository.getData().stateIn(null)
}

State management

Your UI (and consequently, your ViewModel) has its own state. It’s tempting to model it with MutableStateFlow(). But… this state is lost on process death—meaning no state restoration for you.

To handle this, you need to inject SavedStateHandle and properly derive the state from it. Doing this manually is both error-prone and verbose.

Ideally, we want something similar to the saved extension on SavedStateHandle, but with MutableStateFlow support. The resulting usage should look like this:

class MyViewModel(
	state: SavedStateHandle,
) {
	val selectedAge: MutableStateFlow<Int?> by state.savedStateFlow { null }
	val selectedGender: MutableStateFlow<Gender> by state.savedStateFlow { Gender.Female }
}

Until recently, implementing support for non-primitive types was sub-optimal, as there was no support KotlinX.Serialization. But those days are over with lifecycle-viewmodel 2.9.0 (currently in alpha). Now, you can easily handle this using decodeFromSavedState and encodeToSavedState functions MutableStateFlowSerializer.

Edit: My original manual approach solution can be simplified using the new built-in MutableStateFlowSerializer. Thanks to Ian Lake for highlighting this. The original implementation is collapesed bellow the new one.

protected inline fun <reified T> SavedStateHandle.savedStateFlow(
	serializer: KSerializer<T> = serializer<T>(),
	key: String? = null,
	crossinline init: () -> T,
) = saved(
	serializer = MutableStateFlowSerializer(serializer),
	key = key,
	init = { MutableStateFlow(init()) },
)
The original implementation
@JvmName("savedMutableState")
protected inline fun <reified T : Any> SavedStateHandle.saved(
	key: String? = null,
	noinline init: () -> T,
): ReadOnlyProperty<ViewModel, MutableStateFlow<T>> =
	NotNullableMutableStateFlorPropertyDelegate(this, key, serializer(), init)

@JvmName("savedNullableMutableState")
@Suppress("unused")
protected inline fun <W : MutableStateFlow<T?>, reified T : Any> SavedStateHandle.savedNullable(
	key: String? = null,
	noinline init: () -> T? = { null },
): ReadOnlyProperty<ViewModel, MutableStateFlow<T?>> =
	NullableMutableStateFlorPropertyDelegate(this, key, serializer<T>(), init)

protected class NotNullableMutableStateFlorPropertyDelegate<T : Any>(
	private val savedStateHandle: SavedStateHandle,
	private val key: String?,
	private val serializer: KSerializer<T>,
	private val init: () -> T,
) : ReadOnlyProperty<ViewModel, MutableStateFlow<T>> {
	private var stateFlow: MutableStateFlow<T>? = null

	override fun getValue(thisRef: ViewModel, property: KProperty<*>): MutableStateFlow<T> {
		this.stateFlow?.let { return it }
		val qualifiedKey = key ?: (thisRef::class.qualifiedName + "." + property.name)
		val initialState = savedStateHandle.get<SavedState>(qualifiedKey)
		val initialValue = initialState?.let { decodeFromSavedState(serializer, initialState) } ?: init()
		val stateFlow = MutableStateFlow(initialValue)
		savedStateHandle.setSavedStateProvider(qualifiedKey) { encodeToSavedState(serializer, stateFlow.value) }
		return stateFlow.also { this.stateFlow = it }
	}
}

protected class NullableMutableStateFlorPropertyDelegate<T : Any?>(
	private val savedStateHandle: SavedStateHandle,
	private val key: String?,
	private val serializer: KSerializer<T & Any>,
	private val init: () -> T?,
) : ReadOnlyProperty<ViewModel, MutableStateFlow<T?>> {
	private var stateFlow: MutableStateFlow<T?>? = null

	override fun getValue(thisRef: ViewModel, property: KProperty<*>): MutableStateFlow<T?> {
		this.stateFlow?.let { return it }

		val qualifiedKey = key ?: (thisRef::class.qualifiedName + "." + property.name)
		val initialState = savedStateHandle.get<SavedState>(qualifiedKey)
		val initialValue = if (initialState != null) {
			if (initialState.isEmpty) null else decodeFromSavedState(serializer, initialState)
		} else {
			init()
		}
		val stateFlow = MutableStateFlow(initialValue)
		savedStateHandle.setSavedStateProvider(qualifiedKey) {
			val value = stateFlow.value
			if (value == null) savedState() else encodeToSavedState(serializer, value)
		}
		return stateFlow.also { this.stateFlow = it }
	}
}

The final implementation is a bit longer, mainly to tackle all the edge-cases and nullability.

Derived State

To effectively & reactively drive your UI layer, use combine() function. The input state may be local (the UI state) or global (injected via dependency injection). Next, combine the state to perform actions—such as fetching data from an API.

val selectedAge: MutableStateFlow<Int?> by state.savedStateFlow { null }
val selectedGender: MutableStateFlow<Gender> by state.savedStateFlow { Gender.Female }

val result: StateFlow<Data?> = 
	combine(
		selectedAge,
		selectedGender,
	) { age, gender ->
		repository.fetch(age, gender)
	}.stateIn(null)	

You can see that the example already uses the stateIn() helper we defined earlier.

Since the repository takes arguments from the reactive streams in the same order they are passed to the combine method, we can shorten our code even further:

val result: StateFlow<Data?> =
	combine(selectedAge, selectedGender, repository::fetch).stateIn(null)

If the repository returns a Flow stream of data, you can easily use flatMapLatest or combineTransform.

val result: StateFlow<Data?> =
	combine(selectedAge, selectedGender, repository::fetch)
		.flatMapLatest { it }
		.stateIn(null)

Refreshing

The previous section didn’t introduce any new helpers, but it introduced a crucial concept: The ViewModel’s behavior is fully reactive, driven by actual inputs. No need to worry about synchronization issues. But—as always—there’s a catch.

How Do We Model a Refresh Request? Since combine is based on state, how do we force a refresh when the user, for example, performs a pull-to-refresh gesture? This kind of user interaction isn’t state (StateFlow), but an event. Events must be modeled using SharedFlow.

Let’s add one: every time the user performs the gesture, we’ll emit a new Unit value into the event flow.

val refreshEvents = MutableSharedFlow<Unit>(extraBufferCapacity = 1)

To properly force an API refetch, you can use a “resubscription” trick. When a Unit is observed, the code will simply stop collecting and start collecting the upstream again. This works because [email protected] is stopped when collectLatest receives a new value—since collectLatest cancels its block whenever a new value arrives.

fun <T> Flow<T>.refreshOn(refreshSignal: Flow<Unit>): Flow<T> =
	channelFlow {
		refreshSignal
			.onStart { emit(Unit) }
			.collectLatest { [email protected] { send(it) } }
	}

Using this new helper is straightforward.

val result: StateFlow<Data?> =
	combine(selectedAge, selectedGender, repository::fetch)
		.flatMapLatest { it }
		.refreshOn(refreshEvents)
		.stateIn(null)

Conclusion

The goal may not be obvious at first, but we aimed to make our code concise, reactive, and free from race conditions. Most importantly, we wanted to preserve state across process death. This last requirement is often overlooked by Android developers—partly because modern high-end devices make process death less noticeable in daily use.

Who wouldn’t like this concise, fully reactive view model?

class MyViewModel(
	state: SavedStateHandle,
	defaultDispatcher: CoroutineContext,
	private val repository: DataRepository,
) : BaseViewModel() {
	private val refreshEvents = MutableSharedFlow<Unit>(extraBufferCapacity = 1)
	private val selectedAge: MutableStateFlow<Int?> by state.savedStateFlow { null }
	private val selectedGender: MutableStateFlow<Gender> by state.savedStateFlow { Gender.Female }

	val age: StateFlow<Int?> = selectedAge.asStateFlow()
	val gender: StateFlow<Gender> = selectedGender.asStateFlow()
	val result: StateFlow<Data?> =
		combine(selectedAge, selectedGender, repository::fetch)
			.flatMapLatest { it }
			.refreshOn(refreshEvents)
			.flowOn(defaultDispatcher)
			.stateIn(null)

	fun refresh() { refreshEvents.tryEmit(Unit) }
	fun onAge(age: Int) { selectedAge.value = age }
	fun onGender(gender: Gender) { selectedGender.value = gender }
}

Edit: I am publishing unit tests I made for this helpers: https://gist.github.com/hrach/ec3b3d118afd35efc6a03cfd3a187c46