package main

// Parsers for the plain-text files under /proc. Each parser takes the file's
// bytes so the tests can feed fixtures (testdata/); the readers that open the
// real files live next to the code that samples them.

import (
	"bufio"
	"bytes"
	"encoding/hex"
	"errors"
	"fmt"
	"net"
	"strconv"
	"strings"
)

// CPUTimes is the aggregate "cpu" line of /proc/stat, in clock ticks.
type CPUTimes struct {
	User, Nice, System, Idle, IOWait, IRQ, SoftIRQ, Steal uint64
}

func (c CPUTimes) Total() uint64 {
	return c.User + c.Nice + c.System + c.Idle + c.IOWait + c.IRQ + c.SoftIRQ + c.Steal
}

// Busy is everything but idle and iowait.
func (c CPUTimes) Busy() uint64 { return c.Total() - c.Idle - c.IOWait }

// ProcStat is what the sampler needs from /proc/stat.
type ProcStat struct {
	CPU      CPUTimes
	Cores    int
	BootTime int64
}

func ParseProcStat(b []byte) (ProcStat, error) {
	var ps ProcStat
	found := false
	sc := bufio.NewScanner(bytes.NewReader(b))
	sc.Buffer(make([]byte, 64*1024), 1024*1024)
	for sc.Scan() {
		f := strings.Fields(sc.Text())
		if len(f) == 0 {
			continue
		}
		switch {
		case f[0] == "cpu":
			if len(f) < 5 {
				return ps, errors.New("short cpu line")
			}
			v := make([]uint64, 8)
			for i := 0; i < 8 && i+1 < len(f); i++ {
				n, err := strconv.ParseUint(f[i+1], 10, 64)
				if err != nil {
					return ps, fmt.Errorf("cpu field %d: %w", i, err)
				}
				v[i] = n
			}
			ps.CPU = CPUTimes{v[0], v[1], v[2], v[3], v[4], v[5], v[6], v[7]}
			found = true
		case strings.HasPrefix(f[0], "cpu"):
			ps.Cores++
		case f[0] == "btime" && len(f) > 1:
			ps.BootTime, _ = strconv.ParseInt(f[1], 10, 64)
		}
	}
	if !found {
		return ps, errors.New("no cpu line")
	}
	if ps.Cores == 0 {
		ps.Cores = 1
	}
	return ps, nil
}

// CPUPercent between two readings of the aggregate line (0..100).
func CPUPercent(prev, cur CPUTimes) float64 {
	dt := float64(cur.Total()) - float64(prev.Total())
	if dt <= 0 {
		return 0
	}
	db := float64(cur.Busy()) - float64(prev.Busy())
	if db < 0 {
		db = 0
	}
	return clamp(100*db/dt, 0, 100)
}

// MemInfo in bytes.
type MemInfo struct {
	Total, Available, Free, Buffers, Cached, SwapTotal, SwapFree uint64
}

func (m MemInfo) UsedPct() float64 {
	if m.Total == 0 {
		return 0
	}
	return clamp(100*float64(m.Total-minU(m.Available, m.Total))/float64(m.Total), 0, 100)
}

func (m MemInfo) SwapUsed() uint64 {
	if m.SwapFree > m.SwapTotal {
		return 0
	}
	return m.SwapTotal - m.SwapFree
}

func (m MemInfo) SwapPct() float64 {
	if m.SwapTotal == 0 {
		return 0
	}
	return clamp(100*float64(m.SwapUsed())/float64(m.SwapTotal), 0, 100)
}

func ParseMemInfo(b []byte) (MemInfo, error) {
	var m MemInfo
	hasAvail := false
	sc := bufio.NewScanner(bytes.NewReader(b))
	for sc.Scan() {
		line := sc.Text()
		i := strings.IndexByte(line, ':')
		if i < 0 {
			continue
		}
		f := strings.Fields(line[i+1:])
		if len(f) == 0 {
			continue
		}
		n, err := strconv.ParseUint(f[0], 10, 64)
		if err != nil {
			continue
		}
		if len(f) > 1 && f[1] == "kB" {
			n *= 1024
		}
		switch line[:i] {
		case "MemTotal":
			m.Total = n
		case "MemAvailable":
			m.Available, hasAvail = n, true
		case "MemFree":
			m.Free = n
		case "Buffers":
			m.Buffers = n
		case "Cached":
			m.Cached = n
		case "SwapTotal":
			m.SwapTotal = n
		case "SwapFree":
			m.SwapFree = n
		}
	}
	if m.Total == 0 {
		return m, errors.New("no MemTotal")
	}
	if !hasAvail { // kernels before 3.14
		m.Available = m.Free + m.Buffers + m.Cached
	}
	return m, nil
}

// ParseVMStat returns the swap-in / swap-out page counters.
func ParseVMStat(b []byte) (pswpin, pswpout uint64) {
	sc := bufio.NewScanner(bytes.NewReader(b))
	for sc.Scan() {
		f := strings.Fields(sc.Text())
		if len(f) != 2 {
			continue
		}
		switch f[0] {
		case "pswpin":
			pswpin, _ = strconv.ParseUint(f[1], 10, 64)
		case "pswpout":
			pswpout, _ = strconv.ParseUint(f[1], 10, 64)
		}
	}
	return
}

func ParseLoadAvg(b []byte) ([3]float64, error) {
	var l [3]float64
	f := strings.Fields(string(b))
	if len(f) < 3 {
		return l, errors.New("short loadavg")
	}
	for i := 0; i < 3; i++ {
		v, err := strconv.ParseFloat(f[i], 64)
		if err != nil {
			return l, err
		}
		l[i] = v
	}
	return l, nil
}

func ParseUptime(b []byte) (float64, error) {
	f := strings.Fields(string(b))
	if len(f) < 1 {
		return 0, errors.New("empty uptime")
	}
	return strconv.ParseFloat(f[0], 64)
}

// NetCounters are the summed byte counters of the interfaces that carry real traffic.
type NetCounters struct{ RX, TX uint64 }

// virtualIface: loopback and the host side of containers / bridges (their traffic is
// counted again on the physical interface).
func virtualIface(name string) bool {
	if name == "lo" {
		return true
	}
	for _, p := range []string{"veth", "docker", "br-", "virbr", "cni", "flannel", "cali", "vxlan", "tun", "tap", "kube", "lxc", "vnet"} {
		if strings.HasPrefix(name, p) {
			return true
		}
	}
	return false
}

func ParseNetDev(b []byte) NetCounters {
	var n NetCounters
	sc := bufio.NewScanner(bytes.NewReader(b))
	for sc.Scan() {
		line := sc.Text()
		i := strings.IndexByte(line, ':')
		if i < 0 {
			continue
		}
		name := strings.TrimSpace(line[:i])
		if virtualIface(name) {
			continue
		}
		f := strings.Fields(line[i+1:])
		if len(f) < 9 {
			continue
		}
		rx, _ := strconv.ParseUint(f[0], 10, 64)
		tx, _ := strconv.ParseUint(f[8], 10, 64)
		n.RX += rx
		n.TX += tx
	}
	return n
}

// Socket is one listening socket from /proc/net/{tcp,tcp6,udp,udp6}.
type Socket struct {
	Proto string // tcp | udp (the address family is in Addr)
	Addr  string
	Port  int
	UID   int
	Inode uint64
}

// Local: bound to a loopback address, so not reachable from outside.
func (s Socket) Local() bool {
	ip := net.ParseIP(s.Addr)
	return ip != nil && ip.IsLoopback()
}

// ParseNetSockets reads one /proc/net/{tcp,udp}[6] table and returns the listening
// sockets: TCP in state LISTEN (0A), UDP unconnected (07) with no remote peer.
func ParseNetSockets(b []byte, proto string) ([]Socket, error) {
	var out []Socket
	sc := bufio.NewScanner(bytes.NewReader(b))
	first := true
	for sc.Scan() {
		if first { // header
			first = false
			continue
		}
		f := strings.Fields(sc.Text())
		if len(f) < 10 {
			continue
		}
		st := f[3]
		if proto == "tcp" && st != "0A" {
			continue
		}
		if proto == "udp" {
			if st != "07" {
				continue
			}
			if _, rport, err := splitHexAddr(f[2]); err != nil || rport != 0 {
				continue
			}
		}
		addr, port, err := splitHexAddr(f[1])
		if err != nil {
			return nil, err
		}
		uid, _ := strconv.Atoi(f[7])
		inode, _ := strconv.ParseUint(f[9], 10, 64)
		out = append(out, Socket{Proto: proto, Addr: addr, Port: port, UID: uid, Inode: inode})
	}
	return out, nil
}

// splitHexAddr decodes "0100007F:1F90" (IPv4) or the 32-hex-digit IPv6 form. The kernel
// prints each 32-bit word in host byte order (little endian on amd64 and arm64).
func splitHexAddr(s string) (string, int, error) {
	i := strings.IndexByte(s, ':')
	if i < 0 {
		return "", 0, fmt.Errorf("bad address %q", s)
	}
	raw, err := hex.DecodeString(s[:i])
	if err != nil || (len(raw) != 4 && len(raw) != 16) {
		return "", 0, fmt.Errorf("bad address %q", s)
	}
	port, err := strconv.ParseUint(s[i+1:], 16, 16)
	if err != nil {
		return "", 0, fmt.Errorf("bad port %q", s)
	}
	ip := make(net.IP, len(raw))
	for w := 0; w < len(raw); w += 4 {
		ip[w], ip[w+1], ip[w+2], ip[w+3] = raw[w+3], raw[w+2], raw[w+1], raw[w]
	}
	return ip.String(), int(port), nil
}

// PidStat is what we use from /proc/<pid>/stat.
type PidStat struct {
	Comm      string
	State     byte
	PPid      int
	UTime     uint64
	STime     uint64
	StartTime uint64 // clock ticks after boot
	RSSPages  int64
}

// ParsePidStat handles a comm that contains spaces or parentheses: it runs from the
// first '(' to the last ')'.
func ParsePidStat(b []byte) (PidStat, error) {
	var p PidStat
	s := string(b)
	l, r := strings.IndexByte(s, '('), strings.LastIndexByte(s, ')')
	if l < 0 || r < l {
		return p, errors.New("bad stat")
	}
	p.Comm = s[l+1 : r]
	f := strings.Fields(s[r+1:])
	// f[0]=state (field 3) ... utime is field 14 -> f[11], stime f[12], starttime field 22 -> f[19], rss field 24 -> f[21]
	if len(f) < 22 {
		return p, errors.New("short stat")
	}
	p.State = f[0][0]
	p.PPid, _ = strconv.Atoi(f[1])
	p.UTime, _ = strconv.ParseUint(f[11], 10, 64)
	p.STime, _ = strconv.ParseUint(f[12], 10, 64)
	p.StartTime, _ = strconv.ParseUint(f[19], 10, 64)
	p.RSSPages, _ = strconv.ParseInt(f[21], 10, 64)
	return p, nil
}

// ParseOSRelease returns PRETTY_NAME (or NAME VERSION_ID).
func ParseOSRelease(b []byte) string {
	vals := map[string]string{}
	sc := bufio.NewScanner(bytes.NewReader(b))
	for sc.Scan() {
		k, v, ok := strings.Cut(strings.TrimSpace(sc.Text()), "=")
		if !ok {
			continue
		}
		vals[k] = strings.Trim(v, `"'`)
	}
	if v := vals["PRETTY_NAME"]; v != "" {
		return v
	}
	return strings.TrimSpace(vals["NAME"] + " " + vals["VERSION_ID"])
}

// Users from /etc/passwd: uid -> name, plus the accounts that can log in.
type PasswdEntry struct {
	Name  string
	UID   int
	Home  string
	Shell string
}

func ParsePasswd(b []byte) []PasswdEntry {
	var out []PasswdEntry
	sc := bufio.NewScanner(bytes.NewReader(b))
	for sc.Scan() {
		f := strings.Split(sc.Text(), ":")
		if len(f) < 7 || strings.HasPrefix(f[0], "#") {
			continue
		}
		uid, err := strconv.Atoi(f[2])
		if err != nil {
			continue
		}
		out = append(out, PasswdEntry{Name: f[0], UID: uid, Home: f[5], Shell: f[6]})
	}
	return out
}

// LoginShell: an account whose shell lets someone log in.
func LoginShell(shell string) bool {
	if shell == "" {
		return false
	}
	base := shell[strings.LastIndexByte(shell, '/')+1:]
	switch base {
	case "nologin", "false", "sync", "halt", "shutdown", "true":
		return false
	}
	return true
}

func clamp(v, lo, hi float64) float64 {
	if v < lo {
		return lo
	}
	if v > hi {
		return hi
	}
	return v
}

func minU(a, b uint64) uint64 {
	if a < b {
		return a
	}
	return b
}
