some NetBufferUtility revamp.

- introduced base NetBufferFieldReader which deals with obtaining a continuguous field to deal with.
 - renamed Detach() to Sync() to better express the fact that we may be writing to the buffer.
 - Fixed IPv4's and ICMP's checksum calculation as it assumed the underlying header was contiguous.
 - ICMP is now able to reply to ICMP Echo Requests which were split across data nodes.
 - in split_buffer(), make sure we were able to trim() the new buffer before damaging the original one.


git-svn-id: file:///srv/svn/repos/haiku/haiku/trunk@20657 a95241bf-73f2-0310-859d-f6bbb57e9c96
This commit is contained in:
Hugo Santos
2007-04-11 16:40:40 +00:00
parent 453fb26fda
commit 87001e059c
8 changed files with 153 additions and 142 deletions
+90 -87
View File
@@ -16,94 +16,26 @@ class NetBufferModuleGetter {
static net_buffer_module_info *Get() { return gBufferModule; } static net_buffer_module_info *Get() { return gBufferModule; }
}; };
//! A class to retrieve and remove a header from a buffer //! A class to access a field safely across node boundaries
template<typename Type, typename Module = NetBufferModuleGetter > class NetBufferHeader { template<typename Type, int Offset, typename Module = NetBufferModuleGetter>
class NetBufferFieldReader {
public: public:
NetBufferHeader(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,
~NetBufferHeader() sizeof(Type), (void **)&fData);
{ if (fStatus != B_OK) {
Remove(); fData = NULL;
} fStatus = Module::Get()->read(fBuffer, Offset,
&fDataBuffer, sizeof(Type));
status_t }
Status()
{
return fBuffer->size < sizeof(Type) ? B_BAD_VALUE : B_OK;
}
status_t
SetTo(net_buffer *buffer)
{
fBuffer = buffer;
return Status();
}
Type &
Data()
{
Type *data;
if (Module::Get()->direct_access(fBuffer, 0, sizeof(Type),
(void **)&data) == B_OK)
return *data;
Module::Get()->read(fBuffer, 0, &fDataBuffer, sizeof(Type));
return fDataBuffer;
}
void
Remove()
{
Remove(sizeof(Type));
}
void
Remove(size_t bytes)
{
if (fBuffer != NULL) {
Module::Get()->remove_header(fBuffer, bytes);
fBuffer = NULL;
} }
} }
void
Detach()
{
fBuffer = NULL;
}
private:
net_buffer *fBuffer;
Type fDataBuffer;
};
//! A class to access a header safely across data node boundaries
template<typename Type, typename Module = NetBufferModuleGetter>
class NetBufferSafeHeader {
public:
NetBufferSafeHeader(net_buffer *buffer)
:
fBuffer(buffer)
{
fStatus = Module::Get()->direct_access(fBuffer, 0,
sizeof(Type), (void **)&fData);
if (fStatus != B_OK) {
fData = NULL;
fStatus = Module::Get()->read(fBuffer, 0, &fDataBuffer,
sizeof(Type));
}
}
~NetBufferSafeHeader()
{
if (fBuffer != NULL)
Detach();
}
status_t status_t
Status() Status()
{ {
@@ -119,16 +51,33 @@ class NetBufferSafeHeader {
return fDataBuffer; return fDataBuffer;
} }
void Type *
Detach() operator->()
{ {
return &Data();
}
Type &
operator*()
{
return Data();
}
void
Sync()
{
if (fBuffer == NULL)
return;
if (fData == NULL) if (fData == NULL)
Module::Get()->write(fBuffer, 0, &fDataBuffer, sizeof(Type)); Module::Get()->write(fBuffer, Offset, &fDataBuffer,
sizeof(Type));
fBuffer = NULL; fBuffer = NULL;
} }
protected: protected:
NetBufferSafeHeader() {} NetBufferFieldReader() {}
net_buffer *fBuffer; net_buffer *fBuffer;
status_t fStatus; status_t fStatus;
@@ -136,9 +85,58 @@ class NetBufferSafeHeader {
Type fDataBuffer; Type fDataBuffer;
}; };
template<typename Type, int Offset, typename Module = NetBufferModuleGetter>
class NetBufferField : public NetBufferFieldReader<Type, Offset, Module> {
public:
NetBufferField(net_buffer *buffer)
: NetBufferFieldReader<Type, Offset, Module>(buffer)
{}
~NetBufferField()
{
Sync();
}
};
template<typename Type, typename Module = NetBufferModuleGetter>
class NetBufferHeaderReader : public NetBufferFieldReader<Type, 0, Module> {
public:
NetBufferHeaderReader(net_buffer *buffer)
: NetBufferFieldReader<Type, 0, Module>(buffer)
{}
void
Remove()
{
Remove(sizeof(Type));
}
void
Remove(size_t bytes)
{
if (fBuffer != NULL) {
Module::Get()->remove_header(fBuffer, bytes);
fBuffer = NULL;
}
}
};
template<typename Type, typename Module = NetBufferModuleGetter>
class NetBufferHeaderRemover : public NetBufferHeaderReader<Type, Module> {
public:
NetBufferHeaderRemover(net_buffer *buffer)
: NetBufferHeaderReader<Type, Module>(buffer)
{}
~NetBufferHeaderRemover()
{
Remove();
}
};
//! A class to add a header to a buffer //! A class to add a header to a buffer
template<typename Type, typename Module = NetBufferModuleGetter> template<typename Type, typename Module = NetBufferModuleGetter>
class NetBufferPrepend : public NetBufferSafeHeader<Type, 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 = 0)
{ {
@@ -150,6 +148,11 @@ class NetBufferPrepend : public NetBufferSafeHeader<Type, Module> {
fStatus = Module::Get()->prepend_size(buffer, size, (void **)&fData); fStatus = Module::Get()->prepend_size(buffer, size, (void **)&fData);
} }
~NetBufferPrepend()
{
Sync();
}
}; };
#endif // NET_BUFFER_UTILITIES_H #endif // NET_BUFFER_UTILITIES_H
@@ -353,7 +353,7 @@ arp_receive(void *cookie, net_buffer *buffer)
{ {
TRACE(("ARP receive\n")); TRACE(("ARP receive\n"));
NetBufferHeader<arp_header> bufferHeader(buffer); NetBufferHeaderReader<arp_header> bufferHeader(buffer);
if (bufferHeader.Status() < B_OK) if (bufferHeader.Status() < B_OK)
return bufferHeader.Status(); return bufferHeader.Status();
@@ -383,8 +383,6 @@ arp_receive(void *cookie, net_buffer *buffer)
|| header.protocol_length != sizeof(in_addr_t)) || header.protocol_length != sizeof(in_addr_t))
return B_BAD_DATA; return B_BAD_DATA;
bufferHeader.Detach();
// handle packet // handle packet
switch (opcode) { switch (opcode) {
@@ -38,7 +38,7 @@ ethernet_deframe(net_device *device, net_buffer *buffer)
{ {
//dprintf("asked to deframe buffer for device %s\n", device->name); //dprintf("asked to deframe buffer for device %s\n", device->name);
NetBufferHeader<ether_header> bufferHeader(buffer); NetBufferHeaderRemover<ether_header> bufferHeader(buffer);
if (bufferHeader.Status() < B_OK) if (bufferHeader.Status() < B_OK)
return bufferHeader.Status(); return bufferHeader.Status();
@@ -138,7 +138,7 @@ ethernet_frame_send_data(net_datalink_protocol *protocol,
else else
memcpy(header.destination, destination.sdl_data, ETHER_ADDRESS_LENGTH); memcpy(header.destination, destination.sdl_data, ETHER_ADDRESS_LENGTH);
bufferHeader.Detach(); bufferHeader.Sync();
// make sure the framing is already written to the buffer at this point // make sure the framing is already written to the buffer at this point
return protocol->next->module->send_data(protocol->next, buffer); return protocol->next->module->send_data(protocol->next, buffer);
@@ -48,6 +48,8 @@ struct icmp_header {
}; };
}; };
typedef NetBufferField<uint16, offsetof(icmp_header, checksum)> ICMPChecksumField;
#define ICMP_TYPE_ECHO_REPLY 0 #define ICMP_TYPE_ECHO_REPLY 0
#define ICMP_TYPE_UNREACH 3 #define ICMP_TYPE_UNREACH 3
#define ICMP_TYPE_REDIRECT 5 #define ICMP_TYPE_REDIRECT 5
@@ -211,13 +213,11 @@ icmp_receive_data(net_buffer *buffer)
{ {
TRACE(("ICMP received some data, buffer length %lu\n", buffer->size)); TRACE(("ICMP received some data, buffer length %lu\n", buffer->size));
NetBufferHeader<icmp_header> bufferHeader(buffer); NetBufferHeaderReader<icmp_header> bufferHeader(buffer);
if (bufferHeader.Status() < B_OK) if (bufferHeader.Status() < B_OK)
return bufferHeader.Status(); return bufferHeader.Status();
icmp_header &header = bufferHeader.Data(); icmp_header &header = bufferHeader.Data();
bufferHeader.Detach();
// the pointer stays valid after this
TRACE((" got type %u, code %u, checksum %u\n", header.type, header.code, TRACE((" got type %u, code %u, checksum %u\n", header.type, header.code,
ntohs(header.checksum))); ntohs(header.checksum)));
@@ -249,19 +249,18 @@ icmp_receive_data(net_buffer *buffer)
memcpy(&reply->destination, &buffer->source, buffer->source.ss_len); memcpy(&reply->destination, &buffer->source, buffer->source.ss_len);
// There already is an ICMP header, and we'll reuse it // There already is an ICMP header, and we'll reuse it
icmp_header *header; NetBufferHeaderReader<icmp_header> header(reply);
status_t status = gBufferModule->direct_access(reply,
0, sizeof(icmp_header), (void **)&header);
if (status == B_OK) {
header->type = ICMP_TYPE_ECHO_REPLY;
header->code = 0;
header->checksum = 0;
header->checksum = gBufferModule->checksum(reply, 0, reply->size, true);
}
if (status == B_OK) header->type = ICMP_TYPE_ECHO_REPLY;
status = domain->module->send_data(NULL, reply); header->code = 0;
header->checksum = 0;
header.Sync();
ICMPChecksumField checksum(reply);
*checksum = gBufferModule->checksum(reply, 0, reply->size, true);
status_t status = domain->module->send_data(NULL, reply);
if (status < B_OK) { if (status < B_OK) {
gBufferModule->free(reply); gBufferModule->free(reply);
return status; return status;
@@ -74,6 +74,8 @@ struct ipv4_header {
typedef DoublyLinkedList<struct net_buffer, typedef DoublyLinkedList<struct net_buffer,
DoublyLinkedListCLink<struct net_buffer> > FragmentList; DoublyLinkedListCLink<struct net_buffer> > FragmentList;
typedef NetBufferField<uint16, offsetof(ipv4_header, checksum)> IPChecksumField;
struct ipv4_packet_key { struct ipv4_packet_key {
in_addr_t source; in_addr_t source;
in_addr_t destination; in_addr_t destination;
@@ -560,14 +562,11 @@ send_fragments(ipv4_protocol *protocol, struct net_route *route,
TRACE(("ipv4 needs to fragment (size %lu, MTU %lu)...\n", TRACE(("ipv4 needs to fragment (size %lu, MTU %lu)...\n",
buffer->size, mtu)); buffer->size, mtu));
NetBufferHeader<ipv4_header> bufferHeader(buffer); NetBufferHeaderReader<ipv4_header> originalHeader(buffer);
if (bufferHeader.Status() < B_OK) if (originalHeader.Status() < B_OK)
return bufferHeader.Status(); return originalHeader.Status();
ipv4_header *header = &bufferHeader.Data(); uint16 headerLength = originalHeader->HeaderLength();
bufferHeader.Detach();
uint16 headerLength = header->HeaderLength();
uint32 bytesLeft = buffer->size - headerLength; uint32 bytesLeft = buffer->size - headerLength;
uint32 fragmentOffset = 0; uint32 fragmentOffset = 0;
status_t status = B_OK; status_t status = B_OK;
@@ -576,9 +575,10 @@ send_fragments(ipv4_protocol *protocol, struct net_route *route,
if (headerBuffer == NULL) if (headerBuffer == NULL)
return B_NO_MEMORY; return B_NO_MEMORY;
bufferHeader.SetTo(headerBuffer); // TODO we need to make sure ipv4_header is contiguous or
header = &bufferHeader.Data(); // use another construct.
bufferHeader.Detach(); NetBufferHeaderReader<ipv4_header> bufferHeader(headerBuffer);
ipv4_header *header = &bufferHeader.Data();
// adapt MTU to be a multiple of 8 (fragment offsets can only be specified this way) // adapt MTU to be a multiple of 8 (fragment offsets can only be specified this way)
mtu -= headerLength; mtu -= headerLength;
@@ -671,6 +671,15 @@ 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 -
@@ -955,27 +964,27 @@ ipv4_send_routed_data(net_protocol *_protocol, struct net_route *route,
header.destination = ((sockaddr_in *)&buffer->destination)->sin_addr.s_addr; header.destination = ((sockaddr_in *)&buffer->destination)->sin_addr.s_addr;
header.checksum = gBufferModule->checksum(buffer, 0, bufferHeader.Sync();
sizeof(ipv4_header), true);
//dump_ipv4_header(header);
bufferHeader.Detach();
// make sure the IP-header is already written to the // make sure the IP-header is already written to the
// buffer at this point // buffer at this point
update_checksum(buffer);
//dump_ipv4_header(header);
} else { } else {
// if IP_HDRINCL, check if the source address is set // if IP_HDRINCL, check if the source address is set
NetBufferHeader<ipv4_header> bufferHeader(buffer); NetBufferHeaderReader<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(); 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 = gBufferModule->checksum(buffer,
sizeof(ipv4_header), sizeof(ipv4_header), true); header.Sync();
update_checksum(buffer);
} }
bufferHeader.Detach();
} }
if (buffer->size > 0xffff) if (buffer->size > 0xffff)
@@ -1079,12 +1088,11 @@ ipv4_receive_data(net_buffer *buffer)
{ {
TRACE(("IPv4 received a packet (%p) of %ld size!\n", buffer, buffer->size)); TRACE(("IPv4 received a packet (%p) of %ld size!\n", buffer, buffer->size));
NetBufferHeader<ipv4_header> bufferHeader(buffer); NetBufferHeaderReader<ipv4_header> bufferHeader(buffer);
if (bufferHeader.Status() < B_OK) if (bufferHeader.Status() < B_OK)
return bufferHeader.Status(); return bufferHeader.Status();
ipv4_header &header = bufferHeader.Data(); ipv4_header &header = bufferHeader.Data();
bufferHeader.Detach();
//dump_ipv4_header(header); //dump_ipv4_header(header);
if (header.version != IP_VERSION) if (header.version != IP_VERSION)
@@ -39,6 +39,9 @@
#endif #endif
typedef NetBufferField<uint16, offsetof(tcp_header, checksum)> TCPChecksumField;
net_domain *gDomain; net_domain *gDomain;
net_address_module_info *gAddressModule; net_address_module_info *gAddressModule;
net_buffer_module_info *gBufferModule; net_buffer_module_info *gBufferModule;
@@ -150,7 +153,7 @@ add_tcp_header(tcp_segment_header &segment, net_buffer *buffer)
// we must detach before calculating the checksum as we may // we must detach before calculating the checksum as we may
// not have a contiguous buffer. // not have a contiguous buffer.
bufferHeader.Detach(); bufferHeader.Sync();
if (optionsLength > 0) if (optionsLength > 0)
gBufferModule->write(buffer, sizeof(tcp_header), optionsBuffer, optionsLength); gBufferModule->write(buffer, sizeof(tcp_header), optionsBuffer, optionsLength);
@@ -167,9 +170,8 @@ 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);
// we are pretty sure the header is there. TCPChecksumField checksumField(buffer);
NetBufferSafeHeader<tcp_header> headerRef(buffer); *checksumField = checksum;
headerRef.Data().checksum = checksum;
return B_OK; return B_OK;
} }
@@ -507,7 +509,7 @@ tcp_receive_data(net_buffer *buffer)
if (gDomain == NULL && set_domain(buffer->interface) != B_OK) if (gDomain == NULL && set_domain(buffer->interface) != B_OK)
return B_ERROR; return B_ERROR;
NetBufferHeader<tcp_header> bufferHeader(buffer); NetBufferHeaderReader<tcp_header> bufferHeader(buffer);
if (bufferHeader.Status() < B_OK) if (bufferHeader.Status() < B_OK)
return bufferHeader.Status(); return bufferHeader.Status();
@@ -440,7 +440,7 @@ UdpEndpointManager::DemuxIncomingBuffer(net_buffer *buffer)
status_t status_t
UdpEndpointManager::ReceiveData(net_buffer *buffer) UdpEndpointManager::ReceiveData(net_buffer *buffer)
{ {
NetBufferHeader<udp_header> bufferHeader(buffer); NetBufferHeaderReader<udp_header> bufferHeader(buffer);
if (bufferHeader.Status() < B_OK) if (bufferHeader.Status() < B_OK)
return bufferHeader.Status(); return bufferHeader.Status();
@@ -464,9 +464,10 @@ split_buffer(net_buffer *from, uint32 offset)
TRACE(("split_buffer(buffer %p -> %p, offset %ld)\n", from, buffer, offset)); TRACE(("split_buffer(buffer %p -> %p, offset %ld)\n", from, buffer, offset));
if (remove_header(from, offset) == B_OK if (trim_data(buffer, offset) == B_OK) {
&& trim_data(buffer, offset) == B_OK) if (remove_header(from, offset) == B_OK)
return buffer; return buffer;
}
free_buffer(buffer); free_buffer(buffer);
return NULL; return NULL;