From 9c0102a2d5b45050dc7e53cb9a1829b104d72e40 Mon Sep 17 00:00:00 2001 From: hunterinvariants Date: Sat, 6 Jun 2026 16:18:49 +0200 Subject: [PATCH 001/137] Add collectors exports approvals and packaging --- .github/ISSUE_TEMPLATE/bug_report.yml | 27 ++ .github/ISSUE_TEMPLATE/collector_request.yml | 27 ++ .github/ISSUE_TEMPLATE/feature_request.yml | 30 ++ .github/workflows/release.yml | 49 +++ README.md | 27 +- cmd/promtact/main.go | 4 + cmd/promtactl/main.go | 78 +++- docs/architecture.md | 12 + docs/operations.md | 65 ++++ docs/roadmap.md | 6 +- internal/collectors/collectors.go | 375 +++++++++++++++++++ internal/collectors/collectors_test.go | 80 ++++ internal/domain/types.go | 20 +- internal/exporter/webhook.go | 66 ++++ internal/exporter/webhook_test.go | 33 ++ internal/response/planner.go | 55 +-- internal/server/server.go | 77 +++- internal/server/server_test.go | 84 +++++ internal/store/store.go | 19 + internal/store/store_test.go | 26 ++ internal/telemetry/jsonl.go | 1 + packaging/systemd/oadtd.service | 23 ++ packaging/windows/install-service.ps1 | 37 ++ web/app.js | 18 +- 24 files changed, 1191 insertions(+), 48 deletions(-) create mode 100644 .github/ISSUE_TEMPLATE/bug_report.yml create mode 100644 .github/ISSUE_TEMPLATE/collector_request.yml create mode 100644 .github/ISSUE_TEMPLATE/feature_request.yml create mode 100644 .github/workflows/release.yml create mode 100644 docs/operations.md create mode 100644 internal/collectors/collectors.go create mode 100644 internal/collectors/collectors_test.go create mode 100644 internal/exporter/webhook.go create mode 100644 internal/exporter/webhook_test.go create mode 100644 packaging/systemd/oadtd.service create mode 100644 packaging/windows/install-service.ps1 diff --git a/.github/ISSUE_TEMPLATE/bug_report.yml b/.github/ISSUE_TEMPLATE/bug_report.yml new file mode 100644 index 0000000..3f1d7ec --- /dev/null +++ b/.github/ISSUE_TEMPLATE/bug_report.yml @@ -0,0 +1,27 @@ +name: Bug report +description: Report a defect in defensive behavior, API, UI, or packaging. +title: "[bug] " +labels: ["bug"] +body: + - type: textarea + id: summary + attributes: + label: Summary + description: What went wrong? + validations: + required: true + - type: textarea + id: steps + attributes: + label: Reproduction Steps + description: Include commands, sample defensive telemetry, and expected vs actual behavior. + validations: + required: true + - type: textarea + id: safety + attributes: + label: Safety Check + description: Do not include exploit code, credentials, malware samples, or unauthorized target details. + validations: + required: true + diff --git a/.github/ISSUE_TEMPLATE/collector_request.yml b/.github/ISSUE_TEMPLATE/collector_request.yml new file mode 100644 index 0000000..909c1c6 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/collector_request.yml @@ -0,0 +1,27 @@ +name: Collector request +description: Request or improve a defensive telemetry collector. +title: "[collector] " +labels: ["collector"] +body: + - type: input + id: source + attributes: + label: Telemetry Source + placeholder: Sysmon, auditd, Zeek, Suricata, proxy, EDR, SIEM + validations: + required: true + - type: textarea + id: sample + attributes: + label: Redacted Sample + description: Provide redacted defensive log samples only. Remove secrets, usernames if sensitive, public IPs if needed, and customer identifiers. + validations: + required: true + - type: textarea + id: mapping + attributes: + label: Desired Mapping + description: Which Promtact event fields should this populate? + validations: + required: false + diff --git a/.github/ISSUE_TEMPLATE/feature_request.yml b/.github/ISSUE_TEMPLATE/feature_request.yml new file mode 100644 index 0000000..0975751 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/feature_request.yml @@ -0,0 +1,30 @@ +name: Feature request +description: Suggest defensive functionality or operational improvements. +title: "[feature] " +labels: ["enhancement"] +body: + - type: textarea + id: problem + attributes: + label: Problem + description: What defensive workflow or operational gap should this solve? + validations: + required: true + - type: textarea + id: proposal + attributes: + label: Proposal + description: What should Promtact do? + validations: + required: true + - type: dropdown + id: edition + attributes: + label: Likely Edition + options: + - Community core + - Commercial/enterprise + - Unsure + validations: + required: true + diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..5647182 --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,49 @@ +name: release + +on: + push: + tags: + - "v*.*.*" + +permissions: + contents: write + +jobs: + build: + runs-on: ubuntu-latest + strategy: + matrix: + goos: [linux, windows] + goarch: [amd64, arm64] + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.24.x" + - run: go test ./... + - name: Build + shell: bash + run: | + suffix="" + if [ "${{ matrix.goos }}" = "windows" ]; then suffix=".exe"; fi + mkdir -p dist + GOOS=${{ matrix.goos }} GOARCH=${{ matrix.goarch }} go build -o "dist/promtact-${{ matrix.goos }}-${{ matrix.goarch }}${suffix}" ./cmd/promtact + GOOS=${{ matrix.goos }} GOARCH=${{ matrix.goarch }} go build -o "dist/promtactl-${{ matrix.goos }}-${{ matrix.goarch }}${suffix}" ./cmd/promtactl + - uses: actions/upload-artifact@v4 + with: + name: promtact-${{ matrix.goos }}-${{ matrix.goarch }} + path: dist/* + + publish: + needs: build + runs-on: ubuntu-latest + steps: + - uses: actions/download-artifact@v4 + with: + path: dist + merge-multiple: true + - uses: softprops/action-gh-release@v2 + with: + files: dist/* + generate_release_notes: true + diff --git a/README.md b/README.md index 35f7b57..e0254f6 100644 --- a/README.md +++ b/README.md @@ -21,6 +21,8 @@ malware behavior, or autonomous propagation. Demo data generates telemetry only. - `promtactl replay` for safe JSONL telemetry replay into the ingest API. - Browser dashboard with asset risk graph, alerts, events, rules, and response actions. +- Alert webhook export for SIEM-style integrations. +- systemd and Windows service starter packaging. - AGPLv3-or-later community license, commercial dual-license path, and CLA from day 1. @@ -112,12 +114,14 @@ Useful endpoints: ## Runtime Options ```text ---addr HTTP listen address, default :8080 ---web static dashboard directory, default web ---demo load safe demo telemetry at startup ---data optional JSON snapshot path for local persistence ---policy optional JSON policy configuration path ---api-token optional token for POST endpoints, defaults to PROMTACT_API_TOKEN +--addr HTTP listen address, default :8080 +--web static dashboard directory, default web +--demo load safe demo telemetry at startup +--data optional JSON snapshot path for local persistence +--policy optional JSON policy configuration path +--api-token optional token for POST endpoints, defaults to PROMTACT_API_TOKEN +--alert-webhook-url optional SIEM/webhook URL for new alerts +--alert-webhook-token optional bearer token for alert webhook ``` When `--api-token` or `PROMTACT_API_TOKEN` is set, read endpoints remain available @@ -159,6 +163,17 @@ Validate a file without sending it: go run ./cmd/promtactl replay --file examples\demo-events.jsonl --dry-run ``` +Normalize external defensive logs to Promtact JSONL: + +```powershell +go run ./cmd/promtactl collect --source suricata-eve --file eve.json --output events.jsonl +go run ./cmd/promtactl collect --source zeek-conn --file conn.log --output events.jsonl +go run ./cmd/promtactl collect --source sysmon-json --file sysmon.jsonl --output events.jsonl +go run ./cmd/promtactl collect --source auditd --file audit.log --output events.jsonl +``` + +Operations notes are in [docs/operations.md](docs/operations.md). + ## License Community distribution is licensed under AGPL-3.0-or-later. diff --git a/cmd/promtact/main.go b/cmd/promtact/main.go index 286e7d0..d352b35 100644 --- a/cmd/promtact/main.go +++ b/cmd/promtact/main.go @@ -16,6 +16,8 @@ func main() { dataPath := flag.String("data", "", "optional JSON snapshot path for local persistence") policyPath := flag.String("policy", "", "optional JSON policy configuration path") apiToken := flag.String("api-token", os.Getenv("PROMTACT_API_TOKEN"), "optional API token for write endpoints") + alertWebhookURL := flag.String("alert-webhook-url", os.Getenv("PROMTACT_ALERT_WEBHOOK_URL"), "optional SIEM/webhook URL for new alerts") + alertWebhookToken := flag.String("alert-webhook-token", os.Getenv("PROMTACT_ALERT_WEBHOOK_TOKEN"), "optional bearer token for alert webhook") withDemo := flag.Bool("demo", false, "load safe demo telemetry at startup") flag.Parse() @@ -34,6 +36,8 @@ func main() { APIToken: *apiToken, Policy: runtimeConfig.PolicyConfig(), CorrelationWindow: window, + AlertWebhookURL: *alertWebhookURL, + AlertWebhookToken: *alertWebhookToken, }) if err != nil { log.Fatal(err) diff --git a/cmd/promtactl/main.go b/cmd/promtactl/main.go index b6c0578..9bc22e6 100644 --- a/cmd/promtactl/main.go +++ b/cmd/promtactl/main.go @@ -13,6 +13,7 @@ import ( "strings" "time" + "github.com/hunterinvariants/promtact/internal/collectors" "github.com/hunterinvariants/promtact/internal/domain" "github.com/hunterinvariants/promtact/internal/telemetry" ) @@ -25,6 +26,10 @@ func main() { } switch os.Args[1] { + case "collect": + if err := collect(os.Args[2:]); err != nil { + log.Fatal(err) + } case "replay": if err := replay(os.Args[2:]); err != nil { log.Fatal(err) @@ -35,6 +40,50 @@ func main() { } } +func collect(args []string) error { + fs := flag.NewFlagSet("collect", flag.ContinueOnError) + source := fs.String("source", "", "collector source: "+strings.Join(collectors.Sources(), ", ")) + filePath := fs.String("file", "", "source log file, or - for stdin") + outputPath := fs.String("output", "-", "output JSONL file, or - for stdout") + if err := fs.Parse(args); err != nil { + return err + } + if *source == "" { + return errors.New("collect requires --source") + } + if *filePath == "" { + return errors.New("collect requires --file") + } + + input, closeInput, err := openInput(*filePath) + if err != nil { + return err + } + defer closeInput() + + events, err := collectors.Normalize(*source, input) + if err != nil { + return err + } + + output, closeOutput, err := openOutput(*outputPath) + if err != nil { + return err + } + defer closeOutput() + + encoder := json.NewEncoder(output) + for _, event := range events { + if err := encoder.Encode(event); err != nil { + return err + } + } + if *outputPath != "-" { + fmt.Printf("events=%d output=%s\n", len(events), *outputPath) + } + return nil +} + func replay(args []string) error { fs := flag.NewFlagSet("replay", flag.ContinueOnError) filePath := fs.String("file", "", "JSONL event file, or - for stdin") @@ -80,16 +129,34 @@ func replay(args []string) error { } func readEvents(filePath string) ([]domain.Event, error) { - if filePath == "-" { - return telemetry.ReadJSONL(os.Stdin) + input, closeInput, err := openInput(filePath) + if err != nil { + return nil, err } + defer closeInput() + return telemetry.ReadJSONL(input) +} +func openInput(filePath string) (io.Reader, func(), error) { + if filePath == "-" { + return os.Stdin, func() {}, nil + } file, err := os.Open(filePath) if err != nil { - return nil, err + return nil, nil, err + } + return file, func() { _ = file.Close() }, nil +} + +func openOutput(filePath string) (io.Writer, func(), error) { + if filePath == "-" { + return os.Stdout, func() {}, nil + } + file, err := os.Create(filePath) + if err != nil { + return nil, nil, err } - defer file.Close() - return telemetry.ReadJSONL(file) + return file, func() { _ = file.Close() }, nil } func postEvents(client *http.Client, baseURL string, token string, events []domain.Event) (int, error) { @@ -129,5 +196,6 @@ func postEvents(client *http.Client, baseURL string, token string, events []doma func usage() { fmt.Fprintln(os.Stderr, "usage:") + fmt.Fprintln(os.Stderr, " promtactl collect --source suricata-eve --file eve.json --output events.jsonl") fmt.Fprintln(os.Stderr, " promtactl replay --file events.jsonl [--url http://localhost:8080] [--token TOKEN]") } diff --git a/docs/architecture.md b/docs/architecture.md index 840bb86..c7c3bbb 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -25,6 +25,10 @@ static dashboard. This gives teams a safe way to replay logs and simulation traces without adding offensive behavior. +The `collect` command normalizes supported defensive telemetry sources into +Promtact JSONL. Current sources are Sysmon JSON, auditd text, Zeek conn logs, and +Suricata EVE JSON. + ### Domain Model `internal/domain` defines the shared event, alert, asset, rule, and response @@ -63,6 +67,9 @@ appear on the same asset. The window is configurable with `internal/response` creates dry-run response plans. The MVP does not execute containment actions against real systems. +Response actions that would affect hosts, egress, tools, or secrets are marked +as requiring approval before any future execution backend can act on them. + ### Dashboard `web/` provides an operational dashboard for assets, alerts, events, policies, @@ -74,6 +81,11 @@ Write endpoints can be protected with `--api-token` or `PROMTACT_API_TOKEN`. Read endpoints remain available so the dashboard and health checks can load without embedding a token in static assets. +### SIEM/Webhook Export + +When `--alert-webhook-url` is set, newly created alerts are sent to that +endpoint as `promtact.alerts` JSON payloads. + ## Near-Term Production Shape The next architecture step is to split collectors, policy evaluation, durable diff --git a/docs/operations.md b/docs/operations.md new file mode 100644 index 0000000..3914370 --- /dev/null +++ b/docs/operations.md @@ -0,0 +1,65 @@ +# Operations + +## Linux systemd + +Example unit: + +```text +packaging/systemd/promtact.service +``` + +Suggested layout: + +```text +/opt/promtact/promtact +/opt/promtact/promtactl +/etc/promtact/policy.json +/etc/promtact/promtact.env +/var/lib/promtact/state.json +``` + +Create a dedicated user, copy the binaries and policy file, install the unit, +then enable it: + +```bash +sudo useradd --system --home /var/lib/promtact --shell /usr/sbin/nologin promtact +sudo mkdir -p /opt/promtact /etc/promtact /var/lib/promtact +sudo chown -R promtact:promtact /var/lib/promtact +sudo cp packaging/systemd/promtact.service /etc/systemd/system/promtact.service +sudo systemctl daemon-reload +sudo systemctl enable --now promtact +``` + +## Windows Service + +Build or download `promtact.exe`, place it at `C:\Program Files\Promtact\promtact.exe`, +then run PowerShell as Administrator: + +```powershell +.\packaging\windows\install-service.ps1 +``` + +The script registers a Windows service named `Promtact` and stores runtime state +under `C:\ProgramData\Promtact`. + +## Webhook Export + +New alerts can be exported to a SIEM or webhook endpoint: + +```powershell +$env:PROMTACT_ALERT_WEBHOOK_URL="https://siem.example.invalid/promtact" +$env:PROMTACT_ALERT_WEBHOOK_TOKEN="replace-with-token" +go run ./cmd/promtact --demo +``` + +The payload type is `promtact.alerts`. + +## Storage + +Current durable storage is the local JSON snapshot configured with `--data`. +This is suitable for local labs, pilots, and single-node testing. + +SQLite/Postgres is the next storage milestone. It should be implemented behind +the existing store boundary so the API and collectors do not change when the +storage backend changes. + diff --git a/docs/roadmap.md b/docs/roadmap.md index 9ecf507..7f4c03c 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -11,6 +11,10 @@ - Optional token protection for write endpoints. - JSON policy configuration for tool, egress, and correlation-window defaults. - Safe JSONL telemetry replay client. +- Collector normalizers for Sysmon JSON, auditd, Zeek conn, and Suricata EVE. +- Alert webhook export. +- Response approval state for planned actions. +- systemd and Windows service starter packaging. - AGPLv3-or-later plus commercial dual-license path. - CLA requirement from day 1. @@ -22,7 +26,7 @@ Status: implemented in this repository. - Authenticated API. - Policy reload without restart. - Signed tool manifests for AI-agent and MCP surfaces. -- Collector adapters for Sysmon, auditd, Zeek, Suricata, and proxy logs. +- Long-running collector agents for Sysmon, auditd, Zeek, Suricata, and proxy logs. - JSONL replay batching, backoff, and structured import reports. - Export to SIEM via webhook or JSONL. - Basic installer and service wrapper. diff --git a/internal/collectors/collectors.go b/internal/collectors/collectors.go new file mode 100644 index 0000000..3148274 --- /dev/null +++ b/internal/collectors/collectors.go @@ -0,0 +1,375 @@ +package collectors + +import ( + "bufio" + "encoding/csv" + "encoding/json" + "errors" + "fmt" + "io" + "strconv" + "strings" + "time" + + "github.com/hunterinvariants/promtact/internal/domain" +) + +const ( + SourceAuditd = "auditd" + SourceSuricataEVE = "suricata-eve" + SourceSysmonJSON = "sysmon-json" + SourceZeekConn = "zeek-conn" +) + +func Normalize(source string, r io.Reader) ([]domain.Event, error) { + switch source { + case SourceAuditd: + return normalizeAuditd(r) + case SourceSuricataEVE: + return normalizeSuricataEVE(r) + case SourceSysmonJSON: + return normalizeSysmonJSON(r) + case SourceZeekConn: + return normalizeZeekConn(r) + default: + return nil, fmt.Errorf("unsupported collector source %q", source) + } +} + +func Sources() []string { + return []string{SourceAuditd, SourceSuricataEVE, SourceSysmonJSON, SourceZeekConn} +} + +func scanLines(r io.Reader, handle func(lineNumber int, line string) error) error { + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024) + lineNumber := 0 + for scanner.Scan() { + lineNumber++ + line := strings.TrimSpace(scanner.Text()) + line = strings.TrimPrefix(line, "\ufeff") + if line == "" { + continue + } + if err := handle(lineNumber, line); err != nil { + return err + } + } + return scanner.Err() +} + +func normalizeAuditd(r io.Reader) ([]domain.Event, error) { + events := []domain.Event{} + err := scanLines(r, func(lineNumber int, line string) error { + fields := keyValues(line) + eventType := strings.TrimPrefix(fields["type"], "type=") + if eventType == "" { + if strings.Contains(line, "type=SYSCALL") { + eventType = "SYSCALL" + } + if strings.Contains(line, "type=EXECVE") { + eventType = "EXECVE" + } + } + + if eventType != "SYSCALL" && eventType != "EXECVE" && eventType != "USER_AUTH" && eventType != "USER_LOGIN" { + return nil + } + + command := firstNonEmpty(fields["cmd"], fields["comm"], fields["exe"]) + event := domain.Event{ + Kind: domain.EventProcessStart, + AssetID: firstNonEmpty(fields["node"], fields["hostname"], "auditd-host"), + Hostname: firstNonEmpty(fields["node"], fields["hostname"]), + Process: trimQuotes(firstNonEmpty(fields["comm"], fields["exe"])), + Command: trimQuotes(command), + Signal: "auditd " + eventType, + Labels: []string{"auditd", strings.ToLower(eventType)}, + Metadata: map[string]string{ + "collector": SourceAuditd, + "line_number": strconv.Itoa(lineNumber), + "auid": trimQuotes(fields["auid"]), + "uid": trimQuotes(fields["uid"]), + }, + } + if eventType == "USER_AUTH" || eventType == "USER_LOGIN" { + event.Kind = domain.EventAuth + } + events = append(events, event) + return nil + }) + return events, err +} + +func normalizeSuricataEVE(r io.Reader) ([]domain.Event, error) { + events := []domain.Event{} + err := scanLines(r, func(lineNumber int, line string) error { + var eve map[string]any + if err := json.Unmarshal([]byte(line), &eve); err != nil { + return fmt.Errorf("line %d: %w", lineNumber, err) + } + eventType := stringValue(eve, "event_type") + if eventType == "" { + return nil + } + event := domain.Event{ + Timestamp: parseTime(stringValue(eve, "timestamp")), + Kind: domain.EventNetworkFlow, + AssetID: firstNonEmpty(stringValue(eve, "host"), stringValue(eve, "src_ip"), "suricata-sensor"), + Hostname: stringValue(eve, "host"), + SourceIP: stringValue(eve, "src_ip"), + Destination: joinHostPort(stringValue(eve, "dest_ip"), intStringValue(eve, "dest_port")), + Signal: "suricata " + eventType, + Labels: []string{"suricata", eventType}, + Metadata: map[string]string{ + "collector": SourceSuricataEVE, + "line_number": strconv.Itoa(lineNumber), + "proto": stringValue(eve, "proto"), + }, + } + if alert, ok := eve["alert"].(map[string]any); ok { + event.Kind = domain.EventFinding + event.Signal = firstNonEmpty(stringValue(alert, "signature"), event.Signal) + event.Metadata["signature_id"] = intStringValue(alert, "signature_id") + event.Metadata["category"] = stringValue(alert, "category") + event.Metadata["severity"] = intStringValue(alert, "severity") + } + events = append(events, event) + return nil + }) + return events, err +} + +func normalizeSysmonJSON(r io.Reader) ([]domain.Event, error) { + events := []domain.Event{} + err := scanLines(r, func(lineNumber int, line string) error { + var raw map[string]any + if err := json.Unmarshal([]byte(line), &raw); err != nil { + return fmt.Errorf("line %d: %w", lineNumber, err) + } + data := flattenEventData(raw) + eventID := firstNonEmpty(stringValue(raw, "EventID"), stringValue(raw, "event_id"), data["EventID"]) + host := firstNonEmpty(stringValue(raw, "Computer"), stringValue(raw, "Hostname"), data["Computer"], data["Hostname"]) + event := domain.Event{ + Timestamp: parseTime(firstNonEmpty(stringValue(raw, "UtcTime"), stringValue(raw, "TimeCreated"), data["UtcTime"], data["TimeCreated"])), + Kind: domain.EventHostObservation, + AssetID: firstNonEmpty(host, data["SourceIp"], data["DestinationIp"], "sysmon-host"), + Hostname: host, + SourceIP: data["SourceIp"], + Process: firstNonEmpty(data["Image"], data["ProcessName"]), + Command: data["CommandLine"], + Signal: "sysmon event " + eventID, + Labels: []string{"sysmon", "event-" + eventID}, + Metadata: map[string]string{ + "collector": SourceSysmonJSON, + "line_number": strconv.Itoa(lineNumber), + "event_id": eventID, + "rule_name": data["RuleName"], + }, + } + switch eventID { + case "1": + event.Kind = domain.EventProcessStart + case "3": + event.Kind = domain.EventNetworkFlow + event.Destination = joinHostPort(firstNonEmpty(data["DestinationHostname"], data["DestinationIp"]), data["DestinationPort"]) + } + events = append(events, event) + return nil + }) + return events, err +} + +func normalizeZeekConn(r io.Reader) ([]domain.Event, error) { + reader := bufio.NewReader(r) + peek, err := reader.Peek(1) + if err != nil && !errors.Is(err, io.EOF) { + return nil, err + } + if len(peek) == 0 { + return []domain.Event{}, nil + } + if peek[0] == '{' { + return normalizeZeekConnJSON(reader) + } + return normalizeZeekConnTSV(reader) +} + +func normalizeZeekConnJSON(r io.Reader) ([]domain.Event, error) { + events := []domain.Event{} + err := scanLines(r, func(lineNumber int, line string) error { + var raw map[string]any + if err := json.Unmarshal([]byte(line), &raw); err != nil { + return fmt.Errorf("line %d: %w", lineNumber, err) + } + events = append(events, zeekConnEvent(raw, lineNumber)) + return nil + }) + return events, err +} + +func normalizeZeekConnTSV(r io.Reader) ([]domain.Event, error) { + scanner := bufio.NewScanner(r) + fields := []string{} + events := []domain.Event{} + lineNumber := 0 + for scanner.Scan() { + lineNumber++ + line := scanner.Text() + if strings.HasPrefix(line, "#fields") { + fields = strings.Fields(line)[1:] + continue + } + if strings.HasPrefix(line, "#") || strings.TrimSpace(line) == "" { + continue + } + if len(fields) == 0 { + fields = []string{"ts", "uid", "id.orig_h", "id.orig_p", "id.resp_h", "id.resp_p", "proto", "service"} + } + reader := csv.NewReader(strings.NewReader(line)) + reader.Comma = '\t' + reader.FieldsPerRecord = -1 + parts, err := reader.Read() + if err != nil { + return nil, fmt.Errorf("line %d: %w", lineNumber, err) + } + raw := map[string]any{} + for i, field := range fields { + if i < len(parts) { + raw[field] = parts[i] + } + } + events = append(events, zeekConnEvent(raw, lineNumber)) + } + return events, scanner.Err() +} + +func zeekConnEvent(raw map[string]any, lineNumber int) domain.Event { + sourceIP := firstNonEmpty(stringValue(raw, "id.orig_h"), stringValue(raw, "src_ip")) + dest := firstNonEmpty(stringValue(raw, "id.resp_h"), stringValue(raw, "dest_ip")) + port := firstNonEmpty(stringValue(raw, "id.resp_p"), intStringValue(raw, "dest_port")) + return domain.Event{ + Timestamp: parseZeekTime(firstNonEmpty(stringValue(raw, "ts"), stringValue(raw, "timestamp"))), + Kind: domain.EventNetworkFlow, + AssetID: firstNonEmpty(sourceIP, "zeek-sensor"), + SourceIP: sourceIP, + Destination: joinHostPort(dest, port), + Signal: "zeek conn flow", + Labels: []string{"zeek", "conn"}, + Metadata: map[string]string{ + "collector": SourceZeekConn, + "line_number": strconv.Itoa(lineNumber), + "uid": stringValue(raw, "uid"), + "proto": stringValue(raw, "proto"), + "service": stringValue(raw, "service"), + }, + } +} + +func flattenEventData(raw map[string]any) map[string]string { + result := map[string]string{} + for key, value := range raw { + if scalar, ok := value.(string); ok { + result[key] = scalar + } + } + if nested, ok := raw["EventData"].(map[string]any); ok { + for key, value := range nested { + result[key] = fmt.Sprint(value) + } + } + if nested, ok := raw["event_data"].(map[string]any); ok { + for key, value := range nested { + result[key] = fmt.Sprint(value) + } + } + return result +} + +func keyValues(line string) map[string]string { + values := map[string]string{} + for _, part := range strings.Fields(line) { + key, value, ok := strings.Cut(part, "=") + if ok { + values[key] = trimQuotes(value) + } + } + return values +} + +func stringValue(raw map[string]any, key string) string { + value, ok := raw[key] + if !ok || value == nil { + return "" + } + switch typed := value.(type) { + case string: + return typed + case float64: + if typed == float64(int64(typed)) { + return strconv.FormatInt(int64(typed), 10) + } + return strconv.FormatFloat(typed, 'f', -1, 64) + default: + return fmt.Sprint(typed) + } +} + +func intStringValue(raw map[string]any, key string) string { + value := stringValue(raw, key) + if value == "" { + return "" + } + return value +} + +func parseTime(value string) time.Time { + value = strings.TrimSpace(value) + if value == "" { + return time.Time{} + } + formats := []string{time.RFC3339Nano, time.RFC3339, "2006-01-02 15:04:05.999", "2006-01-02 15:04:05"} + for _, format := range formats { + parsed, err := time.Parse(format, value) + if err == nil { + return parsed.UTC() + } + } + return time.Time{} +} + +func parseZeekTime(value string) time.Time { + if value == "" { + return time.Time{} + } + if seconds, err := strconv.ParseFloat(value, 64); err == nil { + whole, fraction := int64(seconds), seconds-float64(int64(seconds)) + return time.Unix(whole, int64(fraction*1e9)).UTC() + } + return parseTime(value) +} + +func joinHostPort(host string, port string) string { + host = strings.TrimSpace(host) + port = strings.TrimSpace(port) + if host == "" { + return "" + } + if port == "" || port == "-" { + return host + } + return host + ":" + port +} + +func firstNonEmpty(values ...string) string { + for _, value := range values { + if strings.TrimSpace(value) != "" && value != "-" { + return strings.TrimSpace(value) + } + } + return "" +} + +func trimQuotes(value string) string { + return strings.Trim(value, `"`) +} diff --git a/internal/collectors/collectors_test.go b/internal/collectors/collectors_test.go new file mode 100644 index 0000000..ece3531 --- /dev/null +++ b/internal/collectors/collectors_test.go @@ -0,0 +1,80 @@ +package collectors + +import ( + "strings" + "testing" + + "github.com/hunterinvariants/promtact/internal/domain" +) + +func TestNormalizeSuricataEVE(t *testing.T) { + input := strings.NewReader(`{"timestamp":"2026-06-06T12:00:00Z","event_type":"alert","src_ip":"10.0.0.5","dest_ip":"203.0.113.9","dest_port":443,"proto":"TCP","alert":{"signature":"Test canary egress","signature_id":1001,"category":"Policy","severity":2}}`) + + events, err := Normalize(SourceSuricataEVE, input) + if err != nil { + t.Fatalf("normalize: %v", err) + } + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].Kind != domain.EventFinding { + t.Fatalf("expected finding, got %s", events[0].Kind) + } + if events[0].Destination != "203.0.113.9:443" { + t.Fatalf("unexpected destination: %s", events[0].Destination) + } +} + +func TestNormalizeZeekConnTSV(t *testing.T) { + input := strings.NewReader("#fields\tts\tuid\tid.orig_h\tid.orig_p\tid.resp_h\tid.resp_p\tproto\tservice\n1717675200.0\tC1\t10.0.0.5\t51512\t203.0.113.9\t443\ttcp\tssl\n") + + events, err := Normalize(SourceZeekConn, input) + if err != nil { + t.Fatalf("normalize: %v", err) + } + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].Kind != domain.EventNetworkFlow { + t.Fatalf("expected network flow, got %s", events[0].Kind) + } + if events[0].SourceIP != "10.0.0.5" || events[0].Destination != "203.0.113.9:443" { + t.Fatalf("unexpected flow: %s -> %s", events[0].SourceIP, events[0].Destination) + } +} + +func TestNormalizeSysmonJSON(t *testing.T) { + input := strings.NewReader(`{"EventID":1,"Computer":"win-01","EventData":{"Image":"C:\\Windows\\System32\\WindowsPowerShell\\v1.0\\powershell.exe","CommandLine":"whoami; ipconfig /all"}}`) + + events, err := Normalize(SourceSysmonJSON, input) + if err != nil { + t.Fatalf("normalize: %v", err) + } + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].Kind != domain.EventProcessStart { + t.Fatalf("expected process start, got %s", events[0].Kind) + } + if events[0].Hostname != "win-01" { + t.Fatalf("unexpected host: %s", events[0].Hostname) + } +} + +func TestNormalizeAuditd(t *testing.T) { + input := strings.NewReader(`type=EXECVE msg=audit(1717675200.0:42): argc=2 a0="curl" a1="https://example.com" node=linux-01`) + + events, err := Normalize(SourceAuditd, input) + if err != nil { + t.Fatalf("normalize: %v", err) + } + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].Kind != domain.EventProcessStart { + t.Fatalf("expected process start, got %s", events[0].Kind) + } + if events[0].AssetID != "linux-01" { + t.Fatalf("unexpected asset: %s", events[0].AssetID) + } +} diff --git a/internal/domain/types.go b/internal/domain/types.go index 8add055..a83d9b6 100644 --- a/internal/domain/types.go +++ b/internal/domain/types.go @@ -79,14 +79,17 @@ type Alert struct { } type ResponseAction struct { - ID string `json:"id"` - Type string `json:"type"` - Mode string `json:"mode"` - AssetID string `json:"asset_id"` - Target string `json:"target"` - Reason string `json:"reason"` - CreatedAt time.Time `json:"created_at"` - Metadata map[string]string `json:"metadata"` + ID string `json:"id"` + Type string `json:"type"` + Mode string `json:"mode"` + AssetID string `json:"asset_id"` + Target string `json:"target"` + Reason string `json:"reason"` + CreatedAt time.Time `json:"created_at"` + ApprovalStatus string `json:"approval_status,omitempty"` + ApprovedBy string `json:"approved_by,omitempty"` + ApprovedAt *time.Time `json:"approved_at,omitempty"` + Metadata map[string]string `json:"metadata"` } type Asset struct { @@ -120,4 +123,5 @@ type Status struct { StorageMode string `json:"storage_mode"` StoragePath string `json:"storage_path,omitempty"` LastStorageError string `json:"last_storage_error,omitempty"` + LastExportError string `json:"last_export_error,omitempty"` } diff --git a/internal/exporter/webhook.go b/internal/exporter/webhook.go new file mode 100644 index 0000000..aa7e337 --- /dev/null +++ b/internal/exporter/webhook.go @@ -0,0 +1,66 @@ +package exporter + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/hunterinvariants/promtact/internal/domain" +) + +type Webhook struct { + URL string + Token string + Client *http.Client +} + +type AlertPayload struct { + Type string `json:"type"` + ExportedAt time.Time `json:"exported_at"` + Alerts []domain.Alert `json:"alerts"` +} + +func (w Webhook) ExportAlerts(alerts []domain.Alert) error { + if w.URL == "" || len(alerts) == 0 { + return nil + } + client := w.Client + if client == nil { + client = &http.Client{Timeout: 10 * time.Second} + } + + payload := AlertPayload{ + Type: "promtact.alerts", + ExportedAt: time.Now().UTC(), + Alerts: alerts, + } + body, err := json.Marshal(payload) + if err != nil { + return err + } + + req, err := http.NewRequest(http.MethodPost, w.URL, bytes.NewReader(body)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + if w.Token != "" { + req.Header.Set("Authorization", "Bearer "+w.Token) + } + + resp, err := client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return fmt.Errorf("webhook returned %s: %s", resp.Status, strings.TrimSpace(string(respBody))) + } + return nil +} diff --git a/internal/exporter/webhook_test.go b/internal/exporter/webhook_test.go new file mode 100644 index 0000000..4fc9cf2 --- /dev/null +++ b/internal/exporter/webhook_test.go @@ -0,0 +1,33 @@ +package exporter + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/hunterinvariants/promtact/internal/domain" +) + +func TestWebhookExportsAlerts(t *testing.T) { + var got AlertPayload + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bearer token" { + t.Fatalf("missing authorization header") + } + if err := json.NewDecoder(r.Body).Decode(&got); err != nil { + t.Fatalf("decode payload: %v", err) + } + w.WriteHeader(http.StatusAccepted) + })) + defer server.Close() + + webhook := Webhook{URL: server.URL, Token: "token", Client: server.Client()} + err := webhook.ExportAlerts([]domain.Alert{{ID: "alert-1", Title: "test"}}) + if err != nil { + t.Fatalf("export alerts: %v", err) + } + if got.Type != "promtact.alerts" || len(got.Alerts) != 1 { + t.Fatalf("unexpected payload: %#v", got) + } +} diff --git a/internal/response/planner.go b/internal/response/planner.go index 6c041f5..9183222 100644 --- a/internal/response/planner.go +++ b/internal/response/planner.go @@ -22,51 +22,56 @@ func (p *Planner) Plan(alert domain.Alert) []domain.ResponseAction { actions := []domain.ResponseAction{ { - Type: "create_incident_ticket", - Mode: mode, - AssetID: alert.AssetID, - Target: alert.ID, - Reason: "Create an audit trail and assign incident ownership.", + Type: "create_incident_ticket", + Mode: mode, + AssetID: alert.AssetID, + Target: alert.ID, + Reason: "Create an audit trail and assign incident ownership.", + ApprovalStatus: "not_required", }, } if alert.Severity.Rank() >= domain.SeverityHigh.Rank() { actions = append(actions, domain.ResponseAction{ - Type: "isolate_host", - Mode: mode, - AssetID: alert.AssetID, - Target: alert.AssetID, - Reason: "Contain high-severity activity before lateral movement expands.", + Type: "isolate_host", + Mode: mode, + AssetID: alert.AssetID, + Target: alert.AssetID, + Reason: "Contain high-severity activity before lateral movement expands.", + ApprovalStatus: "required", }) } if strings.Contains(alert.RuleID, "egress") || strings.Contains(alert.RuleID, "sequence") || strings.Contains(alert.RuleID, "model.runtime") { actions = append(actions, domain.ResponseAction{ - Type: "block_egress", - Mode: mode, - AssetID: alert.AssetID, - Target: firstNonEmpty(alert.Evidence["destination"], alert.Evidence["egress_event"], "unknown"), - Reason: "Stop unexpected external communication while preserving evidence.", + Type: "block_egress", + Mode: mode, + AssetID: alert.AssetID, + Target: firstNonEmpty(alert.Evidence["destination"], alert.Evidence["egress_event"], "unknown"), + Reason: "Stop unexpected external communication while preserving evidence.", + ApprovalStatus: "required", }) } if strings.Contains(alert.RuleID, "agent.tool") { actions = append(actions, domain.ResponseAction{ - Type: "disable_agent_tool", - Mode: mode, - AssetID: alert.AssetID, - Target: firstNonEmpty(alert.Evidence["tool"], "unknown"), - Reason: "Remove unapproved tool access from the agent runtime.", + Type: "disable_agent_tool", + Mode: mode, + AssetID: alert.AssetID, + Target: firstNonEmpty(alert.Evidence["tool"], "unknown"), + Reason: "Remove unapproved tool access from the agent runtime.", + ApprovalStatus: "required", }) } if strings.Contains(alert.RuleID, "secret") || strings.Contains(alert.RuleID, "canary") || strings.Contains(alert.RuleID, "deception") { actions = append(actions, domain.ResponseAction{ - Type: "rotate_related_secrets", - Mode: mode, - AssetID: alert.AssetID, - Target: alert.AssetID, - Reason: "Invalidate credentials that may have been exposed or touched.", + Type: "rotate_related_secrets", + Mode: mode, + AssetID: alert.AssetID, + Target: alert.AssetID, + Reason: "Invalidate credentials that may have been exposed or touched.", + ApprovalStatus: "required", }) } diff --git a/internal/server/server.go b/internal/server/server.go index cb5d605..3a01ebc 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -9,11 +9,13 @@ import ( "net/http" "path/filepath" "strings" + "sync" "sync/atomic" "time" "github.com/hunterinvariants/promtact/internal/correlator" "github.com/hunterinvariants/promtact/internal/domain" + "github.com/hunterinvariants/promtact/internal/exporter" "github.com/hunterinvariants/promtact/internal/policy" "github.com/hunterinvariants/promtact/internal/response" "github.com/hunterinvariants/promtact/internal/store" @@ -28,6 +30,9 @@ type App struct { responder *response.Planner webDir string apiToken string + webhook exporter.Webhook + exportMu sync.RWMutex + exportErr string startedAt time.Time counter atomic.Uint64 } @@ -38,6 +43,8 @@ type Options struct { APIToken string Policy policy.Config CorrelationWindow time.Duration + AlertWebhookURL string + AlertWebhookToken string } func New(webDir string) *App { @@ -66,7 +73,11 @@ func NewWithOptions(options Options) (*App, error) { responder: response.NewDryRun(), webDir: options.WebDir, apiToken: options.APIToken, - startedAt: time.Now().UTC(), + webhook: exporter.Webhook{ + URL: options.AlertWebhookURL, + Token: options.AlertWebhookToken, + }, + startedAt: time.Now().UTC(), }, nil } @@ -76,6 +87,7 @@ func (a *App) Routes() http.Handler { mux.HandleFunc("/api/events", a.handleEvents) mux.HandleFunc("/api/alerts", a.handleAlerts) mux.HandleFunc("/api/assets", a.handleAssets) + mux.HandleFunc("/api/responses/approve", a.handleResponseApproval) mux.HandleFunc("/api/responses", a.handleResponses) mux.HandleFunc("/api/policies", a.handlePolicies) mux.HandleFunc("/api/demo", a.handleDemo) @@ -113,6 +125,7 @@ func (a *App) handleStatus(w http.ResponseWriter, r *http.Request) { StorageMode: storageMode, StoragePath: a.store.PersistencePath(), LastStorageError: a.store.LastPersistenceError(), + LastExportError: a.lastExportError(), }) } @@ -165,6 +178,38 @@ func (a *App) handlePolicies(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusOK, a.policy.Rules()) } +func (a *App) handleResponseApproval(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + methodNotAllowed(w) + return + } + var req struct { + ActionID string `json:"action_id"` + ApprovedBy string `json:"approved_by"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, err) + return + } + if req.ActionID == "" { + writeError(w, http.StatusBadRequest, errors.New("action_id is required")) + return + } + if req.ApprovedBy == "" { + req.ApprovedBy = "operator" + } + action, ok, err := a.store.ApproveAction(req.ActionID, req.ApprovedBy, time.Now().UTC()) + if err != nil { + writeError(w, http.StatusInternalServerError, err) + return + } + if !ok { + writeError(w, http.StatusNotFound, errors.New("action not found")) + return + } + writeJSON(w, http.StatusAccepted, action) +} + func (a *App) handleResponses(w http.ResponseWriter, r *http.Request) { switch r.Method { case http.MethodGet: @@ -228,7 +273,12 @@ func (a *App) ingest(events []domain.Event) ([]domain.Alert, error) { alerts = append(alerts, a.correlator.Evaluate(a.store.ListEvents())...) a.prepareAlerts(alerts) - return a.store.AddAlerts(alerts) + added, err := a.store.AddAlerts(alerts) + if err != nil { + return nil, err + } + a.exportAlerts(added) + return added, nil } func (a *App) prepareEvent(event *domain.Event) { @@ -275,6 +325,29 @@ func (a *App) nextID(prefix string) string { return fmt.Sprintf("%s-%d", prefix, a.counter.Add(1)) } +func (a *App) exportAlerts(alerts []domain.Alert) { + if len(alerts) == 0 || a.webhook.URL == "" { + return + } + if err := a.webhook.ExportAlerts(alerts); err != nil { + a.setExportError(err.Error()) + return + } + a.setExportError("") +} + +func (a *App) setExportError(value string) { + a.exportMu.Lock() + defer a.exportMu.Unlock() + a.exportErr = value +} + +func (a *App) lastExportError() string { + a.exportMu.RLock() + defer a.exportMu.RUnlock() + return a.exportErr +} + func (a *App) staticHandler() http.Handler { root := http.Dir(filepath.Clean(a.webDir)) files := http.FileServer(root) diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 323ed32..f93e47e 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -6,6 +6,8 @@ import ( "net/http/httptest" "strings" "testing" + + "github.com/hunterinvariants/promtact/internal/domain" ) func TestWriteEndpointsRequireTokenWhenConfigured(t *testing.T) { @@ -76,3 +78,85 @@ func TestEmptyListEndpointsReturnArrays(t *testing.T) { } } } + +func TestResponseApprovalEndpoint(t *testing.T) { + app, err := NewWithOptions(Options{}) + if err != nil { + t.Fatalf("new app: %v", err) + } + if _, err := app.LoadDemo(); err != nil { + t.Fatalf("load demo: %v", err) + } + alerts := app.store.ListAlerts() + if len(alerts) == 0 { + t.Fatal("expected demo alerts") + } + + planReq := httptest.NewRequest(http.MethodPost, "/api/responses", strings.NewReader(`{"alert_id":"`+alerts[0].ID+`"}`)) + planRec := httptest.NewRecorder() + app.Routes().ServeHTTP(planRec, planReq) + if planRec.Code != http.StatusAccepted { + t.Fatalf("expected plan 202, got %d: %s", planRec.Code, planRec.Body.String()) + } + actions := app.store.ListActions() + var actionID string + for _, action := range actions { + if action.ApprovalStatus == "required" { + actionID = action.ID + break + } + } + if actionID == "" { + t.Fatal("expected at least one action requiring approval") + } + + approveReq := httptest.NewRequest(http.MethodPost, "/api/responses/approve", strings.NewReader(`{"action_id":"`+actionID+`","approved_by":"alice"}`)) + approveRec := httptest.NewRecorder() + app.Routes().ServeHTTP(approveRec, approveReq) + if approveRec.Code != http.StatusAccepted { + t.Fatalf("expected approve 202, got %d: %s", approveRec.Code, approveRec.Body.String()) + } + var approved struct { + ApprovalStatus string `json:"approval_status"` + ApprovedBy string `json:"approved_by"` + } + if err := json.Unmarshal(approveRec.Body.Bytes(), &approved); err != nil { + t.Fatalf("decode approval: %v", err) + } + if approved.ApprovalStatus != "approved" || approved.ApprovedBy != "alice" { + t.Fatalf("unexpected approval response: %#v", approved) + } +} + +func TestAlertWebhookExportsNewAlerts(t *testing.T) { + exported := 0 + webhook := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var payload struct { + Type string `json:"type"` + Alerts []domain.Alert `json:"alerts"` + } + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Fatalf("decode webhook: %v", err) + } + if payload.Type != "promtact.alerts" { + t.Fatalf("unexpected payload type: %s", payload.Type) + } + exported += len(payload.Alerts) + w.WriteHeader(http.StatusAccepted) + })) + defer webhook.Close() + + app, err := NewWithOptions(Options{AlertWebhookURL: webhook.URL}) + if err != nil { + t.Fatalf("new app: %v", err) + } + if _, err := app.LoadDemo(); err != nil { + t.Fatalf("load demo: %v", err) + } + if exported == 0 { + t.Fatal("expected webhook export") + } + if app.lastExportError() != "" { + t.Fatalf("unexpected export error: %s", app.lastExportError()) + } +} diff --git a/internal/store/store.go b/internal/store/store.go index 618c904..cc05e8d 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -101,6 +101,25 @@ func (s *Store) AddActions(actions []domain.ResponseAction) error { return s.persistLocked() } +func (s *Store) ApproveAction(id string, approvedBy string, approvedAt time.Time) (domain.ResponseAction, bool, error) { + s.mu.Lock() + defer s.mu.Unlock() + + for i := range s.actions { + if s.actions[i].ID != id { + continue + } + s.actions[i].ApprovalStatus = "approved" + s.actions[i].ApprovedBy = approvedBy + s.actions[i].ApprovedAt = &approvedAt + if err := s.persistLocked(); err != nil { + return domain.ResponseAction{}, true, err + } + return s.actions[i], true, nil + } + return domain.ResponseAction{}, false, nil +} + func (s *Store) ListActions() []domain.ResponseAction { s.mu.RLock() defer s.mu.RUnlock() diff --git a/internal/store/store_test.go b/internal/store/store_test.go index 7fb0101..ee7e777 100644 --- a/internal/store/store_test.go +++ b/internal/store/store_test.go @@ -99,3 +99,29 @@ func TestStoreSkipsDuplicateAlertFingerprintsAfterLoad(t *testing.T) { t.Fatalf("expected duplicate alert to be skipped, got %d", len(added)) } } + +func TestStoreApprovesAction(t *testing.T) { + s := New() + action := domain.ResponseAction{ + ID: "act-1", + Type: "isolate_host", + Mode: "dry-run", + AssetID: "asset-1", + ApprovalStatus: "required", + CreatedAt: time.Now().UTC(), + } + if err := s.AddActions([]domain.ResponseAction{action}); err != nil { + t.Fatalf("add action: %v", err) + } + + approved, ok, err := s.ApproveAction("act-1", "alice", time.Now().UTC()) + if err != nil { + t.Fatalf("approve action: %v", err) + } + if !ok { + t.Fatal("expected action to be found") + } + if approved.ApprovalStatus != "approved" || approved.ApprovedBy != "alice" || approved.ApprovedAt == nil { + t.Fatalf("unexpected approved action: %#v", approved) + } +} diff --git a/internal/telemetry/jsonl.go b/internal/telemetry/jsonl.go index 8fcc483..ff20510 100644 --- a/internal/telemetry/jsonl.go +++ b/internal/telemetry/jsonl.go @@ -19,6 +19,7 @@ func ReadJSONL(r io.Reader) ([]domain.Event, error) { for scanner.Scan() { lineNumber++ line := strings.TrimSpace(scanner.Text()) + line = strings.TrimPrefix(line, "\ufeff") if line == "" { continue } diff --git a/packaging/systemd/oadtd.service b/packaging/systemd/oadtd.service new file mode 100644 index 0000000..77f0e0c --- /dev/null +++ b/packaging/systemd/oadtd.service @@ -0,0 +1,23 @@ +[Unit] +Description=Promtact +After=network-online.target +Wants=network-online.target + +[Service] +Type=simple +User=promtact +Group=promtact +WorkingDirectory=/opt/promtact +ExecStart=/opt/promtact/promtact --addr :8080 --data /var/lib/promtact/state.json --policy /etc/promtact/policy.json +EnvironmentFile=-/etc/promtact/promtact.env +Restart=on-failure +RestartSec=5s +NoNewPrivileges=true +PrivateTmp=true +ProtectSystem=full +ProtectHome=true +ReadWritePaths=/var/lib/promtact + +[Install] +WantedBy=multi-user.target + diff --git a/packaging/windows/install-service.ps1 b/packaging/windows/install-service.ps1 new file mode 100644 index 0000000..adc6e8f --- /dev/null +++ b/packaging/windows/install-service.ps1 @@ -0,0 +1,37 @@ +param( + [string]$BinaryPath = "C:\Program Files\Promtact\promtact.exe", + [string]$WorkingDirectory = "C:\ProgramData\Promtact", + [string]$PolicyPath = "C:\ProgramData\Promtact\policy.json", + [string]$DataPath = "C:\ProgramData\Promtact\state.json", + [string]$ListenAddress = ":8080", + [string]$ServiceName = "Promtact" +) + +$ErrorActionPreference = "Stop" + +if (-not (Test-Path -LiteralPath $BinaryPath)) { + throw "Binary not found: $BinaryPath" +} + +New-Item -ItemType Directory -Force -Path $WorkingDirectory | Out-Null + +$arguments = "--addr $ListenAddress --data `"$DataPath`"" +if (Test-Path -LiteralPath $PolicyPath) { + $arguments = "$arguments --policy `"$PolicyPath`"" +} + +$binPath = "`"$BinaryPath`" $arguments" + +$existing = Get-Service -Name $ServiceName -ErrorAction SilentlyContinue +if ($existing) { + sc.exe stop $ServiceName | Out-Null + sc.exe delete $ServiceName | Out-Null + Start-Sleep -Seconds 2 +} + +sc.exe create $ServiceName binPath= $binPath start= auto DisplayName= "Promtact" | Out-Null +sc.exe description $ServiceName "Defensive control plane for agentic threat telemetry, policy, and response planning." | Out-Null +sc.exe start $ServiceName | Out-Null + +Write-Host "Installed and started service $ServiceName" + diff --git a/web/app.js b/web/app.js index 99dee1d..8655b24 100644 --- a/web/app.js +++ b/web/app.js @@ -162,10 +162,11 @@ function renderActions(actions) { node.innerHTML = `
${escapeHtml(action.type)}
- ${escapeHtml(action.mode)} + ${escapeHtml(action.approval_status || action.mode)}
${escapeHtml(action.asset_id || "unknown asset")} · ${escapeHtml(action.target || "-")}

${escapeHtml(action.reason || "")}

+ ${action.approval_status === "required" ? `
approval required
` : ""} `; els.actionsList.append(node); }); @@ -293,6 +294,21 @@ els.alertsList.addEventListener("click", async (event) => { } }); +els.actionsList.addEventListener("click", async (event) => { + const button = event.target.closest("[data-approve]"); + if (!button) return; + button.disabled = true; + try { + await api("/api/responses/approve", { + method: "POST", + body: JSON.stringify({ action_id: button.dataset.approve, approved_by: "dashboard" }) + }); + await refresh(); + } finally { + button.disabled = false; + } +}); + refresh().catch((error) => { console.error(error); }); From edf73e86c34321c240524da6f40047c8e84e5546 Mon Sep 17 00:00:00 2001 From: hunterinvariants Date: Sat, 6 Jun 2026 16:43:44 +0200 Subject: [PATCH 002/137] Add Postgres storage and RBAC --- .github/workflows/ci.yml | 3 +- .github/workflows/release.yml | 3 +- README.md | 54 ++++- cmd/promtact/main.go | 3 + cmd/promtactl/main.go | 23 ++ configs/example.policy.json | 1 - configs/example.rbac.policy.json | 26 +++ docs/architecture.md | 16 +- docs/operations.md | 36 ++- docs/roadmap.md | 11 +- go.mod | 11 +- go.sum | 26 +++ internal/auth/auth.go | 147 ++++++++++++ internal/auth/auth_test.go | 27 +++ internal/config/config.go | 10 +- internal/config/config_test.go | 6 +- internal/server/server.go | 51 +++-- internal/server/server_test.go | 42 ++++ internal/store/persistence.go | 24 +- internal/store/postgres.go | 307 ++++++++++++++++++++++++++ internal/store/store.go | 27 +++ packaging/systemd/oadtd.service | 3 +- packaging/windows/install-service.ps1 | 8 +- 23 files changed, 793 insertions(+), 72 deletions(-) create mode 100644 configs/example.rbac.policy.json create mode 100644 go.sum create mode 100644 internal/auth/auth.go create mode 100644 internal/auth/auth_test.go create mode 100644 internal/store/postgres.go diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index bcb8f7a..a143f6a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -12,6 +12,5 @@ jobs: - uses: actions/checkout@v4 - uses: actions/setup-go@v5 with: - go-version: "1.24.x" + go-version: "1.25.x" - run: go test ./... - diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 5647182..7df9602 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -19,7 +19,7 @@ jobs: - uses: actions/checkout@v4 - uses: actions/setup-go@v5 with: - go-version: "1.24.x" + go-version: "1.25.x" - run: go test ./... - name: Build shell: bash @@ -46,4 +46,3 @@ jobs: with: files: dist/* generate_release_notes: true - diff --git a/README.md b/README.md index e0254f6..5095e42 100644 --- a/README.md +++ b/README.md @@ -9,15 +9,15 @@ malware behavior, or autonomous propagation. Demo data generates telemetry only. ## What Exists Now -- Go HTTP service with in-memory storage or optional local JSON snapshot - persistence. +- Go HTTP service with Postgres persistence for production and JSON snapshot + fallback for local development. - Policy engine for agent-tool abuse, secret exposure, unexpected egress, discovery behavior, deception hits, and suspicious model runtime activity. - Correlator for multi-signal sequences such as discovery, credential touch, agent tool call, and outbound flow. - Dry-run response planner for host isolation, egress blocking, tool disabling, ticket creation, and secret rotation. -- Optional token protection for write endpoints. +- User/token authentication with role-based access control. - `promtactl replay` for safe JSONL telemetry replay into the ingest API. - Browser dashboard with asset risk graph, alerts, events, rules, and response actions. @@ -38,10 +38,16 @@ $env:GOMODCACHE="$PWD\.cache\go-mod" go run ./cmd/promtact --demo ``` -Run with local persistence and write-token protection: +Run with Postgres persistence: + +```powershell +$env:PROMTACT_POSTGRES_DSN="postgres://promtact:promtact@localhost:5432/promtact?sslmode=disable" +go run ./cmd/promtact --demo --policy configs\example.policy.json +``` + +Run with local JSON persistence for development: ```powershell -$env:PROMTACT_API_TOKEN="replace-with-a-local-secret" go run ./cmd/promtact --demo --data .cache\promtact-state.json ``` @@ -117,16 +123,18 @@ Useful endpoints: --addr HTTP listen address, default :8080 --web static dashboard directory, default web --demo load safe demo telemetry at startup ---data optional JSON snapshot path for local persistence +--data optional JSON snapshot path for local development persistence +--postgres-dsn Postgres DSN for production persistence, defaults to PROMTACT_POSTGRES_DSN --policy optional JSON policy configuration path ---api-token optional token for POST endpoints, defaults to PROMTACT_API_TOKEN +--api-token legacy admin token, defaults to PROMTACT_API_TOKEN --alert-webhook-url optional SIEM/webhook URL for new alerts --alert-webhook-token optional bearer token for alert webhook ``` -When `--api-token` or `PROMTACT_API_TOKEN` is set, read endpoints remain available -for the dashboard and health checks, while write endpoints require -`Authorization: Bearer ` or `X-Promtact-Token: `. +When users are configured in the policy file, all API endpoints require +`Authorization: Bearer ` or `X-Promtact-Token: ` and are checked +against RBAC roles. `--api-token` remains a legacy admin-token compatibility +path. ## Policy Configuration @@ -136,11 +144,33 @@ The policy file is JSON: { "approved_tools": ["asset_inventory", "ticket_create", "policy_read", "siem_search"], "approved_egress_hosts": ["api.openai.com", "github.com", "login.microsoftonline.com"], - "correlation_window": "30m" + "correlation_window": "30m", + "users": [ + { + "name": "admin", + "token_sha256": "replace-with-sha256-token-hash", + "roles": ["admin"] + } + ] } ``` -See [configs/example.policy.json](configs/example.policy.json). +See [configs/example.policy.json](configs/example.policy.json) and +[configs/example.rbac.policy.json](configs/example.rbac.policy.json). + +Create a token hash: + +```powershell +go run ./cmd/promtactl token-hash --token "replace-with-secret-token" +``` + +Roles: + +- `viewer`: read-only API access. +- `ingestor`: read API access and event/demo ingestion. +- `analyst`: read API access, ingestion, and response planning. +- `operator`: analyst permissions plus response approvals. +- `admin`: all API operations. ## Telemetry Replay diff --git a/cmd/promtact/main.go b/cmd/promtact/main.go index d352b35..33e001c 100644 --- a/cmd/promtact/main.go +++ b/cmd/promtact/main.go @@ -14,6 +14,7 @@ func main() { addr := flag.String("addr", ":8080", "HTTP listen address") webDir := flag.String("web", "web", "static dashboard directory") dataPath := flag.String("data", "", "optional JSON snapshot path for local persistence") + postgresDSN := flag.String("postgres-dsn", os.Getenv("PROMTACT_POSTGRES_DSN"), "Postgres DSN for production persistence") policyPath := flag.String("policy", "", "optional JSON policy configuration path") apiToken := flag.String("api-token", os.Getenv("PROMTACT_API_TOKEN"), "optional API token for write endpoints") alertWebhookURL := flag.String("alert-webhook-url", os.Getenv("PROMTACT_ALERT_WEBHOOK_URL"), "optional SIEM/webhook URL for new alerts") @@ -33,7 +34,9 @@ func main() { app, err := server.NewWithOptions(server.Options{ WebDir: *webDir, DataPath: *dataPath, + PostgresDSN: *postgresDSN, APIToken: *apiToken, + Users: runtimeConfig.Users, Policy: runtimeConfig.PolicyConfig(), CorrelationWindow: window, AlertWebhookURL: *alertWebhookURL, diff --git a/cmd/promtactl/main.go b/cmd/promtactl/main.go index 9bc22e6..9c2c8fd 100644 --- a/cmd/promtactl/main.go +++ b/cmd/promtactl/main.go @@ -13,6 +13,7 @@ import ( "strings" "time" + "github.com/hunterinvariants/promtact/internal/auth" "github.com/hunterinvariants/promtact/internal/collectors" "github.com/hunterinvariants/promtact/internal/domain" "github.com/hunterinvariants/promtact/internal/telemetry" @@ -34,12 +35,33 @@ func main() { if err := replay(os.Args[2:]); err != nil { log.Fatal(err) } + case "token-hash": + if err := tokenHash(os.Args[2:]); err != nil { + log.Fatal(err) + } default: usage() os.Exit(2) } } +func tokenHash(args []string) error { + fs := flag.NewFlagSet("token-hash", flag.ContinueOnError) + token := fs.String("token", "", "token to hash") + if err := fs.Parse(args); err != nil { + return err + } + value := *token + if value == "" { + value = os.Getenv("PROMTACT_TOKEN") + } + if value == "" { + return errors.New("token-hash requires --token or PROMTACT_TOKEN") + } + fmt.Println(auth.HashToken(value)) + return nil +} + func collect(args []string) error { fs := flag.NewFlagSet("collect", flag.ContinueOnError) source := fs.String("source", "", "collector source: "+strings.Join(collectors.Sources(), ", ")) @@ -198,4 +220,5 @@ func usage() { fmt.Fprintln(os.Stderr, "usage:") fmt.Fprintln(os.Stderr, " promtactl collect --source suricata-eve --file eve.json --output events.jsonl") fmt.Fprintln(os.Stderr, " promtactl replay --file events.jsonl [--url http://localhost:8080] [--token TOKEN]") + fmt.Fprintln(os.Stderr, " promtactl token-hash --token TOKEN") } diff --git a/configs/example.policy.json b/configs/example.policy.json index 8e10eef..edbdafd 100644 --- a/configs/example.policy.json +++ b/configs/example.policy.json @@ -12,4 +12,3 @@ ], "correlation_window": "30m" } - diff --git a/configs/example.rbac.policy.json b/configs/example.rbac.policy.json new file mode 100644 index 0000000..99f24e2 --- /dev/null +++ b/configs/example.rbac.policy.json @@ -0,0 +1,26 @@ +{ + "approved_tools": [ + "asset_inventory", + "ticket_create", + "policy_read", + "siem_search" + ], + "approved_egress_hosts": [ + "api.openai.com", + "github.com", + "login.microsoftonline.com" + ], + "correlation_window": "30m", + "users": [ + { + "name": "admin", + "token_sha256": "replace-with-sha256-token-hash", + "roles": ["admin"] + }, + { + "name": "collector", + "token_sha256": "replace-with-sha256-token-hash", + "roles": ["ingestor"] + } + ] +} diff --git a/docs/architecture.md b/docs/architecture.md index c7c3bbb..7a119b1 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -37,9 +37,9 @@ types. ### Store `internal/store` keeps events, alerts, response actions, and risk ranked assets. -By default it runs in memory. When `--data` is set, it writes a local JSON -snapshot and restores state on startup. This keeps the MVP dependency-free while -leaving room for SQLite or Postgres in the alpha phase. +For production, `--postgres-dsn` stores data in Postgres tables with JSONB +payloads and indexed core columns. For local development, `--data` writes a JSON +snapshot and restores state on startup. ### Policy Engine @@ -75,11 +75,11 @@ as requiring approval before any future execution backend can act on them. `web/` provides an operational dashboard for assets, alerts, events, policies, and dry-run response actions. -### Write Authentication +### Authentication And RBAC -Write endpoints can be protected with `--api-token` or `PROMTACT_API_TOKEN`. -Read endpoints remain available so the dashboard and health checks can load -without embedding a token in static assets. +API access supports user tokens with RBAC roles configured in the policy file. +Token values are not stored in config; only SHA-256 token hashes are stored. +The legacy `--api-token` path behaves as an admin token for compatibility. ### SIEM/Webhook Export @@ -105,7 +105,7 @@ flowchart LR ``` The current file-backed snapshot should be treated as local development storage, -not as a clustered production database. +not as the production database. Production deployments should use Postgres. ## Defensive Boundaries diff --git a/docs/operations.md b/docs/operations.md index 3914370..40152e2 100644 --- a/docs/operations.md +++ b/docs/operations.md @@ -18,6 +18,12 @@ Suggested layout: /var/lib/promtact/state.json ``` +For production, set Postgres in `/etc/promtact/promtact.env`: + +```text +PROMTACT_POSTGRES_DSN=postgres://promtact:promtact@postgres:5432/promtact?sslmode=disable +``` + Create a dedicated user, copy the binaries and policy file, install the unit, then enable it: @@ -56,10 +62,30 @@ The payload type is `promtact.alerts`. ## Storage -Current durable storage is the local JSON snapshot configured with `--data`. -This is suitable for local labs, pilots, and single-node testing. +Production durable storage is Postgres via `--postgres-dsn` or +`PROMTACT_POSTGRES_DSN`. Promtact creates the required tables automatically. + +The local JSON snapshot configured with `--data` remains useful for development +and quick labs, but it is not the production storage path. + +## RBAC -SQLite/Postgres is the next storage milestone. It should be implemented behind -the existing store boundary so the API and collectors do not change when the -storage backend changes. +Define users in the policy file with token hashes: +```json +{ + "users": [ + { + "name": "admin", + "token_sha256": "replace-with-sha256-token-hash", + "roles": ["admin"] + } + ] +} +``` + +Generate a hash: + +```powershell +.\promtactl.exe token-hash --token "replace-with-secret-token" +``` diff --git a/docs/roadmap.md b/docs/roadmap.md index 7f4c03c..a1b0e59 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -7,8 +7,9 @@ - Correlation engine. - Dry-run response planner. - Local dashboard. -- Optional JSON snapshot persistence. -- Optional token protection for write endpoints. +- Postgres persistence for production. +- Optional JSON snapshot persistence for local development. +- User-token authentication with RBAC roles. - JSON policy configuration for tool, egress, and correlation-window defaults. - Safe JSONL telemetry replay client. - Collector normalizers for Sysmon JSON, auditd, Zeek conn, and Suricata EVE. @@ -22,8 +23,8 @@ Status: implemented in this repository. ## 1-2 Weeks: Alpha -- Durable SQLite or Postgres storage. -- Authenticated API. +- Postgres migrations hardening and backup/restore docs. +- Session-based dashboard login on top of token/RBAC API. - Policy reload without restart. - Signed tool manifests for AI-agent and MCP surfaces. - Long-running collector agents for Sysmon, auditd, Zeek, Suricata, and proxy logs. @@ -35,7 +36,7 @@ Status: implemented in this repository. ## 3-6 Weeks: Beta - Multi-tenant control plane. -- RBAC. +- Organization-level RBAC policies. - Policy packs. - Deception token registry. - Response approvals. diff --git a/go.mod b/go.mod index 6389611..a37da31 100644 --- a/go.mod +++ b/go.mod @@ -1,4 +1,13 @@ module github.com/hunterinvariants/promtact -go 1.24 +go 1.25.0 +require github.com/jackc/pgx/v5 v5.10.0 + +require ( + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect + golang.org/x/sync v0.17.0 // indirect + golang.org/x/text v0.29.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..c0e505b --- /dev/null +++ b/go.sum @@ -0,0 +1,26 @@ +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug= +golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk= +golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/auth/auth.go b/internal/auth/auth.go new file mode 100644 index 0000000..f5b97cb --- /dev/null +++ b/internal/auth/auth.go @@ -0,0 +1,147 @@ +package auth + +import ( + "crypto/sha256" + "crypto/subtle" + "encoding/hex" + "net/http" + "strings" +) + +const ( + RoleViewer = "viewer" + RoleIngestor = "ingestor" + RoleAnalyst = "analyst" + RoleOperator = "operator" + RoleAdmin = "admin" +) + +type UserConfig struct { + Name string `json:"name"` + TokenHash string `json:"token_sha256"` + Roles []string `json:"roles"` +} + +type Principal struct { + Name string `json:"name"` + Roles []string `json:"roles"` +} + +type Authenticator struct { + users []UserConfig + legacyHash string +} + +func New(users []UserConfig, legacyToken string) *Authenticator { + authenticator := &Authenticator{users: normalizeUsers(users)} + if legacyToken != "" { + authenticator.legacyHash = HashToken(legacyToken) + } + return authenticator +} + +func HashToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} + +func (a *Authenticator) Enabled() bool { + return len(a.users) > 0 || a.legacyHash != "" +} + +func (a *Authenticator) HasUsers() bool { + return len(a.users) > 0 +} + +func (a *Authenticator) Authenticate(r *http.Request) (Principal, bool) { + token := readToken(r) + if token == "" { + return Principal{}, false + } + tokenHash := HashToken(token) + for _, user := range a.users { + if constantTimeEqual(tokenHash, user.TokenHash) { + return Principal{Name: user.Name, Roles: user.Roles}, true + } + } + if a.legacyHash != "" && constantTimeEqual(tokenHash, a.legacyHash) { + return Principal{Name: "legacy-token", Roles: []string{RoleAdmin}}, true + } + return Principal{}, false +} + +func (p Principal) HasAny(roles ...string) bool { + for _, have := range p.Roles { + for _, want := range roles { + if have == RoleAdmin || have == want { + return true + } + } + } + return false +} + +func RequiredRoles(method string, path string) []string { + if method == http.MethodGet || method == http.MethodHead || method == http.MethodOptions { + return []string{RoleViewer, RoleAnalyst, RoleOperator, RoleIngestor} + } + if path == "/api/events" || path == "/api/demo" { + return []string{RoleIngestor, RoleAnalyst, RoleOperator} + } + if strings.HasPrefix(path, "/api/responses/approve") { + return []string{RoleOperator} + } + if strings.HasPrefix(path, "/api/responses") { + return []string{RoleAnalyst, RoleOperator} + } + return []string{RoleAdmin} +} + +func normalizeUsers(users []UserConfig) []UserConfig { + normalized := make([]UserConfig, 0, len(users)) + for _, user := range users { + user.Name = strings.TrimSpace(user.Name) + user.TokenHash = strings.ToLower(strings.TrimSpace(user.TokenHash)) + if user.Name == "" || user.TokenHash == "" { + continue + } + user.Roles = normalizeRoles(user.Roles) + if len(user.Roles) == 0 { + user.Roles = []string{RoleViewer} + } + normalized = append(normalized, user) + } + return normalized +} + +func normalizeRoles(roles []string) []string { + seen := map[string]struct{}{} + normalized := []string{} + for _, role := range roles { + role = strings.ToLower(strings.TrimSpace(role)) + if role == "" { + continue + } + if _, ok := seen[role]; ok { + continue + } + seen[role] = struct{}{} + normalized = append(normalized, role) + } + return normalized +} + +func readToken(r *http.Request) string { + header := r.Header.Get("Authorization") + if strings.HasPrefix(strings.ToLower(header), "bearer ") { + return strings.TrimSpace(header[len("Bearer "):]) + } + return strings.TrimSpace(r.Header.Get("X-Promtact-Token")) +} + +func constantTimeEqual(got string, want string) bool { + if got == "" || want == "" { + return false + } + return subtle.ConstantTimeCompare([]byte(got), []byte(want)) == 1 +} diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go new file mode 100644 index 0000000..989b83f --- /dev/null +++ b/internal/auth/auth_test.go @@ -0,0 +1,27 @@ +package auth + +import ( + "net/http" + "net/http/httptest" + "testing" +) + +func TestAuthenticateUserToken(t *testing.T) { + a := New([]UserConfig{{Name: "alice", TokenHash: HashToken("secret"), Roles: []string{RoleOperator}}}, "") + req := httptest.NewRequest(http.MethodGet, "/api/status", nil) + req.Header.Set("Authorization", "Bearer secret") + + principal, ok := a.Authenticate(req) + if !ok { + t.Fatal("expected authentication") + } + if principal.Name != "alice" || !principal.HasAny(RoleOperator) { + t.Fatalf("unexpected principal: %#v", principal) + } +} + +func TestRequiredRoles(t *testing.T) { + if roles := RequiredRoles(http.MethodPost, "/api/responses/approve"); len(roles) != 1 || roles[0] != RoleOperator { + t.Fatalf("unexpected approve roles: %#v", roles) + } +} diff --git a/internal/config/config.go b/internal/config/config.go index d60b2d2..9867c7e 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1,19 +1,22 @@ package config import ( + "bytes" "encoding/json" "os" "time" + "github.com/hunterinvariants/promtact/internal/auth" "github.com/hunterinvariants/promtact/internal/policy" ) const DefaultCorrelationWindow = 30 * time.Minute type Config struct { - ApprovedTools []string `json:"approved_tools"` - ApprovedEgressHosts []string `json:"approved_egress_hosts"` - CorrelationWindow string `json:"correlation_window"` + ApprovedTools []string `json:"approved_tools"` + ApprovedEgressHosts []string `json:"approved_egress_hosts"` + CorrelationWindow string `json:"correlation_window"` + Users []auth.UserConfig `json:"users"` } func Load(path string) (Config, error) { @@ -25,6 +28,7 @@ func Load(path string) (Config, error) { if err != nil { return Config{}, err } + data = bytes.TrimPrefix(data, []byte("\xef\xbb\xbf")) var config Config if err := json.Unmarshal(data, &config); err != nil { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 05774b2..47d828a 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -12,7 +12,8 @@ func TestLoadConfig(t *testing.T) { body := `{ "approved_tools": ["asset_inventory", "shell_exec"], "approved_egress_hosts": ["example.com"], - "correlation_window": "45m" + "correlation_window": "45m", + "users": [{"name":"alice","token_sha256":"hash","roles":["operator"]}] }` if err := os.WriteFile(path, []byte(body), 0o600); err != nil { t.Fatalf("write config: %v", err) @@ -32,6 +33,9 @@ func TestLoadConfig(t *testing.T) { if window != 45*time.Minute { t.Fatalf("unexpected window: %s", window) } + if len(loaded.Users) != 1 || loaded.Users[0].Name != "alice" { + t.Fatalf("unexpected users: %#v", loaded.Users) + } } func TestDefaultCorrelationWindow(t *testing.T) { diff --git a/internal/server/server.go b/internal/server/server.go index 3a01ebc..d244edc 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -1,7 +1,6 @@ package server import ( - "crypto/subtle" "encoding/json" "errors" "fmt" @@ -13,6 +12,7 @@ import ( "sync/atomic" "time" + "github.com/hunterinvariants/promtact/internal/auth" "github.com/hunterinvariants/promtact/internal/correlator" "github.com/hunterinvariants/promtact/internal/domain" "github.com/hunterinvariants/promtact/internal/exporter" @@ -29,7 +29,7 @@ type App struct { correlator *correlator.Correlator responder *response.Planner webDir string - apiToken string + auth *auth.Authenticator webhook exporter.Webhook exportMu sync.RWMutex exportErr string @@ -40,7 +40,9 @@ type App struct { type Options struct { WebDir string DataPath string + PostgresDSN string APIToken string + Users []auth.UserConfig Policy policy.Config CorrelationWindow time.Duration AlertWebhookURL string @@ -59,7 +61,13 @@ func NewWithOptions(options Options) (*App, error) { if options.WebDir == "" { options.WebDir = "web" } - st, err := store.NewWithPath(options.DataPath) + var st *store.Store + var err error + if options.PostgresDSN != "" { + st, err = store.NewWithPostgres(options.PostgresDSN) + } else { + st, err = store.NewWithPath(options.DataPath) + } if err != nil { return nil, err } @@ -72,7 +80,7 @@ func NewWithOptions(options Options) (*App, error) { correlator: correlator.New(options.CorrelationWindow), responder: response.NewDryRun(), webDir: options.WebDir, - apiToken: options.APIToken, + auth: auth.New(options.Users, options.APIToken), webhook: exporter.Webhook{ URL: options.AlertWebhookURL, Token: options.AlertWebhookToken, @@ -110,10 +118,6 @@ func (a *App) handleStatus(w http.ResponseWriter, r *http.Request) { } events, alerts, assets, actions := a.store.Counts() - storageMode := "memory" - if a.store.PersistencePath() != "" { - storageMode = "file" - } writeJSON(w, http.StatusOK, domain.Status{ Version: Version, UptimeSeconds: int64(time.Since(a.startedAt).Seconds()), @@ -122,7 +126,7 @@ func (a *App) handleStatus(w http.ResponseWriter, r *http.Request) { AssetCount: assets, ActionCount: actions, StartedAt: a.startedAt, - StorageMode: storageMode, + StorageMode: a.store.PersistenceMode(), StoragePath: a.store.PersistencePath(), LastStorageError: a.store.LastPersistenceError(), LastExportError: a.lastExportError(), @@ -415,14 +419,24 @@ func withSecurityHeaders(next http.Handler) http.Handler { func (a *App) withAuth(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if a.apiToken == "" || !strings.HasPrefix(r.URL.Path, "/api/") || isReadOnly(r.Method) { + if !strings.HasPrefix(r.URL.Path, "/api/") || a.auth == nil || !a.auth.Enabled() { next.ServeHTTP(w, r) return } - if !constantTimeEqual(readToken(r), a.apiToken) { + if !a.auth.HasUsers() && isReadOnly(r.Method) { + next.ServeHTTP(w, r) + return + } + principal, ok := a.auth.Authenticate(r) + if !ok { writeError(w, http.StatusUnauthorized, errors.New("missing or invalid API token")) return } + required := auth.RequiredRoles(r.Method, r.URL.Path) + if !principal.HasAny(required...) { + writeError(w, http.StatusForbidden, errors.New("insufficient role")) + return + } next.ServeHTTP(w, r) }) } @@ -430,18 +444,3 @@ func (a *App) withAuth(next http.Handler) http.Handler { func isReadOnly(method string) bool { return method == http.MethodGet || method == http.MethodHead || method == http.MethodOptions } - -func readToken(r *http.Request) string { - header := r.Header.Get("Authorization") - if strings.HasPrefix(strings.ToLower(header), "bearer ") { - return strings.TrimSpace(header[len("Bearer "):]) - } - return strings.TrimSpace(r.Header.Get("X-Promtact-Token")) -} - -func constantTimeEqual(got string, want string) bool { - if got == "" || want == "" { - return false - } - return subtle.ConstantTimeCompare([]byte(got), []byte(want)) == 1 -} diff --git a/internal/server/server_test.go b/internal/server/server_test.go index f93e47e..5b364c8 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -7,6 +7,7 @@ import ( "strings" "testing" + "github.com/hunterinvariants/promtact/internal/auth" "github.com/hunterinvariants/promtact/internal/domain" ) @@ -56,6 +57,47 @@ func TestReadEndpointsDoNotRequireToken(t *testing.T) { } } +func TestRBACRequiresTokenForReadWhenUsersConfigured(t *testing.T) { + app, err := NewWithOptions(Options{ + Users: []auth.UserConfig{{Name: "viewer", TokenHash: auth.HashToken("view-token"), Roles: []string{auth.RoleViewer}}}, + }) + if err != nil { + t.Fatalf("new app: %v", err) + } + + req := httptest.NewRequest(http.MethodGet, "/api/status", nil) + rec := httptest.NewRecorder() + app.Routes().ServeHTTP(rec, req) + if rec.Code != http.StatusUnauthorized { + t.Fatalf("expected 401, got %d", rec.Code) + } + + req = httptest.NewRequest(http.MethodGet, "/api/status", nil) + req.Header.Set("Authorization", "Bearer view-token") + rec = httptest.NewRecorder() + app.Routes().ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", rec.Code) + } +} + +func TestRBACBlocksInsufficientRole(t *testing.T) { + app, err := NewWithOptions(Options{ + Users: []auth.UserConfig{{Name: "viewer", TokenHash: auth.HashToken("view-token"), Roles: []string{auth.RoleViewer}}}, + }) + if err != nil { + t.Fatalf("new app: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/api/events", strings.NewReader(`{"kind":"finding","asset_id":"a1"}`)) + req.Header.Set("Authorization", "Bearer view-token") + rec := httptest.NewRecorder() + app.Routes().ServeHTTP(rec, req) + if rec.Code != http.StatusForbidden { + t.Fatalf("expected 403, got %d", rec.Code) + } +} + func TestEmptyListEndpointsReturnArrays(t *testing.T) { app, err := NewWithOptions(Options{}) if err != nil { diff --git a/internal/store/persistence.go b/internal/store/persistence.go index 4aeabae..07833b4 100644 --- a/internal/store/persistence.go +++ b/internal/store/persistence.go @@ -27,6 +27,7 @@ func NewWithPath(path string) (*Store, error) { if path == "" { return s, nil } + s.mode = "file" data, err := os.ReadFile(path) if errors.Is(err, os.ErrNotExist) { @@ -57,7 +58,7 @@ func NewWithPath(path string) (*Store, error) { } func (s *Store) persistLocked() error { - if s.path == "" { + if s.path == "" || s.mode == "postgres" { s.lastErr = "" return nil } @@ -96,6 +97,27 @@ func (s *Store) persistLocked() error { return nil } +func (s *Store) persistEventLocked(event domain.Event) error { + if s.db != nil { + return s.postgresPersistEventLocked(event) + } + return nil +} + +func (s *Store) persistAlertsLocked(alerts []domain.Alert) error { + if s.db != nil { + return s.postgresPersistAlertsLocked(alerts) + } + return nil +} + +func (s *Store) persistActionsLocked(actions []domain.ResponseAction) error { + if s.db != nil { + return s.postgresPersistActionsLocked(actions) + } + return nil +} + func replaceFile(src string, dst string) error { if err := os.Rename(src, dst); err == nil { return nil diff --git a/internal/store/postgres.go b/internal/store/postgres.go new file mode 100644 index 0000000..b2432fa --- /dev/null +++ b/internal/store/postgres.go @@ -0,0 +1,307 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "time" + + _ "github.com/jackc/pgx/v5/stdlib" + "github.com/hunterinvariants/promtact/internal/domain" +) + +const postgresTimeout = 10 * time.Second + +func NewWithPostgres(dsn string) (*Store, error) { + db, err := sql.Open("pgx", dsn) + if err != nil { + return nil, err + } + ctx, cancel := context.WithTimeout(context.Background(), postgresTimeout) + defer cancel() + if err := db.PingContext(ctx); err != nil { + _ = db.Close() + return nil, err + } + + s := New() + s.db = db + s.mode = "postgres" + s.path = redactDSN(dsn) + if err := s.postgresMigrate(ctx); err != nil { + _ = db.Close() + return nil, err + } + if err := s.postgresLoad(ctx); err != nil { + _ = db.Close() + return nil, err + } + return s, nil +} + +func (s *Store) Close() error { + s.mu.Lock() + defer s.mu.Unlock() + if s.db == nil { + return nil + } + err := s.db.Close() + s.db = nil + return err +} + +func (s *Store) postgresMigrate(ctx context.Context) error { + _, err := s.db.ExecContext(ctx, ` +CREATE TABLE IF NOT EXISTS promtact_events ( + id TEXT PRIMARY KEY, + occurred_at TIMESTAMPTZ, + asset_id TEXT, + kind TEXT, + data JSONB NOT NULL +); +CREATE INDEX IF NOT EXISTS idx_promtact_events_occurred_at ON promtact_events (occurred_at DESC); +CREATE INDEX IF NOT EXISTS idx_promtact_events_asset_id ON promtact_events (asset_id); + +CREATE TABLE IF NOT EXISTS promtact_alerts ( + id TEXT PRIMARY KEY, + fingerprint TEXT UNIQUE, + created_at TIMESTAMPTZ, + asset_id TEXT, + severity TEXT, + status TEXT, + data JSONB NOT NULL +); +CREATE INDEX IF NOT EXISTS idx_promtact_alerts_created_at ON promtact_alerts (created_at DESC); +CREATE INDEX IF NOT EXISTS idx_promtact_alerts_asset_id ON promtact_alerts (asset_id); +CREATE INDEX IF NOT EXISTS idx_promtact_alerts_status ON promtact_alerts (status); + +CREATE TABLE IF NOT EXISTS promtact_actions ( + id TEXT PRIMARY KEY, + created_at TIMESTAMPTZ, + asset_id TEXT, + approval_status TEXT, + data JSONB NOT NULL +); +CREATE INDEX IF NOT EXISTS idx_promtact_actions_created_at ON promtact_actions (created_at DESC); +CREATE INDEX IF NOT EXISTS idx_promtact_actions_asset_id ON promtact_actions (asset_id); + +CREATE TABLE IF NOT EXISTS promtact_assets ( + id TEXT PRIMARY KEY, + last_seen TIMESTAMPTZ, + risk_score INTEGER, + data JSONB NOT NULL +);`) + return err +} + +func (s *Store) postgresLoad(ctx context.Context) error { + if err := s.postgresLoadEvents(ctx); err != nil { + return err + } + if err := s.postgresLoadAlerts(ctx); err != nil { + return err + } + if err := s.postgresLoadActions(ctx); err != nil { + return err + } + if err := s.postgresLoadAssets(ctx); err != nil { + return err + } + s.rebuildFingerprintsLocked() + if len(s.assets) == 0 { + s.rebuildAssetsLocked() + } + return nil +} + +func (s *Store) postgresLoadEvents(ctx context.Context) error { + rows, err := s.db.QueryContext(ctx, `SELECT data FROM promtact_events ORDER BY occurred_at ASC NULLS LAST, id ASC`) + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var data []byte + if err := rows.Scan(&data); err != nil { + return err + } + var event domain.Event + if err := json.Unmarshal(data, &event); err != nil { + return err + } + s.events = append(s.events, event) + } + return rows.Err() +} + +func (s *Store) postgresLoadAlerts(ctx context.Context) error { + rows, err := s.db.QueryContext(ctx, `SELECT data FROM promtact_alerts ORDER BY created_at ASC NULLS LAST, id ASC`) + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var data []byte + if err := rows.Scan(&data); err != nil { + return err + } + var alert domain.Alert + if err := json.Unmarshal(data, &alert); err != nil { + return err + } + s.alerts = append(s.alerts, alert) + } + return rows.Err() +} + +func (s *Store) postgresLoadActions(ctx context.Context) error { + rows, err := s.db.QueryContext(ctx, `SELECT data FROM promtact_actions ORDER BY created_at ASC NULLS LAST, id ASC`) + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var data []byte + if err := rows.Scan(&data); err != nil { + return err + } + var action domain.ResponseAction + if err := json.Unmarshal(data, &action); err != nil { + return err + } + s.actions = append(s.actions, action) + } + return rows.Err() +} + +func (s *Store) postgresLoadAssets(ctx context.Context) error { + rows, err := s.db.QueryContext(ctx, `SELECT data FROM promtact_assets`) + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var data []byte + if err := rows.Scan(&data); err != nil { + return err + } + var asset domain.Asset + if err := json.Unmarshal(data, &asset); err != nil { + return err + } + if asset.ID != "" { + s.assets[asset.ID] = asset + } + } + return rows.Err() +} + +func (s *Store) postgresPersistEventLocked(event domain.Event) error { + ctx, cancel := context.WithTimeout(context.Background(), postgresTimeout) + defer cancel() + data, err := json.Marshal(event) + if err != nil { + return err + } + if _, err := s.db.ExecContext(ctx, ` +INSERT INTO promtact_events (id, occurred_at, asset_id, kind, data) +VALUES ($1, $2, $3, $4, $5) +ON CONFLICT (id) DO UPDATE SET + occurred_at = EXCLUDED.occurred_at, + asset_id = EXCLUDED.asset_id, + kind = EXCLUDED.kind, + data = EXCLUDED.data`, + event.ID, nullableTime(event.Timestamp), event.AssetID, string(event.Kind), data); err != nil { + return err + } + return s.postgresPersistAssetsLocked(ctx) +} + +func (s *Store) postgresPersistAlertsLocked(alerts []domain.Alert) error { + ctx, cancel := context.WithTimeout(context.Background(), postgresTimeout) + defer cancel() + for _, alert := range alerts { + data, err := json.Marshal(alert) + if err != nil { + return err + } + if _, err := s.db.ExecContext(ctx, ` +INSERT INTO promtact_alerts (id, fingerprint, created_at, asset_id, severity, status, data) +VALUES ($1, $2, $3, $4, $5, $6, $7) +ON CONFLICT (id) DO UPDATE SET + fingerprint = EXCLUDED.fingerprint, + created_at = EXCLUDED.created_at, + asset_id = EXCLUDED.asset_id, + severity = EXCLUDED.severity, + status = EXCLUDED.status, + data = EXCLUDED.data`, + alert.ID, nullEmpty(alert.Fingerprint), nullableTime(alert.CreatedAt), alert.AssetID, string(alert.Severity), string(alert.Status), data); err != nil { + return err + } + } + return s.postgresPersistAssetsLocked(ctx) +} + +func (s *Store) postgresPersistActionsLocked(actions []domain.ResponseAction) error { + ctx, cancel := context.WithTimeout(context.Background(), postgresTimeout) + defer cancel() + for _, action := range actions { + data, err := json.Marshal(action) + if err != nil { + return err + } + if _, err := s.db.ExecContext(ctx, ` +INSERT INTO promtact_actions (id, created_at, asset_id, approval_status, data) +VALUES ($1, $2, $3, $4, $5) +ON CONFLICT (id) DO UPDATE SET + created_at = EXCLUDED.created_at, + asset_id = EXCLUDED.asset_id, + approval_status = EXCLUDED.approval_status, + data = EXCLUDED.data`, + action.ID, nullableTime(action.CreatedAt), action.AssetID, action.ApprovalStatus, data); err != nil { + return err + } + } + return nil +} + +func (s *Store) postgresPersistAssetsLocked(ctx context.Context) error { + for _, asset := range s.assets { + data, err := json.Marshal(asset) + if err != nil { + return err + } + if _, err := s.db.ExecContext(ctx, ` +INSERT INTO promtact_assets (id, last_seen, risk_score, data) +VALUES ($1, $2, $3, $4) +ON CONFLICT (id) DO UPDATE SET + last_seen = EXCLUDED.last_seen, + risk_score = EXCLUDED.risk_score, + data = EXCLUDED.data`, + asset.ID, nullableTime(asset.LastSeen), asset.RiskScore, data); err != nil { + return err + } + } + return nil +} + +func nullableTime(value time.Time) any { + if value.IsZero() { + return nil + } + return value +} + +func nullEmpty(value string) any { + if value == "" { + return nil + } + return value +} + +func redactDSN(dsn string) string { + if dsn == "" { + return "" + } + return "postgres" +} diff --git a/internal/store/store.go b/internal/store/store.go index cc05e8d..1004293 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -1,6 +1,7 @@ package store import ( + "database/sql" "sort" "sync" "time" @@ -10,6 +11,8 @@ import ( type Store struct { mu sync.RWMutex + db *sql.DB + mode string events []domain.Event alerts []domain.Alert actions []domain.ResponseAction @@ -21,6 +24,7 @@ type Store struct { func New() *Store { return &Store{ + mode: "memory", assets: make(map[string]domain.Asset), fingerprints: make(map[string]struct{}), } @@ -32,6 +36,10 @@ func (s *Store) AddEvent(event domain.Event) error { s.events = append(s.events, event) s.upsertAssetLocked(event) + if err := s.persistEventLocked(event); err != nil { + s.lastErr = err.Error() + return err + } return s.persistLocked() } @@ -66,6 +74,10 @@ func (s *Store) AddAlerts(alerts []domain.Alert) ([]domain.Alert, error) { if len(added) == 0 { return added, nil } + if err := s.persistAlertsLocked(added); err != nil { + s.lastErr = err.Error() + return nil, err + } return added, s.persistLocked() } @@ -98,6 +110,10 @@ func (s *Store) AddActions(actions []domain.ResponseAction) error { defer s.mu.Unlock() s.actions = append(s.actions, actions...) + if err := s.persistActionsLocked(actions); err != nil { + s.lastErr = err.Error() + return err + } return s.persistLocked() } @@ -112,6 +128,10 @@ func (s *Store) ApproveAction(id string, approvedBy string, approvedAt time.Time s.actions[i].ApprovalStatus = "approved" s.actions[i].ApprovedBy = approvedBy s.actions[i].ApprovedAt = &approvedAt + if err := s.persistActionsLocked([]domain.ResponseAction{s.actions[i]}); err != nil { + s.lastErr = err.Error() + return domain.ResponseAction{}, true, err + } if err := s.persistLocked(); err != nil { return domain.ResponseAction{}, true, err } @@ -163,6 +183,13 @@ func (s *Store) PersistencePath() string { return s.path } +func (s *Store) PersistenceMode() string { + s.mu.RLock() + defer s.mu.RUnlock() + + return s.mode +} + func (s *Store) LastPersistenceError() string { s.mu.RLock() defer s.mu.RUnlock() diff --git a/packaging/systemd/oadtd.service b/packaging/systemd/oadtd.service index 77f0e0c..d3d4395 100644 --- a/packaging/systemd/oadtd.service +++ b/packaging/systemd/oadtd.service @@ -8,7 +8,7 @@ Type=simple User=promtact Group=promtact WorkingDirectory=/opt/promtact -ExecStart=/opt/promtact/promtact --addr :8080 --data /var/lib/promtact/state.json --policy /etc/promtact/policy.json +ExecStart=/opt/promtact/promtact --addr :8080 --policy /etc/promtact/policy.json EnvironmentFile=-/etc/promtact/promtact.env Restart=on-failure RestartSec=5s @@ -20,4 +20,3 @@ ReadWritePaths=/var/lib/promtact [Install] WantedBy=multi-user.target - diff --git a/packaging/windows/install-service.ps1 b/packaging/windows/install-service.ps1 index adc6e8f..9e07641 100644 --- a/packaging/windows/install-service.ps1 +++ b/packaging/windows/install-service.ps1 @@ -2,7 +2,7 @@ param( [string]$BinaryPath = "C:\Program Files\Promtact\promtact.exe", [string]$WorkingDirectory = "C:\ProgramData\Promtact", [string]$PolicyPath = "C:\ProgramData\Promtact\policy.json", - [string]$DataPath = "C:\ProgramData\Promtact\state.json", + [string]$PostgresDsn = "", [string]$ListenAddress = ":8080", [string]$ServiceName = "Promtact" ) @@ -15,7 +15,10 @@ if (-not (Test-Path -LiteralPath $BinaryPath)) { New-Item -ItemType Directory -Force -Path $WorkingDirectory | Out-Null -$arguments = "--addr $ListenAddress --data `"$DataPath`"" +$arguments = "--addr $ListenAddress" +if ($PostgresDsn -ne "") { + $arguments = "$arguments --postgres-dsn `"$PostgresDsn`"" +} if (Test-Path -LiteralPath $PolicyPath) { $arguments = "$arguments --policy `"$PolicyPath`"" } @@ -34,4 +37,3 @@ sc.exe description $ServiceName "Defensive control plane for agentic threat tele sc.exe start $ServiceName | Out-Null Write-Host "Installed and started service $ServiceName" - From aaa5b23fd56432f6d5a936c04bbc190b4112bc4e Mon Sep 17 00:00:00 2001 From: hunterinvariants Date: Sat, 6 Jun 2026 16:55:25 +0200 Subject: [PATCH 003/137] Add audit logging and Postgres integration setup --- README.md | 14 ++ compose.yaml | 19 +++ docs/architecture.md | 21 ++- docs/operations.md | 24 ++++ docs/roadmap.md | 9 +- internal/auth/auth.go | 3 + internal/auth/auth_test.go | 3 + internal/domain/types.go | 15 +++ internal/server/server.go | 84 +++++++++++- internal/server/server_test.go | 35 +++++ internal/store/persistence.go | 10 ++ internal/store/postgres.go | 61 ++++++++- internal/store/postgres_integration_test.go | 142 ++++++++++++++++++++ internal/store/store.go | 29 +++- internal/store/store_test.go | 37 ++++- web/app.js | 10 +- web/index.html | 5 +- web/styles.css | 3 +- 18 files changed, 501 insertions(+), 23 deletions(-) create mode 100644 compose.yaml create mode 100644 internal/store/postgres_integration_test.go diff --git a/README.md b/README.md index 5095e42..1642715 100644 --- a/README.md +++ b/README.md @@ -18,6 +18,8 @@ malware behavior, or autonomous propagation. Demo data generates telemetry only. - Dry-run response planner for host isolation, egress blocking, tool disabling, ticket creation, and secret rotation. - User/token authentication with role-based access control. +- Audit log for authentication failures, RBAC denials, ingestion, response + planning, and response approvals. - `promtactl replay` for safe JSONL telemetry replay into the ingest API. - Browser dashboard with asset risk graph, alerts, events, rules, and response actions. @@ -41,6 +43,7 @@ go run ./cmd/promtact --demo Run with Postgres persistence: ```powershell +docker compose up -d postgres $env:PROMTACT_POSTGRES_DSN="postgres://promtact:promtact@localhost:5432/promtact?sslmode=disable" go run ./cmd/promtact --demo --policy configs\example.policy.json ``` @@ -73,6 +76,14 @@ $env:GOMODCACHE="$PWD\.cache\go-mod" go test ./... ``` +Run the optional Postgres integration test: + +```powershell +docker compose up -d postgres +$env:PROMTACT_TEST_POSTGRES_DSN="postgres://promtact:promtact@localhost:5432/promtact?sslmode=disable" +go test ./internal/store -run TestPostgresPersistenceIntegration -count=1 +``` + Replay safe JSONL telemetry into a running server: ```powershell @@ -112,6 +123,7 @@ Useful endpoints: - `POST /api/events` - `GET /api/alerts` - `GET /api/assets` +- `GET /api/audit` - `GET /api/policies` - `GET /api/responses` - `POST /api/responses` @@ -172,6 +184,8 @@ Roles: - `operator`: analyst permissions plus response approvals. - `admin`: all API operations. +Audit logs require `analyst`, `operator`, or `admin`. + ## Telemetry Replay `promtactl replay` reads newline-delimited JSON events and posts them to diff --git a/compose.yaml b/compose.yaml new file mode 100644 index 0000000..6a0d702 --- /dev/null +++ b/compose.yaml @@ -0,0 +1,19 @@ +services: + postgres: + image: postgres:16-alpine + environment: + POSTGRES_DB: promtact + POSTGRES_USER: promtact + POSTGRES_PASSWORD: promtact + ports: + - "5432:5432" + healthcheck: + test: ["CMD-SHELL", "pg_isready -U promtact -d promtact"] + interval: 5s + timeout: 3s + retries: 20 + volumes: + - postgres-data:/var/lib/postgresql/data + +volumes: + postgres-data: diff --git a/docs/architecture.md b/docs/architecture.md index 7a119b1..26d238b 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -31,15 +31,15 @@ Suricata EVE JSON. ### Domain Model -`internal/domain` defines the shared event, alert, asset, rule, and response -types. +`internal/domain` defines the shared event, alert, asset, rule, response, and +audit event types. ### Store -`internal/store` keeps events, alerts, response actions, and risk ranked assets. -For production, `--postgres-dsn` stores data in Postgres tables with JSONB -payloads and indexed core columns. For local development, `--data` writes a JSON -snapshot and restores state on startup. +`internal/store` keeps events, alerts, response actions, audit events, and risk +ranked assets. For production, `--postgres-dsn` stores data in Postgres tables +with JSONB payloads and indexed core columns. For local development, `--data` +writes a JSON snapshot and restores state on startup. ### Policy Engine @@ -81,6 +81,13 @@ API access supports user tokens with RBAC roles configured in the policy file. Token values are not stored in config; only SHA-256 token hashes are stored. The legacy `--api-token` path behaves as an admin token for compatibility. +### Audit Logging + +The HTTP layer records audit events for authentication failures, RBAC denials, +event ingestion, demo loads, response planning, and response approvals. Audit +events are first-class store records and are persisted to Postgres table +`promtact_audit_events` in production mode. + ### SIEM/Webhook Export When `--alert-webhook-url` is set, newly created alerts are sent to that @@ -102,6 +109,8 @@ flowchart LR F --> G["Response Orchestrator"] G --> H["SIEM/EDR/Firewall/Ticketing"] F --> I["Dashboard"] + B --> J["Audit Log"] + G --> J ``` The current file-backed snapshot should be treated as local development storage, diff --git a/docs/operations.md b/docs/operations.md index 40152e2..c91762e 100644 --- a/docs/operations.md +++ b/docs/operations.md @@ -24,6 +24,13 @@ For production, set Postgres in `/etc/promtact/promtact.env`: PROMTACT_POSTGRES_DSN=postgres://promtact:promtact@postgres:5432/promtact?sslmode=disable ``` +For local development and integration tests, start the bundled Compose service: + +```powershell +docker compose up -d postgres +$env:PROMTACT_POSTGRES_DSN="postgres://promtact:promtact@localhost:5432/promtact?sslmode=disable" +``` + Create a dedicated user, copy the binaries and policy file, install the unit, then enable it: @@ -68,6 +75,23 @@ Production durable storage is Postgres via `--postgres-dsn` or The local JSON snapshot configured with `--data` remains useful for development and quick labs, but it is not the production storage path. +The optional Postgres integration test is disabled by default and runs only when +`PROMTACT_TEST_POSTGRES_DSN` is set: + +```powershell +$env:PROMTACT_TEST_POSTGRES_DSN="postgres://promtact:promtact@localhost:5432/promtact?sslmode=disable" +go test ./internal/store -run TestPostgresPersistenceIntegration -count=1 +``` + +## Audit Log + +The service records audit events for authentication failures, RBAC denials, +event ingestion, demo loads, response planning, and response approvals. Audit +events are stored in Postgres table `promtact_audit_events` in production mode and +are exposed through `GET /api/audit`. + +`GET /api/audit` requires `analyst`, `operator`, or `admin`. + ## RBAC Define users in the policy file with token hashes: diff --git a/docs/roadmap.md b/docs/roadmap.md index a1b0e59..297a8c8 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -15,6 +15,9 @@ - Collector normalizers for Sysmon JSON, auditd, Zeek conn, and Suricata EVE. - Alert webhook export. - Response approval state for planned actions. +- Audit logging for authentication, RBAC, ingestion, planning, and approval + events. +- Docker Compose Postgres service plus optional Postgres integration test. - systemd and Windows service starter packaging. - AGPLv3-or-later plus commercial dual-license path. - CLA requirement from day 1. @@ -23,7 +26,7 @@ Status: implemented in this repository. ## 1-2 Weeks: Alpha -- Postgres migrations hardening and backup/restore docs. +- Postgres migration versioning and backup/restore docs. - Session-based dashboard login on top of token/RBAC API. - Policy reload without restart. - Signed tool manifests for AI-agent and MCP surfaces. @@ -40,7 +43,7 @@ Status: implemented in this repository. - Policy packs. - Deception token registry. - Response approvals. -- Integration tests with replayed telemetry. +- Integration tests with replayed telemetry and Postgres-backed API smoke tests. - Windows and Linux packaging. - Better asset graph and investigation timeline. @@ -49,7 +52,7 @@ Status: implemented in this repository. - Hardening review. - Threat model. - Signed releases. -- Audit logging. +- Tamper-evident audit log export. - Enterprise connectors. - SSO/SAML for the commercial edition. - Commercial license workflow. diff --git a/internal/auth/auth.go b/internal/auth/auth.go index f5b97cb..3365dbb 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -82,6 +82,9 @@ func (p Principal) HasAny(roles ...string) bool { } func RequiredRoles(method string, path string) []string { + if path == "/api/audit" { + return []string{RoleAnalyst, RoleOperator} + } if method == http.MethodGet || method == http.MethodHead || method == http.MethodOptions { return []string{RoleViewer, RoleAnalyst, RoleOperator, RoleIngestor} } diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go index 989b83f..1498830 100644 --- a/internal/auth/auth_test.go +++ b/internal/auth/auth_test.go @@ -24,4 +24,7 @@ func TestRequiredRoles(t *testing.T) { if roles := RequiredRoles(http.MethodPost, "/api/responses/approve"); len(roles) != 1 || roles[0] != RoleOperator { t.Fatalf("unexpected approve roles: %#v", roles) } + if roles := RequiredRoles(http.MethodGet, "/api/audit"); len(roles) != 2 || roles[0] != RoleAnalyst || roles[1] != RoleOperator { + t.Fatalf("unexpected audit roles: %#v", roles) + } } diff --git a/internal/domain/types.go b/internal/domain/types.go index a83d9b6..248d08f 100644 --- a/internal/domain/types.go +++ b/internal/domain/types.go @@ -104,6 +104,20 @@ type Asset struct { Metadata map[string]string `json:"metadata"` } +type AuditEvent struct { + ID string `json:"id"` + Timestamp time.Time `json:"timestamp"` + Actor string `json:"actor"` + Roles []string `json:"roles"` + Action string `json:"action"` + ResourceType string `json:"resource_type"` + ResourceID string `json:"resource_id,omitempty"` + Outcome string `json:"outcome"` + SourceIP string `json:"source_ip,omitempty"` + UserAgent string `json:"user_agent,omitempty"` + Metadata map[string]string `json:"metadata"` +} + type RuleDescriptor struct { ID string `json:"id"` Name string `json:"name"` @@ -119,6 +133,7 @@ type Status struct { AlertCount int `json:"alert_count"` AssetCount int `json:"asset_count"` ActionCount int `json:"action_count"` + AuditCount int `json:"audit_count"` StartedAt time.Time `json:"started_at"` StorageMode string `json:"storage_mode"` StoragePath string `json:"storage_path,omitempty"` diff --git a/internal/server/server.go b/internal/server/server.go index d244edc..2470d01 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -1,10 +1,12 @@ package server import ( + "context" "encoding/json" "errors" "fmt" "io" + "net" "net/http" "path/filepath" "strings" @@ -95,6 +97,7 @@ func (a *App) Routes() http.Handler { mux.HandleFunc("/api/events", a.handleEvents) mux.HandleFunc("/api/alerts", a.handleAlerts) mux.HandleFunc("/api/assets", a.handleAssets) + mux.HandleFunc("/api/audit", a.handleAudit) mux.HandleFunc("/api/responses/approve", a.handleResponseApproval) mux.HandleFunc("/api/responses", a.handleResponses) mux.HandleFunc("/api/policies", a.handlePolicies) @@ -117,7 +120,7 @@ func (a *App) handleStatus(w http.ResponseWriter, r *http.Request) { return } - events, alerts, assets, actions := a.store.Counts() + events, alerts, assets, actions, audits := a.store.Counts() writeJSON(w, http.StatusOK, domain.Status{ Version: Version, UptimeSeconds: int64(time.Since(a.startedAt).Seconds()), @@ -125,6 +128,7 @@ func (a *App) handleStatus(w http.ResponseWriter, r *http.Request) { AlertCount: alerts, AssetCount: assets, ActionCount: actions, + AuditCount: audits, StartedAt: a.startedAt, StorageMode: a.store.PersistenceMode(), StoragePath: a.store.PersistencePath(), @@ -148,6 +152,10 @@ func (a *App) handleEvents(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusInternalServerError, err) return } + a.recordAudit(r, principalFromRequest(r), "events.ingest", "events", "", "accepted", map[string]string{ + "events": fmt.Sprintf("%d", len(events)), + "alerts": fmt.Sprintf("%d", len(alerts)), + }) writeJSON(w, http.StatusAccepted, map[string]any{ "events_ingested": len(events), "alerts_created": len(alerts), @@ -174,6 +182,14 @@ func (a *App) handleAssets(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusOK, a.store.ListAssets()) } +func (a *App) handleAudit(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + methodNotAllowed(w) + return + } + writeJSON(w, http.StatusOK, a.store.ListAudits()) +} + func (a *App) handlePolicies(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet { methodNotAllowed(w) @@ -211,6 +227,15 @@ func (a *App) handleResponseApproval(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusNotFound, errors.New("action not found")) return } + principal := principalFromRequest(r) + if principal.Name == "" { + principal.Name = req.ApprovedBy + } + a.recordAudit(r, principal, "responses.approve", "response_action", action.ID, "accepted", map[string]string{ + "approved_by": req.ApprovedBy, + "asset_id": action.AssetID, + "action_type": action.Type, + }) writeJSON(w, http.StatusAccepted, action) } @@ -243,6 +268,9 @@ func (a *App) handleResponses(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusInternalServerError, err) return } + a.recordAudit(r, principalFromRequest(r), "responses.plan", "alert", alert.ID, "accepted", map[string]string{ + "actions": fmt.Sprintf("%d", len(actions)), + }) writeJSON(w, http.StatusAccepted, actions) default: methodNotAllowed(w) @@ -259,6 +287,9 @@ func (a *App) handleDemo(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusInternalServerError, err) return } + a.recordAudit(r, principalFromRequest(r), "demo.load", "demo", "", "accepted", map[string]string{ + "alerts": fmt.Sprintf("%d", len(alerts)), + }) writeJSON(w, http.StatusAccepted, map[string]any{ "alerts_created": len(alerts), "alerts": alerts, @@ -329,6 +360,29 @@ func (a *App) nextID(prefix string) string { return fmt.Sprintf("%s-%d", prefix, a.counter.Add(1)) } +func (a *App) recordAudit(r *http.Request, principal auth.Principal, action string, resourceType string, resourceID string, outcome string, metadata map[string]string) { + if metadata == nil { + metadata = make(map[string]string) + } + if principal.Name == "" { + principal.Name = "anonymous" + } + event := domain.AuditEvent{ + ID: a.nextID("aud"), + Timestamp: time.Now().UTC(), + Actor: principal.Name, + Roles: append([]string(nil), principal.Roles...), + Action: action, + ResourceType: resourceType, + ResourceID: resourceID, + Outcome: outcome, + SourceIP: sourceIP(r), + UserAgent: r.UserAgent(), + Metadata: metadata, + } + _ = a.store.AddAudit(event) +} + func (a *App) exportAlerts(alerts []domain.Alert) { if len(alerts) == 0 || a.webhook.URL == "" { return @@ -429,14 +483,21 @@ func (a *App) withAuth(next http.Handler) http.Handler { } principal, ok := a.auth.Authenticate(r) if !ok { + a.recordAudit(r, auth.Principal{Name: "anonymous"}, "auth.authenticate", "http_request", r.URL.Path, "denied", map[string]string{ + "method": r.Method, + }) writeError(w, http.StatusUnauthorized, errors.New("missing or invalid API token")) return } required := auth.RequiredRoles(r.Method, r.URL.Path) if !principal.HasAny(required...) { + a.recordAudit(r, principal, "auth.authorize", "http_request", r.URL.Path, "denied", map[string]string{ + "method": r.Method, + }) writeError(w, http.StatusForbidden, errors.New("insufficient role")) return } + r = r.WithContext(context.WithValue(r.Context(), principalContextKey{}, principal)) next.ServeHTTP(w, r) }) } @@ -444,3 +505,24 @@ func (a *App) withAuth(next http.Handler) http.Handler { func isReadOnly(method string) bool { return method == http.MethodGet || method == http.MethodHead || method == http.MethodOptions } + +type principalContextKey struct{} + +func principalFromRequest(r *http.Request) auth.Principal { + principal, ok := r.Context().Value(principalContextKey{}).(auth.Principal) + if !ok { + return auth.Principal{} + } + return principal +} + +func sourceIP(r *http.Request) string { + if forwarded := strings.TrimSpace(r.Header.Get("X-Forwarded-For")); forwarded != "" { + return strings.TrimSpace(strings.Split(forwarded, ",")[0]) + } + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err == nil { + return host + } + return r.RemoteAddr +} diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 5b364c8..b584b55 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -98,6 +98,27 @@ func TestRBACBlocksInsufficientRole(t *testing.T) { } } +func TestRBACBlocksAuditForViewer(t *testing.T) { + app, err := NewWithOptions(Options{ + Users: []auth.UserConfig{{Name: "viewer", TokenHash: auth.HashToken("view-token"), Roles: []string{auth.RoleViewer}}}, + }) + if err != nil { + t.Fatalf("new app: %v", err) + } + + req := httptest.NewRequest(http.MethodGet, "/api/audit", nil) + req.Header.Set("Authorization", "Bearer view-token") + rec := httptest.NewRecorder() + app.Routes().ServeHTTP(rec, req) + if rec.Code != http.StatusForbidden { + t.Fatalf("expected 403, got %d", rec.Code) + } + audits := app.store.ListAudits() + if len(audits) != 1 || audits[0].Action != "auth.authorize" || audits[0].Outcome != "denied" { + t.Fatalf("unexpected audit log: %#v", audits) + } +} + func TestEmptyListEndpointsReturnArrays(t *testing.T) { app, err := NewWithOptions(Options{}) if err != nil { @@ -168,6 +189,20 @@ func TestResponseApprovalEndpoint(t *testing.T) { if approved.ApprovalStatus != "approved" || approved.ApprovedBy != "alice" { t.Fatalf("unexpected approval response: %#v", approved) } + audits := app.store.ListAudits() + if len(audits) == 0 { + t.Fatal("expected audit events") + } + foundApproval := false + for _, audit := range audits { + if audit.Action == "responses.approve" && audit.ResourceID == actionID && audit.Outcome == "accepted" { + foundApproval = true + break + } + } + if !foundApproval { + t.Fatalf("expected response approval audit event, got %#v", audits) + } } func TestAlertWebhookExportsNewAlerts(t *testing.T) { diff --git a/internal/store/persistence.go b/internal/store/persistence.go index 07833b4..61b2527 100644 --- a/internal/store/persistence.go +++ b/internal/store/persistence.go @@ -18,6 +18,7 @@ type snapshot struct { Events []domain.Event `json:"events"` Alerts []domain.Alert `json:"alerts"` Actions []domain.ResponseAction `json:"actions"` + Audits []domain.AuditEvent `json:"audits"` Assets map[string]domain.Asset `json:"assets"` } @@ -48,6 +49,7 @@ func NewWithPath(path string) (*Store, error) { s.events = append([]domain.Event(nil), snap.Events...) s.alerts = append([]domain.Alert(nil), snap.Alerts...) s.actions = append([]domain.ResponseAction(nil), snap.Actions...) + s.audits = append([]domain.AuditEvent(nil), snap.Audits...) if snap.Assets != nil { s.assets = snap.Assets } else { @@ -69,6 +71,7 @@ func (s *Store) persistLocked() error { Events: append([]domain.Event(nil), s.events...), Alerts: append([]domain.Alert(nil), s.alerts...), Actions: append([]domain.ResponseAction(nil), s.actions...), + Audits: append([]domain.AuditEvent(nil), s.audits...), Assets: cloneAssets(s.assets), } @@ -118,6 +121,13 @@ func (s *Store) persistActionsLocked(actions []domain.ResponseAction) error { return nil } +func (s *Store) persistAuditLocked(event domain.AuditEvent) error { + if s.db != nil { + return s.postgresPersistAuditLocked(event) + } + return nil +} + func replaceFile(src string, dst string) error { if err := os.Rename(src, dst); err == nil { return nil diff --git a/internal/store/postgres.go b/internal/store/postgres.go index b2432fa..5f3de15 100644 --- a/internal/store/postgres.go +++ b/internal/store/postgres.go @@ -90,7 +90,21 @@ CREATE TABLE IF NOT EXISTS promtact_assets ( last_seen TIMESTAMPTZ, risk_score INTEGER, data JSONB NOT NULL -);`) +); + +CREATE TABLE IF NOT EXISTS promtact_audit_events ( + id TEXT PRIMARY KEY, + occurred_at TIMESTAMPTZ NOT NULL, + actor TEXT, + action TEXT NOT NULL, + resource_type TEXT, + resource_id TEXT, + outcome TEXT, + data JSONB NOT NULL +); +CREATE INDEX IF NOT EXISTS idx_promtact_audit_events_occurred_at ON promtact_audit_events (occurred_at DESC); +CREATE INDEX IF NOT EXISTS idx_promtact_audit_events_actor ON promtact_audit_events (actor); +CREATE INDEX IF NOT EXISTS idx_promtact_audit_events_action ON promtact_audit_events (action);`) return err } @@ -107,6 +121,9 @@ func (s *Store) postgresLoad(ctx context.Context) error { if err := s.postgresLoadAssets(ctx); err != nil { return err } + if err := s.postgresLoadAudits(ctx); err != nil { + return err + } s.rebuildFingerprintsLocked() if len(s.assets) == 0 { s.rebuildAssetsLocked() @@ -196,6 +213,26 @@ func (s *Store) postgresLoadAssets(ctx context.Context) error { return rows.Err() } +func (s *Store) postgresLoadAudits(ctx context.Context) error { + rows, err := s.db.QueryContext(ctx, `SELECT data FROM promtact_audit_events ORDER BY occurred_at ASC, id ASC`) + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var data []byte + if err := rows.Scan(&data); err != nil { + return err + } + var event domain.AuditEvent + if err := json.Unmarshal(data, &event); err != nil { + return err + } + s.audits = append(s.audits, event) + } + return rows.Err() +} + func (s *Store) postgresPersistEventLocked(event domain.Event) error { ctx, cancel := context.WithTimeout(context.Background(), postgresTimeout) defer cancel() @@ -265,6 +302,28 @@ ON CONFLICT (id) DO UPDATE SET return nil } +func (s *Store) postgresPersistAuditLocked(event domain.AuditEvent) error { + ctx, cancel := context.WithTimeout(context.Background(), postgresTimeout) + defer cancel() + data, err := json.Marshal(event) + if err != nil { + return err + } + _, err = s.db.ExecContext(ctx, ` +INSERT INTO promtact_audit_events (id, occurred_at, actor, action, resource_type, resource_id, outcome, data) +VALUES ($1, $2, $3, $4, $5, $6, $7, $8) +ON CONFLICT (id) DO UPDATE SET + occurred_at = EXCLUDED.occurred_at, + actor = EXCLUDED.actor, + action = EXCLUDED.action, + resource_type = EXCLUDED.resource_type, + resource_id = EXCLUDED.resource_id, + outcome = EXCLUDED.outcome, + data = EXCLUDED.data`, + event.ID, event.Timestamp, event.Actor, event.Action, event.ResourceType, event.ResourceID, event.Outcome, data) + return err +} + func (s *Store) postgresPersistAssetsLocked(ctx context.Context) error { for _, asset := range s.assets { data, err := json.Marshal(asset) diff --git a/internal/store/postgres_integration_test.go b/internal/store/postgres_integration_test.go new file mode 100644 index 0000000..c585589 --- /dev/null +++ b/internal/store/postgres_integration_test.go @@ -0,0 +1,142 @@ +package store + +import ( + "os" + "strings" + "testing" + "time" + + "github.com/hunterinvariants/promtact/internal/domain" +) + +func TestPostgresPersistenceIntegration(t *testing.T) { + dsn := os.Getenv("PROMTACT_TEST_POSTGRES_DSN") + if dsn == "" { + t.Skip("set PROMTACT_TEST_POSTGRES_DSN to run Postgres integration tests") + } + + suffix := strings.ReplaceAll(time.Now().UTC().Format("20060102150405.000000000"), ".", "") + eventID := "it-evt-" + suffix + alertID := "it-alert-" + suffix + actionID := "it-act-" + suffix + auditID := "it-aud-" + suffix + assetID := "it-asset-" + suffix + + s, err := NewWithPostgres(dsn) + if err != nil { + t.Fatalf("new postgres store: %v", err) + } + defer cleanupPostgresIntegrationRows(t, s, eventID, alertID, actionID, auditID, assetID) + + event := domain.Event{ + ID: eventID, + Timestamp: time.Now().UTC(), + Kind: domain.EventAgentToolCall, + AssetID: assetID, + Hostname: assetID, + ToolName: "shell_exec", + } + if err := s.AddEvent(event); err != nil { + t.Fatalf("add event: %v", err) + } + + alert := domain.Alert{ + ID: alertID, + Fingerprint: "integration:" + eventID, + RuleID: "integration", + Title: "integration alert", + Severity: domain.SeverityHigh, + Status: domain.AlertOpen, + AssetID: assetID, + CreatedAt: time.Now().UTC(), + EventIDs: []string{eventID}, + } + if _, err := s.AddAlerts([]domain.Alert{alert}); err != nil { + t.Fatalf("add alert: %v", err) + } + + action := domain.ResponseAction{ + ID: actionID, + Type: "isolate_host", + Mode: "dry-run", + AssetID: assetID, + ApprovalStatus: "required", + CreatedAt: time.Now().UTC(), + } + if err := s.AddActions([]domain.ResponseAction{action}); err != nil { + t.Fatalf("add action: %v", err) + } + if _, ok, err := s.ApproveAction(actionID, "integration", time.Now().UTC()); err != nil || !ok { + t.Fatalf("approve action ok=%v err=%v", ok, err) + } + + audit := domain.AuditEvent{ + ID: auditID, + Timestamp: time.Now().UTC(), + Actor: "integration", + Action: "responses.approve", + ResourceType: "response_action", + ResourceID: actionID, + Outcome: "accepted", + Metadata: map[string]string{"asset_id": assetID}, + } + if err := s.AddAudit(audit); err != nil { + t.Fatalf("add audit: %v", err) + } + + loaded, err := NewWithPostgres(dsn) + if err != nil { + t.Fatalf("reload postgres store: %v", err) + } + defer loaded.Close() + + if !hasEvent(loaded.ListEvents(), eventID) { + t.Fatalf("expected reloaded event %s", eventID) + } + if !hasAction(loaded.ListActions(), actionID, "approved") { + t.Fatalf("expected reloaded approved action %s", actionID) + } + if !hasAudit(loaded.ListAudits(), auditID) { + t.Fatalf("expected reloaded audit event %s", auditID) + } +} + +func cleanupPostgresIntegrationRows(t *testing.T, s *Store, ids ...string) { + t.Helper() + if s == nil || s.db == nil { + return + } + for _, table := range []string{"promtact_events", "promtact_alerts", "promtact_actions", "promtact_audit_events", "promtact_assets"} { + if _, err := s.db.Exec("DELETE FROM "+table+" WHERE id IN ($1, $2, $3, $4, $5)", ids[0], ids[1], ids[2], ids[3], ids[4]); err != nil { + t.Logf("cleanup %s: %v", table, err) + } + } + _ = s.Close() +} + +func hasEvent(events []domain.Event, id string) bool { + for _, event := range events { + if event.ID == id { + return true + } + } + return false +} + +func hasAction(actions []domain.ResponseAction, id string, approvalStatus string) bool { + for _, action := range actions { + if action.ID == id && action.ApprovalStatus == approvalStatus { + return true + } + } + return false +} + +func hasAudit(audits []domain.AuditEvent, id string) bool { + for _, audit := range audits { + if audit.ID == id { + return true + } + } + return false +} diff --git a/internal/store/store.go b/internal/store/store.go index 1004293..41801d8 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -16,6 +16,7 @@ type Store struct { events []domain.Event alerts []domain.Alert actions []domain.ResponseAction + audits []domain.AuditEvent assets map[string]domain.Asset fingerprints map[string]struct{} path string @@ -152,6 +153,30 @@ func (s *Store) ListActions() []domain.ResponseAction { return actions } +func (s *Store) AddAudit(event domain.AuditEvent) error { + s.mu.Lock() + defer s.mu.Unlock() + + s.audits = append(s.audits, event) + if err := s.persistAuditLocked(event); err != nil { + s.lastErr = err.Error() + return err + } + return s.persistLocked() +} + +func (s *Store) ListAudits() []domain.AuditEvent { + s.mu.RLock() + defer s.mu.RUnlock() + + audits := make([]domain.AuditEvent, len(s.audits)) + copy(audits, s.audits) + sort.Slice(audits, func(i, j int) bool { + return audits[i].Timestamp.After(audits[j].Timestamp) + }) + return audits +} + func (s *Store) ListAssets() []domain.Asset { s.mu.RLock() defer s.mu.RUnlock() @@ -169,11 +194,11 @@ func (s *Store) ListAssets() []domain.Asset { return assets } -func (s *Store) Counts() (events int, alerts int, assets int, actions int) { +func (s *Store) Counts() (events int, alerts int, assets int, actions int, audits int) { s.mu.RLock() defer s.mu.RUnlock() - return len(s.events), len(s.alerts), len(s.assets), len(s.actions) + return len(s.events), len(s.alerts), len(s.assets), len(s.actions), len(s.audits) } func (s *Store) PersistencePath() string { diff --git a/internal/store/store_test.go b/internal/store/store_test.go index ee7e777..d614f43 100644 --- a/internal/store/store_test.go +++ b/internal/store/store_test.go @@ -57,15 +57,46 @@ func TestStorePersistsAndLoadsSnapshot(t *testing.T) { if err != nil { t.Fatalf("load store: %v", err) } - events, alerts, assets, actions := loaded.Counts() - if events != 1 || alerts != 1 || assets != 1 || actions != 1 { - t.Fatalf("unexpected counts: events=%d alerts=%d assets=%d actions=%d", events, alerts, assets, actions) + events, alerts, assets, actions, audits := loaded.Counts() + if events != 1 || alerts != 1 || assets != 1 || actions != 1 || audits != 0 { + t.Fatalf("unexpected counts: events=%d alerts=%d assets=%d actions=%d audits=%d", events, alerts, assets, actions, audits) } if loaded.LastPersistenceError() != "" { t.Fatalf("unexpected persistence error: %s", loaded.LastPersistenceError()) } } +func TestStorePersistsAndLoadsAuditSnapshot(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + s, err := NewWithPath(path) + if err != nil { + t.Fatalf("new store: %v", err) + } + + audit := domain.AuditEvent{ + ID: "aud-1", + Timestamp: time.Now().UTC(), + Actor: "alice", + Action: "responses.approve", + ResourceType: "response_action", + ResourceID: "act-1", + Outcome: "accepted", + Metadata: map[string]string{"asset_id": "asset-1"}, + } + if err := s.AddAudit(audit); err != nil { + t.Fatalf("add audit: %v", err) + } + + loaded, err := NewWithPath(path) + if err != nil { + t.Fatalf("load store: %v", err) + } + audits := loaded.ListAudits() + if len(audits) != 1 || audits[0].Actor != "alice" || audits[0].Action != "responses.approve" { + t.Fatalf("unexpected audit events: %#v", audits) + } +} + func TestStoreSkipsDuplicateAlertFingerprintsAfterLoad(t *testing.T) { path := filepath.Join(t.TempDir(), "state.json") s, err := NewWithPath(path) diff --git a/web/app.js b/web/app.js index 8655b24..f437220 100644 --- a/web/app.js +++ b/web/app.js @@ -6,7 +6,8 @@ const els = { events: document.querySelector("#metric-events"), alerts: document.querySelector("#metric-alerts"), assets: document.querySelector("#metric-assets"), - actions: document.querySelector("#metric-actions") + actions: document.querySelector("#metric-actions"), + audit: document.querySelector("#metric-audit") }, graph: document.querySelector("#asset-graph"), assetsBody: document.querySelector("#assets-body"), @@ -71,6 +72,7 @@ function renderStatus(status) { els.metrics.alerts.textContent = status.alert_count; els.metrics.assets.textContent = status.asset_count; els.metrics.actions.textContent = status.action_count; + els.metrics.audit.textContent = status.audit_count; els.version.textContent = status.version; } @@ -114,7 +116,7 @@ function renderAlerts(alerts) {
${escapeHtml(alert.title)}
${escapeHtml(alert.severity)} -
${escapeHtml(alert.asset_id || "unknown asset")} · ${escapeHtml(alert.rule_id)}
+
${escapeHtml(alert.asset_id || "unknown asset")} - ${escapeHtml(alert.rule_id)}

${escapeHtml(alert.description)}

${renderEvidence(alert.evidence)}
@@ -141,7 +143,7 @@ function renderEvents(events) {
${escapeHtml(event.kind)}
${escapeHtml(event.asset_id || "no asset")}
-
${formatTime(event.timestamp)} · ${escapeHtml(event.hostname || event.source_ip || "-")}
+
${formatTime(event.timestamp)} - ${escapeHtml(event.hostname || event.source_ip || "-")}

${escapeHtml(event.signal || event.command || event.destination || "-")}

${(event.labels || []).map((label) => `${escapeHtml(label)}`).join("")}
`; @@ -164,7 +166,7 @@ function renderActions(actions) {
${escapeHtml(action.type)}
${escapeHtml(action.approval_status || action.mode)} -
${escapeHtml(action.asset_id || "unknown asset")} · ${escapeHtml(action.target || "-")}
+
${escapeHtml(action.asset_id || "unknown asset")} - ${escapeHtml(action.target || "-")}

${escapeHtml(action.reason || "")}

${action.approval_status === "required" ? `
approval required
` : ""} `; diff --git a/web/index.html b/web/index.html index 6cd5f57..6d02b44 100644 --- a/web/index.html +++ b/web/index.html @@ -36,6 +36,10 @@

Agentic Threat Control Plane

Dry Runs 0 +
+ Audit + 0 +
@@ -112,4 +116,3 @@

Rules

- diff --git a/web/styles.css b/web/styles.css index be94d01..65305d8 100644 --- a/web/styles.css +++ b/web/styles.css @@ -103,7 +103,7 @@ h2 { .metrics { display: grid; - grid-template-columns: repeat(4, minmax(0, 1fr)); + grid-template-columns: repeat(5, minmax(0, 1fr)); gap: 12px; margin-bottom: 14px; } @@ -384,4 +384,3 @@ td { font-size: 25px; } } - From e87fe2eabb4302a6b6edccfcb0b68744ff80ba86 Mon Sep 17 00:00:00 2001 From: hunterinvariants Date: Sat, 6 Jun 2026 17:11:06 +0200 Subject: [PATCH 004/137] Run CI against Postgres --- .github/workflows/ci.yml | 43 ++++++++++++++++++++++++++++++++++++++++ README.md | 5 +++++ 2 files changed, 48 insertions(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a143f6a..b852d81 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -5,12 +5,55 @@ on: push: branches: [main] +permissions: + contents: read + jobs: test: runs-on: ubuntu-latest + services: + postgres: + image: postgres:16-alpine + env: + POSTGRES_DB: promtact + POSTGRES_USER: promtact + POSTGRES_PASSWORD: promtact + ports: + - 5432:5432 + options: >- + --health-cmd "pg_isready -U promtact -d promtact" + --health-interval 5s + --health-timeout 3s + --health-retries 20 + env: + PROMTACT_TEST_POSTGRES_DSN: postgres://promtact:promtact@localhost:5432/promtact?sslmode=disable steps: - uses: actions/checkout@v4 - uses: actions/setup-go@v5 with: go-version: "1.25.x" + - run: go mod download + - run: go vet ./... - run: go test ./... + + build: + runs-on: ubuntu-latest + needs: test + strategy: + fail-fast: false + matrix: + goos: [linux, windows] + goarch: [amd64, arm64] + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.25.x" + - name: Build binaries + shell: bash + run: | + suffix="" + if [ "${{ matrix.goos }}" = "windows" ]; then suffix=".exe"; fi + mkdir -p dist + GOOS=${{ matrix.goos }} GOARCH=${{ matrix.goarch }} go build -o "dist/promtact-${{ matrix.goos }}-${{ matrix.goarch }}${suffix}" ./cmd/promtact + GOOS=${{ matrix.goos }} GOARCH=${{ matrix.goarch }} go build -o "dist/promtactl-${{ matrix.goos }}-${{ matrix.goarch }}${suffix}" ./cmd/promtactl diff --git a/README.md b/README.md index 1642715..beb6df2 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,7 @@ # Promtact +[![ci](https://github.com/hunterinvariants/promtact/actions/workflows/ci.yml/badge.svg)](https://github.com/hunterinvariants/promtact/actions/workflows/ci.yml) + Promtact is a defensive control plane for detecting and containing agentic threat behavior across AI-agent tool calls, host telemetry, network egress, deception signals, and response workflows. @@ -76,6 +78,9 @@ $env:GOMODCACHE="$PWD\.cache\go-mod" go test ./... ``` +GitHub CI runs the same test suite with a real Postgres service and builds +Linux/Windows binaries for `amd64` and `arm64`. + Run the optional Postgres integration test: ```powershell From 643bb42689b75ffeff22884bc45b5b9404082e7b Mon Sep 17 00:00:00 2001 From: hunterinvariants Date: Sat, 6 Jun 2026 17:25:34 +0200 Subject: [PATCH 005/137] Add schema migrations and security automation --- .github/dependabot.yml | 23 +++++ .github/workflows/codeql.yml | 25 +++++ README.md | 3 + docs/architecture.md | 5 +- docs/operations.md | 5 +- docs/roadmap.md | 4 +- internal/domain/types.go | 1 + internal/server/server.go | 1 + internal/store/postgres.go | 101 +++++++++++++++++++- internal/store/postgres_integration_test.go | 6 ++ internal/store/store.go | 30 +++--- 11 files changed, 186 insertions(+), 18 deletions(-) create mode 100644 .github/dependabot.yml create mode 100644 .github/workflows/codeql.yml diff --git a/.github/dependabot.yml b/.github/dependabot.yml new file mode 100644 index 0000000..570bfe6 --- /dev/null +++ b/.github/dependabot.yml @@ -0,0 +1,23 @@ +version: 2 +updates: + - package-ecosystem: gomod + directory: / + schedule: + interval: weekly + day: monday + time: "04:00" + open-pull-requests-limit: 5 + labels: + - dependencies + - go + + - package-ecosystem: github-actions + directory: / + schedule: + interval: weekly + day: monday + time: "04:30" + open-pull-requests-limit: 5 + labels: + - dependencies + - github-actions diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml new file mode 100644 index 0000000..1b66dba --- /dev/null +++ b/.github/workflows/codeql.yml @@ -0,0 +1,25 @@ +name: codeql + +on: + pull_request: + push: + branches: [main] + schedule: + - cron: "24 3 * * 1" + +permissions: + actions: read + contents: read + security-events: write + +jobs: + analyze: + name: Analyze Go + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: github/codeql-action/init@v3 + with: + languages: go + - uses: github/codeql-action/autobuild@v3 + - uses: github/codeql-action/analyze@v3 diff --git a/README.md b/README.md index beb6df2..348bc61 100644 --- a/README.md +++ b/README.md @@ -81,6 +81,9 @@ go test ./... GitHub CI runs the same test suite with a real Postgres service and builds Linux/Windows binaries for `amd64` and `arm64`. +GitHub security automation includes CodeQL analysis and Dependabot updates for +Go modules and GitHub Actions. + Run the optional Postgres integration test: ```powershell diff --git a/docs/architecture.md b/docs/architecture.md index 26d238b..ad051a7 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -38,8 +38,9 @@ audit event types. `internal/store` keeps events, alerts, response actions, audit events, and risk ranked assets. For production, `--postgres-dsn` stores data in Postgres tables -with JSONB payloads and indexed core columns. For local development, `--data` -writes a JSON snapshot and restores state on startup. +with JSONB payloads and indexed core columns. Schema changes are applied through +versioned migrations recorded in `promtact_schema_migrations`. For local +development, `--data` writes a JSON snapshot and restores state on startup. ### Policy Engine diff --git a/docs/operations.md b/docs/operations.md index c91762e..3bf7ceb 100644 --- a/docs/operations.md +++ b/docs/operations.md @@ -70,11 +70,14 @@ The payload type is `promtact.alerts`. ## Storage Production durable storage is Postgres via `--postgres-dsn` or -`PROMTACT_POSTGRES_DSN`. Promtact creates the required tables automatically. +`PROMTACT_POSTGRES_DSN`. Promtact creates and upgrades the required tables through +versioned migrations tracked in `promtact_schema_migrations`. The local JSON snapshot configured with `--data` remains useful for development and quick labs, but it is not the production storage path. +`GET /api/status` exposes the active `schema_version` when Postgres is enabled. + The optional Postgres integration test is disabled by default and runs only when `PROMTACT_TEST_POSTGRES_DSN` is set: diff --git a/docs/roadmap.md b/docs/roadmap.md index 297a8c8..aa692fe 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -18,6 +18,8 @@ - Audit logging for authentication, RBAC, ingestion, planning, and approval events. - Docker Compose Postgres service plus optional Postgres integration test. +- Versioned Postgres schema migrations. +- CodeQL and Dependabot GitHub automation. - systemd and Windows service starter packaging. - AGPLv3-or-later plus commercial dual-license path. - CLA requirement from day 1. @@ -26,7 +28,7 @@ Status: implemented in this repository. ## 1-2 Weeks: Alpha -- Postgres migration versioning and backup/restore docs. +- Postgres backup/restore docs. - Session-based dashboard login on top of token/RBAC API. - Policy reload without restart. - Signed tool manifests for AI-agent and MCP surfaces. diff --git a/internal/domain/types.go b/internal/domain/types.go index 248d08f..e9c8c2d 100644 --- a/internal/domain/types.go +++ b/internal/domain/types.go @@ -137,6 +137,7 @@ type Status struct { StartedAt time.Time `json:"started_at"` StorageMode string `json:"storage_mode"` StoragePath string `json:"storage_path,omitempty"` + SchemaVersion int `json:"schema_version,omitempty"` LastStorageError string `json:"last_storage_error,omitempty"` LastExportError string `json:"last_export_error,omitempty"` } diff --git a/internal/server/server.go b/internal/server/server.go index 2470d01..5a4fb1c 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -132,6 +132,7 @@ func (a *App) handleStatus(w http.ResponseWriter, r *http.Request) { StartedAt: a.startedAt, StorageMode: a.store.PersistenceMode(), StoragePath: a.store.PersistencePath(), + SchemaVersion: a.store.SchemaVersion(), LastStorageError: a.store.LastPersistenceError(), LastExportError: a.lastExportError(), }) diff --git a/internal/store/postgres.go b/internal/store/postgres.go index 5f3de15..4689083 100644 --- a/internal/store/postgres.go +++ b/internal/store/postgres.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "encoding/json" + "fmt" "time" _ "github.com/jackc/pgx/v5/stdlib" @@ -51,7 +52,101 @@ func (s *Store) Close() error { } func (s *Store) postgresMigrate(ctx context.Context) error { - _, err := s.db.ExecContext(ctx, ` + if _, err := s.db.ExecContext(ctx, `SELECT pg_advisory_lock(72743001)`); err != nil { + return err + } + defer s.db.ExecContext(context.Background(), `SELECT pg_advisory_unlock(72743001)`) + + if _, err := s.db.ExecContext(ctx, ` +CREATE TABLE IF NOT EXISTS promtact_schema_migrations ( + version INTEGER PRIMARY KEY, + name TEXT NOT NULL, + applied_at TIMESTAMPTZ NOT NULL DEFAULT now() +);`); err != nil { + return err + } + + applied, err := s.postgresAppliedMigrations(ctx) + if err != nil { + return err + } + for _, migration := range postgresMigrations { + if applied[migration.Version] { + continue + } + if err := s.applyPostgresMigration(ctx, migration); err != nil { + return err + } + applied[migration.Version] = true + } + + version, err := s.postgresCurrentSchemaVersion(ctx) + if err != nil { + return err + } + s.schemaVersion = version + return nil +} + +func (s *Store) postgresAppliedMigrations(ctx context.Context) (map[int]bool, error) { + rows, err := s.db.QueryContext(ctx, `SELECT version FROM promtact_schema_migrations`) + if err != nil { + return nil, err + } + defer rows.Close() + + applied := map[int]bool{} + for rows.Next() { + var version int + if err := rows.Scan(&version); err != nil { + return nil, err + } + applied[version] = true + } + return applied, rows.Err() +} + +func (s *Store) applyPostgresMigration(ctx context.Context, migration postgresMigration) error { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + if _, err := tx.ExecContext(ctx, migration.SQL); err != nil { + _ = tx.Rollback() + return fmt.Errorf("apply postgres migration %d %s: %w", migration.Version, migration.Name, err) + } + if _, err := tx.ExecContext(ctx, ` +INSERT INTO promtact_schema_migrations (version, name) +VALUES ($1, $2) +ON CONFLICT (version) DO NOTHING`, migration.Version, migration.Name); err != nil { + _ = tx.Rollback() + return err + } + return tx.Commit() +} + +func (s *Store) postgresCurrentSchemaVersion(ctx context.Context) (int, error) { + var version sql.NullInt64 + if err := s.db.QueryRowContext(ctx, `SELECT max(version) FROM promtact_schema_migrations`).Scan(&version); err != nil { + return 0, err + } + if !version.Valid { + return 0, nil + } + return int(version.Int64), nil +} + +type postgresMigration struct { + Version int + Name string + SQL string +} + +var postgresMigrations = []postgresMigration{ + { + Version: 1, + Name: "initial_schema", + SQL: ` CREATE TABLE IF NOT EXISTS promtact_events ( id TEXT PRIMARY KEY, occurred_at TIMESTAMPTZ, @@ -104,8 +199,8 @@ CREATE TABLE IF NOT EXISTS promtact_audit_events ( ); CREATE INDEX IF NOT EXISTS idx_promtact_audit_events_occurred_at ON promtact_audit_events (occurred_at DESC); CREATE INDEX IF NOT EXISTS idx_promtact_audit_events_actor ON promtact_audit_events (actor); -CREATE INDEX IF NOT EXISTS idx_promtact_audit_events_action ON promtact_audit_events (action);`) - return err +CREATE INDEX IF NOT EXISTS idx_promtact_audit_events_action ON promtact_audit_events (action);`, + }, } func (s *Store) postgresLoad(ctx context.Context) error { diff --git a/internal/store/postgres_integration_test.go b/internal/store/postgres_integration_test.go index c585589..a17d60e 100644 --- a/internal/store/postgres_integration_test.go +++ b/internal/store/postgres_integration_test.go @@ -27,6 +27,9 @@ func TestPostgresPersistenceIntegration(t *testing.T) { t.Fatalf("new postgres store: %v", err) } defer cleanupPostgresIntegrationRows(t, s, eventID, alertID, actionID, auditID, assetID) + if s.SchemaVersion() < 1 { + t.Fatalf("expected postgres schema version, got %d", s.SchemaVersion()) + } event := domain.Event{ ID: eventID, @@ -89,6 +92,9 @@ func TestPostgresPersistenceIntegration(t *testing.T) { t.Fatalf("reload postgres store: %v", err) } defer loaded.Close() + if loaded.SchemaVersion() != s.SchemaVersion() { + t.Fatalf("unexpected reloaded schema version: got %d want %d", loaded.SchemaVersion(), s.SchemaVersion()) + } if !hasEvent(loaded.ListEvents(), eventID) { t.Fatalf("expected reloaded event %s", eventID) diff --git a/internal/store/store.go b/internal/store/store.go index 41801d8..4932368 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -10,17 +10,18 @@ import ( ) type Store struct { - mu sync.RWMutex - db *sql.DB - mode string - events []domain.Event - alerts []domain.Alert - actions []domain.ResponseAction - audits []domain.AuditEvent - assets map[string]domain.Asset - fingerprints map[string]struct{} - path string - lastErr string + mu sync.RWMutex + db *sql.DB + mode string + events []domain.Event + alerts []domain.Alert + actions []domain.ResponseAction + audits []domain.AuditEvent + assets map[string]domain.Asset + fingerprints map[string]struct{} + path string + lastErr string + schemaVersion int } func New() *Store { @@ -222,6 +223,13 @@ func (s *Store) LastPersistenceError() string { return s.lastErr } +func (s *Store) SchemaVersion() int { + s.mu.RLock() + defer s.mu.RUnlock() + + return s.schemaVersion +} + func (s *Store) upsertAssetLocked(event domain.Event) { if event.AssetID == "" { return From be02588a5b1d3fcd51eb98a017900ee5447eedbb Mon Sep 17 00:00:00 2001 From: hunterinvariants Date: Sat, 6 Jun 2026 21:37:39 +0200 Subject: [PATCH 006/137] Add dashboard session login --- README.md | 6 +- docs/architecture.md | 4 +- docs/operations.md | 26 +++++ internal/auth/auth.go | 159 +++++++++++++++++++++++-- internal/auth/auth_test.go | 31 +++++ internal/server/server.go | 101 +++++++++++++++- internal/server/server_test.go | 57 +++++++++ web/app.js | 207 ++++++++++++++++++++++++++++----- web/index.html | 204 +++++++++++++++++--------------- web/styles.css | 83 +++++++++++++ 10 files changed, 750 insertions(+), 128 deletions(-) diff --git a/README.md b/README.md index 348bc61..bd96014 100644 --- a/README.md +++ b/README.md @@ -24,7 +24,7 @@ malware behavior, or autonomous propagation. Demo data generates telemetry only. planning, and response approvals. - `promtactl replay` for safe JSONL telemetry replay into the ingest API. - Browser dashboard with asset risk graph, alerts, events, rules, and response - actions. + actions, plus session-based dashboard login. - Alert webhook export for SIEM-style integrations. - systemd and Windows service starter packaging. - AGPLv3-or-later community license, commercial dual-license path, and CLA from @@ -156,6 +156,10 @@ When users are configured in the policy file, all API endpoints require against RBAC roles. `--api-token` remains a legacy admin-token compatibility path. +The dashboard uses `POST /api/session` to exchange a configured user name and +token for a session cookie. `GET /api/session` reports the current dashboard +state and `DELETE /api/session` logs out. + ## Policy Configuration The policy file is JSON: diff --git a/docs/architecture.md b/docs/architecture.md index ad051a7..3a17abc 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -74,7 +74,9 @@ as requiring approval before any future execution backend can act on them. ### Dashboard `web/` provides an operational dashboard for assets, alerts, events, policies, -and dry-run response actions. +and dry-run response actions. The browser uses `POST /api/session` to create a +session cookie and `GET /api/session` to restore its authenticated state on +reload. ### Authentication And RBAC diff --git a/docs/operations.md b/docs/operations.md index 3bf7ceb..29ca271 100644 --- a/docs/operations.md +++ b/docs/operations.md @@ -116,3 +116,29 @@ Generate a hash: ```powershell .\promtactl.exe token-hash --token "replace-with-secret-token" ``` + +## Dashboard Login + +The dashboard uses a session cookie instead of storing bearer tokens in the +browser. + +Login: + +```http +POST /api/session +Content-Type: application/json + +{"username":"admin","token":"replace-with-secret-token"} +``` + +Check the current session: + +```http +GET /api/session +``` + +Logout: + +```http +DELETE /api/session +``` diff --git a/internal/auth/auth.go b/internal/auth/auth.go index 3365dbb..0c328e7 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -1,11 +1,14 @@ package auth import ( + "crypto/rand" "crypto/sha256" "crypto/subtle" "encoding/hex" "net/http" "strings" + "sync" + "time" ) const ( @@ -27,13 +30,25 @@ type Principal struct { Roles []string `json:"roles"` } +type SessionInfo struct { + Principal Principal `json:"principal"` + ExpiresAt time.Time `json:"expires_at"` +} + type Authenticator struct { users []UserConfig legacyHash string + sessionMu sync.RWMutex + sessions map[string]SessionInfo + sessionTTL time.Duration } func New(users []UserConfig, legacyToken string) *Authenticator { - authenticator := &Authenticator{users: normalizeUsers(users)} + authenticator := &Authenticator{ + users: normalizeUsers(users), + sessions: make(map[string]SessionInfo), + sessionTTL: 12 * time.Hour, + } if legacyToken != "" { authenticator.legacyHash = HashToken(legacyToken) } @@ -54,20 +69,71 @@ func (a *Authenticator) HasUsers() bool { } func (a *Authenticator) Authenticate(r *http.Request) (Principal, bool) { + if principal, ok := a.authenticateSession(r); ok { + return principal, true + } token := readToken(r) if token == "" { return Principal{}, false } tokenHash := HashToken(token) - for _, user := range a.users { - if constantTimeEqual(tokenHash, user.TokenHash) { - return Principal{Name: user.Name, Roles: user.Roles}, true + principal, ok := a.principalForToken(tokenHash) + if ok { + return principal, true + } + return Principal{}, false +} + +func (a *Authenticator) Login(username string, token string) (SessionInfo, string, bool) { + principal, ok := a.principalForCredentials(username, token) + if !ok { + return SessionInfo{}, "", false + } + sessionID := randomSessionID() + if sessionID == "" { + return SessionInfo{}, "", false + } + info := SessionInfo{ + Principal: principal, + ExpiresAt: time.Now().UTC().Add(a.sessionTTL), + } + a.sessionMu.Lock() + a.sessions[sessionID] = info + a.sessionMu.Unlock() + return info, sessionID, true +} + +func (a *Authenticator) Session(r *http.Request) (SessionInfo, bool) { + sessionID := readSessionID(r) + if sessionID == "" { + return SessionInfo{}, false + } + a.sessionMu.RLock() + info, ok := a.sessions[sessionID] + a.sessionMu.RUnlock() + if !ok || time.Now().UTC().After(info.ExpiresAt) { + if ok { + a.sessionMu.Lock() + delete(a.sessions, sessionID) + a.sessionMu.Unlock() } + return SessionInfo{}, false } - if a.legacyHash != "" && constantTimeEqual(tokenHash, a.legacyHash) { - return Principal{Name: "legacy-token", Roles: []string{RoleAdmin}}, true + return info, true +} + +func (a *Authenticator) Logout(r *http.Request) bool { + sessionID := readSessionID(r) + if sessionID == "" { + return false } - return Principal{}, false + a.sessionMu.Lock() + defer a.sessionMu.Unlock() + if _, ok := a.sessions[sessionID]; !ok { + return false + } + delete(a.sessions, sessionID) + return true } func (p Principal) HasAny(roles ...string) bool { @@ -81,6 +147,31 @@ func (p Principal) HasAny(roles ...string) bool { return false } +func (a *Authenticator) SetSessionCookie(w http.ResponseWriter, sessionID string, expiresAt time.Time, secure bool) { + http.SetCookie(w, &http.Cookie{ + Name: sessionCookieName, + Value: sessionID, + Path: "/", + Expires: expiresAt, + MaxAge: int(time.Until(expiresAt).Seconds()), + HttpOnly: true, + Secure: secure, + SameSite: http.SameSiteLaxMode, + }) +} + +func (a *Authenticator) ClearSessionCookie(w http.ResponseWriter) { + http.SetCookie(w, &http.Cookie{ + Name: sessionCookieName, + Value: "", + Path: "/", + Expires: time.Unix(0, 0), + MaxAge: -1, + HttpOnly: true, + SameSite: http.SameSiteLaxMode, + }) +} + func RequiredRoles(method string, path string) []string { if path == "/api/audit" { return []string{RoleAnalyst, RoleOperator} @@ -134,6 +225,34 @@ func normalizeRoles(roles []string) []string { return normalized } +func (a *Authenticator) principalForCredentials(username string, token string) (Principal, bool) { + token = strings.TrimSpace(token) + if token == "" { + return Principal{}, false + } + tokenHash := HashToken(token) + username = strings.TrimSpace(username) + if principal, ok := a.principalForToken(tokenHash); ok { + if username != "" && !strings.EqualFold(username, principal.Name) && principal.Name != "legacy-token" { + return Principal{}, false + } + return principal, true + } + return Principal{}, false +} + +func (a *Authenticator) principalForToken(tokenHash string) (Principal, bool) { + for _, user := range a.users { + if constantTimeEqual(tokenHash, user.TokenHash) { + return Principal{Name: user.Name, Roles: user.Roles}, true + } + } + if a.legacyHash != "" && constantTimeEqual(tokenHash, a.legacyHash) { + return Principal{Name: "legacy-token", Roles: []string{RoleAdmin}}, true + } + return Principal{}, false +} + func readToken(r *http.Request) string { header := r.Header.Get("Authorization") if strings.HasPrefix(strings.ToLower(header), "bearer ") { @@ -142,6 +261,32 @@ func readToken(r *http.Request) string { return strings.TrimSpace(r.Header.Get("X-Promtact-Token")) } +func readSessionID(r *http.Request) string { + cookie, err := r.Cookie(sessionCookieName) + if err != nil { + return "" + } + return strings.TrimSpace(cookie.Value) +} + +func (a *Authenticator) authenticateSession(r *http.Request) (Principal, bool) { + info, ok := a.Session(r) + if !ok { + return Principal{}, false + } + return info.Principal, true +} + +func randomSessionID() string { + var raw [32]byte + if _, err := rand.Read(raw[:]); err != nil { + return "" + } + return hex.EncodeToString(raw[:]) +} + +const sessionCookieName = "promtact_session" + func constantTimeEqual(got string, want string) bool { if got == "" || want == "" { return false diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go index 1498830..02c0756 100644 --- a/internal/auth/auth_test.go +++ b/internal/auth/auth_test.go @@ -20,6 +20,37 @@ func TestAuthenticateUserToken(t *testing.T) { } } +func TestSessionLoginAndAuthenticate(t *testing.T) { + a := New([]UserConfig{{Name: "alice", TokenHash: HashToken("secret"), Roles: []string{RoleOperator}}}, "") + info, sessionID, ok := a.Login("alice", "secret") + if !ok { + t.Fatal("expected login to succeed") + } + if info.Principal.Name != "alice" || info.ExpiresAt.IsZero() { + t.Fatalf("unexpected session info: %#v", info) + } + + req := httptest.NewRequest(http.MethodGet, "/api/status", nil) + req.AddCookie(&http.Cookie{Name: sessionCookieName, Value: sessionID}) + + session, ok := a.Session(req) + if !ok || session.Principal.Name != "alice" { + t.Fatalf("unexpected session lookup: %#v", session) + } + + principal, ok := a.Authenticate(req) + if !ok || principal.Name != "alice" { + t.Fatalf("unexpected authenticated principal: %#v", principal) + } + + if !a.Logout(req) { + t.Fatal("expected logout to remove session") + } + if _, ok := a.Session(req); ok { + t.Fatal("expected session to be removed") + } +} + func TestRequiredRoles(t *testing.T) { if roles := RequiredRoles(http.MethodPost, "/api/responses/approve"); len(roles) != 1 || roles[0] != RoleOperator { t.Fatalf("unexpected approve roles: %#v", roles) diff --git a/internal/server/server.go b/internal/server/server.go index 5a4fb1c..705d4a5 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -94,6 +94,7 @@ func NewWithOptions(options Options) (*App, error) { func (a *App) Routes() http.Handler { mux := http.NewServeMux() mux.HandleFunc("/api/status", a.handleStatus) + mux.HandleFunc("/api/session", a.handleSession) mux.HandleFunc("/api/events", a.handleEvents) mux.HandleFunc("/api/alerts", a.handleAlerts) mux.HandleFunc("/api/assets", a.handleAssets) @@ -138,6 +139,104 @@ func (a *App) handleStatus(w http.ResponseWriter, r *http.Request) { }) } +func (a *App) handleSession(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodGet: + if a.auth == nil || !a.auth.Enabled() { + writeJSON(w, http.StatusOK, map[string]any{ + "authenticated": true, + "mode": "open", + "principal": auth.Principal{ + Name: "anonymous", + Roles: []string{auth.RoleAdmin}, + }, + }) + return + } + if info, ok := a.auth.Session(r); ok { + writeJSON(w, http.StatusOK, map[string]any{ + "authenticated": true, + "mode": "session", + "principal": info.Principal, + "expires_at": info.ExpiresAt, + }) + return + } + if principal, ok := a.auth.Authenticate(r); ok { + writeJSON(w, http.StatusOK, map[string]any{ + "authenticated": true, + "mode": "token", + "principal": principal, + }) + return + } + writeJSON(w, http.StatusOK, map[string]any{ + "authenticated": false, + }) + case http.MethodPost: + if a.auth == nil || !a.auth.Enabled() { + writeJSON(w, http.StatusAccepted, map[string]any{ + "authenticated": true, + "mode": "open", + "principal": auth.Principal{ + Name: "anonymous", + Roles: []string{auth.RoleAdmin}, + }, + }) + return + } + var req struct { + Username string `json:"username"` + Token string `json:"token"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, err) + return + } + info, sessionID, ok := a.auth.Login(req.Username, req.Token) + if !ok { + a.recordAudit(r, auth.Principal{Name: "anonymous"}, "auth.login", "session", "", "denied", map[string]string{ + "username": strings.TrimSpace(req.Username), + }) + writeError(w, http.StatusUnauthorized, errors.New("invalid credentials")) + return + } + a.auth.SetSessionCookie(w, sessionID, info.ExpiresAt, r.TLS != nil) + a.recordAudit(r, info.Principal, "auth.login", "session", "", "accepted", map[string]string{ + "mode": "session", + }) + writeJSON(w, http.StatusAccepted, map[string]any{ + "authenticated": true, + "mode": "session", + "principal": info.Principal, + "expires_at": info.ExpiresAt, + }) + case http.MethodDelete: + if a.auth == nil || !a.auth.Enabled() { + writeJSON(w, http.StatusOK, map[string]any{ + "authenticated": true, + "mode": "open", + "principal": auth.Principal{ + Name: "anonymous", + Roles: []string{auth.RoleAdmin}, + }, + }) + return + } + info, ok := a.auth.Session(r) + if ok { + a.recordAudit(r, info.Principal, "auth.logout", "session", "", "accepted", nil) + _ = a.auth.Logout(r) + } + a.auth.ClearSessionCookie(w) + writeJSON(w, http.StatusOK, map[string]any{ + "authenticated": false, + }) + default: + methodNotAllowed(w) + } +} + func (a *App) handleEvents(w http.ResponseWriter, r *http.Request) { switch r.Method { case http.MethodGet: @@ -474,7 +573,7 @@ func withSecurityHeaders(next http.Handler) http.Handler { func (a *App) withAuth(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if !strings.HasPrefix(r.URL.Path, "/api/") || a.auth == nil || !a.auth.Enabled() { + if !strings.HasPrefix(r.URL.Path, "/api/") || r.URL.Path == "/api/session" || a.auth == nil || !a.auth.Enabled() { next.ServeHTTP(w, r) return } diff --git a/internal/server/server_test.go b/internal/server/server_test.go index b584b55..3bea490 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -119,6 +119,63 @@ func TestRBACBlocksAuditForViewer(t *testing.T) { } } +func TestSessionLoginAndLogout(t *testing.T) { + app, err := NewWithOptions(Options{ + Users: []auth.UserConfig{{Name: "alice", TokenHash: auth.HashToken("secret"), Roles: []string{auth.RoleOperator}}}, + }) + if err != nil { + t.Fatalf("new app: %v", err) + } + + loginReq := httptest.NewRequest(http.MethodPost, "/api/session", strings.NewReader(`{"username":"alice","token":"secret"}`)) + loginRec := httptest.NewRecorder() + app.Routes().ServeHTTP(loginRec, loginReq) + if loginRec.Code != http.StatusAccepted { + t.Fatalf("expected login 202, got %d: %s", loginRec.Code, loginRec.Body.String()) + } + var loginPayload struct { + Authenticated bool `json:"authenticated"` + Mode string `json:"mode"` + Principal struct { + Name string `json:"name"` + } `json:"principal"` + } + if err := json.Unmarshal(loginRec.Body.Bytes(), &loginPayload); err != nil { + t.Fatalf("decode login: %v", err) + } + if !loginPayload.Authenticated || loginPayload.Mode != "session" || loginPayload.Principal.Name != "alice" { + t.Fatalf("unexpected login payload: %#v", loginPayload) + } + cookies := loginRec.Result().Cookies() + if len(cookies) == 0 { + t.Fatal("expected session cookie") + } + + statusReq := httptest.NewRequest(http.MethodGet, "/api/status", nil) + statusReq.AddCookie(cookies[0]) + statusRec := httptest.NewRecorder() + app.Routes().ServeHTTP(statusRec, statusReq) + if statusRec.Code != http.StatusOK { + t.Fatalf("expected authenticated status 200, got %d", statusRec.Code) + } + + logoutReq := httptest.NewRequest(http.MethodDelete, "/api/session", nil) + logoutReq.AddCookie(cookies[0]) + logoutRec := httptest.NewRecorder() + app.Routes().ServeHTTP(logoutRec, logoutReq) + if logoutRec.Code != http.StatusOK { + t.Fatalf("expected logout 200, got %d", logoutRec.Code) + } + + statusReq = httptest.NewRequest(http.MethodGet, "/api/status", nil) + statusReq.AddCookie(cookies[0]) + statusRec = httptest.NewRecorder() + app.Routes().ServeHTTP(statusRec, statusReq) + if statusRec.Code != http.StatusUnauthorized { + t.Fatalf("expected 401 after logout, got %d", statusRec.Code) + } +} + func TestEmptyListEndpointsReturnArrays(t *testing.T) { app, err := NewWithOptions(Options{}) if err != nil { diff --git a/web/app.js b/web/app.js index f437220..818d36a 100644 --- a/web/app.js +++ b/web/app.js @@ -1,7 +1,16 @@ const els = { + loginView: document.querySelector("#login-view"), + appView: document.querySelector("#app-view"), + loginForm: document.querySelector("#login-form"), + loginUsername: document.querySelector("#login-username"), + loginToken: document.querySelector("#login-token"), + loginError: document.querySelector("#login-error"), + loginSubmit: document.querySelector("#login-submit"), + sessionLabel: document.querySelector("#session-label"), version: document.querySelector("#version"), loadDemo: document.querySelector("#load-demo"), refresh: document.querySelector("#refresh"), + logout: document.querySelector("#logout"), metrics: { events: document.querySelector("#metric-events"), alerts: document.querySelector("#metric-alerts"), @@ -18,37 +27,119 @@ const els = { }; const emptyTemplate = document.querySelector("#empty-template"); +const state = { + session: null, + pollHandle: null +}; async function api(path, options = {}) { - const headers = { "Content-Type": "application/json" }; - const token = sessionStorage.getItem("promtact_api_token"); + const { useToken = true, ...fetchOptions } = options; + const headers = { + "Content-Type": "application/json", + ...(fetchOptions.headers || {}) + }; + const token = useToken ? sessionStorage.getItem("promtact_api_token") : ""; if (token) { headers.Authorization = `Bearer ${token}`; } const response = await fetch(path, { - headers, - ...options + credentials: "same-origin", + ...fetchOptions, + headers }); - if (response.status === 401 && isWrite(options.method)) { - const nextToken = window.prompt("API token"); - if (nextToken) { - sessionStorage.setItem("promtact_api_token", nextToken); - return api(path, options); - } - } + const bodyText = await response.text(); + const payload = bodyText ? tryParseJSON(bodyText) : null; if (!response.ok) { - const payload = await response.json().catch(() => ({})); - throw new Error(payload.error || `${response.status} ${response.statusText}`); + const error = new Error((payload && payload.error) || `${response.status} ${response.statusText}`); + error.status = response.status; + error.payload = payload; + throw error; + } + return payload; +} + +function tryParseJSON(value) { + try { + return JSON.parse(value); + } catch { + return value; + } +} + +function isLoggedIn() { + return Boolean(state.session && state.session.authenticated); +} + +function setView(session) { + state.session = session; + const loggedIn = isLoggedIn(); + document.body.classList.toggle("logged-out", !loggedIn); + els.loginView.hidden = loggedIn; + els.appView.hidden = !loggedIn; + if (loggedIn) { + const principal = session.principal || {}; + const mode = session.mode || "session"; + const roles = Array.isArray(principal.roles) ? principal.roles.join(", ") : ""; + els.sessionLabel.textContent = roles ? `${principal.name} - ${mode} - ${roles}` : `${principal.name} - ${mode}`; + } else { + els.sessionLabel.textContent = ""; + } +} + +function showLogin(message = "") { + stopPolling(); + setView({ authenticated: false }); + els.loginError.textContent = message; + els.loginSubmit.disabled = false; + els.loginForm.reset(); + els.loginUsername.focus(); +} + +function showApp(session) { + els.loginError.textContent = ""; + setView(session); + startPolling(); +} + +function startPolling() { + stopPolling(); + state.pollHandle = window.setInterval(() => { + refresh().catch(handleApiFailure); + }, 8000); +} + +function stopPolling() { + if (state.pollHandle !== null) { + window.clearInterval(state.pollHandle); + state.pollHandle = null; + } +} + +function handleApiFailure(error) { + if (error && error.status === 401) { + sessionStorage.removeItem("promtact_api_token"); + showLogin("Session expired."); + return; } - return response.json(); + console.error(error); } -function isWrite(method = "GET") { - return !["GET", "HEAD", "OPTIONS"].includes(method.toUpperCase()); +async function loadSession() { + const session = await api("/api/session"); + if (session && session.authenticated) { + showApp(session); + return true; + } + showLogin(); + return false; } async function refresh() { + if (!isLoggedIn()) { + return; + } + const [status, alerts, assets, events, actions, rules] = await Promise.all([ api("/api/status"), api("/api/alerts"), @@ -221,10 +312,11 @@ function renderGraph(assets, alerts) { .map((node) => ``) .join(""); - const circles = nodes.map((node) => { - const name = escapeHtml(node.asset.hostname || node.asset.id); - const risk = Math.min(99, node.asset.risk_score || 0); - return ` + const circles = nodes + .map((node) => { + const name = escapeHtml(node.asset.hostname || node.asset.id); + const risk = Math.min(99, node.asset.risk_score || 0); + return ` ${risk} @@ -232,7 +324,8 @@ function renderGraph(assets, alerts) { ${node.alerts} alerts `; - }).join(""); + }) + .join(""); els.graph.innerHTML = ` @@ -269,17 +362,67 @@ function escapeHtml(value) { .replaceAll("'", "'"); } +async function login(username, token) { + const session = await api("/api/session", { + method: "POST", + useToken: false, + body: JSON.stringify({ username, token }) + }); + sessionStorage.removeItem("promtact_api_token"); + showApp(session); + await refresh(); +} + +async function logout() { + try { + await api("/api/session", { + method: "DELETE", + useToken: false + }); + } catch (error) { + handleApiFailure(error); + } + sessionStorage.removeItem("promtact_api_token"); + showLogin(); +} + +els.loginForm.addEventListener("submit", async (event) => { + event.preventDefault(); + els.loginSubmit.disabled = true; + els.loginError.textContent = ""; + try { + await login(els.loginUsername.value, els.loginToken.value); + } catch (error) { + if (error && error.status === 401) { + els.loginError.textContent = "Invalid credentials."; + return; + } + handleApiFailure(error); + els.loginError.textContent = "Login failed."; + } finally { + els.loginSubmit.disabled = false; + } +}); + +els.logout.addEventListener("click", () => { + logout().catch(handleApiFailure); +}); + els.loadDemo.addEventListener("click", async () => { els.loadDemo.disabled = true; try { await api("/api/demo", { method: "POST", body: "{}" }); await refresh(); + } catch (error) { + handleApiFailure(error); } finally { els.loadDemo.disabled = false; } }); -els.refresh.addEventListener("click", refresh); +els.refresh.addEventListener("click", () => { + refresh().catch(handleApiFailure); +}); els.alertsList.addEventListener("click", async (event) => { const button = event.target.closest("[data-respond]"); @@ -291,6 +434,8 @@ els.alertsList.addEventListener("click", async (event) => { body: JSON.stringify({ alert_id: button.dataset.respond }) }); await refresh(); + } catch (error) { + handleApiFailure(error); } finally { button.disabled = false; } @@ -306,15 +451,23 @@ els.actionsList.addEventListener("click", async (event) => { body: JSON.stringify({ action_id: button.dataset.approve, approved_by: "dashboard" }) }); await refresh(); + } catch (error) { + handleApiFailure(error); } finally { button.disabled = false; } }); -refresh().catch((error) => { - console.error(error); +els.loginUsername.addEventListener("input", () => { + els.loginError.textContent = ""; }); -setInterval(() => { - refresh().catch((error) => console.error(error)); -}, 8000); +els.loginToken.addEventListener("input", () => { + els.loginError.textContent = ""; +}); + +loadSession().then((loggedIn) => { + if (loggedIn) { + refresh().catch(handleApiFailure); + } +}).catch(handleApiFailure); diff --git a/web/index.html b/web/index.html index 6d02b44..6b0c438 100644 --- a/web/index.html +++ b/web/index.html @@ -7,107 +7,129 @@ -
-
+
-
- - -
-
+

Sign in

+ + + + +

+ +
-
-
-
- Events - 0 + + -
-
-
-
-

Asset Risk Graph

- -
- -
+
+
+
+ Events + 0 +
+
+ Alerts + 0 +
+
+ Assets + 0 +
+
+ Dry Runs + 0 +
+
+ Audit + 0 +
+
-
-
-

Assets

- risk-ranked -
-
- - - - - - - - - - -
AssetIPAgent SurfaceRisk
-
-
+
+
+
+
+

Asset Risk Graph

+ +
+ +
-
-
-

Event Stream

- latest first -
-
-
-
+
+
+

Assets

+ risk-ranked +
+
+ + + + + + + + + + +
AssetIPAgent SurfaceRisk
+
+
+ +
+
+

Event Stream

+ latest first +
+
+
+
-
-
+
+
+

Rules

+ enabled +
+
+
+ + + +