Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
106 changes: 86 additions & 20 deletions containers/container.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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}
Expand All @@ -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
}
Expand All @@ -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
Expand All @@ -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,
Expand All @@ -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()
}
Comment thread
mayankpande88 marked this conversation as resolved.

// canTrackConnection reports whether pid+fd k may be added to
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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)) {
Comment thread
mayankpande88 marked this conversation as resolved.
// 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
Expand Down
193 changes: 193 additions & 0 deletions containers/socket_connection_test.go
Original file line number Diff line number Diff line change
@@ -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.
Expand Down Expand Up @@ -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
}
Loading