/*
 * Copyright (C) 2024 The Android Open Source Project
 *
 * Licensed under the Apache License, Version 2.0 (the "License");
 * you may not use this file except in compliance with the License.
 * You may obtain a copy of the License at
 *
 *      http://www.apache.org/licenses/LICENSE-2.0
 *
 * Unless required by applicable law or agreed to in writing, software
 * distributed under the License is distributed on an "AS IS" BASIS,
 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
 * See the License for the specific language governing permissions and
 * limitations under the License.
 */

package com.android.ondevicepersonalization.services.download;

import android.adservices.ondevicepersonalization.Constants;
import android.adservices.ondevicepersonalization.DownloadCompletedOutputParcel;
import android.adservices.ondevicepersonalization.DownloadInputParcel;
import android.adservices.ondevicepersonalization.aidl.IIsolatedModelService;
import android.annotation.NonNull;
import android.content.ComponentName;
import android.content.Context;
import android.net.Uri;
import android.os.Bundle;

import com.android.odp.module.common.Clock;
import com.android.odp.module.common.MonotonicClock;
import com.android.odp.module.common.PackageUtils;
import com.android.ondevicepersonalization.internal.util.LoggerFactory;
import com.android.ondevicepersonalization.services.FlagsFactory;
import com.android.ondevicepersonalization.services.OnDevicePersonalizationExecutors;
import com.android.ondevicepersonalization.services.data.DataAccessPermission;
import com.android.ondevicepersonalization.services.data.DataAccessServiceImpl;
import com.android.ondevicepersonalization.services.data.vendor.OnDevicePersonalizationVendorDataDao;
import com.android.ondevicepersonalization.services.data.vendor.VendorData;
import com.android.ondevicepersonalization.services.download.mdd.MobileDataDownloadFactory;
import com.android.ondevicepersonalization.services.download.mdd.OnDevicePersonalizationFileGroupPopulator;
import com.android.ondevicepersonalization.services.federatedcompute.FederatedComputeServiceImpl;
import com.android.ondevicepersonalization.services.inference.IsolatedModelServiceProvider;
import com.android.ondevicepersonalization.services.manifest.AppManifestConfigHelper;
import com.android.ondevicepersonalization.services.policyengine.UserDataAccessor;
import com.android.ondevicepersonalization.services.serviceflow.ServiceFlow;
import com.android.ondevicepersonalization.services.util.StatsUtils;

import com.google.android.libraries.mobiledatadownload.GetFileGroupRequest;
import com.google.android.libraries.mobiledatadownload.MobileDataDownload;
import com.google.android.libraries.mobiledatadownload.RemoveFileGroupRequest;
import com.google.android.libraries.mobiledatadownload.file.SynchronousFileStorage;
import com.google.android.libraries.mobiledatadownload.file.openers.ReadStreamOpener;
import com.google.common.util.concurrent.FluentFuture;
import com.google.common.util.concurrent.FutureCallback;
import com.google.common.util.concurrent.Futures;
import com.google.common.util.concurrent.ListenableFuture;
import com.google.common.util.concurrent.ListeningExecutorService;
import com.google.mobiledatadownload.ClientConfigProto;

import java.io.IOException;
import java.io.InputStream;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;

public class DownloadFlow implements ServiceFlow<DownloadCompletedOutputParcel> {
    private static final LoggerFactory.Logger sLogger = LoggerFactory.getLogger();
    private static final String TAG = DownloadFlow.class.getSimpleName();
    private final String mPackageName;
    private final Context mContext;
    private OnDevicePersonalizationVendorDataDao mDao;

    @NonNull
    private IsolatedModelServiceProvider mModelServiceProvider;
    private long mStartServiceTimeMillis;
    private ComponentName mService;
    private ParsedFileContents mParsedFileContents;

    private final Injector mInjector;
    private final FutureCallback<DownloadCompletedOutputParcel> mCallback;

    private static class Injector {
        Clock getClock() {
            return MonotonicClock.getInstance();
        }

        ListeningExecutorService getExecutor() {
            return OnDevicePersonalizationExecutors.getBackgroundExecutor();
        }
    }

    public DownloadFlow(String packageName,
                        Context context, FutureCallback<DownloadCompletedOutputParcel> callback) {
        mPackageName = packageName;
        mContext = context;
        mCallback = callback;
        mInjector = new Injector();
    }

    @Override
    public boolean isServiceFlowReady() {
        try {
            mStartServiceTimeMillis = mInjector.getClock().elapsedRealtime();

            Uri uri = Objects.requireNonNull(getClientFileUri());

            ParsedFileContents fileContents;

            SynchronousFileStorage fileStorage = MobileDataDownloadFactory.getFileStorage(mContext);
            try (InputStream in = fileStorage.open(uri, ReadStreamOpener.create())) {
                fileContents = DownloadedFileParser.parseJson(in);
            } catch (IOException ie) {
                sLogger.e(ie, TAG + mPackageName + " Failed to process downloaded JSON file");
                onSuccess(null);
                return false;
            }

            long syncToken = fileContents.getSyncToken();
            if (syncToken == -1 || !validateSyncToken(syncToken)) {
                sLogger.d(TAG + mPackageName
                        + " downloaded JSON file has invalid syncToken provided");
                onSuccess(null);
                return false;
            }

            var vendorDataMap = fileContents.getVendorDataMap();
            if (vendorDataMap == null || vendorDataMap.isEmpty()) {
                sLogger.d(TAG + mPackageName + " downloaded JSON file has no content provided");
                onSuccess(null);
                return false;
            }

            mDao = OnDevicePersonalizationVendorDataDao.getInstance(mContext, getService(),
                    PackageUtils.getCertDigest(mContext, mPackageName));
            long existingSyncToken = mDao.getSyncToken();

            // If existingToken is greaterThan or equal to the new token, skip as there is
            // no new data. Mark success to upstream caller for reporting purpose
            if (existingSyncToken >= syncToken) {
                sLogger.d(
                        TAG
                                + ": new syncToken value "
                                + syncToken
                                + " is not newer than existing token value "
                                + existingSyncToken);
                onSuccess(null);
                return false;
            }

            mParsedFileContents = fileContents;

            return true;
        } catch (Exception e) {
            mCallback.onFailure(e);
            return false;
        }
    }

    @Override
    public ComponentName getService() {
        if (mService != null) return mService;

        mService = ComponentName.createRelative(mPackageName,
                AppManifestConfigHelper.getServiceNameFromOdpSettings(mContext, mPackageName));
        return mService;
    }

    @Override
    public Bundle getServiceParams() {
        Bundle serviceParams = new Bundle();

        serviceParams.putBinder(Constants.EXTRA_DATA_ACCESS_SERVICE_BINDER,
                new DataAccessServiceImpl(getService(), mContext,
                        /* localDataPermission */ DataAccessPermission.READ_WRITE,
                        /* eventDataPermission */ DataAccessPermission.READ_ONLY));

        serviceParams.putBinder(Constants.EXTRA_FEDERATED_COMPUTE_SERVICE_BINDER,
                new FederatedComputeServiceImpl(getService(), mContext));

        Map<String, byte[]> downloadedContent = new HashMap<>();
        for (String key : mParsedFileContents.getVendorDataMap().keySet()) {
            downloadedContent.put(key, mParsedFileContents.getVendorDataMap().get(key).getData());
        }

        DataAccessServiceImpl downloadedContentBinder = new DataAccessServiceImpl(
                getService(), mContext, /* remoteData */ downloadedContent,
                /* localDataPermission */ DataAccessPermission.DENIED,
                /* eventDataPermission */ DataAccessPermission.DENIED);

        serviceParams.putParcelable(Constants.EXTRA_INPUT,
                new DownloadInputParcel.Builder()
                        .setDataAccessServiceBinder(downloadedContentBinder)
                        .build());

        serviceParams.putParcelable(Constants.EXTRA_USER_DATA,
                new UserDataAccessor().getUserData());

        mModelServiceProvider = new IsolatedModelServiceProvider();
        IIsolatedModelService modelService = mModelServiceProvider.getModelService(mContext);
        serviceParams.putBinder(Constants.EXTRA_MODEL_SERVICE_BINDER, modelService.asBinder());

        return serviceParams;
    }

    @Override
    public void uploadServiceFlowMetrics(ListenableFuture<Bundle> runServiceFuture) {
        var unused =
                FluentFuture.from(runServiceFuture)
                        .transform(
                                val -> {
                                    StatsUtils.writeServiceRequestMetrics(
                                            Constants.API_NAME_SERVICE_ON_DOWNLOAD_COMPLETED,
                                            mService.getPackageName(),
                                            val,
                                            mInjector.getClock(),
                                            Constants.STATUS_SUCCESS,
                                            mStartServiceTimeMillis);
                                    return val;
                                },
                                mInjector.getExecutor())
                        .catchingAsync(
                                Exception.class,
                                e -> {
                                    StatsUtils.writeServiceRequestMetrics(
                                            Constants.API_NAME_SERVICE_ON_DOWNLOAD_COMPLETED,
                                            mService.getPackageName(),
                                            /* result= */ null,
                                            mInjector.getClock(),
                                            Constants.STATUS_INTERNAL_ERROR,
                                            mStartServiceTimeMillis);
                                    return Futures.immediateFailedFuture(e);
                                },
                                mInjector.getExecutor());
    }

    @Override
    public ListenableFuture<DownloadCompletedOutputParcel> getServiceFlowResultFuture(
            ListenableFuture<Bundle> runServiceFuture) {
        return FluentFuture.from(runServiceFuture)
                .transform(
                        result -> {
                            DownloadCompletedOutputParcel downloadResult =
                                    result.getParcelable(Constants.EXTRA_RESULT,
                                            DownloadCompletedOutputParcel.class);

                            List<String> retainedKeys = downloadResult.getRetainedKeys();
                            if (retainedKeys == null) {
                                // TODO(b/270710021): Determine how to correctly handle null
                                //  retainedKeys.
                                return null;
                            }

                            List<VendorData> filteredList = new ArrayList<>();
                            for (String key : retainedKeys) {
                                if (mParsedFileContents.getVendorDataMap().containsKey(key)) {
                                    filteredList.add(
                                            mParsedFileContents.getVendorDataMap().get(key));
                                }
                            }

                            boolean transactionResult =
                                    mDao.batchUpdateOrInsertVendorDataTransaction(filteredList,
                                            retainedKeys, mParsedFileContents.getSyncToken());

                            sLogger.d(TAG + ": filter and store data completed, transaction"
                                    + " successful: "
                                    + transactionResult);

                            return downloadResult;
                        },
                        mInjector.getExecutor())
                .catching(
                        Exception.class,
                        e -> {
                            sLogger.e(TAG + ": Processing failed.", e);
                            return null;
                        },
                        mInjector.getExecutor());
    }

    private ListenableFuture<Boolean> removeFileGroup() throws Exception {
        MobileDataDownload mdd = MobileDataDownloadFactory.getMdd(mContext);
        String fileGroupName =
                OnDevicePersonalizationFileGroupPopulator.createPackageFileGroupName(
                        mPackageName, mContext);

        return mdd.removeFileGroup(RemoveFileGroupRequest.newBuilder()
                .setGroupName(fileGroupName).build());
    }

    @Override
    public void returnResultThroughCallback(
            ListenableFuture<DownloadCompletedOutputParcel> serviceFlowResultFuture) {
        try {
            onSuccess(serviceFlowResultFuture.get());
        } catch (Exception e) {
            mCallback.onFailure(e);
        }
    }

    @Override
    public void cleanUpServiceParams() {
        mModelServiceProvider.unBindFromModelService();
    }

    private Uri getClientFileUri() throws Exception {
        MobileDataDownload mdd = MobileDataDownloadFactory.getMdd(mContext);

        String fileGroupName =
                OnDevicePersonalizationFileGroupPopulator.createPackageFileGroupName(
                        mPackageName, mContext);

        ClientConfigProto.ClientFileGroup cfg = mdd.getFileGroup(
                        GetFileGroupRequest.newBuilder()
                                .setGroupName(fileGroupName)
                                .build())
                .get();

        if (cfg == null || cfg.getStatus() != ClientConfigProto.ClientFileGroup.Status.DOWNLOADED) {
            sLogger.d(TAG + mPackageName + " has no completed downloads.");
            // No completed downloads is a valid case. Mark as success and return null.
            mCallback.onSuccess(null);
            return null;
        }

        // It is currently expected that we will only download a single file per package.
        if (cfg.getFileCount() != 1) {
            sLogger.d(TAG + ": package : "
                    + mPackageName + " has "
                    + cfg.getFileCount() + " files in the fileGroup");
            onFailure(new IllegalArgumentException("Invalid file count."));
            return null;
        }

        ClientConfigProto.ClientFile clientFile = cfg.getFile(0);
        int fileSize = clientFile.getFullSizeInBytes();
        if (fileSize > FlagsFactory.getFlags().getDefaultDownloadRejectCapInMb()) {
            // File size exceeds download limit is a valid case. Mark as success and return null.
            StatsUtils.writeServiceRequestMetrics(
                    Constants.API_NAME_SERVICE_ON_DOWNLOAD_COMPLETED,
                    mService.getPackageName(),
                    /* result= */ null,
                    mInjector.getClock(),
                    Constants.STATUS_DOWNLOAD_SIZE_EXCEED_CAP_ERROR,
                    mStartServiceTimeMillis);
            sLogger.d(TAG + ": File size " + fileSize + " exceed download size cap.");
            mCallback.onSuccess(null);
            return null;
        }
        return Uri.parse(clientFile.getFileUri());
    }

    private static boolean validateSyncToken(long syncToken) {
        // TODO(b/249813538) Add any additional requirements
        return syncToken % 3600 == 0;
    }

    private void onFailure(Exception exception) throws Exception {
        Futures.addCallback(removeFileGroup(),
                new FutureCallback<>() {
                    @Override
                    public void onSuccess(Boolean result) {
                        try {
                            mCallback.onFailure(exception);
                        } catch (Exception e) {
                            mCallback.onFailure(e);
                        }
                    }

                    @Override
                    public void onFailure(Throwable t) {
                        mCallback.onFailure(t);
                    }
                }, mInjector.getExecutor());
    }

    private void onSuccess(DownloadCompletedOutputParcel output) throws Exception {
        Futures.addCallback(removeFileGroup(),
                new FutureCallback<>() {
                    @Override
                    public void onSuccess(Boolean result) {
                        try {
                            mCallback.onSuccess(output);
                        } catch (Exception e) {
                            mCallback.onFailure(e);
                        }
                    }

                    @Override
                    public void onFailure(Throwable t) {
                        mCallback.onFailure(t);
                    }
                }, mInjector.getExecutor());
    }
}
