package notifications

import (
	"context"
	"database/sql"
	"errors"
	"path/filepath"
	"sync/atomic"
	"testing"
	"testing/fstest"

	"github.com/rs/zerolog"

	dbpkg "github.com/operator/command-center/internal/db"
)

func TestColorAndPriority(t *testing.T) {
	if colorForSeverity(SeverityCritical) != 0xDC2626 {
		t.Error("critical color")
	}
	if priorityForSeverity(SeverityCritical) != "5" {
		t.Error("critical priority")
	}
	if priorityForSeverity(SeverityInfo) != "3" {
		t.Error("info priority")
	}
}

type fakeChannel struct {
	name      string
	calls     atomic.Int32
	failEvery int32
}

func (f *fakeChannel) Name() string { return f.name }
func (f *fakeChannel) Send(_ context.Context, _ Notification, _ map[string]any) (int, error) {
	n := f.calls.Add(1)
	if f.failEvery > 0 && n%f.failEvery == 0 {
		return 0, errors.New("synthetic")
	}
	return 1, nil
}

func TestDispatcherFanOut(t *testing.T) {
	db := setupSQLite(t)
	d := NewDispatcher(db, zerolog.Nop())
	a := &fakeChannel{name: "a"}
	b := &fakeChannel{name: "b"}
	d.Register(a)
	d.Register(b)

	d.Dispatch(context.Background(),
		Notification{Title: "T", RuleID: 42, RuleName: "R"},
		[]string{"a", "b"},
		map[string]map[string]any{},
	)
	if a.calls.Load() != 1 || b.calls.Load() != 1 {
		t.Fatalf("expected both channels called once: a=%d b=%d", a.calls.Load(), b.calls.Load())
	}

	rows, err := d.LoadRecent(context.Background(), 10)
	if err != nil {
		t.Fatalf("LoadRecent: %v", err)
	}
	if len(rows) != 3 {
		t.Fatalf("expected 3 log rows (dashboard+a+b), got %d", len(rows))
	}
}

func TestDispatcherChannelFailureIsolated(t *testing.T) {
	db := setupSQLite(t)
	d := NewDispatcher(db, zerolog.Nop())
	good := &fakeChannel{name: "good"}
	bad := &fakeChannel{name: "bad", failEvery: 1}
	d.Register(good)
	d.Register(bad)
	d.Dispatch(context.Background(),
		Notification{Title: "T"},
		[]string{"good", "bad"},
		nil,
	)
	if good.calls.Load() != 1 {
		t.Errorf("good not called: %d", good.calls.Load())
	}
	rows, _ := d.LoadRecent(context.Background(), 10)
	if len(rows) != 3 {
		t.Errorf("expected 3 rows, got %d", len(rows))
	}
	failedFound := false
	for _, r := range rows {
		if r.Channel == "bad" && !r.Delivered && r.ErrorMessage != "" {
			failedFound = true
		}
	}
	if !failedFound {
		t.Error("expected a failed log row for 'bad' channel")
	}
}

func TestDispatchAlwaysLogsToDashboard(t *testing.T) {
	db := setupSQLite(t)
	d := NewDispatcher(db, zerolog.Nop())
	d.Dispatch(context.Background(), Notification{Title: "alone"}, []string{}, nil)
	rows, _ := d.LoadRecent(context.Background(), 10)
	if len(rows) != 1 || rows[0].Channel != "dashboard" {
		t.Errorf("expected 1 dashboard row, got %+v", rows)
	}
}

func TestVAPIDStorePersistsAcrossLoads(t *testing.T) {
	mem := &memSecrets{m: map[string][]byte{}}
	s := NewVAPIDStore(mem)
	k1, err := s.Load(context.Background())
	if err != nil {
		t.Fatalf("first Load: %v", err)
	}
	if k1.Public == "" || k1.Private == "" {
		t.Fatal("empty keys")
	}
	s2 := NewVAPIDStore(mem)
	k2, err := s2.Load(context.Background())
	if err != nil {
		t.Fatalf("second Load: %v", err)
	}
	if k1.Public != k2.Public || k1.Private != k2.Private {
		t.Error("VAPID keys not persisted")
	}
}

type memSecrets struct{ m map[string][]byte }

func (m *memSecrets) Get(_ context.Context, k string) ([]byte, error) {
	if v, ok := m.m[k]; ok {
		return v, nil
	}
	return nil, errors.New("not found")
}
func (m *memSecrets) Set(_ context.Context, k string, v []byte) error {
	m.m[k] = append([]byte(nil), v...)
	return nil
}

func setupSQLite(t *testing.T) *sql.DB {
	t.Helper()
	dbPath := filepath.Join(t.TempDir(), "test.db")
	db, err := dbpkg.OpenSQLite(dbPath)
	if err != nil {
		t.Fatalf("OpenSQLite: %v", err)
	}
	t.Cleanup(func() { _ = db.Close() })
	migs, err := dbpkg.LoadMigrations(fstest.MapFS{
		"007_phase6.sql": {Data: []byte(`
			CREATE TABLE push_subscriptions (
				id INTEGER PRIMARY KEY AUTOINCREMENT,
				endpoint TEXT UNIQUE NOT NULL, p256dh TEXT NOT NULL, auth TEXT NOT NULL,
				user_agent TEXT, device_label TEXT,
				created_at INTEGER NOT NULL, last_delivery_at INTEGER,
				failure_count INTEGER NOT NULL DEFAULT 0
			);
			CREATE TABLE notification_rules (
				id INTEGER PRIMARY KEY AUTOINCREMENT,
				name TEXT NOT NULL, enabled INTEGER NOT NULL DEFAULT 1,
				trigger_type TEXT NOT NULL, trigger_config_json TEXT NOT NULL,
				channels_json TEXT NOT NULL DEFAULT '["push"]',
				cooldown_seconds INTEGER NOT NULL DEFAULT 3600,
				last_fired_at INTEGER, source TEXT NOT NULL DEFAULT 'yaml',
				created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL
			);
			CREATE TABLE notification_log (
				id INTEGER PRIMARY KEY AUTOINCREMENT,
				rule_id INTEGER, title TEXT NOT NULL, body TEXT,
				sent_at INTEGER NOT NULL, delivered INTEGER NOT NULL DEFAULT 0,
				channel TEXT, subscription_id INTEGER, error_message TEXT
			);
		`)},
	})
	if err != nil {
		t.Fatalf("LoadMigrations: %v", err)
	}
	if _, err := dbpkg.Apply(db, migs); err != nil {
		t.Fatalf("Apply: %v", err)
	}
	return db
}
