package main

import (
	"encoding/json"
	"os"
	"path/filepath"
	"runtime"
	"sort"
	"strconv"
	"strings"
	"sync"
	"time"
)

const (
	hogPercent  = 90.0             // one process at this share of a core ...
	hogDuration = 15 * time.Minute // ... for this long, without a break
	fillWindow  = 6 * 3600         // disk fill projection looks at the last 6 hours
)

// ---------------------------------------------------------------------------
// Data shared by the messages (messages.go). Field names are short: they go over the network.
// ---------------------------------------------------------------------------

type Minute struct {
	T      int64   `json:"t"`       // start of the minute, unix seconds (UTC)
	CPU    float64 `json:"cpu"`     // % of all cores, average
	CPUMax float64 `json:"cpu_max"` // highest sample in the minute
	Load1  float64 `json:"load1"`
	Mem    float64 `json:"mem"` // % used (MemAvailable based)
	MemMax float64 `json:"mem_max"`
	Swap   float64 `json:"swap"`    // % of swap used
	SwapIO float64 `json:"swap_io"` // pages swapped in + out per second
	RX     float64 `json:"rx"`      // bytes per second
	TX     float64 `json:"tx"`
	Disk   float64 `json:"disk"` // fullest mount, % used
	N      int     `json:"n"`    // samples in the minute
}

type DiskOut struct {
	Mount     string   `json:"mount"`
	FS        string   `json:"fs"`
	Total     uint64   `json:"total"`
	Used      uint64   `json:"used"`
	Pct       float64  `json:"pct"`
	InodesPct float64  `json:"inodes_pct"`
	RateBPH   *float64 `json:"rate_bph,omitempty"`  // bytes per hour over the last 6 h (least squares)
	FullInH   *float64 `json:"full_in_h,omitempty"` // hours until full at that rate
}

type PortOut struct {
	Proto string `json:"proto"`
	Addr  string `json:"addr"`
	Port  int    `json:"port"`
	Local bool   `json:"local"` // bound to loopback only
	PID   int    `json:"pid,omitempty"`
	Name  string `json:"name,omitempty"`
	User  string `json:"user,omitempty"`
}

func (p PortOut) Key() string {
	k := p.Proto + "/" + strconv.Itoa(p.Port)
	if p.Local {
		k += "/local"
	}
	return k
}

type MemOut struct {
	Total     uint64  `json:"total"`
	Available uint64  `json:"available"`
	UsedPct   float64 `json:"used_pct"`
	SwapTotal uint64  `json:"swap_total"`
	SwapUsed  uint64  `json:"swap_used"`
	SwapPct   float64 `json:"swap_pct"`
}

type Current struct {
	At     int64      `json:"at"`
	CPU    float64    `json:"cpu"`
	Load   [3]float64 `json:"load"`
	Mem    MemOut     `json:"mem"`
	Disks  []DiskOut  `json:"disks"`
	RX     float64    `json:"rx"`
	TX     float64    `json:"tx"`
	Procs  int        `json:"procs"`
	TopCPU []Proc     `json:"top_cpu"`
	TopMem []Proc     `json:"top_mem"`
	Ports  []PortOut  `json:"ports"`
}

type Signal struct {
	Kind   string         `json:"kind"` // miner | tmp_exe | deleted_exe | cpu_hog
	Key    string         `json:"key"`  // stable across samples and restarts of the process
	At     int64          `json:"at"`   // first seen in this report window
	Detail map[string]any `json:"detail"`
}

type WebRootOut struct {
	Roots    []string    `json:"roots"`
	Files    int         `json:"files"`
	LastScan int64       `json:"last_scan,omitempty"`
	Baseline bool        `json:"baseline,omitempty"` // first scan: nothing to compare with yet
	Changes  *WebChanges `json:"changes,omitempty"`
}

type AgentInfo struct {
	Version   string   `json:"version"`
	Root      bool     `json:"root"`
	UID       int      `json:"uid"`
	Skipped   []string `json:"skipped"`
	SampleS   int      `json:"sample_s"`
	ReportS   int      `json:"report_s"`
	StartedAt int64    `json:"started_at"`
}

type HostInfo struct {
	Hostname string `json:"hostname"`
	OS       string `json:"os"`
	Kernel   string `json:"kernel"`
	Arch     string `json:"arch"`
	Cores    int    `json:"cores"`
	BootTime int64  `json:"boot_time"`
	UptimeS  int64  `json:"uptime_s"`
}

// ---------------------------------------------------------------------------
// Collector
// ---------------------------------------------------------------------------

type minuteAgg struct {
	n                                                          int
	cpu, cpuMax, load, mem, memMax, swap, swapIO, rx, tx, disk float64
}

type Collector struct {
	cfg     Config
	version string
	proc    string // /proc
	host    string // / (root of the files read)
	users   *UserCache
	procs   *ProcSampler
	started time.Time

	prevT       time.Time
	prevCPU     CPUTimes
	prevNet     NetCounters
	prevSwapIn  uint64
	prevSwapOut uint64
	havePrev    bool

	fill      map[string]*FillTracker
	minutes   map[int64]*minuteAgg
	sampleSig map[string]Signal    // process signals seen in the latest sample
	hog       map[string]time.Time // pid:start -> first sample over hogPercent
	rules     *Rules
	pending   []Event // tripped since the last event message
	lastSwap  float64

	portInodes map[uint64]bool
	portOwners map[uint64]Proc

	current  Current
	hostInfo HostInfo
	samples  int
	errors   int
	seq      int64

	mu       sync.Mutex // guards the scan results below (filled by the slow scans goroutine)
	setuid   []string   // nil until the first scan
	setuidAt int64
	web      *WebRootOut
	webPrev  map[string]string
	webPend  *WebChanges

	cronUnreadable []string
	sshUnreadable  []string
}

func NewCollector(cfg Config, version string) *Collector {
	users := NewUserCache("/etc/passwd")
	c := &Collector{
		cfg: cfg, version: version, proc: "/proc", host: "/", users: users,
		procs: NewProcSampler("/proc", users), started: time.Now(),
		fill: map[string]*FillTracker{}, minutes: map[int64]*minuteAgg{}, sampleSig: map[string]Signal{},
		hog: map[string]time.Time{}, portInodes: map[uint64]bool{}, portOwners: map[uint64]Proc{}, rules: NewRules(),
	}
	return c
}

func (c *Collector) read(name string) ([]byte, error) {
	return os.ReadFile(filepath.Join(c.proc, name))
}

// Sample takes one reading of everything that is sampled every interval.
func (c *Collector) Sample(now time.Time) {
	c.samples++
	ok := true
	var cur Current
	cur.At = now.Unix()

	b, err := c.read("stat")
	ps, perr := ParseProcStat(b)
	if err != nil || perr != nil {
		ok = false
	}
	b, _ = c.read("meminfo")
	mem, merr := ParseMemInfo(b)
	if merr != nil {
		ok = false
	}
	b, _ = c.read("vmstat")
	swIn, swOut := ParseVMStat(b)
	b, _ = c.read("net/dev")
	netc := ParseNetDev(b)
	b, _ = c.read("loadavg")
	cur.Load, _ = ParseLoadAvg(b)
	b, _ = c.read("uptime")
	up, _ := ParseUptime(b)

	elapsed := now.Sub(c.prevT).Seconds()
	var swapIO float64
	if c.havePrev && elapsed > 0 {
		cur.CPU = round1(CPUPercent(c.prevCPU, ps.CPU))
		if netc.RX >= c.prevNet.RX {
			cur.RX = float64(netc.RX-c.prevNet.RX) / elapsed
		}
		if netc.TX >= c.prevNet.TX {
			cur.TX = float64(netc.TX-c.prevNet.TX) / elapsed
		}
		if swIn >= c.prevSwapIn && swOut >= c.prevSwapOut {
			swapIO = float64(swIn-c.prevSwapIn+swOut-c.prevSwapOut) / elapsed
		}
	}
	cur.RX, cur.TX = float64(int64(cur.RX)), float64(int64(cur.TX))
	c.prevCPU, c.prevNet, c.prevSwapIn, c.prevSwapOut, c.prevT, c.havePrev = ps.CPU, netc, swIn, swOut, now, true

	cur.Mem = MemOut{Total: mem.Total, Available: mem.Available, UsedPct: round1(mem.UsedPct()), SwapTotal: mem.SwapTotal,
		SwapUsed: mem.SwapUsed(), SwapPct: round1(mem.SwapPct())}

	// Disks
	b, _ = c.read("self/mountinfo")
	maxDisk := 0.0
	seenMounts := map[string]bool{}
	for _, m := range ParseMountInfo(b) {
		d, err := statDisk(m.Point)
		if err != nil || d.Total == 0 {
			continue
		}
		seenMounts[m.Point] = true
		ft := c.fill[m.Point]
		if ft == nil {
			ft = NewFillTracker(fillWindow)
			c.fill[m.Point] = ft
		}
		ft.Add(now.Unix(), d.Used)
		o := DiskOut{Mount: m.Point, FS: m.FSType, Total: d.Total, Used: d.Used, Pct: round1(d.UsedPct), InodesPct: round1(d.InodesPct)}
		if r, ok := ft.Rate(); ok {
			r = float64(int64(r))
			o.RateBPH = &r
			if h, ok := ft.FullInHours(d.Avail); ok {
				h = round1(h)
				o.FullInH = &h
			}
		}
		if o.Pct > maxDisk {
			maxDisk = o.Pct
		}
		cur.Disks = append(cur.Disks, o)
	}
	for k := range c.fill {
		if !seenMounts[k] {
			delete(c.fill, k)
		}
	}

	// Processes
	procs := c.procs.Sample(now)
	cur.Procs = len(procs)
	cur.TopCPU = TopBy(procs, 10, func(p Proc) float64 { return p.CPU })
	cur.TopMem = TopBy(procs, 10, func(p Proc) float64 { return float64(p.RSS) })
	c.processSignals(now, procs)

	// Listening sockets
	cur.Ports = c.ports(procs)

	c.current = cur
	c.hostInfo = c.readHost(ps, up)
	c.lastSwap = swapIO

	// Local rules: what trips now goes out in the next event message (and only then).
	evs, _ := c.rules.Update(now, Reading{Mem: cur.Mem.UsedPct, SwapIO: swapIO, SwapPct: cur.Mem.SwapPct, Load1: cur.Load[0],
		Cores: ps.Cores, Disks: cur.Disks, TopCPU: cur.TopCPU, TopMem: cur.TopMem, Signals: c.sampleSig})
	c.pending = append(c.pending, evs...)
	if len(c.pending) > 100 {
		c.pending = c.pending[len(c.pending)-100:]
	}

	if !ok {
		c.errors++
		return
	}
	t := now.Unix() - now.Unix()%60
	a := c.minutes[t]
	if a == nil {
		a = &minuteAgg{}
		c.minutes[t] = a
	}
	a.n++
	a.cpu += cur.CPU
	a.cpuMax = maxF(a.cpuMax, cur.CPU)
	a.load += cur.Load[0]
	a.mem += cur.Mem.UsedPct
	a.memMax = maxF(a.memMax, cur.Mem.UsedPct)
	a.swap += cur.Mem.SwapPct
	a.swapIO += swapIO
	a.rx += cur.RX
	a.tx += cur.TX
	a.disk = maxF(a.disk, maxDisk)
}

func (c *Collector) readHost(ps ProcStat, up float64) HostInfo {
	h := HostInfo{Arch: runtime.GOARCH, Cores: ps.Cores, BootTime: ps.BootTime, UptimeS: int64(up)}
	if b, err := c.read("sys/kernel/hostname"); err == nil {
		h.Hostname = strings.TrimSpace(string(b))
	}
	if b, err := c.read("sys/kernel/osrelease"); err == nil {
		h.Kernel = strings.TrimSpace(string(b))
	}
	if b, err := os.ReadFile(filepath.Join(c.host, "etc/os-release")); err == nil {
		h.OS = ParseOSRelease(b)
	} else if b, err := os.ReadFile(filepath.Join(c.host, "usr/lib/os-release")); err == nil {
		h.OS = ParseOSRelease(b)
	}
	return h
}

func (c *Collector) processSignals(now time.Time, procs []Proc) {
	c.sampleSig = map[string]Signal{}
	live := map[string]bool{}
	for _, p := range procs {
		if p.kernel || p.PID == os.Getpid() {
			continue
		}
		// Crypto miner by name or arguments
		if why := sigs.MinerMatch(p.Name, p.argv0, p.Cmd); why != "" {
			c.addSignal(now, "miner", "miner:"+strings.ToLower(p.Name)+":"+exeOrArgv0(p), p, map[string]any{"match": why})
		}
		// Executable in a temporary directory, or deleted after start
		exe := strings.TrimSuffix(p.Exe, " (deleted)")
		switch {
		case p.exeKnown && InTmp(exe):
			c.addSignal(now, "tmp_exe", "tmp_exe:"+exe, p, map[string]any{"source": "exe"})
		case !p.exeKnown && strings.HasPrefix(p.argv0, "/") && InTmp(p.argv0):
			c.addSignal(now, "tmp_exe", "tmp_exe:"+p.argv0, p, map[string]any{"source": "argv0"})
		}
		if p.exeKnown && SuspiciousDeleted(p.Exe) {
			c.addSignal(now, "deleted_exe", "deleted_exe:"+exe, p, map[string]any{"source": "exe"})
		}
		// Sustained CPU by one process
		id := strconv.Itoa(p.PID) + ":" + strconv.FormatUint(p.startTime, 10)
		live[id] = true
		if p.CPU >= hogPercent {
			first, seen := c.hog[id]
			if !seen {
				c.hog[id] = now
				first = now
			}
			if now.Sub(first) >= hogDuration && !sigs.Busy[strings.ToLower(p.Name)] {
				c.addSignal(now, "cpu_hog", "cpu_hog:"+strings.ToLower(p.Name), p,
					map[string]any{"minutes": int(now.Sub(first).Minutes()), "since": first.Unix()})
			}
		} else {
			delete(c.hog, id)
		}
	}
	for id := range c.hog {
		if !live[id] {
			delete(c.hog, id)
		}
	}
}

func exeOrArgv0(p Proc) string {
	if p.exeKnown {
		return strings.TrimSuffix(p.Exe, " (deleted)")
	}
	return p.argv0
}

func (c *Collector) addSignal(now time.Time, kind, key string, p Proc, detail map[string]any) {
	key = truncate(key, 300)
	detail["process"] = p.Reported()
	c.sampleSig[key] = Signal{Kind: kind, Key: key, At: now.Unix(), Detail: detail}
}

// localPortRange: UDP sockets bound to a port in this range are clients (resolvers,
// time sync), not services.
func (c *Collector) localPortRange() (int, int) {
	b, err := c.read("sys/net/ipv4/ip_local_port_range")
	if err == nil {
		f := strings.Fields(string(b))
		if len(f) == 2 {
			lo, e1 := strconv.Atoi(f[0])
			hi, e2 := strconv.Atoi(f[1])
			if e1 == nil && e2 == nil {
				return lo, hi
			}
		}
	}
	return 32768, 60999
}

func (c *Collector) ports(procs []Proc) []PortOut {
	var socks []Socket
	for _, f := range [][2]string{{"net/tcp", "tcp"}, {"net/tcp6", "tcp"}, {"net/udp", "udp"}, {"net/udp6", "udp"}} {
		b, err := c.read(f[0])
		if err != nil {
			continue
		}
		s, _ := ParseNetSockets(b, f[1])
		socks = append(socks, s...)
	}
	lo, hi := c.localPortRange()
	inodes := map[uint64]bool{}
	var kept []Socket
	for _, s := range socks {
		if s.Proto == "udp" && s.Port >= lo && s.Port <= hi {
			continue
		}
		kept = append(kept, s)
		if s.Inode != 0 {
			inodes[s.Inode] = true
		}
	}
	if !sameSet(inodes, c.portInodes) {
		c.portOwners = SocketOwners(c.proc, inodes, procs)
		c.portInodes = inodes
	}
	seen := map[string]bool{}
	var out []PortOut
	for _, s := range kept {
		po := PortOut{Proto: s.Proto, Addr: s.Addr, Port: s.Port, Local: s.Local()}
		if p, ok := c.portOwners[s.Inode]; ok {
			po.PID, po.Name, po.User = p.PID, p.Name, p.User
		} else {
			po.User = c.users.Name(s.UID)
		}
		k := po.Key() + "|" + po.Addr
		if seen[k] {
			continue
		}
		seen[k] = true
		out = append(out, po)
	}
	sort.Slice(out, func(i, j int) bool {
		if out[i].Proto != out[j].Proto {
			return out[i].Proto < out[j].Proto
		}
		if out[i].Port != out[j].Port {
			return out[i].Port < out[j].Port
		}
		return out[i].Addr < out[j].Addr
	})
	if len(out) > 200 {
		out = out[:200]
	}
	return out
}

func sameSet(a, b map[uint64]bool) bool {
	if len(a) != len(b) {
		return false
	}
	for k := range a {
		if !b[k] {
			return false
		}
	}
	return true
}

// ---------------------------------------------------------------------------
// Slow scans (setuid every 6 h, web roots every 15 min), on their own goroutine
// ---------------------------------------------------------------------------

func (c *Collector) ScanSetuidNow() {
	items, truncated := ScanSetuid(setuidRoots, 200000)
	if len(items) > 500 {
		items = items[:500]
	}
	_ = truncated
	if items == nil {
		items = []string{}
	}
	c.mu.Lock()
	c.setuid, c.setuidAt = items, time.Now().Unix()
	c.mu.Unlock()
}

func (c *Collector) webStatePath() string { return filepath.Join(c.cfg.StateDir, "webroot.json") }

// LoadWebState reads the previous web root scan (kept across restarts).
func (c *Collector) LoadWebState() {
	b, err := os.ReadFile(c.webStatePath())
	if err != nil {
		return
	}
	var m map[string]string
	if json.Unmarshal(b, &m) == nil {
		c.mu.Lock()
		c.webPrev = m
		c.mu.Unlock()
	}
}

// ScanWebNow compares the web roots with the previous scan. persist=false (the `check`
// command) leaves the saved state alone.
func (c *Collector) ScanWebNow(persist bool) {
	if len(c.cfg.WebRoots) == 0 {
		return
	}
	files, truncated := ScanWebRoots(c.cfg.WebRoots, c.cfg.WebRootExclude, c.cfg.WebRootMaxFiles)
	c.mu.Lock()
	defer c.mu.Unlock()
	out := &WebRootOut{Roots: c.cfg.WebRoots, Files: len(files), LastScan: time.Now().Unix()}
	if c.webPrev == nil {
		out.Baseline = true
	} else {
		ch := DiffWebRoots(c.webPrev, files, 50)
		ch.Truncated = truncated
		if !ch.Empty() {
			c.webPend = mergeWeb(c.webPend, &ch)
		}
	}
	c.web = out
	c.webPrev = files
	if persist {
		if b, err := json.Marshal(files); err == nil {
			_ = writeFileAtomic(c.webStatePath(), b, 0o600)
		}
	}
}

func mergeWeb(a, b *WebChanges) *WebChanges {
	if a == nil {
		return b
	}
	m := *a
	m.Added = capList(dedupSorted(sortedCopy(append(m.Added, b.Added...))), 50)
	m.Changed = capList(dedupSorted(sortedCopy(append(m.Changed, b.Changed...))), 50)
	m.Removed = capList(dedupSorted(sortedCopy(append(m.Removed, b.Removed...))), 50)
	m.NAdded += b.NAdded
	m.NChanged += b.NChanged
	m.NRemoved += b.NRemoved
	m.Files = b.Files
	m.Truncated = m.Truncated || b.Truncated
	return &m
}

func capList(xs []string, n int) []string {
	if len(xs) > n {
		return xs[:n]
	}
	return xs
}

// ---------------------------------------------------------------------------
// Building the report
// ---------------------------------------------------------------------------

// Skipped names the checks this agent cannot do with its permissions.
func (c *Collector) Skipped() []string {
	var s []string
	if os.Geteuid() != 0 {
		s = append(s, "exe_paths_other_users", "port_owners_other_users")
	}
	if len(c.cronUnreadable) > 0 {
		s = append(s, "crontabs_unreadable")
	}
	if len(c.sshUnreadable) > 0 {
		s = append(s, "ssh_keys_unreadable")
	}
	if len(c.cfg.WebRoots) == 0 {
		s = append(s, "web_roots_not_configured")
	}
	return s
}

func maxF(a, b float64) float64 {
	if a > b {
		return a
	}
	return b
}

func round2(v float64) float64 { return float64(int64(v*100+0.5)) / 100 }

func writeFileAtomic(path string, b []byte, mode os.FileMode) error {
	if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
		return err
	}
	tmp := path + ".tmp"
	if err := os.WriteFile(tmp, b, mode); err != nil {
		return err
	}
	return os.Rename(tmp, path)
}
