diff --git a/filter.go b/filter.go index 86e3f9d..bc8be51 100644 --- a/filter.go +++ b/filter.go @@ -93,7 +93,7 @@ func (f *filter) MarkMask(mask uint32) Filter { } func (f *filter) Status(status Status) Filter { - f.f[ctaStatus] = netfilter.Uint32Bytes(uint32(status.Value)) + f.f[ctaStatus] = netfilter.Uint32Bytes(uint32(status)) return f } diff --git a/filter_test.go b/filter_test.go index 1d526c4..b356304 100644 --- a/filter_test.go +++ b/filter_test.go @@ -13,7 +13,7 @@ func TestFilterMarshal(t *testing.T) { f := NewFilter(). Mark(0xf0000000).MarkMask(0x0000000f). Zone(42). - Status(Status{StatusDying}).StatusMask(0xdeadbeef) + Status(StatusDying).StatusMask(0xdeadbeef) want := []netfilter.Attribute{ { diff --git a/flow.go b/flow.go index b94afa0..6798842 100644 --- a/flow.go +++ b/flow.go @@ -44,12 +44,12 @@ type Flow struct { // source and destination addresses. srcPort and dstPort are the source and // destination ports. timeout is the non-zero time-to-live of a connection in // seconds. -func NewFlow(proto uint8, status StatusFlag, srcAddr, destAddr netip.Addr, +func NewFlow(proto uint8, status Status, srcAddr, destAddr netip.Addr, srcPort, destPort uint16, timeout, mark uint32) Flow { var f Flow - f.Status.Value = status + f.Status = status f.Timeout = timeout f.Mark = mark @@ -107,7 +107,7 @@ func (f *Flow) unmarshal(ad *netlink.AttributeDecoder) error { // CTA_STATUS is a bitfield of the state of the connection // (eg. if packets are seen in both directions, etc.) case ctaStatus: - f.Status.Value = StatusFlag(ad.Uint32()) + f.Status = Status(ad.Uint32()) } } @@ -211,7 +211,7 @@ func (f Flow) marshal() ([]netfilter.Attribute, error) { attrs = append(attrs, a) } - if f.Status.Value != 0 { + if f.Status != 0 { attrs = append(attrs, f.Status.marshal()) } diff --git a/flow_integration_test.go b/flow_integration_test.go index 713c529..1167c4f 100644 --- a/flow_integration_test.go +++ b/flow_integration_test.go @@ -478,16 +478,16 @@ func TestStatusFilter(t *testing.T) { require.NoError(t, err) assert.Len(t, flows, 2, "expected 2 flows in total") - flows, err = c.DumpFilter(NewFilter().Status(Status{StatusConfirmed}), nil) + flows, err = c.DumpFilter(NewFilter().Status(StatusConfirmed), nil) require.NoError(t, err) assert.Len(t, flows, 2) - flows, err = c.DumpFilter(NewFilter().Status(Status{StatusDying}), nil) + flows, err = c.DumpFilter(NewFilter().Status(StatusDying), nil) require.NoError(t, err) assert.Len(t, flows, 0) // This filter can never return anything since status and mask don't overlap. - flows, err = c.DumpFilter(NewFilter().Status(Status{StatusConfirmed}).StatusMask(0x1), nil) + flows, err = c.DumpFilter(NewFilter().Status(StatusConfirmed).StatusMask(0x1), nil) require.NoError(t, err) assert.Len(t, flows, 0) } diff --git a/flow_test.go b/flow_test.go index 2418c85..55b04a1 100644 --- a/flow_test.go +++ b/flow_test.go @@ -148,7 +148,7 @@ var ( Data: []byte{0xff, 0x00, 0xff, 0x00}, }, }, - flow: Flow{Status: Status{Value: 0xff00ff00}}, + flow: Flow{Status: 0xff00ff00}, }, { name: "protoinfo attribute w/ tcp info", @@ -420,7 +420,7 @@ func TestFlowMarshal(t *testing.T) { attrs, err := Flow{ TupleOrig: flowIPPT, TupleReply: flowIPPT, TupleMaster: flowIPPT, ProtoInfo: ProtoInfo{TCP: &ProtoInfoTCP{State: 42}}, - Timeout: 123, Status: Status{Value: 1234}, Mark: 0x1234, Zone: 2, + Timeout: 123, Status: 1234, Mark: 0x1234, Zone: 2, Helper: Helper{Name: "ftp"}, SeqAdjOrig: SequenceAdjust{Position: 1, OffsetBefore: 2, OffsetAfter: 3}, SeqAdjReply: SequenceAdjust{Position: 5, OffsetBefore: 6, OffsetAfter: 7}, @@ -519,7 +519,7 @@ func TestNewFlow(t *testing.T) { ) want := Flow{ - Status: Status{Value: StatusNATMask}, + Status: StatusNATMask, Timeout: 400, TupleOrig: Tuple{ IP: IPTuple{ diff --git a/status.go b/status.go index 89e6c69..5dc2eb1 100644 --- a/status.go +++ b/status.go @@ -5,12 +5,7 @@ import ( "github.com/ti-mo/netfilter" ) -// Status represents a snapshot of a conntrack connection's state. -type Status struct { - Value StatusFlag -} - -// unmarshal unmarshals a netfilter.Attribute into a Status structure. +// unmarshal unmarshals a Status from ad. func (s *Status) unmarshal(ad *netlink.AttributeDecoder) error { if ad.Len() != 1 { return errNeedSingleChild @@ -24,7 +19,7 @@ func (s *Status) unmarshal(ad *netlink.AttributeDecoder) error { return errIncorrectSize } - s.Value = StatusFlag(ad.Uint32()) + *s = Status(ad.Uint32()) return ad.Err() } @@ -33,106 +28,106 @@ func (s *Status) unmarshal(ad *netlink.AttributeDecoder) error { func (s Status) marshal() netfilter.Attribute { return netfilter.Attribute{ Type: uint16(ctaStatus), - Data: netfilter.Uint32Bytes(uint32(s.Value)), + Data: netfilter.Uint32Bytes(uint32(s)), } } // Expected indicates that this connection is an expected connection, // created by Conntrack helpers based on the state of another, related connection. func (s Status) Expected() bool { - return s.Value&StatusExpected != 0 + return s&StatusExpected != 0 } // SeenReply is set when the flow has seen traffic both ways. func (s Status) SeenReply() bool { - return s.Value&StatusSeenReply != 0 + return s&StatusSeenReply != 0 } // Assured is set when eg. three-way handshake is completed on a TCP flow. func (s Status) Assured() bool { - return s.Value&StatusAssured != 0 + return s&StatusAssured != 0 } // Confirmed is set when the original packet has left the box. func (s Status) Confirmed() bool { - return s.Value&StatusConfirmed != 0 + return s&StatusConfirmed != 0 } // SrcNAT means the connection needs source NAT in the original direction. func (s Status) SrcNAT() bool { - return s.Value&StatusSrcNAT != 0 + return s&StatusSrcNAT != 0 } // DstNAT means the connection needs destination NAT in the original direction. func (s Status) DstNAT() bool { - return s.Value&StatusDstNAT != 0 + return s&StatusDstNAT != 0 } // SeqAdjust means the connection needs its TCP sequence to be adjusted. func (s Status) SeqAdjust() bool { - return s.Value&StatusSeqAdjust != 0 + return s&StatusSeqAdjust != 0 } // SrcNATDone is set when source NAT was applied onto the connection. func (s Status) SrcNATDone() bool { - return s.Value&StatusSrcNATDone != 0 + return s&StatusSrcNATDone != 0 } // DstNATDone is set when destination NAT was applied onto the connection. func (s Status) DstNATDone() bool { - return s.Value&StatusDstNATDone != 0 + return s&StatusDstNATDone != 0 } // Dying means the connection has concluded and needs to be cleaned up by GC. func (s Status) Dying() bool { - return s.Value&StatusDying != 0 + return s&StatusDying != 0 } // FixedTimeout means the connection's timeout value cannot be changed. func (s Status) FixedTimeout() bool { - return s.Value&StatusFixedTimeout != 0 + return s&StatusFixedTimeout != 0 } // Template indicates if the connection is a template. func (s Status) Template() bool { - return s.Value&StatusTemplate != 0 + return s&StatusTemplate != 0 } // Helper is set when a helper was explicitly attached using a Conntrack target. func (s Status) Helper() bool { - return s.Value&StatusHelper != 0 + return s&StatusHelper != 0 } // Offload is set when the connection was offloaded to flow table. func (s Status) Offload() bool { - return s.Value&StatusOffload != 0 + return s&StatusOffload != 0 } -// StatusFlag describes a status bit in a Status structure. -type StatusFlag uint32 +// Status is a bitfield describing the state of a Flow. +type Status uint32 // Conntrack connection's status flags, from enum ip_conntrack_status. // uapi/linux/netfilter/nf_conntrack_common.h const ( - StatusExpected StatusFlag = 1 // IPS_EXPECTED - StatusSeenReply StatusFlag = 1 << 1 // IPS_SEEN_REPLY - StatusAssured StatusFlag = 1 << 2 // IPS_ASSURED - StatusConfirmed StatusFlag = 1 << 3 // IPS_CONFIRMED - StatusSrcNAT StatusFlag = 1 << 4 // IPS_SRC_NAT - StatusDstNAT StatusFlag = 1 << 5 // IPS_DST_NAT + StatusExpected Status = 1 // IPS_EXPECTED + StatusSeenReply Status = 1 << 1 // IPS_SEEN_REPLY + StatusAssured Status = 1 << 2 // IPS_ASSURED + StatusConfirmed Status = 1 << 3 // IPS_CONFIRMED + StatusSrcNAT Status = 1 << 4 // IPS_SRC_NAT + StatusDstNAT Status = 1 << 5 // IPS_DST_NAT StatusNATMask = StatusDstNAT | StatusSrcNAT // IPS_NAT_MASK - StatusSeqAdjust StatusFlag = 1 << 6 // IPS_SEQ_ADJUST - StatusSrcNATDone StatusFlag = 1 << 7 // IPS_SRC_NAT_DONE - StatusDstNATDone StatusFlag = 1 << 8 // IPS_DST_NAT_DONE + StatusSeqAdjust Status = 1 << 6 // IPS_SEQ_ADJUST + StatusSrcNATDone Status = 1 << 7 // IPS_SRC_NAT_DONE + StatusDstNATDone Status = 1 << 8 // IPS_DST_NAT_DONE StatusNATDoneMask = StatusDstNATDone | StatusSrcNATDone // IPS_NAT_DONE_MASK - StatusDying StatusFlag = 1 << 9 - StatusFixedTimeout StatusFlag = 1 << 10 // IPS_FIXED_TIMEOUT - StatusTemplate StatusFlag = 1 << 11 // IPS_TEMPLATE - StatusUntracked StatusFlag = 1 << 12 // IPS_UNTRACKED - StatusHelper StatusFlag = 1 << 13 // IPS_HELPER - StatusOffload StatusFlag = 1 << 14 // IPS_OFFLOAD + StatusDying Status = 1 << 9 + StatusFixedTimeout Status = 1 << 10 // IPS_FIXED_TIMEOUT + StatusTemplate Status = 1 << 11 // IPS_TEMPLATE + StatusUntracked Status = 1 << 12 // IPS_UNTRACKED + StatusHelper Status = 1 << 13 // IPS_HELPER + StatusOffload Status = 1 << 14 // IPS_OFFLOAD ) diff --git a/status_test.go b/status_test.go index 59bba46..0002dec 100644 --- a/status_test.go +++ b/status_test.go @@ -34,12 +34,17 @@ func TestStatusMarshalTwoWay(t *testing.T) { { name: "default values", b: []byte{0x00, 0x00, 0x00, 0x00}, - status: Status{}, + status: 0, + }, + { + name: "assured", + b: []byte{0x00, 0x00, 0x00, 0xc}, + status: StatusAssured | StatusConfirmed, }, { name: "out of range, only highest bits flipped", b: []byte{0xFF, 0xFF, 0x80, 0x00}, - status: Status{Value: 0xFFFF8000}, + status: 0xFFFF8000, }, { name: "error, byte array too short", @@ -54,9 +59,7 @@ func TestStatusMarshalTwoWay(t *testing.T) { } for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - // Wrap in status attribute container nfa := netfilter.Attribute{ Type: uint16(ctaStatus), @@ -64,14 +67,13 @@ func TestStatusMarshalTwoWay(t *testing.T) { } var s Status - err := s.unmarshal(mustDecodeAttribute(nfa)) if err != nil || tt.err != nil { require.ErrorIs(t, err, tt.err) return } - require.Equal(t, tt.status.Value, s.Value, "unexpected unmarshal") + require.Equal(t, tt.status, s, "unexpected unmarshal") ms := s.marshal() assert.Equal(t, nfa, ms, "unexpected marshal") @@ -80,55 +82,24 @@ func TestStatusMarshalTwoWay(t *testing.T) { } func TestStatusFieldTest(t *testing.T) { - - var s Status - - s.Value = StatusExpected - assert.Equal(t, true, s.Expected(), "expected") - - s.Value = StatusSeenReply - assert.Equal(t, true, s.SeenReply(), "seenreply") - - s.Value = StatusAssured - assert.Equal(t, true, s.Assured(), "assured") - - s.Value = StatusConfirmed - assert.Equal(t, true, s.Confirmed(), "confirmed") - - s.Value = StatusSrcNAT - assert.Equal(t, true, s.SrcNAT(), "srcnat") - - s.Value = StatusDstNAT - assert.Equal(t, true, s.DstNAT(), "dstnat") - - s.Value = StatusSeqAdjust - assert.Equal(t, true, s.SeqAdjust(), "seqadjust") - - s.Value = StatusSrcNATDone - assert.Equal(t, true, s.SrcNATDone(), "srcnatdone") - - s.Value = StatusDstNATDone - assert.Equal(t, true, s.DstNATDone(), "dstnatdone") - - s.Value = StatusDying - assert.Equal(t, true, s.Dying(), "dying") - - s.Value = StatusFixedTimeout - assert.Equal(t, true, s.FixedTimeout(), "fixedtimeout") - - s.Value = StatusTemplate - assert.Equal(t, true, s.Template(), "template") - - s.Value = StatusHelper - assert.Equal(t, true, s.Helper(), "helper") - - s.Value = StatusOffload - assert.Equal(t, true, s.Offload(), "offload") + assert.Equal(t, true, StatusExpected.Expected(), "expected") + assert.Equal(t, true, StatusSeenReply.SeenReply(), "seenreply") + assert.Equal(t, true, StatusAssured.Assured(), "assured") + assert.Equal(t, true, StatusConfirmed.Confirmed(), "confirmed") + assert.Equal(t, true, StatusSrcNAT.SrcNAT(), "srcnat") + assert.Equal(t, true, StatusDstNAT.DstNAT(), "dstnat") + assert.Equal(t, true, StatusSeqAdjust.SeqAdjust(), "seqadjust") + assert.Equal(t, true, StatusSrcNATDone.SrcNATDone(), "srcnatdone") + assert.Equal(t, true, StatusDstNATDone.DstNATDone(), "dstnatdone") + assert.Equal(t, true, StatusDying.Dying(), "dying") + assert.Equal(t, true, StatusFixedTimeout.FixedTimeout(), "fixedtimeout") + assert.Equal(t, true, StatusTemplate.Template(), "template") + assert.Equal(t, true, StatusHelper.Helper(), "helper") + assert.Equal(t, true, StatusOffload.Offload(), "offload") } func TestStatusString(t *testing.T) { - full := Status{Value: 0xffffffff} - empty := Status{} + full, empty := Status(0xffffffff), Status(0) wantFull := "EXPECTED|SEEN_REPLY|ASSURED|CONFIRMED|SRC_NAT|DST_NAT|SEQ_ADJUST|SRC_NAT_DONE|DST_NAT_DONE|" + "DYING|FIXED_TIMEOUT|TEMPLATE|UNTRACKED|HELPER|OFFLOAD" diff --git a/string.go b/string.go index 574edb4..e14df4b 100644 --- a/string.go +++ b/string.go @@ -51,7 +51,7 @@ func (s Status) String() string { // Loop over the field's bits for i, name := range names { - if s.Value&(1<