diff --git a/test/trace/BUILD b/test/trace/BUILD index 74597ab6d..bac34d6b2 100644 --- a/test/trace/BUILD +++ b/test/trace/BUILD @@ -22,6 +22,7 @@ go_test( "//pkg/test/testutil", "//test/trace/config", "@org_golang_google_protobuf//proto:go_default_library", + "@org_golang_x_sys//unix:go_default_library", ], ) diff --git a/test/trace/trace_test.go b/test/trace/trace_test.go index 617c7f0b8..ebd3846eb 100644 --- a/test/trace/trace_test.go +++ b/test/trace/trace_test.go @@ -23,6 +23,7 @@ import ( "testing" "time" + "golang.org/x/sys/unix" "google.golang.org/protobuf/proto" "gvisor.dev/gvisor/pkg/sentry/seccheck" "gvisor.dev/gvisor/pkg/sentry/seccheck/checkers/remote/test" @@ -31,6 +32,8 @@ import ( "gvisor.dev/gvisor/test/trace/config" ) +var cutoffTime time.Time + // TestAll enabled all trace points in the system with all optional and context // fields enabled. Then it runs a workload that will trigger those points and // run some basic validation over the points generated. @@ -68,6 +71,8 @@ func TestAll(t *testing.T) { if err != nil { t.Fatal(err) } + // No trace point should have a time lesser than this. + cutoffTime = time.Now() cmd := exec.Command( runsc, "--debug", "--alsologtostderr", // Debug logging for troubleshooting @@ -75,10 +80,10 @@ func TestAll(t *testing.T) { "--pod-init-config", cfgFile.Name(), "do", workload) out, err := cmd.CombinedOutput() + t.Log(string(out)) if err != nil { t.Fatalf("runsc do: %v", err) } - t.Log(string(out)) // Wait until the sandbox disconnects to ensure all points were gathered. server.WaitForNoClients() @@ -91,12 +96,18 @@ func matchPoints(t *testing.T, msgs []test.Message) { checker func(test.Message) error count int }{ - pb.MessageType_MESSAGE_CONTAINER_START: {checker: checkContainerStart}, - pb.MessageType_MESSAGE_SENTRY_TASK_EXIT: {checker: checkSentryTaskExit}, - pb.MessageType_MESSAGE_SYSCALL_RAW: {checker: checkSyscallRaw}, - pb.MessageType_MESSAGE_SYSCALL_OPEN: {checker: checkSyscallOpen}, - pb.MessageType_MESSAGE_SYSCALL_CLOSE: {checker: checkSyscallClose}, - pb.MessageType_MESSAGE_SYSCALL_READ: {checker: checkSyscallRead}, + pb.MessageType_MESSAGE_CONTAINER_START: {checker: checkContainerStart}, + pb.MessageType_MESSAGE_SENTRY_CLONE: {checker: checkSentryClone}, + pb.MessageType_MESSAGE_SENTRY_EXEC: {checker: checkSentryExec}, + pb.MessageType_MESSAGE_SENTRY_EXIT_NOTIFY_PARENT: {checker: checkSentryExitNotifyParent}, + pb.MessageType_MESSAGE_SENTRY_TASK_EXIT: {checker: checkSentryTaskExit}, + pb.MessageType_MESSAGE_SYSCALL_CLOSE: {checker: checkSyscallClose}, + pb.MessageType_MESSAGE_SYSCALL_CONNECT: {checker: checkSyscallConnect}, + pb.MessageType_MESSAGE_SYSCALL_EXECVE: {checker: checkSyscallExecve}, + pb.MessageType_MESSAGE_SYSCALL_OPEN: {checker: checkSyscallOpen}, + pb.MessageType_MESSAGE_SYSCALL_RAW: {checker: checkSyscallRaw}, + pb.MessageType_MESSAGE_SYSCALL_READ: {checker: checkSyscallRead}, + pb.MessageType_MESSAGE_SYSCALL_SOCKET: {checker: checkSyscallSocket}, } for _, msg := range msgs { t.Logf("Processing message type %v", msg.MsgType) @@ -120,7 +131,22 @@ func matchPoints(t *testing.T, msgs []test.Message) { } } +func checkTimeNs(ns int64) error { + if ns <= int64(cutoffTime.Nanosecond()) { + return fmt.Errorf("time should not be less than %d (%v), got: %d (%v)", cutoffTime.Nanosecond(), cutoffTime, ns, time.Unix(0, ns)) + } + return nil +} + +type contextDataOpts struct { + skipCwd bool +} + func checkContextData(data *pb.ContextData) error { + return checkContextDataOpts(data, contextDataOpts{}) +} + +func checkContextDataOpts(data *pb.ContextData, opts contextDataOpts) error { if data == nil { return fmt.Errorf("ContextData should not be nil") } @@ -128,18 +154,17 @@ func checkContextData(data *pb.ContextData) error { return fmt.Errorf("invalid container ID %q", data.ContainerId) } - cutoff := time.Now().Add(-time.Minute) - if data.TimeNs <= int64(cutoff.Nanosecond()) { - return fmt.Errorf("time should not be less than %d (%v), got: %d (%v)", cutoff.Nanosecond(), cutoff, data.TimeNs, time.Unix(0, data.TimeNs)) + if err := checkTimeNs(data.TimeNs); err != nil { + return err } - if data.ThreadStartTimeNs <= int64(cutoff.Nanosecond()) { - return fmt.Errorf("thread_start_time should not be less than %d (%v), got: %d (%v)", cutoff.Nanosecond(), cutoff, data.ThreadStartTimeNs, time.Unix(0, data.ThreadStartTimeNs)) + if err := checkTimeNs(data.ThreadStartTimeNs); err != nil { + return err } if data.ThreadStartTimeNs > data.TimeNs { return fmt.Errorf("thread_start_time should not be greater than point time: %d (%v), got: %d (%v)", data.TimeNs, time.Unix(0, data.TimeNs), data.ThreadStartTimeNs, time.Unix(0, data.ThreadStartTimeNs)) } - if data.ThreadGroupStartTimeNs <= int64(cutoff.Nanosecond()) { - return fmt.Errorf("thread_group_start_time should not be less than %d (%v), got: %d (%v)", cutoff.Nanosecond(), cutoff, data.ThreadGroupStartTimeNs, time.Unix(0, data.ThreadGroupStartTimeNs)) + if err := checkTimeNs(data.ThreadGroupStartTimeNs); err != nil { + return err } if data.ThreadGroupStartTimeNs > data.TimeNs { return fmt.Errorf("thread_group_start_time should not be greater than point time: %d (%v), got: %d (%v)", data.TimeNs, time.Unix(0, data.TimeNs), data.ThreadGroupStartTimeNs, time.Unix(0, data.ThreadGroupStartTimeNs)) @@ -151,7 +176,7 @@ func checkContextData(data *pb.ContextData) error { if data.ThreadGroupId <= 0 { return fmt.Errorf("invalid thread_group_id: %v", data.ThreadGroupId) } - if len(data.Cwd) == 0 { + if !opts.skipCwd && len(data.Cwd) == 0 { return fmt.Errorf("invalid cwd: %v", data.Cwd) } if len(data.ProcessName) == 0 { @@ -262,3 +287,149 @@ func checkSyscallRead(msg test.Message) error { } return nil } + +func checkSentryClone(msg test.Message) error { + p := pb.CloneInfo{} + if err := proto.Unmarshal(msg.Msg, &p); err != nil { + return err + } + if err := checkContextData(p.ContextData); err != nil { + return err + } + if p.CreatedThreadId < 0 { + return fmt.Errorf("invalid TID: %d", p.CreatedThreadId) + } + if p.CreatedThreadGroupId < 0 { + return fmt.Errorf("invalid TGID: %d", p.CreatedThreadGroupId) + } + if p.CreatedThreadStartTimeNs < 0 { + return fmt.Errorf("invalid TID: %d", p.CreatedThreadId) + } + return checkTimeNs(p.CreatedThreadStartTimeNs) +} + +func checkSentryExec(msg test.Message) error { + p := pb.ExecveInfo{} + if err := proto.Unmarshal(msg.Msg, &p); err != nil { + return err + } + if err := checkContextData(p.ContextData); err != nil { + return err + } + if want := "/bin/true"; want != p.BinaryPath { + return fmt.Errorf("wrong BinaryPath, want: %q, got: %q", want, p.BinaryPath) + } + if len(p.Argv) == 0 { + return fmt.Errorf("empty Argv") + } + if p.Argv[0] != p.BinaryPath { + return fmt.Errorf("wrong Argv[0], want: %q, got: %q", p.BinaryPath, p.Argv[0]) + } + if len(p.Env) == 0 { + return fmt.Errorf("empty Env") + } + if want := "TEST=123"; want != p.Env[0] { + return fmt.Errorf("wrong Env[0], want: %q, got: %q", want, p.Env[0]) + } + if (p.BinaryMode & 0111) == 0 { + return fmt.Errorf("executing non-executable file, mode: %#o (%#x)", p.BinaryMode, p.BinaryMode) + } + const nobody = 65534 + if p.BinaryUid != nobody { + return fmt.Errorf("BinaryUid, want: %d, got: %d", nobody, p.BinaryUid) + } + if p.BinaryGid != nobody { + return fmt.Errorf("BinaryGid, want: %d, got: %d", nobody, p.BinaryGid) + } + return nil +} + +func checkSyscallExecve(msg test.Message) error { + p := pb.Execve{} + if err := proto.Unmarshal(msg.Msg, &p); err != nil { + return err + } + if err := checkContextData(p.ContextData); err != nil { + return err + } + if p.Fd < 3 { + return fmt.Errorf("execve invalid FD: %d", p.Fd) + } + if want := "/"; want != p.FdPath { + return fmt.Errorf("wrong FdPath, want: %q, got: %q", want, p.FdPath) + } + if want := "/bin/true"; want != p.Pathname { + return fmt.Errorf("wrong Pathname, want: %q, got: %q", want, p.Pathname) + } + if len(p.Argv) == 0 { + return fmt.Errorf("empty Argv") + } + if p.Argv[0] != p.Pathname { + return fmt.Errorf("wrong Argv[0], want: %q, got: %q", p.Pathname, p.Argv[0]) + } + if len(p.Envv) == 0 { + return fmt.Errorf("empty Envv") + } + if want := "TEST=123"; want != p.Envv[0] { + return fmt.Errorf("wrong Envv[0], want: %q, got: %q", want, p.Envv[0]) + } + return nil +} + +func checkSentryExitNotifyParent(msg test.Message) error { + p := pb.ExitNotifyParentInfo{} + if err := proto.Unmarshal(msg.Msg, &p); err != nil { + return err + } + // cwd is empty because the task has already been destroyed when the point + // fires. + opts := contextDataOpts{skipCwd: true} + if err := checkContextDataOpts(p.ContextData, opts); err != nil { + return err + } + if p.ExitStatus != 0 { + return fmt.Errorf("wrong ExitStatus, want: 0, got: %d", p.ExitStatus) + } + return nil +} + +func checkSyscallConnect(msg test.Message) error { + p := pb.Connect{} + if err := proto.Unmarshal(msg.Msg, &p); err != nil { + return err + } + if err := checkContextData(p.ContextData); err != nil { + return err + } + if p.Fd < 3 { + return fmt.Errorf("invalid FD: %d", p.Fd) + } + if want := "socket:"; !strings.HasPrefix(p.FdPath, want) { + return fmt.Errorf("FdPath should start with %q, got: %q", want, p.FdPath) + } + if len(p.Address) == 0 { + return fmt.Errorf("empty address: %q", string(p.Address)) + } + + return nil +} + +func checkSyscallSocket(msg test.Message) error { + p := pb.Socket{} + if err := proto.Unmarshal(msg.Msg, &p); err != nil { + return err + } + if err := checkContextData(p.ContextData); err != nil { + return err + } + if want := unix.AF_UNIX; int32(want) != p.Domain { + return fmt.Errorf("wrong Domain, want: %v, got: %v", want, p.Domain) + } + if want := unix.SOCK_STREAM; int32(want) != p.Type { + return fmt.Errorf("wrong Type, want: %v, got: %v", want, p.Type) + } + if want := int32(0); want != p.Protocol { + return fmt.Errorf("wrong Protocol, want: %v, got: %v", want, p.Protocol) + } + return nil +} diff --git a/test/trace/workload/BUILD b/test/trace/workload/BUILD index fcca7b93b..7391ac2c8 100644 --- a/test/trace/workload/BUILD +++ b/test/trace/workload/BUILD @@ -10,5 +10,12 @@ cc_binary( ], visibility = ["//test/trace:__pkg__"], deps = [ + "//test/util:file_descriptor", + "//test/util:multiprocess_util", + "//test/util:posix_error", + "//test/util:test_util", + "@com_google_absl//absl/cleanup", + "@com_google_absl//absl/strings", + "@com_google_absl//absl/time", ], ) diff --git a/test/trace/workload/workload.cc b/test/trace/workload/workload.cc index 72eee9bc6..457cafb5b 100644 --- a/test/trace/workload/workload.cc +++ b/test/trace/workload/workload.cc @@ -12,5 +12,115 @@ // See the License for the specific language governing permissions and // limitations under the License. -// Empty for now. Actual workload will be added as more points are covered. -int main(int argc, char** argv) { return 0; } +#include +#include +#include +#include + +#include "absl/cleanup/cleanup.h" +#include "absl/strings/str_cat.h" +#include "absl/time/clock.h" +#include "test/util/file_descriptor.h" +#include "test/util/multiprocess_util.h" +#include "test/util/posix_error.h" +#include "test/util/test_util.h" + +namespace gvisor { +namespace testing { + +void runForkExecve() { + auto root_or_error = Open("/", O_RDONLY, 0); + auto& root = root_or_error.ValueOrDie(); + + pid_t child; + int execve_errno; + ExecveArray argv = {"/bin/true"}; + ExecveArray envv = {"TEST=123"}; + auto kill_or_error = ForkAndExecveat(root.get(), "/bin/true", argv, envv, 0, + nullptr, &child, &execve_errno); + ASSERT_EQ(0, execve_errno); + + // Don't kill child, just wait for gracefully exit. + kill_or_error.ValueOrDie().Release(); + RetryEINTR(waitpid)(child, nullptr, 0); +} + +// Creates a simple UDS in the abstract namespace and send one byte from the +// client to the server. +void runSocket() { + auto path = absl::StrCat(std::string("\0", 1), "trace_test.", getpid(), + absl::GetCurrentTimeNanos()); + + struct sockaddr_un addr; + addr.sun_family = AF_UNIX; + strncpy(addr.sun_path, path.c_str(), path.size() + 1); + + int parent_sock = socket(AF_UNIX, SOCK_STREAM, 0); + if (parent_sock < 0) { + err(1, "socket"); + } + auto sock_closer = absl::MakeCleanup([parent_sock] { close(parent_sock); }); + + if (bind(parent_sock, reinterpret_cast(&addr), + sizeof(addr))) { + err(1, "bind"); + } + if (listen(parent_sock, 5) < 0) { + err(1, "listen"); + } + + pid_t pid = fork(); + if (pid < 0) { + // Fork error. + err(1, "fork"); + + } else if (pid == 0) { + // Child. + close(parent_sock); // ensure it's not mistakely used in child. + + int server = socket(AF_UNIX, SOCK_STREAM, 0); + if (server < 0) { + err(1, "socket"); + } + auto server_closer = absl::MakeCleanup([server] { close(server); }); + + if (connect(server, reinterpret_cast(&addr), + sizeof(addr)) < 0) { + err(1, "connect"); + } + + char buf = 'A'; + int bytes = write(server, &buf, sizeof(buf)); + if (bytes != 1) { + err(1, "write: %d", bytes); + } + exit(0); + + } else { + // Parent. + int client = RetryEINTR(accept)(parent_sock, nullptr, nullptr); + if (client < 0) { + err(1, "accept"); + } + auto client_closer = absl::MakeCleanup([client] { close(client); }); + + char buf; + int bytes = read(client, &buf, sizeof(buf)); + if (bytes != 1) { + err(1, "read: %d", bytes); + } + + // Wait to reap the child. + RetryEINTR(waitpid)(pid, nullptr, 0); + } +} + +} // namespace testing +} // namespace gvisor + +int main(int argc, char** argv) { + ::gvisor::testing::runForkExecve(); + ::gvisor::testing::runSocket(); + + return 0; +}