diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 52e4af61117..b7c88254c93 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -350,6 +350,7 @@ jobs: org.apache.comet.CometSetOpWithGroupBySuite org.apache.comet.CometSparkSessionExtensionsSuite org.apache.spark.CometPluginsSuite + org.apache.spark.CometTaskMemoryManagerSuite org.apache.spark.CometPluginsDefaultSuite org.apache.spark.CometPluginsNonOverrideSuite org.apache.spark.CometPluginsUnifiedModeOverrideSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 5095f8b5493..8210f91b7f9 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -166,6 +166,7 @@ jobs: org.apache.comet.CometSetOpWithGroupBySuite org.apache.comet.CometSparkSessionExtensionsSuite org.apache.spark.CometPluginsSuite + org.apache.spark.CometTaskMemoryManagerSuite org.apache.spark.CometPluginsDefaultSuite org.apache.spark.CometPluginsNonOverrideSuite org.apache.spark.CometPluginsUnifiedModeOverrideSuite diff --git a/spark/src/main/java/org/apache/spark/CometTaskMemoryManager.java b/spark/src/main/java/org/apache/spark/CometTaskMemoryManager.java index e15729fcad0..a0131cb8922 100644 --- a/spark/src/main/java/org/apache/spark/CometTaskMemoryManager.java +++ b/spark/src/main/java/org/apache/spark/CometTaskMemoryManager.java @@ -112,6 +112,12 @@ public long spill(long size, MemoryConsumer trigger) throws IOException { return 0; } + @Override + public long getUsed() { + // Native allocations call TaskMemoryManager directly, bypassing MemoryConsumer.used. + return CometTaskMemoryManager.this.used.get(); + } + @Override public String toString() { return String.format("NativeMemoryConsumer(id=%d)", id); diff --git a/spark/src/test/scala/org/apache/spark/CometTaskMemoryManagerSuite.scala b/spark/src/test/scala/org/apache/spark/CometTaskMemoryManagerSuite.scala new file mode 100644 index 00000000000..548ecfa4888 --- /dev/null +++ b/spark/src/test/scala/org/apache/spark/CometTaskMemoryManagerSuite.scala @@ -0,0 +1,77 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.spark + +import java.util.Properties + +import org.scalatest.funsuite.AnyFunSuite + +import org.apache.spark.executor.TaskMetrics +import org.apache.spark.memory.{MemoryConsumer, TaskMemoryManager, TestMemoryManager} + +class CometTaskMemoryManagerSuite extends AnyFunSuite { + + test("native memory usage is visible to Spark's memory consumer") { + val memoryManager = new TestMemoryManager(new SparkConf()) + memoryManager.limit(1024) + val taskMemoryManager = new TaskMemoryManager(memoryManager, 0L) + val taskContext = new TaskContextImpl( + stageId = 0, + stageAttemptNumber = 0, + partitionId = 0, + numPartitions = 1, + taskAttemptId = 0L, + attemptNumber = 0, + taskMemoryManager = taskMemoryManager, + localProperties = new Properties, + metricsSystem = null, + taskMetrics = TaskMetrics.empty, + cpus = 1, + resources = Map.empty) + + TaskContext.setTaskContext(taskContext) + try { + val manager = new CometTaskMemoryManager(1L, 0L) + val consumer = nativeMemoryConsumer(manager) + + assert(manager.getUsed == 0L) + assert(consumer.getUsed == 0L) + + assert(manager.acquireMemory(128L) == 128L) + assert(manager.getUsed == 128L) + assert(consumer.getUsed == 128L) + assert(taskMemoryManager.getMemoryConsumptionForThisTask == 128L) + + manager.releaseMemory(128L) + assert(manager.getUsed == 0L) + assert(consumer.getUsed == 0L) + assert(taskMemoryManager.getMemoryConsumptionForThisTask == 0L) + } finally { + taskMemoryManager.cleanUpAllAllocatedMemory() + TaskContext.unset() + } + } + + private def nativeMemoryConsumer(manager: CometTaskMemoryManager): MemoryConsumer = { + val field = classOf[CometTaskMemoryManager].getDeclaredField("nativeMemoryConsumer") + field.setAccessible(true) + field.get(manager).asInstanceOf[MemoryConsumer] + } +}