| // Copyright 2020 Google LLC |
| // |
| // Licensed under the Apache License, Version 2.0 (the "License"); |
| // you may not use this file except in compliance with the License. |
| // You may obtain a copy of the License at |
| // |
| // http://www.apache.org/licenses/LICENSE-2.0 |
| // |
| // Unless required by applicable law or agreed to in writing, software |
| // distributed under the License is distributed on an "AS IS" BASIS, |
| // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. |
| // See the License for the specific language governing permissions and |
| // limitations under the License. |
| |
| package ext |
| |
| import ( |
| "encoding/base64" |
| "encoding/json" |
| "fmt" |
| "math" |
| |
| "cel.dev/cel-go/cel" |
| "cel.dev/cel-go/checker" |
| "cel.dev/cel-go/common/cost" |
| "cel.dev/cel-go/common/types" |
| "cel.dev/cel-go/common/types/ref" |
| "cel.dev/cel-go/interpreter" |
| "google.golang.org/protobuf/encoding/protojson" |
| "google.golang.org/protobuf/types/known/structpb" |
| ) |
| |
| // Encoders returns a cel.EnvOption to configure extended functions for string, byte, and object |
| // encodings. |
| // |
| // # Base64.Decode |
| // |
| // Decodes base64-encoded string to bytes. |
| // |
| // This function will return an error if the string input is not base64-encoded. |
| // |
| // base64.decode(<string>) -> <bytes> |
| // |
| // Examples: |
| // |
| // base64.decode('aGVsbG8=') // return b'hello' |
| // base64.decode('aGVsbG8') // return b'hello' |
| // |
| // # Base64.Encode |
| // |
| // Encodes bytes to a base64-encoded string. |
| // |
| // base64.encode(<bytes>) -> <string> |
| // |
| // Examples: |
| // |
| // base64.encode(b'hello') // return b'aGVsbG8=' |
| // |
| // # JSON.Encode |
| // |
| // Introduced at version: 1 |
| // |
| // Encodes a CEL value to a JSON string. |
| // |
| // json.encode(<dyn>) -> <string> |
| // |
| // Examples: |
| // |
| // json.encode({'hello': 'world'}) // return '{"hello":"world"}' |
| func Encoders(options ...EncodersOption) cel.EnvOption { |
| l := &encoderLib{version: math.MaxUint32} |
| for _, o := range options { |
| l = o(l) |
| } |
| return cel.Lib(l) |
| } |
| |
| // EncodersOption declares a functional operator for configuring encoder extensions. |
| type EncodersOption func(*encoderLib) *encoderLib |
| |
| // EncodersVersion sets the library version for encoder extensions. |
| func EncodersVersion(version uint32) EncodersOption { |
| return func(lib *encoderLib) *encoderLib { |
| lib.version = version |
| return lib |
| } |
| } |
| |
| type encoderLib struct { |
| version uint32 |
| } |
| |
| func (*encoderLib) LibraryName() string { |
| return "cel.lib.ext.encoders" |
| } |
| |
| func (lib *encoderLib) CompileOptions() []cel.EnvOption { |
| opts := []cel.EnvOption{ |
| cel.Function("base64.decode", |
| cel.Overload("base64_decode_string", []*cel.Type{cel.StringType}, cel.BytesType, |
| cel.UnaryBinding(func(str ref.Val) ref.Val { |
| s := str.(types.String) |
| return bytesOrError(base64DecodeString(string(s))) |
| }))), |
| cel.Function("base64.encode", |
| cel.Overload("base64_encode_bytes", []*cel.Type{cel.BytesType}, cel.StringType, |
| cel.UnaryBinding(func(bytes ref.Val) ref.Val { |
| b := bytes.(types.Bytes) |
| return stringOrError(base64EncodeBytes([]byte(b))) |
| }))), |
| } |
| if lib.version >= 1 { |
| estimators := []checker.CostOption{ |
| checker.OverloadCostEstimate("base64_decode_string", estimateDecode), |
| checker.OverloadCostEstimate("base64_encode_bytes", estimateEncode), |
| checker.OverloadCostEstimate("json_encode_dyn", estimateJSONEncode), |
| } |
| opts = append(opts, cel.CostEstimatorOptions(estimators...)) |
| opts = append(opts, |
| cel.Function("json.encode", |
| cel.Overload("json_encode_dyn", []*cel.Type{cel.DynType}, cel.StringType, |
| cel.UnaryBinding(func(val ref.Val) ref.Val { |
| return stringOrError(jsonEncodeValue(val)) |
| }))), |
| ) |
| } |
| return opts |
| } |
| |
| func (lib *encoderLib) ProgramOptions() []cel.ProgramOption { |
| var opts []cel.ProgramOption |
| if lib.version >= 1 { |
| trackers := []interpreter.CostTrackerOption{ |
| interpreter.OverloadCostTracker("base64_decode_string", trackDecode), |
| interpreter.OverloadCostTracker("base64_encode_bytes", trackEncode), |
| interpreter.OverloadCostTracker("json_encode_dyn", trackJSONEncode), |
| } |
| opts = append(opts, cel.CostTrackerOptions(trackers...)) |
| } |
| return opts |
| } |
| |
| func base64DecodeString(str string) ([]byte, error) { |
| b, err := base64.StdEncoding.DecodeString(str) |
| if err == nil { |
| return b, nil |
| } |
| if _, tryAltEncoding := err.(base64.CorruptInputError); tryAltEncoding { |
| return base64.RawStdEncoding.DecodeString(str) |
| } |
| return nil, err |
| } |
| |
| func base64EncodeBytes(bytes []byte) (string, error) { |
| return base64.StdEncoding.EncodeToString(bytes), nil |
| } |
| |
| func estimateEncode(estimator checker.CostEstimator, target *checker.AstNode, args []checker.AstNode) *checker.CallEstimate { |
| if len(args) != 1 { |
| return nil |
| } |
| sz := estimateSize(estimator, args[0]) |
| cost := sz.MultiplyByCostFactor(stringCostFactor).Add(callCostEstimate) |
| resSize := estimateEncodeSize(sz) |
| return &checker.CallEstimate{CostEstimate: cost, ResultSize: &resSize} |
| } |
| |
| func estimateJSONEncode(estimator checker.CostEstimator, target *checker.AstNode, args []checker.AstNode) *checker.CallEstimate { |
| if len(args) != 1 { |
| return nil |
| } |
| size := estimateJSONEncodeSize() |
| return &checker.CallEstimate{CostEstimate: checker.UnknownCostEstimate(), ResultSize: &size} |
| } |
| |
| func estimateDecode(estimator checker.CostEstimator, target *checker.AstNode, args []checker.AstNode) *checker.CallEstimate { |
| if len(args) != 1 { |
| return nil |
| } |
| sz := estimateSize(estimator, args[0]) |
| cost := sz.MultiplyByCostFactor(stringCostFactor).Add(callCostEstimate) |
| resSize := estimateDecodeSize(sz) |
| return &checker.CallEstimate{CostEstimate: cost, ResultSize: &resSize} |
| } |
| |
| func trackEncode(args []ref.Val, _ ref.Val) *uint64 { |
| sz := actualSize(args[0]) |
| total := cost.SafeAdd(cost.SafeMultiplyByFactor(sz, stringCostFactor), callCost) |
| return &total |
| } |
| |
| func trackJSONEncode(args []ref.Val, _ ref.Val) *uint64 { |
| maxCost := uint64(math.MaxUint64) |
| return &maxCost |
| } |
| |
| func trackDecode(args []ref.Val, _ ref.Val) *uint64 { |
| sz := actualSize(args[0]) |
| total := cost.SafeAdd(cost.SafeMultiplyByFactor(sz, stringCostFactor), callCost) |
| return &total |
| } |
| |
| func estimateEncodeSize(sz checker.SizeEstimate) checker.SizeEstimate { |
| minVal := (sz.Min*4 + 2) / 3 |
| maxVal := (sz.Max*4 + 2) / 3 |
| if sz.Max > math.MaxUint64/4 { |
| maxVal = math.MaxUint64 |
| } |
| return checker.SizeEstimate{Min: minVal, Max: maxVal} |
| } |
| |
| func estimateJSONEncodeSize() checker.SizeEstimate { |
| // TODO: provide a more sophisticated size estimate based on the CEL value's type. |
| return checker.UnknownSizeEstimate() |
| } |
| |
| func estimateDecodeSize(sz checker.SizeEstimate) checker.SizeEstimate { |
| minVal := sz.Min * 3 / 4 |
| maxVal := sz.Max * 3 / 4 |
| return checker.SizeEstimate{Min: minVal, Max: maxVal} |
| } |
| |
| func jsonEncodeValue(val ref.Val) (string, error) { |
| native, err := val.ConvertToNative(types.JSONValueType) |
| if err != nil { |
| return "", err |
| } |
| jsonValue, ok := native.(*structpb.Value) |
| if !ok { |
| return "", fmt.Errorf("cannot convert %T to JSON value", native) |
| } |
| jsonBytes, err := protojson.Marshal(jsonValue) |
| if err != nil { |
| return "", err |
| } |
| var obj interface{} |
| if err := json.Unmarshal(jsonBytes, &obj); err != nil { |
| return "", fmt.Errorf("unmarshaling protojson: %w", err) |
| } |
| // Re-marshal with standard json.Marshal for deterministic compact output |
| jsonBytes, err = json.Marshal(obj) |
| if err != nil { |
| return "", fmt.Errorf("re-marshaling value: %w", err) |
| } |
| return string(jsonBytes), nil |
| } |