Structural taint engine: numeric-typed param source classification (assume-safe step 2)

ober

4f9f2744cc7a4335f0a76647affdd7bf95bfcab6

diff --git a/lib/semgrep/scan.sls b/lib/semgrep/scan.sls
index d348915..c983a75 100644
--- a/lib/semgrep/scan.sls
+++ b/lib/semgrep/scan.sls
@@ -31925,6 +31925,45 @@
            (taint-assume-safe-booleans? rule)))
   (def (taint-assume-safe-numbers? rule)
        (rule-option-enabled? rule "taint_assume_safe_numbers"))
+  (def (java-numeric-type-name? t)
+       (let ([tt (string-trim t)])
+         (and (member
+                tt
+                '("int" "long" "short" "byte" "float" "double" "char"
+                   "Integer" "Long" "Short" "Byte" "Float" "Double"
+                   "Character"))
+              #t)))
+  (def (param-type-char? c)
+       (or (char-alphabetic? c)
+           (char-numeric? c)
+           (memv c '(#\_ #\$ #\[ #\] #\< #\> #\.))))
+  (def (param-type-before-name source name-start)
+       (let ([len (string-length source)])
+         (if (or (<= name-start 0) (> name-start len))
+             ""
+             (let skip-ws ([i (- name-start 1)])
+               (cond
+                 [(< i 0) ""]
+                 [(char-whitespace? (string-ref source i))
+                  (skip-ws (- i 1))]
+                 [(not (param-type-char? (string-ref source i))) ""]
+                 [else
+                  (let collect ([j i])
+                    (cond
+                      [(< j 0) (substring source 0 (+ i 1))]
+                      [(param-type-char? (string-ref source j))
+                       (collect (- j 1))]
+                      [else (substring source (+ j 1) (+ i 1))]))])))))
+  (def (source-state-numeric-typed-param?
+         source-state
+         source-text)
+       (let ([finding (and source-state
+                           (taint-state-finding source-state))])
+         (and finding
+              (java-numeric-type-name?
+                (param-type-before-name
+                  source-text
+                  (finding-start-offset finding))))))
   (def (taint-assume-safe-indexes? rule)
        (rule-option-enabled? rule "taint_assume_safe_indexes"))
   (def (taint-assume-safe-functions? rule)
@@ -33011,8 +33050,16 @@
   (def (sink-output-findings rule sink-spec sources sanitizers
          assignment-kills source-text)
        (let* ([requires (alist-ref/default sink-spec 'requires #f)]
-              [reaching (source-states-reaching-sink rule sink-spec sources sanitizers
-                          assignment-kills source-text)]
+              [reaching0 (source-states-reaching-sink rule sink-spec sources sanitizers
+                           assignment-kills source-text)]
+              [reaching (if (taint-assume-safe-numbers? rule)
+                            (sg-filter
+                              (lambda (ss)
+                                (not (source-state-numeric-typed-param?
+                                       ss
+                                       source-text)))
+                              reaching0)
+                            reaching0)]
               [labels (label-set-for-source-states reaching)])
          (if (and (not (null? labels))
                   (requires-satisfied? labels requires))
diff --git a/src/semgrep/scan.ss b/src/semgrep/scan.ss
index 953ff74..452fea9 100644
--- a/src/semgrep/scan.ss
+++ b/src/semgrep/scan.ss
@@ -31867,6 +31867,40 @@
 (def (taint-assume-safe-numbers? rule)
   (rule-option-enabled? rule "taint_assume_safe_numbers"))
 
+;; --- assume_safe_numbers: a numeric-typed parameter source carries no taint ---
+(def (java-numeric-type-name? t)
+  (let ([tt (string-trim t)])
+    (and (member tt
+                 '("int" "long" "short" "byte" "float" "double" "char"
+                   "Integer" "Long" "Short" "Byte" "Float" "Double" "Character"))
+         #t)))
+
+(def (param-type-char? c)
+  (or (char-alphabetic? c) (char-numeric? c)
+      (memv c '(#\_ #\$ #\[ #\] #\< #\> #\.))))
+
+(def (param-type-before-name source name-start)
+  (let ([len (string-length source)])
+    (if (or (<= name-start 0) (> name-start len))
+        ""
+        (let skip-ws ([i (- name-start 1)])
+          (cond
+            [(< i 0) ""]
+            [(char-whitespace? (string-ref source i)) (skip-ws (- i 1))]
+            [(not (param-type-char? (string-ref source i))) ""]
+            [else
+             (let collect ([j i])
+               (cond
+                 [(< j 0) (substring source 0 (+ i 1))]
+                 [(param-type-char? (string-ref source j)) (collect (- j 1))]
+                 [else (substring source (+ j 1) (+ i 1))]))])))))
+
+(def (source-state-numeric-typed-param? source-state source-text)
+  (let ([finding (and source-state (taint-state-finding source-state))])
+    (and finding
+         (java-numeric-type-name?
+           (param-type-before-name source-text (finding-start-offset finding))))))
+
 (def (taint-assume-safe-indexes? rule)
   (rule-option-enabled? rule "taint_assume_safe_indexes"))
 
@@ -32920,7 +32954,7 @@
        assignment-kills
        source-text)
   (let* ([requires (alist-ref/default sink-spec 'requires #f)]
-         [reaching
+         [reaching0
           (source-states-reaching-sink
             rule
             sink-spec
@@ -32928,6 +32962,13 @@
             sanitizers
             assignment-kills
             source-text)]
+         [reaching
+          (if (taint-assume-safe-numbers? rule)
+              (sg-filter
+                (lambda (ss)
+                  (not (source-state-numeric-typed-param? ss source-text)))
+                reaching0)
+              reaching0)]
          [labels (label-set-for-source-states reaching)])
     (if (and (not (null? labels))
              (requires-satisfied? labels requires))