diff --git a/lualib-src/sproto/lsproto.c b/lualib-src/sproto/lsproto.c index c33118d4..99448487 100644 --- a/lualib-src/sproto/lsproto.c +++ b/lualib-src/sproto/lsproto.c @@ -573,6 +573,55 @@ lloadproto(lua_State *L) { return 1; } +static int +encode_default(const struct sproto_arg *args) { + lua_State *L = args->ud; + lua_pushstring(L, args->tagname); + if (args->index > 0) { + lua_newtable(L); + } else { + switch(args->type) { + case SPROTO_TINTEGER: + lua_pushinteger(L, 0); + break; + case SPROTO_TBOOLEAN: + lua_pushboolean(L, 0); + break; + case SPROTO_TSTRING: + lua_pushliteral(L, ""); + break; + case SPROTO_TSTRUCT: + lua_createtable(L, 0, 1); + lua_pushstring(L, sproto_name(args->subtype)); + lua_setfield(L, -2, "__type"); + break; + } + } + lua_rawset(L, -3); + return 0; +} + +/* + lightuserdata sproto_type + return default table + */ +static int +ldefault(lua_State *L) { + int ret; + // 32 is enough for dummy buffer, because ldefault encode nothing but the header. + char dummy[32]; + struct sproto_type * st = lua_touserdata(L, 1); + if (st == NULL) { + return luaL_argerror(L, 1, "Need a sproto_type object"); + } + lua_newtable(L); + ret = sproto_encode(st, dummy, sizeof(dummy), encode_default, L); + if (ret<0) { + return luaL_error(L, "dummy buffer (%d) is too small", (int)sizeof(dummy)); + } + return 1; +} + int luaopen_sproto_core(lua_State *L) { #ifdef luaL_checkversion @@ -587,6 +636,7 @@ luaopen_sproto_core(lua_State *L) { { "protocol", lprotocol }, { "loadproto", lloadproto }, { "saveproto", lsaveproto }, + { "default", ldefault }, { NULL, NULL }, }; luaL_newlib(L,l); diff --git a/lualib-src/sproto/sproto.c b/lualib-src/sproto/sproto.c index c37fdc1b..11756796 100644 --- a/lualib-src/sproto/sproto.c +++ b/lualib-src/sproto/sproto.c @@ -500,7 +500,11 @@ sproto_dump(struct sproto *s) { printf("=== %d protocol ===\n", s->protocol_n); for (i=0;iprotocol_n;i++) { struct protocol *p = &s->proto[i]; - printf("\t%s (%d) request:%s", p->name, p->tag, p->p[SPROTO_REQUEST]->name); + if (p->p[SPROTO_REQUEST]) { + printf("\t%s (%d) request:%s", p->name, p->tag, p->p[SPROTO_REQUEST]->name); + } else { + printf("\t%s (%d) request:(null)", p->name, p->tag); + } if (p->p[SPROTO_RESPONSE]) { printf(" response:%s", p->p[SPROTO_RESPONSE]->name); } diff --git a/lualib/sproto.lua b/lualib/sproto.lua index a9466f9c..911dc877 100644 --- a/lualib/sproto.lua +++ b/lualib/sproto.lua @@ -101,27 +101,64 @@ end function sproto:request_encode(protoname, tbl) local p = queryproto(self, protoname) - return core.encode(p.request,tbl) , p.tag + local request = p.request + if request then + return core.encode(request,tbl) , p.tag + else + return "" , p.tag + end end function sproto:response_encode(protoname, tbl) local p = queryproto(self, protoname) - return core.encode(p.response,tbl) + local response = p.response + if response then + return core.encode(response,tbl) + else + return "" + end end function sproto:request_decode(protoname, ...) local p = queryproto(self, protoname) - return core.decode(p.request,...) , p.name + local request = p.request + if request then + return core.decode(request,...) , p.name + else + return nil, p.name + end end function sproto:response_decode(protoname, ...) local p = queryproto(self, protoname) - return core.decode(p.response,...) + local response = p.response + if response then + return core.decode(response,...) + end end sproto.pack = core.pack sproto.unpack = core.unpack +function sproto:default(typename, type) + if type == nil then + return core.default(querytype(self, typename)) + else + local p = queryproto(self, typename) + if type == "REQUEST" then + if p.request then + return core.default(p.request) + end + elseif type == "RESPONSE" then + if p.response then + return core.default(p.response) + end + else + error "Invalid type" + end + end +end + local header_tmp = {} local function gen_response(self, response, session) diff --git a/lualib/sprotoparser.lua b/lualib/sprotoparser.lua index e7da7fa0..7b0b04e2 100644 --- a/lualib/sprotoparser.lua +++ b/lualib/sprotoparser.lua @@ -194,7 +194,7 @@ local function check_protocol(r) local request = v.request local response = v.response local p = map[tag] - + if p then error(string.format("redefined protocol tag %d at %s", tag, name)) end @@ -340,9 +340,6 @@ local function packtype(name, t, alltypes) end local function packproto(name, p, alltypes) --- if p.request == nil then --- error(string.format("Protocol %s need request", name)) --- end if p.request then local request = alltypes[p.request] if request == nil then