Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
58 changes: 56 additions & 2 deletions ext/encoders_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -249,8 +249,10 @@ func TestEncodersCosts(t *testing.T) {
"x": 100,
},
estimatedCost: checker.CostEstimate{Min: 2, Max: math.MaxUint64},
actualCost: 1,
version: 1,
// json.encode reports an unbounded actual cost, which saturates the
// running total rather than overflowing it back to a small value.
actualCost: math.MaxUint64,
version: 1,
},
}
for _, tc := range tests {
Expand Down Expand Up @@ -337,3 +339,55 @@ func TestJSONEncodeCostUnbounded(t *testing.T) {
}
}

// TestJSONEncodeCostUnboundedWithAccruedCost checks that the unbounded cost of json.encode
// survives being combined with cost accrued elsewhere in the expression, and that it still
// trips a cost limit. A total which overflowed back to a small value would not.
func TestJSONEncodeCostUnboundedWithAccruedCost(t *testing.T) {
env, err := cel.NewEnv(Encoders(EncodersVersion(1)), cel.Variable("v", cel.StringType))
if err != nil {
t.Fatalf("cel.NewEnv() failed: %v", err)
}
exprs := []string{
"json.encode(v)",
"json.encode(v) == json.encode(v)",
"size(v) > 0 && json.encode(v) != ''",
}
for _, expr := range exprs {
t.Run(expr, func(t *testing.T) {
ast, iss := env.Compile(expr)
if iss.Err() != nil {
t.Fatalf("env.Compile(%q) failed: %v", expr, iss.Err())
}
est, err := env.EstimateCost(ast, testCostHintEstimator{})
if err != nil {
t.Fatalf("env.EstimateCost() failed: %v", err)
}
if est.Max != math.MaxUint64 {
t.Errorf("env.EstimateCost() got max %d, wanted %d", est.Max, uint64(math.MaxUint64))
}

prg, err := env.Program(ast, cel.CostTracking(nil))
if err != nil {
t.Fatalf("env.Program() failed: %v", err)
}
_, det, err := prg.Eval(map[string]any{"v": "hello"})
if err != nil {
t.Fatalf("prg.Eval() failed: %v", err)
}
if det.ActualCost() == nil {
t.Fatal("det.ActualCost() got nil, wanted a value")
}
if *det.ActualCost() != math.MaxUint64 {
t.Errorf("det.ActualCost() got %d, wanted %d", *det.ActualCost(), uint64(math.MaxUint64))
}

limited, err := env.Program(ast, cel.CostLimit(1000))
if err != nil {
t.Fatalf("env.Program() failed: %v", err)
}
if _, _, err := limited.Eval(map[string]any{"v": "hello"}); err == nil {
t.Error("prg.Eval() got nil error, wanted cost limit exceeded")
}
})
}
}
14 changes: 7 additions & 7 deletions interpreter/runtimecost.go
Original file line number Diff line number Diff line change
Expand Up @@ -108,7 +108,7 @@ func (ct *costTrackerFactory) Observe(vars Activation, id int64, programStep any
case ConstantQualifier:
// TODO: Push identifiers on to the stack before observing constant qualifiers that apply to them
// and enable the below pop. Once enabled this can case can be collapsed into the Qualifier case.
tracker.cost++
tracker.cost = cost.SafeAdd(tracker.cost, 1)
case InterpretableConst:
// zero cost
case InterpretableAttribute:
Expand All @@ -118,7 +118,7 @@ func (ct *costTrackerFactory) Observe(vars Activation, id int64, programStep any
tracker.stack.drop(a.falsy.ID(), a.truthy.ID(), a.expr.ID())
default:
tracker.stack.drop(t.Attr().ID())
tracker.cost += common.SelectAndIdentCost
tracker.cost = cost.SafeAdd(tracker.cost, common.SelectAndIdentCost)
}
if !tracker.presenceTestHasCost {
if _, isTestOnly := programStep.(*evalTestOnly); isTestOnly {
Expand Down Expand Up @@ -150,20 +150,20 @@ func (ct *costTrackerFactory) Observe(vars Activation, id int64, programStep any
case *evalFold:
tracker.stack.drop(t.iterRange.ID())
case Qualifier:
tracker.cost++
tracker.cost = cost.SafeAdd(tracker.cost, 1)
case InterpretableCall:
if argVals, ok := tracker.stack.dropArgs(t.Args()); ok {
tracker.cost += tracker.costCall(t, argVals, val)
tracker.cost = cost.SafeAdd(tracker.cost, tracker.costCall(t, argVals, val))
}
case InterpretableConstructor:
tracker.stack.dropArgs(t.InitVals())
switch t.Type() {
case types.ListType:
tracker.cost += common.ListCreateBaseCost
tracker.cost = cost.SafeAdd(tracker.cost, common.ListCreateBaseCost)
case types.MapType:
tracker.cost += common.MapCreateBaseCost
tracker.cost = cost.SafeAdd(tracker.cost, common.MapCreateBaseCost)
default:
tracker.cost += common.StructCreateBaseCost
tracker.cost = cost.SafeAdd(tracker.cost, common.StructCreateBaseCost)
}
}
tracker.stack.push(val, id)
Expand Down
49 changes: 49 additions & 0 deletions interpreter/runtimecost_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -904,3 +904,52 @@ func TestRuntimeCost(t *testing.T) {
})
}
}

// TestRuntimeCostUnboundedOverloadSaturates verifies that an overload which reports an unbounded
// actual cost saturates the running total rather than overflowing it back to a small value. A
// wrapped total would silently defeat CostTrackerLimit.
func TestRuntimeCostUnboundedOverloadSaturates(t *testing.T) {
unbounded := func(args []ref.Val, result ref.Val) *uint64 {
maxCost := uint64(math.MaxUint64)
return &maxCost
}
vars := []*decls.VariableDecl{
decls.NewVariable("str1", types.StringType),
decls.NewVariable("str2", types.StringType),
}
in := map[string]any{"str1": "val1", "str2": "val2222222"}

tests := []struct {
name string
expr string
}{
// Each expression accrues cost from the variable references before the unbounded
// call is observed, and continues to accrue cost after it in the last two cases.
{name: "unbounded call", expr: `"abcdefg".contains(str1 + str2)`},
{name: "unbounded call in conjunction", expr: `str1 != "" && "abcdefg".contains(str1 + str2)`},
{name: "unbounded call then comparison", expr: `"abcdefg".contains(str1 + str2) == true`},
{name: "unbounded call in list", expr: `["abcdefg".contains(str1 + str2), true]`},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
opts := []CostTrackerOption{OverloadCostTracker(overloads.ContainsString, unbounded)}
actualCost, _, err := computeCost(t, tc.expr, vars, constructActivation(t, in), opts)
if err != nil {
t.Fatalf("computeCost() failed: %v", err)
}
if actualCost != math.MaxUint64 {
t.Errorf("computeCost() got cost %d, wanted %d", actualCost, uint64(math.MaxUint64))
}

// The saturated cost must also trip a cost limit.
limitOpts := []CostTrackerOption{
OverloadCostTracker(overloads.ContainsString, unbounded),
CostTrackerLimit(1000),
}
_, _, err = computeCost(t, tc.expr, vars, constructActivation(t, in), limitOpts)
if err == nil {
t.Error("computeCost() with a cost limit got nil error, wanted cost limit exceeded")
}
})
}
}