// Package nvdcheck cross-references the classifier's CVSS score against
// NVD's canonical score for the same CVE. Large mismatches (> 2.0 CVSS
// points) downgrade the finding's confidence to Unverified.
//
// Phase 4.3.5 of Wave 4. Results are cached locally so NVD's rate limit
// (5 req / 30s without a key, 50 / 30s with one) doesn't throttle a
// large campaign.
package nvdcheck

import (
	"context"
	"encoding/json"
	"fmt"
	"io"
	"math"
	"net/http"
	"os"
	"path/filepath"
	"sync"
	"time"

	"github.com/Armur-Ai/Pentest-Swarm-AI/internal/pipeline"
)

// Client talks to NVD and caches results on disk.
type Client struct {
	http     *http.Client
	baseURL  string
	apiKey   string
	cacheDir string

	mu    sync.Mutex
	cache map[string]*Entry // cveID -> entry
}

// Entry is one cached CVSS lookup.
type Entry struct {
	CVEID      string    `json:"cve_id"`
	BaseScore  float64   `json:"base_score"`
	BaseVector string    `json:"base_vector"`
	Severity   string    `json:"severity"`
	FetchedAt  time.Time `json:"fetched_at"`
}

// Config customizes a Client.
type Config struct {
	APIKey   string        // optional NVD API key — https://nvd.nist.gov/developers/request-an-api-key
	CacheDir string        // default: ~/.pentestswarm/nvd-cache
	Timeout  time.Duration // default 15s
}

// NewClient builds a client with disk cache loaded.
func NewClient(cfg Config) (*Client, error) {
	if cfg.CacheDir == "" {
		home, _ := os.UserHomeDir()
		cfg.CacheDir = filepath.Join(home, ".pentestswarm", "nvd-cache")
	}
	if cfg.Timeout == 0 {
		cfg.Timeout = 15 * time.Second
	}
	c := &Client{
		http:     &http.Client{Timeout: cfg.Timeout},
		baseURL:  "https://services.nvd.nist.gov/rest/json/cves/2.0",
		apiKey:   cfg.APIKey,
		cacheDir: cfg.CacheDir,
		cache:    map[string]*Entry{},
	}
	_ = os.MkdirAll(cfg.CacheDir, 0o755)
	return c, nil
}

// Lookup fetches CVSS details for a CVE id, using cache first.
func (c *Client) Lookup(ctx context.Context, cveID string) (*Entry, error) {
	c.mu.Lock()
	if e, ok := c.cache[cveID]; ok {
		c.mu.Unlock()
		return e, nil
	}
	c.mu.Unlock()

	// Try disk cache.
	if e, ok := c.loadFromDisk(cveID); ok {
		c.mu.Lock()
		c.cache[cveID] = e
		c.mu.Unlock()
		return e, nil
	}

	// Fetch.
	req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+"?cveId="+cveID, nil)
	if err != nil {
		return nil, err
	}
	if c.apiKey != "" {
		req.Header.Set("apiKey", c.apiKey)
	}
	resp, err := c.http.Do(req)
	if err != nil {
		return nil, fmt.Errorf("nvd transport: %w", err)
	}
	defer resp.Body.Close()
	body, _ := io.ReadAll(resp.Body)
	if resp.StatusCode >= 400 {
		return nil, fmt.Errorf("nvd %d: %s", resp.StatusCode, string(body))
	}

	entry, err := parseNVDResponse(cveID, body)
	if err != nil {
		return nil, err
	}
	c.mu.Lock()
	c.cache[cveID] = entry
	c.mu.Unlock()
	c.saveToDisk(entry)
	return entry, nil
}

// SanityCheck compares the classifier's CVSS against NVD's authoritative
// number. Returns (nvdScore, delta, ok) — ok=false signals either "not in
// NVD" (for novel findings with no CVE) or "big mismatch so downgrade".
//
// Callers that get ok=false + a non-zero delta should treat the finding
// as Unverified. ok=false + delta=0 just means "we couldn't check" —
// leave the finding alone.
func (c *Client) SanityCheck(ctx context.Context, f pipeline.ClassifiedFinding) (float64, float64, bool) {
	if len(f.CVEIDs) == 0 {
		return 0, 0, true // nothing to check against — not a failure
	}
	e, err := c.Lookup(ctx, f.CVEIDs[0])
	if err != nil || e == nil {
		return 0, 0, true // couldn't check; leave the finding alone
	}
	delta := math.Abs(f.CVSSScore - e.BaseScore)
	return e.BaseScore, delta, delta <= 2.0
}

func (c *Client) diskPath(cveID string) string {
	return filepath.Join(c.cacheDir, cveID+".json")
}

func (c *Client) loadFromDisk(cveID string) (*Entry, bool) {
	data, err := os.ReadFile(c.diskPath(cveID))
	if err != nil {
		return nil, false
	}
	var e Entry
	if err := json.Unmarshal(data, &e); err != nil {
		return nil, false
	}
	return &e, true
}

func (c *Client) saveToDisk(e *Entry) {
	data, _ := json.Marshal(e)
	_ = os.WriteFile(c.diskPath(e.CVEID), data, 0o644)
}

// parseNVDResponse extracts the first v3 (preferred) or v2 CVSS metric.
// NVD's API shape changes occasionally; we parse defensively.
func parseNVDResponse(cveID string, body []byte) (*Entry, error) {
	var resp struct {
		Vulnerabilities []struct {
			CVE struct {
				Metrics struct {
					CVSSMetricV31 []struct {
						CVSSData struct {
							BaseScore    float64 `json:"baseScore"`
							VectorString string  `json:"vectorString"`
							BaseSeverity string  `json:"baseSeverity"`
						} `json:"cvssData"`
					} `json:"cvssMetricV31"`
					CVSSMetricV30 []struct {
						CVSSData struct {
							BaseScore    float64 `json:"baseScore"`
							VectorString string  `json:"vectorString"`
							BaseSeverity string  `json:"baseSeverity"`
						} `json:"cvssData"`
					} `json:"cvssMetricV30"`
				} `json:"metrics"`
			} `json:"cve"`
		} `json:"vulnerabilities"`
	}
	if err := json.Unmarshal(body, &resp); err != nil {
		return nil, fmt.Errorf("parse nvd: %w", err)
	}
	if len(resp.Vulnerabilities) == 0 {
		return nil, fmt.Errorf("nvd: cve %s not found", cveID)
	}
	m := resp.Vulnerabilities[0].CVE.Metrics
	entry := &Entry{CVEID: cveID, FetchedAt: time.Now()}
	switch {
	case len(m.CVSSMetricV31) > 0:
		d := m.CVSSMetricV31[0].CVSSData
		entry.BaseScore = d.BaseScore
		entry.BaseVector = d.VectorString
		entry.Severity = d.BaseSeverity
	case len(m.CVSSMetricV30) > 0:
		d := m.CVSSMetricV30[0].CVSSData
		entry.BaseScore = d.BaseScore
		entry.BaseVector = d.VectorString
		entry.Severity = d.BaseSeverity
	default:
		return nil, fmt.Errorf("nvd: %s has no CVSS v3 metric", cveID)
	}
	return entry, nil
}
