diff --git a/cmd/gomodfs/portmapper.go b/cmd/gomodfs/portmapper.go index 6e09133..a052f2c 100644 --- a/cmd/gomodfs/portmapper.go +++ b/cmd/gomodfs/portmapper.go @@ -1,258 +1,34 @@ package main import ( - "bufio" - "bytes" - "encoding/binary" - "fmt" - "io" "log" "net" - "github.com/tailscale/gomodfs/temp-dev-fork/willscott/go-nfs" + "github.com/tailscale/gomodfs/portmap" ) -func startPortmapper() error { - tcpLn, err := net.Listen("tcp", rpcBindAddr) - if err != nil { - return err - } - go runPortMapperTCP(tcpLn) - - udpAddr := mustResolveUDP(rpcBindAddr) - udpConn, err := net.ListenUDP("udp", udpAddr) - if err != nil { - return err - } - go runPortMapperUDP(udpConn) - return nil -} - -func runPortMapperTCP(ln net.Listener) { - defer ln.Close() - - log.Printf("portmap-static: listening on TCP %s", ln.Addr()) - - for { - conn, err := ln.Accept() - if err != nil { - log.Printf("runPortMapperTCP.Accept: %v", err) - return - } - handleRPCBindConn(conn) - } -} - -func runPortMapperUDP(conn *net.UDPConn) { - defer conn.Close() - log.Printf("portmap-static: listening on UDP %s", conn.LocalAddr()) - - buf := make([]byte, 64*1024) - for { - n, addr, err := conn.ReadFromUDP(buf) - if err != nil { - log.Printf("runPortMapperUDP: %v", err) - return - } - handleUDPPacket(conn, addr, buf[:n]) - } -} - const ( rpcBindAddr = ":111" staticPort = 2049 - - rpcCall = 0 - rpcReply = 1 - - replyMsgAccepted = 0 - authNull = 0 - acceptSuccess = 0 - - portmapProg = 100000 - portmapVers = 2 - - pmapProcNull = 0 - pmapprocGetport = 3 - - protTCP = 6 - protUDP = 17 - - nfsProg = 100003 - mountProg = 100005 - nlmProg = 100021 ) -func mustResolveUDP(addr string) *net.UDPAddr { - ua, err := net.ResolveUDPAddr("udp", addr) +func startPortmapper() error { + tcpLn, err := net.Listen("tcp", rpcBindAddr) if err != nil { - log.Fatal(err) - } - return ua -} - -func handleRPCBindConn(c net.Conn) { - log.Printf("rpcbind: accepted TCP connection from %s", c.RemoteAddr()) - defer c.Close() - br := bufio.NewReader(c) - for { - rec, err := nfs.ReadRPCRecord(br) - log.Printf("rpcbind: TCP conn %s sent rec % 02x", c.RemoteAddr(), rec) - if err != nil { - if err != io.EOF { - log.Printf("rpcbind: readRPCRecord: %v", err) - } - return - } - handleTCPPacket(c, rec) - } -} - -func handleTCPPacket(c net.Conn, data []byte) { - log.Printf("rpcbind: received %d bytes from TCP %s: % 02x", len(data), c.RemoteAddr(), data) - handlePacket(data, func(resp []byte) error { - var buf [4]byte - const lastFragMask = 1 << 31 - binary.BigEndian.PutUint32(buf[:], uint32(len(resp))|lastFragMask) - _, err := fmt.Fprintf(c, "%s%s", buf[:], resp) return err - }) -} - -func handleUDPPacket(conn *net.UDPConn, addr *net.UDPAddr, data []byte) { - log.Printf("rpcbind: received %d bytes from UDP %s: % 02x", len(data), addr, data) - handlePacket(data, func(resp []byte) error { - _, err := conn.WriteToUDP(resp, addr) - return err - }) -} - -func handlePacket(data []byte, sendResp func([]byte) error) { - r := bytes.NewReader(data) - - var xid, mtype, rpcvers, prog, vers, proc uint32 - if err := binary.Read(r, binary.BigEndian, &xid); err != nil { - return - } - if err := binary.Read(r, binary.BigEndian, &mtype); err != nil || mtype != rpcCall { - return - } - if err := binary.Read(r, binary.BigEndian, &rpcvers); err != nil || rpcvers != 2 { - return - } - if err := binary.Read(r, binary.BigEndian, &prog); err != nil || prog != portmapProg { - log.Printf("portmap-static: unknown program %02x; want %02x", prog, portmapProg) - return - } - if err := binary.Read(r, binary.BigEndian, &vers); err != nil || vers != portmapVers { - return - } - if err := binary.Read(r, binary.BigEndian, &proc); err != nil { - return - } - - log.Printf("portmap-static: RPC call rpcver=%v xid=%d prog=%d vers=%d proc=%d", rpcvers, xid, prog, vers, proc) - - if !skipOpaqueAuth(r) { // cred - log.Printf("portmap-static: failed to skip cred") - return - } - if !skipOpaqueAuth(r) { // verf - log.Printf("portmap-static: failed to skip cred") - return } - - if proc == pmapProcNull { - // Reply with empty response - var buf bytes.Buffer - writeReplySuccess(&buf, xid) - resp := buf.Bytes() - log.Printf("rpcbind: replying to null ping") - if err := sendResp(resp); err != nil { - log.Printf("rpcbind: sendResp: %v", err) - } - return - } - - if proc != pmapprocGetport { - log.Printf("portmap-static: unknown procedure %08x; want %08x", proc, pmapprocGetport) - return - } - - // Portmap v2 pmap args: prog, vers, prot, port - var argProg, argVers, argProt, argPort uint32 - if err := binary.Read(r, binary.BigEndian, &argProg); err != nil { - return - } - if err := binary.Read(r, binary.BigEndian, &argVers); err != nil { - return - } - if err := binary.Read(r, binary.BigEndian, &argProt); err != nil { - return - } - if err := binary.Read(r, binary.BigEndian, &argPort); err != nil { - return - } - _ = argVers - _ = argProt - _ = argPort - - log.Printf("rpcbind: GetPort prog=%d vers=%d prot=%d port=%d", argProg, argVers, argProt, argPort) - - var port uint32 - switch argProt { - case protUDP: - port = 0 // don't advertise UDP support - case protTCP: - switch argProg { - case nfsProg, mountProg, nlmProg: - port = staticPort - } - } - - resp := buildPortReply(xid, port) - log.Printf("rpcbind: replying with port %d: % 02x", port, resp) - if err := sendResp(resp); err != nil { - log.Printf("rpcbind: sendResp: %v", err) - } -} - -func skipOpaqueAuth(r *bytes.Reader) bool { - var flavor, length uint32 - if err := binary.Read(r, binary.BigEndian, &flavor); err != nil { - return false - } - if err := binary.Read(r, binary.BigEndian, &length); err != nil { - return false - } - _ = flavor - pad := (length + 3) &^ 3 - if pad == 0 { - return true + udpConn, err := net.ListenPacket("udp", rpcBindAddr) + if err != nil { + tcpLn.Close() + return err } - _, err := r.Seek(int64(pad), 1) - return err == nil -} - -func writeReplySuccess(buf *bytes.Buffer, xid uint32) { - // rpc_msg - _ = binary.Write(buf, binary.BigEndian, xid) // xid - _ = binary.Write(buf, binary.BigEndian, uint32(rpcReply)) // mtype = REPLY - _ = binary.Write(buf, binary.BigEndian, uint32(replyMsgAccepted)) // reply_stat = MSG_ACCEPTED - - // verf: AUTH_NULL, length 0 - _ = binary.Write(buf, binary.BigEndian, uint32(authNull)) - _ = binary.Write(buf, binary.BigEndian, uint32(0)) - - _ = binary.Write(buf, binary.BigEndian, uint32(acceptSuccess)) -} - -func buildPortReply(xid, port uint32) []byte { - var buf bytes.Buffer - writeReplySuccess(&buf, xid) - - // result: uint32 port - _ = binary.Write(&buf, binary.BigEndian, port) - - return buf.Bytes() + pm := &portmap.Server{Port: staticPort, Logf: log.Printf} + log.Printf("portmap-static: listening on TCP and UDP %s", rpcBindAddr) + go func() { + log.Printf("portmap-static: TCP: %v", pm.ServeTCP(tcpLn)) + }() + go func() { + log.Printf("portmap-static: UDP: %v", pm.ServeUDP(udpConn)) + }() + return nil } diff --git a/portmap/portmap.go b/portmap/portmap.go new file mode 100644 index 0000000..4e57608 --- /dev/null +++ b/portmap/portmap.go @@ -0,0 +1,202 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + +// Package portmap implements a static ONC RPC portmapper (rpcbind version 2, +// RFC 1833) that tells NFS clients which port gomodfs's NFS server is on. +// +// Linux and macOS NFS clients can be told the NFS and MOUNT ports directly, +// but the Windows NFS client always asks the portmapper on port 111 first, so +// serving gomodfs to Windows clients requires one. +package portmap + +import ( + "bufio" + "encoding/binary" + "fmt" + "io" + "net" +) + +const ( + rpcCall = 0 + rpcReply = 1 + rpcVersion = 2 + replyMsgAccepted = 0 + acceptSuccess = 0 + + portmapProg = 100000 + portmapVers = 2 + pmapProcNull = 0 + pmapProcGetport = 3 + ipprotoTCP = 6 + + nfsProg = 100003 + mountProg = 100005 + nlmProg = 100021 +) + +// maxRecord bounds the size of an RPC call that a Server reads. +const maxRecord = 64 << 10 + +// Server is a static portmapper. It answers GETPORT queries for the NFS, +// MOUNT, and NLM programs over TCP with Port, and all other GETPORT queries, +// including any over UDP, with port 0, meaning "not registered". +type Server struct { + // Port is the TCP port of the NFS server. + Port int + + // Logf, if non-nil, logs each query and reply. + Logf func(format string, args ...any) +} + +func (s *Server) logf(format string, args ...any) { + if s.Logf != nil { + s.Logf(format, args...) + } +} + +// ServeTCP serves the portmapper on ln until ln.Accept fails, and returns +// that error. +func (s *Server) ServeTCP(ln net.Listener) error { + for { + c, err := ln.Accept() + if err != nil { + return err + } + go s.serveConn(c) + } +} + +// ServeUDP serves the portmapper on pc until reading from it fails, and +// returns that error. +func (s *Server) ServeUDP(pc net.PacketConn) error { + buf := make([]byte, maxRecord) + for { + n, addr, err := pc.ReadFrom(buf) + if err != nil { + return err + } + if resp := s.reply(buf[:n]); resp != nil { + if _, err := pc.WriteTo(resp, addr); err != nil { + s.logf("portmap: replying to %v: %v", addr, err) + } + } + } +} + +// serveConn answers RPC calls on c, which uses RPC record marking (RFC 5531, +// section 11): each record is a series of fragments that each start with a +// 4-byte header holding the fragment length and a last-fragment bit. +func (s *Server) serveConn(c net.Conn) { + defer c.Close() + br := bufio.NewReader(c) + for { + rec, err := readRecord(br) + if err != nil { + if err != io.EOF { + s.logf("portmap: reading from %v: %v", c.RemoteAddr(), err) + } + return + } + resp := s.reply(rec) + if resp == nil { + return + } + out := binary.BigEndian.AppendUint32(nil, uint32(len(resp))|1<<31) + if _, err := c.Write(append(out, resp...)); err != nil { + return + } + } +} + +func readRecord(br *bufio.Reader) ([]byte, error) { + var rec []byte + for { + var hdr [4]byte + if _, err := io.ReadFull(br, hdr[:]); err != nil { + return nil, err + } + h := binary.BigEndian.Uint32(hdr[:]) + n := int(h &^ (1 << 31)) + if len(rec)+n > maxRecord { + return nil, fmt.Errorf("record larger than %d bytes", maxRecord) + } + start := len(rec) + rec = append(rec, make([]byte, n)...) + if _, err := io.ReadFull(br, rec[start:]); err != nil { + return nil, err + } + if h&(1<<31) != 0 { + return rec, nil + } + } +} + +// xdrReader reads big-endian uint32s from b, recording any short read in bad. +type xdrReader struct { + b []byte + bad bool +} + +func (r *xdrReader) u32() uint32 { + if len(r.b) < 4 { + r.bad = true + return 0 + } + v := binary.BigEndian.Uint32(r.b) + r.b = r.b[4:] + return v +} + +// skipOpaqueAuth skips an RPC opaque_auth: a flavor and padded opaque body. +func (r *xdrReader) skipOpaqueAuth() { + r.u32() // flavor + n := (uint64(r.u32()) + 3) &^ 3 + if uint64(len(r.b)) < n { + r.bad = true + return + } + r.b = r.b[n:] +} + +// reply returns the reply to the RPC call in msg, or nil if msg isn't a +// portmapper call that s answers. +func (s *Server) reply(msg []byte) []byte { + r := &xdrReader{b: msg} + xid := r.u32() + if r.u32() != rpcCall || r.u32() != rpcVersion || r.u32() != portmapProg || r.u32() != portmapVers { + return nil + } + proc := r.u32() + r.skipOpaqueAuth() // credentials + r.skipOpaqueAuth() // verifier + var prog, vers, prot, port uint32 + switch proc { + case pmapProcNull: + case pmapProcGetport: + prog, vers, prot = r.u32(), r.u32(), r.u32() + if prot == ipprotoTCP && (prog == nfsProg || prog == mountProg || prog == nlmProg) { + port = uint32(s.Port) + } + default: + s.logf("portmap: ignoring call to procedure %d", proc) + return nil + } + if r.bad { + return nil + } + + if proc == pmapProcGetport { + s.logf("portmap: GETPORT prog=%d vers=%d prot=%d: port %d", prog, vers, prot, port) + } + b := binary.BigEndian.AppendUint32(nil, xid) + b = binary.BigEndian.AppendUint32(b, rpcReply) + b = binary.BigEndian.AppendUint32(b, replyMsgAccepted) + b = binary.BigEndian.AppendUint32(b, 0) // verifier flavor: AUTH_NONE + b = binary.BigEndian.AppendUint32(b, 0) // verifier length + b = binary.BigEndian.AppendUint32(b, acceptSuccess) + if proc == pmapProcGetport { + b = binary.BigEndian.AppendUint32(b, port) + } + return b +} diff --git a/portmap/portmap_test.go b/portmap/portmap_test.go new file mode 100644 index 0000000..49d8933 --- /dev/null +++ b/portmap/portmap_test.go @@ -0,0 +1,139 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + +package portmap + +import ( + "encoding/binary" + "io" + "net" + "testing" +) + +// call returns a portmapper call for procedure proc, with GETPORT arguments +// for prog and prot. It has an AUTH_UNIX credential, as Windows sends, to check +// that the credential is skipped correctly. +func call(xid, proc, prog, prot uint32) []byte { + b := binary.BigEndian.AppendUint32(nil, xid) + for _, v := range []uint32{rpcCall, rpcVersion, portmapProg, portmapVers, proc} { + b = binary.BigEndian.AppendUint32(b, v) + } + cred := []byte("unix credential!") // 16 bytes, already padded + b = binary.BigEndian.AppendUint32(b, 1) // AUTH_UNIX + b = binary.BigEndian.AppendUint32(b, uint32(len(cred))) + b = append(b, cred...) + b = binary.BigEndian.AppendUint32(b, 0) // verifier: AUTH_NONE + b = binary.BigEndian.AppendUint32(b, 0) + if proc == pmapProcGetport { + for _, v := range []uint32{prog, 3, prot, 0} { + b = binary.BigEndian.AppendUint32(b, v) + } + } + return b +} + +func TestReply(t *testing.T) { + s := &Server{Port: 2050, Logf: t.Logf} + for _, tc := range []struct { + name string + msg []byte + wantNil bool + wantPort int // -1 for a reply with no port (NULL) + }{ + {"nfs-tcp", call(7, pmapProcGetport, nfsProg, ipprotoTCP), false, 2050}, + {"mount-tcp", call(7, pmapProcGetport, mountProg, ipprotoTCP), false, 2050}, + {"nlm-tcp", call(7, pmapProcGetport, nlmProg, ipprotoTCP), false, 2050}, + {"nfs-udp", call(7, pmapProcGetport, nfsProg, 17), false, 0}, + {"other-prog", call(7, pmapProcGetport, 100024, ipprotoTCP), false, 0}, + {"null", call(7, pmapProcNull, 0, 0), false, -1}, + {"other-proc", call(7, 4, 0, 0), true, 0}, + {"truncated", call(7, pmapProcGetport, nfsProg, ipprotoTCP)[:50], true, 0}, + {"empty", nil, true, 0}, + } { + t.Run(tc.name, func(t *testing.T) { + got := s.reply(tc.msg) + if tc.wantNil { + if got != nil { + t.Fatalf("got reply % x; want none", got) + } + return + } + wantLen := 24 + if tc.wantPort >= 0 { + wantLen += 4 + } + if len(got) != wantLen { + t.Fatalf("reply is %d bytes; want %d", len(got), wantLen) + } + if xid := binary.BigEndian.Uint32(got); xid != 7 { + t.Errorf("xid = %d; want 7", xid) + } + if mtype := binary.BigEndian.Uint32(got[4:]); mtype != rpcReply { + t.Errorf("message type = %d; want %d", mtype, rpcReply) + } + if tc.wantPort >= 0 { + if port := binary.BigEndian.Uint32(got[24:]); port != uint32(tc.wantPort) { + t.Errorf("port = %d; want %d", port, tc.wantPort) + } + } + }) + } +} + +func TestServeConn(t *testing.T) { + s := &Server{Port: 2050} + c1, c2 := net.Pipe() + defer c1.Close() + go s.serveConn(c2) + + // Send the call split across two record fragments. + msg := call(9, pmapProcGetport, nfsProg, ipprotoTCP) + go func() { + c1.Write(binary.BigEndian.AppendUint32(nil, 10)) + c1.Write(msg[:10]) + c1.Write(binary.BigEndian.AppendUint32(nil, uint32(len(msg)-10)|1<<31)) + c1.Write(msg[10:]) + }() + + var hdr [4]byte + if _, err := io.ReadFull(c1, hdr[:]); err != nil { + t.Fatal(err) + } + h := binary.BigEndian.Uint32(hdr[:]) + if h&(1<<31) == 0 { + t.Errorf("reply record header %#x lacks the last-fragment bit", h) + } + resp := make([]byte, h&^(1<<31)) + if _, err := io.ReadFull(c1, resp); err != nil { + t.Fatal(err) + } + if len(resp) != 28 || binary.BigEndian.Uint32(resp[24:]) != 2050 { + t.Errorf("reply % x; want 28 bytes ending in port 2050", resp) + } +} + +func TestServeUDP(t *testing.T) { + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer pc.Close() + go (&Server{Port: 2050}).ServeUDP(pc) + + c, err := net.Dial("udp", pc.LocalAddr().String()) + if err != nil { + t.Fatal(err) + } + defer c.Close() + if _, err := c.Write(call(3, pmapProcGetport, mountProg, ipprotoTCP)); err != nil { + t.Fatal(err) + } + resp := make([]byte, 100) + n, err := c.Read(resp) + if err != nil { + t.Fatal(err) + } + if n != 28 || binary.BigEndian.Uint32(resp[24:]) != 2050 { + t.Errorf("reply % x; want 28 bytes ending in port 2050", resp[:n]) + } +}