#include "socket.h"
#include <sys/types.h>       // For data types
#include <sys/socket.h>      // For socket(), connect(), send(), and recv()
#include <netdb.h>           // For gethostbyname()
#include <arpa/inet.h>       // For inet_addr()
#include <unistd.h>          // For close()
#include <netinet/in.h>      // For sockaddr_in
typedef void raw_type;       // Type used for raw data on this platform

#include <error.h>
#include <cstring>
#include <cerrno>
#include <cstdlib>

using namespace std;

// SocketException Code

SocketException::SocketException(const string &message, bool inclSysMsg)
  throw() : userMessage(message)
{
	if (inclSysMsg) {
		userMessage.append(": ");
		userMessage.append(strerror(errno));
	}
}

SocketException::~SocketException() throw()
{
}

const char *SocketException::what() const throw()
{
	return userMessage.c_str();
}

// Function to fill in address structure given an address and port
static void fillAddr(const string &address, uint16_t port,
                     sockaddr_in &addr)
{
	memset(&addr, 0, sizeof(addr));  // Zero out address structure
	addr.sin_family = AF_INET;       // Internet address
	hostent *host;  // Resolve name
	if ((host = gethostbyname(address.c_str())) == NULL) {
		// strerror() will not work for gethostbyname() and hstrerror()
		// is supposedly obsolete
		throw SocketException("Failed to resolve name (gethostbyname())");
	}
	addr.sin_addr.s_addr = *((unsigned long *) host->h_addr_list[0]);
	addr.sin_port = htons(port);     // Assign port in network byte order
}

// Socket Code

Socket::Socket(int type, int protocol) throw(SocketException)
{
	// Make a new socket
	if ((sockDesc = socket(PF_INET, type, protocol)) < 0) {
		throw SocketException("Socket creation failed (socket())", true);
	}
}

Socket::Socket(int sockDesc)
{
	this->sockDesc = sockDesc;
}

Socket::~Socket()
{
	::close(sockDesc);
	sockDesc = -1;
}

string Socket::getLocalAddress() throw(SocketException)
{
	sockaddr_in addr;
	unsigned int addr_len = sizeof(addr);

	if (getsockname(sockDesc, (sockaddr *) &addr, (socklen_t *) &addr_len) < 0) {
		throw SocketException("Fetch of local address failed (getsockname())", true);
	}
	return inet_ntoa(addr.sin_addr);
}

uint16_t Socket::getLocalPort() throw(SocketException)
{
	sockaddr_in addr;
	unsigned int addr_len = sizeof(addr);

	if (getsockname(sockDesc, (sockaddr *) &addr, (socklen_t *) &addr_len) < 0) {
		throw SocketException("Fetch of local port failed (getsockname())", true);
	}
	return ntohs(addr.sin_port);
}

void Socket::setLocalPort(uint16_t localPort) throw(SocketException)
{
	// Bind the socket to its port
	sockaddr_in localAddr;
	memset(&localAddr, 0, sizeof(localAddr));
	localAddr.sin_family = AF_INET;
	localAddr.sin_addr.s_addr = htonl(INADDR_ANY);
	localAddr.sin_port = htons(localPort);

	if (bind(sockDesc, (sockaddr *) &localAddr, sizeof(sockaddr_in)) < 0) {
		throw SocketException("Set of local port failed (bind())", true);
	}
}

void Socket::setLocalAddressAndPort(const string &localAddress,
    uint16_t localPort) throw(SocketException)
{
	// Get the address of the requested host
	sockaddr_in localAddr;
	fillAddr(localAddress, localPort, localAddr);

	if (bind(sockDesc, (sockaddr *) &localAddr, sizeof(sockaddr_in)) < 0) {
		throw SocketException("Set of local address and port failed (bind())", true);
	}
}

void Socket::cleanUp() throw(SocketException)
{
}

uint16_t Socket::resolveService(const string &service,
                                      const string &protocol)
{
	struct servent *serv;        /* Structure containing service information */

	if ((serv = getservbyname(service.c_str(), protocol.c_str())) == NULL)
		return atoi(service.c_str());  /* Service is port number */
	else
		return ntohs(serv->s_port);    /* Found port (network byte order) by name */
}

// CommunicatingSocket Code

CommunicatingSocket::CommunicatingSocket(int32_t type, int32_t protocol)
    throw(SocketException) : Socket(type, protocol)
{
}

CommunicatingSocket::CommunicatingSocket(int32_t newConnSD) : Socket(newConnSD)
{
}

void CommunicatingSocket::connect(const string &foreignAddress,
    uint16_t foreignPort) throw(SocketException)
{
	// Get the address of the requested host
	sockaddr_in destAddr;
	fillAddr(foreignAddress, foreignPort, destAddr);

	// Try to connect to the given port
	if (::connect(sockDesc, (sockaddr *) &destAddr, sizeof(destAddr)) < 0) {
		throw SocketException("Connect failed (connect())", true);
	}
}

void CommunicatingSocket::send(const void *buffer, uint32_t bufferLen)
    throw(SocketException)
{
	if (::send(sockDesc, (raw_type *) buffer, bufferLen, 0) < 0) {
		throw SocketException("Send failed (send())", true);
	}
}

int32_t CommunicatingSocket::recv(void *buffer, uint32_t bufferLen)
    throw(SocketException)
{
	int32_t rtn;

	if ((rtn = ::recv(sockDesc, (raw_type *) buffer, bufferLen, 0)) < 0) {
		throw SocketException("Received failed (recv())", true);
	}
	return rtn;
}

string CommunicatingSocket::getForeignAddress()
    throw(SocketException)
{
	sockaddr_in addr;
	uint32_t addr_len = sizeof(addr);

	if (getpeername(sockDesc, (sockaddr *) &addr,(socklen_t *) &addr_len) < 0) {
		throw SocketException("Fetch of foreign address failed (getpeername())", true);
	}
	return inet_ntoa(addr.sin_addr);
}

uint16_t CommunicatingSocket::getForeignPort() throw(SocketException)
{
	sockaddr_in addr;
	uint32_t addr_len = sizeof(addr);

	if (getpeername(sockDesc, (sockaddr *) &addr, (socklen_t *) &addr_len) < 0) {
		throw SocketException("Fetch of foreign port failed (getpeername())", true);
	}
	return ntohs(addr.sin_port);
}

// TCP_Socket Code

TCP_Socket::TCP_Socket()
    throw(SocketException) : CommunicatingSocket(SOCK_STREAM, IPPROTO_TCP)
{
}

TCP_Socket::TCP_Socket(const string &foreignAddress, uint16_t foreignPort)
    throw(SocketException) : CommunicatingSocket(SOCK_STREAM, IPPROTO_TCP)
{
    connect(foreignAddress, foreignPort);
}

TCP_Socket::TCP_Socket(int32_t newConnSD) : CommunicatingSocket(newConnSD)
{
}

// TCPServerSocket Code
TCP_ServerSocket::TCP_ServerSocket(uint16_t localPort, int32_t queueLen)
    throw(SocketException) : Socket(SOCK_STREAM, IPPROTO_TCP)
{
	setLocalPort(localPort);
	setListen(queueLen);
}

TCP_ServerSocket::TCP_ServerSocket(const string &localAddress,
    uint16_t localPort, int32_t queueLen)
    throw(SocketException) : Socket(SOCK_STREAM, IPPROTO_TCP)
{
	setLocalAddressAndPort(localAddress, localPort);
	setListen(queueLen);
}

TCP_Socket *TCP_ServerSocket::accept() throw(SocketException)
{
	int32_t newConnSD;
	if ((newConnSD = ::accept(sockDesc, NULL, 0)) < 0) {
		throw SocketException("Accept failed (accept())", true);
	}
	return new TCP_Socket(newConnSD);
}

void TCP_ServerSocket::setListen(int32_t queueLen) throw(SocketException)
{
	if (listen(sockDesc, queueLen) < 0) {
		throw SocketException("Set listening socket failed (listen())", true);
	}
}

// UDP_Socket Code
UDP_Socket::UDP_Socket() throw(SocketException) : CommunicatingSocket(SOCK_DGRAM,
    IPPROTO_UDP)
{
	setBroadcast();
}

UDP_Socket::UDP_Socket(uint16_t localPort)  throw(SocketException) :
    CommunicatingSocket(SOCK_DGRAM, IPPROTO_UDP)
{
	setLocalPort(localPort);
	setBroadcast();
}

UDP_Socket::UDP_Socket(const string &localAddress, uint16_t localPort)
     throw(SocketException) : CommunicatingSocket(SOCK_DGRAM, IPPROTO_UDP)
{
	setLocalAddressAndPort(localAddress, localPort);
	setBroadcast();
}

void UDP_Socket::setBroadcast()
{
	// If this fails, we'll hear about it when we try to send.  This will allow
	// system that cannot broadcast to continue if they don't plan to broadcast
	int32_t broadcastPermission = 1;
	setsockopt(sockDesc, SOL_SOCKET, SO_BROADCAST,
			 (raw_type *) &broadcastPermission, sizeof(broadcastPermission));
}

void UDP_Socket::disconnect() throw(SocketException)
{
	sockaddr_in nullAddr;
	memset(&nullAddr, 0, sizeof(nullAddr));
	nullAddr.sin_family = AF_UNSPEC;

	// Try to disconnect
	if (::connect(sockDesc, (sockaddr *) &nullAddr, sizeof(nullAddr)) < 0) {
		if (errno != EAFNOSUPPORT) {
			throw SocketException("Disconnect failed (connect())", true);
		}
	}
}

void UDP_Socket::sendTo(const void *buffer, uint32_t bufferLen,
    const string &foreignAddress, uint16_t foreignPort)
    throw(SocketException)
{
	sockaddr_in destAddr;
	fillAddr(foreignAddress, foreignPort, destAddr);

	// Write out the whole buffer as a single message.
	if (sendto(sockDesc, (raw_type *) buffer, bufferLen, 0,
			 (sockaddr *) &destAddr, sizeof(destAddr)) != bufferLen) {
		throw SocketException("Send failed (sendto())", true);
	}
}

int32_t UDP_Socket::recvFrom(void *buffer, uint32_t bufferLen, string &sourceAddress,
    uint16_t&sourcePort) throw(SocketException)
{
	sockaddr_in clntAddr;
	socklen_t addrLen = sizeof(clntAddr);
	int32_t rtn;
	if ((rtn = recvfrom(sockDesc, (raw_type *) buffer, bufferLen, 0,
					  (sockaddr *) &clntAddr, (socklen_t *) &addrLen)) < 0) {
		throw SocketException("Receive failed (recvfrom())", true);
	}
	sourceAddress = inet_ntoa(clntAddr.sin_addr);
	sourcePort = ntohs(clntAddr.sin_port);

	return rtn;
}

void UDP_Socket::setMulticastTTL(uint8_t* multicastTTL) throw(SocketException)
{
	if (setsockopt(sockDesc, IPPROTO_IP, IP_MULTICAST_TTL,
				 (raw_type *) &multicastTTL, sizeof(multicastTTL)) < 0) {
		throw SocketException("Multicast TTL set failed (setsockopt())", true);
	}
}

void UDP_Socket::joinGroup(const string &multicastGroup) throw(SocketException)
{
	struct ip_mreq multicastRequest;

	multicastRequest.imr_multiaddr.s_addr = inet_addr(multicastGroup.c_str());
	multicastRequest.imr_interface.s_addr = htonl(INADDR_ANY);
	if (setsockopt(sockDesc, IPPROTO_IP, IP_ADD_MEMBERSHIP,
				 (raw_type *) &multicastRequest,
				 sizeof(multicastRequest)) < 0) {
		throw SocketException("Multicast group join failed (setsockopt())", true);
	}
}

void UDP_Socket::leaveGroup(const string &multicastGroup) throw(SocketException)
{
	struct ip_mreq multicastRequest;

	multicastRequest.imr_multiaddr.s_addr = inet_addr(multicastGroup.c_str());
	multicastRequest.imr_interface.s_addr = htonl(INADDR_ANY);
	if (setsockopt(sockDesc, IPPROTO_IP, IP_DROP_MEMBERSHIP,
				 (raw_type *) &multicastRequest,
				 sizeof(multicastRequest)) < 0) {
		throw SocketException("Multicast group leave failed (setsockopt())", true);
	}
}
