package main

import (
	"bufio"
	"bytes"
	"errors"
	"fmt"
	"net"
	"net/url"
	"os"
	"strconv"
	"strings"
	"time"
)

const DefaultConfigPath = "/etc/approvalens-agent.yaml"

// Config is /etc/approvalens-agent.yaml. Only a small YAML subset is read: `key: value`
// lines and lists written as `- item` lines under a key (or `[a, b]`).
type Config struct {
	Token             string
	URL               string
	WebRoots          []string
	WebRootExclude    []string
	WebRootMaxFiles   int
	SampleInterval    time.Duration
	HeartbeatInterval time.Duration
	RollupInterval    time.Duration
	EventGap          time.Duration
	SetuidInterval    time.Duration
	WebRootInterval   time.Duration
	StateDir          string
	path              string
}

func DefaultConfig() Config {
	return Config{
		URL:               "https://approvalens.com",
		WebRootMaxFiles:   20000,
		SampleInterval:    60 * time.Second,
		HeartbeatInterval: 5 * time.Minute,
		RollupInterval:    15 * time.Minute,
		EventGap:          60 * time.Second,
		SetuidInterval:    6 * time.Hour,
		WebRootInterval:   15 * time.Minute,
		StateDir:          "/var/lib/approvalens-agent",
	}
}

func LoadConfig(path string) (Config, error) {
	c := DefaultConfig()
	c.path = path
	b, err := os.ReadFile(path)
	if err != nil {
		return c, err
	}
	if err := c.parse(b); err != nil {
		return c, fmt.Errorf("%s: %w", path, err)
	}
	if env := os.Getenv("APPROVALENS_AGENT_TOKEN"); env != "" {
		c.Token = env
	}
	return c, c.Validate()
}

func (c *Config) parse(b []byte) error {
	sc := bufio.NewScanner(bytes.NewReader(b))
	listKey := ""
	n := 0
	for sc.Scan() {
		n++
		raw := sc.Text()
		line := strings.TrimSpace(stripComment(raw))
		if line == "" {
			continue
		}
		if strings.HasPrefix(line, "- ") || line == "-" {
			if listKey == "" {
				return fmt.Errorf("line %d: list item without a key", n)
			}
			if err := c.set(listKey, unquote(strings.TrimSpace(strings.TrimPrefix(line, "-"))), true); err != nil {
				return fmt.Errorf("line %d: %w", n, err)
			}
			continue
		}
		k, v, ok := strings.Cut(line, ":")
		if !ok {
			return fmt.Errorf("line %d: expected `key: value`", n)
		}
		k, v = strings.TrimSpace(k), strings.TrimSpace(v)
		listKey = ""
		if v == "" {
			listKey = k
			continue
		}
		if strings.HasPrefix(v, "[") && strings.HasSuffix(v, "]") {
			for _, item := range strings.Split(strings.Trim(v, "[]"), ",") {
				if item = unquote(strings.TrimSpace(item)); item != "" {
					if err := c.set(k, item, true); err != nil {
						return fmt.Errorf("line %d: %w", n, err)
					}
				}
			}
			continue
		}
		if err := c.set(k, unquote(v), false); err != nil {
			return fmt.Errorf("line %d: %w", n, err)
		}
	}
	return sc.Err()
}

func stripComment(s string) string {
	inQ := byte(0)
	for i := 0; i < len(s); i++ {
		ch := s[i]
		switch {
		case inQ != 0 && ch == inQ:
			inQ = 0
		case inQ == 0 && (ch == '"' || ch == '\''):
			inQ = ch
		case inQ == 0 && ch == '#' && (i == 0 || s[i-1] == ' ' || s[i-1] == '\t'):
			return s[:i]
		}
	}
	return s
}

func unquote(s string) string {
	if len(s) >= 2 && (s[0] == '"' && s[len(s)-1] == '"' || s[0] == '\'' && s[len(s)-1] == '\'') {
		return s[1 : len(s)-1]
	}
	return s
}

func (c *Config) set(k, v string, list bool) error {
	switch k {
	case "token":
		c.Token = v
	case "url":
		c.URL = strings.TrimRight(v, "/")
	case "web_roots":
		c.WebRoots = append(c.WebRoots, v)
	case "web_root_exclude":
		c.WebRootExclude = append(c.WebRootExclude, v)
	case "web_root_max_files":
		n, err := strconv.Atoi(v)
		if err != nil || n < 1 {
			return fmt.Errorf("web_root_max_files: %q", v)
		}
		c.WebRootMaxFiles = n
	case "sample_interval", "heartbeat_interval", "rollup_interval", "event_min_gap", "setuid_interval", "web_root_interval":
		d, err := parseDuration(v)
		if err != nil {
			return fmt.Errorf("%s: %w", k, err)
		}
		switch k {
		case "sample_interval":
			c.SampleInterval = d
		case "heartbeat_interval":
			c.HeartbeatInterval = d
		case "rollup_interval":
			c.RollupInterval = d
		case "event_min_gap":
			c.EventGap = d
		case "setuid_interval":
			c.SetuidInterval = d
		case "web_root_interval":
			c.WebRootInterval = d
		}
	case "state_dir":
		c.StateDir = v
	default:
		return fmt.Errorf("unknown key %q", k)
	}
	_ = list
	return nil
}

// parseDuration accepts Go durations ("90s", "5m", "6h") or plain seconds.
func parseDuration(v string) (time.Duration, error) {
	if n, err := strconv.Atoi(v); err == nil {
		return time.Duration(n) * time.Second, nil
	}
	return time.ParseDuration(v)
}

func (c Config) Validate() error {
	if c.Token == "" {
		return errors.New("token is missing")
	}
	for _, r := range c.Token {
		if !(r >= 'a' && r <= 'z' || r >= 'A' && r <= 'Z' || r >= '0' && r <= '9' || r == '_' || r == '-') {
			return errors.New("token has unexpected characters")
		}
	}
	u, err := url.Parse(c.URL)
	if err != nil || u.Host == "" {
		return fmt.Errorf("url %q is not valid", c.URL)
	}
	// Plain HTTP only to this machine (local testing); anything else must be HTTPS.
	if u.Scheme != "https" {
		host := u.Hostname()
		ip := net.ParseIP(host)
		if u.Scheme != "http" || !(host == "localhost" || (ip != nil && ip.IsLoopback())) {
			return fmt.Errorf("url must be https:// (got %q)", c.URL)
		}
	}
	if c.SampleInterval < 5*time.Second || c.SampleInterval > 10*time.Minute {
		return errors.New("sample_interval must be between 5s and 10m")
	}
	if c.HeartbeatInterval < 30*time.Second || c.HeartbeatInterval > 10*time.Minute || c.HeartbeatInterval < c.SampleInterval {
		return errors.New("heartbeat_interval must be between 30s and 10m, and not shorter than sample_interval")
	}
	if c.RollupInterval < c.HeartbeatInterval || c.RollupInterval > time.Hour {
		return errors.New("rollup_interval must be between heartbeat_interval and 1h")
	}
	if c.EventGap < 10*time.Second || c.EventGap > 10*time.Minute {
		return errors.New("event_min_gap must be between 10s and 10m")
	}
	if c.SetuidInterval < time.Minute {
		return errors.New("setuid_interval must be at least 1m")
	}
	if c.WebRootInterval < time.Minute {
		return errors.New("web_root_interval must be at least 1m")
	}
	for _, r := range c.WebRoots {
		if !strings.HasPrefix(r, "/") {
			return fmt.Errorf("web_roots: %q is not an absolute path", r)
		}
	}
	return nil
}

// Endpoint the reports go to.
func (c Config) Endpoint() string { return c.URL + "/api/agent/v1/report" }
