#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <unistd.h>
#include <sys/types.h>
#include <sys/socket.h>
#include <netinet/in.h>
#include <arpa/inet.h>
#include <netdb.h>
#include <poll.h>

// Get sockaddr, IPv4 or IPv6:
void *get_in_addr(struct sockaddr *sa)
{
    if (sa->sa_family == AF_INET) {
        return &(((struct sockaddr_in*)sa)->sin_addr);
    }

    return &(((struct sockaddr_in6*)sa)->sin6_addr);
}

// Return a listening socket
int get_listener_socket(int argc, char * argv[])
{
    int listener;     // Listening socket descriptor
    int yes=1;        // For setsockopt() SO_REUSEADDR, below
    int rv;

    struct addrinfo hints, *ai, *p;

    // Get us a socket and bind it
    memset(&hints, 0, sizeof hints);
    hints.ai_family = AF_UNSPEC;
    if (strcmp(argv[3], "UDP") == 0) 
      hints.ai_socktype = SOCK_DGRAM;
    else
      hints.ai_socktype = SOCK_STREAM;
    hints.ai_flags = AI_PASSIVE;
    if ((rv = getaddrinfo(argv[1], argv[2], &hints, &ai)) != 0) {
        fprintf(stderr, "selectserver: %s\n", gai_strerror(rv));
        exit(1);
    }
    
    for(p = ai; p != NULL; p = p->ai_next) {
        listener = socket(p->ai_family, p->ai_socktype, p->ai_protocol);
        if (listener < 0) { 
            continue;
        }
        
        // Lose the pesky "address already in use" error message
        setsockopt(listener, SOL_SOCKET, SO_REUSEADDR, &yes, sizeof(int));

        if (bind(listener, p->ai_addr, p->ai_addrlen) < 0) {
            close(listener);
            continue;
        }

        break;
    }

    freeaddrinfo(ai); // All done with this

    // If we got here, it means we didn't get bound
    if (p == NULL) {
        return -1;
    }

    
    // Listen, if we're TCP...
    if (hints.ai_socktype == SOCK_STREAM ) {
      if (listen(listener, 10) == -1) {
        return -1;
      }
    }

    return listener;
}


// Main
int main(int argc, char *argv[])
{
  if (argc < 4) {
    printf("Usage: ./recieve_file.o hostname port TCP|UDP\n");
    exit(0);
  }
  
    int listener;     // Listening socket descriptor

    int newfd;        // Newly accept()ed socket descriptor
    char buf[256];    // Buffer for client data

    // Set up and get a listening socket (if TCP, otherwise just a socket)
    listener = get_listener_socket(argc, argv); // getaddrinfo, socket, bind, listen

    if (listener == -1) {
        fprintf(stderr, "error getting listening socket\n");
        exit(1);
    }

    if (strcmp(argv[3], "UDP") != 0) {
      newfd = accept(listener,
		   NULL, NULL);
    }
    else {
      newfd = listener;
    }
    
    if (newfd == -1) {
      perror("accept");
    } else {
      ssize_t result;
      long int total_bytes = -1; // -1 means keep going forever (ie TCP)
      long int total_read = 0;
      // if UDP, get the total # of bytes expected
      if (strcmp(argv[3], "UDP") == 0) {
	result = recv(newfd, &total_bytes, sizeof(total_bytes), 0);
	if (result != sizeof(total_bytes)) {
	  fprintf(stderr, "Did not receive enough bytes...\n");
	  exit(0);
	}
      }
      
      // recieve data until EOF
      while (1)  {
	result = recv(newfd, buf, 255, 0);
	if (result > 0) {
	  for(int i=0; i < result; i++)
	    putc(buf[i], stdout);

	  total_read += result;

	  if (total_bytes > 0 && total_read >= total_bytes) break;
	}
	else {
	  break;
	}
      }
    }
    
    return 0;
}