Merge: java parameter-pattern preprocessor for param-source taint matching

ober

b58b22f3fe57fa1aaeacb1c7e39d6d927a188eff

diff --git a/lib/semgrep/match/structural.sls b/lib/semgrep/match/structural.sls
index ebd58e8..8ded0c2 100644
--- a/lib/semgrep/match/structural.sls
+++ b/lib/semgrep/match/structural.sls
@@ -620,6 +620,7 @@
              (string=?
                text
                (string-append "func " ellipsis-token-name "()"))
+             (java-synthetic-param-ellipsis? n)
              (and (string=? (node-type n) "expression_statement")
                   (let ([child (single-named-child n)])
                     (and child
@@ -1235,6 +1236,28 @@
                 (when pdecl (node-close! pdecl))
                 (when tdecl (node-close! tdecl))
                 (and result (not (null? result)) result)))))
+  (def (java-param-metavar-matches
+         language
+         pattern
+         target
+         bindings)
+       (and (string=? language "java")
+            (string=? (node-type pattern) "formal_parameter")
+            (string=? (node-type target) "formal_parameter")
+            (let ([name (java-synthetic-param-metavar-name pattern)])
+              (and name
+                   (let ([count (node-named-child-count target)])
+                     (and (> count 0)
+                          (let ([tname (node-named-child
+                                         target
+                                         (- count 1))])
+                            (let ([next (and tname
+                                             (bind-metavariable
+                                               name
+                                               tname
+                                               bindings))])
+                              (when tname (node-close! tname))
+                              (and next (list next))))))))))
   (def (structural-node-matches
          language
          pattern
@@ -1247,6 +1270,12 @@
             (lambda (matches) matches)]
            [(java-field-decl-matches language pattern target bindings) =>
             (lambda (matches) matches)]
+           [(java-param-metavar-matches
+              language
+              pattern
+              target
+              bindings) =>
+            (lambda (matches) matches)]
            [(and (php-language? language)
                  (php-metavariable-name-from-node pattern)) =>
             (lambda (name)
@@ -1715,40 +1744,173 @@
          (if chain-ellipsis?
              (drop-contained-structural-matches raw-matches)
              raw-matches)))
+  (define java-synthetic-param-type "__sg_ptype__")
+  (def (string-all-identifier-chars? s)
+       (let ([len (string-length s)])
+         (and (> len 0)
+              (let loop ([i 0])
+                (cond
+                  [(= i len) #t]
+                  [(identifier-rest? (string-ref s i)) (loop (+ i 1))]
+                  [else #f])))))
+  (def (first-balanced-paren-content text)
+       (let ([len (string-length text)])
+         (let loop ([i 0])
+           (cond
+             [(>= i len) #f]
+             [(char=? (string-ref text i) #\()
+              (let scan ([j (+ i 1)] [depth 1])
+                (cond
+                  [(>= j len) #f]
+                  [(char=? (string-ref text j) #\()
+                   (scan (+ j 1) (+ depth 1))]
+                  [(char=? (string-ref text j) #\))
+                   (if (= depth 1)
+                       (cons (+ i 1) j)
+                       (scan (+ j 1) (- depth 1)))]
+                  [else (scan (+ j 1) depth)]))]
+             [else (loop (+ i 1))]))))
+  (def (split-top-level-commas-str s)
+       (let ([len (string-length s)])
+         (let loop ([i 0] [depth 0] [start 0] [acc '()])
+           (if (= i len)
+               (reverse (cons (substring s start i) acc))
+               (let ([c (string-ref s i)])
+                 (cond
+                   [(or (char=? c #\() (char=? c #\[) (char=? c #\<))
+                    (loop (+ i 1) (+ depth 1) start acc)]
+                   [(or (char=? c #\)) (char=? c #\]) (char=? c #\>))
+                    (loop (+ i 1) (max 0 (- depth 1)) start acc)]
+                   [(and (char=? c #\,) (= depth 0))
+                    (loop
+                      (+ i 1)
+                      0
+                      (+ i 1)
+                      (cons (substring s start i) acc))]
+                   [else (loop (+ i 1) depth start acc)]))))))
+  (def (java-bare-param-token? p)
+       (let ([t (string-trim p)])
+         (or (string=? t ellipsis-token-name)
+             (and (sg-string-prefix? metavariable-prefix t)
+                  (string-all-identifier-chars? t)))))
+  (def (java-has-bare-param-token? text)
+       (let ([grp (first-balanced-paren-content text)])
+         (and grp
+              (let loop ([ps (split-top-level-commas-str
+                               (substring text (car grp) (cdr grp)))])
+                (and (not (null? ps))
+                     (or (java-bare-param-token? (car ps))
+                         (loop (cdr ps))))))))
+  (def (join-with-commas parts)
+       (cond
+         [(null? parts) ""]
+         [(null? (cdr parts)) (car parts)]
+         [else
+          (string-append
+            (car parts)
+            ","
+            (join-with-commas (cdr parts)))]))
+  (def (java-rewrite-param-metavars text)
+       (let ([grp (first-balanced-paren-content text)])
+         (if (not grp)
+             text
+             (let* ([cstart (car grp)]
+                    [cend (cdr grp)]
+                    [content (substring text cstart cend)]
+                    [params (split-top-level-commas-str content)]
+                    [rewritten (map (lambda (p)
+                                      (if (java-bare-param-token? p)
+                                          (string-append
+                                            java-synthetic-param-type
+                                            " "
+                                            (string-trim p))
+                                          p))
+                                    params)]
+                    [new-content (join-with-commas rewritten)])
+               (string-append
+                 (substring text 0 cstart)
+                 new-content
+                 (substring text cend (string-length text)))))))
+  (def (java-param-name-text n)
+       (and (string=? (node-type n) "formal_parameter")
+            (let ([count (node-named-child-count n)])
+              (and (> count 0)
+                   (let ([c (node-named-child n (- count 1))])
+                     (and c
+                          (let ([t (node-text c)]) (node-close! c) t)))))))
+  (def (java-synthetic-param-ellipsis? n)
+       (let ([name (java-param-name-text n)])
+         (and name (string=? name ellipsis-token-name))))
+  (def (java-synthetic-param-metavar-name n)
+       (let ([name (java-param-name-text n)])
+         (and name (metavariable-name-from-text name))))
   (define java-pattern-scaffolds
     (list
       (cons "class _C{Object _f=" ";}")
       (cons "class _C{void _m(){" ";}}")
       (cons "class _C{void _m(){" "}}")))
+  (def (java-parse-method-pattern trimmed)
+       (if (not (java-has-bare-param-token? trimmed))
+           (values #f #f)
+           (let* ([rewritten (java-rewrite-param-metavars trimmed)]
+                  [prefix "class _C{ "]
+                  [suffix " }"]
+                  [wrapped (string-append prefix rewritten suffix)]
+                  [result (parse-target-string "java" wrapped)]
+                  [root (parse-result-root result)])
+             (if (not root)
+                 (begin
+                   (tree-close! (parse-result-tree result))
+                   (values #f #f))
+                 (let* ([hstart (utf8-length prefix)]
+                        [hend (+ hstart (utf8-length rewritten))]
+                        [hole (and (<= (node-start-byte root) hstart)
+                                   (>= (node-end-byte root) hend)
+                                   (go-deepest-node-containing
+                                     root
+                                     hstart
+                                     hend))])
+                   (node-close! root)
+                   (if hole
+                       (values hole result)
+                       (begin
+                         (tree-close! (parse-result-tree result))
+                         (values #f #f))))))))
   (def (java-parse-pattern rewritten-pattern)
        (let ([trimmed (string-trim rewritten-pattern)])
-         (let loop ([scaffolds java-pattern-scaffolds])
-           (if (null? scaffolds)
-               (values #f #f)
-               (let* ([prefix (caar scaffolds)]
-                      [suffix (cdar scaffolds)]
-                      [wrapped (string-append prefix trimmed suffix)]
-                      [result (parse-target-string "java" wrapped)]
-                      [root (parse-result-root result)])
-                 (if (or (not root) (parse-result-has-errors? result))
-                     (begin
-                       (when root (node-close! root))
-                       (tree-close! (parse-result-tree result))
-                       (loop (cdr scaffolds)))
-                     (let* ([hstart (utf8-length prefix)]
-                            [hend (+ hstart (utf8-length trimmed))]
-                            [hole (and (<= (node-start-byte root) hstart)
-                                       (>= (node-end-byte root) hend)
-                                       (go-deepest-node-containing
-                                         root
-                                         hstart
-                                         hend))])
-                       (node-close! root)
-                       (if hole
-                           (values hole result)
+         (let-values ([(mnode mresult)
+                       (java-parse-method-pattern trimmed)])
+           (if mnode
+               (values mnode mresult)
+               (let loop ([scaffolds java-pattern-scaffolds])
+                 (if (null? scaffolds)
+                     (values #f #f)
+                     (let* ([prefix (caar scaffolds)]
+                            [suffix (cdar scaffolds)]
+                            [wrapped (string-append prefix trimmed suffix)]
+                            [result (parse-target-string "java" wrapped)]
+                            [root (parse-result-root result)])
+                       (if (or (not root)
+                               (parse-result-has-errors? result))
                            (begin
+                             (when root (node-close! root))
                              (tree-close! (parse-result-tree result))
-                             (loop (cdr scaffolds)))))))))))
+                             (loop (cdr scaffolds)))
+                           (let* ([hstart (utf8-length prefix)]
+                                  [hend (+ hstart (utf8-length trimmed))]
+                                  [hole (and (<= (node-start-byte root)
+                                                 hstart)
+                                             (>= (node-end-byte root) hend)
+                                             (go-deepest-node-containing
+                                               root
+                                               hstart
+                                               hend))])
+                             (node-close! root)
+                             (if hole
+                                 (values hole result)
+                                 (begin
+                                   (tree-close! (parse-result-tree result))
+                                   (loop (cdr scaffolds)))))))))))))
   (def (structural-pattern-matches-with-bindings
          language
          pattern-source
diff --git a/src/semgrep/match/structural.ss b/src/semgrep/match/structural.ss
index 4565319..fc23fd6 100644
--- a/src/semgrep/match/structural.ss
+++ b/src/semgrep/match/structural.ss
@@ -644,6 +644,8 @@
         ;; Go statement / top-level ellipsis encodings
         (string=? text (string-append ellipsis-token-name "()"))
         (string=? text (string-append "func " ellipsis-token-name "()"))
+        ;; java param-list ellipsis: a synthetic formal_parameter named __sg_ellipsis__
+        (java-synthetic-param-ellipsis? n)
         (and (string=? (node-type n) "expression_statement")
              (let ([child (single-named-child n)])
                (and child
@@ -1251,6 +1253,23 @@
            (when tdecl (node-close! tdecl))
            (and result (not (null? result)) result)))))
 
+;; A synthetic `__sg_ptype__ $X` parameter pattern matches any target
+;; formal_parameter, binding $X to the parameter's NAME (how it is referenced
+;; in the body, e.g. the source param in `public void f(..., $X, ...)`).
+(def (java-param-metavar-matches language pattern target bindings)
+  (and (string=? language "java")
+       (string=? (node-type pattern) "formal_parameter")
+       (string=? (node-type target) "formal_parameter")
+       (let ([name (java-synthetic-param-metavar-name pattern)])
+         (and name
+              (let ([count (node-named-child-count target)])
+                (and (> count 0)
+                     (let ([tname (node-named-child target (- count 1))])
+                       (let ([next (and tname
+                                        (bind-metavariable name tname bindings))])
+                         (when tname (node-close! tname))
+                         (and next (list next))))))))))
+
 (def (structural-node-matches language pattern target bindings)
   (let ([text (node-text pattern)])
     (cond
@@ -1259,6 +1278,8 @@
        => (lambda (matches) matches)]
       [(java-field-decl-matches language pattern target bindings)
        => (lambda (matches) matches)]
+      [(java-param-metavar-matches language pattern target bindings)
+       => (lambda (matches) matches)]
       ;; PHP metavariable `$NAME` (uppercase) -> bind to the target node.
       [(and (php-language? language)
             (php-metavariable-name-from-node pattern))
@@ -1703,6 +1724,116 @@
         (drop-contained-structural-matches raw-matches)
         raw-matches)))
 
+;; --- Java method-parameter pattern preprocessing -------------------------
+;; A param-source pattern like `public void $F(..., $X, ...) { ... }` cannot be
+;; parsed by tree-sitter-java because java parameters must be typed: `$X` and
+;; `...` are bare identifiers in parameter position. We give such params a
+;; synthetic type so they parse as formal_parameters, then the matcher treats a
+;; formal_parameter named `__sg_ellipsis__` as an ellipsis and one named
+;; `__sg_mvar_X` as a metavariable bound to the whole target parameter.
+(define java-synthetic-param-type "__sg_ptype__")
+
+(def (string-all-identifier-chars? s)
+  (let ([len (string-length s)])
+    (and (> len 0)
+         (let loop ([i 0])
+           (cond
+             [(= i len) #t]
+             [(identifier-rest? (string-ref s i)) (loop (+ i 1))]
+             [else #f])))))
+
+;; Content byte range (start . end) of the first balanced (...) group, or #f.
+(def (first-balanced-paren-content text)
+  (let ([len (string-length text)])
+    (let loop ([i 0])
+      (cond
+        [(>= i len) #f]
+        [(char=? (string-ref text i) #\()
+         (let scan ([j (+ i 1)] [depth 1])
+           (cond
+             [(>= j len) #f]
+             [(char=? (string-ref text j) #\() (scan (+ j 1) (+ depth 1))]
+             [(char=? (string-ref text j) #\))
+              (if (= depth 1) (cons (+ i 1) j) (scan (+ j 1) (- depth 1)))]
+             [else (scan (+ j 1) depth)]))]
+        [else (loop (+ i 1))]))))
+
+;; Split on top-level commas (ignoring nested ()/[]/<>).
+(def (split-top-level-commas-str s)
+  (let ([len (string-length s)])
+    (let loop ([i 0] [depth 0] [start 0] [acc '()])
+      (if (= i len)
+          (reverse (cons (substring s start i) acc))
+          (let ([c (string-ref s i)])
+            (cond
+              [(or (char=? c #\() (char=? c #\[) (char=? c #\<))
+               (loop (+ i 1) (+ depth 1) start acc)]
+              [(or (char=? c #\)) (char=? c #\]) (char=? c #\>))
+               (loop (+ i 1) (max 0 (- depth 1)) start acc)]
+              [(and (char=? c #\,) (= depth 0))
+               (loop (+ i 1) 0 (+ i 1) (cons (substring s start i) acc))]
+              [else (loop (+ i 1) depth start acc)]))))))
+
+;; A parameter that is a bare `...` or a bare `$X` (single metavar identifier).
+(def (java-bare-param-token? p)
+  (let ([t (string-trim p)])
+    (or (string=? t ellipsis-token-name)
+        (and (sg-string-prefix? metavariable-prefix t)
+             (string-all-identifier-chars? t)))))
+
+(def (java-has-bare-param-token? text)
+  (let ([grp (first-balanced-paren-content text)])
+    (and grp
+         (let loop ([ps (split-top-level-commas-str
+                          (substring text (car grp) (cdr grp)))])
+           (and (not (null? ps))
+                (or (java-bare-param-token? (car ps)) (loop (cdr ps))))))))
+
+(def (join-with-commas parts)
+  (cond
+    [(null? parts) ""]
+    [(null? (cdr parts)) (car parts)]
+    [else (string-append (car parts) "," (join-with-commas (cdr parts)))]))
+
+;; Give bare `...`/`$X` params a synthetic type so the param list parses.
+(def (java-rewrite-param-metavars text)
+  (let ([grp (first-balanced-paren-content text)])
+    (if (not grp)
+        text
+        (let* ([cstart (car grp)]
+               [cend (cdr grp)]
+               [content (substring text cstart cend)]
+               [params (split-top-level-commas-str content)]
+               [rewritten
+                (map (lambda (p)
+                       (if (java-bare-param-token? p)
+                           (string-append java-synthetic-param-type
+                                          " "
+                                          (string-trim p))
+                           p))
+                     params)]
+               [new-content (join-with-commas rewritten)])
+          (string-append (substring text 0 cstart)
+                         new-content
+                         (substring text cend (string-length text)))))))
+
+;; A formal_parameter's name = its last named child's text, or #f.
+(def (java-param-name-text n)
+  (and (string=? (node-type n) "formal_parameter")
+       (let ([count (node-named-child-count n)])
+         (and (> count 0)
+              (let ([c (node-named-child n (- count 1))])
+                (and c
+                     (let ([t (node-text c)]) (node-close! c) t)))))))
+
+(def (java-synthetic-param-ellipsis? n)
+  (let ([name (java-param-name-text n)])
+    (and name (string=? name ellipsis-token-name))))
+
+(def (java-synthetic-param-metavar-name n)
+  (let ([name (java-param-name-text n)])
+    (and name (metavariable-name-from-text name))))
+
 ;; Java bare expressions/identifiers (e.g. a `tainted` source) error at the
 ;; top level, so wrap them in a class context and extract the spanned node,
 ;; like the Go/PHP scaffolds. Used only as a fallback when the bare parse errors
@@ -1712,8 +1843,37 @@
         (cons "class _C{void _m(){" ";}}")   ; bare statements needing a `;`
         (cons "class _C{void _m(){" "}}")))  ; statements already terminated
 
+;; Method-declaration param-source patterns: param-rewrite + wrap in a class.
+;; The body `{ ... }` recovers (a bare `__sg_ellipsis__` statement errors but the
+;; method_declaration + formal_parameters are present), so errors are tolerated
+;; here and the spanned method node is extracted.
+(def (java-parse-method-pattern trimmed)
+  (if (not (java-has-bare-param-token? trimmed))
+      (values #f #f)
+      (let* ([rewritten (java-rewrite-param-metavars trimmed)]
+             [prefix "class _C{ "]
+             [suffix " }"]
+             [wrapped (string-append prefix rewritten suffix)]
+             [result (parse-target-string "java" wrapped)]
+             [root (parse-result-root result)])
+        (if (not root)
+            (begin (tree-close! (parse-result-tree result)) (values #f #f))
+            (let* ([hstart (utf8-length prefix)]
+                   [hend (+ hstart (utf8-length rewritten))]
+                   [hole (and (<= (node-start-byte root) hstart)
+                              (>= (node-end-byte root) hend)
+                              (go-deepest-node-containing root hstart hend))])
+              (node-close! root)
+              (if hole
+                  (values hole result)
+                  (begin (tree-close! (parse-result-tree result))
+                         (values #f #f))))))))
+
 (def (java-parse-pattern rewritten-pattern)
   (let ([trimmed (string-trim rewritten-pattern)])
+    (let-values ([(mnode mresult) (java-parse-method-pattern trimmed)])
+      (if mnode
+          (values mnode mresult)
     (let loop ([scaffolds java-pattern-scaffolds])
       (if (null? scaffolds)
           (values #f #f)
@@ -1737,7 +1897,7 @@
                       (values hole result)
                       (begin
                         (tree-close! (parse-result-tree result))
-                        (loop (cdr scaffolds)))))))))))
+                        (loop (cdr scaffolds)))))))))))))
 
 (def (structural-pattern-matches-with-bindings
        language