#include <sys/cdefs.h>
__KERNEL_RCSID(0, "$NetBSD: in_cksum.c,v 1.15 2011/07/10 23:13:22 matt Exp $");
#include <sys/param.h>
#include <sys/endian.h>
#include <sys/mbuf.h>
#include <sys/systm.h>
#include <netinet/in_systm.h>
#include <netinet/in.h>
#include <netinet/ip.h>
#include <netinet/ip_var.h>
union memptr {
uint32_t *l;
uintptr_t u;
uint16_t *s;
uint8_t *c;
};
static inline uint32_t fastsum(union memptr, int, unsigned int, int);
static inline uint32_t
fastsum(union memptr buf, int n, unsigned int oldsum, int odd_aligned)
{
unsigned long hilo = 0, high = 0;
unsigned long w0, w1;
unsigned int sum = 0;
if (buf.u & 0x3) {
if (buf.u & 0x1) {
#if BYTE_ORDER == BIG_ENDIAN
sum += *(buf.c++);
#else
sum += (*(buf.c++) << 8);
#endif
n -= 1;
odd_aligned = !odd_aligned;
}
if (n <= 2)
goto postunaligned;
if (buf.u & 0x2) {
sum += *(buf.s++);
n -= 2;
}
}
if (n < 64 + 8)
goto notmuchleft;
w0 = buf.l[0];
w1 = buf.l[1];
do {
hilo += w0;
high += w0 >> 16;
w0 = buf.l[2];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[3];
hilo += w0;
high += w0 >> 16;
w0 = buf.l[4];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[5];
hilo += w0;
high += w0 >> 16;
w0 = buf.l[6];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[7];
hilo += w0;
high += w0 >> 16;
w0 = buf.l[8];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[9];
hilo += w0;
high += w0 >> 16;
w0 = buf.l[10];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[11];
hilo += w0;
high += w0 >> 16;
w0 = buf.l[12];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[13];
hilo += w0;
high += w0 >> 16;
w0 = buf.l[14];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[15];
hilo += w0;
high += w0 >> 16;
w0 = buf.l[16];
hilo += w1;
high += w1 >> 16;
w1 = buf.l[17];
n -= 64;
buf.c += 64;
} while (n >= 64 + 8);
hilo -= (high << 16);
sum += hilo;
sum += high;
notmuchleft:
high = hilo = 0;
while (n >= sizeof(uint32_t)) {
w0 = *(buf.l++);
hilo += w0;
high += w0 >> 16;
n -= 4;
}
hilo -= (high << 16);
sum += hilo;
sum += high;
postunaligned:
if (n >= sizeof(uint16_t)) {
sum += *(buf.s++);
n -= sizeof(uint16_t);
}
if (n > 0) {
#if BYTE_ORDER == BIG_ENDIAN
sum += *(buf.c++) << 8;
#else
sum += *(buf.c++);
#endif
n = 0;
}
if (odd_aligned) {
sum = (sum & 0xffff) + (sum >> 16);
sum = (sum & 0xffff) + (sum >> 16);
sum = oldsum + ((sum >> 8) & 0xff) + ((sum & 0xff) << 8);
} else {
sum = oldsum + sum;
sum = (sum & 0xffff) + (sum >> 16);
}
sum = (sum & 0xffff) + (sum >> 16);
return(sum);
}
static inline int
in_cksum_internal(struct mbuf *m, int off, int len, uint32_t sum)
{
union memptr w;
int mlen;
int odd_aligned = 0;
for (; m && len; m = m->m_next) {
if (m->m_len == 0)
continue;
w.c = mtod(m, u_char *) + off;
mlen = m->m_len - off;
off = 0;
if (len < mlen)
mlen = len;
len -= mlen;
sum = fastsum(w, mlen, sum, odd_aligned);
odd_aligned = (odd_aligned + mlen) & 0x01;
}
if (len != 0) {
printf("cksum: out of data, %d\n", len);
}
return (~sum & 0xffff);
}
int
in_cksum(struct mbuf *m, int len)
{
return (in_cksum_internal(m, 0, len, 0));
}
int
in4_cksum(struct mbuf *m, uint8_t nxt, int off, int len)
{
uint sum = 0;
if (nxt != 0) {
uint16_t *w;
union {
struct ipovly ipov;
u_int16_t w[10];
} u;
memset(&u.ipov, 0, sizeof(u.ipov));
u.ipov.ih_len = htons(len);
u.ipov.ih_pr = nxt;
u.ipov.ih_src = mtod(m, struct ip *)->ip_src;
u.ipov.ih_dst = mtod(m, struct ip *)->ip_dst;
w = u.w;
sum += w[0]; sum += w[1]; sum += w[2]; sum += w[3]; sum += w[4];
sum += w[5]; sum += w[6]; sum += w[7]; sum += w[8]; sum += w[9];
}
while (m && off > 0) {
if (m->m_len > off)
break;
off -= m->m_len;
m = m->m_next;
}
return (in_cksum_internal(m, off, len, sum));
}