package exploit

import (
	"context"
	"encoding/base64"
	"encoding/json"
	"net/http"
	"net/http/httptest"
	"strings"
	"testing"

	"github.com/google/uuid"

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

func TestJWT_ForgeProducesValidThreeSegmentToken(t *testing.T) {
	e := NewExecutor(loopbackScope(), nil, false)
	step := pipeline.AttackStep{
		ID:   uuid.New(),
		Name: "forge",
		Command: "jwt --action forge --secret crapi " +
			`--claims '{"sub":"victim@example.com","role":"admin"}' --capture jwt=$.token`,
	}

	res, err := e.Execute(context.Background(), step, uuid.New())
	if err != nil {
		t.Fatalf("jwt forge errored: %v", err)
	}
	if !res.Success {
		t.Fatalf("expected success; got %+v", res)
	}

	var out struct {
		Token  string         `json:"token"`
		Header map[string]any `json:"header"`
		Claims map[string]any `json:"claims"`
	}
	jsonStart := strings.Index(res.Output, "{")
	if jsonStart < 0 {
		t.Fatalf("no JSON in output: %q", res.Output)
	}
	if err := json.Unmarshal([]byte(res.Output[jsonStart:]), &out); err != nil {
		t.Fatalf("output not valid JSON: %v\n%s", err, res.Output)
	}

	segs := strings.Split(out.Token, ".")
	if len(segs) != 3 {
		t.Fatalf("token has %d segments, want 3: %q", len(segs), out.Token)
	}
	if segs[2] == "" {
		t.Errorf("forge signature segment is empty, want a real HMAC signature")
	}

	// The payload round-trips: decoding segs[1] gives back the claims we asked for.
	payloadBytes, err := base64.RawURLEncoding.DecodeString(segs[1])
	if err != nil {
		t.Fatalf("payload segment not valid base64url: %v", err)
	}
	var payload map[string]any
	if err := json.Unmarshal(payloadBytes, &payload); err != nil {
		t.Fatalf("payload segment not valid JSON: %v", err)
	}
	if payload["sub"] != "victim@example.com" || payload["role"] != "admin" {
		t.Errorf("payload = %+v, want sub/role claims to round-trip", payload)
	}

	if out.Header["alg"] != "HS256" {
		t.Errorf("header alg = %v, want HS256 (default)", out.Header["alg"])
	}
}

// The forged-token analogue of httpreq's marquee capture-then-replay test:
// step 1 forges a token with a known secret; step 2 replays it as an
// Authorization header — proving the capture wires the jwt verb's own JSON
// output into the chain variable store just like an httpreq response does.
func TestJWT_CaptureForgedTokenThenReplay(t *testing.T) {
	var replayedAuth string
	srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		replayedAuth = r.Header.Get("Authorization")
		w.WriteHeader(http.StatusOK)
	}))
	defer srv.Close()

	e := NewExecutor(loopbackScope(), nil, false)
	steps := []pipeline.AttackStep{
		{
			ID:   uuid.New(),
			Name: "forge",
			Command: "jwt --action forge --secret crapi " +
				`--claims '{"sub":"victim@example.com","role":"admin"}' --capture jwt=$.token`,
		},
		{
			ID:      uuid.New(),
			Name:    "replay",
			Command: "httpreq --url " + srv.URL + `/identity/api/v2/user/dashboard --header 'Authorization: Bearer {{jwt}}'`,
		},
	}

	results, err := e.ExecuteChain(context.Background(), steps, uuid.New())
	if err != nil {
		t.Fatalf("chain errored: %v", err)
	}
	if len(results) != 2 {
		t.Fatalf("want 2 results, got %d", len(results))
	}
	if !results[0].Success {
		t.Fatalf("forge step did not succeed: %+v", results[0])
	}
	if replayedAuth == "" || replayedAuth == "Bearer {{jwt}}" || !strings.HasPrefix(replayedAuth, "Bearer ") {
		t.Fatalf("forged token not captured/replayed; server saw auth=%q", replayedAuth)
	}
	if got := strings.Count(strings.TrimPrefix(replayedAuth, "Bearer "), "."); got != 2 {
		t.Errorf("replayed token has %d dots, want 2 (3 segments): %q", got, replayedAuth)
	}
}

func TestJWT_NoneProducesAlgNoneWithEmptySignature(t *testing.T) {
	e := NewExecutor(loopbackScope(), nil, false)
	step := pipeline.AttackStep{
		ID:      uuid.New(),
		Name:    "none",
		Command: `jwt --action none --claims '{"sub":"victim@example.com"}' --capture jwt=$.token`,
	}

	res, err := e.Execute(context.Background(), step, uuid.New())
	if err != nil {
		t.Fatalf("jwt none errored: %v", err)
	}
	if !res.Success {
		t.Fatalf("expected success; got %+v", res)
	}

	jsonStart := strings.Index(res.Output, "{")
	var out struct {
		Token  string         `json:"token"`
		Header map[string]any `json:"header"`
		Claims map[string]any `json:"claims"`
	}
	if err := json.Unmarshal([]byte(res.Output[jsonStart:]), &out); err != nil {
		t.Fatalf("output not valid JSON: %v\n%s", err, res.Output)
	}

	if out.Header["alg"] != "none" {
		t.Errorf("header alg = %v, want none", out.Header["alg"])
	}

	segs := strings.Split(out.Token, ".")
	if len(segs) != 3 {
		t.Fatalf("token has %d segments, want 3: %q", len(segs), out.Token)
	}
	if segs[2] != "" {
		t.Errorf("alg:none signature segment = %q, want empty", segs[2])
	}
}

func TestJWT_DecodeExtractsClaimsWithoutVerification(t *testing.T) {
	forgeExec := NewExecutor(loopbackScope(), nil, false)
	forgeStep := pipeline.AttackStep{
		ID:      uuid.New(),
		Name:    "forge",
		Command: `jwt --action forge --secret s3cr3t --claims '{"sub":"alice","role":"user"}'`,
	}
	forgeRes, err := forgeExec.Execute(context.Background(), forgeStep, uuid.New())
	if err != nil {
		t.Fatalf("setup forge errored: %v", err)
	}
	var forged struct {
		Token string `json:"token"`
	}
	jsonStart := strings.Index(forgeRes.Output, "{")
	if err := json.Unmarshal([]byte(forgeRes.Output[jsonStart:]), &forged); err != nil {
		t.Fatalf("setup forge output not valid JSON: %v", err)
	}
	token := forged.Token
	if token == "" {
		t.Fatalf("setup forge did not produce a token")
	}

	e := NewExecutor(loopbackScope(), nil, false)
	step := pipeline.AttackStep{
		ID:      uuid.New(),
		Name:    "decode",
		Command: "jwt --action decode --token " + token,
	}
	res, err := e.Execute(context.Background(), step, uuid.New())
	if err != nil {
		t.Fatalf("jwt decode errored: %v", err)
	}
	if !res.Success {
		t.Fatalf("expected success; got %+v", res)
	}

	decodeJSONStart := strings.Index(res.Output, "{")
	var out struct {
		Header map[string]any `json:"header"`
		Claims map[string]any `json:"claims"`
	}
	if err := json.Unmarshal([]byte(res.Output[decodeJSONStart:]), &out); err != nil {
		t.Fatalf("output not valid JSON: %v\n%s", err, res.Output)
	}
	if out.Claims["sub"] != "alice" || out.Claims["role"] != "user" {
		t.Errorf("decoded claims = %+v, want sub=alice role=user", out.Claims)
	}
	if out.Header["alg"] != "HS256" {
		t.Errorf("decoded header alg = %v, want HS256", out.Header["alg"])
	}
}

// jwt is a built-in verb: it must bypass the binary allowlist exactly like
// httpreq, since there is no binary to allow.
func TestJWT_BypassesBinaryAllowlist(t *testing.T) {
	e := NewExecutor(loopbackScope(), nil, false).
		WithAllowedExecutables([]string{"nuclei", "sqlmap"}) // deliberately no jwt
	step := pipeline.AttackStep{
		ID:      uuid.New(),
		Name:    "forge",
		Command: `jwt --action forge --secret x --claims '{}'`,
	}

	res, err := e.Execute(context.Background(), step, uuid.New())
	if err != nil {
		t.Fatalf("built-in should bypass allowlist; got %v", err)
	}
	if !res.Success {
		t.Fatalf("expected success; got %+v", res)
	}
}

func TestJWT_ForgeRequiresSecret(t *testing.T) {
	e := NewExecutor(loopbackScope(), nil, false)
	step := pipeline.AttackStep{
		ID:      uuid.New(),
		Name:    "forge-no-secret",
		Command: `jwt --action forge --claims '{}'`,
	}

	res, err := e.Execute(context.Background(), step, uuid.New())
	if err == nil {
		t.Fatalf("expected an error for missing --secret")
	}
	if res.Success {
		t.Fatalf("expected failure; got %+v", res)
	}
	if !strings.Contains(res.Output, "BLOCKED") || !strings.Contains(res.Output, "--secret") {
		t.Errorf("output = %q, want a BLOCKED message mentioning --secret", res.Output)
	}
}
