#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <unistd.h>
#include <fcntl.h>
#include <sys/types.h>
#include <sys/socket.h>
#include <sys/ioctl.h>
#include <netinet/in.h>
#include <linux/if_packet.h>
#include <linux/if.h>

#include <errno.h>

#include "raw_sock.h"

/****************************************************************************
 * Function   : 	set_socket_nonblocking
 * Description: 	将套接字设置为非阻塞模式

 * Input      : 	socket    : 套接字文件描述符
 				
 * Output     : 	N/A
 * Return     :	成功返回OK，出错返回ERROR
****************************************************************************/
STATUS set_socket_nonblocking(int socket)
{
	int val;

	val = fcntl(socket, F_GETFL);
	val |= O_NONBLOCK;

	if (fcntl(socket, F_SETFL, val) < 0)
	{
		perror("fcntl");
		return ERROR;
	}

	return OK;
}


/****************************************************************************
 * Function   : 	set_socket_timeout_val
 * Description: 	将套接字设置为非阻塞模式

 * Input      : 	socket : 套接字文件描述符
 				type   : 类型(发送超时或者接收超时)
 				sec    : 超时时间(秒)
 				usec   : 超时时间(微秒)
 				
 * Output     : 	N/A
 * Return     :	成功返回OK，出错返回ERROR
****************************************************************************/
STATUS set_socket_timeout_val(int socket, int type, int sec, int usec)
{
	struct timeval tv;
	int ret;

	if ((socket < 0) || (sec < 0) || (usec < 0) 
		|| (type != SO_RCVTIMEO && type != SO_SNDTIMEO))
	{
		return ERROR;
	}

	tv.tv_sec = sec;
	tv.tv_usec = usec;
	
	ret = setsockopt(socket, SOL_SOCKET, type, &tv, sizeof(struct timeval));
	if (ret < 0)
	{
		DPRINT_ERR("set sock timeout error, socket: %d, type: %s, sec: %d, usec: %d\n", 
			socket, (type == SO_RCVTIMEO) ? "SO_RCVTIMEO" : "SO_SNDTIMEO", sec, usec);
		return ERROR;
	}

	return OK;
}

/****************************************************************************
 * Function   : 	bind_to_device
 * Description: 	将套接字绑定到指定设备(仅用于协议族PF_PACKET)

 * Input      : 	socket   : 套接字文件描述符
				protocol : 协议类型
 				dev      : 设备名称
 				
 * Output     : 	N/A
 * Return     :	成功返回OK，出错返回ERROR
****************************************************************************/
STATUS bind_to_device(int socket, int protocol, char dev[])
{
	struct sockaddr_ll sa;
	struct ifreq ifstruct;

	memset(&sa, 0, sizeof(sa));
	memset(&ifstruct, 0, sizeof(ifstruct));
	
	dstrlcpy(ifstruct.ifr_name, dev, sizeof(ifstruct.ifr_name));
	if (ioctl(socket, SIOCGIFINDEX, &ifstruct) < 0)
	{
		perror("ioctl SIOCGIFINDEX");
		return ERROR;
	}

	sa.sll_family = PF_PACKET;
	sa.sll_protocol = htons(protocol);
	sa.sll_ifindex = ifstruct.ifr_ifindex;

	/* 将socket与设备绑定 */
	if (-1 == bind(socket, (struct sockaddr  *)&sa, sizeof(struct sockaddr_ll)))
	{
		perror("bind");
		return ERROR;
	}
	
	return OK;
}


/****************************************************************************
 * Function   : 	raw_socket_init
 * Description: 	创建原始套接字，并设置为非阻塞模式，并绑定到指定设备
 				仅用于协议族PF_PACKET

 * Input      : 	protocol : 协议类型
 				dev      : 设备名称 
 				
 * Output     : 	N/A
 * Return     :	成功返回socket文件描述符，出错返回ERROR
****************************************************************************/
int raw_socket_init(int protocol, char *dev)
{
	int rawsock;

	int ret;
	
	/* 创建原始套接字 */
	if ((rawsock = socket(PF_PACKET, SOCK_RAW, htons(protocol))) < 0)
	{
		perror("socket");
		return ERROR;
	}

	/* 将套接字与设备绑定 */
	ret = bind_to_device(rawsock, protocol, dev);
	if (ERROR == ret)
	{
		DPRINT_ERR("bind to device failed, dev:%s\n",dev);
		
		close(rawsock);
		return ERROR;
	}
	
	return rawsock;
}


/****************************************************************************
 * Function   : 	raw_data_send
 * Description: 	用原始套接字将数据发送出去
 				约定数据长度不得大于1500字节，不得小于14字节

 * Input      : 	rawsock : 套接字
 				rawdata : 数据缓存结构指针
 				
 * Output     : 	N/A
 * Return     :	成功返回OK，出错返回ERROR
****************************************************************************/
STATUS raw_data_send(int rawsock, RAW_DATA *rawdata)
{
	int ret;
	
	if (rawdata->len < RAW_PACKET_LEN_MIN || rawdata->len > RAW_PACKET_LEN_MAX)
	{
		DPRINT_ERR("invalid data len : %d, min = %d, max = %d\n", 
			rawdata->len, RAW_PACKET_LEN_MIN, RAW_PACKET_LEN_MAX);
		return ERROR;
	}
	
	ret = send(rawsock, rawdata->data, rawdata->len, 0);
	if (ret != rawdata->len)
	{
		DPRINT_ERR("raw data send failed, ret = %d\n", ret);
		return ERROR;
	}
	
	return OK;
}


/****************************************************************************
 * Function   : 	raw_data_recv
 * Description: 	从指定套接字中读取数据.由于待读取数据长度未知，所以缓存
 				空间由此函数分配，由调用者释放

 * Input      : 	rawsock : 套接字
 				rawdata : 数据缓存结构指针
 				
 * Output     : 	N/A
 * Return     :	成功返回OK，出错返回ERROR
****************************************************************************/
STATUS raw_data_recv(int rawsock, RAW_DATA *rawdata)
{
	void *start, *tmp;
	
	int buf_len, ret;
	
	if (!rawdata)
	{
		DPRINT_ERR("invalid argument : rawdata = %p\n", rawdata);
		return ERROR;
	}

	buf_len = RAW_PACKET_LEN_MAX;
	
	rawdata->data = dcalloc(buf_len, sizeof(char));
	if (!rawdata->data)
	{
		perror("malloc");
		return ERROR;
	}

	rawdata->len = 0;
	start = rawdata->data;
	
	while (1)
	{
		ret = recv(rawsock, start, RAW_PACKET_LEN_MAX, 0);

		if (ret < 0)
		{
			/* 所有数据已接收完成 */
			if (errno == EAGAIN)
			{
				break;
			}
			else
			{
				perror ("recv");
				dfree(rawdata->data);
				return ERROR;
			}
		}
		else if (0 == ret)
		{
			DPRINT_ERR("socket %d disconnected\n", rawsock);
			dfree(rawdata->data);
			return ERROR;
		}

		rawdata->len += ret;

		/* 此时所有数据同样读取完成 */
		if (ret < RAW_PACKET_LEN_MAX)
		{
			break;
		}
		/* 还有数据未读取，继续读数据 */
		else
		{
			DPRINT_INFO("recv data length is more than %d\n ", RAW_PACKET_LEN_MAX);
			
			buf_len += RAW_PACKET_LEN_MAX;
			
			tmp = realloc(rawdata->data, buf_len);
			if (!tmp)
			{
				perror("realloc");
				dfree(rawdata->data);
				return ERROR;
			}

			rawdata->data = tmp;
			start = rawdata->data + buf_len - RAW_PACKET_LEN_MAX;
		}
	}
	
	return OK;
}
 

 
