package main

import (
	"bytes"
	"os"
	"path/filepath"
	"regexp"
	"sort"
	"strconv"
	"strings"
	"sync"
	"syscall"
	"time"
)

const clockTicks = 100 // USER_HZ: 100 on every Linux build for amd64 and arm64

var pageSize = int64(os.Getpagesize())

// Proc is one process as reported (top lists and signal evidence).
type Proc struct {
	PID  int     `json:"pid"`
	Name string  `json:"name"`
	User string  `json:"user"`
	Exe  string  `json:"exe,omitempty"`
	Cmd  string  `json:"cmd,omitempty"`
	CPU  float64 `json:"cpu"` // percent of one core over the last sample interval
	RSS  int64   `json:"rss"` // bytes
	// internal
	uid       int
	argv0     string
	exeKnown  bool
	startTime uint64
	kernel    bool
}

type procPrev struct {
	ticks uint64
	start uint64
}

// ProcSampler reads every /proc/<pid> once per sample and keeps the CPU ticks of the
// previous sample to turn them into a percentage.
type ProcSampler struct {
	root  string
	prev  map[int]procPrev
	prevT time.Time
	users *UserCache
}

func NewProcSampler(root string, users *UserCache) *ProcSampler {
	return &ProcSampler{root: root, prev: map[int]procPrev{}, users: users}
}

func (s *ProcSampler) Sample(now time.Time) []Proc {
	ents, err := os.ReadDir(s.root)
	if err != nil {
		return nil
	}
	elapsed := now.Sub(s.prevT).Seconds()
	next := make(map[int]procPrev, len(s.prev))
	out := make([]Proc, 0, len(ents))
	for _, e := range ents {
		pid, err := strconv.Atoi(e.Name())
		if err != nil || pid <= 0 {
			continue
		}
		dir := filepath.Join(s.root, e.Name())
		b, err := os.ReadFile(filepath.Join(dir, "stat"))
		if err != nil {
			continue
		}
		st, err := ParsePidStat(b)
		if err != nil {
			continue
		}
		p := Proc{PID: pid, Name: st.Comm, RSS: st.RSSPages * pageSize, startTime: st.StartTime}
		p.kernel = pid == 2 || st.PPid == 2
		if fi, err := os.Stat(dir); err == nil {
			p.uid = fileUID(fi)
		}
		p.User = s.users.Name(p.uid)
		ticks := st.UTime + st.STime
		if pp, ok := s.prev[pid]; ok && pp.start == st.StartTime && elapsed > 0 && ticks >= pp.ticks {
			p.CPU = round1(100 * float64(ticks-pp.ticks) / clockTicks / elapsed)
		}
		next[pid] = procPrev{ticks: ticks, start: st.StartTime}
		if !p.kernel {
			if exe, err := os.Readlink(filepath.Join(dir, "exe")); err == nil {
				p.Exe, p.exeKnown = exe, true
			}
			if cb, err := readSmall(filepath.Join(dir, "cmdline"), 4096); err == nil && len(cb) > 0 {
				args := bytes.Split(bytes.TrimRight(cb, "\x00"), []byte{0})
				if len(args) > 0 {
					p.argv0 = string(args[0])
				}
				p.Cmd = string(bytes.Join(args, []byte{' '}))
			}
		}
		out = append(out, p)
	}
	s.prev = next
	s.prevT = now
	return out
}

// Reported strips what the report must not carry raw: the command line is redacted
// and cut to 200 characters.
func (p Proc) Reported() Proc {
	q := p
	q.Cmd = truncate(Redact(p.Cmd), 200)
	q.Exe = truncate(p.Exe, 200)
	return q
}

var (
	redactKV   = regexp.MustCompile(`(?i)((?:pass(?:word|wd)?|pwd|secret|token|api[_-]?key|apikey|access[_-]?key|auth|credentials?|private[_-]?key)[=:])\S+`)
	redactFlag = regexp.MustCompile(`(?i)(--(?:pass(?:word|wd)?|secret|token|api-?key|auth)[ =])\S+`)
	redactURL  = regexp.MustCompile(`(://[^/\s:@]+:)[^@\s/]+@`)
	redactMy   = regexp.MustCompile(`(\s-p)[^\s]+`)
)

// Redact hides values that look like secrets in a command line: key=value pairs named
// like passwords or tokens, --password values, user:password@ in URLs and mysql-style -pSECRET.
func Redact(cmd string) string {
	cmd = redactKV.ReplaceAllString(cmd, "${1}***")
	cmd = redactFlag.ReplaceAllString(cmd, "${1}***")
	cmd = redactURL.ReplaceAllString(cmd, "${1}***@")
	if strings.Contains(cmd, "mysql") || strings.Contains(cmd, "mariadb") {
		cmd = redactMy.ReplaceAllString(cmd, "${1}***")
	}
	return cmd
}

// TopBy returns the n processes with the highest key (ties: lower pid first).
func TopBy(ps []Proc, n int, key func(Proc) float64) []Proc {
	cp := append([]Proc(nil), ps...)
	sort.SliceStable(cp, func(i, j int) bool {
		ki, kj := key(cp[i]), key(cp[j])
		if ki != kj {
			return ki > kj
		}
		return cp[i].PID < cp[j].PID
	})
	if len(cp) > n {
		cp = cp[:n]
	}
	out := make([]Proc, len(cp))
	for i, p := range cp {
		out[i] = p.Reported()
	}
	return out
}

// UserCache maps uids to names from /etc/passwd, reloading when the file changes.
type UserCache struct {
	path  string
	mu    sync.Mutex
	mtime time.Time
	names map[int]string
	all   []PasswdEntry
}

func NewUserCache(path string) *UserCache { return &UserCache{path: path, names: map[int]string{}} }

func (u *UserCache) refresh() {
	fi, err := os.Stat(u.path)
	if err != nil || fi.ModTime().Equal(u.mtime) {
		return
	}
	b, err := os.ReadFile(u.path)
	if err != nil {
		return
	}
	u.mtime = fi.ModTime()
	u.all = ParsePasswd(b)
	u.names = make(map[int]string, len(u.all))
	for _, e := range u.all {
		if _, ok := u.names[e.UID]; !ok {
			u.names[e.UID] = e.Name
		}
	}
}

func (u *UserCache) Name(uid int) string {
	u.mu.Lock()
	defer u.mu.Unlock()
	u.refresh()
	if n, ok := u.names[uid]; ok {
		return n
	}
	return strconv.Itoa(uid)
}

func (u *UserCache) All() []PasswdEntry {
	u.mu.Lock()
	defer u.mu.Unlock()
	u.refresh()
	return append([]PasswdEntry(nil), u.all...)
}

func fileUID(fi os.FileInfo) int {
	if st, ok := fi.Sys().(*syscall.Stat_t); ok {
		return int(st.Uid)
	}
	return -1
}

// SocketOwners maps socket inodes to the process holding them, by reading the fd links
// of every process it may read (all of them as root, only its own otherwise).
func SocketOwners(root string, inodes map[uint64]bool, procs []Proc) map[uint64]Proc {
	out := map[uint64]Proc{}
	if len(inodes) == 0 {
		return out
	}
	for _, p := range procs {
		if p.kernel {
			continue
		}
		fdDir := filepath.Join(root, strconv.Itoa(p.PID), "fd")
		ents, err := os.ReadDir(fdDir)
		if err != nil {
			continue
		}
		for _, e := range ents {
			l, err := os.Readlink(filepath.Join(fdDir, e.Name()))
			if err != nil || !strings.HasPrefix(l, "socket:[") {
				continue
			}
			ino, err := strconv.ParseUint(strings.TrimSuffix(strings.TrimPrefix(l, "socket:["), "]"), 10, 64)
			if err == nil && inodes[ino] {
				if _, seen := out[ino]; !seen {
					out[ino] = p
				}
			}
		}
		if len(out) == len(inodes) {
			break
		}
	}
	return out
}

func round1(v float64) float64 { return float64(int64(v*10+0.5)) / 10 }

func itoa(n int) string     { return strconv.Itoa(n) }
func itoa64(n int64) string { return strconv.FormatInt(n, 10) }
