diff --git a/internal/api/blastradius_test.go b/internal/api/blastradius_test.go new file mode 100644 index 0000000..ddd6e5a --- /dev/null +++ b/internal/api/blastradius_test.go @@ -0,0 +1,181 @@ +// Copyright 2026 OpenSourceOM +// SPDX-License-Identifier: Apache-2.0 + +package api + +import ( + "context" + "crypto/rand" + "encoding/hex" + "encoding/json" + "net/http" + "net/http/httptest" + "net/url" + "os" + "path/filepath" + "runtime" + "strings" + "testing" + + "github.com/OpenSourceOM/core/internal/graph" + "github.com/OpenSourceOM/core/internal/migrate" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestHandleBlastRadiusRequiresIdentity(t *testing.T) { + server := NewServer(nil, "") + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/v1/identity/blast-radius", nil) + + server.Handler().ServeHTTP(recorder, request) + + if recorder.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want %d", recorder.Code, http.StatusBadRequest) + } + var body map[string]string + if err := json.NewDecoder(recorder.Body).Decode(&body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["error"] != "identity_id or name required" { + t.Fatalf("error = %q", body["error"]) + } +} + +func TestHandleBlastRadius(t *testing.T) { + ctx := context.Background() + store := openBlastRadiusStore(t) + + const ( + user = "identity:dev" + role = "identity:admin" + db = "datastore:prod" + net = "network:public" + ) + batch := graph.Batch{ + Nodes: []graph.Node{ + {ID: user, Type: graph.NodeIdentity, Name: "dev"}, + {ID: role, Type: graph.NodeIdentity, Name: "AdminRole"}, + {ID: db, Type: graph.NodeDatastore, Name: "prod-db"}, + {ID: net, Type: graph.NodeNetwork, Name: "public"}, + }, + Edges: []graph.Edge{ + {ID: user + "|" + role + "|" + graph.EdgeAssumes, SourceID: user, TargetID: role, Type: graph.EdgeAssumes}, + {ID: role + "|" + db + "|" + graph.EdgeCanAccess, SourceID: role, TargetID: db, Type: graph.EdgeCanAccess}, + {ID: user + "|" + net + "|" + graph.EdgeReachable, SourceID: user, TargetID: net, Type: graph.EdgeReachable}, + }, + } + if err := store.UpsertBatch(ctx, batch); err != nil { + t.Fatalf("seed: %v", err) + } + + handler := NewServer(store, "").Handler() + wantIDs := db + "," + role + + t.Run("identity id", func(t *testing.T) { + result := getBlastRadius(t, handler, "/v1/identity/blast-radius?identity_id="+url.QueryEscape(user)) + if result.Identity.Name != "dev" || result.MaxDepth != 6 { + t.Fatalf("identity = %q depth %d", result.Identity.Name, result.MaxDepth) + } + if got := blastRadiusIDs(result); got != wantIDs { + t.Fatalf("reachable = %s, want %s", got, wantIDs) + } + }) + + t.Run("name", func(t *testing.T) { + result := getBlastRadius(t, handler, "/v1/identity/blast-radius?name=dev") + if result.Identity.ID != user { + t.Fatalf("identity id = %s, want %s", result.Identity.ID, user) + } + if got := blastRadiusIDs(result); got != wantIDs { + t.Fatalf("reachable = %s, want %s", got, wantIDs) + } + }) + + t.Run("unknown name", func(t *testing.T) { + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/v1/identity/blast-radius?name=missing", nil) + handler.ServeHTTP(recorder, request) + if recorder.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want %d", recorder.Code, http.StatusBadRequest) + } + }) +} + +func getBlastRadius(t *testing.T, handler http.Handler, path string) graph.BlastRadiusResult { + t.Helper() + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, path, nil) + handler.ServeHTTP(recorder, request) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, body %s", recorder.Code, recorder.Body.String()) + } + var result graph.BlastRadiusResult + if err := json.NewDecoder(recorder.Body).Decode(&result); err != nil { + t.Fatalf("decode: %v", err) + } + return result +} + +func blastRadiusIDs(result graph.BlastRadiusResult) string { + ids := make([]string, len(result.Reachable)) + for i, node := range result.Reachable { + ids[i] = node.ID + } + return strings.Join(ids, ",") +} + +func openBlastRadiusStore(t *testing.T) *graph.Store { + t.Helper() + adminURL := os.Getenv("TEST_DATABASE_URL") + if adminURL == "" { + t.Skip("TEST_DATABASE_URL is not set") + } + + ctx := context.Background() + admin, err := pgxpool.New(ctx, adminURL) + if err != nil { + t.Fatalf("connect postgres: %v", err) + } + t.Cleanup(admin.Close) + + var raw [8]byte + if _, err := rand.Read(raw[:]); err != nil { + t.Fatal(err) + } + dbName := "ombr" + hex.EncodeToString(raw[:]) + ident := pgx.Identifier{dbName}.Sanitize() + if _, err := admin.Exec(ctx, "CREATE DATABASE "+ident); err != nil { + t.Fatalf("create database: %v", err) + } + t.Cleanup(func() { + if _, err := admin.Exec(context.Background(), "DROP DATABASE "+ident+" WITH (FORCE)"); err != nil { + t.Errorf("drop database %s: %v", dbName, err) + } + }) + + databaseURL, err := url.Parse(adminURL) + if err != nil { + t.Fatal(err) + } + databaseURL.Path = "/" + dbName + + if err := migrate.Run(ctx, databaseURL.String(), blastRadiusMigrationsDir(t)); err != nil { + t.Fatalf("migrate: %v", err) + } + store, err := graph.NewStore(ctx, databaseURL.String()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(store.Close) + return store +} + +func blastRadiusMigrationsDir(t *testing.T) string { + t.Helper() + _, file, _, ok := runtime.Caller(0) + if !ok { + t.Fatal("locate test file") + } + return filepath.Join(filepath.Dir(file), "..", "..", "migrations") +} diff --git a/internal/export/findings_test.go b/internal/export/findings_test.go new file mode 100644 index 0000000..89c0758 --- /dev/null +++ b/internal/export/findings_test.go @@ -0,0 +1,251 @@ +// Copyright 2026 OpenSourceOM +// SPDX-License-Identifier: Apache-2.0 + +package export + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/OpenSourceOM/core/internal/graph" +) + +func TestWriteSIEM(t *testing.T) { + when := time.Date(2026, 9, 30, 15, 4, 5, 0, time.UTC) + records := []FindingRecord{ + { + Timestamp: when, + Finding: graph.Node{ + ID: "finding:cve:web", + Type: graph.NodeFinding, + Name: "CVE-2021-44228", + Properties: map[string]any{ + "severity": "critical", + "title": "Log4Shell", + }, + }, + AffectedID: "workload:web", + AffectedName: "web-1", + AffectedType: graph.NodeWorkload, + Path: []string{"internet:global", "workload:web", "datastore:prod"}, + }, + { + Timestamp: when, + Finding: graph.Node{ + ID: "finding:public:bucket", + Type: graph.NodeFinding, + Name: "Public bucket", + }, + AffectedID: "datastore:logs", + AffectedName: "logs", + AffectedType: graph.NodeDatastore, + }, + } + + var buf bytes.Buffer + if err := WriteSIEM(&buf, records); err != nil { + t.Fatalf("WriteSIEM: %v", err) + } + lines := strings.Split(strings.TrimSuffix(buf.String(), "\n"), "\n") + if len(lines) != len(records) { + t.Fatalf("lines = %d, want %d", len(lines), len(records)) + } + for i, line := range lines { + var got FindingRecord + if err := json.Unmarshal([]byte(line), &got); err != nil { + t.Fatalf("line %d: %v", i, err) + } + if !got.Timestamp.Equal(records[i].Timestamp) { + t.Fatalf("line %d timestamp = %s", i, got.Timestamp) + } + if got.Finding.ID != records[i].Finding.ID || got.AffectedName != records[i].AffectedName { + t.Fatalf("line %d = finding %s affected %s", i, got.Finding.ID, got.AffectedName) + } + if strings.Join(got.Path, ",") != strings.Join(records[i].Path, ",") { + t.Fatalf("line %d path = %v", i, got.Path) + } + } + + buf.Reset() + if err := WriteSIEM(&buf, nil); err != nil { + t.Fatalf("WriteSIEM empty: %v", err) + } + if buf.Len() != 0 { + t.Fatalf("empty export wrote %q", buf.String()) + } +} + +func TestPostSlackTruncatesAfterTwenty(t *testing.T) { + if err := PostSlack(context.Background(), "", nil); err == nil || !strings.Contains(err.Error(), "SLACK_WEBHOOK_URL") { + t.Fatalf("empty webhook error = %v", err) + } + + var text string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Content-Type") != "application/json" { + t.Errorf("content-type = %s", r.Header.Get("Content-Type")) + } + var payload map[string]string + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Errorf("decode: %v", err) + } + text = payload["text"] + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(server.Close) + + records := make([]FindingRecord, 21) + for i := range records { + records[i] = FindingRecord{ + Finding: graph.Node{ + Name: fmt.Sprintf("name-%02d", i), + Properties: map[string]any{ + "severity": "high", + "title": fmt.Sprintf("title-%02d", i), + }, + }, + AffectedName: fmt.Sprintf("res-%02d", i), + } + } + records[0].Finding.Properties["title"] = "" + + if err := PostSlack(context.Background(), server.URL, records); err != nil { + t.Fatalf("PostSlack: %v", err) + } + lines := strings.Split(text, "\n") + if len(lines) != 22 { + t.Fatalf("lines = %d, want header, 20 findings, and the remainder", len(lines)) + } + if lines[0] != "*OpenSourceOM findings export*" { + t.Fatalf("header = %q", lines[0]) + } + if lines[1] != "• [HIGH] name-00 — res-00" { + t.Fatalf("first finding = %q", lines[1]) + } + if lines[20] != "• [HIGH] title-19 — res-19" { + t.Fatalf("twentieth finding = %q", lines[20]) + } + if lines[21] != "…and 1 more" { + t.Fatalf("remainder = %q", lines[21]) + } + if strings.Contains(text, "title-20") { + t.Fatalf("message includes the finding past the cutoff:\n%s", text) + } + + fail := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "nope", http.StatusInternalServerError) + })) + t.Cleanup(fail.Close) + err := PostSlack(context.Background(), fail.URL, records[:1]) + if err == nil || !strings.Contains(err.Error(), "500") || !strings.Contains(err.Error(), "nope") { + t.Fatalf("webhook error = %v", err) + } +} + +func TestCreateJiraIssues(t *testing.T) { + _, err := CreateJiraIssues(context.Background(), JiraConfig{}, []FindingRecord{{}}) + if err == nil || !strings.Contains(err.Error(), "JIRA_URL") { + t.Fatalf("missing config error = %v", err) + } + + var ( + user, pass string + bodies []map[string]any + ) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/rest/api/3/issue" { + t.Errorf("path = %s", r.URL.Path) + } + user, pass, _ = r.BasicAuth() + body, err := io.ReadAll(r.Body) + if err != nil { + t.Errorf("read body: %v", err) + } + var payload map[string]any + if err := json.Unmarshal(body, &payload); err != nil { + t.Errorf("decode: %v", err) + } + bodies = append(bodies, payload) + w.WriteHeader(http.StatusCreated) + })) + t.Cleanup(server.Close) + + record := FindingRecord{ + Finding: graph.Node{ + Name: "fallback", + Properties: map[string]any{ + "severity": "critical", + "title": "Log4Shell", + "description": "Remote code execution", + }, + }, + AffectedName: "web-1", + AffectedType: graph.NodeWorkload, + } + cfg := JiraConfig{ + BaseURL: server.URL + "/", + Email: "bot@example.com", + APIToken: "tok", + Project: "SEC", + } + created, err := CreateJiraIssues(context.Background(), cfg, []FindingRecord{record}) + if err != nil { + t.Fatalf("CreateJiraIssues: %v", err) + } + if created != 1 { + t.Fatalf("created = %d, want 1", created) + } + if user != cfg.Email || pass != cfg.APIToken { + t.Fatalf("basic auth = %s:%s", user, pass) + } + if len(bodies) != 1 { + t.Fatalf("requests = %d", len(bodies)) + } + fields, _ := bodies[0]["fields"].(map[string]any) + project, _ := fields["project"].(map[string]any) + if project["key"] != "SEC" { + t.Fatalf("project = %v", fields["project"]) + } + if fields["summary"] != "[CRITICAL] Log4Shell" { + t.Fatalf("summary = %v", fields["summary"]) + } + description := jiraDescriptionText(t, fields["description"]) + if description != "Remote code execution\nAffected: web-1 (Workload)" { + t.Fatalf("description = %q", description) + } + + denied := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusForbidden) + })) + t.Cleanup(denied.Close) + cfg.BaseURL = denied.URL + created, err = CreateJiraIssues(context.Background(), cfg, []FindingRecord{record}) + if err == nil || !strings.Contains(err.Error(), "403") || created != 0 { + t.Fatalf("denied create = (%d, %v)", created, err) + } +} + +func jiraDescriptionText(t *testing.T, description any) string { + t.Helper() + doc, _ := description.(map[string]any) + content, _ := doc["content"].([]any) + if len(content) != 1 { + t.Fatalf("description content = %v", description) + } + paragraph, _ := content[0].(map[string]any) + parts, _ := paragraph["content"].([]any) + if len(parts) != 1 { + t.Fatalf("paragraph = %v", paragraph) + } + part, _ := parts[0].(map[string]any) + text, _ := part["text"].(string) + return text +} diff --git a/internal/graph/blastradius_test.go b/internal/graph/blastradius_test.go new file mode 100644 index 0000000..1873f0d --- /dev/null +++ b/internal/graph/blastradius_test.go @@ -0,0 +1,130 @@ +// Copyright 2026 OpenSourceOM +// SPDX-License-Identifier: Apache-2.0 + +package graph_test + +import ( + "context" + "strings" + "testing" + + "github.com/OpenSourceOM/core/internal/graph" +) + +func TestBlastRadiusFollowsAssumedRole(t *testing.T) { + ctx := context.Background() + store := openQueryFixture(t) + + const ( + user = "identity:dev" + role = "identity:admin" + db = "datastore:prod" + net = "network:public" + ) + batch := graph.Batch{ + Nodes: []graph.Node{ + {ID: user, Type: graph.NodeIdentity, Name: "dev"}, + {ID: role, Type: graph.NodeIdentity, Name: "AdminRole"}, + {ID: db, Type: graph.NodeDatastore, Name: "prod-db"}, + {ID: net, Type: graph.NodeNetwork, Name: "public"}, + }, + Edges: []graph.Edge{ + blastEdge(user, role, graph.EdgeAssumes), + blastEdge(role, db, graph.EdgeCanAccess), + blastEdge(user, net, graph.EdgeReachable), + }, + } + if err := store.UpsertBatch(ctx, batch); err != nil { + t.Fatalf("seed: %v", err) + } + + result, err := store.BlastRadius(ctx, user, 6) + if err != nil { + t.Fatalf("BlastRadius: %v", err) + } + if result.Identity.ID != user || result.Identity.Name != "dev" { + t.Fatalf("identity = %s %q", result.Identity.ID, result.Identity.Name) + } + if result.MaxDepth != 6 { + t.Fatalf("MaxDepth = %d, want 6", result.MaxDepth) + } + got := blastIDs(result.Reachable) + want := []string{db, role} + if strings.Join(got, ",") != strings.Join(want, ",") { + t.Fatalf("reachable = %v, want datastore via the assumed role, not the network", got) + } + wantSummary := "Identity dev can reach 2 resources (0 workloads, 1 datastores, 0 networks)" + if result.Summary != wantSummary { + t.Fatalf("summary = %q, want %q", result.Summary, wantSummary) + } +} + +func TestBlastRadiusStopsAtMaxDepth(t *testing.T) { + ctx := context.Background() + store := openQueryFixture(t) + + const ( + walker = "identity:walker" + hopA = "workload:a" + hopB = "workload:b" + hopC = "datastore:c" + net = "network:ignored" + ) + batch := graph.Batch{ + Nodes: []graph.Node{ + {ID: walker, Type: graph.NodeIdentity, Name: "walker"}, + {ID: hopA, Type: graph.NodeWorkload, Name: "a"}, + {ID: hopB, Type: graph.NodeWorkload, Name: "b"}, + {ID: hopC, Type: graph.NodeDatastore, Name: "c"}, + {ID: net, Type: graph.NodeNetwork, Name: "ignored"}, + }, + Edges: []graph.Edge{ + blastEdge(walker, hopA, graph.EdgeCanAccess), + blastEdge(hopA, hopB, graph.EdgeCanAccess), + blastEdge(hopB, hopC, graph.EdgeCanAccess), + blastEdge(walker, net, graph.EdgeReachable), + }, + } + if err := store.UpsertBatch(ctx, batch); err != nil { + t.Fatalf("seed: %v", err) + } + + shallow, err := store.BlastRadius(ctx, walker, 2) + if err != nil { + t.Fatalf("BlastRadius depth 2: %v", err) + } + if shallow.MaxDepth != 2 { + t.Fatalf("MaxDepth = %d, want 2", shallow.MaxDepth) + } + if got := strings.Join(blastIDs(shallow.Reachable), ","); got != hopA+","+hopB { + t.Fatalf("depth 2 reachable = %s, want %s and %s", got, hopA, hopB) + } + + full, err := store.BlastRadius(ctx, walker, 0) + if err != nil { + t.Fatalf("BlastRadius default depth: %v", err) + } + if full.MaxDepth != 6 { + t.Fatalf("default MaxDepth = %d, want 6", full.MaxDepth) + } + if got := strings.Join(blastIDs(full.Reachable), ","); got != hopC+","+hopA+","+hopB { + t.Fatalf("default reachable = %s, want the chain and not the network", got) + } +} + +func blastEdge(source, target, edgeType string) graph.Edge { + return graph.Edge{ + ID: source + "|" + target + "|" + edgeType, + SourceID: source, + TargetID: target, + Type: edgeType, + } +} + +func blastIDs(nodes []graph.Node) []string { + ids := make([]string, len(nodes)) + for i, node := range nodes { + ids[i] = node.ID + } + return ids +}