Changeset: 80aa14b12a24 for MonetDB
URL: https://dev.monetdb.org/hg/MonetDB?cmd=changeset;node=80aa14b12a24
Modified Files:
        sql/backends/monet5/sql_rank.c
        sql/common/sql_types.c
        sql/common/sql_types.h
        sql/server/rel_select.c
        sql/test/analytics/Tests/analytics00.sql
        sql/test/analytics/Tests/analytics00.stable.out
        sql/test/analytics/Tests/analytics01.sql
        sql/test/analytics/Tests/analytics01.stable.out
Branch: analytics
Log Message:

Handle null values in window function parameters.


diffs (truncated from 599 to 300 lines):

diff --git a/sql/backends/monet5/sql_rank.c b/sql/backends/monet5/sql_rank.c
--- a/sql/backends/monet5/sql_rank.c
+++ b/sql/backends/monet5/sql_rank.c
@@ -720,10 +720,11 @@ SQLlast_value(Client cntxt, MalBlkPtr mb
 
 #define NTH_VALUE_SINGLE_IMP(TPE)                                              
                 \
        do {                                                                    
                    \
-               TPE val = *(TPE*) nth, *rres = (TPE*) res;                      
                        \
-               if(!is_##TPE##_nil(val) && val < 1)                             
                        \
+               TPE val = *(TPE*) VALget(nth), *toset;                          
                        \
+               if(!VALisnil(nth) && val < 1)                                   
                        \
                        throw(SQL, "sql.nth_value", SQLSTATE(42000) "nth_value 
must be greater than zero"); \
-               *rres = (is_##TPE##_nil(val) || val > 1) ? TPE##_nil : *(TPE*) 
in;                      \
+               toset = (VALisnil(nth) || val > 1) ? (TPE*) ATOMnilptr(tp1) : 
(TPE*) in;                \
+               VALset(res, tp1, toset);                                        
                        \
        } while(0);
 
 str
@@ -781,9 +782,9 @@ SQLnth_value(Client cntxt, MalBlkPtr mb,
                else
                        throw(SQL, "sql.nth_value", SQLSTATE(HY001) 
MAL_MALLOC_FAIL);
        } else {
-               ptr res = getArgReference_ptr(stk, pci, 0);
-               ptr in = getArgReference_ptr(stk, pci, 1);
-               ptr nth = getArgReference_ptr(stk, pci, 2);
+               ValRecord *res = &(stk)->stk[(pci)->argv[0]];
+               ValRecord *in = &(stk)->stk[(pci)->argv[1]];
+               ValRecord *nth = &(stk)->stk[(pci)->argv[2]];
 
                switch (tp2) {
                        case TYPE_bte:
@@ -826,7 +827,6 @@ do_lead_lag(Client cntxt, MalBlkPtr mb, 
        int tp1, tp2, tp3, base = 2;
        BUN l_value = 1;
        const void *restrict default_value;
-       size_t default_value_size = 0;
 
        (void)cntxt;
        if (pci->argc < 4 || pci->argc > 6)
@@ -867,15 +867,13 @@ do_lead_lag(Client cntxt, MalBlkPtr mb, 
                tp3 = getArgType(mb, pci, 3);
                if (isaBatType(tp3))
                        throw(SQL, op, SQLSTATE(42000) "%s third argument must 
a single atom", desc);
-               default_value = vin->val.pval;
-               default_value_size = vin->len;
+               default_value = VALget(vin);
                base = 4;
        } else {
                int tpe = tp1;
                if (isaBatType(tpe))
                        tpe = getBatType(tp1);
                default_value = ATOMnilptr(tpe);
-               default_value_size = ATOMlen(tpe, default_value);
        }
 
        assert(default_value); //default value must be set
@@ -916,12 +914,13 @@ do_lead_lag(Client cntxt, MalBlkPtr mb, 
                else
                        throw(SQL, op, SQLSTATE(HY001) MAL_MALLOC_FAIL);
        } else {
-               ptr res = getArgReference_ptr(stk, pci, 0);
+               ValRecord *res = &(stk)->stk[(pci)->argv[0]];
                ValRecord *vin = &(stk)->stk[(pci)->argv[1]];
                if(l_value == 0) {
-                       memcpy(res, vin->val.pval, vin->len);
+                       if(!VALcopy(res, vin))
+                               throw(SQL, op, SQLSTATE(HY001) MAL_MALLOC_FAIL);
                } else {
-                       memcpy(res, default_value, default_value_size);
+                       VALset(res, tp1, (ptr) default_value);
                }
        }
        return MAL_SUCCEED;
diff --git a/sql/common/sql_types.c b/sql/common/sql_types.c
--- a/sql/common/sql_types.c
+++ b/sql/common/sql_types.c
@@ -764,6 +764,29 @@ func_cmp(sql_allocator *sa, sql_func *f,
 }
 
 sql_subfunc *
+sql_find_func_by_name(sql_allocator *sa, sql_schema *s, const char *name, int 
nrargs, int type)
+{
+       if (s && s->funcs.set)
+               for (node *n=s->funcs.set->h; n; n = n->next) {
+                       sql_func *f = n->data;
+
+                       if ((f->type != type || !f->res || list_length(f->ops) 
!= nrargs))
+                               continue;
+                       if (strcmp(f->base.name, name) == 0)
+                               return sql_dup_subfunc(sa, f, NULL, NULL);
+               }
+       for (node *n=funcs->h; n; n = n->next) {
+               sql_func *f = n->data;
+
+               if ((f->type != type || !f->res || list_length(f->ops) != 
nrargs))
+                       continue;
+               if (strcmp(f->base.name, name) == 0)
+                       return sql_dup_subfunc(sa, f, NULL, NULL);
+       }
+       return NULL;
+}
+
+sql_subfunc *
 sql_find_func(sql_allocator *sa, sql_schema *s, const char *sqlfname, int 
nrargs, int type, sql_subfunc *prev)
 {
        sql_subfunc *fres;
diff --git a/sql/common/sql_types.h b/sql/common/sql_types.h
--- a/sql/common/sql_types.h
+++ b/sql/common/sql_types.h
@@ -96,6 +96,7 @@ extern sql_subaggr *sql_find_aggr(sql_al
 extern int subaggr_cmp( sql_subaggr *a1, sql_subaggr *a2);
 
 extern int subfunc_cmp( sql_subfunc *f1, sql_subfunc *f2);
+extern sql_subfunc *sql_find_func_by_name(sql_allocator *sa, sql_schema *s, 
const char *name, int nrargs, int type);
 extern sql_subfunc *sql_find_func(sql_allocator *sa, sql_schema *s, const char 
*name, int nrargs, int type, sql_subfunc *prev);
 extern list *sql_find_funcs(sql_allocator *sa, sql_schema *s, const char 
*name, int nrargs, int type);
 extern sql_subfunc *sql_bind_member(sql_allocator *sa, sql_schema *s, const 
char *name, sql_subtype *tp, int nrargs, sql_subfunc *prev);
diff --git a/sql/server/rel_select.c b/sql/server/rel_select.c
--- a/sql/server/rel_select.c
+++ b/sql/server/rel_select.c
@@ -3649,7 +3649,7 @@ static sql_exp *
                                }
                                if (a && list_length(nexps))  /* count(col) has 
|exps| != |nexps| */
                                        exps = nexps;
-                               }
+                       }
                } else {
                        sql_exp *l = exps->h->data, *ol = l;
                        sql_exp *r = exps->h->next->data, *or = r;
@@ -4559,8 +4559,10 @@ rel_rankop(mvc *sql, sql_rel **rel, symb
        fargs = sa_list(sql->sa);
        if (!aggr) { //rank function call
                dlist* dnn = window_function->data.lval->h->next->data.lval;
-
-               if(!dnn || (strcmp(s->base.name, "sys") == 0 && strcmp(aname, 
"ntile") == 0)) {
+               bool is_ntile = (strcmp(s->base.name, "sys") == 0 && 
strcmp(aname, "ntile") == 0),
+                        is_nth_value = (strcmp(s->base.name, "sys") == 0 && 
strcmp(aname, "nth_value") == 0);
+
+               if(!dnn || is_ntile) {
                        e = p->exps->h->data;
                        e = exp_column(sql->sa, exp_relname(e), exp_name(e), 
exp_subtype(e), exp_card(e), has_nil(e), is_intern(e));
                        append(fargs, e);
@@ -4569,7 +4571,31 @@ rel_rankop(mvc *sql, sql_rel **rel, symb
                        for(dnode *nn = dnn->h ; nn ; nn = nn->next) {
                                is_last = 0;
                                exp_kind ek = {type_value, card_column, FALSE};
-                               append(fargs, rel_value_exp2(sql, &p, 
nn->data.sym, f, ek, &is_last));
+                               e = rel_value_exp2(sql, &p, nn->data.sym, f, 
ek, &is_last);
+
+                               if(is_ntile) { /* ntile only has one argument 
and in null case this cast should be done */
+                                       sql_subtype *empty = 
sql_bind_localtype("void");
+                                       if(subtype_cmp(&(e->tpe), empty) == 0) {
+                                               sql_subtype *to = 
sql_bind_localtype("bte");
+                                               e = exp_convert(sql->sa, e, 
empty, to);
+                                       }
+                               } else if(is_nth_value && dnn->h && nn == 
dnn->h->next) { /* corner case for nth_value */
+                                       sql_subtype *empty = 
sql_bind_localtype("void");
+                                       if(subtype_cmp(&(e->tpe), empty) == 0) {
+                                               sql_exp *ep = p->exps->h->data;
+                                               e = exp_convert(sql->sa, e, 
empty, &(ep->tpe));
+                                       }
+                               }
+                               append(fargs, e);
+                       }
+                       if(is_nth_value && fargs->h) { /* another corner case 
for nth_value */
+                               /* TODO this triggers a warning for a bulk 
execution in the MAL layer that I cannot fix */
+                               sql_subtype *empty = sql_bind_localtype("void");
+                               e = fargs->h->data;
+                               if(subtype_cmp(&(e->tpe), empty) == 0) {
+                                       sql_exp *ep = p->exps->h->data;
+                                       fargs->h->data = exp_convert(sql->sa, 
e, empty, &(ep->tpe));
+                               }
                        }
                }
        } else { //aggregation function call
@@ -4593,9 +4619,10 @@ rel_rankop(mvc *sql, sql_rel **rel, symb
                                distinct = n->data.i_val;
                                /*
                                 * all aggregations implemented in a window 
have 1 and only 1 argument only, so for now no further
-                                * checking is needed
+                                * symbol compilation is required
                                 */
-                               append(fargs, rel_value_exp2(sql, &p, 
n->next->data.sym, f, ek1, &is_last));
+                               e = rel_value_exp2(sql, &p, n->next->data.sym, 
f, ek1, &is_last);
+                               append(fargs, e);
                                if(strcmp(s->base.name, "sys") == 0 && 
strcmp(aname, "count") == 0)
                                        append(fargs, exp_atom_bool(sql->sa, 
1)); //ignore nills
                        }
@@ -4655,19 +4682,39 @@ rel_rankop(mvc *sql, sql_rel **rel, symb
 
        if (!pe || !oe)
                return NULL;
-       types = sa_list(sql->sa);
-       for(node *nn = fargs->h ; nn ; nn = nn->next)
-               append(types, exp_subtype((sql_exp*) nn->data));
-       append(types, exp_subtype(pe));
-       append(types, exp_subtype(oe));
+
+       append(fargs, pe);
+       append(fargs, oe);
+       types = exp_types(sql->sa, fargs);
        wf = bind_func_(sql, s, aname, types, F_ANALYTIC);
-       if (!wf)
-               return sql_error(sql, 02, SQLSTATE(42000) "SELECT: function 
'%s' not found", aname );
+       if (!wf) {
+               wf = sql_find_func_by_name(sql->sa, NULL, aname, 
list_length(types), F_ANALYTIC);
+               if (wf) {
+                       node *op = wf->func->ops->h;
+                       list *nexps = sa_list(sql->sa);
+
+                       for (n = fargs->h ; wf && op && n; op = op->next, n = 
n->next ) {
+                               sql_arg *arg = op->data;
+                               e = n->data;
+
+                               e = rel_check_type(sql, &arg->type, e, 
type_equal);
+                               if (!e) {
+                                       wf = NULL;
+                                       break;
+                               }
+                               list_append(nexps, e);
+                       }
+                       if (wf && list_length(nexps))
+                               fargs = nexps;
+                       else
+                               return sql_error(sql, 02, SQLSTATE(42000) 
"SELECT: function '%s' not found", aname );
+               } else {
+                       return sql_error(sql, 02, SQLSTATE(42000) "SELECT: 
function '%s' not found", aname );
+               }
+       }
        args = sa_list(sql->sa);
        for(node *nn = fargs->h ; nn ; nn = nn->next)
                append(args, (sql_exp*) nn->data);
-       append(args, pe);
-       append(args, oe);
        if (fbe) {
                append(args, list_fetch(fbe, 0)); /*units */
                append(args, list_fetch(fbe, 1)); /*start */
diff --git a/sql/test/analytics/Tests/analytics00.sql 
b/sql/test/analytics/Tests/analytics00.sql
--- a/sql/test/analytics/Tests/analytics00.sql
+++ b/sql/test/analytics/Tests/analytics00.sql
@@ -75,6 +75,13 @@ select count(aa) over () from analytics;
 select count(*) over () from analytics;
 select count(*) over ();
 
+select min(null) over () from analytics;
+select max(null) over () from analytics;
+select cast(sum(null) over () as bigint) from analytics;
+select cast(prod(null) over () as bigint) from analytics;
+select avg(null) over () from analytics;
+select count(null) over () from analytics;
+
 create table stressme (aa varchar(64), bb int);
 insert into stressme values ('one', 1), ('another', 1), ('stress', 1), (NULL, 
2), ('ok', 2), ('check', 3), ('me', 3), ('please', 3), (NULL, 4);
 
diff --git a/sql/test/analytics/Tests/analytics00.stable.out 
b/sql/test/analytics/Tests/analytics00.stable.out
--- a/sql/test/analytics/Tests/analytics00.stable.out
+++ b/sql/test/analytics/Tests/analytics00.stable.out
@@ -944,6 +944,96 @@ Ready.
 % bigint # type
 % 1 # length
 [ 1    ]
+#select min(null) over () from analytics;
+% .L4 # table_name
+% L4 # name
+% char # type
+% 0 # length
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+#select max(null) over () from analytics;
+% .L4 # table_name
+% L4 # name
+% char # type
+% 0 # length
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
+#select cast(sum(null) over () as bigint) from analytics;
+% .L5 # table_name
+% L5 # name
+% bigint # type
+% 1 # length
+[ NULL ]
+[ NULL ]
+[ NULL ]
+[ NULL ]
_______________________________________________
checkin-list mailing list
[email protected]
https://www.monetdb.org/mailman/listinfo/checkin-list

Reply via email to