diff --git a/python/pylibcudf/tests/test_string_contains.py b/python/pylibcudf/tests/test_string_contains.py index 9807e78bf680..ec0d1b6c67db 100644 --- a/python/pylibcudf/tests/test_string_contains.py +++ b/python/pylibcudf/tests/test_string_contains.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024-2025, NVIDIA CORPORATION. +# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 import pyarrow as pa @@ -9,44 +9,26 @@ import pylibcudf as plc -@pytest.fixture(scope="module") -def target_col(): - pa_array = pa.array( - ["AbC", "de", "FGHI", "j", "kLm", "nOPq", None, "RsT", None, "uVw"] - ) - return pa_array, plc.Column.from_arrow(pa_array) - - -@pytest.fixture( - params=[ - "A", - "de", - ".*", - "^a", - "^A", - "[^a-z]", - "[a-z]{3,}", - "^[A-Z]{2,}", - "j|u", - ], - scope="module", -) -def pa_target_scalar(request): - return pa.scalar(request.param, type=pa.string()) +def _make_prog(pattern): + flags = plc.strings.regex_flags.RegexFlags.DEFAULT + return plc.strings.regex_program.RegexProgram.create(pattern, flags) -@pytest.fixture(scope="module") -def plc_target_pat(pa_target_scalar): - prog = plc.strings.regex_program.RegexProgram.create( - pa_target_scalar.as_py(), plc.strings.regex_flags.RegexFlags.DEFAULT +@pytest.mark.parametrize( + "input_strings", + [["AbC", "de", "FGHI", "j", "kLm", "nOPq", None, "RsT", None, "uVw"]], +) +@pytest.mark.parametrize( + "pattern", + ["A", "de", ".*", "^a", "^A", "[^a-z]", "[a-z]{3,}", "^[A-Z]{2,}", "j|u"], +) +def test_contains_re(input_strings, pattern): + input = pa.array(input_strings) + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(input), + _make_prog(pattern), ) - return prog - - -def test_contains_re(target_col, pa_target_scalar, plc_target_pat): - pa_target_col, plc_target_col = target_col - got = plc.strings.contains.contains_re(plc_target_col, plc_target_pat) - expect = pc.match_substring_regex(pa_target_col, pa_target_scalar.as_py()) + expect = pc.match_substring_regex(input, pattern) assert_column_eq(expect, got) @@ -85,3 +67,214 @@ def test_like(): ) expect = pc.match_like(arr, pattern) assert_column_eq(expect, got) + + +# Tests derived from cudf-spark integration tests + + +@pytest.fixture(scope="module") +def spark_strings(): + """Rich string array that exercises the cudf-spark regex patterns.""" + return pa.array( + [ + "abc", + "aabbc", + "123abc", + "abc123def", + "a1b2c3", + "boo:and:foo", + "foo:boo:", + "TEST", + "TESTaaa", + "TEST123", + "abcd", + "aaa", + "bbb", + "ccc", + "abb", + "ab", + "aab", + "a\nb", + "a\tb", + "a b", + "", + "aaabbb", + "abcabc", + "aa|bb", + "foobar", + "12345", + "abcdef", + "ABCDEF", + "AbCdEf", + None, + ] + ) + + +# Basic quantifiers (test_rlike, test_regexp, test_regexp_like) +@pytest.mark.parametrize( + "pattern", + [ + "a{2}", + "a{1,3}", + "a{1,}", + "a[bc]d", + ], +) +def test_contains_re_basic(spark_strings, pattern): + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + ) + expect = pc.match_substring_regex(spark_strings, pattern) + assert_column_eq(expect, got) + + +# Alternation patterns (test_regexp_choice, test_rlike_rewrite_optimization) +@pytest.mark.parametrize( + "pattern", + [ + "aaa|bbb|ccc", + "1|2|3|4|5|6", + "[abcd]|[123]", + "aaa|bbb", + "aaa|(bbb|ccc)", + ".*.*(aaa|bbb).*.*", + "^.*(aaa|bbb|ccc)", + "abd1a$|^ab2a", + "[ab]+|^cd1", + ], +) +def test_contains_re_alternation(spark_strings, pattern): + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + ) + expect = pc.match_substring_regex(spark_strings, pattern) + assert_column_eq(expect, got) + + +# Anchor + wildcard patterns (test_rlike_rewrite_optimization) +@pytest.mark.parametrize( + "pattern", + [ + "^abb", + "^.*(aaa)", + "^(abb)(.*)", + "abb(.*)", + "(.*)(abb)(.*)", + "ab(.*)cd", + "(.*)(.*)abb", + ], +) +def test_contains_re_anchors_wildcards(spark_strings, pattern): + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + ) + expect = pc.match_substring_regex(spark_strings, pattern) + assert_column_eq(expect, got) + + +# Bounded repetition (test_rlike_rewrite_optimization, test_regexp) +@pytest.mark.parametrize( + "pattern", + [ + "ab[a-c]{3}", + "a[a-c]{1,3}", + "a[a-c]{1,}", + "a[a-c]+", + "(ab)([a-c]{1})", + "(ab[a-c]{1})", + "a{6}", + "a{1,6}", + ], +) +def test_contains_re_bounded_repetition(spark_strings, pattern): + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + ) + expect = pc.match_substring_regex(spark_strings, pattern) + assert_column_eq(expect, got) + + +# Non-capturing groups / complex quantifiers (test_regexp_memory_ok) +@pytest.mark.parametrize( + "pattern", + [ + "(?:12345)+", + "(?:aa)+", + "abcdef", + "(1)(2)(3)", + ], +) +def test_contains_re_groups(spark_strings, pattern): + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + ) + expect = pc.match_substring_regex(spark_strings, pattern) + assert_column_eq(expect, got) + + +# Predefined character classes (test_character_classes, test_regexp_whitespace, +# test_regexp_replace_digit, test_regexp_replace_word) +@pytest.mark.parametrize( + "pattern", + [ + r"\d", + r"\D", + r"[0-9]", + r"[^0-9]", + r"\w", + r"[a-zA-Z_0-9]", + r"\s", + r"\S", + r"[abcd]+\s+[0-9]+", + r"\S{3}", + r"[^\n\r]", + ], +) +def test_contains_re_char_classes(spark_strings, pattern): + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + ) + expect = pc.match_substring_regex(spark_strings, pattern) + assert_column_eq(expect, got) + + +# \W: cuDF uses ASCII semantics — only [^a-zA-Z0-9_] are non-word chars, +# unlike pyarrow/RE2 which is Unicode-aware. Test with hardcoded expectations. +def test_contains_re_nonword(): + arr = pa.array( + ["abc", "abc123", "abc!", "a b", "aa|bb", "12345", "", None] + ) + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(arr), + _make_prog(r"\W"), + ) + # cuDF \W matches non-ASCII-word characters: space, !, | + expect = pa.array([False, False, True, True, True, False, False, None]) + assert_column_eq(expect, got) + + +# Escape / character class edge cases (test_rlike_escape, test_rlike_missing_escape) +@pytest.mark.parametrize( + "pattern", + [ + r"a[\-]", + r"a[+-]", + r"a[a-b-]", + r"[a-z]{3,}", + r"^[A-Z]{2,}", + ], +) +def test_contains_re_escape_edge_cases(spark_strings, pattern): + got = plc.strings.contains.contains_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + ) + expect = pc.match_substring_regex(spark_strings, pattern) + assert_column_eq(expect, got) diff --git a/python/pylibcudf/tests/test_string_replace_re.py b/python/pylibcudf/tests/test_string_replace_re.py index 798c95bb7ce7..7303f2d040c8 100644 --- a/python/pylibcudf/tests/test_string_replace_re.py +++ b/python/pylibcudf/tests/test_string_replace_re.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION. +# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 import pyarrow as pa @@ -9,6 +9,11 @@ import pylibcudf as plc +def _make_prog(pattern): + flags = plc.strings.regex_flags.RegexFlags.DEFAULT + return plc.strings.regex_program.RegexProgram.create(pattern, flags) + + @pytest.mark.parametrize("max_replace_count", [-1, 1]) def test_replace_re_regex_program_scalar(max_replace_count): arr = pa.array(["foo", "fuz", None]) @@ -16,9 +21,7 @@ def test_replace_re_regex_program_scalar(max_replace_count): repl = "ba" got = plc.strings.replace_re.replace_re( plc.Column.from_arrow(arr), - plc.strings.regex_program.RegexProgram.create( - pat, plc.strings.regex_flags.RegexFlags.DEFAULT - ), + _make_prog(pat), plc.Scalar.from_arrow(pa.scalar(repl)), max_replace_count=max_replace_count, ) @@ -37,10 +40,204 @@ def test_replace_with_backrefs(): arr = pa.array(["Z756", None]) got = plc.strings.replace_re.replace_with_backrefs( plc.Column.from_arrow(arr), - plc.strings.regex_program.RegexProgram.create( - "(\\d)(\\d)", plc.strings.regex_flags.RegexFlags.DEFAULT - ), + _make_prog("(\\d)(\\d)"), "V\\2\\1", ) expect = pa.array(["ZV576", None]) assert_column_eq(expect, got) + + +# Tests derived from cudf-spark integration tests + + +@pytest.fixture(scope="module") +def spark_strings(): + """Rich string array that exercises the cudf-spark regex patterns.""" + return pa.array( + [ + "abc", + "aabbc", + "123abc", + "abc123def", + "a1b2c3", + "boo:and:foo", + "foo:boo:", + "TEST", + "TESTaaa", + "TEST123", + "abcd", + "aaa", + "bbb", + "ccc", + "abb", + "ab", + "aab", + "a\nb", + "a\tb", + "a b", + "", + "aaabbb", + "abcabc", + "aa|bb", + "foobar", + "12345", + "abcdef", + "ABCDEF", + "AbCdEf", + None, + ] + ) + + +REPLACE_REPL = "X" # replacement string used in all replace tests + + +# Basic replace patterns (test_re_replace, test_regexp_replace) +@pytest.mark.parametrize( + "pattern", + [ + "TEST", + "[A-Z]+", + "a", + "[^xyz]", + "a|b|c", + ], +) +def test_replace_re_basic(spark_strings, pattern): + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pc.replace_substring_regex(spark_strings, pattern, REPLACE_REPL) + assert_column_eq(expect, got) + + +# Repetition quantifiers (test_re_replace_repetition) +@pytest.mark.parametrize( + "pattern", + [ + "[E]+", + "[A]+", + "[A-Z]+", + ], +) +def test_replace_re_quantifiers(spark_strings, pattern): + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pc.replace_substring_regex(spark_strings, pattern, REPLACE_REPL) + assert_column_eq(expect, got) + + +# Negated character classes (test_regexp_replace_character_set_negated) +@pytest.mark.parametrize( + "pattern", + [ + "[^a]", + r"[^a\r\n]", + r"[^\r\n]", + r"[^\r]", + r"[^\n]", + ], +) +def test_replace_re_negated_classes(spark_strings, pattern): + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pc.replace_substring_regex(spark_strings, pattern, REPLACE_REPL) + assert_column_eq(expect, got) + + +# Digit and word classes (test_regexp_replace_digit, test_regexp_replace_word) +# \D and \W are excluded from pyarrow comparison: cuDF uses ASCII semantics for +# \w/\W (only [a-zA-Z0-9_]), while pyarrow/RE2 is Unicode-aware. +@pytest.mark.parametrize( + "pattern", + [ + r"\d", + r"[0-9]", + r"[^0-9]", + r"\w", + r"[a-zA-Z_0-9]", + r"[^a-zA-Z_0-9]", + ], +) +def test_replace_re_digit_word(spark_strings, pattern): + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pc.replace_substring_regex(spark_strings, pattern, REPLACE_REPL) + assert_column_eq(expect, got) + + +# \D replace: cuDF ASCII semantics — non-digits include letters, spaces, punctuation +def test_replace_re_nondigit(): + arr = pa.array(["abc", "a1b2", "123", "a b", "", None]) + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(arr), + _make_prog(r"\D"), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pa.array(["XXX", "X1X2", "123", "XXX", "", None]) + assert_column_eq(expect, got) + + +# \W replace: cuDF ASCII semantics — non-word chars are non-[a-zA-Z0-9_] +def test_replace_re_nonword(): + arr = pa.array(["abc", "a b", "a!b", "aa|bb", "123", "", None]) + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(arr), + _make_prog(r"\W"), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pa.array(["abc", "aXb", "aXb", "aaXbb", "123", "", None]) + assert_column_eq(expect, got) + + +# Multi-alternation (test_regexp_replace_multi_optimization) +@pytest.mark.parametrize( + "pattern", + [ + "aa|bb", + "aa|bb|cc", + "aa|bb|cc|dd", + "aa|bb|cc|dd|ee", + "aa|bb|cc|dd|ee|ff", + "(aa)|(bb)", + "(aa)|(bb)|(cc)", + "(aa|bb)|(cc|dd)", + ], +) +def test_replace_re_multi_alternation(spark_strings, pattern): + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pc.replace_substring_regex(spark_strings, pattern, REPLACE_REPL) + assert_column_eq(expect, got) + + +# Non-capturing groups (test_regexp_replace) +@pytest.mark.parametrize( + "pattern", + [ + "(?:aa)+", + "([^x])|([^y])", + ], +) +def test_replace_re_noncapturing(spark_strings, pattern): + got = plc.strings.replace_re.replace_re( + plc.Column.from_arrow(spark_strings), + _make_prog(pattern), + plc.Scalar.from_arrow(pa.scalar(REPLACE_REPL)), + ) + expect = pc.replace_substring_regex(spark_strings, pattern, REPLACE_REPL) + assert_column_eq(expect, got)