diff --git a/app/src/include/firebase/internal/future_impl.h b/app/src/include/firebase/internal/future_impl.h index 59e7771d01..37941f11d9 100644 --- a/app/src/include/firebase/internal/future_impl.h +++ b/app/src/include/firebase/internal/future_impl.h @@ -152,21 +152,54 @@ class CompletionCallbackHandle { void (*user_data_delete_fn_)(void*); }; +template +struct TypedCompletionCallbackData { + typename Future::TypedCompletionCallback callback; + void* user_data; +}; + +template +inline void TypedCompletionCallbackTrampoline(const FutureBase& future, + void* data_ptr) { + auto* data = static_cast*>(data_ptr); + if (data != nullptr && data->callback != nullptr) { + data->callback(static_cast&>(future), data->user_data); + } +} + +template +inline void DeleteTypedCompletionCallbackData(void* data_ptr) { + delete static_cast*>(data_ptr); +} + } // namespace detail template void Future::OnCompletion(TypedCompletionCallback callback, void* user_data) const { - FutureBase::OnCompletion(reinterpret_cast(callback), - user_data); + MutexLock lock(mutex_); + if (api_ != nullptr) { + if (callback == nullptr) { + api_->AddCompletionCallback(handle_, nullptr, nullptr, nullptr, + /*clear_existing_callbacks=*/true); + } else { + auto* data = + new detail::TypedCompletionCallbackData{callback, user_data}; + api_->AddCompletionCallback( + handle_, detail::TypedCompletionCallbackTrampoline, data, + detail::DeleteTypedCompletionCallbackData, + /*clear_existing_callbacks=*/true); + } + } } #if defined(FIREBASE_USE_STD_FUNCTION) template inline void Future::OnCompletion( std::function&)> callback) const { - FutureBase::OnCompletion( - *reinterpret_cast*>(&callback)); + FutureBase::OnCompletion([callback](const FutureBase& future) { + callback(static_cast&>(future)); + }); } #endif // defined(FIREBASE_USE_STD_FUNCTION) @@ -174,16 +207,28 @@ inline void Future::OnCompletion( template FutureBase::CompletionCallbackHandle Future::AddOnCompletion( TypedCompletionCallback callback, void* user_data) const { - return FutureBase::AddOnCompletion( - reinterpret_cast(callback), user_data); + MutexLock lock(mutex_); + if (api_ != nullptr) { + if (callback == nullptr) { + return CompletionCallbackHandle(); + } + auto* data = + new detail::TypedCompletionCallbackData{callback, user_data}; + return api_->AddCompletionCallback( + handle_, detail::TypedCompletionCallbackTrampoline, data, + detail::DeleteTypedCompletionCallbackData, + /*clear_existing_callbacks=*/false); + } + return CompletionCallbackHandle(); } #if defined(FIREBASE_USE_STD_FUNCTION) template inline FutureBase::CompletionCallbackHandle Future::AddOnCompletion( std::function&)> callback) const { - return FutureBase::AddOnCompletion( - *reinterpret_cast*>(&callback)); + return FutureBase::AddOnCompletion([callback](const FutureBase& future) { + callback(static_cast&>(future)); + }); } #endif // defined(FIREBASE_USE_STD_FUNCTION)