diff --git a/pkg/tcpip/stack/iptables.go b/pkg/tcpip/stack/iptables.go index c3a06b7c9..809109d00 100644 --- a/pkg/tcpip/stack/iptables.go +++ b/pkg/tcpip/stack/iptables.go @@ -219,6 +219,11 @@ func EmptyNATTable() Table { func (it *IPTables) GetTable(id TableID, ipv6 bool) Table { it.mu.RLock() defer it.mu.RUnlock() + return it.getTableRLocked(id, ipv6) +} + +// +checklocksread:it.mu +func (it *IPTables) getTableRLocked(id TableID, ipv6 bool) Table { if ipv6 { return it.v6Tables[id] } @@ -260,6 +265,40 @@ const ( chainReturn ) +type checkTable struct { + fn checkTableFn + tableID TableID + table Table +} + +// shouldSkipOrPopulateTables returns true iff IPTables should be skipped. +// +// If IPTables should not be skipped, tables will be updated with the +// specified table. +func (it *IPTables) shouldSkipOrPopulateTables(tables []checkTable, pkt *PacketBuffer) bool { + switch pkt.NetworkProtocolNumber { + case header.IPv4ProtocolNumber, header.IPv6ProtocolNumber: + default: + // IPTables only supports IPv4/IPv6. + return true + } + + it.mu.RLock() + defer it.mu.RUnlock() + + if !it.modified { + // Many users never configure iptables. Spare them the cost of rule + // traversal if rules have never been set. + return true + } + + for i := range tables { + table := &tables[i] + table.table = it.getTableRLocked(table.tableID, pkt.NetworkProtocolNumber == header.IPv6ProtocolNumber) + } + return false +} + // CheckPrerouting performs the prerouting hook on the packet. // // Returns true iff the packet may continue traversing the stack; the packet @@ -267,20 +306,25 @@ const ( // // Precondition: The packet's network and transport header must be set. func (it *IPTables) CheckPrerouting(pkt *PacketBuffer, addressEP AddressableEndpoint, inNicName string) bool { - it.mu.RLock() - defer it.mu.RUnlock() + tables := [...]checkTable{ + { + fn: it.check, + tableID: MangleID, + }, + { + fn: it.checkNAT, + tableID: NATID, + }, + } - if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { + if it.shouldSkipOrPopulateTables(tables[:], pkt) { return true } pkt.tuple = it.connections.getConnAndUpdate(pkt) - for _, check := range [...]checkTableFn{ - it.checkMangleRLocked, - it.checkNATRLocked, - } { - if !check(Prerouting, pkt, nil /* route */, addressEP, inNicName, "" /* outNicName */) { + for _, table := range tables { + if !table.fn(table.table, Prerouting, pkt, nil /* route */, addressEP, inNicName, "" /* outNicName */) { return false } } @@ -295,18 +339,23 @@ 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 { - it.mu.RLock() - defer it.mu.RUnlock() + tables := [...]checkTable{ + { + fn: it.checkNAT, + tableID: NATID, + }, + { + fn: it.check, + tableID: FilterID, + }, + } - if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { + if it.shouldSkipOrPopulateTables(tables[:], pkt) { return true } - for _, check := range [...]checkTableFn{ - it.checkNATRLocked, - it.checkFilterRLocked, - } { - if !check(Input, pkt, nil /* route */, nil /* addressEP */, inNicName, "" /* outNicName */) { + for _, table := range tables { + if !table.fn(table.table, Input, pkt, nil /* route */, nil /* addressEP */, inNicName, "" /* outNicName */) { return false } } @@ -325,14 +374,24 @@ 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 { - it.mu.RLock() - defer it.mu.RUnlock() + tables := [...]checkTable{ + { + fn: it.check, + tableID: FilterID, + }, + } - if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { + if it.shouldSkipOrPopulateTables(tables[:], pkt) { return true } - return it.checkFilterRLocked(Forward, pkt, nil /* route */, nil /* addressEP */, inNicName, outNicName) + for _, table := range tables { + if !table.fn(table.table, Forward, pkt, nil /* route */, nil /* addressEP */, inNicName, outNicName) { + return false + } + } + + return true } // CheckOutput performs the output hook on the packet. @@ -342,21 +401,29 @@ 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 { - it.mu.RLock() - defer it.mu.RUnlock() + tables := [...]checkTable{ + { + fn: it.check, + tableID: MangleID, + }, + { + fn: it.checkNAT, + tableID: NATID, + }, + { + fn: it.check, + tableID: FilterID, + }, + } - if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { + if it.shouldSkipOrPopulateTables(tables[:], pkt) { return true } pkt.tuple = it.connections.getConnAndUpdate(pkt) - for _, check := range [...]checkTableFn{ - it.checkMangleRLocked, - it.checkNATRLocked, - it.checkFilterRLocked, - } { - if !check(Output, pkt, r, nil /* addressEP */, "" /* inNicName */, outNicName) { + for _, table := range tables { + if !table.fn(table.table, Output, pkt, r, nil /* addressEP */, "" /* inNicName */, outNicName) { return false } } @@ -371,18 +438,23 @@ 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 { - it.mu.RLock() - defer it.mu.RUnlock() + tables := [...]checkTable{ + { + fn: it.check, + tableID: MangleID, + }, + { + fn: it.checkNAT, + tableID: NATID, + }, + } - if it.shouldSkipRLocked(pkt.NetworkProtocolNumber) { + if it.shouldSkipOrPopulateTables(tables[:], pkt) { return true } - for _, check := range [...]checkTableFn{ - it.checkMangleRLocked, - it.checkNATRLocked, - } { - if !check(Postrouting, pkt, r, addressEP, "" /* inNicName */, outNicName) { + for _, table := range tables { + if !table.fn(table.table, Postrouting, pkt, r, addressEP, "" /* inNicName */, outNicName) { return false } } @@ -394,43 +466,18 @@ func (it *IPTables) CheckPostrouting(pkt *PacketBuffer, r *Route, addressEP Addr return true } -// +checklocksread:it.mu -func (it *IPTables) shouldSkipRLocked(netProto tcpip.NetworkProtocolNumber) bool { - switch netProto { - case header.IPv4ProtocolNumber, header.IPv6ProtocolNumber: - default: - // IPTables only supports IPv4/IPv6. - return true - } +type checkTableFn func(table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool - // Many users never configure iptables. Spare them the cost of rule - // traversal if rules have never been set. - return !it.modified -} - -type checkTableFn func(hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool - -// checkMangleRLocked runs the packet through the mangle table. +// checkNAT runs the packet through the NAT table. // -// 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) -} - -// 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 { +// See check. +func (it *IPTables) checkNAT(table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { t := pkt.tuple if t != nil && t.conn.handlePacket(pkt, hook, r) { return true } - if !it.checkRLocked(NATID, hook, pkt, r, addressEP, inNicName, outNicName) { + if !it.check(table, hook, pkt, r, addressEP, inNicName, outNicName) { return false } @@ -462,29 +509,12 @@ func (it *IPTables) checkNATRLocked(hook Hook, pkt *PacketBuffer, r *Route, addr return true } -// 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 +// check 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] - } +func (it *IPTables) check(table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { 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. diff --git a/pkg/tcpip/stack/iptables_types.go b/pkg/tcpip/stack/iptables_types.go index 778f78496..25b51acaa 100644 --- a/pkg/tcpip/stack/iptables_types.go +++ b/pkg/tcpip/stack/iptables_types.go @@ -90,8 +90,11 @@ type IPTables struct { // v4Tables and v6tables map tableIDs to tables. They hold builtin // tables only, not user tables. // + // mu protects the array of tables, but not the tables themselves. // +checklocks:mu v4Tables [NumTables]Table + // + // mu protects the array of tables, but not the tables themselves. // +checklocks:mu v6Tables [NumTables]Table // modified is whether tables have been modified at least once. It is