| 1 | -- |
| 2 | -- server.lua by blastbeat of the luadch project |
| 3 | -- Re-used here under the MIT/X Consortium License |
| 4 | -- |
| 5 | -- Modifications (C) 2008-2010 Matthew Wild, Waqas Hussain |
| 6 | -- |
| 7 | |
| 8 | -- // wrapping luadch stuff // -- |
| 9 | |
| 10 | local use = function( what ) |
| 11 | return _G[ what ] |
| 12 | end |
| 13 | local clean = function( tbl ) |
| 14 | for i, k in pairs( tbl ) do |
| 15 | tbl[ i ] = nil |
| 16 | end |
| 17 | end |
| 18 | |
| 19 | local log, table_concat = require ("util.logger").init("socket"), table.concat; |
| 20 | local out_put = function (...) return log("debug", table_concat{...}); end |
| 21 | local out_error = function (...) return log("warn", table_concat{...}); end |
| 22 | local mem_free = collectgarbage |
| 23 | |
| 24 | ----------------------------------// DECLARATION //-- |
| 25 | |
| 26 | --// constants //-- |
| 27 | |
| 28 | local STAT_UNIT = 1 -- byte |
| 29 | |
| 30 | --// lua functions //-- |
| 31 | |
| 32 | local type = use "type" |
| 33 | local pairs = use "pairs" |
| 34 | local ipairs = use "ipairs" |
| 35 | local tostring = use "tostring" |
| 36 | local collectgarbage = use "collectgarbage" |
| 37 | |
| 38 | --// lua libs //-- |
| 39 | |
| 40 | local os = use "os" |
| 41 | local table = use "table" |
| 42 | local string = use "string" |
| 43 | local coroutine = use "coroutine" |
| 44 | |
| 45 | --// lua lib methods //-- |
| 46 | |
| 47 | local os_time = os.time |
| 48 | local os_difftime = os.difftime |
| 49 | local table_concat = table.concat |
| 50 | local table_remove = table.remove |
| 51 | local string_len = string.len |
| 52 | local string_sub = string.sub |
| 53 | local coroutine_wrap = coroutine.wrap |
| 54 | local coroutine_yield = coroutine.yield |
| 55 | |
| 56 | --// extern libs //-- |
| 57 | |
| 58 | local luasec = use "ssl" |
| 59 | local luasocket = use "socket" or require "socket" |
| 60 | |
| 61 | --// extern lib methods //-- |
| 62 | |
| 63 | local ssl_wrap = ( luasec and luasec.wrap ) |
| 64 | local socket_bind = luasocket.bind |
| 65 | local socket_sleep = luasocket.sleep |
| 66 | local socket_select = luasocket.select |
| 67 | local ssl_newcontext = ( luasec and luasec.newcontext ) |
| 68 | |
| 69 | --// functions //-- |
| 70 | |
| 71 | local id |
| 72 | local loop |
| 73 | local stats |
| 74 | local idfalse |
| 75 | local addtimer |
| 76 | local closeall |
| 77 | local addserver |
| 78 | local getserver |
| 79 | local wrapserver |
| 80 | local getsettings |
| 81 | local closesocket |
| 82 | local removesocket |
| 83 | local removeserver |
| 84 | local changetimeout |
| 85 | local wrapconnection |
| 86 | local changesettings |
| 87 | |
| 88 | --// tables //-- |
| 89 | |
| 90 | local _server |
| 91 | local _readlist |
| 92 | local _timerlist |
| 93 | local _sendlist |
| 94 | local _socketlist |
| 95 | local _closelist |
| 96 | local _readtimes |
| 97 | local _writetimes |
| 98 | |
| 99 | --// simple data types //-- |
| 100 | |
| 101 | local _ |
| 102 | local _readlistlen |
| 103 | local _sendlistlen |
| 104 | local _timerlistlen |
| 105 | |
| 106 | local _sendtraffic |
| 107 | local _readtraffic |
| 108 | |
| 109 | local _selecttimeout |
| 110 | local _sleeptime |
| 111 | |
| 112 | local _starttime |
| 113 | local _currenttime |
| 114 | |
| 115 | local _maxsendlen |
| 116 | local _maxreadlen |
| 117 | |
| 118 | local _checkinterval |
| 119 | local _sendtimeout |
| 120 | local _readtimeout |
| 121 | |
| 122 | local _cleanqueue |
| 123 | |
| 124 | local _timer |
| 125 | |
| 126 | local _maxclientsperserver |
| 127 | |
| 128 | ----------------------------------// DEFINITION //-- |
| 129 | |
| 130 | _server = { } -- key = port, value = table; list of listening servers |
| 131 | _readlist = { } -- array with sockets to read from |
| 132 | _sendlist = { } -- arrary with sockets to write to |
| 133 | _timerlist = { } -- array of timer functions |
| 134 | _socketlist = { } -- key = socket, value = wrapped socket (handlers) |
| 135 | _readtimes = { } -- key = handler, value = timestamp of last data reading |
| 136 | _writetimes = { } -- key = handler, value = timestamp of last data writing/sending |
| 137 | _closelist = { } -- handlers to close |
| 138 | |
| 139 | _readlistlen = 0 -- length of readlist |
| 140 | _sendlistlen = 0 -- length of sendlist |
| 141 | _timerlistlen = 0 -- lenght of timerlist |
| 142 | |
| 143 | _sendtraffic = 0 -- some stats |
| 144 | _readtraffic = 0 |
| 145 | |
| 146 | _selecttimeout = 1 -- timeout of socket.select |
| 147 | _sleeptime = 0 -- time to wait at the end of every loop |
| 148 | |
| 149 | _maxsendlen = 51000 * 1024 -- max len of send buffer |
| 150 | _maxreadlen = 25000 * 1024 -- max len of read buffer |
| 151 | |
| 152 | _checkinterval = 1200000 -- interval in secs to check idle clients |
| 153 | _sendtimeout = 60000 -- allowed send idle time in secs |
| 154 | _readtimeout = 6 * 60 * 60 -- allowed read idle time in secs |
| 155 | |
| 156 | _cleanqueue = false -- clean bufferqueue after using |
| 157 | |
| 158 | _maxclientsperserver = 1000 |
| 159 | |
| 160 | _maxsslhandshake = 30 -- max handshake round-trips |
| 161 | |
| 162 | ----------------------------------// PRIVATE //-- |
| 163 | |
| 164 | wrapserver = function( listeners, socket, ip, serverport, pattern, sslctx, maxconnections ) -- this function wraps a server |
| 165 | |
| 166 | maxconnections = maxconnections or _maxclientsperserver |
| 167 | |
| 168 | local connections = 0 |
| 169 | |
| 170 | local dispatch, disconnect = listeners.onconnect or listeners.onincoming, listeners.ondisconnect |
| 171 | |
| 172 | local accept = socket.accept |
| 173 | |
| 174 | --// public methods of the object //-- |
| 175 | |
| 176 | local handler = { } |
| 177 | |
| 178 | handler.shutdown = function( ) end |
| 179 | |
| 180 | handler.ssl = function( ) |
| 181 | return sslctx ~= nil |
| 182 | end |
| 183 | handler.sslctx = function( ) |
| 184 | return sslctx |
| 185 | end |
| 186 | handler.remove = function( ) |
| 187 | connections = connections - 1 |
| 188 | end |
| 189 | handler.close = function( ) |
| 190 | for _, handler in pairs( _socketlist ) do |
| 191 | if handler.serverport == serverport then |
| 192 | handler.disconnect( handler, "server closed" ) |
| 193 | handler:close( true ) |
| 194 | end |
| 195 | end |
| 196 | socket:close( ) |
| 197 | _sendlistlen = removesocket( _sendlist, socket, _sendlistlen ) |
| 198 | _readlistlen = removesocket( _readlist, socket, _readlistlen ) |
| 199 | _socketlist[ socket ] = nil |
| 200 | handler = nil |
| 201 | socket = nil |
| 202 | --mem_free( ) |
| 203 | out_put "server.lua: closed server handler and removed sockets from list" |
| 204 | end |
| 205 | handler.ip = function( ) |
| 206 | return ip |
| 207 | end |
| 208 | handler.serverport = function( ) |
| 209 | return serverport |
| 210 | end |
| 211 | handler.socket = function( ) |
| 212 | return socket |
| 213 | end |
| 214 | handler.readbuffer = function( ) |
| 215 | if connections > maxconnections then |
| 216 | out_put( "server.lua: refused new client connection: server full" ) |
| 217 | return false |
| 218 | end |
| 219 | local client, err = accept( socket ) -- try to accept |
| 220 | if client then |
| 221 | local ip, clientport = client:getpeername( ) |
| 222 | client:settimeout( 0 ) |
| 223 | local handler, client, err = wrapconnection( handler, listeners, client, ip, serverport, clientport, pattern, sslctx ) -- wrap new client socket |
| 224 | if err then -- error while wrapping ssl socket |
| 225 | return false |
| 226 | end |
| 227 | connections = connections + 1 |
| 228 | out_put( "server.lua: accepted new client connection from ", tostring(ip), ":", tostring(clientport), " to ", tostring(serverport)) |
| 229 | return dispatch( handler ) |
| 230 | elseif err then -- maybe timeout or something else |
| 231 | out_put( "server.lua: error with new client connection: ", tostring(err) ) |
| 232 | return false |
| 233 | end |
| 234 | end |
| 235 | return handler |
| 236 | end |
| 237 | |
| 238 | wrapconnection = function( server, listeners, socket, ip, serverport, clientport, pattern, sslctx ) -- this function wraps a client to a handler object |
| 239 | |
| 240 | socket:settimeout( 0 ) |
| 241 | |
| 242 | --// local import of socket methods //-- |
| 243 | |
| 244 | local send |
| 245 | local receive |
| 246 | local shutdown |
| 247 | |
| 248 | --// private closures of the object //-- |
| 249 | |
| 250 | local ssl |
| 251 | |
| 252 | local dispatch = listeners.onincoming |
| 253 | local status = listeners.onstatus |
| 254 | local disconnect = listeners.ondisconnect |
| 255 | local drain = listeners.ondrain |
| 256 | |
| 257 | local bufferqueue = { } -- buffer array |
| 258 | local bufferqueuelen = 0 -- end of buffer array |
| 259 | |
| 260 | local toclose |
| 261 | local fatalerror |
| 262 | local needtls |
| 263 | |
| 264 | local bufferlen = 0 |
| 265 | |
| 266 | local noread = false |
| 267 | local nosend = false |
| 268 | |
| 269 | local sendtraffic, readtraffic = 0, 0 |
| 270 | |
| 271 | local maxsendlen = _maxsendlen |
| 272 | local maxreadlen = _maxreadlen |
| 273 | |
| 274 | --// public methods of the object //-- |
| 275 | |
| 276 | local handler = bufferqueue -- saves a table ^_^ |
| 277 | |
| 278 | handler.dispatch = function( ) |
| 279 | return dispatch |
| 280 | end |
| 281 | handler.disconnect = function( ) |
| 282 | return disconnect |
| 283 | end |
| 284 | handler.setlistener = function( self, listeners ) |
| 285 | dispatch = listeners.onincoming |
| 286 | disconnect = listeners.ondisconnect |
| 287 | status = listeners.onstatus |
| 288 | drain = listeners.ondrain |
| 289 | end |
| 290 | handler.getstats = function( ) |
| 291 | return readtraffic, sendtraffic |
| 292 | end |
| 293 | handler.ssl = function( ) |
| 294 | return ssl |
| 295 | end |
| 296 | handler.sslctx = function ( ) |
| 297 | return sslctx |
| 298 | end |
| 299 | handler.send = function( _, data, i, j ) |
| 300 | return send( socket, data, i, j ) |
| 301 | end |
| 302 | handler.receive = function( pattern, prefix ) |
| 303 | return receive( socket, pattern, prefix ) |
| 304 | end |
| 305 | handler.shutdown = function( pattern ) |
| 306 | return shutdown( socket, pattern ) |
| 307 | end |
| 308 | handler.setoption = function (self, option, value) |
| 309 | if socket.setoption then |
| 310 | return socket:setoption(option, value); |
| 311 | end |
| 312 | return false, "setoption not implemented"; |
| 313 | end |
| 314 | handler.close = function( self, forced ) |
| 315 | if not handler then return true; end |
| 316 | _readlistlen = removesocket( _readlist, socket, _readlistlen ) |
| 317 | _readtimes[ handler ] = nil |
| 318 | if bufferqueuelen ~= 0 then |
| 319 | if not ( forced or fatalerror ) then |
| 320 | handler.sendbuffer( ) |
| 321 | if bufferqueuelen ~= 0 then -- try again... |
| 322 | if handler then |
| 323 | handler.write = nil -- ... but no further writing allowed |
| 324 | end |
| 325 | toclose = true |
| 326 | return false |
| 327 | end |
| 328 | else |
| 329 | send( socket, table_concat( bufferqueue, "", 1, bufferqueuelen ), 1, bufferlen ) -- forced send |
| 330 | end |
| 331 | end |
| 332 | if socket then |
| 333 | _ = shutdown and shutdown( socket ) |
| 334 | socket:close( ) |
| 335 | _sendlistlen = removesocket( _sendlist, socket, _sendlistlen ) |
| 336 | _socketlist[ socket ] = nil |
| 337 | socket = nil |
| 338 | else |
| 339 | out_put "server.lua: socket already closed" |
| 340 | end |
| 341 | if handler then |
| 342 | _writetimes[ handler ] = nil |
| 343 | _closelist[ handler ] = nil |
| 344 | handler = nil |
| 345 | end |
| 346 | if server then |
| 347 | server.remove( ) |
| 348 | end |
| 349 | out_put "server.lua: closed client handler and removed socket from list" |
| 350 | return true |
| 351 | end |
| 352 | handler.ip = function( ) |
| 353 | return ip |
| 354 | end |
| 355 | handler.serverport = function( ) |
| 356 | return serverport |
| 357 | end |
| 358 | handler.clientport = function( ) |
| 359 | return clientport |
| 360 | end |
| 361 | local write = function( self, data ) |
| 362 | bufferlen = bufferlen + string_len( data ) |
| 363 | if bufferlen > maxsendlen then |
| 364 | _closelist[ handler ] = "send buffer exceeded" -- cannot close the client at the moment, have to wait to the end of the cycle |
| 365 | handler.write = idfalse -- dont write anymore |
| 366 | return false |
| 367 | elseif socket and not _sendlist[ socket ] then |
| 368 | _sendlistlen = addsocket(_sendlist, socket, _sendlistlen) |
| 369 | end |
| 370 | bufferqueuelen = bufferqueuelen + 1 |
| 371 | bufferqueue[ bufferqueuelen ] = data |
| 372 | if handler then |
| 373 | _writetimes[ handler ] = _writetimes[ handler ] or _currenttime |
| 374 | end |
| 375 | return true |
| 376 | end |
| 377 | handler.write = write |
| 378 | handler.bufferqueue = function( self ) |
| 379 | return bufferqueue |
| 380 | end |
| 381 | handler.socket = function( self ) |
| 382 | return socket |
| 383 | end |
| 384 | handler.set_mode = function( self, new ) |
| 385 | pattern = new or pattern |
| 386 | return pattern |
| 387 | end |
| 388 | handler.set_send = function ( self, newsend ) |
| 389 | send = newsend or send |
| 390 | return send |
| 391 | end |
| 392 | handler.bufferlen = function( self, readlen, sendlen ) |
| 393 | maxsendlen = sendlen or maxsendlen |
| 394 | maxreadlen = readlen or maxreadlen |
| 395 | return bufferlen, maxreadlen, maxsendlen |
| 396 | end |
| 397 | --TODO: Deprecate |
| 398 | handler.lock_read = function (self, switch) |
| 399 | if switch == true then |
| 400 | local tmp = _readlistlen |
| 401 | _readlistlen = removesocket( _readlist, socket, _readlistlen ) |
| 402 | _readtimes[ handler ] = nil |
| 403 | if _readlistlen ~= tmp then |
| 404 | noread = true |
| 405 | end |
| 406 | elseif switch == false then |
| 407 | if noread then |
| 408 | noread = false |
| 409 | _readlistlen = addsocket(_readlist, socket, _readlistlen) |
| 410 | _readtimes[ handler ] = _currenttime |
| 411 | end |
| 412 | end |
| 413 | return noread |
| 414 | end |
| 415 | handler.pause = function (self) |
| 416 | return self:lock_read(true); |
| 417 | end |
| 418 | handler.resume = function (self) |
| 419 | return self:lock_read(false); |
| 420 | end |
| 421 | handler.lock = function( self, switch ) |
| 422 | handler.lock_read (switch) |
| 423 | if switch == true then |
| 424 | handler.write = idfalse |
| 425 | local tmp = _sendlistlen |
| 426 | _sendlistlen = removesocket( _sendlist, socket, _sendlistlen ) |
| 427 | _writetimes[ handler ] = nil |
| 428 | if _sendlistlen ~= tmp then |
| 429 | nosend = true |
| 430 | end |
| 431 | elseif switch == false then |
| 432 | handler.write = write |
| 433 | if nosend then |
| 434 | nosend = false |
| 435 | write( "" ) |
| 436 | end |
| 437 | end |
| 438 | return noread, nosend |
| 439 | end |
| 440 | local _readbuffer = function( ) -- this function reads data |
| 441 | local buffer, err, part = receive( socket, pattern ) -- receive buffer with "pattern" |
| 442 | if not err or (err == "wantread" or err == "timeout") then -- received something |
| 443 | local buffer = buffer or part or "" |
| 444 | local len = string_len( buffer ) |
| 445 | if len > maxreadlen then |
| 446 | disconnect( handler, "receive buffer exceeded" ) |
| 447 | handler:close( true ) |
| 448 | return false |
| 449 | end |
| 450 | local count = len * STAT_UNIT |
| 451 | readtraffic = readtraffic + count |
| 452 | _readtraffic = _readtraffic + count |
| 453 | _readtimes[ handler ] = _currenttime |
| 454 | --out_put( "server.lua: read data '", buffer:gsub("[^%w%p ]", "."), "', error: ", err ) |
| 455 | return dispatch( handler, buffer, err ) |
| 456 | else -- connections was closed or fatal error |
| 457 | out_put( "server.lua: client ", tostring(ip), ":", tostring(clientport), " read error: ", tostring(err) ) |
| 458 | fatalerror = true |
| 459 | disconnect( handler, err ) |
| 460 | _ = handler and handler:close( ) |
| 461 | return false |
| 462 | end |
| 463 | end |
| 464 | local _sendbuffer = function( ) -- this function sends data |
| 465 | local succ, err, byte, buffer, count; |
| 466 | local count; |
| 467 | if socket then |
| 468 | buffer = table_concat( bufferqueue, "", 1, bufferqueuelen ) |
| 469 | succ, err, byte = send( socket, buffer, 1, bufferlen ) |
| 470 | count = ( succ or byte or 0 ) * STAT_UNIT |
| 471 | sendtraffic = sendtraffic + count |
| 472 | _sendtraffic = _sendtraffic + count |
| 473 | _ = _cleanqueue and clean( bufferqueue ) |
| 474 | --out_put( "server.lua: sended '", buffer, "', bytes: ", tostring(succ), ", error: ", tostring(err), ", part: ", tostring(byte), ", to: ", tostring(ip), ":", tostring(clientport) ) |
| 475 | else |
| 476 | succ, err, count = false, "closed", 0; |
| 477 | end |
| 478 | if succ then -- sending succesful |
| 479 | bufferqueuelen = 0 |
| 480 | bufferlen = 0 |
| 481 | _sendlistlen = removesocket( _sendlist, socket, _sendlistlen ) -- delete socket from writelist |
| 482 | _writetimes[ handler ] = nil |
| 483 | if drain then |
| 484 | drain(handler) |
| 485 | end |
| 486 | _ = needtls and handler:starttls(nil) |
| 487 | _ = toclose and handler:close( ) |
| 488 | return true |
| 489 | elseif byte and ( err == "timeout" or err == "wantwrite" ) then -- want write |
| 490 | buffer = string_sub( buffer, byte + 1, bufferlen ) -- new buffer |
| 491 | bufferqueue[ 1 ] = buffer -- insert new buffer in queue |
| 492 | bufferqueuelen = 1 |
| 493 | bufferlen = bufferlen - byte |
| 494 | _writetimes[ handler ] = _currenttime |
| 495 | return true |
| 496 | else -- connection was closed during sending or fatal error |
| 497 | out_put( "server.lua: client ", tostring(ip), ":", tostring(clientport), " write error: ", tostring(err) ) |
| 498 | fatalerror = true |
| 499 | disconnect( handler, err ) |
| 500 | _ = handler and handler:close( ) |
| 501 | return false |
| 502 | end |
| 503 | end |
| 504 | |
| 505 | -- Set the sslctx |
| 506 | local handshake; |
| 507 | function handler.set_sslctx(self, new_sslctx) |
| 508 | ssl = true |
| 509 | sslctx = new_sslctx; |
| 510 | local wrote |
| 511 | local read |
| 512 | handshake = coroutine_wrap( function( client ) -- create handshake coroutine |
| 513 | local err |
| 514 | for i = 1, _maxsslhandshake do |
| 515 | _sendlistlen = ( wrote and removesocket( _sendlist, client, _sendlistlen ) ) or _sendlistlen |
| 516 | _readlistlen = ( read and removesocket( _readlist, client, _readlistlen ) ) or _readlistlen |
| 517 | read, wrote = nil, nil |
| 518 | _, err = client:dohandshake( ) |
| 519 | if not err then |
| 520 | out_put( "server.lua: ssl handshake done" ) |
| 521 | handler.readbuffer = _readbuffer -- when handshake is done, replace the handshake function with regular functions |
| 522 | handler.sendbuffer = _sendbuffer |
| 523 | _ = status and status( handler, "ssl-handshake-complete" ) |
| 524 | _readlistlen = addsocket(_readlist, client, _readlistlen) |
| 525 | return true |
| 526 | else |
| 527 | out_put( "server.lua: error during ssl handshake: ", tostring(err) ) |
| 528 | if err == "wantwrite" and not wrote then |
| 529 | _sendlistlen = addsocket(_sendlist, client, _sendlistlen) |
| 530 | wrote = true |
| 531 | elseif err == "wantread" and not read then |
| 532 | _readlistlen = addsocket(_readlist, client, _readlistlen) |
| 533 | read = true |
| 534 | else |
| 535 | break; |
| 536 | end |
| 537 | --coroutine_yield( handler, nil, err ) -- handshake not finished |
| 538 | coroutine_yield( ) |
| 539 | end |
| 540 | end |
| 541 | disconnect( handler, "ssl handshake failed" ) |
| 542 | _ = handler and handler:close( true ) -- forced disconnect |
| 543 | return false -- handshake failed |
| 544 | end |
| 545 | ) |
| 546 | end |
| 547 | if luasec then |
| 548 | if sslctx then -- ssl? |
| 549 | handler:set_sslctx(sslctx); |
| 550 | out_put("server.lua: ", "starting ssl handshake") |
| 551 | local err |
| 552 | socket, err = ssl_wrap( socket, sslctx ) -- wrap socket |
| 553 | if err then |
| 554 | out_put( "server.lua: ssl error: ", tostring(err) ) |
| 555 | --mem_free( ) |
| 556 | return nil, nil, err -- fatal error |
| 557 | end |
| 558 | socket:settimeout( 0 ) |
| 559 | handler.readbuffer = handshake |
| 560 | handler.sendbuffer = handshake |
| 561 | handshake( socket ) -- do handshake |
| 562 | if not socket then |
| 563 | return nil, nil, "ssl handshake failed"; |
| 564 | end |
| 565 | else |
| 566 | local sslctx; |
| 567 | handler.starttls = function( self, _sslctx) |
| 568 | if _sslctx then |
| 569 | sslctx = _sslctx; |
| 570 | handler:set_sslctx(sslctx); |
| 571 | end |
| 572 | if bufferqueuelen > 0 then |
| 573 | out_put "server.lua: we need to do tls, but delaying until send buffer empty" |
| 574 | needtls = true |
| 575 | return |
| 576 | end |
| 577 | out_put( "server.lua: attempting to start tls on " .. tostring( socket ) ) |
| 578 | local oldsocket, err = socket |
| 579 | socket, err = ssl_wrap( socket, sslctx ) -- wrap socket |
| 580 | --out_put( "server.lua: sslwrapped socket is " .. tostring( socket ) ) |
| 581 | if err then |
| 582 | out_put( "server.lua: error while starting tls on client: ", tostring(err) ) |
| 583 | return nil, err -- fatal error |
| 584 | end |
| 585 | |
| 586 | socket:settimeout( 0 ) |
| 587 | |
| 588 | -- add the new socket to our system |
| 589 | |
| 590 | send = socket.send |
| 591 | receive = socket.receive |
| 592 | shutdown = id |
| 593 | |
| 594 | _socketlist[ socket ] = handler |
| 595 | _readlistlen = addsocket(_readlist, socket, _readlistlen) |
| 596 | |
| 597 | -- remove traces of the old socket |
| 598 | |
| 599 | _readlistlen = removesocket( _readlist, oldsocket, _readlistlen ) |
| 600 | _sendlistlen = removesocket( _sendlist, oldsocket, _sendlistlen ) |
| 601 | _socketlist[ oldsocket ] = nil |
| 602 | |
| 603 | handler.starttls = nil |
| 604 | needtls = nil |
| 605 | |
| 606 | -- Secure now |
| 607 | ssl = true |
| 608 | |
| 609 | handler.readbuffer = handshake |
| 610 | handler.sendbuffer = handshake |
| 611 | handshake( socket ) -- do handshake |
| 612 | end |
| 613 | handler.readbuffer = _readbuffer |
| 614 | handler.sendbuffer = _sendbuffer |
| 615 | end |
| 616 | else |
| 617 | handler.readbuffer = _readbuffer |
| 618 | handler.sendbuffer = _sendbuffer |
| 619 | end |
| 620 | send = socket.send |
| 621 | receive = socket.receive |
| 622 | shutdown = ( ssl and id ) or socket.shutdown |
| 623 | |
| 624 | _socketlist[ socket ] = handler |
| 625 | _readlistlen = addsocket(_readlist, socket, _readlistlen) |
| 626 | return handler, socket |
| 627 | end |
| 628 | |
| 629 | id = function( ) |
| 630 | end |
| 631 | |
| 632 | idfalse = function( ) |
| 633 | return false |
| 634 | end |
| 635 | |
| 636 | addsocket = function( list, socket, len ) |
| 637 | if not list[ socket ] then |
| 638 | len = len + 1 |
| 639 | list[ len ] = socket |
| 640 | list[ socket ] = len |
| 641 | end |
| 642 | return len; |
| 643 | end |
| 644 | |
| 645 | removesocket = function( list, socket, len ) -- this function removes sockets from a list ( copied from copas ) |
| 646 | local pos = list[ socket ] |
| 647 | if pos then |
| 648 | list[ socket ] = nil |
| 649 | local last = list[ len ] |
| 650 | list[ len ] = nil |
| 651 | if last ~= socket then |
| 652 | list[ last ] = pos |
| 653 | list[ pos ] = last |
| 654 | end |
| 655 | return len - 1 |
| 656 | end |
| 657 | return len |
| 658 | end |
| 659 | |
| 660 | closesocket = function( socket ) |
| 661 | _sendlistlen = removesocket( _sendlist, socket, _sendlistlen ) |
| 662 | _readlistlen = removesocket( _readlist, socket, _readlistlen ) |
| 663 | _socketlist[ socket ] = nil |
| 664 | socket:close( ) |
| 665 | --mem_free( ) |
| 666 | end |
| 667 | |
| 668 | local function link(sender, receiver, buffersize) |
| 669 | sender:set_mode(buffersize); |
| 670 | local sender_locked; |
| 671 | local _sendbuffer = receiver.sendbuffer; |
| 672 | function receiver.sendbuffer() |
| 673 | _sendbuffer(); |
| 674 | if sender_locked and receiver.bufferlen() < buffersize then |
| 675 | sender:lock_read(false); -- Unlock now |
| 676 | sender_locked = nil; |
| 677 | end |
| 678 | end |
| 679 | |
| 680 | local _readbuffer = sender.readbuffer; |
| 681 | function sender.readbuffer() |
| 682 | _readbuffer(); |
| 683 | if not sender_locked and receiver.bufferlen() >= buffersize then |
| 684 | sender_locked = true; |
| 685 | sender:lock_read(true); |
| 686 | end |
| 687 | end |
| 688 | end |
| 689 | |
| 690 | ----------------------------------// PUBLIC //-- |
| 691 | |
| 692 | addserver = function( addr, port, listeners, pattern, sslctx ) -- this function provides a way for other scripts to reg a server |
| 693 | local err |
| 694 | if type( listeners ) ~= "table" then |
| 695 | err = "invalid listener table" |
| 696 | end |
| 697 | if type( port ) ~= "number" or not ( port >= 0 and port <= 65535 ) then |
| 698 | err = "invalid port" |
| 699 | elseif _server[ port ] then |
| 700 | err = "listeners on port '" .. port .. "' already exist" |
| 701 | elseif sslctx and not luasec then |
| 702 | err = "luasec not found" |
| 703 | end |
| 704 | if err then |
| 705 | out_error( "server.lua, port ", port, ": ", err ) |
| 706 | return nil, err |
| 707 | end |
| 708 | addr = addr or "*" |
| 709 | local server, err = socket_bind( addr, port ) |
| 710 | if err then |
| 711 | out_error( "server.lua, port ", port, ": ", err ) |
| 712 | return nil, err |
| 713 | end |
| 714 | local handler, err = wrapserver( listeners, server, addr, port, pattern, sslctx, _maxclientsperserver ) -- wrap new server socket |
| 715 | if not handler then |
| 716 | server:close( ) |
| 717 | return nil, err |
| 718 | end |
| 719 | server:settimeout( 0 ) |
| 720 | _readlistlen = addsocket(_readlist, server, _readlistlen) |
| 721 | _server[ port ] = handler |
| 722 | _socketlist[ server ] = handler |
| 723 | out_put( "server.lua: new "..(sslctx and "ssl " or "").."server listener on '", addr, ":", port, "'" ) |
| 724 | return handler |
| 725 | end |
| 726 | |
| 727 | getserver = function ( port ) |
| 728 | return _server[ port ]; |
| 729 | end |
| 730 | |
| 731 | removeserver = function( port ) |
| 732 | local handler = _server[ port ] |
| 733 | if not handler then |
| 734 | return nil, "no server found on port '" .. tostring( port ) .. "'" |
| 735 | end |
| 736 | handler:close( ) |
| 737 | _server[ port ] = nil |
| 738 | return true |
| 739 | end |
| 740 | |
| 741 | closeall = function( ) |
| 742 | for _, handler in pairs( _socketlist ) do |
| 743 | handler:close( ) |
| 744 | _socketlist[ _ ] = nil |
| 745 | end |
| 746 | _readlistlen = 0 |
| 747 | _sendlistlen = 0 |
| 748 | _timerlistlen = 0 |
| 749 | _server = { } |
| 750 | _readlist = { } |
| 751 | _sendlist = { } |
| 752 | _timerlist = { } |
| 753 | _socketlist = { } |
| 754 | --mem_free( ) |
| 755 | end |
| 756 | |
| 757 | getsettings = function( ) |
| 758 | return _selecttimeout, _sleeptime, _maxsendlen, _maxreadlen, _checkinterval, _sendtimeout, _readtimeout, _cleanqueue, _maxclientsperserver, _maxsslhandshake |
| 759 | end |
| 760 | |
| 761 | changesettings = function( new ) |
| 762 | if type( new ) ~= "table" then |
| 763 | return nil, "invalid settings table" |
| 764 | end |
| 765 | _selecttimeout = tonumber( new.timeout ) or _selecttimeout |
| 766 | _sleeptime = tonumber( new.sleeptime ) or _sleeptime |
| 767 | _maxsendlen = tonumber( new.maxsendlen ) or _maxsendlen |
| 768 | _maxreadlen = tonumber( new.maxreadlen ) or _maxreadlen |
| 769 | _checkinterval = tonumber( new.checkinterval ) or _checkinterval |
| 770 | _sendtimeout = tonumber( new.sendtimeout ) or _sendtimeout |
| 771 | _readtimeout = tonumber( new.readtimeout ) or _readtimeout |
| 772 | _cleanqueue = new.cleanqueue |
| 773 | _maxclientsperserver = new._maxclientsperserver or _maxclientsperserver |
| 774 | _maxsslhandshake = new._maxsslhandshake or _maxsslhandshake |
| 775 | return true |
| 776 | end |
| 777 | |
| 778 | addtimer = function( listener ) |
| 779 | if type( listener ) ~= "function" then |
| 780 | return nil, "invalid listener function" |
| 781 | end |
| 782 | _timerlistlen = _timerlistlen + 1 |
| 783 | _timerlist[ _timerlistlen ] = listener |
| 784 | return true |
| 785 | end |
| 786 | |
| 787 | stats = function( ) |
| 788 | return _readtraffic, _sendtraffic, _readlistlen, _sendlistlen, _timerlistlen |
| 789 | end |
| 790 | |
| 791 | local dontstop = true; -- thinking about tomorrow, ... |
| 792 | |
| 793 | setquitting = function (quit) |
| 794 | dontstop = not quit; |
| 795 | return; |
| 796 | end |
| 797 | |
| 798 | loop = function( ) -- this is the main loop of the program |
| 799 | while dontstop do |
| 800 | local read, write, err = socket_select( _readlist, _sendlist, _selecttimeout ) |
| 801 | for i, socket in ipairs( write ) do -- send data waiting in writequeues |
| 802 | local handler = _socketlist[ socket ] |
| 803 | if handler then |
| 804 | handler.sendbuffer( ) |
| 805 | else |
| 806 | closesocket( socket ) |
| 807 | out_put "server.lua: found no handler and closed socket (writelist)" -- this should not happen |
| 808 | end |
| 809 | end |
| 810 | for i, socket in ipairs( read ) do -- receive data |
| 811 | local handler = _socketlist[ socket ] |
| 812 | if handler then |
| 813 | handler.readbuffer( ) |
| 814 | else |
| 815 | closesocket( socket ) |
| 816 | out_put "server.lua: found no handler and closed socket (readlist)" -- this can happen |
| 817 | end |
| 818 | end |
| 819 | for handler, err in pairs( _closelist ) do |
| 820 | handler.disconnect( )( handler, err ) |
| 821 | handler:close( true ) -- forced disconnect |
| 822 | end |
| 823 | clean( _closelist ) |
| 824 | _currenttime = os_time( ) |
| 825 | if os_difftime( _currenttime - _timer ) >= 1 then |
| 826 | for i = 1, _timerlistlen do |
| 827 | _timerlist[ i ]( _currenttime ) -- fire timers |
| 828 | end |
| 829 | _timer = _currenttime |
| 830 | end |
| 831 | socket_sleep( _sleeptime ) -- wait some time |
| 832 | --collectgarbage( ) |
| 833 | end |
| 834 | return "quitting" |
| 835 | end |
| 836 | |
| 837 | local function get_backend() |
| 838 | return "select"; |
| 839 | end |
| 840 | |
| 841 | --// EXPERIMENTAL //-- |
| 842 | |
| 843 | local wrapclient = function( socket, ip, serverport, listeners, pattern, sslctx ) |
| 844 | local handler = wrapconnection( nil, listeners, socket, ip, serverport, "clientport", pattern, sslctx ) |
| 845 | _socketlist[ socket ] = handler |
| 846 | _sendlistlen = addsocket(_sendlist, socket, _sendlistlen) |
| 847 | if listeners.onconnect then |
| 848 | -- When socket is writeable, call onconnect |
| 849 | local _sendbuffer = handler.sendbuffer; |
| 850 | handler.sendbuffer = function () |
| 851 | listeners.onconnect(handler); |
| 852 | handler.sendbuffer = _sendbuffer; |
| 853 | -- If there was data with the incoming packet, handle it now. |
| 854 | if #handler:bufferqueue() > 0 then |
| 855 | return _sendbuffer(); |
| 856 | end |
| 857 | end |
| 858 | end |
| 859 | return handler, socket |
| 860 | end |
| 861 | |
| 862 | local addclient = function( address, port, listeners, pattern, sslctx ) |
| 863 | local client, err = luasocket.tcp( ) |
| 864 | if err then |
| 865 | return nil, err |
| 866 | end |
| 867 | client:settimeout( 0 ) |
| 868 | _, err = client:connect( address, port ) |
| 869 | if err then -- try again |
| 870 | local handler = wrapclient( client, address, port, listeners ) |
| 871 | else |
| 872 | wrapconnection( nil, listeners, client, address, port, "clientport", pattern, sslctx ) |
| 873 | end |
| 874 | end |
| 875 | |
| 876 | --// EXPERIMENTAL //-- |
| 877 | |
| 878 | ----------------------------------// BEGIN //-- |
| 879 | |
| 880 | use "setmetatable" ( _socketlist, { __mode = "k" } ) |
| 881 | use "setmetatable" ( _readtimes, { __mode = "k" } ) |
| 882 | use "setmetatable" ( _writetimes, { __mode = "k" } ) |
| 883 | |
| 884 | _timer = os_time( ) |
| 885 | _starttime = os_time( ) |
| 886 | |
| 887 | addtimer( function( ) |
| 888 | local difftime = os_difftime( _currenttime - _starttime ) |
| 889 | if difftime > _checkinterval then |
| 890 | _starttime = _currenttime |
| 891 | for handler, timestamp in pairs( _writetimes ) do |
| 892 | if os_difftime( _currenttime - timestamp ) > _sendtimeout then |
| 893 | --_writetimes[ handler ] = nil |
| 894 | handler.disconnect( )( handler, "send timeout" ) |
| 895 | handler:close( true ) -- forced disconnect |
| 896 | end |
| 897 | end |
| 898 | for handler, timestamp in pairs( _readtimes ) do |
| 899 | if os_difftime( _currenttime - timestamp ) > _readtimeout then |
| 900 | --_readtimes[ handler ] = nil |
| 901 | handler.disconnect( )( handler, "read timeout" ) |
| 902 | handler:close( ) -- forced disconnect? |
| 903 | end |
| 904 | end |
| 905 | end |
| 906 | end |
| 907 | ) |
| 908 | |
| 909 | local function setlogger(new_logger) |
| 910 | local old_logger = log; |
| 911 | if new_logger then |
| 912 | log = new_logger; |
| 913 | end |
| 914 | return old_logger; |
| 915 | end |
| 916 | |
| 917 | ----------------------------------// PUBLIC INTERFACE //-- |
| 918 | |
| 919 | return { |
| 920 | |
| 921 | addclient = addclient, |
| 922 | wrapclient = wrapclient, |
| 923 | |
| 924 | loop = loop, |
| 925 | link = link, |
| 926 | stats = stats, |
| 927 | closeall = closeall, |
| 928 | addtimer = addtimer, |
| 929 | addserver = addserver, |
| 930 | getserver = getserver, |
| 931 | setlogger = setlogger, |
| 932 | getsettings = getsettings, |
| 933 | setquitting = setquitting, |
| 934 | removeserver = removeserver, |
| 935 | get_backend = get_backend, |
| 936 | changesettings = changesettings, |
| 937 | } |