a6b480a3
feat: add per-host rule exemptions to scanner
a73x 2026-03-31 06:00
Commit message
proxy/proxy.go
| Old | New | ||
|---|---|---|---|
| @@ -214,7 +214,7 @@ func (p *Proxy) handleConnectMITM(w http.ResponseWriter, r *http.Request) { | |||
| 214 | return | 214 | return |
| 215 | } | 215 | } |
| 216 | 216 | ||
| 217 | findings := p.scanRequest(req) | 217 | findings := p.scanRequest(req, host) |
| 218 | if len(findings) > 0 { | 218 | if len(findings) > 0 { |
| 219 | rules := make([]string, 0, len(findings)) | 219 | rules := make([]string, 0, len(findings)) |
| 220 | for _, f := range findings { | 220 | for _, f := range findings { |
| @@ -294,7 +294,7 @@ func (p *Proxy) getOrCreateLeaf(host string) (*tls.Certificate, error) { | |||
| 294 | return &leaf, nil | 294 | return &leaf, nil |
| 295 | } | 295 | } |
| 296 | 296 | ||
| 297 | func (p *Proxy) scanRequest(r *http.Request) []scanner.Finding { | 297 | func (p *Proxy) scanRequest(r *http.Request, host string) []scanner.Finding { |
| 298 | if p.scanner == nil { | 298 | if p.scanner == nil { |
| 299 | return nil | 299 | return nil |
| 300 | } | 300 | } |
| @@ -319,11 +319,12 @@ func (p *Proxy) scanRequest(r *http.Request) []scanner.Finding { | |||
| 319 | r.Body = io.NopCloser(bytes.NewReader(body)) | 319 | r.Body = io.NopCloser(bytes.NewReader(body)) |
| 320 | } | 320 | } |
| 321 | 321 | ||
| 322 | return p.scanner.Scan(buf.Bytes()) | 322 | return p.scanner.Scan(buf.Bytes(), host) |
| 323 | } | 323 | } |
| 324 | 324 | ||
| 325 | func (p *Proxy) handleHTTP(w http.ResponseWriter, r *http.Request) { | 325 | func (p *Proxy) handleHTTP(w http.ResponseWriter, r *http.Request) { |
| 326 | findings := p.scanRequest(r) | 326 | host := extractHost(r.Host) |
| 327 | findings := p.scanRequest(r, host) | ||
| 327 | if len(findings) > 0 { | 328 | if len(findings) > 0 { |
| 328 | rules := make([]string, 0, len(findings)) | 329 | rules := make([]string, 0, len(findings)) |
| 329 | for _, f := range findings { | 330 | for _, f := range findings { |
scanner/scanner.go
| Old | New | ||
|---|---|---|---|
| @@ -15,8 +15,9 @@ type Finding struct { | |||
| 15 | } | 15 | } |
| 16 | 16 | ||
| 17 | type rule struct { | 17 | type rule struct { |
| 18 | name string | 18 | name string |
| 19 | pattern *regexp.Regexp | 19 | pattern *regexp.Regexp |
| 20 | exemptHosts []string | ||
| 20 | } | 21 | } |
| 21 | 22 | ||
| 22 | // Scanner holds compiled rules for scanning request bodies. | 23 | // Scanner holds compiled rules for scanning request bodies. |
| @@ -25,8 +26,9 @@ type Scanner struct { | |||
| 25 | } | 26 | } |
| 26 | 27 | ||
| 27 | type yamlRule struct { | 28 | type yamlRule struct { |
| 28 | Name string `yaml:"name"` | 29 | Name string `yaml:"name"` |
| 29 | Pattern string `yaml:"pattern"` | 30 | Pattern string `yaml:"pattern"` |
| 31 | ExemptHosts []string `yaml:"exempt_hosts"` | ||
| 30 | } | 32 | } |
| 31 | 33 | ||
| 32 | type yamlConfig struct { | 34 | type yamlConfig struct { |
| @@ -52,7 +54,7 @@ func New(path string) (*Scanner, error) { | |||
| 52 | if err != nil { | 54 | if err != nil { |
| 53 | return nil, fmt.Errorf("compiling pattern for rule %q: %w", yr.Name, err) | 55 | return nil, fmt.Errorf("compiling pattern for rule %q: %w", yr.Name, err) |
| 54 | } | 56 | } |
| 55 | rules = append(rules, rule{name: yr.Name, pattern: re}) | 57 | rules = append(rules, rule{name: yr.Name, pattern: re, exemptHosts: yr.ExemptHosts}) |
| 56 | } | 58 | } |
| 57 | 59 | ||
| 58 | return &Scanner{rules: rules}, nil | 60 | return &Scanner{rules: rules}, nil |
| @@ -65,9 +67,13 @@ func (s *Scanner) RuleCount() int { | |||
| 65 | 67 | ||
| 66 | // Scan checks body against all rules and returns any findings. | 68 | // Scan checks body against all rules and returns any findings. |
| 67 | // Match snippets are truncated to 40 characters. | 69 | // Match snippets are truncated to 40 characters. |
| 68 | func (s *Scanner) Scan(body []byte) []Finding { | 70 | // Rules with exempt_hosts are skipped when host matches. |
| 71 | func (s *Scanner) Scan(body []byte, host string) []Finding { | ||
| 69 | var findings []Finding | 72 | var findings []Finding |
| 70 | for _, r := range s.rules { | 73 | for _, r := range s.rules { |
| 74 | if r.isExempt(host) { | ||
| 75 | continue | ||
| 76 | } | ||
| 71 | match := r.pattern.Find(body) | 77 | match := r.pattern.Find(body) |
| 72 | if match == nil { | 78 | if match == nil { |
| 73 | continue | 79 | continue |
| @@ -81,6 +87,15 @@ func (s *Scanner) Scan(body []byte) []Finding { | |||
| 81 | return findings | 87 | return findings |
| 82 | } | 88 | } |
| 83 | 89 | ||
| 90 | func (r *rule) isExempt(host string) bool { | ||
| 91 | for _, h := range r.exemptHosts { | ||
| 92 | if h == host { | ||
| 93 | return true | ||
| 94 | } | ||
| 95 | } | ||
| 96 | return false | ||
| 97 | } | ||
| 98 | |||
| 84 | const defaultRulesYAML = `rules: | 99 | const defaultRulesYAML = `rules: |
| 85 | - name: ssh-private-key | 100 | - name: ssh-private-key |
| 86 | pattern: "-----BEGIN (OPENSSH|RSA|DSA|EC|ED25519) PRIVATE KEY-----" | 101 | pattern: "-----BEGIN (OPENSSH|RSA|DSA|EC|ED25519) PRIVATE KEY-----" |
| @@ -90,6 +105,8 @@ const defaultRulesYAML = `rules: | |||
| 90 | pattern: "Authorization:\\s*Basic\\s+" | 105 | pattern: "Authorization:\\s*Basic\\s+" |
| 91 | - name: bearer-token | 106 | - name: bearer-token |
| 92 | pattern: "Authorization:\\s*Bearer\\s+" | 107 | pattern: "Authorization:\\s*Bearer\\s+" |
| 108 | exempt_hosts: | ||
| 109 | - api.anthropic.com | ||
| 93 | - name: aws-access-key | 110 | - name: aws-access-key |
| 94 | pattern: "AKIA[0-9A-Z]{16}" | 111 | pattern: "AKIA[0-9A-Z]{16}" |
| 95 | - name: github-token | 112 | - name: github-token |
scanner/scanner_test.go
| Old | New | ||
|---|---|---|---|
| @@ -48,7 +48,7 @@ func writeRules(t *testing.T, content string) string { | |||
| 48 | func TestScan_DetectsSSHPrivateKey(t *testing.T) { | 48 | func TestScan_DetectsSSHPrivateKey(t *testing.T) { |
| 49 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN (OPENSSH|RSA|DSA|EC|ED25519) PRIVATE KEY-----\"\n") | 49 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN (OPENSSH|RSA|DSA|EC|ED25519) PRIVATE KEY-----\"\n") |
| 50 | s, _ := scanner.New(path) | 50 | s, _ := scanner.New(path) |
| 51 | findings := s.Scan([]byte("some data\n-----BEGIN RSA PRIVATE KEY-----\nMIIE...")) | 51 | findings := s.Scan([]byte("some data\n-----BEGIN RSA PRIVATE KEY-----\nMIIE..."), "") |
| 52 | if len(findings) != 1 { | 52 | if len(findings) != 1 { |
| 53 | t.Fatalf("expected 1 finding, got %d", len(findings)) | 53 | t.Fatalf("expected 1 finding, got %d", len(findings)) |
| 54 | } | 54 | } |
| @@ -60,7 +60,7 @@ func TestScan_DetectsSSHPrivateKey(t *testing.T) { | |||
| 60 | func TestScan_DetectsAWSKey(t *testing.T) { | 60 | func TestScan_DetectsAWSKey(t *testing.T) { |
| 61 | path := writeRules(t, "rules:\n - name: aws-access-key\n pattern: \"AKIA[0-9A-Z]{16}\"\n") | 61 | path := writeRules(t, "rules:\n - name: aws-access-key\n pattern: \"AKIA[0-9A-Z]{16}\"\n") |
| 62 | s, _ := scanner.New(path) | 62 | s, _ := scanner.New(path) |
| 63 | findings := s.Scan([]byte("{\"key\": \"AKIAIOSFODNN7EXAMPLE\"}")) | 63 | findings := s.Scan([]byte("{\"key\": \"AKIAIOSFODNN7EXAMPLE\"}"), "") |
| 64 | if len(findings) != 1 { | 64 | if len(findings) != 1 { |
| 65 | t.Fatalf("expected 1 finding, got %d", len(findings)) | 65 | t.Fatalf("expected 1 finding, got %d", len(findings)) |
| 66 | } | 66 | } |
| @@ -73,7 +73,7 @@ func TestScan_ReturnsMultipleFindings(t *testing.T) { | |||
| 73 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN RSA PRIVATE KEY-----\"\n - name: aws-access-key\n pattern: \"AKIA[0-9A-Z]{16}\"\n") | 73 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN RSA PRIVATE KEY-----\"\n - name: aws-access-key\n pattern: \"AKIA[0-9A-Z]{16}\"\n") |
| 74 | s, _ := scanner.New(path) | 74 | s, _ := scanner.New(path) |
| 75 | body := []byte("-----BEGIN RSA PRIVATE KEY-----\nkey\nAKIAIOSFODNN7EXAMPLE") | 75 | body := []byte("-----BEGIN RSA PRIVATE KEY-----\nkey\nAKIAIOSFODNN7EXAMPLE") |
| 76 | findings := s.Scan(body) | 76 | findings := s.Scan(body, "") |
| 77 | if len(findings) != 2 { | 77 | if len(findings) != 2 { |
| 78 | t.Fatalf("expected 2 findings, got %d", len(findings)) | 78 | t.Fatalf("expected 2 findings, got %d", len(findings)) |
| 79 | } | 79 | } |
| @@ -82,7 +82,7 @@ func TestScan_ReturnsMultipleFindings(t *testing.T) { | |||
| 82 | func TestScan_ReturnsEmptyForCleanBody(t *testing.T) { | 82 | func TestScan_ReturnsEmptyForCleanBody(t *testing.T) { |
| 83 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN RSA PRIVATE KEY-----\"\n") | 83 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN RSA PRIVATE KEY-----\"\n") |
| 84 | s, _ := scanner.New(path) | 84 | s, _ := scanner.New(path) |
| 85 | findings := s.Scan([]byte("just some normal POST data")) | 85 | findings := s.Scan([]byte("just some normal POST data"), "") |
| 86 | if len(findings) != 0 { | 86 | if len(findings) != 0 { |
| 87 | t.Errorf("expected 0 findings, got %d", len(findings)) | 87 | t.Errorf("expected 0 findings, got %d", len(findings)) |
| 88 | } | 88 | } |
| @@ -91,7 +91,7 @@ func TestScan_ReturnsEmptyForCleanBody(t *testing.T) { | |||
| 91 | func TestScan_TruncatesMatchSnippet(t *testing.T) { | 91 | func TestScan_TruncatesMatchSnippet(t *testing.T) { |
| 92 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN RSA PRIVATE KEY-----\"\n") | 92 | path := writeRules(t, "rules:\n - name: ssh-private-key\n pattern: \"-----BEGIN RSA PRIVATE KEY-----\"\n") |
| 93 | s, _ := scanner.New(path) | 93 | s, _ := scanner.New(path) |
| 94 | findings := s.Scan([]byte("-----BEGIN RSA PRIVATE KEY-----")) | 94 | findings := s.Scan([]byte("-----BEGIN RSA PRIVATE KEY-----"), "") |
| 95 | if len(findings) != 1 { | 95 | if len(findings) != 1 { |
| 96 | t.Fatalf("expected 1 finding, got %d", len(findings)) | 96 | t.Fatalf("expected 1 finding, got %d", len(findings)) |
| 97 | } | 97 | } |
| @@ -131,3 +131,25 @@ func TestWriteDefaultRules_DoesNotOverwrite(t *testing.T) { | |||
| 131 | t.Errorf("expected 1 rule (not overwritten), got %d", s.RuleCount()) | 131 | t.Errorf("expected 1 rule (not overwritten), got %d", s.RuleCount()) |
| 132 | } | 132 | } |
| 133 | } | 133 | } |
| 134 | |||
| 135 | func TestScan_SkipsExemptHost(t *testing.T) { | ||
| 136 | path := writeRules(t, `rules: | ||
| 137 | - name: bearer-token | ||
| 138 | pattern: "Authorization:\\s*Bearer\\s+" | ||
| 139 | exempt_hosts: | ||
| 140 | - api.anthropic.com | ||
| 141 | `) | ||
| 142 | s, _ := scanner.New(path) | ||
| 143 | |||
| 144 | body := []byte("Authorization: Bearer sk-ant-123") | ||
| 145 | |||
| 146 | findings := s.Scan(body, "api.anthropic.com") | ||
| 147 | if len(findings) != 0 { | ||
| 148 | t.Errorf("expected 0 findings for exempt host, got %d", len(findings)) | ||
| 149 | } | ||
| 150 | |||
| 151 | findings = s.Scan(body, "evil.com") | ||
| 152 | if len(findings) != 1 { | ||
| 153 | t.Errorf("expected 1 finding for non-exempt host, got %d", len(findings)) | ||
| 154 | } | ||
| 155 | } | ||