diff --git a/Makefile b/Makefile index 1bcb3603..ed8d1c7d 100644 --- a/Makefile +++ b/Makefile @@ -43,7 +43,7 @@ jemalloc : $(MALLOC_STATICLIB) CSERVICE = snlua logger gate harbor LUA_CLIB = skynet socketdriver int64 bson mongo md5 netpack \ cjson clientsocket memory profile multicast \ - cluster crypt sharedata + cluster crypt sharedata stm SKYNET_SRC = skynet_main.c skynet_handle.c skynet_module.c skynet_mq.c \ skynet_server.c skynet_start.c skynet_timer.c skynet_error.c \ @@ -116,6 +116,9 @@ $(LUA_CLIB_PATH)/crypt.so : lualib-src/lua-crypt.c | $(LUA_CLIB_PATH) $(LUA_CLIB_PATH)/sharedata.so : lualib-src/lua-sharedata.c | $(LUA_CLIB_PATH) $(CC) $(CFLAGS) $(SHARED) $^ -o $@ +$(LUA_CLIB_PATH)/stm.so : lualib-src/lua-stm.c | $(LUA_CLIB_PATH) + $(CC) $(CFLAGS) $(SHARED) -Iskynet-src $^ -o $@ + clean : rm -f $(SKYNET_BUILD_PATH)/skynet $(CSERVICE_PATH)/*.so $(LUA_CLIB_PATH)/*.so diff --git a/lualib-src/lua-stm.c b/lualib-src/lua-stm.c new file mode 100644 index 00000000..afe26f07 --- /dev/null +++ b/lualib-src/lua-stm.c @@ -0,0 +1,244 @@ +#include +#include +#include +#include +#include + +#include "rwlock.h" +#include "skynet_malloc.h" + +struct stm_object { + struct rwlock lock; + int reference; + struct stm_copy * copy; +}; + +struct stm_copy { + int reference; + uint32_t sz; + void * msg; +}; + +// msg should alloc by skynet_malloc +static struct stm_copy * +stm_newcopy(void * msg, int32_t sz) { + struct stm_copy * copy = skynet_malloc(sizeof(*copy)); + copy->reference = 1; + copy->sz = sz; + copy->msg = msg; + + return copy; +} + +static struct stm_object * +stm_new(void * msg, int32_t sz) { + struct stm_object * obj = skynet_malloc(sizeof(*obj)); + rwlock_init(&obj->lock); + obj->reference = 1; + obj->copy = stm_newcopy(msg, sz); + + return obj; +} + +static void +stm_releasecopy(struct stm_copy *copy) { + if (copy == NULL) + return; + if (__sync_sub_and_fetch(©->reference, 1) == 0) { + skynet_free(copy->msg); + skynet_free(copy); + } +} + +static void +stm_release(struct stm_object *obj) { + assert(obj->copy); + rwlock_wlock(&obj->lock); + // writer release the stm object, so release the last copy . + stm_releasecopy(obj->copy); + obj->copy = NULL; + if (--obj->reference > 0) { + // stm object grab by readers, reset the copy to NULL. + rwlock_wunlock(&obj->lock); + return; + } + // no one grab the stm object, no need to unlock wlock. + skynet_free(obj); +} + +static void +stm_releasereader(struct stm_object *obj) { + rwlock_rlock(&obj->lock); + if (__sync_sub_and_fetch(&obj->reference,1) == 0) { + // last reader, no writer. so no need to unlock + assert(obj->copy == NULL); + skynet_free(obj); + return; + } + rwlock_runlock(&obj->lock); +} + +static void +stm_grab(struct stm_object *obj) { + rwlock_rlock(&obj->lock); + int ref = __sync_fetch_and_add(&obj->reference,1); + rwlock_runlock(&obj->lock); + assert(ref > 0); +} + +static struct stm_copy * +stm_copy(struct stm_object *obj) { + rwlock_rlock(&obj->lock); + struct stm_copy * ret = obj->copy; + if (ret) { + int ref = __sync_fetch_and_add(&ret->reference,1); + assert(ref > 0); + } + rwlock_runlock(&obj->lock); + + return ret; +} + +static void +stm_update(struct stm_object *obj, void *msg, int32_t sz) { + struct stm_copy *copy = stm_newcopy(msg, sz); + rwlock_wlock(&obj->lock); + struct stm_copy *oldcopy = obj->copy; + obj->copy = copy; + rwlock_wunlock(&obj->lock); + + stm_releasecopy(oldcopy); +} + +// lua binding + +struct boxstm { + struct stm_object * obj; +}; + +static int +lcopy(lua_State *L) { + struct boxstm * box = lua_touserdata(L, 1); + stm_grab(box->obj); + lua_pushlightuserdata(L, box->obj); + return 1; +} + +static int +lnewwriter(lua_State *L) { + void * msg = lua_touserdata(L, 1); + uint32_t sz = luaL_checkunsigned(L, 2); + struct boxstm * box = lua_newuserdata(L, sizeof(*box)); + box->obj = stm_new(msg,sz); + lua_pushvalue(L, lua_upvalueindex(1)); + lua_setmetatable(L, -2); + + return 1; +} + +static int +ldeletewriter(lua_State *L) { + struct boxstm * box = lua_touserdata(L, 1); + stm_release(box->obj); + box->obj = NULL; + + return 0; +} + +static int +lupdate(lua_State *L) { + struct boxstm * box = lua_touserdata(L, 1); + void * msg = lua_touserdata(L, 2); + uint32_t sz = luaL_checkunsigned(L, 3); + stm_update(box->obj, msg, sz); + + return 0; +} + +struct boxreader { + struct stm_object *obj; + struct stm_copy *lastcopy; +}; + +static int +lnewreader(lua_State *L) { + struct boxreader * box = lua_newuserdata(L, sizeof(*box)); + box->obj = lua_touserdata(L, 1); + box->lastcopy = NULL; + lua_pushvalue(L, lua_upvalueindex(1)); + lua_setmetatable(L, -2); + + return 1; +} + +static int +ldeletereader(lua_State *L) { + struct boxreader * box = lua_touserdata(L, 1); + stm_releasereader(box->obj); + box->obj = NULL; + stm_releasecopy(box->lastcopy); + box->lastcopy = NULL; + + return 0; +} + +static int +lread(lua_State *L) { + struct boxreader * box = lua_touserdata(L, 1); + luaL_checktype(L, 2, LUA_TFUNCTION); + struct stm_copy * copy = stm_copy(box->obj); + if (copy == box->lastcopy) { + // not update + stm_releasecopy(copy); + lua_pushboolean(L, 0); + return 1; + } + + stm_releasecopy(box->lastcopy); + box->lastcopy = copy; + if (copy) { + lua_settop(L, 2); + lua_pushlightuserdata(L, copy->msg); + lua_pushunsigned(L, copy->sz); + lua_call(L, 2, LUA_MULTRET); + lua_pushboolean(L, 1); + lua_replace(L, 1); + return lua_gettop(L); + } else { + lua_pushboolean(L, 0); + return 1; + } +} + +int +luaopen_stm(lua_State *L) { + luaL_checkversion(L); + lua_createtable(L, 0, 3); + + lua_pushcfunction(L, lcopy); + lua_setfield(L, -2, "copy"); + + luaL_Reg writer[] = { + { "new", lnewwriter }, + { NULL, NULL }, + }; + lua_createtable(L, 0, 2); + lua_pushcfunction(L, ldeletewriter), + lua_setfield(L, -2, "__gc"); + lua_pushcfunction(L, lupdate), + lua_setfield(L, -2, "__call"); + luaL_setfuncs(L, writer, 1); + + luaL_Reg reader[] = { + { "newcopy", lnewreader }, + { NULL, NULL }, + }; + lua_createtable(L, 0, 2); + lua_pushcfunction(L, ldeletereader), + lua_setfield(L, -2, "__gc"); + lua_pushcfunction(L, lread), + lua_setfield(L, -2, "__call"); + luaL_setfuncs(L, reader, 1); + + return 1; +} diff --git a/test/teststm.lua b/test/teststm.lua new file mode 100644 index 00000000..7416a9f2 --- /dev/null +++ b/test/teststm.lua @@ -0,0 +1,36 @@ +local skynet = require "skynet" +local stm = require "stm" + +local mode = ... + +if mode == "slave" then + +skynet.start(function() + skynet.dispatch("lua", function (_,_, obj) + local obj = stm.newcopy(obj) + print("read:", obj(skynet.unpack)) + skynet.ret() + skynet.error("sleep and read") + for i=1,10 do + skynet.sleep(10) + print("read:", obj(skynet.unpack)) + end + skynet.exit() + end) +end) + +else + +skynet.start(function() + local slave = skynet.newservice(SERVICE_NAME, "slave") + local obj = stm.new(skynet.pack(1,2,3,4,5)) + local copy = stm.copy(obj) + skynet.call(slave, "lua", copy) + for i=1,5 do + skynet.sleep(20) + print("write", i) + obj(skynet.pack("hello world", i)) + end + skynet.exit() +end) +end