package exploit

import (
	"context"
	"encoding/json"
	"fmt"
	"io"
	"net/http"
	"regexp"
	"strconv"
	"strings"
	"sync"
	"time"

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

// builtinHTTPReq is the executor verb that performs an authenticated,
// arbitrary HTTP request in-process. It exists because API business-logic
// vulnerabilities — BOLA/IDOR, mass assignment, broken function-level auth,
// JWT abuse — cannot be reached by the recon tool binaries (nuclei/sqlmap/
// dalfox only fingerprint and template-match). They require crafting a
// specific request: a chosen method, an Authorization header carrying a
// captured token, a JSON body with an extra field. httpreq is that primitive.
//
// It is handled in-process rather than by shelling out to curl for three
// reasons: curl is deliberately barred from the executor allowlist (a
// prompt-injection bridge to RCE), an in-process request is scope-validated
// against the exact parsed URL host, and the response is returned in a
// structured, capped form the LLM/classifier can reason over.
const builtinHTTPReq = "httpreq"

// maxRaceRequests caps --race so a business-logic race test can't be turned into
// a denial-of-service flood against the target.
const maxRaceRequests = 30

// builtinVerbs are executor commands handled in-process instead of by
// launching a binary. They bypass the binary allowlist (there is no binary)
// but still pass scope validation and safe-mode.
var builtinVerbs = map[string]struct{}{
	builtinHTTPReq: {},
	builtinJWT:     {},
}

// isBuiltinVerb reports whether cmd's first token is an in-process verb.
func isBuiltinVerb(exe string) bool {
	_, ok := builtinVerbs[strings.ToLower(exe)]
	return ok
}

// firstToken returns the first whitespace-separated token of a command line
// (the executable/verb), or "" if the command is empty. Used to decide, before
// full parsing, whether a command is a builtin verb.
func firstToken(cmd string) string {
	fields := strings.Fields(cmd)
	if len(fields) == 0 {
		return ""
	}
	return fields[0]
}

// httpReqBodyCap bounds how much of a response body is captured into the
// step result. API responses are usually small JSON; this keeps a runaway
// endpoint (a file download, an error page dump) from flooding the report
// and the downstream LLM context.
const httpReqBodyCap = 16 * 1024

// httpReqOptions is the parsed form of an `httpreq` command line.
type httpReqOptions struct {
	method   string
	url      string
	headers  []string // each "Key: Value"
	body     string
	follow   bool
	timeout  time.Duration
	captures []captureSpec // --capture name=selector
	race     int           // --race N: fire N identical requests concurrently (business-logic race / TOCTOU)
}

// captureSpec extracts a value from a response into a chain variable that
// later steps reference as {{name}}. This is what turns independent requests
// into a real attack chain: log in, capture the JWT, replay it as the victim.
type captureSpec struct {
	name     string
	selector string // "header:Name" or a JSON path ("$.a.b" / "a.b")
}

// runHTTPReq executes an in-process HTTP request. Command syntax:
//
//	httpreq --method GET --url <URL> [--header "K: V"]... [--body <string>]
//	        [--follow] [--timeout <seconds>]
//
// --header may repeat. --follow enables redirect following (off by default so
// the raw 30x is visible). The request is scope-validated against the parsed
// URL host — a target outside scope is a hard block, identical to the binary
// tool path.
// vars, when non-nil, is the chain variable store: --capture directives write
// extracted response values into it for later steps. It is nil on the
// single-step Execute path (nothing to capture into).
func (e *Executor) runHTTPReq(ctx 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 := parseHTTPReqArgs(parts[1:])
	if err != nil {
		return blocked(err.Error()), fmt.Errorf("httpreq: %w", err)
	}
	if opts.url == "" {
		return blocked("httpreq requires --url"), fmt.Errorf("httpreq: missing --url")
	}

	// Scope validation against the exact URL host — defence in depth on top of
	// the ValidateCommand check Execute already ran on the whole command line.
	if e.scopeDef != nil {
		if serr := scope.ValidateAndLog(builtinHTTPReq, opts.url, *e.scopeDef); serr != nil {
			return blocked(serr.Error()), fmt.Errorf("scope violation: %w", serr)
		}
	}

	reqCtx := ctx
	if opts.timeout > 0 {
		var cancel context.CancelFunc
		reqCtx, cancel = context.WithTimeout(ctx, opts.timeout)
		defer cancel()
	}

	req, err := buildHTTPReq(reqCtx, opts)
	if err != nil {
		return blocked(err.Error()), fmt.Errorf("httpreq: %w", err)
	}

	client := &http.Client{}
	if !opts.follow {
		client.CheckRedirect = func(*http.Request, []*http.Request) error {
			return http.ErrUseLastResponse
		}
	}

	// --race N fires N identical requests concurrently to probe for
	// business-logic race conditions (TOCTOU) — e.g. redeeming a single-use
	// coupon or withdrawing a balance twice. The representative response feeds
	// the normal capture/pattern logic; a race summary is prepended to output.
	var raceSummary string
	var resp *http.Response
	if opts.race > 1 {
		resp, raceSummary, err = fireRace(reqCtx, client, opts)
	} else {
		resp, err = client.Do(req)
	}
	if err != nil {
		// A transport error (connection refused, timeout) is a real, reportable
		// outcome, not an executor failure — record it as an unsuccessful step.
		return &pipeline.ExecutionResult{
			StepID:          step.ID,
			CampaignID:      campaignID,
			CommandExecuted: step.Command,
			Output:          fmt.Sprintf("%s %s\nrequest error: %s", opts.method, opts.url, err),
			Success:         false,
			ExecutedAt:      start,
			DurationMs:      int(time.Since(start).Milliseconds()),
		}, nil
	}
	defer resp.Body.Close()

	body, _ := io.ReadAll(io.LimitReader(resp.Body, httpReqBodyCap+1))
	truncated := false
	if len(body) > httpReqBodyCap {
		body = body[:httpReqBodyCap]
		truncated = true
	}

	// Capture response values into chain variables for later steps.
	if vars != nil {
		for _, c := range opts.captures {
			if v, ok := captureValue(c.selector, resp, body); ok {
				vars[c.name] = v
			}
		}
	}

	output := formatHTTPResponse(opts, resp, string(body), truncated)
	if raceSummary != "" {
		output = raceSummary + output
	}

	// Success = the request completed and (if the step declared an expected
	// pattern) the response matches it. Status codes are intentionally NOT
	// treated as success/failure here: a 401/403 is a meaningful result an
	// access-control test wants to observe, not an executor error.
	success := true
	if step.ExpectedOutputPattern != "" {
		success, _ = regexp.MatchString(step.ExpectedOutputPattern, output)
	}

	return &pipeline.ExecutionResult{
		StepID:          step.ID,
		CampaignID:      campaignID,
		CommandExecuted: step.Command,
		Output:          output,
		Success:         success,
		ExecutedAt:      start,
		DurationMs:      int(time.Since(start).Milliseconds()),
		Evidence: []pipeline.Evidence{{
			Type:        "http_response",
			Content:     output,
			Timestamp:   time.Now(),
			Description: fmt.Sprintf("%s %s → %s", opts.method, opts.url, resp.Status),
		}},
	}, nil
}

// buildHTTPReq constructs a fresh *http.Request from the parsed options. Shared
// by the single-request path and each goroutine of a --race burst (an
// *http.Request can't be reused across concurrent sends).
func buildHTTPReq(ctx context.Context, opts httpReqOptions) (*http.Request, error) {
	var bodyReader io.Reader
	if opts.body != "" {
		bodyReader = strings.NewReader(opts.body)
	}
	req, err := http.NewRequestWithContext(ctx, opts.method, opts.url, bodyReader)
	if err != nil {
		return nil, fmt.Errorf("building request: %w", err)
	}
	for _, h := range opts.headers {
		k, v, ok := splitHeader(h)
		if !ok {
			return nil, fmt.Errorf("malformed --header %q (want \"Key: Value\")", h)
		}
		req.Header.Set(k, v)
	}
	// Default the Content-Type for a request that carries a body. Go's client
	// sends no Content-Type of its own, and JSON APIs (crAPI, most REST back
	// ends) reject a bodied POST/PUT without one — a 415 that silently breaks
	// the first step of an authenticated chain. If the step set its own
	// Content-Type header, we never override it.
	if opts.body != "" && req.Header.Get("Content-Type") == "" {
		req.Header.Set("Content-Type", "application/json")
	}
	return req, nil
}

// fireRace sends opts.race identical requests concurrently, released together
// for maximum simultaneity, and returns a representative response (first 2xx,
// else the first) plus a summary of how many succeeded — the signal for a
// business-logic race condition. Non-representative response bodies are closed.
func fireRace(ctx context.Context, client *http.Client, opts httpReqOptions) (*http.Response, string, error) {
	n := opts.race
	type outcome struct {
		resp *http.Response
		err  error
	}
	outcomes := make([]outcome, n)
	start := make(chan struct{})
	var wg sync.WaitGroup
	for i := 0; i < n; i++ {
		wg.Add(1)
		go func(i int) {
			defer wg.Done()
			r, berr := buildHTTPReq(ctx, opts)
			if berr != nil {
				outcomes[i] = outcome{err: berr}
				return
			}
			<-start // block until every goroutine is ready, then fire together
			resp, derr := client.Do(r)
			outcomes[i] = outcome{resp: resp, err: derr}
		}(i)
	}
	close(start)
	wg.Wait()

	var resps []*http.Response
	var firstErr error
	ok2xx := 0
	for i := range outcomes {
		if outcomes[i].err != nil {
			if firstErr == nil {
				firstErr = outcomes[i].err
			}
			continue
		}
		if outcomes[i].resp == nil {
			continue
		}
		resps = append(resps, outcomes[i].resp)
		if sc := outcomes[i].resp.StatusCode; sc >= 200 && sc < 300 {
			ok2xx++
		}
	}
	if len(resps) == 0 {
		return nil, "", firstErr
	}
	rep := resps[0]
	for _, r := range resps {
		if r.StatusCode >= 200 && r.StatusCode < 300 {
			rep = r
			break
		}
	}
	for _, r := range resps {
		if r != rep {
			r.Body.Close()
		}
	}
	summary := fmt.Sprintf("[race] fired %d concurrent identical requests — %d returned 2xx. "+
		"A single-use or limited action succeeding more than once here is a business-logic race condition (TOCTOU).\n\n", n, ok2xx)
	return rep, summary, nil
}

// parseHTTPReqArgs parses the argv (excluding the leading "httpreq") into
// httpReqOptions. Unknown flags are an error so a malformed step fails loudly
// instead of silently sending the wrong request.
func parseHTTPReqArgs(args []string) (httpReqOptions, error) {
	opts := httpReqOptions{method: http.MethodGet}
	for i := 0; i < len(args); i++ {
		switch args[i] {
		case "--method", "-X":
			v, err := nextArg(args, &i, "--method")
			if err != nil {
				return opts, err
			}
			opts.method = strings.ToUpper(v)
		case "--url", "-u":
			v, err := nextArg(args, &i, "--url")
			if err != nil {
				return opts, err
			}
			opts.url = v
		case "--header", "-H":
			v, err := nextArg(args, &i, "--header")
			if err != nil {
				return opts, err
			}
			opts.headers = append(opts.headers, v)
		case "--body", "-d":
			v, err := nextArg(args, &i, "--body")
			if err != nil {
				return opts, err
			}
			opts.body = v
		case "--follow":
			opts.follow = true
		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})
		case "--timeout":
			v, err := nextArg(args, &i, "--timeout")
			if err != nil {
				return opts, err
			}
			secs := 0
			if _, serr := fmt.Sscanf(v, "%d", &secs); serr != nil || secs <= 0 {
				return opts, fmt.Errorf("--timeout wants a positive integer number of seconds, got %q", v)
			}
			opts.timeout = time.Duration(secs) * time.Second
		case "--race":
			v, err := nextArg(args, &i, "--race")
			if err != nil {
				return opts, err
			}
			n := 0
			if _, serr := fmt.Sscanf(v, "%d", &n); serr != nil || n < 2 {
				return opts, fmt.Errorf("--race wants an integer >= 2 (concurrent requests), got %q", v)
			}
			if n > maxRaceRequests {
				n = maxRaceRequests
			}
			opts.race = n
		default:
			return opts, fmt.Errorf("unknown httpreq flag %q", args[i])
		}
	}
	return opts, nil
}

// nextArg returns args[*i+1] and advances *i, or errors if the flag has no value.
func nextArg(args []string, i *int, flag string) (string, error) {
	if *i+1 >= len(args) {
		return "", fmt.Errorf("%s requires a value", flag)
	}
	*i++
	return args[*i], nil
}

// splitHeader splits a "Key: Value" (or "Key:Value") header into its parts.
func splitHeader(h string) (key, value string, ok bool) {
	idx := strings.Index(h, ":")
	if idx <= 0 {
		return "", "", false
	}
	key = strings.TrimSpace(h[:idx])
	value = strings.TrimSpace(h[idx+1:])
	if key == "" {
		return "", "", false
	}
	return key, value, true
}

// captureValue extracts a value from a response per a capture selector:
//
//	header:Name  → the response header "Name" (e.g. header:Location, header:Set-Cookie)
//	$.a.b / a.b  → a dotted path into the JSON response body; numeric segments
//	               index into arrays (e.g. $.tokens.0.value)
//
// Returns ok=false when the selector doesn't resolve, so a missing value
// leaves the chain variable unset rather than binding it to "".
func captureValue(selector string, resp *http.Response, body []byte) (string, bool) {
	if rest, ok := strings.CutPrefix(selector, "header:"); ok {
		v := resp.Header.Get(strings.TrimSpace(rest))
		return v, v != ""
	}
	path := strings.TrimPrefix(selector, "$.")
	path = strings.TrimPrefix(path, "$")
	path = strings.Trim(path, ".")
	if path == "" {
		return "", false
	}
	var doc any
	if err := json.Unmarshal(body, &doc); err != nil {
		return "", false
	}
	return jsonPathString(doc, strings.Split(path, "."))
}

// jsonPathString walks a decoded JSON document along the given segments and
// returns the leaf rendered as a string. Object keys and numeric array indices
// are both supported.
func jsonPathString(node any, segments []string) (string, bool) {
	cur := node
	for _, seg := range segments {
		switch typed := cur.(type) {
		case map[string]any:
			next, ok := typed[seg]
			if !ok {
				return "", false
			}
			cur = next
		case []any:
			idx, err := strconv.Atoi(seg)
			if err != nil || idx < 0 || idx >= len(typed) {
				return "", false
			}
			cur = typed[idx]
		default:
			return "", false
		}
	}
	switch leaf := cur.(type) {
	case string:
		return leaf, true
	case float64:
		return strconv.FormatFloat(leaf, 'f', -1, 64), true
	case bool:
		return strconv.FormatBool(leaf), true
	case nil:
		return "", false
	default:
		// Object/array leaf — re-encode so the whole subtree is still usable.
		b, err := json.Marshal(leaf)
		if err != nil {
			return "", false
		}
		return string(b), true
	}
}

// varPattern matches a {{name}} chain-variable placeholder.
var varPattern = regexp.MustCompile(`\{\{\s*([A-Za-z0-9_]+)\s*\}\}`)

// substituteVars replaces {{name}} placeholders in cmd with captured values.
// Unknown placeholders are left untouched so a missing capture is visible in
// the command output (and the step simply fails) rather than silently becoming
// an empty string that might, e.g., send an unauthenticated request that looks
// like a successful access-control bypass.
func substituteVars(cmd string, vars map[string]string) string {
	if len(vars) == 0 || !strings.Contains(cmd, "{{") {
		return cmd
	}
	return varPattern.ReplaceAllStringFunc(cmd, func(m string) string {
		name := varPattern.FindStringSubmatch(m)[1]
		if v, ok := vars[name]; ok {
			return v
		}
		return m
	})
}

// formatHTTPResponse renders a request/response into the compact, greppable
// text form recorded as the step output and evidence.
func formatHTTPResponse(opts httpReqOptions, resp *http.Response, body string, truncated bool) string {
	var b strings.Builder
	fmt.Fprintf(&b, "%s %s\n", opts.method, opts.url)
	fmt.Fprintf(&b, "HTTP %s\n", resp.Status)
	// A small, stable set of response headers that matter for API testing.
	for _, h := range []string{"Content-Type", "Content-Length", "Location", "Set-Cookie", "WWW-Authenticate"} {
		if v := resp.Header.Get(h); v != "" {
			fmt.Fprintf(&b, "%s: %s\n", h, v)
		}
	}
	b.WriteString("\n")
	b.WriteString(body)
	if truncated {
		fmt.Fprintf(&b, "\n… [response body truncated at %d bytes]", httpReqBodyCap)
	}
	return b.String()
}
