package db

import (
	"context"
	"database/sql"
	"path/filepath"
	"testing"
	"testing/fstest"
	"time"
)

func openTestSQLiteWithSchema(t *testing.T) *sql.DB {
	t.Helper()
	dbPath := filepath.Join(t.TempDir(), "test.db")
	db, err := OpenSQLite(dbPath)
	if err != nil {
		t.Fatalf("OpenSQLite: %v", err)
	}
	t.Cleanup(func() { _ = db.Close() })

	// Apply() creates schema_migrations itself via IF NOT EXISTS, so the
	// test fixture skips it.
	migs, err := LoadMigrations(fstest.MapFS{
		"002_phase1.sql": {Data: []byte(`
			CREATE TABLE trackers (
				id TEXT PRIMARY KEY, name TEXT NOT NULL, type TEXT NOT NULL,
				base_url TEXT NOT NULL, config_json TEXT NOT NULL,
				enabled INTEGER NOT NULL DEFAULT 1,
				scrape_interval_seconds INTEGER NOT NULL DEFAULT 300,
				scrape_jitter_seconds INTEGER NOT NULL DEFAULT 60,
				use_byparr INTEGER NOT NULL DEFAULT 0,
				created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL
			);
			CREATE TABLE ratio_snapshots (
				id INTEGER PRIMARY KEY AUTOINCREMENT,
				tracker_id TEXT NOT NULL REFERENCES trackers(id),
				timestamp INTEGER NOT NULL,
				simulation_id INTEGER,
				real_uploaded_bytes INTEGER, real_downloaded_bytes INTEGER, real_ratio REAL,
				displayed_uploaded_bytes INTEGER, displayed_downloaded_bytes INTEGER, displayed_ratio REAL,
				bonus_points INTEGER, unsat_count INTEGER, unsat_limit INTEGER,
				class_or_rank TEXT, raw_json TEXT
			);
		`)},
	})
	if err != nil {
		t.Fatalf("LoadMigrations: %v", err)
	}
	if _, err := Apply(db, migs); err != nil {
		t.Fatalf("Apply: %v", err)
	}
	return db
}

func TestUpsertTrackerAndWriteRatio(t *testing.T) {
	db := openTestSQLiteWithSchema(t)
	w := NewSnapshotWriter(db, nil)
	ctx := context.Background()

	if err := w.UpsertTracker(ctx, TrackerRow{
		ID: "mam", Name: "MAM", Type: "mam",
		BaseURL: "https://x", ConfigJSON: "{}",
		Enabled: true,
		ScrapeIntervalSeconds: 600, ScrapeJitterSeconds: 60,
	}); err != nil {
		t.Fatalf("UpsertTracker: %v", err)
	}

	up := int64(123)
	dn := int64(45)
	r := 2.73
	if err := w.WriteRatio(ctx, RatioSnapshotRow{
		TrackerID:                "mam",
		Timestamp:                time.Now(),
		RealUploadedBytes:        &up,
		RealDownloadedBytes:      &dn,
		RealRatio:                &r,
		DisplayedUploadedBytes:   &up,
		DisplayedDownloadedBytes: &dn,
		DisplayedRatio:           &r,
		ClassOrRank:              "Vip",
		RawJSON:                  `{"x":1}`,
	}); err != nil {
		t.Fatalf("WriteRatio: %v", err)
	}

	var count int
	if err := db.QueryRow(`SELECT COUNT(*) FROM ratio_snapshots WHERE tracker_id = ?`, "mam").Scan(&count); err != nil {
		t.Fatalf("count: %v", err)
	}
	if count != 1 {
		t.Errorf("snapshot count = %d, want 1", count)
	}
}

func TestUpsertTrackerIsIdempotent(t *testing.T) {
	db := openTestSQLiteWithSchema(t)
	w := NewSnapshotWriter(db, nil)
	ctx := context.Background()

	row := TrackerRow{
		ID: "mam", Name: "MAM v1", Type: "mam",
		BaseURL: "https://x", ConfigJSON: "{}",
		Enabled: true, ScrapeIntervalSeconds: 600,
	}
	for i := 0; i < 3; i++ {
		if err := w.UpsertTracker(ctx, row); err != nil {
			t.Fatalf("upsert %d: %v", i, err)
		}
	}
	row.Name = "MAM v2"
	if err := w.UpsertTracker(ctx, row); err != nil {
		t.Fatalf("upsert v2: %v", err)
	}

	var name string
	if err := db.QueryRow(`SELECT name FROM trackers WHERE id = ?`, "mam").Scan(&name); err != nil {
		t.Fatalf("select: %v", err)
	}
	if name != "MAM v2" {
		t.Errorf("name = %q, want MAM v2", name)
	}
}

func TestWriteRatioPreservesNullColumns(t *testing.T) {
	db := openTestSQLiteWithSchema(t)
	w := NewSnapshotWriter(db, nil)
	ctx := context.Background()

	if err := w.UpsertTracker(ctx, TrackerRow{
		ID: "x", Name: "X", Type: "mam", BaseURL: "https://x", ConfigJSON: "{}", Enabled: true,
	}); err != nil {
		t.Fatalf("upsert: %v", err)
	}

	if err := w.WriteRatio(ctx, RatioSnapshotRow{
		TrackerID: "x",
		Timestamp: time.Now(),
	}); err != nil {
		t.Fatalf("WriteRatio: %v", err)
	}

	row := db.QueryRow(`SELECT real_uploaded_bytes, displayed_ratio, bonus_points
	                    FROM ratio_snapshots WHERE tracker_id = ?`, "x")
	var u sql.NullInt64
	var rratio sql.NullFloat64
	var bonus sql.NullInt64
	if err := row.Scan(&u, &rratio, &bonus); err != nil {
		t.Fatalf("scan: %v", err)
	}
	if u.Valid || rratio.Valid || bonus.Valid {
		t.Errorf("expected NULLs, got %v / %v / %v", u, rratio, bonus)
	}
}
