dbd/db2/statement.c

1
#include "dbd_db2.h"
2
 
3
#define BIND_BUFFER_SIZE    1024
4
 
5
static lua_push_type_t db2_to_lua_push(unsigned int db2_type, int len) {
6
    lua_push_type_t lua_type;
7
 
8
    if (len == SQL_NULL_DATA)
9
	return LUA_PUSH_NIL;
10
 
11
    switch(db2_type) {
12
    case SQL_SMALLINT:
13
    case SQL_INTEGER:
14
	lua_type = LUA_PUSH_INTEGER; 
15
	break;
16
    case SQL_DECIMAL:
17
	lua_type = LUA_PUSH_NUMBER;
18
	break;
19
    default:
20
        lua_type = LUA_PUSH_STRING;
21
    }
22
 
23
    return lua_type;
24
}
25
 
26
/*
27
 * num_affected_rows = statement:affected()
28
 */
29
static int statement_affected(lua_State *L) {
30
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_DB2_STATEMENT);
31
    SQLRETURN rc = SQL_SUCCESS;
32
    SQLINTEGER affected;
33
 
34
    if (!statement->stmt) {
35
        luaL_error(L, DBI_ERR_INVALID_STATEMENT);
36
    }
37
 
38
    rc = SQLRowCount(statement->stmt, &affected);
39
 
40
 
41
    lua_pushinteger(L, affected);
42
 
43
    return 1;
44
}
45
 
46
/*
47
 * success = statement:close()
48
 */
49
static int statement_close(lua_State *L) {
50
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_DB2_STATEMENT);
51
 
52
    if (statement->stmt) {
53
        SQLFreeHandle(SQL_HANDLE_STMT, statement->stmt);
54
 
55
	if (statement->resultset) {
56
	    free(statement->resultset);
57
	    statement->resultset = NULL;
58
	}
59
 
60
	if (statement->bind) {
61
	    int i;
62
 
63
	    for (i = 0; i < statement->num_result_columns; i++) {
64
		free(statement->bind[i].buffer);
65
	    }
66
 
67
	    free(statement->bind);
68
	    statement->bind = NULL;
69
	}
70
 
71
	statement->num_result_columns = 0;
72
    }
73
 
74
    return 0;    
75
}
76
 
77
/*
78
 *  column_names = statement:columns()
79
 */
80
static int statement_columns(lua_State *L) {
81
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_DB2_STATEMENT);
82
 
83
    int i;
84
    int d;
85
 
86
    SQLRETURN rc = SQL_SUCCESS;
87
 
88
    if (!statement->resultset || !statement->bind) {
89
	lua_pushnil(L);
90
	return 1;
91
    }
92
 
93
    d = 1; 
94
    lua_newtable(L);
95
    for (i = 0; i < statement->num_result_columns; i++) {
96
	const char *name = strlower(statement->resultset[i].name);
97
	LUA_PUSH_ARRAY_STRING(d, name);
98
    }
99
 
100
    return 1;
101
}
102
 
103
/*
104
 * success = statement:execute(...)
105
 */
106
static int statement_execute(lua_State *L) {
107
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_DB2_STATEMENT);
108
    int n = lua_gettop(L);
109
    int p;
110
    int i;
111
    int errflag = 0;
112
    const char *errstr = NULL;
113
    SQLRETURN rc = SQL_SUCCESS;
114
    unsigned char *buffer = NULL;
115
    int offset = 0;
116
    resultset_t *resultset = NULL; 
117
    bindparams_t *bind; /* variable to read the results */
118
    SQLSMALLINT num_params;
119
 
120
    SQLCHAR message[SQL_MAX_MESSAGE_LENGTH + 1];
121
    SQLCHAR sqlstate[SQL_SQLSTATE_SIZE + 1];
122
    SQLINTEGER sqlcode;
123
    SQLSMALLINT length;	
124
 
125
    if (!statement->stmt) {
126
	lua_pushboolean(L, 0);
127
	lua_pushstring(L, DBI_ERR_EXECUTE_INVALID);
128
	return 2;
129
    }
130
 
131
    rc = SQLNumParams(statement->stmt, &num_params);
132
    if (rc != SQL_SUCCESS) {
133
        SQLGetDiagRec(SQL_HANDLE_STMT, statement->stmt, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
134
 
135
        lua_pushboolean(L, 0);
136
        lua_pushfstring(L, DBI_ERR_PREP_STATEMENT, message);
137
        return 2;
138
    }
139
 
140
    if (num_params != n-1) {
141
        /*
142
	 * SQLExecute does not handle this condition,
143
 	 * and the client library will fill unset params
144
	 * with NULLs
145
	 */
146
	lua_pushboolean(L, 0);
147
        lua_pushfstring(L, DBI_ERR_PARAM_MISCOUNT, num_params, n-1);
148
	return 2;
149
    }
150
 
151
    if (num_params > 0) {
152
        buffer = (unsigned char *)malloc(sizeof(double) * num_params);
153
    }
154
 
155
    for (p = 2; p <= n; p++) {
156
	int i = p - 1;
157
	int type = lua_type(L, p);
158
	char err[64];
159
	const char *str = NULL;
160
	size_t len = 0;
161
	double *num;
162
	int *boolean;
163
	const static SQLLEN nullvalue = SQL_NULL_DATA;
164
 
165
	switch(type) {
166
	case LUA_TNIL:
167
	    rc = SQLBindParameter(statement->stmt, i, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, (SQLPOINTER)0, 0, (SQLPOINTER)&nullvalue);
168
	    errflag = rc != SQL_SUCCESS && rc != SQL_SUCCESS_WITH_INFO;
169
	    break;
170
	case LUA_TNUMBER:
171
	    num = (double *)(buffer + offset);
172
	    *num = lua_tonumber(L, p);
173
	    offset += sizeof(double);
174
	    rc = SQLBindParameter(statement->stmt, i, SQL_PARAM_INPUT, SQL_C_DOUBLE, SQL_DECIMAL, 10, 0, (SQLPOINTER)num, 0, NULL);
175
	    errflag = rc != SQL_SUCCESS && rc != SQL_SUCCESS_WITH_INFO;
176
	    break;
177
	case LUA_TSTRING:
178
	    str = lua_tolstring(L, p, &len);
179
	    rc = SQLBindParameter(statement->stmt, i, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 0, 0, (SQLPOINTER)str, len, NULL);
180
	    errflag = rc != SQL_SUCCESS && rc != SQL_SUCCESS_WITH_INFO;
181
	    break;
182
	case LUA_TBOOLEAN:
183
	    boolean = (int *)(buffer + offset);
184
	    *boolean = lua_toboolean(L, p);
185
	    offset += sizeof(int);
186
	    rc = SQLBindParameter(statement->stmt, i, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, (SQLPOINTER)boolean, len, NULL);
187
	    errflag = rc != SQL_SUCCESS && rc != SQL_SUCCESS_WITH_INFO;
188
	    break;
189
	default:
190
	    /*
191
	     * Unknown/unsupported value type
192
	     */
193
	    errflag = 1;
194
            snprintf(err, sizeof(err)-1, DBI_ERR_BINDING_TYPE_ERR, lua_typename(L, type));
195
            errstr = err;
196
	}
197
 
198
	if (errflag)
199
	    break;
200
    }
201
 
202
    if (errflag) {
203
        if (buffer) 
204
            free(buffer);
205
 
206
	lua_pushboolean(L, 0);
207
 
208
	if (errstr) {
209
	    lua_pushfstring(L, DBI_ERR_BINDING_PARAMS, errstr);
210
	} else {
211
	    SQLGetDiagRec(SQL_HANDLE_STMT, statement->stmt, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
212
 
213
	    lua_pushfstring(L, DBI_ERR_BINDING_PARAMS, message);
214
	}
215
    
216
	return 2;
217
    }
218
 
219
    rc = SQLExecute(statement->stmt);
220
    if (rc != SQL_SUCCESS) {
221
        if (buffer) 
222
            free(buffer);
223
 
224
	SQLGetDiagRec(SQL_HANDLE_STMT, statement->stmt, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
225
 
226
	lua_pushnil(L);
227
	lua_pushfstring(L, DBI_ERR_PREP_STATEMENT, message);
228
	return 2;
229
    }
230
 
231
    /* 
232
     * identify the number of output columns 
233
     */
234
    rc = SQLNumResultCols(statement->stmt, &statement->num_result_columns);
235
 
236
    if (statement->num_result_columns > 0) {
237
	resultset = (resultset_t *)malloc(sizeof(resultset_t) * statement->num_result_columns);
238
	memset(resultset, 0, sizeof(resultset_t) * statement->num_result_columns);
239
 
240
	bind = (bindparams_t *)malloc(sizeof(bindparams_t) * statement->num_result_columns);
241
	memset(bind, 0, sizeof(bindparams_t) * statement->num_result_columns);
242
 
243
	for (i = 0; i < statement->num_result_columns; i++) {
244
	    /* 
245
	     * return a set of attributes for a column 
246
	     */
247
	    rc = SQLDescribeCol(statement->stmt,
248
                        (SQLSMALLINT)(i + 1),
249
                        resultset[i].name,
250
                        sizeof(resultset[i].name),
251
                        &resultset[i].name_len,
252
                        &resultset[i].type,
253
                        &resultset[i].size,
254
                        &resultset[i].scale,
255
                        NULL);
256
 
257
	    if (rc != SQL_SUCCESS) {
258
                if (buffer) 
259
                    free(buffer);
260
 
261
		SQLGetDiagRec(SQL_HANDLE_STMT, statement->stmt, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
262
 
263
		lua_pushnil(L);
264
		lua_pushfstring(L, DBI_ERR_DESC_RESULT, message);
265
		return 2;
266
	    }
267
 
268
	    bind[i].buffer_len = resultset[i].size+1;
269
 
270
	    /* 
271
	     *allocate memory to bind a column 
272
	     */
273
	    bind[i].buffer = (SQLCHAR *)malloc((int)bind[i].buffer_len);
274
 
275
	    rc = SQLBindCol(statement->stmt,
276
                       (SQLSMALLINT)(i + 1),
277
                       SQL_C_CHAR,
278
                       bind[i].buffer,
279
                       bind[i].buffer_len,
280
                       &bind[i].len);
281
 
282
	    if (rc != SQL_SUCCESS) {
283
                if (buffer) 
284
                    free(buffer);
285
 
286
		SQLGetDiagRec(SQL_HANDLE_STMT, statement->stmt, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
287
 
288
		lua_pushnil(L);
289
		lua_pushfstring(L, DBI_ERR_ALLOC_RESULT, message);
290
		return 2;
291
	    }
292
	}
293
 
294
	statement->resultset = resultset;
295
	statement->bind = bind;
296
    }
297
 
298
    if (buffer) 
299
        free(buffer);
300
 
301
    lua_pushboolean(L, 1);
302
    return 1;
303
}
304
 
305
/*
306
 * must be called after an execute
307
 */
308
static int statement_fetch_impl(lua_State *L, statement_t *statement, int named_columns) {
309
    int i;
310
    int d;
311
 
312
    SQLCHAR message[SQL_MAX_MESSAGE_LENGTH + 1];
313
    SQLCHAR sqlstate[SQL_SQLSTATE_SIZE + 1];
314
    SQLINTEGER sqlcode;
315
    SQLSMALLINT length;
316
 
317
    SQLRETURN rc = SQL_SUCCESS;
318
 
319
    if (!statement->resultset || !statement->bind) {
320
	lua_pushnil(L);
321
	return 1;
322
    }
323
 
324
    /* fetch each row, and display */
325
    rc = SQLFetch(statement->stmt);
326
    if (rc == SQL_NO_DATA_FOUND) {
327
        SQLFreeStmt(statement->stmt, SQL_RESET_PARAMS);
328
	lua_pushnil(L);
329
	return 1;
330
    }
331
 
332
    if (rc != SQL_SUCCESS) {
333
        SQLGetDiagRec(SQL_HANDLE_STMT, statement->stmt, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
334
 
335
        luaL_error(L, DBI_ERR_FETCH_FAILED, message);
336
    }
337
 
338
    d = 1; 
339
    lua_newtable(L);
340
    for (i = 0; i < statement->num_result_columns; i++) {
341
	lua_push_type_t lua_push = db2_to_lua_push(statement->resultset[i].type, statement->bind[i].len);
342
	const char *name = strlower(statement->resultset[i].name);
343
	double val;
344
	char *value = (char *)statement->bind[i].buffer;
345
 
346
	switch (lua_push) {
347
	case LUA_PUSH_NIL:
348
	    if (named_columns) {
349
		LUA_PUSH_ATTRIB_NIL(name);
350
	    } else {
351
		LUA_PUSH_ARRAY_NIL(d);
352
	    }
353
	    break;
354
	case LUA_PUSH_INTEGER:
355
            if (named_columns) {
356
                LUA_PUSH_ATTRIB_INT(name, atoi(value));
357
            } else {
358
                LUA_PUSH_ARRAY_INT(d, atoi(value));
359
            }
360
	    break;
361
	case LUA_PUSH_NUMBER:
362
	    val = strtod(value, NULL);
363
 
364
            if (named_columns) {
365
		LUA_PUSH_ATTRIB_FLOAT(name, val);
366
            } else {
367
		LUA_PUSH_ARRAY_FLOAT(d, val);
368
            }
369
	    break;
370
	case LUA_PUSH_BOOLEAN:
371
            if (named_columns) {
372
                LUA_PUSH_ATTRIB_BOOL(name, atoi(value));
373
            } else {
374
                LUA_PUSH_ARRAY_BOOL(d, atoi(value));
375
            }
376
            break;	    
377
	case LUA_PUSH_STRING:
378
	    if (named_columns) {
379
		LUA_PUSH_ATTRIB_STRING(name, value);
380
	    } else {
381
		LUA_PUSH_ARRAY_STRING(d, value);
382
	    }    
383
	    break;
384
	default:
385
	    luaL_error(L, DBI_ERR_UNKNOWN_PUSH);
386
	}
387
    }
388
 
389
    return 1;    
390
}
391
 
392
 
393
static int next_iterator(lua_State *L) {
394
    statement_t *statement = (statement_t *)luaL_checkudata(L, lua_upvalueindex(1), DBD_DB2_STATEMENT);
395
    int named_columns = lua_toboolean(L, lua_upvalueindex(2));
396
 
397
    return statement_fetch_impl(L, statement, named_columns);
398
}
399
 
400
/*
401
 * table = statement:fetch(named_indexes)
402
 */
403
static int statement_fetch(lua_State *L) {
404
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_DB2_STATEMENT);
405
    int named_columns = lua_toboolean(L, 2);
406
 
407
    return statement_fetch_impl(L, statement, named_columns);
408
}
409
 
410
/*
411
 * num_rows = statement:rowcount()
412
 */
413
static int statement_rowcount(lua_State *L) {
414
    luaL_error(L, DBI_ERR_NOT_IMPLEMENTED, DBD_DB2_STATEMENT, "rowcount");
415
 
416
    return 0;
417
}
418
 
419
/*
420
 * iterfunc = statement:rows(named_indexes)
421
 */
422
static int statement_rows(lua_State *L) {
423
    if (lua_gettop(L) == 1) {
424
        lua_pushvalue(L, 1);
425
        lua_pushboolean(L, 0);
426
    } else {
427
        lua_pushvalue(L, 1);
428
        lua_pushboolean(L, lua_toboolean(L, 2));
429
    }
430
 
431
    lua_pushcclosure(L, next_iterator, 2);
432
    return 1;
433
}
434
 
435
/*
436
 * __gc
437
 */
438
static int statement_gc(lua_State *L) {
439
    /* always free the handle */
440
    statement_close(L);
441
 
442
    return 0;
443
}
444
 
445
/*
446
 * __tostring
447
 */
448
static int statement_tostring(lua_State *L) {
449
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_DB2_STATEMENT);
450
 
451
    lua_pushfstring(L, "%s: %p", DBD_DB2_STATEMENT, statement);
452
 
453
    return 1;
454
}
455
 
456
int dbd_db2_statement_create(lua_State *L, connection_t *conn, const char *sql_query) { 
457
    SQLRETURN rc = SQL_SUCCESS;
458
    statement_t *statement = NULL;
459
    SQLHANDLE stmt;
460
 
461
    SQLCHAR message[SQL_MAX_MESSAGE_LENGTH + 1];
462
    SQLCHAR sqlstate[SQL_SQLSTATE_SIZE + 1];
463
    SQLINTEGER sqlcode;
464
    SQLSMALLINT length;	
465
 
466
    rc = SQLAllocHandle(SQL_HANDLE_STMT, conn->db2, &stmt);
467
    if (rc != SQL_SUCCESS) {
468
	SQLGetDiagRec(SQL_HANDLE_DBC, conn->db2, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
469
 
470
        lua_pushnil(L);
471
        lua_pushfstring(L, DBI_ERR_ALLOC_STATEMENT, message);
472
        return 2;
473
    }
474
 
475
    /*
476
     * turn off deferred prepare
477
     * statements will be sent to the server at prepare time,
478
     * and therefore we can catch errors now rather 
479
     * than at execute time
480
     */
481
    rc = SQLSetStmtAttr(stmt,SQL_ATTR_DEFERRED_PREPARE,(SQLPOINTER)SQL_DEFERRED_PREPARE_OFF,0);
482
 
483
    rc = SQLPrepare(stmt, (SQLCHAR *)sql_query, SQL_NTS);
484
    if (rc != SQL_SUCCESS) {
485
	SQLGetDiagRec(SQL_HANDLE_STMT, stmt, 1, sqlstate, &sqlcode, message, SQL_MAX_MESSAGE_LENGTH + 1, &length);
486
 
487
	lua_pushnil(L);
488
	lua_pushfstring(L, DBI_ERR_PREP_STATEMENT, message);
489
	return 2;
490
    }
491
 
492
    statement = (statement_t *)lua_newuserdata(L, sizeof(statement_t));
493
    statement->stmt = stmt;
494
    statement->db2 = conn->db2;
495
    statement->resultset = NULL;
496
    statement->bind = NULL;
497
 
498
    luaL_getmetatable(L, DBD_DB2_STATEMENT);
499
    lua_setmetatable(L, -2);
500
 
501
    return 1;
502
} 
503
 
504
int dbd_db2_statement(lua_State *L) {
505
    static const luaL_Reg statement_methods[] = {
506
	{"affected", statement_affected},
507
	{"close", statement_close},
508
	{"columns", statement_columns},
509
	{"execute", statement_execute},
510
	{"fetch", statement_fetch},
511
	{"rowcount", statement_rowcount},
512
	{"rows", statement_rows},
513
	{NULL, NULL}
514
    };
515
 
516
    static const luaL_Reg statement_class_methods[] = {
517
	{NULL, NULL}
518
    };
519
 
520
    luaL_newmetatable(L, DBD_DB2_STATEMENT);
521
    luaL_register(L, 0, statement_methods);
522
    lua_pushvalue(L,-1);
523
    lua_setfield(L, -2, "__index");
524
 
525
    lua_pushcfunction(L, statement_gc);
526
    lua_setfield(L, -2, "__gc");
527
 
528
    lua_pushcfunction(L, statement_tostring);
529
    lua_setfield(L, -2, "__tostring");
530
 
531
    luaL_register(L, DBD_DB2_STATEMENT, statement_class_methods);
532
 
533
    return 1;    
534
}