Skip to content
Open
Show file tree
Hide file tree
Changes from 4 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
209 changes: 201 additions & 8 deletions ebpf/connections/factory.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@ func NewFactory() *Factory {
}
}

func convertToSingleByteArr(bufMap map[int][]byte) []byte {
func convertToSingleByteArr(bufMap map[int][]byte, isWebsocket bool) []byte {

if len(bufMap) == 0 {
return make([]byte, 0)
Expand All @@ -59,17 +59,17 @@ func convertToSingleByteArr(bufMap map[int][]byte) []byte {
for _, k := range keys {
if kPrev == -1 {
// C sets read, write event count=0 only on new connection open
// For requests arriving after a time gap on the same underlying connection the
// For requests arriving after a time gap on the same underlying connection the
// read,write count will not be 1, they will simply continue from the last request
// This can only be replicated when there is a time gap/inactivityThreshold between requests
// on the same underlying connection
if !sequenceCheckSkip && k != 1 {
if !isWebsocket && !sequenceCheckSkip && k != 1 {
utils.LogProcessing("Bad start sequence", "key", k, "value", string(bufMap[k]))
break
}
kPrev = k
} else {
if kPrev+1 != k {
if !isWebsocket && kPrev+1 != k {
utils.LogProcessing("Missing sequence", "prev", kPrev, "current", k, "value", string(bufMap[k]), "prevValue", string(bufMap[kPrev]))
break
}
Expand All @@ -93,6 +93,10 @@ var (
trackerDataProcessInterval = 100

socketDataEventBytesThreshold = 10 * 1024 * 1024

// WebSocket connection tracking
wsConnectionsMutex sync.RWMutex
wsConnections = make(map[string]bool) // keyed by "IP:Port:Pid:Fd"
)

func init() {
Expand All @@ -105,15 +109,73 @@ func init() {
utils.InitVar("SOCKET_DATA_EVENT_BYTES_THRESHOLD", &socketDataEventBytesThreshold)
}

func buildWebSocketConnectionKey(connID structs.ConnID) string {
// Use only ConnID fields which are stable kernel identifiers
return fmt.Sprintf("%d:%d", connID.Id, connID.Fd)
}

func isWebSocketConnection(connID structs.ConnID) bool {
wsConnectionsMutex.RLock()
defer wsConnectionsMutex.RUnlock()
key := buildWebSocketConnectionKey(connID)
return wsConnections[key]
}

func markAsWebSocketConnection(connID structs.ConnID) {
wsConnectionsMutex.Lock()
defer wsConnectionsMutex.Unlock()
key := buildWebSocketConnectionKey(connID)
wsConnections[key] = true
slog.Info("Marked connection as WebSocket", "key", key)
}

func unmarkWebSocketConnection(connID structs.ConnID) {
wsConnectionsMutex.Lock()
defer wsConnectionsMutex.Unlock()
key := buildWebSocketConnectionKey(connID)
delete(wsConnections, key)
slog.Info("Unmarked WebSocket connection (closed)", "key", key)
}

// hasCloseFrame checks if any of the frames contain a WebSocket close frame (opcode 8)
func hasCloseFrame(frames []kafkaUtil.WebSocketMessage) bool {
for _, f := range frames {
if f.Opcode == 0x8 {
return true
}
}
return false
}

// extractHTTPHeadersOnly returns the buffer up to and including \r\n\r\n (HTTP headers only)
// This strips any WebSocket frames that may follow the headers
func extractHTTPHeadersOnly(buffer []byte) []byte {
if len(buffer) < 4 {
return buffer
}

headerEnd := bytes.Index(buffer, []byte("\r\n\r\n"))
if headerEnd == -1 {

return buffer
}

return buffer[:headerEnd+4]
}

func ProcessTrackerData(connID structs.ConnID, tracker *Tracker, isComplete bool) {
tracker.mutex.Lock()
defer tracker.mutex.Unlock()

if len(tracker.sentBuf) == 0 || len(tracker.recvBuf) == 0 {
return
}
receiveBuffer := convertToSingleByteArr(tracker.recvBuf)
sentBuffer := convertToSingleByteArr(tracker.sentBuf)

// Check if this is already a known WebSocket connection - if so, skip sequence checks
// since binary WebSocket frames won't have sequential packet numbering
isKnownWebSocket := isWebSocketConnection(connID)
receiveBuffer := convertToSingleByteArr(tracker.recvBuf, isKnownWebSocket)
sentBuffer := convertToSingleByteArr(tracker.sentBuf, isKnownWebSocket)

originalInt := uint32(connID.Ip)
// Convert integer to little-endian byte slice
Expand All @@ -134,17 +196,140 @@ func ProcessTrackerData(connID structs.ConnID, tracker *Tracker, isComplete bool
hostName = kafkaUtil.PodInformerInstance.GetPodNameByProcessId(int32(connID.Id >> 32))
}

if len(sentBuffer) >= len(httpBytes) && (bytes.Equal(sentBuffer[:len(httpBytes)], httpBytes)) {
if isWebSocketConnection(connID) {
processWebSocketConnection(connID, receiveBuffer, sentBuffer, isComplete, hostName)
return
}

isWebSocket := isWebSocketUpgradeData(sentBuffer, receiveBuffer)

if isWebSocket {
markAsWebSocketConnection(connID)

headers := extractHeadersFromHTTPBuffer(receiveBuffer)

kafkaUtil.WSConnectionManager.RegisterConnectionNumeric(connID.Id, connID.Fd, connID.Ip, connID.Port, tracker.srcIp, tracker.srcPort, headers)

// Send the initial handshake as a normal HTTP request/response pair
// Strip WebSocket frames from buffers first - only pass HTTP headers to tryReadFromBD
httpReceiveBuffer := extractHTTPHeadersOnly(receiveBuffer)
httpSentBuffer := extractHTTPHeadersOnly(sentBuffer)
tryReadFromBD(destIpStr, srcIpStr, httpReceiveBuffer, httpSentBuffer, isComplete, 1, connID.Id, connID.Fd, uniqueDaemonsetId, hostName)
} else if len(sentBuffer) >= len(httpBytes) && (bytes.Equal(sentBuffer[:len(httpBytes)], httpBytes)) {
tryReadFromBD(destIpStr, srcIpStr, receiveBuffer, sentBuffer, isComplete, 1, connID.Id, connID.Fd, uniqueDaemonsetId, hostName)
}
if !disableEgress {
if !disableEgress && !isWebSocket {
// attempt to parse the egress as well by switching the recv and sent buffers.
if len(receiveBuffer) >= len(httpBytes) && (bytes.Equal(receiveBuffer[:len(httpBytes)], httpBytes)) {
tryReadFromBD(srcIpStr, destIpStr, sentBuffer, receiveBuffer, isComplete, 2, connID.Id, connID.Fd, uniqueDaemonsetId, hostName)
}
}
}

func isWebSocketUpgradeData(sentBuffer, receiveBuffer []byte) bool {
if len(sentBuffer) == 0 || len(receiveBuffer) == 0 {
return false
}

if bytes.Contains(sentBuffer, []byte(":9092")) || bytes.Contains(receiveBuffer, []byte(":9092")) {
return false
}

// Check for 101 Switching Protocols in sent buffer (response)
has101 := bytes.Contains(sentBuffer, []byte("101 Switching Protocols"))

hasUpgradeResponseHeader := bytes.Contains(sentBuffer, []byte("Upgrade: websocket")) ||
bytes.Contains(sentBuffer, []byte("upgrade: websocket"))

hasUpgradeRequestHeader := bytes.Contains(receiveBuffer, []byte("Upgrade: websocket")) ||
bytes.Contains(receiveBuffer, []byte("upgrade: websocket"))

hasConnectionResponseHeader := bytes.Contains(sentBuffer, []byte("Connection: Upgrade")) ||
bytes.Contains(sentBuffer, []byte("Connection: upgrade")) ||
bytes.Contains(sentBuffer, []byte("connection: Upgrade")) ||
bytes.Contains(sentBuffer, []byte("connection: upgrade"))

hasConnectionRequestHeader := bytes.Contains(receiveBuffer, []byte("Connection: Upgrade")) ||
bytes.Contains(receiveBuffer, []byte("Connection: upgrade")) ||
bytes.Contains(receiveBuffer, []byte("connection: Upgrade")) ||
bytes.Contains(receiveBuffer, []byte("connection: upgrade"))

return has101 && hasUpgradeResponseHeader && hasUpgradeRequestHeader &&
hasConnectionResponseHeader && hasConnectionRequestHeader
}

func extractHeadersFromHTTPBuffer(buffer []byte) map[string]string {
headers := make(map[string]string)
if len(buffer) == 0 {
return headers
}

lines := bytes.Split(buffer, []byte("\r\n"))

for i := 1; i < len(lines); i++ {
line := lines[i]

if len(line) == 0 {
break
}

parts := bytes.SplitN(line, []byte(":"), 2)
if len(parts) == 2 {
key := string(bytes.TrimSpace(parts[0]))
value := string(bytes.TrimSpace(parts[1]))
headers[key] = value
}
}

return headers
}

func processWebSocketConnection(connID structs.ConnID, receiveBuffer, sentBuffer []byte, isComplete bool, hostName string) {
connectionClosed := false

if len(sentBuffer) > 0 {
frames := kafkaUtil.ParseWebSocketFrames(sentBuffer, "outgoing")
if len(frames) > 0 {
if hasCloseFrame(frames) {
connectionClosed = true
}
err := kafkaUtil.WSConnectionManager.AccumulateMessagesNumeric(
connID.Id,
connID.Fd,
frames,
)
if err != nil {
slog.Debug("Failed to accumulate WebSocket frames from sent buffer", "error", err)
}
}
}

if len(receiveBuffer) > 0 {
frames := kafkaUtil.ParseWebSocketFrames(receiveBuffer, "incoming")
if len(frames) > 0 {
if hasCloseFrame(frames) {
connectionClosed = true
}
err := kafkaUtil.WSConnectionManager.AccumulateMessagesNumeric(
connID.Id,
connID.Fd,
frames,
)
if err != nil {
slog.Debug("Failed to accumulate WebSocket frames from receive buffer", "error", err)
}
}
}

// If a close frame was seen, unmark so the fd can be reused cleanly by a new connection
if connectionClosed {
unmarkWebSocketConnection(connID)
kafkaUtil.WSConnectionManager.RemoveConnectionByConnID(connID.Id, connID.Fd)
}

slog.Debug("WebSocket connection processed", "connID", connID, "isComplete", isComplete)
}

func (factory *Factory) CanBeFilled() bool {
factory.mutex.RLock()
defer factory.mutex.RUnlock()
Expand Down Expand Up @@ -310,6 +495,11 @@ func (factory *Factory) DeleteWorker(connectionID structs.ConnID) {
}

if _, exists := factory.connections[connectionID]; exists {

if _, ok := factory.connections[connectionID]; ok {
// Don't unmark WebSocket connections - they maintain their type across the connection lifetime
// even if factory.connections is cleaned up during inactivity
}
delete(factory.connections, connectionID)
utils.LogProcessing("Deleted connection", "fd", connectionID.Fd, "id", connectionID.Id, "timestamp", connectionID.Conn_start_ns, "ip", connectionID.Ip, "port", connectionID.Port)
requestProcessCount++
Expand All @@ -332,6 +522,9 @@ func (factory *Factory) DeleteWorker(connectionID structs.ConnID) {
close(ch)
delete(factory.processor, key)
}
if _, ok := factory.connections[key]; ok {
// Don't unmark WebSocket connections - they maintain their type across the connection lifetime
}
delete(factory.connections, key)
}
}
Expand Down
94 changes: 94 additions & 0 deletions trafficUtil/kafkaUtil/kafka.go
Original file line number Diff line number Diff line change
Expand Up @@ -143,11 +143,105 @@ func InitKafka() {
// Start heartbeat routine
go sendKafkaHeartbeat()
slog.Info("Started Kafka heartbeat routine", "interval_seconds", heartbeatIntervalSeconds)

go sendWebSocketBatches()
slog.Info("Started WebSocket batch sender routine", "interval_minutes", 1)
break
}
}
}

func sendWebSocketBatches() {
ticker := time.NewTicker(1 * time.Minute)
defer ticker.Stop()

for range ticker.C {
if WSConnectionManager == nil {
continue
}

connections := WSConnectionManager.GetAllConnections()
slog.Debug("WebSocket batch sender tick", "active_connections", len(connections))

for _, key := range connections {

parts := strings.Split(key, ":")
if len(parts) != 2 {
slog.Warn("Invalid WebSocket connection key format", "key", key)
continue
}

connIDId, err := strconv.ParseUint(parts[0], 10, 64)
if err != nil {
slog.Warn("Failed to parse ConnID Id", "key", key, "error", err)
continue
}
connIDFd, err := strconv.ParseUint(parts[1], 10, 32)
if err != nil {
slog.Warn("Failed to parse ConnID Fd", "key", key, "error", err)
continue
}

batch, err := WSConnectionManager.GetAndClearBatchByConnID(connIDId, uint32(connIDFd))
if err != nil {
slog.Debug("Failed to get WebSocket batch", "key", key, "error", err)
continue
}

if len(batch.Messages) == 0 {
continue
}

eventsJSON, _ := json.Marshal(convertMessagesToEventList(batch.Messages))
headersJSON, _ := json.Marshal(batch.Connection.Headers)

payload := map[string]string{
"ip": batch.Connection.SourceIP,
"destIp": batch.Connection.DestIP,
"time": fmt.Sprint(batch.BatchTime.Unix()),
"akto_account_id": fmt.Sprint(1000000),
"akto_vxlan_id": fmt.Sprint(0),
"source": "MIRRORING",
"connection_type": "WEBSOCKET",
"headers": string(headersJSON),
"events": string(eventsJSON),
}

out, err := json.Marshal(payload)
if err != nil {
slog.Error("Failed to marshal WebSocket batch payload", "error", err)
continue
}

ctx := context.Background()
sourceIP := batch.Connection.SourceIP
err = ProduceStr(ctx, string(out), key, sourceIP, "WEBSOCKET")
if err != nil {
slog.Error("Failed to write WebSocket batch to Kafka", "error", err)
} else {
slog.Debug("WebSocket batch sent to Kafka", "messages", len(batch.Messages), "source", sourceIP)
}
}
}
}

// convertMessagesToEventList converts WebSocketMessage slice to a list of event maps
func convertMessagesToEventList(messages []WebSocketMessage) []map[string]interface{} {
eventList := make([]map[string]interface{}, 0)
for _, msg := range messages {
eventList = append(eventList, map[string]interface{}{
"event_type": msg.EventType,
"payload": msg.Payload,
"direction": msg.Direction,
"timestamp": msg.Timestamp.Unix(),
"fin": msg.FIN,
"opcode": msg.Opcode,
"masked": msg.Masked,
})
}
return eventList
}

func kafkaCompletion() func(messages []kafka.Message, err error) {
return func(messages []kafka.Message, err error) {
if err != nil {
Expand Down
Loading