Change ioctl_sniffer to use sockets instead of pipes.

This allows the tool to work with programs that use concurrent subprocesses.

Fixes #10542.

PiperOrigin-RevId: 647058279
This commit is contained in:
Anthony Cui
2024-06-26 13:27:59 -07:00
committed by gVisor bot
parent 344f19745f
commit b0abb4e7fa
6 changed files with 235 additions and 43 deletions
+2 -1
View File
@@ -24,6 +24,7 @@
#include <unistd.h>
#include <cstring>
#include <iostream>
#include <string>
#include "absl/strings/match.h"
@@ -42,7 +43,7 @@ void init_libc_ioctl_handle() {
libc_ioctl_handle = (libc_ioctl)dlsym(RTLD_NEXT, "ioctl");
if (!libc_ioctl_handle) {
printf("Failed to hook ioctl: %s\n", dlerror());
std::cerr << "Failed to hook ioctl: " << dlerror() << "\n";
exit(1);
}
}
+26 -17
View File
@@ -16,6 +16,7 @@
package main
import (
"context"
"flag"
"fmt"
"os"
@@ -53,7 +54,7 @@ func createSharedObject() (*os.File, error) {
}
// Main is our main function.
func Main() error {
func Main(ctx context.Context) error {
flag.Parse()
if len(flag.Args()) == 0 {
return fmt.Errorf("no command specified")
@@ -78,47 +79,55 @@ func Main() error {
}
}()
// Create a pipe to read the output of the command.
r, w, err := os.Pipe()
if err != nil {
return fmt.Errorf("failed to create pipe: %w", err)
// Start the sniffer server.
server := sniffer.NewServer()
if err := server.Listen(); err != nil {
return fmt.Errorf("failed to start sniffer server: %w", err)
}
serveCtx, serveCancel := context.WithCancel(ctx)
defer serveCancel()
go func() {
if err := server.Serve(serveCtx); err != nil {
log.Warningf("failed to serve sniffer server: %w", err)
}
}()
// Set up command from flags
cmd := exec.Command(flag.Arg(0), flag.Args()[1:]...)
cmd.ExtraFiles = []*os.File{w}
cmd.Stdin = os.Stdin
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
// Refer to the hook file by file descriptor here as its named file no
// longer exists.
cmd.Env = append(os.Environ(), fmt.Sprintf("LD_PRELOAD=/proc/%d/fd/%d", os.Getpid(), hookFile.Fd()))
cmd.Env = append(cmd.Env, fmt.Sprintf("GVISOR_IOCTL_SNIFFER_SOCKET_PATH=%v", server.Addr()))
// Run the command and start reading the output.
if err := cmd.Start(); err != nil {
return fmt.Errorf("failed to run command: %w", err)
}
w.Close()
results := sniffer.ReadHookOutput(r)
// Once our command is done, we can close the sniffer server and print the
// results.
cmdErr := cmd.Wait()
serveCancel()
if *enforceCompatability && results.HasUnsupportedIoctl() {
return fmt.Errorf("unsupported ioctls found: %v", results)
}
// Once we've read all the output, print the list of missing ioctls.
// Merge results from each connection.
finalResults := server.AllResults()
log.Infof("============== Unsupported ioctls ==============")
log.Infof("%s", results)
log.Infof("%v", finalResults)
if err := cmd.Wait(); err != nil {
return fmt.Errorf("command exited with error: %w", err)
if cmdErr != nil {
return fmt.Errorf("command exited with error: %w", cmdErr)
}
return nil
}
func main() {
if err := Main(); err != nil {
ctx := context.Background()
if err := Main(ctx); err != nil {
log.Warningf("%v", err)
os.Exit(1)
}
+12 -2
View File
@@ -16,6 +16,7 @@
package sniffer
import (
"context"
"errors"
"fmt"
"io"
@@ -162,6 +163,15 @@ func (r *Results) HasUnsupportedIoctl() bool {
return false
}
// Merge merges the results from another Results object into this one.
func (r *Results) Merge(other *Results) {
for class := ioctlClass(0); class < _numClasses; class++ {
for _, ioctl := range other.unsupported[class] {
r.AddUnsupportedIoctl(ioctl)
}
}
}
// Init reads from nvproxy and sets up the supported ioctl maps.
func Init() error {
nvproxy.Init()
@@ -194,10 +204,10 @@ func Init() error {
}
// ReadHookOutput reads the output of the ioctl hook until an EOF is reached.
func ReadHookOutput(r io.Reader) *Results {
func (c Connection) ReadHookOutput(ctx context.Context) *Results {
res := NewResults()
for {
ioctlPB, err := ReadIoctlProto(r)
ioctlPB, err := c.ReadIoctlProto(ctx)
if err != nil {
if !errors.Is(err, io.EOF) {
log.Warningf("Error reading ioctl proto: %v", err)
+142 -12
View File
@@ -15,46 +15,176 @@
package sniffer
import (
"context"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"os"
"sync"
"time"
"google.golang.org/protobuf/proto"
pb "gvisor.dev/gvisor/tools/ioctl_sniffer/ioctl_go_proto"
)
var (
// Connection is a connection to the sniffer hook.
type Connection struct {
protoBytesBuf []byte
)
conn net.Conn
}
// ReadIoctlProto reads a single ioctl proto from the given reader. Our format is:
// readFullWithContext tries to fill the buffer with data from the connection. It returns an error
// once the context is cancelled and the read would block, or if the read fails.
func (c *Connection) readFullWithContext(ctx context.Context, buf []byte) error {
nread := 0
for {
// Don't block for long if we're cancelled.
timeout := time.Second
if ctx.Err() != nil {
timeout = time.Millisecond
}
if err := c.conn.SetDeadline(time.Now().Add(timeout)); err != nil {
return fmt.Errorf("failed to set deadline: %w", err)
}
n, err := c.conn.Read(buf[nread:])
if err != nil {
// Only retry if we're not cancelled.
if errors.Is(err, os.ErrDeadlineExceeded) && ctx.Err() == nil {
continue
}
return fmt.Errorf("failed to read from connection: %w", err)
}
nread += n
if nread == len(buf) {
break
}
}
return nil
}
// ReadIoctlProto reads a single ioctl proto from this connection. Our format is:
// - 8 byte little endian uint64 containing the size of the proto.
// - The proto bytes.
//
// This should match the format in sniffer_bridge.h.
func ReadIoctlProto(r io.Reader) (*pb.Ioctl, error) {
// Read next proto from pipe.
func (c *Connection) ReadIoctlProto(ctx context.Context) (*pb.Ioctl, error) {
// First read in proto size
var protoSizeBuf [8]byte
if _, err := io.ReadFull(r, protoSizeBuf[:]); err != nil {
if err := c.readFullWithContext(ctx, protoSizeBuf[:]); err != nil {
return nil, fmt.Errorf("failed to read proto size: %w", err)
}
protoSize := binary.LittleEndian.Uint64(protoSizeBuf[:])
// See if we need to reallocate the buffer.
if cap(protoBytesBuf) < int(protoSize) {
protoBytesBuf = make([]byte, protoSize)
if cap(c.protoBytesBuf) < int(protoSize) {
c.protoBytesBuf = make([]byte, protoSize)
} else {
protoBytesBuf = protoBytesBuf[:protoSize]
c.protoBytesBuf = c.protoBytesBuf[:protoSize]
}
if _, err := io.ReadFull(r, protoBytesBuf); err != nil {
// Read the proto data.
if err := c.readFullWithContext(ctx, c.protoBytesBuf); err != nil {
return nil, fmt.Errorf("failed to read proto data: %w", err)
}
// Unmarshal and parse proto.
ioctl := &pb.Ioctl{}
if err := proto.Unmarshal(protoBytesBuf, ioctl); err != nil {
if err := proto.Unmarshal(c.protoBytesBuf, ioctl); err != nil {
return nil, fmt.Errorf("failed to unmarshal proto: %w", err)
}
return ioctl, nil
}
// Server is a server that accepts connections from the sniffer hook. It reads ioctl protos from
// each connection and sends them to the results channel.
type Server struct {
resultsChan chan *Results
connectionsWG sync.WaitGroup
listener net.Listener
}
// NewServer creates a new Server.
func NewServer() *Server {
return &Server{
resultsChan: make(chan *Results),
}
}
// Listen opens a new socket server.
func (s *Server) Listen() error {
// Create a unique socket path for this process.
// Go will automatically delete the file when the socket is closed.
addr := fmt.Sprintf("/tmp/sniffer_bridge_%d.sock", os.Getpid())
l, err := net.Listen("unix", addr)
if err != nil {
return fmt.Errorf("failed to listen on socket: %w", err)
}
s.listener = l
return nil
}
// Serve opens a new socket server, continually accepts connections from the socket and
// reads ioctl protos from each connection. It blocks until the context is cancelled.
func (s *Server) Serve(ctx context.Context) error {
// Accept connections from the socket and read ioctl protos from each connection.
errChan := make(chan error)
go func() {
defer close(errChan)
for ctx.Err() == nil {
conn, err := s.listener.Accept()
if err != nil {
errChan <- fmt.Errorf("failed to accept connection: %w", err)
return
}
s.connectionsWG.Add(1)
go func() {
conn := Connection{conn: conn}
s.resultsChan <- conn.ReadHookOutput(ctx)
s.connectionsWG.Done()
}()
}
}()
// Wait for the context cancellation.
<-ctx.Done()
if err := s.listener.Close(); err != nil {
return fmt.Errorf("failed to close socket: %w", err)
}
for err := range errChan {
if errors.Is(err, net.ErrClosed) {
continue
}
return fmt.Errorf("failed to accept connection: %w", err)
}
return nil
}
// AllResults blocks until all connections have closed and returns an aggregate of all the results.
func (s *Server) AllResults() *Results {
// Wait for all connections to close.
// Do this in a separate goroutine so we can start reading from the results channel.
go func() {
s.connectionsWG.Wait()
close(s.resultsChan)
}()
finalResults := NewResults()
for results := range s.resultsChan {
finalResults.Merge(results)
}
return finalResults
}
// Addr returns the address of the socket.
func (s *Server) Addr() string {
return s.listener.Addr().String()
}
+48 -8
View File
@@ -16,18 +16,58 @@
#include <stdint.h>
#include <stdio.h>
#include <sys/socket.h>
#include <sys/un.h>
#include <unistd.h>
#include <cerrno>
#include <cstdlib>
#include <cstring>
#include <iostream>
#include <vector>
#include "tools/ioctl_sniffer/ioctl.pb.h"
#include "google/protobuf/io/zero_copy_stream_impl.h"
thread_local pid_t socket_owner_tid = -1;
thread_local int socket_fd = -1;
void InitializeSocket() {
int sfd = socket(AF_UNIX, SOCK_STREAM, 0);
if (sfd < 0) {
std::cerr << "Failed to create socket: " << strerror(errno) << "\n";
exit(1);
}
struct sockaddr_un addr;
addr.sun_family = AF_UNIX;
strncpy(addr.sun_path, std::getenv("GVISOR_IOCTL_SNIFFER_SOCKET_PATH"),
sizeof(addr.sun_path));
// Ensure the path is null terminated.
addr.sun_path[sizeof(addr.sun_path) - 1] = '\0';
if (connect(sfd, reinterpret_cast<sockaddr *>(&addr), sizeof(addr)) < 0) {
std::cerr << "Failed to connect to socket " << addr.sun_path << ": "
<< strerror(errno) << "\n";
exit(1);
}
socket_owner_tid = gettid();
socket_fd = sfd;
}
void WriteIoctlProto(gvisor::Ioctl &ioctl) {
// Write size of the proto message.
uint64_t size = ioctl.ByteSizeLong();
write(LOG_OUTPUT_FD, &size, sizeof(size));
if (socket_owner_tid != gettid()) {
InitializeSocket();
}
// Write the proto message.
google::protobuf::io::FileOutputStream os(LOG_OUTPUT_FD);
ioctl.SerializeToZeroCopyStream(&os);
os.Flush();
static thread_local std::vector<char> buffer;
uint64_t size = ioctl.ByteSizeLong();
buffer.resize(size + sizeof(size));
// Write size of the proto message first.
memcpy(buffer.data(), &size, sizeof(size));
ioctl.SerializeToArray(buffer.data() + sizeof(size), size);
write(socket_fd, buffer.data(), buffer.size());
}
+5 -3
View File
@@ -15,11 +15,11 @@
#ifndef TOOLS_IOCTL_SNIFFER_SNIFFER_BRIDGE_H_
#define TOOLS_IOCTL_SNIFFER_SNIFFER_BRIDGE_H_
#include <syscall.h>
#include "tools/ioctl_sniffer/ioctl.pb.h"
// The file descriptor to write the ioctl data to. Go's os.exec will
// always make this available.
constexpr int LOG_OUTPUT_FD = 3;
inline pid_t gettid() { return syscall(SYS_gettid); }
// Write the ioctl proto to the log output file descriptor. Our format is:
// - 8 byte little endian uint64 containing the size of the proto.
@@ -27,4 +27,6 @@ constexpr int LOG_OUTPUT_FD = 3;
// This should match the format in sniffer_bridge.go.
void WriteIoctlProto(gvisor::Ioctl &ioctl);
void InitializeSocket();
#endif // TOOLS_IOCTL_SNIFFER_SNIFFER_BRIDGE_H_