This commit is contained in:
Dominik Picheta 2014-12-26 19:13:27 +00:00
commit c35182aca7
2 changed files with 42 additions and 26 deletions

View file

@ -69,13 +69,13 @@ type
# TODO: I would prefer to just do: # TODO: I would prefer to just do:
# AsyncSocket* {.borrow: `.`.} = distinct Socket. But that doesn't work. # AsyncSocket* {.borrow: `.`.} = distinct Socket. But that doesn't work.
AsyncSocketDesc = object AsyncSocketDesc = object
fd*: SocketHandle fd: SocketHandle
closed*: bool ## determines whether this socket has been closed closed: bool ## determines whether this socket has been closed
case isBuffered*: bool ## determines whether this socket is buffered. case isBuffered: bool ## determines whether this socket is buffered.
of true: of true:
buffer*: array[0..BufferSize, char] buffer: array[0..BufferSize, char]
currPos*: int # current index in buffer currPos: int # current index in buffer
bufLen*: int # current length of buffer bufLen: int # current length of buffer
of false: nil of false: nil
case isSsl: bool case isSsl: bool
of true: of true:
@ -91,7 +91,8 @@ type
# TODO: Save AF, domain etc info and reuse it in procs which need it like connect. # TODO: Save AF, domain etc info and reuse it in procs which need it like connect.
proc newSocket(fd: TAsyncFD, isBuff: bool): AsyncSocket = proc newAsyncSocket*(fd: TAsyncFD, isBuff: bool): AsyncSocket =
## Creates a new ``AsyncSocket`` based on the supplied params.
assert fd != osInvalidSocket.TAsyncFD assert fd != osInvalidSocket.TAsyncFD
new(result) new(result)
result.fd = fd.SocketHandle result.fd = fd.SocketHandle
@ -102,11 +103,17 @@ proc newSocket(fd: TAsyncFD, isBuff: bool): AsyncSocket =
proc newAsyncSocket*(domain: Domain = AF_INET, typ: SockType = SOCK_STREAM, proc newAsyncSocket*(domain: Domain = AF_INET, typ: SockType = SOCK_STREAM,
protocol: Protocol = IPPROTO_TCP, buffered = true): AsyncSocket = protocol: Protocol = IPPROTO_TCP, buffered = true): AsyncSocket =
## Creates a new asynchronous socket. ## Creates a new asynchronous socket.
result = newSocket(newAsyncRawSocket(domain, typ, protocol), buffered) ##
## This procedure will also create a brand new file descriptor for
## this socket.
result = newAsyncSocket(newAsyncRawSocket(domain, typ, protocol), buffered)
proc newAsyncSocket*(domain, typ, protocol: cint, buffered = true): AsyncSocket = proc newAsyncSocket*(domain, typ, protocol: cint, buffered = true): AsyncSocket =
## Creates a new asynchronous socket. ## Creates a new asynchronous socket.
result = newSocket(newAsyncRawSocket(domain, typ, protocol), buffered) ##
## This procedure will also create a brand new file descriptor for
## this socket.
result = newAsyncSocket(newAsyncRawSocket(domain, typ, protocol), buffered)
when defined(ssl): when defined(ssl):
proc getSslError(handle: SslPtr, err: cint): cint = proc getSslError(handle: SslPtr, err: cint): cint =
@ -275,7 +282,7 @@ proc acceptAddr*(socket: AsyncSocket, flags = {SocketFlag.SafeDisconn}):
retFuture.fail(future.readError) retFuture.fail(future.readError)
else: else:
let resultTup = (future.read.address, let resultTup = (future.read.address,
newSocket(future.read.client, socket.isBuffered)) newAsyncSocket(future.read.client, socket.isBuffered))
retFuture.complete(resultTup) retFuture.complete(resultTup)
return retFuture return retFuture
@ -439,6 +446,14 @@ proc setSockOpt*(socket: AsyncSocket, opt: SOBool, value: bool,
var valuei = cint(if value: 1 else: 0) var valuei = cint(if value: 1 else: 0)
setSockOptInt(socket.fd, cint(level), toCInt(opt), valuei) setSockOptInt(socket.fd, cint(level), toCInt(opt), valuei)
proc isSsl*(socket: AsyncSocket): bool =
## Determines whether ``socket`` is a SSL socket.
socket.isSsl
proc getFd*(socket: AsyncSocket): SocketHandle =
## Returns the socket's file descriptor.
return socket.fd
when isMainModule: when isMainModule:
type type
TestCases = enum TestCases = enum

View file

@ -44,21 +44,21 @@ const
type type
SocketImpl* = object ## socket type SocketImpl* = object ## socket type
fd*: SocketHandle fd: SocketHandle
case isBuffered*: bool # determines whether this socket is buffered. case isBuffered: bool # determines whether this socket is buffered.
of true: of true:
buffer*: array[0..BufferSize, char] buffer: array[0..BufferSize, char]
currPos*: int # current index in buffer currPos: int # current index in buffer
bufLen*: int # current length of buffer bufLen: int # current length of buffer
of false: nil of false: nil
when defined(ssl): when defined(ssl):
case isSsl*: bool case isSsl: bool
of true: of true:
sslHandle*: SSLPtr sslHandle: SSLPtr
sslContext*: SSLContext sslContext: SSLContext
sslNoHandshake*: bool # True if needs handshake. sslNoHandshake: bool # True if needs handshake.
sslHasPeekChar*: bool sslHasPeekChar: bool
sslPeekChar*: char sslPeekChar: char
of false: nil of false: nil
Socket* = ref SocketImpl Socket* = ref SocketImpl
@ -100,7 +100,8 @@ proc toOSFlags*(socketFlags: set[SocketFlag]): cint =
result = result or MSG_PEEK result = result or MSG_PEEK
of SocketFlag.SafeDisconn: continue of SocketFlag.SafeDisconn: continue
proc createSocket(fd: SocketHandle, isBuff: bool): Socket = proc newSocket(fd: SocketHandle, isBuff: bool): Socket =
## Creates a new socket as specified by the params.
assert fd != osInvalidSocket assert fd != osInvalidSocket
new(result) new(result)
result.fd = fd result.fd = fd
@ -115,7 +116,7 @@ proc newSocket*(domain, typ, protocol: cint, buffered = true): Socket =
let fd = newRawSocket(domain, typ, protocol) let fd = newRawSocket(domain, typ, protocol)
if fd == osInvalidSocket: if fd == osInvalidSocket:
raiseOSError(osLastError()) raiseOSError(osLastError())
result = createSocket(fd, buffered) result = newSocket(fd, buffered)
proc newSocket*(domain: Domain = AF_INET, typ: SockType = SOCK_STREAM, proc newSocket*(domain: Domain = AF_INET, typ: SockType = SOCK_STREAM,
protocol: Protocol = IPPROTO_TCP, buffered = true): Socket = protocol: Protocol = IPPROTO_TCP, buffered = true): Socket =
@ -125,7 +126,7 @@ proc newSocket*(domain: Domain = AF_INET, typ: SockType = SOCK_STREAM,
let fd = newRawSocket(domain, typ, protocol) let fd = newRawSocket(domain, typ, protocol)
if fd == osInvalidSocket: if fd == osInvalidSocket:
raiseOSError(osLastError()) raiseOSError(osLastError())
result = createSocket(fd, buffered) result = newSocket(fd, buffered)
when defined(ssl): when defined(ssl):
CRYPTO_malloc_init() CRYPTO_malloc_init()
@ -937,10 +938,10 @@ proc connect*(socket: Socket, address: string, port = Port(0), timeout: int,
doAssert socket.handshake() doAssert socket.handshake()
socket.fd.setBlocking(true) socket.fd.setBlocking(true)
proc isSSL*(socket: Socket): bool = return socket.isSSL proc isSsl*(socket: Socket): bool = return socket.isSSL
## Determines whether ``socket`` is a SSL socket. ## Determines whether ``socket`` is a SSL socket.
proc getFD*(socket: Socket): SocketHandle = return socket.fd proc getFd*(socket: Socket): SocketHandle = return socket.fd
## Returns the socket's file descriptor ## Returns the socket's file descriptor
type type