package main

import (
	"context"
	"encoding/json"
	"fmt"
	"io"
	"net/http"
	"os"
	"path/filepath"
	"runtime"
	"sort"
	"strconv"
	"strings"
	"sync"
	"time"
)

type Resources struct {
	CPUTicks  uint64 `json:"cpu_ticks"`
	RSSKB     uint64 `json:"rss_kb"`
	PSSKB     uint64 `json:"pss_kb"`
	Processes int    `json:"processes"`
}

func resources(pid int) Resources {
	result := Resources{}
	todo := []int{pid}
	seen := map[int]bool{}
	for len(todo) > 0 {
		p := todo[0]
		todo = todo[1:]
		if p <= 0 || seen[p] {
			continue
		}
		seen[p] = true
		raw, err := os.ReadFile(fmt.Sprintf("/proc/%d/stat", p))
		if err != nil {
			continue
		}
		fields := strings.Fields(string(raw)[strings.LastIndex(string(raw), ")")+1:])
		if len(fields) < 22 {
			continue
		}
		ut, _ := strconv.ParseUint(fields[11], 10, 64)
		st, _ := strconv.ParseUint(fields[12], 10, 64)
		result.CPUTicks += ut + st
		result.Processes++
		smaps, _ := os.ReadFile(fmt.Sprintf("/proc/%d/smaps_rollup", p))
		for _, line := range strings.Split(string(smaps), "\n") {
			f := strings.Fields(line)
			if len(f) < 2 {
				continue
			}
			v, _ := strconv.ParseUint(f[1], 10, 64)
			if f[0] == "Rss:" {
				result.RSSKB += v
			}
			if f[0] == "Pss:" {
				result.PSSKB += v
			}
		}
		children, _ := os.ReadFile(fmt.Sprintf("/proc/%d/task/%d/children", p, p))
		for _, s := range strings.Fields(string(children)) {
			id, _ := strconv.Atoi(s)
			todo = append(todo, id)
		}
	}
	return result
}

type Sample struct {
	Workload    string    `json:"workload"`
	Repeat      int       `json:"repeat"`
	Concurrency int       `json:"concurrency"`
	Requests    int       `json:"requests"`
	Errors      int       `json:"errors"`
	Seconds     float64   `json:"seconds"`
	RPS         float64   `json:"rps"`
	P50MS       float64   `json:"p50_ms"`
	P95MS       float64   `json:"p95_ms"`
	Bytes       int64     `json:"bytes"`
	Latencies   []float64 `json:"latencies_ms"`
	Before      Resources `json:"before"`
	After       Resources `json:"after"`
	Peak        Resources `json:"peak"`
}
type BenchReport struct {
	Preparation string    `json:"preparation"`
	Target      string    `json:"target"`
	Dataset     string    `json:"dataset"`
	Started     time.Time `json:"started"`
	Go          string    `json:"harness_go"`
	CPUCount    int       `json:"logical_cpus"`
	ClockTicks  int       `json:"clock_ticks_per_second"`
	Samples     []Sample  `json:"samples"`
}

func benchmark(ctx context.Context, dsn, url, output string, n, concurrency, repeats, pid int, filter string, full bool) error {
	if n < 1 || concurrency < 1 || repeats < 1 || pid < 1 {
		return fmt.Errorf("positive requests, concurrency, repeats and a server PID are required")
	}
	if resources(pid).Processes == 0 {
		return fmt.Errorf("cannot inspect server PID %d in this process namespace", pid)
	}
	db, err := openDB(ctx, dsn)
	if err != nil {
		return err
	}
	defer db.Close(ctx)
	if err = guard(ctx, db); err != nil {
		return err
	}
	workloads := []Case{{"posts", "GET", "/api/posts", ""}, {"search", "GET", "/api/posts/search?q=python", ""}, {"tag", "GET", "/api/posts/by-tag/python", ""}, {"detail", "GET", "/api/posts/1", ""}, {"user", "GET", "/api/users/1", ""}, {"email", "GET", "/api/users/find?email=alice%40example.com", ""}, {"create", "POST", "/api/posts", `{"author_id":1,"title":"Bench","body":"Body","tag_slugs":["python"]}`}, {"comment", "POST", "/api/posts/1/comments", `{"author_id":1,"body":"Bench"}`}}
	if full {
		workloads[5].Path = "/api/users/find?email=user00000%40example.com"
	}
	client := &http.Client{Timeout: 120 * time.Second, Transport: &http.Transport{MaxIdleConns: concurrency, MaxIdleConnsPerHost: concurrency, MaxConnsPerHost: concurrency, DisableCompression: true}}
	defer client.CloseIdleConnections()
	report := BenchReport{Target: url, Dataset: fixtureVersion, Started: time.Now().UTC(), Go: runtime.Version(), CPUCount: runtime.NumCPU(), ClockTicks: 100, Samples: []Sample{}}
	if full {
		report.Dataset = "synthetic-full-v1"
	}
	report.Preparation = "vacuum-analyze-v2"
	persist := func() error {
		if err := os.MkdirAll(filepath.Dir(output), 0755); err != nil {
			return err
		}
		b, _ := json.MarshalIndent(report, "", "  ")
		return os.WriteFile(output, append(b, '\n'), 0644)
	}
	for _, tc := range workloads {
		if filter != "" && filter != tc.Name {
			continue
		}
		for rep := 0; rep < repeats; rep++ {
			if !full {
				if err = fixture(ctx, db); err != nil {
					return err
				}
			} else {
				// Restore only rows touched by benchmark writes; preserve the full fixture.
				if _, err = db.Exec(ctx, `DELETE FROM blog_comment WHERE id>500000; DELETE FROM blog_post_tags WHERE post_id>100000; DELETE FROM blog_post WHERE id>100000; SELECT setval(pg_get_serial_sequence('blog_post','id'),100000); SELECT setval(pg_get_serial_sequence('blog_comment','id'),500000); UPDATE blog_post p SET view_count=s.view_count,updated_at=s.updated_at FROM harness_seed_state s WHERE p.id=s.id`); err != nil {
					return err
				}
			}
			if err = prepareTables(ctx, db); err != nil {
				return err
			}
			// Warm up connections and code paths; reset counters before measurement.
			request := func() (float64, int64, bool) {
				t := time.Now()
				req, _ := http.NewRequestWithContext(ctx, tc.Method, url+tc.Path, strings.NewReader(tc.Body))
				req.Header.Set("Content-Type", "application/json")
				resp, err := client.Do(req)
				if err != nil {
					return float64(time.Since(t)) / 1e6, 0, true
				}
				size, err := io.Copy(io.Discard, resp.Body)
				resp.Body.Close()
				return float64(time.Since(t)) / 1e6, size, err != nil || resp.StatusCode != 200
			}
			for range min(3, n) {
				request()
			}
			if !full {
				if err = fixture(ctx, db); err != nil {
					return err
				}
			}
			if err = prepareTables(ctx, db); err != nil {
				return err
			}
			sample := Sample{Workload: tc.Name, Repeat: rep, Concurrency: concurrency, Requests: n, Latencies: make([]float64, n), Before: resources(pid)}
			stop := make(chan struct{})
			done := make(chan Resources, 1)
			go func() {
				peak := resources(pid)
				ticker := time.NewTicker(25 * time.Millisecond)
				defer ticker.Stop()
				for {
					select {
					case <-ticker.C:
						r := resources(pid)
						peak.RSSKB = max(peak.RSSKB, r.RSSKB)
						peak.PSSKB = max(peak.PSSKB, r.PSSKB)
						peak.Processes = max(peak.Processes, r.Processes)
					case <-stop:
						done <- peak
						return
					}
				}
			}()
			jobs := make(chan int)
			var wg sync.WaitGroup
			var mu sync.Mutex
			t := time.Now()
			for range concurrency {
				wg.Add(1)
				go func() {
					defer wg.Done()
					for i := range jobs {
						lat, size, failed := request()
						sample.Latencies[i] = lat
						mu.Lock()
						sample.Bytes += size
						if failed {
							sample.Errors++
						}
						mu.Unlock()
					}
				}()
			}
			for i := 0; i < n; i++ {
				jobs <- i
			}
			close(jobs)
			wg.Wait()
			sample.Seconds = time.Since(t).Seconds()
			sample.After = resources(pid)
			close(stop)
			sample.Peak = <-done
			sample.Peak.RSSKB = max(sample.Peak.RSSKB, sample.After.RSSKB)
			sample.Peak.PSSKB = max(sample.Peak.PSSKB, sample.After.PSSKB)
			sorted := append([]float64(nil), sample.Latencies...)
			sort.Float64s(sorted)
			sample.P50MS = sorted[(n-1)/2]
			sample.P95MS = sorted[min(n-1, int(float64(n)*0.95))]
			sample.RPS = float64(n-sample.Errors) / sample.Seconds
			report.Samples = append(report.Samples, sample)
			if err = persist(); err != nil {
				return err
			}
			fmt.Printf("%s rep=%d c=%d p50=%.3fms p95=%.3fms rps=%.1f errors=%d pss=%dKiB\n", tc.Name, rep, concurrency, sample.P50MS, sample.P95MS, sample.RPS, sample.Errors, sample.Peak.PSSKB)
		}
	}
	return persist()
}
