@@ -635,10 +635,10 @@ func TestGenerateRecursiveRecordSpecialization(t *testing.T) {
635635 expected , generated )
636636 }
637637 }
638- if got := strings .Count (generated , "var aotRecordFn" ); got != 3 {
638+ if got := strings .Count (generated , "var aotRecordFn" ); got != 2 {
639639 t .Fatalf (
640- "generated %d record-specialized functions, want 3 ; " +
641- "record equality must retain generic semantics:\n %s" ,
640+ "generated %d record-specialized functions, want 2 ; " +
641+ "nullable recursive fields must retain generic semantics:\n %s" ,
642642 got ,
643643 generated ,
644644 )
@@ -670,6 +670,40 @@ func TestGenerateRecursiveRecordSpecialization(t *testing.T) {
670670 }
671671}
672672
673+ func TestNullableRecordProducerUsesGenericFallback (t * testing.T ) {
674+ ns := lang .FindOrCreateNamespace (
675+ lang .NewSymbol ("codegen.nullable-record-producer" ),
676+ )
677+ ns .ReferAllSnapshot (lang .NSCore , nil )
678+ lang .PushThreadBindings (lang .NewMap (lang .VarCurrentNS , ns ))
679+ defer lang .PopThreadBindings ()
680+
681+ ReadEval (`
682+ (defrecord MaybeLink [value next])
683+ (defn maybe-link [value]
684+ (if (zero? value)
685+ nil
686+ (->MaybeLink value (maybe-link (dec value)))))
687+ (defn sample-maybe-link []
688+ (maybe-link 0))` )
689+
690+ var output bytes.Buffer
691+ if err := NewGenerator (& output ).Generate (ns ); err != nil {
692+ t .Fatalf ("generate nullable record producer: %v" , err )
693+ }
694+ if strings .Contains (output .String (), "var aotRecordFn" ) {
695+ t .Fatalf (
696+ "nullable record producer received an unsound pointer-returning specialization:\n %s" ,
697+ output .String (),
698+ )
699+ }
700+ if got := ns .FindInternedVar (
701+ lang .NewSymbol ("sample-maybe-link" ),
702+ ).Invoke (); got != nil {
703+ t .Fatalf ("nullable record producer returned %#v, want nil" , got )
704+ }
705+ }
706+
673707func TestGenerateBooleanRecordSpecialization (t * testing.T ) {
674708 ns := lang .FindOrCreateNamespace (lang .NewSymbol ("codegen.record-bool-specialization" ))
675709 ns .ReferAllSnapshot (lang .NSCore , nil )
@@ -2958,14 +2992,18 @@ func TestGenerateOwnedNestedVectorUpdateRegion(t *testing.T) {
29582992 (mapv
29592993 (fn [row]
29602994 (assoc row 0 (+ (nth row 0) delta)))
2961- updated)))` )
2995+ updated)))
2996+ (defn ordered-assoc [values divisor]
2997+ (let [row (nth values 0)
2998+ updated (assoc row -777777 99 1 (quot 888888 divisor))]
2999+ (assoc-in values [0 0] (nth updated 0))))` )
29623000
29633001 var output bytes.Buffer
29643002 generator := NewGenerator (& output )
29653003 if err := generator .Generate (ns ); err != nil {
29663004 t .Fatalf ("generate owned nested vector region: %v" , err )
29673005 }
2968- for _ , name := range []string {"update-cell" , "update-all" } {
3006+ for _ , name := range []string {"update-cell" , "update-all" , "ordered-assoc" } {
29693007 vr := ns .FindInternedVar (lang .NewSymbol (name ))
29703008 target := generator .aotCallTargets [vr ]
29713009 if target == nil || target .ownedVectorAnalysis == nil {
@@ -2997,6 +3035,23 @@ func TestGenerateOwnedNestedVectorUpdateRegion(t *testing.T) {
29973035 expected , generated )
29983036 }
29993037 }
3038+ valueIndex := strings .Index (
3039+ generated ,
3040+ "lang.Numbers.Quotient(int64(888888)" ,
3041+ )
3042+ assocIndex := strings .Index (
3043+ generated ,
3044+ ".AssocCopy(lang.IntCast(int64(-777777))" ,
3045+ )
3046+ if valueIndex < 0 || assocIndex < 0 || valueIndex > assocIndex {
3047+ t .Fatalf (
3048+ "owned vector assoc did not evaluate all operands before mutation " +
3049+ "(value=%d assoc=%d):\n %s" ,
3050+ valueIndex ,
3051+ assocIndex ,
3052+ generated ,
3053+ )
3054+ }
30003055
30013056 update := ns .FindInternedVar (lang .NewSymbol ("update-all" ))
30023057 original := lang .NewVector (
0 commit comments