diff --git a/pkg/sentry/socket/netfilter/dnat.go b/pkg/sentry/socket/netfilter/dnat.go index 41f1486e0..ecbc9a6cf 100644 --- a/pkg/sentry/socket/netfilter/dnat.go +++ b/pkg/sentry/socket/netfilter/dnat.go @@ -32,12 +32,14 @@ const DNATTargetName = "DNAT" type dnatTarget struct { stack.DNATTarget + revision uint8 } -func (st *dnatTarget) id() targetID { +func (dt *dnatTarget) id() targetID { return targetID{ name: DNATTargetName, - networkProtocol: st.NetworkProtocol, + networkProtocol: dt.NetworkProtocol, + revision: dt.revision, } } @@ -45,15 +47,15 @@ type dnatTargetMakerV4 struct { NetworkProtocol tcpip.NetworkProtocolNumber } -func (st *dnatTargetMakerV4) id() targetID { +func (dt *dnatTargetMakerV4) id() targetID { return targetID{ name: DNATTargetName, - networkProtocol: st.NetworkProtocol, + networkProtocol: dt.NetworkProtocol, } } func (*dnatTargetMakerV4) marshal(target target) []byte { - st := target.(*dnatTarget) + dt := target.(*dnatTarget) // This is a dnat target named dnat. xt := linux.XTNATTargetV0{ Target: linux.XTEntryTarget{ @@ -62,12 +64,17 @@ func (*dnatTargetMakerV4) marshal(target target) []byte { } copy(xt.Target.Name[:], DNATTargetName) + if dt.ChangeAddress { + xt.NfRange.RangeIPV4.Flags |= linux.NF_NAT_RANGE_MAP_IPS + } + if dt.ChangePort { + xt.NfRange.RangeIPV4.Flags |= linux.NF_NAT_RANGE_PROTO_SPECIFIED + } xt.NfRange.RangeSize = 1 - xt.NfRange.RangeIPV4.Flags |= linux.NF_NAT_RANGE_MAP_IPS | linux.NF_NAT_RANGE_PROTO_SPECIFIED - xt.NfRange.RangeIPV4.MinPort = htons(st.Port) + xt.NfRange.RangeIPV4.MinPort = htons(dt.Port) xt.NfRange.RangeIPV4.MaxPort = xt.NfRange.RangeIPV4.MinPort - copy(xt.NfRange.RangeIPV4.MinIP[:], st.Addr.AsSlice()) - copy(xt.NfRange.RangeIPV4.MaxIP[:], st.Addr.AsSlice()) + copy(xt.NfRange.RangeIPV4.MinIP[:], dt.Addr.AsSlice()) + copy(xt.NfRange.RangeIPV4.MaxIP[:], dt.Addr.AsSlice()) return marshal.Marshal(&xt) } @@ -82,8 +89,8 @@ func (*dnatTargetMakerV4) unmarshal(buf []byte, filter stack.IPHeaderFilter) (ta return nil, syserr.ErrInvalidArgument } - var st linux.XTNATTargetV0 - st.UnmarshalUnsafe(buf) + var dt linux.XTNATTargetV0 + dt.UnmarshalUnsafe(buf) // Copy linux.XTNATTargetV0 to stack.DNATTarget. target := dnatTarget{DNATTarget: stack.DNATTarget{ @@ -91,7 +98,7 @@ func (*dnatTargetMakerV4) unmarshal(buf []byte, filter stack.IPHeaderFilter) (ta }} // RangeSize should be 1. - nfRange := st.NfRange + nfRange := dt.NfRange if nfRange.RangeSize != 1 { nflog("dnatTargetMakerV4: bad rangesize %d", nfRange.RangeSize) return nil, syserr.ErrInvalidArgument @@ -112,7 +119,13 @@ func (*dnatTargetMakerV4) unmarshal(buf []byte, filter stack.IPHeaderFilter) (ta nflog("dnatTargetMakerV4: MinIP != MaxIP (%d, %d)", nfRange.RangeIPV4.MinPort, nfRange.RangeIPV4.MaxPort) return nil, syserr.ErrInvalidArgument } + if nfRange.RangeIPV4.Flags&^(linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED) != 0 { + nflog("dnatTargetMakerV4: unknown flags used (%x)", nfRange.RangeIPV4.Flags) + return nil, syserr.ErrInvalidArgument + } + target.ChangeAddress = nfRange.RangeIPV4.Flags&linux.NF_NAT_RANGE_MAP_IPS != 0 + target.ChangePort = nfRange.RangeIPV4.Flags&linux.NF_NAT_RANGE_PROTO_SPECIFIED != 0 target.Addr = tcpip.AddrFrom4(nfRange.RangeIPV4.MinIP) target.Port = ntohs(nfRange.RangeIPV4.MinPort) @@ -123,34 +136,40 @@ type dnatTargetMakerR1 struct { NetworkProtocol tcpip.NetworkProtocolNumber } -func (st *dnatTargetMakerR1) id() targetID { +func (dt *dnatTargetMakerR1) id() targetID { return targetID{ name: DNATTargetName, - networkProtocol: st.NetworkProtocol, + networkProtocol: dt.NetworkProtocol, revision: 1, } } func (*dnatTargetMakerR1) marshal(target target) []byte { - st := target.(*dnatTarget) + dt := target.(*dnatTarget) nt := linux.XTNATTargetV1{ Target: linux.XTEntryTarget{ TargetSize: linux.SizeOfXTNATTargetV1, - }, - Range: linux.NFNATRange{ - Flags: linux.NF_NAT_RANGE_MAP_IPS | linux.NF_NAT_RANGE_PROTO_SPECIFIED, + Revision: 1, }, } copy(nt.Target.Name[:], DNATTargetName) - copy(nt.Range.MinAddr[:], st.Addr.AsSlice()) - copy(nt.Range.MaxAddr[:], st.Addr.AsSlice()) - nt.Range.MinProto = htons(st.Port) + + if dt.ChangeAddress { + nt.Range.Flags |= linux.NF_NAT_RANGE_MAP_IPS + } + if dt.ChangePort { + nt.Range.Flags |= linux.NF_NAT_RANGE_PROTO_SPECIFIED + } + + copy(nt.Range.MinAddr[:], dt.Addr.AsSlice()) + copy(nt.Range.MaxAddr[:], dt.Addr.AsSlice()) + nt.Range.MinProto = htons(dt.Port) nt.Range.MaxProto = nt.Range.MinProto return marshal.Marshal(&nt) } -func (st *dnatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) (target, *syserr.Error) { +func (dt *dnatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) (target, *syserr.Error) { if size := linux.SizeOfXTNATTargetV1; len(buf) < size { nflog("dnatTargetMakerR1: buf has insufficient size (%d) for DNAT target (%d)", len(buf), size) return nil, syserr.ErrInvalidArgument @@ -164,7 +183,6 @@ func (st *dnatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) var natRange linux.NFNATRange natRange.UnmarshalUnsafe(buf[linux.SizeOfXTEntryTarget:]) - // TODO(gvisor.dev/issue/5697): Support port or address ranges. if natRange.MinAddr != natRange.MaxAddr { nflog("dnatTargetMakerR1: MinAddr and MaxAddr are different") return nil, syserr.ErrInvalidArgument @@ -174,9 +192,8 @@ func (st *dnatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) return nil, syserr.ErrInvalidArgument } - // TODO(gvisor.dev/issue/5698): Support other NF_NAT_RANGE flags. - if natRange.Flags != linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED { - nflog("dnatTargetMakerR1: invalid range flags %d", natRange.Flags) + if natRange.Flags&^(linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED) != 0 { + nflog("dnatTargetMakerR1: invalid flags used (%x)", natRange.Flags) return nil, syserr.ErrInvalidArgument } @@ -184,15 +201,18 @@ func (st *dnatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) DNATTarget: stack.DNATTarget{ NetworkProtocol: filter.NetworkProtocol(), Port: ntohs(natRange.MinProto), + ChangeAddress: natRange.Flags&linux.NF_NAT_RANGE_MAP_IPS != 0, + ChangePort: natRange.Flags&linux.NF_NAT_RANGE_PROTO_SPECIFIED != 0, }, + revision: 1, } - switch st.NetworkProtocol { + switch dt.NetworkProtocol { case header.IPv4ProtocolNumber: target.DNATTarget.Addr = tcpip.AddrFrom4Slice(natRange.MinAddr[:4]) case header.IPv6ProtocolNumber: target.DNATTarget.Addr = tcpip.AddrFrom16(natRange.MinAddr) default: - panic(fmt.Sprintf("invalid protocol number: %d", st.NetworkProtocol)) + panic(fmt.Sprintf("invalid protocol number: %d", dt.NetworkProtocol)) } return &target, nil @@ -202,34 +222,40 @@ type dnatTargetMakerR2 struct { NetworkProtocol tcpip.NetworkProtocolNumber } -func (st *dnatTargetMakerR2) id() targetID { +func (dt *dnatTargetMakerR2) id() targetID { return targetID{ name: DNATTargetName, - networkProtocol: st.NetworkProtocol, + networkProtocol: dt.NetworkProtocol, revision: 2, } } func (*dnatTargetMakerR2) marshal(target target) []byte { - st := target.(*dnatTarget) + dt := target.(*dnatTarget) nt := linux.XTNATTargetV2{ Target: linux.XTEntryTarget{ TargetSize: linux.SizeOfXTNATTargetV1, - }, - Range: linux.NFNATRange2{ - Flags: linux.NF_NAT_RANGE_MAP_IPS | linux.NF_NAT_RANGE_PROTO_SPECIFIED, + Revision: 2, }, } copy(nt.Target.Name[:], DNATTargetName) - copy(nt.Range.MinAddr[:], st.Addr.AsSlice()) - copy(nt.Range.MaxAddr[:], st.Addr.AsSlice()) - nt.Range.MinProto = htons(st.Port) + + if dt.ChangeAddress { + nt.Range.Flags |= linux.NF_NAT_RANGE_MAP_IPS + } + if dt.ChangePort { + nt.Range.Flags |= linux.NF_NAT_RANGE_PROTO_SPECIFIED + } + copy(nt.Range.MinAddr[:], dt.Addr.AsSlice()) + copy(nt.Range.MaxAddr[:], dt.Addr.AsSlice()) + nt.Range.MinProto = htons(dt.Port) nt.Range.MaxProto = nt.Range.MinProto return marshal.Marshal(&nt) } -func (st *dnatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) (target, *syserr.Error) { +func (dt *dnatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) (target, *syserr.Error) { + nflog("dnatTargetMakerR2 unmarshal") if size := linux.SizeOfXTNATTargetV2; len(buf) < size { nflog("dnatTargetMakerR2: buf has insufficient size (%d) for DNAT target (%d)", len(buf), size) return nil, syserr.ErrInvalidArgument @@ -243,7 +269,6 @@ func (st *dnatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) var natRange linux.NFNATRange2 natRange.UnmarshalUnsafe(buf[linux.SizeOfXTEntryTarget:]) - // TODO(gvisor.dev/issue/5697): Support port or address ranges. if natRange.MinAddr != natRange.MaxAddr { nflog("dnatTargetMakerR2: MinAddr and MaxAddr are different") return nil, syserr.ErrInvalidArgument @@ -257,9 +282,8 @@ func (st *dnatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) return nil, syserr.ErrInvalidArgument } - // TODO(gvisor.dev/issue/5698): Support other NF_NAT_RANGE flags. - if natRange.Flags != linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED { - nflog("dnatTargetMakerR2: invalid range flags %d", natRange.Flags) + if natRange.Flags&^(linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED) != 0 { + nflog("dnatTargetMakerR2: invalid flags used (%x)", natRange.Flags) return nil, syserr.ErrInvalidArgument } @@ -267,15 +291,18 @@ func (st *dnatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) DNATTarget: stack.DNATTarget{ NetworkProtocol: filter.NetworkProtocol(), Port: ntohs(natRange.MinProto), + ChangeAddress: natRange.Flags&linux.NF_NAT_RANGE_MAP_IPS != 0, + ChangePort: natRange.Flags&linux.NF_NAT_RANGE_PROTO_SPECIFIED != 0, }, + revision: 2, } - switch st.NetworkProtocol { + switch dt.NetworkProtocol { case header.IPv4ProtocolNumber: target.DNATTarget.Addr = tcpip.AddrFrom4Slice(natRange.MinAddr[:4]) case header.IPv6ProtocolNumber: target.DNATTarget.Addr = tcpip.AddrFrom16(natRange.MinAddr) default: - panic(fmt.Sprintf("invalid protocol number: %d", st.NetworkProtocol)) + panic(fmt.Sprintf("invalid protocol number: %d", dt.NetworkProtocol)) } return &target, nil diff --git a/pkg/sentry/socket/netfilter/snat.go b/pkg/sentry/socket/netfilter/snat.go index 92d672535..d3e9a9c32 100644 --- a/pkg/sentry/socket/netfilter/snat.go +++ b/pkg/sentry/socket/netfilter/snat.go @@ -32,12 +32,14 @@ const SNATTargetName = "SNAT" type snatTarget struct { stack.SNATTarget + revision uint8 } func (st *snatTarget) id() targetID { return targetID{ name: SNATTargetName, networkProtocol: st.NetworkProtocol, + revision: st.revision, } } @@ -62,8 +64,13 @@ func (*snatTargetMakerV4) marshal(target target) []byte { } copy(xt.Target.Name[:], SNATTargetName) + if st.ChangeAddress { + xt.NfRange.RangeIPV4.Flags |= linux.NF_NAT_RANGE_MAP_IPS + } + if st.ChangePort { + xt.NfRange.RangeIPV4.Flags |= linux.NF_NAT_RANGE_PROTO_SPECIFIED + } xt.NfRange.RangeSize = 1 - xt.NfRange.RangeIPV4.Flags |= linux.NF_NAT_RANGE_MAP_IPS | linux.NF_NAT_RANGE_PROTO_SPECIFIED xt.NfRange.RangeIPV4.MinPort = htons(st.Port) xt.NfRange.RangeIPV4.MaxPort = xt.NfRange.RangeIPV4.MinPort copy(xt.NfRange.RangeIPV4.MinIP[:], st.Addr.AsSlice()) @@ -112,7 +119,13 @@ func (*snatTargetMakerV4) unmarshal(buf []byte, filter stack.IPHeaderFilter) (ta nflog("snatTargetMakerV4: MinIP != MaxIP (%d, %d)", nfRange.RangeIPV4.MinPort, nfRange.RangeIPV4.MaxPort) return nil, syserr.ErrInvalidArgument } + if nfRange.RangeIPV4.Flags&^(linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED) != 0 { + nflog("snatTargetMakerV4: unknown flags used (%x)", nfRange.RangeIPV4.Flags) + return nil, syserr.ErrInvalidArgument + } + target.ChangeAddress = nfRange.RangeIPV4.Flags&linux.NF_NAT_RANGE_MAP_IPS != 0 + target.ChangePort = nfRange.RangeIPV4.Flags&linux.NF_NAT_RANGE_PROTO_SPECIFIED != 0 target.Addr = tcpip.AddrFrom4(nfRange.RangeIPV4.MinIP) target.Port = ntohs(nfRange.RangeIPV4.MinPort) @@ -136,12 +149,17 @@ func (*snatTargetMakerR1) marshal(target target) []byte { nt := linux.XTNATTargetV1{ Target: linux.XTEntryTarget{ TargetSize: linux.SizeOfXTNATTargetV1, - }, - Range: linux.NFNATRange{ - Flags: linux.NF_NAT_RANGE_MAP_IPS | linux.NF_NAT_RANGE_PROTO_SPECIFIED, + Revision: 1, }, } copy(nt.Target.Name[:], SNATTargetName) + + if st.ChangeAddress { + nt.Range.Flags |= linux.NF_NAT_RANGE_MAP_IPS + } + if st.ChangePort { + nt.Range.Flags |= linux.NF_NAT_RANGE_PROTO_SPECIFIED + } copy(nt.Range.MinAddr[:], st.Addr.AsSlice()) copy(nt.Range.MaxAddr[:], st.Addr.AsSlice()) nt.Range.MinProto = htons(st.Port) @@ -164,7 +182,6 @@ func (st *snatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) var natRange linux.NFNATRange natRange.UnmarshalUnsafe(buf[linux.SizeOfXTEntryTarget:]) - // TODO(gvisor.dev/issue/5697): Support port or address ranges. if natRange.MinAddr != natRange.MaxAddr { nflog("snatTargetMakerR1: MinAddr and MaxAddr are different") return nil, syserr.ErrInvalidArgument @@ -174,9 +191,8 @@ func (st *snatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) return nil, syserr.ErrInvalidArgument } - // TODO(gvisor.dev/issue/5698): Support other NF_NAT_RANGE flags. - if natRange.Flags != linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED { - nflog("snatTargetMakerR1: invalid range flags %d", natRange.Flags) + if natRange.Flags&^(linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED) != 0 { + nflog("snatTargetMakerR1: unknown flags used (%x)", natRange.Flags) return nil, syserr.ErrInvalidArgument } @@ -184,7 +200,10 @@ func (st *snatTargetMakerR1) unmarshal(buf []byte, filter stack.IPHeaderFilter) SNATTarget: stack.SNATTarget{ NetworkProtocol: filter.NetworkProtocol(), Port: ntohs(natRange.MinProto), + ChangeAddress: natRange.Flags&linux.NF_NAT_RANGE_MAP_IPS != 0, + ChangePort: natRange.Flags&linux.NF_NAT_RANGE_PROTO_SPECIFIED != 0, }, + revision: 1, } switch st.NetworkProtocol { case header.IPv4ProtocolNumber: @@ -215,12 +234,17 @@ func (*snatTargetMakerR2) marshal(target target) []byte { nt := linux.XTNATTargetV2{ Target: linux.XTEntryTarget{ TargetSize: linux.SizeOfXTNATTargetV1, - }, - Range: linux.NFNATRange2{ - Flags: linux.NF_NAT_RANGE_MAP_IPS | linux.NF_NAT_RANGE_PROTO_SPECIFIED, + Revision: 2, }, } copy(nt.Target.Name[:], SNATTargetName) + + if st.ChangeAddress { + nt.Range.Flags |= linux.NF_NAT_RANGE_MAP_IPS + } + if st.ChangePort { + nt.Range.Flags |= linux.NF_NAT_RANGE_PROTO_SPECIFIED + } copy(nt.Range.MinAddr[:], st.Addr.AsSlice()) copy(nt.Range.MaxAddr[:], st.Addr.AsSlice()) nt.Range.MinProto = htons(st.Port) @@ -243,7 +267,6 @@ func (st *snatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) var natRange linux.NFNATRange2 natRange.UnmarshalUnsafe(buf[linux.SizeOfXTEntryTarget:]) - // TODO(gvisor.dev/issue/5697): Support port or address ranges. if natRange.MinAddr != natRange.MaxAddr { nflog("snatTargetMakerR2: MinAddr and MaxAddr are different") return nil, syserr.ErrInvalidArgument @@ -257,9 +280,8 @@ func (st *snatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) return nil, syserr.ErrInvalidArgument } - // TODO(gvisor.dev/issue/5698): Support other NF_NAT_RANGE flags. - if natRange.Flags != linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED { - nflog("snatTargetMakerR2: invalid range flags %d", natRange.Flags) + if natRange.Flags&^(linux.NF_NAT_RANGE_MAP_IPS|linux.NF_NAT_RANGE_PROTO_SPECIFIED) != 0 { + nflog("snatTargetMakerR1: unknown flags used (%x)", natRange.Flags) return nil, syserr.ErrInvalidArgument } @@ -267,7 +289,10 @@ func (st *snatTargetMakerR2) unmarshal(buf []byte, filter stack.IPHeaderFilter) SNATTarget: stack.SNATTarget{ NetworkProtocol: filter.NetworkProtocol(), Port: ntohs(natRange.MinProto), + ChangeAddress: natRange.Flags&linux.NF_NAT_RANGE_MAP_IPS != 0, + ChangePort: natRange.Flags&linux.NF_NAT_RANGE_PROTO_SPECIFIED != 0, }, + revision: 2, } switch st.NetworkProtocol { case header.IPv4ProtocolNumber: diff --git a/pkg/tcpip/stack/conntrack.go b/pkg/tcpip/stack/conntrack.go index 02bce8704..215eb3612 100644 --- a/pkg/tcpip/stack/conntrack.go +++ b/pkg/tcpip/stack/conntrack.go @@ -725,7 +725,7 @@ type portOrIdentRange struct { // // Generally, only the first packet of a connection reaches this method; other // packets will be manipulated without needing to modify the connection. -func (cn *conn) performNAT(pkt PacketBufferPtr, hook Hook, r *Route, portsOrIdents portOrIdentRange, natAddress tcpip.Address, dnat bool) { +func (cn *conn) performNAT(pkt PacketBufferPtr, hook Hook, r *Route, portsOrIdents portOrIdentRange, natAddress tcpip.Address, dnat, changePort, changeAddress bool) { lastPortOrIdent := func() uint16 { lastPortOrIdent := uint32(portsOrIdents.start) + portsOrIdents.size - 1 if lastPortOrIdent > math.MaxUint16 { @@ -762,7 +762,14 @@ func (cn *conn) performNAT(pkt PacketBufferPtr, hook Hook, r *Route, portsOrIden return } *manip = manipPerformed - *address = natAddress + if changeAddress { + *address = natAddress + } + + // Everything below here is port-fiddling. + if !changePort { + return + } // Does the current port/ident fit in the range? if portsOrIdents.start <= *portOrIdent && *portOrIdent <= lastPortOrIdent { diff --git a/pkg/tcpip/stack/iptables_targets.go b/pkg/tcpip/stack/iptables_targets.go index 4ba1f3e8d..e3cedaf0d 100644 --- a/pkg/tcpip/stack/iptables_targets.go +++ b/pkg/tcpip/stack/iptables_targets.go @@ -182,6 +182,16 @@ type DNATTarget struct { // // Immutable. NetworkProtocol tcpip.NetworkProtocolNumber + + // ChangeAddress indicates whether we should check addresses. + // + // Immutable. + ChangeAddress bool + + // ChangePort indicates whether we should check ports. + // + // Immutable. + ChangePort bool } // Action implements Target.Action. @@ -201,7 +211,7 @@ func (rt *DNATTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, addressEP panic(fmt.Sprintf("%s unrecognized", hook)) } - return dnatAction(pkt, hook, r, rt.Port, rt.Addr) + return dnatAction(pkt, hook, r, rt.Port, rt.Addr, rt.ChangePort, rt.ChangeAddress) } @@ -244,7 +254,7 @@ func (rt *RedirectTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, addre panic("redirect target is supported only on output and prerouting hooks") } - return dnatAction(pkt, hook, r, rt.Port, address) + return dnatAction(pkt, hook, r, rt.Port, address, true /* changePort */, true /* changeAddress */) } // SNATTarget modifies the source port/IP in the outgoing packets. @@ -255,10 +265,20 @@ type SNATTarget struct { // NetworkProtocol is the network protocol the target is used with. It // is immutable. NetworkProtocol tcpip.NetworkProtocolNumber + + // ChangeAddress indicates whether we should check addresses. + // + // Immutable. + ChangeAddress bool + + // ChangePort indicates whether we should check ports. + // + // Immutable. + ChangePort bool } -func dnatAction(pkt PacketBufferPtr, hook Hook, r *Route, port uint16, address tcpip.Address) (RuleVerdict, int) { - return natAction(pkt, hook, r, portOrIdentRange{start: port, size: 1}, address, true /* dnat */) +func dnatAction(pkt PacketBufferPtr, hook Hook, r *Route, port uint16, address tcpip.Address, changePort, changeAddress bool) (RuleVerdict, int) { + return natAction(pkt, hook, r, portOrIdentRange{start: port, size: 1}, address, true /* dnat */, changePort, changeAddress) } func targetPortRangeForTCPAndUDP(originalSrcPort uint16) portOrIdentRange { @@ -278,7 +298,7 @@ func targetPortRangeForTCPAndUDP(originalSrcPort uint16) portOrIdentRange { } } -func snatAction(pkt PacketBufferPtr, hook Hook, r *Route, port uint16, address tcpip.Address) (RuleVerdict, int) { +func snatAction(pkt PacketBufferPtr, hook Hook, r *Route, port uint16, address tcpip.Address, changePort, changeAddress bool) (RuleVerdict, int) { portsOrIdents := portOrIdentRange{start: port, size: 1} switch pkt.TransportProtocolNumber { @@ -298,17 +318,17 @@ func snatAction(pkt PacketBufferPtr, hook Hook, r *Route, port uint16, address t portsOrIdents = portOrIdentRange{start: 0, size: math.MaxUint16 + 1} } - return natAction(pkt, hook, r, portsOrIdents, address, false /* dnat */) + return natAction(pkt, hook, r, portsOrIdents, address, false /* dnat */, changePort, changeAddress) } -func natAction(pkt PacketBufferPtr, hook Hook, r *Route, portsOrIdents portOrIdentRange, address tcpip.Address, dnat bool) (RuleVerdict, int) { +func natAction(pkt PacketBufferPtr, hook Hook, r *Route, portsOrIdents portOrIdentRange, address tcpip.Address, dnat, changePort, changeAddress bool) (RuleVerdict, int) { // Drop the packet if network and transport header are not set. if len(pkt.NetworkHeader().Slice()) == 0 || len(pkt.TransportHeader().Slice()) == 0 { return RuleDrop, 0 } if t := pkt.tuple; t != nil { - t.conn.performNAT(pkt, hook, r, portsOrIdents, address, dnat) + t.conn.performNAT(pkt, hook, r, portsOrIdents, address, dnat, changePort, changeAddress) return RuleAccept, 0 } @@ -332,7 +352,7 @@ func (st *SNATTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, _ Address panic(fmt.Sprintf("%s unrecognized", hook)) } - return snatAction(pkt, hook, r, st.Port, st.Addr) + return snatAction(pkt, hook, r, st.Port, st.Addr, st.ChangePort, st.ChangeAddress) } // MasqueradeTarget modifies the source port/IP in the outgoing packets. @@ -368,7 +388,7 @@ func (mt *MasqueradeTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, add address := ep.AddressWithPrefix().Address ep.DecRef() - return snatAction(pkt, hook, r, 0 /* port */, address) + return snatAction(pkt, hook, r, 0 /* port */, address, true /* changePort */, true /* changeAddress */) } func rewritePacket(n header.Network, t header.Transport, updateSRCFields, fullChecksum, updatePseudoHeader bool, newPortOrIdent uint16, newAddr tcpip.Address) { diff --git a/pkg/tcpip/stack/iptables_test.go b/pkg/tcpip/stack/iptables_test.go index 654803eb7..1265cbeb0 100644 --- a/pkg/tcpip/stack/iptables_test.go +++ b/pkg/tcpip/stack/iptables_test.go @@ -82,7 +82,7 @@ func TestNATedConnectionReap(t *testing.T) { Rules: []Rule{ // Prerouting { - Target: &DNATTarget{NetworkProtocol: netProto, Addr: nattedAddr, Port: nattedPort}, + Target: &DNATTarget{NetworkProtocol: netProto, Addr: nattedAddr, Port: nattedPort, ChangePort: true, ChangeAddress: true}, }, { Target: &AcceptTarget{}, @@ -381,7 +381,7 @@ func TestNATConflict(t *testing.T) { // Input { - Target: &SNATTarget{NetworkProtocol: header.IPv6ProtocolNumber, Addr: nattedAddr, Port: nattedPort}, + Target: &SNATTarget{NetworkProtocol: header.IPv6ProtocolNumber, Addr: nattedAddr, Port: nattedPort, ChangeAddress: true, ChangePort: true}, }, { Target: &AcceptTarget{}, @@ -399,7 +399,7 @@ func TestNATConflict(t *testing.T) { // Postrouting { - Target: &SNATTarget{NetworkProtocol: header.IPv6ProtocolNumber, Addr: nattedAddr, Port: nattedPort}, + Target: &SNATTarget{NetworkProtocol: header.IPv6ProtocolNumber, Addr: nattedAddr, Port: nattedPort, ChangeAddress: true, ChangePort: true}, }, { Target: &AcceptTarget{}, diff --git a/pkg/tcpip/tests/integration/iptables_test.go b/pkg/tcpip/tests/integration/iptables_test.go index c61a6fecf..497449e8d 100644 --- a/pkg/tcpip/tests/integration/iptables_test.go +++ b/pkg/tcpip/tests/integration/iptables_test.go @@ -1324,7 +1324,7 @@ var ( setupNAT: func(t *testing.T, s *stack.Stack, netProto tcpip.NetworkProtocolNumber, transProto tcpip.TransportProtocolNumber, snatAddr, _ tcpip.Address, _ uint16) { t.Helper() - setupSNAT(t, s, netProto, transProto, &stack.SNATTarget{NetworkProtocol: netProto, Addr: snatAddr}) + setupSNAT(t, s, netProto, transProto, &stack.SNATTarget{NetworkProtocol: netProto, Addr: snatAddr, ChangeAddress: true, ChangePort: true}) }, }, { @@ -1342,7 +1342,7 @@ var ( setupNAT: func(t *testing.T, s *stack.Stack, netProto tcpip.NetworkProtocolNumber, transProto tcpip.TransportProtocolNumber, _, dnatAddr tcpip.Address, dnatPort uint16) { t.Helper() - setupDNAT(t, s, netProto, transProto, &stack.DNATTarget{NetworkProtocol: netProto, Addr: dnatAddr, Port: dnatPort}) + setupDNAT(t, s, netProto, transProto, &stack.DNATTarget{NetworkProtocol: netProto, Addr: dnatAddr, Port: dnatPort, ChangeAddress: true, ChangePort: true}) }, } @@ -1364,7 +1364,7 @@ var ( setupNAT: func(t *testing.T, s *stack.Stack, netProto tcpip.NetworkProtocolNumber, transProto tcpip.TransportProtocolNumber, snatAddr, dnatAddr tcpip.Address, dnatPort uint16) { t.Helper() - setupTwiceNAT(t, s, netProto, transProto, dnatAddr, &stack.DNATTarget{NetworkProtocol: netProto, Addr: dnatAddr, Port: dnatPort}, &stack.MasqueradeTarget{NetworkProtocol: netProto}) + setupTwiceNAT(t, s, netProto, transProto, dnatAddr, &stack.DNATTarget{NetworkProtocol: netProto, Addr: dnatAddr, Port: dnatPort, ChangeAddress: true, ChangePort: true}, &stack.MasqueradeTarget{NetworkProtocol: netProto}) }, }, { @@ -1372,7 +1372,7 @@ var ( setupNAT: func(t *testing.T, s *stack.Stack, netProto tcpip.NetworkProtocolNumber, transProto tcpip.TransportProtocolNumber, snatAddr, dnatAddr tcpip.Address, dnatPort uint16) { t.Helper() - setupTwiceNAT(t, s, netProto, transProto, dnatAddr, &stack.DNATTarget{NetworkProtocol: netProto, Addr: dnatAddr, Port: dnatPort}, &stack.SNATTarget{NetworkProtocol: netProto, Addr: snatAddr}) + setupTwiceNAT(t, s, netProto, transProto, dnatAddr, &stack.DNATTarget{NetworkProtocol: netProto, Addr: dnatAddr, Port: dnatPort, ChangeAddress: true, ChangePort: true}, &stack.SNATTarget{NetworkProtocol: netProto, Addr: snatAddr, ChangeAddress: true, ChangePort: true}) }, }, } @@ -2484,7 +2484,7 @@ func TestNATICMPError(t *testing.T) { CheckProtocol: true, InputInterface: utils.RouterNIC2Name, }, - Target: &stack.DNATTarget{NetworkProtocol: test.netProto, Addr: test.host1Addr, Port: dstPort}, + Target: &stack.DNATTarget{NetworkProtocol: test.netProto, Addr: test.host1Addr, Port: dstPort, ChangeAddress: true, ChangePort: true}, }, { Target: &stack.AcceptTarget{}, @@ -2826,7 +2826,7 @@ func TestSNATHandlePortOrIdentConflicts(t *testing.T) { { name: "SNAT", target: func(netProto tcpip.NetworkProtocolNumber, addr tcpip.Address) stack.Target { - return &stack.SNATTarget{NetworkProtocol: netProto, Addr: addr} + return &stack.SNATTarget{NetworkProtocol: netProto, Addr: addr, ChangeAddress: true, ChangePort: true} }, }, } diff --git a/test/iptables/iptables_test.go b/test/iptables/iptables_test.go index 1accdaab7..6827314b2 100644 --- a/test/iptables/iptables_test.go +++ b/test/iptables/iptables_test.go @@ -368,6 +368,14 @@ func TestNATOutDNAT(t *testing.T) { singleTest(t, &NATOutDNAT{}) } +func TestNATOutDNATAddrOnly(t *testing.T) { + singleTest(t, &NATOutDNATAddrOnly{}) +} + +func TestNATOutDNATPortOnly(t *testing.T) { + singleTest(t, &NATOutDNATPortOnly{}) +} + func TestNATPreRedirectIP(t *testing.T) { singleTest(t, &NATPreRedirectIP{}) } diff --git a/test/iptables/nat.go b/test/iptables/nat.go index 0c2ac5fc0..d1ea2c926 100644 --- a/test/iptables/nat.go +++ b/test/iptables/nat.go @@ -52,6 +52,8 @@ func init() { RegisterTestCase(&NATPostSNATUDP{}) RegisterTestCase(&NATPostSNATTCP{}) RegisterTestCase(&NATOutDNAT{}) + RegisterTestCase(&NATOutDNATAddrOnly{}) + RegisterTestCase(&NATOutDNATPortOnly{}) } // NATPreRedirectUDPPort tests that packets are redirected to different port. @@ -672,13 +674,19 @@ func listenForRedirectedConn(ctx context.Context, ipv6 bool, originalDsts []net. // loopbackTests runs an iptables rule and ensures that packets sent to // dest:dropPort are received by localhost:acceptPort. func loopbackTest(ctx context.Context, ipv6 bool, dest net.IP, args ...string) error { + return loopbackTestPort(ctx, ipv6, dest, dropPort, args...) +} + +// loopbackTests runs an iptables rule and ensures that packets sent to +// dest:port are received by localhost:acceptPort. +func loopbackTestPort(ctx context.Context, ipv6 bool, dest net.IP, port int, args ...string) error { if err := natTable(ipv6, args...); err != nil { return err } sendCh := make(chan error, 1) listenCh := make(chan error, 1) go func() { - sendCh <- sendUDPLoop(ctx, dest, dropPort, ipv6) + sendCh <- sendUDPLoop(ctx, dest, port, ipv6) }() go func() { listenCh <- listenUDP(ctx, acceptPort, ipv6) @@ -1063,3 +1071,56 @@ func (*NATOutDNAT) ContainerAction(ctx context.Context, ip net.IP, ipv6 bool) er func (*NATOutDNAT) LocalAction(ctx context.Context, ip net.IP, ipv6 bool) error { return nil } + +// NATOutDNATAddrOnly tests that the source IP only in the packets are modified +// as expected. It tests the latest-implemented revision of the DNAT target. +type NATOutDNATAddrOnly struct{ containerCase } + +var _ TestCase = (*NATOutDNATAddrOnly)(nil) + +// Name implements TestCase.Name. +func (*NATOutDNATAddrOnly) Name() string { + return "NATOutDNATAddrOnly" +} + +// ContainerAction implements TestCase.ContainerAction. +func (*NATOutDNATAddrOnly) ContainerAction(ctx context.Context, ip net.IP, ipv6 bool) error { + dst := nowhereIP(ipv6) + return loopbackTestPort(ctx, ipv6, net.ParseIP(dst), acceptPort, + "-A", "OUTPUT", + "-d", dst, + "-p", "udp", "-m", "udp", + "-j", "DNAT", "--to-destination", "127.0.0.1") +} + +// LocalAction implements TestCase.LocalAction. +func (*NATOutDNATAddrOnly) LocalAction(ctx context.Context, ip net.IP, ipv6 bool) error { + return nil +} + +// NATOutDNATPortOnly tests that the source port only in the packets are +// modified as expected. It tests the latest-implemented revision of the DNAT +// target. +type NATOutDNATPortOnly struct{ containerCase } + +var _ TestCase = (*NATOutDNATPortOnly)(nil) + +// Name implements TestCase.Name. +func (*NATOutDNATPortOnly) Name() string { + return "NATOutDNATPortOnly" +} + +// ContainerAction implements TestCase.ContainerAction. +func (*NATOutDNATPortOnly) ContainerAction(ctx context.Context, ip net.IP, ipv6 bool) error { + const dst = "127.0.0.1" + return loopbackTest(ctx, ipv6, net.ParseIP(dst), + "-A", "OUTPUT", + "-d", dst, + "-p", "udp", "-m", "udp", + "-j", "DNAT", "--to-destination", fmt.Sprintf(":%d", acceptPort)) +} + +// LocalAction implements TestCase.LocalAction. +func (*NATOutDNATPortOnly) LocalAction(ctx context.Context, ip net.IP, ipv6 bool) error { + return nil +}