diff --git a/.github/workflows/beam_PostRelease_NightlySnapshot.yml b/.github/workflows/beam_PostRelease_NightlySnapshot.yml index d4e2e0690cfb..d119dfb7754a 100644 --- a/.github/workflows/beam_PostRelease_NightlySnapshot.yml +++ b/.github/workflows/beam_PostRelease_NightlySnapshot.yml @@ -19,12 +19,13 @@ on: workflow_dispatch: inputs: RELEASE: - description: Beam version of current release (e.g. 2.XX.0) - required: true + description: Beam version of current release (pass in empty string for nightly SNAPSHOT) + required: false default: '2.XX.0' SNAPSHOT_URL: - description: Location of the staged artifacts in Maven central (https://repository.apache.org/content/repositories/orgapachebeam-NNNN/). - required: true + description: Location of the staged artifacts in Maven central (https://repository.apache.org/content/repositories/orgapachebeam-NNNN/ or leave empty for snapshots). + required: false + default: '' schedule: - cron: '15 16 * * *' diff --git a/.github/workflows/build_release_candidate.yml b/.github/workflows/build_release_candidate.yml index 28d24e986e52..978b546b8068 100644 --- a/.github/workflows/build_release_candidate.yml +++ b/.github/workflows/build_release_candidate.yml @@ -65,7 +65,7 @@ jobs: with: ref: "v${{ github.event.inputs.RELEASE }}-RC${{ github.event.inputs.RC }}" repository: apache/beam - persist-credentials: false + persist-credentials: true - name: Install Java 11 uses: actions/setup-java@v5 with: @@ -184,7 +184,7 @@ jobs: - name: Checkout uses: actions/checkout@v7 with: - persist-credentials: false + persist-credentials: true - name: Mask Apache Password run: | # Workaround for Actions bug - https://github.com/actions/runner/issues/643 @@ -294,7 +294,7 @@ jobs: with: ref: "v${{ github.event.inputs.RELEASE }}-RC${{ github.event.inputs.RC }}" repository: apache/beam - persist-credentials: false + persist-credentials: true - name: Free Disk Space (Ubuntu) uses: jlumbroso/free-disk-space@v1.3.1 - name: Install Java @@ -346,7 +346,7 @@ jobs: ref: "v${{ github.event.inputs.RELEASE }}-RC${{ github.event.inputs.RC }}" repository: apache/beam path: beam - persist-credentials: false + persist-credentials: true - name: Checkout Beam Site Repo uses: actions/checkout@v7 with: @@ -354,7 +354,7 @@ jobs: path: beam-site token: ${{ github.event.inputs.REPO_TOKEN }} ref: release-docs - persist-credentials: false + persist-credentials: true - name: Install Python 3.10 uses: actions/setup-python@v7 with: @@ -467,7 +467,7 @@ jobs: with: ref: "v${{ github.event.inputs.RELEASE }}-RC${{ github.event.inputs.RC }}" repository: apache/beam - persist-credentials: false + persist-credentials: true - name: Mask Apache Password run: | # Workaround for Actions bug - https://github.com/actions/runner/issues/643 diff --git a/CHANGES.md b/CHANGES.md index 6f61ef415b4b..e2c3f2625361 100644 --- a/CHANGES.md +++ b/CHANGES.md @@ -123,6 +123,7 @@ * (Python) Added `Watch`, a transform that polls a growing set of outputs for each input element, deduplicates outputs across poll rounds, and stops per a user-supplied termination condition ([#21521](https://github.com/apache/beam/issues/21521)). * (Python) Added support to analyze core dumps created after python worker segmentation faults with `pystack` (or `gdb` if installed) using the `--profiler_agent=coredump` pipeline option. ([#39484](https://github.com/apache/beam/issues/39484)). +* (Python) Added `Sample.Any`, the Python equivalent of Java's `Sample.any`, which returns up to n arbitrary elements from a PCollection ([#18552](https://github.com/apache/beam/issues/18552)). * (Java) Added per-element OpenTelemetry trace propagation across stages in the Dataflow Streaming Runner. Enable it with `--experiments=enable_otel_defaults,element_metadata_supported,disable_portable_worker`. Cloud Trace incurs additional cost. ([#33176](https://github.com/apache/beam/issues/33176)) * (Java) Added OpenTelemetry header propagation support for both reads and writes in KafkaIO and PubSubIO. ([#33176](https://github.com/apache/beam/issues/33176)) * (Java) Added OpenTelemetry tracing support for SpannerIO change streams ([#33176](https://github.com/apache/beam/issues/33176)) diff --git a/release/build.gradle.kts b/release/build.gradle.kts index 54165dc49654..6cf3540a7f88 100644 --- a/release/build.gradle.kts +++ b/release/build.gradle.kts @@ -41,7 +41,7 @@ task("runJavaExamplesValidationTask") { dependsOn(":runners:spark:3:runQuickstartJavaSpark") dependsOn(":runners:flink:2.2:runQuickstartJavaFlinkLocal") dependsOn(":runners:direct-java:runMobileGamingJavaDirect") - if (project.hasProperty("ver") || !project.version.toString().endsWith("SNAPSHOT")) { + if ((project.findProperty("ver")?.toString()?.isNotEmpty() == true && project.findProperty("ver") != "2.XX.0") || !project.version.toString().endsWith("SNAPSHOT")) { // only run one variant of MobileGaming on Dataflow for nightly dependsOn(":runners:google-cloud-dataflow-java:runMobileGamingJavaDataflow") } diff --git a/release/src/main/groovy/TestScripts.groovy b/release/src/main/groovy/TestScripts.groovy index dc2438007ac1..e0e9cf454495 100644 --- a/release/src/main/groovy/TestScripts.groovy +++ b/release/src/main/groovy/TestScripts.groovy @@ -25,6 +25,10 @@ import groovy.util.CliBuilder */ class TestScripts { + class BackgroundProcessInfo { + String cmd + } + // Global state to maintain when running the steps class var { static File startDir @@ -37,6 +41,8 @@ class TestScripts { static String bqDataset static String pubsubTopic static String mavenLocalPath + static List backgroundProcesses = Collections.synchronizedList(new ArrayList()) + static Map backgroundProcessInfo = Collections.synchronizedMap(new HashMap()) } def TestScripts(String[] args) { @@ -79,6 +85,10 @@ class TestScripts { var.mavenLocalPath = options.mavenLocalPath println "Maven local path: ${var.mavenLocalPath}" } + + Runtime.getRuntime().addShutdownHook(new Thread({ + stopAllBackgroundProcesses() + })) } def ver() { @@ -135,6 +145,75 @@ class TestScripts { } } + // Run a command in the background, returning the Process object. + public Process runBackground(String cmd) { + println cmd + if (cmd.startsWith("mvn ")) { + return _mvnBackground(cmd.substring(4)) + } else { + return _executeBackground(cmd) + } + } + + // Check whether any background processes exited unexpectedly with a non-zero exit code + public void checkBackgroundProcesses() { + def procs = new ArrayList<>(var.backgroundProcesses) + for (Process proc : procs) { + if (proc != null && !proc.isAlive()) { + int exitVal = proc.exitValue() + if (exitVal != 0) { + def info = var.backgroundProcessInfo.get(proc) + String cmd = info ? info.cmd : "unknown command" + error("Background command failed with exit code ${exitVal}: ${cmd}") + } + } + } + } + + // Stop/kill a background process and all its descendants. + public void stopProcess(Process proc) { + if (proc != null) { + if (!proc.isAlive()) { + int exitVal = proc.exitValue() + var.backgroundProcesses.remove(proc) + def info = var.backgroundProcessInfo.remove(proc) + if (exitVal != 0) { + String cmd = info ? info.cmd : "unknown command" + error("Background command failed with exit code ${exitVal}: ${cmd}") + } + } else { + try { + proc.descendants().forEach { it.destroyForcibly() } + } catch (Throwable ignored) { + } + proc.destroyForcibly() + proc.waitFor(10, java.util.concurrent.TimeUnit.SECONDS) + var.backgroundProcesses.remove(proc) + var.backgroundProcessInfo.remove(proc) + } + } + } + + // Stop all active background processes. + public void stopAllBackgroundProcesses() { + def procs = new ArrayList<>(var.backgroundProcesses) + procs.each { proc -> + if (proc != null && proc.isAlive()) { + try { + proc.descendants().forEach { it.destroyForcibly() } + } catch (Throwable ignored) { + } + proc.destroyForcibly() + try { + proc.waitFor(10, java.util.concurrent.TimeUnit.SECONDS) + } catch (Throwable ignored) { + } + } + var.backgroundProcesses.remove(proc) + var.backgroundProcessInfo.remove(proc) + } + } + // Check for expected results in actual stdout from previous command, if fails, log errors then exit. public void see(String expected, String actual) { if (!actual.contains(expected)) { @@ -159,6 +238,8 @@ class TestScripts { // Cleanup and print success public void done() { + checkBackgroundProcesses() + stopAllBackgroundProcesses() var.startDir.deleteDir() println "[SUCCESS]" System.exit(0) @@ -166,6 +247,7 @@ class TestScripts { // Run a single command, capture output, verify return code is 0 private String _execute(String cmd) { + checkBackgroundProcesses() def shell = "sh -c cmd".split(' ') shell[2] = cmd def pb = new ProcessBuilder(shell) @@ -187,6 +269,27 @@ class TestScripts { return output_text } + // Run a single command asynchronously in the background + private Process _executeBackground(String cmd) { + def shell = "sh -c cmd".split(' ') + shell[2] = cmd + def pb = new ProcessBuilder(shell) + pb.directory(var.curDir) + pb.redirectErrorStream(true) + def proc = pb.start() + var.backgroundProcesses.add(proc) + var.backgroundProcessInfo.put(proc, new BackgroundProcessInfo(cmd: cmd)) + Thread.startDaemon { + try { + proc.inputStream.eachLine { + println it + } + } catch (Throwable ignored) { + } + } + return proc + } + // Change directory private void _chdir(String subdir) { var.curDir = new File(var.curDir.absolutePath, subdir) @@ -195,8 +298,8 @@ class TestScripts { } } - // Run a maven command, setting up a new local repository and a settings.xml with a custom repository if needed - private String _mvn(String args) { + // Build the maven command string with custom repository and settings.xml + private String _buildMvnCmd(String args) { String mvnlocalPath = var.mavenLocalPath if (!(var.mavenLocalPath)) { mvnlocalPath = var.startDir @@ -204,37 +307,48 @@ class TestScripts { def m2 = new File(mvnlocalPath, ".m2/repository") m2.mkdirs() def settings = new File(mvnlocalPath, "settings.xml") - if(!settings.exists()) { - settings.write """ - - ${m2.absolutePath} - - - testrel - - - test.release - ${var.repoUrl} - - - - - - """ + if (!settings.exists()) { + settings.write """ + + ${m2.absolutePath} + + + testrel + + + test.release + ${var.repoUrl} + + + + + + """ } def cmd = "mvn ${args} -s ${settings.absolutePath} -Ptestrel -B" - String path = System.getenv("PATH"); + String path = System.getenv("PATH") // Set the path on jenkins executors to use a recent maven // MAVEN_HOME is not set on some executors, so default to 3.5.2 String maven_home = System.getenv("MAVEN_HOME") ?: '/usr/local/maven' println "Using maven ${maven_home}" def mvnPath = "${maven_home}/bin" def setPath = "export PATH=\"${mvnPath}:${path}\" && " - return _execute(setPath + cmd) + return setPath + cmd + } + + // Run a maven command, setting up a new local repository and a settings.xml with a custom repository if needed + private String _mvn(String args) { + return _execute(_buildMvnCmd(args)) + } + + // Run a maven command in the background + private Process _mvnBackground(String args) { + return _executeBackground(_buildMvnCmd(args)) } // Clean up and report error public void error(String text) { + stopAllBackgroundProcesses() var.startDir.deleteDir() println "[ERROR] $text" System.exit(1) diff --git a/release/src/main/groovy/mobilegaming-java-dataflow.groovy b/release/src/main/groovy/mobilegaming-java-dataflow.groovy index 51ea528a7638..31b4f0670f98 100644 --- a/release/src/main/groovy/mobilegaming-java-dataflow.groovy +++ b/release/src/main/groovy/mobilegaming-java-dataflow.groovy @@ -120,37 +120,53 @@ class LeaderBoardRunner { ].join(",") String tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") - - if (!tables.contains(userTable)) { - t.intent("Creating table: ${userTable}") - t.run("bq mk --table ${dataset}.${userTable} ${userSchema}") + if (tables.contains(userTable)) { + t.run("bq rm -f -t ${dataset}.${userTable}") + } + if (tables.contains(teamTable)) { + t.run("bq rm -f -t ${dataset}.${teamTable}") } - if (!tables.contains(teamTable)) { - t.intent("Creating table: ${teamTable}") - t.run("bq mk --table ${dataset}.${teamTable} ${teamSchema}") + int retries = 10 + boolean deleted = false + for (int i = 0; i < retries; i++) { + tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") + if (!tables.contains(userTable) && !tables.contains(teamTable)) { + deleted = true + break + } + sleep(3000) } + if (!deleted) { + t.error("Timed out waiting for tables ${userTable} / ${teamTable} to be deleted.") + } + + t.intent("Creating table: ${userTable}") + t.run("bq mk --table ${dataset}.${userTable} ${userSchema}") + t.intent("Creating table: ${teamTable}") + t.run("bq mk --table ${dataset}.${teamTable} ${teamSchema}") // Verify that the tables have been created successfully - tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") - while (!tables.contains(userTable) || !tables.contains(teamTable)) { - sleep(3000) + boolean created = false + for (int i = 0; i < retries; i++) { tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") + if (tables.contains(userTable) && tables.contains(teamTable)) { + created = true + break + } + sleep(3000) + } + if (!created) { + t.error("Timed out waiting for tables ${userTable} / ${teamTable} to be created.") } println "Tables ${userTable} and ${teamTable} created successfully." - def InjectorThread = Thread.start() { - t.run(mobileGamingCommands.createInjectorCommand()) - } + def injectorProcess = t.runBackground(mobileGamingCommands.createInjectorCommand()) String jobName = "leaderboard-validation-" + new Date().getTime() + "-" + new Random().nextInt(1000) - def LeaderBoardThread = Thread.start() { - if (useStreamingEngine) { - t.run(mobileGamingCommands.createPipelineCommand( - "LeaderBoardWithStreamingEngine", runner, jobName, "LeaderBoard")) - } else { - t.run(mobileGamingCommands.createPipelineCommand("LeaderBoard", runner, jobName)) - } - } + def leaderBoardProcess = useStreamingEngine ? + t.runBackground(mobileGamingCommands.createPipelineCommand( + "LeaderBoardWithStreamingEngine", runner, jobName, "LeaderBoard")) : + t.runBackground(mobileGamingCommands.createPipelineCommand("LeaderBoard", runner, jobName)) t.run("gcloud dataflow jobs list | grep pyflow-wordstream-candidate | grep Running | cut -d' ' -f1") @@ -175,8 +191,8 @@ class LeaderBoardRunner { println "Waiting for pipeline to produce more results..." sleep(60000) // wait for 1 min } - InjectorThread.stop() - LeaderBoardThread.stop() + t.stopProcess(injectorProcess) + t.stopProcess(leaderBoardProcess) t.run("""RUNNING_JOB=`gcloud dataflow jobs list | grep ${jobName} | grep Running | cut -d' ' -f1` if [ ! -z "\${RUNNING_JOB}" ] then @@ -202,10 +218,17 @@ fi // It will take couple seconds to clean up tables. // This loop makes sure tables are completely deleted before running the pipeline - tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") - while (tables.contains(userTable) || tables.contains(teamTable)) { - sleep(3000) + deleted = false + for (int i = 0; i < retries; i++) { tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") + if (!tables.contains(userTable) && !tables.contains(teamTable)) { + deleted = true + break + } + sleep(3000) + } + if (!deleted) { + println "Warning: Timed out waiting for tables ${userTable} / ${teamTable} to be deleted." } } } diff --git a/release/src/main/groovy/mobilegaming-java-direct.groovy b/release/src/main/groovy/mobilegaming-java-direct.groovy index 34eab4c00768..398822a9a2ce 100644 --- a/release/src/main/groovy/mobilegaming-java-direct.groovy +++ b/release/src/main/groovy/mobilegaming-java-direct.groovy @@ -80,32 +80,51 @@ def teamSchema = [ ].join(",") String tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") +if (tables.contains(userTable)) { + t.run("bq rm -f -t ${dataset}.${userTable}") +} +if (tables.contains(teamTable)) { + t.run("bq rm -f -t ${dataset}.${teamTable}") +} -if (!tables.contains(userTable)) { - t.intent("Creating table: ${userTable}") - t.run("bq mk --table ${dataset}.${userTable} ${userSchema}") +int retries = 10 +boolean deleted = false +for (int i = 0; i < retries; i++) { + tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") + if (!tables.contains(userTable) && !tables.contains(teamTable)) { + deleted = true + break + } + sleep(3000) } -if (!tables.contains(teamTable)) { - t.intent("Creating table: ${teamTable}") - t.run("bq mk --table ${dataset}.${teamTable} ${teamSchema}") +if (!deleted) { + t.error("Timed out waiting for tables ${userTable} / ${teamTable} to be deleted.") } +t.intent("Creating table: ${userTable}") +t.run("bq mk --table ${dataset}.${userTable} ${userSchema}") +t.intent("Creating table: ${teamTable}") +t.run("bq mk --table ${dataset}.${teamTable} ${teamSchema}") + // Verify that the tables have been created -tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") -while (!tables.contains(userTable) || !tables.contains(teamTable)) { - sleep(3000) +boolean created = false +for (int i = 0; i < retries; i++) { tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") + if (tables.contains(userTable) && tables.contains(teamTable)) { + created = true + break + } + sleep(3000) +} +if (!created) { + t.error("Timed out waiting for tables ${userTable} / ${teamTable} to be created.") } println "Tables ${userTable} and ${teamTable} created successfully." -def InjectorThread = Thread.start() { - t.run(mobileGamingCommands.createInjectorCommand()) -} +def injectorProcess = t.runBackground(mobileGamingCommands.createInjectorCommand()) jobName = "leaderboard-validation-" + new Date().getTime() + "-" + new Random().nextInt(1000) -def LeaderBoardThread = Thread.start() { - t.run(mobileGamingCommands.createPipelineCommand("LeaderBoard", runner, jobName)) -} +def leaderBoardProcess = t.runBackground(mobileGamingCommands.createPipelineCommand("LeaderBoard", runner, jobName)) // verify outputs in BQ tables def startTime = System.currentTimeMillis() @@ -128,12 +147,32 @@ while ((System.currentTimeMillis() - startTime)/60000 < mobileGamingCommands.EXE println "Waiting for pipeline to produce more results..." sleep(60000) // wait for 1 min } -InjectorThread.stop() -LeaderBoardThread.stop() +t.stopProcess(injectorProcess) +t.stopProcess(leaderBoardProcess) if(!isSuccess){ t.error("FAILED: Failed running LeaderBoard on DirectRunner") } t.success("LeaderBoard successfully run on DirectRunner.") +tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") +if (tables.contains(userTable)) { + t.run("bq rm -f -t ${dataset}.${userTable}") +} +if (tables.contains(teamTable)) { + t.run("bq rm -f -t ${dataset}.${teamTable}") +} +deleted = false +for (int i = 0; i < retries; i++) { + tables = t.run("bq query --use_legacy_sql=false 'SELECT table_name FROM ${dataset}.INFORMATION_SCHEMA.TABLES'") + if (!tables.contains(userTable) && !tables.contains(teamTable)) { + deleted = true + break + } + sleep(3000) +} +if (!deleted) { + println "Warning: Timed out waiting for tables ${userTable} / ${teamTable} to be deleted." +} + t.done() diff --git a/release/src/main/groovy/quickstart-java-flinklocal.groovy b/release/src/main/groovy/quickstart-java-flinklocal.groovy index 36c6ddd38354..3cd59270c04d 100644 --- a/release/src/main/groovy/quickstart-java-flinklocal.groovy +++ b/release/src/main/groovy/quickstart-java-flinklocal.groovy @@ -41,9 +41,9 @@ t.describe 'Run Apache Beam Java SDK Quickstart - Flink Local' -Dhttp.keepAlive=false \ -Pflink-runner""" - def cp = "target/classes:${deps}" + def cp = "target/classes:${deps.trim()}" t.run """mvn exec:exec -q -Dexec.executable=java \ - -Dexec.args="-cp ${cp} org.apache.beam.examples.WordCount \ + -Dexec.args="-cp '${cp}' org.apache.beam.examples.WordCount \ --inputFile=pom.xml --output=counts --runner=FlinkRunner" """ // Verify text from the pom.xml input file diff --git a/release/src/main/groovy/quickstart-java-spark.groovy b/release/src/main/groovy/quickstart-java-spark.groovy index 3c5be754daab..248e85dbcc59 100644 --- a/release/src/main/groovy/quickstart-java-spark.groovy +++ b/release/src/main/groovy/quickstart-java-spark.groovy @@ -30,10 +30,22 @@ t.describe 'Run Apache Beam Java SDK Quickstart - Spark' t.intent 'Runs the WordCount Code with Spark runner' // Run the wordcount example with the spark runner - t.run """mvn compile exec:java -q \ - -Dexec.mainClass=org.apache.beam.examples.WordCount \ - -Dexec.args="--inputFile=pom.xml --output=counts \ - --runner=SparkRunner" -Pspark-runner""" + + // Retrieve classpath + def deps = t.run """mvn compile dependency:build-classpath -q \ + -Dmdep.outputFile=/dev/stdout \ + -Dmaven.wagon.http.retryHandler.class=default \ + -Dmaven.wagon.http.retryHandler.count=5 \ + -Dmaven.wagon.http.pool=false \ + -Dmaven.wagon.httpconnectionManager.ttlSeconds=120 \ + -Dhttp.keepAlive=false \ + -Pspark-runner""" + + def cp = "target/classes:${deps.trim()}" + def jvmArgs = "--add-opens=java.base/sun.nio.ch=ALL-UNNAMED --add-opens=java.base/java.nio=ALL-UNNAMED --add-opens=java.base/java.util=ALL-UNNAMED --add-opens=java.base/java.lang.invoke=ALL-UNNAMED --add-opens=java.base/java.lang=ALL-UNNAMED" + t.run """mvn exec:exec -q -Dexec.executable=java \ + -Dexec.args="${jvmArgs} -cp '${cp}' org.apache.beam.examples.WordCount \ + --inputFile=pom.xml --output=counts --runner=SparkRunner" """ // Verify text from the pom.xml input file String result = t.run "grep Foundation counts*" diff --git a/sdks/python/apache_beam/transforms/combiners.py b/sdks/python/apache_beam/transforms/combiners.py index 8d35405f3fff..c45ba4e89b9a 100644 --- a/sdks/python/apache_beam/transforms/combiners.py +++ b/sdks/python/apache_beam/transforms/combiners.py @@ -597,6 +597,35 @@ def display_data(self): def default_label(self): return 'FixedSizePerKey(%d)' % self._n + @with_input_types(T) + @with_output_types(T) + class Any(ptransform.PTransform): + """Returns up to n arbitrary elements from the input PCollection. + + This is the Python equivalent of Java's ``Sample.any``. Unlike + ``FixedSizeGlobally`` it does not sample uniformly at random, and it returns + the selected elements rather than a single list. If the input has fewer than + n elements, all of them are returned. + """ + def __init__(self, n): + if n < 0: + raise ValueError('Expected non-negative n, received %s.' % n) + self._n = n + + def expand(self, pcoll): + return ( + pcoll + | core.CombineGlobally(_SampleAnyCombineFn( + self._n)).without_defaults() + | core.FlatMap(lambda elements: elements).with_input_types( + list[T]).with_output_types(T)) + + def display_data(self): + return {'n': self._n} + + def default_label(self): + return 'Any(%d)' % self._n + @with_input_types(T) @with_output_types(list[T]) @@ -636,6 +665,35 @@ def teardown(self): self._top_combiner.teardown() +@with_input_types(T) +@with_output_types(list[T]) +class _SampleAnyCombineFn(core.CombineFn): + """CombineFn that keeps up to n arbitrary elements (no random sampling).""" + def __init__(self, n): + super().__init__() + self._n = n + + def create_accumulator(self): + return [] + + def add_input(self, accumulator, element): + if len(accumulator) < self._n: + accumulator.append(element) + return accumulator + + def merge_accumulators(self, accumulators): + result = [] + for accumulator in accumulators: + for element in accumulator: + if len(result) >= self._n: + return result + result.append(element) + return result + + def extract_output(self, accumulator): + return accumulator + + class _TupleCombineFnBase(core.CombineFn): def __init__(self, *combiners, merge_accumulators_batch_size=None): self._combiners = [core.CombineFn.maybe_from_callable(c) for c in combiners] diff --git a/sdks/python/apache_beam/transforms/combiners_test.py b/sdks/python/apache_beam/transforms/combiners_test.py index a7f357719617..14348bb8ce78 100644 --- a/sdks/python/apache_beam/transforms/combiners_test.py +++ b/sdks/python/apache_beam/transforms/combiners_test.py @@ -253,6 +253,7 @@ def individual_test_per_key_dd(sampleFn, n): individual_test_per_key_dd(combine.Sample.FixedSizePerKey, 5) individual_test_per_key_dd(combine.Sample.FixedSizeGlobally, 5) + individual_test_per_key_dd(combine.Sample.Any, 5) def test_combine_globally_display_data(self): transform = beam.CombineGlobally(combine.Smallest(5)) @@ -359,6 +360,59 @@ def match(actual): assert_that(result, matcher()) + def test_sample_any(self): + with TestPipeline() as pipeline: + pcoll = pipeline | 'start' >> Create([1, 2, 3, 4, 5]) + result = pcoll | 'sample-any' >> combine.Sample.Any(3) + + def check(actual): + assert len(actual) == 3, actual + for element in actual: + assert element in [1, 2, 3, 4, 5], element + + assert_that(result, check) + + def test_sample_any_at_most_input_size(self): + with TestPipeline() as pipeline: + pcoll = pipeline | 'start' >> Create([1, 2]) + result = pcoll | 'sample-any' >> combine.Sample.Any(5) + assert_that(result, equal_to([1, 2])) + + def test_sample_any_windowed(self): + with TestPipeline() as pipeline: + pcoll = ( + pipeline + | 'start' >> Create([1, 2, 3, 4]) + | 'timestamp' >> Map(lambda x: TimestampedValue(x, x * 10)) + | 'window' >> WindowInto(FixedWindows(15))) + result = pcoll | 'sample-any' >> combine.Sample.Any(1) + + def check(actual): + # Timestamps 10, 20, 30, 40 fall into fixed windows [0, 15), [15, 30) + # and [30, 45), holding {1}, {2} and {3, 4}. One element is sampled from + # each window that has elements. + assert len(actual) == 3, actual + for element in actual: + assert element in [1, 2, 3, 4], element + + assert_that(result, check) + + def test_sample_any_empty(self): + with TestPipeline() as pipeline: + pcoll = pipeline | 'start' >> Create([]) + result = pcoll | 'sample-any' >> combine.Sample.Any(3) + assert_that(result, equal_to([])) + + def test_sample_any_zero(self): + with TestPipeline() as pipeline: + pcoll = pipeline | 'start' >> Create([1, 2, 3]) + result = pcoll | 'sample-any' >> combine.Sample.Any(0) + assert_that(result, equal_to([])) + + def test_sample_any_negative_n(self): + with self.assertRaises(ValueError): + combine.Sample.Any(-1) + def test_tuple_combine_fn(self): with TestPipeline() as p: result = (