#include "qnn_meta.hpp"

#include "System/QnnSystemContext.h"

uint32_t qnn_data_type_size(Qnn_DataType_t t) {
    switch (t) {
        case QNN_DATATYPE_INT_8:
        case QNN_DATATYPE_UINT_8:
        case QNN_DATATYPE_SFIXED_POINT_8:
        case QNN_DATATYPE_UFIXED_POINT_8:
        case QNN_DATATYPE_BOOL_8:
            return 1;
        case QNN_DATATYPE_INT_16:
        case QNN_DATATYPE_UINT_16:
        case QNN_DATATYPE_SFIXED_POINT_16:
        case QNN_DATATYPE_UFIXED_POINT_16:
        case QNN_DATATYPE_FLOAT_16:
        case QNN_DATATYPE_BFLOAT_16:
            return 2;
        case QNN_DATATYPE_INT_32:
        case QNN_DATATYPE_UINT_32:
        case QNN_DATATYPE_SFIXED_POINT_32:
        case QNN_DATATYPE_UFIXED_POINT_32:
        case QNN_DATATYPE_FLOAT_32:
            return 4;
        case QNN_DATATYPE_INT_64:
        case QNN_DATATYPE_UINT_64:
        case QNN_DATATYPE_FLOAT_64:
            return 8;
        default:
            return 0;
    }
}

namespace {

// Version-tolerant accessors for the Qnn_Tensor_t v1/v2 union. The fields
// used here have identical names in both versions.
template <typename F>
auto with_tensor(const Qnn_Tensor_t &t, F f) {
    if (t.version == QNN_TENSOR_VERSION_2) {
        return f(t.v2);
    }
    return f(t.v1);
}

TensorMeta tensor_meta_from(const Qnn_Tensor_t &t) {
    TensorMeta meta;
    with_tensor(t, [&](const auto &v) {
        meta.name = v.name ? v.name : "";
        meta.id = v.id;
        meta.data_type = v.dataType;
        meta.dims.assign(v.dimensions, v.dimensions + v.rank);

        if (v.quantizeParams.encodingDefinition == QNN_DEFINITION_DEFINED) {
            switch (v.quantizeParams.quantizationEncoding) {
                case QNN_QUANTIZATION_ENCODING_SCALE_OFFSET:
                    meta.quant = TensorMeta::Quant::PerTensor;
                    meta.scale_offset = {v.quantizeParams.scaleOffsetEncoding.scale,
                                         v.quantizeParams.scaleOffsetEncoding.offset};
                    break;
                case QNN_QUANTIZATION_ENCODING_AXIS_SCALE_OFFSET: {
                    const auto &enc = v.quantizeParams.axisScaleOffsetEncoding;
                    meta.quant = TensorMeta::Quant::PerAxis;
                    meta.axis = enc.axis;
                    meta.axis_scale_offsets.reserve(enc.numScaleOffsets);
                    for (uint32_t i = 0; i < enc.numScaleOffsets; i++) {
                        meta.axis_scale_offsets.push_back(
                            {enc.scaleOffset[i].scale, enc.scaleOffset[i].offset});
                    }
                    break;
                }
                default:
                    // Other encodings (blockwise etc.) are exposed as :none;
                    // callers must pass pre-quantized data for those tensors.
                    meta.quant = TensorMeta::Quant::None;
                    break;
            }
        }
        return 0;
    });

    uint64_t elems = 1;
    for (uint32_t d : meta.dims) elems *= d;
    meta.byte_size = elems * qnn_data_type_size(meta.data_type);
    return meta;
}

// GraphInfo v1/v2/v3 share the graphName/inputs/outputs field names.
template <typename G>
GraphMeta graph_meta_from(const G &g) {
    GraphMeta meta;
    meta.name = g.graphName ? g.graphName : "";
    meta.inputs.reserve(g.numGraphInputs);
    for (uint32_t i = 0; i < g.numGraphInputs; i++) {
        meta.inputs.push_back(tensor_meta_from(g.graphInputs[i]));
    }
    meta.outputs.reserve(g.numGraphOutputs);
    for (uint32_t i = 0; i < g.numGraphOutputs; i++) {
        meta.outputs.push_back(tensor_meta_from(g.graphOutputs[i]));
    }
    return meta;
}

GraphMeta graph_meta_from_info(const QnnSystemContext_GraphInfo_t &info) {
    switch (info.version) {
        case QNN_SYSTEM_CONTEXT_GRAPH_INFO_VERSION_1:
            return graph_meta_from(info.graphInfoV1);
        case QNN_SYSTEM_CONTEXT_GRAPH_INFO_VERSION_2:
            return graph_meta_from(info.graphInfoV2);
        case QNN_SYSTEM_CONTEXT_GRAPH_INFO_VERSION_3:
            return graph_meta_from(info.graphInfoV3);
        default:
            return {};
    }
}

template <typename B>
void collect_graphs(const B &bin, std::vector<GraphMeta> &out) {
    out.reserve(bin.numGraphs);
    for (uint32_t i = 0; i < bin.numGraphs; i++) {
        out.push_back(graph_meta_from_info(bin.graphs[i]));
    }
}

} // namespace

std::string qnn_parse_binary_info(const QnnApi &api,
                                  const void *buf,
                                  uint64_t size,
                                  std::vector<GraphMeta> &out) {
    if (!api.sys_fn.systemContextCreate || !api.sys_fn.systemContextGetBinaryInfo ||
        !api.sys_fn.systemContextFree) {
        return "system_interface_unavailable";
    }

    QnnSystemContext_Handle_t handle = nullptr;
    if (api.sys_fn.systemContextCreate(&handle) != QNN_SUCCESS) {
        return "system_context_create_failed";
    }

    const QnnSystemContext_BinaryInfo_t *info = nullptr;
    Qnn_ContextBinarySize_t info_size = 0;
    Qnn_ErrorHandle_t err = api.sys_fn.systemContextGetBinaryInfo(
        handle, const_cast<void *>(buf), size, &info, &info_size);
    if (err != QNN_SUCCESS || !info) {
        api.sys_fn.systemContextFree(handle);
        return "get_binary_info_failed";
    }

    switch (info->version) {
        case QNN_SYSTEM_CONTEXT_BINARY_INFO_VERSION_1:
            collect_graphs(info->contextBinaryInfoV1, out);
            break;
        case QNN_SYSTEM_CONTEXT_BINARY_INFO_VERSION_2:
            collect_graphs(info->contextBinaryInfoV2, out);
            break;
        case QNN_SYSTEM_CONTEXT_BINARY_INFO_VERSION_3:
            collect_graphs(info->contextBinaryInfoV3, out);
            break;
        default:
            api.sys_fn.systemContextFree(handle);
            return "unknown_binary_info_version";
    }

    api.sys_fn.systemContextFree(handle);
    return "";
}
