nfs4: Make networking code IP version agnostic

This commit is contained in:
Pawel Dziepak
2012-07-12 00:22:35 +02:00
parent 367b07d54a
commit 137884e0d2
5 changed files with 156 additions and 87 deletions
@@ -19,6 +19,8 @@
#include <net/dns_resolver.h>
#define NFS4_PORT 2049
#define LAST_FRAGMENT 0x80000000
#define MAX_PACKET_SIZE 65535
@@ -26,17 +28,15 @@
bool
ServerAddress::operator==(const ServerAddress& address)
{
return fAddress == address.fAddress && fPort == address.fPort
&& fProtocol == address.fProtocol;
return memcmp(&fAddress, &address.fAddress, sizeof(fAddress)) == 0
&& fProtocol == address.fProtocol;
}
bool
ServerAddress::operator<(const ServerAddress& address)
{
return fAddress < address.fAddress ||
(fAddress == address.fAddress && fPort < address.fPort) ||
(fAddress == address.fAddress && fPort == address.fPort &&
fProtocol < address.fProtocol);
int compare = memcmp(&fAddress, &address.fAddress, sizeof(fAddress));
return compare < 0 || (compare == 0 && fProtocol < address.fProtocol);
}
@@ -44,21 +44,34 @@ ServerAddress&
ServerAddress::operator=(const ServerAddress& address)
{
fAddress = address.fAddress;
fPort = address.fPort;
fProtocol = address.fProtocol;
return *this;
}
ServerAddress::ServerAddress()
:
fProtocol(0)
{
memset(&fAddress, 0, sizeof(fAddress));
}
status_t
ServerAddress::ResolveName(const char* name, ServerAddress* address)
{
address->fPort = 2049;
address->fProtocol = IPPROTO_UDP;
struct in_addr iaddr;
if (inet_aton(name, &iaddr) != 0) {
address->fAddress = ntohl(iaddr.s_addr);
// getaddrinfo() is very expensive when called from kernel, so we do not
// want to call it unless there is no other choice.
struct sockaddr_in addr;
memset(&addr, 0, sizeof(addr));
if (inet_aton(name, &addr.sin_addr) == 1) {
addr.sin_family = AF_INET;
addr.sin_port = htons(NFS4_PORT);
memcpy(&address->fAddress, &addr, sizeof(addr));
return B_OK;
}
@@ -70,16 +83,20 @@ ServerAddress::ResolveName(const char* name, ServerAddress* address)
addrinfo* current = ai;
while (current != NULL) {
if (current->ai_family == AF_INET) {
sockaddr_in* sin =
reinterpret_cast<sockaddr_in*>(current->ai_addr);
address->fAddress = ntohl(sin->sin_addr.s_addr);
freeaddrinfo(ai);
return B_OK;
memcpy(&address->fAddress, current->ai_addr, sizeof(sockaddr_in));
reinterpret_cast<sockaddr_in*>(&address->fAddress)->sin_port
= htons(NFS4_PORT);
} else if (current->ai_family == AF_INET6) {
memcpy(&address->fAddress, current->ai_addr, sizeof(sockaddr_in6));
reinterpret_cast<sockaddr_in6*>(&address->fAddress)->sin6_port
= htons(NFS4_PORT);
} else {
current = current->ai_next;
continue;
}
current = current->ai_next;
freeaddrinfo(ai);
return B_OK;
}
freeaddrinfo(ai);
@@ -87,27 +104,26 @@ ServerAddress::ResolveName(const char* name, ServerAddress* address)
}
Connection::Connection(const sockaddr_in& address, int protocol)
Connection::Connection(const ServerAddress& address)
:
fWaitCancel(create_sem(0, NULL)),
fSocket(-1),
fProtocol(protocol),
fServerAddress(address)
{
mutex_init(&fSocketLock, NULL);
}
ConnectionStream::ConnectionStream(const sockaddr_in& address, int protocol)
ConnectionStream::ConnectionStream(const ServerAddress& address)
:
Connection(address, protocol)
Connection(address)
{
}
ConnectionPacket::ConnectionPacket(const sockaddr_in& address, int protocol)
ConnectionPacket::ConnectionPacket(const ServerAddress& address)
:
Connection(address, protocol)
Connection(address)
{
}
@@ -124,17 +140,22 @@ Connection::~Connection()
status_t
Connection::GetLocalAddress(ServerAddress* address)
{
struct sockaddr_in saddr;
socklen_t slen = sizeof(saddr);
status_t result = getsockname(fSocket, (struct sockaddr*)&saddr, &slen);
if (result != B_OK)
return result;
address->fProtocol = fServerAddress.fProtocol;
address->fProtocol = fProtocol;
address->fPort = ntohs(saddr.sin_port);
address->fAddress = ntohl(saddr.sin_addr.s_addr);
socklen_t addressSize;
switch (reinterpret_cast<const sockaddr*>(&fServerAddress)->sa_family) {
case AF_INET:
addressSize = sizeof(sockaddr_in);
break;
case AF_INET6:
addressSize = sizeof(sockaddr_in6);
break;
default:
return B_BAD_VALUE;
}
return B_OK;
return getsockname(fSocket,
(struct sockaddr*)&address->fAddress, &addressSize);
}
@@ -315,21 +336,13 @@ ConnectionPacket::Receive(void** _buffer, uint32* _size)
status_t
Connection::Connect(Connection **_connection, const ServerAddress& address)
{
struct sockaddr_in addr;
memset(&addr, 0, sizeof(addr));
addr.sin_len = sizeof(struct sockaddr_in);
addr.sin_family = AF_INET;
addr.sin_addr.s_addr = htonl(address.fAddress);
addr.sin_port = htons(address.fPort);
Connection* conn;
switch (address.fProtocol) {
case IPPROTO_TCP:
conn = new(std::nothrow) ConnectionStream(addr, address.fProtocol);
conn = new(std::nothrow) ConnectionStream(address);
break;
case IPPROTO_UDP:
conn = new(std::nothrow) ConnectionPacket(addr, address.fProtocol);
conn = new(std::nothrow) ConnectionPacket(address);
break;
default:
return B_BAD_VALUE;
@@ -352,12 +365,15 @@ Connection::Connect(Connection **_connection, const ServerAddress& address)
status_t
Connection::Connect()
{
switch (fProtocol) {
const sockaddr& address =
*reinterpret_cast<const sockaddr*>(&fServerAddress);
switch (fServerAddress.fProtocol) {
case IPPROTO_TCP:
fSocket = socket(AF_INET, SOCK_STREAM, IPPROTO_TCP);
fSocket = socket(address.sa_family, SOCK_STREAM, IPPROTO_TCP);
break;
case IPPROTO_UDP:
fSocket = socket(AF_INET, SOCK_DGRAM, IPPROTO_UDP);
fSocket = socket(address.sa_family, SOCK_DGRAM, IPPROTO_UDP);
break;
default:
return B_BAD_VALUE;
@@ -365,8 +381,19 @@ Connection::Connect()
if (fSocket < 0)
return errno;
status_t result = connect(fSocket, (struct sockaddr*)&fServerAddress,
fServerAddress.sin_len);
socklen_t addressSize;
switch (address.sa_family) {
case AF_INET:
addressSize = sizeof(sockaddr_in);
break;
case AF_INET6:
addressSize = sizeof(sockaddr_in6);
break;
default:
return B_BAD_VALUE;
}
status_t result = connect(fSocket, &address, addressSize);
if (result < 0) {
result = errno;
close(fSocket);
@@ -16,8 +16,7 @@
struct ServerAddress {
uint32 fAddress;
uint16 fPort;
sockaddr_storage fAddress;
int fProtocol;
bool operator==(const ServerAddress& address);
@@ -25,6 +24,8 @@ struct ServerAddress {
ServerAddress& operator=(const ServerAddress& address);
ServerAddress();
static status_t ResolveName(const char* name,
ServerAddress* address);
};
@@ -44,22 +45,19 @@ public:
void Disconnect();
protected:
Connection(const sockaddr_in& address,
int protocol);
Connection(const ServerAddress& address);
status_t Connect();
sem_id fWaitCancel;
int fSocket;
mutex fSocketLock;
const int fProtocol;
const sockaddr_in fServerAddress;
const ServerAddress fServerAddress;
};
class ConnectionStream : public Connection {
public:
ConnectionStream(const sockaddr_in& address,
int protocol);
ConnectionStream(const ServerAddress& address);
virtual status_t Send(const void* buffer, uint32 size);
virtual status_t Receive(void** buffer, uint32* size);
@@ -67,8 +65,7 @@ public:
class ConnectionPacket : public Connection {
public:
ConnectionPacket(const sockaddr_in& address,
int protocol);
ConnectionPacket(const ServerAddress& address);
virtual status_t Send(const void* buffer, uint32 size);
virtual status_t Receive(void** buffer, uint32* size);
@@ -9,6 +9,7 @@
#include "RequestBuilder.h"
#include <errno.h>
#include <string.h>
#include "Cookie.h"
@@ -592,7 +593,7 @@ RequestBuilder::SetAttr(const uint32* id, uint32 stateSeq, AttrValue* attr,
status_t
RequestBuilder::SetClientID(const RPC::Server* serv)
RequestBuilder::SetClientID(const RPC::Server* server)
{
if (fProcedure != ProcCompound)
return B_BAD_VALUE;
@@ -604,19 +605,9 @@ RequestBuilder::SetClientID(const RPC::Server* serv)
verifier = verifier << 32 | rand();
fRequest->Stream().AddUHyper(verifier);
char id[128] = "HAIKU:kernel:";
int pos = strlen(id);
*(uint32*)(id + pos) = serv->ID().fAddress;
pos += sizeof(uint32);
*(uint16*)(id + pos) = serv->ID().fPort;
pos += sizeof(uint16);
*(uint16*)(id + pos) = serv->ID().fProtocol;
pos += sizeof(uint16);
*(uint32*)(id + pos) = serv->LocalID().fAddress;
pos += sizeof(uint32);
fRequest->Stream().AddOpaque(id, pos);
status_t result = _GenerateClientId(fRequest->Stream(), server);
if (result != B_OK)
return result;
// Callbacks are currently not supported
fRequest->Stream().AddUInt(0);
@@ -630,6 +621,64 @@ RequestBuilder::SetClientID(const RPC::Server* serv)
}
status_t
RequestBuilder::_GenerateClientId(XDR::WriteStream& stream,
const RPC::Server* server)
{
char id[512] = "HAIKU:kernel:";
int pos = strlen(id);
const sockaddr* remoteAddress =
reinterpret_cast<const sockaddr*>(&server->ID().fAddress);
ServerAddress local = server->LocalID();
const sockaddr* localAddress = reinterpret_cast<sockaddr*>(&local.fAddress);
const sockaddr_in* address4;
const sockaddr_in6* address6;
switch (remoteAddress->sa_family) {
case AF_INET:
address4 = reinterpret_cast<const sockaddr_in*>(remoteAddress);
memcpy(id + pos, &address4->sin_addr, sizeof(address4->sin_addr));
pos += sizeof(address4->sin_addr);
memcpy(id + pos,
&reinterpret_cast<const sockaddr_in*>(localAddress)->sin_addr,
sizeof(address4->sin_addr));
pos += sizeof(address4->sin_addr);
*(uint16*)(id + pos) = address4->sin_port;
break;
case AF_INET6:
address6 = reinterpret_cast<const sockaddr_in6*>(remoteAddress);
memcpy(id + pos, &address6->sin6_addr, sizeof(address6->sin6_addr));
pos += sizeof(address6->sin6_addr);
memcpy(id + pos,
&reinterpret_cast<const sockaddr_in6*>(localAddress)->sin6_addr,
sizeof(address6->sin6_addr));
pos += sizeof(address6->sin6_addr);
*(uint16*)(id + pos) = address6->sin6_port;
break;
default:
return B_BAD_VALUE;
}
pos += sizeof(uint16);
*(uint16*)(id + pos) = server->ID().fProtocol;
pos += sizeof(uint16);
stream.AddOpaque(id, pos);
return B_OK;
}
status_t
RequestBuilder::SetClientIDConfirm(uint64 id, uint64 ver)
{
@@ -68,7 +68,7 @@ public:
status_t SaveFH();
status_t SetAttr(const uint32* id, uint32 stateSeq,
AttrValue* attr, uint32 count);
status_t SetClientID(const RPC::Server* serv);
status_t SetClientID(const RPC::Server* server);
status_t SetClientIDConfirm(uint64 id, uint64 ver);
status_t Verify(AttrValue* attr, uint32 count);
status_t Write(const uint32* id, uint32 stateSeq,
@@ -85,6 +85,8 @@ private:
void _GenerateLockOwner(XDR::WriteStream& stream,
OpenFileCookie* cookie,
LockOwner* owner);
status_t _GenerateClientId(XDR::WriteStream& stream,
const RPC::Server* server);
void _EncodeAttrs(XDR::WriteStream& stream,
AttrValue* attr, uint32 count);
@@ -39,8 +39,10 @@ CreateNFS4Server(RPC::Server* serv)
}
// TODO: IPv6 address will cause problems
static status_t
ParseArguments(const char* _args, uint32* _ip, char* _path)
ParseArguments(const char* _args, ServerAddress* address, char* _path)
{
if (_args == NULL)
return B_BAD_VALUE;
@@ -53,13 +55,10 @@ ParseArguments(const char* _args, uint32* _ip, char* _path)
}
*path++ = '\0';
ServerAddress addr;
status_t result = ServerAddress::ResolveName(args, &addr);
status_t result = ServerAddress::ResolveName(args, address);
if (result != B_OK)
return result;
*_ip = addr.fAddress;
_path[255] = '\0';
strncpy(_path, path, 255);
@@ -74,19 +73,14 @@ nfs4_mount(fs_volume* volume, const char* device, uint32 flags,
{
status_t result;
uint32 ip;
ServerAddress address;
char path[256];
result = ParseArguments(args, &ip, path);
result = ParseArguments(args, &address, path);
if (result != B_OK)
return result;
ServerAddress id;
id.fAddress = ip;
id.fPort = 2049;
id.fProtocol = IPPROTO_UDP;
RPC::Server *server;
result = gRPCServerManager->Acquire(&server, id, CreateNFS4Server);
result = gRPCServerManager->Acquire(&server, address, CreateNFS4Server);
if (result != B_OK)
return result;