diff --git a/core/java/com/android/internal/app/AbstractResolverComparator.java b/core/java/com/android/internal/app/AbstractResolverComparator.java index 975954035c176..930f6e0c29111 100644 --- a/core/java/com/android/internal/app/AbstractResolverComparator.java +++ b/core/java/com/android/internal/app/AbstractResolverComparator.java @@ -30,11 +30,16 @@ import android.os.UserHandle; import android.util.Log; import com.android.internal.app.ResolverActivity.ResolvedComponentInfo; +import com.android.internal.app.chooser.TargetInfo; + +import com.google.android.collect.Lists; import java.text.Collator; import java.util.ArrayList; import java.util.Comparator; +import java.util.HashMap; import java.util.List; +import java.util.Map; /** * Used to sort resolved activities in {@link ResolverListController}. @@ -48,8 +53,8 @@ public abstract class AbstractResolverComparator implements Comparator mPmMap = new HashMap<>(); + protected final Map mUsmMap = new HashMap<>(); protected String[] mAnnotations; protected String mContentType; @@ -98,14 +103,28 @@ public abstract class AbstractResolverComparator implements Comparator targetUserSpaceList) { String scheme = intent.getScheme(); mHttp = "http".equals(scheme) || "https".equals(scheme); mContentType = intent.getType(); getContentAnnotations(intent); - mPm = context.getPackageManager(); - mUsm = (UsageStatsManager) context.getSystemService(Context.USAGE_STATS_SERVICE); - mAzComparator = new AzInfoComparator(context); + for (UserHandle user : targetUserSpaceList) { + Context userContext = launchedFromContext.createContextAsUser(user, 0); + mPmMap.put(user, userContext.getPackageManager()); + mUsmMap.put(user, + (UsageStatsManager) userContext.getSystemService(Context.USAGE_STATS_SERVICE)); + } + mAzComparator = new AzInfoComparator(launchedFromContext); } // get annotations of content from intent. @@ -208,8 +227,8 @@ public abstract class AbstractResolverComparator implements ComparatorDefault implementation does nothing, as we could have simple model that does not train * online. * - * @param componentName the component that the user clicked + * @param targetInfo the target that the user clicked. */ - void updateModel(ComponentName componentName) { + void updateModel(TargetInfo targetInfo) { } /** Called before {@link #doCompute(List)}. Sets up 500ms timeout. */ diff --git a/core/java/com/android/internal/app/AppPredictionServiceResolverComparator.java b/core/java/com/android/internal/app/AppPredictionServiceResolverComparator.java index 115a9d772658f..b9f02365bbe74 100644 --- a/core/java/com/android/internal/app/AppPredictionServiceResolverComparator.java +++ b/core/java/com/android/internal/app/AppPredictionServiceResolverComparator.java @@ -31,6 +31,9 @@ import android.os.UserHandle; import android.util.Log; import com.android.internal.app.ResolverActivity.ResolvedComponentInfo; +import com.android.internal.app.chooser.TargetInfo; + +import com.google.android.collect.Lists; import java.util.ArrayList; import java.util.Collections; @@ -70,7 +73,7 @@ class AppPredictionServiceResolverComparator extends AbstractResolverComparator AppPredictor appPredictor, UserHandle user, ChooserActivityLogger chooserActivityLogger) { - super(context, intent); + super(context, intent, Lists.newArrayList(user)); mContext = context; mIntent = intent; mAppPredictor = appPredictor; @@ -99,13 +102,13 @@ class AppPredictionServiceResolverComparator extends AbstractResolverComparator } @Override - float getScore(ComponentName name) { - return mComparatorModel.getScore(name); + float getScore(TargetInfo targetInfo) { + return mComparatorModel.getScore(targetInfo); } @Override - void updateModel(ComponentName componentName) { - mComparatorModel.notifyOnTargetSelected(componentName); + void updateModel(TargetInfo targetInfo) { + mComparatorModel.notifyOnTargetSelected(targetInfo); } @Override @@ -158,9 +161,12 @@ class AppPredictionServiceResolverComparator extends AbstractResolverComparator private void setupFallbackModel(List targets) { mResolverRankerService = new ResolverRankerServiceResolverComparator( - mContext, mIntent, mReferrerPackage, + mContext, + mIntent, + mReferrerPackage, () -> mHandler.sendEmptyMessage(RANKER_SERVICE_RESULT), - getChooserActivityLogger()); + getChooserActivityLogger(), + mUser); mComparatorModel = mModelBuilder.buildFallbackModel(mResolverRankerService); mResolverRankerService.compute(targets); } @@ -224,13 +230,13 @@ class AppPredictionServiceResolverComparator extends AbstractResolverComparator } @Override - public float getScore(ComponentName componentName) { - return comparator.getScore(componentName); + public float getScore(TargetInfo targetInfo) { + return comparator.getScore(targetInfo); } @Override - public void notifyOnTargetSelected(ComponentName componentName) { - comparator.updateModel(componentName); + public void notifyOnTargetSelected(TargetInfo targetInfo) { + comparator.updateModel(targetInfo); } }; } @@ -271,8 +277,8 @@ class AppPredictionServiceResolverComparator extends AbstractResolverComparator } @Override - public float getScore(ComponentName name) { - Integer rank = mTargetRanks.get(name); + public float getScore(TargetInfo targetInfo) { + Integer rank = mTargetRanks.get(targetInfo.getResolvedComponentName()); if (rank == null) { Log.w(TAG, "Score requested for unknown component. Did you call compute yet?"); return 0f; @@ -282,13 +288,14 @@ class AppPredictionServiceResolverComparator extends AbstractResolverComparator } @Override - public void notifyOnTargetSelected(ComponentName componentName) { + public void notifyOnTargetSelected(TargetInfo targetInfo) { mAppPredictor.notifyAppTargetEvent( new AppTargetEvent.Builder( new AppTarget.Builder( - new AppTargetId(componentName.toString()), - componentName.getPackageName(), mUser) - .setClassName(componentName.getClassName()).build(), + new AppTargetId(targetInfo.getResolvedComponentName().toString()), + targetInfo.getResolvedComponentName().getPackageName(), mUser) + .setClassName(targetInfo.getResolvedComponentName() + .getClassName()).build(), ACTION_LAUNCH).build()); } } diff --git a/core/java/com/android/internal/app/ChooserActivity.java b/core/java/com/android/internal/app/ChooserActivity.java index 791819b18e131..f257f1c9b7cec 100644 --- a/core/java/com/android/internal/app/ChooserActivity.java +++ b/core/java/com/android/internal/app/ChooserActivity.java @@ -2268,9 +2268,11 @@ public class ChooserActivity extends ResolverActivity implements mChooserMultiProfilePagerAdapter.getActiveListAdapter(); if (currentListAdapter != null) { sendImpressionToAppPredictor(info, currentListAdapter); - currentListAdapter.updateModel(info.getResolvedComponentName()); - currentListAdapter.updateChooserCounts(ri.activityInfo.packageName, - targetIntent.getAction()); + currentListAdapter.updateModel(info); + currentListAdapter.updateChooserCounts( + ri.activityInfo.packageName, + targetIntent.getAction(), + ri.userHandle); } if (DEBUG) { Log.d(TAG, "ResolveInfo Package is " + ri.activityInfo.packageName); @@ -2395,7 +2397,10 @@ public class ChooserActivity extends ResolverActivity implements */ @Nullable private AppPredictor getAppPredictorForShareActivitiesIfEnabled(UserHandle userHandle) { - return USE_PREDICTION_MANAGER_FOR_SHARE_ACTIVITIES ? createAppPredictor(userHandle) : null; + // We cannot use APS service when clone profile is present as APS service cannot sort + // cross profile targets as of now. + return USE_PREDICTION_MANAGER_FOR_SHARE_ACTIVITIES && getCloneProfileUserHandle() == null + ? createAppPredictor(userHandle) : null; } void onRefinementResult(TargetInfo selectedTarget, Intent matchingIntent) { @@ -2548,8 +2553,13 @@ public class ChooserActivity extends ResolverActivity implements getReferrerPackageName(), appPredictor, userHandle, getChooserActivityLogger()); } else { resolverComparator = - new ResolverRankerServiceResolverComparator(this, getTargetIntent(), - getReferrerPackageName(), null, getChooserActivityLogger()); + new ResolverRankerServiceResolverComparator( + this, + getTargetIntent(), + getReferrerPackageName(), + null, + getChooserActivityLogger(), + getResolverRankerServiceUserHandleList(userHandle)); } UserHandle queryIntentsUser = getQueryIntentsUser(userHandle); diff --git a/core/java/com/android/internal/app/ResolverActivity.java b/core/java/com/android/internal/app/ResolverActivity.java index ddfc238034523..992e243060fc4 100644 --- a/core/java/com/android/internal/app/ResolverActivity.java +++ b/core/java/com/android/internal/app/ResolverActivity.java @@ -1617,6 +1617,14 @@ public class ResolverActivity extends Activity implements @VisibleForTesting protected ResolverListController createListController(UserHandle userHandle) { UserHandle queryIntentsUser = getQueryIntentsUser(userHandle); + ResolverRankerServiceResolverComparator resolverComparator = + new ResolverRankerServiceResolverComparator( + this, + getTargetIntent(), + getReferrerPackageName(), + null, + null, + getResolverRankerServiceUserHandleList(userHandle)); return new ResolverListController( this, mPm, @@ -1624,6 +1632,7 @@ public class ResolverActivity extends Activity implements getReferrerPackageName(), mLaunchedFromUid, userHandle, + resolverComparator, queryIntentsUser); } @@ -2534,4 +2543,21 @@ public class ResolverActivity extends Activity implements UserHandle predictedHandle) { return resolveInfo.userHandle; } + + /** + * Returns the {@link List} of {@link UserHandle} to pass on to the + * {@link ResolverRankerServiceResolverComparator} as per the provided {@code userHandle}. + */ + protected final List getResolverRankerServiceUserHandleList(UserHandle userHandle) { + List userList = new ArrayList<>(); + userList.add(userHandle); + // Add clonedProfileUserHandle to the list only if we are: + // a. Building the Personal Tab. + // b. CloneProfile exists on the device. + if (userHandle.equals(getPersonalProfileUserHandle()) + && getCloneProfileUserHandle() != null) { + userList.add(getCloneProfileUserHandle()); + } + return userList; + } } diff --git a/core/java/com/android/internal/app/ResolverComparatorModel.java b/core/java/com/android/internal/app/ResolverComparatorModel.java index 3e8f64bf4ed37..a3900166821d0 100644 --- a/core/java/com/android/internal/app/ResolverComparatorModel.java +++ b/core/java/com/android/internal/app/ResolverComparatorModel.java @@ -16,11 +16,11 @@ package com.android.internal.app; -import android.content.ComponentName; import android.content.pm.ResolveInfo; +import com.android.internal.app.chooser.TargetInfo; + import java.util.Comparator; -import java.util.List; /** * A ranking model for resolver targets, providing ordering and (optionally) numerical scoring. @@ -45,7 +45,7 @@ interface ResolverComparatorModel { * likelihood that the user will select that component as the target. Implementations that don't * assign numerical scores are recommended to return a value of 0 for all components. */ - float getScore(ComponentName name); + float getScore(TargetInfo targetInfo); /** * Notify the model that the user selected a target. (Models may log this information, use it as @@ -53,5 +53,5 @@ interface ResolverComparatorModel { * {@code ResolverComparatorModel} instance is immutable, clients will need to get an up-to-date * instance in order to see any changes in the ranking that might result from this feedback. */ - void notifyOnTargetSelected(ComponentName componentName); + void notifyOnTargetSelected(TargetInfo targetInfo); } diff --git a/core/java/com/android/internal/app/ResolverListAdapter.java b/core/java/com/android/internal/app/ResolverListAdapter.java index e3236ba5ed63b..2df2b2bec2f80 100644 --- a/core/java/com/android/internal/app/ResolverListAdapter.java +++ b/core/java/com/android/internal/app/ResolverListAdapter.java @@ -21,7 +21,6 @@ import static android.content.Context.ACTIVITY_SERVICE; import android.annotation.NonNull; import android.annotation.Nullable; import android.app.ActivityManager; -import android.content.ComponentName; import android.content.Context; import android.content.Intent; import android.content.PermissionChecker; @@ -157,17 +156,17 @@ public class ResolverListAdapter extends BaseAdapter { /** * Returns the app share score of the given {@code componentName}. */ - public float getScore(ComponentName componentName) { - return mResolverListController.getScore(componentName); + public float getScore(TargetInfo targetInfo) { + return mResolverListController.getScore(targetInfo); } - public void updateModel(ComponentName componentName) { - mResolverListController.updateModel(componentName); + public void updateModel(TargetInfo targetInfo) { + mResolverListController.updateModel(targetInfo); } - public void updateChooserCounts(String packageName, String action) { + public void updateChooserCounts(String packageName, String action, UserHandle userHandle) { mResolverListController.updateChooserCounts( - packageName, getUserHandle().getIdentifier(), action); + packageName, userHandle, action); } List getUnfilteredResolveList() { diff --git a/core/java/com/android/internal/app/ResolverListController.java b/core/java/com/android/internal/app/ResolverListController.java index 5f90e7386bd2e..d9a19b0f0283a 100644 --- a/core/java/com/android/internal/app/ResolverListController.java +++ b/core/java/com/android/internal/app/ResolverListController.java @@ -33,6 +33,7 @@ import android.util.Log; import com.android.internal.annotations.VisibleForTesting; import com.android.internal.app.chooser.DisplayResolveInfo; +import com.android.internal.app.chooser.TargetInfo; import java.util.ArrayList; import java.util.Collections; @@ -72,7 +73,13 @@ public class ResolverListController { UserHandle queryIntentsAsUser) { this(context, pm, targetIntent, referrerPackage, launchedFromUid, userHandle, new ResolverRankerServiceResolverComparator( - context, targetIntent, referrerPackage, null, null), queryIntentsAsUser); + context, + targetIntent, + referrerPackage, + null, + null, + userHandle), + queryIntentsAsUser); } public ResolverListController( @@ -397,22 +404,22 @@ public class ResolverListController { @VisibleForTesting public float getScore(DisplayResolveInfo target) { - return mResolverComparator.getScore(target.getResolvedComponentName()); + return mResolverComparator.getScore(target); } /** * Returns the app share score of the given {@code componentName}. */ - public float getScore(ComponentName componentName) { - return mResolverComparator.getScore(componentName); + public float getScore(TargetInfo targetInfo) { + return mResolverComparator.getScore(targetInfo); } - public void updateModel(ComponentName componentName) { - mResolverComparator.updateModel(componentName); + public void updateModel(TargetInfo targetInfo) { + mResolverComparator.updateModel(targetInfo); } - public void updateChooserCounts(String packageName, int userId, String action) { - mResolverComparator.updateChooserCounts(packageName, userId, action); + public void updateChooserCounts(String packageName, UserHandle user, String action) { + mResolverComparator.updateChooserCounts(packageName, user, action); } public void destroy() { diff --git a/core/java/com/android/internal/app/ResolverRankerServiceResolverComparator.java b/core/java/com/android/internal/app/ResolverRankerServiceResolverComparator.java index e7f80a7f60710..78c453dd41055 100644 --- a/core/java/com/android/internal/app/ResolverRankerServiceResolverComparator.java +++ b/core/java/com/android/internal/app/ResolverRankerServiceResolverComparator.java @@ -17,11 +17,13 @@ package com.android.internal.app; +import android.annotation.Nullable; import android.app.usage.UsageStats; import android.content.ComponentName; import android.content.Context; import android.content.Intent; import android.content.ServiceConnection; +import android.content.pm.ActivityInfo; import android.content.pm.ApplicationInfo; import android.content.pm.PackageManager; import android.content.pm.PackageManager.NameNotFoundException; @@ -38,12 +40,16 @@ import android.service.resolver.ResolverTarget; import android.util.Log; import com.android.internal.app.ResolverActivity.ResolvedComponentInfo; +import com.android.internal.app.chooser.TargetInfo; import com.android.internal.logging.MetricsLogger; import com.android.internal.logging.nano.MetricsProto.MetricsEvent; +import com.google.android.collect.Lists; + import java.text.Collator; import java.util.ArrayList; import java.util.Comparator; +import java.util.HashMap; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; @@ -69,10 +75,10 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator private static final int CONNECTION_COST_TIMEOUT_MILLIS = 200; private final Collator mCollator; - private final Map mStats; + private final Map> mStatsPerUser; private final long mCurrentTime; private final long mSinceTime; - private final LinkedHashMap mTargetsDict = new LinkedHashMap<>(); + private final Map> mTargetsDictPerUser; private final String mReferrerPackage; private final Object mLock = new Object(); private ArrayList mTargets; @@ -85,17 +91,34 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator private CountDownLatch mConnectSignal; private ResolverRankerServiceComparatorModel mComparatorModel; - public ResolverRankerServiceResolverComparator(Context context, Intent intent, + // context here refers to the activity calling this comparator. + // targetUserSpace refers to the userSpace in which the targets to be ranked lie. + public ResolverRankerServiceResolverComparator(Context launchedFromContext, Intent intent, String referrerPackage, AfterCompute afterCompute, - ChooserActivityLogger chooserActivityLogger) { - super(context, intent); - mCollator = Collator.getInstance(context.getResources().getConfiguration().locale); - mReferrerPackage = referrerPackage; - mContext = context; + ChooserActivityLogger chooserActivityLogger, UserHandle targetUserSpace) { + this(launchedFromContext, intent, referrerPackage, afterCompute, chooserActivityLogger, + Lists.newArrayList(targetUserSpace)); + } + // context here refers to the activity calling this comparator. + // targetUserSpaceList refers to the userSpace(s) in which the targets to be ranked lie. + public ResolverRankerServiceResolverComparator(Context launchedFromContext, Intent intent, + String referrerPackage, AfterCompute afterCompute, + ChooserActivityLogger chooserActivityLogger, List targetUserSpaceList) { + super(launchedFromContext, intent, targetUserSpaceList); + mCollator = Collator.getInstance(launchedFromContext + .getResources().getConfiguration().locale); + mReferrerPackage = referrerPackage; + mContext = launchedFromContext; mCurrentTime = System.currentTimeMillis(); mSinceTime = mCurrentTime - USAGE_STATS_PERIOD; - mStats = mUsm.queryAndAggregateUsageStats(mSinceTime, mCurrentTime); + mStatsPerUser = new HashMap<>(); + mTargetsDictPerUser = new HashMap<>(); + for (UserHandle user : targetUserSpaceList) { + mStatsPerUser.put(user, mUsmMap.get(user) + .queryAndAggregateUsageStats(mSinceTime, mCurrentTime)); + mTargetsDictPerUser.put(user, new LinkedHashMap<>()); + } mAction = intent.getAction(); mRankerServiceName = new ComponentName(mContext, this.getClass()); setCallBack(afterCompute); @@ -147,57 +170,63 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator for (ResolvedComponentInfo target : targets) { final ResolverTarget resolverTarget = new ResolverTarget(); - mTargetsDict.put(target.name, resolverTarget); - final UsageStats pkStats = mStats.get(target.name.getPackageName()); - if (pkStats != null) { - // Only count recency for apps that weren't the caller - // since the caller is always the most recent. - // Persistent processes muck this up, so omit them too. - if (!target.name.getPackageName().equals(mReferrerPackage) - && !isPersistentProcess(target)) { - final float recencyScore = - (float) Math.max(pkStats.getLastTimeUsed() - recentSinceTime, 0); - resolverTarget.setRecencyScore(recencyScore); - if (recencyScore > mostRecencyScore) { - mostRecencyScore = recencyScore; - } - } - final float timeSpentScore = (float) pkStats.getTotalTimeInForeground(); - resolverTarget.setTimeSpentScore(timeSpentScore); - if (timeSpentScore > mostTimeSpentScore) { - mostTimeSpentScore = timeSpentScore; - } - final float launchScore = (float) pkStats.mLaunchCount; - resolverTarget.setLaunchScore(launchScore); - if (launchScore > mostLaunchScore) { - mostLaunchScore = launchScore; - } - - float chooserScore = 0.0f; - if (pkStats.mChooserCounts != null && mAction != null - && pkStats.mChooserCounts.get(mAction) != null) { - chooserScore = (float) pkStats.mChooserCounts.get(mAction) - .getOrDefault(mContentType, 0); - if (mAnnotations != null) { - final int size = mAnnotations.length; - for (int i = 0; i < size; i++) { - chooserScore += (float) pkStats.mChooserCounts.get(mAction) - .getOrDefault(mAnnotations[i], 0); + final LinkedHashMap targetsDict = mTargetsDictPerUser + .get(target.getResolveInfoAt(0).userHandle); + final Map stats = mStatsPerUser + .get(target.getResolveInfoAt(0).userHandle); + if (targetsDict != null && stats != null) { + targetsDict.put(target.name, resolverTarget); + final UsageStats pkStats = stats.get(target.name.getPackageName()); + if (pkStats != null) { + // Only count recency for apps that weren't the caller + // since the caller is always the most recent. + // Persistent processes muck this up, so omit them too. + if (!target.name.getPackageName().equals(mReferrerPackage) + && !isPersistentProcess(target)) { + final float recencyScore = + (float) Math.max(pkStats.getLastTimeUsed() - recentSinceTime, 0); + resolverTarget.setRecencyScore(recencyScore); + if (recencyScore > mostRecencyScore) { + mostRecencyScore = recencyScore; } } - } - if (DEBUG) { - if (mAction == null) { - Log.d(TAG, "Action type is null"); - } else { - Log.d(TAG, "Chooser Count of " + mAction + ":" + - target.name.getPackageName() + " is " + - Float.toString(chooserScore)); + final float timeSpentScore = (float) pkStats.getTotalTimeInForeground(); + resolverTarget.setTimeSpentScore(timeSpentScore); + if (timeSpentScore > mostTimeSpentScore) { + mostTimeSpentScore = timeSpentScore; + } + final float launchScore = (float) pkStats.mLaunchCount; + resolverTarget.setLaunchScore(launchScore); + if (launchScore > mostLaunchScore) { + mostLaunchScore = launchScore; + } + + float chooserScore = 0.0f; + if (pkStats.mChooserCounts != null && mAction != null + && pkStats.mChooserCounts.get(mAction) != null) { + chooserScore = (float) pkStats.mChooserCounts.get(mAction) + .getOrDefault(mContentType, 0); + if (mAnnotations != null) { + final int size = mAnnotations.length; + for (int i = 0; i < size; i++) { + chooserScore += (float) pkStats.mChooserCounts.get(mAction) + .getOrDefault(mAnnotations[i], 0); + } + } + } + if (DEBUG) { + if (mAction == null) { + Log.d(TAG, "Action type is null"); + } else { + Log.d(TAG, "Chooser Count of " + mAction + ":" + + target.name.getPackageName() + " is " + + Float.toString(chooserScore)); + } + } + resolverTarget.setChooserScore(chooserScore); + if (chooserScore > mostChooserScore) { + mostChooserScore = chooserScore; } - } - resolverTarget.setChooserScore(chooserScore); - if (chooserScore > mostChooserScore) { - mostChooserScore = chooserScore; } } } @@ -209,7 +238,11 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator + " mostChooserScore: " + mostChooserScore); } - mTargets = new ArrayList<>(mTargetsDict.values()); + mTargets = new ArrayList<>(); + for (UserHandle u : mTargetsDictPerUser.keySet()) { + mTargets.addAll(mTargetsDictPerUser.get(u).values()); + } + for (ResolverTarget target : mTargets) { final float recency = target.getRecencyScore() / mostRecencyScore; setFeatures(target, recency * recency * RECENCY_MULTIPLIER, @@ -232,15 +265,15 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator } @Override - public float getScore(ComponentName name) { - return mComparatorModel.getScore(name); + public float getScore(TargetInfo targetInfo) { + return mComparatorModel.getScore(targetInfo); } // update ranking model when the connection to it is valid. @Override - public void updateModel(ComponentName componentName) { + public void updateModel(TargetInfo targetInfo) { synchronized (mLock) { - mComparatorModel.notifyOnTargetSelected(componentName); + mComparatorModel.notifyOnTargetSelected(targetInfo); } } @@ -281,7 +314,8 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator // resolve the service for ranking. private Intent resolveRankerService() { Intent intent = new Intent(ResolverRankerService.SERVICE_INTERFACE); - final List resolveInfos = mPm.queryIntentServices(intent, 0); + final List resolveInfos = mContext.getPackageManager() + .queryIntentServices(intent, 0); for (ResolveInfo resolveInfo : resolveInfos) { if (resolveInfo == null || resolveInfo.serviceInfo == null || resolveInfo.serviceInfo.applicationInfo == null) { @@ -294,7 +328,8 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator resolveInfo.serviceInfo.applicationInfo.packageName, resolveInfo.serviceInfo.name); try { - final String perm = mPm.getServiceInfo(componentName, 0).permission; + final String perm = mContext.getPackageManager() + .getServiceInfo(componentName, 0).permission; if (!ResolverRankerService.BIND_PERMISSION.equals(perm)) { Log.w(TAG, "ResolverRankerService " + componentName + " does not require" + " permission " + ResolverRankerService.BIND_PERMISSION @@ -305,9 +340,9 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator + " in the manifest."); continue; } - if (PackageManager.PERMISSION_GRANTED != mPm.checkPermission( - ResolverRankerService.HOLD_PERMISSION, - resolveInfo.serviceInfo.packageName)) { + if (PackageManager.PERMISSION_GRANTED != mContext.getPackageManager() + .checkPermission(ResolverRankerService.HOLD_PERMISSION, + resolveInfo.serviceInfo.packageName)) { Log.w(TAG, "ResolverRankerService " + componentName + " does not hold" + " permission " + ResolverRankerService.HOLD_PERMISSION + " - this service will not be queried for " @@ -385,7 +420,9 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator @Override void beforeCompute() { super.beforeCompute(); - mTargetsDict.clear(); + for (UserHandle userHandle : mTargetsDictPerUser.keySet()) { + mTargetsDictPerUser.get(userHandle).clear(); + } mTargets = null; mRankerServiceName = new ComponentName(mContext, this.getClass()); mComparatorModel = buildUpdatedModel(); @@ -465,14 +502,14 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator // so the ResolverComparatorModel may provide inconsistent results. We should make immutable // copies of the data (waiting for any necessary remaining data before creating the model). return new ResolverRankerServiceComparatorModel( - mStats, - mTargetsDict, + mStatsPerUser, + mTargetsDictPerUser, mTargets, mCollator, mRanker, mRankerServiceName, (mAnnotations != null), - mPm); + mPmMap); } /** @@ -481,35 +518,36 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator * removing the complex legacy API. */ static class ResolverRankerServiceComparatorModel implements ResolverComparatorModel { - private final Map mStats; // Treat as immutable. - private final Map mTargetsDict; // Treat as immutable. + private final Map> mStatsPerUser; // Treat as immutable. + private final Map> mTargetsDictPerUser; // Treat as immutable. private final List mTargets; // Treat as immutable. private final Collator mCollator; private final IResolverRankerService mRanker; private final ComponentName mRankerServiceName; private final boolean mAnnotationsUsed; - private final PackageManager mPm; + private final Map mPmMap; // TODO: it doesn't look like we should have to pass both targets and targetsDict, but it's // not written in a way that makes it clear whether we can derive one from the other (at // least in this constructor). ResolverRankerServiceComparatorModel( - Map stats, - Map targetsDict, + Map> statsPerUser, + Map> targetsDictPerUser, List targets, Collator collator, IResolverRankerService ranker, ComponentName rankerServiceName, boolean annotationsUsed, - PackageManager pm) { - mStats = stats; - mTargetsDict = targetsDict; + Map pmMap) { + mStatsPerUser = statsPerUser; + mTargetsDictPerUser = targetsDictPerUser; mTargets = targets; mCollator = collator; mRanker = ranker; mRankerServiceName = rankerServiceName; mAnnotationsUsed = annotationsUsed; - mPm = pm; + mPmMap = pmMap; } @Override @@ -518,25 +556,29 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator // a bug there, or do we have a way of knowing it will be non-null under certain // conditions? return (lhs, rhs) -> { - if (mStats != null) { - final ResolverTarget lhsTarget = mTargetsDict.get(new ComponentName( - lhs.activityInfo.packageName, lhs.activityInfo.name)); - final ResolverTarget rhsTarget = mTargetsDict.get(new ComponentName( - rhs.activityInfo.packageName, rhs.activityInfo.name)); + final ResolverTarget lhsTarget = getActivityResolverTargetForUser(lhs.activityInfo, + lhs.userHandle); + final ResolverTarget rhsTarget = getActivityResolverTargetForUser(rhs.activityInfo, + rhs.userHandle); - if (lhsTarget != null && rhsTarget != null) { - final int selectProbabilityDiff = Float.compare( - rhsTarget.getSelectProbability(), lhsTarget.getSelectProbability()); + if (lhsTarget != null && rhsTarget != null) { + final int selectProbabilityDiff = Float.compare( + rhsTarget.getSelectProbability(), lhsTarget.getSelectProbability()); - if (selectProbabilityDiff != 0) { - return selectProbabilityDiff > 0 ? 1 : -1; - } + if (selectProbabilityDiff != 0) { + return selectProbabilityDiff > 0 ? 1 : -1; } } - CharSequence sa = lhs.loadLabel(mPm); + CharSequence sa = null; + if (mPmMap.containsKey(lhs.userHandle)) { + sa = lhs.loadLabel(mPmMap.get(lhs.userHandle)); + } if (sa == null) sa = lhs.activityInfo.name; - CharSequence sb = rhs.loadLabel(mPm); + CharSequence sb = null; + if (mPmMap.containsKey(rhs.userHandle)) { + sb = rhs.loadLabel(mPmMap.get(rhs.userHandle)); + } if (sb == null) sb = rhs.activityInfo.name; return mCollator.compare(sa.toString().trim(), sb.toString().trim()); @@ -544,22 +586,28 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator } @Override - public float getScore(ComponentName name) { - final ResolverTarget target = mTargetsDict.get(name); - if (target != null) { - return target.getSelectProbability(); + public float getScore(TargetInfo targetInfo) { + if (mTargetsDictPerUser.containsKey(targetInfo.getResolveInfo().userHandle) + && mTargetsDictPerUser.get(targetInfo.getResolveInfo().userHandle) + .get(targetInfo.getResolvedComponentName()) != null) { + return mTargetsDictPerUser.get(targetInfo.getResolveInfo().userHandle) + .get(targetInfo.getResolvedComponentName()).getSelectProbability(); } return 0; } @Override - public void notifyOnTargetSelected(ComponentName componentName) { + public void notifyOnTargetSelected(TargetInfo targetInfo) { if (mRanker != null) { try { - int selectedPos = new ArrayList(mTargetsDict.keySet()) - .indexOf(componentName); + int selectedPos = -1; + if (mTargetsDictPerUser.containsKey(targetInfo.getResolveInfo().userHandle)) { + selectedPos = new ArrayList<>(mTargetsDictPerUser + .get(targetInfo.getResolveInfo().userHandle).keySet()) + .indexOf(targetInfo.getResolvedComponentName()); + } if (selectedPos >= 0 && mTargets != null) { - final float selectedProbability = getScore(componentName); + final float selectedProbability = getScore(targetInfo); int order = 0; for (ResolverTarget target : mTargets) { if (target.getSelectProbability() > selectedProbability) { @@ -570,7 +618,8 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator mRanker.train(mTargets, selectedPos); } else { if (DEBUG) { - Log.d(TAG, "Selected a unknown component: " + componentName); + Log.d(TAG, "Selected a unknown component: " + targetInfo + .getResolvedComponentName()); } } } catch (RemoteException e) { @@ -594,5 +643,16 @@ class ResolverRankerServiceResolverComparator extends AbstractResolverComparator metricsLogger.write(log); } } + + @Nullable + private ResolverTarget getActivityResolverTargetForUser( + ActivityInfo activity, UserHandle user) { + if ((mStatsPerUser == null) || !mTargetsDictPerUser.containsKey(user)) { + return null; + } + return mTargetsDictPerUser + .get(user) + .get(new ComponentName(activity.packageName, activity.name)); + } } } diff --git a/core/tests/coretests/src/com/android/internal/app/AbstractResolverComparatorTest.java b/core/tests/coretests/src/com/android/internal/app/AbstractResolverComparatorTest.java index 3e640c1bad39c..bdf42aca8edf4 100644 --- a/core/tests/coretests/src/com/android/internal/app/AbstractResolverComparatorTest.java +++ b/core/tests/coretests/src/com/android/internal/app/AbstractResolverComparatorTest.java @@ -18,6 +18,7 @@ package com.android.internal.app; import static junit.framework.Assert.assertEquals; +import android.app.Instrumentation; import android.content.ComponentName; import android.content.Context; import android.content.Intent; @@ -27,6 +28,10 @@ import android.os.Message; import androidx.test.InstrumentationRegistry; +import com.android.internal.app.chooser.TargetInfo; + +import com.google.android.collect.Lists; + import org.junit.Test; import java.util.List; @@ -96,7 +101,8 @@ public class AbstractResolverComparatorTest { Intent intent = new Intent(); AbstractResolverComparator testComparator = - new AbstractResolverComparator(context, intent) { + new AbstractResolverComparator(context, intent, + Lists.newArrayList(context.getUser())) { @Override int compare(ResolveInfo lhs, ResolveInfo rhs) { @@ -109,7 +115,7 @@ public class AbstractResolverComparatorTest { void doCompute(List targets) {} @Override - float getScore(ComponentName name) { + float getScore(TargetInfo targetInfo) { return 0; } diff --git a/core/tests/coretests/src/com/android/internal/app/ChooserActivityTest.java b/core/tests/coretests/src/com/android/internal/app/ChooserActivityTest.java index 78a75dbceaa13..c06df9455e07f 100644 --- a/core/tests/coretests/src/com/android/internal/app/ChooserActivityTest.java +++ b/core/tests/coretests/src/com/android/internal/app/ChooserActivityTest.java @@ -545,13 +545,17 @@ public class ChooserActivityTest { return true; }; ResolveInfo toChoose = resolvedComponentInfos.get(0).getResolveInfoAt(0); + DisplayResolveInfo testDri = + activity.createTestDisplayResolveInfo(sendIntent, toChoose, "testLabel", "testInfo", + sendIntent,/* resolveInfoPresentationGetter */ null); onView(withText(toChoose.activityInfo.name)) .perform(click()); waitForIdle(); verify(ChooserActivityOverrideData.getInstance().resolverListController, times(1)) - .updateChooserCounts(Mockito.anyString(), anyInt(), Mockito.anyString()); + .updateChooserCounts(Mockito.anyString(), any(UserHandle.class), + Mockito.anyString()); verify(ChooserActivityOverrideData.getInstance().resolverListController, times(1)) - .updateModel(toChoose.activityInfo.getComponentName()); + .updateModel(testDri); assertThat(activity.getIsSelected(), is(true)); } diff --git a/core/tests/coretests/src/com/android/internal/app/FakeResolverComparatorModel.java b/core/tests/coretests/src/com/android/internal/app/FakeResolverComparatorModel.java index fbbe57c8e3252..573135ffa5a1f 100644 --- a/core/tests/coretests/src/com/android/internal/app/FakeResolverComparatorModel.java +++ b/core/tests/coretests/src/com/android/internal/app/FakeResolverComparatorModel.java @@ -19,6 +19,8 @@ package com.android.internal.app; import android.content.ComponentName; import android.content.pm.ResolveInfo; +import com.android.internal.app.chooser.TargetInfo; + import java.util.ArrayList; import java.util.Comparator; import java.util.List; @@ -45,14 +47,15 @@ public class FakeResolverComparatorModel implements ResolverComparatorModel { } @Override - public float getScore(ComponentName name) { + public float getScore(TargetInfo targetInfo) { return 0.0f; // Models are not required to provide numerical scores. } @Override - public void notifyOnTargetSelected(ComponentName componentName) { + public void notifyOnTargetSelected(TargetInfo targetInfo) { System.out.println( - "User selected " + componentName + " under model " + System.identityHashCode(this)); + "User selected " + targetInfo.getResolvedComponentName() + " under model " + + System.identityHashCode(this)); } private FakeResolverComparatorModel(Comparator comparator) { diff --git a/core/tests/coretests/src/com/android/internal/app/ResolverListControllerTest.java b/core/tests/coretests/src/com/android/internal/app/ResolverListControllerTest.java index b0f1e047f5d00..8f6cee33078e0 100644 --- a/core/tests/coretests/src/com/android/internal/app/ResolverListControllerTest.java +++ b/core/tests/coretests/src/com/android/internal/app/ResolverListControllerTest.java @@ -86,6 +86,7 @@ public class ResolverListControllerTest { config.locale = Locale.getDefault(); List services = new ArrayList<>(); mUsm = new UsageStatsManager(mMockContext, mMockService); + when(mMockContext.createContextAsUser(any(), anyInt())).thenReturn(mMockContext); when(mMockContext.getSystemService(Context.USAGE_STATS_SERVICE)).thenReturn(mUsm); when(mMockPackageManager.queryIntentServices(any(), anyInt())).thenReturn(services); when(mMockResources.getConfiguration()).thenReturn(config); @@ -126,7 +127,7 @@ public class ResolverListControllerTest { UserHandle.SYSTEM); mController.sort(new ArrayList()); long beforeReport = getCount(mUsm, packageName, action, annotation); - mController.updateChooserCounts(packageName, UserHandle.USER_CURRENT, action); + mController.updateChooserCounts(packageName, UserHandle.SYSTEM, action); long afterReport = getCount(mUsm, packageName, action, annotation); assertThat(afterReport, is(beforeReport + 1l)); }