Compare commits

..

89 Commits

Author SHA1 Message Date
Cloud Wu
77d53cf5f1 release v0.5.2 2014-08-11 09:45:37 +08:00
云风
5621c52a44 Merge pull request #150 from cloudwu/http
Http
2014-08-11 09:43:10 +08:00
Cloud Wu
f6fa6c7ded timer support more than 497 days 2014-08-08 20:33:38 +08:00
Cloud Wu
e739df0a32 fix log 2014-08-04 18:17:35 +08:00
Cloud Wu
b18962929b bugfix: chunked mode 2014-08-04 18:11:11 +08:00
Cloud Wu
dec2e6fe4c add httpc 2014-08-04 18:06:16 +08:00
Cloud Wu
7beed39b1d bugfix: use temp array for each request 2014-08-04 16:34:29 +08:00
Cloud Wu
4673344e1e update readme, add wiki link 2014-08-02 22:02:17 +08:00
Cloud Wu
b49fc291a7 bugfix: delete local channel 2014-08-02 18:32:54 +08:00
Cloud Wu
11c528987c bugfix: socket.read(fd, nil) should clear buffer size to 0 2014-07-31 17:22:59 +08:00
Cloud Wu
25a8b3a179 release 0.5.0 2014-07-28 15:45:49 +08:00
云风
c323adb8d8 Merge pull request #144 from cloudwu/dev
pre release 0.5.0
2014-07-28 15:44:46 +08:00
Cloud Wu
7a4c7bc8a2 Merge branch 'master' of github.com:cloudwu/skynet into dev 2014-07-25 21:11:09 +08:00
Cloud Wu
a56b6aa425 bugfix: don't use string.format to concat binary string 2014-07-25 21:10:11 +08:00
云风
eaaf57ee0e Merge pull request #141 from duidui007/master
Update lua-bson.c
2014-07-24 11:56:04 +08:00
zeYang
7babdd3a97 Update lua-bson.c
//8  thx!
2014-07-23 21:26:25 +08:00
Cloud Wu
19e6462376 httpd.write_response capture socket error 2014-07-23 14:27:42 +08:00
云风
861f53858f Merge pull request #140 from cloudwu/http
Http
2014-07-23 14:08:02 +08:00
Cloud Wu
db25d1acc6 update history 2014-07-23 14:04:32 +08:00
Cloud Wu
0188e263dd recv header all, and then split with \r\n 2014-07-23 13:34:30 +08:00
Cloud Wu
18f20425a0 httpd don't need readline now 2014-07-23 12:19:48 +08:00
Cloud Wu
551d5048c4 simple url parser 2014-07-22 23:04:37 +08:00
Cloud Wu
d2ad63da81 read_request must return header table 2014-07-22 21:52:18 +08:00
Cloud Wu
932b2943dc simple httpd 2014-07-22 21:38:44 +08:00
Cloud Wu
362d6822bf remove dup line 2014-07-21 14:37:22 +08:00
Cloud Wu
7bdbdcd054 update history 2014-07-21 14:33:29 +08:00
云风
cb7058b563 Merge pull request #138 from cloudwu/config
config can read ENV
2014-07-21 14:33:01 +08:00
Cloud Wu
27ac1642e4 update history 2014-07-21 14:29:11 +08:00
Cloud Wu
d9d2a7d2ce Merge branch 'mongors' into dev 2014-07-21 14:28:09 +08:00
Cloud Wu
39ff6eb2a6 space to tab 2014-07-21 14:27:50 +08:00
云风
b933046fb9 Merge pull request #137 from bttscut/mongors
mongo driver support rs
2014-07-21 14:23:33 +08:00
btt
e59b6b44a5 mongo driver support rs 2014-07-21 11:47:38 +08:00
Cloud Wu
a8b80cd73e simplify clientsocket lib 2014-07-19 14:22:24 +08:00
Cloud Wu
d41c9db019 bugfix: update the response fd, when the client reconnect 2014-07-18 12:09:05 +08:00
Cloud Wu
e1ca48603e simplify code 2014-07-17 22:00:48 +08:00
Cloud Wu
8e4e9ed5c1 bugfix: connect auth 2014-07-17 21:58:09 +08:00
Cloud Wu
a0b536718c config can read ENV 2014-07-17 21:41:08 +08:00
Cloud Wu
6bafd05c65 don't support ip:port 2014-07-17 17:35:36 +08:00
Cloud Wu
7d2d107518 bugfix: socketchannel auth 2014-07-17 17:06:47 +08:00
Cloud Wu
75b2feff73 fix typo 2014-07-16 17:53:53 +08:00
Cloud Wu
9421435ac5 remove dup code 2014-07-16 17:53:18 +08:00
Cloud Wu
1181730302 socketchannel backup support default port 2014-07-16 17:21:39 +08:00
Cloud Wu
acc175e821 socketchannel backup support host/port 2014-07-16 17:18:56 +08:00
Cloud Wu
0ae1fcf7cd add socketchannel backup host 2014-07-16 13:46:44 +08:00
Cloud Wu
37a46607ce add socketchannel:changhost 2014-07-16 12:17:55 +08:00
Cloud Wu
691bb28163 update comment 2014-07-16 12:08:42 +08:00
Cloud Wu
dbf492b761 add snax.gateserver snax.loginserver snax.msgserver 2014-07-15 12:05:58 +08:00
Cloud Wu
43719de101 Merge branch 'master' into dev 2014-07-14 14:20:48 +08:00
云风
a8a683b48e Merge pull request #135 from cloudwu/dev
release v0.4.2
2014-07-14 14:16:17 +08:00
Cloud Wu
4cdf034f15 release 0.4.2 2014-07-14 14:14:42 +08:00
Cloud Wu
9feaf15b3d add netpack.pack_padding 2014-07-14 14:14:05 +08:00
Cloud Wu
ad58ed8a26 add subid 2014-07-13 21:25:44 +08:00
Cloud Wu
0d1c3a64d1 bugfix 2014-07-13 21:02:23 +08:00
Cloud Wu
e963ee6aba Merge branch 'dev' into sconn 2014-07-13 19:50:03 +08:00
Cloud Wu
68ddeab8fa bugfix: skynet.localname return nil not 0 2014-07-13 19:44:59 +08:00
Cloud Wu
3611bfe1af create gamefw 2014-07-13 19:44:23 +08:00
Cloud Wu
ec50e02777 Merge branch 'dev' into sconn 2014-07-13 19:37:21 +08:00
Cloud Wu
3bc7304609 skynet.localname return nil when the name not exist 2014-07-13 19:37:08 +08:00
Cloud Wu
0cdd71c0c2 msggate example 2014-07-13 19:36:44 +08:00
Cloud Wu
e1674f04c3 add gateserver 2014-07-13 14:54:53 +08:00
Cloud Wu
0dba8ed385 Merge branch 'dev' into sconn 2014-07-13 11:58:45 +08:00
Cloud Wu
2e35f405e7 update history 2014-07-13 11:56:28 +08:00
Cloud Wu
6813fd9ef7 bugfix: skynet.exit redirect error, and add datacenter.wait 2014-07-13 11:55:25 +08:00
Cloud Wu
04cb72d1a8 config.name 2014-07-13 11:50:30 +08:00
Cloud Wu
4284cfc372 loginserver and example 2014-07-13 10:32:21 +08:00
Cloud Wu
e8397348dd Merge branch 'dev' into sconn 2014-07-12 23:25:33 +08:00
Cloud Wu
7b5e62b896 fix typo 2014-07-12 23:24:29 +08:00
Cloud Wu
123d942819 Merge branch 'dev' into sconn 2014-07-12 23:14:25 +08:00
Cloud Wu
a090899bce improve cluster 2014-07-12 23:14:05 +08:00
Cloud Wu
68f4f45168 Merge branch 'dev' into sconn 2014-07-12 20:35:17 +08:00
Cloud Wu
3a5de32ad0 add socket.limit for defence 2014-07-12 20:30:02 +08:00
Cloud Wu
c3e758dc87 add crypt lib 2014-07-12 20:28:36 +08:00
Cloud Wu
54e6d03d36 update readme 2014-07-10 16:57:41 +08:00
Cloud Wu
83289b7612 bugfix: Issue #133 2014-07-09 21:07:26 +08:00
Cloud Wu
b7c12846a0 bugfix: socket channel 2014-07-09 15:59:06 +08:00
Cloud Wu
b5e10b8f9f add skynet.queue 2014-07-09 15:26:34 +08:00
Cloud Wu
0855bf30d9 update history 2014-07-09 14:33:53 +08:00
Cloud Wu
d05859b766 gate support TCP_NODELAY 2014-07-09 12:03:44 +08:00
Cloud Wu
fe5c73b1e8 socket id can't be less then 0 2014-07-09 11:25:28 +08:00
Cloud Wu
35775685c3 add worker dispatch weight 2014-07-08 19:56:30 +08:00
Cloud Wu
e60fb1d722 bugfix: remote send handle destination 2014-07-08 11:01:26 +08:00
云风
4e4dacbb4d Merge pull request #130 from cloudwu/dev
Dev
2014-07-07 20:53:38 +08:00
Cloud Wu
8f8b844bde ready for release 2014-07-07 20:52:34 +08:00
Cloud Wu
54f4d94ba2 bugfix: create queue first 2014-07-07 19:50:31 +08:00
Cloud Wu
4967dc2fce skynet.task return session:traceback 2014-07-07 19:00:22 +08:00
Cloud Wu
ece89a1b49 add new api skynet.task() 2014-07-07 18:29:27 +08:00
Cloud Wu
f874fdc618 throw error when skynet.exit 2014-07-03 17:40:07 +08:00
Cloud Wu
711c04e6a9 bugfix: redirect should pass session (0) 2014-07-03 17:13:47 +08:00
Cloud Wu
1e0189962b bugfix: dead lock when service_harbor exit 2014-06-30 11:31:24 +08:00
65 changed files with 3369 additions and 418 deletions

8
.gitignore vendored
View File

@@ -1,10 +1,10 @@
*.o *.o
*.a *.a
./skynet /skynet
./skynet.pid /skynet.pid
3rd/lua/lua 3rd/lua/lua
3rd/lua/luac 3rd/lua/luac
./cservice /cservice
./luaclib /luaclib
*.so *.so
*.dSYM *.dSYM

View File

@@ -1,3 +1,44 @@
v0.5.2 (2014-8-11)
-----------
* Bugfix : httpd request
* Bugifx : http chunked mode
* Add : httpc
* timer support more than 497 days
v0.5.1 (2014-8-4)
-----------
* Bugfix : http module
* Bugfix : multicast local channel delete
* Bugfix : socket.read(fd)
v0.5.0 (2014-7-28)
-----------
* skynet.exit will quit service immediately.
* Add snax.gateserver, snax.loginserver, snax.msgserver
* Simplify clientsocket lib
* mongo driver support replica set
* config file support read from ENV
* add simple httpd (see examples/simpleweb.lua)
v0.4.2 (2014-7-14)
-----------
* Bugfix : invalid negative socket id
* Add optional TCP_NODELAY support
* Add worker thread weight
* Add skynet.queue
* Bugfix: socketchannel
* cluster can throw error
* Add readline and writeline to clientsocket lib
* Add cluster.reload to reload config file
* Add datacenter.wait
v0.4.1 (2014-7-7)
-----------
* Add SERVICE_NAME in loader
* Throw error back when skynet.error
* Add skynet.task
* Bugfix for last version (harbor service bugs)
v0.4.0 (2014-6-30) v0.4.0 (2014-6-30)
----------- -----------
* Optimize redis driver `compose_message`. * Optimize redis driver `compose_message`.

View File

@@ -43,7 +43,7 @@ jemalloc : $(MALLOC_STATICLIB)
CSERVICE = snlua logger gate harbor CSERVICE = snlua logger gate harbor
LUA_CLIB = skynet socketdriver int64 bson mongo md5 netpack \ LUA_CLIB = skynet socketdriver int64 bson mongo md5 netpack \
cjson clientsocket memory profile multicast \ cjson clientsocket memory profile multicast \
cluster cluster crypt
SKYNET_SRC = skynet_main.c skynet_handle.c skynet_module.c skynet_mq.c \ 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 \ skynet_server.c skynet_start.c skynet_timer.c skynet_error.c \
@@ -110,6 +110,9 @@ $(LUA_CLIB_PATH)/multicast.so : lualib-src/lua-multicast.c | $(LUA_CLIB_PATH)
$(LUA_CLIB_PATH)/cluster.so : lualib-src/lua-cluster.c | $(LUA_CLIB_PATH) $(LUA_CLIB_PATH)/cluster.so : lualib-src/lua-cluster.c | $(LUA_CLIB_PATH)
$(CC) $(CFLAGS) $(SHARED) -Iskynet-src $^ -o $@ $(CC) $(CFLAGS) $(SHARED) -Iskynet-src $^ -o $@
$(LUA_CLIB_PATH)/crypt.so : lualib-src/lua-crypt.c | $(LUA_CLIB_PATH)
$(CC) $(CFLAGS) $(SHARED) $^ -o $@
clean : clean :
rm -f $(SKYNET_BUILD_PATH)/skynet $(CSERVICE_PATH)/*.so $(LUA_CLIB_PATH)/*.so rm -f $(SKYNET_BUILD_PATH)/skynet $(CSERVICE_PATH)/*.so $(LUA_CLIB_PATH)/*.so

View File

@@ -3,7 +3,7 @@
For linux, install autoconf first for jemalloc For linux, install autoconf first for jemalloc
``` ```
git clone git@github.com:cloudwu/skynet.git git clone https://github.com/cloudwu/skynet.git
cd skynet cd skynet
make 'PLATFORM' # PLATFORM can be linux, macosx, freebsd now make 'PLATFORM' # PLATFORM can be linux, macosx, freebsd now
``` ```
@@ -34,9 +34,6 @@ Each lua file only load once and cache it in memory during skynet start . so if
You can also use the offical lua version , edit the makefile by yourself . You can also use the offical lua version , edit the makefile by yourself .
## Blog (in Chinese) ## How To (in Chinese)
* http://blog.codingnow.com/2012/09/the_design_of_skynet.html * Read Wiki https://github.com/cloudwu/skynet/wiki
* http://blog.codingnow.com/2012/08/skynet.html
* http://blog.codingnow.com/2012/08/skynet_harbor_rpc.html
* http://blog.codingnow.com/eo/skynet/

View File

@@ -2,30 +2,46 @@ package.cpath = "luaclib/?.so"
local socket = require "clientsocket" local socket = require "clientsocket"
local cjson = require "cjson" local cjson = require "cjson"
local bit32 = require "bit32"
local fd = socket.connect("127.0.0.1", 8888) local fd = assert(socket.connect("127.0.0.1", 8888))
local last local function send_package(fd, pack)
local result = {} local size = #pack
local package = string.char(bit32.extract(size,8,8)) ..
string.char(bit32.extract(size,0,8))..
pack
local function dispatch() socket.send(fd, package)
while true do end
local status
status, last = socket.recv(fd, last, result) local function unpack_package(text)
if status == nil then local size = #text
error "Server closed" if size < 2 then
end return nil, text
if not status then
break
end
for _, v in ipairs(result) do
local session,t,str = string.match(v, "(%d+)(.)(.*)")
assert(t == '-' or t == '+')
session = tonumber(session)
local result = cjson.decode(str)
print("Response:",session, result[1], result[2])
end
end end
local s = text:byte(1) * 256 + text:byte(2)
if size < s+2 then
return nil, text
end
return text:sub(3,2+s), text:sub(3+s)
end
local function recv_package(last)
local result
result, last = unpack_package(last)
if result then
return result, last
end
local r = socket.recv(fd)
if not r then
return nil, last
end
if r == "" then
error "Server closed"
end
return unpack_package(last .. r)
end end
local session = 0 local session = 0
@@ -33,13 +49,26 @@ local session = 0
local function send_request(v) local function send_request(v)
session = session + 1 session = session + 1
local str = string.format("%d+%s",session, cjson.encode(v)) local str = string.format("%d+%s",session, cjson.encode(v))
socket.send(fd, str) send_package(fd, str)
print("Request:", session) print("Request:", session)
end end
local last = ""
while true do while true do
dispatch() while true do
local cmd = socket.readline() local v
v, last = recv_package(last)
if not v then
break
end
local session,t,str = string.match(v, "(%d+)(.)(.*)")
assert(t == '-' or t == '+')
session = tonumber(session)
local result = cjson.decode(str)
print("Response:",session, result[1], result[2])
end
local cmd = socket.readstdin()
if cmd then if cmd then
local args = {} local args = {}
string.gsub(cmd, '[^ ]+', function(v) table.insert(args, v) end ) string.gsub(cmd, '[^ ]+', function(v) table.insert(args, v) end )

8
examples/config.login Normal file
View File

@@ -0,0 +1,8 @@
thread = 8
logger = nil
harbor = 0
start = "main"
bootstrap = "snlua bootstrap" -- The service for bootstrap
luaservice = "./service/?.lua;./examples/login/?.lua"
lualoader = "lualib/loader.lua"
cpath = "./cservice/?.so"

184
examples/login/client.lua Normal file
View File

@@ -0,0 +1,184 @@
package.cpath = "luaclib/?.so"
local socket = require "clientsocket"
local crypt = require "crypt"
local bit32 = require "bit32"
local fd = assert(socket.connect("127.0.0.1", 8001))
local function writeline(fd, text)
socket.send(fd, text .. "\n")
end
local function unpack_line(text)
local from = text:find("\n", 1, true)
if from then
return text:sub(1, from-1), text:sub(from+1)
end
return nil, text
end
local last = ""
local function unpack_f(f)
local function try_recv(fd, last)
local result
result, last = f(last)
if result then
return result, last
end
local r = socket.recv(fd)
if not r then
return nil, last
end
if r == "" then
error "Server closed"
end
return f(last .. r)
end
return function()
while true do
local result
result, last = try_recv(fd, last)
if result then
return result
end
socket.usleep(100)
end
end
end
local readline = unpack_f(unpack_line)
local challenge = crypt.base64decode(readline())
local clientkey = crypt.randomkey()
writeline(fd, crypt.base64encode(crypt.dhexchange(clientkey)))
local secret = crypt.dhsecret(crypt.base64decode(readline()), clientkey)
print("sceret is ", crypt.hexencode(secret))
local hmac = crypt.hmac64(challenge, secret)
writeline(fd, crypt.base64encode(hmac))
local token = {
server = "sample",
user = "hello",
pass = "password",
}
local function encode_token(token)
return string.format("%s@%s:%s",
crypt.base64encode(token.user),
crypt.base64encode(token.server),
crypt.base64encode(token.pass))
end
local etoken = crypt.desencode(secret, encode_token(token))
local b = crypt.base64encode(etoken)
writeline(fd, crypt.base64encode(etoken))
local result = readline()
print(result)
local code = tonumber(string.sub(result, 1, 3))
assert(code == 200)
socket.close(fd)
local subid = crypt.base64decode(string.sub(result, 5))
print("login ok, subid=", subid)
----- connect to game server
local function send_request(v, session)
local size = #v + 4
local package = string.char(bit32.extract(size,8,8))..
string.char(bit32.extract(size,0,8))..
v..
string.char(bit32.extract(session,24,8))..
string.char(bit32.extract(session,16,8))..
string.char(bit32.extract(session,8,8))..
string.char(bit32.extract(session,0,8))
socket.send(fd, package)
return v, session
end
local function recv_response(v)
local content = v:sub(1,-6)
local ok = v:sub(-5,-5):byte()
local session = 0
for i=-4,-1 do
local c = v:byte(i)
session = session + bit32.lshift(c,(-1-i) * 8)
end
return ok ~=0 , content, session
end
local function unpack_package(text)
local size = #text
if size < 2 then
return nil, text
end
local s = text:byte(1) * 256 + text:byte(2)
if size < s+2 then
return nil, text
end
return text:sub(3,2+s), text:sub(3+s)
end
local readpackage = unpack_f(unpack_package)
local function send_package(fd, pack)
local size = #pack
local package = string.char(bit32.extract(size,8,8))..
string.char(bit32.extract(size,0,8))..
pack
socket.send(fd, package)
end
local text = "echo"
local index = 1
print("connect")
local fd = assert(socket.connect("127.0.0.1", 8888))
last = ""
local handshake = string.format("%s@%s#%s:%d", crypt.base64encode(token.user), crypt.base64encode(token.server),crypt.base64encode(subid) , index)
local hmac = crypt.hmac64(crypt.hashkey(handshake), secret)
send_package(fd, handshake .. ":" .. crypt.base64encode(hmac))
print(readpackage())
print("===>",send_request(text,0))
-- don't recv response
-- print("<===",recv_response(readpackage()))
print("disconnect")
socket.close(fd)
index = index + 1
print("connect again")
local fd = assert(socket.connect("127.0.0.1", 8888))
last = ""
local handshake = string.format("%s@%s#%s:%d", crypt.base64encode(token.user), crypt.base64encode(token.server),crypt.base64encode(subid) , index)
local hmac = crypt.hmac64(crypt.hashkey(handshake), secret)
send_package(fd, handshake .. ":" .. crypt.base64encode(hmac))
print(readpackage())
print("===>",send_request("fake",0)) -- request again (use last session 0, so the request message is fake)
print("===>",send_request("again",1)) -- request again (use new session)
print("<===",recv_response(readpackage()))
print("<===",recv_response(readpackage()))
print("disconnect")
socket.close(fd)

88
examples/login/gated.lua Normal file
View File

@@ -0,0 +1,88 @@
local msgserver = require "snax.msgserver"
local crypt = require "crypt"
local skynet = require "skynet"
local loginservice = tonumber(...)
local server = {}
local users = {}
local username_map = {}
local internal_id = 0
-- login server disallow multi login, so login_handler never be reentry
-- call by login server
function server.login_handler(uid, secret)
if users[uid] then
error(string.format("%s is already login", uid))
end
internal_id = internal_id + 1
local username = msgserver.username(uid, internal_id, servername)
-- you can use a pool to alloc new agent
local agent = skynet.newservice "msgagent"
local u = {
username = username,
agent = agent,
uid = uid,
subid = internal_id,
}
-- trash subid (no used)
skynet.call(agent, "lua", "login", uid, internal_id, secret)
users[uid] = u
username_map[username] = u
msgserver.login(username, secret)
-- you should return unique subid
return internal_id
end
-- call by agent
function server.logout_handler(uid, subid)
local u = users[uid]
if u then
local username = msgserver.username(uid, subid, servername)
assert(u.username == username)
msgserver.logout(u.username)
users[uid] = nil
username_map[u.username] = nil
skynet.call(loginservice, "lua", "logout",uid, subid)
end
end
-- call by login server
function server.kick_handler(uid, subid)
local u = users[uid]
if u then
local username = msgserver.username(uid, subid, servername)
assert(u.username == username)
-- NOTICE: logout may call skynet.exit, so you should use pcall.
pcall(skynet.call, u.agent, "lua", "logout")
end
end
-- call by self (when socket disconnect)
function server.disconnect_handler(username)
local u = username_map[username]
if u then
skynet.call(u.agent, "lua", "afk")
end
end
-- call by self (when recv a request from client)
function server.request_handler(username, msg, sz)
local u = username_map[username]
return skynet.tostring(skynet.rawcall(u.agent, "client", msg, sz))
end
-- call by self (when gate open)
function server.register_handler(name)
servername = name
skynet.call(loginservice, "lua", "register_gate", servername, skynet.self())
end
msgserver.start(server)

62
examples/login/logind.lua Normal file
View File

@@ -0,0 +1,62 @@
local login = require "snax.loginserver"
local crypt = require "crypt"
local skynet = require "skynet"
local server = {
host = "127.0.0.1",
port = 8001,
multilogin = false, -- disallow multilogin
name = "login_master",
}
local server_list = {}
local user_online = {}
local user_login = {}
function server.auth_handler(token)
-- the token is base64(user)@base64(server):base64(password)
local user, server, password = token:match("([^@]+)@([^:]+):(.+)")
user = crypt.base64decode(user)
server = crypt.base64decode(server)
password = crypt.base64decode(password)
assert(password == "password")
return server, user
end
function server.login_handler(server, uid, secret)
print(string.format("%s@%s is login, secret is %s", uid, server, crypt.hexencode(secret)))
local gameserver = assert(server_list[server], "Unknown server")
-- only one can login, because disallow multilogin
local last = user_online[uid]
if last then
skynet.call(last.address, "lua", "kick", uid, last.subid)
end
if user_online[uid] then
error(string.format("user %s is already online", uid))
end
local subid = tostring(skynet.call(gameserver, "lua", "login", uid, secret))
user_online[uid] = { address = gameserver, subid = subid , server = server}
return subid
end
local CMD = {}
function CMD.register_gate(server, address)
server_list[server] = address
end
function CMD.logout(uid, subid)
local u = user_online[uid]
if u then
print(string.format("%s@%s is logout", uid, u.server))
user_online[uid] = nil
end
end
function server.command_handler(command, source, ...)
local f = assert(CMD[command])
return f(source, ...)
end
login(server)

12
examples/login/main.lua Normal file
View File

@@ -0,0 +1,12 @@
local skynet = require "skynet"
skynet.start(function()
local loginserver = skynet.newservice("logind")
local gate = skynet.newservice("gated", loginserver)
skynet.call(gate, "lua", "open" , {
port = 8888,
maxclient = 64,
servername = "sample",
})
end)

View File

@@ -0,0 +1,53 @@
local skynet = require "skynet"
skynet.register_protocol {
name = "client",
id = skynet.PTYPE_CLIENT,
unpack = skynet.tostring,
}
local gate
local userid, subid
local CMD = {}
function CMD.login(source, uid, sid, secret)
-- you may use secret to make a encrypted data stream
skynet.error(string.format("%s is login", uid))
gate = source
userid = uid
subid = sid
-- you may load user data from database
end
local function logout()
if gate then
skynet.call(gate, "lua", "logout", userid, subid)
end
skynet.exit()
end
function CMD.logout(source)
-- NOTICE: The logout MAY be reentry
skynet.error(string.format("%s is logout", userid))
logout()
end
function CMD.afk(source)
-- the connection is broken, but the user may back
skynet.error(string.format("AFK"))
end
skynet.start(function()
-- If you want to fork a work thread , you MUST do it in CMD.login
skynet.dispatch("lua", function(session, source, command, ...)
local f = assert(CMD[command])
skynet.ret(skynet.pack(f(source, ...)))
end)
skynet.dispatch("client", function(_,_, msg)
-- the simple ehco service
skynet.sleep(10) -- sleep a while
skynet.ret(msg)
end)
end)

View File

@@ -11,7 +11,9 @@ skynet.start(function()
skynet.call(watchdog, "lua", "start", { skynet.call(watchdog, "lua", "start", {
port = 8888, port = 8888,
maxclient = max_client, maxclient = max_client,
nodelay = true,
}) })
print("Watchdog listen on ", 8888)
skynet.exit() skynet.exit()
end) end)

View File

@@ -16,7 +16,11 @@ end
skynet.start(function() skynet.start(function()
skynet.dispatch("lua", function(session, address, cmd, ...) skynet.dispatch("lua", function(session, address, cmd, ...)
local f = command[string.upper(cmd)] local f = command[string.upper(cmd)]
skynet.ret(skynet.pack(f(...))) if f then
skynet.ret(skynet.pack(f(...)))
else
error(string.format("Unknown command %s", tostring(cmd)))
end
end) end)
skynet.register "SIMPLEDB" skynet.register "SIMPLEDB"
end) end)

View File

@@ -12,7 +12,7 @@ skynet.register_protocol {
local w = service_map[address] local w = service_map[address]
if w then if w then
for watcher in pairs(w) do for watcher in pairs(w) do
skynet.redirect(watcher, address, "error", "") skynet.redirect(watcher, address, "error", 0, "")
end end
service_map[address] = false service_map[address] = false
end end

79
examples/simpleweb.lua Normal file
View File

@@ -0,0 +1,79 @@
local skynet = require "skynet"
local socket = require "socket"
local httpd = require "http.httpd"
local sockethelper = require "http.sockethelper"
local urllib = require "http.url"
local table = table
local string = string
local mode = ...
if mode == "agent" then
local function response(id, ...)
local ok, err = httpd.write_response(sockethelper.writefunc(id), ...)
if not ok then
-- if err == sockethelper.socket_error , that means socket closed.
skynet.error(string.format("fd = %d, %s", id, err))
end
end
skynet.start(function()
skynet.dispatch("lua", function (_,_,id)
socket.start(id)
-- limit request body size to 8192 (you can pass nil to unlimit)
local code, url, method, header, body = httpd.read_request(sockethelper.readfunc(id), 8192)
if code then
if code ~= 200 then
response(id, code)
else
local tmp = {}
if header.host then
table.insert(tmp, string.format("host: %s", header.host))
end
local path, query = urllib.parse(url)
table.insert(tmp, string.format("path: %s", path))
if query then
local q = urllib.parse_query(query)
for k, v in pairs(q) do
table.insert(tmp, string.format("query: %s= %s", k,v))
end
end
table.insert(tmp, "-----header----")
for k,v in pairs(header) do
table.insert(tmp, string.format("%s = %s",k,v))
end
table.insert(tmp, "-----body----\n" .. body)
response(id, code, table.concat(tmp,"\n"))
end
else
if url == sockethelper.socket_error then
skynet.error("socket closed")
else
skynet.error(url)
end
end
socket.close(id)
end)
end)
else
skynet.start(function()
local agent = {}
for i= 1, 20 do
agent[i] = skynet.newservice(SERVICE_NAME, "agent")
end
local balance = 1
local id = socket.listen("0.0.0.0", 8001)
socket.start(id , function(id, addr)
skynet.error(string.format("%s connected, pass it to agent :%08x", addr, agent[balance]))
skynet.send(agent[balance], "lua", id)
balance = balance + 1
if balance > #agent then
balance = 1
end
end)
end)
end

View File

@@ -1083,7 +1083,7 @@ typeclosure(lua_State *L) {
"string", // 5 "string", // 5
"binary", // 6 "binary", // 6
"objectid", // 7 "objectid", // 7
"timestamp",// 8 "timestamp", // 8
"date", // 9 "date", // 9
"regex", // 10 "regex", // 10
"minkey", // 11 "minkey", // 11
@@ -1180,7 +1180,6 @@ luaopen_bson(lua_State *L) {
{ "timestamp", ltimestamp }, { "timestamp", ltimestamp },
{ "regex", lregex }, { "regex", lregex },
{ "binary", lbinary }, { "binary", lbinary },
{ "regex", lregex },
{ "objectid", lobjectid }, { "objectid", lobjectid },
{ "decode", ldecode }, { "decode", ldecode },
{ NULL, NULL }, { NULL, NULL },

View File

@@ -75,47 +75,12 @@ lsend(lua_State *L) {
size_t sz = 0; size_t sz = 0;
int fd = luaL_checkinteger(L,1); int fd = luaL_checkinteger(L,1);
const char * msg = luaL_checklstring(L, 2, &sz); const char * msg = luaL_checklstring(L, 2, &sz);
uint8_t tmp[sz + 2];
if (sz >= 0x10000) {
return luaL_error(L, "package too long %d (16bit limited)", (int)sz);
}
tmp[0] = (sz >> 8) & 0xff;
tmp[1] = sz & 0xff;
memcpy(tmp+2, msg, sz);
block_send(L, fd, (const char *)tmp, (int)sz+2); block_send(L, fd, msg, (int)sz);
return 0; return 0;
} }
static int
unpack(lua_State *L, uint8_t *buffer, int sz, int n) {
int size = 0;
if (sz >= 2) {
size = buffer[0] << 8 | buffer[1];
if (size > sz - 2) {
goto _block;
}
} else {
goto _block;
}
++n;
lua_pushlstring(L, (const char *)buffer+2, size);
lua_rawseti(L, 3, n);
buffer += size + 2;
sz -= size + 2;
return unpack(L, buffer, sz, n);
_block:
lua_pushboolean(L, n==0 ? 0:1);
if (sz == 0) {
lua_pushnil(L);
} else {
lua_pushlstring(L, (const char *)buffer, sz);
}
return 2;
}
/* /*
intger fd intger fd
string last string last
@@ -125,46 +90,31 @@ _block:
boolean (true: data, false: block, nil: close) boolean (true: data, false: block, nil: close)
string last string last
*/ */
struct socket_buffer {
void * buffer;
int sz;
};
static int static int
lrecv(lua_State *L) { lrecv(lua_State *L) {
int fd = luaL_checkinteger(L,1); int fd = luaL_checkinteger(L,1);
size_t sz = 0;
const char * last = lua_tolstring(L,2,&sz);
luaL_checktype(L, 3, LUA_TTABLE);
char tmp[CACHE_SIZE]; char buffer[CACHE_SIZE];
char * buffer; int r = recv(fd, buffer, CACHE_SIZE, 0);
int r = recv(fd, tmp, CACHE_SIZE, 0);
if (r == 0) { if (r == 0) {
lua_pushliteral(L, "");
// close // close
return 0; return 1;
} }
if (r < 0) { if (r < 0) {
if (errno == EAGAIN || errno == EINTR) { if (errno == EAGAIN || errno == EINTR) {
lua_pushboolean(L, 0); return 0;
lua_pushvalue(L, 2);
return 2;
} }
luaL_error(L, "socket error: %s", strerror(errno)); luaL_error(L, "socket error: %s", strerror(errno));
} }
if (sz + r <= CACHE_SIZE) { lua_pushlstring(L, buffer, r);
buffer = tmp; return 1;
memmove(buffer + sz, buffer, r);
memcpy(buffer, last, sz);
} else {
buffer = lua_newuserdata(L, r + sz);
memcpy(buffer, last, sz);
memcpy(buffer + sz, tmp, r);
}
int i;
int n = lua_rawlen(L, 3);
for (i=1;i<=n;i++) {
lua_pushnil(L);
lua_rawseti(L, 3, i);
}
return unpack(L, (uint8_t *)buffer, r+sz, 0);
} }
static int static int
@@ -219,7 +169,7 @@ readline_stdin(void * arg) {
} }
static int static int
lreadline(lua_State *L) { lreadstdin(lua_State *L) {
struct queue *q = lua_touserdata(L, lua_upvalueindex(1)); struct queue *q = lua_touserdata(L, lua_upvalueindex(1));
LOCK(q); LOCK(q);
if (q->head == q->tail) { if (q->head == q->tail) {
@@ -251,8 +201,8 @@ luaopen_clientsocket(lua_State *L) {
struct queue * q = lua_newuserdata(L, sizeof(*q)); struct queue * q = lua_newuserdata(L, sizeof(*q));
memset(q, 0, sizeof(*q)); memset(q, 0, sizeof(*q));
lua_pushcclosure(L, lreadline, 1); lua_pushcclosure(L, lreadstdin, 1);
lua_setfield(L, -2, "readline"); lua_setfield(L, -2, "readstdin");
pthread_t pid ; pthread_t pid ;
pthread_create(&pid, NULL, readline_stdin, q); pthread_create(&pid, NULL, readline_stdin, q);

View File

@@ -15,7 +15,7 @@
uint32_t next_session uint32_t next_session
*/ */
#define TEMP_LENGTH 0x10002 #define TEMP_LENGTH 0x10007
static void static void
fill_uint32(uint8_t * buf, uint32_t n) { fill_uint32(uint8_t * buf, uint32_t n) {
@@ -146,6 +146,7 @@ lunpackrequest(lua_State *L) {
/* /*
int session int session
boolean ok
lightuserdata msg lightuserdata msg
int sz int sz
return string response return string response
@@ -155,15 +156,27 @@ lpackresponse(lua_State *L) {
uint32_t session = luaL_checkunsigned(L,1); uint32_t session = luaL_checkunsigned(L,1);
// clusterd.lua:command.socket call lpackresponse, // clusterd.lua:command.socket call lpackresponse,
// and the msg/sz is return by skynet.rawcall , so don't free(msg) // and the msg/sz is return by skynet.rawcall , so don't free(msg)
void * msg = lua_touserdata(L,2); int ok = lua_toboolean(L,2);
size_t sz = luaL_checkunsigned(L, 3); void * msg;
size_t sz;
if (lua_type(L,3) == LUA_TSTRING) {
msg = (void *)lua_tolstring(L, 3, &sz);
if (sz > 0x1000) {
sz = 0x1000;
}
} else {
msg = lua_touserdata(L,3);
sz = luaL_checkunsigned(L, 4);
}
uint8_t buf[TEMP_LENGTH]; uint8_t buf[TEMP_LENGTH];
fill_header(L, buf, sz+4, msg); fill_header(L, buf, sz+5, msg);
fill_uint32(buf+2, session); fill_uint32(buf+2, session);
memcpy(buf+6,msg,sz); buf[6] = ok;
memcpy(buf+7,msg,sz);
lua_pushlstring(L, (const char *)buf, sz+6); lua_pushlstring(L, (const char *)buf, sz+7);
return 1; return 1;
} }
@@ -178,13 +191,13 @@ static int
lunpackresponse(lua_State *L) { lunpackresponse(lua_State *L) {
size_t sz; size_t sz;
const char * buf = luaL_checklstring(L, 1, &sz); const char * buf = luaL_checklstring(L, 1, &sz);
if (sz < 4) { if (sz < 5) {
return 0; return 0;
} }
uint32_t session = unpack_uint32((const uint8_t *)buf); uint32_t session = unpack_uint32((const uint8_t *)buf);
lua_pushunsigned(L, session); lua_pushunsigned(L, session);
lua_pushboolean(L, 1); lua_pushboolean(L, buf[4]);
lua_pushlstring(L, buf+4, sz-4); lua_pushlstring(L, buf+5, sz-5);
return 3; return 3;
} }

859
lualib-src/lua-crypt.c Normal file
View File

@@ -0,0 +1,859 @@
#include <lua.h>
#include <lauxlib.h>
#include <time.h>
#include <stdint.h>
#include <string.h>
#include <stdlib.h>
#define SMALL_CHUNK 256
/* the eight DES S-boxes */
uint32_t SB1[64] = {
0x01010400, 0x00000000, 0x00010000, 0x01010404,
0x01010004, 0x00010404, 0x00000004, 0x00010000,
0x00000400, 0x01010400, 0x01010404, 0x00000400,
0x01000404, 0x01010004, 0x01000000, 0x00000004,
0x00000404, 0x01000400, 0x01000400, 0x00010400,
0x00010400, 0x01010000, 0x01010000, 0x01000404,
0x00010004, 0x01000004, 0x01000004, 0x00010004,
0x00000000, 0x00000404, 0x00010404, 0x01000000,
0x00010000, 0x01010404, 0x00000004, 0x01010000,
0x01010400, 0x01000000, 0x01000000, 0x00000400,
0x01010004, 0x00010000, 0x00010400, 0x01000004,
0x00000400, 0x00000004, 0x01000404, 0x00010404,
0x01010404, 0x00010004, 0x01010000, 0x01000404,
0x01000004, 0x00000404, 0x00010404, 0x01010400,
0x00000404, 0x01000400, 0x01000400, 0x00000000,
0x00010004, 0x00010400, 0x00000000, 0x01010004
};
static uint32_t SB2[64] = {
0x80108020, 0x80008000, 0x00008000, 0x00108020,
0x00100000, 0x00000020, 0x80100020, 0x80008020,
0x80000020, 0x80108020, 0x80108000, 0x80000000,
0x80008000, 0x00100000, 0x00000020, 0x80100020,
0x00108000, 0x00100020, 0x80008020, 0x00000000,
0x80000000, 0x00008000, 0x00108020, 0x80100000,
0x00100020, 0x80000020, 0x00000000, 0x00108000,
0x00008020, 0x80108000, 0x80100000, 0x00008020,
0x00000000, 0x00108020, 0x80100020, 0x00100000,
0x80008020, 0x80100000, 0x80108000, 0x00008000,
0x80100000, 0x80008000, 0x00000020, 0x80108020,
0x00108020, 0x00000020, 0x00008000, 0x80000000,
0x00008020, 0x80108000, 0x00100000, 0x80000020,
0x00100020, 0x80008020, 0x80000020, 0x00100020,
0x00108000, 0x00000000, 0x80008000, 0x00008020,
0x80000000, 0x80100020, 0x80108020, 0x00108000
};
static uint32_t SB3[64] = {
0x00000208, 0x08020200, 0x00000000, 0x08020008,
0x08000200, 0x00000000, 0x00020208, 0x08000200,
0x00020008, 0x08000008, 0x08000008, 0x00020000,
0x08020208, 0x00020008, 0x08020000, 0x00000208,
0x08000000, 0x00000008, 0x08020200, 0x00000200,
0x00020200, 0x08020000, 0x08020008, 0x00020208,
0x08000208, 0x00020200, 0x00020000, 0x08000208,
0x00000008, 0x08020208, 0x00000200, 0x08000000,
0x08020200, 0x08000000, 0x00020008, 0x00000208,
0x00020000, 0x08020200, 0x08000200, 0x00000000,
0x00000200, 0x00020008, 0x08020208, 0x08000200,
0x08000008, 0x00000200, 0x00000000, 0x08020008,
0x08000208, 0x00020000, 0x08000000, 0x08020208,
0x00000008, 0x00020208, 0x00020200, 0x08000008,
0x08020000, 0x08000208, 0x00000208, 0x08020000,
0x00020208, 0x00000008, 0x08020008, 0x00020200
};
static uint32_t SB4[64] = {
0x00802001, 0x00002081, 0x00002081, 0x00000080,
0x00802080, 0x00800081, 0x00800001, 0x00002001,
0x00000000, 0x00802000, 0x00802000, 0x00802081,
0x00000081, 0x00000000, 0x00800080, 0x00800001,
0x00000001, 0x00002000, 0x00800000, 0x00802001,
0x00000080, 0x00800000, 0x00002001, 0x00002080,
0x00800081, 0x00000001, 0x00002080, 0x00800080,
0x00002000, 0x00802080, 0x00802081, 0x00000081,
0x00800080, 0x00800001, 0x00802000, 0x00802081,
0x00000081, 0x00000000, 0x00000000, 0x00802000,
0x00002080, 0x00800080, 0x00800081, 0x00000001,
0x00802001, 0x00002081, 0x00002081, 0x00000080,
0x00802081, 0x00000081, 0x00000001, 0x00002000,
0x00800001, 0x00002001, 0x00802080, 0x00800081,
0x00002001, 0x00002080, 0x00800000, 0x00802001,
0x00000080, 0x00800000, 0x00002000, 0x00802080
};
static uint32_t SB5[64] = {
0x00000100, 0x02080100, 0x02080000, 0x42000100,
0x00080000, 0x00000100, 0x40000000, 0x02080000,
0x40080100, 0x00080000, 0x02000100, 0x40080100,
0x42000100, 0x42080000, 0x00080100, 0x40000000,
0x02000000, 0x40080000, 0x40080000, 0x00000000,
0x40000100, 0x42080100, 0x42080100, 0x02000100,
0x42080000, 0x40000100, 0x00000000, 0x42000000,
0x02080100, 0x02000000, 0x42000000, 0x00080100,
0x00080000, 0x42000100, 0x00000100, 0x02000000,
0x40000000, 0x02080000, 0x42000100, 0x40080100,
0x02000100, 0x40000000, 0x42080000, 0x02080100,
0x40080100, 0x00000100, 0x02000000, 0x42080000,
0x42080100, 0x00080100, 0x42000000, 0x42080100,
0x02080000, 0x00000000, 0x40080000, 0x42000000,
0x00080100, 0x02000100, 0x40000100, 0x00080000,
0x00000000, 0x40080000, 0x02080100, 0x40000100
};
static uint32_t SB6[64] = {
0x20000010, 0x20400000, 0x00004000, 0x20404010,
0x20400000, 0x00000010, 0x20404010, 0x00400000,
0x20004000, 0x00404010, 0x00400000, 0x20000010,
0x00400010, 0x20004000, 0x20000000, 0x00004010,
0x00000000, 0x00400010, 0x20004010, 0x00004000,
0x00404000, 0x20004010, 0x00000010, 0x20400010,
0x20400010, 0x00000000, 0x00404010, 0x20404000,
0x00004010, 0x00404000, 0x20404000, 0x20000000,
0x20004000, 0x00000010, 0x20400010, 0x00404000,
0x20404010, 0x00400000, 0x00004010, 0x20000010,
0x00400000, 0x20004000, 0x20000000, 0x00004010,
0x20000010, 0x20404010, 0x00404000, 0x20400000,
0x00404010, 0x20404000, 0x00000000, 0x20400010,
0x00000010, 0x00004000, 0x20400000, 0x00404010,
0x00004000, 0x00400010, 0x20004010, 0x00000000,
0x20404000, 0x20000000, 0x00400010, 0x20004010
};
static uint32_t SB7[64] = {
0x00200000, 0x04200002, 0x04000802, 0x00000000,
0x00000800, 0x04000802, 0x00200802, 0x04200800,
0x04200802, 0x00200000, 0x00000000, 0x04000002,
0x00000002, 0x04000000, 0x04200002, 0x00000802,
0x04000800, 0x00200802, 0x00200002, 0x04000800,
0x04000002, 0x04200000, 0x04200800, 0x00200002,
0x04200000, 0x00000800, 0x00000802, 0x04200802,
0x00200800, 0x00000002, 0x04000000, 0x00200800,
0x04000000, 0x00200800, 0x00200000, 0x04000802,
0x04000802, 0x04200002, 0x04200002, 0x00000002,
0x00200002, 0x04000000, 0x04000800, 0x00200000,
0x04200800, 0x00000802, 0x00200802, 0x04200800,
0x00000802, 0x04000002, 0x04200802, 0x04200000,
0x00200800, 0x00000000, 0x00000002, 0x04200802,
0x00000000, 0x00200802, 0x04200000, 0x00000800,
0x04000002, 0x04000800, 0x00000800, 0x00200002
};
static uint32_t SB8[64] = {
0x10001040, 0x00001000, 0x00040000, 0x10041040,
0x10000000, 0x10001040, 0x00000040, 0x10000000,
0x00040040, 0x10040000, 0x10041040, 0x00041000,
0x10041000, 0x00041040, 0x00001000, 0x00000040,
0x10040000, 0x10000040, 0x10001000, 0x00001040,
0x00041000, 0x00040040, 0x10040040, 0x10041000,
0x00001040, 0x00000000, 0x00000000, 0x10040040,
0x10000040, 0x10001000, 0x00041040, 0x00040000,
0x00041040, 0x00040000, 0x10041000, 0x00001000,
0x00000040, 0x10040040, 0x00001000, 0x00041040,
0x10001000, 0x00000040, 0x10000040, 0x10040000,
0x10040040, 0x10000000, 0x00040000, 0x10001040,
0x00000000, 0x10041040, 0x00040040, 0x10000040,
0x10040000, 0x10001000, 0x10001040, 0x00000000,
0x10041040, 0x00041000, 0x00041000, 0x00001040,
0x00001040, 0x00040040, 0x10000000, 0x10041000
};
/* PC1: left and right halves bit-swap */
static uint32_t LHs[16] = {
0x00000000, 0x00000001, 0x00000100, 0x00000101,
0x00010000, 0x00010001, 0x00010100, 0x00010101,
0x01000000, 0x01000001, 0x01000100, 0x01000101,
0x01010000, 0x01010001, 0x01010100, 0x01010101
};
static uint32_t RHs[16] = {
0x00000000, 0x01000000, 0x00010000, 0x01010000,
0x00000100, 0x01000100, 0x00010100, 0x01010100,
0x00000001, 0x01000001, 0x00010001, 0x01010001,
0x00000101, 0x01000101, 0x00010101, 0x01010101,
};
/* platform-independant 32-bit integer manipulation macros */
#define GET_UINT32(n,b,i) \
{ \
(n) = ( (uint32_t) (b)[(i) ] << 24 ) \
| ( (uint32_t) (b)[(i) + 1] << 16 ) \
| ( (uint32_t) (b)[(i) + 2] << 8 ) \
| ( (uint32_t) (b)[(i) + 3] ); \
}
#define PUT_UINT32(n,b,i) \
{ \
(b)[(i) ] = (uint8_t) ( (n) >> 24 ); \
(b)[(i) + 1] = (uint8_t) ( (n) >> 16 ); \
(b)[(i) + 2] = (uint8_t) ( (n) >> 8 ); \
(b)[(i) + 3] = (uint8_t) ( (n) ); \
}
/* Initial Permutation macro */
#define DES_IP(X,Y) \
{ \
T = ((X >> 4) ^ Y) & 0x0F0F0F0F; Y ^= T; X ^= (T << 4); \
T = ((X >> 16) ^ Y) & 0x0000FFFF; Y ^= T; X ^= (T << 16); \
T = ((Y >> 2) ^ X) & 0x33333333; X ^= T; Y ^= (T << 2); \
T = ((Y >> 8) ^ X) & 0x00FF00FF; X ^= T; Y ^= (T << 8); \
Y = ((Y << 1) | (Y >> 31)) & 0xFFFFFFFF; \
T = (X ^ Y) & 0xAAAAAAAA; Y ^= T; X ^= T; \
X = ((X << 1) | (X >> 31)) & 0xFFFFFFFF; \
}
/* Final Permutation macro */
#define DES_FP(X,Y) \
{ \
X = ((X << 31) | (X >> 1)) & 0xFFFFFFFF; \
T = (X ^ Y) & 0xAAAAAAAA; X ^= T; Y ^= T; \
Y = ((Y << 31) | (Y >> 1)) & 0xFFFFFFFF; \
T = ((Y >> 8) ^ X) & 0x00FF00FF; X ^= T; Y ^= (T << 8); \
T = ((Y >> 2) ^ X) & 0x33333333; X ^= T; Y ^= (T << 2); \
T = ((X >> 16) ^ Y) & 0x0000FFFF; Y ^= T; X ^= (T << 16); \
T = ((X >> 4) ^ Y) & 0x0F0F0F0F; Y ^= T; X ^= (T << 4); \
}
/* DES round macro */
#define DES_ROUND(X,Y) \
{ \
T = *SK++ ^ X; \
Y ^= SB8[ (T ) & 0x3F ] ^ \
SB6[ (T >> 8) & 0x3F ] ^ \
SB4[ (T >> 16) & 0x3F ] ^ \
SB2[ (T >> 24) & 0x3F ]; \
\
T = *SK++ ^ ((X << 28) | (X >> 4)); \
Y ^= SB7[ (T ) & 0x3F ] ^ \
SB5[ (T >> 8) & 0x3F ] ^ \
SB3[ (T >> 16) & 0x3F ] ^ \
SB1[ (T >> 24) & 0x3F ]; \
}
/* DES key schedule */
static void
des_main_ks( uint32_t SK[32], const uint8_t key[8] ) {
int i;
uint32_t X, Y, T;
GET_UINT32( X, key, 0 );
GET_UINT32( Y, key, 4 );
/* Permuted Choice 1 */
T = ((Y >> 4) ^ X) & 0x0F0F0F0F; X ^= T; Y ^= (T << 4);
T = ((Y ) ^ X) & 0x10101010; X ^= T; Y ^= (T );
X = (LHs[ (X ) & 0xF] << 3) | (LHs[ (X >> 8) & 0xF ] << 2)
| (LHs[ (X >> 16) & 0xF] << 1) | (LHs[ (X >> 24) & 0xF ] )
| (LHs[ (X >> 5) & 0xF] << 7) | (LHs[ (X >> 13) & 0xF ] << 6)
| (LHs[ (X >> 21) & 0xF] << 5) | (LHs[ (X >> 29) & 0xF ] << 4);
Y = (RHs[ (Y >> 1) & 0xF] << 3) | (RHs[ (Y >> 9) & 0xF ] << 2)
| (RHs[ (Y >> 17) & 0xF] << 1) | (RHs[ (Y >> 25) & 0xF ] )
| (RHs[ (Y >> 4) & 0xF] << 7) | (RHs[ (Y >> 12) & 0xF ] << 6)
| (RHs[ (Y >> 20) & 0xF] << 5) | (RHs[ (Y >> 28) & 0xF ] << 4);
X &= 0x0FFFFFFF;
Y &= 0x0FFFFFFF;
/* calculate subkeys */
for( i = 0; i < 16; i++ )
{
if( i < 2 || i == 8 || i == 15 )
{
X = ((X << 1) | (X >> 27)) & 0x0FFFFFFF;
Y = ((Y << 1) | (Y >> 27)) & 0x0FFFFFFF;
}
else
{
X = ((X << 2) | (X >> 26)) & 0x0FFFFFFF;
Y = ((Y << 2) | (Y >> 26)) & 0x0FFFFFFF;
}
*SK++ = ((X << 4) & 0x24000000) | ((X << 28) & 0x10000000)
| ((X << 14) & 0x08000000) | ((X << 18) & 0x02080000)
| ((X << 6) & 0x01000000) | ((X << 9) & 0x00200000)
| ((X >> 1) & 0x00100000) | ((X << 10) & 0x00040000)
| ((X << 2) & 0x00020000) | ((X >> 10) & 0x00010000)
| ((Y >> 13) & 0x00002000) | ((Y >> 4) & 0x00001000)
| ((Y << 6) & 0x00000800) | ((Y >> 1) & 0x00000400)
| ((Y >> 14) & 0x00000200) | ((Y ) & 0x00000100)
| ((Y >> 5) & 0x00000020) | ((Y >> 10) & 0x00000010)
| ((Y >> 3) & 0x00000008) | ((Y >> 18) & 0x00000004)
| ((Y >> 26) & 0x00000002) | ((Y >> 24) & 0x00000001);
*SK++ = ((X << 15) & 0x20000000) | ((X << 17) & 0x10000000)
| ((X << 10) & 0x08000000) | ((X << 22) & 0x04000000)
| ((X >> 2) & 0x02000000) | ((X << 1) & 0x01000000)
| ((X << 16) & 0x00200000) | ((X << 11) & 0x00100000)
| ((X << 3) & 0x00080000) | ((X >> 6) & 0x00040000)
| ((X << 15) & 0x00020000) | ((X >> 4) & 0x00010000)
| ((Y >> 2) & 0x00002000) | ((Y << 8) & 0x00001000)
| ((Y >> 14) & 0x00000808) | ((Y >> 9) & 0x00000400)
| ((Y ) & 0x00000200) | ((Y << 7) & 0x00000100)
| ((Y >> 7) & 0x00000020) | ((Y >> 3) & 0x00000011)
| ((Y << 2) & 0x00000004) | ((Y >> 21) & 0x00000002);
}
}
/* DES 64-bit block encryption/decryption */
static void
des_crypt( const uint32_t SK[32], const uint8_t input[8], uint8_t output[8] ) {
uint32_t X, Y, T;
GET_UINT32( X, input, 0 );
GET_UINT32( Y, input, 4 );
DES_IP( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_ROUND( Y, X ); DES_ROUND( X, Y );
DES_FP( Y, X );
PUT_UINT32( Y, output, 0 );
PUT_UINT32( X, output, 4 );
}
static int
lrandomkey(lua_State *L) {
char tmp[8];
int i;
for (i=0;i<8;i++) {
tmp[i] = random() & 0xff;
}
lua_pushlstring(L, tmp, 8);
return 1;
}
static void
des_key(lua_State *L, uint32_t SK[32]) {
size_t keysz = 0;
const void * key = luaL_checklstring(L, 1, &keysz);
if (keysz != 8) {
luaL_error(L, "Invalid key size %d, need 8 bytes", (int)keysz);
}
des_main_ks(SK, key);
}
static int
ldesencode(lua_State *L) {
uint32_t SK[32];
des_key(L, SK);
size_t textsz = 0;
const uint8_t * text = (const uint8_t *)luaL_checklstring(L, 2, &textsz);
size_t chunksz = (textsz + 8) & ~7;
uint8_t tmp[SMALL_CHUNK];
uint8_t *buffer = tmp;
if (chunksz > SMALL_CHUNK) {
buffer = lua_newuserdata(L, chunksz);
}
int i;
for (i=0;i<(int)textsz-7;i+=8) {
des_crypt(SK, text+i, buffer+i);
}
int bytes = textsz - i;
uint8_t tail[8];
int j;
for (j=0;j<8;j++) {
if (j < bytes) {
tail[j] = text[i+j];
} else if (j==bytes) {
tail[j] = 0x80;
} else {
tail[j] = 0;
}
}
des_crypt(SK, tail, buffer+i);
lua_pushlstring(L, (const char *)buffer, chunksz);
return 1;
}
static int
ldesdecode(lua_State *L) {
uint32_t ESK[32];
des_key(L, ESK);
uint32_t SK[32];
int i;
for( i = 0; i < 32; i += 2 ) {
SK[i] = ESK[30 - i];
SK[i + 1] = ESK[31 - i];
}
size_t textsz = 0;
const uint8_t *text = (const uint8_t *)luaL_checklstring(L, 2, &textsz);
if ((textsz & 7) || textsz == 0) {
return luaL_error(L, "Invalid des crypt text length %d", (int)textsz);
}
uint8_t tmp[SMALL_CHUNK];
uint8_t *buffer = tmp;
if (textsz > SMALL_CHUNK) {
buffer = lua_newuserdata(L, textsz);
}
for (i=0;i<textsz;i+=8) {
des_crypt(SK, text+i, buffer+i);
}
int padding = 1;
for (i=textsz-1;i>=textsz-8;i--) {
if (buffer[i] == 0) {
padding++;
} else if (buffer[i] == 0x80) {
break;
} else {
return luaL_error(L, "Invalid des crypt text");
}
}
if (padding > 8) {
return luaL_error(L, "Invalid des crypt text");
}
lua_pushlstring(L, (const char *)buffer, textsz - padding);
return 1;
}
static void
Hash(const char * str, int sz, uint8_t key[8]) {
uint32_t djb_hash = 5381L;
uint32_t js_hash = 1315423911L;
int i;
for (i=0;i<sz;i++) {
uint8_t c = (uint8_t)str[i];
djb_hash += (djb_hash << 5) + c;
js_hash ^= ((js_hash << 5) + c + (js_hash >> 2));
}
key[0] = djb_hash & 0xff;
key[1] = (djb_hash >> 8) & 0xff;
key[2] = (djb_hash >> 16) & 0xff;
key[3] = (djb_hash >> 24) & 0xff;
key[4] = js_hash & 0xff;
key[5] = (js_hash >> 8) & 0xff;
key[6] = (js_hash >> 16) & 0xff;
key[7] = (js_hash >> 24) & 0xff;
}
static int
lhashkey(lua_State *L) {
size_t sz = 0;
const char * key = luaL_checklstring(L, 1, &sz);
uint8_t realkey[8];
Hash(key,(int)sz,realkey);
lua_pushlstring(L, (const char *)realkey, 8);
return 1;
}
static int
ltohex(lua_State *L) {
static char hex[] = "0123456789abcdef";
size_t sz = 0;
const uint8_t * text = (const uint8_t *)luaL_checklstring(L, 1, &sz);
char tmp[SMALL_CHUNK];
char *buffer = tmp;
if (sz > SMALL_CHUNK/2) {
buffer = lua_newuserdata(L, sz * 2);
}
int i;
for (i=0;i<sz;i++) {
buffer[i*2] = hex[text[i] >> 4];
buffer[i*2+1] = hex[text[i] & 0xf];
}
lua_pushlstring(L, buffer, sz * 2);
return 1;
}
#define HEX(v,c) { char tmp = (char) c; if (tmp >= '0' && tmp <= '9') { v = tmp-'0'; } else { v = tmp - 'a' + 10; } }
static int
lfromhex(lua_State *L) {
size_t sz = 0;
const char * text = luaL_checklstring(L, 1, &sz);
if (sz & 2) {
return luaL_error(L, "Invalid hex text size %d", (int)sz);
}
char tmp[SMALL_CHUNK];
char *buffer = tmp;
if (sz > SMALL_CHUNK*2) {
buffer = lua_newuserdata(L, sz / 2);
}
int i;
for (i=0;i<sz;i+=2) {
uint8_t hi,low;
HEX(hi, text[i]);
HEX(low, text[i+1]);
if (hi > 16 || low > 16) {
return luaL_error(L, "Invalid hex text", text);
}
buffer[i/2] = hi<<4 | low;
}
lua_pushlstring(L, buffer, i/2);
return 1;
}
// Constants are the integer part of the sines of integers (in radians) * 2^32.
const uint32_t k[64] = {
0xd76aa478, 0xe8c7b756, 0x242070db, 0xc1bdceee ,
0xf57c0faf, 0x4787c62a, 0xa8304613, 0xfd469501 ,
0x698098d8, 0x8b44f7af, 0xffff5bb1, 0x895cd7be ,
0x6b901122, 0xfd987193, 0xa679438e, 0x49b40821 ,
0xf61e2562, 0xc040b340, 0x265e5a51, 0xe9b6c7aa ,
0xd62f105d, 0x02441453, 0xd8a1e681, 0xe7d3fbc8 ,
0x21e1cde6, 0xc33707d6, 0xf4d50d87, 0x455a14ed ,
0xa9e3e905, 0xfcefa3f8, 0x676f02d9, 0x8d2a4c8a ,
0xfffa3942, 0x8771f681, 0x6d9d6122, 0xfde5380c ,
0xa4beea44, 0x4bdecfa9, 0xf6bb4b60, 0xbebfbc70 ,
0x289b7ec6, 0xeaa127fa, 0xd4ef3085, 0x04881d05 ,
0xd9d4d039, 0xe6db99e5, 0x1fa27cf8, 0xc4ac5665 ,
0xf4292244, 0x432aff97, 0xab9423a7, 0xfc93a039 ,
0x655b59c3, 0x8f0ccc92, 0xffeff47d, 0x85845dd1 ,
0x6fa87e4f, 0xfe2ce6e0, 0xa3014314, 0x4e0811a1 ,
0xf7537e82, 0xbd3af235, 0x2ad7d2bb, 0xeb86d391 };
// r specifies the per-round shift amounts
const uint32_t r[] = {7, 12, 17, 22, 7, 12, 17, 22, 7, 12, 17, 22, 7, 12, 17, 22,
5, 9, 14, 20, 5, 9, 14, 20, 5, 9, 14, 20, 5, 9, 14, 20,
4, 11, 16, 23, 4, 11, 16, 23, 4, 11, 16, 23, 4, 11, 16, 23,
6, 10, 15, 21, 6, 10, 15, 21, 6, 10, 15, 21, 6, 10, 15, 21};
// leftrotate function definition
#define LEFTROTATE(x, c) (((x) << (c)) | ((x) >> (32 - (c))))
static void
hmac(uint32_t x[2], uint32_t y[2], uint32_t result[2]) {
uint32_t w[16];
uint32_t a, b, c, d, f, g, temp;
int i;
a = 0x67452301u;
b = 0xefcdab89u;
c = 0x98badcfeu;
d = 0x10325476u;
for (i=0;i<16;i+=4) {
w[i] = x[1];
w[i+1] = x[0];
w[i+2] = y[1];
w[i+3] = y[0];
}
for(i = 0; i<64; i++) {
if (i < 16) {
f = (b & c) | ((~b) & d);
g = i;
} else if (i < 32) {
f = (d & b) | ((~d) & c);
g = (5*i + 1) % 16;
} else if (i < 48) {
f = b ^ c ^ d;
g = (3*i + 5) % 16;
} else {
f = c ^ (b | (~d));
g = (7*i) % 16;
}
temp = d;
d = c;
c = b;
b = b + LEFTROTATE((a + f + k[i] + w[g]), r[i]);
a = temp;
}
result[0] = c^d;
result[1] = a^b;
}
static void
read64(lua_State *L, uint32_t xx[2], uint32_t yy[2]) {
size_t sz = 0;
const uint8_t *x = (const uint8_t *)luaL_checklstring(L, 1, &sz);
if (sz != 8) {
luaL_error(L, "Invalid hmac x");
}
const uint8_t *y = (const uint8_t *)luaL_checklstring(L, 2, &sz);
if (sz != 8) {
luaL_error(L, "Invalid hmac y");
}
xx[0] = x[0] | x[1]<<8 | x[2]<<16 | x[3]<<24;
xx[1] = x[4] | x[5]<<8 | x[6]<<16 | x[7]<<24;
yy[0] = y[0] | y[1]<<8 | y[2]<<16 | y[3]<<24;
yy[1] = y[4] | y[5]<<8 | y[6]<<16 | y[7]<<24;
}
static int
lhmac64(lua_State *L) {
uint32_t x[2], y[2];
read64(L, x, y);
uint32_t result[2];
hmac(x,y,result);
uint8_t tmp[8];
tmp[0] = result[0] & 0xff;
tmp[1] = (result[0] >> 8 )& 0xff;
tmp[2] = (result[0] >> 16 )& 0xff;
tmp[3] = (result[0] >> 24 )& 0xff;
tmp[4] = result[1] & 0xff;
tmp[5] = (result[1] >> 8 )& 0xff;
tmp[6] = (result[1] >> 16 )& 0xff;
tmp[7] = (result[1] >> 24 )& 0xff;
lua_pushlstring(L, (const char *)tmp, 8);
return 1;
}
// powmodp64 for DH-key exchange
// The biggest 64bit prime
#define P 0xffffffffffffffc5ull
static inline uint64_t
mul_mod_p(uint64_t a, uint64_t b) {
uint64_t m = 0;
while(b) {
if(b&1) {
uint64_t t = P-a;
if ( m >= t) {
m -= t;
} else {
m += a;
}
}
if (a >= P - a) {
a = a * 2 - P;
} else {
a = a * 2;
}
b>>=1;
}
return m;
}
static inline uint64_t
pow_mod_p(uint64_t a, uint64_t b) {
if (b==1) {
return a;
}
uint64_t t = pow_mod_p(a, b>>1);
t = mul_mod_p(t,t);
if (b % 2) {
t = mul_mod_p(t, a);
}
return t;
}
// calc a^b % p
static uint64_t
powmodp(uint64_t a, uint64_t b) {
if (a > P)
a%=P;
return pow_mod_p(a,b);
}
static void
push64(lua_State *L, uint64_t r) {
uint8_t tmp[8];
tmp[0] = r & 0xff;
tmp[1] = (r >> 8 )& 0xff;
tmp[2] = (r >> 16 )& 0xff;
tmp[3] = (r >> 24 )& 0xff;
tmp[4] = (r >> 32 )& 0xff;
tmp[5] = (r >> 40 )& 0xff;
tmp[6] = (r >> 48 )& 0xff;
tmp[7] = (r >> 56 )& 0xff;
lua_pushlstring(L, (const char *)tmp, 8);
}
static int
ldhsecret(lua_State *L) {
uint32_t x[2], y[2];
read64(L, x, y);
uint64_t r = powmodp((uint64_t)x[0] | (uint64_t)x[1]<<32,
(uint64_t)y[0] | (uint64_t)y[1]<<32);
push64(L, r);
return 1;
}
#define G 5
static int
ldhexchange(lua_State *L) {
size_t sz = 0;
const uint8_t *x = (const uint8_t *)luaL_checklstring(L, 1, &sz);
if (sz != 8) {
luaL_error(L, "Invalid hmac x");
}
uint32_t xx[2];
xx[0] = x[0] | x[1]<<8 | x[2]<<16 | x[3]<<24;
xx[1] = x[4] | x[5]<<8 | x[6]<<16 | x[7]<<24;
uint64_t r = powmodp(5, (uint64_t)xx[0] | (uint64_t)xx[1]<<32);
push64(L, r);
return 1;
}
// base64
static int
lb64encode(lua_State *L) {
static const char* encoding = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
size_t sz = 0;
const uint8_t * text = (const uint8_t *)luaL_checklstring(L, 1, &sz);
int encode_sz = (sz + 2)/3*4;
char tmp[SMALL_CHUNK];
char *buffer = tmp;
if (encode_sz > SMALL_CHUNK) {
buffer = lua_newuserdata(L, encode_sz);
}
int i,j;
j=0;
for (i=0;i<(int)sz-2;i+=3) {
uint32_t v = text[i] << 16 | text[i+1] << 8 | text[i+2];
buffer[j] = encoding[v >> 18];
buffer[j+1] = encoding[(v >> 12) & 0x3f];
buffer[j+2] = encoding[(v >> 6) & 0x3f];
buffer[j+3] = encoding[(v) & 0x3f];
j+=4;
}
int padding = sz-i;
uint32_t v;
switch(padding) {
case 1 :
v = text[i];
buffer[j] = encoding[v >> 2];
buffer[j+1] = encoding[(v & 3) << 4];
buffer[j+2] = '=';
buffer[j+3] = '=';
break;
case 2 :
v = text[i] << 8 | text[i+1];
buffer[j] = encoding[v >> 10];
buffer[j+1] = encoding[(v >> 4) & 0x3f];
buffer[j+2] = encoding[(v & 0xf) << 2];
buffer[j+3] = '=';
break;
}
lua_pushlstring(L, buffer, encode_sz);
return 1;
}
static inline int
b64index(uint8_t c) {
static const int decoding[] = {62,-1,-1,-1,63,52,53,54,55,56,57,58,59,60,61,-1,-1,-1,-2,-1,-1,-1,0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,-1,-1,-1,-1,-1,-1,26,27,28,29,30,31,32,33,34,35,36,37,38,39,40,41,42,43,44,45,46,47,48,49,50,51};
int decoding_size = sizeof(decoding)/sizeof(decoding[0]);
if (c<43) {
return -1;
}
c -= 43;
if (c>=decoding_size)
return -1;
return decoding[c];
}
static int
lb64decode(lua_State *L) {
size_t sz = 0;
const uint8_t * text = (const uint8_t *)luaL_checklstring(L, 1, &sz);
int decode_sz = (sz+3)/4*3;
char tmp[SMALL_CHUNK];
char *buffer = tmp;
if (decode_sz > SMALL_CHUNK) {
buffer = lua_newuserdata(L, decode_sz);
}
int i,j;
int output = 0;
for (i=0;i<sz;) {
int padding = 0;
int c[4];
for (j=0;j<4;) {
if (i>=sz) {
return luaL_error(L, "Invalid base64 text");
}
c[j] = b64index(text[i]);
if (c[j] == -1) {
++i;
continue;
}
if (c[j] == -2) {
++padding;
}
++i;
++j;
}
uint32_t v;
switch (padding) {
case 0:
v = (unsigned)c[0] << 18 | c[1] << 12 | c[2] << 6 | c[3];
buffer[output] = v >> 16;
buffer[output+1] = (v >> 8) & 0xff;
buffer[output+2] = v & 0xff;
output += 3;
break;
case 1:
if (c[3] != -2 || (c[2] & 3)!=0) {
return luaL_error(L, "Invalid base64 text");
}
v = (unsigned)c[0] << 10 | c[1] << 4 | c[2] >> 2 ;
buffer[output] = v >> 8;
buffer[output+1] = v & 0xff;
output += 2;
break;
case 2:
if (c[3] != -2 || c[2] != -2 || (c[1] & 0xf) !=0) {
return luaL_error(L, "Invalid base64 text");
}
v = (unsigned)c[0] << 2 | c[1] >> 4;
buffer[output] = v;
++ output;
break;
default:
return luaL_error(L, "Invalid base64 text");
}
}
lua_pushlstring(L, buffer, output);
return 1;
}
int
luaopen_crypt(lua_State *L) {
luaL_checkversion(L);
srandom(time(NULL));
luaL_Reg l[] = {
{ "hashkey", lhashkey },
{ "randomkey", lrandomkey },
{ "desencode", ldesencode },
{ "desdecode", ldesdecode },
{ "hexencode", ltohex },
{ "hexdecode", lfromhex },
{ "hmac64", lhmac64 },
{ "dhexchange", ldhexchange },
{ "dhsecret", ldhsecret },
{ "base64encode", lb64encode },
{ "base64decode", lb64decode },
{ NULL, NULL },
};
luaL_newlib(L,l);
return 1;
}

View File

@@ -506,7 +506,7 @@ op_insert(lua_State *L) {
int i; int i;
for (i=1;i<=s;i++) { for (i=1;i<=s;i++) {
lua_rawgeti(L,3,i); lua_rawgeti(L,3,i);
document doc = lua_touserdata(L,3); document doc = lua_touserdata(L,-1);
luaL_addlstring(&b, (const char *)doc, get_length(doc)); luaL_addlstring(&b, (const char *)doc, get_length(doc));
lua_pop(L,1); lua_pop(L,1);
} }

View File

@@ -393,13 +393,13 @@ lpop(lua_State *L) {
*/ */
static const char * static const char *
tolstring(lua_State *L, size_t *sz) { tolstring(lua_State *L, size_t *sz, int index) {
const char * ptr; const char * ptr;
if (lua_isuserdata(L,1)) { if (lua_isuserdata(L,index)) {
ptr = (const char *)lua_touserdata(L,1); ptr = (const char *)lua_touserdata(L,index);
*sz = (size_t)luaL_checkinteger(L, 2); *sz = (size_t)luaL_checkinteger(L, index+1);
} else { } else {
ptr = luaL_checklstring(L, 1, sz); ptr = luaL_checklstring(L, index, sz);
} }
return ptr; return ptr;
} }
@@ -413,7 +413,7 @@ write_size(uint8_t * buffer, int len) {
static int static int
lpack(lua_State *L) { lpack(lua_State *L) {
size_t len; size_t len;
const char * ptr = tolstring(L, &len); const char * ptr = tolstring(L, &len, 1);
if (len > 0x10000) { if (len > 0x10000) {
return luaL_error(L, "Invalid size (too long) of data : %d", (int)len); return luaL_error(L, "Invalid size (too long) of data : %d", (int)len);
} }
@@ -433,7 +433,7 @@ lpack_string(lua_State *L) {
uint8_t tmp[SMALLSTRING+2]; uint8_t tmp[SMALLSTRING+2];
size_t len; size_t len;
uint8_t *buffer; uint8_t *buffer;
const char * ptr = tolstring(L, &len); const char * ptr = tolstring(L, &len, 1);
if (len > 0x10000) { if (len > 0x10000) {
return luaL_error(L, "Invalid size (too long) of data : %d", (int)len); return luaL_error(L, "Invalid size (too long) of data : %d", (int)len);
} }
@@ -451,6 +451,34 @@ lpack_string(lua_State *L) {
return 1; return 1;
} }
static int
lpack_padding(lua_State *L) {
uint8_t tmp[SMALLSTRING+2];
size_t content_sz;
uint8_t *buffer;
const char * ptr = tolstring(L, &content_sz, 2);
size_t cookie_sz = 0;
const char * cookie = luaL_checklstring(L,1,&cookie_sz);
size_t len = cookie_sz + content_sz;
if (len > 0x10000) {
return luaL_error(L, "Invalid size (too long) of data : %d", (int)len);
}
if (len <= SMALLSTRING) {
buffer = tmp;
} else {
buffer = lua_newuserdata(L, len + 2);
}
write_size(buffer, len);
memcpy(buffer+2, ptr, content_sz);
memcpy(buffer+2+content_sz, cookie, cookie_sz);
lua_pushlstring(L, (const char *)buffer, len+2);
return 1;
}
static int static int
ltostring(lua_State *L) { ltostring(lua_State *L) {
void * ptr = lua_touserdata(L, 1); void * ptr = lua_touserdata(L, 1);
@@ -458,8 +486,19 @@ ltostring(lua_State *L) {
if (ptr == NULL) { if (ptr == NULL) {
lua_pushliteral(L, ""); lua_pushliteral(L, "");
} else { } else {
lua_pushlstring(L, (const char *)ptr, size); if (lua_isnumber(L, 3)) {
skynet_free(ptr); int offset = lua_tointeger(L, 3);
if (offset < 0) {
return luaL_error(L, "Invalid offset %d", offset);
}
if (offset > size) {
offset = size;
}
lua_pushlstring(L, (const char *)ptr + offset, size-offset);
} else {
lua_pushlstring(L, (const char *)ptr, size);
skynet_free(ptr);
}
} }
return 1; return 1;
} }
@@ -471,6 +510,7 @@ luaopen_netpack(lua_State *L) {
{ "pop", lpop }, { "pop", lpop },
{ "pack", lpack }, { "pack", lpack },
{ "pack_string", lpack_string }, { "pack_string", lpack_string },
{ "pack_padding", lpack_padding },
{ "clear", lclear }, { "clear", lclear },
{ "tostring", ltostring }, { "tostring", ltostring },
{ NULL, NULL }, { NULL, NULL },

View File

@@ -14,6 +14,7 @@
#define BACKLOG 32 #define BACKLOG 32
// 2 ** 12 == 4096 // 2 ** 12 == 4096
#define LARGE_PAGE_NODE 12 #define LARGE_PAGE_NODE 12
#define BUFFER_LIMIT (256 * 1024)
struct buffer_node { struct buffer_node {
char * msg; char * msg;
@@ -245,6 +246,7 @@ lclearbuffer(lua_State *L) {
while(sb->head) { while(sb->head) {
return_free_node(L,2,sb); return_free_node(L,2,sb);
} }
sb->size = 0;
return 0; return 0;
} }
@@ -263,6 +265,7 @@ lreadall(lua_State *L) {
return_free_node(L,2,sb); return_free_node(L,2,sb);
} }
luaL_pushresult(&b); luaL_pushresult(&b);
sb->size = 0;
return 1; return 1;
} }
@@ -476,6 +479,14 @@ lstart(lua_State *L) {
return 0; return 0;
} }
static int
lnodelay(lua_State *L) {
struct skynet_context * ctx = lua_touserdata(L, lua_upvalueindex(1));
int id = luaL_checkinteger(L, 1);
skynet_socket_nodelay(ctx,id);
return 0;
}
int int
luaopen_socketdriver(lua_State *L) { luaopen_socketdriver(lua_State *L) {
luaL_checkversion(L); luaL_checkversion(L);
@@ -502,6 +513,7 @@ luaopen_socketdriver(lua_State *L) {
{ "lsend", lsendlow }, { "lsend", lsendlow },
{ "bind", lbind }, { "bind", lbind },
{ "start", lstart }, { "start", lstart },
{ "nodelay", lnodelay },
{ NULL, NULL }, { NULL, NULL },
}; };
lua_getfield(L, LUA_REGISTRYINDEX, "skynet_context"); lua_getfield(L, LUA_REGISTRYINDEX, "skynet_context");

View File

@@ -16,6 +16,10 @@ function cluster.open(port)
end end
end end
function cluster.reload()
skynet.call(clusterd, "lua", "reload")
end
skynet.init(function() skynet.init(function()
clusterd = skynet.uniqueservice("clusterd") clusterd = skynet.uniqueservice("clusterd")
end) end)

View File

@@ -10,5 +10,9 @@ function datacenter.set(...)
return skynet.call("DATACENTER", "lua", "UPDATE", ...) return skynet.call("DATACENTER", "lua", "UPDATE", ...)
end end
function datacenter.wait(...)
return skynet.call("DATACENTER", "lua", "WAIT", ...)
end
return datacenter return datacenter

114
lualib/http/httpc.lua Normal file
View File

@@ -0,0 +1,114 @@
local socket = require "http.sockethelper"
local url = require "http.url"
local internal = require "http.internal"
local string = string
local table = table
local httpc = {}
local function request(fd, method, host, url, recvheader, header, content)
local read = socket.readfunc(fd)
local write = socket.writefunc(fd)
local header_content = ""
if header then
for k,v in pairs(header) do
header_content = string.format("%s%s:%s\r\n", header_content, k, v)
end
end
if content then
local data = string.format("%s %s HTTP/1.1\r\nhost:%s\r\ncontent-length:%d\r\n%s\r\n%s", method, url, host, #content, header_content, content)
write(data)
else
local request_header = string.format("%s %s HTTP/1.1\r\nhost:%s\r\ncontent-length:0\r\n%s\r\n", method, url, host, header_content)
write(request_header)
end
local tmpline = {}
local body = internal.recvheader(read, tmpline, "")
if not body then
error(socket.socket_error)
end
local statusline = tmpline[1]
local code, info = statusline:match "HTTP/[%d%.]+%s+([%d]+)%s+(.*)$"
code = assert(tonumber(code))
local header = internal.parseheader(tmpline,2,recvheader or {})
if not header then
error("Invalid HTTP response header")
end
local length = header["content-length"]
if length then
length = tonumber(length)
end
local mode = header["transfer-encoding"]
if mode then
if mode ~= "identity" and mode ~= "chunked" then
error ("Unsupport transfer-encoding")
end
end
if mode == "chunked" then
body, header = internal.recvchunkedbody(read, nil, header, body)
if not body then
error("Invalid response body")
end
else
-- identity mode
if length then
if #body >= length then
body = body:sub(1,length)
else
local padding = read(length - #body)
body = body .. padding
end
else
body = nil
end
end
return code, body
end
function httpc.request(method, host, url, recvheader, header, content)
local hostname, port = host:match"([^:]+):?(%d*)$"
if port == "" then
port = 80
else
port = tonumber(port)
end
local fd = socket.connect(hostname, port)
local ok , statuscode, body = pcall(request, fd,method, host, url, recvheader, header, content)
if ok then
return statuscode, body
else
socket.close(fd)
error(statuscode)
end
end
function httpc.get(...)
return httpc.request("GET", ...)
end
local function escape(s)
return (string.gsub(s, "([^A-Za-z0-9_])", function(c)
return string.format("%%%02X", string.byte(c))
end))
end
function httpc.post(host, url, form, recvheader)
local header = {
["content-type"] = "application/x-www-form-urlencoded"
}
local body = {}
for k,v in pairs(form) do
table.insert(body, string.format("%s=%s",escape(k),escape(v)))
end
return httpc.request("POST", host, url, recvheader, header, table.concat(body , "&"))
end
return httpc

145
lualib/http/httpd.lua Normal file
View File

@@ -0,0 +1,145 @@
local internal = require "http.internal"
local table = table
local httpd = {}
local http_status_msg = {
[100] = "Continue",
[101] = "Switching Protocols",
[200] = "OK",
[201] = "Created",
[202] = "Accepted",
[203] = "Non-Authoritative Information",
[204] = "No Content",
[205] = "Reset Content",
[206] = "Partial Content",
[300] = "Multiple Choices",
[301] = "Moved Permanently",
[302] = "Found",
[303] = "See Other",
[304] = "Not Modified",
[305] = "Use Proxy",
[307] = "Temporary Redirect",
[400] = "Bad Request",
[401] = "Unauthorized",
[402] = "Payment Required",
[403] = "Forbidden",
[404] = "Not Found",
[405] = "Method Not Allowed",
[406] = "Not Acceptable",
[407] = "Proxy Authentication Required",
[408] = "Request Time-out",
[409] = "Conflict",
[410] = "Gone",
[411] = "Length Required",
[412] = "Precondition Failed",
[413] = "Request Entity Too Large",
[414] = "Request-URI Too Large",
[415] = "Unsupported Media Type",
[416] = "Requested range not satisfiable",
[417] = "Expectation Failed",
[500] = "Internal Server Error",
[501] = "Not Implemented",
[502] = "Bad Gateway",
[503] = "Service Unavailable",
[504] = "Gateway Time-out",
[505] = "HTTP Version not supported",
}
local function readall(readbytes, bodylimit)
local tmpline = {}
local body = internal.recvheader(readbytes, tmpline, "")
if not body then
return 413 -- Request Entity Too Large
end
local request = assert(tmpline[1])
local method, url, httpver = request:match "^(%a+)%s+(.-)%s+HTTP/([%d%.]+)$"
assert(method and url and httpver)
httpver = assert(tonumber(httpver))
if httpver < 1.0 or httpver > 1.1 then
return 505 -- HTTP Version not supported
end
local header = internal.parseheader(tmpline,2,{})
if not header then
return 400 -- Bad request
end
local length = header["content-length"]
if length then
length = tonumber(length)
end
local mode = header["transfer-encoding"]
if mode then
if mode ~= "identity" and mode ~= "chunked" then
return 501 -- Not Implemented
end
end
if mode == "chunked" then
body, header = internal.recvchunkedbody(readbytes, bodylimit, header, body)
if not body then
return 413
end
else
-- identity mode
if length then
if length > bodylimit then
return 413
end
if #body >= length then
body = body:sub(1,length)
else
local padding = readbytes(length - #body)
body = body .. padding
end
end
end
return 200, url, method, header, body
end
function httpd.read_request(...)
local ok, code, url, method, header, body = pcall(readall, ...)
if ok then
return code, url, method, header, body
else
return nil, code
end
end
local function writeall(writefunc, statuscode, bodyfunc, header)
local statusline = string.format("HTTP/1.1 %03d %s\r\n", statuscode, http_status_msg[statuscode] or "")
writefunc(statusline)
if header then
for k,v in pairs(header) do
writefunc(string.format("%s: %s\r\n", k,v))
end
end
local t = type(bodyfunc)
if t == "string" then
writefunc(string.format("content-length: %d\r\n\r\n", #bodyfunc))
writefunc(bodyfunc)
elseif t == "function" then
writefunc("transfer-encoding: chunked\r\n")
while true do
local s = bodyfunc()
if s then
if s ~= "" then
writefunc(string.format("\r\n%x\r\n", #s))
writefunc(s)
end
else
writefunc("\r\n0\r\n\r\n")
end
end
else
assert(t == "nil")
writefunc("\r\n")
end
end
function httpd.write_response(...)
return pcall(writeall, ...)
end
return httpd

135
lualib/http/internal.lua Normal file
View File

@@ -0,0 +1,135 @@
local M = {}
local LIMIT = 8192
local function chunksize(readbytes, body)
while true do
if #body > 128 then
return
end
body = body .. readbytes()
local f,e = body:find("\r\n",1,true)
if f then
return tonumber(body:sub(1,f-1),16), body:sub(e+1)
end
end
end
local function readcrln(readbytes, body)
if #body >= 2 then
if body:sub(1,2) ~= "\r\n" then
return
end
return body:sub(3)
else
body = body .. readbytes(2-#body)
if body ~= "\r\n" then
return
end
return ""
end
end
function M.recvheader(readbytes, lines, header)
if #header >= 2 then
if header:find "^\r\n" then
return header:sub(3)
end
end
local result
local e = header:find("\r\n\r\n", 1, true)
if e then
result = header:sub(e+4)
else
while true do
local bytes = readbytes()
header = header .. bytes
if #header > LIMIT then
return
end
e = header:find("\r\n\r\n", -#bytes-3, true)
if e then
result = header:sub(e+4)
break
end
if header:find "^\r\n" then
return header:sub(3)
end
end
end
for v in header:gmatch("(.-)\r\n") do
if v == "" then
break
end
table.insert(lines, v)
end
return result
end
function M.parseheader(lines, from, header)
local name, value
for i=from,#lines do
local line = lines[i]
if line:byte(1) == 9 then -- tab, append last line
if name == nil then
return
end
header[name] = header[name] .. line:sub(2)
else
name, value = line:match "^(.-):%s*(.*)"
if name == nil or value == nil then
return
end
name = name:lower()
if header[name] then
header[name] = header[name] .. ", " .. value
else
header[name] = value
end
end
end
return header
end
function M.recvchunkedbody(readbytes, bodylimit, header, body)
local result = ""
local size = 0
while true do
local sz
sz , body = chunksize(readbytes, body)
if not sz then
return
end
if sz == 0 then
break
end
size = size + sz
if bodylimit and size > bodylimit then
return
end
if #body >= sz then
result = result .. body:sub(1,sz)
body = body:sub(sz+1)
else
result = result .. body .. readbytes(sz - #body)
body = ""
end
body = readcrln(readbytes, body)
if not body then
return
end
end
local tmpline = {}
body = M.recvheader(readbytes, tmpline, body)
if not body then
return
end
header = M.parseheader(tmpline,1,header)
return result, header
end
return M

View File

@@ -0,0 +1,43 @@
local socket = require "socket"
local readbytes = socket.read
local writebytes = socket.write
local sockethelper = {}
local socket_error = setmetatable({} , { __tostring = function() return "[Socket Error]" end })
sockethelper.socket_error = socket_error
function sockethelper.readfunc(fd)
return function (sz)
local ret = readbytes(fd, sz)
if ret then
return ret
else
error(socket_error)
end
end
end
function sockethelper.writefunc(fd)
return function(content)
local ok = writebytes(fd, content)
if not ok then
error(socket_error)
end
end
end
function sockethelper.connect(host, port)
local fd = socket.open(host, port)
if fd then
return fd
end
error(socket_error)
end
function sockethelper.close(fd)
socket.close(fd)
end
return sockethelper

28
lualib/http/url.lua Normal file
View File

@@ -0,0 +1,28 @@
local url = {}
local function decode_func(c)
return string.char(tonumber(c, 16))
end
local function decode(str)
local str = str:gsub('+', ' ')
return str:gsub("%%(..)", decode_func)
end
function url.parse(u)
local path,query = u:match "([^?]*)%??(.*)"
if path then
path = decode(path)
end
return path, query
end
function url.parse_query(q)
local r = {}
for k,v in q:gmatch "(.-)=([^&]*)&?" do
r[decode(k)] = decode(v)
end
return r
end
return url

View File

@@ -3,11 +3,13 @@ for word in string.gmatch(..., "%S+") do
table.insert(args, word) table.insert(args, word)
end end
SERVICE_NAME = args[1]
local main, pattern local main, pattern
local err = {} local err = {}
for pat in string.gmatch(LUA_SERVICE, "([^;]+);*") do for pat in string.gmatch(LUA_SERVICE, "([^;]+);*") do
local filename = string.gsub(pat, "?", args[1]) local filename = string.gsub(pat, "?", SERVICE_NAME)
local f, msg = loadfile(filename) local f, msg = loadfile(filename)
if not f then if not f then
table.insert(err, msg) table.insert(err, msg)

View File

@@ -1,96 +1,125 @@
local bson = require "bson" local bson = require "bson"
local socket = require "socket" local socket = require "socket"
local socketchannel = require "socketchannel" local socketchannel = require "socketchannel"
local skynet = require "skynet" local skynet = require "skynet"
local driver = require "mongo.driver" local driver = require "mongo.driver"
local md5 = require "md5" local md5 = require "md5"
local rawget = rawget local rawget = rawget
local assert = assert local assert = assert
local bson_encode = bson.encode local bson_encode = bson.encode
local bson_encode_order = bson.encode_order local bson_encode_order = bson.encode_order
local bson_decode = bson.decode local bson_decode = bson.decode
local empty_bson = bson_encode {} local empty_bson = bson_encode {}
local mongo = {} local mongo = {}
mongo.null = assert(bson.null) mongo.null = assert(bson.null)
mongo.maxkey = assert(bson.maxkey) mongo.maxkey = assert(bson.maxkey)
mongo.minkey = assert(bson.minkey) mongo.minkey = assert(bson.minkey)
mongo.type = assert(bson.type) mongo.type = assert(bson.type)
local mongo_cursor = {} local mongo_cursor = {}
local cursor_meta = { local cursor_meta = {
__index = mongo_cursor, __index = mongo_cursor,
} }
local mongo_client = {} local mongo_client = {}
local client_meta = { local client_meta = {
__index = function(self, key) __index = function(self, key)
return rawget(mongo_client, key) or self:getDB(key) return rawget(mongo_client, key) or self:getDB(key)
end, end,
__tostring = function (self) __tostring = function (self)
local port_string local port_string
if self.port then if self.port then
port_string = ":" .. tostring(self.port) port_string = ":" .. tostring(self.port)
else else
port_string = "" port_string = ""
end end
return "[mongo client : " .. self.host .. port_string .."]" return "[mongo client : " .. self.host .. port_string .."]"
end, end,
-- DO NOT need disconnect, because channel will shutdown during gc -- DO NOT need disconnect, because channel will shutdown during gc
} }
local mongo_db = {} local mongo_db = {}
local db_meta = { local db_meta = {
__index = function (self, key) __index = function (self, key)
return rawget(mongo_db, key) or self:getCollection(key) return rawget(mongo_db, key) or self:getCollection(key)
end, end,
__tostring = function (self) __tostring = function (self)
return "[mongo db : " .. self.name .. "]" return "[mongo db : " .. self.name .. "]"
end end
} }
local mongo_collection = {} local mongo_collection = {}
local collection_meta = { local collection_meta = {
__index = function(self, key) __index = function(self, key)
return rawget(mongo_collection, key) or self:getCollection(key) return rawget(mongo_collection, key) or self:getCollection(key)
end , end ,
__tostring = function (self) __tostring = function (self)
return "[mongo collection : " .. self.full_name .. "]" return "[mongo collection : " .. self.full_name .. "]"
end end
} }
local function dispatch_reply(so) local function dispatch_reply(so)
local len_reply = so:read(4) local len_reply = so:read(4)
local reply = so:read(driver.length(len_reply)) local reply = so:read(driver.length(len_reply))
local result = { result = {} } local result = { result = {} }
local succ, reply_id, document, cursor_id, startfrom = driver.reply(reply, result.result) local succ, reply_id, document, cursor_id, startfrom = driver.reply(reply, result.result)
result.document = document result.document = document
result.cursor_id = cursor_id result.cursor_id = cursor_id
result.startfrom = startfrom result.startfrom = startfrom
result.data = reply result.data = reply
return reply_id, succ, result return reply_id, succ, result
end end
local function mongo_auth(mongoc) local function __parse_addr(addr)
local user = rawget(mongoc, "username") local host, port = string.match(addr, "([^:]+):(.+)")
local pass = rawget(mongoc, "password") return host, tonumber(port)
end
local function mongo_auth(mongoc)
local user = rawget(mongoc, "username")
local pass = rawget(mongoc, "password")
if user == nil or pass == nil then
return
end
return function() return function()
assert(mongoc:auth(user, pass)) if user ~= nil and pass ~= nil then
assert(mongoc:auth(user, pass))
end
local rs_data = mongoc:runCommand("ismaster")
if rs_data.ok == 1 then
if rs_data.hosts then
local backup = {}
for _, v in ipairs(rs_data.hosts) do
local host, port = __parse_addr(v)
table.insert(backup, {host = host, port = port})
end
mongoc.__sock:changebackup(backup)
end
if rs_data.ismaster then
return
else
local host, port = __parse_addr(rs_data.primary)
mongoc.host = host
mongoc.port = port
mongoc.__sock:changehost(host, port)
end
end
end end
end end
function mongo.client( conf ) function mongo.client( conf )
local obj = { local first = conf
host = conf.host, local backup = nil
port = conf.port or 27017, if conf.rs then
first = conf.rs[1]
backup = conf.rs
end
local obj = {
host = first.host,
port = first.port or 27017,
} }
obj.__id = 0 obj.__id = 0
@@ -99,9 +128,10 @@ function mongo.client( conf )
port = obj.port, port = obj.port,
response = dispatch_reply, response = dispatch_reply,
auth = mongo_auth(obj), auth = mongo_auth(obj),
backup = backup,
} }
setmetatable(obj, client_meta) setmetatable(obj, client_meta)
obj.__sock:connect(true) -- try connect only once obj.__sock:connect(true) -- try connect only once
return obj return obj
end end
@@ -109,32 +139,32 @@ function mongo_client:getDB(dbname)
local db = { local db = {
connection = self, connection = self,
name = dbname, name = dbname,
full_name = dbname, full_name = dbname,
database = false, database = false,
__cmd = dbname .. "." .. "$cmd", __cmd = dbname .. "." .. "$cmd",
} }
db.database = db db.database = db
return setmetatable(db, db_meta) return setmetatable(db, db_meta)
end end
function mongo_client:disconnect() function mongo_client:disconnect()
if self.__sock then if self.__sock then
local so = self.__sock local so = self.__sock
self.__sock = false self.__sock = false
so:close() so:close()
end end
end end
function mongo_client:genId() function mongo_client:genId()
local id = self.__id + 1 local id = self.__id + 1
self.__id = id self.__id = id
return id return id
end end
function mongo_client:runCommand(...) function mongo_client:runCommand(...)
if not self.admin then if not self.admin then
self.admin = self:getDB "admin" self.admin = self:getDB "admin"
end end
return self.admin:runCommand(...) return self.admin:runCommand(...)
end end
@@ -146,14 +176,14 @@ function mongo_client:auth(user,password)
return false return false
end end
local key = md5.sumhexa(string.format("%s%s%s",result.nonce,user,password)) local key = md5.sumhexa(string.format("%s%s%s",result.nonce,user,password))
local result= self:runCommand ("authenticate",1,"user",user,"nonce",result.nonce,"key",key) local result= self:runCommand ("authenticate",1,"user",user,"nonce",result.nonce,"key",key)
return result.ok == 1 return result.ok == 1
end end
function mongo_client:logout() function mongo_client:logout()
local result = self:runCommand "logout" local result = self:runCommand "logout"
return result.ok == 1 return result.ok == 1
end end
function mongo_db:runCommand(cmd,cmd_v,...) function mongo_db:runCommand(cmd,cmd_v,...)
@@ -166,18 +196,18 @@ function mongo_db:runCommand(cmd,cmd_v,...)
else else
bson_cmd = bson_encode_order(cmd,cmd_v,...) bson_cmd = bson_encode_order(cmd,cmd_v,...)
end end
local pack = driver.query(request_id, 0, self.__cmd, 0, 1, bson_cmd) local pack = driver.query(request_id, 0, self.__cmd, 0, 1, bson_cmd)
-- we must hold req (req.data), because req.document is a lightuserdata, it's a pointer to the string (req.data) -- we must hold req (req.data), because req.document is a lightuserdata, it's a pointer to the string (req.data)
local req = sock:request(pack, request_id) local req = sock:request(pack, request_id)
local doc = req.document local doc = req.document
return bson_decode(doc) return bson_decode(doc)
end end
function mongo_db:getCollection(collection) function mongo_db:getCollection(collection)
local col = { local col = {
connection = self.connection, connection = self.connection,
name = collection, name = collection,
full_name = self.full_name .. "." .. collection, full_name = self.full_name .. "." .. collection,
database = self.database, database = self.database,
} }
self[collection] = setmetatable(col, collection_meta) self[collection] = setmetatable(col, collection_meta)
@@ -188,20 +218,20 @@ mongo_collection.getCollection = mongo_db.getCollection
function mongo_collection:insert(doc) function mongo_collection:insert(doc)
if doc._id == nil then if doc._id == nil then
doc._id = bson.objectid() doc._id = bson.objectid()
end end
local sock = self.connection.__sock local sock = self.connection.__sock
local pack = driver.insert(0, self.full_name, bson_encode(doc)) local pack = driver.insert(0, self.full_name, bson_encode(doc))
-- flags support 1: ContinueOnError -- flags support 1: ContinueOnError
sock:request(pack) sock:request(pack)
end end
function mongo_collection:batch_insert(docs) function mongo_collection:batch_insert(docs)
for i=1,#docs do for i=1,#docs do
if docs[i]._id == nil then if docs[i]._id == nil then
docs[i]._id = bson.objectid() docs[i]._id = bson.objectid()
end end
docs[i] = bson_encode(docs[i]) docs[i] = bson_encode(docs[i])
end end
local sock = self.connection.__sock local sock = self.connection.__sock
local pack = driver.insert(0, self.full_name, docs) local pack = driver.insert(0, self.full_name, docs)
@@ -209,7 +239,7 @@ function mongo_collection:batch_insert(docs)
end end
function mongo_collection:update(selector,update,upsert,multi) function mongo_collection:update(selector,update,upsert,multi)
local flags = (upsert and 1 or 0) + (multi and 2 or 0) local flags = (upsert and 1 or 0) + (multi and 2 or 0)
local sock = self.connection.__sock local sock = self.connection.__sock
local pack = driver.update(self.full_name, flags, bson_encode(selector), bson_encode(update)) local pack = driver.update(self.full_name, flags, bson_encode(selector), bson_encode(update))
sock:request(pack) sock:request(pack)
@@ -225,24 +255,24 @@ function mongo_collection:findOne(query, selector)
local conn = self.connection local conn = self.connection
local request_id = conn:genId() local request_id = conn:genId()
local sock = conn.__sock local sock = conn.__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)) 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))
-- we must hold req (req.data), because req.document is a lightuserdata, it's a pointer to the string (req.data) -- we must hold req (req.data), because req.document is a lightuserdata, it's a pointer to the string (req.data)
local req = sock:request(pack, request_id) local req = sock:request(pack, request_id)
local doc = req.document local doc = req.document
return bson_decode(doc) return bson_decode(doc)
end end
function mongo_collection:find(query, selector) function mongo_collection:find(query, selector)
return setmetatable( { return setmetatable( {
__collection = self, __collection = self,
__query = query and bson_encode(query) or empty_bson, __query = query and bson_encode(query) or empty_bson,
__selector = selector and bson_encode(selector), __selector = selector and bson_encode(selector),
__ptr = nil, __ptr = nil,
__data = nil, __data = nil,
__cursor = nil, __cursor = nil,
__document = {}, __document = {},
__flags = 0, __flags = 0,
} , cursor_meta) } , cursor_meta)
end end
function mongo_cursor:hasNext() function mongo_cursor:hasNext()
@@ -255,42 +285,42 @@ function mongo_cursor:hasNext()
local sock = conn.__sock local sock = conn.__sock
local pack local pack
if self.__data == nil then if self.__data == nil then
pack = driver.query(request_id, self.__flags, self.__collection.full_name,0,0,self.__query,self.__selector) pack = driver.query(request_id, self.__flags, self.__collection.full_name,0,0,self.__query,self.__selector)
else else
if self.__cursor then if self.__cursor then
pack = driver.more(request_id, self.__collection.full_name,0,self.__cursor) pack = driver.more(request_id, self.__collection.full_name,0,self.__cursor)
else else
-- no more -- no more
self.__document = nil self.__document = nil
self.__data = nil self.__data = nil
return false return false
end end
end end
local ok, result = pcall(sock.request,sock,pack, request_id) local ok, result = pcall(sock.request,sock,pack, request_id)
local doc = result.document local doc = result.document
local cursor = result.cursor_id local cursor = result.cursor_id
if ok then if ok then
if doc then if doc then
self.__document = result.result self.__document = result.result
self.__data = result.data self.__data = result.data
self.__ptr = 1 self.__ptr = 1
self.__cursor = cursor self.__cursor = cursor
return true return true
else else
self.__document = nil self.__document = nil
self.__data = nil self.__data = nil
self.__cursor = nil self.__cursor = nil
return false return false
end end
else else
self.__document = nil self.__document = nil
self.__data = nil self.__data = nil
self.__cursor = nil self.__cursor = nil
if doc then if doc then
local err = bson_decode(doc) local err = bson_decode(doc)
error(err["$err"]) error(err["$err"])
else else
error("Reply from mongod error") error("Reply from mongod error")
@@ -303,11 +333,11 @@ end
function mongo_cursor:next() function mongo_cursor:next()
if self.__ptr == nil then if self.__ptr == nil then
error "Call hasNext first" error "Call hasNext first"
end end
local r = bson_decode(self.__document[self.__ptr]) local r = bson_decode(self.__document[self.__ptr])
self.__ptr = self.__ptr + 1 self.__ptr = self.__ptr + 1
if self.__ptr > #self.__document then if self.__ptr > #self.__document then
self.__ptr = nil self.__ptr = nil
end end

View File

@@ -1,3 +1,5 @@
-- This is a deprecated module, use skynet.queue instead.
local skynet = require "skynet" local skynet = require "skynet"
local c = require "skynet.c" local c = require "skynet.c"

View File

@@ -22,7 +22,7 @@ local skynet = {
PTYPE_HARBOR = 5, PTYPE_HARBOR = 5,
PTYPE_SOCKET = 6, PTYPE_SOCKET = 6,
PTYPE_ERROR = 7, PTYPE_ERROR = 7,
PTYPE_QUEUE = 8, PTYPE_QUEUE = 8, -- use in deprecated mqueue, use skynet.queue instead
PTYPE_DEBUG = 9, PTYPE_DEBUG = 9,
PTYPE_LUA = 10, PTYPE_LUA = 10,
PTYPE_SNAX = 11, PTYPE_SNAX = 11,
@@ -127,7 +127,7 @@ function suspend(co, result, command, param, size)
if not result then if not result then
local session = session_coroutine_id[co] local session = session_coroutine_id[co]
local addr = session_coroutine_address[co] local addr = session_coroutine_address[co]
if session and session ~= 0 then if session then
c.send(addr, skynet.PTYPE_ERROR, session, "") c.send(addr, skynet.PTYPE_ERROR, session, "")
end end
session_coroutine_id[co] = nil session_coroutine_id[co] = nil
@@ -151,6 +151,9 @@ function suspend(co, result, command, param, size)
-- coroutine exit -- coroutine exit
session_coroutine_id[co] = nil session_coroutine_id[co] = nil
session_coroutine_address[co] = nil session_coroutine_address[co] = nil
elseif command == "QUIT" then
-- service exit
return
else else
error("Unknown command : " .. command .. "\n" .. debug.traceback(co)) error("Unknown command : " .. command .. "\n" .. debug.traceback(co))
end end
@@ -229,7 +232,10 @@ function skynet.self()
end end
function skynet.localname(name) function skynet.localname(name)
return string_to_handle(c.command("QUERY", name)) local addr = c.command("QUERY", name)
if addr then
return string_to_handle(addr)
end
end end
function skynet.launch(...) function skynet.launch(...)
@@ -253,7 +259,16 @@ end
function skynet.exit() function skynet.exit()
skynet.send(".launcher","lua","REMOVE",skynet.self()) skynet.send(".launcher","lua","REMOVE",skynet.self())
for co, session in pairs(session_coroutine_id) do
local address = session_coroutine_address[co]
local self = skynet.self()
if session~=0 and address then
skynet.redirect(address, self, "error", session, "")
end
end
c.command("EXIT") c.command("EXIT")
-- quit service
coroutine_yield "QUIT"
end end
function skynet.kill(name) function skynet.kill(name)
@@ -318,6 +333,9 @@ end
function skynet.rawcall(addr, typename, msg, sz) function skynet.rawcall(addr, typename, msg, sz)
local p = proto[typename] local p = proto[typename]
if watching_service[addr] == false then
error("Service is dead")
end
local session = assert(c.send(addr, p.id , nil , msg, sz), "call to invalid address") local session = assert(c.send(addr, p.id , nil , msg, sz), "call to invalid address")
return yield_call(addr, session) return yield_call(addr, session)
end end
@@ -569,12 +587,24 @@ function skynet.monitor(service, query)
end end
assert(monitor, "Monitor launch failed") assert(monitor, "Monitor launch failed")
c.command("MONITOR", string.format(":%08x", monitor)) c.command("MONITOR", string.format(":%08x", monitor))
return monitor
end end
function skynet.mqlen() function skynet.mqlen()
return tonumber(c.command "MQLEN") return tonumber(c.command "MQLEN")
end end
function skynet.task(ret)
local t = 0
for session,co in pairs(session_id_coroutine) do
if ret then
ret[session] = debug.traceback(co)
end
t = t + 1
end
return t
end
-- Inject internal debug framework -- Inject internal debug framework
local debug = require "skynet.debug" local debug = require "skynet.debug"
debug(skynet) debug(skynet)

View File

@@ -21,9 +21,16 @@ end
function dbgcmd.STAT() function dbgcmd.STAT()
local stat = {} local stat = {}
stat.mqlen = skynet.mqlen() stat.mqlen = skynet.mqlen()
stat.task = skynet.task()
skynet.ret(skynet.pack(stat)) skynet.ret(skynet.pack(stat))
end end
function dbgcmd.TASK()
local task = {}
skynet.task(task)
skynet.ret(skynet.pack(task))
end
function dbgcmd.INFO() function dbgcmd.INFO()
if internal_info_func then if internal_info_func then
skynet.ret(skynet.pack(internal_info_func())) skynet.ret(skynet.pack(internal_info_func()))

33
lualib/skynet/queue.lua Normal file
View File

@@ -0,0 +1,33 @@
local skynet = require "skynet"
local coroutine = coroutine
local pcall = pcall
local table = table
function skynet.queue()
local current_thread
local ref = 0
local thread_queue = {}
return function(f, ...)
local thread = coroutine.running()
if ref == 0 then
current_thread = thread
elseif current_thread ~= thread then
table.insert(thread_queue, thread)
skynet.wait()
assert(ref == 0)
end
ref = ref + 1
local ok, err = pcall(f, ...)
ref = ref - 1
if ref == 0 then
current_thread = nil
local co = table.remove(thread_queue,1)
if co then
skynet.wakeup(co)
end
end
assert(ok,err)
end
end
return skynet.queue

130
lualib/snax/gateserver.lua Normal file
View File

@@ -0,0 +1,130 @@
local skynet = require "skynet"
local netpack = require "netpack"
local socketdriver = require "socketdriver"
local gateserver = {}
local socket -- listen socket
local queue -- message queue
local maxclient -- max client
local client_number = 0
local CMD = setmetatable({}, { __gc = function() netpack.clear(queue) end })
local nodelay = false
local connection = {}
function gateserver.openclient(fd)
if connection[fd] then
socketdriver.start(fd)
end
end
function gateserver.closeclient(fd)
local c = connection[fd]
if c then
connection[fd] = false
socketdriver.close(fd)
end
end
function gateserver.start(handler)
assert(handler.message)
assert(handler.connect)
function CMD.open( source, conf )
assert(not socket)
local address = conf.address or "0.0.0.0"
local port = assert(conf.port)
maxclient = conf.maxclient or 1024
nodelay = conf.nodelay
socket = socketdriver.listen(address, port)
socketdriver.start(socket)
if handler.open then
return handler.open(source, conf)
end
end
function CMD.close()
assert(socket)
socketdriver.close(socket)
socket = nil
end
local MSG = {}
function MSG.data(fd, msg, sz)
if connection[fd] then
handler.message(fd, msg, sz)
end
end
function MSG.more()
for fd, msg, sz in netpack.pop, queue do
if connection[fd] then
handler.message(fd, msg, sz)
end
end
end
function MSG.open(fd, msg)
if client_number >= maxclient then
socketdriver.close(fd)
return
end
if nodelay then
socketdriver.nodelay(fd)
end
connection[fd] = true
client_number = client_number + 1
handler.connect(fd, msg)
end
local function close_fd(fd)
local c = connection[fd]
if c ~= nil then
connection[fd] = nil
client_number = client_number - 1
end
end
function MSG.close(fd)
if handler.disconnect then
handler.disconnect(fd)
end
close_fd(fd)
end
function MSG.error(fd, msg)
if handler.error then
handler.error(fd)
end
close_fd(fd)
end
skynet.register_protocol {
name = "socket",
id = skynet.PTYPE_SOCKET, -- PTYPE_SOCKET = 6
unpack = function ( msg, sz )
return netpack.filter( queue, msg, sz)
end,
dispatch = function (_, _, q, type, ...)
queue = q
if type then
MSG[type](...)
end
end
}
skynet.start(function()
skynet.dispatch("lua", function (_, address, cmd, ...)
local f = CMD[cmd]
if f then
skynet.ret(skynet.pack(f(address, ...)))
else
skynet.ret(skynet.pack(handler.command(cmd, address, ...)))
end
end)
end)
end
return gateserver

197
lualib/snax/loginserver.lua Normal file
View File

@@ -0,0 +1,197 @@
local skynet = require "skynet"
local socket = require "socket"
local crypt = require "crypt"
local table = table
local string = string
local assert = assert
--[[
Protocol:
line (\n) based text protocol
1. Server->Client : base64(8bytes random challenge)
2. Client->Server : base64(8bytes handshake client key)
3. Server: Gen a 8bytes handshake server key
4. Server->Client : base64(DH-Exchange(server key))
5. Server/Client secret := DH-Secret(client key/server key)
6. Client->Server : base64(HMAC(challenge, secret))
7. Client->Server : DES(secret, base64(token))
8. Server : call auth_handler(token) -> server, uid (A user defined method)
9. Server : call login_handler(server, uid, secret) ->subid (A user defined method)
10. Server->Client : 200 base64(subid)
Error Code:
400 Bad Request . challenge failed
401 Unauthorized . unauthorized by auth_handler
403 Forbidden . login_handler failed
406 Not Acceptable . already in login (disallow multi login)
Success:
200 base64(subid)
]]
local socket_error = {}
local function assert_socket(v, fd)
if v then
return v
else
skynet.error(string.format("auth failed: socket (fd = %d) closed", fd))
error(socket_error)
end
end
local function write(fd, text)
assert_socket(socket.write(fd, text), fd)
end
local function launch_slave(auth_handler)
local function auth(fd, addr)
fd = assert(tonumber(fd))
skynet.error(string.format("connect from %s (fd = %d)", addr, fd))
socket.start(fd)
-- set socket buffer limit (8K)
-- If the attacker send large package, close the socket
socket.limit(fd, 8192)
local challenge = crypt.randomkey()
write(fd, crypt.base64encode(challenge).."\n")
local handshake = assert_socket(socket.readline(fd), fd)
local clientkey = crypt.base64decode(handshake)
if #clientkey ~= 8 then
error "Invalid client key"
end
local serverkey = crypt.randomkey()
write(fd, crypt.base64encode(crypt.dhexchange(serverkey)).."\n")
local secret = crypt.dhsecret(clientkey, serverkey)
local response = assert_socket(socket.readline(fd), fd)
local hmac = crypt.hmac64(challenge, secret)
if hmac ~= crypt.base64decode(response) then
write(fd, "400 Bad Request\n")
error "challenge failed"
end
local etoken = assert_socket(socket.readline(fd),fd)
local token = crypt.desdecode(secret, crypt.base64decode(etoken))
local ok, server, uid = pcall(auth_handler,token)
socket.abandon(fd)
return ok, server, uid, secret
end
local function ret_pack(ok, err, ...)
if ok then
skynet.ret(skynet.pack(err, ...))
elseif err ~= socket_error then
error(err)
end
end
skynet.dispatch("lua", function(_,_,...)
ret_pack(pcall(auth, ...))
end)
end
local user_login = {}
local function accept(conf, s, fd, addr)
-- call slave auth
local ok, server, uid, secret = skynet.call(s, "lua", fd, addr)
socket.start(fd)
if not ok then
write(fd, "401 Unauthorized\n")
error(server)
end
if not conf.multilogin then
if user_login[uid] then
write(fd, "406 Not Acceptable\n")
error(string.format("User %s is already login", uid))
end
user_login[uid] = true
end
local ok, err = pcall(conf.login_handler, server, uid, secret)
-- unlock login
user_login[uid] = nil
if ok then
err = err or ""
write(fd, "200 "..crypt.base64encode(err).."\n")
else
write(fd, "403 Forbidden\n")
error(err)
end
end
local function launch_master(conf)
local instance = conf.instance or 8
assert(instance > 0)
local host = conf.host or "0.0.0.0"
local port = assert(tonumber(conf.port))
local slave = {}
local balance = 1
skynet.dispatch("lua", function(_,source,command, ...)
if command == "register_slave" then
table.insert(slave, source)
skynet.ret(skynet.pack(nil))
else
skynet.ret(skynet.pack(conf.command_handler(command, ...)))
end
end)
for i=1,instance do
skynet.newservice(SERVICE_NAME)
end
skynet.error(string.format("login server listen at : %s %d", host, port))
local id = socket.listen(host, port)
socket.start(id , function(fd, addr)
local s = slave[balance]
balance = balance + 1
if balance > #slave then
balance = 1
end
local ok, err = pcall(accept, conf, s, fd, addr)
if not ok then
if err ~= socket_error then
skynet.error(string.format("invalid client (fd = %d) error = %s", fd, err))
end
end
socket.close(fd)
end)
end
local function login(conf)
local name = "." .. (conf.name or "login")
skynet.start(function()
local loginmaster = skynet.localname(name)
if loginmaster then
skynet.call(loginmaster, "lua", "register_slave")
local auth_handler = assert(conf.auth_handler)
launch_master = nil
conf = nil
launch_slave(auth_handler)
else
launch_slave = nil
conf.auth_handler = nil
assert(conf.login_handler)
assert(conf.command_handler)
skynet.register(name)
launch_master(conf)
end
end)
end
return login

318
lualib/snax/msgserver.lua Normal file
View File

@@ -0,0 +1,318 @@
local skynet = require "skynet"
local gateserver = require "snax.gateserver"
local netpack = require "netpack"
local crypt = require "crypt"
local socketdriver = require "socketdriver"
local assert = assert
local b64encode = crypt.base64encode
local b64decode = crypt.base64decode
--[[
Protocol:
All the number type is big-endian
Shakehands (The first package)
Client -> Server :
base64(uid)@base64(server)#base64(subid):index:base64(hmac)
Server -> Client
XXX ErrorCode
404 User Not Found
403 Index Expired
401 Unauthorized
400 Bad Request
200 OK
Req-Resp
Client -> Server : Request
word size (Not include self)
string content (size-4)
dword session
Server -> Client : Response
word size (Not include self)
string content (size-5)
byte ok (1 is ok, 0 is error)
dword session
API:
server.userid(username)
return uid, subid, server
server.username(uid, subid, server)
return username
server.login(username, secret)
update user secret
server.logout(username)
user logout
server.ip(username)
return ip when connection establish, or nil
server.start(conf)
start server
Supported skynet command:
kick username (may used by loginserver)
login username secret (used by loginserver)
logout username (used by agent)
Config for server.start:
conf.expired_number : the number of the response message cached after sending out (default is 128)
conf.login_handler(uid, secret) -> subid : the function when a new user login, alloc a subid for it. (may call by login server)
conf.logout_handler(uid, subid) : the functon when a user logout. (may call by agent)
conf.kick_handler(uid, subid) : the functon when a user logout. (may call by login server)
conf.request_handler(username, session, msg, sz) : the function when recv a new request.
conf.register_handler(servername) : call when gate open
conf.disconnect_handler(username) : call when a connection disconnect (afk)
]]
local server = {}
skynet.register_protocol {
name = "client",
id = skynet.PTYPE_CLIENT,
}
local user_online = {}
local handshake = {}
local connection = {}
function server.userid(username)
-- base64(uid)@base64(server)#base64(subid)
local uid, servername, subid = username:match "([^@]*)@([^#]*)#(.*)"
return b64decode(uid), b64decode(subid), b64decode(servername)
end
function server.username(uid, subid, servername)
return string.format("%s@%s#%s", b64encode(uid), b64encode(servername), b64encode(tostring(subid)))
end
function server.logout(username)
local u = user_online[username]
user_online[username] = nil
if u.fd then
gateserver.closeclient(u.fd)
connection[u.fd] = nil
end
end
function server.login(username, secret)
assert(user_online[username] == nil)
user_online[username] = {
secret = secret,
version = 0,
index = 0,
username = username,
response = {}, -- response cache
}
end
function server.ip(username)
local u = user_online[username]
if u and u.fd then
return u.ip
end
end
function server.start(conf)
local expired_number = conf.expired_number or 128
local handler = {}
local CMD = {
login = assert(conf.login_handler),
logout = assert(conf.logout_handler),
kick = assert(conf.kick_handler),
}
function handler.command(cmd, source, ...)
local f = assert(CMD[cmd])
return f(...)
end
function handler.open(source, gateconf)
local servername = assert(gateconf.servername)
return conf.register_handler(servername)
end
function handler.connect(fd, addr)
handshake[fd] = addr
gateserver.openclient(fd)
end
function handler.disconnect(fd)
handshake[fd] = nil
local c = connection[fd]
if c then
c.fd = nil
connection[fd] = nil
if conf.disconnect_handler then
conf.disconnect_handler(c.username)
end
end
end
handler.error = handler.disconnect
-- atomic , no yield
local function do_auth(fd, message, addr)
local username, index, hmac = string.match(message, "([^:]*):([^:]*):([^:]*)")
local u = user_online[username]
if u == nil then
return "404 User Not Found"
end
local idx = assert(tonumber(index))
hmac = b64decode(hmac)
if idx <= u.version then
return "403 Index Expired"
end
local text = string.format("%s:%s", username, index)
local v = crypt.hmac64(crypt.hashkey(text), u.secret)
if v ~= hmac then
return "401 Unauthorized"
end
u.version = idx
u.fd = fd
u.ip = addr
connection[fd] = u
end
local function auth(fd, addr, msg, sz)
local message = netpack.tostring(msg, sz)
local ok, result = pcall(do_auth, fd, message, addr)
if not ok then
skynet.error(result)
result = "400 Bad Request"
end
local close = result ~= nil
if result == nil then
result = "200 OK"
end
socketdriver.send(fd, netpack.pack(result))
if close then
gateserver.closeclient(fd)
end
end
local request_handler = assert(conf.request_handler)
-- u.response is a struct { return_fd , response, version, index }
local function retire_response(u)
if u.index >= expired_number * 2 then
local max = 0
local response = u.response
for k,p in pairs(response) do
if p[1] == nil then
-- request complete, check expired
if p[4] < expired_number then
response[k] = nil
else
p[4] = p[4] - expired_number
if p[4] > max then
max = p[4]
end
end
end
end
u.index = max + 1
end
end
local function do_request(fd, msg, sz)
local u = assert(connection[fd], "invalid fd")
local msg_sz = sz - 4
local session = netpack.tostring(msg, sz, msg_sz)
local p = u.response[session]
if p then
-- session can be reuse in the same connection
if p[3] == u.version then
local last = u.response[session]
u.response[session] = nil
p = nil
if last[2] == nil then
local error_msg = string.format("Conflict session %s", crypt.hexencode(session))
skynet.error(error_msg)
error(error_msg)
end
end
end
if p == nil then
p = { fd }
u.response[session] = p
local ok, result = pcall(conf.request_handler, u.username, msg, msg_sz)
result = result or ""
-- NOTICE: YIELD here, socket may close.
if not ok then
skynet.error(result)
result = "\0" .. session
else
result = result .. '\1' .. session
end
p[2] = netpack.pack_string(result)
p[3] = u.version
p[4] = u.index
else
netpack.tostring(msg, sz) -- request before, so free msg
-- update version/index, change return fd.
-- resend response.
p[1] = fd
p[3] = u.version
p[4] = u.index
if p[2] == nil then
-- already request, but response is not ready
return
end
end
u.index = u.index + 1
-- the return fd is p[1] (fd may change by multi request) check connect
fd = p[1]
if connection[fd] then
socketdriver.send(fd, p[2])
end
p[1] = nil
retire_response(u)
end
local function request(fd, msg, sz)
local ok, err = pcall(do_request, fd, msg, sz)
-- not atomic, may yield
if not ok then
skynet.error(string.format("Invalid package %s : %s", err, netpack.tostring(msg, sz)))
if connection[fd] then
gateserver.closeclient(fd)
end
end
end
function handler.message(fd, msg, sz)
local addr = handshake[fd]
if addr then
auth(fd,addr,msg,sz)
handshake[fd] = nil
else
request(fd, msg, sz)
end
end
return gateserver.start(handler)
end
return server

View File

@@ -56,11 +56,19 @@ socket_message[1] = function(id, size, data)
s.read_required = nil s.read_required = nil
wakeup(s) wakeup(s)
end end
elseif rrt == "string" then else
-- read line if s.buffer_limit and sz > s.buffer_limit then
if driver.readline(s.buffer,nil,rr) then skynet.error(string.format("socket buffer overlow: fd=%d size=%d", id , sz))
s.read_required = nil driver.clear(s.buffer,buffer_pool)
wakeup(s) driver.close(id)
return
end
if rrt == "string" then
-- read line
if driver.readline(s.buffer,nil,rr) then
s.read_required = nil
wakeup(s)
end
end end
end end
end end
@@ -199,6 +207,27 @@ end
function socket.read(id, sz) function socket.read(id, sz)
local s = socket_pool[id] local s = socket_pool[id]
assert(s) assert(s)
if sz == nil then
-- read some bytes
local ret = driver.readall(s.buffer, buffer_pool)
if ret ~= "" then
return ret
end
if not s.connected then
return false, ret
end
assert(not s.read_required)
s.read_required = 0
suspend(s)
ret = driver.readall(s.buffer, buffer_pool)
if ret ~= "" then
return ret
else
return false, ret
end
end
local ret = driver.pop(s.buffer, buffer_pool, sz) local ret = driver.pop(s.buffer, buffer_pool, sz)
if ret then if ret then
return ret return ret
@@ -318,4 +347,9 @@ function socket.abandon(id)
socket_pool[id] = nil socket_pool[id] = nil
end end
function socket.limit(id, limit)
local s = assert(socket_pool[id])
s.buffer_limit = limit
end
return socket return socket

View File

@@ -26,6 +26,7 @@ function socket_channel.channel(desc)
local c = { local c = {
__host = assert(desc.host), __host = assert(desc.host),
__port = assert(desc.port), __port = assert(desc.port),
__backup = desc.backup,
__auth = desc.auth, __auth = desc.auth,
__response = desc.response, -- It's for session mode __response = desc.response, -- It's for session mode
__request = {}, -- request seq { response func or session } -- It's for order mode __request = {}, -- request seq { response func or session } -- It's for order mode
@@ -135,6 +136,9 @@ local function dispatch_by_order(self)
if result ~= socket_error then if result ~= socket_error then
errmsg = result_ok errmsg = result_ok
end end
self.__result[co] = socket_error
self.__result_data[co] = errmsg
skynet.wakeup(co)
wakeup_all(self, errmsg) wakeup_all(self, errmsg)
end end
end end
@@ -149,34 +153,66 @@ local function dispatch_function(self)
end end
end end
local function connect_backup(self)
if self.__backup then
for _, addr in ipairs(self.__backup) do
local host, port
if type(addr) == "table" then
host, port = addr.host, addr.port
else
host = addr
port = self.__port
end
skynet.error("socket: connect to backup host", host, port)
local fd = socket.open(host, port)
if fd then
self.__host = host
self.__port = port
return fd
end
end
end
end
local function connect_once(self) local function connect_once(self)
if self.__closed then
return false
end
assert(not self.__sock and not self.__authcoroutine) assert(not self.__sock and not self.__authcoroutine)
local fd = socket.open(self.__host, self.__port) local fd = socket.open(self.__host, self.__port)
if not fd then if not fd then
return false fd = connect_backup(self)
if not fd then
return false
end
end end
self.__authcoroutine = coroutine.running()
self.__sock = setmetatable( {fd} , channel_socket_meta ) self.__sock = setmetatable( {fd} , channel_socket_meta )
skynet.fork(dispatch_function(self), self) skynet.fork(dispatch_function(self), self)
if self.__auth then if self.__auth then
self.__authcoroutine = coroutine.running()
local ok , message = pcall(self.__auth, self) local ok , message = pcall(self.__auth, self)
if not ok then if not ok then
close_channel_socket(self) close_channel_socket(self)
if message ~= socket_error then if message ~= socket_error then
self.__authcoroutine = false
skynet.error("socket: auth failed", message) skynet.error("socket: auth failed", message)
end end
end end
self.__authcoroutine = false self.__authcoroutine = false
if ok and not self.__sock then
-- auth may change host, so connect again
return connect_once(self)
end
return ok return ok
end end
self.__authcoroutine = false
return true return true
end end
local function try_connect(self , once) local function try_connect(self , once)
local t = 100 local t = 0
while not self.__closed do while not self.__closed do
if connect_once(self) then if connect_once(self) then
if not once then if not once then
@@ -289,6 +325,20 @@ function channel:close()
end end
end end
function channel:changehost(host, port)
self.__host = host
if port then
self.__port = port
end
if not self.__closed then
close_channel_socket(self)
end
end
function channel:changebackup(backup)
self.__backup = backup
end
channel_meta.__gc = channel.close channel_meta.__gc = channel.close
local function wrapper_socket_function(f) local function wrapper_socket_function(f)

View File

@@ -262,7 +262,8 @@ harbor_release(struct harbor *h) {
struct slave *s = &h->s[i]; struct slave *s = &h->s[i];
if (s->fd && s->status != STATUS_DOWN) { if (s->fd && s->status != STATUS_DOWN) {
close_harbor(h,i); close_harbor(h,i);
report_harbor_down(h,i); // don't call report_harbor_down.
// never call skynet_send during module exit, because of dead lock
} }
} }
hash_delete(h->map); hash_delete(h->map);
@@ -314,7 +315,9 @@ forward_local_messsage(struct harbor *h, void *msg, int sz) {
int type = (destination >> HANDLE_REMOTE_SHIFT) | PTYPE_TAG_DONTCOPY; int type = (destination >> HANDLE_REMOTE_SHIFT) | PTYPE_TAG_DONTCOPY;
destination = (destination & HANDLE_MASK) | ((uint32_t)h->id << HANDLE_REMOTE_SHIFT); destination = (destination & HANDLE_MASK) | ((uint32_t)h->id << HANDLE_REMOTE_SHIFT);
skynet_send(h->ctx, header.source, destination, type, (int)header.session, (void *)msg, sz-HEADER_COOKIE_LENGTH); if (skynet_send(h->ctx, header.source, destination, type, (int)header.session, (void *)msg, sz-HEADER_COOKIE_LENGTH) < 0) {
skynet_error(h->ctx, "Unknown destination :%x from :%x", destination, header.source);
}
} }
static void static void
@@ -509,9 +512,12 @@ remote_send_handle(struct harbor *h, uint32_t source, uint32_t destination, int
skynet_send(context, destination, source, PTYPE_ERROR, 0 , NULL, 0); skynet_send(context, destination, source, PTYPE_ERROR, 0 , NULL, 0);
skynet_error(context, "Drop message to harbor %d from %x to %x (session = %d, msgsz = %d)",harbor_id, source, destination,session,(int)sz); skynet_error(context, "Drop message to harbor %d from %x to %x (session = %d, msgsz = %d)",harbor_id, source, destination,session,(int)sz);
} else { } else {
if (s->queue == NULL) {
s->queue = new_queue();
}
struct remote_message_header header; struct remote_message_header header;
header.source = source; header.source = source;
header.destination = type << HANDLE_REMOTE_SHIFT; header.destination = (type << HANDLE_REMOTE_SHIFT) | (destination & HANDLE_MASK);
header.session = (uint32_t)session; header.session = (uint32_t)session;
push_queue(s->queue, (void *)msg, sz, &header); push_queue(s->queue, (void *)msg, sz, &header);
return 1; return 1;

View File

@@ -5,7 +5,14 @@ local cluster = require "cluster.c"
local config_name = skynet.getenv "cluster" local config_name = skynet.getenv "cluster"
local node_address = {} local node_address = {}
assert(loadfile(config_name, "t", node_address))()
local function loadconfig()
local f = assert(io.open(config_name))
local source = f:read "*a"
f:close()
assert(load(source, "@"..config_name, "t", node_address))()
end
local node_session = {} local node_session = {}
local command = {} local command = {}
@@ -30,6 +37,11 @@ end
local node_channel = setmetatable({}, { __index = open_channel }) local node_channel = setmetatable({}, { __index = open_channel })
function command.reload()
loadconfig()
skynet.ret(skynet.pack(nil))
end
function command.listen(source, addr, port) function command.listen(source, addr, port)
local gate = skynet.newservice("gate") local gate = skynet.newservice("gate")
if port == nil then if port == nil then
@@ -53,8 +65,13 @@ local request_fd = {}
function command.socket(source, subcmd, fd, msg) function command.socket(source, subcmd, fd, msg)
if subcmd == "data" then if subcmd == "data" then
local addr, session, msg = cluster.unpackrequest(msg) local addr, session, msg = cluster.unpackrequest(msg)
local msg, sz = skynet.rawcall(addr, "lua", msg) local ok , msg, sz = pcall(skynet.rawcall, addr, "lua", msg)
local response = cluster.packresponse(session, msg, sz) local response
if ok then
response = cluster.packresponse(session, true, msg, sz)
else
response = cluster.packresponse(session, false, msg)
end
socket.write(fd, response) socket.write(fd, response)
elseif subcmd == "open" then elseif subcmd == "open" then
skynet.error(string.format("socket accept from %s", msg)) skynet.error(string.format("socket accept from %s", msg))
@@ -65,7 +82,8 @@ function command.socket(source, subcmd, fd, msg)
end end
skynet.start(function() skynet.start(function()
skynet.dispatch("lua", function(_, source, cmd, ...) loadconfig()
skynet.dispatch("lua", function(session , source, cmd, ...)
local f = assert(command[cmd]) local f = assert(command[cmd])
f(source, ...) f(source, ...)
end) end)

View File

@@ -2,6 +2,8 @@ local skynet = require "skynet"
local command = {} local command = {}
local database = {} local database = {}
local wait_queue = {}
local mode = {}
local function query(db, key, ...) local function query(db, key, ...)
if key == nil then if key == nil then
@@ -22,7 +24,7 @@ local function update(db, key, value, ...)
if select("#",...) == 0 then if select("#",...) == 0 then
local ret = db[key] local ret = db[key]
db[key] = value db[key] = value
return ret return ret, value
else else
if db[key] == nil then if db[key] == nil then
db[key] = {} db[key] = {}
@@ -31,13 +33,82 @@ local function update(db, key, value, ...)
end end
end end
local function wakeup(db, key1, key2, value, ...)
if key1 == nil then
return
end
local q = db[key1]
if q == nil then
return
end
if q[mode] == "queue" then
db[key1] = nil
if value then
-- throw error because can't wake up a branch
for _,v in ipairs(q) do
local session = v[1]
local source = v[2]
skynet.redirect(source, 0, "error", session, "")
end
else
return q
end
else
-- it's branch
return wakeup(q , key2, value, ...)
end
end
function command.UPDATE(...) function command.UPDATE(...)
return update(database, ...) local ret, value = update(database, ...)
if ret or value == nil then
return ret
end
local q = wakeup(wait_queue, ...)
if q then
for _, v in ipairs(q) do
local session = v[1]
local source = v[2]
skynet.redirect(source, 0, "response", session, skynet.pack(value))
end
end
end
local function waitfor(session, source, db, key1, key2, ...)
if key2 == nil then
-- push queue
local q = db[key1]
if q == nil then
q = { [mode] = "queue" }
db[key1] = q
else
assert(q[mode] == "queue")
end
table.insert(q, { session, source })
else
local q = db[key1]
if q == nil then
q = { [mode] = "branch" }
db[key1] = q
else
assert(q[mode] == "branch")
end
return waitfor(session, source, q, key2, ...)
end
end end
skynet.start(function() skynet.start(function()
skynet.dispatch("lua", function (_, source, cmd, ...) skynet.dispatch("lua", function (session, source, cmd, ...)
local f = assert(command[cmd]) if cmd == "WAIT" then
skynet.ret(skynet.pack(f(...))) local ret = command.QUERY(...)
if ret then
skynet.ret(skynet.pack(ret))
else
waitfor(session, source, wait_queue, ...)
end
else
local f = assert(command[cmd])
skynet.ret(skynet.pack(f(...)))
end
end) end)
end) end)

View File

@@ -113,6 +113,7 @@ function COMMAND.help()
snax = "lanuch a new snax service", snax = "lanuch a new snax service",
clearcache = "clear lua code cache", clearcache = "clear lua code cache",
service = "List unique service", service = "List unique service",
task = "task address : show service task detail",
} }
end end

View File

@@ -1,84 +1,23 @@
local skynet = require "skynet" local skynet = require "skynet"
local gateserver = require "snax.gateserver"
local netpack = require "netpack" local netpack = require "netpack"
local socketdriver = require "socketdriver"
local socket
local queue
local watchdog local watchdog
local maxclient
local client_number = 0
local CMD = setmetatable({}, { __gc = function() netpack.clear(queue) end })
local connection = {} -- fd -> connection : { fd , client, agent , ip, mode } local connection = {} -- fd -> connection : { fd , client, agent , ip, mode }
local forwarding = {} -- agent -> connection local forwarding = {} -- agent -> connection
function CMD.open( source , conf ) skynet.register_protocol {
assert(not socket) name = "client",
local address = conf.address or "0.0.0.0" id = skynet.PTYPE_CLIENT,
local port = assert(conf.port) }
maxclient = conf.maxclient or 1024
local handler = {}
function handler.open(source, conf)
watchdog = conf.watchdog or source watchdog = conf.watchdog or source
socket = socketdriver.listen(address, port)
socketdriver.start(socket)
end end
function CMD.close() function handler.message(fd, msg, sz)
assert(socket)
socketdriver.close(socket)
socket = nil
end
local function unforward(c)
if c.agent then
forwarding[c.agent] = nil
c.agent = nil
c.client = nil
end
end
local function start(c)
if not c.mode then
c.mode = "open"
socketdriver.start(c.fd)
end
end
function CMD.forward(source, fd, client, address)
local c = assert(connection[fd])
unforward(c)
start(c)
c.client = client or 0
c.agent = address or source
forwarding[c.agent] = c
end
function CMD.accept(source, fd)
local c = assert(connection[fd])
unforward(c)
start(c)
end
function CMD.kick(source, fd)
local c
if fd then
c = connection[fd]
else
c = forwarding[source]
end
assert(c)
if c.mode ~= "close" then
c.mode = "close"
socketdriver.close(c.fd)
end
end
local MSG = {}
function MSG.data(fd, msg, sz)
-- recv a package, forward it -- recv a package, forward it
local c = connection[fd] local c = connection[fd]
local agent = c.agent local agent = c.agent
@@ -89,67 +28,65 @@ function MSG.data(fd, msg, sz)
end end
end end
function MSG.more() function handler.connect(fd, addr)
for fd, msg, sz in netpack.pop, queue do
MSG.data(fd, msg, sz)
end
end
function MSG.open(fd, msg)
if client_number >= maxclient then
socketdriver.close(fd)
return
end
local c = { local c = {
fd = fd, fd = fd,
ip = msg, ip = msg,
} }
connection[fd] = c connection[fd] = c
client_number = client_number + 1 skynet.send(watchdog, "lua", "socket", "open", fd, addr)
skynet.send(watchdog, "lua", "socket", "open", fd, msg)
end end
local function close_fd(fd, message) local function unforward(c)
if c.agent then
forwarding[c.agent] = nil
c.agent = nil
c.client = nil
end
end
local function close_fd(fd)
local c = connection[fd] local c = connection[fd]
if c then if c then
unforward(c) unforward(c)
connection[fd] = nil connection[fd] = nil
client_number = client_number - 1
end end
end end
function MSG.close(fd) function handler.disconnect(fd)
close_fd(fd) close_fd(fd)
skynet.send(watchdog, "lua", "socket", "close", fd) skynet.send(watchdog, "lua", "socket", "close", fd)
end end
function MSG.error(fd, msg) function handler.error(fd, msg)
close_fd(fd) close_fd(fd)
skynet.send(watchdog, "lua", "socket", "error", fd, msg) skynet.send(watchdog, "lua", "socket", "error", fd, msg)
end end
skynet.register_protocol { local CMD = {}
name = "socket",
id = skynet.PTYPE_SOCKET, -- PTYPE_SOCKET = 6
unpack = function ( msg, sz )
return netpack.filter( queue, msg, sz)
end,
dispatch = function (_, _, q, type, ...)
queue = q
if type then
MSG[type](...)
end
end
}
skynet.register_protocol { function CMD.forward(source, fd, client, address)
name = "client", local c = assert(connection[fd])
id = skynet.PTYPE_CLIENT, unforward(c)
} c.client = client or 0
c.agent = address or source
forwarding[c.agent] = c
gateserver.openclient(fd)
end
skynet.start(function() function CMD.accept(source, fd)
skynet.dispatch("lua", function (_, address, cmd, ...) local c = assert(connection[fd])
local f = assert(CMD[cmd]) unforward(c)
skynet.ret(skynet.pack(f(address, ...))) gateserver.openclient(fd)
end) end
end)
function CMD.kick(source, fd)
gateserver.closeclient(fd)
end
function handler.command(cmd, source, ...)
local f = assert(CMD[cmd])
return f(source, ...)
end
gateserver.start(handler)

View File

@@ -38,6 +38,16 @@ function command.INFO(_, _, handle)
end end
end end
function command.TASK(_, _, handle)
handle = handle_to_address(handle)
if services[handle] == nil then
return
else
local result = skynet.call(handle,"debug","TASK")
return result
end
end
function command.KILL(_, _, handle) function command.KILL(_, _, handle)
handle = handle_to_address(handle) handle = handle_to_address(handle)
skynet.kill(handle) skynet.kill(handle)

View File

@@ -50,8 +50,10 @@ function command.DEL(source, c)
channel[c] = nil channel[c] = nil
channel_n[c] = nil channel_n[c] = nil
channel_remote[c] = nil channel_remote[c] = nil
for node in pairs(remote) do if remote then
skynet.send(node_address[node], "lua", "DELR", c) for node in pairs(remote) do
skynet.send(node_address[node], "lua", "DELR", c)
end
end end
return NORET return NORET
end end

View File

@@ -176,7 +176,7 @@ skynet.start(function()
end end
end) end)
local handle = skynet.localname ".service" local handle = skynet.localname ".service"
if handle ~= 0 then if handle then
skynet.error(".service is already register by ", skynet.address(handle)) skynet.error(".service is already register by ", skynet.address(handle))
skynet.exit() skynet.exit()
else else

View File

@@ -11,6 +11,7 @@
#include <lualib.h> #include <lualib.h>
#include <lauxlib.h> #include <lauxlib.h>
#include <signal.h> #include <signal.h>
#include <assert.h>
static int static int
optint(const char *key, int opt) { optint(const char *key, int opt) {
@@ -51,7 +52,6 @@ optstring(const char *key,const char * opt) {
static void static void
_init_env(lua_State *L) { _init_env(lua_State *L) {
lua_pushglobaltable(L);
lua_pushnil(L); /* first key */ lua_pushnil(L); /* first key */
while (lua_next(L, -2) != 0) { while (lua_next(L, -2) != 0) {
int keyt = lua_type(L, -2); int keyt = lua_type(L, -2);
@@ -83,6 +83,18 @@ int sigign() {
return 0; return 0;
} }
static const char * load_config = "\
local config_name = ...\
local f = assert(io.open(config_name))\
local code = assert(f:read \'*a\')\
local function getenv(name) return assert(os.getenv(name), name) end\
code = string.gsub(code, \'%$([%w_%d]+)\', getenv)\
f:close()\
local result = {}\
assert(load(code,\'=(load)\',\'t\',result))()\
return result\
";
int int
main(int argc, char *argv[]) { main(int argc, char *argv[]) {
const char * config_file = "config"; const char * config_file = "config";
@@ -98,11 +110,12 @@ main(int argc, char *argv[]) {
struct lua_State *L = lua_newstate(skynet_lalloc, NULL); struct lua_State *L = lua_newstate(skynet_lalloc, NULL);
luaL_openlibs(L); // link lua lib luaL_openlibs(L); // link lua lib
lua_close(L);
L = luaL_newstate(); int err = luaL_loadstring(L, load_config);
assert(err == LUA_OK);
lua_pushstring(L, config_file);
int err = luaL_dofile(L, config_file); err = lua_pcall(L, 1, 1, 0);
if (err) { if (err) {
fprintf(stderr,"%s\n",lua_tostring(L,-1)); fprintf(stderr,"%s\n",lua_tostring(L,-1));
lua_close(L); lua_close(L);

View File

@@ -245,7 +245,7 @@ skynet_context_dispatchall(struct skynet_context * ctx) {
} }
struct message_queue * struct message_queue *
skynet_context_message_dispatch(struct skynet_monitor *sm, struct message_queue *q) { skynet_context_message_dispatch(struct skynet_monitor *sm, struct message_queue *q, int weight) {
if (q == NULL) { if (q == NULL) {
q = skynet_globalmq_pop(); q = skynet_globalmq_pop();
if (q==NULL) if (q==NULL)
@@ -261,18 +261,27 @@ skynet_context_message_dispatch(struct skynet_monitor *sm, struct message_queue
return skynet_globalmq_pop(); return skynet_globalmq_pop();
} }
int i,n=1;
struct skynet_message msg; struct skynet_message msg;
if (skynet_mq_pop(q,&msg)) {
skynet_context_release(ctx);
return skynet_globalmq_pop();
}
skynet_monitor_trigger(sm, msg.source , handle); for (i=0;i<n;i++) {
if (skynet_mq_pop(q,&msg)) {
skynet_context_release(ctx);
return skynet_globalmq_pop();
} else if (i==0 && weight >= 0) {
n = skynet_mq_length(q);
n >>= weight;
}
if (ctx->cb == NULL) { skynet_monitor_trigger(sm, msg.source , handle);
skynet_free(msg.data);
} else { if (ctx->cb == NULL) {
_dispatch_message(ctx, &msg); skynet_free(msg.data);
} else {
_dispatch_message(ctx, &msg);
}
skynet_monitor_trigger(sm, 0,0);
} }
assert(q == ctx->queue); assert(q == ctx->queue);
@@ -285,8 +294,6 @@ skynet_context_message_dispatch(struct skynet_monitor *sm, struct message_queue
} }
skynet_context_release(ctx); skynet_context_release(ctx);
skynet_monitor_trigger(sm, 0,0);
return q; return q;
} }
@@ -361,8 +368,10 @@ static const char *
cmd_query(struct skynet_context * context, const char * param) { cmd_query(struct skynet_context * context, const char * param) {
if (param[0] == '.') { if (param[0] == '.') {
uint32_t handle = skynet_handle_findname(param+1); uint32_t handle = skynet_handle_findname(param+1);
sprintf(context->result, ":%x", handle); if (handle) {
return context->result; sprintf(context->result, ":%x", handle);
return context->result;
}
} }
return NULL; return NULL;
} }

View File

@@ -16,7 +16,7 @@ void skynet_context_init(struct skynet_context *, uint32_t handle);
int skynet_context_push(uint32_t handle, struct skynet_message *message); int skynet_context_push(uint32_t handle, struct skynet_message *message);
void skynet_context_send(struct skynet_context * context, void * msg, size_t sz, uint32_t source, int type, int session); void skynet_context_send(struct skynet_context * context, void * msg, size_t sz, uint32_t source, int type, int session);
int skynet_context_newsession(struct skynet_context *); int skynet_context_newsession(struct skynet_context *);
struct message_queue * skynet_context_message_dispatch(struct skynet_monitor *, struct message_queue *); // return next queue struct message_queue * skynet_context_message_dispatch(struct skynet_monitor *, struct message_queue *, int weight); // return next queue
int skynet_context_total(); int skynet_context_total();
void skynet_context_dispatchall(struct skynet_context * context); // for skynet_error output before exit void skynet_context_dispatchall(struct skynet_context * context); // for skynet_error output before exit

View File

@@ -150,3 +150,8 @@ skynet_socket_start(struct skynet_context *ctx, int id) {
uint32_t source = skynet_context_handle(ctx); uint32_t source = skynet_context_handle(ctx);
socket_server_start(SOCKET_SERVER, source, id); socket_server_start(SOCKET_SERVER, source, id);
} }
void
skynet_socket_nodelay(struct skynet_context *ctx, int id) {
socket_server_nodelay(SOCKET_SERVER, id);
}

View File

@@ -28,5 +28,6 @@ int skynet_socket_connect(struct skynet_context *ctx, const char *host, int port
int skynet_socket_bind(struct skynet_context *ctx, int fd); int skynet_socket_bind(struct skynet_context *ctx, int fd);
void skynet_socket_close(struct skynet_context *ctx, int id); void skynet_socket_close(struct skynet_context *ctx, int id);
void skynet_socket_start(struct skynet_context *ctx, int id); void skynet_socket_start(struct skynet_context *ctx, int id);
void skynet_socket_nodelay(struct skynet_context *ctx, int id);
#endif #endif

View File

@@ -27,6 +27,7 @@ struct monitor {
struct worker_parm { struct worker_parm {
struct monitor *m; struct monitor *m;
int id; int id;
int weight;
}; };
#define CHECK_ABORT if (skynet_context_total()==0) break; #define CHECK_ABORT if (skynet_context_total()==0) break;
@@ -118,12 +119,13 @@ static void *
_worker(void *p) { _worker(void *p) {
struct worker_parm *wp = p; struct worker_parm *wp = p;
int id = wp->id; int id = wp->id;
int weight = wp->weight;
struct monitor *m = wp->m; struct monitor *m = wp->m;
struct skynet_monitor *sm = m->m[id]; struct skynet_monitor *sm = m->m[id];
skynet_initthread(THREAD_WORKER); skynet_initthread(THREAD_WORKER);
struct message_queue * q = NULL; struct message_queue * q = NULL;
for (;;) { for (;;) {
q = skynet_context_message_dispatch(sm, q); q = skynet_context_message_dispatch(sm, q, weight);
if (q == NULL) { if (q == NULL) {
CHECK_ABORT CHECK_ABORT
if (pthread_mutex_lock(&m->mutex) == 0) { if (pthread_mutex_lock(&m->mutex) == 0) {
@@ -169,10 +171,20 @@ _start(int thread) {
create_thread(&pid[1], _timer, m); create_thread(&pid[1], _timer, m);
create_thread(&pid[2], _socket, m); create_thread(&pid[2], _socket, m);
static int weight[] = {
-1, -1, -1, -1, 0, 0, 0, 0,
1, 1, 1, 1, 1, 1, 1, 1,
2, 2, 2, 2, 2, 2, 2, 2,
3, 3, 3, 3, 3, 3, 3, 3, };
struct worker_parm wp[thread]; struct worker_parm wp[thread];
for (i=0;i<thread;i++) { for (i=0;i<thread;i++) {
wp[i].m = m; wp[i].m = m;
wp[i].id = i; wp[i].id = i;
if (i < sizeof(weight)/sizeof(weight[0])) {
wp[i].weight= weight[i];
} else {
wp[i].weight = 0;
}
create_thread(&pid[i+3], _worker, &wp[i]); create_thread(&pid[i+3], _worker, &wp[i]);
} }

View File

@@ -9,6 +9,7 @@
#include <assert.h> #include <assert.h>
#include <string.h> #include <string.h>
#include <stdlib.h> #include <stdlib.h>
#include <stdint.h>
#if defined(__APPLE__) #if defined(__APPLE__)
#include <sys/time.h> #include <sys/time.h>
@@ -33,7 +34,7 @@ struct timer_event {
struct timer_node { struct timer_node {
struct timer_node *next; struct timer_node *next;
int expire; uint32_t expire;
}; };
struct link_list { struct link_list {
@@ -43,9 +44,9 @@ struct link_list {
struct timer { struct timer {
struct link_list near[TIME_NEAR]; struct link_list near[TIME_NEAR];
struct link_list t[4][TIME_LEVEL-1]; struct link_list t[4][TIME_LEVEL];
int lock; int lock;
int time; uint32_t time;
uint32_t current; uint32_t current;
uint32_t starttime; uint32_t starttime;
uint64_t current_point; uint64_t current_point;
@@ -72,21 +73,22 @@ link(struct link_list *list,struct timer_node *node) {
static void static void
add_node(struct timer *T,struct timer_node *node) { add_node(struct timer *T,struct timer_node *node) {
int time=node->expire; uint32_t time=node->expire;
int current_time=T->time; uint32_t current_time=T->time;
if ((time|TIME_NEAR_MASK)==(current_time|TIME_NEAR_MASK)) { if ((time|TIME_NEAR_MASK)==(current_time|TIME_NEAR_MASK)) {
link(&T->near[time&TIME_NEAR_MASK],node); link(&T->near[time&TIME_NEAR_MASK],node);
} else { } else {
int i; int i;
int mask=TIME_NEAR << TIME_LEVEL_SHIFT; uint32_t mask=TIME_NEAR << TIME_LEVEL_SHIFT;
for (i=0;i<3;i++) { for (i=0;i<3;i++) {
if ((time|(mask-1))==(current_time|(mask-1))) { if ((time|(mask-1))==(current_time|(mask-1))) {
break; break;
} }
mask <<= TIME_LEVEL_SHIFT; mask <<= TIME_LEVEL_SHIFT;
} }
link(&T->t[i][((time>>(TIME_NEAR_SHIFT + i*TIME_LEVEL_SHIFT)) & TIME_LEVEL_MASK)-1],node);
link(&T->t[i][((time>>(TIME_NEAR_SHIFT + i*TIME_LEVEL_SHIFT)) & TIME_LEVEL_MASK)],node);
} }
} }
@@ -103,28 +105,37 @@ timer_add(struct timer *T,void *arg,size_t sz,int time) {
UNLOCK(T); UNLOCK(T);
} }
static void
move_list(struct timer *T, int level, int idx) {
struct timer_node *current = link_clear(&T->t[level][idx]);
while (current) {
struct timer_node *temp=current->next;
add_node(T,current);
current=temp;
}
}
static void static void
timer_shift(struct timer *T) { timer_shift(struct timer *T) {
LOCK(T); LOCK(T);
int mask = TIME_NEAR; int mask = TIME_NEAR;
int time = (++T->time) >> TIME_NEAR_SHIFT; uint32_t ct = ++T->time;
int i=0; if (ct == 0) {
move_list(T, 3, 0);
} else {
uint32_t time = ct >> TIME_NEAR_SHIFT;
int i=0;
while ((T->time & (mask-1))==0) { while ((ct & (mask-1))==0) {
int idx=time & TIME_LEVEL_MASK; int idx=time & TIME_LEVEL_MASK;
if (idx!=0) { if (idx!=0) {
--idx; move_list(T, i, idx);
struct timer_node *current = link_clear(&T->t[i][idx]); break;
while (current) {
struct timer_node *temp=current->next;
add_node(T,current);
current=temp;
} }
break; mask <<= TIME_LEVEL_SHIFT;
time >>= TIME_LEVEL_SHIFT;
++i;
} }
mask <<= TIME_LEVEL_SHIFT;
time >>= TIME_LEVEL_SHIFT;
++i;
} }
UNLOCK(T); UNLOCK(T);
} }
@@ -187,7 +198,7 @@ timer_create_timer() {
} }
for (i=0;i<4;i++) { for (i=0;i<4;i++) {
for (j=0;j<TIME_LEVEL-1;j++) { for (j=0;j<TIME_LEVEL;j++) {
link_clear(&r->t[i][j]); link_clear(&r->t[i][j]);
} }
} }
@@ -265,6 +276,7 @@ skynet_updatetime(void) {
uint64_t cp = gettime(); uint64_t cp = gettime();
if(cp < TI->current_point) { if(cp < TI->current_point) {
skynet_error(NULL, "time diff error: change from %lld to %lld", cp, TI->current_point); skynet_error(NULL, "time diff error: change from %lld to %lld", cp, TI->current_point);
TI->current_point = cp;
} else if (cp != TI->current_point) { } else if (cp != TI->current_point) {
uint32_t diff = (uint32_t)(cp - TI->current_point); uint32_t diff = (uint32_t)(cp - TI->current_point);
TI->current_point = cp; TI->current_point = cp;

View File

@@ -5,6 +5,7 @@
#include <sys/types.h> #include <sys/types.h>
#include <sys/socket.h> #include <sys/socket.h>
#include <netinet/tcp.h>
#include <unistd.h> #include <unistd.h>
#include <errno.h> #include <errno.h>
#include <stdlib.h> #include <stdlib.h>
@@ -34,6 +35,8 @@
#define PRIORITY_HIGH 0 #define PRIORITY_HIGH 0
#define PRIORITY_LOW 1 #define PRIORITY_LOW 1
#define HASH_ID(id) (((unsigned)id) % MAX_SOCKET)
struct write_buffer { struct write_buffer {
struct write_buffer * next; struct write_buffer * next;
char *ptr; char *ptr;
@@ -107,6 +110,12 @@ struct request_start {
uintptr_t opaque; uintptr_t opaque;
}; };
struct request_setopt {
int id;
int what;
int value;
};
struct request_package { struct request_package {
uint8_t header[8]; // 6 bytes dummy uint8_t header[8]; // 6 bytes dummy
union { union {
@@ -117,6 +126,7 @@ struct request_package {
struct request_listen listen; struct request_listen listen;
struct request_bind bind; struct request_bind bind;
struct request_start start; struct request_start start;
struct request_setopt setopt;
} u; } u;
uint8_t dummy[256]; uint8_t dummy[256];
}; };
@@ -144,7 +154,7 @@ reserve_id(struct socket_server *ss) {
if (id < 0) { if (id < 0) {
id = __sync_and_and_fetch(&(ss->alloc_id), 0x7fffffff); id = __sync_and_and_fetch(&(ss->alloc_id), 0x7fffffff);
} }
struct socket *s = &ss->slot[id % MAX_SOCKET]; struct socket *s = &ss->slot[HASH_ID(id)];
if (s->type == SOCKET_TYPE_INVALID) { if (s->type == SOCKET_TYPE_INVALID) {
if (__sync_bool_compare_and_swap(&s->type, SOCKET_TYPE_INVALID, SOCKET_TYPE_RESERVE)) { if (__sync_bool_compare_and_swap(&s->type, SOCKET_TYPE_INVALID, SOCKET_TYPE_RESERVE)) {
s->id = id; s->id = id;
@@ -267,7 +277,7 @@ check_wb_list(struct wb_list *s) {
static struct socket * static struct socket *
new_fd(struct socket_server *ss, int id, int fd, uintptr_t opaque, bool add) { new_fd(struct socket_server *ss, int id, int fd, uintptr_t opaque, bool add) {
struct socket * s = &ss->slot[id % MAX_SOCKET]; struct socket * s = &ss->slot[HASH_ID(id)];
assert(s->type == SOCKET_TYPE_RESERVE); assert(s->type == SOCKET_TYPE_RESERVE);
if (add) { if (add) {
@@ -357,7 +367,7 @@ open_socket(struct socket_server *ss, struct request_open * request, struct sock
return -1; return -1;
_failed: _failed:
freeaddrinfo( ai_list ); freeaddrinfo( ai_list );
ss->slot[id % MAX_SOCKET].type = SOCKET_TYPE_INVALID; ss->slot[HASH_ID(id)].type = SOCKET_TYPE_INVALID;
return SOCKET_ERROR; return SOCKET_ERROR;
} }
@@ -502,7 +512,7 @@ send_buffer_empty(struct socket *s) {
static int static int
send_socket(struct socket_server *ss, struct request_send * request, struct socket_message *result, int priority) { send_socket(struct socket_server *ss, struct request_send * request, struct socket_message *result, int priority) {
int id = request->id; int id = request->id;
struct socket * s = &ss->slot[id % MAX_SOCKET]; struct socket * s = &ss->slot[HASH_ID(id)];
if (s->type == SOCKET_TYPE_INVALID || s->id != id if (s->type == SOCKET_TYPE_INVALID || s->id != id
|| s->type == SOCKET_TYPE_HALFCLOSE || s->type == SOCKET_TYPE_HALFCLOSE
|| s->type == SOCKET_TYPE_PACCEPT) { || s->type == SOCKET_TYPE_PACCEPT) {
@@ -556,7 +566,7 @@ _failed:
result->id = id; result->id = id;
result->ud = 0; result->ud = 0;
result->data = NULL; result->data = NULL;
ss->slot[id % MAX_SOCKET].type = SOCKET_TYPE_INVALID; ss->slot[HASH_ID(id)].type = SOCKET_TYPE_INVALID;
return SOCKET_ERROR; return SOCKET_ERROR;
} }
@@ -564,7 +574,7 @@ _failed:
static int static int
close_socket(struct socket_server *ss, struct request_close *request, struct socket_message *result) { close_socket(struct socket_server *ss, struct request_close *request, struct socket_message *result) {
int id = request->id; int id = request->id;
struct socket * s = &ss->slot[id % MAX_SOCKET]; struct socket * s = &ss->slot[HASH_ID(id)];
if (s->type == SOCKET_TYPE_INVALID || s->id != id) { if (s->type == SOCKET_TYPE_INVALID || s->id != id) {
result->id = id; result->id = id;
result->opaque = request->opaque; result->opaque = request->opaque;
@@ -612,7 +622,7 @@ start_socket(struct socket_server *ss, struct request_start *request, struct soc
result->opaque = request->opaque; result->opaque = request->opaque;
result->ud = 0; result->ud = 0;
result->data = NULL; result->data = NULL;
struct socket *s = &ss->slot[id % MAX_SOCKET]; struct socket *s = &ss->slot[HASH_ID(id)];
if (s->type == SOCKET_TYPE_INVALID || s->id !=id) { if (s->type == SOCKET_TYPE_INVALID || s->id !=id) {
return SOCKET_ERROR; return SOCKET_ERROR;
} }
@@ -633,6 +643,17 @@ start_socket(struct socket_server *ss, struct request_start *request, struct soc
return -1; return -1;
} }
static void
setopt_socket(struct socket_server *ss, struct request_setopt *request) {
int id = request->id;
struct socket *s = &ss->slot[HASH_ID(id)];
if (s->type == SOCKET_TYPE_INVALID || s->id !=id) {
return;
}
int v = request->value;
setsockopt(s->fd, IPPROTO_TCP, request->what, &v, sizeof(v));
}
static void static void
block_readpipe(int pipefd, void *buffer, int sz) { block_readpipe(int pipefd, void *buffer, int sz) {
for (;;) { for (;;) {
@@ -696,6 +717,9 @@ ctrl_cmd(struct socket_server *ss, struct socket_message *result) {
return send_socket(ss, (struct request_send *)buffer, result, PRIORITY_HIGH); return send_socket(ss, (struct request_send *)buffer, result, PRIORITY_HIGH);
case 'P': case 'P':
return send_socket(ss, (struct request_send *)buffer, result, PRIORITY_LOW); return send_socket(ss, (struct request_send *)buffer, result, PRIORITY_LOW);
case 'T':
setopt_socket(ss, (struct request_setopt *)buffer);
return -1;
default: default:
fprintf(stderr, "socket-server: Unknown ctrl %c.\n",type); fprintf(stderr, "socket-server: Unknown ctrl %c.\n",type);
return -1; return -1;
@@ -942,7 +966,7 @@ socket_server_connect(struct socket_server *ss, uintptr_t opaque, const char * a
// return -1 when error // return -1 when error
int64_t int64_t
socket_server_send(struct socket_server *ss, int id, const void * buffer, int sz) { socket_server_send(struct socket_server *ss, int id, const void * buffer, int sz) {
struct socket * s = &ss->slot[id % MAX_SOCKET]; struct socket * s = &ss->slot[HASH_ID(id)];
if (s->id != id || s->type == SOCKET_TYPE_INVALID) { if (s->id != id || s->type == SOCKET_TYPE_INVALID) {
return -1; return -1;
} }
@@ -958,7 +982,7 @@ socket_server_send(struct socket_server *ss, int id, const void * buffer, int sz
void void
socket_server_send_lowpriority(struct socket_server *ss, int id, const void * buffer, int sz) { socket_server_send_lowpriority(struct socket_server *ss, int id, const void * buffer, int sz) {
struct socket * s = &ss->slot[id % MAX_SOCKET]; struct socket * s = &ss->slot[HASH_ID(id)];
if (s->id != id || s->type == SOCKET_TYPE_INVALID) { if (s->id != id || s->type == SOCKET_TYPE_INVALID) {
return; return;
} }
@@ -1053,4 +1077,11 @@ socket_server_start(struct socket_server *ss, uintptr_t opaque, int id) {
send_request(ss, &request, 'S', sizeof(request.u.start)); send_request(ss, &request, 'S', sizeof(request.u.start));
} }
void
socket_server_nodelay(struct socket_server *ss, int id) {
struct request_package request;
request.u.setopt.id = id;
request.u.setopt.what = TCP_NODELAY;
request.u.setopt.value = 1;
send_request(ss, &request, 'T', sizeof(request.u.setopt));
}

View File

@@ -36,4 +36,6 @@ int socket_server_listen(struct socket_server *, uintptr_t opaque, const char *
int socket_server_connect(struct socket_server *, uintptr_t opaque, const char * addr, int port); int socket_server_connect(struct socket_server *, uintptr_t opaque, const char * addr, int port);
int socket_server_bind(struct socket_server *, uintptr_t opaque, int fd); int socket_server_bind(struct socket_server *, uintptr_t opaque, int fd);
void socket_server_nodelay(struct socket_server *, int id);
#endif #endif

View File

@@ -1,4 +1,5 @@
local skynet = require "skynet" local skynet = require "skynet"
local queue = require "skynet.queue"
local i = 0 local i = 0
local hello = "hello" local hello = "hello"
@@ -8,9 +9,27 @@ function response.ping(hello)
return hello return hello
end end
-- response.sleep and accept.hello share one lock
local lock
function accept.sleep(queue, n)
if queue then
lock(
function()
print("queue=",queue, n)
skynet.sleep(n)
end)
else
print("queue=",queue, n)
skynet.sleep(n)
end
end
function accept.hello() function accept.hello()
lock(function()
i = i + 1 i = i + 1
print (i, hello) print (i, hello)
end)
end end
function response.error() function response.error()
@@ -19,6 +38,8 @@ end
function init( ... ) function init( ... )
print ("ping server start:", ...) print ("ping server start:", ...)
-- init queue
lock = queue()
-- You can return "queue" for queue service mode -- You can return "queue" for queue service mode
-- return "queue" -- return "queue"

23
test/testdatacenter.lua Normal file
View File

@@ -0,0 +1,23 @@
local skynet = require "skynet"
local datacenter = require "datacenter"
local function f1()
print("====1==== wait hello")
print("\t1>",datacenter.wait ("hello"))
print("====1==== wait key.foobar")
print("\t1>", pcall(datacenter.wait,"key")) -- will failed, because "key" is a branch
print("\t1>",datacenter.wait ("key", "foobar"))
end
local function f2()
skynet.sleep(10)
print("====2==== set key.foobar")
datacenter.set("key", "foobar", "bingo")
end
skynet.start(function()
datacenter.set("hello", "world")
print(datacenter.get "hello")
skynet.fork(f1)
skynet.fork(f2)
end)

View File

@@ -14,12 +14,10 @@ end)
else else
skynet.start(function() skynet.start(function()
local test = skynet.newservice("testdeadcall", "test") -- launch self in test mode local test = skynet.newservice(SERVICE_NAME, "test") -- launch self in test mode
print(pcall(function() print(pcall(function()
skynet.send(test,"lua", "hello world") skynet.call(test,"lua", "dead call")
skynet.send(test,"lua", "never get there")
skynet.call(test,"lua", "fake call")
end)) end))
skynet.exit() skynet.exit()

16
test/testhttp.lua Normal file
View File

@@ -0,0 +1,16 @@
local skynet = require "skynet"
local httpc = require "http.httpc"
skynet.start(function()
print("GET baidu.com")
local header = {}
local status, body = httpc.get("baidu.com", "/", header)
print("[header] =====>")
for k,v in pairs(header) do
print(k,v)
end
print("[body] =====>", status)
print(body)
skynet.exit()
end)

View File

@@ -27,7 +27,7 @@ skynet.start(function()
local channel = mc.new() local channel = mc.new()
print("New channel", channel) print("New channel", channel)
for i=1,10 do for i=1,10 do
local sub = skynet.newservice("testmulticast", "sub") local sub = skynet.newservice(SERVICE_NAME, "sub")
skynet.call(sub, "lua", "init", channel.channel) skynet.call(sub, "lua", "init", channel.channel)
end end

View File

@@ -4,9 +4,13 @@ local mc = require "multicast"
skynet.start(function() skynet.start(function()
print("remote start") print("remote start")
skynet.monitor("simplemonitor", true)
local console = skynet.newservice("console") local console = skynet.newservice("console")
local channel = dc.get "MCCHANNEL" local channel = dc.get "MCCHANNEL"
if channel then
print("remote channel", channel)
else
print("create local channel")
end
for i=1,10 do for i=1,10 do
local sub = skynet.newservice("testmulticast", "sub") local sub = skynet.newservice("testmulticast", "sub")
skynet.call(sub, "lua", "init", channel) skynet.call(sub, "lua", "init", channel)

18
test/testqueue.lua Normal file
View File

@@ -0,0 +1,18 @@
local skynet = require "skynet"
local snax = require "snax"
skynet.start(function()
local ps = snax.uniqueservice ("pingserver", "test queue")
for i=1, 10 do
ps.post.sleep(true,i*10)
ps.post.hello()
end
for i=1, 10 do
ps.post.sleep(false,i*10)
ps.post.hello()
end
skynet.exit()
end)

View File

@@ -7,9 +7,9 @@ local function echo(id)
socket.start(id) socket.start(id)
while true do while true do
local str = socket.readline(id,"\n") local str = socket.read(id)
if str then if str then
socket.write(id, str .. "\n") socket.write(id, str)
else else
socket.close(id) socket.close(id)
return return
@@ -30,7 +30,7 @@ else
local function accept(id) local function accept(id)
socket.start(id) socket.start(id)
socket.write(id, "Hello Skynet\n") socket.write(id, "Hello Skynet\n")
skynet.newservice("testsocket", "agent", id) skynet.newservice(SERVICE_NAME, "agent", id)
-- notice: Some data on this connection(id) may lost before new service start. -- notice: Some data on this connection(id) may lost before new service start.
-- So, be careful when you want to use start / abandon / start . -- So, be careful when you want to use start / abandon / start .
socket.abandon(id) socket.abandon(id)