package config

import (
	"bytes"
	"os"
	"path/filepath"
	"strings"
	"sync/atomic"
	"testing"
	"time"

	"github.com/rs/zerolog"
)

const validYAML = `
dev_mode: true
listen:
  port: 8443
  tailscale_interface_pattern: "tailscale*"
  read_timeout_seconds: 15
  write_timeout_seconds: 30
database:
  sqlite_path: "./data/command-center.db"
  duckdb_path: "./data/command-center.duckdb"
secrets:
  age_identity_file: ""
logging:
  level: "info"
  format: "json"
`

func TestLoadFromBytesValid(t *testing.T) {
	cfg, err := LoadFromBytes([]byte(validYAML))
	if err != nil {
		t.Fatalf("LoadFromBytes: %v", err)
	}
	if !cfg.DevMode {
		t.Errorf("dev_mode not parsed")
	}
	if cfg.Listen.Port != 8443 {
		t.Errorf("port: got %d", cfg.Listen.Port)
	}
	if cfg.Database.SQLitePath != "./data/command-center.db" {
		t.Errorf("sqlite path: got %q", cfg.Database.SQLitePath)
	}
}

func TestLoadFromBytesAppliesDefaults(t *testing.T) {
	cfg, err := LoadFromBytes([]byte(`
dev_mode: true
listen:
  port: 9000
database:
  sqlite_path: "./x.db"
  duckdb_path: "./x.duckdb"
`))
	if err != nil {
		t.Fatalf("LoadFromBytes: %v", err)
	}
	if cfg.Listen.TailscaleInterfacePattern != "tailscale*" {
		t.Errorf("tailscale pattern default not applied: %q", cfg.Listen.TailscaleInterfacePattern)
	}
	if cfg.Logging.Level != "info" {
		t.Errorf("logging.level default not applied: %q", cfg.Logging.Level)
	}
	if cfg.Auth.SessionLifetimeSeconds == 0 {
		t.Errorf("auth.session_lifetime_seconds default not applied")
	}
	// ResolvedAuth in dev_mode supplies localhost RPID + 127.0.0.1 origin.
	a := cfg.ResolvedAuth()
	if a.RPID != "localhost" {
		t.Errorf("dev_mode auth.rp_id: %q, want localhost", a.RPID)
	}
	if len(a.RPOrigins) == 0 {
		t.Errorf("dev_mode auth.rp_origins not auto-populated")
	}
}

func TestAuthValidationRequiresRPIDInProduction(t *testing.T) {
	_, err := LoadFromBytes([]byte(`
dev_mode: false
listen:
  port: 8443
database:
  sqlite_path: a
  duckdb_path: b
`))
	if err == nil {
		t.Fatal("expected validation error: prod requires auth.rp_id")
	}
}

func TestValidationRejectsBadYAML(t *testing.T) {
	cases := []struct {
		name string
		yaml string
		want string
	}{
		{"port too low", `listen: {port: 0}` + "\n" + `database: {sqlite_path: a, duckdb_path: b}`, "listen.port"},
		{"port too high", `listen: {port: 999999}` + "\n" + `database: {sqlite_path: a, duckdb_path: b}`, "listen.port"},
		{"empty sqlite path", `listen: {port: 80}` + "\n" + `database: {sqlite_path: "", duckdb_path: b}`, "sqlite_path"},
		{"empty duckdb path", `listen: {port: 80}` + "\n" + `database: {sqlite_path: a, duckdb_path: ""}`, "duckdb_path"},
		{"bad log level", `listen: {port: 80}` + "\n" + `database: {sqlite_path: a, duckdb_path: b}` + "\n" + `logging: {level: shouty}`, "logging.level"},
		{"empty interface pattern", `listen: {port: 80, tailscale_interface_pattern: ""}` + "\n" + `database: {sqlite_path: a, duckdb_path: b}`, "tailscale_interface_pattern"},
	}
	for _, tc := range cases {
		t.Run(tc.name, func(t *testing.T) {
			_, err := LoadFromBytes([]byte(tc.yaml))
			if err == nil {
				t.Fatalf("expected error containing %q", tc.want)
			}
			if !strings.Contains(err.Error(), tc.want) {
				t.Errorf("error %q missing expected substring %q", err.Error(), tc.want)
			}
		})
	}
}

func TestValidationRejectsMalformedYAML(t *testing.T) {
	_, err := LoadFromBytes([]byte("dev_mode: not_a_bool\nlisten: {port: 80}"))
	if err == nil {
		t.Fatalf("expected yaml parse error")
	}
}

func writeConfig(t *testing.T, dir, body string) {
	t.Helper()
	if err := os.WriteFile(filepath.Join(dir, SystemFileName), []byte(body), 0o644); err != nil {
		t.Fatalf("write config: %v", err)
	}
}

func TestManagerHotReload(t *testing.T) {
	dir := t.TempDir()
	writeConfig(t, dir, validYAML)

	var buf bytes.Buffer
	logger := zerolog.New(&buf)

	m, err := NewManager(dir, logger)
	if err != nil {
		t.Fatalf("NewManager: %v", err)
	}
	defer m.Stop()

	// Shorten debounce for the test.
	m.debounceWindow = 50 * time.Millisecond

	if got := m.Current().Listen.Port; got != 8443 {
		t.Fatalf("initial port: got %d", got)
	}

	var reloaded atomic.Int32
	m.OnReload(func(c *SystemConfig) { reloaded.Add(1) })

	// Valid change: port 8443 -> 9001.
	writeConfig(t, dir, strings.Replace(validYAML, "port: 8443", "port: 9001", 1))

	if !waitFor(func() bool { return reloaded.Load() >= 1 }, 3*time.Second) {
		t.Fatalf("reload callback never fired; log=%s", buf.String())
	}
	if got := m.Current().Listen.Port; got != 9001 {
		t.Errorf("after reload port: got %d, want 9001", got)
	}
}

func TestManagerRejectsBadReload(t *testing.T) {
	dir := t.TempDir()
	writeConfig(t, dir, validYAML)

	var buf bytes.Buffer
	logger := zerolog.New(&buf)

	m, err := NewManager(dir, logger)
	if err != nil {
		t.Fatalf("NewManager: %v", err)
	}
	defer m.Stop()
	m.debounceWindow = 50 * time.Millisecond

	initialPort := m.Current().Listen.Port

	// Write a config that fails validation: bad port.
	writeConfig(t, dir, strings.Replace(validYAML, "port: 8443", "port: 0", 1))

	// Give the watcher time to react and reject.
	time.Sleep(400 * time.Millisecond)

	if got := m.Current().Listen.Port; got != initialPort {
		t.Errorf("invalid reload was accepted: port now %d, want %d", got, initialPort)
	}
	if !strings.Contains(buf.String(), "config reload rejected") {
		t.Errorf("expected rejection log; got: %s", buf.String())
	}
}

func waitFor(cond func() bool, timeout time.Duration) bool {
	deadline := time.Now().Add(timeout)
	for time.Now().Before(deadline) {
		if cond() {
			return true
		}
		time.Sleep(20 * time.Millisecond)
	}
	return cond()
}
