From 940abe6dffc7254f07cc3712a782c81e478ee08b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=BA=91=E9=A3=8E?= Date: Tue, 30 Jul 2013 14:51:26 +0800 Subject: [PATCH] add mongo driver from lua-mongo project (wip) --- Makefile | 8 + lualib-src/lua-mongo.c | 540 +++++++++++++++++++++++++++++++++++++++++ lualib/mongo.lua | 286 ++++++++++++++++++++++ lualib/socket.lua | 7 + 4 files changed, 841 insertions(+) create mode 100644 lualib-src/lua-mongo.c create mode 100644 lualib/mongo.lua diff --git a/Makefile b/Makefile index f50cb8dc..aa031817 100644 --- a/Makefile +++ b/Makefile @@ -31,6 +31,8 @@ all : \ luaclib/socketbuffer.so \ luaclib/int64.so \ luaclib/mcast.so \ + luaclib/bson.so \ + luaclib/mongo.so \ client skynet : \ @@ -95,6 +97,12 @@ luaclib/int64.so : lua-int64/int64.c | luaclib luaclib/mcast.so : lualib-src/lua-localcast.c | luaclib gcc $(CFLAGS) $(SHARED) -Iluacompat $^ -o $@ -Iskynet-src -Iservice-src +luaclib/bson.so : lualib-src/lua-bson.c | luaclib + gcc $(CFLAGS) $(SHARED) -Iluacompat $^ -o $@ + +luaclib/mongo.so : lualib-src/lua-mongo.c | luaclib + gcc $(CFLAGS) $(SHARED) -Iluacompat $^ -o $@ -Iskynet-src + client : client-src/client.c gcc $(CFLAGS) $^ -o $@ -lpthread diff --git a/lualib-src/lua-mongo.c b/lualib-src/lua-mongo.c new file mode 100644 index 00000000..5045186b --- /dev/null +++ b/lualib-src/lua-mongo.c @@ -0,0 +1,540 @@ +#include +#include + +#include +#include +#include +#include + +#define OP_REPLY 1 +#define OP_MSG 1000 +#define OP_UPDATE 2001 +#define OP_INSERT 2002 +#define OP_QUERY 2004 +#define OP_GET_MORE 2005 +#define OP_DELETE 2006 +#define OP_KILL_CURSORS 2007 + +#define REPLY_CURSORNOTFOUND 1 +#define REPLY_QUERYFAILURE 2 +#define REPLY_AWAITCAPABLE 8 // ignore because mongo 1.6+ always set it + +#define DEFAULT_CAP 128 + +struct connection { + int sock; + int id; +}; + +struct response { + int flags; + int32_t cursor_id[2]; + int starting_from; + int number; +}; + +struct buffer { + int size; + int cap; + uint8_t * ptr; + uint8_t buffer[DEFAULT_CAP]; +}; + +static inline uint32_t +to_little_endian(uint32_t v) { + union { + uint32_t v; + uint8_t b[4]; + } u; + u.v = v; + return u.b[0] | u.b[1] << 8 | u.b[2] << 16 | u.b[3] << 24; +} + +typedef void * document; + +static inline uint32_t +get_length(const document buffer) { + union { + uint32_t v; + uint8_t b[4]; + } u; + memcpy(&u.v, buffer, 4); + return u.b[0] | u.b[1] << 8 | u.b[2] << 16 | u.b[3] << 24; +} + +static inline void +buffer_destroy(struct buffer *b) { + if (b->ptr != b->buffer) { + free(b->ptr); + } +} + +static inline void +buffer_create(struct buffer *b) { + b->size = 0; + b->cap = DEFAULT_CAP; + b->ptr = b->buffer; +} + +static inline void +buffer_reserve(struct buffer *b, int sz) { + if (b->size + sz <= b->cap) + return; + do { + b->cap *= 2; + } while (b->cap <= b->size + sz); + + if (b->ptr == b->buffer) { + b->ptr = malloc(b->cap); + memcpy(b->ptr, b->buffer, b->size); + } else { + b->ptr = realloc(b->ptr, b->cap); + } +} + +static inline void +write_int32(struct buffer *b, int32_t v) { + uint32_t uv = (uint32_t)v; + buffer_reserve(b,4); + b->ptr[b->size++] = uv & 0xff; + b->ptr[b->size++] = (uv >> 8)&0xff; + b->ptr[b->size++] = (uv >> 16)&0xff; + b->ptr[b->size++] = (uv >> 24)&0xff; +} + +static inline void +write_bytes(struct buffer *b, const void * buf, int sz) { + buffer_reserve(b,sz); + memcpy(b->ptr + b->size, buf, sz); + b->size += sz; +} + +static void +write_string(struct buffer *b, const char *key, size_t sz) { + buffer_reserve(b,sz+1); + memcpy(b->ptr + b->size, key, sz); + b->ptr[b->size+sz] = '\0'; + b->size+=sz+1; +} + +static inline int +reserve_length(struct buffer *b) { + int sz = b->size; + buffer_reserve(b,4); + b->size +=4; + return sz; +} + +static inline void +write_length(struct buffer *b, int32_t v, int off) { + uint32_t uv = (uint32_t)v; + b->ptr[off++] = uv & 0xff; + b->ptr[off++] = (uv >> 8)&0xff; + b->ptr[off++] = (uv >> 16)&0xff; + b->ptr[off++] = (uv >> 24)&0xff; +} + +// 1 integer id +// 2 integer flags +// 3 string collection name +// 4 integer skip +// 5 integer return number +// 6 document query +// 7 document selector (optional) +// return string package +static int +op_query(lua_State *L) { + int id = luaL_checkinteger(L,1); + document query = lua_touserdata(L,6); + if (query == NULL) { + return luaL_error(L, "require query document"); + } + document selector = lua_touserdata(L,7); + int flags = luaL_checkinteger(L, 2); + size_t nsz = 0; + const char *name = luaL_checklstring(L,3,&nsz); + int skip = luaL_checkinteger(L, 4); + int number = luaL_checkinteger(L, 5); + + luaL_Buffer b; + luaL_buffinit(L,&b); + + struct buffer buf; + buffer_create(&buf); + int len = reserve_length(&buf); + write_int32(&buf, id); + write_int32(&buf, 0); + write_int32(&buf, OP_QUERY); + write_int32(&buf, flags); + write_string(&buf, name, nsz); + write_int32(&buf, skip); + write_int32(&buf, number); + + int32_t query_len = get_length(query); + int total = buf.size + query_len; + int32_t selector_len = 0; + if (selector) { + selector_len = get_length(selector); + total += selector_len; + } + + write_length(&buf, total, len); + luaL_addlstring(&b, (const char *)buf.ptr, buf.size); + buffer_destroy(&buf); + + luaL_addlstring(&b, (const char *)query, query_len); + + if (selector) { + luaL_addlstring(&b, (const char *)selector, selector_len); + } + + luaL_pushresult(&b); + + return 1; +} + +// 1 string data +// 2 result document table +// return boolean succ (false -> request id, error document) +// number request_id +// document first +// string cursor_id +// integer startfrom +static int +op_reply(lua_State *L) { + size_t data_len = 0; + const char * data = luaL_checklstring(L,1,&data_len); + struct { +// int32_t length; // total message size, including this + int32_t request_id; // identifier for this message + int32_t response_id; // requestID from the original request + // (used in reponses from db) + int32_t opcode; // request type + int32_t flags; + int32_t cursor_id[2]; + int32_t starting; + int32_t number; + } const *reply = (const void *)data; + + if (data_len < sizeof(*reply)) { + lua_pushboolean(L, 0); + return 1; + } + + int id = to_little_endian(reply->response_id); + int flags = to_little_endian(reply->flags); + if (flags & REPLY_QUERYFAILURE) { + lua_pushboolean(L,0); + lua_pushinteger(L, id); + lua_pushlightuserdata(L, (void *)(reply+1)); + return 3; + } + + int starting_from = to_little_endian(reply->starting); + int number = to_little_endian(reply->number); + int sz = (int)data_len - sizeof(*reply); + const uint8_t * doc = (const uint8_t *)(reply+1); + + if (lua_istable(L,2)) { + int i = 1; + while (sz > 4) { + lua_pushlightuserdata(L, (void *)doc); + lua_rawseti(L, 2, i); + + int32_t doc_len = get_length((const document)doc); + + doc += doc_len; + sz -= doc_len; + + ++i; + } + if (i != number + 1) { + lua_pushboolean(L,0); + lua_pushinteger(L, id); + return 2; + } + int c = lua_rawlen(L, 2); + for (;i<=c;i++) { + lua_pushnil(L); + lua_rawseti(L, 2, i); + } + } + lua_pushboolean(L,1); + lua_pushinteger(L, id); + if (number == 0) + lua_pushnil(L); + else + lua_pushlightuserdata(L, (void *)(reply+1)); + if (reply->cursor_id[0] == 0 && reply->cursor_id[1]==0) { + // closed cursor + lua_pushnil(L); + } else { + lua_pushlstring(L, (const char *)(reply->cursor_id), 8); + } + lua_pushinteger(L, starting_from); + + return 5; +} + +/* + 1 string cursor_id + return string package + */ +static int +op_kill(lua_State *L) { + size_t cursor_len = 0; + const char * cursor_id = luaL_tolstring(L, 1, &cursor_len); + if (cursor_len != 8) { + return luaL_error(L, "Invalid cursor id"); + } + + struct buffer buf; + buffer_create(&buf); + + int len = reserve_length(&buf); + write_int32(&buf, 0); + write_int32(&buf, 0); + write_int32(&buf, OP_KILL_CURSORS); + + write_int32(&buf, 0); + write_int32(&buf, 1); + write_bytes(&buf, cursor_id, 8); + + write_length(&buf, buf.size, len); + + lua_pushlstring(L, (const char *)buf.ptr, buf.size); + buffer_destroy(&buf); + + return 1; +} + +/* + 1 string collection + 2 integer single remove + 3 document selector + + return string package + */ +static int +op_delete(lua_State *L) { + document selector = lua_touserdata(L,3); + if (selector == NULL) { + luaL_error(L, "Invalid param"); + } + size_t sz = 0; + const char * name = luaL_checklstring(L,1,&sz); + + luaL_Buffer b; + luaL_buffinit(L,&b); + + struct buffer buf; + buffer_create(&buf); + int len = reserve_length(&buf); + write_int32(&buf, 0); + write_int32(&buf, 0); + write_int32(&buf, OP_DELETE); + write_int32(&buf, 0); + write_string(&buf, name, sz); + write_int32(&buf, lua_tointeger(L,2)); + + int32_t selector_len = get_length(selector); + int total = buf.size + selector_len; + write_length(&buf, total, len); + + luaL_addlstring(&b, (const char *)buf.ptr, buf.size); + buffer_destroy(&buf); + + luaL_addlstring(&b, (const char *)selector, selector_len); + luaL_pushresult(&b); + + return 0; +} + +/* + 1 integer id + 2 string collection + 3 integer number + 4 cursor_id (8 bytes string/ 64bit) + + return string package + */ +static int +op_get_more(lua_State *L) { + int id = luaL_checkinteger(L, 1); + size_t sz = 0; + const char * name = luaL_checklstring(L,2,&sz); + int number = luaL_checkinteger(L, 3); + size_t cursor_len = 0; + const char * cursor_id = luaL_tolstring(L, 4, &cursor_len); + if (cursor_len != 8) { + return luaL_error(L, "Invalid cursor id"); + } + + struct buffer buf; + buffer_create(&buf); + int len = reserve_length(&buf); + write_int32(&buf, id); + write_int32(&buf, 0); + write_int32(&buf, OP_GET_MORE); + write_int32(&buf, 0); + write_string(&buf, name, sz); + write_int32(&buf, number); + write_bytes(&buf, cursor_id, 8); + write_length(&buf, buf.size, len); + + lua_pushlstring(L, (const char *)buf.ptr, buf.size); + buffer_destroy(&buf); + + return 1; +} + +// 1 string collection +// 2 integer flags +// 3 document selector +// 4 document update +// return string package +static int +op_update(lua_State *L) { + document selector = lua_touserdata(L,3); + document update = lua_touserdata(L,4); + if (selector == NULL || update == NULL) { + luaL_error(L, "Invalid param"); + } + size_t sz = 0; + const char * name = luaL_checklstring(L,1,&sz); + + luaL_Buffer b; + luaL_buffinit(L, &b); + + struct buffer buf; + buffer_create(&buf); + // make package header, don't raise L error + int len = reserve_length(&buf); + write_int32(&buf, 0); + write_int32(&buf, 0); + write_int32(&buf, OP_UPDATE); + write_int32(&buf, 0); + write_string(&buf, name, sz); + write_int32(&buf, lua_tointeger(L,2)); + + int32_t selector_len = get_length(selector); + int32_t update_len = get_length(update); + + int total = buf.size + selector_len + update_len; + write_length(&buf, total, len); + + luaL_addlstring(&b, (const char *)buf.ptr, buf.size); + buffer_destroy(&buf); + + luaL_addlstring(&b, (const char *)selector, selector_len); + luaL_addlstring(&b, (const char *)update, update_len); + + luaL_pushresult(&b); + + return 1; +} + +static int +document_length(lua_State *L) { + if (lua_isuserdata(L, 3)) { + document doc = lua_touserdata(L,3); + return get_length(doc); + } + if (lua_istable(L,3)) { + int total = 0; + int s = lua_rawlen(L,3); + int i; + for (i=1;i<=s;i++) { + lua_rawgeti(L, 3, i); + document doc = lua_touserdata(L,-1); + if (doc == NULL) { + lua_pop(L,1); + return luaL_error(L, "Invalid document at %d", i); + } else { + total += get_length(doc); + lua_pop(L,1); + } + } + return total; + } + return luaL_error(L, "Insert need documents"); +} + +// 1 integer flags +// 2 string collection +// 3 documents +// return string package +static int +op_insert(lua_State *L) { + size_t sz = 0; + const char * name = luaL_checklstring(L,2,&sz); + int dsz = document_length(L); + + luaL_Buffer b; + luaL_buffinit(L, &b); + + struct buffer buf; + buffer_create(&buf); + // make package header, don't raise L error + int len = reserve_length(&buf); + write_int32(&buf, 0); + write_int32(&buf, 0); + write_int32(&buf, OP_INSERT); + write_int32(&buf, lua_tointeger(L,1)); + write_string(&buf, name, sz); + + int total = buf.size + dsz; + write_length(&buf, total, len); + + luaL_addlstring(&b, (const char *)buf.ptr, buf.size); + buffer_destroy(&buf); + + if (lua_isuserdata(L,3)) { + document doc = lua_touserdata(L,3); + luaL_addlstring(&b, (const char *)doc, get_length(doc)); + } else { + int s = lua_rawlen(L, 3); + int i; + for (i=1;i<=s;i++) { + lua_rawgeti(L,3,i); + document doc = lua_touserdata(L,3); + luaL_addlstring(&b, (const char *)doc, get_length(doc)); + lua_pop(L,1); + } + } + + luaL_pushresult(&b); + + return 1; +} + +// string 4 bytes length +// return integer +static int +reply_length(lua_State *L) { + const char * rawlen_str = luaL_checkstring(L, 1); + int rawlen = 0; + memcpy(&rawlen, rawlen_str, sizeof(int)); + int length = to_little_endian(rawlen); + lua_pushinteger(L, length - 4); + return 1; +} + +int +luaopen_mongo_driver(lua_State *L) { + luaL_checkversion(L); + luaL_Reg l[] ={ + { "query", op_query }, + { "reply", op_reply }, + { "kill", op_kill }, + { "delete", op_delete }, + { "more", op_get_more }, + { "update", op_update }, + { "insert", op_insert }, + { "length", reply_length }, + { NULL, NULL }, + }; + + luaL_newlib(L,l); + return 1; +} diff --git a/lualib/mongo.lua b/lualib/mongo.lua new file mode 100644 index 00000000..9060f3b6 --- /dev/null +++ b/lualib/mongo.lua @@ -0,0 +1,286 @@ +local bson = require "bson" +local socket = require "socket" +local driver = require "mongo.driver" +local rawget = rawget +local assert = assert + +local bson_encode = bson.encode +local bson_decode = bson.decode +local empty_bson = bson_encode {} + +local mongo = {} +mongo.null = assert(bson.null) +mongo.maxkey = assert(bson.maxkey) +mongo.minkey = assert(bson.minkey) +mongo.type = assert(bson.type) + +local mongo_cursor = {} +local cursor_meta = { + __index = mongo_cursor, +} + +local mongo_client = {} + +local client_meta = { + __index = function(self, key) + return rawget(mongo_client, key) or self:getDB(key) + end, + __tostring = function (self) + local port_string + if self.port then + port_string = ":" .. tostring(self.port) + else + port_string = "" + end + + return "[mongo client : " .. self.host .. port_string .."]" + end, + __gc = function(self) + self:disconnect() + end +} + +local mongo_db = {} + +local db_meta = { + __index = function (self, key) + return rawget(mongo_db, key) or self:getCollection(key) + end, + __tostring = function (self) + return "[mongo db : " .. self.name .. "]" + end +} + +local mongo_collection = {} +local collection_meta = { + __index = function(self, key) + return rawget(mongo_collection, key) or self:getCollection(key) + end , + __tostring = function (self) + return "[mongo collection : " .. self.full_name .. "]" + end +} + +function mongo.client( obj ) + obj.port = obj.port or 27017 + obj.__id = 0 + obj.__sock = assert(socket.open(obj.host, obj.port),"Connect failed") + return setmetatable(obj, client_meta) +end + +function mongo_client:getDB(dbname) + local db = { + connection = self, + name = dbname, + full_name = dbname, + database = false, + __cmd = dbname .. "." .. "$cmd", + } + db.database = db + + return setmetatable(db, db_meta) +end + +function mongo_client:disconnect() + if self.__sock then + socket.close(self.__sock) + self.__sock = nil + end +end + +function mongo_client:genId() + local id = self.__id + 1 + self.__id = id + return id +end + +function mongo_client:runCommand(cmd) + if not self.admin then + self.admin = self:getDB "admin" + end + return self.admin:runCommand(cmd) +end + +local function get_reply(sock, result) + local length = driver.length(socket.read(sock, 4)) + local reply = socket.read(sock, length) + return reply, driver.reply(reply, result) +end + +function mongo_db:runCommand(cmd) + local request_id = self.connection:genId() + local sock = self.connection.__sock + socket.lock(sock) + local pack = driver.query(request_id, 0, self.__cmd, 0, 1, bson_encode(cmd)) + -- todo: check send + socket.write(sock, pack) + + local _, succ, reply_id, doc = get_reply(sock) + socket.unlock(sock) + assert(request_id == reply_id, "Reply from mongod error") + -- todo: check succ + return bson_decode(doc) +end + +function mongo_db:getCollection(collection) + local col = { + connection = self.connection, + name = collection, + full_name = self.full_name .. "." .. collection, + database = self.database, + } + self[collection] = setmetatable(col, collection_meta) + return col +end + +mongo_collection.getCollection = mongo_db.getCollection + +function mongo_collection:insert(doc) + if doc._id == nil then + doc._id = bson.objectid() + end + local sock = self.connection.__sock + local pack = driver.insert(0, self.full_name, bson_encode(doc)) + -- todo: check send + -- flags support 1: ContinueOnError + socket.write(sock, pack) +end + +function mongo_collection:batch_insert(docs) + for i=1,#docs do + if docs[i]._id == nil then + docs[i]._id = bson.objectid() + end + docs[i] = bson_encode(docs[i]) + end + local sock = self.connection.__sock + local pack = driver.insert(0, self.full_name, docs) + -- todo: check send + socket.write(sock, pack) +end + +function mongo_collection:update(selector,update,upsert,multi) + local flags = (upsert and 1 or 0) + (multi and 2 or 0) + local sock = self.connection.__sock + local pack = driver.update(self.full_name, flags, bson_encode(selector), bson_encode(update)) + -- todo: check send + socket.write(sock, pack) +end + +function mongo_collection:delete(selector, single) + local sock = self.connection.__sock + local pack = driver.delete(self.full_name, single, bson_encode(selector)) + -- todo: check send + socket.write(sock, pack) +end + +function mongo_collection:findOne(query, selector) + local request_id = self.connection:genId() + local sock = self.connection.__sock + socket.lock(sock) + local pack = driver.query(request_id, 0, self.full_name, 0, 1, query and bson_encode(query) or empty_bson, selector and bson_encode(selector)) + + -- todo: check send + socket.write(sock, pack) + + local _, succ, reply_id, doc = get_reply(sock) + socket.unlock(sock) + assert(request_id == reply_id, "Reply from mongod error") + -- todo: check succ + return bson_decode(doc) +end + +function mongo_collection:find(query, selector) + return setmetatable( { + __collection = self, + __query = query and bson_encode(query) or empty_bson, + __selector = selector and bson_encode(selector), + __ptr = nil, + __data = nil, + __cursor = nil, + __document = {}, + __flags = 0, + } , cursor_meta) +end + +function mongo_cursor:hasNext() + if self.__ptr == nil then + if self.__document == nil then + return false + end + local conn = self.__collection.connection + local request_id = conn:genId() + local sock = conn.__sock + local pack + if self.__data == nil then + pack = driver.query(request_id, self.__flags, self.__collection.full_name,0,0,self.__query,self.__selector) + else + if self.__cursor then + pack = driver.more(request_id, self.__collection.full_name,0,self.__cursor) + else + -- no more + self.__document = nil + self.__data = nil + return false + end + end + + socket.lock(sock) + --todo: check send + socket.write(sock, pack) + + local data, succ, reply_id, doc, cursor = get_reply(sock, self.__document) + socket.unlock(sock) + assert(request_id == reply_id, "Reply from mongod error") + if succ then + if doc then + self.__data = data + self.__ptr = 1 + self.__cursor = cursor + return true + else + self.__document = nil + self.__data = nil + self.__cursor = nil + return false + end + else + self.__document = nil + self.__data = nil + self.__cursor = nil + if doc then + local err = bson_decode(doc) + error(err["$err"]) + else + error("Reply from mongod error") + end + end + end + + return true +end + +function mongo_cursor:next() + if self.__ptr == nil then + error "Call hasNext first" + end + local r = bson_decode(self.__document[self.__ptr]) + self.__ptr = self.__ptr + 1 + if self.__ptr > #self.__document then + self.__ptr = nil + end + + return r +end + +function mongo_cursor:close() + -- todo: warning hasNext after close + if self.__cursor then + local sock = self.__collection.connection.__sock + local pack = driver.kill(self.__cursor) + -- todo: check send + socket.write(sock, pack) + end +end + +return mongo \ No newline at end of file diff --git a/lualib/socket.lua b/lualib/socket.lua index 7965ac2b..2e242591 100644 --- a/lualib/socket.lua +++ b/lualib/socket.lua @@ -226,7 +226,14 @@ function socket.write(fd, msg, sz) return true end +function socket.invalid(fd) + return CLOSED[fd] or not READBUF[fd] +end + function socket.lock(fd) + if CLOSED[fd] or not READBUF[fd] then + return + end local locked = READTHREAD[fd] if locked then -- lock fd