recvLine now works with unbuffered ssl sockets.

Added higher level recv functions.
This commit is contained in:
Dominik Picheta 2012-12-22 23:03:28 +00:00
commit 6cb8edfce9

View file

@ -62,6 +62,8 @@ type
sslHandle: PSSL sslHandle: PSSL
sslContext: PSSLContext sslContext: PSSLContext
sslNoHandshake: bool # True if needs handshake. sslNoHandshake: bool # True if needs handshake.
sslHasPeekChar: bool
sslPeekChar: char
of false: nil of false: nil
TSocket* = ref TSocketImpl TSocket* = ref TSocketImpl
@ -291,6 +293,7 @@ when defined(ssl):
socket.sslContext = ctx socket.sslContext = ctx
socket.sslHandle = SSLNew(PSSLCTX(socket.sslContext)) socket.sslHandle = SSLNew(PSSLCTX(socket.sslContext))
socket.sslNoHandshake = false socket.sslNoHandshake = false
socket.sslHasPeekChar = false
if socket.sslHandle == nil: if socket.sslHandle == nil:
SSLError() SSLError()
@ -849,10 +852,7 @@ proc checkBuffer(readfds: var seq[TSocket]): int =
var res: seq[TSocket] = @[] var res: seq[TSocket] = @[]
result = 0 result = 0
for s in readfds: for s in readfds:
if s.isBuffered: if hasDataBuffered(s):
if s.bufLen <= 0 or s.currPos == s.bufLen:
res.add(s)
else:
inc(result) inc(result)
else: else:
res.add(s) res.add(s)
@ -975,11 +975,11 @@ template retRead(flags, readBytes: int) =
proc recv*(socket: TSocket, data: pointer, size: int): int {.tags: [FReadIO].} = proc recv*(socket: TSocket, data: pointer, size: int): int {.tags: [FReadIO].} =
## receives data from a socket ## receives data from a socket
if size == 0: return
if socket.isBuffered: if socket.isBuffered:
if socket.bufLen == 0: if socket.bufLen == 0:
retRead(0'i32, 0) retRead(0'i32, 0)
when true:
var read = 0 var read = 0
while read < size: while read < size:
if socket.currPos >= socket.bufLen: if socket.currPos >= socket.bufLen:
@ -990,27 +990,31 @@ proc recv*(socket: TSocket, data: pointer, size: int): int {.tags: [FReadIO].} =
copyMem(addr(d[read]), addr(socket.buffer[socket.currPos]), chunk) copyMem(addr(d[read]), addr(socket.buffer[socket.currPos]), chunk)
read.inc(chunk) read.inc(chunk)
socket.currPos.inc(chunk) socket.currPos.inc(chunk)
else:
var read = 0
while read < size:
if socket.currPos >= socket.bufLen:
retRead(0'i32, read)
var d = cast[cstring](data)
d[read] = socket.buffer[socket.currPos]
read.inc(1)
socket.currPos.inc(1)
result = read result = read
else: else:
when defined(ssl): when defined(ssl):
if socket.isSSL: if socket.isSSL:
if socket.sslHasPeekChar:
copyMem(data, addr(socket.sslPeekChar), 1)
socket.sslHasPeekChar = false
if size-1 > 0:
var d = cast[cstring](data)
result = SSLRead(socket.sslHandle, addr(d[1]), size-1) + 1
else:
result = 1
else:
result = SSLRead(socket.sslHandle, data, size) result = SSLRead(socket.sslHandle, data, size)
else: else:
result = recv(socket.fd, data, size.cint, 0'i32) result = recv(socket.fd, data, size.cint, 0'i32)
else: else:
result = recv(socket.fd, data, size.cint, 0'i32) result = recv(socket.fd, data, size.cint, 0'i32)
proc recv*(socket: TSocket, data: var string, size: int): int =
## higher-level version of the above
data.setLen(size)
result = recv(socket, cstring(data), size)
proc waitFor(socket: TSocket, waited: var float, timeout: int): int {. proc waitFor(socket: TSocket, waited: var float, timeout: int): int {.
tags: [FTime].} = tags: [FTime].} =
## returns the number of characters available to be read. In unbuffered ## returns the number of characters available to be read. In unbuffered
@ -1045,6 +1049,11 @@ proc recv*(socket: TSocket, data: pointer, size: int, timeout: int): int {.
result = read result = read
proc recv*(socket: TSocket, data: var string, size: int, timeout: int): int =
# higher-level version of the above
data.setLen(size)
result = recv(socket, cstring(data), size, timeout)
proc peekChar(socket: TSocket, c: var char): int {.tags: [FReadIO].} = proc peekChar(socket: TSocket, c: var char): int {.tags: [FReadIO].} =
if socket.isBuffered: if socket.isBuffered:
result = 1 result = 1
@ -1057,8 +1066,12 @@ proc peekChar(socket: TSocket, c: var char): int {.tags: [FReadIO].} =
else: else:
when defined(ssl): when defined(ssl):
if socket.isSSL: if socket.isSSL:
raise newException(ESSL, "Sorry, you cannot use recvLine on an unbuffered SSL socket.") if not socket.sslHasPeekChar:
result = SSLRead(socket.sslHandle, addr(socket.sslPeekChar), 1)
socket.sslHasPeekChar = true
c = socket.sslPeekChar
return
result = recv(socket.fd, addr(c), 1, MSG_PEEK) result = recv(socket.fd, addr(c), 1, MSG_PEEK)
proc recvLine*(socket: TSocket, line: var TaintedString): bool {. proc recvLine*(socket: TSocket, line: var TaintedString): bool {.
@ -1068,13 +1081,11 @@ proc recvLine*(socket: TSocket, line: var TaintedString): bool {.
## will be set to it. ## will be set to it.
## ##
## ``True`` is returned if data is available. ``False`` usually suggests an ## ``True`` is returned if data is available. ``False`` usually suggests an
## error, EOS exceptions are not raised in favour of this. ## error, EOS exceptions are not raised and ``False`` is simply returned
## instead.
## ##
## If the socket is disconnected, ``line`` will be set to ``""`` and ``True`` ## If the socket is disconnected, ``line`` will be set to ``""`` and ``True``
## will be returned. ## will be returned.
##
## **Warning:** Using this function on a unbuffered ssl socket will result
## in an error.
template addNLIfEmpty(): stmt = template addNLIfEmpty(): stmt =
if line.len == 0: if line.len == 0:
line.add("\c\L") line.add("\c\L")