package com.android.internal.protolog;

import static android.internal.perfetto.protos.Protolog.ProtoLogViewerConfig.GROUPS;
import static android.internal.perfetto.protos.Protolog.ProtoLogViewerConfig.Group.ID;
import static android.internal.perfetto.protos.Protolog.ProtoLogViewerConfig.Group.NAME;
import static android.internal.perfetto.protos.Protolog.ProtoLogViewerConfig.MESSAGES;
import static android.internal.perfetto.protos.Protolog.ProtoLogViewerConfig.MessageData.MESSAGE;
import static android.internal.perfetto.protos.Protolog.ProtoLogViewerConfig.MessageData.MESSAGE_ID;
import static android.internal.perfetto.protos.Protolog.ProtoLogViewerConfig.MessageData.GROUP_ID;

import android.annotation.NonNull;
import android.annotation.Nullable;
import android.util.LongSparseArray;
import android.util.proto.ProtoInputStream;

import com.android.internal.protolog.common.ILogger;

import java.io.IOException;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
import java.util.TreeMap;

public class ProtoLogViewerConfigReader {
    @NonNull
    private final ViewerConfigInputStreamProvider mViewerConfigInputStreamProvider;
    @NonNull
    private final Map<String, Set<Long>> mGroupHashes = new TreeMap<>();
    @NonNull
    private final LongSparseArray<String> mLogMessageMap = new LongSparseArray<>();

    public ProtoLogViewerConfigReader(
            @NonNull ViewerConfigInputStreamProvider viewerConfigInputStreamProvider) {
        this.mViewerConfigInputStreamProvider = viewerConfigInputStreamProvider;
    }

    /**
     * Data class for a log message from the viewer config.
     */
    public static class MessageData {
        @NonNull
        public final String message;
        @NonNull
        public final String group;

        public MessageData(@NonNull String message, @NonNull String group) {
            this.message = message;
            this.group = group;
        }
    }

    /**
     * Returns message format string for its hash or null if unavailable
     * or the viewer config is not loaded into memory.
     */
    @Nullable
    public String getViewerString(long messageHash) {
        return mLogMessageMap.get(messageHash);
    }

    /**
     * Load the viewer configs for the target groups into memory.
     * Only viewer configs loaded into memory can be required. So this must be called for all groups
     * we want to query before we query their viewer config.
     *
     * @param groups Groups to load the viewer configs from file into memory.
     */
    public synchronized void loadViewerConfig(@NonNull String[] groups) {
        loadViewerConfig(groups, (message) -> {});
    }

    /**
     * Loads the viewer config into memory. No-op if already loaded in memory.
     */
    public synchronized void loadViewerConfig(@NonNull String[] groups, @NonNull ILogger logger) {
        for (String group : groups) {
            if (mGroupHashes.containsKey(group)) {
                continue;
            }

            try {
                Map<Long, String> mappings = loadViewerConfigMappingForGroup(group);
                mGroupHashes.put(group, mappings.keySet());
                for (Long key : mappings.keySet()) {
                    mLogMessageMap.put(key, mappings.get(key));
                }

                logger.log("Loaded " + mLogMessageMap.size() + " log definitions");
            } catch (IOException e) {
                logger.log("Unable to load log definitions: "
                        + "IOException while processing viewer config" + e);
            }
        }
    }

    public synchronized void unloadViewerConfig(@NonNull String[] groups) {
        unloadViewerConfig(groups, (message) -> {});
    }

    /**
     * Unload the viewer config from memory.
     */
    public synchronized void unloadViewerConfig(@NonNull String[] groups, @NonNull ILogger logger) {
        for (String group : groups) {
            if (!mGroupHashes.containsKey(group)) {
                continue;
            }

            final Set<Long> hashes = mGroupHashes.get(group);
            for (Long hash : hashes) {
                logger.log("Unloading viewer config hash " + hash);
                mLogMessageMap.remove(hash);
            }
            mGroupHashes.remove(group);
        }
    }

    /**
     * Return whether or not the viewer config file contains a message with the specified hash.
     * @param messageHash The hash message we are looking for in the viewer config file
     * @return True iff the message with message hash is contained in the viewer config.
     * @throws IOException if there was an issue reading the viewer config file.
     */
    public boolean messageHashIsAvailableInFile(long messageHash)
            throws IOException {
        try (var pisWrapper = mViewerConfigInputStreamProvider.getInputStream()) {
            final var pis = pisWrapper.get();
            while (pis.nextField() != ProtoInputStream.NO_MORE_FIELDS) {
                if (pis.getFieldNumber() == (int) MESSAGES) {
                    final long inMessageToken = pis.start(MESSAGES);

                    while (pis.nextField() != ProtoInputStream.NO_MORE_FIELDS) {
                        if (pis.getFieldNumber() == (int) MESSAGE_ID) {
                            if (pis.readLong(MESSAGE_ID) == messageHash) {
                                return true;
                            }
                        }
                    }

                    pis.end(inMessageToken);
                }
            }
        }

        return false;
    }

    /**
     * Returns the message data for a given message hash from the viewer config file.
     *
     * @param messageHash The hash of the message we are looking for in the viewer config file.
     * @return The {@link MessageData} if the message is found, null otherwise.
     * @throws IOException if there was an issue reading the viewer config file.
     */
    @Nullable
    public MessageData getMessageDataForHashFromFile(long messageHash)
            throws IOException {
        try (var pisWrapper = mViewerConfigInputStreamProvider.getInputStream()) {
            final var pis = pisWrapper.get();

            String foundMessage = null;
            long foundGroupId = -1;
            final LongSparseArray<String> groupMap = new LongSparseArray<>();

            while (pis.nextField() != ProtoInputStream.NO_MORE_FIELDS) {
                if (pis.getFieldNumber() == (int) MESSAGES) {
                    final long inMessageToken = pis.start(MESSAGES);
                    ParsedMessage parsedMessage = readMessage(pis);
                    if (parsedMessage.messageId == messageHash) {
                        foundMessage = parsedMessage.message;
                        foundGroupId = parsedMessage.groupId;
                    }

                    pis.end(inMessageToken);
                } else if (pis.getFieldNumber() == (int) GROUPS) {
                    final long inMessageToken = pis.start(GROUPS);

                    long groupId = 0;
                    ParsedGroup parsedGroup = readGroup(pis);
                    if (parsedGroup.groupName != null) {
                        groupMap.put(parsedGroup.groupId, parsedGroup.groupName);
                    }
                    pis.end(inMessageToken);
                }
            }

            if (foundMessage != null) {
                String groupName = groupMap.get(foundGroupId);
                if (groupName != null) {
                    return new MessageData(foundMessage, groupName);
                }
            }
        }

        return null;
    }

    @NonNull
    private Map<Long, String> loadViewerConfigMappingForGroup(@NonNull String group)
            throws IOException {
        long targetGroupId = loadGroupId(group);

        final Map<Long, String> hashesForGroup = new TreeMap<>();
        try (var pisWrapper = mViewerConfigInputStreamProvider.getInputStream()) {
            final var pis = pisWrapper.get();
            while (pis.nextField() != ProtoInputStream.NO_MORE_FIELDS) {
                if (pis.getFieldNumber() == (int) MESSAGES) {
                    final long inMessageToken = pis.start(MESSAGES);
                    ParsedMessage parsedMessage = readMessage(pis);

                    if (parsedMessage.groupId == 0) {
                        throw new IOException("Failed to get group id");
                    }

                    if (parsedMessage.messageId == 0) {
                        throw new IOException("Failed to get message id");
                    }

                    if (parsedMessage.message == null) {
                        throw new IOException("Failed to get message string");
                    }

                    if (parsedMessage.groupId == targetGroupId) {
                        hashesForGroup.put(parsedMessage.messageId, parsedMessage.message);
                    }

                    pis.end(inMessageToken);
                }
            }
        }

        return hashesForGroup;
    }

    private long loadGroupId(@NonNull String group) throws IOException {
        try (var pisWrapper = mViewerConfigInputStreamProvider.getInputStream()) {
            final var pis = pisWrapper.get();

            while (pis.nextField() != ProtoInputStream.NO_MORE_FIELDS) {
                if (pis.getFieldNumber() == (int) GROUPS) {
                    final long inMessageToken = pis.start(GROUPS);
                    ParsedGroup parsedGroup = readGroup(pis);
                    if (Objects.equals(parsedGroup.groupName, group)) {
                        return parsedGroup.groupId;
                    }

                    pis.end(inMessageToken);
                }
            }
        }

        throw new RuntimeException("Group " + group + " not found in viewer config");
    }

    private static class ParsedMessage {
        long messageId = 0;
        String message = null;
        int groupId = 0;
    }

    private static ParsedMessage readMessage(ProtoInputStream pis) throws IOException {
        final ParsedMessage result = new ParsedMessage();
        while (pis.nextField() != ProtoInputStream.NO_MORE_FIELDS) {
            switch (pis.getFieldNumber()) {
                case (int) MESSAGE_ID:
                    result.messageId = pis.readLong(MESSAGE_ID);
                    break;
                case (int) MESSAGE:
                    result.message = pis.readString(MESSAGE);
                    break;
                case (int) GROUP_ID:
                    result.groupId = pis.readInt(GROUP_ID);
                    break;
            }
        }
        return result;
    }

    private static class ParsedGroup {
        long groupId = 0;
        String groupName = null;
    }

    private static ParsedGroup readGroup(ProtoInputStream pis) throws IOException {
        final ParsedGroup result = new ParsedGroup();
        while (pis.nextField() != ProtoInputStream.NO_MORE_FIELDS) {
            switch (pis.getFieldNumber()) {
                case (int) ID:
                    result.groupId = pis.readInt(ID);
                    break;
                case (int) NAME:
                    result.groupName = pis.readString(NAME);
                    break;
            }
        }
        return result;
    }
}
