diff --git a/record.ml b/record.ml index c3d5c66..1aa8a02 100644 --- a/record.ml +++ b/record.ml @@ -28,3 +28,13 @@ let () = let original = Mutable.{ count = 1; label = "grey" } in let changed = Mutable.set_count 2 original in assert (original.Mutable.count = 1 && changed.Mutable.count = 2) + +let () = + let calls = ref 0 in + let original = Colour.{ red = 3; green = 4; blue = 5 } in + let changed = Colour.update_green ~f:(fun n -> incr calls; n + 1) original in + assert (Colour.green changed = 5 && !calls = 1); + assert (Box.update_contents ~f:string_of_int Box.{ contents = 9 } = Box.{ contents = "9" }); + (try ignore (Colour.update_red ~f:(fun _ -> failwith "callback") original); assert false + with Failure message -> assert (message = "callback")); + assert (Colour.red original = 3) diff --git a/rewrite.ml b/rewrite.ml index 4e607e1..0c45f82 100644 --- a/rewrite.ml +++ b/rewrite.ml @@ -31,13 +31,17 @@ let generate declaration = let name = field.pld_name.txt in let getter = lambda ~loc Nolabel source_pattern (Exp.field ~loc source (ident ~loc name)) in - let replacement = Exp.record ~loc [ident ~loc name, value ~loc "__fg_value"] + let replace expression = Exp.record ~loc [ident ~loc name, expression] (if List.length fields = 1 then None else Some source) in let setter = lambda ~loc Nolabel (variable ~loc "__fg_value") - (lambda ~loc Nolabel source_pattern replacement) in + (lambda ~loc Nolabel source_pattern (replace (value ~loc "__fg_value"))) in + let updater = lambda ~loc (Labelled "f") (variable ~loc "__fg_function") + (lambda ~loc Nolabel source_pattern + (replace (Exp.apply ~loc (value ~loc "__fg_function") + [Nolabel, Exp.field ~loc source (ident ~loc name)]))) in List.map (fun (name, body) -> Str.value ~loc Nonrecursive [Vb.mk ~loc (variable ~loc name) body]) - [name, getter; "set_" ^ name, setter]) fields) + [name, getter; "set_" ^ name, setter; "update_" ^ name, updater]) fields) let structure mapper items = List.concat (List.map (fun item ->