diff --git a/pkg/tcpip/stack/iptables.go b/pkg/tcpip/stack/iptables.go index 370bc59a6..303157180 100644 --- a/pkg/tcpip/stack/iptables.go +++ b/pkg/tcpip/stack/iptables.go @@ -175,13 +175,6 @@ func DefaultTables(clock tcpip.Clock, rand *rand.Rand) *IPTables { }, }, }, - priorities: [NumHooks][]TableID{ - Prerouting: {MangleID, NATID}, - Input: {NATID, FilterID}, - Forward: {FilterID}, - Output: {MangleID, NATID, FilterID}, - Postrouting: {MangleID, NATID}, - }, connections: ConnTrack{ seed: rand.Uint32(), clock: clock, @@ -259,7 +252,7 @@ const ( // chainAccept indicates the packet should continue through netstack. chainAccept chainVerdict = iota - // chainAccept indicates the packet should be dropped. + // chainDrop indicates the packet should be dropped. chainDrop // chainReturn indicates the packet should return to the calling chain @@ -274,15 +267,25 @@ const ( // // Precondition: The packet's network and transport header must be set. func (it *IPTables) CheckPrerouting(pkt *PacketBuffer, addressEP AddressableEndpoint, inNicName string) bool { - const hook = Prerouting + it.mu.RLock() + defer it.mu.RUnlock() - if it.shouldSkip(pkt.NetworkProtocolNumber) { + if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { return true } pkt.tuple = it.connections.getConnOrMaybeInsertNoop(pkt) - return it.check(hook, pkt, nil /* route */, addressEP, inNicName, "" /* outNicName */) + for _, check := range [...]checkTableFn{ + it.checkMangleRLocked, + it.checkNATRLocked, + } { + if !check(Prerouting, pkt, nil /* route */, addressEP, inNicName, "" /* outNicName */) { + return false + } + } + + return true } // CheckInput performs the input hook on the packet. @@ -292,18 +295,27 @@ func (it *IPTables) CheckPrerouting(pkt *PacketBuffer, addressEP AddressableEndp // // Precondition: The packet's network and transport header must be set. func (it *IPTables) CheckInput(pkt *PacketBuffer, inNicName string) bool { - const hook = Input + it.mu.RLock() + defer it.mu.RUnlock() - if it.shouldSkip(pkt.NetworkProtocolNumber) { + if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { return true } - ret := it.check(hook, pkt, nil /* route */, nil /* addressEP */, inNicName, "" /* outNicName */) + for _, check := range [...]checkTableFn{ + it.checkNATRLocked, + it.checkFilterRLocked, + } { + if !check(Input, pkt, nil /* route */, nil /* addressEP */, inNicName, "" /* outNicName */) { + return false + } + } + if t := pkt.tuple; t != nil { t.conn.finalize() } pkt.tuple = nil - return ret + return true } // CheckForward performs the forward hook on the packet. @@ -313,10 +325,14 @@ func (it *IPTables) CheckInput(pkt *PacketBuffer, inNicName string) bool { // // Precondition: The packet's network and transport header must be set. func (it *IPTables) CheckForward(pkt *PacketBuffer, inNicName, outNicName string) bool { - if it.shouldSkip(pkt.NetworkProtocolNumber) { + it.mu.RLock() + defer it.mu.RUnlock() + + if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { return true } - return it.check(Forward, pkt, nil /* route */, nil /* addressEP */, inNicName, outNicName) + + return it.checkFilterRLocked(Forward, pkt, nil /* route */, nil /* addressEP */, inNicName, outNicName) } // CheckOutput performs the output hook on the packet. @@ -326,15 +342,26 @@ func (it *IPTables) CheckForward(pkt *PacketBuffer, inNicName, outNicName string // // Precondition: The packet's network and transport header must be set. func (it *IPTables) CheckOutput(pkt *PacketBuffer, r *Route, outNicName string) bool { - const hook = Output + it.mu.RLock() + defer it.mu.RUnlock() - if it.shouldSkip(pkt.NetworkProtocolNumber) { + if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { return true } pkt.tuple = it.connections.getConnOrMaybeInsertNoop(pkt) - return it.check(hook, pkt, r, nil /* addressEP */, "" /* inNicName */, outNicName) + for _, check := range [...]checkTableFn{ + it.checkMangleRLocked, + it.checkNATRLocked, + it.checkFilterRLocked, + } { + if !check(Output, pkt, r, nil /* addressEP */, "" /* inNicName */, outNicName) { + return false + } + } + + return true } // CheckPostrouting performs the postrouting hook on the packet. @@ -344,21 +371,31 @@ func (it *IPTables) CheckOutput(pkt *PacketBuffer, r *Route, outNicName string) // // Precondition: The packet's network and transport header must be set. func (it *IPTables) CheckPostrouting(pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, outNicName string) bool { - const hook = Postrouting + it.mu.RLock() + defer it.mu.RUnlock() - if it.shouldSkip(pkt.NetworkProtocolNumber) { + if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { return true } - ret := it.check(hook, pkt, r, addressEP, "" /* inNicName */, outNicName) + for _, check := range [...]checkTableFn{ + it.checkMangleRLocked, + it.checkNATRLocked, + } { + if !check(Postrouting, pkt, r, addressEP, "" /* inNicName */, outNicName) { + return false + } + } + if t := pkt.tuple; t != nil { t.conn.finalize() } pkt.tuple = nil - return ret + return true } -func (it *IPTables) shouldSkip(netProto tcpip.NetworkProtocolNumber) bool { +// +checklocksread:it.mu +func (it *IPTables) shouldSkipRLocked(netProto tcpip.NetworkProtocolNumber) bool { switch netProto { case header.IPv4ProtocolNumber, header.IPv6ProtocolNumber: default: @@ -366,63 +403,83 @@ func (it *IPTables) shouldSkip(netProto tcpip.NetworkProtocolNumber) bool { return true } - it.mu.RLock() - defer it.mu.RUnlock() // Many users never configure iptables. Spare them the cost of rule // traversal if rules have never been set. return !it.modified } -// check runs pkt through the rules for hook. It returns true when the packet -// should continue traversing the network stack and false when it should be -// dropped. +type checkTableFn func(hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool + +// checkMangleRLocked runs the packet through the mangle table. // -// Precondition: The packet's network and transport header must be set. -func (it *IPTables) check(hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { - it.mu.RLock() - defer it.mu.RUnlock() +// See checkRLocked. +// +// +checklocksread:it.mu +func (it *IPTables) checkMangleRLocked(hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { + return it.checkRLocked(MangleID, hook, pkt, r, addressEP, inNicName, outNicName) +} - // Go through each table containing the hook. - priorities := it.priorities[hook] - for _, tableID := range priorities { - if t := pkt.tuple; t != nil && tableID == NATID && t.conn.handlePacket(pkt, hook, r) { - continue - } - var table Table - if pkt.NetworkProtocolNumber == header.IPv6ProtocolNumber { - table = it.v6Tables[tableID] - } else { - table = it.v4Tables[tableID] - } - ruleIdx := table.BuiltinChains[hook] - switch verdict := it.checkChain(hook, pkt, table, ruleIdx, r, addressEP, inNicName, outNicName); verdict { - // If the table returns Accept, move on to the next table. - case chainAccept: - continue - // The Drop verdict is final. - case chainDrop: - return false - case chainReturn: - // Any Return from a built-in chain means we have to - // call the underflow. - underflow := table.Rules[table.Underflows[hook]] - switch v, _ := underflow.Target.Action(pkt, hook, r, addressEP); v { - case RuleAccept: - continue - case RuleDrop: - return false - case RuleJump, RuleReturn: - panic("Underflows should only return RuleAccept or RuleDrop.") - default: - panic(fmt.Sprintf("Unknown verdict: %d", v)) - } - - default: - panic(fmt.Sprintf("Unknown verdict %v.", verdict)) - } +// checkNATRLocked runs the packet through the NAT table. +// +// See checkRLocked. +// +// +checklocksread:it.mu +func (it *IPTables) checkNATRLocked(hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { + if t := pkt.tuple; t != nil && t.conn.handlePacket(pkt, hook, r) { + return true } - return true + return it.checkRLocked(NATID, hook, pkt, r, addressEP, inNicName, outNicName) +} + +// checkFilterRLocked runs the packet through the filter table. +// +// See checkRLocked. +// +// +checklocksread:it.mu +func (it *IPTables) checkFilterRLocked(hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { + return it.checkRLocked(FilterID, hook, pkt, r, addressEP, inNicName, outNicName) +} + +// checkRLocked runs the packet through the rules in the specified table for the +// hook. It returns true if the packet should continue to traverse through the +// network stack or tables, or false when it must be dropped. +// +// Precondition: The packet's network and transport header must be set. +// +// +checklocksread:it.mu +func (it *IPTables) checkRLocked(tableID TableID, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { + var table Table + if pkt.NetworkProtocolNumber == header.IPv6ProtocolNumber { + table = it.v6Tables[tableID] + } else { + table = it.v4Tables[tableID] + } + ruleIdx := table.BuiltinChains[hook] + switch verdict := it.checkChain(hook, pkt, table, ruleIdx, r, addressEP, inNicName, outNicName); verdict { + // If the table returns Accept, move on to the next table. + case chainAccept: + return true + // The Drop verdict is final. + case chainDrop: + return false + case chainReturn: + // Any Return from a built-in chain means we have to + // call the underflow. + underflow := table.Rules[table.Underflows[hook]] + switch v, _ := underflow.Target.Action(pkt, hook, r, addressEP); v { + case RuleAccept: + return true + case RuleDrop: + return false + case RuleJump, RuleReturn: + panic("Underflows should only return RuleAccept or RuleDrop.") + default: + panic(fmt.Sprintf("Unknown verdict: %d", v)) + } + default: + panic(fmt.Sprintf("Unknown verdict %v.", verdict)) + } } // beforeSave is invoked by stateify.