diff --git a/core/java/android/window/ImeOnBackInvokedDispatcher.java b/core/java/android/window/ImeOnBackInvokedDispatcher.java index 34b75a4788c40..9ef68807419a9 100644 --- a/core/java/android/window/ImeOnBackInvokedDispatcher.java +++ b/core/java/android/window/ImeOnBackInvokedDispatcher.java @@ -82,8 +82,13 @@ public class ImeOnBackInvokedDispatcher implements OnBackInvokedDispatcher, Parc @NonNull OnBackInvokedCallback callback) { final Bundle bundle = new Bundle(); // Always invoke back for ime without checking the window focus. + // We use strong reference in the binder wrapper to avoid accidentally GC the callback. + // This is necessary because the callback is sent to and registered from + // the app process, which may treat the IME callback as weakly referenced. This will not + // cause a memory leak because the app side already clears the reference correctly. final IOnBackInvokedCallback iCallback = - new WindowOnBackInvokedDispatcher.OnBackInvokedCallbackWrapper(callback); + new WindowOnBackInvokedDispatcher.OnBackInvokedCallbackWrapper( + callback, false /* useWeakRef */); bundle.putBinder(RESULT_KEY_CALLBACK, iCallback.asBinder()); bundle.putInt(RESULT_KEY_PRIORITY, priority); bundle.putInt(RESULT_KEY_ID, callback.hashCode()); diff --git a/core/java/android/window/WindowOnBackInvokedDispatcher.java b/core/java/android/window/WindowOnBackInvokedDispatcher.java index d34ece9027782..a8c2b2f28df4a 100644 --- a/core/java/android/window/WindowOnBackInvokedDispatcher.java +++ b/core/java/android/window/WindowOnBackInvokedDispatcher.java @@ -254,10 +254,34 @@ public class WindowOnBackInvokedDispatcher implements OnBackInvokedDispatcher { } static class OnBackInvokedCallbackWrapper extends IOnBackInvokedCallback.Stub { - private final WeakReference mCallback; + static class CallbackRef { + final WeakReference mWeakRef; + final OnBackInvokedCallback mStrongRef; + CallbackRef(@NonNull OnBackInvokedCallback callback, boolean useWeakRef) { + if (useWeakRef) { + mWeakRef = new WeakReference<>(callback); + mStrongRef = null; + } else { + mStrongRef = callback; + mWeakRef = null; + } + } + + OnBackInvokedCallback get() { + if (mStrongRef != null) { + return mStrongRef; + } + return mWeakRef.get(); + } + } + final CallbackRef mCallbackRef; OnBackInvokedCallbackWrapper(@NonNull OnBackInvokedCallback callback) { - mCallback = new WeakReference<>(callback); + mCallbackRef = new CallbackRef(callback, true /* useWeakRef */); + } + + OnBackInvokedCallbackWrapper(@NonNull OnBackInvokedCallback callback, boolean useWeakRef) { + mCallbackRef = new CallbackRef(callback, useWeakRef); } @Override @@ -300,8 +324,9 @@ public class WindowOnBackInvokedDispatcher implements OnBackInvokedDispatcher { public void onBackInvoked() throws RemoteException { Handler.getMain().post(() -> { mProgressAnimator.reset(); - final OnBackInvokedCallback callback = mCallback.get(); + final OnBackInvokedCallback callback = mCallbackRef.get(); if (callback == null) { + Log.d(TAG, "Trying to call onBackInvoked() on a null callback reference."); return; } callback.onBackInvoked(); @@ -310,7 +335,7 @@ public class WindowOnBackInvokedDispatcher implements OnBackInvokedDispatcher { @Nullable private OnBackAnimationCallback getBackAnimationCallback() { - OnBackInvokedCallback callback = mCallback.get(); + OnBackInvokedCallback callback = mCallbackRef.get(); return callback instanceof OnBackAnimationCallback ? (OnBackAnimationCallback) callback : null; }