From d4a6e33aa662a4aa64461f53ff04832071861c90 Mon Sep 17 00:00:00 2001 From: PJ Date: Fri, 17 Apr 2026 23:54:05 +0700 Subject: [PATCH] feat(runner): pause-snapshot-evaluate-resume loop Wires agent.Conn + driver.Driver + verifier.Verifier + trace.Writer into the v0.1 step cycle: snapshot the SDK, push to verifier, evaluate properties, write the trace step (with violations), release the SDK pause, apply the next action via the driver, wait for idle. Driver.Launch happens once before the loop and Terminate runs in defer so even an early error tears down the app cleanly. Summary returns step count and per-step violation records for the caller to print or persist. --- internal/runner/runner.go | 203 +++++++++++++++++++++++ internal/runner/runner_test.go | 295 +++++++++++++++++++++++++++++++++ 2 files changed, 498 insertions(+) create mode 100644 internal/runner/runner.go create mode 100644 internal/runner/runner_test.go 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 +}