如何编写Java聚合函数计算Cassandra中MAP列各键的MIN、MAX值?
没问题,我来带你一步步实现Cassandra中Map类型的自定义聚合函数,这确实是个实用的需求,咱们分两种情况来讲解——先处理整数类型的Map,再扩展到浮点数的统计计算。
实现Cassandra Map类型的自定义聚合函数
一、处理Map<String, Integer>的min/max聚合
1. 定义状态UDT和Java状态类
首先,我们需要一个UDT(用户定义类型)来存储每个键的min和max值,对应的Java状态类必须实现Serializable,因为Cassandra需要序列化聚合状态在节点间传递。
先创建Cassandra UDT:
CREATE TYPE int_stats (min int, max int);
对应的Java状态类:
import java.io.Serializable; public class IntStats implements Serializable { private int min; private int max; public IntStats(int min, int max) { this.min = min; this.max = max; } // 必须添加getter和setter方法,供聚合函数调用 public int getMin() { return min; } public void setMin(int min) { this.min = min; } public int getMax() { return max; } public void setMax(int max) { this.max = max; } }
2. 编写聚合函数类
Cassandra的自定义聚合函数(UDA)需要实现AggregateFunction接口,核心包含四个方法:newState(初始化空状态)、update(处理单行输入的Map)、merge(合并多节点的聚合状态)、finalize(将状态转换为你需要的输出格式)。
import com.datastax.oss.driver.api.core.type.DataTypes; import org.apache.cassandra.cql3.functions.AggregateFunction; import org.apache.cassandra.cql3.functions.FunctionParameter; import org.apache.cassandra.cql3.functions.UDAggregate; import java.util.HashMap; import java.util.List; import java.util.Map; import java.util.stream.Collectors; @UDAggregate( name = "map_int_min_max", inputTypes = { @FunctionParameter(type = DataTypes.MAP, parameters = {DataTypes.TEXT, DataTypes.INT}) }, stateType = @FunctionParameter(type = DataTypes.MAP, parameters = {DataTypes.TEXT, DataTypes.UDT}), udtName = "int_stats" ) public class MapIntMinMaxAggregate implements AggregateFunction<Map<String, IntStats>, Map<String, Integer>, List<Map<String, Map<String, Integer>>>> { // 初始化空的聚合状态 @Override public Map<String, IntStats> newState() { return new HashMap<>(); } // 处理单个Map,更新每个键的min/max值 @Override public Map<String, IntStats> update(Map<String, IntStats> state, Map<String, Integer> input) { if (input == null || input.isEmpty()) return state; for (Map.Entry<String, Integer> entry : input.entrySet()) { String key = entry.getKey(); int value = entry.getValue(); IntStats stats = state.get(key); if (stats == null) { // 首次遇到该键,直接用当前值初始化min和max state.put(key, new IntStats(value, value)); } else { // 已存在的键,更新为更小的min和更大的max stats.setMin(Math.min(stats.getMin(), value)); stats.setMax(Math.max(stats.getMax(), value)); } } return state; } // 合并两个节点的聚合状态,保证分布式场景下计算正确 @Override public Map<String, IntStats> merge(Map<String, IntStats> state1, Map<String, IntStats> state2) { if (state2.isEmpty()) return state1; if (state1.isEmpty()) return state2; for (Map.Entry<String, IntStats> entry : state2.entrySet()) { String key = entry.getKey(); IntStats stats2 = entry.getValue(); IntStats stats1 = state1.get(key); if (stats1 == null) { state1.put(key, stats2); } else { stats1.setMin(Math.min(stats1.getMin(), stats2.getMin())); stats1.setMax(Math.max(stats1.getMax(), stats2.getMax())); } } return state1; } // 将最终状态转换为你需要的输出格式 @Override public List<Map<String, Map<String, Integer>>> finalize(Map<String, IntStats> state) { return state.entrySet().stream() .map(entry -> { Map<String, Map<String, Integer>> resultEntry = new HashMap<>(); Map<String, Integer> statsMap = new HashMap<>(); statsMap.put("min", entry.getValue().getMin()); statsMap.put("max", entry.getValue().getMax()); resultEntry.put(entry.getKey(), statsMap); return resultEntry; }) .collect(Collectors.toList()); } }
3. 创建并使用聚合函数
- 把编译好的Java类(包括
IntStats和聚合函数类)打包成jar文件 - 将jar放到所有Cassandra节点的
lib目录下 - 重启所有Cassandra服务
- 执行CQL创建聚合函数:
CREATE AGGREGATE map_int_min_max(map<text, int>) SFUNC update STYPE map<text, int_stats> FINALFUNC finalize INITCOND {};
现在就可以用这个函数查询了:
SELECT map_int_min_max(elements) FROM your_table_name;
二、处理Map<String, Float>的min/max/avg聚合
这个需求只需要在上面的基础上扩展状态,增加sum和count字段来计算平均值。
1. 定义状态UDT和Java状态类
先创建UDT:
CREATE TYPE float_stats (min float, max float, sum float, count int);
对应的Java状态类:
import java.io.Serializable; public class FloatStats implements Serializable { private float min; private float max; private float sum; private int count; public FloatStats(float min, float max, float sum, int count) { this.min = min; this.max = max; this.sum = sum; this.count = count; } // Getter和setter方法 public float getMin() { return min; } public void setMin(float min) { this.min = min; } public float getMax() { return max; } public void setMax(float max) { this.max = max; } public float getSum() { return sum; } public void setSum(float sum) { this.sum = sum; } public int getCount() { return count; } public void setCount(int count) { this.count = count; } }
2. 编写聚合函数类
import com.datastax.oss.driver.api.core.type.DataTypes; import org.apache.cassandra.cql3.functions.AggregateFunction; import org.apache.cassandra.cql3.functions.FunctionParameter; import org.apache.cassandra.cql3.functions.UDAggregate; import java.util.HashMap; import java.util.List; import java.util.Map; import java.util.stream.Collectors; @UDAggregate( name = "map_float_min_max_avg", inputTypes = { @FunctionParameter(type = DataTypes.MAP, parameters = {DataTypes.TEXT, DataTypes.FLOAT}) }, stateType = @FunctionParameter(type = DataTypes.MAP, parameters = {DataTypes.TEXT, DataTypes.UDT}), udtName = "float_stats" ) public class MapFloatMinMaxAvgAggregate implements AggregateFunction<Map<String, FloatStats>, Map<String, Float>, List<Map<String, Map<String, Object>>>> { @Override public Map<String, FloatStats> newState() { return new HashMap<>(); } @Override public Map<String, FloatStats> update(Map<String, FloatStats> state, Map<String, Float> input) { if (input == null || input.isEmpty()) return state; for (Map.Entry<String, Float> entry : input.entrySet()) { String key = entry.getKey(); float value = entry.getValue(); FloatStats stats = state.get(key); if (stats == null) { state.put(key, new FloatStats(value, value, value, 1)); } else { stats.setMin(Math.min(stats.getMin(), value)); stats.setMax(Math.max(stats.getMax(), value)); stats.setSum(stats.getSum() + value); stats.setCount(stats.getCount() + 1); } } return state; } @Override public Map<String, FloatStats> merge(Map<String, FloatStats> state1, Map<String, FloatStats> state2) { if (state2.isEmpty()) return state1; if (state1.isEmpty()) return state2; for (Map.Entry<String, FloatStats> entry : state2.entrySet()) { String key = entry.getKey(); FloatStats stats2 = entry.getValue(); FloatStats stats1 = state1.get(key); if (stats1 == null) { state1.put(key, stats2); } else { stats1.setMin(Math.min(stats1.getMin(), stats2.getMin())); stats1.setMax(Math.max(stats1.getMax(), stats2.getMax())); stats1.setSum(stats1.getSum() + stats2.getSum()); stats1.setCount(stats1.getCount() + stats2.getCount()); } } return state1; } @Override public List<Map<String, Map<String, Object>>> finalize(Map<String, FloatStats> state) { return state.entrySet().stream() .map(entry -> { Map<String, Map<String, Object>> resultEntry = new HashMap<>(); Map<String, Object> statsMap = new HashMap<>(); FloatStats stats = entry.getValue(); statsMap.put("min", stats.getMin()); statsMap.put("max", stats.getMax()); // 避免除以0的异常,没有数据时avg设为null statsMap.put("avg", stats.getCount() > 0 ? stats.getSum() / stats.getCount() : null); resultEntry.put(entry.getKey(), statsMap); return resultEntry; }) .collect(Collectors.toList()); } }
3. 创建并使用聚合函数
同样按照之前的步骤打包jar、部署到所有节点、重启服务,然后创建聚合函数:
CREATE AGGREGATE map_float_min_max_avg(map<text, float>) SFUNC update STYPE map<text, float_stats> FINALFUNC finalize INITCOND {};
使用方式:
SELECT map_float_min_max_avg(elements) FROM your_table_name;
注意事项
- 所有Cassandra节点必须部署相同的jar包,否则会出现函数不存在的错误
- 如果你的Map存在null值,可以在
update方法里添加判断逻辑跳过null值 - UDT和聚合函数的名称必须和Java注解里的配置完全一致
内容的提问来源于stack exchange,提问作者Bony
相关产品推荐
相关产品推荐

