diff --git a/containers/container.go b/containers/container.go index 29b61ef..3ad9ede 100644 --- a/containers/container.go +++ b/containers/container.go @@ -102,6 +102,10 @@ type ActiveConnection struct { Timestamp uint64 Closed time.Time + // dst is the address the socket is connected to, before NAT. It tells + // this connection apart from another socket that had the same fd number. + dst netaddr.IPPort + BytesSent uint64 BytesReceived uint64 Protocol uint8 @@ -843,6 +847,7 @@ func (c *Container) onConnectionOpen(pid uint32, fd uint64, src, dst, actualDst Fd: fd, Timestamp: timestamp, srcWorkload: srcWorkload, + dst: dst, } c.activeConnections[ConnectionKey{src: src, dst: dst}] = connection k := PidFd{Pid: pid, Fd: fd} @@ -860,6 +865,29 @@ func (c *Container) onConnectionOpen(pid uint32, fd uint64, src, dst, actualDst // This is used when TCP connection tracking fails (common for Go TLS due to goroutine thread switching) // but we have socket tuple info extracted directly from the fd func (c *Container) createConnectionFromSocketInfo(pid uint32, fd uint64, timestamp uint64, socketInfo *ebpftracer.SocketInfo) (conn *ActiveConnection, filtered bool) { + connection, filtered := c.connectionFromSocketInfo(pid, fd, timestamp, socketInfo) + if connection == nil { + return nil, filtered + } + + // Store in connectionsByPidFd for future L7 events on same connection + k := PidFd{Pid: pid, Fd: fd} + if !c.canTrackConnection(k) { + ConnectionCapDropsTotal.Inc() + return nil, false + } + c.connectionsByPidFd[k] = connection + + klog.V(3).Infof("L7_CONN_CREATED_FROM_SOCKET: pid=%d fd=%d dst=%s actual_dst=%s", + pid, fd, connection.dst, connection.DestinationKey.ActualDestinationIfKnown()) + + return connection, false +} + +// connectionFromSocketInfo builds the connection an L7 event's socket tuple +// describes, without tracking it. filtered is true for a connection the agent +// does not track. +func (c *Container) connectionFromSocketInfo(pid uint32, fd uint64, timestamp uint64, socketInfo *ebpftracer.SocketInfo) (conn *ActiveConnection, filtered bool) { if socketInfo == nil || !socketInfo.Valid { return nil, false } @@ -871,21 +899,18 @@ func (c *Container) createConnectionFromSocketInfo(pid uint32, fd uint64, timest return nil, true } - // Parse destination IP - dstIP, err := netaddr.ParseIP(socketInfo.DstIP) - if err != nil { - klog.V(2).Infof("createConnectionFromSocketInfo: failed to parse dst IP %s: %v", socketInfo.DstIP, err) + dst, ok := socketDestination(socketInfo) + if !ok { + klog.V(2).Infof("connectionFromSocketInfo: failed to parse dst IP %s", socketInfo.DstIP) return nil, false } // Parse source IP srcIP, err := netaddr.ParseIP(socketInfo.SrcIP) if err != nil { - klog.V(2).Infof("createConnectionFromSocketInfo: failed to parse src IP %s: %v", socketInfo.SrcIP, err) + klog.V(2).Infof("connectionFromSocketInfo: failed to parse src IP %s: %v", socketInfo.SrcIP, err) return nil, false } - - dst := netaddr.IPPortFrom(dstIP, socketInfo.DstPort) src := netaddr.IPPortFrom(srcIP, socketInfo.SrcPort) // The socket holds the address the application connected to, before any @@ -903,8 +928,7 @@ func (c *Container) createConnectionFromSocketInfo(pid uint32, fd uint64, timest return nil, true } - // Create connection - connection := &ActiveConnection{ + return &ActiveConnection{ DestinationKey: key, Pid: pid, Fd: fd, @@ -913,20 +937,41 @@ func (c *Container) createConnectionFromSocketInfo(pid uint32, fd uint64, timest // in onL7RequestWithResult and was dropped. Timestamp: timestamp, srcWorkload: srcWorkload, - } + dst: dst, + }, false +} - // Store in connectionsByPidFd for future L7 events on same connection - k := PidFd{Pid: pid, Fd: fd} - if !c.canTrackConnection(k) { - ConnectionCapDropsTotal.Inc() - return nil, false +// socketDestination returns the destination of an L7 event's socket tuple. +func socketDestination(si *ebpftracer.SocketInfo) (netaddr.IPPort, bool) { + if si == nil || !si.Valid { + return netaddr.IPPort{}, false } - c.connectionsByPidFd[k] = connection - - klog.V(3).Infof("L7_CONN_CREATED_FROM_SOCKET: pid=%d fd=%d src=%s dst=%s actual_dst=%s", - pid, fd, src, dst, key.ActualDestinationIfKnown()) + ip, err := netaddr.ParseIP(si.DstIP) + if err != nil { + return netaddr.IPPort{}, false + } + return netaddr.IPPortFrom(ip, si.DstPort), true +} - return connection, false +// isSocket reports whether conn is the socket an L7 event's tuple, read from +// the fd as the event happened, describes. The kernel reuses an fd number as +// soon as its socket is closed, so the entry tracked for a pid and fd can be +// an earlier socket's: L7 events are handled as they come, connection events +// later. It is true when there is nothing to compare, with no tuple or with +// a connection whose address is unknown. +func (conn *ActiveConnection) isSocket(si *ebpftracer.SocketInfo) bool { + if si == nil || !si.Valid || conn.dst.IP().IsZero() { + return true + } + // The port is enough to tell most reuses apart without parsing the IP. + if si.DstPort != conn.dst.Port() { + return false + } + dst, ok := socketDestination(si) + if !ok { + return true + } + return dst.IP().Unmap() == conn.dst.IP().Unmap() } // canTrackConnection reports whether pid+fd k may be added to @@ -1146,6 +1191,12 @@ func (c *Container) onL7RequestWithResult(pid uint32, fd uint64, timestamp uint6 } pidFd := PidFd{Pid: pid, Fd: fd} conn := c.connectionsByPidFd[pidFd] + if conn != nil && !conn.isSocket(socketInfo) { + // An earlier socket on this fd. Naming its address after this + // ClientHello's host is how a resolver came to be named after + // whatever host was connected to right after resolving it. + conn = nil + } destIP := netaddr.IP{} if conn != nil { destIP = conn.DestinationKey.ActualDestinationIfKnown().IP() @@ -1184,6 +1235,21 @@ func (c *Container) onL7RequestWithResult(pid uint32, fd uint64, timestamp uint6 c.detectLLMEndpoint(pid, fd, timestamp, r, socketInfo) conn := c.connectionsByPidFd[PidFd{Pid: pid, Fd: fd}] + if r.Protocol == l7.ProtocolDNS && socketInfo != nil && socketInfo.Valid && + (conn == nil || !conn.isSocket(socketInfo)) { + // DNS over UDP is never tracked: a UDP socket has no close event, so + // its entry would outlive it and be taken for the next socket on its + // fd, usually the connection to the address just resolved. For the + // same reason an entry on the fd that is not this socket is an + // earlier socket's. The query takes its connection from its own tuple. + var filtered bool + if conn, filtered = c.connectionFromSocketInfo(pid, fd, timestamp, socketInfo); conn == nil { + if !filtered { + dropL7Event(c.id, "unknown_connection", pid, fd, r, socketInfo) + } + return nil, L7RequestProcessed + } + } if conn == nil { // TCP connection tracking failed - common for Go TLS due to goroutine thread switching // Try to create connection from socket info extracted directly from fd in eBPF diff --git a/containers/socket_connection_test.go b/containers/socket_connection_test.go index 7aa71d5..3cd0f51 100644 --- a/containers/socket_connection_test.go +++ b/containers/socket_connection_test.go @@ -1,12 +1,19 @@ package containers import ( + "crypto/tls" + "net" "testing" + "time" "github.com/coroot/coroot-node-agent/common" "github.com/coroot/coroot-node-agent/ebpftracer" "github.com/coroot/coroot-node-agent/ebpftracer/l7" "github.com/coroot/coroot-node-agent/flags" + "github.com/prometheus/client_golang/prometheus" + dto "github.com/prometheus/client_model/go" + "golang.org/x/net/dns/dnsmessage" + "inet.af/netaddr" ) // stubResolver names every IP after a fixed table, falling back to the IP. @@ -118,3 +125,189 @@ func TestDropStaleHTTP2Parser(t *testing.T) { } (&Container{}).dropStaleHTTP2Parser(k, 1) // no parsers yet: must not panic } + +// isSocket tells the tracked entry of an fd apart from a later socket that got +// the same fd number, and gives the benefit of the doubt when it cannot tell. +func TestConnectionIsSocket(t *testing.T) { + conn := &ActiveConnection{dst: netaddr.MustParseIPPort("192.0.2.53:53")} + for _, tc := range []struct { + name string + info *ebpftracer.SocketInfo + want bool + }{ + {"same socket", socketInfo("10.0.0.2", 40000, "192.0.2.53", 53), true}, + {"same socket, IPv4-mapped", socketInfo("::ffff:10.0.0.2", 40000, "::ffff:192.0.2.53", 53), true}, + {"other address", socketInfo("10.0.0.2", 40001, "203.0.113.9", 443), false}, + {"other port", socketInfo("10.0.0.2", 40001, "192.0.2.53", 853), false}, + {"no tuple", nil, true}, + {"invalid tuple", &ebpftracer.SocketInfo{DstIP: "203.0.113.9", DstPort: 443}, true}, + } { + if got := conn.isSocket(tc.info); got != tc.want { + t.Errorf("%s: isSocket = %v, want %v", tc.name, got, tc.want) + } + } + if !(&ActiveConnection{}).isSocket(socketInfo("10.0.0.2", 40001, "203.0.113.9", 443)) { + t.Error("a connection with no known address must not be taken for another socket") + } +} + +// A resolver must keep its own name. An application resolves a host over UDP, +// closes the socket, and connects to the host on the same fd number. When the +// UDP socket was tracked, its entry outlived it and the ClientHello on the new +// connection named the resolver after the host; every later query to the +// resolver then carried that host as its destination, a new set of series for +// every host the container connected to. +func TestDNSResolverKeepsItsNameAcrossReusedFd(t *testing.T) { + c := newL7TestContainer(t) + const ( + resolver = "192.0.2.53" + hostIP = "203.0.113.9" + ) + dnsTuple := socketInfo("10.0.0.2", 40000, resolver, 53) + tlsTuple := socketInfo("10.0.0.2", 40001, hostIP, 443) + + c.handleL7(t, 1, 7, 0, dnsResponse(t, "api.example.com", hostIP), dnsTuple) + if _, ok := c.connectionsByPidFd[PidFd{Pid: 1, Fd: 7}]; ok { + t.Error("the UDP socket of a DNS query was tracked") + } + + ip2fqdn := c.handleL7(t, 1, 7, 100, clientHello(t, "api.example.com"), tlsTuple) + if d := ip2fqdn[netaddr.MustParseIP(hostIP)]; d == nil || d.FQDN != "api.example.com" { + t.Errorf("ClientHello did not name %s: %v", hostIP, ip2fqdn) + } + + c.handleL7(t, 1, 7, 0, dnsResponse(t, "www.example.org", "198.51.100.7"), dnsTuple) + if d := c.registry.getDomain(netaddr.MustParseIP(resolver)); d != nil { + t.Errorf("resolver %s named %q", resolver, d.FQDN) + } + if got := c.dnsDestinations(t); len(got) != 1 || !got[resolver+":53"] { + t.Errorf("DNS destinations = %v, want only %s:53", got, resolver) + } +} + +// The tracked entry of an fd may be an earlier socket's: this socket's open +// event is handled after its first write. A ClientHello names the address of +// its own socket, and a DNS query on that fd is not counted against the +// earlier connection. +func TestEarlierSocketOnFdIsNotUsed(t *testing.T) { + c := newL7TestContainer(t) + earlier, _ := c.createConnectionFromSocketInfo(1, 7, 0, socketInfo("10.0.0.2", 40000, "192.0.2.53", 53)) + if earlier == nil { + t.Fatal("no connection for the earlier socket") + } + + ip2fqdn := c.handleL7(t, 1, 7, 100, clientHello(t, "api.example.com"), socketInfo("10.0.0.2", 40001, "203.0.113.9", 443)) + if _, ok := ip2fqdn[netaddr.MustParseIP("192.0.2.53")]; ok { + t.Errorf("ClientHello named the earlier socket's address: %v", ip2fqdn) + } + if d := ip2fqdn[netaddr.MustParseIP("203.0.113.9")]; d == nil || d.FQDN != "api.example.com" { + t.Errorf("ClientHello did not name its own socket's address: %v", ip2fqdn) + } + if got := earlier.DestinationKey.DestinationLabelValue(); got != "192.0.2.53:53" { + t.Errorf("earlier connection renamed to %q", got) + } + + c.connectionsByPidFd[PidFd{Pid: 1, Fd: 7}], _ = c.connectionFromSocketInfo(1, 7, 100, socketInfo("10.0.0.2", 40001, "203.0.113.9", 443)) + c.handleL7(t, 1, 7, 0, dnsResponse(t, "www.example.org", "198.51.100.7"), socketInfo("10.0.0.2", 40002, "192.0.2.53", 53)) + if got := c.dnsDestinations(t); len(got) != 1 || !got["192.0.2.53:53"] { + t.Errorf("DNS destinations = %v, want only 192.0.2.53:53", got) + } +} + +func newL7TestContainer(t *testing.T) *Container { + t.Helper() + // Without a parsed command line no public network is tracked. + common.ConnectionFilter.WhitelistPrefix(netaddr.MustParseIPPrefix("0.0.0.0/0")) + prevMax := *flags.MaxFQDNsPerContainer + *flags.MaxFQDNsPerContainer = 50 + t.Cleanup(func() { *flags.MaxFQDNsPerContainer = prevMax }) + c := newSocketTestContainer(t, nil) + c.l7Stats = NewL7Stats(nil) + c.googleHTTP2Parsers = map[PidFd]*l7.Http2Parser{} + c.llmCaptures = map[PidFd]*llmCapture{} + return c +} + +// handleL7 processes an L7 event as the registry does: the IP names it +// returns are recorded before the next event. +func (c *Container) handleL7(t *testing.T, pid uint32, fd uint64, ts uint64, r *l7.RequestData, si *ebpftracer.SocketInfo) map[netaddr.IP]*common.Domain { + t.Helper() + ip2fqdn, result := c.onL7RequestWithResult(pid, fd, ts, r, si) + if result != L7RequestProcessed { + t.Fatalf("protocol %d: result %v, want processed", r.Protocol, result) + } + for ip, d := range ip2fqdn { + c.registry.ip2fqdn.Put(ip, d) + } + return ip2fqdn +} + +// dnsDestinations returns the destination label values of the DNS counter. +func (c *Container) dnsDestinations(t *testing.T) map[string]bool { + t.Helper() + ch := make(chan prometheus.Metric, 100) + c.l7Stats.requests[l7.ProtocolDNS].Collect(ch) + close(ch) + got := map[string]bool{} + for m := range ch { + var pb dto.Metric + if err := m.Write(&pb); err != nil { + t.Fatal(err) + } + for _, l := range pb.GetLabel() { + if l.GetName() == "destination" { + got[l.GetValue()] = true + } + } + } + return got +} + +func dnsResponse(t *testing.T, name, ip string) *l7.RequestData { + t.Helper() + b := dnsmessage.NewBuilder(nil, dnsmessage.Header{Response: true}) + b.EnableCompression() + q := dnsmessage.Question{Name: dnsmessage.MustNewName(name + "."), Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET} + if err := b.StartQuestions(); err != nil { + t.Fatal(err) + } + if err := b.Question(q); err != nil { + t.Fatal(err) + } + if err := b.StartAnswers(); err != nil { + t.Fatal(err) + } + if err := b.AResource(dnsmessage.ResourceHeader{Name: q.Name, Class: dnsmessage.ClassINET, TTL: 60}, dnsmessage.AResource{A: netaddr.MustParseIP(ip).As4()}); err != nil { + t.Fatal(err) + } + payload, err := b.Finish() + if err != nil { + t.Fatal(err) + } + return &l7.RequestData{Protocol: l7.ProtocolDNS, Payload: payload, Duration: time.Millisecond} +} + +// clientHello returns a TLS ClientHello event for serverName, as crypto/tls +// writes it. +func clientHello(t *testing.T, serverName string) *l7.RequestData { + t.Helper() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + hello := make(chan []byte, 1) + go func() { + buf := make([]byte, 4096) + n, _ := server.Read(buf) + hello <- buf[:n] + }() + tc := tls.Client(client, &tls.Config{ServerName: serverName, InsecureSkipVerify: true}) + _ = tc.SetDeadline(time.Now().Add(time.Second)) + go func() { _ = tc.Handshake() }() + select { + case payload := <-hello: + return &l7.RequestData{Protocol: l7.ProtocolTLSClientHello, Payload: payload} + case <-time.After(3 * time.Second): + t.Fatal("no ClientHello written") + } + return nil +}