Added experimental version of a Socket API with SSL support.

* Each class has a Socket() method to retrieve the underlaying file descriptor
  to be able to do the more advanced stuff, if necessary.
* A server socket is yet missing, but the rest is pretty much covered.
This commit is contained in:
Axel Dörfler
2011-11-21 22:07:52 +01:00
parent db528c0065
commit 0e478f5aec
9 changed files with 956 additions and 0 deletions
+66
View File
@@ -0,0 +1,66 @@
/*
* Copyright 2011, Haiku, Inc. All Rights Reserved.
* Distributed under the terms of the MIT License.
*/
#ifndef _ABSTRACT_SOCKET_H
#define _ABSTRACT_SOCKET_H
#include <DataIO.h>
#include <NetworkAddress.h>
#include <sys/socket.h>
class BAbstractSocket : public BDataIO {
public:
BAbstractSocket();
BAbstractSocket(const BAbstractSocket& other);
virtual ~BAbstractSocket();
status_t InitCheck() const;
virtual status_t Bind(const BNetworkAddress& local) = 0;
virtual bool IsBound() const;
virtual status_t Connect(const BNetworkAddress& peer,
bigtime_t timeout = B_INFINITE_TIMEOUT) = 0;
virtual bool IsConnected() const;
virtual void Disconnect();
virtual status_t SetTimeout(bigtime_t timeout);
virtual bigtime_t Timeout() const;
virtual const BNetworkAddress& Local() const;
virtual const BNetworkAddress& Peer() const;
virtual size_t MaxTransmissionSize() const;
virtual status_t WaitForReadable(bigtime_t timeout
= B_INFINITE_TIMEOUT) const;
virtual status_t WaitForWritable(bigtime_t timeout
= B_INFINITE_TIMEOUT) const;
int Socket() const;
protected:
status_t Bind(const BNetworkAddress& local, int type);
status_t Connect(const BNetworkAddress& peer, int type,
bigtime_t timeout = B_INFINITE_TIMEOUT);
private:
status_t _OpenIfNeeded(int family, int type);
status_t _UpdateLocalAddress();
status_t _WaitFor(int flags, bigtime_t timeout) const;
protected:
status_t fInitStatus;
int fSocket;
BNetworkAddress fLocal;
BNetworkAddress fPeer;
bool fIsBound;
bool fIsConnected;
};
#endif // _ABSTRACT_SOCKET_H
+41
View File
@@ -0,0 +1,41 @@
/*
* Copyright 2011, Haiku, Inc. All Rights Reserved.
* Distributed under the terms of the MIT License.
*/
#ifndef _DATAGRAM_SOCKET_H
#define _DATAGRAM_SOCKET_H
#include <AbstractSocket.h>
class BDatagramSocket : public BAbstractSocket {
public:
BDatagramSocket();
BDatagramSocket(const BNetworkAddress& peer,
bigtime_t timeout = -1);
BDatagramSocket(const BDatagramSocket& other);
virtual ~BDatagramSocket();
virtual status_t Bind(const BNetworkAddress& peer);
virtual status_t Connect(const BNetworkAddress& peer,
bigtime_t timeout = B_INFINITE_TIMEOUT);
status_t SetBroadcast(bool broadcast);
void SetPeer(const BNetworkAddress& peer);
virtual size_t MaxTransmissionSize() const;
virtual size_t SendTo(const BNetworkAddress& address,
const void* buffer, size_t size);
virtual size_t ReceiveFrom(void* buffer, size_t bufferSize,
BNetworkAddress& from);
// BDataIO implementation
virtual ssize_t Read(void* buffer, size_t size);
virtual ssize_t Write(const void* buffer, size_t size);
};
#endif // _DATAGRAM_SOCKET_H
+38
View File
@@ -0,0 +1,38 @@
/*
* Copyright 2011, Haiku, Inc. All Rights Reserved.
* Distributed under the terms of the MIT License.
*/
#ifndef _SECURE_SOCKET_H
#define _SECURE_SOCKET_H
#include <Socket.h>
class BSecureSocket : public BSocket {
public:
BSecureSocket();
BSecureSocket(const BNetworkAddress& peer,
bigtime_t timeout = B_INFINITE_TIMEOUT);
BSecureSocket(const BSecureSocket& other);
virtual ~BSecureSocket();
virtual status_t Connect(const BNetworkAddress& peer,
bigtime_t timeout = B_INFINITE_TIMEOUT);
virtual void Disconnect();
virtual status_t WaitForReadable(bigtime_t timeout
= B_INFINITE_TIMEOUT) const;
// BDataIO implementation
virtual ssize_t Read(void* buffer, size_t size);
virtual ssize_t Write(const void* buffer, size_t size);
private:
class Private;
Private* fPrivate;
};
#endif // _SECURE_SOCKET_H
+37
View File
@@ -0,0 +1,37 @@
/*
* Copyright 2011, Haiku, Inc. All Rights Reserved.
* Distributed under the terms of the MIT License.
*/
#ifndef _SOCKET_H
#define _SOCKET_H
#include <AbstractSocket.h>
class BSocket : public BAbstractSocket {
public:
BSocket();
BSocket(const BNetworkAddress& peer,
bigtime_t timeout = B_INFINITE_TIMEOUT);
BSocket(const BSocket& other);
virtual ~BSocket();
virtual status_t Bind(const BNetworkAddress& peer);
virtual status_t Connect(const BNetworkAddress& peer,
bigtime_t timeout = B_INFINITE_TIMEOUT);
// BDataIO implementation
virtual ssize_t Read(void* buffer, size_t size);
virtual ssize_t Write(const void* buffer, size_t size);
private:
friend class BServerSocket;
void _SetTo(int fd, const BNetworkAddress& local,
const BNetworkAddress& peer);
};
#endif // _SOCKET_H
@@ -0,0 +1,268 @@
/*
* Copyright 2011, Axel Dörfler, [email protected].
* Distributed under the terms of the MIT License.
*/
#include <AbstractSocket.h>
#include <arpa/inet.h>
#include <fcntl.h>
#include <netinet/in.h>
#include <sys/poll.h>
//#define TRACE_SOCKET
#ifdef TRACE_SOCKET
# define TRACE(x...) printf(x)
#else
# define TRACE(x...) ;
#endif
BAbstractSocket::BAbstractSocket()
:
fInitStatus(B_NO_INIT),
fSocket(-1),
fIsBound(false),
fIsConnected(false)
{
}
BAbstractSocket::BAbstractSocket(const BAbstractSocket& other)
:
fInitStatus(other.fInitStatus),
fLocal(other.fLocal),
fPeer(other.fPeer),
fIsConnected(other.fIsConnected)
{
fSocket = dup(other.fSocket);
if (fSocket < 0)
fInitStatus = errno;
}
BAbstractSocket::~BAbstractSocket()
{
Disconnect();
}
status_t
BAbstractSocket::InitCheck() const
{
return fInitStatus;
}
bool
BAbstractSocket::IsBound() const
{
return fIsBound;
}
bool
BAbstractSocket::IsConnected() const
{
return fIsConnected;
}
void
BAbstractSocket::Disconnect()
{
if (fSocket < 0)
return;
TRACE("%p: BAbstractSocket::Disconnect()\n", this);
close(fSocket);
fSocket = -1;
fIsConnected = false;
fIsBound = false;
}
status_t
BAbstractSocket::SetTimeout(bigtime_t timeout)
{
if (timeout < 0)
timeout = 0;
struct timeval tv;
tv.tv_sec = timeout / 1000000LL;
tv.tv_usec = timeout % 1000000LL;
if (setsockopt(fSocket, SOL_SOCKET, SO_SNDTIMEO, &tv, sizeof(timeval)) != 0
|| setsockopt(fSocket, SOL_SOCKET, SO_RCVTIMEO, &tv,
sizeof(timeval)) != 0)
return errno;
return B_OK;
}
bigtime_t
BAbstractSocket::Timeout() const
{
struct timeval tv;
socklen_t size = sizeof(tv);
if (getsockopt(fSocket, SOL_SOCKET, SO_SNDTIMEO, &tv, &size) != 0)
return B_INFINITE_TIMEOUT;
return tv.tv_sec * 1000000LL + tv.tv_usec;
}
const BNetworkAddress&
BAbstractSocket::Local() const
{
return fLocal;
}
const BNetworkAddress&
BAbstractSocket::Peer() const
{
return fPeer;
}
size_t
BAbstractSocket::MaxTransmissionSize() const
{
return SSIZE_MAX;
}
status_t
BAbstractSocket::WaitForReadable(bigtime_t timeout) const
{
return _WaitFor(POLLIN, timeout);
}
status_t
BAbstractSocket::WaitForWritable(bigtime_t timeout) const
{
return _WaitFor(POLLOUT, timeout);
}
int
BAbstractSocket::Socket() const
{
return fSocket;
}
// #pragma mark - protected
status_t
BAbstractSocket::Bind(const BNetworkAddress& local, int type)
{
fInitStatus = _OpenIfNeeded(local.Family(), type);
if (fInitStatus != B_OK)
return fInitStatus;
if (bind(fSocket, local, local.Length()) != 0)
return fInitStatus = errno;
fIsBound = true;
_UpdateLocalAddress();
return B_OK;
}
status_t
BAbstractSocket::Connect(const BNetworkAddress& peer, int type,
bigtime_t timeout)
{
Disconnect();
fInitStatus = _OpenIfNeeded(peer.Family(), type);
if (fInitStatus == B_OK)
fInitStatus = SetTimeout(timeout);
if (fInitStatus == B_OK && !IsBound()) {
BNetworkAddress local;
local.SetToWildcard(peer.Family());
fInitStatus = Bind(local);
}
if (fInitStatus != B_OK)
return fInitStatus;
BNetworkAddress normalized = peer;
if (connect(fSocket, normalized, normalized.Length()) != 0) {
TRACE("%p: connecting to %s: %s\n", this,
normalized.ToString().c_str(), strerror(errno));
return fInitStatus = errno;
}
fIsConnected = true;
fPeer = normalized;
_UpdateLocalAddress();
TRACE("%p: connected to %s (local %s)\n", this, peer.ToString().c_str(),
fLocal.ToString().c_str());
return fInitStatus = B_OK;
}
// #pragma mark - private
status_t
BAbstractSocket::_OpenIfNeeded(int family, int type)
{
if (fSocket >= 0)
return B_OK;
fSocket = socket(family, type, 0);
if (fSocket < 0)
return errno;
TRACE("%p: socket opened FD %d\n", this, fSocket);
return B_OK;
}
status_t
BAbstractSocket::_UpdateLocalAddress()
{
socklen_t localLength = sizeof(sockaddr_storage);
if (getsockname(fSocket, fLocal, &localLength) != 0)
return errno;
return B_OK;
}
status_t
BAbstractSocket::_WaitFor(int flags, bigtime_t timeout) const
{
if (fInitStatus != B_OK)
return fInitStatus;
int millis = 0;
if (timeout == B_INFINITE_TIMEOUT)
millis = -1;
if (timeout > 0)
millis = timeout / 1000;
struct pollfd entry;
entry.fd = Socket();
entry.events = flags;
int result = poll(&entry, 1, -1);
if (result < 0)
return errno;
if (result == 0)
return millis > 0 ? B_TIMED_OUT : B_WOULD_BLOCK;
return B_OK;
}
@@ -0,0 +1,142 @@
/*
* Copyright 2011, Axel Dörfler, [email protected].
* Distributed under the terms of the MIT License.
*/
#include <DatagramSocket.h>
//#define TRACE_SOCKET
#ifdef TRACE_SOCKET
# define TRACE(x...) printf(x)
#else
# define TRACE(x...) ;
#endif
BDatagramSocket::BDatagramSocket()
{
}
BDatagramSocket::BDatagramSocket(const BNetworkAddress& peer, bigtime_t timeout)
{
Connect(peer, timeout);
}
BDatagramSocket::BDatagramSocket(const BDatagramSocket& other)
:
BAbstractSocket(other)
{
}
BDatagramSocket::~BDatagramSocket()
{
}
status_t
BDatagramSocket::Bind(const BNetworkAddress& local)
{
return BAbstractSocket::Bind(local, SOCK_DGRAM);
}
status_t
BDatagramSocket::Connect(const BNetworkAddress& peer, bigtime_t timeout)
{
return BAbstractSocket::Connect(peer, SOCK_DGRAM, timeout);
}
status_t
BDatagramSocket::SetBroadcast(bool broadcast)
{
int value = broadcast ? 1 : 0;
if (setsockopt(fSocket, SOL_SOCKET, SO_BROADCAST, &value, sizeof(value))
!= 0)
return errno;
return B_OK;
}
void
BDatagramSocket::SetPeer(const BNetworkAddress& peer)
{
fPeer = peer;
}
size_t
BDatagramSocket::MaxTransmissionSize() const
{
// TODO: might vary on family!
return 32768;
}
size_t
BDatagramSocket::SendTo(const BNetworkAddress& address, const void* buffer,
size_t size)
{
ssize_t bytesSent = sendto(fSocket, buffer, size, 0, address,
address.Length());
if (bytesSent < 0)
return errno;
return bytesSent;
}
size_t
BDatagramSocket::ReceiveFrom(void* buffer, size_t bufferSize,
BNetworkAddress& from)
{
socklen_t fromLength = sizeof(sockaddr_storage);
ssize_t bytesReceived = recvfrom(fSocket, buffer, bufferSize, 0,
from, &fromLength);
if (bytesReceived < 0)
return errno;
return bytesReceived;
}
// #pragma mark - BDataIO implementation
ssize_t
BDatagramSocket::Read(void* buffer, size_t size)
{
ssize_t bytesReceived = recv(Socket(), buffer, size, 0);
if (bytesReceived < 0) {
TRACE("%p: BSocket::Read() error: %s\n", this, strerror(errno));
return errno;
}
return bytesReceived;
}
ssize_t
BDatagramSocket::Write(const void* buffer, size_t size)
{
ssize_t bytesSent;
if (!fIsConnected)
bytesSent = sendto(Socket(), buffer, size, 0, fPeer, fPeer.Length());
else
bytesSent = send(Socket(), buffer, size, 0);
if (bytesSent < 0) {
TRACE("%p: BDatagramSocket::Write() error: %s\n", this,
strerror(errno));
return errno;
}
return bytesSent;
}
+31
View File
@@ -0,0 +1,31 @@
/*
* Copyright 2011, Axel Dörfler, [email protected].
* Distributed under the terms of the MIT License.
*/
#include <OS.h>
#include <openssl/ssl.h>
#include <openssl/rand.h>
namespace BPrivate {
class SSL {
public:
SSL()
{
SSL_library_init();
int64 seed = find_thread(NULL) ^ system_time();
RAND_seed(&seed, sizeof(seed));
}
};
static SSL sSSL;
} // namespace BPrivate
+233
View File
@@ -0,0 +1,233 @@
/*
* Copyright 2011, Axel Dörfler, [email protected].
* Copyright 2010, Clemens Zeidler <[email protected]>
* Distributed under the terms of the MIT License.
*/
#include <SecureSocket.h>
#ifdef OPENSSL_ENABLED
# include <openssl/ssl.h>
#endif
//#define TRACE_SOCKET
#ifdef TRACE_SOCKET
# define TRACE(x...) printf(x)
#else
# define TRACE(x...) ;
#endif
#ifdef OPENSSL_ENABLED
class BSecureSocket::Private {
public:
SSL_CTX* fCTX;
SSL* fSSL;
BIO* fBIO;
};
BSecureSocket::BSecureSocket()
:
fPrivate(NULL)
{
}
BSecureSocket::BSecureSocket(const BNetworkAddress& peer, bigtime_t timeout)
:
fPrivate(NULL)
{
Connect(peer, timeout);
}
BSecureSocket::BSecureSocket(const BSecureSocket& other)
:
BSocket(other)
{
// TODO: this won't work this way!
fPrivate = (BSecureSocket::Private*)malloc(sizeof(BSecureSocket::Private));
if (fPrivate != NULL)
memcpy(fPrivate, other.fPrivate, sizeof(BSecureSocket::Private));
else
fInitStatus = B_NO_MEMORY;
}
BSecureSocket::~BSecureSocket()
{
free(fPrivate);
}
status_t
BSecureSocket::Connect(const BNetworkAddress& peer, bigtime_t timeout)
{
if (fPrivate == NULL) {
fPrivate = (BSecureSocket::Private*)calloc(1,
sizeof(BSecureSocket::Private));
if (fPrivate == NULL)
return B_NO_MEMORY;
}
status_t status = BSocket::Connect(peer, timeout);
if (status != B_OK)
return status;
fPrivate->fCTX = SSL_CTX_new(SSLv23_method());
fPrivate->fSSL = SSL_new(fPrivate->fCTX);
fPrivate->fBIO = BIO_new_socket(fSocket, BIO_NOCLOSE);
SSL_set_bio(fPrivate->fSSL, fPrivate->fBIO, fPrivate->fBIO);
if (SSL_connect(fPrivate->fSSL) <= 0) {
TRACE("SSLConnection can't connect\n");
BSocket::Disconnect();
// TODO: translate ssl to Haiku error
return B_ERROR;
}
return B_OK;
}
void
BSecureSocket::Disconnect()
{
if (IsConnected()) {
if (fPrivate->fSSL != NULL) {
SSL_shutdown(fPrivate->fSSL);
fPrivate->fSSL = NULL;
}
if (fPrivate->fCTX != NULL) {
SSL_CTX_free(fPrivate->fCTX);
fPrivate->fCTX = NULL;
}
if (fPrivate->fBIO != NULL) {
BIO_free(fPrivate->fBIO);
fPrivate->fBIO = NULL;
}
}
return BSocket::Disconnect();
}
status_t
BSecureSocket::WaitForReadable(bigtime_t timeout) const
{
if (fInitStatus != B_OK)
return fInitStatus;
if (!IsConnected())
return B_ERROR;
if (SSL_pending(fPrivate->fSSL) > 0)
return B_OK;
return BSocket::WaitForReadable(timeout);
}
// #pragma mark - BDataIO implementation
ssize_t
BSecureSocket::Read(void* buffer, size_t size)
{
if (!IsConnected())
return B_ERROR;
int bytesRead = SSL_read(fPrivate->fSSL, buffer, size);
if (bytesRead > 0)
return bytesRead;
// TODO: translate SSL error codes!
return B_ERROR;
}
ssize_t
BSecureSocket::Write(const void* buffer, size_t size)
{
if (!IsConnected())
return B_ERROR;
int bytesWritten = SSL_write(fPrivate->fSSL, buffer, size);
if (bytesWritten > 0)
return bytesWritten;
// TODO: translate SSL error codes!
return B_ERROR;
}
#else // OPENSSL_ENABLED
// #pragma mark - No-SSL stubs
BSecureSocket::BSecureSocket()
{
}
BSecureSocket::BSecureSocket(const BNetworkAddress& peer, bigtime_t timeout)
{
fInitStatus = B_UNSUPPORTED;
}
BSecureSocket::BSecureSocket(const BSecureSocket& other)
:
BSocket(other)
{
}
BSecureSocket::~BSecureSocket()
{
}
status_t
BSecureSocket::Connect(const BNetworkAddress& peer, bigtime_t timeout)
{
return fInitStatus = B_UNSUPPORTED;
}
void
BSecureSocket::Disconnect()
{
}
status_t
BSecureSocket::WaitForReadable(bigtime_t timeout) const
{
return B_UNSUPPORTED;
}
// #pragma mark - BDataIO implementation
ssize_t
BSecureSocket::Read(void* buffer, size_t size)
{
return B_UNSUPPORTED;
}
ssize_t
BSecureSocket::Write(const void* buffer, size_t size)
{
return B_UNSUPPORTED;
}
#endif // !OPENSSL_ENABLED
+100
View File
@@ -0,0 +1,100 @@
/*
* Copyright 2011, Axel Dörfler, axeld@pinc-software.de.
* Distributed under the terms of the MIT License.
*/
#include <Socket.h>
//#define TRACE_SOCKET
#ifdef TRACE_SOCKET
# define TRACE(x...) printf(x)
#else
# define TRACE(x...) ;
#endif
BSocket::BSocket()
{
}
BSocket::BSocket(const BNetworkAddress& peer, bigtime_t timeout)
{
Connect(peer, timeout);
}
BSocket::BSocket(const BSocket& other)
:
BAbstractSocket(other)
{
}
BSocket::~BSocket()
{
}
status_t
BSocket::Bind(const BNetworkAddress& local)
{
return BAbstractSocket::Bind(local, SOCK_STREAM);
}
status_t
BSocket::Connect(const BNetworkAddress& peer, bigtime_t timeout)
{
return BAbstractSocket::Connect(peer, SOCK_STREAM, timeout);
}
// #pragma mark - BDataIO implementation
ssize_t
BSocket::Read(void* buffer, size_t size)
{
ssize_t bytesReceived = recv(Socket(), buffer, size, 0);
if (bytesReceived < 0) {
TRACE("%p: BSocket::Read() error: %s\n", this, strerror(errno));
return errno;
}
return bytesReceived;
}
ssize_t
BSocket::Write(const void* buffer, size_t size)
{
ssize_t bytesSent = send(Socket(), buffer, size, 0);
if (bytesSent < 0) {
TRACE("%p: BSocket::Write() error: %s\n", this, strerror(errno));
return errno;
}
return bytesSent;
}
// #pragma mark - private
void
BSocket::_SetTo(int fd, const BNetworkAddress& local,
const BNetworkAddress& peer)
{
Disconnect();
fInitStatus = B_OK;
fSocket = fd;
fLocal = local;
fPeer = peer;
TRACE("%p: accepted from %s to %s\n", this, local.ToString().c_str(),
peer.ToString().c_str());
}