diff --git a/mlf-codegen/src/lib.rs b/mlf-codegen/src/lib.rs index a3e0465..f29d0a7 100644 --- a/mlf-codegen/src/lib.rs +++ b/mlf-codegen/src/lib.rs @@ -378,6 +378,85 @@ fn build_param_properties( (properties, required) } +/// Shared iteration for records and inline object types. Returns +/// `(properties, required, nullable)`: +/// +/// * `properties` — field-name → emitted JSON type, in declaration order. +/// * `required` — non-optional field names. +/// * `nullable` — fields whose MLF type is `T | null` (ATProto spec encodes +/// nullability as a sibling `nullable` array on the enclosing object, +/// not as a union member). The null is stripped from the emitted type. +fn collect_object_fields( + fields: &[Field], + usage_counts: &HashMap, + workspace: &Workspace, + current_namespace: &str, +) -> (Map, Vec, Vec) { + let mut properties = Map::new(); + let mut required = Vec::new(); + let mut nullable = Vec::new(); + for field in fields { + if !field.optional { + required.push(field.name.name.clone()); + } + let (effective_ty, is_nullable) = match strip_nullable(&field.ty) { + Some(stripped) => (stripped, true), + None => (field.ty.clone(), false), + }; + if is_nullable { + nullable.push(field.name.name.clone()); + } + let mut field_json = generate_type_json(&effective_ty, usage_counts, workspace, current_namespace); + add_description_from_docs(&mut field_json, &field.docs); + properties.insert(field.name.name.clone(), field_json); + } + (properties, required, nullable) +} + +/// If `ty` is a union that includes a `null` primitive, return the type +/// with the null stripped — the ATProto-side lowering of MLF's `T | null` +/// pattern, which is how we express field nullability in the surface +/// language. Returns `None` when `ty` isn't a nullable union. +/// +/// Unwraps parenthesized types transparently so `(T | null)` behaves the +/// same as `T | null`. Collapses a single-remaining-member union to just +/// that member. +fn strip_nullable(ty: &Type) -> Option { + if let Type::Parenthesized { inner, .. } = ty { + return strip_nullable(inner); + } + let Type::Union { types, closed, span } = ty else { + return None; + }; + let has_null = types.iter().any(is_null_primitive); + if !has_null { + return None; + } + let remaining: Vec = types + .iter() + .filter(|t| !is_null_primitive(t)) + .cloned() + .collect(); + // A single-member union degenerates to the member itself — the same + // invariant our Lexicon→MLF pass maintains, kept symmetric here. + if remaining.len() == 1 { + return Some(remaining.into_iter().next().unwrap()); + } + Some(Type::Union { + types: remaining, + closed: *closed, + span: *span, + }) +} + +fn is_null_primitive(ty: &Type) -> bool { + match ty { + Type::Primitive { kind: PrimitiveType::Null, .. } => true, + Type::Parenthesized { inner, .. } => is_null_primitive(inner), + _ => false, + } +} + /// Build a `type: "params"` JSON object from the output of /// [`build_param_properties`]. `required` is omitted when empty, and an /// empty `properties` object is still emitted (queries always carry an @@ -436,21 +515,13 @@ fn unresolved_ref_fallback(path: &Path) -> String { } fn generate_record_json(record: &Record, usage_counts: &HashMap, workspace: &Workspace, current_namespace: &str) -> Value { - let mut required = Vec::new(); - let mut properties = Map::new(); - - for field in &record.fields { - if !field.optional { - required.push(field.name.name.clone()); - } - let mut field_json = generate_type_json(&field.ty, usage_counts, workspace, current_namespace); - add_description_from_docs(&mut field_json, &field.docs); - properties.insert(field.name.name.clone(), field_json); - } + let (properties, required, nullable) = + collect_object_fields(&record.fields, usage_counts, workspace, current_namespace); let mut record_obj = Map::new(); record_obj.insert("type".to_string(), json!("object")); insert_opt_list(&mut record_obj, "required", &required); + insert_opt_list(&mut record_obj, "nullable", &nullable); record_obj.insert("properties".to_string(), Value::Object(properties)); // Check for @key annotation, default to "tid" @@ -674,20 +745,13 @@ fn generate_type_json(ty: &Type, usage_counts: &HashMap, workspac Value::Object(union_obj) } Type::Object { fields, .. } => { - let mut required = Vec::new(); - let mut properties = Map::new(); - for field in fields { - if !field.optional { - required.push(field.name.name.clone()); - } - let mut field_json = generate_type_json(&field.ty, usage_counts, workspace, current_namespace); - add_description_from_docs(&mut field_json, &field.docs); - properties.insert(field.name.name.clone(), field_json); - } + let (properties, required, nullable) = + collect_object_fields(fields, usage_counts, workspace, current_namespace); let mut obj = Map::new(); obj.insert("type".to_string(), json!("object")); insert_opt_list(&mut obj, "required", &required); + insert_opt_list(&mut obj, "nullable", &nullable); obj.insert("properties".to_string(), Value::Object(properties)); Value::Object(obj) }