diff --git a/internal/runner/runner.go b/internal/runner/runner.go new file mode 100644 index 0000000..af04f0c --- /dev/null +++ b/internal/runner/runner.go @@ -0,0 +1,203 @@ +package runner + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "time" + + "github.com/priyanshujain/uatu/internal/agent" + "github.com/priyanshujain/uatu/internal/driver" + "github.com/priyanshujain/uatu/internal/ltl" + "github.com/priyanshujain/uatu/internal/trace" + "github.com/priyanshujain/uatu/internal/verifier" +) + +type Options struct { + BundleID string + ClearState bool + Duration time.Duration + SnapshotTimeout time.Duration + IdleTimeout time.Duration + + Connection *agent.Conn + Driver driver.Driver + Verifier *verifier.Verifier + TraceWriter *trace.Writer +} + +type Summary struct { + StartTime time.Time + EndTime time.Time + Steps int + Violations []ViolationRecord +} + +type ViolationRecord struct { + StepIndex int + Properties []string +} + +func Run(ctx context.Context, options Options) (Summary, error) { + if err := validate(options); err != nil { + return Summary{}, err + } + + summary := Summary{StartTime: time.Now()} + + if err := options.Driver.Launch(ctx, options.BundleID, options.ClearState); err != nil { + return summary, fmt.Errorf("launch: %w", err) + } + defer func() { + _ = options.Driver.Terminate(context.Background()) + }() + + deadline := summary.StartTime.Add(options.Duration) + stepIndex := 0 + for time.Now().Before(deadline) { + if err := ctx.Err(); err != nil { + break + } + stepIndex++ + stepStart := time.Now() + + snapshot, err := snapshotStep(ctx, options) + if err != nil { + return summary, fmt.Errorf("step %d snapshot: %w", stepIndex, err) + } + + if err := options.Verifier.PushSnapshot(verifier.Snapshots(snapshot.Snapshots)); err != nil { + return summary, fmt.Errorf("step %d push: %w", stepIndex, err) + } + verdicts := options.Verifier.EvaluateProperties() + violations := violationNames(verdicts) + + nextAction, nextErr := options.Verifier.NextAction() + var traceAction *trace.Action + if nextErr == nil { + traceAction = traceActionFor(nextAction) + } else if !errors.Is(nextErr, verifier.ErrNoAction) { + return summary, fmt.Errorf("step %d next action: %w", stepIndex, nextErr) + } + + step := trace.Step{ + Index: stepIndex, + Timestamp: stepStart, + Screen: screenFromSnapshot(snapshot.Snapshots), + Snapshots: snapshot.Snapshots, + Action: traceAction, + Violations: violations, + } + if err := options.TraceWriter.WriteStep(step); err != nil { + return summary, fmt.Errorf("step %d trace: %w", stepIndex, err) + } + summary.Steps = stepIndex + if len(violations) > 0 { + summary.Violations = append(summary.Violations, ViolationRecord{ + StepIndex: stepIndex, + Properties: violations, + }) + } + + if err := options.Connection.Release(ctx); err != nil { + return summary, fmt.Errorf("step %d release: %w", stepIndex, err) + } + + if nextErr == nil { + if err := applyAction(ctx, options.Driver, nextAction); err != nil { + return summary, fmt.Errorf("step %d apply: %w", stepIndex, err) + } + } + + idleCtx, idleCancel := context.WithTimeout(ctx, options.IdleTimeout) + _ = options.Driver.WaitForIdle(idleCtx, options.IdleTimeout) + idleCancel() + } + + summary.EndTime = time.Now() + return summary, nil +} + +func validate(options Options) error { + if options.Connection == nil { + return errors.New("runner: Connection is required") + } + if options.Driver == nil { + return errors.New("runner: Driver is required") + } + if options.Verifier == nil { + return errors.New("runner: Verifier is required") + } + if options.TraceWriter == nil { + return errors.New("runner: TraceWriter is required") + } + if options.Duration <= 0 { + return errors.New("runner: Duration must be positive") + } + if options.SnapshotTimeout <= 0 { + options.SnapshotTimeout = 5 * time.Second + } + if options.IdleTimeout <= 0 { + options.IdleTimeout = 2 * time.Second + } + return nil +} + +func snapshotStep(ctx context.Context, options Options) (agent.Message, error) { + snapshotTimeout := options.SnapshotTimeout + if snapshotTimeout <= 0 { + snapshotTimeout = 5 * time.Second + } + snapshotCtx, snapshotCancel := context.WithTimeout(ctx, snapshotTimeout) + defer snapshotCancel() + return options.Connection.Snapshot(snapshotCtx) +} + +func violationNames(verdicts map[string]ltl.Verdict) []string { + var names []string + for name, verdict := range verdicts { + if verdict == ltl.VerdictViolated { + names = append(names, name) + } + } + return names +} + +func screenFromSnapshot(snapshots map[string]json.RawMessage) string { + raw, ok := snapshots["screen"] + if !ok { + return "" + } + var screen string + _ = json.Unmarshal(raw, &screen) + return screen +} + +func applyAction(ctx context.Context, drv driver.Driver, action verifier.Action) error { + switch action.Kind { + case verifier.ActionKindTap: + if action.On == "" { + return nil + } + return drv.TapSelector(ctx, action.On) + case verifier.ActionKindInputText: + return drv.InputText(ctx, action.Text) + default: + return fmt.Errorf("unknown action kind %q", action.Kind) + } +} + +func traceActionFor(action verifier.Action) *trace.Action { + traceAction := &trace.Action{Kind: string(action.Kind)} + switch action.Kind { + case verifier.ActionKindTap: + // Selector lives in the trace step's action.text field for now — + // trace.Action only has X/Y/Text and we don't resolve coordinates + // at the runner layer. + traceAction.Text = action.On + case verifier.ActionKindInputText: + traceAction.Text = action.Text + } + return traceAction +} diff --git a/internal/runner/runner_test.go b/internal/runner/runner_test.go new file mode 100644 index 0000000..78e040f --- /dev/null +++ b/internal/runner/runner_test.go @@ -0,0 +1,295 @@ +package runner + +import ( + "context" + "encoding/json" + "net" + "os" + "path/filepath" + "slices" + "strings" + "sync" + "testing" + "time" + + "github.com/priyanshujain/uatu/internal/agent" + mockdriver "github.com/priyanshujain/uatu/internal/driver/mock" + "github.com/priyanshujain/uatu/internal/trace" + "github.com/priyanshujain/uatu/internal/verifier" +) + +const fixtureSpec = ` +const screen = __uatu__.extract(state => state.snapshots.screen ?? ""); +const balance = __uatu__.extract(state => state.snapshots.balance ?? 0); +globalThis.properties = { + balanceNonNegative: __uatu__.always(() => balance.current >= 0), +}; +globalThis.actions = __uatu__.actions(() => [__uatu__.tap({ on: "id:next" })]); +` + +type harness struct { + server *agent.Server + listener net.Listener + clientWG sync.WaitGroup + conn *agent.Conn + mock *mockdriver.Driver + verifier *verifier.Verifier + writer *trace.Writer + snapshot []map[string]json.RawMessage +} + +func newHarness(t *testing.T, snapshots []map[string]json.RawMessage) *harness { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := agent.NewServer(listener) + directory := t.TempDir() + writer, err := trace.NewWriter(directory) + if err != nil { + t.Fatal(err) + } + verifierInstance, err := verifier.New() + if err != nil { + t.Fatal(err) + } + if err := verifierInstance.Load(fixtureSpec); err != nil { + t.Fatal(err) + } + state := &harness{ + server: server, + listener: listener, + mock: mockdriver.New(), + verifier: verifierInstance, + writer: writer, + snapshot: snapshots, + } + t.Cleanup(func() { + _ = listener.Close() + _ = writer.Close() + }) + return state +} + +func (h *harness) startSDK(t *testing.T) { + t.Helper() + h.clientWG.Go(func() { + conn, err := net.Dial("tcp", h.listener.Addr().String()) + if err != nil { + t.Errorf("dial: %v", err) + return + } + defer conn.Close() + if err := agent.WriteMessage(conn, agent.Hello("0.0.1", "android", "com.fixture")); err != nil { + t.Errorf("hello: %v", err) + return + } + index := 0 + for { + message, err := agent.ReadMessage(conn) + if err != nil { + return + } + if message.Type == agent.MessageTypePause { + snapshots := map[string]json.RawMessage{} + if index < len(h.snapshot) { + snapshots = h.snapshot[index] + } + if err := agent.WriteMessage(conn, agent.State(message.ID, snapshots)); err != nil { + return + } + index++ + } + } + }) +} + +func (h *harness) acceptConnection(t *testing.T) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + connection, err := h.server.Accept(ctx) + if err != nil { + t.Fatalf("Accept: %v", err) + } + h.conn = connection +} + +func TestRunner_HappyPathStepsAndTraces(t *testing.T) { + snapshots := []map[string]json.RawMessage{ + {"screen": json.RawMessage(`"home"`), "balance": json.RawMessage(`100`)}, + {"screen": json.RawMessage(`"home"`), "balance": json.RawMessage(`200`)}, + {"screen": json.RawMessage(`"home"`), "balance": json.RawMessage(`300`)}, + } + state := newHarness(t, snapshots) + state.startSDK(t) + state.acceptConnection(t) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + summary, err := Run(ctx, Options{ + BundleID: "com.fixture", + Duration: 100 * time.Millisecond, + SnapshotTimeout: 2 * time.Second, + IdleTimeout: 50 * time.Millisecond, + Connection: state.conn, + Driver: state.mock, + Verifier: state.verifier, + TraceWriter: state.writer, + }) + if err != nil { + t.Fatalf("Run: %v", err) + } + if summary.Steps == 0 { + t.Errorf("expected at least one step, got 0") + } + if len(summary.Violations) != 0 { + t.Errorf("no violations expected, got %v", summary.Violations) + } + + actions := state.mock.Actions() + if !containsAction(actions, mockdriver.ActionLaunch, "com.fixture") { + t.Errorf("expected Launch with com.fixture, got %v", actions) + } + if !containsAction(actions, mockdriver.ActionTapSelector, "id:next") { + t.Errorf("expected TapSelector with id:next, got %v", actions) + } + if !containsAction(actions, mockdriver.ActionTerminate, "") { + t.Errorf("expected Terminate, got %v", actions) + } +} + +func TestRunner_ViolationSurfacesInSummary(t *testing.T) { + snapshots := []map[string]json.RawMessage{ + {"balance": json.RawMessage(`100`)}, + {"balance": json.RawMessage(`-1`)}, + {"balance": json.RawMessage(`50`)}, + } + state := newHarness(t, snapshots) + state.startSDK(t) + state.acceptConnection(t) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + summary, err := Run(ctx, Options{ + BundleID: "com.fixture", + Duration: 100 * time.Millisecond, + SnapshotTimeout: 2 * time.Second, + IdleTimeout: 50 * time.Millisecond, + Connection: state.conn, + Driver: state.mock, + Verifier: state.verifier, + TraceWriter: state.writer, + }) + if err != nil { + t.Fatalf("Run: %v", err) + } + if len(summary.Violations) == 0 { + t.Errorf("expected at least one violation, got %v", summary.Violations) + } + if !containsProperty(summary.Violations, "balanceNonNegative") { + t.Errorf("expected balanceNonNegative in violations: %v", summary.Violations) + } +} + +func TestRunner_RejectsMissingFields(t *testing.T) { + _, err := Run(context.Background(), Options{Duration: time.Second}) + if err == nil || !strings.Contains(err.Error(), "Connection") { + t.Errorf("expected Connection-required error, got %v", err) + } +} + +func TestRunner_RejectsZeroDuration(t *testing.T) { + _, err := Run(context.Background(), Options{ + Connection: &agent.Conn{}, + Driver: mockdriver.New(), + Verifier: mustNewVerifier(t), + TraceWriter: mustNewTraceWriter(t), + }) + if err == nil || !strings.Contains(err.Error(), "Duration") { + t.Errorf("expected Duration-required error, got %v", err) + } +} + +func TestRunner_RecordsScreenFieldFromSnapshot(t *testing.T) { + snapshots := []map[string]json.RawMessage{ + {"screen": json.RawMessage(`"customer_ledger"`), "balance": json.RawMessage(`1`)}, + } + state := newHarness(t, snapshots) + state.startSDK(t) + state.acceptConnection(t) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if _, err := Run(ctx, Options{ + BundleID: "com.fixture", + Duration: 100 * time.Millisecond, + SnapshotTimeout: 2 * time.Second, + IdleTimeout: 50 * time.Millisecond, + Connection: state.conn, + Driver: state.mock, + Verifier: state.verifier, + TraceWriter: state.writer, + }); err != nil { + t.Fatal(err) + } + body, err := os.ReadFile(filepath.Join(state.writer.Directory(), "trace.jsonl")) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(body), `"screen":"customer_ledger"`) { + t.Errorf("screen field not in trace: %s", body) + } +} + +func mustNewVerifier(t *testing.T) *verifier.Verifier { + t.Helper() + verifierInstance, err := verifier.New() + if err != nil { + t.Fatal(err) + } + return verifierInstance +} + +func mustNewTraceWriter(t *testing.T) *trace.Writer { + t.Helper() + writer, err := trace.NewWriter(t.TempDir()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = writer.Close() }) + return writer +} + +func containsAction(actions []mockdriver.Action, kind mockdriver.ActionKind, payload string) bool { + for _, action := range actions { + if action.Kind != kind { + continue + } + switch kind { + case mockdriver.ActionLaunch: + if action.BundleID == payload { + return true + } + case mockdriver.ActionTapSelector: + if action.Selector == payload { + return true + } + case mockdriver.ActionTerminate: + return true + default: + return true + } + } + return false +} + +func containsProperty(records []ViolationRecord, property string) bool { + for _, record := range records { + if slices.Contains(record.Properties, property) { + return true + } + } + return false +}