package main

import (
	"bufio"
	"bytes"
	"sort"
	"strconv"
	"strings"
	"syscall"
)

// Mount is one real filesystem from /proc/self/mountinfo.
type Mount struct {
	Point  string
	FSType string
	Source string
}

// Filesystems that hold data. Pseudo filesystems, tmpfs, overlays, snap images
// (squashfs, always 100% full) and network filesystems (statfs can hang) are left out.
var realFS = map[string]bool{
	"ext2": true, "ext3": true, "ext4": true, "xfs": true, "btrfs": true, "zfs": true, "f2fs": true,
	"jfs": true, "reiserfs": true, "vfat": true, "exfat": true, "ntfs": true, "ntfs3": true,
	"bcachefs": true, "hfsplus": true, "ufs": true,
}

// unescapeMount decodes the octal escapes mountinfo uses (\040 = space).
func unescapeMount(s string) string {
	if !strings.Contains(s, `\`) {
		return s
	}
	var b strings.Builder
	for i := 0; i < len(s); i++ {
		if s[i] == '\\' && i+3 < len(s) {
			if n, err := strconv.ParseUint(s[i+1:i+4], 8, 8); err == nil {
				b.WriteByte(byte(n))
				i += 3
				continue
			}
		}
		b.WriteByte(s[i])
	}
	return b.String()
}

// ParseMountInfo returns one mount per real filesystem: a device mounted at several
// places (bind mounts, btrfs subvolumes, the read-only view systemd gives the agent)
// is listed once, at its shortest mount point.
func ParseMountInfo(b []byte) []Mount {
	best := map[string]Mount{}
	sc := bufio.NewScanner(bytes.NewReader(b))
	sc.Buffer(make([]byte, 64*1024), 1024*1024)
	for sc.Scan() {
		line := sc.Text()
		pre, post, ok := strings.Cut(line, " - ")
		if !ok {
			continue
		}
		pf, qf := strings.Fields(pre), strings.Fields(post)
		if len(pf) < 5 || len(qf) < 2 {
			continue
		}
		fstype, source := qf[0], qf[1]
		if !realFS[fstype] {
			continue
		}
		point := unescapeMount(pf[4])
		if skipMountPoint(point) {
			continue
		}
		key := source
		if !strings.HasPrefix(source, "/dev/") { // zfs datasets, odd sources: the device numbers
			key = pf[2] + "|" + source
		}
		cur, seen := best[key]
		if !seen || len(point) < len(cur.Point) {
			best[key] = Mount{Point: point, FSType: fstype, Source: source}
		}
	}
	out := make([]Mount, 0, len(best))
	for _, m := range best {
		out = append(out, m)
	}
	sort.Slice(out, func(i, j int) bool { return out[i].Point < out[j].Point })
	return out
}

func skipMountPoint(p string) bool {
	for _, pre := range []string{"/proc", "/sys", "/dev", "/run", "/snap", "/var/lib/docker", "/var/lib/containers", "/var/snap"} {
		if p == pre || strings.HasPrefix(p, pre+"/") {
			return true
		}
	}
	return false
}

// DiskUsage of one mount, in bytes and inodes.
type DiskUsage struct {
	Total, Used, Avail uint64
	Inodes, InodesFree uint64
	UsedPct, InodesPct float64
}

func statDisk(path string) (DiskUsage, error) {
	var st syscall.Statfs_t
	if err := syscall.Statfs(path, &st); err != nil {
		return DiskUsage{}, err
	}
	bs := uint64(st.Bsize)
	d := DiskUsage{Total: st.Blocks * bs, Avail: st.Bavail * bs, Inodes: st.Files, InodesFree: st.Ffree}
	free := st.Bfree * bs
	if d.Total >= free {
		d.Used = d.Total - free
	}
	// As df: used / (used + available to unprivileged users), so root's reserve counts as full.
	if den := d.Used + d.Avail; den > 0 {
		d.UsedPct = 100 * float64(d.Used) / float64(den)
	}
	if d.Inodes > 0 && d.Inodes >= d.InodesFree {
		d.InodesPct = 100 * float64(d.Inodes-d.InodesFree) / float64(d.Inodes)
	}
	return d, nil
}

// FillTracker keeps the used bytes of one mount over the last window (6 h) and projects
// when the disk is full if the trend goes on.
type FillTracker struct {
	Window int64 // seconds
	ts     []int64
	used   []float64
}

func NewFillTracker(window int64) *FillTracker { return &FillTracker{Window: window} }

func (f *FillTracker) Add(t int64, used uint64) {
	f.ts = append(f.ts, t)
	f.used = append(f.used, float64(used))
	cut := 0
	for cut < len(f.ts) && f.ts[cut] < t-f.Window {
		cut++
	}
	if cut > 0 {
		f.ts = append(f.ts[:0], f.ts[cut:]...)
		f.used = append(f.used[:0], f.used[cut:]...)
	}
}

// Rate is the least-squares slope in bytes per hour (0 without enough data: at least
// 30 minutes and 10 samples).
func (f *FillTracker) Rate() (float64, bool) {
	n := len(f.ts)
	if n < 10 || f.ts[n-1]-f.ts[0] < 1800 {
		return 0, false
	}
	t0 := float64(f.ts[0])
	var sx, sy, sxx, sxy float64
	for i := range f.ts {
		x := (float64(f.ts[i]) - t0) / 3600
		y := f.used[i]
		sx += x
		sy += y
		sxx += x * x
		sxy += x * y
	}
	fn := float64(n)
	den := fn*sxx - sx*sx
	if den == 0 {
		return 0, false
	}
	return (fn*sxy - sx*sy) / den, true
}

// FullInHours projects when `avail` bytes are used up at the current rate. ok=false
// when the disk is not filling (or not fast enough to matter: under 1 MB per hour).
func (f *FillTracker) FullInHours(avail uint64) (float64, bool) {
	rate, ok := f.Rate()
	if !ok || rate < 1<<20 {
		return 0, false
	}
	return float64(avail) / rate, true
}
