package exploit

import (
	"fmt"
	"sort"
	"strings"

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

// relationship defines how one finding type can lead to another.
type relationship struct {
	from string
	to   []string
}

// predefined attack chain relationships
var chainRules = []relationship{
	{from: "sqli", to: []string{"data_extraction", "auth_bypass", "file_read"}},
	{from: "path_traversal", to: []string{"file_read", "source_code_disclosure", "credential_access"}},
	{from: "open_redirect", to: []string{"phishing_amplification", "oauth_bypass"}},
	{from: "exposed_git", to: []string{"source_code_disclosure", "credential_extraction", "secret_discovery"}},
	{from: "weak_credentials", to: []string{"authentication_bypass", "privilege_escalation"}},
	{from: "outdated_software", to: []string{"known_exploit"}},
	{from: "misconfigured_cors", to: []string{"cross_origin_data_theft"}},
	{from: "xxe", to: []string{"ssrf", "file_read", "internal_network_scan"}},
	{from: "ssrf", to: []string{"internal_network_scan", "cloud_metadata_access", "port_scan"}},
	{from: "xss", to: []string{"session_hijack", "credential_theft", "phishing"}},
	{from: "rce", to: []string{"full_system_access", "lateral_movement", "data_exfiltration"}},
	{from: "lfi", to: []string{"file_read", "rce_via_log_poisoning", "credential_access"}},
	{from: "idor", to: []string{"data_extraction", "privilege_escalation"}},
}

// killChainBridges maps an "enabler" keyword to the keywords of findings it can
// feed into. It's what lets findings from DIFFERENT agents compose into a single
// cross-engagement kill-chain: an information leak feeding a credential, feeding
// an authorization bypass, feeding account takeover. Matched case-insensitively
// against a finding's category + title + description.
var killChainBridges = map[string][]string{
	"exposure":   {"credential", "token", "key", "auth"},
	"leak":       {"credential", "token", "key", "auth"},
	"verbose":    {"credential", "token", "path", "internal"},
	"git":        {"credential", "secret", "source"},
	"secret":     {"credential", "auth", "cloud"},
	"credential": {"privilege", "admin", "account", "auth", "cloud"},
	"weak":       {"privilege", "admin", "account"},
	"ssrf":       {"cloud", "metadata", "credential", "internal"},
	"sqli":       {"credential", "data", "auth"},
	"injection":  {"credential", "data", "auth"},
	"xss":        {"session", "account", "credential"},
	"bola":       {"account", "data", "privilege"},
	"idor":       {"account", "data", "privilege"},
	"jwt":        {"account", "admin", "privilege"},
	"bfla":       {"admin", "privilege", "account"},
}

// maxKillChains caps how many composed cross-finding chains we emit so a large
// finding set can't produce a quadratic blow-up.
const maxKillChains = 12

// BuildKillChains composes findings from different agents into cross-engagement
// kill-chains: for each ordered pair where the first finding is a known enabler
// of the second, it emits a combined AttackPath (finding A → finding B) tagged
// with both finding IDs and the higher of the two severities. Pure + deterministic
// (sorted, capped) so it's testable without a live run.
func (p *PathBuilder) BuildKillChains(findings []pipeline.ClassifiedFinding) []pipeline.AttackPath {
	text := func(f pipeline.ClassifiedFinding) string {
		return strings.ToLower(f.AttackCategory + " " + f.Title + " " + f.Description)
	}
	var chains []pipeline.AttackPath
	for i := range findings {
		at := text(findings[i])
		enabled := map[string]bool{}
		for enabler, targets := range killChainBridges {
			if strings.Contains(at, enabler) {
				for _, t := range targets {
					enabled[t] = true
				}
			}
		}
		if len(enabled) == 0 {
			continue
		}
		for j := range findings {
			if i == j {
				continue
			}
			bt := text(findings[j])
			matched := ""
			for kw := range enabled {
				if strings.Contains(bt, kw) {
					matched = kw
					break
				}
			}
			if matched == "" {
				continue
			}
			a, b := findings[i], findings[j]
			chains = append(chains, pipeline.AttackPath{
				ID:          uuid.New(),
				Name:        a.Title + " → " + b.Title,
				Description: fmt.Sprintf("Cross-finding kill-chain: %q enables %q (via %q).", a.Title, b.Title, matched),
				Steps: []pipeline.AttackStep{
					{ID: uuid.New(), Name: "Leverage: " + a.Title},
					{ID: uuid.New(), Name: "Escalate to: " + b.Title, TechniqueID: killChainTechnique(matched)},
				},
				TargetFindingIDs:            []uuid.UUID{a.ID, b.ID},
				EstimatedSuccessProbability: killChainProbability(a, b),
				ExpectedImpact:              higherSeverity(a.Severity, b.Severity),
			})
		}
	}
	// Deterministic order: strongest impact first, then name.
	sort.SliceStable(chains, func(i, j int) bool {
		si, sj := severityRank(chains[i].ExpectedImpact), severityRank(chains[j].ExpectedImpact)
		if si != sj {
			return si > sj
		}
		return chains[i].Name < chains[j].Name
	})
	if len(chains) > maxKillChains {
		chains = chains[:maxKillChains]
	}
	return chains
}

// killChainTechnique maps the bridge keyword that connected two findings to the
// MITRE ATT&CK technique the escalation step represents.
func killChainTechnique(keyword string) string {
	switch keyword {
	case "privilege", "admin":
		return "T1068"
	case "credential", "token", "key", "secret":
		return "T1552"
	case "account", "auth":
		return "T1078"
	case "cloud", "metadata":
		return "T1552"
	case "session":
		return "T1539"
	case "data":
		return "T1213"
	default:
		return ""
	}
}

func killChainProbability(a, b pipeline.ClassifiedFinding) float64 {
	// A chain is bounded by its weaker link; average the two CVSS-derived
	// priors and damp slightly (composition is speculative).
	avg := (a.CVSSScore + b.CVSSScore) / 20.0 // CVSS 0..10 → 0..1 each
	p := avg * 0.8
	if p > 0.9 {
		p = 0.9
	}
	if p < 0.1 {
		p = 0.1
	}
	return p
}

func severityRank(s any) int {
	switch strings.ToLower(fmt.Sprintf("%v", s)) {
	case "critical":
		return 4
	case "high":
		return 3
	case "medium":
		return 2
	case "low":
		return 1
	default:
		return 0
	}
}

func higherSeverity(a, b pipeline.Severity) string {
	if severityRank(a) >= severityRank(b) {
		return string(a)
	}
	return string(b)
}

// PathBuilder constructs attack chains from classified findings using predefined rules.
type PathBuilder struct{}

// NewPathBuilder creates a new path builder.
func NewPathBuilder() *PathBuilder {
	return &PathBuilder{}
}

// BuildChains constructs all valid attack chains from the given findings.
func (p *PathBuilder) BuildChains(findings []pipeline.ClassifiedFinding) []pipeline.AttackPath {
	var paths []pipeline.AttackPath

	for _, finding := range findings {
		category := finding.AttackCategory
		if category == "" {
			continue
		}

		for _, rule := range chainRules {
			if rule.from != category {
				continue
			}

			// Build a path for each possible chain
			for _, target := range rule.to {
				path := pipeline.AttackPath{
					ID:          uuid.New(),
					Name:        category + " → " + target,
					Description: "Chain from " + finding.Title + " to " + target,
					Steps: []pipeline.AttackStep{
						{
							ID:   uuid.New(),
							Name: "Exploit " + finding.Title,
						},
						{
							ID:   uuid.New(),
							Name: "Achieve " + target,
						},
					},
					TargetFindingIDs:            []uuid.UUID{finding.ID},
					EstimatedSuccessProbability: estimateSuccess(finding, target),
					ExpectedImpact:              impactLevel(target),
				}
				paths = append(paths, path)
			}
		}
	}

	// Sort by estimated success probability descending
	sort.Slice(paths, func(i, j int) bool {
		return paths[i].EstimatedSuccessProbability > paths[j].EstimatedSuccessProbability
	})

	return paths
}

func estimateSuccess(finding pipeline.ClassifiedFinding, target string) float64 {
	// Base probability from finding confidence
	base := 0.5
	switch finding.Confidence {
	case pipeline.ConfidenceHigh:
		base = 0.8
	case pipeline.ConfidenceMedium:
		base = 0.6
	case pipeline.ConfidenceLow:
		base = 0.3
	}

	// Adjust based on chain complexity
	complexTargets := map[string]float64{
		"full_system_access":    0.6,
		"lateral_movement":      0.5,
		"privilege_escalation":  0.7,
		"cloud_metadata_access": 0.8,
		"data_extraction":       0.9,
		"session_hijack":        0.7,
	}

	if modifier, ok := complexTargets[target]; ok {
		base *= modifier
	}

	return base
}

func impactLevel(target string) string {
	highImpact := map[string]bool{
		"full_system_access": true, "rce_via_log_poisoning": true,
		"lateral_movement": true, "data_exfiltration": true,
		"credential_access": true, "privilege_escalation": true,
	}
	if highImpact[target] {
		return "high"
	}
	return "medium"
}
