moved address selection logic to a new 'get_buffer_route'.

- IPv4 now assumes the addresses it is supplied in send_routed_data are already the appropriate ones.
 - made the Data(), operator* and operator-> methods in NetBufferFieldReader const so we can use them in the same expression as the constructor.
 - fixed an issue with UDP where the wrong source address could be used in the calculating the checksum.
 - changed ipv4_print_address to use the more common 0.0.0.0 format.


git-svn-id: file:///srv/svn/repos/haiku/haiku/trunk@20660 a95241bf-73f2-0310-859d-f6bbb57e9c96
This commit is contained in:
Hugo Santos
2007-04-11 22:24:12 +00:00
parent 41d7b9a54c
commit 6c35350908
10 changed files with 196 additions and 145 deletions
+22 -30
View File
@@ -22,56 +22,49 @@ class NetBufferFieldReader {
public: public:
NetBufferFieldReader(net_buffer *buffer) NetBufferFieldReader(net_buffer *buffer)
: :
fBuffer(buffer), fBuffer(buffer)
fStatus(B_BAD_VALUE)
{ {
if ((Offset + sizeof(Type)) <= buffer->size) { fStatus = Module::Get()->direct_access(fBuffer, Offset,
fStatus = Module::Get()->direct_access(fBuffer, Offset, sizeof(Type), (void **)&fData);
sizeof(Type), (void **)&fData); if (fStatus != B_OK) {
if (fStatus != B_OK) { fStatus = Module::Get()->read(fBuffer, Offset,
fData = NULL; &fDataBuffer, sizeof(Type));
fStatus = Module::Get()->read(fBuffer, Offset, fData = &fDataBuffer;
&fDataBuffer, sizeof(Type));
}
} }
} }
status_t status_t
Status() Status() const
{ {
return fStatus; return fStatus;
} }
Type & Type &
Data() Data() const
{ {
if (fData != NULL) return *fData;
return *fData;
return fDataBuffer;
} }
Type * Type *
operator->() operator->() const
{ {
return &Data(); return fData;
} }
Type & Type &
operator*() operator*() const
{ {
return Data(); return *fData;
} }
void void
Sync() Sync()
{ {
if (fBuffer == NULL) if (fBuffer == NULL || fStatus < B_OK)
return; return;
if (fData == NULL) if (fData == &fDataBuffer)
Module::Get()->write(fBuffer, Offset, &fDataBuffer, Module::Get()->write(fBuffer, Offset, fData, sizeof(Type));
sizeof(Type));
fBuffer = NULL; fBuffer = NULL;
} }
@@ -138,15 +131,14 @@ class NetBufferHeaderRemover : public NetBufferHeaderReader<Type, Module> {
template<typename Type, typename Module = NetBufferModuleGetter> template<typename Type, typename Module = NetBufferModuleGetter>
class NetBufferPrepend : public NetBufferFieldReader<Type, 0, Module> { class NetBufferPrepend : public NetBufferFieldReader<Type, 0, Module> {
public: public:
NetBufferPrepend(net_buffer *buffer, size_t size = 0) NetBufferPrepend(net_buffer *buffer, size_t size = sizeof(Type))
{ {
fBuffer = buffer; fBuffer = buffer;
fData = NULL;
if (size == 0) fStatus = Module::Get()->prepend_size(buffer, size,
size = sizeof(Type); (void **)&fData);
if (fStatus == B_OK && fData == NULL)
fStatus = Module::Get()->prepend_size(buffer, size, (void **)&fData); fData = &fDataBuffer;
} }
~NetBufferPrepend() ~NetBufferPrepend()
+4
View File
@@ -79,6 +79,8 @@ struct net_datalink_module_info {
const struct net_route *route); const struct net_route *route);
struct net_route *(*get_route)(struct net_domain *domain, struct net_route *(*get_route)(struct net_domain *domain,
const struct sockaddr *address); const struct sockaddr *address);
status_t (*get_buffer_route)(struct net_domain *domain,
struct net_buffer *buffer, struct net_route **_route);
void (*put_route)(struct net_domain *domain, struct net_route *route); void (*put_route)(struct net_domain *domain, struct net_route *route);
status_t (*register_route_info)(struct net_domain *domain, status_t (*register_route_info)(struct net_domain *domain,
@@ -118,6 +120,8 @@ struct net_address_module_info {
status_t (*set_to)(sockaddr *address, const sockaddr *from); status_t (*set_to)(sockaddr *address, const sockaddr *from);
status_t (*set_to_empty_address)(sockaddr *address); status_t (*set_to_empty_address)(sockaddr *address);
status_t (*update_to)(sockaddr *address, const sockaddr *from);
uint32 (*hash_address_pair)(const sockaddr *ourAddress, uint32 (*hash_address_pair)(const sockaddr *ourAddress,
const sockaddr *peerAddress); const sockaddr *peerAddress);
@@ -257,8 +257,8 @@ icmp_receive_data(net_buffer *buffer)
header.Sync(); header.Sync();
ICMPChecksumField checksum(reply); *ICMPChecksumField(reply) = gBufferModule->checksum(reply, 0,
*checksum = gBufferModule->checksum(reply, 0, reply->size, true); reply->size, true);
status_t status = domain->module->send_data(NULL, reply); status_t status = domain->module->send_data(NULL, reply);
if (status < B_OK) { if (status < B_OK) {
@@ -671,15 +671,6 @@ receiving_protocol(uint8 protocol)
} }
static void
update_checksum(net_buffer *buffer)
{
IPChecksumField checksum(buffer);
*checksum = gBufferModule->checksum(buffer, 0, sizeof(ipv4_header), true);
}
// #pragma mark - // #pragma mark -
@@ -928,49 +919,31 @@ ipv4_send_routed_data(net_protocol *_protocol, struct net_route *route,
TRACE(("someone tries to send some actual routed data!\n")); TRACE(("someone tries to send some actual routed data!\n"));
sockaddr_in &source = *(sockaddr_in *)&buffer->source; sockaddr_in &source = *(sockaddr_in *)&buffer->source;
if (source.sin_addr.s_addr == INADDR_ANY && route->interface->address != NULL) { sockaddr_in &destination = *(sockaddr_in *)&buffer->destination;
// replace an unbound source address with the address of the interface
// TODO: couldn't we replace all addresses here?
source.sin_addr.s_addr = ((sockaddr_in *)route->interface->address)->sin_addr.s_addr;
}
bool headerIncluded = false; bool headerIncluded = false, checksumNeeded = true;
if (protocol != NULL) if (protocol != NULL)
headerIncluded = (protocol->flags & IP_FLAG_HEADER_INCLUDED) != 0; headerIncluded = (protocol->flags & IP_FLAG_HEADER_INCLUDED) != 0;
// Add IP header (if needed) // Add IP header (if needed)
if (!headerIncluded) { if (!headerIncluded) {
NetBufferPrepend<ipv4_header> bufferHeader(buffer); NetBufferPrepend<ipv4_header> header(buffer);
if (bufferHeader.Status() < B_OK) if (header.Status() < B_OK)
return bufferHeader.Status(); return header.Status();
ipv4_header &header = bufferHeader.Data(); header->version = IP_VERSION;
header->header_length = sizeof(ipv4_header) / 4;
header.version = IP_VERSION; header->service_type = protocol ? protocol->service_type : 0;
header.header_length = sizeof(ipv4_header) >> 2; header->total_length = htons(buffer->size);
header.service_type = protocol ? protocol->service_type : 0; header->id = htons(atomic_add(&sPacketID, 1));
header.total_length = htons(buffer->size); header->fragment_offset = 0;
header.id = htons(atomic_add(&sPacketID, 1)); header->time_to_live = protocol ? protocol->time_to_live : 254;
header.fragment_offset = 0; header->protocol = protocol ? protocol->socket->protocol : buffer->protocol;
header.time_to_live = protocol ? protocol->time_to_live : 254; header->checksum = 0;
header.protocol = protocol ? protocol->socket->protocol : buffer->protocol;
header.checksum = 0;
if (route->interface->address != NULL) {
header.source = ((sockaddr_in *)route->interface->address)->sin_addr.s_addr;
// always use the actual used source address
} else
header.source = 0;
header.destination = ((sockaddr_in *)&buffer->destination)->sin_addr.s_addr;
bufferHeader.Sync();
// make sure the IP-header is already written to the
// buffer at this point
update_checksum(buffer);
//dump_ipv4_header(header);
header->source = source.sin_addr.s_addr;
header->destination = destination.sin_addr.s_addr;
} else { } else {
// if IP_HDRINCL, check if the source address is set // if IP_HDRINCL, check if the source address is set
NetBufferHeaderReader<ipv4_header> header(buffer); NetBufferHeaderReader<ipv4_header> header(buffer);
@@ -980,22 +953,24 @@ ipv4_send_routed_data(net_protocol *_protocol, struct net_route *route,
if (header->source == 0) { if (header->source == 0) {
header->source = source.sin_addr.s_addr; header->source = source.sin_addr.s_addr;
header->checksum = 0; header->checksum = 0;
header.Sync(); header.Sync();
} else
update_checksum(buffer); checksumNeeded = false;
}
} }
if (buffer->size > 0xffff) if (buffer->size > 0xffff)
return EMSGSIZE; return EMSGSIZE;
if (checksumNeeded)
*IPChecksumField(buffer) = gBufferModule->checksum(buffer, 0,
sizeof(ipv4_header), true);
TRACE(("header chksum: %ld, buffer checksum: %ld\n", TRACE(("header chksum: %ld, buffer checksum: %ld\n",
gBufferModule->checksum(buffer, 0, sizeof(ipv4_header), true), gBufferModule->checksum(buffer, 0, sizeof(ipv4_header), true),
gBufferModule->checksum(buffer, 0, buffer->size, true))); gBufferModule->checksum(buffer, 0, buffer->size, true)));
TRACE(("destination-IP: buffer=%p addr=%p %08lx\n", buffer, &buffer->destination, TRACE(("destination-IP: buffer=%p addr=%p %08lx\n", buffer, &buffer->destination,
ntohl(((sockaddr_in *)&buffer->destination)->sin_addr.s_addr))); ntohl(destination->sin_addr.s_addr)));
uint32 mtu = route->mtu ? route->mtu : interface->mtu; uint32 mtu = route->mtu ? route->mtu : interface->mtu;
if (buffer->size > mtu) { if (buffer->size > mtu) {
@@ -1012,14 +987,13 @@ ipv4_send_data(net_protocol *protocol, net_buffer *buffer)
{ {
TRACE(("someone tries to send some actual data!\n")); TRACE(("someone tries to send some actual data!\n"));
// find route net_route *route = NULL;
struct net_route *route = sDatalinkModule->get_route(sDomain, status_t status = sDatalinkModule->get_buffer_route(sDomain, buffer,
(sockaddr *)&buffer->destination); &route);
if (route == NULL) if (status >= B_OK) {
return ENETUNREACH; status = ipv4_send_routed_data(protocol, route, buffer);
sDatalinkModule->put_route(sDomain, route);
status_t status = ipv4_send_routed_data(protocol, route, buffer); }
sDatalinkModule->put_route(sDomain, route);
return status; return status;
} }
@@ -1,5 +1,5 @@
/* /*
* Copyright 2006, Haiku, Inc. All Rights Reserved. * Copyright 2006-2007, Haiku, Inc. All Rights Reserved.
* Distributed under the terms of the MIT License. * Distributed under the terms of the MIT License.
* *
* Authors: * Authors:
@@ -251,24 +251,35 @@ ipv4_check_mask(const sockaddr *_mask)
\return B_NO_MEMORY if the buffer could not be allocated \return B_NO_MEMORY if the buffer could not be allocated
*/ */
static status_t static status_t
ipv4_print_address(const sockaddr *address, char **_buffer, bool printPort) ipv4_print_address(const sockaddr *_address, char **_buffer, bool printPort)
{ {
const sockaddr_in *address = (const sockaddr_in *)_address;
if (_buffer == NULL) if (_buffer == NULL)
return B_BAD_VALUE; return B_BAD_VALUE;
int bufLen = printPort ? 15 : 9; char tmp[64];
char *buffer = (char *)malloc(bufLen);
if (buffer == NULL)
return B_NO_MEMORY;
if (address == NULL) if (address == NULL)
strcpy(buffer, "<none>"); strcpy(tmp, "<none>");
else if (printPort) { else {
sprintf(buffer, "%08lx:%u", ntohl(((sockaddr_in *)address)->sin_addr.s_addr), unsigned int addr = ntohl(address->sin_addr.s_addr);
ntohs(((sockaddr_in *)address)->sin_port));
} else if (printPort)
sprintf(buffer, "%08lx", ntohl(((sockaddr_in *)address)->sin_addr.s_addr)); sprintf(tmp, "%u.%u.%u.%u:%u", (addr >> 24) & 0xff,
*_buffer = buffer; (addr >> 16) & 0xff, (addr >> 8) & 0xff, addr & 0xff,
ntohs(address->sin_port));
else
sprintf(tmp, "%u.%u.%u.%u", (addr >> 24) & 0xff,
(addr >> 16) & 0xff, (addr >> 8) & 0xff, addr & 0xff);
}
*_buffer = strdup(tmp);
if (*_buffer == NULL)
return B_NO_MEMORY;
return B_OK; return B_OK;
} }
@@ -324,6 +335,31 @@ ipv4_set_to(sockaddr *address, const sockaddr *from)
} }
static status_t
ipv4_update_to(sockaddr *_address, const sockaddr *_from)
{
sockaddr_in *address = (sockaddr_in *)_address;
const sockaddr_in *from = (const sockaddr_in *)_from;
if (address == NULL || from == NULL)
return B_BAD_VALUE;
if (from->sin_family != AF_INET)
return B_BAD_VALUE;
address->sin_family = AF_INET;
address->sin_len = sizeof(sockaddr_in);
if (address->sin_port == 0)
address->sin_port = from->sin_port;
if (address->sin_addr.s_addr == INADDR_ANY)
address->sin_addr.s_addr = from->sin_addr.s_addr;
return B_OK;
}
/*! /*!
Sets \a address to the empty address (0.0.0.0). Sets \a address to the empty address (0.0.0.0).
\return B_OK if \a address has been set \return B_OK if \a address has been set
@@ -394,6 +430,7 @@ net_address_module_info gIPv4AddressModule = {
ipv4_set_port, ipv4_set_port,
ipv4_set_to, ipv4_set_to,
ipv4_set_to_empty_address, ipv4_set_to_empty_address,
ipv4_update_to,
ipv4_hash_address_pair, ipv4_hash_address_pair,
ipv4_checksum_address, ipv4_checksum_address,
NULL // ipv4_matches_broadcast_address, NULL // ipv4_matches_broadcast_address,
@@ -170,8 +170,7 @@ add_tcp_header(tcp_segment_header &segment, net_buffer *buffer)
<< (uint16)htons(buffer->size) << (uint16)htons(buffer->size)
<< Checksum::BufferHelper(buffer, gBufferModule); << Checksum::BufferHelper(buffer, gBufferModule);
TCPChecksumField checksumField(buffer); *TCPChecksumField(buffer) = checksum;
*checksumField = checksum;
return B_OK; return B_OK;
} }
@@ -47,6 +47,10 @@ struct udp_header {
} _PACKED; } _PACKED;
typedef NetBufferField<uint16, offsetof(udp_header, udp_checksum)>
UDPChecksumField;
class UdpEndpoint : public net_protocol { class UdpEndpoint : public net_protocol {
public: public:
UdpEndpoint(net_socket *socket); UdpEndpoint(net_socket *socket);
@@ -766,43 +770,44 @@ UdpEndpoint::SendData(net_buffer *buffer, net_route *route)
{ {
if (buffer->size > (0xffff - sizeof(udp_header))) if (buffer->size > (0xffff - sizeof(udp_header)))
return EMSGSIZE; return EMSGSIZE;
buffer->protocol = IPPROTO_UDP; buffer->protocol = IPPROTO_UDP;
{ // scope for lifetime of bufferHeader // add and fill UDP-specific header:
NetBufferPrepend<udp_header> header(buffer);
if (header.Status() < B_OK)
return header.Status();
// add and fill UDP-specific header: header->source_port = sAddressModule->get_port((sockaddr *)&buffer->source);
NetBufferPrepend<udp_header> bufferHeader(buffer); header->destination_port = sAddressModule->get_port(
if (bufferHeader.Status() < B_OK) (sockaddr *)&buffer->destination);
return bufferHeader.Status(); header->udp_length = htons(buffer->size);
// the udp-header is already included in the buffer-size
udp_header &header = bufferHeader.Data(); header->udp_checksum = 0;
header.source_port = sAddressModule->get_port((sockaddr *)&buffer->source); header.Sync();
header.destination_port = sAddressModule->get_port(
(sockaddr *)&buffer->destination); // generate UDP-checksum (simulating a so-called "pseudo-header"):
header.udp_length = htons(buffer->size); Checksum udpChecksum;
// the udp-header is already included in the buffer-size sAddressModule->checksum_address(&udpChecksum,
header.udp_checksum = 0; (sockaddr *)route->interface->address);
sAddressModule->checksum_address(&udpChecksum,
(sockaddr *)&buffer->destination);
udpChecksum
<< (uint16)htons(IPPROTO_UDP)
<< (uint16)htons(buffer->size)
// peculiar but correct: UDP-len is used twice for checksum
// (as it is already contained in udp_header)
<< Checksum::BufferHelper(buffer, gBufferModule);
uint16 calculatedChecksum = udpChecksum;
if (calculatedChecksum == 0)
calculatedChecksum = 0xffff;
*UDPChecksumField(buffer) = calculatedChecksum;
TRACE_BLOCK(((char*)&header, sizeof(udp_header), "udp-hdr: "));
// generate UDP-checksum (simulating a so-called "pseudo-header"):
Checksum udpChecksum;
sAddressModule->checksum_address(&udpChecksum,
(sockaddr *)route->interface->address);
sAddressModule->checksum_address(&udpChecksum,
(sockaddr *)&buffer->destination);
udpChecksum
<< (uint16)htons(IPPROTO_UDP)
<< (uint16)htons(buffer->size)
// peculiar but correct: UDP-len is used twice for checksum
// (as it is already contained in udp_header)
<< Checksum::BufferHelper(buffer, gBufferModule);
header.udp_checksum = udpChecksum;
if (header.udp_checksum == 0)
header.udp_checksum = 0xFFFF;
TRACE_BLOCK(((char*)&header, sizeof(udp_header), "udp-hdr: "));
}
return next->module->send_routed_data(next, route, buffer); return next->module->send_routed_data(next, route, buffer);
} }
@@ -966,8 +971,8 @@ udp_send_routed_data(net_protocol *protocol, struct net_route *route,
net_buffer *buffer) net_buffer *buffer)
{ {
TRACE(("udp_send_routed_data(%p) size=%lu\n", protocol, buffer->size)); TRACE(("udp_send_routed_data(%p) size=%lu\n", protocol, buffer->size));
UdpEndpoint *udpEndpoint = (UdpEndpoint *)protocol;
return udpEndpoint->SendData(buffer, route); return ((UdpEndpoint *)protocol)->SendData(buffer, route);
} }
@@ -976,14 +981,14 @@ udp_send_data(net_protocol *protocol, net_buffer *buffer)
{ {
TRACE(("udp_send_data(%p) size=%lu\n", protocol, buffer->size)); TRACE(("udp_send_data(%p) size=%lu\n", protocol, buffer->size));
struct net_route *route = sDatalinkModule->get_route(sDomain, net_route *route = NULL;
(sockaddr *)&buffer->destination); status_t status = sDatalinkModule->get_buffer_route(sDomain, buffer,
if (route == NULL) &route);
return ENETUNREACH; if (status >= B_OK) {
status = udp_send_routed_data(protocol, route, buffer);
sDatalinkModule->put_route(sDomain, route);
}
UdpEndpoint *udpEndpoint = (UdpEndpoint *)protocol;
status_t status = udpEndpoint->SendData(buffer, route);
sDatalinkModule->put_route(sDomain, route);
return status; return status;
} }
@@ -749,6 +749,7 @@ net_datalink_module_info gNetDatalinkModule = {
add_route, add_route,
remove_route, remove_route,
get_route, get_route,
get_buffer_route,
put_route, put_route,
register_route_info, register_route_info,
unregister_route_info, unregister_route_info,
@@ -535,6 +535,43 @@ get_route(struct net_domain *_domain, const struct sockaddr *address)
} }
status_t
get_buffer_route(net_domain *_domain, net_buffer *buffer, net_route **_route)
{
net_domain_private *domain = (net_domain_private *)_domain;
BenaphoreLocker _(domain->lock);
net_route *route = get_route_internal(domain,
(sockaddr *)&buffer->destination);
if (route == NULL)
return ENETUNREACH;
status_t status = B_OK;
sockaddr *source = (sockaddr *)&buffer->source;
// TODO we are quite relaxed in the address checking here
// as we might proceed with srcaddr=INADDR_ANY.
if (route->interface && route->interface->address) {
sockaddr *interfaceAddress = route->interface->address;
net_address_module_info *addressModule = domain->address_module;
if (addressModule->is_empty_address(source, true))
addressModule->set_to(source, interfaceAddress);
else
status = addressModule->update_to(source, interfaceAddress);
}
if (status != B_OK)
put_route_internal(domain, route);
else
*_route = route;
return status;
}
void void
put_route(struct net_domain *_domain, net_route *route) put_route(struct net_domain *_domain, net_route *route)
{ {
@@ -42,6 +42,8 @@ status_t get_route_information(struct net_domain *domain, void *buffer,
void invalidate_routes(net_domain *, net_interface *); void invalidate_routes(net_domain *, net_interface *);
struct net_route *get_route(struct net_domain *domain, struct net_route *get_route(struct net_domain *domain,
const struct sockaddr *address); const struct sockaddr *address);
status_t get_buffer_route(struct net_domain *domain,
struct net_buffer *buffer, struct net_route **_route);
void put_route(struct net_domain *domain, struct net_route *route); void put_route(struct net_domain *domain, struct net_route *route);
status_t register_route_info(struct net_domain *domain, status_t register_route_info(struct net_domain *domain,