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

PostgreSQL C开发支持任意数量任意类型参数的自定义聚合函数咨询

需求说明

需要实现一款PostgreSQL自定义聚合函数,满足以下要求:

  • 调用方式:select f(rc, col1, col2,..., coln) from table
  • 功能逻辑:按col1到coln构成的全序取最大值对应的rc值,本质为折叠(fold)操作,状态转换函数伪代码如下:
f(_state, rc, col1, col2,..., coln)
{
  nargs = extract_variadic_args(fcinfo, 0, false, &args, &types, &nulls);

  for (i = 0; i < nargs; i++) {
    if (_state.coli > coli) return (_state.rc, _state.col1, ..., _state.colnargs)
    else if (_state.coli == coli) continue;
    else return (rc, col1, ..., colnargs);
  }

  return (rc, col1, ..., colnargs);
}

例如测试表T(rc text, col1 int, col2 real, col3 timestamp)调用该聚合后,会返回全序最高的行对应的rc值。

现有方案局限性
  • 原生PGSQL多态类型anyelement要求所有使用场景类型一致,无法直接支持任意数量、任意类型的参数输入
  • 现有PLPGSQL实现存在两个固有缺陷:
    • 状态类型(stype)结构强制与表行结构绑定,无法存储额外折叠状态
    • 状态类型仅支持单表列作为排序维度,无法适配多表join场景下跨表列排序的需求
核心问题拆解
  1. 如何设计可存储任意数量、任意类型值的自定义聚合状态类型,用于与入参排序列做比较
  2. 如何无需为所有rc类型写分支判断,直接通过PG_RETURN_DATUM返回运行时才能确定类型的rc值
C语言实现方案

核心设计思路

  1. 自定义变长状态类型,动态存储rc值、类型信息,以及所有排序列的值、类型信息、比较函数缓存,无需预定义结构,适配任意数量任意类型的排序参数
  2. 使用PostgreSQL内置的"any"伪类型作为可变参数类型,支持任意数量任意类型的输入参数;使用anynonarray多态类型作为返回值,最终直接返回存储的rc Datum即可,无需类型分支判断

完整代码实现

1. C扩展代码

#include "postgres.h"
#include "fmgr.h"
#include "utils/builtins.h"
#include "utils/lsyscache.h"
#include "utils/datum.h"
#include "access/tupmacs.h"

PG_MODULE_MAGIC;

/* 排序列条目结构 */
typedef struct SortColEntry {
    Oid         typid;
    int32       typmod;
    Oid         coll;
    Datum       val;
    bool        isnull;
    FmgrInfo    cmp_func;
} SortColEntry;

/* 自定义聚合状态结构(变长) */
typedef struct PriorityState {
    char        vl_len_[4];
    Oid         rc_typid;
    int32       rc_typmod;
    Datum       rc_val;
    bool        rc_isnull;
    int         n_sort_cols;
    SortColEntry sort_cols[FLEXIBLE_ARRAY_MEMBER];
} PriorityState;

/* 状态输入函数(仅占位,无需外部调用) */
PG_FUNCTION_INFO_V1(priority_state_in);
Datum
priority_state_in(PG_FUNCTION_ARGS)
{
    ereport(ERROR, (errcode(ERRCODE_FEATURE_NOT_SUPPORTED),
                    errmsg("priority_state cannot be input from string")));
    PG_RETURN_NULL();
}

/* 状态输出函数(仅占位,无需外部调用) */
PG_FUNCTION_INFO_V1(priority_state_out);
Datum
priority_state_out(PG_FUNCTION_ARGS)
{
    PG_RETURN_CSTRING(cstring_to_text("priority_state"));
}

/* 状态转换函数 */
PG_FUNCTION_INFO_V1(priority_sfunc);
Datum
priority_sfunc(PG_FUNCTION_ARGS)
{
    PriorityState *state;
    Datum rc_val;
    bool rc_isnull;
    Oid rc_typid;
    int32 rc_typmod;
    int n_sort_cols;
    int i;
    bool keep_old = false;

    /* 首次调用,状态为空,初始化新状态 */
    if (PG_ARGISNULL(0)) {
        rc_val = PG_GETARG_DATUM(1);
        rc_isnull = PG_ARGISNULL(1);
        rc_typid = get_fn_expr_argtype(fcinfo->flinfo, 1);
        rc_typmod = get_fn_expr_argtypmod(fcinfo->flinfo, 1);
        n_sort_cols = PG_NARGS() - 2;

        /* 分配状态内存 */
        state = (PriorityState *) palloc0(offsetof(PriorityState, sort_cols) + n_sort_cols * sizeof(SortColEntry));
        SET_VARSIZE(state, offsetof(PriorityState, sort_cols) + n_sort_cols * sizeof(SortColEntry));

        state->rc_typid = rc_typid;
        state->rc_typmod = rc_typmod;
        state->rc_val = rc_val;
        state->rc_isnull = rc_isnull;
        state->n_sort_cols = n_sort_cols;

        /* 存储所有排序列信息 */
        for (i = 0; i < n_sort_cols; i++) {
            Oid typid = get_fn_expr_argtype(fcinfo->flinfo, i + 2);
            int32 typmod = get_fn_expr_argtypmod(fcinfo->flinfo, i + 2);
            Oid coll = fcinfo->fncollation;
            Oid cmp_func_oid;

            state->sort_cols[i].typid = typid;
            state->sort_cols[i].typmod = typmod;
            state->sort_cols[i].coll = coll;
            state->sort_cols[i].val = PG_GETARG_DATUM(i + 2);
            state->sort_cols[i].isnull = PG_ARGISNULL(i + 2);

            /* 缓存排序比较函数 */
            get_sort_function_for_type(typid, false, &cmp_func_oid, NULL);
            fmgr_info(cmp_func_oid, &state->sort_cols[i].cmp_func);
        }

        PG_RETURN_POINTER(state);
    }

    /* 非首次调用,比较当前行与状态值 */
    state = (PriorityState *) PG_GETARG_POINTER(0);
    n_sort_cols = state->n_sort_cols;

    for (i = 0; i < n_sort_cols; i++) {
        Datum old_val = state->sort_cols[i].val;
        bool old_isnull = state->sort_cols[i].isnull;
        Datum new_val = PG_GETARG_DATUM(i + 2);
        bool new_isnull = PG_ARGISNULL(i + 2);
        int cmp_res;

        /* 空值默认按最小值处理,可根据需求调整逻辑 */
        if (old_isnull && new_isnull) continue;
        if (old_isnull) {
            keep_old = false;
            break;
        }
        if (new_isnull) {
            keep_old = true;
            break;
        }

        /* 调用对应类型的比较函数 */
        cmp_res = DatumGetInt32(FunctionCall2(&state->sort_cols[i].cmp_func, old_val, new_val));
        if (cmp_res > 0) {
            keep_old = true;
            break;
        } else if (cmp_res < 0) {
            keep_old = false;
            break;
        }
        /* 相等则继续比较下一列 */
    }

    /* 保留旧状态直接返回 */
    if (keep_old) PG_RETURN_POINTER(state);

    /* 更新状态为当前行值 */
    state->rc_val = PG_GETARG_DATUM(1);
    state->rc_isnull = PG_ARGISNULL(1);
    for (i = 0; i < n_sort_cols; i++) {
        state->sort_cols[i].val = PG_GETARG_DATUM(i + 2);
        state->sort_cols[i].isnull = PG_ARGISNULL(i + 2);
    }

    PG_RETURN_POINTER(state);
}

/* 最终结果返回函数 */
PG_FUNCTION_INFO_V1(priority_final);
Datum
priority_final(PG_FUNCTION_ARGS)
{
    PriorityState *state = (PriorityState *) PG_GETARG_POINTER(0);

    if (state->rc_isnull) PG_RETURN_NULL();
    PG_RETURN_DATUM(state->rc_val);
}

2. SQL注册语句

将C代码编译为扩展后,执行以下SQL注册聚合:

-- 自定义状态类型
CREATE TYPE priority_state;
CREATE OR REPLACE FUNCTION priority_state_in(cstring) RETURNS priority_state AS 'MODULE_PATHNAME' LANGUAGE C STRICT;
CREATE OR REPLACE FUNCTION priority_state_out(priority_state) RETURNS cstring AS 'MODULE_PATHNAME' LANGUAGE C STRICT;
CREATE TYPE priority_state (
    INPUT = priority_state_in,
    OUTPUT = priority_state_out,
    INTERNALLENGTH = VARIABLE,
    ALIGNMENT = double
);

-- 注册状态转换函数与最终函数
CREATE OR REPLACE FUNCTION priority_sfunc(priority_state, anynonarray, VARIADIC "any")
RETURNS priority_state AS 'MODULE_PATHNAME', 'priority_sfunc' LANGUAGE C STRICT;
CREATE OR REPLACE FUNCTION priority_final(priority_state) RETURNS anynonarray AS 'MODULE_PATHNAME', 'priority_final' LANGUAGE C STRICT;

-- 注册聚合函数,命名为max_by
CREATE AGGREGATE max_by(anynonarray, VARIADIC "any") (
    SFUNC = priority_sfunc,
    STYPE = priority_state,
    FINALFUNC = priority_final
);

3. 测试用例

-- 构造测试表
CREATE TABLE t(rc text, col1 int, col2 real, col3 timestamp);
INSERT INTO t VALUES 
('Andy', 3, 1.5, '2024-01-01'),
('Bob', 2, 2.5, '2024-02-01'),
('Charlie', 3, 2.0, '2023-12-31');

-- 调用聚合,按col1降序、col2降序、col3降序取最高行的rc值
SELECT max_by(rc, col1, col2, col3) FROM t;
-- 返回结果:Charlie

内容的提问来源于stack exchange,提问作者docjosh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 20:39:03