dbd/mysql/statement.c

1
#include "dbd_mysql.h"
2
 
3
static lua_push_type_t mysql_to_lua_push(unsigned int mysql_type) {
4
    lua_push_type_t lua_type;
5
 
6
    switch(mysql_type) {
7
    case MYSQL_TYPE_NULL:
8
	lua_type = LUA_PUSH_NIL;
9
	break;
10
 
11
    case MYSQL_TYPE_TINY:
12
    case MYSQL_TYPE_YEAR:
13
    case MYSQL_TYPE_SHORT:
14
    case MYSQL_TYPE_LONG:	
15
	lua_type =  LUA_PUSH_INTEGER;
16
	break;
17
 
18
    case MYSQL_TYPE_DOUBLE:
19
    case MYSQL_TYPE_LONGLONG:
20
	lua_type = LUA_PUSH_NUMBER;
21
	break;
22
 
23
    default:
24
	lua_type = LUA_PUSH_STRING;
25
    }
26
 
27
    return lua_type;
28
} 
29
 
30
static size_t mysql_buffer_size(MYSQL_FIELD *field) {
31
    unsigned int mysql_type = field->type;
32
    size_t size = 0;
33
    
34
    switch (mysql_type) {
35
	case MYSQL_TYPE_TINY:
36
	    size = 1;
37
	    break;
38
	case MYSQL_TYPE_YEAR:
39
	case MYSQL_TYPE_SHORT:
40
	    size = 2;
41
	    break;
42
	case MYSQL_TYPE_INT24:
43
	    size = 4;
44
	    break;
45
	case MYSQL_TYPE_LONG:
46
	    size = 4;
47
	    break;
48
	case MYSQL_TYPE_LONGLONG:
49
	    size = 8;
50
	    break;
51
	case MYSQL_TYPE_FLOAT:
52
	    size = 4;
53
	    break;
54
	case MYSQL_TYPE_DOUBLE:
55
	    size = 8;
56
	    break;
57
	case MYSQL_TYPE_TIME:
58
	case MYSQL_TYPE_DATE:
59
	case MYSQL_TYPE_DATETIME:
60
	case MYSQL_TYPE_TIMESTAMP:
61
	    size = sizeof(MYSQL_TIME);	
62
	    break;
63
	default:
64
	    size = field->length;
65
    }
66
 
67
    return size;
68
}
69
 
70
/*
71
 * num_affected_rows = statement:affected()
72
 */
73
static int statement_affected(lua_State *L) {
74
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_MYSQL_STATEMENT);
75
 
76
    if (!statement->stmt) {
77
        luaL_error(L, DBI_ERR_INVALID_STATEMENT);
78
    }
79
 
80
    lua_pushinteger(L, mysql_stmt_affected_rows(statement->stmt));
81
 
82
    return 1;
83
}
84
 
85
/*
86
 * success = statement:close()
87
 */
88
static int statement_close(lua_State *L) {
89
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_MYSQL_STATEMENT);
90
 
91
    if (statement->metadata) {
92
	mysql_free_result(statement->metadata);
93
	statement->metadata = NULL;
94
    }
95
 
96
    if (statement->stmt) {
97
	mysql_stmt_close(statement->stmt);
98
	statement->stmt = NULL;
99
    }
100
 
101
    lua_pushboolean(L, 1);
102
    return 1;    
103
}
104
 
105
/*
106
 * column_names = statement:columns()
107
 */
108
static int statement_columns(lua_State *L) {
109
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_MYSQL_STATEMENT);
110
 
111
    MYSQL_FIELD *fields;
112
    int i;
113
    int num_columns;
114
    int d = 1;
115
 
116
    if (!statement->stmt) {
117
        luaL_error(L, DBI_ERR_INVALID_STATEMENT);
118
        return 0;
119
    }
120
 
121
    fields = mysql_fetch_fields(statement->metadata);
122
    num_columns = mysql_num_fields(statement->metadata);
123
    lua_newtable(L);
124
    for (i = 0; i < num_columns; i++) {
125
	const char *name = fields[i].name;
126
 
127
        LUA_PUSH_ARRAY_STRING(d, name);
128
    }
129
 
130
    return 1;
131
}
132
 
133
 
134
/*
135
 * success,err = statement:execute(...)
136
 */
137
static int statement_execute(lua_State *L) {
138
    int n = lua_gettop(L);
139
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_MYSQL_STATEMENT); 
140
    int num_bind_params = n - 1;
141
    int expected_params;
142
 
143
    unsigned char *buffer = NULL;
144
    int offset = 0;
145
    
146
    MYSQL_BIND *bind = NULL;
147
    MYSQL_RES *metadata = NULL;
148
 
149
    char *error_message = NULL;
150
    char *errstr = NULL;
151
 
152
    int p;
153
 
154
    if (statement->metadata) {
155
	/*
156
	 * free existing metadata from any previous executions
157
         */
158
	mysql_free_result(statement->metadata);
159
	statement->metadata = NULL;
160
    }
161
 
162
    if (!statement->stmt) {
163
	lua_pushboolean(L, 0);
164
	lua_pushstring(L, DBI_ERR_EXECUTE_INVALID);
165
	return 2;
166
    }
167
 
168
    expected_params = mysql_stmt_param_count(statement->stmt);
169
 
170
    if (expected_params != num_bind_params) {
171
	/*
172
         * mysql_stmt_bind_param does not handle this condition,
173
         * and the client library will segfault if these do no match
174
         */ 
175
	lua_pushboolean(L, 0);
176
	lua_pushfstring(L, DBI_ERR_PARAM_MISCOUNT, expected_params, num_bind_params); 
177
	return 2;
178
    }
179
 
180
    if (num_bind_params > 0) {
181
        bind = malloc(sizeof(MYSQL_BIND) * num_bind_params);
182
        if (bind == NULL) {
183
            luaL_error(L, "Could not alloc bind params\n");
184
        }
185
 
186
        buffer = (unsigned char *)malloc(num_bind_params * sizeof(double));
187
        memset(bind, 0, sizeof(MYSQL_BIND) * num_bind_params);
188
    }
189
 
190
    for (p = 2; p <= n; p++) {
191
	int type = lua_type(L, p);
192
	int i = p - 2;
193
 
194
	const char *str = NULL;
195
	size_t *str_len = NULL;
196
	double *num = NULL;
197
	int *boolean = NULL;
198
	char err[64];
199
 
200
	switch(type) {
201
	    case LUA_TNIL:
202
		bind[i].buffer_type = MYSQL_TYPE_NULL;
203
		bind[i].is_null = (my_bool*)1;
204
		break;
205
 
206
	    case LUA_TBOOLEAN:
207
		boolean = (int *)(buffer + offset);
208
		offset += sizeof(int);
209
		*boolean = lua_toboolean(L, p);
210
 
211
		bind[i].buffer_type = MYSQL_TYPE_LONG;
212
		bind[i].is_null = (my_bool*)0;
213
		bind[i].buffer = (char *)boolean;
214
		bind[i].length = 0;
215
		break;
216
 
217
	    case LUA_TNUMBER:
218
		/*
219
		 * num needs to be it's own 
220
		 * memory here
221
                 */
222
		num = (double *)(buffer + offset);
223
		offset += sizeof(double);
224
		*num = lua_tonumber(L, p);
225
 
226
		bind[i].buffer_type = MYSQL_TYPE_DOUBLE;
227
		bind[i].is_null = (my_bool*)0;
228
		bind[i].buffer = (char *)num;
229
		bind[i].length = 0;
230
		break;
231
 
232
	    case LUA_TSTRING:
233
		str_len = (size_t *)(buffer + offset);
234
		offset += sizeof(size_t);
235
		str = lua_tolstring(L, p, str_len);
236
 
237
		bind[i].buffer_type = MYSQL_TYPE_STRING;
238
		bind[i].is_null = (my_bool*)0;
239
		bind[i].buffer = (char *)str;
240
		bind[i].length = str_len;
241
		break;
242
 
243
	    default:
244
		snprintf(err, sizeof(err)-1, DBI_ERR_BINDING_TYPE_ERR, lua_typename(L, type));
245
		errstr = err;
246
		error_message = DBI_ERR_BINDING_PARAMS;
247
		goto cleanup;
248
	}
249
    }
250
 
251
    if (mysql_stmt_bind_param(statement->stmt, bind)) {
252
	error_message = DBI_ERR_BINDING_PARAMS;
253
	goto cleanup;
254
    }
255
 
256
    if (mysql_stmt_execute(statement->stmt)) {
257
	error_message = DBI_ERR_BINDING_EXEC;
258
	goto cleanup;
259
    }
260
 
261
    metadata = mysql_stmt_result_metadata(statement->stmt);
262
 
263
    if (metadata) {
264
        mysql_stmt_store_result(statement->stmt);
265
    }
266
 
267
cleanup:
268
    if (bind) { 
269
	free(bind);
270
    }
271
 
272
    if (buffer) {
273
        free(buffer);
274
    }
275
 
276
    if (error_message) {
277
	lua_pushboolean(L, 0);
278
	lua_pushfstring(L, error_message, errstr ? errstr : mysql_stmt_error(statement->stmt));
279
	return 2;
280
    }
281
 
282
    statement->metadata = metadata;
283
 
284
    lua_pushboolean(L, 1);
285
    return 1;
286
}
287
 
288
static int statement_fetch_impl(lua_State *L, statement_t *statement, int named_columns) {
289
    int column_count, fetch_result_ok;
290
    MYSQL_BIND *bind = NULL;
291
    unsigned long *real_length = NULL;
292
    const char *error_message = NULL;
293
 
294
    if (!statement->stmt) {
295
	luaL_error(L, DBI_ERR_FETCH_INVALID);
296
	return 0;
297
    }
298
 
299
    if (!statement->metadata) {
300
	luaL_error(L, DBI_ERR_FETCH_NO_EXECUTE);
301
	return 0;
302
    }
303
 
304
    column_count = mysql_num_fields(statement->metadata);
305
 
306
    if (column_count > 0) {
307
	int i;
308
	MYSQL_FIELD *fields;
309
 
310
	real_length = calloc(column_count, sizeof(unsigned long));
311
 
312
        bind = malloc(sizeof(MYSQL_BIND) * column_count);
313
        memset(bind, 0, sizeof(MYSQL_BIND) * column_count);
314
 
315
	fields = mysql_fetch_fields(statement->metadata);
316
 
317
	for (i = 0; i < column_count; i++) {
318
	    unsigned int length = mysql_buffer_size(&fields[i]);
319
	    if (length > sizeof(MYSQL_TIME)) {
320
		bind[i].buffer = NULL;
321
		bind[i].buffer_length = 0;
322
	    } else {
323
		char *buffer = (char *)malloc(length);
324
		memset(buffer, 0, length);
325
 
326
		bind[i].buffer = buffer;
327
		bind[i].buffer_length = length;
328
	    }
329
 
330
	    bind[i].buffer_type = fields[i].type; 
331
	    bind[i].length = &real_length[i];
332
	}
333
 
334
	if (mysql_stmt_bind_result(statement->stmt, bind)) {
335
	    error_message = DBI_ERR_BINDING_RESULTS;
336
	    goto cleanup;
337
	}
338
 
339
	fetch_result_ok = mysql_stmt_fetch(statement->stmt);
340
	if (fetch_result_ok == 0 || fetch_result_ok == MYSQL_DATA_TRUNCATED) {
341
	    int d = 1;
342
 
343
	    lua_newtable(L);
344
	    for (i = 0; i < column_count; i++) {
345
		lua_push_type_t lua_push = mysql_to_lua_push(fields[i].type);
346
		const char *name = fields[i].name;
347
 
348
		if (bind[i].buffer == NULL) {
349
		    char *buffer = (char *)calloc(real_length[i]+1, sizeof(char));
350
		    bind[i].buffer = buffer;
351
		    bind[i].buffer_length = real_length[i];
352
		    mysql_stmt_fetch_column(statement->stmt, &bind[i], i, 0);
353
		}
354
 
355
		if (lua_push == LUA_PUSH_NIL) {
356
		    if (named_columns) {
357
			LUA_PUSH_ATTRIB_NIL(name);
358
		    } else {
359
			LUA_PUSH_ARRAY_NIL(d);
360
		    }
361
		} else if (lua_push == LUA_PUSH_INTEGER) {
362
		    if (fields[i].type == MYSQL_TYPE_YEAR || fields[i].type == MYSQL_TYPE_SHORT) {
363
			if (named_columns) {
364
			    LUA_PUSH_ATTRIB_INT(name, *(short *)(bind[i].buffer)); 
365
			} else {
366
			    LUA_PUSH_ARRAY_INT(d, *(short *)(bind[i].buffer)); 
367
			}
368
		    } else if (fields[i].type == MYSQL_TYPE_TINY) {
369
			if (named_columns) {
370
			    LUA_PUSH_ATTRIB_INT(name, (int)*(char *)(bind[i].buffer)); 
371
			} else {
372
			    LUA_PUSH_ARRAY_INT(d, (int)*(char *)(bind[i].buffer)); 
373
			}
374
		    } else {
375
			if (named_columns) {
376
			    LUA_PUSH_ATTRIB_INT(name, *(int *)(bind[i].buffer)); 
377
			} else {
378
			    LUA_PUSH_ARRAY_INT(d, *(int *)(bind[i].buffer)); 
379
			}
380
		    }
381
		} else if (lua_push == LUA_PUSH_NUMBER) {
382
		    if (named_columns) {
383
			LUA_PUSH_ATTRIB_FLOAT(name, *(double *)(bind[i].buffer));
384
		    } else {
385
			LUA_PUSH_ARRAY_FLOAT(d, *(double *)(bind[i].buffer));
386
		    }
387
		} else if (lua_push == LUA_PUSH_STRING) {
388
 
389
		    if (fields[i].type == MYSQL_TYPE_TIMESTAMP || fields[i].type == MYSQL_TYPE_DATETIME) {
390
			char str[20];
391
			struct st_mysql_time *t = bind[i].buffer;
392
 
393
			snprintf(str, 20, "%d-%02d-%02d %02d:%02d:%02d", t->year, t->month, t->day, t->hour, t->minute, t->second);
394
 
395
			if (named_columns) {
396
			    LUA_PUSH_ATTRIB_STRING(name, str);
397
			} else {
398
			    LUA_PUSH_ARRAY_STRING(d, str);
399
			}
400
		    } else if (fields[i].type == MYSQL_TYPE_TIME) {
401
			char str[9];
402
			struct st_mysql_time *t = bind[i].buffer;
403
 
404
			snprintf(str, 9, "%02d:%02d:%02d", t->hour, t->minute, t->second);
405
 
406
			if (named_columns) {
407
			    LUA_PUSH_ATTRIB_STRING(name, str);
408
			} else {
409
			    LUA_PUSH_ARRAY_STRING(d, str);
410
			}
411
		    } else if (fields[i].type == MYSQL_TYPE_DATE) {
412
			char str[20];
413
			struct st_mysql_time *t = bind[i].buffer;
414
 
415
			snprintf(str, 11, "%d-%02d-%02d", t->year, t->month, t->day);
416
 
417
			if (named_columns) {
418
			    LUA_PUSH_ATTRIB_STRING(name, str);
419
			} else {
420
			    LUA_PUSH_ARRAY_STRING(d, str);
421
			}
422
 
423
		    } else {
424
			if (named_columns) {
425
			    LUA_PUSH_ATTRIB_STRING(name, bind[i].buffer);
426
			} else {
427
			    LUA_PUSH_ARRAY_STRING(d, bind[i].buffer);
428
			}
429
		    }
430
		} else if (lua_push == LUA_PUSH_BOOLEAN) {
431
		    if (named_columns) {
432
			LUA_PUSH_ATTRIB_BOOL(name, *(int *)(bind[i].buffer));
433
		    } else {
434
			LUA_PUSH_ARRAY_BOOL(d, *(int *)(bind[i].buffer));
435
		    }
436
		} else {
437
		    luaL_error(L, DBI_ERR_UNKNOWN_PUSH);
438
		}
439
	    }
440
	} else {
441
	    lua_pushnil(L);	    
442
	}
443
    }
444
 
445
cleanup:
446
    free(real_length);
447
 
448
    if (bind) {
449
	int i;
450
 
451
	for (i = 0; i < column_count; i++) {
452
	    free(bind[i].buffer);
453
	}
454
 
455
	free(bind);
456
    }
457
 
458
    if (error_message) {
459
        luaL_error(L, error_message, mysql_stmt_error(statement->stmt));
460
        return 0;
461
    }
462
 
463
    return 1;    
464
}
465
 
466
static int next_iterator(lua_State *L) {
467
    statement_t *statement = (statement_t *)luaL_checkudata(L, lua_upvalueindex(1), DBD_MYSQL_STATEMENT);
468
    int named_columns = lua_toboolean(L, lua_upvalueindex(2));
469
 
470
    return statement_fetch_impl(L, statement, named_columns);
471
}
472
 
473
/*
474
 * table = statement:fetch(named_indexes)
475
 */
476
static int statement_fetch(lua_State *L) {
477
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_MYSQL_STATEMENT);
478
    int named_columns = lua_toboolean(L, 2);
479
 
480
    return statement_fetch_impl(L, statement, named_columns);
481
}
482
 
483
/*
484
 * num_rows = statement:rowcount()
485
 */
486
static int statement_rowcount(lua_State *L) {
487
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_MYSQL_STATEMENT);
488
 
489
    if (!statement->stmt) {
490
        luaL_error(L, DBI_ERR_INVALID_STATEMENT);
491
    }
492
 
493
    lua_pushinteger(L, mysql_stmt_num_rows(statement->stmt));
494
 
495
    return 1;
496
}
497
 
498
/*
499
 * iterfunc = statement:rows(named_indexes)
500
 */
501
static int statement_rows(lua_State *L) {
502
    if (lua_gettop(L) == 1) {			
503
	lua_pushvalue(L, 1);
504
	lua_pushboolean(L, 0);
505
    } else {
506
        lua_pushvalue(L, 1);
507
	lua_pushboolean(L, lua_toboolean(L, 2));
508
    }
509
 
510
    lua_pushcclosure(L, next_iterator, 2);
511
    return 1;
512
}
513
 
514
/*
515
 * __gc
516
 */
517
static int statement_gc(lua_State *L) {
518
    /* always free the handle */
519
    statement_close(L);
520
 
521
    return 0;
522
}
523
 
524
/*
525
 * __tostring
526
 */
527
static int statement_tostring(lua_State *L) {
528
    statement_t *statement = (statement_t *)luaL_checkudata(L, 1, DBD_MYSQL_STATEMENT);
529
 
530
    lua_pushfstring(L, "%s: %p", DBD_MYSQL_STATEMENT, statement);
531
 
532
    return 1;
533
}
534
 
535
int dbd_mysql_statement_create(lua_State *L, connection_t *conn, const char *sql_query) { 
536
    unsigned long sql_len = strlen(sql_query);
537
 
538
    statement_t *statement = NULL;
539
 
540
    MYSQL_STMT *stmt = mysql_stmt_init(conn->mysql);
541
 
542
    if (!stmt) {
543
	lua_pushnil(L);
544
	lua_pushfstring(L, DBI_ERR_ALLOC_STATEMENT, mysql_error(conn->mysql));
545
	return 2;
546
    }
547
 
548
    if (mysql_stmt_prepare(stmt, sql_query, sql_len)) {
549
	lua_pushnil(L);
550
	lua_pushfstring(L, DBI_ERR_PREP_STATEMENT, mysql_stmt_error(stmt));
551
	return 2;
552
    }
553
 
554
    statement = (statement_t *)lua_newuserdata(L, sizeof(statement_t));
555
    statement->mysql = conn->mysql;
556
    statement->stmt = stmt;
557
    statement->metadata = NULL;
558
 
559
    /*
560
    mysql_stmt_attr_set(stmt, STMT_ATTR_UPDATE_MAX_LENGTH, (my_bool*)0);
561
    */
562
 
563
    luaL_getmetatable(L, DBD_MYSQL_STATEMENT);
564
    lua_setmetatable(L, -2);
565
 
566
    return 1;
567
} 
568
 
569
int dbd_mysql_statement(lua_State *L) {
570
    static const luaL_Reg statement_methods[] = {
571
        {"affected", statement_affected},
572
	{"close", statement_close},
573
	{"columns", statement_columns},
574
	{"execute", statement_execute},
575
	{"fetch", statement_fetch},
576
	{"rowcount", statement_rowcount},
577
	{"rows", statement_rows},
578
	{NULL, NULL}
579
    };
580
 
581
    static const luaL_Reg statement_class_methods[] = {
582
	{NULL, NULL}
583
    };
584
 
585
    luaL_newmetatable(L, DBD_MYSQL_STATEMENT);
586
    luaL_register(L, 0, statement_methods);
587
    lua_pushvalue(L,-1);
588
    lua_setfield(L, -2, "__index");
589
 
590
    lua_pushcfunction(L, statement_gc);
591
    lua_setfield(L, -2, "__gc");
592
 
593
    lua_pushcfunction(L, statement_tostring);
594
    lua_setfield(L, -2, "__tostring");
595
 
596
    luaL_register(L, DBD_MYSQL_STATEMENT, statement_class_methods);
597
 
598
    return 1;    
599
}