diff --git a/headers/private/net/NetBufferUtilities.h b/headers/private/net/NetBufferUtilities.h index 290d496c95..7d8b467c02 100644 --- a/headers/private/net/NetBufferUtilities.h +++ b/headers/private/net/NetBufferUtilities.h @@ -16,94 +16,26 @@ class NetBufferModuleGetter { static net_buffer_module_info *Get() { return gBufferModule; } }; -//! A class to retrieve and remove a header from a buffer -template class NetBufferHeader { +//! A class to access a field safely across node boundaries +template +class NetBufferFieldReader { public: - NetBufferHeader(net_buffer *buffer) + NetBufferFieldReader(net_buffer *buffer) : - fBuffer(buffer) + fBuffer(buffer), + fStatus(B_BAD_VALUE) { - } - - ~NetBufferHeader() - { - Remove(); - } - - 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; + if ((Offset + sizeof(Type)) <= buffer->size) { + fStatus = Module::Get()->direct_access(fBuffer, Offset, + sizeof(Type), (void **)&fData); + if (fStatus != B_OK) { + fData = NULL; + fStatus = Module::Get()->read(fBuffer, Offset, + &fDataBuffer, sizeof(Type)); + } } } - void - Detach() - { - fBuffer = NULL; - } - - private: - net_buffer *fBuffer; - Type fDataBuffer; -}; - -//! A class to access a header safely across data node boundaries -template -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() { @@ -119,16 +51,33 @@ class NetBufferSafeHeader { return fDataBuffer; } - void - Detach() + Type * + operator->() { + return &Data(); + } + + Type & + operator*() + { + return Data(); + } + + void + Sync() + { + if (fBuffer == NULL) + return; + if (fData == NULL) - Module::Get()->write(fBuffer, 0, &fDataBuffer, sizeof(Type)); + Module::Get()->write(fBuffer, Offset, &fDataBuffer, + sizeof(Type)); + fBuffer = NULL; } protected: - NetBufferSafeHeader() {} + NetBufferFieldReader() {} net_buffer *fBuffer; status_t fStatus; @@ -136,9 +85,58 @@ class NetBufferSafeHeader { Type fDataBuffer; }; +template +class NetBufferField : public NetBufferFieldReader { + public: + NetBufferField(net_buffer *buffer) + : NetBufferFieldReader(buffer) + {} + + ~NetBufferField() + { + Sync(); + } +}; + +template +class NetBufferHeaderReader : public NetBufferFieldReader { + public: + NetBufferHeaderReader(net_buffer *buffer) + : NetBufferFieldReader(buffer) + {} + + void + Remove() + { + Remove(sizeof(Type)); + } + + void + Remove(size_t bytes) + { + if (fBuffer != NULL) { + Module::Get()->remove_header(fBuffer, bytes); + fBuffer = NULL; + } + } +}; + +template +class NetBufferHeaderRemover : public NetBufferHeaderReader { + public: + NetBufferHeaderRemover(net_buffer *buffer) + : NetBufferHeaderReader(buffer) + {} + + ~NetBufferHeaderRemover() + { + Remove(); + } +}; + //! A class to add a header to a buffer template -class NetBufferPrepend : public NetBufferSafeHeader { +class NetBufferPrepend : public NetBufferFieldReader { public: NetBufferPrepend(net_buffer *buffer, size_t size = 0) { @@ -150,6 +148,11 @@ class NetBufferPrepend : public NetBufferSafeHeader { fStatus = Module::Get()->prepend_size(buffer, size, (void **)&fData); } + + ~NetBufferPrepend() + { + Sync(); + } }; #endif // NET_BUFFER_UTILITIES_H diff --git a/src/add-ons/kernel/network/datalink_protocols/arp/arp.cpp b/src/add-ons/kernel/network/datalink_protocols/arp/arp.cpp index 0c8a90ae08..7298935c12 100644 --- a/src/add-ons/kernel/network/datalink_protocols/arp/arp.cpp +++ b/src/add-ons/kernel/network/datalink_protocols/arp/arp.cpp @@ -353,7 +353,7 @@ arp_receive(void *cookie, net_buffer *buffer) { TRACE(("ARP receive\n")); - NetBufferHeader bufferHeader(buffer); + NetBufferHeaderReader bufferHeader(buffer); if (bufferHeader.Status() < B_OK) return bufferHeader.Status(); @@ -383,8 +383,6 @@ arp_receive(void *cookie, net_buffer *buffer) || header.protocol_length != sizeof(in_addr_t)) return B_BAD_DATA; - bufferHeader.Detach(); - // handle packet switch (opcode) { diff --git a/src/add-ons/kernel/network/datalink_protocols/ethernet_frame/ethernet_frame.cpp b/src/add-ons/kernel/network/datalink_protocols/ethernet_frame/ethernet_frame.cpp index 002aa99dfd..9668e6fbf7 100644 --- a/src/add-ons/kernel/network/datalink_protocols/ethernet_frame/ethernet_frame.cpp +++ b/src/add-ons/kernel/network/datalink_protocols/ethernet_frame/ethernet_frame.cpp @@ -38,7 +38,7 @@ ethernet_deframe(net_device *device, net_buffer *buffer) { //dprintf("asked to deframe buffer for device %s\n", device->name); - NetBufferHeader bufferHeader(buffer); + NetBufferHeaderRemover bufferHeader(buffer); if (bufferHeader.Status() < B_OK) return bufferHeader.Status(); @@ -138,7 +138,7 @@ ethernet_frame_send_data(net_datalink_protocol *protocol, else 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 return protocol->next->module->send_data(protocol->next, buffer); diff --git a/src/add-ons/kernel/network/protocols/icmp/icmp.cpp b/src/add-ons/kernel/network/protocols/icmp/icmp.cpp index 8fd6a64925..7299466831 100644 --- a/src/add-ons/kernel/network/protocols/icmp/icmp.cpp +++ b/src/add-ons/kernel/network/protocols/icmp/icmp.cpp @@ -48,6 +48,8 @@ struct icmp_header { }; }; +typedef NetBufferField ICMPChecksumField; + #define ICMP_TYPE_ECHO_REPLY 0 #define ICMP_TYPE_UNREACH 3 #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)); - NetBufferHeader bufferHeader(buffer); + NetBufferHeaderReader bufferHeader(buffer); if (bufferHeader.Status() < B_OK) return bufferHeader.Status(); 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, ntohs(header.checksum))); @@ -249,19 +249,18 @@ icmp_receive_data(net_buffer *buffer) memcpy(&reply->destination, &buffer->source, buffer->source.ss_len); // There already is an ICMP header, and we'll reuse it - icmp_header *header; - 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); - } + NetBufferHeaderReader header(reply); - if (status == B_OK) - status = domain->module->send_data(NULL, reply); + header->type = ICMP_TYPE_ECHO_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) { gBufferModule->free(reply); return status; diff --git a/src/add-ons/kernel/network/protocols/ipv4/ipv4.cpp b/src/add-ons/kernel/network/protocols/ipv4/ipv4.cpp index 3afa5df732..c56e792585 100644 --- a/src/add-ons/kernel/network/protocols/ipv4/ipv4.cpp +++ b/src/add-ons/kernel/network/protocols/ipv4/ipv4.cpp @@ -74,6 +74,8 @@ struct ipv4_header { typedef DoublyLinkedList > FragmentList; +typedef NetBufferField IPChecksumField; + struct ipv4_packet_key { in_addr_t source; 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", buffer->size, mtu)); - NetBufferHeader bufferHeader(buffer); - if (bufferHeader.Status() < B_OK) - return bufferHeader.Status(); + NetBufferHeaderReader originalHeader(buffer); + if (originalHeader.Status() < B_OK) + return originalHeader.Status(); - ipv4_header *header = &bufferHeader.Data(); - bufferHeader.Detach(); - - uint16 headerLength = header->HeaderLength(); + uint16 headerLength = originalHeader->HeaderLength(); uint32 bytesLeft = buffer->size - headerLength; uint32 fragmentOffset = 0; status_t status = B_OK; @@ -576,9 +575,10 @@ send_fragments(ipv4_protocol *protocol, struct net_route *route, if (headerBuffer == NULL) return B_NO_MEMORY; - bufferHeader.SetTo(headerBuffer); - header = &bufferHeader.Data(); - bufferHeader.Detach(); + // TODO we need to make sure ipv4_header is contiguous or + // use another construct. + NetBufferHeaderReader bufferHeader(headerBuffer); + ipv4_header *header = &bufferHeader.Data(); // adapt MTU to be a multiple of 8 (fragment offsets can only be specified this way) 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 - @@ -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.checksum = gBufferModule->checksum(buffer, 0, - sizeof(ipv4_header), true); - //dump_ipv4_header(header); - - bufferHeader.Detach(); + bufferHeader.Sync(); // make sure the IP-header is already written to the // buffer at this point + + update_checksum(buffer); + //dump_ipv4_header(header); + } else { // if IP_HDRINCL, check if the source address is set - NetBufferHeader bufferHeader(buffer); - if (bufferHeader.Status() < B_OK) - return bufferHeader.Status(); + NetBufferHeaderReader header(buffer); + if (header.Status() < B_OK) + return header.Status(); - ipv4_header &header = bufferHeader.Data(); - if (header.source == 0) { - header.source = source.sin_addr.s_addr; - header.checksum = gBufferModule->checksum(buffer, - sizeof(ipv4_header), sizeof(ipv4_header), true); + if (header->source == 0) { + header->source = source.sin_addr.s_addr; + header->checksum = 0; + + header.Sync(); + + update_checksum(buffer); } - - bufferHeader.Detach(); } 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)); - NetBufferHeader bufferHeader(buffer); + NetBufferHeaderReader bufferHeader(buffer); if (bufferHeader.Status() < B_OK) return bufferHeader.Status(); ipv4_header &header = bufferHeader.Data(); - bufferHeader.Detach(); //dump_ipv4_header(header); if (header.version != IP_VERSION) diff --git a/src/add-ons/kernel/network/protocols/tcp/tcp.cpp b/src/add-ons/kernel/network/protocols/tcp/tcp.cpp index e9fb65f8b0..2ea58b91e1 100644 --- a/src/add-ons/kernel/network/protocols/tcp/tcp.cpp +++ b/src/add-ons/kernel/network/protocols/tcp/tcp.cpp @@ -39,6 +39,9 @@ #endif +typedef NetBufferField TCPChecksumField; + + net_domain *gDomain; net_address_module_info *gAddressModule; 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 // not have a contiguous buffer. - bufferHeader.Detach(); + bufferHeader.Sync(); if (optionsLength > 0) 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) << Checksum::BufferHelper(buffer, gBufferModule); - // we are pretty sure the header is there. - NetBufferSafeHeader headerRef(buffer); - headerRef.Data().checksum = checksum; + TCPChecksumField checksumField(buffer); + *checksumField = checksum; return B_OK; } @@ -507,7 +509,7 @@ tcp_receive_data(net_buffer *buffer) if (gDomain == NULL && set_domain(buffer->interface) != B_OK) return B_ERROR; - NetBufferHeader bufferHeader(buffer); + NetBufferHeaderReader bufferHeader(buffer); if (bufferHeader.Status() < B_OK) return bufferHeader.Status(); diff --git a/src/add-ons/kernel/network/protocols/udp/udp.cpp b/src/add-ons/kernel/network/protocols/udp/udp.cpp index e64d8bc117..6f73a4ab0b 100644 --- a/src/add-ons/kernel/network/protocols/udp/udp.cpp +++ b/src/add-ons/kernel/network/protocols/udp/udp.cpp @@ -440,7 +440,7 @@ UdpEndpointManager::DemuxIncomingBuffer(net_buffer *buffer) status_t UdpEndpointManager::ReceiveData(net_buffer *buffer) { - NetBufferHeader bufferHeader(buffer); + NetBufferHeaderReader bufferHeader(buffer); if (bufferHeader.Status() < B_OK) return bufferHeader.Status(); diff --git a/src/add-ons/kernel/network/stack/net_buffer.cpp b/src/add-ons/kernel/network/stack/net_buffer.cpp index 9f0163bef9..e9fd2f8bee 100644 --- a/src/add-ons/kernel/network/stack/net_buffer.cpp +++ b/src/add-ons/kernel/network/stack/net_buffer.cpp @@ -464,9 +464,10 @@ split_buffer(net_buffer *from, uint32 offset) TRACE(("split_buffer(buffer %p -> %p, offset %ld)\n", from, buffer, offset)); - if (remove_header(from, offset) == B_OK - && trim_data(buffer, offset) == B_OK) - return buffer; + if (trim_data(buffer, offset) == B_OK) { + if (remove_header(from, offset) == B_OK) + return buffer; + } free_buffer(buffer); return NULL;