// Copyright 2024 Google Inc. All rights reserved.
//
// 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 find_input_delta_lib

import (
	"errors"
	"fmt"
	"io/fs"
	"os"
	"regexp"
	"slices"

	fid_proto "android/soong/cmd/find_input_delta/find_input_delta_proto_internal"
	"android/soong/third_party/zip"
	"github.com/google/blueprint/pathtools"
	"google.golang.org/protobuf/proto"
)

// Load the internal state from a file.
// If the file does not exist, an empty state is returned.
func loadState(filename string, fsys fs.ReadFileFS) (*fid_proto.PartialCompileInputs, error) {
	var message = &fid_proto.PartialCompileInputs{}
	data, err := fsys.ReadFile(filename)
	if err != nil && !errors.Is(err, fs.ErrNotExist) {
		return message, err
	}
	proto.Unmarshal(data, message)
	return message, nil
}

type StatReadFileFS interface {
	fs.StatFS
	fs.ReadFileFS
}

// Create the internal state by examining the inputs.
func createState(version string, tools, inputs []string, inspect_contents bool, fsys StatReadFileFS) (*fid_proto.PartialCompileInputs, error) {
	ret := &fid_proto.PartialCompileInputs{}
	if version != "" {
		ret.Version = proto.String(version)
	}
	statFile := func(input string, inspect bool) (pci *fid_proto.PartialCompileInput, err error) {
		stat, err := fs.Stat(fsys, input)
		if err != nil {
			return nil, err
		}
		pci = &fid_proto.PartialCompileInput{
			Name:      proto.String(input),
			MtimeNsec: proto.Int64(stat.ModTime().UnixNano()),
			// If we ever have an easy hash, assign it here.
		}
		if inspect {
			// NOTE: When we find it useful, we can parallelize the file inspection for speed.
			contents, err := inspectFileContents(input)
			if err != nil {
				return pci, err
			}
			if contents != nil {
				pci.Contents = contents
			}
		}
		return pci, nil
	}

	slices.Sort(tools)
	for _, tool := range tools {
		pci, err := statFile(tool, false)
		if err != nil {
			return ret, err
		}
		ret.Tools = append(ret.Tools, pci)
	}

	slices.Sort(inputs)
	for _, input := range inputs {
		pci, err := statFile(input, inspect_contents)
		if err != nil {
			if errors.Is(err, fs.ErrNotExist) {
				continue
			}
			return ret, err
		}
		ret.InputFiles = append(ret.InputFiles, pci)
	}
	return ret, nil
}

// We ignore any suffix digit caused by sharding.
var InspectExtsZipRegexp = regexp.MustCompile("\\.(jar|apex|apk)[0-9]*$")

// Inspect the file and extract the state of the elements in the archive.
// If this is not an archive of some sort, nil is returned.
func inspectFileContents(name string) ([]*fid_proto.PartialCompileInput, error) {
	if InspectExtsZipRegexp.Match([]byte(name)) {
		return inspectZipFileContents(name)
	}
	return nil, nil
}

func inspectZipFileContents(name string) ([]*fid_proto.PartialCompileInput, error) {
	rc, err := zip.OpenReader(name)
	if err != nil {
		return nil, err
	}
	ret := []*fid_proto.PartialCompileInput{}
	for _, v := range rc.File {
		// Only include timestamp when there is no CRC.
		timeNsec := proto.Int64(v.ModTime().UnixNano())
		if v.CRC32 != 0 {
			timeNsec = nil
		}
		pci := &fid_proto.PartialCompileInput{
			Name:      proto.String(v.Name),
			MtimeNsec: timeNsec,
			Hash:      proto.String(fmt.Sprintf("%08x", v.CRC32)),
		}
		ret = append(ret, pci)
		// We do not support nested inspection.
	}
	return ret, nil
}

func writeState(s *fid_proto.PartialCompileInputs, path string) error {
	data, err := proto.Marshal(s)
	if err != nil {
		return err
	}
	return pathtools.WriteFileIfChanged(path, data, 0644)
}

func compareInternalState(prior, other *fid_proto.PartialCompileInputs, target string) *FileList {
	// If Version or any Tools are different, then claim every input is new.
	nilInputs := []*fid_proto.PartialCompileInput{}
	if prior.Version != other.Version {
		return compareInputFiles(nilInputs, other.GetInputFiles(), target)
	} else {
		_priorToolsMap := make(map[string]*fid_proto.PartialCompileInput)
		for _, v := range prior.GetTools() {
			_priorToolsMap[v.GetName()] = v
		}
		for _, v := range other.GetTools() {
			tool := v.GetName()
			if _, ok := _priorToolsMap[tool]; !ok || !proto.Equal(_priorToolsMap[tool], v) {
				return compareInputFiles(nilInputs, other.GetInputFiles(), target)
			}
		}
	}
	// Check the inputs for changes.
	return compareInputFiles(prior.GetInputFiles(), other.GetInputFiles(), target)
}

func compareInputFiles(prior, other []*fid_proto.PartialCompileInput, name string) *FileList {
	fl := FileListFactory(name)
	PriorMap := make(map[string]*fid_proto.PartialCompileInput, len(prior))
	for _, v := range prior {
		PriorMap[v.GetName()] = v
	}
	otherMap := make(map[string]*fid_proto.PartialCompileInput, len(other))
	for _, v := range other {
		name = v.GetName()
		otherMap[name] = v
		if _, ok := PriorMap[name]; !ok {
			// Added file
			fl.addFile(name)
		} else if !proto.Equal(PriorMap[name], v) {
			// Changed file
			fl.changeFile(name, compareInputFiles(PriorMap[name].GetContents(), v.GetContents(), name))
		}
	}
	for _, v := range prior {
		name := v.GetName()
		if _, ok := otherMap[name]; !ok {
			// Deleted file
			fl.deleteFile(name)
		}
	}
	return fl
}

func GenerateFileList(target, prior_state_file, new_state_file, version string, tools, inputs []string, inspect bool, fsys StatReadFileFS) (file_list *FileList, err error) {
	// Read the prior state
	prior_state, err := loadState(prior_state_file, fsys)
	if err != nil {
		return
	}
	// Create the new state
	new_state, err := createState(version, tools, inputs, inspect, fsys)
	if err != nil {
		return
	}
	if err = writeState(new_state, new_state_file); err != nil {
		return
	}

	file_list = compareInternalState(prior_state, new_state, target)

	metrics_dir := os.Getenv("SOONG_METRICS_AGGREGATION_DIR")
	out_dir := os.Getenv("OUT_DIR")
	if metrics_dir != "" {
		if err = file_list.WriteMetrics(metrics_dir, out_dir); err != nil {
			return
		}
	}
	return
}
