/*
 *------------------------------------------------------------------
 *
 * $Source: /mit/cgw/src/tftp.mar/RCS/tftp.c,v $
 * $Revision: 1.9 $
 * $Date: 89/04/11 15:18:03 $
 * $State: Exp $
 * $Author: mar $
 * $Locker:  $
 *
 * $Log:	tftp.c,v $
 * Revision 1.9  89/04/11  15:18:03  mar
 * removed lots of unused code.
 * set ACKEARLY for all machine types
 * set a long timeout on connection failures
 * 
 * Revision 1.8  89/04/10  22:52:18  mar
 * (probe) fixed checksums, moved ack sending around
 * 
 * Revision 1.7  89/04/04  20:02:38  probe
 * Working RT version.  Includes UDP checksum code (#ifdef UDPCKSUM).
 * 
 * Revision 1.6  89/03/29  18:27:36  probe
 * Forgot a couple of parentheses.
 * 
 * Revision 1.5  89/03/29  17:41:37  probe
 * Ported to the ibm032 architecture
 * 
 * Revision 1.4  89/03/15  18:07:18  mar
 * Put in a timeout on the initial packet
 * Cleaned up print messages & #ifdefed some on DEBUG & TIMING
 * 
 * Revision 1.3  88/07/11  19:08:24  jon
 * Sommerfeld's changes.
 * 
 * Revision 1.2  88/05/31  21:12:20  wesommer
 * [jon] random changes.
 * 
 * 
 *------------------------------------------------------------------
 */

#ifndef lint
static char *rcsid_tftp_c = "$Header: tftp.c,v 1.9 89/04/11 15:18:03 mar Exp $";
#endif	lint

#include <types.h>
#include <sys.h>
#include <a.out.h>

#ifdef ETHER
#include "ether.h"
#else
#include "vii.h"
#endif

#define u_char unsigned char
#define u_short unsigned short

#define ACKEARLY			/* send an ACK before processing pkt */

/* Memory management */
#ifdef vax
#   define ENTRYMASK	0x7fffffff	/* mask for calculating entry point */
#   define CLBYTES	1024		/* cluster size */
#endif vax
#ifdef ibm032
#   define BASEMASK	0x0ff80000	/* mask for calculating base address */
#   define ENTRYMASK	0x0fffffff	/* mask for calculating entry point */
#   define BSS_SLOP	(8 * 512)	/* extra bss to clear */
#   define CLBYTES	2048		/* cluster size */
#   define POST_START	0x800
#   define POST_SIZE	0x800
#   define POST_END	(POST_START+POST_SIZE)
#endif

#define CLOFSET		(CLBYTES-1)


/* Packet structures */
struct udp {
	ln_hdr_type ln_hdr;	/* local net header */
	char ip_vhl;    	/* Internet header length in 32 bit words */
	char ip_tsrv;		/* Type of service */
	short ip_len;		/* Total packet length including header */
	short ip_id;		/* ID for fragmentation */
	short ip_flgs;		/* flags and fragment offset */
/*	short ip_foff : 13;	/* Fragment offset */
	char ip_time;		/* Time to live (secs) */
	char ip_prot;		/* protocol */
	short ip_chksum;	/* Header checksum */
	long ip_src;		/* Source name */
	long ip_dst;		/* Destination name */
	short ud_srcp;		/* source port */
	short ud_dstp;		/* dest port */
	short ud_len;		/* length of UDP packet */
	short ud_cksum;		/* UDP checksum */
	short tf_op;		/* tftp opcode */
	short tf_block;		/* block or error code */
	short tf_data;		/* take the address of this */
	};

struct ph {
	long ph_src;		/* source address */
	long ph_dest;	/* dest address */
	char	ph_zero;	/* zero (reserved) */
	char	ph_prot;	/* protocol */
	short	ph_len;		/* udp length */
	};

/* Some goodly constants, macros and an external */
#define	UDPPROT	17	/* UDP Internet protocol number */
#define	UDPHDRSIZE	(sizeof(struct udp)-sizeof(struct ip))

/* TFTP port number and opcode values */
#ifdef BIG_ENDIAN
#define	TFTPPORT	0x0045		/* TFTP's well known port (69) */
#define	RRQ		0x0001		/* read  request */
#define	WRQ		0x0002		/* write request */
#define	TDATA		0x0003		/* data packet */
#define	ACK		0x0004		/* acknowledgement packet */
#define	ERROR		0x0005		/* error packet */
#else
#define	TFTPPORT	0x4500		/* TFTP's well known port (69) */
#define	RRQ		0x0100		/* read  request */
#define	WRQ		0x0200		/* write request */
#define	TDATA		0x0300		/* data packet */
#define	ACK		0x0400		/* acknowledgement packet */
#define	ERROR		0x0500		/* error packet */
#endif

/* TFTP error codes */
#define	ERRTXT		0	/* see the enclosed text */
#define	FNOTFOUND	1	/* file not found */
#define	ACCESS		2	/* access violation */
#define	DISKFULL	3	/* don't even ask. */
#define	ILLTFTP		4	/* illegal TFTP operation */
#define	BADTID		5	/* unkown transfer ID */
#define	FEXISTS		6	/* file already exists */
#define	NOUSER		7	/* no such user */

/* TFTP states */
#define	DATAWAIT	1
#define	ACKWAIT		2
#define	DEAD		3
#define	TIMEOUT		4
#define	RCVERR		5
#define	RCVACK		6
#define	RCVDATA		7
#define	RCVLASTDATA	8
#define	TERMINATED	9

#define	TFTPTRIES	50	/* # of retries on packet transmission */
#define	REQTRIES	4	/* # of retries on initial request */
#define SHORTTMO	10000	/* this is the only real timeout */
#define	NORMLEN		512	/* normal length of received packet */

#define	max(a,b)	((a) > (b) ? (a) : (b))
#define	min(a,b)	((a) < (b) ? (a) : (b))

/* This is the source code for TFTP. I've tried to write it so that it
	may be called from inside a program with minimal hassle... */

char pin[1600];
char pout[1600];

struct udp *pip;
struct udp *pop;

int tf_expected, tf_lport, tf_fport, tf_fhost;
int ip_uid = 1;

extern char *load_start;
char *tf_data;
int nbytes, nskip, ntbytes;
int tf_bss, tf_entry;

extern long in_me;
extern long gateway;


tftp_use(fhost, rmfile)
	long fhost;
	char *rmfile;
{
	unsigned len;
	register char *data;
	struct exec *x;
	register int i, progress, retries;

	pip = (struct udp *)pin;
	pop = (struct udp *)pout;
	tf_fhost = fhost;
	tf_expected = 1;
	tf_fport = 0;
	tf_lport = 0x2020;

	pop->ud_srcp = 0x2020;
	pop->ud_dstp = TFTPPORT;

	pop->ip_prot = UDPPROT;
	pop->ip_vhl = 0x45;
	pop->ip_time = 0xff;
	pop->ip_flgs = 0;
	pop->ip_id = ip_uid++;
	pop->ip_src = in_me;
	pop->ip_dst = fhost;
	pop->ip_tsrv = 0;

	tfsndreq(rmfile);
	progress = retries = 0;

	while(1) {
		while((len = net_read(pin, sizeof(pin))) == 0)
		  if (progress++ > SHORTTMO) {
		      if (tf_expected == 1) {
			  if (retries++ > REQTRIES) {
			      printf("TFTP: no response from server\n");
			      return(0);
			  }
			  printf("TFTP: no response for file request... still trying...\n");
			  tfsndreq(rmfile);
			  progress = 0;
			  continue;
		      } else if (retries++ > TFTPTRIES) {
			  if (nbytes == 0)
			    goto golaunch;
			  printf("TFTP: lost connection with server\n");
			  return(0);
		      }
		  }

		if ((pip->ip_vhl != 0x45) ||
		    (pip->ip_dst != in_me) ||
		    (pip->ip_prot != UDPPROT) ||
		    (pip->ud_dstp != tf_lport) ||
		    (tf_fport && pip->ud_srcp != tf_fport)) {
#ifdef TIMING
			puts("!");
#endif TIMING
			progress++;
			continue;
		}

		len = swab(pip->ip_len) - 28;
		*((char *)&pip->ud_srcp + swab(pip->ud_len)) = '\0';

		if ((pip->ip_chksum && (u_short)~cksum(&pip->ip_vhl, 10, 0)) ||
		    (pip->ud_cksum && (u_short)~cksum(&pip->ip_src,
						      ((swab(pip->ud_len)+1)>>1) + 4,
						      swab((u_char)pip->ip_prot) + (u_short)pip->ud_len))) {
#ifdef TIMING
			if(pip->ip_chksum&&(u_short)~cksum(&pip->ip_vhl,10,0))
				puts("i");
			else
				puts("u");
#endif TIMING
			continue;
		}

		progress = retries = 0;
		if(pip->tf_op == TDATA) {
			if(tf_fport == 0)
				tf_fport = pop->ud_dstp = pip->ud_srcp;

			if(len < 4) {
				puts("TFDODATA: Died of CSR disease.\n");
				return 0;
			}

			len -= 4; /* sizeof(tftp header).  BAD. */

			if(swab(pip->tf_block) != tf_expected) {
#ifdef TIMING
				puts("@");
#endif TIMING
				tfsndack(tf_expected - 1);
				continue;
			}

#ifdef TIMING
			puts("#");
#else
			if (tf_expected % 32 == 0) puts(".");
#endif
#ifdef ACKEARLY
			/*
			 * Send an ack if we do not mind receiving packets
			 * while we are still processing the current one.
			 * 
			 * Do not do this if you do not have enough receive
			 * buffers.
			 */
			tfsndack(tf_expected);
#endif ACKEARLY

			/* if it's block 1, get info from the exec header */
			if(tf_expected == 1) {
				x = (struct exec *)&pip->tf_data;
				if (x->a_magic != ZMAGIC
				    && x->a_magic != NMAGIC) {
					printf("Not a ZMAGIC or NMAGIC file: %x\n",
					       x->a_magic);
					return 0;
				}
				tf_entry = x->a_entry & ENTRYMASK;
				tf_bss = x->a_bss;
#ifdef ibm032
				load_start = (char *)(x->a_entry & BASEMASK);
				tf_bss += BSS_SLOP;
#endif ibm032
				tf_data = load_start;
				nskip = N_TXTOFF(*x);
				ntbytes = x->a_text;
				nbytes = ntbytes + x->a_data;
			}
			if (tf_expected > 0  && nbytes > 0) {
				data = (char *)&pip->tf_data;
				for(i=0; i<NORMLEN; i++) {
					if (nskip-- > 0) data++;
					else {
						if (nbytes-- == 0) break;
						if (ntbytes-- == 0)
							tf_data = (char *)(((int)tf_data + CLOFSET) & ~CLOFSET);
#ifdef ibm032
						if ((int)tf_data>=POST_START &&
						    (int)tf_data<POST_END) {
							tf_data++; data++;
							continue;
						}
#endif ibm032
						*tf_data++ = *data++;
					}
				}
			}
#ifdef ACKEARLY
			tf_expected++;
#else
			tfsndack(tf_expected++);
#endif

			/* and here we blithely drop the rest of the packets */
			if(len == NORMLEN) continue;
			else {
			golaunch:
				while (tf_bss--)
					*tf_data++ = 0;
				printf("\nFile loaded (%d bytes), entry point is 0x%X\n",
				       tf_data-load_start, tf_entry);
				launch (tf_entry, tf_data-load_start);
				return 1;
			}
		}

		else if(pip->tf_op == ERROR) {
			tfdoerr(len);
			return 0;
		}
		else {
			printf("TFTPRCV: Got bad opcode %d.\n", pip->tf_op);
			continue;
		}

	}
}


/* Utility routines */

/* Format up and send out an initial request for a tftp connection. */

tfsndreq(fname)
	char *fname; {

	pop->tf_op = RRQ;

	strcpy((char *)&(pop->tf_block), fname);
	strcpy((char *)&(pop->tf_block)+strlen(fname)+1, "image");

#ifdef DEBUG
	puts("TFTP:  sending initial request\n");
#endif
	tf_write(strlen(fname)+9);
	}

/* Process an incoming error packet */

tfdoerr(len)
	unsigned len; {

	printf("TFTP: Error from foreign host: %s\n", &(pip->tf_data));
	}

/* ack a certain block number */

tfsndack(number)
	unsigned number; {

	pop->tf_op = ACK;
	pop->tf_block = swab(number);
#ifdef DEBUG
	printf("ack %d\n", number);
#endif
	return tf_write(4);
	}

/* write a tftp packet */

tf_write(len)
	unsigned len; {

#ifdef DEBUG
	printf("tfwrite(%d)\n", len);
#endif

	len += 8;
	if(len & 1) ((char *)&pop->ud_srcp)[len] = 0;

	pop->ud_len = swab(len);
	pop->ud_cksum = swab((u_char)pop->ip_prot) + (u_short)pop->ud_len;
	pop->ud_cksum = (u_short)~cksum(&pop->ip_src, ((len+1)>>1) + 4, 0);

	len += 20;
	pop->ip_len = swab(len);
	pop->ip_id = ip_uid++;
	pop->ip_chksum = 0;
	pop->ip_chksum = ~cksum(&pop->ip_vhl, 10, 0);

#ifdef TIMING
	puts("w");
#endif TIMING
	return net_write(pop, gateway ? gateway : tf_fhost, len);
}
