diff --git a/internal/driver/maestro/client.go b/internal/driver/maestro/client.go new file mode 100644 index 0000000..9b0dd33 --- /dev/null +++ b/internal/driver/maestro/client.go @@ -0,0 +1,113 @@ +package maestro + +import ( + "context" + "fmt" + "time" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + + "github.com/priyanshujain/uatu/internal/driver" + driverpb "github.com/priyanshujain/uatu/proto/driverpb" +) + +type Client struct { + connection *grpc.ClientConn + stub driverpb.DriverClient +} + +// Dial connects to the sidecar gRPC server at the given address. +// Address must be a host:port pair, typically "127.0.0.1:". +func Dial(address string) (*Client, error) { + connection, err := grpc.NewClient(address, grpc.WithTransportCredentials(insecure.NewCredentials())) + if err != nil { + return nil, fmt.Errorf("dial sidecar: %w", err) + } + return &Client{connection: connection, stub: driverpb.NewDriverClient(connection)}, nil +} + +func (c *Client) Close() error { return c.connection.Close() } + +// WaitForHealth polls the sidecar's Health RPC until it returns Ready=true +// or the context is canceled. +func (c *Client) WaitForHealth(ctx context.Context, pollInterval time.Duration) error { + if pollInterval <= 0 { + pollInterval = 100 * time.Millisecond + } + for { + response, err := c.stub.Health(ctx, &driverpb.Empty{}) + if err == nil && response.GetReady() { + return nil + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(pollInterval): + } + } +} + +func (c *Client) Launch(ctx context.Context, bundleID string, clearState bool) error { + _, err := c.stub.Launch(ctx, &driverpb.LaunchRequest{BundleId: bundleID, ClearState: clearState}) + return err +} + +func (c *Client) Terminate(ctx context.Context) error { + _, err := c.stub.Terminate(ctx, &driverpb.Empty{}) + return err +} + +func (c *Client) Tap(ctx context.Context, x, y int) error { + _, err := c.stub.Tap(ctx, &driverpb.Point{X: int32(x), Y: int32(y)}) + return err +} + +func (c *Client) TapSelector(ctx context.Context, selector string) error { + _, err := c.stub.TapSelector(ctx, &driverpb.Selector{Value: selector}) + return err +} + +func (c *Client) InputText(ctx context.Context, text string) error { + _, err := c.stub.InputText(ctx, &driverpb.Text{Value: text}) + return err +} + +func (c *Client) Hierarchy(ctx context.Context) (string, error) { + response, err := c.stub.Hierarchy(ctx, &driverpb.Empty{}) + if err != nil { + return "", err + } + return response.GetJson(), nil +} + +func (c *Client) Screenshot(ctx context.Context) (driver.Image, error) { + response, err := c.stub.Screenshot(ctx, &driverpb.Empty{}) + if err != nil { + return driver.Image{}, err + } + return driver.Image{ + PNG: response.GetPng(), + Width: int(response.GetWidth()), + Height: int(response.GetHeight()), + }, nil +} + +func (c *Client) WaitForIdle(ctx context.Context, duration time.Duration) error { + _, err := c.stub.WaitForIdle(ctx, &driverpb.Duration{Millis: duration.Milliseconds()}) + return err +} + +func (c *Client) Health(ctx context.Context) (driver.Health, error) { + response, err := c.stub.Health(ctx, &driverpb.Empty{}) + if err != nil { + return driver.Health{}, err + } + return driver.Health{ + Ready: response.GetReady(), + Version: response.GetVersion(), + Platform: response.GetPlatform(), + }, nil +} + +var _ driver.Driver = (*Client)(nil) diff --git a/internal/driver/maestro/client_test.go b/internal/driver/maestro/client_test.go new file mode 100644 index 0000000..ab5a590 --- /dev/null +++ b/internal/driver/maestro/client_test.go @@ -0,0 +1,271 @@ +package maestro + +import ( + "context" + "net" + "strings" + "sync" + "testing" + "time" + + "google.golang.org/grpc" + + driverpb "github.com/priyanshujain/uatu/proto/driverpb" +) + +type fakeServer struct { + driverpb.UnimplementedDriverServer + mutex sync.Mutex + + healthReady bool + healthCalls int + + launchedBundleID string + clearState bool + terminateCalls int + taps []int32 + tapSelectors []string + inputs []string + idleMillis []int64 + hierarchy string + imagePNG []byte + imageWidth int32 + imageHeight int32 + + healthError error +} + +func (s *fakeServer) Health(_ context.Context, _ *driverpb.Empty) (*driverpb.HealthStatus, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + s.healthCalls++ + if s.healthError != nil { + return nil, s.healthError + } + return &driverpb.HealthStatus{Ready: s.healthReady, Version: "test", Platform: "android"}, nil +} + +func (s *fakeServer) Launch(_ context.Context, request *driverpb.LaunchRequest) (*driverpb.Empty, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + s.launchedBundleID = request.GetBundleId() + s.clearState = request.GetClearState() + return &driverpb.Empty{}, nil +} + +func (s *fakeServer) Terminate(_ context.Context, _ *driverpb.Empty) (*driverpb.Empty, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + s.terminateCalls++ + return &driverpb.Empty{}, nil +} + +func (s *fakeServer) Tap(_ context.Context, point *driverpb.Point) (*driverpb.Empty, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + s.taps = append(s.taps, point.GetX(), point.GetY()) + return &driverpb.Empty{}, nil +} + +func (s *fakeServer) TapSelector(_ context.Context, selector *driverpb.Selector) (*driverpb.Empty, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + s.tapSelectors = append(s.tapSelectors, selector.GetValue()) + return &driverpb.Empty{}, nil +} + +func (s *fakeServer) InputText(_ context.Context, text *driverpb.Text) (*driverpb.Empty, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + s.inputs = append(s.inputs, text.GetValue()) + return &driverpb.Empty{}, nil +} + +func (s *fakeServer) WaitForIdle(_ context.Context, duration *driverpb.Duration) (*driverpb.Empty, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + s.idleMillis = append(s.idleMillis, duration.GetMillis()) + return &driverpb.Empty{}, nil +} + +func (s *fakeServer) Hierarchy(_ context.Context, _ *driverpb.Empty) (*driverpb.HierarchyJSON, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + return &driverpb.HierarchyJSON{Json: s.hierarchy}, nil +} + +func (s *fakeServer) Screenshot(_ context.Context, _ *driverpb.Empty) (*driverpb.Image, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + return &driverpb.Image{Png: s.imagePNG, Width: s.imageWidth, Height: s.imageHeight}, nil +} + +type harness struct { + server *grpc.Server + fake *fakeServer + address string +} + +func newHarness(t *testing.T) *harness { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := grpc.NewServer() + fake := &fakeServer{healthReady: true, hierarchy: `{"x":1}`, imagePNG: []byte{0xFF}, imageWidth: 1080, imageHeight: 2340} + driverpb.RegisterDriverServer(server, fake) + go func() { _ = server.Serve(listener) }() + t.Cleanup(func() { + server.Stop() + _ = listener.Close() + }) + return &harness{server: server, fake: fake, address: listener.Addr().String()} +} + +func TestClient_HealthRoundTrip(t *testing.T) { + state := newHarness(t) + client, err := Dial(state.address) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + got, err := client.Health(context.Background()) + if err != nil { + t.Fatal(err) + } + if !got.Ready || got.Version != "test" || got.Platform != "android" { + t.Errorf("unexpected health: %+v", got) + } +} + +func TestClient_WaitForHealth_PollsUntilReady(t *testing.T) { + state := newHarness(t) + state.fake.healthReady = false + client, err := Dial(state.address) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + go func() { + time.Sleep(50 * time.Millisecond) + state.fake.mutex.Lock() + state.fake.healthReady = true + state.fake.mutex.Unlock() + }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + if err := client.WaitForHealth(ctx, 25*time.Millisecond); err != nil { + t.Fatalf("WaitForHealth: %v", err) + } + if state.fake.healthCalls < 2 { + t.Errorf("expected at least 2 health polls, got %d", state.fake.healthCalls) + } +} + +func TestClient_WaitForHealth_HonorsContext(t *testing.T) { + state := newHarness(t) + state.fake.healthReady = false + client, err := Dial(state.address) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + err = client.WaitForHealth(ctx, 25*time.Millisecond) + if err == nil || !strings.Contains(err.Error(), "context") { + t.Fatalf("expected context error, got %v", err) + } +} + +func TestClient_LaunchAndTerminate(t *testing.T) { + state := newHarness(t) + client, _ := Dial(state.address) + defer client.Close() + + if err := client.Launch(context.Background(), "com.example", true); err != nil { + t.Fatal(err) + } + if state.fake.launchedBundleID != "com.example" || !state.fake.clearState { + t.Errorf("launch payload wrong: %+v", state.fake) + } + if err := client.Terminate(context.Background()); err != nil { + t.Fatal(err) + } + if state.fake.terminateCalls != 1 { + t.Errorf("terminate calls: %d", state.fake.terminateCalls) + } +} + +func TestClient_TapAndTapSelector(t *testing.T) { + state := newHarness(t) + client, _ := Dial(state.address) + defer client.Close() + + if err := client.Tap(context.Background(), 100, 250); err != nil { + t.Fatal(err) + } + if len(state.fake.taps) != 2 || state.fake.taps[0] != 100 || state.fake.taps[1] != 250 { + t.Errorf("tap coordinates wrong: %v", state.fake.taps) + } + + if err := client.TapSelector(context.Background(), "id:home"); err != nil { + t.Fatal(err) + } + if len(state.fake.tapSelectors) != 1 || state.fake.tapSelectors[0] != "id:home" { + t.Errorf("selectors wrong: %v", state.fake.tapSelectors) + } +} + +func TestClient_InputText(t *testing.T) { + state := newHarness(t) + client, _ := Dial(state.address) + defer client.Close() + + if err := client.InputText(context.Background(), "hello world"); err != nil { + t.Fatal(err) + } + if len(state.fake.inputs) != 1 || state.fake.inputs[0] != "hello world" { + t.Errorf("inputs wrong: %v", state.fake.inputs) + } +} + +func TestClient_HierarchyAndScreenshot(t *testing.T) { + state := newHarness(t) + client, _ := Dial(state.address) + defer client.Close() + + hierarchy, err := client.Hierarchy(context.Background()) + if err != nil { + t.Fatal(err) + } + if hierarchy != `{"x":1}` { + t.Errorf("hierarchy wrong: %q", hierarchy) + } + + image, err := client.Screenshot(context.Background()) + if err != nil { + t.Fatal(err) + } + if image.Width != 1080 || image.Height != 2340 || len(image.PNG) != 1 { + t.Errorf("image wrong: %+v", image) + } +} + +func TestClient_WaitForIdleForwardsMillis(t *testing.T) { + state := newHarness(t) + client, _ := Dial(state.address) + defer client.Close() + + if err := client.WaitForIdle(context.Background(), 250*time.Millisecond); err != nil { + t.Fatal(err) + } + if len(state.fake.idleMillis) != 1 || state.fake.idleMillis[0] != 250 { + t.Errorf("idleMillis wrong: %v", state.fake.idleMillis) + } +}