package main

// Security signals. All of them are heuristics: a match is a reason to look, never proof,
// and the absence of a match proves nothing either. The agent only reads: it never kills,
// quarantines or changes anything.

import (
	"bufio"
	"bytes"
	"crypto/sha256"
	_ "embed"
	"encoding/base64"
	"errors"
	"io/fs"
	"os"
	"path/filepath"
	"sort"
	"strings"
)

//go:embed signatures.txt
var signaturesTxt []byte

// Signatures from signatures.txt.
type Signatures struct {
	Names map[string]bool
	Args  []string
	Busy  map[string]bool
}

func ParseSignatures(b []byte) Signatures {
	s := Signatures{Names: map[string]bool{}, Busy: map[string]bool{}}
	sc := bufio.NewScanner(bytes.NewReader(b))
	for sc.Scan() {
		line := strings.TrimSpace(sc.Text())
		if line == "" || strings.HasPrefix(line, "#") {
			continue
		}
		kind, val, ok := strings.Cut(line, " ")
		val = strings.ToLower(strings.TrimSpace(val))
		if !ok || val == "" {
			continue
		}
		switch kind {
		case "name":
			s.Names[val] = true
		case "arg":
			s.Args = append(s.Args, val)
		case "busy":
			s.Busy[val] = true
		}
	}
	return s
}

var sigs = ParseSignatures(signaturesTxt)

// Programs that show other programs' names in their arguments (someone looking for a
// miner with `grep xmrig` is not running one).
var lookers = map[string]bool{
	"grep": true, "egrep": true, "fgrep": true, "rg": true, "ag": true, "pgrep": true, "pkill": true, "ps": true,
	"less": true, "more": true, "tail": true, "head": true, "cat": true, "vi": true, "vim": true, "nano": true,
	"emacs": true, "journalctl": true, "man": true, "find": true, "locate": true, "htop": true, "top": true,
	"approvalens-age": true, "approvalens-agent": true,
}

// MinerMatch: why a process looks like a crypto miner ("" = it does not).
func (s Signatures) MinerMatch(comm, argv0, cmdline string) string {
	c := strings.ToLower(comm)
	a0 := strings.ToLower(filepath.Base(argv0))
	if s.Names[c] {
		return "name:" + c
	}
	if argv0 != "" && s.Names[a0] {
		return "name:" + a0
	}
	if lookers[c] || lookers[a0] {
		return ""
	}
	lc := strings.ToLower(cmdline)
	for _, a := range s.Args {
		if strings.Contains(lc, a) {
			return "arg:" + a
		}
	}
	return ""
}

var tmpDirs = []string{"/tmp/", "/var/tmp/", "/dev/shm/", "/run/shm/"}

// InTmp: an executable that runs from a world-writable temporary directory.
func InTmp(path string) bool {
	for _, d := range tmpDirs {
		if strings.HasPrefix(path, d) {
			return true
		}
	}
	return false
}

// SuspiciousDeleted: the process runs a binary that was deleted after it started. A
// package upgrade does that to every running daemon, so binaries under the system
// directories are left out; a memfd (a program that never touched the disk) is not.
func SuspiciousDeleted(exe string) bool {
	if !strings.HasSuffix(exe, " (deleted)") {
		return false
	}
	p := strings.TrimSuffix(exe, " (deleted)")
	if strings.HasPrefix(p, "/memfd:") {
		return true
	}
	for _, d := range []string{"/usr/", "/bin/", "/sbin/", "/lib/", "/lib64/", "/snap/", "/opt/", "/nix/", "/var/lib/docker/"} {
		if strings.HasPrefix(p, d) {
			return false
		}
	}
	return true
}

// ---------------------------------------------------------------------------
// crontabs
// ---------------------------------------------------------------------------

// ParseCrontab returns the active lines of one crontab ("<source>: <line>"), whitespace
// collapsed, comments and blank lines dropped. Variable lines (PATH=..., MAILTO=...) are
// kept: a changed PATH is a known way to hijack a job.
func ParseCrontab(content []byte, source string) []string {
	var out []string
	sc := bufio.NewScanner(bytes.NewReader(content))
	sc.Buffer(make([]byte, 64*1024), 1024*1024)
	for sc.Scan() {
		line := strings.Join(strings.Fields(sc.Text()), " ")
		if line == "" || strings.HasPrefix(line, "#") {
			continue
		}
		out = append(out, truncate(source+": "+line, 300))
	}
	return out
}

// DiffSets returns the items of cur that are not in prev, and the items of prev that are
// not in cur (both sorted).
func DiffSets(prev, cur []string) (added, removed []string) {
	p := make(map[string]bool, len(prev))
	for _, x := range prev {
		p[x] = true
	}
	c := make(map[string]bool, len(cur))
	for _, x := range cur {
		c[x] = true
		if !p[x] {
			added = append(added, x)
		}
	}
	for _, x := range prev {
		if !c[x] {
			removed = append(removed, x)
		}
	}
	sort.Strings(added)
	sort.Strings(removed)
	return dedupSorted(added), dedupSorted(removed)
}

func dedupSorted(xs []string) []string {
	if len(xs) < 2 {
		return xs
	}
	out := xs[:1]
	for _, x := range xs[1:] {
		if x != out[len(out)-1] {
			out = append(out, x)
		}
	}
	return out
}

// ReadCrontabs collects every system and user crontab it can read. unreadable lists the
// places it could not (other users' crontabs need root).
func ReadCrontabs(root string) (entries []string, unreadable []string) {
	files := []string{filepath.Join(root, "etc/crontab")}
	for _, dir := range []string{"etc/cron.d", "var/spool/cron/crontabs", "var/spool/cron"} {
		d := filepath.Join(root, dir)
		ents, err := os.ReadDir(d)
		if err != nil {
			if errors.Is(err, fs.ErrPermission) {
				unreadable = append(unreadable, "/"+dir)
			}
			continue
		}
		for _, e := range ents {
			if e.Type().IsRegular() && !strings.HasPrefix(e.Name(), ".") {
				files = append(files, filepath.Join(d, e.Name()))
			}
		}
	}
	for _, f := range files {
		b, err := readSmall(f, 256*1024)
		if err != nil {
			if errors.Is(err, fs.ErrPermission) {
				unreadable = append(unreadable, strings.TrimPrefix(f, strings.TrimSuffix(root, "/")))
			}
			continue
		}
		entries = append(entries, ParseCrontab(b, strings.TrimPrefix(f, strings.TrimSuffix(root, "/")))...)
	}
	sort.Strings(entries)
	entries = dedupSorted(entries)
	if len(entries) > 300 {
		entries = entries[:300]
	}
	return entries, unreadable
}

// ---------------------------------------------------------------------------
// SSH authorized keys
// ---------------------------------------------------------------------------

// ParseAuthorizedKeys returns "<user> <type> SHA256:<fingerprint> <comment>" per key.
func ParseAuthorizedKeys(b []byte, user string) []string {
	var out []string
	sc := bufio.NewScanner(bytes.NewReader(b))
	sc.Buffer(make([]byte, 64*1024), 1024*1024)
	for sc.Scan() {
		line := strings.TrimSpace(sc.Text())
		if line == "" || strings.HasPrefix(line, "#") {
			continue
		}
		f := strings.Fields(line)
		idx := -1
		for i, x := range f {
			if keyType(x) {
				idx = i
				break
			}
		}
		if idx < 0 || idx+1 >= len(f) {
			continue
		}
		blob, err := base64.StdEncoding.DecodeString(f[idx+1])
		if err != nil {
			continue
		}
		sum := sha256.Sum256(blob)
		item := user + " " + f[idx] + " SHA256:" + base64.RawStdEncoding.EncodeToString(sum[:])
		if idx > 0 {
			item += " [options]"
		}
		if idx+2 < len(f) {
			item += " " + truncate(strings.Join(f[idx+2:], " "), 60)
		}
		out = append(out, item)
	}
	return out
}

func keyType(s string) bool {
	return strings.HasPrefix(s, "ssh-") || strings.HasPrefix(s, "ecdsa-sha2-") || strings.HasPrefix(s, "sk-ssh-") ||
		strings.HasPrefix(s, "sk-ecdsa-")
}

// ReadSSHKeys reads authorized_keys of root and of every account with a login shell.
func ReadSSHKeys(root string, users []PasswdEntry) (keys []string, unreadable []string) {
	seen := map[string]bool{}
	for _, u := range users {
		if u.UID != 0 && !LoginShell(u.Shell) {
			continue
		}
		if u.Home == "" || u.Home == "/" || seen[u.Home] {
			continue
		}
		seen[u.Home] = true
		for _, name := range []string{"authorized_keys", "authorized_keys2"} {
			p := filepath.Join(root, u.Home, ".ssh", name)
			b, err := readSmall(p, 512*1024)
			if err != nil {
				if errors.Is(err, fs.ErrPermission) {
					unreadable = append(unreadable, filepath.Join(u.Home, ".ssh", name))
				}
				continue
			}
			keys = append(keys, ParseAuthorizedKeys(b, u.Name)...)
		}
	}
	sort.Strings(keys)
	keys = dedupSorted(keys)
	if len(keys) > 200 {
		keys = keys[:200]
	}
	return keys, dedupSorted(sortedCopy(unreadable))
}

// ---------------------------------------------------------------------------
// setuid / setgid binaries
// ---------------------------------------------------------------------------

var setuidRoots = []string{"/bin", "/sbin", "/usr/bin", "/usr/sbin", "/usr/local/bin", "/usr/local/sbin", "/usr/lib",
	"/usr/libexec", "/usr/local/lib", "/opt", "/tmp", "/var/tmp", "/dev/shm"}

// ScanSetuid walks the usual binary directories (and the temporary ones) for setuid or
// setgid files: "<path> <mode> uid=<n>". It stays light: no symlinks are followed and
// it stops after maxFiles entries.
func ScanSetuid(roots []string, maxFiles int) (items []string, truncated bool) {
	seen := map[string]bool{}
	count := 0
	for _, r := range roots {
		real, err := filepath.EvalSymlinks(r)
		if err != nil || seen[real] {
			continue
		}
		seen[real] = true
		_ = filepath.WalkDir(real, func(p string, d fs.DirEntry, err error) error {
			if err != nil {
				if d != nil && d.IsDir() {
					return fs.SkipDir
				}
				return nil
			}
			count++
			if count > maxFiles {
				truncated = true
				return fs.SkipAll
			}
			if d.IsDir() {
				if strings.Count(strings.TrimPrefix(p, real), "/") > 8 {
					return fs.SkipDir
				}
				return nil
			}
			if !d.Type().IsRegular() {
				return nil
			}
			info, err := d.Info()
			if err != nil {
				return nil
			}
			m := info.Mode()
			if m&(fs.ModeSetuid|fs.ModeSetgid) == 0 {
				return nil
			}
			items = append(items, p+" "+modeString(m)+" uid="+itoa(fileUID(info)))
			return nil
		})
		if truncated {
			break
		}
	}
	sort.Strings(items)
	return items, truncated
}

func modeString(m fs.FileMode) string {
	v := uint32(m.Perm())
	if m&fs.ModeSetuid != 0 {
		v |= 0o4000
	}
	if m&fs.ModeSetgid != 0 {
		v |= 0o2000
	}
	if m&fs.ModeSticky != 0 {
		v |= 0o1000
	}
	s := []byte("0000")
	for i := 3; i >= 0; i-- {
		s[i] = byte('0' + v&7)
		v >>= 3
	}
	return string(s)
}

// ---------------------------------------------------------------------------
// web roots
// ---------------------------------------------------------------------------

var webExts = map[string]bool{".php": true, ".phtml": true, ".php3": true, ".php4": true, ".php5": true, ".php7": true,
	".phar": true, ".inc": true, ".js": true, ".mjs": true, ".html": true, ".htm": true, ".shtml": true, ".htaccess": true,
	".user.ini": true}

var webSkipDirs = map[string]bool{"node_modules": true, ".git": true, ".svn": true, ".hg": true, "cache": true, ".cache": true}

// WebFile key: modification time and size, never the content.
func webWatched(name string) bool {
	if name == ".htaccess" || name == ".user.ini" {
		return true
	}
	return webExts[strings.ToLower(filepath.Ext(name))]
}

// ScanWebRoots maps each PHP / JS / HTML file (and .htaccess) under the roots to
// "<mtime>:<size>". Stops at maxFiles (truncated=true).
func ScanWebRoots(roots []string, exclude []string, maxFiles int) (files map[string]string, truncated bool) {
	files = map[string]string{}
	skip := map[string]bool{}
	for k := range webSkipDirs {
		skip[k] = true
	}
	for _, x := range exclude {
		skip[x] = true
	}
	for _, r := range roots {
		_ = filepath.WalkDir(r, func(p string, d fs.DirEntry, err error) error {
			if err != nil {
				if d != nil && d.IsDir() {
					return fs.SkipDir
				}
				return nil
			}
			if d.IsDir() {
				if p != r && (skip[d.Name()] || skip[p]) {
					return fs.SkipDir
				}
				return nil
			}
			if !d.Type().IsRegular() || !webWatched(d.Name()) {
				return nil
			}
			if len(files) >= maxFiles {
				truncated = true
				return fs.SkipAll
			}
			info, err := d.Info()
			if err != nil {
				return nil
			}
			files[p] = itoa64(info.ModTime().Unix()) + ":" + itoa64(info.Size())
			return nil
		})
		if truncated {
			break
		}
	}
	return files, truncated
}

// WebChanges between two scans: new, changed and removed file names (capped).
type WebChanges struct {
	Added     []string `json:"added"`
	Changed   []string `json:"changed"`
	Removed   []string `json:"removed"`
	NAdded    int      `json:"n_added"`
	NChanged  int      `json:"n_changed"`
	NRemoved  int      `json:"n_removed"`
	Files     int      `json:"files"`
	Truncated bool     `json:"truncated,omitempty"`
}

func (w WebChanges) Empty() bool { return w.NAdded+w.NChanged+w.NRemoved == 0 }

func DiffWebRoots(prev, cur map[string]string, limit int) WebChanges {
	var w WebChanges
	w.Files = len(cur)
	for p, v := range cur {
		old, ok := prev[p]
		switch {
		case !ok:
			w.NAdded++
			w.Added = append(w.Added, p)
		case old != v:
			w.NChanged++
			w.Changed = append(w.Changed, p)
		}
	}
	for p := range prev {
		if _, ok := cur[p]; !ok {
			w.NRemoved++
			w.Removed = append(w.Removed, p)
		}
	}
	for _, xs := range []*[]string{&w.Added, &w.Changed, &w.Removed} {
		sort.Strings(*xs)
		if len(*xs) > limit {
			*xs = (*xs)[:limit]
		}
	}
	return w
}

// ---------------------------------------------------------------------------
// helpers
// ---------------------------------------------------------------------------

func readSmall(path string, max int64) ([]byte, error) {
	f, err := os.Open(path)
	if err != nil {
		return nil, err
	}
	defer f.Close()
	buf := make([]byte, 0, 4096)
	tmp := make([]byte, 32*1024)
	for int64(len(buf)) < max {
		n, err := f.Read(tmp)
		buf = append(buf, tmp[:n]...)
		if err != nil {
			break
		}
	}
	if int64(len(buf)) > max {
		buf = buf[:max]
	}
	return buf, nil
}

func truncate(s string, n int) string {
	if len(s) <= n {
		return s
	}
	// keep valid UTF-8: back off to a rune start
	cut := n
	for cut > 0 && (s[cut]&0xC0) == 0x80 {
		cut--
	}
	return s[:cut] + "…"
}

func sortedCopy(xs []string) []string {
	out := append([]string(nil), xs...)
	sort.Strings(out)
	return out
}
