package exploit

import (
	"context"
	"errors"
	"strings"
	"testing"

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

type fakeScorer struct {
	scores map[string]float64
	err    error
}

func (f fakeScorer) ScoreStrategies(_ context.Context, _ string, _ map[string]string) (map[string]float64, error) {
	return f.scores, f.err
}

func paths() (pipeline.AttackPath, pipeline.AttackPath) {
	// weak has the higher HEURISTIC prior; strong should still win once Jev scores.
	strong := pipeline.AttackPath{ID: uuid.New(), Name: "BOLA takeover chain", EstimatedSuccessProbability: 0.3}
	weak := pipeline.AttackPath{ID: uuid.New(), Name: "header retry", EstimatedSuccessProbability: 0.6}
	return strong, weak
}

func TestRankPaths_JevOverridesHeuristic(t *testing.T) {
	strong, weak := paths()
	sc := &JevScorer{jev: fakeScorer{scores: map[string]float64{
		strong.ID.String(): 0.95,
		weak.ID.String():   0.2,
	}}}
	ranked := sc.RankPaths(context.Background(), "state", []pipeline.AttackPath{weak, strong})
	if ranked[0].Path.ID != strong.ID {
		t.Fatalf("expected strong first, got %q", ranked[0].Path.Name)
	}
	if ranked[0].Source != "jev" || ranked[0].Score != 0.95 {
		t.Fatalf("expected jev source 0.95, got %s %v", ranked[0].Source, ranked[0].Score)
	}
}

func TestRankPaths_FailsOpenToHeuristic(t *testing.T) {
	strong, weak := paths()
	sc := &JevScorer{jev: fakeScorer{err: errors.New("jev down")}}
	ranked := sc.RankPaths(context.Background(), "state", []pipeline.AttackPath{strong, weak})
	// Jev failed → heuristic ranking: weak (0.6) beats strong (0.3).
	if ranked[0].Path.ID != weak.ID || ranked[0].Source != "heuristic" {
		t.Fatalf("expected heuristic weak first, got %s %s", ranked[0].Path.Name, ranked[0].Source)
	}
}

func TestRankPaths_NilScorerSafe(t *testing.T) {
	strong, weak := paths()
	var sc *JevScorer // nil
	ranked := sc.RankPaths(context.Background(), "state", []pipeline.AttackPath{strong, weak})
	if len(ranked) != 2 || ranked[0].Path.ID != weak.ID {
		t.Fatalf("nil scorer should heuristic-rank, got %+v", ranked)
	}
}

func TestFormatScoreboard(t *testing.T) {
	if FormatScoreboard(nil) != nil {
		t.Error("empty set should yield nil")
	}
	sp := []ScoredPath{
		{Path: pipeline.AttackPath{Name: "BOLA chain"}, Score: 0.92, Source: "jev"},
		{Path: pipeline.AttackPath{Name: "JWT forge"}, Score: 0.71, Source: "jev"},
	}
	lines := FormatScoreboard(sp)
	if len(lines) != 3 { // header + 2 rows
		t.Fatalf("expected header+2 rows, got %d: %v", len(lines), lines)
	}
	if !strings.Contains(lines[0], "2 candidate strategies") {
		t.Errorf("header wrong: %q", lines[0])
	}
	if !strings.Contains(lines[1], "▶") || !strings.Contains(lines[1], "0.92") || !strings.Contains(lines[1], "BOLA chain") {
		t.Errorf("top row wrong: %q", lines[1])
	}
}
