This commit is contained in:
@@ -0,0 +1,355 @@
|
||||
package rules
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"example.com/siem-greenfield/internal/clickhouse"
|
||||
)
|
||||
|
||||
type RuleSet struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Version int `json:"version"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Rules []Rule `json:"rules"`
|
||||
}
|
||||
|
||||
type Rule struct {
|
||||
ID string `json:"id"`
|
||||
Title string `json:"title"`
|
||||
Description string `json:"description"`
|
||||
Severity string `json:"severity"`
|
||||
Score float64 `json:"score"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Kind string `json:"kind"`
|
||||
Channels []string `json:"channels,omitempty"`
|
||||
EventCodes []uint32 `json:"event_codes,omitempty"`
|
||||
Conditions []Condition `json:"conditions,omitempty"`
|
||||
GroupBy []string `json:"group_by,omitempty"`
|
||||
Threshold int64 `json:"threshold,omitempty"`
|
||||
DistinctField string `json:"distinct_field,omitempty"`
|
||||
DistinctThreshold int64 `json:"distinct_threshold,omitempty"`
|
||||
WindowSeconds int `json:"window_seconds"`
|
||||
SuppressSeconds int `json:"suppress_seconds,omitempty"`
|
||||
Summary string `json:"summary"`
|
||||
Tags []string `json:"tags,omitempty"`
|
||||
MITRE []string `json:"mitre,omitempty"`
|
||||
}
|
||||
|
||||
type Condition struct {
|
||||
Field string `json:"field"`
|
||||
Operator string `json:"operator"`
|
||||
Value string `json:"value,omitempty"`
|
||||
}
|
||||
|
||||
type StoredRule struct {
|
||||
RuleSetID string `json:"rule_set_id"`
|
||||
RuleSetName string `json:"rule_set_name"`
|
||||
RuleSetOn bool `json:"rule_set_enabled"`
|
||||
Source string `json:"source"`
|
||||
Rule
|
||||
}
|
||||
|
||||
func LoadDir(dir string) ([]RuleSet, error) {
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out []RuleSet
|
||||
for _, ent := range entries {
|
||||
if ent.IsDir() || !strings.HasSuffix(strings.ToLower(ent.Name()), ".json") {
|
||||
continue
|
||||
}
|
||||
b, err := os.ReadFile(filepath.Join(dir, ent.Name()))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var rs RuleSet
|
||||
if err := json.Unmarshal(b, &rs); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", ent.Name(), err)
|
||||
}
|
||||
if err := ValidateSet(rs); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", ent.Name(), err)
|
||||
}
|
||||
out = append(out, rs)
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].ID < out[j].ID })
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func ValidateSet(rs RuleSet) error {
|
||||
if !validID(rs.ID) || strings.TrimSpace(rs.Name) == "" {
|
||||
return fmt.Errorf("invalid rule-set id/name")
|
||||
}
|
||||
if rs.Version <= 0 {
|
||||
return fmt.Errorf("rule-set %s: version must be positive", rs.ID)
|
||||
}
|
||||
seen := map[string]bool{}
|
||||
for i := range rs.Rules {
|
||||
if err := Validate(rs.Rules[i]); err != nil {
|
||||
return fmt.Errorf("rule %d: %w", i+1, err)
|
||||
}
|
||||
if seen[rs.Rules[i].ID] {
|
||||
return fmt.Errorf("duplicate rule id %q", rs.Rules[i].ID)
|
||||
}
|
||||
seen[rs.Rules[i].ID] = true
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func Validate(r Rule) error {
|
||||
if !validID(r.ID) || strings.TrimSpace(r.Title) == "" {
|
||||
return fmt.Errorf("invalid id/title")
|
||||
}
|
||||
switch r.Severity {
|
||||
case "info", "low", "medium", "high", "critical":
|
||||
default:
|
||||
return fmt.Errorf("invalid severity %q", r.Severity)
|
||||
}
|
||||
if r.Score < 0 || r.Score > 100 {
|
||||
return fmt.Errorf("score must be between 0 and 100")
|
||||
}
|
||||
switch r.Kind {
|
||||
case "event", "threshold", "distinct":
|
||||
default:
|
||||
return fmt.Errorf("invalid kind %q", r.Kind)
|
||||
}
|
||||
if len(r.EventCodes) == 0 && len(r.Conditions) == 0 {
|
||||
return fmt.Errorf("rule must constrain event_codes or conditions")
|
||||
}
|
||||
if r.WindowSeconds < 30 || r.WindowSeconds > 86400 {
|
||||
return fmt.Errorf("window_seconds must be between 30 and 86400")
|
||||
}
|
||||
if r.Threshold <= 0 {
|
||||
r.Threshold = 1
|
||||
}
|
||||
if r.Kind == "distinct" {
|
||||
if _, ok := fieldExpr(r.DistinctField); !ok {
|
||||
return fmt.Errorf("invalid distinct_field %q", r.DistinctField)
|
||||
}
|
||||
if r.DistinctThreshold <= 0 {
|
||||
return fmt.Errorf("distinct_threshold must be positive")
|
||||
}
|
||||
}
|
||||
for _, g := range r.GroupBy {
|
||||
if _, ok := groupExpr(g); !ok {
|
||||
return fmt.Errorf("invalid group_by field %q", g)
|
||||
}
|
||||
}
|
||||
for _, c := range r.Conditions {
|
||||
if _, ok := fieldExpr(c.Field); !ok {
|
||||
return fmt.Errorf("invalid condition field %q", c.Field)
|
||||
}
|
||||
switch c.Operator {
|
||||
case "equals", "not_equals", "contains", "not_contains", "regex", "exists", "not_exists", "in":
|
||||
default:
|
||||
return fmt.Errorf("invalid operator %q", c.Operator)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func Compile(r Rule, db, tenant string, end time.Time) (string, time.Time, error) {
|
||||
if err := Validate(r); err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
start := end.Add(-time.Duration(r.WindowSeconds) * time.Second)
|
||||
table := clickhouse.Ident(db) + ".events"
|
||||
where := []string{
|
||||
"tenant_id=" + clickhouse.Q(tenant),
|
||||
"event_time>=" + clickhouse.Q(ts(start)),
|
||||
"event_time<" + clickhouse.Q(ts(end)),
|
||||
}
|
||||
if len(r.Channels) > 0 {
|
||||
where = append(where, "channel IN ("+quoteStrings(r.Channels)+")")
|
||||
}
|
||||
if len(r.EventCodes) > 0 {
|
||||
vals := make([]string, 0, len(r.EventCodes))
|
||||
for _, id := range r.EventCodes {
|
||||
vals = append(vals, strconv.FormatUint(uint64(id), 10))
|
||||
}
|
||||
where = append(where, "event_code IN ("+strings.Join(vals, ",")+")")
|
||||
}
|
||||
for _, c := range r.Conditions {
|
||||
x, err := compileCondition(c)
|
||||
if err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
where = append(where, x)
|
||||
}
|
||||
|
||||
groupFields := make([]string, 0, len(r.GroupBy))
|
||||
selectFields := make([]string, 0, len(r.GroupBy)+8)
|
||||
selected := map[string]bool{}
|
||||
for _, g := range r.GroupBy {
|
||||
expr, _ := groupExpr(g)
|
||||
alias := groupAlias(g)
|
||||
selectFields = append(selectFields, expr+" AS "+alias)
|
||||
groupFields = append(groupFields, expr)
|
||||
selected[g] = true
|
||||
}
|
||||
addContext := func(key, alias string) {
|
||||
if !selected[key] {
|
||||
selectFields = append(selectFields, "'' AS "+alias)
|
||||
}
|
||||
}
|
||||
addContext("host", "host_name")
|
||||
addContext("user", "user_name")
|
||||
addContext("source_ip", "source_ip")
|
||||
addContext("workstation", "workstation")
|
||||
addContext("process_path", "process_path")
|
||||
if selected["event_code"] {
|
||||
// already selected under event_code
|
||||
} else if len(r.EventCodes) == 1 {
|
||||
selectFields = append(selectFields, strconv.FormatUint(uint64(r.EventCodes[0]), 10)+" AS event_code")
|
||||
} else {
|
||||
selectFields = append(selectFields, "toUInt32(0) AS event_code")
|
||||
}
|
||||
selectFields = append(selectFields,
|
||||
"uniqExact(event_uid) AS cnt",
|
||||
"min(event_time) AS window_start",
|
||||
"max(event_time) AS window_end",
|
||||
)
|
||||
if r.Kind == "distinct" {
|
||||
expr, _ := fieldExpr(r.DistinctField)
|
||||
selectFields = append(selectFields, "uniqExact("+expr+") AS distinct_cnt")
|
||||
}
|
||||
having := fmt.Sprintf("cnt >= %d", max64(1, r.Threshold))
|
||||
if r.Kind == "distinct" {
|
||||
having += fmt.Sprintf(" AND distinct_cnt >= %d", r.DistinctThreshold)
|
||||
}
|
||||
q := "SELECT " + strings.Join(selectFields, ", ") + " FROM " + table + " WHERE " + strings.Join(where, " AND ")
|
||||
if len(groupFields) > 0 {
|
||||
q += " GROUP BY " + strings.Join(groupFields, ", ")
|
||||
}
|
||||
q += " HAVING " + having
|
||||
return q, start, nil
|
||||
}
|
||||
|
||||
func RenderSummary(tpl string, row map[string]any) string {
|
||||
vals := map[string]string{
|
||||
"host": value(row["host_name"]),
|
||||
"user": value(row["user_name"]),
|
||||
"source_ip": value(row["source_ip"]),
|
||||
"workstation": value(row["workstation"]),
|
||||
"process": value(row["process_path"]),
|
||||
"event_code": value(row["event_code"]),
|
||||
"count": value(row["cnt"]),
|
||||
"distinct": value(row["distinct_cnt"]),
|
||||
}
|
||||
out := tpl
|
||||
for k, v := range vals {
|
||||
out = strings.ReplaceAll(out, "{"+k+"}", v)
|
||||
}
|
||||
if strings.TrimSpace(out) == "" {
|
||||
out = "Rule matched on " + vals["host"]
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func fieldExpr(name string) (string, bool) {
|
||||
fields := map[string]string{
|
||||
"host": "host_name", "host_name": "host_name",
|
||||
"user": "multiIf(target_user!='',target_user,user_name!='',user_name,subject_user)",
|
||||
"user_name": "user_name", "target_user": "target_user", "subject_user": "subject_user",
|
||||
"source_ip": "source_ip", "destination_ip": "destination_ip", "workstation": "workstation",
|
||||
"process_path": "process_path", "parent_process_path": "parent_process_path", "command_line": "command_line",
|
||||
"message": "message", "channel": "channel", "provider": "provider", "category": "category",
|
||||
"action": "action", "outcome": "outcome", "logon_type": "logon_type",
|
||||
"authentication_package": "authentication_package", "status_code": "status_code",
|
||||
"failure_reason": "failure_reason", "event_code": "event_code",
|
||||
}
|
||||
v, ok := fields[name]
|
||||
return v, ok
|
||||
}
|
||||
|
||||
func groupExpr(name string) (string, bool) { return fieldExpr(name) }
|
||||
func groupAlias(name string) string {
|
||||
switch name {
|
||||
case "host", "host_name":
|
||||
return "host_name"
|
||||
case "user", "user_name", "target_user", "subject_user":
|
||||
return "user_name"
|
||||
default:
|
||||
return name
|
||||
}
|
||||
}
|
||||
|
||||
func compileCondition(c Condition) (string, error) {
|
||||
expr, ok := fieldExpr(c.Field)
|
||||
if !ok {
|
||||
return "", fmt.Errorf("invalid field %q", c.Field)
|
||||
}
|
||||
q := clickhouse.Q(c.Value)
|
||||
switch c.Operator {
|
||||
case "equals":
|
||||
return expr + "=" + q, nil
|
||||
case "not_equals":
|
||||
return expr + "!=" + q, nil
|
||||
case "contains":
|
||||
return "positionCaseInsensitiveUTF8(" + expr + "," + q + ")>0", nil
|
||||
case "not_contains":
|
||||
return "positionCaseInsensitiveUTF8(" + expr + "," + q + ")=0", nil
|
||||
case "regex":
|
||||
return "match(" + expr + "," + q + ")", nil
|
||||
case "exists":
|
||||
return expr + "!=''", nil
|
||||
case "not_exists":
|
||||
return expr + "=''", nil
|
||||
case "in":
|
||||
parts := strings.Split(c.Value, ",")
|
||||
vals := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
if x := strings.TrimSpace(p); x != "" {
|
||||
vals = append(vals, clickhouse.Q(x))
|
||||
}
|
||||
}
|
||||
if len(vals) == 0 {
|
||||
return "", fmt.Errorf("empty in condition")
|
||||
}
|
||||
return expr + " IN (" + strings.Join(vals, ",") + ")", nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported operator %q", c.Operator)
|
||||
}
|
||||
}
|
||||
|
||||
func quoteStrings(v []string) string {
|
||||
out := make([]string, 0, len(v))
|
||||
for _, x := range v {
|
||||
out = append(out, clickhouse.Q(x))
|
||||
}
|
||||
return strings.Join(out, ",")
|
||||
}
|
||||
func validID(s string) bool {
|
||||
if s == "" || len(s) > 128 {
|
||||
return false
|
||||
}
|
||||
for _, r := range s {
|
||||
if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '-' || r == '_' || r == '.' {
|
||||
continue
|
||||
}
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
func ts(t time.Time) string { return t.UTC().Format("2006-01-02 15:04:05.000") }
|
||||
func value(v any) string {
|
||||
if v == nil {
|
||||
return ""
|
||||
}
|
||||
return fmt.Sprint(v)
|
||||
}
|
||||
func max64(a, b int64) int64 {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
package rules
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCompileThreshold(t *testing.T) {
|
||||
r := Rule{ID: "failed", Title: "Failed", Severity: "high", Kind: "threshold", Enabled: true, EventCodes: []uint32{4625}, GroupBy: []string{"host", "user", "source_ip"}, Threshold: 20, WindowSeconds: 300, Summary: "{count} failures for {user}"}
|
||||
q, _, err := Compile(r, "siem", "default", time.Date(2026, 7, 23, 12, 0, 0, 0, time.UTC))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, want := range []string{"event_code IN (4625)", "GROUP BY host_name", "HAVING cnt >= 20"} {
|
||||
if !strings.Contains(q, want) {
|
||||
t.Fatalf("query missing %q: %s", want, q)
|
||||
}
|
||||
}
|
||||
}
|
||||
func TestRejectUnknownField(t *testing.T) {
|
||||
r := Rule{ID: "bad", Title: "Bad", Severity: "high", Kind: "event", EventCodes: []uint32{1}, Conditions: []Condition{{Field: "DROP TABLE", Operator: "equals", Value: "x"}}, WindowSeconds: 300}
|
||||
if err := Validate(r); err == nil {
|
||||
t.Fatal("expected validation error")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user