package main

import (
	"bytes"
	"compress/gzip"
	"encoding/json"
	"errors"
	"fmt"
	"io"
	"log"
	"net/http"
	"os"
	"path/filepath"
	"sort"
	"strconv"
	"strings"
	"sync"
	"sync/atomic"
	"time"
)

const spoolMaxAge = time.Hour // messages older than this are dropped while offline

// Sender posts reports (gzip JSON) and keeps the ones that could not be delivered in a
// spool (memory plus files in the state directory) for at most an hour. It only ever
// sends: the answer's body is discarded and nothing in it is acted upon.
type Sender struct {
	cfg      Config
	version  string
	client   *http.Client
	mu       sync.Mutex
	flushing sync.Mutex
	// OnResync is called when the server answers 205 Reset Content: its copy of the inventory
	// differs from ours, so the agent sends the full lists once. The status code is the only
	// thing read from an answer.
	OnResync  func()
	sentBytes atomic.Int64
	sentMsgs  atomic.Int64
	queue     []spooled
	lastErr   time.Time
}

type spooled struct {
	at   time.Time
	body []byte // gzip JSON
	file string
}

func NewSender(cfg Config, version string) *Sender {
	return &Sender{cfg: cfg, version: version, client: &http.Client{
		Timeout: 20 * time.Second,
		// Redirects are not followed: the report goes to the configured URL or nowhere.
		CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse },
	}}
}

func (s *Sender) spoolDir() string { return filepath.Join(s.cfg.StateDir, "spool") }

// LoadSpool picks up reports a previous run could not send.
func (s *Sender) LoadSpool() {
	ents, err := os.ReadDir(s.spoolDir())
	if err != nil {
		return
	}
	s.mu.Lock()
	defer s.mu.Unlock()
	for _, e := range ents {
		name := e.Name()
		ts, err := strconv.ParseInt(strings.TrimSuffix(name, ".json.gz"), 10, 64)
		p := filepath.Join(s.spoolDir(), name)
		if err != nil || time.Since(time.Unix(0, ts)) > spoolMaxAge {
			_ = os.Remove(p)
			continue
		}
		b, err := os.ReadFile(p)
		if err != nil {
			continue
		}
		s.queue = append(s.queue, spooled{at: time.Unix(0, ts), body: b, file: p})
	}
	sort.Slice(s.queue, func(i, j int) bool { return s.queue[i].at.Before(s.queue[j].at) })
}

func (s *Sender) Pending() int {
	s.mu.Lock()
	defer s.mu.Unlock()
	return len(s.queue)
}

func Encode(r any) ([]byte, error) {
	raw, err := json.Marshal(r)
	if err != nil {
		return nil, err
	}
	var buf bytes.Buffer
	zw, _ := gzip.NewWriterLevel(&buf, gzip.BestCompression)
	if _, err := zw.Write(raw); err != nil {
		return nil, err
	}
	if err := zw.Close(); err != nil {
		return nil, err
	}
	return buf.Bytes(), nil
}

// Enqueue adds a message and tries to deliver everything queued, oldest first.
func (s *Sender) Enqueue(r any) {
	body, err := Encode(r)
	if err != nil {
		log.Printf("encode report: %v", err)
		return
	}
	now := time.Now()
	item := spooled{at: now, body: body}
	if err := os.MkdirAll(s.spoolDir(), 0o700); err == nil {
		p := filepath.Join(s.spoolDir(), strconv.FormatInt(now.UnixNano(), 10)+".json.gz")
		if os.WriteFile(p, body, 0o600) == nil {
			item.file = p
		}
	}
	s.mu.Lock()
	s.queue = append(s.queue, item)
	s.mu.Unlock()
	go s.Flush()
}

var errRetry = errors.New("retry later")

// retryWaits: one try at once, then after 5 s and 20 s.
var retryWaits = []time.Duration{0, 5 * time.Second, 20 * time.Second}

// Flush sends the queue in order, on its own goroutine (sampling never waits for the
// network). A network error or a 5xx / 429 answer is retried with backoff (5 s, 20 s),
// then left for the next report; anything older than an hour is dropped.
func (s *Sender) Flush() {
	if !s.flushing.TryLock() {
		return // another flush is running and will pick the new report up
	}
	defer s.flushing.Unlock()
	for {
		s.mu.Lock()
		for len(s.queue) > 0 && time.Since(s.queue[0].at) > spoolMaxAge {
			s.drop(s.queue[0])
			s.queue = s.queue[1:]
		}
		if len(s.queue) == 0 {
			s.mu.Unlock()
			return
		}
		item := s.queue[0]
		s.mu.Unlock()
		var err error
		for _, wait := range retryWaits {
			time.Sleep(wait)
			if err = s.post(item.body); err == nil || !errors.Is(err, errRetry) {
				break
			}
		}
		s.mu.Lock()
		if err != nil && errors.Is(err, errRetry) {
			if time.Since(s.lastErr) > 15*time.Minute {
				log.Printf("report not delivered, kept in the spool (%d queued): %v", len(s.queue), err)
				s.lastErr = time.Now()
			}
			s.mu.Unlock()
			return
		}
		if err != nil {
			log.Printf("report dropped: %v", err)
		}
		s.drop(item)
		if len(s.queue) > 0 && s.queue[0].at.Equal(item.at) {
			s.queue = s.queue[1:]
		}
		s.mu.Unlock()
	}
}

func (s *Sender) drop(item spooled) {
	if item.file != "" {
		_ = os.Remove(item.file)
	}
}

func (s *Sender) post(body []byte) error {
	req, err := http.NewRequest(http.MethodPost, s.cfg.Endpoint(), bytes.NewReader(body))
	if err != nil {
		return err
	}
	req.Header.Set("Authorization", "Bearer "+s.cfg.Token)
	req.Header.Set("Content-Type", "application/json")
	req.Header.Set("Content-Encoding", "gzip")
	req.Header.Set("User-Agent", "approvalens-agent/"+s.version)
	res, err := s.client.Do(req)
	if err != nil {
		return fmt.Errorf("%w: %v", errRetry, err)
	}
	_, _ = io.Copy(io.Discard, io.LimitReader(res.Body, 4096))
	res.Body.Close()
	switch {
	case res.StatusCode >= 200 && res.StatusCode < 300:
		s.sentBytes.Add(int64(len(body)))
		s.sentMsgs.Add(1)
		if res.StatusCode == http.StatusResetContent && s.OnResync != nil {
			s.OnResync()
		}
		return nil
	case res.StatusCode == 401 || res.StatusCode == 403:
		return fmt.Errorf("token rejected (HTTP %d): check the token in %s", res.StatusCode, s.cfg.path)
	case res.StatusCode == 413 || res.StatusCode == 400:
		return fmt.Errorf("report refused (HTTP %d)", res.StatusCode)
	default:
		return fmt.Errorf("%w: HTTP %d", errRetry, res.StatusCode)
	}
}

// Sent returns the messages and bytes (compressed, as posted) delivered since start.
func (s *Sender) Sent() (msgs, bytes int64) { return s.sentMsgs.Load(), s.sentBytes.Load() }
