You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何编写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. 创建并使用聚合函数

  1. 把编译好的Java类(包括IntStats和聚合函数类)打包成jar文件
  2. 将jar放到所有Cassandra节点的lib目录下
  3. 重启所有Cassandra服务
  4. 执行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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 09:18:03