package localcontrol import ( "context" "encoding/binary" "errors" "io" "net" "os" "path/filepath" "testing" "time" proto_socket "git.toki-labs.com/toki/proto-socket/go" "git.toki-labs.com/toki/proto-socket/go/packets" iop "iop/proto/gen/iop" "google.golang.org/protobuf/proto" ) type fixedPeerCredentials struct { supported bool uid uint32 err error } func (f fixedPeerCredentials) Supported() bool { return f.supported } func (f fixedPeerCredentials) UID(net.Conn) (uint32, error) { return f.uid, f.err } func TestServerSameUserProtoSocket(t *testing.T) { t.Parallel() host := newRecordingHost() stateRoot := filepath.Join(t.TempDir(), "agent-state") service, _ := newTestService( t, filepath.Join(stateRoot, "manager.json"), "daemon-server", 8, host, ) server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service) if err != nil { t.Fatalf("NewServer: %v", err) } ctx, cancel := context.WithCancel(context.Background()) defer cancel() if err := server.Start(ctx); err != nil { t.Fatalf("Start: %v", err) } defer func() { if err := server.Stop(); err != nil { t.Errorf("Stop: %v", err) } }() rootInfo, err := os.Lstat(stateRoot) if err != nil { t.Fatalf("Lstat state root: %v", err) } if rootInfo.Mode().Perm() != 0o700 { t.Fatalf("state root mode = %o, want 700", rootInfo.Mode().Perm()) } socketInfo, err := os.Lstat(server.Path()) if err != nil { t.Fatalf("Lstat socket: %v", err) } if socketInfo.Mode()&os.ModeSocket == 0 || socketInfo.Mode().Perm() != 0o600 { t.Fatalf("socket mode = %v, want socket 0600", socketInfo.Mode()) } conn, err := net.DialTimeout("unix", server.Path(), time.Second) if err != nil { t.Fatalf("Dial: %v", err) } client := proto_socket.NewTcpClient( conn, 0, 0, localControlParserMap(), ) defer client.Close() response, err := proto_socket.SendRequestTyped[ *iop.AgentLocalEnvelope, *iop.AgentLocalEnvelope, ]( &client.Communicator, runtimeStatusRequest("message-real-socket"), 2*time.Second, ) if err != nil { t.Fatalf("SendRequestTyped: %v", err) } if response.GetResponse().GetSnapshot().GetDaemonId() != "daemon-server" { t.Fatalf("response = %#v", response) } if calls := host.callCount(OperationRuntimeStatus); calls != 1 { t.Fatalf("runtime.status calls = %d, want 1", calls) } } func TestPeerUIDMismatchDeniedBeforeDispatch(t *testing.T) { t.Parallel() host := newRecordingHost() stateRoot := filepath.Join(t.TempDir(), "agent-state") service, _ := newTestService( t, filepath.Join(stateRoot, "manager.json"), "daemon-denied", 8, host, ) server, err := newServer( ServerConfig{StateRoot: stateRoot}, service, fixedPeerCredentials{ supported: true, uid: effectiveUID() + 1, }, effectiveUID, ) if err != nil { t.Fatalf("newServer: %v", err) } if err := server.Start(context.Background()); err != nil { t.Fatalf("Start: %v", err) } defer server.Stop() conn, err := net.DialTimeout("unix", server.Path(), time.Second) if err != nil { t.Fatalf("Dial: %v", err) } defer conn.Close() response := readEnvelopePacket(t, conn) if response.GetError().GetCode() != ErrorPermissionDenied { t.Fatalf("denied response = %#v", response) } if calls := host.totalCalls(); calls != 0 { t.Fatalf("denied peer host calls = %d, want 0", calls) } } func TestServerBroadcastsCommittedEventToConcurrentClients(t *testing.T) { t.Parallel() host := newRecordingHost() stateRoot := filepath.Join(t.TempDir(), "agent-state") service, _ := newTestService( t, filepath.Join(stateRoot, "manager.json"), "daemon-broadcast", 8, host, ) server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service) if err != nil { t.Fatalf("NewServer: %v", err) } if err := server.Start(context.Background()); err != nil { t.Fatalf("Start: %v", err) } defer server.Stop() clients := make([]*proto_socket.TcpClient, 0, 2) eventChannels := make([]chan *iop.AgentLocalEnvelope, 0, 2) for index := 0; index < 2; index++ { conn, err := net.DialTimeout("unix", server.Path(), time.Second) if err != nil { t.Fatalf("Dial client %d: %v", index, err) } client := proto_socket.NewTcpClient(conn, 0, 0, localControlParserMap()) events := make(chan *iop.AgentLocalEnvelope, 1) proto_socket.AddListenerTyped( &client.Communicator, func(envelope *iop.AgentLocalEnvelope) { if envelope.GetKind() == iop.AgentLocalKind_AGENT_LOCAL_KIND_EVENT { events <- envelope } }, ) if _, err := proto_socket.SendRequestTyped[ *iop.AgentLocalEnvelope, *iop.AgentLocalEnvelope, ]( &client.Communicator, runtimeStatusRequest("message-client-ready-"+string(rune('0'+index))), 2*time.Second, ); err != nil { t.Fatalf("client %d readiness request: %v", index, err) } clients = append(clients, client) eventChannels = append(eventChannels, events) } defer func() { for _, client := range clients { _ = client.Close() } }() response, err := proto_socket.SendRequestTyped[ *iop.AgentLocalEnvelope, *iop.AgentLocalEnvelope, ]( &clients[0].Communicator, projectRequest( "message-broadcast", "command-broadcast", OperationProjectStart, "project-a", ), 2*time.Second, ) if err != nil { t.Fatalf("SendRequestTyped: %v", err) } if response.GetResponse().GetMutation() == nil { t.Fatalf("response = %#v", response) } for index, events := range eventChannels { select { case event := <-events: if event.GetEvent().GetEventSequence() != 1 { t.Fatalf("client %d event = %#v", index, event) } case <-time.After(2 * time.Second): t.Fatalf("client %d did not receive live event", index) } } } func TestServerRejectsUnsafePaths(t *testing.T) { t.Parallel() host := newRecordingHost() tests := []struct { name string prepare func(*testing.T, string) want error }{ { name: "broad state root", prepare: func(t *testing.T, root string) { t.Helper() if err := os.MkdirAll(root, 0o755); err != nil { t.Fatalf("MkdirAll: %v", err) } if err := os.Chmod(root, 0o755); err != nil { t.Fatalf("Chmod: %v", err) } }, want: ErrUnsafeStateRoot, }, { name: "state root symlink", prepare: func(t *testing.T, root string) { t.Helper() target := root + "-target" if err := os.MkdirAll(target, 0o700); err != nil { t.Fatalf("MkdirAll target: %v", err) } if err := os.Symlink(target, root); err != nil { t.Fatalf("Symlink: %v", err) } }, want: ErrUnsafeStateRoot, }, { name: "existing socket path replacement", prepare: func(t *testing.T, root string) { t.Helper() if err := os.MkdirAll(root, 0o700); err != nil { t.Fatalf("MkdirAll: %v", err) } if err := os.WriteFile( filepath.Join(root, defaultSocketName), []byte("replacement"), 0o600, ); err != nil { t.Fatalf("WriteFile: %v", err) } }, want: ErrUnsafeSocketPath, }, } for _, test := range tests { test := test t.Run(test.name, func(t *testing.T) { t.Parallel() root := filepath.Join(t.TempDir(), "state") test.prepare(t, root) service, _ := newTestService( t, filepath.Join(t.TempDir(), "manager.json"), "daemon-path-"+test.name, 8, host, ) server, err := NewServer(ServerConfig{StateRoot: root}, service) if err != nil { t.Fatalf("NewServer: %v", err) } if err := server.Start(context.Background()); !errors.Is(err, test.want) { t.Fatalf("Start error = %v, want %v", err, test.want) } }) } } func TestServerStopPreservesReplacedSocketPath(t *testing.T) { t.Parallel() host := newRecordingHost() stateRoot := filepath.Join(t.TempDir(), "state") service, _ := newTestService( t, filepath.Join(t.TempDir(), "manager.json"), "daemon-replacement", 8, host, ) server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service) if err != nil { t.Fatalf("NewServer: %v", err) } if err := server.Start(context.Background()); err != nil { t.Fatalf("Start: %v", err) } if err := os.Remove(server.Path()); err != nil { t.Fatalf("Remove socket: %v", err) } replacement := []byte("do-not-remove") if err := os.WriteFile(server.Path(), replacement, 0o600); err != nil { t.Fatalf("WriteFile replacement: %v", err) } if err := server.Stop(); !errors.Is(err, ErrUnsafeSocketPath) { t.Fatalf("Stop error = %v, want ErrUnsafeSocketPath", err) } got, err := os.ReadFile(server.Path()) if err != nil { t.Fatalf("ReadFile replacement: %v", err) } if string(got) != string(replacement) { t.Fatalf("replacement content = %q", got) } } func TestServerFailsBeforeListeningWithoutPeerCredentials(t *testing.T) { t.Parallel() host := newRecordingHost() stateRoot := filepath.Join(t.TempDir(), "state") service, _ := newTestService( t, filepath.Join(t.TempDir(), "manager.json"), "daemon-unsupported", 8, host, ) server, err := newServer( ServerConfig{StateRoot: stateRoot}, service, fixedPeerCredentials{supported: false}, effectiveUID, ) if err != nil { t.Fatalf("newServer: %v", err) } if err := server.Start(context.Background()); !errors.Is( err, ErrPeerCredentialsUnsupported, ) { t.Fatalf("Start error = %v", err) } if _, err := os.Lstat(stateRoot); !errors.Is(err, os.ErrNotExist) { t.Fatalf("state root exists after unsupported preflight: %v", err) } } func TestServerOversizeFrameHasZeroMutation(t *testing.T) { t.Parallel() host := newRecordingHost() stateRoot := filepath.Join(t.TempDir(), "state") service, _ := newTestService( t, filepath.Join(t.TempDir(), "manager.json"), "daemon-frame-bound", 8, host, ) server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service) if err != nil { t.Fatalf("NewServer: %v", err) } if err := server.Start(context.Background()); err != nil { t.Fatalf("Start: %v", err) } defer server.Stop() conn, err := net.DialTimeout("unix", server.Path(), time.Second) if err != nil { t.Fatalf("Dial: %v", err) } defer conn.Close() header := make([]byte, 4) binary.BigEndian.PutUint32(header, MaxFrameBytes+1) if _, err := conn.Write(header); err != nil { t.Fatalf("Write header: %v", err) } if err := conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { t.Fatalf("SetReadDeadline: %v", err) } var one [1]byte if _, err := conn.Read(one[:]); err == nil { t.Fatal("oversize frame connection remained open") } if calls := host.totalCalls(); calls != 0 { t.Fatalf("oversize frame host calls = %d, want 0", calls) } } func localControlParserMap() proto_socket.ParserMap { return proto_socket.ParserMap{ proto_socket.TypeNameOf(&iop.AgentLocalEnvelope{}): func( payload []byte, ) (proto.Message, error) { return ParseEnvelope(payload) }, } } func readEnvelopePacket(t *testing.T, conn net.Conn) *iop.AgentLocalEnvelope { t.Helper() if err := conn.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil { t.Fatalf("SetReadDeadline: %v", err) } header := make([]byte, 4) if _, err := io.ReadFull(conn, header); err != nil { t.Fatalf("ReadFull header: %v", err) } length := binary.BigEndian.Uint32(header) if length == 0 || length > MaxFrameBytes { t.Fatalf("packet length = %d", length) } payload := make([]byte, int(length)) if _, err := io.ReadFull(conn, payload); err != nil { t.Fatalf("ReadFull payload: %v", err) } packet := &packets.PacketBase{} if err := proto.Unmarshal(payload, packet); err != nil { t.Fatalf("Unmarshal packet: %v", err) } envelope, err := ParseEnvelope(packet.GetData()) if err != nil { t.Fatalf("ParseEnvelope: %v", err) } return envelope }