feat(verifier): wire hierarchy into state.ax and carry element coords in Action

This commit is contained in:
pj committed 2026-04-18 01:23:55 +07:00
1 parent e539e89438
commit 5f2a2d52ab
4 files changed
+105 -24

No files matched your search

+5
View File
@@ -17,6 +17,10 @@ type Action struct {
Kind ActionKind Kind ActionKind
On string On string
Text string Text string
// X, Y hold the element center when the spec passed an ax element to
// Tap/InputText. Zero means the runner must resolve On against the
// current hierarchy.
X, Y int
} }
type extractorState struct { type extractorState struct {
@@ -32,6 +36,7 @@ const (
tagFormula = "__uatuFormula" tagFormula = "__uatuFormula"
tagActionGenerator = "__uatuActionGenerator" tagActionGenerator = "__uatuActionGenerator"
tagInternalKind = "__uatuKind" tagInternalKind = "__uatuKind"
tagSelector = "__uatuSelector"
internalKindActions = "actions" internalKindActions = "actions"
internalKindWeighted = "weighted" internalKindWeighted = "weighted"
internalKindBuiltinTaps = "taps" internalKindBuiltinTaps = "taps"
+87 -13
View File
@@ -5,14 +5,17 @@ import (
"fmt" "fmt"
"github.com/dop251/goja" "github.com/dop251/goja"
"github.com/priyanshujain/uatu/internal/hierarchy"
) )
// Snapshots is the per-step extractor output forwarded by the SDK. // Snapshots is the per-step extractor output forwarded by the SDK.
type Snapshots map[string]json.RawMessage type Snapshots map[string]json.RawMessage
// stateObject builds a JS-side `{ snapshots, ax }` matching the State type // stateObject builds a JS-side `{ snapshots, ax }` matching the State type
// from pkg/spec-api. ax is currently a stub returning nothing. // from pkg/spec-api. ax is backed by the parsed uiautomator hierarchy when
func stateObject(runtime *goja.Runtime, snapshots Snapshots) (*goja.Object, error) { // one is provided.
func stateObject(runtime *goja.Runtime, snapshots Snapshots, tree *hierarchy.Tree) (*goja.Object, error) {
state := runtime.NewObject() state := runtime.NewObject()
snapshotsObject := runtime.NewObject() snapshotsObject := runtime.NewObject()
for key, raw := range snapshots { for key, raw := range snapshots {
@@ -27,20 +30,59 @@ func stateObject(runtime *goja.Runtime, snapshots Snapshots) (*goja.Object, erro
if err := state.Set("snapshots", snapshotsObject); err != nil { if err := state.Set("snapshots", snapshotsObject); err != nil {
return nil, err return nil, err
} }
if err := state.Set("ax", accessibilityObject(runtime, tree)); err != nil {
accessibility := runtime.NewObject()
if err := accessibility.Set("find", runtime.ToValue(func(string) goja.Value { return goja.Undefined() })); err != nil {
return nil, err
}
if err := accessibility.Set("findAll", runtime.ToValue(func(string) []goja.Value { return nil })); err != nil {
return nil, err
}
if err := state.Set("ax", accessibility); err != nil {
return nil, err return nil, err
} }
return state, nil return state, nil
} }
func accessibilityObject(runtime *goja.Runtime, tree *hierarchy.Tree) *goja.Object {
accessibility := runtime.NewObject()
find := func(selector string) goja.Value {
if tree == nil {
return goja.Undefined()
}
element := tree.Find(selector)
if element == nil {
return goja.Undefined()
}
return elementObject(runtime, element, selector)
}
findAll := func(selector string) []goja.Value {
if tree == nil {
return nil
}
elements := tree.FindAll(selector)
result := make([]goja.Value, len(elements))
for index, element := range elements {
result[index] = elementObject(runtime, element, selector)
}
return result
}
_ = accessibility.Set("find", runtime.ToValue(find))
_ = accessibility.Set("findAll", runtime.ToValue(findAll))
return accessibility
}
func elementObject(runtime *goja.Runtime, element *hierarchy.Element, selector string) goja.Value {
object := runtime.NewObject()
centerX, centerY := element.Bounds.Center()
_ = object.Set("id", element.ResourceID)
_ = object.Set("text", element.Text)
_ = object.Set("desc", element.Description)
_ = object.Set("class", element.Class)
_ = object.Set("x", centerX)
_ = object.Set("y", centerY)
_ = object.Set(tagSelector, selector)
bounds := runtime.NewObject()
_ = bounds.Set("left", element.Bounds.Left)
_ = bounds.Set("top", element.Bounds.Top)
_ = bounds.Set("right", element.Bounds.Right)
_ = bounds.Set("bottom", element.Bounds.Bottom)
_ = object.Set("bounds", bounds)
return object
}
func jsonToJSValue(runtime *goja.Runtime, raw json.RawMessage) (goja.Value, error) { func jsonToJSValue(runtime *goja.Runtime, raw json.RawMessage) (goja.Value, error) {
if len(raw) == 0 { if len(raw) == 0 {
return goja.Undefined(), nil return goja.Undefined(), nil
@@ -66,16 +108,48 @@ func jsValueToAction(runtime *goja.Runtime, value goja.Value) (Action, error) {
switch kind { switch kind {
case "Tap": case "Tap":
on := object.Get("on") on := object.Get("on")
return Action{Kind: ActionKindTap, On: stringOf(on)}, nil x, y := coordinatesOf(runtime, on)
return Action{Kind: ActionKindTap, On: selectorOf(runtime, on), X: x, Y: y}, nil
case "InputText": case "InputText":
into := object.Get("into") into := object.Get("into")
text := object.Get("text") text := object.Get("text")
return Action{Kind: ActionKindInputText, On: stringOf(into), Text: stringOf(text)}, nil x, y := coordinatesOf(runtime, into)
return Action{Kind: ActionKindInputText, On: selectorOf(runtime, into), Text: stringOf(text), X: x, Y: y}, nil
default: default:
return Action{}, fmt.Errorf("unknown action kind %q", kind) return Action{}, fmt.Errorf("unknown action kind %q", kind)
} }
} }
func selectorOf(runtime *goja.Runtime, value goja.Value) string {
if value == nil || goja.IsNull(value) || goja.IsUndefined(value) {
return ""
}
object := value.ToObject(runtime)
if object == nil {
return value.String()
}
if tag := object.Get(tagSelector); tag != nil && !goja.IsUndefined(tag) {
return tag.String()
}
return value.String()
}
func coordinatesOf(runtime *goja.Runtime, value goja.Value) (int, int) {
if value == nil || goja.IsNull(value) || goja.IsUndefined(value) {
return 0, 0
}
object := value.ToObject(runtime)
if object == nil {
return 0, 0
}
xValue := object.Get("x")
yValue := object.Get("y")
if xValue == nil || yValue == nil || goja.IsUndefined(xValue) || goja.IsUndefined(yValue) {
return 0, 0
}
return int(xValue.ToInteger()), int(yValue.ToInteger())
}
func stringOf(value goja.Value) string { func stringOf(value goja.Value) string {
if value == nil || goja.IsNull(value) || goja.IsUndefined(value) { if value == nil || goja.IsNull(value) || goja.IsUndefined(value) {
return "" return ""
+8 -8
View File
@@ -63,7 +63,7 @@ func TestPushSnapshot_UpdatesExtractorCurrentAndPrevious(t *testing.T) {
if err := verifier.PushSnapshot(Snapshots{ if err := verifier.PushSnapshot(Snapshots{
"screen": json.RawMessage(`"customer_ledger"`), "screen": json.RawMessage(`"customer_ledger"`),
"ledger.balance": json.RawMessage(`1500`), "ledger.balance": json.RawMessage(`1500`),
}); err != nil { }, nil); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -78,7 +78,7 @@ func TestPushSnapshot_UpdatesExtractorCurrentAndPrevious(t *testing.T) {
} }
// Push again: previous should mirror the prior current. // Push again: previous should mirror the prior current.
if err := verifier.PushSnapshot(Snapshots{"ledger.balance": json.RawMessage(`2000`)}); err != nil { if err := verifier.PushSnapshot(Snapshots{"ledger.balance": json.RawMessage(`2000`)}, nil); err != nil {
t.Fatal(err) t.Fatal(err)
} }
balanceValue = verifier.runtime.GlobalObject().Get("balance").ToObject(verifier.runtime) balanceValue = verifier.runtime.GlobalObject().Get("balance").ToObject(verifier.runtime)
@@ -105,7 +105,7 @@ func TestEvaluateProperties_HoldsThenViolates(t *testing.T) {
} }
for index, testCase := range cases { for index, testCase := range cases {
raw, _ := json.Marshal(testCase.balance) raw, _ := json.Marshal(testCase.balance)
if err := verifier.PushSnapshot(Snapshots{"ledger.balance": raw}); err != nil { if err := verifier.PushSnapshot(Snapshots{"ledger.balance": raw}, nil); err != nil {
t.Fatal(err) t.Fatal(err)
} }
verdicts := verifier.EvaluateProperties() verdicts := verifier.EvaluateProperties()
@@ -118,7 +118,7 @@ func TestEvaluateProperties_HoldsThenViolates(t *testing.T) {
func TestNextAction_FromActionsGenerator(t *testing.T) { func TestNextAction_FromActionsGenerator(t *testing.T) {
verifier := newVerifier(t) verifier := newVerifier(t)
mustLoad(t, verifier, helloSpec) mustLoad(t, verifier, helloSpec)
_ = verifier.PushSnapshot(Snapshots{}) _ = verifier.PushSnapshot(Snapshots{}, nil)
action, err := verifier.NextAction() action, err := verifier.NextAction()
if err != nil { if err != nil {
@@ -142,7 +142,7 @@ func TestNextAction_WeightedSelectsByWeight(t *testing.T) {
[99, tapAway], [99, tapAway],
); );
`) `)
_ = verifier.PushSnapshot(Snapshots{}) _ = verifier.PushSnapshot(Snapshots{}, nil)
awayCount := 0 awayCount := 0
homeCount := 0 homeCount := 0
@@ -168,7 +168,7 @@ func TestNextAction_EmptyGeneratorReturnsErrNoAction(t *testing.T) {
mustLoad(t, verifier, ` mustLoad(t, verifier, `
globalThis.actions = __uatu__.actions(() => []); globalThis.actions = __uatu__.actions(() => []);
`) `)
_ = verifier.PushSnapshot(Snapshots{}) _ = verifier.PushSnapshot(Snapshots{}, nil)
_, err := verifier.NextAction() _, err := verifier.NextAction()
if !errors.Is(err, ErrNoAction) { if !errors.Is(err, ErrNoAction) {
@@ -183,7 +183,7 @@ func TestInputText_RoundTrip(t *testing.T) {
__uatu__.inputText({ into: "id:phone", text: "+919876543210" }), __uatu__.inputText({ into: "id:phone", text: "+919876543210" }),
]); ]);
`) `)
_ = verifier.PushSnapshot(Snapshots{}) _ = verifier.PushSnapshot(Snapshots{}, nil)
action, err := verifier.NextAction() action, err := verifier.NextAction()
if err != nil { if err != nil {
@@ -202,7 +202,7 @@ func TestPushSnapshot_FeedsSnapshotsToExtractorState(t *testing.T) {
mustLoad(t, verifier, ` mustLoad(t, verifier, `
globalThis.captured = __uatu__.extract(state => state.snapshots["k"]); globalThis.captured = __uatu__.extract(state => state.snapshots["k"]);
`) `)
if err := verifier.PushSnapshot(Snapshots{"k": json.RawMessage(`"hello"`)}); err != nil { if err := verifier.PushSnapshot(Snapshots{"k": json.RawMessage(`"hello"`)}, nil); err != nil {
t.Fatal(err) t.Fatal(err)
} }
value := verifier.runtime.GlobalObject().Get("captured").ToObject(verifier.runtime).Get("current") value := verifier.runtime.GlobalObject().Get("captured").ToObject(verifier.runtime).Get("current")
+5 -3
View File
@@ -7,6 +7,7 @@ import (
"github.com/dop251/goja" "github.com/dop251/goja"
"github.com/priyanshujain/uatu/internal/hierarchy"
"github.com/priyanshujain/uatu/internal/ltl" "github.com/priyanshujain/uatu/internal/ltl"
) )
@@ -79,9 +80,10 @@ func (v *Verifier) Load(source string) error {
} }
// PushSnapshot updates the JS-side state and refreshes every extractor's // PushSnapshot updates the JS-side state and refreshes every extractor's
// current/previous values in registration order. // current/previous values in registration order. Passing a nil tree is
func (v *Verifier) PushSnapshot(snapshots Snapshots) error { // allowed and yields an empty ax scope.
state, err := stateObject(v.runtime, snapshots) func (v *Verifier) PushSnapshot(snapshots Snapshots, tree *hierarchy.Tree) error {
state, err := stateObject(v.runtime, snapshots, tree)
if err != nil { if err != nil {
return fmt.Errorf("build state: %w", err) return fmt.Errorf("build state: %w", err)
} }