From e9073fa3cb60b3cda0b73d3f09276db9ce7a46b3 Mon Sep 17 00:00:00 2001 From: Darin Kishore Date: Wed, 4 Feb 2026 11:25:56 -0800 Subject: [PATCH] feat: replace wiptype with bamltype and vendor facet --- .github/workflows/qol.yaml | 16 +- .gitignore | 5 +- Cargo.lock | 247 +++- crates/baml-bridge-derive/src/lib.rs | 3 +- crates/baml-bridge/tests/golden.rs | 6 +- crates/bamltype-derive/Cargo.toml | 17 + crates/bamltype-derive/src/lib.rs | 1148 ++++++++++++++++ crates/bamltype/Cargo.toml | 50 + crates/bamltype/src/compat.rs | 385 ++++++ crates/bamltype/src/convert.rs | 871 ++++++++++++ crates/bamltype/src/facet_ext.rs | 82 ++ crates/bamltype/src/lib.rs | 462 +++++++ crates/bamltype/src/schema_builder.rs | 660 +++++++++ .../bamltype/tests/contract_bridge_oracle.rs | 1202 +++++++++++++++++ .../tests/contract_bridge_ui_messages.rs | 36 + crates/bamltype/tests/integration.rs | 1190 ++++++++++++++++ crates/bamltype/tests/parity_bridge_api.rs | 166 +++ crates/bamltype/tests/ui.rs | 5 + crates/bamltype/tests/ui/as_enum_data_enum.rs | 9 + .../tests/ui/as_enum_data_enum.stderr | 8 + crates/bamltype/tests/ui/function_type.rs | 8 + crates/bamltype/tests/ui/function_type.stderr | 5 + .../tests/ui/large_int_without_repr.rs | 8 + .../tests/ui/large_int_without_repr.stderr | 5 + .../bamltype/tests/ui/map_key_non_string.rs | 9 + .../tests/ui/map_key_non_string.stderr | 5 + .../bamltype/tests/ui/map_key_repr_non_map.rs | 9 + .../tests/ui/map_key_repr_non_map.stderr | 5 + .../tests/ui/non_string_literal_attr.rs | 9 + .../tests/ui/non_string_literal_attr.stderr | 5 + .../bamltype/tests/ui/serde_default_path.rs | 13 + .../tests/ui/serde_default_path.stderr | 5 + crates/bamltype/tests/ui/serde_flatten.rs | 10 + crates/bamltype/tests/ui/serde_flatten.stderr | 5 + crates/bamltype/tests/ui/serde_json_value.rs | 8 + .../bamltype/tests/ui/serde_json_value.stderr | 5 + .../bamltype/tests/ui/serde_skip_variant.rs | 10 + .../tests/ui/serde_skip_variant.stderr | 5 + crates/bamltype/tests/ui/serde_untagged.rs | 10 + .../bamltype/tests/ui/serde_untagged.stderr | 5 + crates/bamltype/tests/ui/trait_object.rs | 8 + crates/bamltype/tests/ui/trait_object.stderr | 5 + .../bamltype/tests/ui/tuple_enum_variant.rs | 8 + .../tests/ui/tuple_enum_variant.stderr | 5 + crates/bamltype/tests/ui/tuple_field.rs | 8 + crates/bamltype/tests/ui/tuple_field.stderr | 5 + crates/bamltype/tests/ui/tuple_struct.rs | 6 + crates/bamltype/tests/ui/tuple_struct.stderr | 5 + crates/bamltype/tests/ui/unit_struct.rs | 6 + crates/bamltype/tests/ui/unit_struct.stderr | 5 + .../tests/ui/unsupported_baml_attr.rs | 9 + .../tests/ui/unsupported_baml_attr.stderr | 5 + crates/dspy-rs/Cargo.toml | 3 +- .../examples/16-insurance-claim-prompt.rs | 30 +- crates/dspy-rs/src/adapter/chat.rs | 98 +- crates/dspy-rs/src/core/lm/client_registry.rs | 4 +- crates/dspy-rs/src/core/signature.rs | 4 +- crates/dspy-rs/src/lib.rs | 62 +- crates/dspy-rs/src/optimizer/copro.rs | 1 - crates/dspy-rs/src/optimizer/gepa.rs | 1 - crates/dspy-rs/src/optimizer/mipro.rs | 1 - crates/dspy-rs/src/predictors/predict.rs | 4 +- .../tests/test_bamltype_attr_contract.rs | 24 + .../tests/test_bamltype_docs_contract.rs | 158 +++ crates/dspy-rs/tests/test_input_format.rs | 22 +- .../dspy-rs/tests/test_typed_prompt_format.rs | 6 +- crates/dsrs-macros/Cargo.toml | 6 +- crates/dsrs-macros/src/lib.rs | 309 +++-- crates/dsrs-macros/src/optim.rs | 58 +- crates/dsrs-macros/src/runtime_path.rs | 19 + .../tests/optim/derive_optimizable.rs | 24 + crates/dsrs-macros/tests/signature_derive.rs | 26 +- .../dsrs-macros/tests/ui/input_unknown_arg.rs | 12 + .../tests/ui/input_unknown_arg.stderr | 5 + .../tests/ui/output_unknown_arg.rs | 12 + .../tests/ui/output_unknown_arg.stderr | 5 + docs/docs/building-blocks/signature.mdx | 11 +- docs/docs/building-blocks/types.mdx | 283 ++-- docs/docs/getting-started/quickstart.mdx | 5 +- .../src/output_format/types.rs | 46 +- .../coercer/ir_ref/coerce_class.rs | 38 +- .../coercer/ir_ref/coerce_enum.rs | 51 +- 82 files changed, 7580 insertions(+), 535 deletions(-) create mode 100644 crates/bamltype-derive/Cargo.toml create mode 100644 crates/bamltype-derive/src/lib.rs create mode 100644 crates/bamltype/Cargo.toml create mode 100644 crates/bamltype/src/compat.rs create mode 100644 crates/bamltype/src/convert.rs create mode 100644 crates/bamltype/src/facet_ext.rs create mode 100644 crates/bamltype/src/lib.rs create mode 100644 crates/bamltype/src/schema_builder.rs create mode 100644 crates/bamltype/tests/contract_bridge_oracle.rs create mode 100644 crates/bamltype/tests/contract_bridge_ui_messages.rs create mode 100644 crates/bamltype/tests/integration.rs create mode 100644 crates/bamltype/tests/parity_bridge_api.rs create mode 100644 crates/bamltype/tests/ui.rs create mode 100644 crates/bamltype/tests/ui/as_enum_data_enum.rs create mode 100644 crates/bamltype/tests/ui/as_enum_data_enum.stderr create mode 100644 crates/bamltype/tests/ui/function_type.rs create mode 100644 crates/bamltype/tests/ui/function_type.stderr create mode 100644 crates/bamltype/tests/ui/large_int_without_repr.rs create mode 100644 crates/bamltype/tests/ui/large_int_without_repr.stderr create mode 100644 crates/bamltype/tests/ui/map_key_non_string.rs create mode 100644 crates/bamltype/tests/ui/map_key_non_string.stderr create mode 100644 crates/bamltype/tests/ui/map_key_repr_non_map.rs create mode 100644 crates/bamltype/tests/ui/map_key_repr_non_map.stderr create mode 100644 crates/bamltype/tests/ui/non_string_literal_attr.rs create mode 100644 crates/bamltype/tests/ui/non_string_literal_attr.stderr create mode 100644 crates/bamltype/tests/ui/serde_default_path.rs create mode 100644 crates/bamltype/tests/ui/serde_default_path.stderr create mode 100644 crates/bamltype/tests/ui/serde_flatten.rs create mode 100644 crates/bamltype/tests/ui/serde_flatten.stderr create mode 100644 crates/bamltype/tests/ui/serde_json_value.rs create mode 100644 crates/bamltype/tests/ui/serde_json_value.stderr create mode 100644 crates/bamltype/tests/ui/serde_skip_variant.rs create mode 100644 crates/bamltype/tests/ui/serde_skip_variant.stderr create mode 100644 crates/bamltype/tests/ui/serde_untagged.rs create mode 100644 crates/bamltype/tests/ui/serde_untagged.stderr create mode 100644 crates/bamltype/tests/ui/trait_object.rs create mode 100644 crates/bamltype/tests/ui/trait_object.stderr create mode 100644 crates/bamltype/tests/ui/tuple_enum_variant.rs create mode 100644 crates/bamltype/tests/ui/tuple_enum_variant.stderr create mode 100644 crates/bamltype/tests/ui/tuple_field.rs create mode 100644 crates/bamltype/tests/ui/tuple_field.stderr create mode 100644 crates/bamltype/tests/ui/tuple_struct.rs create mode 100644 crates/bamltype/tests/ui/tuple_struct.stderr create mode 100644 crates/bamltype/tests/ui/unit_struct.rs create mode 100644 crates/bamltype/tests/ui/unit_struct.stderr create mode 100644 crates/bamltype/tests/ui/unsupported_baml_attr.rs create mode 100644 crates/bamltype/tests/ui/unsupported_baml_attr.stderr create mode 100644 crates/dspy-rs/tests/test_bamltype_attr_contract.rs create mode 100644 crates/dspy-rs/tests/test_bamltype_docs_contract.rs create mode 100644 crates/dsrs-macros/src/runtime_path.rs create mode 100644 crates/dsrs-macros/tests/optim/derive_optimizable.rs create mode 100644 crates/dsrs-macros/tests/ui/input_unknown_arg.rs create mode 100644 crates/dsrs-macros/tests/ui/input_unknown_arg.stderr create mode 100644 crates/dsrs-macros/tests/ui/output_unknown_arg.rs create mode 100644 crates/dsrs-macros/tests/ui/output_unknown_arg.stderr diff --git a/.github/workflows/qol.yaml b/.github/workflows/qol.yaml index 7dd31937..5ada22d4 100644 --- a/.github/workflows/qol.yaml +++ b/.github/workflows/qol.yaml @@ -8,7 +8,9 @@ jobs: runs-on: ubuntu-latest steps: - uses: actions/checkout@v4 - - uses: dtolnay/rust-toolchain@nightly + - uses: dtolnay/rust-toolchain@stable + with: + toolchain: 1.90.0 - run: cargo build test: @@ -16,7 +18,9 @@ jobs: runs-on: ubuntu-latest steps: - uses: actions/checkout@v4 - - uses: dtolnay/rust-toolchain@nightly + - uses: dtolnay/rust-toolchain@stable + with: + toolchain: 1.90.0 - run: cargo test --all-features miri-test: @@ -26,9 +30,10 @@ jobs: MIRIFLAGS: -Zmiri-disable-isolation steps: - uses: actions/checkout@v4 - - uses: dtolnay/rust-toolchain@miri + - uses: dtolnay/rust-toolchain@nightly with: - toolchain: nightly-2025-05-16 # https://github.com/rust-lang/miri/issues/4323 + toolchain: nightly-2026-01-25 + components: miri - run: cargo miri setup - run: cargo miri test --all-features @@ -39,8 +44,9 @@ jobs: timeout-minutes: 5 steps: - uses: actions/checkout@v4 - - uses: dtolnay/rust-toolchain@nightly + - uses: dtolnay/rust-toolchain@stable with: + toolchain: 1.90.0 components: rustfmt, clippy - run: cargo fmt --check - run: cargo clippy --all-targets --all-features diff --git a/.gitignore b/.gitignore index 59896d50..5aac4815 100644 --- a/.gitignore +++ b/.gitignore @@ -2,7 +2,7 @@ # will have compiled files and executables debug/ target/ - +.workspaces/ .beads .jj/ @@ -17,6 +17,9 @@ target/ # Contains mutation testing data **/mutants.out*/ +# LLVM profile/coverage raw output +*.profraw + # RustRover # JetBrains specific template is maintained in a separate JetBrains.gitignore that can # be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore diff --git a/Cargo.lock b/Cargo.lock index 0eec455c..5b4d727b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,15 +2,6 @@ # It is not intended for manual editing. version = 4 -[[package]] -name = "addr2line" -version = "0.24.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dfbe277e56a376000877090da837660b4427aad530e3028d44e0bffe4f89a1c1" -dependencies = [ - "gimli", -] - [[package]] name = "adler2" version = "2.0.1" @@ -134,9 +125,9 @@ dependencies = [ [[package]] name = "anyhow" -version = "1.0.99" +version = "1.0.101" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b0674a1ddeecb70197781e945de4b3b8ffb61fa939a5597bcf48503737663100" +checksum = "5f0e0fee31ef5ed1ba1316088939cea399010ed7731dba877ed44aeb407a75ea" [[package]] name = "arc-swap" @@ -203,7 +194,7 @@ dependencies = [ "arrow-schema", "chrono", "half", - "hashbrown", + "hashbrown 0.15.5", "num", ] @@ -467,21 +458,6 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" -[[package]] -name = "backtrace" -version = "0.3.75" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6806a6321ec58106fea15becdad98371e28d92ccbc7c8f1b3b6dd724fe8f1002" -dependencies = [ - "addr2line", - "cfg-if", - "libc", - "miniz_oxide", - "object", - "rustc-demangle", - "windows-targets 0.52.6", -] - [[package]] name = "baml-bridge" version = "0.1.0" @@ -543,6 +519,37 @@ dependencies = [ "web-time", ] +[[package]] +name = "bamltype" +version = "0.1.0" +dependencies = [ + "anyhow", + "baml-bridge", + "baml-types", + "bamltype-derive", + "facet", + "facet-reflect", + "indexmap", + "internal-baml-jinja", + "jsonish", + "minijinja", + "serde_json", + "sha2", + "thiserror", + "trybuild", +] + +[[package]] +name = "bamltype-derive" +version = "0.1.0" +dependencies = [ + "convert_case", + "proc-macro-crate", + "proc-macro2", + "quote", + "syn 2.0.106", +] + [[package]] name = "base64" version = "0.22.1" @@ -843,6 +850,12 @@ dependencies = [ "web-sys", ] +[[package]] +name = "const-fnv1a-hash" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32b13ea120a812beba79e34316b3942a857c86ec1593cb34f27bb28272ce2cca" + [[package]] name = "const-random" version = "0.1.18" @@ -1166,11 +1179,12 @@ dependencies = [ "anyhow", "arrow", "async-trait", - "baml-bridge", + "bamltype", "bon", "csv", "dsrs_macros", "enum_dispatch", + "facet", "foyer", "futures", "hf-hub", @@ -1197,14 +1211,10 @@ dependencies = [ name = "dsrs_macros" version = "0.7.2" dependencies = [ - "anyhow", - "baml-bridge", "dspy-rs", - "indexmap", + "proc-macro-crate", "proc-macro2", "quote", - "schemars", - "serde", "serde_json", "syn 2.0.106", "trybuild", @@ -1328,6 +1338,82 @@ dependencies = [ "once_cell", ] +[[package]] +name = "facet" +version = "0.43.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e338357cf598728b41e45744d024bdc063338214992361766928a1421bd7541d" +dependencies = [ + "autocfg", + "facet-core", + "facet-macros", +] + +[[package]] +name = "facet-core" +version = "0.43.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a63e0ade4c53b40220614b8fc2a0a0ce21975941b553081521a195c848b2e9c2" +dependencies = [ + "autocfg", + "const-fnv1a-hash", + "iddqd", + "impls", +] + +[[package]] +name = "facet-macro-parse" +version = "0.43.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83ea29147986d0e184600cec533c41d6065c3c3d4b5b5745a8403494ca216b09" +dependencies = [ + "facet-macro-types", + "proc-macro2", + "quote", +] + +[[package]] +name = "facet-macro-types" +version = "0.43.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31b0035cf41c0d4eeee82effc9161512d216d1378dd89c4d8721258429e38597" +dependencies = [ + "proc-macro2", + "quote", + "unsynn", +] + +[[package]] +name = "facet-macros" +version = "0.43.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77a784f2fa36d3165b95639af790249dee0d8efdef7d53f9417cace91697e2e3" +dependencies = [ + "facet-macros-impl", +] + +[[package]] +name = "facet-macros-impl" +version = "0.43.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf8f45c6380398bf74e59b97a20012de571502c609e580d84579d1140e491c1c" +dependencies = [ + "facet-macro-parse", + "facet-macro-types", + "proc-macro2", + "quote", + "unsynn", +] + +[[package]] +name = "facet-reflect" +version = "0.43.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4418c9fceaac9adcd055cc3732954d79b5d67ef04fb855dd219f2b314ba26cff" +dependencies = [ + "facet-core", +] + [[package]] name = "fastant" version = "0.1.10" @@ -1395,6 +1481,12 @@ version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" +[[package]] +name = "foldhash" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb" + [[package]] name = "foreign-types" version = "0.3.2" @@ -1479,7 +1571,7 @@ dependencies = [ "equivalent", "foyer-common", "foyer-intrusive-collections", - "hashbrown", + "hashbrown 0.15.5", "itertools 0.14.0", "madsim-tokio", "mixtrics", @@ -1509,7 +1601,7 @@ dependencies = [ "fs4", "futures-core", "futures-util", - "hashbrown", + "hashbrown 0.15.5", "io-uring", "itertools 0.14.0", "libc", @@ -1666,12 +1758,6 @@ dependencies = [ "wasi 0.14.4+wasi-0.2.4", ] -[[package]] -name = "gimli" -version = "0.31.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "07e28edb80900c19c28f1072f2e8aeca7fa06b23cd4169cefe1af5aa3260783f" - [[package]] name = "glob" version = "0.3.3" @@ -1716,7 +1802,16 @@ checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" dependencies = [ "allocator-api2", "equivalent", - "foldhash", + "foldhash 0.1.5", +] + +[[package]] +name = "hashbrown" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +dependencies = [ + "allocator-api2", ] [[package]] @@ -1985,6 +2080,19 @@ dependencies = [ "zerovec", ] +[[package]] +name = "iddqd" +version = "0.3.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b215e67ed1d1a4b1702acd787c487d16e4c977c5dcbcc4587bdb5ea26b6ce06" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash 0.2.0", + "hashbrown 0.16.1", + "rustc-hash", +] + [[package]] name = "ident_case" version = "1.0.1" @@ -2012,6 +2120,12 @@ dependencies = [ "icu_properties", ] +[[package]] +name = "impls" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a46645bbd70538861a90d0f26c31537cdf1e44aae99a794fb75a664b70951bc" + [[package]] name = "indexmap" version = "2.11.0" @@ -2019,7 +2133,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f2481980430f9f78649238835720ddccc57e52df14ffce1c6f37391d61b563e9" dependencies = [ "equivalent", - "hashbrown", + "hashbrown 0.15.5", "serde", ] @@ -2511,6 +2625,12 @@ dependencies = [ "parking_lot", ] +[[package]] +name = "mutants" +version = "0.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc0287524726960e07b119cebd01678f852f147742ae0d925e6a520dca956126" + [[package]] name = "naive-timer" version = "0.2.0" @@ -2658,15 +2778,6 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "830b246a0e5f20af87141b25c173cd1b609bd7779a4617d6ec582abaf90870f3" -[[package]] -name = "object" -version = "0.36.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "62948e14d923ea95ea2c7c86c71013138b66525b86bdc08d2dcc262bdb497b87" -dependencies = [ - "memchr", -] - [[package]] name = "once_cell" version = "1.21.3" @@ -2808,7 +2919,7 @@ dependencies = [ "chrono", "flate2", "half", - "hashbrown", + "hashbrown 0.15.5", "lz4_flex", "num", "num-bigint", @@ -3381,10 +3492,10 @@ dependencies = [ ] [[package]] -name = "rustc-demangle" -version = "0.1.26" +name = "rustc-hash" +version = "2.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "56f7d92ca342cea22a06f2121d944b4fd82af56988c270852495420f961d4ace" +checksum = "357703d41365b4b27c590e3ed91eabb1b663f07c4c084095e60cbed4362dff0d" [[package]] name = "rustc_version" @@ -4042,29 +4153,26 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.47.1" +version = "1.49.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "89e49afdadebb872d3145a5638b59eb0691ea23e46ca484037cfab3b76b95038" +checksum = "72a2903cd7736441aac9df9d7688bd0ce48edccaadf181c3b90be801e81d3d86" dependencies = [ - "backtrace", "bytes", - "io-uring", "libc", "mio", "parking_lot", "pin-project-lite", "signal-hook-registry", - "slab", "socket2", "tokio-macros", - "windows-sys 0.59.0", + "windows-sys 0.61.0", ] [[package]] name = "tokio-macros" -version = "2.5.0" +version = "2.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e06d43f1345a3bcd39f6a56dbb7dcab2ba47e68e8ac134855e7e2bdbaf8cab8" +checksum = "af407857209536a95c8e56f8231ef2c2e2aff839b22e07a1ffcbc617e9db9fa5" dependencies = [ "proc-macro2", "quote", @@ -4393,6 +4501,17 @@ version = "0.2.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861" +[[package]] +name = "unsynn" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "501a7adf1a4bd9951501e5c66621e972ef8874d787628b7f90e64f936ef7ec0a" +dependencies = [ + "mutants", + "proc-macro2", + "rustc-hash", +] + [[package]] name = "untrusted" version = "0.9.0" diff --git a/crates/baml-bridge-derive/src/lib.rs b/crates/baml-bridge-derive/src/lib.rs index 475657e5..47b636c4 100644 --- a/crates/baml-bridge-derive/src/lib.rs +++ b/crates/baml-bridge-derive/src/lib.rs @@ -1517,8 +1517,7 @@ fn type_ir_for_field( attrs: &FieldAttrs, map_entry: Option<&MapEntryInfo>, ) -> syn::Result { - if attrs.with.is_some() { - let adapter = attrs.with.as_ref().unwrap(); + if let Some(adapter) = attrs.with.as_ref() { return Ok(quote! { <#adapter as ::baml_bridge::BamlAdapter<#ty>>::type_ir() }); } diff --git a/crates/baml-bridge/tests/golden.rs b/crates/baml-bridge/tests/golden.rs index 8024a641..9b9aa22e 100644 --- a/crates/baml-bridge/tests/golden.rs +++ b/crates/baml-bridge/tests/golden.rs @@ -45,19 +45,19 @@ fn schema_snapshot_shape_hoisted() { .expect("render failed") .unwrap_or_default(); - expect![[r#"golden::GoldenShape__Circle { + expect![[r#"GoldenShape_Circle { // A circle. kind: "Circle", radius: float, } -golden::GoldenShape__Square { +GoldenShape_Square { kind: "Square", side: float, } Answer in JSON using any of these schemas: -golden::GoldenShape__Circle or golden::GoldenShape__Square"#]] +GoldenShape_Circle or GoldenShape_Square"#]] .assert_eq(&schema); } diff --git a/crates/bamltype-derive/Cargo.toml b/crates/bamltype-derive/Cargo.toml new file mode 100644 index 00000000..639e2fce --- /dev/null +++ b/crates/bamltype-derive/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "bamltype-derive" +version = "0.1.0" +edition = "2024" +license = "Apache-2.0" +publish = false +description = "Attribute macro for facet-based BAML type generation (proc-macro crate)" + +[lib] +proc-macro = true + +[dependencies] +proc-macro2 = "1.0" +quote = "1.0" +syn = { version = "2.0", features = ["full", "extra-traits"] } +proc-macro-crate = "3.2" +convert_case = "0.6" diff --git a/crates/bamltype-derive/src/lib.rs b/crates/bamltype-derive/src/lib.rs new file mode 100644 index 00000000..d83e291c --- /dev/null +++ b/crates/bamltype-derive/src/lib.rs @@ -0,0 +1,1148 @@ +use convert_case::{Case, Casing}; +use proc_macro::TokenStream; +use proc_macro_crate::{FoundCrate, crate_name}; +use proc_macro2::Span; +use quote::{format_ident, quote}; +use syn::spanned::Spanned; +use syn::{ + Attribute, Data, DeriveInput, Expr, ExprLit, Field, Fields, Lit, Meta, Path, Type, + parse_macro_input, +}; + +#[derive(Default)] +struct ContainerCompatAttrs { + rename: Option, + rename_all: Option, + tag: Option, + description: Option, + internal_name: Option, + constraints: Vec, + as_union: bool, + as_enum: bool, +} + +#[derive(Default)] +struct FieldCompatAttrs { + rename: Option, + preserve_original_name: bool, + skip: bool, + default: bool, + with_adapter: Option, + description: Option, + int_repr: Option, + map_key_repr: Option, + constraints: Vec, +} + +#[derive(Default)] +struct VariantCompatAttrs { + rename: Option, + description: Option, +} + +#[derive(Clone)] +struct ConstraintCompatAttr { + kind: ConstraintKind, + label: String, + expr: String, +} + +#[derive(Clone, Copy)] +enum ConstraintKind { + Check, + Assert, +} + +#[derive(Clone, Copy)] +enum RenameRule { + Camel, + Snake, + Pascal, + Kebab, + ScreamingSnake, + Lower, + Upper, + ScreamingKebab, +} + +impl RenameRule { + fn facet_value(self) -> Option<&'static str> { + match self { + RenameRule::Camel => Some("camelCase"), + RenameRule::Snake => Some("snake_case"), + RenameRule::Pascal => Some("PascalCase"), + RenameRule::Kebab => Some("kebab-case"), + RenameRule::ScreamingSnake => Some("SCREAMING_SNAKE_CASE"), + RenameRule::ScreamingKebab => Some("SCREAMING-KEBAB-CASE"), + RenameRule::Lower | RenameRule::Upper => None, + } + } + + fn apply(self, name: &str) -> String { + let case = match self { + RenameRule::Camel => Case::Camel, + RenameRule::Snake => Case::Snake, + RenameRule::Pascal => Case::Pascal, + RenameRule::Kebab => Case::Kebab, + RenameRule::ScreamingSnake => Case::UpperSnake, + RenameRule::Lower => Case::Lower, + RenameRule::Upper => Case::Upper, + RenameRule::ScreamingKebab => Case::UpperKebab, + }; + name.to_case(case) + } +} + +/// Attribute macro that makes a struct or enum usable with BAML. +/// +/// This macro normalizes `#[baml(...)]` (and relevant `#[serde(...)]`) attributes +/// onto the facet model, then derives `facet::Facet` and implements `BamlSchema`. +#[allow(non_snake_case)] +#[proc_macro_attribute] +pub fn BamlType(_attr: TokenStream, item: TokenStream) -> TokenStream { + let mut input = parse_macro_input!(item as DeriveInput); + let name = input.ident.clone(); + + let runtime_crate = match resolve_runtime_crate() { + Ok(path) => path, + Err(err) => return TokenStream::from(err.into_compile_error()), + }; + + let runtime_ns_ident = format_ident!("__bamltype_runtime_{}", name); + let runtime_ns_path: Path = syn::parse_quote!(#runtime_ns_ident); + + if let Err(err) = validate_input(&input) { + return TokenStream::from(err.into_compile_error()); + } + + if let Err(err) = normalize_attrs(&mut input, &runtime_crate, &runtime_ns_path) { + return TokenStream::from(err.into_compile_error()); + } + + let is_enum = matches!(input.data, Data::Enum(_)); + let has_repr = input.attrs.iter().any(|attr| attr.path().is_ident("repr")); + + // Enums need an explicit repr for facet — add #[repr(u8)] if missing. + if is_enum && !has_repr { + input.attrs.insert(0, syn::parse_quote!(#[repr(u8)])); + } + + // Put derive first so helper attrs (`#[facet(...)]`) are recognized without + // tripping legacy helper ordering lints in downstream crates. + let mut reordered_attrs = Vec::with_capacity(input.attrs.len() + 2); + reordered_attrs.push(syn::parse_quote!(#[derive(#runtime_crate::facet::Facet)])); + reordered_attrs.push(syn::parse_quote!(#[facet(crate = #runtime_crate::facet)])); + reordered_attrs.extend(std::mem::take(&mut input.attrs)); + input.attrs = reordered_attrs; + + let name = &input.ident; + let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); + + let runtime_ns_import = quote! { + #[allow(unused_imports)] + use #runtime_crate as #runtime_ns_ident; + }; + + let expanded = quote! { + #runtime_ns_import + #input + + impl #impl_generics #runtime_crate::BamlSchema for #name #ty_generics #where_clause { + fn baml_schema() -> &'static #runtime_crate::SchemaBundle { + static SCHEMA: ::std::sync::OnceLock<#runtime_crate::SchemaBundle> = + ::std::sync::OnceLock::new(); + SCHEMA.get_or_init(|| { + #runtime_crate::SchemaBundle::from_shape( + >::SHAPE + ) + }) + } + } + }; + + TokenStream::from(expanded) +} + +fn resolve_runtime_crate() -> syn::Result { + if let Some(path) = find_crate_path("bamltype") { + return Ok(path); + } + + if let Some(dspy_path) = find_crate_path("dspy-rs") { + return Ok(syn::parse_quote!(#dspy_path::bamltype)); + } + + Err(syn::Error::new( + Span::call_site(), + "could not resolve bamltype runtime crate; expected dependency on `bamltype` or `dspy-rs`", + )) +} + +fn find_crate_path(package_name: &str) -> Option { + match crate_name(package_name).ok()? { + FoundCrate::Itself => Some(syn::parse_quote!(crate)), + FoundCrate::Name(name) => { + let ident = syn::Ident::new(&name.replace('-', "_"), Span::call_site()); + Some(syn::parse_quote!(::#ident)) + } + } +} + +fn normalize_attrs( + input: &mut DeriveInput, + runtime_crate: &Path, + runtime_ns: &Path, +) -> syn::Result<()> { + let keep_serde_attrs = has_serde_derive(&input.attrs)?; + let manual_rename_all = + normalize_container_attrs(&mut input.attrs, keep_serde_attrs, runtime_ns)?; + + match &mut input.data { + Data::Struct(data) => normalize_fields( + &mut data.fields, + keep_serde_attrs, + manual_rename_all, + runtime_crate, + runtime_ns, + )?, + Data::Enum(data) => { + for variant in &mut data.variants { + let variant_name = variant.ident.to_string(); + normalize_variant_attrs( + &mut variant.attrs, + &variant_name, + keep_serde_attrs, + manual_rename_all, + )?; + normalize_fields_with_name_strategy( + &mut variant.fields, + |field, index| field_name_for_alias(field, index, &variant_name), + keep_serde_attrs, + manual_rename_all, + runtime_crate, + runtime_ns, + )?; + } + } + Data::Union(union) => { + return Err(syn::Error::new( + union.union_token.span(), + "BamlType does not support `union` items; hint: use a struct or enum instead", + )); + } + } + + Ok(()) +} + +fn normalize_fields( + fields: &mut Fields, + keep_serde_attrs: bool, + rename_all: Option, + runtime_crate: &Path, + runtime_ns: &Path, +) -> syn::Result<()> { + normalize_fields_with_name_strategy( + fields, + |field, index| field_name_for_alias(field, index, "field"), + keep_serde_attrs, + rename_all, + runtime_crate, + runtime_ns, + ) +} + +fn normalize_fields_with_name_strategy( + fields: &mut Fields, + mut name_of: F, + keep_serde_attrs: bool, + rename_all: Option, + runtime_crate: &Path, + runtime_ns: &Path, +) -> syn::Result<()> +where + F: FnMut(&Field, usize) -> String, +{ + match fields { + Fields::Named(named) => { + for (index, field) in named.named.iter_mut().enumerate() { + let original_name = name_of(field, index); + normalize_field_attrs( + &mut field.attrs, + &field.ty, + &original_name, + keep_serde_attrs, + rename_all, + runtime_crate, + runtime_ns, + )?; + } + } + Fields::Unit => {} + Fields::Unnamed(_) => { + return Err(syn::Error::new( + Span::call_site(), + "tuple fields are not supported for BamlType; use named fields", + )); + } + } + + Ok(()) +} + +fn field_name_for_alias(field: &Field, index: usize, fallback_prefix: &str) -> String { + field + .ident + .as_ref() + .map(std::string::ToString::to_string) + .unwrap_or_else(|| format!("{fallback_prefix}_{index}")) +} + +fn validate_input(input: &DeriveInput) -> syn::Result<()> { + let mut container = ContainerCompatAttrs::default(); + for attr in &input.attrs { + if attr.path().is_ident("baml") { + parse_baml_container_meta(attr, &mut container)?; + } + if attr.path().is_ident("serde") { + parse_serde_container_meta(attr, &mut container)?; + } + } + + match &input.data { + Data::Struct(data) => validate_struct(input, data), + Data::Enum(data) => validate_enum(input, data, &container), + Data::Union(union) => Err(syn::Error::new( + union.union_token.span(), + "BamlType does not support `union` items; hint: use a struct or enum instead", + )), + } +} + +fn validate_struct(input: &DeriveInput, data: &syn::DataStruct) -> syn::Result<()> { + match &data.fields { + Fields::Named(fields) => { + for field in &fields.named { + validate_field(field)?; + } + Ok(()) + } + Fields::Unit => Err(syn::Error::new_spanned( + input, + "Unit structs are not supported for BAML outputs; hint: use a named-field struct or enum", + )), + Fields::Unnamed(_) => Err(syn::Error::new_spanned( + input, + "Tuple structs are not supported for BAML outputs; hint: use a named-field struct", + )), + } +} + +fn validate_enum( + input: &DeriveInput, + data: &syn::DataEnum, + container: &ContainerCompatAttrs, +) -> syn::Result<()> { + let mut has_data_variant = false; + + for variant in &data.variants { + for attr in &variant.attrs { + if attr.path().is_ident("baml") { + let mut out = VariantCompatAttrs::default(); + parse_baml_variant_meta(attr, &mut out)?; + } + if attr.path().is_ident("serde") { + let mut out = VariantCompatAttrs::default(); + parse_serde_variant_meta(attr, &mut out)?; + } + } + + match &variant.fields { + Fields::Unit => {} + Fields::Unnamed(_) => { + return Err(syn::Error::new_spanned( + variant, + "Tuple enum variants are not supported; hint: use a unit or struct-like variant", + )); + } + Fields::Named(fields) => { + if !fields.named.is_empty() { + has_data_variant = true; + } + for field in &fields.named { + validate_field(field)?; + } + } + } + } + + if container.as_enum && has_data_variant { + return Err(syn::Error::new_spanned( + input, + "as_enum is only valid for unit enums; hint: remove #[baml(as_enum)] or convert variants to unit", + )); + } + + Ok(()) +} + +fn validate_field(field: &Field) -> syn::Result<()> { + let mut field_attrs = FieldCompatAttrs::default(); + for attr in &field.attrs { + if attr.path().is_ident("baml") { + parse_baml_field_meta(attr, &mut field_attrs)?; + } + if attr.path().is_ident("serde") { + parse_serde_field_meta(attr, &mut field_attrs)?; + } + } + + if field_attrs.with_adapter.is_some() { + return Ok(()); + } + + if let Some(ty) = find_type_match(&field.ty, &|ty| matches!(ty, Type::BareFn(_))) { + return Err(syn::Error::new_spanned( + ty, + "function types are not supported in BAML outputs; hint: remove the field or use #[baml(with = \"...\")] to adapt it", + )); + } + + if let Some(ty) = find_type_match(&field.ty, &|ty| matches!(ty, Type::Tuple(_))) { + return Err(syn::Error::new_spanned( + ty, + "tuple types are not supported in BAML outputs; hint: use a struct with named fields or a list", + )); + } + + if let Some(ty) = find_type_match(&field.ty, &|ty| matches!(ty, Type::TraitObject(_))) { + return Err(syn::Error::new_spanned( + ty, + "trait objects are not supported in BAML outputs; hint: use a concrete type or a custom adapter", + )); + } + + if let Some(ty) = find_type_match(&field.ty, &is_serde_json_value) { + return Err(syn::Error::new_spanned( + ty, + "serde_json::Value is not supported without a #[baml(with = \"...\")] adapter; hint: use a concrete type or provide a custom adapter", + )); + } + + if field_attrs.map_key_repr.is_some() && map_types_for_repr(&field.ty).is_none() { + let span_ty = map_key_repr_error_span_type(&field.ty); + return Err(syn::Error::new_spanned( + span_ty, + "map_key_repr only applies to map fields (HashMap/BTreeMap), optionally wrapped in Option/Vec/Box/Arc/Rc; hint: remove the attribute or change the field type", + )); + } + + if field_attrs.map_key_repr.is_none() + && let Some((key, _)) = map_types_for_repr(&field.ty) + && !is_string_type(key) + { + return Err(syn::Error::new_spanned( + &field.ty, + "map keys must be String for object maps; hint: use HashMap or add #[baml(map_key_repr = \"string\"|\"pairs\")], or use a custom adapter", + )); + } + + if field_attrs.int_repr.is_none() + && let Some(ty) = find_type_match(&field.ty, &is_large_int_type) + { + return Err(syn::Error::new_spanned( + ty, + "unsupported integer width for BAML outputs; hint: use #[baml(int_repr = \"string\"|\"i64\")] or a smaller integer type", + )); + } + + Ok(()) +} + +fn find_type_match<'a, F>(ty: &'a Type, predicate: &F) -> Option<&'a Type> +where + F: Fn(&Type) -> bool, +{ + if predicate(ty) { + return Some(ty); + } + + match ty { + Type::Array(array) => find_type_match(&array.elem, predicate), + Type::Group(group) => find_type_match(&group.elem, predicate), + Type::Paren(paren) => find_type_match(&paren.elem, predicate), + Type::Ptr(ptr) => find_type_match(&ptr.elem, predicate), + Type::Reference(reference) => find_type_match(&reference.elem, predicate), + Type::Slice(slice) => find_type_match(&slice.elem, predicate), + Type::Tuple(tuple) => { + for elem in &tuple.elems { + if let Some(found) = find_type_match(elem, predicate) { + return Some(found); + } + } + None + } + Type::Path(path) => { + for segment in &path.path.segments { + if let syn::PathArguments::AngleBracketed(args) = &segment.arguments { + for arg in &args.args { + if let syn::GenericArgument::Type(inner) = arg + && let Some(found) = find_type_match(inner, predicate) + { + return Some(found); + } + } + } + } + None + } + _ => None, + } +} + +fn type_ident(ty: &Type) -> Option<&syn::Ident> { + match ty { + Type::Path(path) if path.qself.is_none() => path.path.segments.last().map(|s| &s.ident), + _ => None, + } +} + +fn unwrap_repr_wrapper(ty: &Type) -> Option<&Type> { + extract_single_arg(ty, "Option") + .or_else(|| extract_single_arg(ty, "Vec")) + .or_else(|| extract_single_arg(ty, "Box")) + .or_else(|| extract_single_arg(ty, "Arc")) + .or_else(|| extract_single_arg(ty, "Rc")) +} + +fn extract_single_arg<'a>(ty: &'a Type, ident: &str) -> Option<&'a Type> { + if let Type::Path(path) = ty + && let Some(segment) = path.path.segments.last() + && segment.ident == ident + && let syn::PathArguments::AngleBracketed(args) = &segment.arguments + && let Some(syn::GenericArgument::Type(inner)) = args.args.first() + { + return Some(inner); + } + None +} + +fn map_types(ty: &Type) -> Option<(&Type, &Type)> { + if let Type::Path(path) = ty + && let Some(segment) = path.path.segments.last() + && (segment.ident == "HashMap" || segment.ident == "BTreeMap") + && let syn::PathArguments::AngleBracketed(args) = &segment.arguments + { + let mut iter = args.args.iter(); + let key = match iter.next() { + Some(syn::GenericArgument::Type(t)) => t, + _ => return None, + }; + let value = match iter.next() { + Some(syn::GenericArgument::Type(t)) => t, + _ => return None, + }; + return Some((key, value)); + } + None +} + +fn map_types_for_repr(ty: &Type) -> Option<(&Type, &Type)> { + let mut current = ty; + loop { + if let Some((key, value)) = map_types(current) { + return Some((key, value)); + } + + let inner = unwrap_repr_wrapper(current)?; + current = inner; + } +} + +fn map_key_repr_error_span_type(ty: &Type) -> &Type { + let mut current = ty; + while let Some(inner) = unwrap_repr_wrapper(current) { + current = inner; + } + current +} + +fn is_string_type(ty: &Type) -> bool { + type_ident(ty) + .map(|ident| ident == "String") + .unwrap_or(false) +} + +fn is_large_int_type(ty: &Type) -> bool { + match type_ident(ty).map(|ident| ident.to_string()) { + Some(name) => matches!(name.as_str(), "u64" | "usize" | "i128" | "u128"), + None => false, + } +} + +fn is_serde_json_value(ty: &Type) -> bool { + if let Type::Path(path) = ty + && let Some(segment) = path.path.segments.last() + && segment.ident == "Value" + { + return path + .path + .segments + .iter() + .any(|seg| seg.ident == "serde_json"); + } + false +} + +fn has_serde_derive(attrs: &[Attribute]) -> syn::Result { + for attr in attrs { + if !attr.path().is_ident("derive") { + continue; + } + + let derives = attr.parse_args_with( + syn::punctuated::Punctuated::::parse_terminated, + )?; + + for derive in derives { + if derive.is_ident("Serialize") || derive.is_ident("Deserialize") { + return Ok(true); + } + } + } + + Ok(false) +} + +fn normalize_container_attrs( + attrs: &mut Vec, + keep_serde_attrs: bool, + runtime_ns: &Path, +) -> syn::Result> { + let mut out = Vec::with_capacity(attrs.len()); + let mut compat = ContainerCompatAttrs::default(); + + for attr in std::mem::take(attrs) { + if attr.path().is_ident("baml") { + parse_baml_container_meta(&attr, &mut compat)?; + continue; + } + + if attr.path().is_ident("serde") { + parse_serde_container_meta(&attr, &mut compat)?; + if keep_serde_attrs { + out.push(attr); + } + continue; + } + + out.push(attr); + } + + if let Some(rename) = compat.rename { + let lit = syn::LitStr::new(&rename, Span::call_site()); + out.push(syn::parse_quote!(#[facet(rename = #lit)])); + } + let mut manual_rename_all = None; + if let Some(rename_all) = compat.rename_all { + if let Some(facet_value) = rename_all.facet_value() { + let lit = syn::LitStr::new(facet_value, Span::call_site()); + out.push(syn::parse_quote!(#[facet(rename_all = #lit)])); + } else { + manual_rename_all = Some(rename_all); + } + } + if compat.as_union { + out.push(syn::parse_quote!(#[facet(untagged)])); + } + if let Some(tag) = compat.tag { + let lit = syn::LitStr::new(&tag, Span::call_site()); + out.push(syn::parse_quote!(#[facet(tag = #lit)])); + } + if let Some(internal_name) = compat.internal_name { + let lit = syn::LitStr::new(&internal_name, Span::call_site()); + out.push(syn::parse_quote!(#[facet(#runtime_ns::internal_name = #lit)])); + } + push_constraint_attrs(&mut out, &compat.constraints, runtime_ns); + if let Some(description) = compat.description { + replace_doc_attrs(&mut out, &description); + } + + *attrs = out; + Ok(manual_rename_all) +} + +fn normalize_field_attrs( + attrs: &mut Vec, + field_ty: &Type, + original_name: &str, + keep_serde_attrs: bool, + rename_all: Option, + runtime_crate: &Path, + runtime_ns: &Path, +) -> syn::Result<()> { + let mut out = Vec::with_capacity(attrs.len()); + let mut compat = FieldCompatAttrs::default(); + + for attr in std::mem::take(attrs) { + if attr.path().is_ident("baml") { + parse_baml_field_meta(&attr, &mut compat)?; + continue; + } + + if attr.path().is_ident("serde") { + parse_serde_field_meta(&attr, &mut compat)?; + if keep_serde_attrs { + out.push(attr); + } + continue; + } + + out.push(attr); + } + + apply_rename_all(&mut compat.rename, rename_all, original_name); + + if let Some(rename) = compat.rename { + let lit = syn::LitStr::new(&rename, Span::call_site()); + out.push(syn::parse_quote!(#[facet(rename = #lit)])); + + // Old BAML alias behavior accepted both the alias and the original field name. + if compat.preserve_original_name && rename != original_name { + let original = syn::LitStr::new(original_name, Span::call_site()); + out.push(syn::parse_quote!(#[facet(alias = #original)])); + } + } + + if compat.skip { + out.push(syn::parse_quote!(#[facet(skip)])); + } + // Keep baml-bridge compatibility: skipped fields deserialize from Default. + if compat.default || compat.skip { + out.push(syn::parse_quote!(#[facet(default)])); + } + if let Some(adapter) = compat.with_adapter { + let with_expr = quote! { + &#runtime_crate::facet_ext::WithAdapterFns { + type_ir: || <#adapter as #runtime_crate::compat::BamlAdapter<#field_ty>>::type_ir(), + register: |reg| <#adapter as #runtime_crate::compat::BamlAdapter<#field_ty>>::register(reg), + apply: |partial, value, path| { + let converted = <#adapter as #runtime_crate::compat::BamlAdapter<#field_ty>>::try_from_baml(value, path)?; + partial.set(converted).map_err(|err| { + #runtime_crate::compat::BamlConvertError::new( + ::std::vec::Vec::new(), + "compatible type", + err.to_string(), + err.to_string(), + ) + }) + }, + } + }; + out.push(syn::parse_quote!(#[facet(#runtime_ns::with = #with_expr)])); + } + if let Some(int_repr) = compat.int_repr { + let lit = syn::LitStr::new(&int_repr, Span::call_site()); + out.push(syn::parse_quote!(#[facet(#runtime_ns::int_repr = #lit)])); + } + if let Some(map_key_repr) = compat.map_key_repr { + let lit = syn::LitStr::new(&map_key_repr, Span::call_site()); + out.push(syn::parse_quote!(#[facet(#runtime_ns::map_key_repr = #lit)])); + } + push_constraint_attrs(&mut out, &compat.constraints, runtime_ns); + if let Some(description) = compat.description { + replace_doc_attrs(&mut out, &description); + } + + *attrs = out; + Ok(()) +} + +fn normalize_variant_attrs( + attrs: &mut Vec, + variant_name: &str, + keep_serde_attrs: bool, + rename_all: Option, +) -> syn::Result<()> { + let mut out = Vec::with_capacity(attrs.len()); + let mut compat = VariantCompatAttrs::default(); + + for attr in std::mem::take(attrs) { + if attr.path().is_ident("baml") { + parse_baml_variant_meta(&attr, &mut compat)?; + continue; + } + + if attr.path().is_ident("serde") { + parse_serde_variant_meta(&attr, &mut compat)?; + if keep_serde_attrs { + out.push(attr); + } + continue; + } + + out.push(attr); + } + + apply_rename_all(&mut compat.rename, rename_all, variant_name); + + if let Some(rename) = compat.rename { + let lit = syn::LitStr::new(&rename, Span::call_site()); + out.push(syn::parse_quote!(#[facet(rename = #lit)])); + } + if let Some(description) = compat.description { + replace_doc_attrs(&mut out, &description); + } + + *attrs = out; + Ok(()) +} + +fn apply_rename_all( + rename: &mut Option, + rename_all: Option, + original_name: &str, +) { + if rename.is_none() + && let Some(rule) = rename_all + { + *rename = Some(rule.apply(original_name)); + } +} + +fn push_constraint_attrs( + out: &mut Vec, + constraints: &[ConstraintCompatAttr], + runtime_ns: &Path, +) { + for constraint in constraints { + let label = syn::LitStr::new(&constraint.label, Span::call_site()); + let expr = syn::LitStr::new(&constraint.expr, Span::call_site()); + match constraint.kind { + ConstraintKind::Check => { + out.push( + syn::parse_quote!(#[facet(#runtime_ns::check(label = #label, expr = #expr))]), + ); + } + ConstraintKind::Assert => { + out.push( + syn::parse_quote!(#[facet(#runtime_ns::assert(label = #label, expr = #expr))]), + ); + } + } + } +} + +const UNSUPPORTED_BAML_ATTR_HINT: &str = + "unsupported #[baml(...)] attribute; hint: check the supported keys in the bridge docs"; + +fn parse_baml_container_meta(attr: &Attribute, out: &mut ContainerCompatAttrs) -> syn::Result<()> { + for meta in parse_meta_list(attr)? { + match meta { + Meta::NameValue(meta) if meta.path.is_ident("name") => { + out.rename = Some(parse_string_expr(&meta.value, meta.span())?); + } + Meta::NameValue(meta) if meta.path.is_ident("rename_all") => { + out.rename_all = Some(parse_rename_rule(&meta.value, meta.span())?); + } + Meta::NameValue(meta) if meta.path.is_ident("tag") => { + out.tag = Some(parse_string_expr(&meta.value, meta.span())?); + } + Meta::NameValue(meta) if meta.path.is_ident("description") => { + out.description = parse_optional_string(&meta.value, meta.span())?; + } + + // Accepted for source compatibility (runtime behavior parity is handled + // in bamltype runtime code paths). + Meta::NameValue(meta) if meta.path.is_ident("internal_name") => { + out.internal_name = Some(parse_string_expr(&meta.value, meta.span())?); + } + Meta::List(meta) if meta.path.is_ident("check") => { + out.constraints + .push(parse_constraint_meta(&meta, ConstraintKind::Check)?); + } + Meta::List(meta) if meta.path.is_ident("assert") => { + out.constraints + .push(parse_constraint_meta(&meta, ConstraintKind::Assert)?); + } + Meta::Path(path) if path.is_ident("as_union") => { + out.as_union = true; + } + Meta::Path(path) if path.is_ident("as_enum") => { + out.as_enum = true; + } + + _ => { + return Err(syn::Error::new_spanned(meta, UNSUPPORTED_BAML_ATTR_HINT)); + } + } + } + + Ok(()) +} + +fn parse_baml_field_meta(attr: &Attribute, out: &mut FieldCompatAttrs) -> syn::Result<()> { + for meta in parse_meta_list(attr)? { + match meta { + Meta::NameValue(meta) if meta.path.is_ident("alias") => { + out.rename = Some(parse_string_expr(&meta.value, meta.span())?); + out.preserve_original_name = true; + } + Meta::NameValue(meta) if meta.path.is_ident("description") => { + out.description = parse_optional_string(&meta.value, meta.span())?; + } + Meta::Path(path) if path.is_ident("skip") => { + out.skip = true; + } + Meta::Path(path) if path.is_ident("default") => { + out.default = true; + } + + Meta::NameValue(meta) if meta.path.is_ident("with") => { + let path_str = parse_string_expr(&meta.value, meta.span())?; + out.with_adapter = Some(syn::parse_str::(&path_str)?); + } + Meta::NameValue(meta) if meta.path.is_ident("int_repr") => { + out.int_repr = Some(parse_int_repr(&meta.value, meta.span())?); + } + Meta::NameValue(meta) if meta.path.is_ident("map_key_repr") => { + out.map_key_repr = Some(parse_map_key_repr(&meta.value, meta.span())?); + } + Meta::List(meta) if meta.path.is_ident("check") => { + out.constraints + .push(parse_constraint_meta(&meta, ConstraintKind::Check)?); + } + Meta::List(meta) if meta.path.is_ident("assert") => { + out.constraints + .push(parse_constraint_meta(&meta, ConstraintKind::Assert)?); + } + + _ => { + return Err(syn::Error::new_spanned(meta, UNSUPPORTED_BAML_ATTR_HINT)); + } + } + } + + Ok(()) +} + +fn parse_baml_variant_meta(attr: &Attribute, out: &mut VariantCompatAttrs) -> syn::Result<()> { + for meta in parse_meta_list(attr)? { + match meta { + Meta::NameValue(meta) if meta.path.is_ident("alias") => { + out.rename = Some(parse_string_expr(&meta.value, meta.span())?); + } + Meta::NameValue(meta) if meta.path.is_ident("description") => { + out.description = parse_optional_string(&meta.value, meta.span())?; + } + _ => { + return Err(syn::Error::new_spanned(meta, UNSUPPORTED_BAML_ATTR_HINT)); + } + } + } + Ok(()) +} + +fn parse_serde_container_meta(attr: &Attribute, out: &mut ContainerCompatAttrs) -> syn::Result<()> { + for meta in parse_meta_list(attr)? { + match meta { + Meta::NameValue(meta) if meta.path.is_ident("rename") => { + if out.rename.is_none() { + out.rename = Some(parse_string_expr(&meta.value, meta.span())?); + } + } + Meta::NameValue(meta) if meta.path.is_ident("rename_all") => { + if out.rename_all.is_none() { + out.rename_all = Some(parse_rename_rule(&meta.value, meta.span())?); + } + } + Meta::NameValue(meta) if meta.path.is_ident("tag") => { + if out.tag.is_none() { + out.tag = Some(parse_string_expr(&meta.value, meta.span())?); + } + } + Meta::Path(path) if path.is_ident("untagged") => { + return Err(syn::Error::new_spanned( + path, + "serde(untagged) is not supported; hint: use #[baml(tag = \"...\")] for data enums", + )); + } + Meta::Path(path) if path.is_ident("flatten") => { + return Err(syn::Error::new_spanned( + path, + "serde(flatten) is not supported; hint: model fields explicitly", + )); + } + _ => {} + } + } + Ok(()) +} + +fn parse_serde_field_meta(attr: &Attribute, out: &mut FieldCompatAttrs) -> syn::Result<()> { + for meta in parse_meta_list(attr)? { + match meta { + Meta::NameValue(meta) if meta.path.is_ident("rename") => { + if out.rename.is_none() { + out.rename = Some(parse_string_expr(&meta.value, meta.span())?); + out.preserve_original_name = true; + } + } + Meta::Path(path) if path.is_ident("skip") => { + out.skip = true; + } + Meta::Path(path) if path.is_ident("default") => { + out.default = true; + } + Meta::NameValue(meta) if meta.path.is_ident("default") => { + return Err(syn::Error::new_spanned( + meta, + "serde(default = \"path\") is not supported; hint: use #[baml(default)] or Default::default", + )); + } + Meta::Path(path) if path.is_ident("flatten") => { + return Err(syn::Error::new_spanned( + path, + "serde(flatten) is not supported; hint: model fields explicitly", + )); + } + _ => {} + } + } + Ok(()) +} + +fn parse_serde_variant_meta(attr: &Attribute, out: &mut VariantCompatAttrs) -> syn::Result<()> { + for meta in parse_meta_list(attr)? { + match meta { + Meta::NameValue(meta) if meta.path.is_ident("rename") => { + if out.rename.is_none() { + out.rename = Some(parse_string_expr(&meta.value, meta.span())?); + } + } + Meta::Path(path) if path.is_ident("skip") => { + return Err(syn::Error::new_spanned( + path, + "serde(skip) is not supported on enum variants; hint: remove the variant or use a separate enum", + )); + } + _ => {} + } + } + Ok(()) +} + +fn parse_meta_list(attr: &Attribute) -> syn::Result> { + let metas = attr + .parse_args_with(syn::punctuated::Punctuated::::parse_terminated)?; + Ok(metas.into_iter().collect()) +} + +fn parse_string_expr(expr: &Expr, span: Span) -> syn::Result { + match expr { + Expr::Lit(ExprLit { + lit: Lit::Str(value), + .. + }) => Ok(value.value()), + _ => Err(syn::Error::new( + span, + "expected string literal; hint: wrap the value in quotes", + )), + } +} + +fn parse_optional_string(expr: &Expr, span: Span) -> syn::Result> { + if let Expr::Path(path) = expr + && path.path.is_ident("None") + { + return Ok(None); + } + Ok(Some(parse_string_expr(expr, span)?)) +} + +fn parse_rename_rule(expr: &Expr, span: Span) -> syn::Result { + let value = parse_string_expr(expr, span)?; + match value.as_str() { + "camelCase" => Ok(RenameRule::Camel), + "snake_case" => Ok(RenameRule::Snake), + "PascalCase" => Ok(RenameRule::Pascal), + "kebab-case" => Ok(RenameRule::Kebab), + "SCREAMING_SNAKE_CASE" => Ok(RenameRule::ScreamingSnake), + "lowercase" => Ok(RenameRule::Lower), + "UPPERCASE" => Ok(RenameRule::Upper), + "SCREAMING-KEBAB-CASE" => Ok(RenameRule::ScreamingKebab), + _ => Err(syn::Error::new(span, "unsupported rename_all value")), + } +} + +fn parse_int_repr(expr: &Expr, span: Span) -> syn::Result { + let value = parse_string_expr(expr, span)?; + match value.as_str() { + "string" | "i64" => Ok(value), + _ => Err(syn::Error::new( + span, + "int_repr must be \"string\" or \"i64\"", + )), + } +} + +fn parse_map_key_repr(expr: &Expr, span: Span) -> syn::Result { + let value = parse_string_expr(expr, span)?; + match value.as_str() { + "string" | "pairs" => Ok(value), + _ => Err(syn::Error::new( + span, + "map_key_repr must be \"string\" or \"pairs\"", + )), + } +} + +fn parse_constraint_meta( + meta: &syn::MetaList, + kind: ConstraintKind, +) -> syn::Result { + let nested = meta + .parse_args_with(syn::punctuated::Punctuated::::parse_terminated)?; + let mut label = None; + let mut expr = None; + + for item in nested { + match item { + Meta::NameValue(m) if m.path.is_ident("label") => { + label = Some(parse_string_expr(&m.value, m.span())?); + } + Meta::NameValue(m) if m.path.is_ident("expr") => { + expr = Some(parse_string_expr(&m.value, m.span())?); + } + other => { + return Err(syn::Error::new_spanned( + other, + "unsupported constraint attribute; hint: use #[baml(check(...))] or #[baml(assert(...))]", + )); + } + } + } + + let Some(label) = label else { + return Err(syn::Error::new( + meta.span(), + "constraint missing label; hint: use label = \"...\"", + )); + }; + let Some(expr) = expr else { + return Err(syn::Error::new( + meta.span(), + "constraint missing expr; hint: use expr = \"...\"", + )); + }; + + Ok(ConstraintCompatAttr { kind, label, expr }) +} + +fn replace_doc_attrs(attrs: &mut Vec, description: &str) { + attrs.retain(|attr| !attr.path().is_ident("doc")); + for line in description.lines() { + let lit = syn::LitStr::new(line, Span::call_site()); + attrs.push(syn::parse_quote!(#[doc = #lit])); + } +} diff --git a/crates/bamltype/Cargo.toml b/crates/bamltype/Cargo.toml new file mode 100644 index 00000000..255ff736 --- /dev/null +++ b/crates/bamltype/Cargo.toml @@ -0,0 +1,50 @@ +[package] +name = "bamltype" +version = "0.1.0" +edition = "2024" +license = "Apache-2.0" +publish = false +description = "Facet-based BAML type generation" + +[dependencies] +# Facet for reflection +facet = { version = "0.43.2", default-features = false, features = ["std", "doc"] } +facet-reflect = { version = "0.43.2", default-features = false, features = ["std"] } + +# BAML crates for schema/parsing +anyhow = "1.0" +baml-types = { path = "../../vendor/baml/crates/baml-types" } +internal-baml-jinja = { path = "../../vendor/baml/crates/internal-baml-jinja" } +jsonish = { path = "../../vendor/baml/crates/jsonish" } +minijinja = { git = "https://github.com/boundaryml/minijinja.git", branch = "main", default-features = false, features = [ + "macros", + "builtins", + "debug", + "preserve_order", + "adjacent_loop_items", + "unicode", + "json", + "unstable_machinery", + "unstable_machinery_serde", + "custom_syntax", + "internal_debug", + "deserialization", + "serde", +] } + +# Attribute macro (re-exported) +bamltype-derive = { path = "../bamltype-derive" } + +# Utils +indexmap = "2.1" +sha2 = "0.10.9" +thiserror = "2.0" + +[features] +default = ["derive"] +derive = [] + +[dev-dependencies] +baml-bridge = { path = "../baml-bridge" } +trybuild = "1.0.96" +serde_json = { version = "1", features = ["float_roundtrip", "preserve_order"] } diff --git a/crates/bamltype/src/compat.rs b/crates/bamltype/src/compat.rs new file mode 100644 index 00000000..d229113c --- /dev/null +++ b/crates/bamltype/src/compat.rs @@ -0,0 +1,385 @@ +//! Compatibility layer providing baml-bridge-compatible traits and helpers. + +use std::collections::{HashMap, HashSet}; + +use baml_types::{BamlMap, BamlValue, Constraint, StreamingMode, TypeIR, type_meta}; +use facet::Facet; +use indexmap::{IndexMap, IndexSet}; +use internal_baml_jinja::types::{Class, Enum, OutputFormatContent}; + +use crate::BamlSchema; +use crate::convert; +use crate::schema_builder::build_type_ir_from_shape; + +/// Error during BamlValue ↔ Rust conversion (baml-bridge compatible). +#[derive(Debug, Clone)] +pub struct BamlConvertError { + pub path: Vec, + pub expected: &'static str, + pub got: String, + pub message: String, +} + +impl BamlConvertError { + pub fn new( + path: Vec, + expected: &'static str, + got: impl Into, + message: impl Into, + ) -> Self { + Self { + path, + expected, + got: got.into(), + message: message.into(), + } + } + + pub fn with_path(mut self, segment: impl Into) -> Self { + self.path.push(segment.into()); + self + } + + pub fn path_string(&self) -> String { + if self.path.is_empty() { + "".to_string() + } else { + self.path.join(".") + } + } +} + +impl std::fmt::Display for BamlConvertError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!( + f, + "{} (expected {}, got {}) at {}", + self.message, + self.expected, + self.got, + self.path_string() + ) + } +} + +impl std::error::Error for BamlConvertError {} + +impl From for BamlConvertError { + fn from(err: convert::ConvertError) -> Self { + match err { + convert::ConvertError::Adapter(inner) => inner, + other => Self { + path: Vec::new(), + expected: "compatible type", + got: other.to_string(), + message: other.to_string(), + }, + } + } +} + +/// Registry for schema elements (baml-bridge compatible). +#[derive(Debug, Default)] +pub struct Registry { + enums: IndexMap, + classes: IndexMap<(String, StreamingMode), Class>, + class_deps: IndexMap>, + structural_recursive_aliases: IndexMap, + registered: HashSet, +} + +impl Registry { + pub fn new() -> Self { + Self::default() + } + + pub fn mark_type(&mut self, name: &str) -> bool { + if self.registered.contains(name) { + return false; + } + self.registered.insert(name.to_string()); + true + } + + pub fn register_enum(&mut self, r#enum: Enum) { + let name = r#enum.name.real_name().to_string(); + self.enums.entry(name).or_insert(r#enum); + } + + pub fn register_class(&mut self, class: Class) { + let name = class.name.real_name().to_string(); + let mode = class.namespace; + let key = (name.clone(), mode); + let entry = self.classes.entry(key).or_insert(class); + + let deps = self.class_deps.entry(name).or_default(); + for (_, field_type, _, _) in &entry.fields { + collect_class_refs(field_type, deps); + } + } + + pub fn register_structural_alias(&mut self, name: String, alias: TypeIR) { + self.structural_recursive_aliases.insert(name, alias); + } + + pub fn build(self, target: TypeIR) -> OutputFormatContent { + let recursive_classes = compute_recursive_classes(&self.class_deps); + + let mut enums = self.enums.into_iter().collect::>(); + enums.sort_by(|(a, _), (b, _)| a.cmp(b)); + let enums = enums.into_iter().map(|(_, v)| v).collect::>(); + + let mut classes = self.classes.into_iter().collect::>(); + classes.sort_by(|(a, _), (b, _)| { + let (a_name, a_mode) = a; + let (b_name, b_mode) = b; + match a_name.cmp(b_name) { + std::cmp::Ordering::Equal => mode_rank(*a_mode).cmp(&mode_rank(*b_mode)), + other => other, + } + }); + let classes = classes.into_iter().map(|(_, v)| v).collect::>(); + + OutputFormatContent::target(target) + .enums(enums) + .classes(classes) + .recursive_classes(recursive_classes) + .structural_recursive_aliases(self.structural_recursive_aliases) + .build() + } +} + +/// Internal type metadata (baml-bridge compatible). +pub trait BamlTypeInternal { + fn baml_internal_name() -> &'static str; + fn baml_type_ir() -> TypeIR; + fn register(_reg: &mut Registry) {} +} + +/// Convert from BamlValue to Rust (baml-bridge compatible). +pub trait BamlValueConvert: Sized { + fn try_from_baml_value(value: BamlValue, path: Vec) -> Result; +} + +/// Convert from Rust to BamlValue (baml-bridge compatible). +pub trait ToBamlValue { + fn to_baml_value(&self) -> BamlValue; +} + +/// Adapter for custom field-level conversion / schema representation. +pub trait BamlAdapter { + fn type_ir() -> TypeIR; + fn register(_reg: &mut Registry) {} + fn try_from_baml(value: BamlValue, path: Vec) -> Result; +} + +/// Full BamlType trait (baml-bridge compatible). +/// +/// Named `BamlTypeTrait` to avoid collision with the `#[BamlType]` attribute macro. +pub trait BamlTypeTrait: BamlTypeInternal + BamlValueConvert + Sized + 'static { + fn baml_output_format() -> &'static OutputFormatContent; + + fn baml_internal_name() -> &'static str { + ::baml_internal_name() + } + + fn baml_type_ir() -> TypeIR { + ::baml_type_ir() + } +} + +impl> BamlTypeInternal for T { + fn baml_internal_name() -> &'static str { + for attr in T::SHAPE.attributes { + if attr.ns != Some("bamltype") || attr.key != "internal_name" { + continue; + } + + if let Some(name) = attr.get_as::<&'static str>() { + return name; + } + } + + if let Some(name) = T::SHAPE.get_builtin_attr_value::<&'static str>("internal_name") { + name + } else { + std::any::type_name::() + } + } + + fn baml_type_ir() -> TypeIR { + build_type_ir_from_shape(T::SHAPE) + } +} + +impl> BamlValueConvert for T { + fn try_from_baml_value(value: BamlValue, _path: Vec) -> Result { + convert::from_baml_value(value).map_err(BamlConvertError::from) + } +} + +impl> ToBamlValue for T { + fn to_baml_value(&self) -> BamlValue { + convert::to_baml_value(self).unwrap_or(BamlValue::Null) + } +} + +impl BamlTypeTrait for T { + fn baml_output_format() -> &'static OutputFormatContent { + &T::baml_schema().output_format + } +} + +/// Add constraints to a TypeIR (baml-bridge compatible). +pub fn with_constraints(mut type_ir: TypeIR, constraints: Vec) -> TypeIR { + type_ir.meta_mut().constraints.extend(constraints); + type_ir +} + +/// Default streaming behavior helper (baml-bridge compatible). +pub fn default_streaming_behavior() -> type_meta::base::StreamingBehavior { + type_meta::base::StreamingBehavior::default() +} + +/// Lookup helper matching baml-bridge semantics (`name` then optional alias). +pub fn get_field<'a>( + map: &'a BamlMap, + name: &str, + alias: Option<&str>, +) -> Option<&'a BamlValue> { + map.get(name) + .or_else(|| alias.and_then(|alias| map.get(alias))) +} + +fn collect_class_refs(field_type: &TypeIR, deps: &mut IndexSet) { + match field_type { + TypeIR::Class { name, .. } => { + deps.insert(name.clone()); + } + TypeIR::RecursiveTypeAlias { name, .. } => { + deps.insert(name.clone()); + } + TypeIR::List(inner, _) => collect_class_refs(inner, deps), + TypeIR::Map(key, value, _) => { + collect_class_refs(key, deps); + collect_class_refs(value, deps); + } + TypeIR::Union(union, _) => { + for item in union.iter_include_null() { + collect_class_refs(item, deps); + } + } + TypeIR::Tuple(items, _) => { + for item in items { + collect_class_refs(item, deps); + } + } + TypeIR::Arrow(arrow, _) => { + for param in &arrow.param_types { + collect_class_refs(param, deps); + } + collect_class_refs(&arrow.return_type, deps); + } + TypeIR::Primitive(..) | TypeIR::Enum { .. } | TypeIR::Literal(..) | TypeIR::Top(..) => {} + } +} + +fn compute_recursive_classes(class_deps: &IndexMap>) -> IndexSet { + struct Tarjan<'a> { + next_index: usize, + indices: HashMap, + lowlink: HashMap, + stack: Vec, + on_stack: HashSet, + deps: &'a IndexMap>, + recursive: HashSet, + } + + impl<'a> Tarjan<'a> { + fn new(deps: &'a IndexMap>) -> Self { + Self { + next_index: 0, + indices: HashMap::new(), + lowlink: HashMap::new(), + stack: Vec::new(), + on_stack: HashSet::new(), + deps, + recursive: HashSet::new(), + } + } + + fn strongconnect(&mut self, node: &str) { + let node_key = node.to_string(); + self.indices.insert(node_key.clone(), self.next_index); + self.lowlink.insert(node_key.clone(), self.next_index); + self.next_index += 1; + self.stack.push(node_key.clone()); + self.on_stack.insert(node_key.clone()); + + let child_names = self + .deps + .get(node) + .map(|children| children.iter().cloned().collect::>()) + .unwrap_or_default(); + for child in child_names { + if !self.indices.contains_key(&child) { + self.strongconnect(&child); + let lowlink_child = self.lowlink.get(&child).copied().unwrap_or(0); + let lowlink_node = self.lowlink.get(&node_key).copied().unwrap_or(0); + self.lowlink + .insert(node_key.clone(), lowlink_node.min(lowlink_child)); + } else if self.on_stack.contains(&child) { + let index_child = self.indices.get(&child).copied().unwrap_or(0); + let lowlink_node = self.lowlink.get(&node_key).copied().unwrap_or(0); + self.lowlink + .insert(node_key.clone(), lowlink_node.min(index_child)); + } + } + + let node_index = self.indices.get(&node_key).copied().unwrap_or(0); + let node_lowlink = self.lowlink.get(&node_key).copied().unwrap_or(0); + if node_lowlink == node_index { + let mut scc = Vec::new(); + while let Some(w) = self.stack.pop() { + self.on_stack.remove(&w); + scc.push(w.clone()); + if w == node_key { + break; + } + } + + if scc.len() > 1 { + for name in scc { + self.recursive.insert(name); + } + } else if let Some(name) = scc.first() + && self + .deps + .get(name) + .map(|edges| edges.contains(name)) + .unwrap_or(false) + { + self.recursive.insert(name.clone()); + } + } + } + } + + let mut tarjan = Tarjan::new(class_deps); + for node in class_deps.keys() { + if !tarjan.indices.contains_key(node) { + tarjan.strongconnect(node); + } + } + + let mut sorted = tarjan.recursive.into_iter().collect::>(); + sorted.sort(); + IndexSet::from_iter(sorted) +} + +fn mode_rank(mode: StreamingMode) -> u8 { + match mode { + StreamingMode::NonStreaming => 0, + StreamingMode::Streaming => 1, + } +} diff --git a/crates/bamltype/src/convert.rs b/crates/bamltype/src/convert.rs new file mode 100644 index 00000000..8e2a18f0 --- /dev/null +++ b/crates/bamltype/src/convert.rs @@ -0,0 +1,871 @@ +//! BamlValue ↔ Rust conversion using direct facet reflection +//! +//! This module provides bidirectional conversion between BamlValue and any +//! Rust type that implements Facet, using facet's Partial (for building) +//! and Peek (for reading) APIs. + +use baml_types::{BamlMap, BamlValue}; +use facet::{Def, Facet, Shape, Type, UserType}; +use facet_reflect::{HasFields, HeapValue, Partial, Peek, ReflectError, ScalarType, VariantError}; +use indexmap::IndexMap; + +use crate::BamlValueWithFlags; +use crate::compat::BamlConvertError; +use crate::schema_builder::internal_name_for_shape; + +/// Error during BamlValue conversion. +#[derive(Debug, thiserror::Error)] +pub enum ConvertError { + /// Error from facet reflection operations. + #[error("Reflection error: {0}")] + Reflect(#[from] ReflectError), + + /// Error from enum variant access. + #[error("Variant error: {0}")] + Variant(#[from] VariantError), + + /// Type mismatch during conversion. + #[error("Type mismatch: expected {expected}, got {actual}")] + TypeMismatch { + expected: &'static str, + actual: String, + }, + + /// Missing required field. + #[error("Missing field: {0}")] + MissingField(String), + + /// Unknown enum variant. + #[error("Unknown variant: {0}")] + UnknownVariant(String), + + /// Unsupported type for conversion. + #[error("Unsupported type: {0}")] + Unsupported(String), + + /// Adapter conversion error. + #[error("{0}")] + Adapter(#[from] BamlConvertError), +} + +impl ConvertError { + fn with_path_prefix(self, segment: impl Into) -> Self { + let segment = segment.into(); + match self { + ConvertError::Adapter(mut inner) => { + inner.path.insert(0, segment); + ConvertError::Adapter(inner) + } + other => { + let message = other.to_string(); + let mut inner = + BamlConvertError::new(Vec::new(), "compatible type", message.clone(), message); + inner.path.insert(0, segment); + ConvertError::Adapter(inner) + } + } + } +} + +#[derive(Debug, Clone, Copy, Default)] +struct ReprHints { + int_repr: Option, + map_key_repr: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum IntReprHint { + String, + I64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum MapKeyReprHint { + String, + Pairs, +} + +// ============================================================================ +// BamlValue → Rust (using Partial API) +// ============================================================================ + +/// Convert a BamlValue to a Rust type using facet reflection. +pub fn from_baml_value>(value: BamlValue) -> Result { + let partial = Partial::alloc::()?; + let partial = build_from_baml_value(partial, &value)?; + let heap_value: HeapValue<'static> = partial.build()?; + Ok(heap_value.materialize::()?) +} + +/// Convert a BamlValueWithFlags to a Rust type. +/// +/// This is the primary entry point for converting parsed LLM output to Rust types. +pub fn from_baml_value_with_flags>( + value: &BamlValueWithFlags, +) -> Result { + let baml_value: BamlValue = value.clone().into(); + from_baml_value(baml_value) +} + +fn build_from_baml_value( + partial: Partial<'static>, + value: &BamlValue, +) -> Result, ConvertError> { + build_from_baml_value_with_hints(partial, value, ReprHints::default()) +} + +/// Recursive helper to build a Partial from BamlValue. +fn build_from_baml_value_with_hints( + partial: Partial<'static>, + value: &BamlValue, + hints: ReprHints, +) -> Result, ConvertError> { + let target_shape = partial.shape(); + + // Smart pointers (Box, Arc, Rc) - enter the pointer, build inner, exit. + if let Def::Pointer(_) = target_shape.def { + let p = partial.begin_smart_ptr()?; + let p = build_from_baml_value_with_hints(p, value, hints)?; + return Ok(p.end()?); + } + + // Option - wrap non-null values in Some. + if let Def::Option(_) = target_shape.def { + if matches!(value, BamlValue::Null) { + return Ok(partial.set_default()?); + } + let p = partial.begin_some()?; + let p = build_from_baml_value_with_hints(p, value, hints)?; + return Ok(p.end()?); + } + + match value { + // Strings map to either textual primitives or parseable scalar types. + BamlValue::String(s) => { + if matches!(partial.shape().ty, Type::User(UserType::Enum(_))) { + return select_enum_variant(partial, s); + } + if is_string_target_shape(partial.shape()) { + Ok(partial.set(s.clone())?) + } else if hints.int_repr == Some(IntReprHint::String) { + Ok(partial.parse_from_str(s)?) + } else { + Err(ConvertError::Adapter(BamlConvertError::new( + Vec::new(), + expected_kind_for_shape(partial.shape()), + format!("{value:?}"), + format!("expected a {}", expected_kind_for_shape(partial.shape())), + ))) + } + } + BamlValue::Int(i) => { + if let Some((expected, min, max)) = integer_expected_and_bounds(partial.shape()) { + let value = *i as i128; + if value < min || value > max { + return Err(ConvertError::Adapter(BamlConvertError::new( + Vec::new(), + expected, + i.to_string(), + "integer out of range", + ))); + } + } + Ok(partial.parse_from_str(&i.to_string())?) + } + BamlValue::Float(f) => Ok(partial.parse_from_str(&f.to_string())?), + BamlValue::Bool(b) => Ok(partial.set(*b)?), + BamlValue::Null => Ok(partial.set_default()?), + + // Class input: either enum object form, struct object form, or map object form. + BamlValue::Class(_type_name, fields) => { + if matches!(partial.shape().ty, Type::User(UserType::Enum(_))) { + return build_enum_from_object(partial, fields); + } + + match partial.shape().def { + Def::Map(_) => build_map_from_object_pairs(partial, fields), + _ => build_struct_from_object(partial, fields), + } + } + + // List: either list-like values or map-entry pairs. + BamlValue::List(items) => { + if matches!(partial.shape().def, Def::Map(_)) { + if hints.map_key_repr == Some(MapKeyReprHint::Pairs) { + return build_map_from_pair_entries(partial, items); + } + return Err(ConvertError::Adapter(BamlConvertError::new( + Vec::new(), + "map", + format!("{value:?}"), + "expected a map", + ))); + } + + let mut p = partial.init_list()?; + for (idx, item) in items.iter().enumerate() { + p = p.begin_list_item()?; + p = build_from_baml_value_with_hints(p, item, hints) + .map_err(|err| err.with_path_prefix(idx.to_string()))?; + p = p.end()?; + } + Ok(p) + } + + // Object map. + BamlValue::Map(map) => { + if matches!(partial.shape().ty, Type::User(UserType::Enum(_))) { + return build_enum_from_object(partial, map); + } + + if matches!(partial.shape().ty, Type::User(UserType::Struct(_))) { + return build_struct_from_object(partial, map); + } + + match hints.map_key_repr { + Some(MapKeyReprHint::Pairs) => Err(ConvertError::TypeMismatch { + expected: "list", + actual: "map".to_string(), + }), + Some(MapKeyReprHint::String) | None => build_map_from_object_pairs(partial, map), + } + } + + // Enum variant (unit-like representation). + BamlValue::Enum(_type_name, variant_name) => select_enum_variant(partial, variant_name), + + // Media - not yet supported. + BamlValue::Media(_media) => Err(ConvertError::Unsupported( + "Media type conversion not yet implemented".into(), + )), + } +} + +fn build_struct_from_object( + partial: Partial<'static>, + fields: &BamlMap, +) -> Result, ConvertError> { + build_object_fields(partial, fields, None, true) +} + +fn build_enum_from_object( + partial: Partial<'static>, + fields: &BamlMap, +) -> Result, ConvertError> { + let tag_name = partial.shape().get_tag_attr().unwrap_or("type"); + + let tag_value = fields + .get(tag_name) + .ok_or_else(|| ConvertError::MissingField(tag_name.to_string()))?; + + let variant_name = match tag_value { + BamlValue::String(v) => v.as_str(), + BamlValue::Enum(_, v) => v.as_str(), + other => { + return Err(ConvertError::TypeMismatch { + expected: "string", + actual: baml_value_kind(other), + }); + } + }; + + let p = select_enum_variant(partial, variant_name)?; + build_object_fields(p, fields, Some(tag_name), false) +} + +fn build_object_fields( + mut partial: Partial<'static>, + fields: &BamlMap, + skip_field_name: Option<&str>, + skip_struct_deserialize_skips: bool, +) -> Result, ConvertError> { + for (field_name, field_value) in fields { + if skip_field_name == Some(field_name.as_str()) { + continue; + } + + if skip_struct_deserialize_skips + && should_skip_struct_field_deserializing(partial.shape(), field_name) + { + continue; + } + + // Unknown fields are ignored for compatibility. + let Some(index) = resolve_field_index(&partial, field_name) else { + continue; + }; + + let field = current_field(&partial, index); + if let Some(field) = field + && let Some(with) = crate::facet_ext::with_adapter_fns(field.attributes) + { + let field_path = vec![field_name.to_string()]; + partial = partial.begin_nth_field(index)?; + partial = (with.apply)(partial, field_value.clone(), field_path) + .map_err(ConvertError::Adapter)?; + partial = partial.end()?; + continue; + } + + let hints = field.map(field_hints).unwrap_or_default(); + partial = partial.begin_nth_field(index)?; + partial = build_from_baml_value_with_hints(partial, field_value, hints) + .map_err(|err| err.with_path_prefix(field_name.to_string()))?; + partial = partial.end()?; + } + + Ok(partial) +} + +fn build_map_from_object_pairs( + partial: Partial<'static>, + map: &BamlMap, +) -> Result, ConvertError> { + let mut p = partial.init_map()?; + for (key, value) in map { + p = p.begin_key()?; + if is_string_target_shape(p.shape()) { + p = p.set(key.clone())?; + } else { + p = p + .parse_from_str(key) + .map_err(ConvertError::from) + .map_err(|err| err.with_path_prefix(key.clone()))?; + } + p = p.end()?; + + p = p.begin_value()?; + p = build_from_baml_value_with_hints(p, value, ReprHints::default()) + .map_err(|err| err.with_path_prefix(key.clone()))?; + p = p.end()?; + } + Ok(p) +} + +fn build_map_from_pair_entries( + partial: Partial<'static>, + items: &[BamlValue], +) -> Result, ConvertError> { + let mut p = partial.init_map()?; + + for (idx, item) in items.iter().enumerate() { + let entry_map = match item { + BamlValue::Class(_, map) | BamlValue::Map(map) => map, + other => { + return Err(ConvertError::TypeMismatch { + expected: "object", + actual: baml_value_kind(other), + }); + } + }; + + let key_value = entry_map + .get("key") + .ok_or_else(|| ConvertError::MissingField("key".to_string()))?; + let value_value = entry_map + .get("value") + .ok_or_else(|| ConvertError::MissingField("value".to_string()))?; + + p = p.begin_key()?; + p = build_from_baml_value_with_hints(p, key_value, ReprHints::default()).map_err( + |err| { + err.with_path_prefix("key") + .with_path_prefix(idx.to_string()) + }, + )?; + p = p.end()?; + + p = p.begin_value()?; + p = build_from_baml_value_with_hints(p, value_value, ReprHints::default()).map_err( + |err| { + err.with_path_prefix("value") + .with_path_prefix(idx.to_string()) + }, + )?; + p = p.end()?; + } + + Ok(p) +} + +fn resolve_field_index(partial: &Partial<'static>, field_name: &str) -> Option { + if let Some(index) = partial.field_index(field_name) { + return Some(index); + } + + let shape = partial.shape(); + if let Type::User(UserType::Struct(struct_type)) = &shape.ty { + return struct_type + .fields + .iter() + .enumerate() + .find_map(|(i, field)| field_matches_name(field, field_name).then_some(i)); + } + + if let Type::User(UserType::Enum(_)) = &shape.ty + && let Some(variant) = partial.selected_variant() + { + return variant + .data + .fields + .iter() + .enumerate() + .find_map(|(i, field)| field_matches_name(field, field_name).then_some(i)); + } + + None +} + +fn current_field(partial: &Partial<'static>, index: usize) -> Option { + let shape = partial.shape(); + if let Type::User(UserType::Struct(struct_type)) = &shape.ty { + return struct_type.fields.get(index).copied(); + } + + if let Type::User(UserType::Enum(_)) = &shape.ty + && let Some(variant) = partial.selected_variant() + { + return variant.data.fields.get(index).copied(); + } + + None +} + +fn field_matches_name(field: &facet::Field, input_name: &str) -> bool { + field.name == input_name + || field.effective_name() == input_name + || field.alias == Some(input_name) +} + +fn field_hints(field: facet::Field) -> ReprHints { + ReprHints { + int_repr: bamltype_attr_static_str(field.attributes, "int_repr") + .and_then(parse_int_repr_hint), + map_key_repr: bamltype_attr_static_str(field.attributes, "map_key_repr") + .and_then(parse_map_key_repr_hint), + } +} + +fn bamltype_attr_static_str(attrs: &'static [facet::Attr], key: &str) -> Option<&'static str> { + for attr in attrs { + if attr.ns != Some("bamltype") || attr.key != key { + continue; + } + + if let Some(value) = attr.get_as::<&'static str>() { + return Some(*value); + } + } + + None +} + +fn parse_int_repr_hint(value: &'static str) -> Option { + match value { + "string" => Some(IntReprHint::String), + "i64" => Some(IntReprHint::I64), + _ => None, + } +} + +fn parse_map_key_repr_hint(value: &'static str) -> Option { + match value { + "string" => Some(MapKeyReprHint::String), + "pairs" => Some(MapKeyReprHint::Pairs), + _ => None, + } +} + +fn integer_expected_and_bounds(shape: &'static Shape) -> Option<(&'static str, i128, i128)> { + use facet::{NumericType, PrimitiveType}; + + let Type::Primitive(PrimitiveType::Numeric(NumericType::Integer { .. })) = shape.ty else { + return None; + }; + + match shape.type_identifier { + "i8" => Some(("i8", i8::MIN as i128, i8::MAX as i128)), + "i16" => Some(("i16", i16::MIN as i128, i16::MAX as i128)), + "i32" => Some(("i32", i32::MIN as i128, i32::MAX as i128)), + "i64" => Some(("i64", i64::MIN as i128, i64::MAX as i128)), + "isize" => Some(("isize", isize::MIN as i128, isize::MAX as i128)), + "u8" => Some(("u8", 0, u8::MAX as i128)), + "u16" => Some(("u16", 0, u16::MAX as i128)), + "u32" => Some(("u32", 0, u32::MAX as i128)), + "u64" => Some(("u64", 0, u64::MAX as i128)), + "usize" => Some(("usize", 0, usize::MAX as i128)), + _ => None, + } +} + +fn baml_value_kind(value: &BamlValue) -> String { + match value { + BamlValue::String(_) => "string", + BamlValue::Int(_) => "int", + BamlValue::Float(_) => "float", + BamlValue::Bool(_) => "bool", + BamlValue::Map(_) => "map", + BamlValue::List(_) => "list", + BamlValue::Class(_, _) => "class", + BamlValue::Enum(_, _) => "enum", + BamlValue::Null => "null", + BamlValue::Media(_) => "media", + } + .to_string() +} + +// ============================================================================ +// Rust → BamlValue (using Peek API) +// ============================================================================ + +/// Convert a Rust value to BamlValue using facet reflection. +pub fn to_baml_value>(value: &T) -> Result { + let peek = Peek::new(value); + peek_to_baml_value(peek) +} + +/// Recursive helper to convert Peek to BamlValue. +fn peek_to_baml_value(peek: Peek<'_, '_>) -> Result { + peek_to_baml_value_with_hints(peek, ReprHints::default()) +} + +fn peek_to_baml_value_with_hints( + peek: Peek<'_, '_>, + hints: ReprHints, +) -> Result { + let peek = peek.innermost_peek(); // unwrap transparent wrappers + + // Option - check before list/map/struct. + if let Ok(opt) = peek.into_option() { + return match opt.value() { + Some(inner) => peek_to_baml_value_with_hints(inner, hints), + None => Ok(BamlValue::Null), + }; + } + + // Handle explicit int_repr first. + if let Some(int_repr) = hints.int_repr { + if let Some(value) = peek_signed_i128(peek) { + return match int_repr { + IntReprHint::String => Ok(BamlValue::String(value.to_string())), + IntReprHint::I64 => i64::try_from(value) + .map(BamlValue::Int) + .map_err(|_| ConvertError::Unsupported("integer out of range for i64".into())), + }; + } + if let Some(value) = peek_unsigned_u128(peek) { + return match int_repr { + IntReprHint::String => Ok(BamlValue::String(value.to_string())), + IntReprHint::I64 => i64::try_from(value) + .map(BamlValue::Int) + .map_err(|_| ConvertError::Unsupported("integer out of range for i64".into())), + }; + } + } + + // Try scalar types first. + if let Some(s) = peek.as_str() { + return Ok(BamlValue::String(s.to_string())); + } + if let Ok(b) = peek.get::() { + return Ok(BamlValue::Bool(*b)); + } + + if let Some(value) = peek_signed_i128(peek) { + return i64::try_from(value) + .map(BamlValue::Int) + .map_err(|_| ConvertError::Unsupported("integer out of range for i64".into())); + } + if let Some(value) = peek_unsigned_u128(peek) { + return i64::try_from(value) + .map(BamlValue::Int) + .map_err(|_| ConvertError::Unsupported("integer out of range for i64".into())); + } + + if let Ok(f) = peek.get::() { + return Ok(BamlValue::Float(*f)); + } + if let Ok(f) = peek.get::() { + return Ok(BamlValue::Float(*f as f64)); + } + + // List/Array. + if let Ok(list) = peek.into_list_like() { + let items: Result, _> = list + .iter() + .map(|item| peek_to_baml_value_with_hints(item, hints)) + .collect(); + return Ok(BamlValue::List(items?)); + } + + // Map. + if let Ok(map) = peek.into_map() { + return map_to_baml_value(map, hints); + } + + // Struct. + if let Ok(struct_peek) = peek.into_struct() { + let type_name = internal_name_for_shape(peek.shape()); + let mut fields = IndexMap::new(); + + for (field_item, field_peek) in struct_peek.fields_for_serialize() { + let field_name = field_item.effective_name().to_string(); + let field_hints = field_item.field.map(field_hints).unwrap_or_default(); + fields.insert( + field_name, + peek_to_baml_value_with_hints(field_peek, field_hints)?, + ); + } + return Ok(BamlValue::Class(type_name, fields)); + } + + // Enum. + if let Ok(enum_peek) = peek.into_enum() { + let type_name = internal_name_for_shape(peek.shape()); + let variant = enum_peek.active_variant()?; + + if enum_has_data_variants(peek.shape()) { + let tag_name = peek.shape().get_tag_attr().unwrap_or("type"); + let mut fields = IndexMap::new(); + fields.insert( + tag_name.to_string(), + BamlValue::String(variant.effective_name().to_string()), + ); + + for (field_item, field_peek) in enum_peek.fields_for_serialize() { + let field_name = field_item.effective_name().to_string(); + let field_hints = field_item.field.map(field_hints).unwrap_or_default(); + fields.insert( + field_name, + peek_to_baml_value_with_hints(field_peek, field_hints)?, + ); + } + + return Ok(BamlValue::Class(type_name, fields)); + } + + return Ok(BamlValue::Enum( + type_name, + variant.effective_name().to_string(), + )); + } + + Err(ConvertError::Unsupported(format!( + "Cannot convert type {} to BamlValue", + peek.shape() + ))) +} + +fn map_to_baml_value( + map: facet_reflect::PeekMap<'_, '_>, + hints: ReprHints, +) -> Result { + match hints.map_key_repr { + Some(MapKeyReprHint::Pairs) => { + let mut entries = Vec::with_capacity(map.len()); + for (key, value) in map.iter() { + let mut entry = IndexMap::new(); + entry.insert( + "key".to_string(), + peek_to_baml_value_with_hints(key, ReprHints::default())?, + ); + entry.insert( + "value".to_string(), + peek_to_baml_value_with_hints(value, ReprHints::default())?, + ); + entries.push(BamlValue::Map(entry)); + } + Ok(BamlValue::List(entries)) + } + Some(MapKeyReprHint::String) | None => { + let mut result = IndexMap::new(); + for (key, value) in map.iter() { + let key_str = key_to_string(key)?; + result.insert( + key_str, + peek_to_baml_value_with_hints(value, ReprHints::default())?, + ); + } + Ok(BamlValue::Map(result)) + } + } +} + +fn key_to_string(key: Peek<'_, '_>) -> Result { + if let Some(s) = key.as_str() { + return Ok(s.to_string()); + } + + if let Some(value) = peek_signed_i128(key) { + return Ok(value.to_string()); + } + + if let Some(value) = peek_unsigned_u128(key) { + return Ok(value.to_string()); + } + + if let Ok(value) = key.get::() { + return Ok(value.to_string()); + } + + Ok(format!("{}", key)) +} + +fn enum_has_data_variants(shape: &'static Shape) -> bool { + let Type::User(UserType::Enum(enum_type)) = &shape.ty else { + return false; + }; + enum_type + .variants + .iter() + .any(|variant| !variant.data.fields.is_empty()) +} + +fn peek_signed_i128(peek: Peek<'_, '_>) -> Option { + if let Ok(v) = peek.get::() { + return Some(*v); + } + if let Ok(v) = peek.get::() { + return Some(*v as i128); + } + if let Ok(v) = peek.get::() { + return Some(*v as i128); + } + if let Ok(v) = peek.get::() { + return Some(*v as i128); + } + if let Ok(v) = peek.get::() { + return Some(*v as i128); + } + if let Ok(v) = peek.get::() { + return Some(*v as i128); + } + None +} + +fn peek_unsigned_u128(peek: Peek<'_, '_>) -> Option { + if let Ok(v) = peek.get::() { + return Some(*v); + } + if let Ok(v) = peek.get::() { + return Some(*v as u128); + } + if let Ok(v) = peek.get::() { + return Some(*v as u128); + } + if let Ok(v) = peek.get::() { + return Some(*v as u128); + } + if let Ok(v) = peek.get::() { + return Some(*v as u128); + } + if let Ok(v) = peek.get::() { + return Some(*v as u128); + } + None +} + +fn is_string_target_shape(shape: &'static Shape) -> bool { + matches!( + shape.scalar_type(), + Some(ScalarType::Str | ScalarType::String | ScalarType::CowStr | ScalarType::Char) + ) +} + +fn expected_kind_for_shape(shape: &'static Shape) -> &'static str { + use facet::{NumericType, PrimitiveType, TextualType}; + + match &shape.ty { + Type::Primitive(PrimitiveType::Boolean) => "bool", + Type::Primitive(PrimitiveType::Numeric(NumericType::Integer { .. })) => "int", + Type::Primitive(PrimitiveType::Numeric(NumericType::Float)) => "float", + Type::Primitive(PrimitiveType::Textual(TextualType::Str | TextualType::Char)) => "string", + Type::User(UserType::Struct(_)) | Type::User(UserType::Enum(_)) => "object", + _ => match shape.def { + Def::Map(_) => "map", + Def::List(_) | Def::Array(_) | Def::Set(_) => "list", + _ => "compatible type", + }, + } +} + +fn should_skip_struct_field_deserializing(shape: &'static facet::Shape, input_name: &str) -> bool { + let Type::User(UserType::Struct(struct_type)) = &shape.ty else { + return false; + }; + + struct_type.fields.iter().any(|field| { + let matches_name = field_matches_name(field, input_name); + matches_name && field.should_skip_deserializing() + }) +} + +fn select_enum_variant( + partial: Partial<'static>, + variant_name: &str, +) -> Result, ConvertError> { + let shape = partial.shape(); + let Type::User(UserType::Enum(enum_type)) = &shape.ty else { + return Ok(partial.select_variant_named(variant_name)?); + }; + + if let Some((index, _)) = enum_type.variants.iter().enumerate().find(|(_, variant)| { + variant.effective_name() == variant_name || variant.name == variant_name + }) { + return Ok(partial.select_nth_variant(index)?); + } + + Err(ConvertError::UnknownVariant(variant_name.to_string())) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_primitives_to_baml() { + let pi = std::f64::consts::PI; + assert_eq!( + to_baml_value(&"hello".to_string()).unwrap(), + BamlValue::String("hello".into()) + ); + assert_eq!(to_baml_value(&42i64).unwrap(), BamlValue::Int(42)); + assert_eq!(to_baml_value(&pi).unwrap(), BamlValue::Float(pi)); + assert_eq!(to_baml_value(&true).unwrap(), BamlValue::Bool(true)); + } + + #[test] + fn test_primitives_from_baml() { + let s: String = from_baml_value(BamlValue::String("world".into())).unwrap(); + assert_eq!(s, "world"); + + let i: i64 = from_baml_value(BamlValue::Int(100)).unwrap(); + assert_eq!(i, 100); + + let e = std::f64::consts::E; + let f: f64 = from_baml_value(BamlValue::Float(e)).unwrap(); + assert!((f - e).abs() < 0.001); + + let b: bool = from_baml_value(BamlValue::Bool(false)).unwrap(); + assert!(!b); + } + + #[test] + fn test_list_round_trip() { + let original = vec![1i64, 2, 3]; + let baml = to_baml_value(&original).unwrap(); + assert!(matches!(baml, BamlValue::List(_))); + + let restored: Vec = from_baml_value(baml).unwrap(); + assert_eq!(original, restored); + } + + #[test] + fn test_option_to_baml() { + let some_val: Option = Some(42); + let none_val: Option = None; + + assert_eq!(to_baml_value(&some_val).unwrap(), BamlValue::Int(42)); + assert_eq!(to_baml_value(&none_val).unwrap(), BamlValue::Null); + } +} diff --git a/crates/bamltype/src/facet_ext.rs b/crates/bamltype/src/facet_ext.rs new file mode 100644 index 00000000..1bc8e84e --- /dev/null +++ b/crates/bamltype/src/facet_ext.rs @@ -0,0 +1,82 @@ +//! Facet extension attributes used by bamltype. + +use baml_types::{BamlValue, TypeIR}; + +use crate::compat::{BamlConvertError, Registry}; + +/// Field-level adapter application function. +pub type AdapterApplyFn = fn( + facet_reflect::Partial<'static>, + BamlValue, + Vec, +) -> Result, BamlConvertError>; + +/// Runtime hooks for `#[baml(with = "...")]`. +#[derive(Clone, Copy, Debug, facet::Facet)] +#[facet(opaque)] +pub struct WithAdapterFns { + /// Schema type representation callback. + pub type_ir: fn() -> TypeIR, + /// Schema registration callback. + pub register: fn(&mut Registry), + /// Value conversion callback. + pub apply: AdapterApplyFn, +} + +impl PartialEq for WithAdapterFns { + fn eq(&self, other: &Self) -> bool { + std::ptr::fn_addr_eq(self.type_ir, other.type_ir) + && std::ptr::fn_addr_eq(self.register, other.register) + && std::ptr::fn_addr_eq(self.apply, other.apply) + } +} + +impl Eq for WithAdapterFns {} + +/// Resolve `#[facet(bamltype::with = ...)]` payload from a field attribute list. +pub fn with_adapter_fns(attrs: &'static [facet::Attr]) -> Option<&'static WithAdapterFns> { + for attr in attrs { + if attr.ns != Some("bamltype") || attr.key != "with" { + continue; + } + + let Some(ext_attr) = attr.get_as::() else { + continue; + }; + + if let Attr::With(Some(fns)) = ext_attr { + return Some(*fns); + } + } + + None +} + +facet::define_attr_grammar! { + ns "bamltype"; + crate_path $crate::facet_ext; + + /// Constraint payload for BAML-compatible `check` / `assert` attributes. + pub struct BamlConstraint { + /// Constraint label. + pub label: &'static str, + /// Constraint expression. + pub expr: &'static str, + } + + /// bamltype extension attrs consumed by derive and runtime conversion/schema. + pub enum Attr { + /// Container-level override for internal type name. + InternalName(&'static str), + /// Field-level integer representation strategy. + IntRepr(&'static str), + /// Field-level map key representation strategy. + MapKeyRepr(&'static str), + /// Field-level adapter payload (`#[facet(bamltype::with = ...)]`). + With(Option<&'static WithAdapterFns>), + /// BAML-compatible check constraint (`#[facet(bamltype::check(...))]`). + Check(BamlConstraint), + /// BAML-compatible assert constraint (`#[facet(bamltype::assert(...))]`). + Assert(BamlConstraint), + } +} diff --git a/crates/bamltype/src/lib.rs b/crates/bamltype/src/lib.rs new file mode 100644 index 00000000..a7da9ebc --- /dev/null +++ b/crates/bamltype/src/lib.rs @@ -0,0 +1,462 @@ +//! BamlType - Facet-based BAML type generation +//! +//! This crate provides automatic BAML schema generation and LLM output parsing +//! for Rust types using facet's compile-time reflection. +//! +//! # Usage +//! +//! ```ignore +//! use bamltype::BamlType; +//! +//! #[BamlType] +//! struct Response { +//! /// The user's name +//! name: String, +//! /// Age in years +//! age: u32, +//! } +//! +//! // Render schema for LLM prompt +//! let schema = bamltype::render_schema::(bamltype::RenderOptions::default())?; +//! +//! // Parse LLM output +//! let parsed = bamltype::parse_llm_output::(llm_output, true)?; +//! ``` + +use std::collections::HashSet; + +use baml_types::{BamlValue, Constraint, ConstraintLevel, ResponseCheck, TypeIR}; +use internal_baml_jinja::types::OutputFormatContent; +use jsonish::deserializer::{ + coercer::{ParsingError, run_user_checks}, + deserialize_flags::{DeserializerConditions, Flag}, +}; +use sha2::{Digest, Sha256}; + +// Re-export underlying crates for consumers (replaces baml-bridge re-exports) +pub use baml_types; +pub use internal_baml_jinja; +pub use internal_baml_jinja::types::{HoistClasses, MapStyle, RenderOptions}; +pub use jsonish; +pub use jsonish::BamlValueWithFlags; + +// Re-export facet for users +pub use facet; +pub use facet::Shape; +pub use facet_reflect; + +// Re-export the attribute macro +#[cfg(feature = "derive")] +pub use bamltype_derive::BamlType; + +mod schema_builder; +pub use schema_builder::*; + +mod convert; +pub use convert::{ConvertError, from_baml_value, from_baml_value_with_flags, to_baml_value}; + +pub mod compat; +pub mod facet_ext; +pub use compat::{ + BamlAdapter, BamlConvertError, BamlTypeInternal, BamlTypeTrait, BamlValueConvert, Registry, + ToBamlValue, default_streaming_behavior, get_field, with_constraints, +}; + +/// A bundle containing everything needed to render schemas and parse LLM output. +#[derive(Debug, Clone)] +pub struct SchemaBundle { + /// The TypeIR for the root type + pub target: TypeIR, + /// The output format content (classes, enums, etc.) + pub output_format: OutputFormatContent, +} + +impl SchemaBundle { + /// Build a SchemaBundle from a facet Shape. + /// + /// This walks the type graph starting from the given shape, + /// building BAML class/enum definitions for all reachable types. + pub fn from_shape(shape: &'static facet::Shape) -> Self { + schema_builder::build_schema_bundle(shape) + } +} + +/// Trait for types that can generate BAML schemas. +/// +/// Implemented automatically by `#[BamlType]`. +pub trait BamlSchema: for<'a> facet::Facet<'a> { + /// Get the schema bundle for this type. + /// + /// This is lazily initialized and cached. + fn baml_schema() -> &'static SchemaBundle; +} + +/// Back-compat trait mirroring baml-bridge's `BamlType`. +/// +/// This sits alongside the `#[BamlType]` attribute macro and offers the same +/// runtime trait entry points users expect from the old API. +pub trait BamlType: BamlTypeInternal + BamlValueConvert + Sized + 'static { + fn baml_output_format() -> &'static OutputFormatContent; + + fn baml_internal_name() -> &'static str { + ::baml_internal_name() + } + + fn baml_type_ir() -> TypeIR { + ::baml_type_ir() + } + + fn try_from_baml_value(value: BamlValue) -> Result { + ::try_from_baml_value(value, Vec::new()) + } +} + +impl BamlType for T { + fn baml_output_format() -> &'static OutputFormatContent { + ::baml_output_format() + } +} + +/// Parsed output bundle matching baml-bridge behavior. +#[derive(Debug, Clone)] +pub struct Parsed { + pub value: T, + pub baml_value: BamlValue, + pub flags: Vec, + pub checks: Vec, + pub explanations: Vec, +} + +/// Error type for parse + conversion failures. +#[derive(Debug, thiserror::Error)] +pub enum BamlParseError { + #[error("jsonish parse error: {0}")] + Jsonish(#[from] anyhow::Error), + + #[error("constraint asserts failed")] + ConstraintAssertsFailed { failed: Vec }, + + #[error("conversion error: {0}")] + Convert(#[from] BamlConvertError), +} + +/// Error type for parsing failures. +#[derive(Debug, thiserror::Error)] +pub enum ParseError { + #[error("Schema rendering failed: {0}")] + RenderError(String), + + #[error("Parse error: {0}")] + ParseError(String), + + #[error("Coercion error: {0}")] + CoercionError(String), +} + +/// Render the BAML schema for a type, matching baml-bridge signature. +pub fn render_schema( + options: RenderOptions, +) -> Result, minijinja::Error> { + T::baml_output_format().render(options) +} + +/// Convenience helper equivalent to `render_schema(RenderOptions::default())`. +pub fn render_schema_default() -> Result { + render_schema_string::(RenderOptions::default()) +} + +/// Parse LLM output into a BamlValue. +/// +/// This uses jsonish for flexible JSON parsing that handles +/// common LLM output quirks (markdown code blocks, trailing commas, etc). +pub fn parse(raw: &str) -> Result { + parse_with_mode::(raw, true) +} + +/// Parse streaming LLM output (partial results). +pub fn parse_partial(raw: &str) -> Result { + parse_with_mode::(raw, false) +} + +/// Parse raw LLM output and convert into `T`, collecting flags/checks/explanations. +pub fn parse_llm_output( + raw: &str, + is_done: bool, +) -> Result, BamlParseError> { + let output_format = T::baml_output_format(); + let parsed = match parse_raw::(raw, is_done) { + Ok(parsed) => parsed, + Err(err) => { + if has_assert_failure(&err) { + let failed = collect_assert_constraints(output_format); + return Err(BamlParseError::ConstraintAssertsFailed { failed }); + } + return Err(BamlParseError::Jsonish(err)); + } + }; + + let baml_value_with_meta: baml_types::BamlValueWithMeta = + parsed.clone().into(); + let baml_value: BamlValue = baml_value_with_meta.into(); + + let value = ::try_from_baml_value(baml_value.clone())?; + + let mut flags = Vec::new(); + collect_flags(&parsed, &mut flags); + + let mut checks = Vec::new(); + let mut failed_asserts = Vec::new(); + collect_checks_and_assert_failures(&parsed, &mut checks, &mut failed_asserts)?; + + let mut explanations = Vec::new(); + parsed.explanation_impl(vec!["".to_string()], &mut explanations); + + if !failed_asserts.is_empty() { + return Err(BamlParseError::ConstraintAssertsFailed { + failed: failed_asserts, + }); + } + + Ok(Parsed { + value, + baml_value, + flags, + checks, + explanations, + }) +} + +/// Stable schema fingerprint used by cache keys and regression tests. +pub fn schema_fingerprint( + output_format: &OutputFormatContent, + options: RenderOptions, +) -> Result { + let rendered = output_format.render(options)?.unwrap_or_default(); + let mut hasher = Sha256::new(); + hasher.update(rendered.as_bytes()); + hasher.update(output_format.target.to_string().as_bytes()); + Ok(format!("{:x}", hasher.finalize())) +} + +fn parse_raw(raw: &str, is_done: bool) -> Result { + let output_format = T::baml_output_format(); + jsonish::from_str(output_format, &output_format.target, raw, is_done) +} + +fn render_schema_string(options: RenderOptions) -> Result { + render_schema::(options) + .map(Option::unwrap_or_default) + .map_err(|e| ParseError::RenderError(e.to_string())) +} + +fn parse_with_mode( + raw: &str, + is_done: bool, +) -> Result { + parse_raw::(raw, is_done).map_err(|e| ParseError::ParseError(e.to_string())) +} + +fn collect_flags_recursive(value: &BamlValueWithFlags, flags: &mut Vec) { + match value { + BamlValueWithFlags::String(v) => collect_from_conditions(&v.flags, flags), + BamlValueWithFlags::Int(v) => collect_from_conditions(&v.flags, flags), + BamlValueWithFlags::Float(v) => collect_from_conditions(&v.flags, flags), + BamlValueWithFlags::Bool(v) => collect_from_conditions(&v.flags, flags), + BamlValueWithFlags::Enum(_, _, v) => collect_from_conditions(&v.flags, flags), + BamlValueWithFlags::Media(_, v) => collect_from_conditions(&v.flags, flags), + BamlValueWithFlags::List(conds, _, items) => { + collect_from_conditions(conds, flags); + for item in items { + collect_flags_recursive(item, flags); + } + } + BamlValueWithFlags::Map(conds, _, items) => { + collect_from_conditions(conds, flags); + for (_, (entry_flags, entry_value)) in items { + collect_from_conditions(entry_flags, flags); + collect_flags_recursive(entry_value, flags); + } + } + BamlValueWithFlags::Class(_, conds, _, fields) => { + collect_from_conditions(conds, flags); + for (_, field_value) in fields { + collect_flags_recursive(field_value, flags); + } + } + BamlValueWithFlags::Null(_, conds) => collect_from_conditions(conds, flags), + } +} + +fn collect_from_conditions(conditions: &DeserializerConditions, flags: &mut Vec) { + flags.extend(conditions.flags.iter().cloned()); +} + +fn collect_flags(value: &BamlValueWithFlags, flags: &mut Vec) { + collect_flags_recursive(value, flags); +} + +fn collect_checks_and_assert_failures( + value: &BamlValueWithFlags, + checks: &mut Vec, + failed_asserts: &mut Vec, +) -> Result<(), BamlParseError> { + let baml_value = baml_value_from_flags(value); + let results = run_user_checks(&baml_value, value.field_type()).map_err(BamlParseError::from)?; + + for (constraint, ok) in results { + if constraint.level == ConstraintLevel::Assert { + if !ok { + failed_asserts.push(ResponseCheck { + name: constraint.label.unwrap_or_else(|| "assert".to_string()), + expression: constraint.expression.0, + status: "failed".to_string(), + }); + } + continue; + } + + if let Some(check) = ResponseCheck::from_check_result((constraint, ok)) { + checks.push(check); + } + } + + match value { + BamlValueWithFlags::List(_, _, items) => { + for item in items { + collect_checks_and_assert_failures(item, checks, failed_asserts)?; + } + } + BamlValueWithFlags::Map(_, _, items) => { + for (_, (_, entry_value)) in items { + collect_checks_and_assert_failures(entry_value, checks, failed_asserts)?; + } + } + BamlValueWithFlags::Class(_, _, _, fields) => { + for (_, field_value) in fields { + collect_checks_and_assert_failures(field_value, checks, failed_asserts)?; + } + } + _ => {} + } + + Ok(()) +} + +fn collect_assert_constraints(output_format: &OutputFormatContent) -> Vec { + let mut failed = Vec::new(); + let mut seen = HashSet::new(); + + collect_assert_constraints_in_type(&output_format.target, &mut failed, &mut seen); + + for class in output_format.classes.values() { + for constraint in &class.constraints { + push_assert_constraint(constraint, &mut failed, &mut seen); + } + for (_, field_type, _, _) in &class.fields { + collect_assert_constraints_in_type(field_type, &mut failed, &mut seen); + } + } + + for r#enum in output_format.enums.values() { + for constraint in &r#enum.constraints { + push_assert_constraint(constraint, &mut failed, &mut seen); + } + } + + failed +} + +fn collect_assert_constraints_in_type( + r#type: &TypeIR, + failed: &mut Vec, + seen: &mut HashSet<(String, String)>, +) { + for constraint in &r#type.meta().constraints { + push_assert_constraint(constraint, failed, seen); + } + + match r#type { + TypeIR::List(inner, _) => collect_assert_constraints_in_type(inner, failed, seen), + TypeIR::Map(key, value, _) => { + collect_assert_constraints_in_type(key, failed, seen); + collect_assert_constraints_in_type(value, failed, seen); + } + TypeIR::Union(union, _) => { + for item in union.iter_include_null() { + collect_assert_constraints_in_type(item, failed, seen); + } + } + TypeIR::Tuple(items, _) => { + for item in items { + collect_assert_constraints_in_type(item, failed, seen); + } + } + TypeIR::Arrow(arrow, _) => { + for param in &arrow.param_types { + collect_assert_constraints_in_type(param, failed, seen); + } + collect_assert_constraints_in_type(&arrow.return_type, failed, seen); + } + TypeIR::Top(_) + | TypeIR::Primitive(..) + | TypeIR::Enum { .. } + | TypeIR::Literal(..) + | TypeIR::Class { .. } + | TypeIR::RecursiveTypeAlias { .. } => {} + } +} + +fn push_assert_constraint( + constraint: &Constraint, + failed: &mut Vec, + seen: &mut HashSet<(String, String)>, +) { + if constraint.level != ConstraintLevel::Assert { + return; + } + + let name = constraint + .label + .clone() + .unwrap_or_else(|| "assert".to_string()); + let expr = constraint.expression.0.clone(); + if seen.insert((name.clone(), expr.clone())) { + failed.push(ResponseCheck { + name, + expression: expr, + status: "failed".to_string(), + }); + } +} + +fn has_assert_failure(err: &anyhow::Error) -> bool { + err.to_string().contains("Assertions failed.") +} + +fn baml_value_from_flags(value: &BamlValueWithFlags) -> BamlValue { + match value { + BamlValueWithFlags::String(v) => BamlValue::String(v.value.clone()), + BamlValueWithFlags::Int(v) => BamlValue::Int(v.value), + BamlValueWithFlags::Float(v) => BamlValue::Float(v.value), + BamlValueWithFlags::Bool(v) => BamlValue::Bool(v.value), + BamlValueWithFlags::Enum(name, _, v) => BamlValue::Enum(name.clone(), v.value.clone()), + BamlValueWithFlags::Media(_, v) => BamlValue::Media(v.value.clone()), + BamlValueWithFlags::List(_, _, items) => { + BamlValue::List(items.iter().map(baml_value_from_flags).collect()) + } + BamlValueWithFlags::Map(_, _, items) => BamlValue::Map( + items + .iter() + .map(|(k, (_, v))| (k.clone(), baml_value_from_flags(v))) + .collect(), + ), + BamlValueWithFlags::Class(name, _, _, fields) => BamlValue::Class( + name.clone(), + fields + .iter() + .map(|(k, v)| (k.clone(), baml_value_from_flags(v))) + .collect(), + ), + BamlValueWithFlags::Null(_, _) => BamlValue::Null, + } +} diff --git a/crates/bamltype/src/schema_builder.rs b/crates/bamltype/src/schema_builder.rs new file mode 100644 index 00000000..8b38aeac --- /dev/null +++ b/crates/bamltype/src/schema_builder.rs @@ -0,0 +1,660 @@ +//! Schema builder - walks facet Shapes to build BAML schemas. +//! +//! This module implements the core reflection-to-BAML translation. + +use std::collections::HashMap; + +use baml_types::{Constraint, StreamingMode, TypeIR, type_meta}; +use facet::{Attr, ConstTypeId, Def, Field, Shape, Type, UserType}; +use internal_baml_jinja::types::{Class, Enum, Name, OutputFormatContent}; + +use crate::SchemaBundle; +use crate::compat::Registry; +use crate::facet_ext; + +/// Build a SchemaBundle from a facet Shape. +pub fn build_schema_bundle(shape: &'static Shape) -> SchemaBundle { + let mut builder = SchemaBuilder::new(); + let target = builder.build_type_ir(shape); + let output_format = builder.into_output_format(target.clone()); + SchemaBundle { + target, + output_format, + } +} + +/// Build a TypeIR from a facet Shape (without building full schema). +/// +/// Used by the compat layer to provide `BamlTypeInternal::baml_type_ir()` +/// for any `Facet` type. +pub fn build_type_ir_from_shape(shape: &'static Shape) -> TypeIR { + let mut builder = SchemaBuilder::new(); + builder.build_type_ir(shape) +} + +/// Compute the BAML internal name for a shape. +/// +/// This mirrors bridge behavior: explicit `#[baml(internal_name = ...)]` takes +/// precedence, otherwise module path + type identifier when available. +pub fn internal_name_for_shape(shape: &'static Shape) -> String { + if let Some(name) = bamltype_internal_name(shape.attributes) { + return name.to_string(); + } + + match shape.module_path { + Some(module) if !module.is_empty() => format!("{module}::{}", shape.type_identifier), + _ => shape.type_identifier.to_string(), + } +} + +/// Compute the rendered/display name for a shape. +/// +/// Facet currently stores container rename in `Shape::rename` for some type kinds, +/// but for others it may only be present in builtin attrs. Prefer the explicit +/// shape field, then fall back to builtin attr lookup for parity with baml-bridge. +fn rendered_name_for_shape(shape: &'static Shape) -> String { + if shape.rename.is_some() { + return shape.effective_name().to_string(); + } + + if let Some(name) = shape.get_builtin_attr_value::<&'static str>("rename") { + return name.to_string(); + } + + shape.type_identifier.to_string() +} + +/// Internal builder state for schema construction. +struct SchemaBuilder { + /// Memoization: shape id -> (internal_name, TypeIR) + visited: HashMap, + + /// Collected schema elements. + registry: Registry, + + /// Track which internal names are already used (for collision handling) + used_internal_names: HashMap, + + /// Collision suffix counter + name_counter: usize, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum IntReprMode { + String, + I64, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum MapKeyReprMode { + String, + Pairs, +} + +#[derive(Clone, Debug)] +struct MapEntryContext { + owner_internal_name: String, + field_name: String, + rendered_field: String, + variant_name: Option, + variant_rendered: Option, +} + +impl SchemaBuilder { + fn new() -> Self { + Self { + visited: HashMap::new(), + registry: Registry::new(), + used_internal_names: HashMap::new(), + name_counter: 0, + } + } + + /// Build TypeIR from a facet Shape. + fn build_type_ir(&mut self, shape: &'static Shape) -> TypeIR { + // Check if already visited (handles recursion) + if let Some((_, type_ir)) = self.visited.get(&shape.id) { + return type_ir.clone(); + } + + // Handle based on semantic Def first. + match &shape.def { + Def::Scalar => self.build_scalar_ir(shape), + Def::Option(option_def) => { + let inner_ir = self.build_type_ir(option_def.t); + TypeIR::optional(inner_ir) + } + Def::List(list_def) => self.build_list_ir(list_def.t), + Def::Array(arr_def) => self.build_list_ir(arr_def.t), + Def::Map(map_def) => { + let key_ir = self.build_type_ir(map_def.k); + let value_ir = self.build_type_ir(map_def.v); + TypeIR::map(key_ir, value_ir) + } + Def::Set(set_def) => self.build_list_ir(set_def.t), + Def::Pointer(ptr_def) => { + // Smart pointers - unwrap to inner type when available + if let Some(pointee) = ptr_def.pointee { + self.build_type_ir(pointee) + } else { + TypeIR::string() + } + } + Def::Undefined => { + if let Some(inner) = shape.inner { + return self.build_type_ir(inner); + } + self.build_from_type(shape) + } + _ => self.build_from_type(shape), + } + } + + fn build_list_ir(&mut self, item_shape: &'static Shape) -> TypeIR { + TypeIR::list(self.build_type_ir(item_shape)) + } + + /// Build TypeIR from the shape's Type field (for user-defined types). + fn build_from_type(&mut self, shape: &'static Shape) -> TypeIR { + match &shape.ty { + Type::User(UserType::Struct(struct_type)) => self.build_struct_ir(shape, struct_type), + Type::User(UserType::Enum(enum_type)) => self.build_enum_ir(shape, enum_type), + Type::Primitive(primitive) => self.build_primitive_ir(primitive), + _ => TypeIR::string(), + } + } + + /// Build TypeIR for scalar/primitive shapes. + fn build_scalar_ir(&self, shape: &'static Shape) -> TypeIR { + match &shape.ty { + Type::Primitive(primitive) => self.build_primitive_ir(primitive), + _ => TypeIR::string(), + } + } + + /// Build TypeIR for primitive types. + fn build_primitive_ir(&self, primitive: &facet::PrimitiveType) -> TypeIR { + use facet::{NumericType, PrimitiveType, TextualType}; + + match primitive { + PrimitiveType::Boolean => TypeIR::bool(), + PrimitiveType::Numeric(NumericType::Integer { .. }) => TypeIR::int(), + PrimitiveType::Numeric(NumericType::Float) => TypeIR::float(), + PrimitiveType::Textual(TextualType::Str) => TypeIR::string(), + PrimitiveType::Textual(TextualType::Char) => TypeIR::string(), + PrimitiveType::Never => TypeIR::string(), + } + } + + /// Build TypeIR for struct types, registering the class. + fn build_struct_ir( + &mut self, + shape: &'static Shape, + struct_type: &facet::StructType, + ) -> TypeIR { + let internal_name = self.generate_internal_name(shape); + let display_name = rendered_name_for_shape(shape); + let type_constraints = constraints_from_attrs(shape.attributes); + + // Register as visited BEFORE recursing (handles cycles). + let mut type_ir = TypeIR::class(&internal_name); + type_ir + .meta_mut() + .constraints + .extend(type_constraints.clone()); + self.visited + .insert(shape.id, (internal_name.clone(), type_ir.clone())); + + // Build fields. + let mut fields = Vec::new(); + for field in struct_type.fields.iter() { + if field.should_skip_deserializing() { + continue; + } + + let mut field_ir = self.build_field_type_ir(field, &internal_name, None, None); + field_ir + .meta_mut() + .constraints + .extend(constraints_from_attrs(field.attributes)); + if field.has_default() && !field_ir.is_optional() { + field_ir = TypeIR::optional(field_ir); + } + + let field_name = field.name.to_string(); + let rendered_name = field.effective_name().to_string(); + let description = doc_to_description(field.doc); + let name = name_with_optional_alias(field_name, rendered_name); + fields.push((name, field_ir, description, false)); + } + + let description = doc_to_description(shape.doc); + + let class = Class { + name: name_with_optional_alias(internal_name.clone(), display_name), + description, + namespace: StreamingMode::NonStreaming, + fields, + constraints: type_constraints, + streaming_behavior: Default::default(), + }; + + self.register_class(class); + + type_ir + } + + /// Build TypeIR for enum types. + fn build_enum_ir(&mut self, shape: &'static Shape, enum_type: &facet::EnumType) -> TypeIR { + let internal_name = self.generate_internal_name(shape); + let display_name = rendered_name_for_shape(shape); + let type_constraints = constraints_from_attrs(shape.attributes); + + let is_data_enum = enum_type + .variants + .iter() + .any(|variant| !variant.data.fields.is_empty()); + + if !is_data_enum { + if shape.is_untagged() { + let literals: Vec = enum_type + .variants + .iter() + .map(|variant| TypeIR::literal_string(variant.effective_name().to_string())) + .collect(); + + let mut type_ir = TypeIR::union_with_meta(literals, type_meta::IR::default()); + type_ir + .meta_mut() + .constraints + .extend(type_constraints.clone()); + + self.visited + .insert(shape.id, (internal_name.clone(), type_ir.clone())); + return type_ir; + } + + let mut type_ir = TypeIR::r#enum(&internal_name); + type_ir + .meta_mut() + .constraints + .extend(type_constraints.clone()); + + self.visited + .insert(shape.id, (internal_name.clone(), type_ir.clone())); + + let mut values = Vec::new(); + for variant in enum_type.variants.iter() { + let variant_name = variant.name.to_string(); + let rendered_name = variant.effective_name().to_string(); + let description = doc_to_description(variant.doc); + let name = name_with_optional_alias(variant_name, rendered_name); + values.push((name, description)); + } + + let description = doc_to_description(shape.doc); + let enm = Enum { + name: name_with_optional_alias(internal_name.clone(), display_name), + description, + values, + constraints: type_constraints, + }; + + self.registry.register_enum(enm); + return type_ir; + } + + // Data enum: represented as union of generated variant classes. + let variant_class_names: Vec = enum_type + .variants + .iter() + .map(|variant| format!("{internal_name}__{}", variant.name)) + .collect(); + + let union_variants: Vec = variant_class_names.iter().map(TypeIR::class).collect(); + + let mut type_ir = TypeIR::union_with_meta(union_variants, type_meta::IR::default()); + type_ir + .meta_mut() + .constraints + .extend(type_constraints.clone()); + + // Register as visited before variant field recursion for cycle handling. + self.visited + .insert(shape.id, (internal_name.clone(), type_ir.clone())); + + let tag_name = shape.get_tag_attr().unwrap_or("type").to_string(); + + for variant in enum_type.variants.iter() { + let variant_internal_name = format!("{internal_name}__{}", variant.name); + let rendered_variant = variant.effective_name().to_string(); + let variant_rendered_name = format!("{display_name}_{rendered_variant}"); + let variant_description = doc_to_description(variant.doc); + + let mut fields = Vec::new(); + fields.push(( + Name::new(tag_name.clone()), + TypeIR::literal_string(rendered_variant.clone()), + None, + false, + )); + + for field in variant.data.fields.iter() { + if field.should_skip_deserializing() { + continue; + } + + let mut field_ir = self.build_field_type_ir( + field, + &internal_name, + Some(variant.name), + Some(variant.effective_name()), + ); + field_ir + .meta_mut() + .constraints + .extend(constraints_from_attrs(field.attributes)); + if field.has_default() && !field_ir.is_optional() { + field_ir = TypeIR::optional(field_ir); + } + + let field_name = field.name.to_string(); + let rendered_field = field.effective_name().to_string(); + let field_description = doc_to_description(field.doc); + let field_name = name_with_optional_alias(field_name, rendered_field); + fields.push((field_name, field_ir, field_description, false)); + } + + let class = Class { + name: name_with_optional_alias(variant_internal_name, variant_rendered_name), + description: variant_description, + namespace: StreamingMode::NonStreaming, + fields, + constraints: Vec::new(), + streaming_behavior: Default::default(), + }; + + self.register_class(class); + } + + type_ir + } + + fn build_field_type_ir( + &mut self, + field: &Field, + owner_internal_name: &str, + variant_name: Option<&str>, + variant_rendered: Option<&str>, + ) -> TypeIR { + if let Some(with) = facet_ext::with_adapter_fns(field.attributes) { + (with.register)(&mut self.registry); + return (with.type_ir)(); + } + + if let Some(int_repr) = field_int_repr(field) { + return Self::build_int_repr_ir(field.shape(), int_repr); + } + + if let Some(map_repr) = field_map_key_repr(field) { + let entry_ctx = MapEntryContext { + owner_internal_name: owner_internal_name.to_string(), + field_name: field.name.to_string(), + rendered_field: field.effective_name().to_string(), + variant_name: variant_name.map(std::string::ToString::to_string), + variant_rendered: variant_rendered.map(std::string::ToString::to_string), + }; + return self.build_map_key_repr_ir(field.shape(), map_repr, Some(entry_ctx)); + } + + self.build_type_ir(field.shape()) + } + + fn build_int_repr_ir(shape: &'static Shape, repr: IntReprMode) -> TypeIR { + match &shape.def { + Def::Option(option_def) => { + TypeIR::optional(Self::build_int_repr_ir(option_def.t, repr)) + } + Def::List(list_def) => TypeIR::list(Self::build_int_repr_ir(list_def.t, repr)), + Def::Array(arr_def) => TypeIR::list(Self::build_int_repr_ir(arr_def.t, repr)), + Def::Pointer(ptr_def) => { + if let Some(pointee) = ptr_def.pointee { + Self::build_int_repr_ir(pointee, repr) + } else { + TypeIR::string() + } + } + _ => match repr { + IntReprMode::String => TypeIR::string(), + IntReprMode::I64 => TypeIR::int(), + }, + } + } + + fn build_map_key_repr_ir( + &mut self, + shape: &'static Shape, + repr: MapKeyReprMode, + entry_ctx: Option, + ) -> TypeIR { + match &shape.def { + Def::Option(option_def) => { + TypeIR::optional(self.build_map_key_repr_ir(option_def.t, repr, entry_ctx)) + } + Def::List(list_def) => { + TypeIR::list(self.build_map_key_repr_ir(list_def.t, repr, entry_ctx)) + } + Def::Array(arr_def) => { + TypeIR::list(self.build_map_key_repr_ir(arr_def.t, repr, entry_ctx)) + } + Def::Pointer(ptr_def) => { + if let Some(pointee) = ptr_def.pointee { + self.build_map_key_repr_ir(pointee, repr, entry_ctx) + } else { + TypeIR::string() + } + } + Def::Map(map_def) => match repr { + MapKeyReprMode::String => { + let value_ir = self.build_type_ir(map_def.v); + TypeIR::map(TypeIR::string(), value_ir) + } + MapKeyReprMode::Pairs => { + let ctx = entry_ctx.unwrap_or_else(|| MapEntryContext { + owner_internal_name: "MapEntry".to_string(), + field_name: "entries".to_string(), + rendered_field: "entries".to_string(), + variant_name: None, + variant_rendered: None, + }); + + let (entry_internal_name, rendered_entry_name) = map_entry_names(&ctx); + self.ensure_map_entry_class( + &entry_internal_name, + Some(rendered_entry_name), + map_def.k, + map_def.v, + ); + + TypeIR::list(TypeIR::class(entry_internal_name)) + } + }, + _ => self.build_type_ir(shape), + } + } + + fn ensure_map_entry_class( + &mut self, + internal_name: &str, + rendered_name: Option, + key_shape: &'static Shape, + value_shape: &'static Shape, + ) { + let key_ir = self.build_type_ir(key_shape); + let value_ir = self.build_type_ir(value_shape); + + let class = Class { + name: match rendered_name { + Some(rendered) => name_with_optional_alias(internal_name.to_string(), rendered), + None => Name::new(internal_name.to_string()), + }, + description: None, + namespace: StreamingMode::NonStreaming, + fields: vec![ + (Name::new("key".to_string()), key_ir, None, false), + (Name::new("value".to_string()), value_ir, None, false), + ], + constraints: Vec::new(), + streaming_behavior: Default::default(), + }; + + self.register_class(class); + } + + fn register_class(&mut self, class: Class) { + self.registry.register_class(class); + } + + /// Generate a unique internal name for a type. + fn generate_internal_name(&mut self, shape: &'static Shape) -> String { + let base = internal_name_for_shape(shape); + + if let Some(existing_shape_id) = self.used_internal_names.get(&base) { + if *existing_shape_id == shape.id { + return base; + } + + let mut candidate; + loop { + self.name_counter += 1; + candidate = format!("{base}__{}", self.name_counter); + if !self.used_internal_names.contains_key(&candidate) { + break; + } + } + self.used_internal_names.insert(candidate.clone(), shape.id); + return candidate; + } + + self.used_internal_names.insert(base.clone(), shape.id); + base + } + + /// Finalize and produce the OutputFormatContent. + fn into_output_format(self, target: TypeIR) -> OutputFormatContent { + self.registry.build(target) + } +} + +fn field_int_repr(field: &Field) -> Option { + let repr = bamltype_int_repr(field.attributes)?; + + match repr { + "string" => Some(IntReprMode::String), + "i64" => Some(IntReprMode::I64), + _ => None, + } +} + +fn field_map_key_repr(field: &Field) -> Option { + let repr = bamltype_map_key_repr(field.attributes)?; + + match repr { + "string" => Some(MapKeyReprMode::String), + "pairs" => Some(MapKeyReprMode::Pairs), + _ => None, + } +} + +fn bamltype_internal_name(attrs: &'static [Attr]) -> Option<&'static str> { + bamltype_attr_static_str(attrs, "internal_name") +} + +fn bamltype_int_repr(attrs: &'static [Attr]) -> Option<&'static str> { + bamltype_attr_static_str(attrs, "int_repr") +} + +fn bamltype_map_key_repr(attrs: &'static [Attr]) -> Option<&'static str> { + bamltype_attr_static_str(attrs, "map_key_repr") +} + +fn bamltype_attr_static_str(attrs: &'static [Attr], key: &str) -> Option<&'static str> { + for attr in attrs { + if attr.ns != Some("bamltype") || attr.key != key { + continue; + } + + if let Some(value) = attr.get_as::<&'static str>() { + return Some(*value); + } + } + + None +} + +fn map_entry_names(ctx: &MapEntryContext) -> (String, String) { + let suffix = match &ctx.variant_name { + Some(variant_name) => format!("{variant_name}__{}__Entry", ctx.field_name), + None => format!("{}__Entry", ctx.field_name), + }; + + let internal_name = format!("{}::{suffix}", ctx.owner_internal_name); + + let rendered = match &ctx.variant_rendered { + Some(variant_rendered) => { + format!("{variant_rendered}{}Entry", ctx.rendered_field) + } + None => format!("{}Entry", ctx.rendered_field), + }; + + (internal_name, rendered) +} + +fn constraints_from_attrs(attrs: &'static [Attr]) -> Vec { + let mut out = Vec::new(); + + for attr in attrs { + if attr.ns != Some("bamltype") { + continue; + } + + if attr.key == "check" { + let Some(ext_attr) = attr.get_as::() else { + continue; + }; + if let facet_ext::Attr::Check(payload) = ext_attr { + out.push(Constraint::new_check(payload.label, payload.expr)); + } + } else if attr.key == "assert" { + let Some(ext_attr) = attr.get_as::() else { + continue; + }; + if let facet_ext::Attr::Assert(payload) = ext_attr { + out.push(Constraint::new_assert(payload.label, payload.expr)); + } + } + } + + out +} + +fn doc_to_description(doc: &'static [&'static str]) -> Option { + if doc.is_empty() { + return None; + } + + Some( + doc.iter() + .map(|line| line.trim()) + .collect::>() + .join("\n"), + ) +} + +fn name_with_optional_alias(real_name: String, rendered_name: String) -> Name { + if real_name == rendered_name { + Name::new(real_name) + } else { + Name::new_with_alias(real_name, Some(rendered_name)) + } +} diff --git a/crates/bamltype/tests/contract_bridge_oracle.rs b/crates/bamltype/tests/contract_bridge_oracle.rs new file mode 100644 index 00000000..2033c869 --- /dev/null +++ b/crates/bamltype/tests/contract_bridge_oracle.rs @@ -0,0 +1,1202 @@ +use std::collections::{BTreeMap, HashMap}; +use std::rc::Rc; +use std::sync::Arc; + +use baml_bridge as legacy; +use baml_bridge::baml_types::{BamlValue, LiteralValue, StreamingMode, TypeIR}; +use bamltype as facet_runtime; + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(name = "ContractUser")] +#[baml(internal_name = "contract::User")] +struct BridgeUser { + #[baml(alias = "fullName")] + name: String, + age: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(name = "ContractUser")] +#[baml(internal_name = "contract::User")] +struct FacetUser { + #[baml(alias = "fullName")] + name: String, + age: i64, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(name = "ContractShape")] +#[baml(internal_name = "contract::Shape")] +#[baml(tag = "kind")] +enum BridgeShape { + Circle { radius: f64 }, + Square { side: f64 }, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(name = "ContractShape")] +#[baml(internal_name = "contract::Shape")] +#[baml(tag = "kind")] +enum FacetShape { + Circle { radius: f64 }, + Square { side: f64 }, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::Checked")] +struct BridgeChecked { + #[baml(check(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::Checked")] +struct FacetChecked { + #[baml(check(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::Asserted")] +struct BridgeAsserted { + #[baml(assert(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::Asserted")] +struct FacetAsserted { + #[baml(assert(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::MapPairs")] +struct BridgeMapPairs { + #[baml(map_key_repr = "pairs")] + values: HashMap, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::MapPairs")] +struct FacetMapPairs { + #[baml(map_key_repr = "pairs")] + values: HashMap, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(name = "OrderB")] +#[baml(internal_name = "contract::OrderB")] +struct BridgeOrderB { + value: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(name = "OrderB")] +#[baml(internal_name = "contract::OrderB")] +struct FacetOrderB { + value: i64, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(name = "OrderA")] +#[baml(internal_name = "contract::OrderA")] +struct BridgeOrderA { + value: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(name = "OrderA")] +#[baml(internal_name = "contract::OrderA")] +struct FacetOrderA { + value: i64, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(name = "OrderRoot")] +#[baml(internal_name = "contract::OrderRoot")] +struct BridgeOrderRoot { + b: BridgeOrderB, + a: BridgeOrderA, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(name = "OrderRoot")] +#[baml(internal_name = "contract::OrderRoot")] +struct FacetOrderRoot { + b: FacetOrderB, + a: FacetOrderA, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::BigIntString")] +struct BridgeBigIntString { + #[baml(int_repr = "string")] + id: u64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::BigIntString")] +struct FacetBigIntString { + #[baml(int_repr = "string")] + id: u64, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::Unsigned32")] +struct BridgeUnsigned32 { + value: u32, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::Unsigned32")] +struct FacetUnsigned32 { + value: u32, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::FuzzUser")] +struct BridgeFuzzUser { + name: String, + age: u32, + nickname: Option, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::FuzzUser")] +struct FacetFuzzUser { + name: String, + age: u32, + nickname: Option, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::FuzzExplain")] +struct BridgeFuzzExplain { + values: HashMap, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::FuzzExplain")] +struct FacetFuzzExplain { + values: HashMap, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::ColorAlias")] +enum BridgeColorAlias { + Red, + #[baml(alias = "green")] + Green, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::ColorAlias")] +enum FacetColorAlias { + Red, + #[baml(alias = "green")] + Green, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(name = "ContractDocShape")] +#[baml(internal_name = "contract::DocShape")] +#[baml(tag = "type")] +enum BridgeDocShape { + /// A circle, defined by its radius. + Circle { + /// Radius in meters. + radius: f64, + }, + Rectangle { + width: f64, + height: f64, + }, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(name = "ContractDocShape")] +#[baml(internal_name = "contract::DocShape")] +#[baml(tag = "type")] +enum FacetDocShape { + /// A circle, defined by its radius. + Circle { + /// Radius in meters. + radius: f64, + }, + Rectangle { + width: f64, + height: f64, + }, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::BigIntOption")] +struct BridgeBigIntOption { + #[baml(int_repr = "i64")] + id: Option, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::BigIntOption")] +struct FacetBigIntOption { + #[baml(int_repr = "i64")] + id: Option, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::MapKeys")] +struct BridgeMapKeys { + #[baml(map_key_repr = "string")] + values: HashMap, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::MapKeys")] +struct FacetMapKeys { + #[baml(map_key_repr = "string")] + values: HashMap, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::MapKeysOption")] +struct BridgeMapKeysOption { + #[baml(map_key_repr = "string")] + values: Option>, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::MapKeysOption")] +struct FacetMapKeysOption { + #[baml(map_key_repr = "string")] + values: Option>, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::Node")] +struct BridgeNode { + value: i64, + next: Option>, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::Node")] +struct FacetNode { + value: i64, + next: Option>, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(rename_all = "camelCase")] +#[baml(internal_name = "contract::RenameAllUser")] +struct BridgeRenameAllUser { + full_name: String, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(rename_all = "camelCase")] +#[baml(internal_name = "contract::RenameAllUser")] +struct FacetRenameAllUser { + full_name: String, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::RoundtripStruct")] +struct BridgeRoundtripStruct { + name: String, + count: i32, + tags: Vec, + meta: Option, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::RoundtripStruct")] +struct FacetRoundtripStruct { + name: String, + count: i32, + tags: Vec, + meta: Option, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::RoundtripUnitEnum")] +enum BridgeRoundtripUnitEnum { + Alpha, + #[baml(alias = "beta")] + Beta, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::RoundtripUnitEnum")] +enum FacetRoundtripUnitEnum { + Alpha, + #[baml(alias = "beta")] + Beta, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::RoundtripDataEnum")] +#[baml(tag = "kind")] +enum BridgeRoundtripDataEnum { + Message { body: String, count: i64 }, + Empty, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::RoundtripDataEnum")] +#[baml(tag = "kind")] +enum FacetRoundtripDataEnum { + Message { body: String, count: i64 }, + Empty, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::NestedStruct")] +struct BridgeNestedStruct { + title: String, + items: Option>, + metadata: HashMap>, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::NestedStruct")] +struct FacetNestedStruct { + title: String, + items: Option>, + metadata: HashMap>, +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::AsUnionColor")] +#[baml(as_union)] +enum BridgeAsUnionColor { + Red, + #[baml(alias = "green")] + Green, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::AsUnionColor")] +#[baml(as_union)] +enum FacetAsUnionColor { + Red, + #[baml(alias = "green")] + Green, +} + +struct LegacyI64ObjectAdapter; + +impl legacy::BamlAdapter for LegacyI64ObjectAdapter { + fn type_ir() -> TypeIR { + TypeIR::class("AdapterI64Wrapper") + } + + fn register(reg: &mut legacy::Registry) { + use legacy::internal_baml_jinja::types::{Class, Name}; + + if !reg.mark_type("AdapterI64Wrapper") { + return; + } + + reg.register_class(Class { + name: Name::new("AdapterI64Wrapper".to_string()), + description: None, + namespace: StreamingMode::NonStreaming, + fields: vec![( + Name::new("value".to_string()), + TypeIR::string(), + None, + false, + )], + constraints: Vec::new(), + streaming_behavior: Default::default(), + }); + } + + fn try_from_baml( + value: BamlValue, + mut path: Vec, + ) -> Result { + let map = match value { + BamlValue::Class(_, fields) | BamlValue::Map(fields) => fields, + other => { + return Err(legacy::BamlConvertError::new( + path, + "object", + format!("{other:?}"), + "expected object adapter payload", + )); + } + }; + + let raw_value = map.get("value").ok_or_else(|| { + legacy::BamlConvertError::new( + path.clone(), + "value", + "", + "missing required field", + ) + })?; + path.push("value".to_string()); + + match raw_value { + BamlValue::String(s) => s.parse::().map_err(|_| { + legacy::BamlConvertError::new(path.clone(), "i64", s.clone(), "failed to parse i64") + }), + BamlValue::Int(i) => Ok(*i), + other => Err(legacy::BamlConvertError::new( + path, + "i64", + format!("{other:?}"), + "expected value field to be string or int", + )), + } + } +} + +struct FacetI64ObjectAdapter; + +impl facet_runtime::BamlAdapter for FacetI64ObjectAdapter { + fn type_ir() -> TypeIR { + TypeIR::class("AdapterI64Wrapper") + } + + fn register(reg: &mut facet_runtime::Registry) { + use facet_runtime::internal_baml_jinja::types::{Class, Name}; + + if !reg.mark_type("AdapterI64Wrapper") { + return; + } + + reg.register_class(Class { + name: Name::new("AdapterI64Wrapper".to_string()), + description: None, + namespace: StreamingMode::NonStreaming, + fields: vec![( + Name::new("value".to_string()), + TypeIR::string(), + None, + false, + )], + constraints: Vec::new(), + streaming_behavior: Default::default(), + }); + } + + fn try_from_baml( + value: BamlValue, + mut path: Vec, + ) -> Result { + let map = match value { + BamlValue::Class(_, fields) | BamlValue::Map(fields) => fields, + other => { + return Err(facet_runtime::BamlConvertError::new( + path, + "object", + format!("{other:?}"), + "expected object adapter payload", + )); + } + }; + + let raw_value = map.get("value").ok_or_else(|| { + facet_runtime::BamlConvertError::new( + path.clone(), + "value", + "", + "missing required field", + ) + })?; + path.push("value".to_string()); + + match raw_value { + BamlValue::String(s) => s.parse::().map_err(|_| { + facet_runtime::BamlConvertError::new( + path.clone(), + "i64", + s.clone(), + "failed to parse i64", + ) + }), + BamlValue::Int(i) => Ok(*i), + other => Err(facet_runtime::BamlConvertError::new( + path, + "i64", + format!("{other:?}"), + "expected value field to be string or int", + )), + } + } +} + +#[derive(Debug, Clone, PartialEq, legacy::BamlType)] +#[baml(internal_name = "contract::WithAdapter")] +struct BridgeWithAdapter { + #[baml(with = "LegacyI64ObjectAdapter")] + id: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(internal_name = "contract::WithAdapter")] +struct FacetWithAdapter { + #[baml(with = "FacetI64ObjectAdapter")] + id: i64, +} + +mod legacy_collision_a { + use super::legacy; + + #[derive(Debug, Clone, PartialEq, legacy::BamlType)] + pub struct User { + pub name: String, + } +} + +mod legacy_collision_b { + use super::legacy; + + #[derive(Debug, Clone, PartialEq, legacy::BamlType)] + pub struct User { + pub name: String, + } +} + +mod facet_collision_a { + #[derive(Debug, Clone, PartialEq)] + #[bamltype::BamlType] + pub struct User { + pub name: String, + } +} + +mod facet_collision_b { + #[derive(Debug, Clone, PartialEq)] + #[bamltype::BamlType] + pub struct User { + pub name: String, + } +} + +fn sorted_union_string_literals(type_ir: TypeIR) -> Vec { + let TypeIR::Union(union, _) = type_ir else { + panic!("expected union type IR"); + }; + + let mut literals = union + .iter_skip_null() + .into_iter() + .map(|item| match item { + TypeIR::Literal(LiteralValue::String(value), _) => value.clone(), + other => panic!("expected string literal in union, got {other:?}"), + }) + .collect::>(); + literals.sort(); + literals +} + +#[test] +fn contract_render_schema_default_matches_legacy() { + let old = legacy::render_schema::(legacy::RenderOptions::default()) + .expect("legacy render") + .unwrap_or_default(); + let new = facet_runtime::render_schema::(facet_runtime::RenderOptions::default()) + .expect("facet render") + .unwrap_or_default(); + + assert_eq!(old, new); +} + +#[test] +fn contract_render_schema_hoisted_matches_legacy() { + let opts = legacy::RenderOptions::hoist_classes(legacy::HoistClasses::All); + let old = legacy::render_schema::(opts) + .expect("legacy render") + .unwrap_or_default(); + + let opts = facet_runtime::RenderOptions::hoist_classes(facet_runtime::HoistClasses::All); + let new = facet_runtime::render_schema::(opts) + .expect("facet render") + .unwrap_or_default(); + + assert_eq!(old, new); +} + +#[test] +fn contract_render_schema_ordering_matches_legacy() { + let opts = legacy::RenderOptions::hoist_classes(legacy::HoistClasses::All); + let old = legacy::render_schema::(opts) + .expect("legacy render") + .unwrap_or_default(); + + let opts = facet_runtime::RenderOptions::hoist_classes(facet_runtime::HoistClasses::All); + let new = facet_runtime::render_schema::(opts) + .expect("facet render") + .unwrap_or_default(); + + assert_eq!(old, new); +} + +#[test] +fn contract_parse_envelope_matches_legacy() { + let raw = r#"{ "kind": "Circle", "radius": 2.5 }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + + assert_eq!(old.baml_value, new.baml_value); + assert_eq!(format!("{:?}", old.flags), format!("{:?}", new.flags)); + assert_eq!(format!("{:?}", old.checks), format!("{:?}", new.checks)); + assert_eq!( + old.explanations + .iter() + .map(ToString::to_string) + .collect::>(), + new.explanations + .iter() + .map(ToString::to_string) + .collect::>() + ); +} + +#[test] +fn contract_parse_streaming_mode_matches_legacy() { + let raw = r#"{ "name": "Ada", "age": 36 }"#; + let old = legacy::parse_llm_output::(raw, false).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, false).expect("facet parse"); + + assert_eq!(old.value.name, new.value.name); + assert_eq!(old.value.age, new.value.age); + assert_eq!(old.value.nickname, new.value.nickname); + assert_eq!(old.baml_value, new.baml_value); + assert_eq!(format!("{:?}", old.flags), format!("{:?}", new.flags)); + assert_eq!(format!("{:?}", old.checks), format!("{:?}", new.checks)); + assert_eq!( + old.explanations + .iter() + .map(ToString::to_string) + .collect::>(), + new.explanations + .iter() + .map(ToString::to_string) + .collect::>() + ); +} + +#[test] +fn contract_int_repr_string_matches_legacy() { + let raw = r#"{ "id": "18446744073709551615" }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + + assert_eq!(old.value.id, new.value.id); + assert_eq!(old.baml_value, new.baml_value); +} + +#[test] +fn contract_map_key_pairs_parse_and_registration_match_legacy() { + let raw = r#"{ "values": [ { "key": 1, "value": "a" }, { "key": 2, "value": "b" } ] }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!(old.baml_value, new.baml_value); + + let old_entry = format!( + "{}::values__Entry", + ::baml_internal_name() + ); + let old_class = ::baml_output_format() + .classes + .get(&(old_entry, StreamingMode::NonStreaming)) + .expect("legacy entry class"); + + let new_entry = format!( + "{}::values__Entry", + ::baml_internal_name() + ); + let new_class = ::baml_output_format() + .classes + .get(&(new_entry, StreamingMode::NonStreaming)) + .expect("facet entry class"); + + assert_eq!( + old_class.name.rendered_name(), + new_class.name.rendered_name() + ); +} + +#[test] +fn contract_check_results_match_legacy() { + let raw = r#"{ "value": 3 }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + + assert_eq!(format!("{:?}", old.checks), format!("{:?}", new.checks)); +} + +#[test] +fn contract_assert_error_shape_matches_legacy() { + let raw = r#"{ "value": -1 }"#; + let old = legacy::parse_llm_output::(raw, true).expect_err("legacy err"); + let new = facet_runtime::parse_llm_output::(raw, true).expect_err("facet err"); + + let old_failed = match old { + legacy::BamlParseError::ConstraintAssertsFailed { failed } => failed, + other => panic!("legacy expected assert failure, got {other:?}"), + }; + let new_failed = match new { + facet_runtime::BamlParseError::ConstraintAssertsFailed { failed } => failed, + other => panic!("facet expected assert failure, got {other:?}"), + }; + + assert_eq!(format!("{:?}", old_failed), format!("{:?}", new_failed)); +} + +#[test] +fn contract_markdown_flags_and_explanations_match_legacy() { + let raw = "```json\n{ \"name\": \"Ada\", \"age\": 36 }\n```"; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!(format!("{:?}", old.flags), format!("{:?}", new.flags)); + + let raw = r#"{ "values": { "ok": 1, "bad": "oops" } }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!( + old.explanations + .iter() + .map(ToString::to_string) + .collect::>(), + new.explanations + .iter() + .map(ToString::to_string) + .collect::>() + ); +} + +#[test] +fn contract_unit_enum_alias_parse_matches_legacy() { + let raw = r#""green""#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + + assert_eq!(old.baml_value, new.baml_value); + assert!(matches!(old.value, BridgeColorAlias::Green)); + assert!(matches!(new.value, FacetColorAlias::Green)); +} + +#[test] +fn contract_as_union_type_ir_and_parse_match_legacy() { + let old_literals = + sorted_union_string_literals(::baml_type_ir()); + let new_literals = sorted_union_string_literals( + ::baml_type_ir(), + ); + assert_eq!(old_literals, new_literals); + + let raw = r#""green""#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + + assert_eq!(old.baml_value, new.baml_value); + assert!(matches!(old.value, BridgeAsUnionColor::Green)); + assert!(matches!(new.value, FacetAsUnionColor::Green)); +} + +#[test] +fn contract_data_enum_docs_render_match_legacy() { + let old = legacy::render_schema::(legacy::RenderOptions::default()) + .expect("legacy render") + .unwrap_or_default(); + let new = + facet_runtime::render_schema::(facet_runtime::RenderOptions::default()) + .expect("facet render") + .unwrap_or_default(); + + assert_eq!(old, new); +} + +#[test] +fn contract_int_repr_option_matches_legacy() { + let raw = r#"{ "id": 42 }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!(old.baml_value, new.baml_value); + assert_eq!(old.value.id, new.value.id); + + let raw = r#"{ "id": null }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!(old.baml_value, new.baml_value); + assert_eq!(old.value.id, new.value.id); +} + +#[test] +fn contract_map_key_repr_string_and_option_match_legacy() { + let raw = r#"{ "values": { "1": "a", "2": "b" } }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!(old.baml_value, new.baml_value); + assert_eq!(old.value.values, new.value.values); + + let raw = r#"{ "values": { "10": "x" } }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = + facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!(old.baml_value, new.baml_value); + assert_eq!(old.value.values, new.value.values); +} + +#[test] +fn contract_rename_all_matches_legacy() { + let raw = r#"{ "fullName": "Ada" }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = + facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + + assert_eq!(old.baml_value, new.baml_value); + assert_eq!(old.value.full_name, new.value.full_name); +} + +#[test] +fn contract_recursion_class_detection_matches_legacy() { + let old_of = ::baml_output_format(); + let new_of = ::baml_output_format(); + + assert_eq!(old_of.recursive_classes, new_of.recursive_classes); + assert!( + old_of + .recursive_classes + .contains(::baml_internal_name()) + ); + assert!( + new_of + .recursive_classes + .contains(::baml_internal_name()) + ); +} + +#[test] +fn contract_internal_name_collision_behavior_matches_legacy() { + let old_a = ::baml_internal_name(); + let old_b = ::baml_internal_name(); + let new_a = ::baml_internal_name(); + let new_b = ::baml_internal_name(); + + assert_ne!(old_a, old_b); + assert_ne!(new_a, new_b); + + let old_a_class = ::baml_output_format() + .classes + .get(&(old_a.to_string(), StreamingMode::NonStreaming)) + .expect("legacy class a missing"); + let old_b_class = ::baml_output_format() + .classes + .get(&(old_b.to_string(), StreamingMode::NonStreaming)) + .expect("legacy class b missing"); + let new_a_class = ::baml_output_format() + .classes + .get(&(new_a.to_string(), StreamingMode::NonStreaming)) + .expect("facet class a missing"); + let new_b_class = ::baml_output_format() + .classes + .get(&(new_b.to_string(), StreamingMode::NonStreaming)) + .expect("facet class b missing"); + + assert_eq!(old_a_class.name.rendered_name(), "User"); + assert_eq!(old_b_class.name.rendered_name(), "User"); + assert_eq!(new_a_class.name.rendered_name(), "User"); + assert_eq!(new_b_class.name.rendered_name(), "User"); +} + +#[test] +fn contract_with_adapter_schema_and_parse_match_legacy() { + let old_of = ::baml_output_format(); + let new_of = ::baml_output_format(); + + let old_adapter = old_of + .classes + .get(&("AdapterI64Wrapper".to_string(), StreamingMode::NonStreaming)) + .expect("legacy adapter class missing"); + let new_adapter = new_of + .classes + .get(&("AdapterI64Wrapper".to_string(), StreamingMode::NonStreaming)) + .expect("facet adapter class missing"); + assert_eq!(old_adapter.name.real_name(), new_adapter.name.real_name()); + assert_eq!( + old_adapter + .fields + .iter() + .map(|(name, _, _, _)| name.real_name()) + .collect::>(), + new_adapter + .fields + .iter() + .map(|(name, _, _, _)| name.real_name()) + .collect::>() + ); + + let old_owner_name = ::baml_internal_name().to_string(); + let new_owner_name = + ::baml_internal_name().to_string(); + let old_owner = old_of + .classes + .get(&(old_owner_name, StreamingMode::NonStreaming)) + .expect("legacy owner class missing"); + let new_owner = new_of + .classes + .get(&(new_owner_name, StreamingMode::NonStreaming)) + .expect("facet owner class missing"); + let (_, old_field_ir, _, _) = old_owner + .fields + .iter() + .find(|(name, _, _, _)| name.real_name() == "id") + .expect("legacy id field missing"); + let (_, new_field_ir, _, _) = new_owner + .fields + .iter() + .find(|(name, _, _, _)| name.real_name() == "id") + .expect("facet id field missing"); + assert_eq!(format!("{old_field_ir:?}"), format!("{new_field_ir:?}")); + + let raw = r#"{ "id": { "value": "9223372036854775807" } }"#; + let old = legacy::parse_llm_output::(raw, true).expect("legacy parse"); + let new = facet_runtime::parse_llm_output::(raw, true).expect("facet parse"); + assert_eq!(old.value.id, new.value.id); + assert_eq!(old.baml_value, new.baml_value); +} + +fn assert_cross_runtime_roundtrip(old_value: Old, new_value: New) +where + Old: Clone + std::fmt::Debug + PartialEq + legacy::ToBamlValue + legacy::BamlValueConvert, + New: Clone + + std::fmt::Debug + + PartialEq + + facet_runtime::compat::ToBamlValue + + facet_runtime::compat::BamlValueConvert, +{ + let old_baml = legacy::ToBamlValue::to_baml_value(&old_value); + let new_baml = facet_runtime::compat::ToBamlValue::to_baml_value(&new_value); + assert_eq!(old_baml, new_baml); + + let old_back = ::try_from_baml_value(old_baml, Vec::new()) + .expect("legacy roundtrip"); + let new_back = + ::try_from_baml_value(new_baml, Vec::new()) + .expect("facet roundtrip"); + assert_eq!(old_back, old_value); + assert_eq!(new_back, new_value); +} + +fn assert_default_adapter_parity(value: T) +where + T: Clone + + std::fmt::Debug + + PartialEq + + legacy::ToBamlValue + + legacy::BamlValueConvert + + facet_runtime::compat::ToBamlValue + + facet_runtime::compat::BamlValueConvert, +{ + let old_baml = legacy::ToBamlValue::to_baml_value(&value); + let new_baml = facet_runtime::compat::ToBamlValue::to_baml_value(&value); + assert_eq!(old_baml, new_baml); + + let old_back = ::try_from_baml_value(old_baml, Vec::new()) + .expect("legacy convert"); + let new_back = + ::try_from_baml_value(new_baml, Vec::new()) + .expect("facet convert"); + assert_eq!(old_back, new_back); + assert_eq!(old_back, value); +} + +#[test] +fn contract_default_adapter_complex_roundtrips_match_legacy() { + let old_struct = BridgeRoundtripStruct { + name: "example".to_string(), + count: 7, + tags: vec!["tag".to_string()], + meta: Some("meta".to_string()), + }; + let new_struct = FacetRoundtripStruct { + name: "example".to_string(), + count: 7, + tags: vec!["tag".to_string()], + meta: Some("meta".to_string()), + }; + assert_cross_runtime_roundtrip(old_struct, new_struct); + + assert_cross_runtime_roundtrip(BridgeRoundtripUnitEnum::Beta, FacetRoundtripUnitEnum::Beta); + + assert_cross_runtime_roundtrip( + BridgeRoundtripDataEnum::Message { + body: "hello".to_string(), + count: 3, + }, + FacetRoundtripDataEnum::Message { + body: "hello".to_string(), + count: 3, + }, + ); + + let mut old_metadata = HashMap::new(); + old_metadata.insert("alpha".to_string(), Some(1)); + old_metadata.insert("beta".to_string(), None); + + let mut new_metadata = HashMap::new(); + new_metadata.insert("alpha".to_string(), Some(1)); + new_metadata.insert("beta".to_string(), None); + + let old_nested = BridgeNestedStruct { + title: "nested".to_string(), + items: Some(vec![BridgeRoundtripStruct { + name: "child".to_string(), + count: 2, + tags: vec!["x".to_string(), "y".to_string()], + meta: None, + }]), + metadata: old_metadata, + }; + let new_nested = FacetNestedStruct { + title: "nested".to_string(), + items: Some(vec![FacetRoundtripStruct { + name: "child".to_string(), + count: 2, + tags: vec!["x".to_string(), "y".to_string()], + meta: None, + }]), + metadata: new_metadata, + }; + assert_cross_runtime_roundtrip(old_nested, new_nested); +} + +#[test] +fn contract_default_adapter_roundtrips_match_legacy() { + assert_default_adapter_parity("hello".to_string()); + assert_default_adapter_parity(true); + assert_default_adapter_parity(123i32); + assert_default_adapter_parity(-99i64); + assert_default_adapter_parity(3.5f32); + assert_default_adapter_parity(9.75f64); + + assert_default_adapter_parity(Some(42i32)); + assert_default_adapter_parity(None::); + assert_default_adapter_parity(vec!["a".to_string(), "b".to_string()]); + assert_default_adapter_parity(Box::new("boxed".to_string())); + assert_default_adapter_parity(Arc::new("arc".to_string())); + assert_default_adapter_parity(Rc::new("rc".to_string())); + + let mut hm = HashMap::new(); + hm.insert("answer".to_string(), 42i32); + assert_default_adapter_parity(hm); + + let mut bt = BTreeMap::new(); + bt.insert("left".to_string(), 1i64); + bt.insert("right".to_string(), 2i64); + assert_default_adapter_parity(bt); +} + +#[test] +fn contract_default_adapter_error_messages_match_legacy() { + let invalid = BamlValue::Class( + "contract::Unsigned32".to_string(), + [("value".to_string(), BamlValue::Int(-1))] + .into_iter() + .collect(), + ); + + let old_err = ::try_from_baml_value( + invalid.clone(), + Vec::new(), + ) + .expect_err("legacy should fail"); + let new_err = + ::try_from_baml_value( + invalid, + Vec::new(), + ) + .expect_err("facet should fail"); + + assert_eq!(old_err.expected, new_err.expected); + assert_eq!(old_err.path, new_err.path); + assert_eq!(old_err.to_string(), new_err.to_string()); +} + +#[test] +fn contract_direct_integer_string_conversion_matches_legacy() { + let old_err = ::try_from_baml_value( + BamlValue::String("123".into()), + Vec::new(), + ) + .expect_err("legacy should reject string->int direct conversion"); + let new_err = ::try_from_baml_value( + BamlValue::String("123".into()), + Vec::new(), + ) + .expect_err("facet should reject string->int direct conversion"); + + assert_eq!(old_err.expected, new_err.expected); +} + +#[test] +fn contract_direct_map_pairs_conversion_matches_legacy() { + let pair_entry = BamlValue::Map( + [ + ("key".to_string(), BamlValue::String("k".to_string())), + ("value".to_string(), BamlValue::Int(1)), + ] + .into_iter() + .collect(), + ); + let raw = BamlValue::List(vec![pair_entry]); + + let old_err = as legacy::BamlValueConvert>::try_from_baml_value( + raw.clone(), + Vec::new(), + ) + .expect_err("legacy should reject map pair-list direct conversion"); + let new_err = + as facet_runtime::compat::BamlValueConvert>::try_from_baml_value( + raw, + Vec::new(), + ) + .expect_err("facet should reject map pair-list direct conversion"); + + assert_eq!(old_err.expected, new_err.expected); +} + +#[test] +fn contract_schema_fingerprint_matches_legacy() { + let old_of = ::baml_output_format(); + let new_of = ::baml_output_format(); + + let old_fp = legacy::schema_fingerprint(old_of, legacy::RenderOptions::default()) + .expect("legacy fingerprint"); + let new_fp = facet_runtime::schema_fingerprint(new_of, facet_runtime::RenderOptions::default()) + .expect("facet fingerprint"); + + assert_eq!(old_fp, new_fp); +} diff --git a/crates/bamltype/tests/contract_bridge_ui_messages.rs b/crates/bamltype/tests/contract_bridge_ui_messages.rs new file mode 100644 index 00000000..b8fb307c --- /dev/null +++ b/crates/bamltype/tests/contract_bridge_ui_messages.rs @@ -0,0 +1,36 @@ +use std::fs; +use std::path::Path; + +#[test] +fn contract_ui_error_messages_match_legacy() { + let root = Path::new(env!("CARGO_MANIFEST_DIR")); + let legacy_ui = root.join("../baml-bridge/tests/ui"); + let facet_ui = root.join("tests/ui"); + + let mut compared = 0usize; + for entry in fs::read_dir(&legacy_ui).expect("read legacy ui dir") { + let entry = entry.expect("read dir entry"); + let path = entry.path(); + + if path.extension().and_then(|ext| ext.to_str()) != Some("stderr") { + continue; + } + + let file_name = path + .file_name() + .and_then(|name| name.to_str()) + .expect("stderr file name"); + let facet_path = facet_ui.join(file_name); + assert!( + facet_path.exists(), + "missing facet stderr fixture for {file_name}" + ); + + let legacy = fs::read_to_string(&path).expect("read legacy stderr"); + let facet = fs::read_to_string(&facet_path).expect("read facet stderr"); + assert_eq!(legacy, facet, "stderr mismatch for {file_name}"); + compared += 1; + } + + assert!(compared > 0, "no stderr fixtures compared"); +} diff --git a/crates/bamltype/tests/integration.rs b/crates/bamltype/tests/integration.rs new file mode 100644 index 00000000..901314a3 --- /dev/null +++ b/crates/bamltype/tests/integration.rs @@ -0,0 +1,1190 @@ +//! Integration tests for bamltype. + +use std::collections::{BTreeMap, HashMap}; +use std::sync::Arc; + +use baml_types::{BamlValue, ConstraintLevel, LiteralValue, StreamingMode, TypeIR}; +use bamltype::{ + BamlSchema, + compat::{BamlAdapter, BamlTypeInternal, Registry}, + from_baml_value, from_baml_value_with_flags, parse, render_schema_default, to_baml_value, +}; +use indexmap::IndexMap; + +/// A simple response struct for testing. +#[bamltype::BamlType] +struct Response { + /// The user's name + name: String, + /// Age in years + age: u32, + /// Whether the user is active + active: bool, +} + +/// A nested struct for testing. +#[bamltype::BamlType] +struct UserProfile { + /// User information + user: Response, + /// Optional email address + email: Option, + /// List of tags + tags: Vec, +} + +#[test] +fn test_simple_struct_schema() { + let schema = render_schema_default::().expect("Should render schema"); + + // The schema should contain the field names + assert!( + schema.contains("name"), + "Schema should mention 'name' field" + ); + assert!(schema.contains("age"), "Schema should mention 'age' field"); + assert!( + schema.contains("active"), + "Schema should mention 'active' field" + ); +} + +#[test] +fn test_nested_struct_schema() { + let schema = render_schema_default::().expect("Should render schema"); + + println!("UserProfile schema:\n{}", schema); + + // The schema should contain nested type information + assert!( + schema.contains("user"), + "Schema should mention 'user' field" + ); + assert!( + schema.contains("email"), + "Schema should mention 'email' field" + ); + assert!( + schema.contains("tags"), + "Schema should mention 'tags' field" + ); +} + +#[test] +fn test_schema_bundle_caching() { + // Get schema bundle twice - should return the same static reference + let bundle1 = Response::baml_schema(); + let bundle2 = Response::baml_schema(); + + // These should be the same pointer (cached) + assert!( + std::ptr::eq(bundle1, bundle2), + "Schema bundle should be cached" + ); +} + +#[test] +fn test_schema_output_format() { + let schema = render_schema_default::().expect("Should render schema"); + + // Print the schema for debugging + println!("Generated schema:\n{}", schema); + + // Schema should not be empty + assert!(!schema.is_empty(), "Schema should not be empty"); +} + +#[test] +fn test_parse_llm_output() { + // Simulate LLM output (with markdown code block, which jsonish handles) + let llm_output = r#"```json +{ + "name": "Alice", + "age": 30, + "active": true +} +```"#; + + let parsed = parse::(llm_output).expect("Should parse LLM output"); + + // BamlValueWithFlags should have the right structure + println!("Parsed: {:?}", parsed); +} + +#[test] +fn test_parse_nested_struct() { + let llm_output = r#"{ + "user": { + "name": "Bob", + "age": 25, + "active": false + }, + "email": "bob@example.com", + "tags": ["admin", "user"] + }"#; + + let parsed = parse::(llm_output).expect("Should parse nested struct"); + println!("Parsed nested: {:?}", parsed); +} + +#[test] +fn test_parse_with_optional_null() { + // First check the schema to see how email is typed + let schema = render_schema_default::().expect("render"); + println!("Schema for optional test:\n{}", schema); + + let llm_output = r#"{ + "user": {"name": "Charlie", "age": 35, "active": true}, + "email": null, + "tags": [] + }"#; + + let parsed = parse::(llm_output).expect("Should handle null optional"); + println!("Parsed with null: {:?}", parsed); +} + +// ============================================================================ +// BamlValue conversion tests +// ============================================================================ + +#[test] +fn test_from_baml_value_simple_struct() { + let mut fields = IndexMap::new(); + fields.insert("name".to_string(), BamlValue::String("Alice".into())); + fields.insert("age".to_string(), BamlValue::Int(30)); + fields.insert("active".to_string(), BamlValue::Bool(true)); + + let baml_value = BamlValue::Class("Response".into(), fields); + + let response: Response = from_baml_value(baml_value).expect("Should convert to Response"); + + assert_eq!(response.name, "Alice"); + assert_eq!(response.age, 30); + assert!(response.active); +} + +#[test] +fn test_to_baml_value_simple_struct() { + let name = "Bob".to_string(); + let baml = to_baml_value(&name).expect("Should convert String"); + assert_eq!(baml, BamlValue::String("Bob".into())); + + let num: i64 = 42; + let baml = to_baml_value(&num).expect("Should convert i64"); + assert_eq!(baml, BamlValue::Int(42)); + + let flag: bool = true; + let baml = to_baml_value(&flag).expect("Should convert bool"); + assert_eq!(baml, BamlValue::Bool(true)); +} + +#[test] +fn test_from_baml_value_with_list() { + let items = vec![BamlValue::Int(1), BamlValue::Int(2), BamlValue::Int(3)]; + let baml_value = BamlValue::List(items); + + let result: Vec = from_baml_value(baml_value).expect("Should convert to Vec"); + assert_eq!(result, vec![1i64, 2, 3]); +} + +#[test] +fn test_from_baml_value_nested_struct() { + let mut user_fields = IndexMap::new(); + user_fields.insert("name".to_string(), BamlValue::String("Test".into())); + user_fields.insert("age".to_string(), BamlValue::Int(25)); + user_fields.insert("active".to_string(), BamlValue::Bool(false)); + + let mut profile_fields = IndexMap::new(); + profile_fields.insert( + "user".to_string(), + BamlValue::Class("Response".into(), user_fields), + ); + profile_fields.insert( + "email".to_string(), + BamlValue::String("test@example.com".into()), + ); + profile_fields.insert( + "tags".to_string(), + BamlValue::List(vec![ + BamlValue::String("tag1".into()), + BamlValue::String("tag2".into()), + ]), + ); + + let baml_value = BamlValue::Class("UserProfile".into(), profile_fields); + + let profile: UserProfile = from_baml_value(baml_value).expect("Should convert nested struct"); + + assert_eq!(profile.user.name, "Test"); + assert_eq!(profile.user.age, 25); + assert!(!profile.user.active); + assert_eq!(profile.email, Some("test@example.com".into())); + assert_eq!(profile.tags, vec!["tag1", "tag2"]); +} + +#[test] +fn test_from_baml_value_with_null_optional() { + let mut user_fields = IndexMap::new(); + user_fields.insert("name".to_string(), BamlValue::String("NoEmail".into())); + user_fields.insert("age".to_string(), BamlValue::Int(20)); + user_fields.insert("active".to_string(), BamlValue::Bool(true)); + + let mut profile_fields = IndexMap::new(); + profile_fields.insert( + "user".to_string(), + BamlValue::Class("Response".into(), user_fields), + ); + profile_fields.insert("email".to_string(), BamlValue::Null); + profile_fields.insert("tags".to_string(), BamlValue::List(vec![])); + + let baml_value = BamlValue::Class("UserProfile".into(), profile_fields); + + let profile: UserProfile = from_baml_value(baml_value).expect("Should handle null optional"); + + assert_eq!(profile.email, None); +} + +#[test] +fn test_round_trip_list() { + let original: Vec = vec!["a".into(), "b".into(), "c".into()]; + let baml = to_baml_value(&original).expect("to_baml_value"); + let restored: Vec = from_baml_value(baml).expect("from_baml_value"); + assert_eq!(original, restored); +} + +#[test] +fn test_from_baml_value_with_flags() { + let llm_output = r#"{"name": "FlagsTest", "age": 40, "active": false}"#; + + let parsed = parse::(llm_output).expect("Should parse"); + let response: Response = + from_baml_value_with_flags(&parsed).expect("Should convert from flags"); + + assert_eq!(response.name, "FlagsTest"); + assert_eq!(response.age, 40); + assert!(!response.active); +} + +#[test] +fn test_parse_and_convert_nested() { + let llm_output = r#"{ + "user": {"name": "Integration", "age": 50, "active": true}, + "email": "int@test.com", + "tags": ["rust", "baml"] + }"#; + + let parsed = parse::(llm_output).expect("Should parse"); + let profile: UserProfile = from_baml_value_with_flags(&parsed).expect("Should convert"); + + assert_eq!(profile.user.name, "Integration"); + assert_eq!(profile.user.age, 50); + assert!(profile.user.active); + assert_eq!(profile.email, Some("int@test.com".into())); + assert_eq!(profile.tags, vec!["rust", "baml"]); +} + +// ============================================================================ +// Smart pointer tests +// ============================================================================ + +#[test] +fn test_from_baml_value_box() { + let val: Box = from_baml_value(BamlValue::String("boxed".into())).unwrap(); + assert_eq!(*val, "boxed"); +} + +#[test] +fn test_from_baml_value_arc_struct() { + let mut fields = IndexMap::new(); + fields.insert("name".to_string(), BamlValue::String("ArcUser".into())); + fields.insert("age".to_string(), BamlValue::Int(28)); + fields.insert("active".to_string(), BamlValue::Bool(true)); + + let val: Arc = from_baml_value(BamlValue::Class("Response".into(), fields)).unwrap(); + assert_eq!(val.name, "ArcUser"); + assert_eq!(val.age, 28); + assert!(val.active); +} + +#[test] +fn test_to_baml_value_box() { + let val = Box::new(42i64); + let baml = to_baml_value(&val).unwrap(); + assert_eq!(baml, BamlValue::Int(42)); +} + +#[test] +fn test_to_baml_value_arc() { + let val = Arc::new("hello".to_string()); + let baml = to_baml_value(&val).unwrap(); + assert_eq!(baml, BamlValue::String("hello".into())); +} + +#[test] +fn test_round_trip_box() { + let mut fields = IndexMap::new(); + fields.insert("name".to_string(), BamlValue::String("Boxed".into())); + fields.insert("age".to_string(), BamlValue::Int(33)); + fields.insert("active".to_string(), BamlValue::Bool(false)); + + let boxed: Box = + from_baml_value(BamlValue::Class("Response".into(), fields)).unwrap(); + + let baml = to_baml_value(&boxed).unwrap(); + match &baml { + BamlValue::Class(name, fields) => { + assert_eq!(name, ::baml_internal_name()); + assert_eq!(fields.get("name"), Some(&BamlValue::String("Boxed".into()))); + assert_eq!(fields.get("age"), Some(&BamlValue::Int(33))); + assert_eq!(fields.get("active"), Some(&BamlValue::Bool(false))); + } + other => panic!("Expected Class, got {:?}", other), + } +} + +// ============================================================================ +// Integer narrowing tests +// ============================================================================ + +#[test] +fn test_from_baml_int_to_u32() { + let val: u32 = from_baml_value(BamlValue::Int(30)).unwrap(); + assert_eq!(val, 30); +} + +#[test] +fn test_from_baml_int_to_i32() { + let val: i32 = from_baml_value(BamlValue::Int(-5)).unwrap(); + assert_eq!(val, -5); +} + +#[test] +fn test_from_baml_int_overflow_u32() { + let result = from_baml_value::(BamlValue::Int(-1)); + assert!( + result.is_err(), + "Expected error for -1 → u32, got {:?}", + result + ); +} + +#[test] +fn test_from_baml_int_overflow_i32() { + let result = from_baml_value::(BamlValue::Int(i64::MAX)); + assert!( + result.is_err(), + "Expected error for i64::MAX → i32, got {:?}", + result + ); +} + +#[test] +fn test_to_baml_value_u32() { + let baml = to_baml_value(&42u32).unwrap(); + assert_eq!(baml, BamlValue::Int(42)); +} + +#[test] +fn test_to_baml_value_i32() { + let baml = to_baml_value(&-7i32).unwrap(); + assert_eq!(baml, BamlValue::Int(-7)); +} + +// ============================================================================ +// Float coercion tests +// ============================================================================ + +#[test] +fn test_from_baml_int_to_f64() { + let val: f64 = from_baml_value(BamlValue::Int(42)).unwrap(); + assert!((val - 42.0).abs() < f64::EPSILON); +} + +#[test] +fn test_from_baml_int_to_f32() { + let val: f32 = from_baml_value(BamlValue::Int(42)).unwrap(); + assert!((val - 42.0).abs() < f32::EPSILON); +} + +#[test] +fn test_to_baml_value_f32() { + let pi = std::f32::consts::PI; + let baml = to_baml_value(&pi).unwrap(); + match baml { + BamlValue::Float(f) => assert!((f - f64::from(pi)).abs() < 0.01), + other => panic!("Expected Float, got {:?}", other), + } +} + +// ============================================================================ +// Map tests +// ============================================================================ + +#[test] +fn test_from_baml_map_to_hashmap() { + let mut map = IndexMap::new(); + map.insert("a".to_string(), BamlValue::Int(1)); + map.insert("b".to_string(), BamlValue::Int(2)); + + let result: HashMap = from_baml_value(BamlValue::Map(map)).unwrap(); + assert_eq!(result.get("a"), Some(&1i64)); + assert_eq!(result.get("b"), Some(&2i64)); +} + +#[test] +fn test_from_baml_map_to_btreemap() { + let mut map = IndexMap::new(); + map.insert("x".to_string(), BamlValue::String("hello".into())); + map.insert("y".to_string(), BamlValue::String("world".into())); + + let result: BTreeMap = from_baml_value(BamlValue::Map(map)).unwrap(); + assert_eq!(result.get("x"), Some(&"hello".to_string())); + assert_eq!(result.get("y"), Some(&"world".to_string())); +} + +#[test] +fn test_to_baml_value_hashmap() { + let mut map = HashMap::new(); + map.insert("key".to_string(), 99i64); + + let baml = to_baml_value(&map).unwrap(); + match baml { + BamlValue::Map(m) => { + assert_eq!(m.get("key"), Some(&BamlValue::Int(99))); + } + other => panic!("Expected Map, got {:?}", other), + } +} + +#[test] +fn test_round_trip_hashmap() { + let mut original = HashMap::new(); + original.insert("one".to_string(), 1i64); + original.insert("two".to_string(), 2i64); + + let baml = to_baml_value(&original).unwrap(); + let restored: HashMap = from_baml_value(baml).unwrap(); + assert_eq!(original, restored); +} + +// ============================================================================ +// Nested container tests +// ============================================================================ + +#[test] +fn test_option_vec_some() { + let items = vec![BamlValue::String("a".into()), BamlValue::String("b".into())]; + let baml = BamlValue::List(items); + + let result: Option> = from_baml_value(baml).unwrap(); + assert_eq!(result, Some(vec!["a".to_string(), "b".to_string()])); +} + +#[test] +fn test_option_vec_none() { + let result: Option> = from_baml_value(BamlValue::Null).unwrap(); + assert_eq!(result, None); +} + +#[test] +fn test_vec_option() { + let items = vec![BamlValue::Int(1), BamlValue::Null, BamlValue::Int(3)]; + let baml = BamlValue::List(items); + + let result: Vec> = from_baml_value(baml).unwrap(); + assert_eq!(result, vec![Some(1i64), None, Some(3i64)]); +} + +#[test] +fn test_nested_option_some_some() { + let result: Option> = from_baml_value(BamlValue::String("x".into())).unwrap(); + assert_eq!(result, Some(Some("x".to_string()))); +} + +#[test] +fn test_nested_option_none() { + let result: Option> = from_baml_value(BamlValue::Null).unwrap(); + assert_eq!(result, None); +} + +// ============================================================================ +// Enum tests +// ============================================================================ + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +enum Color { + Red, + Green, + Blue, +} + +#[test] +fn test_from_baml_enum() { + let val: Color = from_baml_value(BamlValue::Enum("Color".into(), "Red".into())).unwrap(); + assert_eq!(val, Color::Red); +} + +#[test] +fn test_to_baml_enum() { + let baml = to_baml_value(&Color::Green).unwrap(); + assert_eq!( + baml, + BamlValue::Enum( + ::baml_internal_name().into(), + "Green".into() + ) + ); +} + +#[test] +fn test_round_trip_enum() { + let original = Color::Blue; + let baml = to_baml_value(&original).unwrap(); + let restored: Color = from_baml_value(baml).unwrap(); + assert_eq!(original, restored); +} + +#[test] +fn test_unknown_variant_errors() { + let result = from_baml_value::(BamlValue::Enum("Color".into(), "Yellow".into())); + assert!( + result.is_err(), + "Expected error for unknown variant 'Yellow', got {:?}", + result + ); +} + +// ============================================================================ +// Edge cases +// ============================================================================ + +#[test] +fn test_extra_fields_ignored() { + let mut fields = IndexMap::new(); + fields.insert("name".to_string(), BamlValue::String("Alice".into())); + fields.insert("age".to_string(), BamlValue::Int(30)); + fields.insert("active".to_string(), BamlValue::Bool(true)); + fields.insert( + "extra_field".to_string(), + BamlValue::String("ignored".into()), + ); + + let response: Response = from_baml_value(BamlValue::Class("Response".into(), fields)).unwrap(); + assert_eq!(response.name, "Alice"); + assert_eq!(response.age, 30); + assert!(response.active); +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct Empty {} + +#[test] +fn test_empty_struct() { + let fields = IndexMap::new(); + let val: Empty = from_baml_value(BamlValue::Class("Empty".into(), fields)).unwrap(); + let baml = to_baml_value(&val).unwrap(); + match baml { + BamlValue::Class(name, fields) => { + assert_eq!(name, ::baml_internal_name()); + assert!(fields.is_empty()); + } + other => panic!("Expected Class, got {:?}", other), + } +} + +// ============================================================================ +// #[baml(...)] compatibility tests +// ============================================================================ + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct CompatStruct { + #[baml(alias = "fullName")] + full_name: String, + #[baml(skip)] + internal: String, + #[baml(default)] + note: Option, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +enum CompatEnum { + #[baml(alias = "go")] + Start, + Stop, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct SerdeRenameCompat { + #[serde(rename = "nickName")] + nickname: String, +} + +#[test] +fn test_baml_alias_and_skip_to_baml() { + let value = CompatStruct { + full_name: "Alice".into(), + internal: "secret".into(), + note: None, + }; + + let baml = to_baml_value(&value).unwrap(); + match baml { + BamlValue::Class(_, fields) => { + assert_eq!( + fields.get("fullName"), + Some(&BamlValue::String("Alice".into())) + ); + assert!(!fields.contains_key("full_name")); + assert!(!fields.contains_key("internal")); + } + other => panic!("Expected Class, got {:?}", other), + } +} + +#[test] +fn test_baml_alias_and_skip_from_baml() { + let mut aliased = IndexMap::new(); + aliased.insert("fullName".to_string(), BamlValue::String("Bob".into())); + aliased.insert( + "internal".to_string(), + BamlValue::String("should_skip".into()), + ); + let parsed_alias: CompatStruct = + from_baml_value(BamlValue::Class("CompatStruct".into(), aliased)).unwrap(); + assert_eq!(parsed_alias.full_name, "Bob"); + assert_eq!(parsed_alias.internal, ""); + assert_eq!(parsed_alias.note, None); + + let mut original = IndexMap::new(); + original.insert("full_name".to_string(), BamlValue::String("Charlie".into())); + let parsed_original: CompatStruct = + from_baml_value(BamlValue::Class("CompatStruct".into(), original)).unwrap(); + assert_eq!(parsed_original.full_name, "Charlie"); + assert_eq!(parsed_original.internal, ""); +} + +#[test] +fn test_baml_skip_field_excluded_from_schema() { + let schema = render_schema_default::().expect("schema"); + assert!(schema.contains("fullName")); + assert!(!schema.contains("internal")); +} + +#[test] +fn test_baml_enum_alias_round_trip() { + let as_baml = to_baml_value(&CompatEnum::Start).unwrap(); + assert_eq!( + as_baml, + BamlValue::Enum( + ::baml_internal_name().into(), + "go".into() + ) + ); + + let from_alias: CompatEnum = + from_baml_value(BamlValue::Enum("CompatEnum".into(), "go".into())).unwrap(); + assert_eq!(from_alias, CompatEnum::Start); + + let from_original: CompatEnum = + from_baml_value(BamlValue::Enum("CompatEnum".into(), "Start".into())).unwrap(); + assert_eq!(from_original, CompatEnum::Start); +} + +#[test] +fn test_serde_rename_is_accepted_by_bamltype() { + let value = SerdeRenameCompat { + nickname: "D".into(), + }; + + let baml = to_baml_value(&value).unwrap(); + match baml { + BamlValue::Class(_, fields) => { + assert_eq!(fields.get("nickName"), Some(&BamlValue::String("D".into()))); + } + other => panic!("Expected Class, got {:?}", other), + } + + let mut aliased = IndexMap::new(); + aliased.insert("nickName".to_string(), BamlValue::String("E".into())); + let parsed_alias: SerdeRenameCompat = + from_baml_value(BamlValue::Class("SerdeRenameCompat".into(), aliased)).unwrap(); + assert_eq!(parsed_alias.nickname, "E"); + + let mut original = IndexMap::new(); + original.insert("nickname".to_string(), BamlValue::String("F".into())); + let parsed_original: SerdeRenameCompat = + from_baml_value(BamlValue::Class("SerdeRenameCompat".into(), original)).unwrap(); + assert_eq!(parsed_original.nickname, "F"); +} + +// ============================================================================ +// Bridge parity tests +// ============================================================================ + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +#[baml(tag = "type")] +enum TaggedShapeParity { + /// A circle, defined by its radius. + Circle { + /// Radius in meters. + radius: f64, + }, + Rectangle { + width: f64, + height: f64, + }, + Empty, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +#[baml(as_union)] +enum UnitAsUnionParity { + Red, + #[baml(alias = "green")] + Green, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +#[baml(rename_all = "lowercase")] +enum LowercaseEnumParity { + Done, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +#[baml(rename_all = "UPPERCASE")] +struct UppercaseFieldParity { + value: i64, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct BigIntStringParity { + #[baml(int_repr = "string")] + id: u64, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct MapKeysStringParity { + #[baml(map_key_repr = "string")] + values: HashMap, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct MapKeysPairsParity { + #[baml(map_key_repr = "pairs")] + values: HashMap, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct BigIntOptionParity { + #[baml(int_repr = "i64")] + id: Option, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct MapKeysOptionParity { + #[baml(map_key_repr = "string")] + values: Option>, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct RecursiveNodeParity { + value: i64, + next: Option>, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct CheckedValueParity { + #[baml(check(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct AssertedValueParity { + #[baml(assert(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +#[baml(rename_all = "camelCase")] +struct RenameAllParity { + full_name: String, +} + +struct U64ObjectAdapter; + +impl BamlAdapter for U64ObjectAdapter { + fn type_ir() -> TypeIR { + TypeIR::class("AdapterU64Wrapper") + } + + fn register(reg: &mut Registry) { + use bamltype::internal_baml_jinja::types::{Class, Name}; + + if !reg.mark_type("AdapterU64Wrapper") { + return; + } + + reg.register_class(Class { + name: Name::new("AdapterU64Wrapper".to_string()), + description: None, + namespace: StreamingMode::NonStreaming, + fields: vec![( + Name::new("value".to_string()), + TypeIR::string(), + None, + false, + )], + constraints: Vec::new(), + streaming_behavior: Default::default(), + }); + } + + fn try_from_baml( + value: BamlValue, + path: Vec, + ) -> Result { + let map = match value { + BamlValue::Class(_, map) | BamlValue::Map(map) => map, + other => { + return Err(bamltype::compat::BamlConvertError::new( + path, + "object", + format!("{other:?}"), + "expected adapter payload object", + )); + } + }; + + let value = bamltype::compat::get_field(&map, "value", None).ok_or_else(|| { + bamltype::compat::BamlConvertError::new( + path, + "value", + "", + "missing required adapter field", + ) + })?; + + match value { + BamlValue::String(raw) => raw.parse::().map_err(|err| { + bamltype::compat::BamlConvertError::new( + Vec::new(), + "u64", + raw.clone(), + format!("invalid integer string: {err}"), + ) + }), + other => Err(bamltype::compat::BamlConvertError::new( + Vec::new(), + "string", + format!("{other:?}"), + "adapter value must be a string", + )), + } + } +} + +#[derive(Debug, PartialEq)] +#[bamltype::BamlType] +struct WithAdapterParity { + #[baml(with = "U64ObjectAdapter")] + id: u64, +} + +mod parity_collision_a { + #[derive(Debug, PartialEq)] + #[bamltype::BamlType] + pub struct User { + pub name: String, + } +} + +mod parity_collision_b { + #[derive(Debug, PartialEq)] + #[bamltype::BamlType] + pub struct User { + pub name: String, + } +} + +#[test] +fn data_enum_tagged_parses() { + let raw = r#"{ "type": "Circle", "radius": 2.5 }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: TaggedShapeParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value, TaggedShapeParity::Circle { radius: 2.5 }); +} + +#[test] +fn as_union_unit_enum_type_ir_is_literal_union() { + let type_ir = ::baml_type_ir(); + let TypeIR::Union(union, _) = type_ir else { + panic!("expected union type IR"); + }; + + let mut literals = union + .iter_skip_null() + .into_iter() + .map(|item| match item { + TypeIR::Literal(LiteralValue::String(value), _) => value.clone(), + other => panic!("expected string literal in union, got {other:?}"), + }) + .collect::>(); + literals.sort(); + + assert_eq!(literals, vec!["Red".to_string(), "green".to_string()]); +} + +#[test] +fn as_union_unit_enum_alias_parses() { + let parsed = parse::(r#""green""#).expect("parse failed"); + let value: UnitAsUnionParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value, UnitAsUnionParity::Green); +} + +#[test] +fn rename_all_lowercase_variant_parses() { + let parsed = parse::(r#""done""#).expect("parse failed"); + let value: LowercaseEnumParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value, LowercaseEnumParity::Done); +} + +#[test] +fn rename_all_uppercase_field_parses() { + let parsed = parse::(r#"{ "VALUE": 7 }"#).expect("parse failed"); + let value: UppercaseFieldParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value.value, 7); + + let roundtrip = + to_baml_value(&UppercaseFieldParity { value: 9 }).expect("to_baml_value failed"); + match roundtrip { + BamlValue::Class(_, fields) | BamlValue::Map(fields) => { + assert_eq!(fields.get("VALUE"), Some(&BamlValue::Int(9))); + } + other => panic!("expected object-like value, got {other:?}"), + } +} + +#[test] +fn to_baml_data_enum_shape() { + let value = TaggedShapeParity::Circle { radius: 1.25 }; + let baml = to_baml_value(&value).expect("to_baml_value failed"); + match baml { + BamlValue::Class(_, fields) | BamlValue::Map(fields) => { + assert_eq!( + fields.get("type"), + Some(&BamlValue::String("Circle".into())) + ); + assert_eq!(fields.get("radius"), Some(&BamlValue::Float(1.25))); + } + other => panic!("expected class-like value, got {other:?}"), + } +} + +#[test] +fn int_repr_string_parses() { + let raw = r#"{ "id": "18446744073709551615" }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: BigIntStringParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value.id, u64::MAX); +} + +#[test] +fn int_repr_option_parses() { + let raw = r#"{ "id": 42 }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: BigIntOptionParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value.id, Some(42)); + + let raw = r#"{ "id": null }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: BigIntOptionParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value.id, None); +} + +#[test] +fn map_key_repr_string_parses() { + let raw = r#"{ "values": { "1": "a", "2": "b" } }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: MapKeysStringParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + + let mut expected = HashMap::new(); + expected.insert(1_u32, "a".to_string()); + expected.insert(2_u32, "b".to_string()); + assert_eq!(value.values, expected); +} + +#[test] +fn map_key_repr_option_parses() { + let raw = r#"{ "values": { "10": "x" } }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: MapKeysOptionParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + + let mut expected = HashMap::new(); + expected.insert(10_u32, "x".to_string()); + assert_eq!(value.values, Some(expected)); +} + +#[test] +fn map_key_repr_pairs_parses() { + let raw = r#"{ "values": [ { "key": 1, "value": "a" }, { "key": 2, "value": "b" } ] }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: MapKeysPairsParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + + let mut expected = HashMap::new(); + expected.insert(1_u32, "a".to_string()); + expected.insert(2_u32, "b".to_string()); + assert_eq!(value.values, expected); +} + +#[test] +fn map_key_repr_pairs_registers_entry_class() { + let entry_name = format!( + "{}::values__Entry", + ::baml_internal_name() + ); + let of = &MapKeysPairsParity::baml_schema().output_format; + let class = of + .classes + .get(&(entry_name, StreamingMode::NonStreaming)) + .expect("entry class missing"); + + assert_eq!(class.name.rendered_name(), "valuesEntry"); +} + +#[test] +fn recursion_is_detected() { + let of = &RecursiveNodeParity::baml_schema().output_format; + assert!( + of.recursive_classes + .contains(::baml_internal_name()) + ); +} + +#[test] +fn rename_all_applies() { + let raw = r#"{ "fullName": "Ada" }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: RenameAllParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value.full_name, "Ada"); +} + +#[test] +fn field_constraints_are_registered() { + let of = &CheckedValueParity::baml_schema().output_format; + let internal_name = ::baml_internal_name().to_string(); + let class = of + .classes + .get(&(internal_name, StreamingMode::NonStreaming)) + .expect("class missing"); + + let (_, field_type, _, _) = class + .fields + .iter() + .find(|(name, _, _, _)| name.real_name() == "value") + .expect("value field missing"); + + assert!(field_type.meta().constraints.iter().any(|constraint| { + constraint.level == ConstraintLevel::Check + && constraint.label.as_deref() == Some("positive") + })); +} + +#[test] +fn field_assert_constraints_are_registered() { + let of = &AssertedValueParity::baml_schema().output_format; + let internal_name = ::baml_internal_name().to_string(); + let class = of + .classes + .get(&(internal_name, StreamingMode::NonStreaming)) + .expect("class missing"); + + let (_, field_type, _, _) = class + .fields + .iter() + .find(|(name, _, _, _)| name.real_name() == "value") + .expect("value field missing"); + + assert!(field_type.meta().constraints.iter().any(|constraint| { + constraint.level == ConstraintLevel::Assert + && constraint.label.as_deref() == Some("positive") + })); +} + +#[test] +fn internal_names_are_unique() { + let a_name = ::baml_internal_name(); + let b_name = ::baml_internal_name(); + assert_ne!(a_name, b_name); + + let a_class = parity_collision_a::User::baml_schema() + .output_format + .classes + .get(&(a_name.to_string(), StreamingMode::NonStreaming)) + .expect("class missing"); + let b_class = parity_collision_b::User::baml_schema() + .output_format + .classes + .get(&(b_name.to_string(), StreamingMode::NonStreaming)) + .expect("class missing"); + + assert_eq!(a_class.name.rendered_name(), "User"); + assert_eq!(b_class.name.rendered_name(), "User"); +} + +#[test] +fn with_adapter_schema_and_registration_are_used() { + let of = &WithAdapterParity::baml_schema().output_format; + + let adapter_class = of + .classes + .get(&("AdapterU64Wrapper".to_string(), StreamingMode::NonStreaming)) + .expect("adapter class missing"); + assert_eq!(adapter_class.name.real_name(), "AdapterU64Wrapper"); + assert!( + adapter_class + .fields + .iter() + .any(|(name, _, _, _)| name.real_name() == "value") + ); + + let owner_name = ::baml_internal_name().to_string(); + let owner = of + .classes + .get(&(owner_name, StreamingMode::NonStreaming)) + .expect("owner class missing"); + let (_, field_ir, _, _) = owner + .fields + .iter() + .find(|(name, _, _, _)| name.real_name() == "id") + .expect("id field missing"); + assert!(matches!( + field_ir, + TypeIR::Class { name, .. } if name == "AdapterU64Wrapper" + )); +} + +#[test] +fn with_adapter_parse_uses_custom_converter() { + let raw = r#"{ "id": { "value": "18446744073709551615" } }"#; + let parsed = parse::(raw).expect("parse failed"); + let value: WithAdapterParity = from_baml_value_with_flags(&parsed).expect("convert failed"); + assert_eq!(value.id, u64::MAX); +} diff --git a/crates/bamltype/tests/parity_bridge_api.rs b/crates/bamltype/tests/parity_bridge_api.rs new file mode 100644 index 00000000..a49d4450 --- /dev/null +++ b/crates/bamltype/tests/parity_bridge_api.rs @@ -0,0 +1,166 @@ +use std::collections::HashMap; + +use bamltype::RenderOptions; +use bamltype::jsonish::deserializer::deserialize_flags::Flag; +use bamltype::{BamlParseError, parse_llm_output, render_schema, schema_fingerprint}; + +/// A user profile returned by the model. +/// +/// ## Notes +/// - `fullName` should be the display name. +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +struct DocUser { + /// Full name as displayed in the UI. + #[baml(alias = "fullName")] + name: String, + age: i64, +} + +/// This should be ignored. +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +#[baml(description = "Override description.")] +struct OverrideUser { + value: String, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +enum Color { + /// Red hot. + Red, + #[baml(alias = "green")] + Green, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +struct CheckedValue { + #[baml(check(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +struct AssertedValue { + #[baml(assert(label = "positive", expr = "this > 0"))] + value: i64, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +struct Unsigned32 { + value: u32, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +struct FuzzUser { + name: String, + age: u32, + nickname: Option, +} + +#[derive(Debug, Clone, PartialEq)] +#[bamltype::BamlType] +struct FuzzExplain { + values: HashMap, +} + +#[test] +fn doc_comments_render() { + let schema = render_schema::(RenderOptions::default()) + .expect("render failed") + .unwrap_or_default(); + + assert!(schema.contains("A user profile returned by the model.")); + assert!(schema.contains("## Notes")); + assert!(schema.contains("Full name as displayed in the UI.")); +} + +#[test] +fn description_override_wins() { + let schema = render_schema::(RenderOptions::default()) + .expect("render failed") + .unwrap_or_default(); + + assert!(schema.contains("Override description.")); + assert!(!schema.contains("This should be ignored.")); +} + +#[test] +fn enum_variant_descriptions_render() { + let schema = render_schema::(RenderOptions::default()) + .expect("render failed") + .unwrap_or_default(); + + assert!(schema.contains("Red: Red hot.")); +} + +#[test] +fn constraint_checks_are_reported() { + let raw = r#"{ "value": 3 }"#; + let parsed = parse_llm_output::(raw, true).expect("parse failed"); + assert!( + parsed + .checks + .iter() + .any(|check| check.name == "positive" && check.status == "succeeded") + ); +} + +#[test] +fn constraint_asserts_fail() { + let raw = r#"{ "value": -1 }"#; + let err = parse_llm_output::(raw, true).expect_err("expected assert failure"); + match err { + BamlParseError::ConstraintAssertsFailed { failed } => { + assert!(failed.iter().any(|check| check.name == "positive")); + } + other => panic!("unexpected error: {other:?}"), + } +} + +#[test] +fn unsigned_bounds_enforced() { + let raw = r#"{ "value": -1 }"#; + let err = parse_llm_output::(raw, true).expect_err("expected range error"); + assert!(matches!(err, BamlParseError::Convert(_))); + + let raw = r#"{ "value": 4294967296 }"#; + let err = parse_llm_output::(raw, true).expect_err("expected range error"); + assert!(matches!(err, BamlParseError::Convert(_))); +} + +#[test] +fn markdown_fence_parses_and_sets_flag() { + let raw = "```json\n{ \"name\": \"Ada\", \"age\": 36 }\n```"; + let parsed = parse_llm_output::(raw, true).expect("parse"); + + assert_eq!(parsed.value.name, "Ada"); + assert!( + parsed + .flags + .iter() + .any(|flag| matches!(flag, Flag::ObjectFromMarkdown(_))) + ); +} + +#[test] +fn explanations_surface_on_map_parse_error() { + let raw = r#"{ "values": { "ok": 1, "bad": "oops" } }"#; + let parsed = parse_llm_output::(raw, true).expect("parse"); + assert_eq!(parsed.value.values.get("ok"), Some(&1)); + assert!(!parsed.explanations.is_empty()); +} + +#[test] +fn schema_fingerprint_is_stable() { + let output = ::baml_output_format(); + let a = schema_fingerprint(output, RenderOptions::default()) + .expect("schema fingerprint should work"); + let b = schema_fingerprint(output, RenderOptions::default()) + .expect("schema fingerprint should work"); + assert_eq!(a, b); +} diff --git a/crates/bamltype/tests/ui.rs b/crates/bamltype/tests/ui.rs new file mode 100644 index 00000000..365834bd --- /dev/null +++ b/crates/bamltype/tests/ui.rs @@ -0,0 +1,5 @@ +#[test] +fn ui_compile_failures() { + let t = trybuild::TestCases::new(); + t.compile_fail("tests/ui/*.rs"); +} diff --git a/crates/bamltype/tests/ui/as_enum_data_enum.rs b/crates/bamltype/tests/ui/as_enum_data_enum.rs new file mode 100644 index 00000000..2b46ebf2 --- /dev/null +++ b/crates/bamltype/tests/ui/as_enum_data_enum.rs @@ -0,0 +1,9 @@ +use bamltype::BamlType; + +#[BamlType] +#[baml(as_enum)] +enum Bad { + A { value: i64 }, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/as_enum_data_enum.stderr b/crates/bamltype/tests/ui/as_enum_data_enum.stderr new file mode 100644 index 00000000..9ae230f7 --- /dev/null +++ b/crates/bamltype/tests/ui/as_enum_data_enum.stderr @@ -0,0 +1,8 @@ +error: as_enum is only valid for unit enums; hint: remove #[baml(as_enum)] or convert variants to unit + --> tests/ui/as_enum_data_enum.rs:4:1 + | +4 | / #[baml(as_enum)] +5 | | enum Bad { +6 | | A { value: i64 }, +7 | | } + | |_^ diff --git a/crates/bamltype/tests/ui/function_type.rs b/crates/bamltype/tests/ui/function_type.rs new file mode 100644 index 00000000..8f773fff --- /dev/null +++ b/crates/bamltype/tests/ui/function_type.rs @@ -0,0 +1,8 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad { + callback: fn(i32) -> i32, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/function_type.stderr b/crates/bamltype/tests/ui/function_type.stderr new file mode 100644 index 00000000..56869b0f --- /dev/null +++ b/crates/bamltype/tests/ui/function_type.stderr @@ -0,0 +1,5 @@ +error: function types are not supported in BAML outputs; hint: remove the field or use #[baml(with = "...")] to adapt it + --> tests/ui/function_type.rs:5:15 + | +5 | callback: fn(i32) -> i32, + | ^^^^^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/large_int_without_repr.rs b/crates/bamltype/tests/ui/large_int_without_repr.rs new file mode 100644 index 00000000..0eb185a3 --- /dev/null +++ b/crates/bamltype/tests/ui/large_int_without_repr.rs @@ -0,0 +1,8 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad { + value: u64, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/large_int_without_repr.stderr b/crates/bamltype/tests/ui/large_int_without_repr.stderr new file mode 100644 index 00000000..9d62076e --- /dev/null +++ b/crates/bamltype/tests/ui/large_int_without_repr.stderr @@ -0,0 +1,5 @@ +error: unsupported integer width for BAML outputs; hint: use #[baml(int_repr = "string"|"i64")] or a smaller integer type + --> tests/ui/large_int_without_repr.rs:5:12 + | +5 | value: u64, + | ^^^ diff --git a/crates/bamltype/tests/ui/map_key_non_string.rs b/crates/bamltype/tests/ui/map_key_non_string.rs new file mode 100644 index 00000000..78ae94dd --- /dev/null +++ b/crates/bamltype/tests/ui/map_key_non_string.rs @@ -0,0 +1,9 @@ +use bamltype::BamlType; +type HashMap = std::collections::HashMap; + +#[BamlType] +struct Bad { + values: HashMap, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/map_key_non_string.stderr b/crates/bamltype/tests/ui/map_key_non_string.stderr new file mode 100644 index 00000000..53044558 --- /dev/null +++ b/crates/bamltype/tests/ui/map_key_non_string.stderr @@ -0,0 +1,5 @@ +error: map keys must be String for object maps; hint: use HashMap or add #[baml(map_key_repr = "string"|"pairs")], or use a custom adapter + --> tests/ui/map_key_non_string.rs:6:13 + | +6 | values: HashMap, + | ^^^^^^^^^^^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/map_key_repr_non_map.rs b/crates/bamltype/tests/ui/map_key_repr_non_map.rs new file mode 100644 index 00000000..7590879a --- /dev/null +++ b/crates/bamltype/tests/ui/map_key_repr_non_map.rs @@ -0,0 +1,9 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad { + #[baml(map_key_repr = "string")] + value: Vec, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/map_key_repr_non_map.stderr b/crates/bamltype/tests/ui/map_key_repr_non_map.stderr new file mode 100644 index 00000000..e91f3f89 --- /dev/null +++ b/crates/bamltype/tests/ui/map_key_repr_non_map.stderr @@ -0,0 +1,5 @@ +error: map_key_repr only applies to map fields (HashMap/BTreeMap), optionally wrapped in Option/Vec/Box/Arc/Rc; hint: remove the attribute or change the field type + --> tests/ui/map_key_repr_non_map.rs:6:16 + | +6 | value: Vec, + | ^^^^^^ diff --git a/crates/bamltype/tests/ui/non_string_literal_attr.rs b/crates/bamltype/tests/ui/non_string_literal_attr.rs new file mode 100644 index 00000000..d2bc878b --- /dev/null +++ b/crates/bamltype/tests/ui/non_string_literal_attr.rs @@ -0,0 +1,9 @@ +use bamltype::BamlType; + +#[BamlType] +#[baml(name = 123)] +struct NonStringName { + value: String, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/non_string_literal_attr.stderr b/crates/bamltype/tests/ui/non_string_literal_attr.stderr new file mode 100644 index 00000000..40feadfe --- /dev/null +++ b/crates/bamltype/tests/ui/non_string_literal_attr.stderr @@ -0,0 +1,5 @@ +error: expected string literal; hint: wrap the value in quotes + --> tests/ui/non_string_literal_attr.rs:4:8 + | +4 | #[baml(name = 123)] + | ^^^^ diff --git a/crates/bamltype/tests/ui/serde_default_path.rs b/crates/bamltype/tests/ui/serde_default_path.rs new file mode 100644 index 00000000..edc91047 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_default_path.rs @@ -0,0 +1,13 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad { + #[serde(default = "default_age")] + age: i64, +} + +fn default_age() -> i64 { + 0 +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/serde_default_path.stderr b/crates/bamltype/tests/ui/serde_default_path.stderr new file mode 100644 index 00000000..8aaa9b07 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_default_path.stderr @@ -0,0 +1,5 @@ +error: serde(default = "path") is not supported; hint: use #[baml(default)] or Default::default + --> tests/ui/serde_default_path.rs:5:13 + | +5 | #[serde(default = "default_age")] + | ^^^^^^^^^^^^^^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/serde_flatten.rs b/crates/bamltype/tests/ui/serde_flatten.rs new file mode 100644 index 00000000..49b64be1 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_flatten.rs @@ -0,0 +1,10 @@ +use bamltype::BamlType; +type HashMap = std::collections::HashMap; + +#[BamlType] +struct Bad { + #[serde(flatten)] + extras: HashMap, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/serde_flatten.stderr b/crates/bamltype/tests/ui/serde_flatten.stderr new file mode 100644 index 00000000..afbab798 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_flatten.stderr @@ -0,0 +1,5 @@ +error: serde(flatten) is not supported; hint: model fields explicitly + --> tests/ui/serde_flatten.rs:6:13 + | +6 | #[serde(flatten)] + | ^^^^^^^ diff --git a/crates/bamltype/tests/ui/serde_json_value.rs b/crates/bamltype/tests/ui/serde_json_value.rs new file mode 100644 index 00000000..521bfa8b --- /dev/null +++ b/crates/bamltype/tests/ui/serde_json_value.rs @@ -0,0 +1,8 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad { + value: serde_json::Value, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/serde_json_value.stderr b/crates/bamltype/tests/ui/serde_json_value.stderr new file mode 100644 index 00000000..ee40cf12 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_json_value.stderr @@ -0,0 +1,5 @@ +error: serde_json::Value is not supported without a #[baml(with = "...")] adapter; hint: use a concrete type or provide a custom adapter + --> tests/ui/serde_json_value.rs:5:12 + | +5 | value: serde_json::Value, + | ^^^^^^^^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/serde_skip_variant.rs b/crates/bamltype/tests/ui/serde_skip_variant.rs new file mode 100644 index 00000000..d71b0fb3 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_skip_variant.rs @@ -0,0 +1,10 @@ +use bamltype::BamlType; + +#[BamlType] +enum Bad { + #[serde(skip)] + A, + B, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/serde_skip_variant.stderr b/crates/bamltype/tests/ui/serde_skip_variant.stderr new file mode 100644 index 00000000..98818a2f --- /dev/null +++ b/crates/bamltype/tests/ui/serde_skip_variant.stderr @@ -0,0 +1,5 @@ +error: serde(skip) is not supported on enum variants; hint: remove the variant or use a separate enum + --> tests/ui/serde_skip_variant.rs:5:13 + | +5 | #[serde(skip)] + | ^^^^ diff --git a/crates/bamltype/tests/ui/serde_untagged.rs b/crates/bamltype/tests/ui/serde_untagged.rs new file mode 100644 index 00000000..59ac8c83 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_untagged.rs @@ -0,0 +1,10 @@ +use bamltype::BamlType; + +#[BamlType] +#[serde(untagged)] +enum Bad { + A { x: i64 }, + B { y: i64 }, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/serde_untagged.stderr b/crates/bamltype/tests/ui/serde_untagged.stderr new file mode 100644 index 00000000..ab19cbb0 --- /dev/null +++ b/crates/bamltype/tests/ui/serde_untagged.stderr @@ -0,0 +1,5 @@ +error: serde(untagged) is not supported; hint: use #[baml(tag = "...")] for data enums + --> tests/ui/serde_untagged.rs:4:9 + | +4 | #[serde(untagged)] + | ^^^^^^^^ diff --git a/crates/bamltype/tests/ui/trait_object.rs b/crates/bamltype/tests/ui/trait_object.rs new file mode 100644 index 00000000..ce46d599 --- /dev/null +++ b/crates/bamltype/tests/ui/trait_object.rs @@ -0,0 +1,8 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad { + value: Box, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/trait_object.stderr b/crates/bamltype/tests/ui/trait_object.stderr new file mode 100644 index 00000000..a1ca7d07 --- /dev/null +++ b/crates/bamltype/tests/ui/trait_object.stderr @@ -0,0 +1,5 @@ +error: trait objects are not supported in BAML outputs; hint: use a concrete type or a custom adapter + --> tests/ui/trait_object.rs:5:16 + | +5 | value: Box, + | ^^^^^^^^^^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/tuple_enum_variant.rs b/crates/bamltype/tests/ui/tuple_enum_variant.rs new file mode 100644 index 00000000..1271832e --- /dev/null +++ b/crates/bamltype/tests/ui/tuple_enum_variant.rs @@ -0,0 +1,8 @@ +use bamltype::BamlType; + +#[BamlType] +enum Bad { + One(u32), +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/tuple_enum_variant.stderr b/crates/bamltype/tests/ui/tuple_enum_variant.stderr new file mode 100644 index 00000000..d4aeeff0 --- /dev/null +++ b/crates/bamltype/tests/ui/tuple_enum_variant.stderr @@ -0,0 +1,5 @@ +error: Tuple enum variants are not supported; hint: use a unit or struct-like variant + --> tests/ui/tuple_enum_variant.rs:5:5 + | +5 | One(u32), + | ^^^^^^^^ diff --git a/crates/bamltype/tests/ui/tuple_field.rs b/crates/bamltype/tests/ui/tuple_field.rs new file mode 100644 index 00000000..d4eae10a --- /dev/null +++ b/crates/bamltype/tests/ui/tuple_field.rs @@ -0,0 +1,8 @@ +use bamltype::BamlType; + +#[BamlType] +struct TupleFieldRejected { + pair: (i32, i32), +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/tuple_field.stderr b/crates/bamltype/tests/ui/tuple_field.stderr new file mode 100644 index 00000000..a21110de --- /dev/null +++ b/crates/bamltype/tests/ui/tuple_field.stderr @@ -0,0 +1,5 @@ +error: tuple types are not supported in BAML outputs; hint: use a struct with named fields or a list + --> tests/ui/tuple_field.rs:5:11 + | +5 | pair: (i32, i32), + | ^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/tuple_struct.rs b/crates/bamltype/tests/ui/tuple_struct.rs new file mode 100644 index 00000000..1a02b279 --- /dev/null +++ b/crates/bamltype/tests/ui/tuple_struct.rs @@ -0,0 +1,6 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad(u32); + +fn main() {} diff --git a/crates/bamltype/tests/ui/tuple_struct.stderr b/crates/bamltype/tests/ui/tuple_struct.stderr new file mode 100644 index 00000000..4aff7977 --- /dev/null +++ b/crates/bamltype/tests/ui/tuple_struct.stderr @@ -0,0 +1,5 @@ +error: Tuple structs are not supported for BAML outputs; hint: use a named-field struct + --> tests/ui/tuple_struct.rs:4:1 + | +4 | struct Bad(u32); + | ^^^^^^^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/unit_struct.rs b/crates/bamltype/tests/ui/unit_struct.rs new file mode 100644 index 00000000..29147057 --- /dev/null +++ b/crates/bamltype/tests/ui/unit_struct.rs @@ -0,0 +1,6 @@ +use bamltype::BamlType; + +#[BamlType] +struct Bad; + +fn main() {} diff --git a/crates/bamltype/tests/ui/unit_struct.stderr b/crates/bamltype/tests/ui/unit_struct.stderr new file mode 100644 index 00000000..7504516d --- /dev/null +++ b/crates/bamltype/tests/ui/unit_struct.stderr @@ -0,0 +1,5 @@ +error: Unit structs are not supported for BAML outputs; hint: use a named-field struct or enum + --> tests/ui/unit_struct.rs:4:1 + | +4 | struct Bad; + | ^^^^^^^^^^^ diff --git a/crates/bamltype/tests/ui/unsupported_baml_attr.rs b/crates/bamltype/tests/ui/unsupported_baml_attr.rs new file mode 100644 index 00000000..7641b452 --- /dev/null +++ b/crates/bamltype/tests/ui/unsupported_baml_attr.rs @@ -0,0 +1,9 @@ +use bamltype::BamlType; + +#[BamlType] +#[baml(unknown = "x")] +struct UnsupportedAttr { + value: String, +} + +fn main() {} diff --git a/crates/bamltype/tests/ui/unsupported_baml_attr.stderr b/crates/bamltype/tests/ui/unsupported_baml_attr.stderr new file mode 100644 index 00000000..3ab8add9 --- /dev/null +++ b/crates/bamltype/tests/ui/unsupported_baml_attr.stderr @@ -0,0 +1,5 @@ +error: unsupported #[baml(...)] attribute; hint: check the supported keys in the bridge docs + --> tests/ui/unsupported_baml_attr.rs:4:8 + | +4 | #[baml(unknown = "x")] + | ^^^^^^^^^^^^^ diff --git a/crates/dspy-rs/Cargo.toml b/crates/dspy-rs/Cargo.toml index 04803da5..16f99ec3 100644 --- a/crates/dspy-rs/Cargo.toml +++ b/crates/dspy-rs/Cargo.toml @@ -25,7 +25,8 @@ tokio = { version = "1.46.1", features = ["full"] } async-trait = "0.1.83" anyhow = "1.0.99" bon = "3.7.0" -baml-bridge = { path = "../baml-bridge", features = ["derive"] } +bamltype = { path = "../bamltype" } +facet = { version = "0.43.2", default-features = false, features = ["std"] } thiserror = "2.0.17" dsrs_macros = { version = "0.7.2", path = "../dsrs-macros" } csv = { version = "1.3.1" } diff --git a/crates/dspy-rs/examples/16-insurance-claim-prompt.rs b/crates/dspy-rs/examples/16-insurance-claim-prompt.rs index 33d180bf..da712dc9 100644 --- a/crates/dspy-rs/examples/16-insurance-claim-prompt.rs +++ b/crates/dspy-rs/examples/16-insurance-claim-prompt.rs @@ -11,7 +11,8 @@ use dspy_rs::{BamlType, ChatAdapter, Signature, init_tracing}; type NaiveDate = String; /// Basic claim information (metadata about the claim intake). -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub struct ClaimHeader { /// Claim ID in format `CLM-XXXXXX`, where `X` is a digit. pub claim_id: Option, @@ -30,7 +31,8 @@ pub struct ClaimHeader { } /// Channel used to report a claim. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub enum ClaimChannel { Email, Phone, @@ -39,7 +41,8 @@ pub enum ClaimChannel { } /// Policy information if available. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub struct PolicyDetails { /// Policy number in format `POL-XXXXXXXXX`, where `X` is a digit. pub policy_number: Option, @@ -58,7 +61,8 @@ pub struct PolicyDetails { } /// Type of insurance coverage. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub enum CoverageType { Property, Auto, @@ -69,7 +73,8 @@ pub enum CoverageType { } /// An insured object involved in the claim (vehicle, building, person, etc.). -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub struct InsuredObject { /// Unique identifier for insured object. /// @@ -98,7 +103,8 @@ pub struct InsuredObject { } /// Type of insured object. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub enum InsuredObjectType { Vehicle, Building, @@ -107,7 +113,8 @@ pub enum InsuredObjectType { } /// Structured incident details. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub struct IncidentDescription { /// Specific standardized incident type. pub incident_type: IncidentType, @@ -123,7 +130,8 @@ pub struct IncidentDescription { } /// Specific standardized incident type. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub enum IncidentType { RearEndCollision, SideImpactCollision, @@ -143,7 +151,8 @@ pub enum IncidentType { } /// Standardized location type where incident occurred. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub enum LocationType { Intersection, Highway, @@ -157,7 +166,8 @@ pub enum LocationType { } /// Top-level insurance claim object aggregating all extracted fields. -#[derive(Debug, Clone, PartialEq, Eq, BamlType)] +#[derive(Debug, Clone, PartialEq, Eq)] +#[BamlType] pub struct InsuranceClaim { /// Basic claim information. pub header: ClaimHeader, diff --git a/crates/dspy-rs/src/adapter/chat.rs b/crates/dspy-rs/src/adapter/chat.rs index 228a6c68..a05a9286 100644 --- a/crates/dspy-rs/src/adapter/chat.rs +++ b/crates/dspy-rs/src/adapter/chat.rs @@ -8,13 +8,11 @@ use std::sync::{Arc, LazyLock}; use tracing::{Instrument, debug, trace}; use super::Adapter; -use crate::baml_bridge::BamlType; -use crate::baml_bridge::BamlValueConvert; -use crate::baml_bridge::ToBamlValue; -use crate::baml_bridge::jsonish; -use crate::baml_bridge::jsonish::BamlValueWithFlags; -use crate::baml_bridge::jsonish::deserializer::coercer::run_user_checks; -use crate::baml_bridge::jsonish::deserializer::deserialize_flags::DeserializerConditions; +use crate::bamltype::compat::{BamlTypeTrait, BamlValueConvert, ToBamlValue}; +use crate::bamltype::jsonish; +use crate::bamltype::jsonish::BamlValueWithFlags; +use crate::bamltype::jsonish::deserializer::coercer::run_user_checks; +use crate::bamltype::jsonish::deserializer::deserialize_flags::DeserializerConditions; use crate::serde_utils::get_iter_from_value; use crate::utils::cache::CacheEntry; use crate::{ @@ -29,10 +27,6 @@ pub struct ChatAdapter; static FIELD_HEADER_PATTERN: LazyLock = LazyLock::new(|| Regex::new(r"^\[\[ ## (\w+) ## \]\]").unwrap()); -fn get_type_hint(_field: &Value) -> String { - String::new() -} - fn render_field_type_schema( parent_format: &OutputFormatContent, type_ir: &TypeIR, @@ -52,7 +46,25 @@ fn render_field_type_schema( Ok(schema) } -fn simplify_type_name(raw: &str) -> String { +fn resolve_rendered_type_token(token: &str, output_format: Option<&OutputFormatContent>) -> String { + if let Some(output_format) = output_format { + if let Some(class) = output_format + .classes + .iter() + .find_map(|((name, _), class)| (name == token).then_some(class)) + { + return class.name.rendered_name().to_string(); + } + + if let Some(enm) = output_format.enums.get(token) { + return enm.name.rendered_name().to_string(); + } + } + + token.rsplit("::").next().unwrap_or(token).to_string() +} + +fn simplify_type_name(raw: &str, output_format: Option<&OutputFormatContent>) -> String { let mut result = String::with_capacity(raw.len()); let mut chars = raw.chars(); while let Some(ch) = chars.next() { @@ -64,8 +76,8 @@ fn simplify_type_name(raw: &str) -> String { } token.push(next); } - let simplified = token.rsplit("::").next().unwrap_or(&token); - result.push_str(simplified); + let rendered = resolve_rendered_type_token(&token, output_format); + result.push_str(&rendered); } else { result.push(ch); } @@ -73,9 +85,12 @@ fn simplify_type_name(raw: &str) -> String { result } -fn render_type_name_for_prompt(type_ir: &TypeIR) -> String { +fn render_type_name_for_prompt( + type_ir: &TypeIR, + output_format: Option<&OutputFormatContent>, +) -> String { let raw = type_ir.diagnostic_repr().to_string(); - let simplified = simplify_type_name(&raw); + let simplified = simplify_type_name(&raw, output_format); simplified .replace("class ", "") .replace("enum ", "") @@ -366,23 +381,11 @@ impl ChatAdapter { .next() .unwrap() .clone(); - let first_output_field_value = signature - .output_fields() - .as_object() - .unwrap() - .get(&first_output_field) - .unwrap() - .clone(); - - let type_hint = get_type_hint(&first_output_field_value); - let mut user_message = format!( - "Respond with the corresponding output fields, starting with the field `[[ ## {first_output_field} ## ]]`{type_hint}," + "Respond with the corresponding output fields, starting with the field `[[ ## {first_output_field} ## ]]`," ); - for (field_name, field) in get_iter_from_value(&signature.output_fields()).skip(1) { - user_message.push_str( - format!(" then `[[ ## {field_name} ## ]]`{},", get_type_hint(&field)).as_str(), - ); + for (field_name, _) in get_iter_from_value(&signature.output_fields()).skip(1) { + user_message.push_str(format!(" then `[[ ## {field_name} ## ]]`,").as_str()); } user_message.push_str(" and then ending with the marker for `[[ ## completed ## ]]`."); @@ -450,10 +453,13 @@ impl ChatAdapter { } fn format_field_descriptions_typed(&self) -> String { + let input_format = ::baml_output_format(); + let output_format = S::output_format_content(); + let mut lines = Vec::new(); lines.push("Your input fields are:".to_string()); for (i, field) in S::input_fields().iter().enumerate() { - let type_name = render_type_name_for_prompt(&(field.type_ir)()); + let type_name = render_type_name_for_prompt(&(field.type_ir)(), Some(input_format)); let mut line = format!("{}. `{}` ({type_name})", i + 1, field.name); if !field.description.is_empty() { line.push_str(": "); @@ -465,7 +471,7 @@ impl ChatAdapter { lines.push(String::new()); lines.push("Your output fields are:".to_string()); for (i, field) in S::output_fields().iter().enumerate() { - let type_name = render_type_name_for_prompt(&(field.type_ir)()); + let type_name = render_type_name_for_prompt(&(field.type_ir)(), Some(output_format)); let mut line = format!("{}. `{}` ({type_name})", i + 1, field.name); if !field.description.is_empty() { line.push_str(": "); @@ -492,7 +498,7 @@ impl ChatAdapter { let parent_format = S::output_format_content(); for field in S::output_fields() { let type_ir = (field.type_ir)(); - let type_name = render_type_name_for_prompt(&type_ir); + let type_name = render_type_name_for_prompt(&type_ir, Some(parent_format)); let schema = render_field_type_schema(parent_format, &type_ir)?; lines.push(format!("[[ ## {} ## ]]", field.name)); lines.push(format!( @@ -519,7 +525,7 @@ impl ChatAdapter { let Some(fields) = baml_value_fields(&baml_value) else { return String::new(); }; - let input_output_format = ::baml_output_format(); + let input_output_format = ::baml_output_format(); let mut result = String::new(); for field_spec in S::input_fields() { @@ -593,7 +599,7 @@ impl ChatAdapter { let mut metas = IndexMap::new(); let mut errors = Vec::new(); - let mut output_map = crate::baml_bridge::baml_types::BamlMap::new(); + let mut output_map = crate::bamltype::baml_types::BamlMap::new(); let mut checks_total = 0usize; let mut checks_failed = 0usize; let mut asserts_failed = 0usize; @@ -642,7 +648,7 @@ impl ChatAdapter { let baml_value: BamlValue = parsed.clone().into(); let mut flags = Vec::new(); - collect_flags(&parsed, &mut flags); + collect_flags_recursive(&parsed, &mut flags); let mut checks = Vec::new(); match run_user_checks(&baml_value, &type_ir) { @@ -723,13 +729,19 @@ impl ChatAdapter { let partial = if output_map.is_empty() { None } else { - Some(BamlValue::Map(output_map)) + Some(BamlValue::Class( + ::baml_internal_name().to_string(), + output_map, + )) }; return Err(ParseError::Multiple { errors, partial }); } let typed_output = ::try_from_baml_value( - BamlValue::Map(output_map), + BamlValue::Class( + ::baml_internal_name().to_string(), + output_map, + ), Vec::new(), ) .map_err(|err| ParseError::ExtractionFailed { @@ -852,7 +864,7 @@ fn parse_sections(content: &str) -> IndexMap { fn baml_value_fields( value: &BamlValue, -) -> Option<&crate::baml_bridge::baml_types::BamlMap> { +) -> Option<&crate::bamltype::baml_types::BamlMap> { match value { BamlValue::Class(_, fields) => Some(fields), BamlValue::Map(fields) => Some(fields), @@ -883,14 +895,10 @@ fn format_baml_value_for_prompt_typed( } }; - crate::baml_bridge::internal_baml_jinja::format_baml_value(value, output_format, format) + crate::bamltype::internal_baml_jinja::format_baml_value(value, output_format, format) .unwrap_or_else(|_| "".to_string()) } -fn collect_flags(value: &BamlValueWithFlags, flags: &mut Vec) { - collect_flags_recursive(value, flags); -} - fn collect_flags_recursive(value: &BamlValueWithFlags, flags: &mut Vec) { match value { BamlValueWithFlags::String(v) => { diff --git a/crates/dspy-rs/src/core/lm/client_registry.rs b/crates/dspy-rs/src/core/lm/client_registry.rs index 202a6dde..cc52682d 100644 --- a/crates/dspy-rs/src/core/lm/client_registry.rs +++ b/crates/dspy-rs/src/core/lm/client_registry.rs @@ -320,8 +320,8 @@ impl LMClient { let (provider, model_id) = model_str.split_once(':').ok_or(anyhow::anyhow!( "Model string must be in format 'provider:model_name'" ))?; - tracing::Span::current().record("provider", &tracing::field::display(provider)); - tracing::Span::current().record("model_id", &tracing::field::display(model_id)); + tracing::Span::current().record("provider", tracing::field::display(provider)); + tracing::Span::current().record("model_id", tracing::field::display(model_id)); match provider { "openai" => { diff --git a/crates/dspy-rs/src/core/signature.rs b/crates/dspy-rs/src/core/signature.rs index f3401013..b91d4a10 100644 --- a/crates/dspy-rs/src/core/signature.rs +++ b/crates/dspy-rs/src/core/signature.rs @@ -37,8 +37,8 @@ pub trait MetaSignature: Send + Sync { } pub trait Signature: Send + Sync + 'static { - type Input: baml_bridge::BamlType + Send + Sync; - type Output: baml_bridge::BamlType + Send + Sync; + type Input: bamltype::compat::BamlTypeTrait + Send + Sync; + type Output: bamltype::compat::BamlTypeTrait + Send + Sync; fn instruction() -> &'static str; fn input_fields() -> &'static [FieldSpec]; diff --git a/crates/dspy-rs/src/lib.rs b/crates/dspy-rs/src/lib.rs index d68bca09..1f4991f9 100644 --- a/crates/dspy-rs/src/lib.rs +++ b/crates/dspy-rs/src/lib.rs @@ -1,3 +1,5 @@ +extern crate self as dspy_rs; + pub mod adapter; pub mod core; pub mod data; @@ -8,25 +10,31 @@ pub mod trace; pub mod utils; pub use adapter::chat::*; +pub use anyhow; pub use core::*; -pub use core::{ - CallResult, ConstraintKind, ConstraintResult, ConstraintSpec, ConversionError, ErrorClass, - FieldMeta, FieldSpec, JsonishError, LmError, ParseError, PredictError, Signature, -}; pub use data::*; pub use evaluate::*; +pub use indexmap; pub use optimizer::*; pub use predictors::*; +pub use schemars; +pub use serde; +pub use serde_json; pub use utils::*; -pub use baml_bridge; -pub use baml_bridge::BamlConvertError; -pub use baml_bridge::BamlType; -pub use baml_bridge::baml_types::{ +pub use bamltype; +pub use bamltype::BamlType; // attribute macro +pub use bamltype::baml_types::{ BamlValue, Constraint, ConstraintLevel, ResponseCheck, StreamingMode, TypeIR, }; -pub use baml_bridge::internal_baml_jinja::types::{OutputFormatContent, RenderOptions}; -pub use baml_bridge::jsonish::deserializer::deserialize_flags::Flag; +pub use bamltype::compat::BamlConvertError; +pub use bamltype::compat::{ + BamlAdapter, BamlTypeInternal, BamlTypeTrait, BamlValueConvert, Registry, ToBamlValue, + with_constraints, +}; +pub use bamltype::facet; +pub use bamltype::internal_baml_jinja::types::{OutputFormatContent, RenderOptions}; +pub use bamltype::jsonish::deserializer::deserialize_flags::Flag; pub use dsrs_macros::*; #[deprecated( @@ -38,8 +46,8 @@ macro_rules! example { // Pattern: { "key": <__dsrs_field_type>: "value", ... } { $($key:literal : $field_type:literal => $value:expr),* $(,)? } => {{ use std::collections::HashMap; - use dspy_rs::data::example::Example; - use dspy_rs::trace::{NodeType, record_node}; + use $crate::data::example::Example; + use $crate::trace::{NodeType, record_node}; let mut input_keys = vec![]; let mut output_keys = vec![]; @@ -55,7 +63,7 @@ macro_rules! example { } let tracked = { - use dspy_rs::trace::IntoTracked; + use $crate::trace::IntoTracked; $value.into_tracked() }; @@ -108,11 +116,11 @@ macro_rules! example { macro_rules! prediction { { $($key:literal => $value:expr),* $(,)? } => {{ use std::collections::HashMap; - use dspy_rs::{Prediction, LmUsage}; + use $crate::{Prediction, LmUsage}; let mut fields = HashMap::new(); $( - fields.insert($key.to_string(), serde_json::to_value($value).unwrap()); + fields.insert($key.to_string(), $crate::serde_json::to_value($value).unwrap()); )* Prediction::new(fields, LmUsage::default()) @@ -138,15 +146,15 @@ macro_rules! field { // Pattern for field definitions with descriptions { $($field_type:ident[$desc:literal] => $field_name:ident : $field_ty:ty),* $(,)? } => {{ - use serde_json::json; + use $crate::serde_json::json; - let mut result = serde_json::Map::new(); + let mut result = $crate::serde_json::Map::new(); $( let type_str = stringify!($field_ty); let schema = { - let schema = schemars::schema_for!($field_ty); - let schema_json = serde_json::to_value(schema).unwrap(); + let schema = $crate::schemars::schema_for!($field_ty); + let schema_json = $crate::serde_json::to_value(schema).unwrap(); // Extract just the properties if it's an object schema if let Some(obj) = schema_json.as_object() { if obj.contains_key("properties") { @@ -169,20 +177,20 @@ macro_rules! field { ); )* - serde_json::Value::Object(result) + $crate::serde_json::Value::Object(result) }}; // Pattern for field definitions without descriptions { $($field_type:ident => $field_name:ident : $field_ty:ty),* $(,)? } => {{ - use serde_json::json; + use $crate::serde_json::json; - let mut result = serde_json::Map::new(); + let mut result = $crate::serde_json::Map::new(); $( let type_str = stringify!($field_ty); let schema = { - let schema = schemars::schema_for!($field_ty); - let schema_json = serde_json::to_value(schema).unwrap(); + let schema = $crate::schemars::schema_for!($field_ty); + let schema_json = $crate::serde_json::to_value(schema).unwrap(); // Extract just the properties if it's an object schema if let Some(obj) = schema_json.as_object() { if obj.contains_key("properties") { @@ -205,7 +213,7 @@ macro_rules! field { ); )* - serde_json::Value::Object(result) + $crate::serde_json::Value::Object(result) }}; } @@ -231,7 +239,7 @@ macro_rules! sign { // Pattern: input fields -> output fields { ($($input_name:ident : $input_type:ty),* $(,)?) -> $($output_name:ident : $output_type:ty),* $(,)? } => {{ - #[derive(::dspy_rs::Signature, Clone)] + #[derive($crate::Signature, Clone)] struct __InlineSignature { $( #[input] @@ -243,7 +251,7 @@ macro_rules! sign { )* } - ::dspy_rs::Predict::<__InlineSignature>::new() + $crate::Predict::<__InlineSignature>::new() }}; } diff --git a/crates/dspy-rs/src/optimizer/copro.rs b/crates/dspy-rs/src/optimizer/copro.rs index b1d357d9..83ad0dd6 100644 --- a/crates/dspy-rs/src/optimizer/copro.rs +++ b/crates/dspy-rs/src/optimizer/copro.rs @@ -1,6 +1,5 @@ #![allow(deprecated)] -use crate as dspy_rs; use crate::{ Evaluator, Example, LM, LegacyPredict, Module, Optimizable, Optimizer, Prediction, Predictor, example, get_lm, diff --git a/crates/dspy-rs/src/optimizer/gepa.rs b/crates/dspy-rs/src/optimizer/gepa.rs index 14662987..2d51f14a 100644 --- a/crates/dspy-rs/src/optimizer/gepa.rs +++ b/crates/dspy-rs/src/optimizer/gepa.rs @@ -15,7 +15,6 @@ use bon::Builder; use serde::{Deserialize, Serialize}; use std::sync::Arc; -use crate as dspy_rs; use crate::{ Example, LM, LegacyPredict, Module, Optimizable, Optimizer, Prediction, Predictor, evaluate::FeedbackEvaluator, example, diff --git a/crates/dspy-rs/src/optimizer/mipro.rs b/crates/dspy-rs/src/optimizer/mipro.rs index 1a4d2f94..d3b760e3 100644 --- a/crates/dspy-rs/src/optimizer/mipro.rs +++ b/crates/dspy-rs/src/optimizer/mipro.rs @@ -1,6 +1,5 @@ #![allow(deprecated)] -use crate as dspy_rs; /// MIPROv2 Optimizer Implementation /// /// Multi-prompt Instruction Proposal Optimizer (MIPROv2) is an advanced optimizer diff --git a/crates/dspy-rs/src/predictors/predict.rs b/crates/dspy-rs/src/predictors/predict.rs index 394ba5c4..cb280cac 100644 --- a/crates/dspy-rs/src/predictors/predict.rs +++ b/crates/dspy-rs/src/predictors/predict.rs @@ -8,8 +8,8 @@ use std::sync::Arc; use tracing::{debug, trace}; use crate::adapter::Adapter; -use crate::baml_bridge::baml_types::BamlMap; -use crate::baml_bridge::{BamlValueConvert, ToBamlValue}; +use crate::bamltype::baml_types::BamlMap; +use crate::bamltype::compat::{BamlValueConvert, ToBamlValue}; use crate::core::{FieldSpec, MetaSignature, Module, Optimizable, Signature}; use crate::{ BamlValue, CallResult, Chat, ChatAdapter, Example, GLOBAL_SETTINGS, LM, LmError, LmUsage, diff --git a/crates/dspy-rs/tests/test_bamltype_attr_contract.rs b/crates/dspy-rs/tests/test_bamltype_attr_contract.rs new file mode 100644 index 00000000..37046b52 --- /dev/null +++ b/crates/dspy-rs/tests/test_bamltype_attr_contract.rs @@ -0,0 +1,24 @@ +use dspy_rs::{BamlType, BamlTypeTrait, RenderOptions}; + +#[derive(Debug, Clone, PartialEq)] +#[BamlType] +#[baml(internal_name = "contract::DsrsUser")] +struct DsrsUser { + #[baml(alias = "fullName")] + name: String, + age: u32, +} + +#[test] +fn bamltype_attribute_macro_works_from_dspy_rs() { + let schema = ::baml_output_format() + .render(RenderOptions::default()) + .expect("render schema") + .unwrap_or_default(); + assert!(schema.contains("fullName")); + + let raw = r#"{ "fullName": "Ada", "age": 36 }"#; + let parsed = dspy_rs::bamltype::parse_llm_output::(raw, true).expect("parse"); + assert_eq!(parsed.value.name, "Ada"); + assert_eq!(parsed.value.age, 36); +} diff --git a/crates/dspy-rs/tests/test_bamltype_docs_contract.rs b/crates/dspy-rs/tests/test_bamltype_docs_contract.rs new file mode 100644 index 00000000..8605cd33 --- /dev/null +++ b/crates/dspy-rs/tests/test_bamltype_docs_contract.rs @@ -0,0 +1,158 @@ +use dspy_rs::bamltype::HoistClasses; +use dspy_rs::{BamlType, BamlTypeTrait, ChatAdapter, RenderOptions, Signature}; + +#[derive(Clone, Debug)] +#[BamlType] +struct AliasPayload { + #[baml(alias = "fullName")] + full_name: String, +} + +#[derive(Clone, Debug)] +#[BamlType] +#[baml(rename_all = "camelCase")] +struct RenamedFieldsPayload { + user_name: String, + created_at: String, +} + +#[derive(Clone, Debug)] +#[BamlType] +struct SkipDefaultPayload { + content: String, + #[baml(skip)] + internal_id: i64, + #[baml(default)] + retries: i32, +} + +#[derive(Clone, Debug)] +#[BamlType] +struct BigIdPayload { + #[baml(int_repr = "string")] + large_id: u64, +} + +#[derive(Clone, Debug)] +#[BamlType] +#[baml(name = "UserProfile")] +struct NamedPayload { + name: String, +} + +#[derive(Signature, Clone, Debug)] +/// Contract test for docs-visible type behavior in prompts. +struct DocsTypeEffectsSig { + #[input] + question: String, + + #[output] + alias_payload: AliasPayload, + + #[output] + renamed_payload: RenamedFieldsPayload, + + #[output] + skip_default_payload: SkipDefaultPayload, + + #[output] + big_id_payload: BigIdPayload, + + #[output] + named_payload: NamedPayload, +} + +fn system_message() -> String { + let adapter = ChatAdapter; + adapter + .format_system_message_typed::() + .expect("system message") +} + +fn extract_field_block(message: &str, field_name: &str) -> String { + let marker = format!("[[ ## {field_name} ## ]]"); + let start = message + .find(&marker) + .unwrap_or_else(|| panic!("missing marker: {field_name}")); + let after = start + marker.len(); + let remaining = &message[after..]; + let end = remaining.find("[[ ##").unwrap_or(remaining.len()); + remaining[..end].trim().to_string() +} + +fn find_line<'a>(block: &'a str, needle: &str) -> &'a str { + block + .lines() + .find(|line| line.contains(needle)) + .unwrap_or_else(|| panic!("missing line containing {needle:?} in:\n{block}")) +} + +#[test] +fn alias_is_visible_to_model_schema() { + let block = extract_field_block(&system_message(), "alias_payload"); + assert!(block.contains("fullName")); + assert!(!block.contains("full_name")); +} + +#[test] +fn rename_all_is_visible_to_model_schema() { + let block = extract_field_block(&system_message(), "renamed_payload"); + assert!(block.contains("userName")); + assert!(block.contains("createdAt")); + assert!(!block.contains("user_name")); + assert!(!block.contains("created_at")); +} + +#[test] +fn skip_hides_field_and_default_marks_optional() { + let block = extract_field_block(&system_message(), "skip_default_payload"); + assert!(block.contains("content")); + assert!(!block.contains("internal_id")); + + let retries_line = find_line(&block, "retries"); + assert!( + retries_line.contains('?') || retries_line.contains("null"), + "expected optional retries marker in line: {retries_line}" + ); +} + +#[test] +fn int_repr_changes_prompt_type_to_string() { + let block = extract_field_block(&system_message(), "big_id_payload"); + let id_line = find_line(&block, "large_id"); + assert!( + id_line.contains("string"), + "expected string type for int_repr=\"string\": {id_line}" + ); +} + +#[test] +fn name_changes_type_label_in_prompt_and_hoisted_render() { + let block = extract_field_block(&system_message(), "named_payload"); + assert!( + block.contains("Output field `named_payload` should be of type: UserProfile"), + "expected renamed type label in prompt:\n{block}" + ); + + let rendered = ::baml_output_format() + .render(RenderOptions::hoist_classes(HoistClasses::All)) + .expect("render") + .unwrap_or_default(); + assert!( + rendered.contains("UserProfile"), + "expected renamed class in hoisted render:\n{rendered}" + ); +} + +#[test] +fn parse_behavior_matches_skip_and_default_claims() { + let parsed = dspy_rs::bamltype::parse_llm_output::( + r#"{ "content": "hello" }"#, + true, + ) + .expect("parse"); + + assert_eq!(parsed.value.content, "hello"); + assert_eq!(parsed.value.internal_id, 0); + assert_eq!(parsed.value.retries, 0); +} diff --git a/crates/dspy-rs/tests/test_input_format.rs b/crates/dspy-rs/tests/test_input_format.rs index e9442819..129ee475 100644 --- a/crates/dspy-rs/tests/test_input_format.rs +++ b/crates/dspy-rs/tests/test_input_format.rs @@ -1,9 +1,9 @@ -use dspy_rs::baml_bridge::ToBamlValue; -use dspy_rs::{BamlType, BamlValue, ChatAdapter, Signature}; +use dspy_rs::bamltype::compat::ToBamlValue; +use dspy_rs::{BamlType, BamlTypeTrait, BamlValue, ChatAdapter, Signature}; -#[derive(BamlType, Clone, Debug)] +#[derive(Clone, Debug)] +#[BamlType] struct Document { - #[baml(alias = "docText")] text: String, } @@ -83,7 +83,7 @@ fn extract_baml_field<'a>(value: &'a BamlValue, field_name: &str) -> &'a BamlVal } #[test] -fn typed_input_format_yaml_preserves_aliases() { +fn typed_input_format_yaml_renders_field_names() { let adapter = ChatAdapter; let input = FormatSigInput { question: "What is YAML?".to_string(), @@ -96,12 +96,12 @@ fn typed_input_format_yaml_preserves_aliases() { let context_value = extract_field(&message, "context"); let question_value = extract_field(&message, "question"); - assert!(context_value.contains("docText: Hello")); + assert!(context_value.contains("text: Hello")); assert_eq!(question_value, "What is YAML?"); } #[test] -fn typed_input_format_json_is_parsable_with_aliases() { +fn typed_input_format_json_is_parsable() { let adapter = ChatAdapter; let input = FormatJsonSigInput { question: "What is JSON?".to_string(), @@ -119,7 +119,7 @@ fn typed_input_format_json_is_parsable_with_aliases() { .and_then(|items| items.first()) .and_then(|value| value.as_object()) .expect("expected array with object"); - assert_eq!(first.get("docText").and_then(|v| v.as_str()), Some("Hello")); + assert_eq!(first.get("text").and_then(|v| v.as_str()), Some("Hello")); } #[test] @@ -137,8 +137,8 @@ fn typed_input_format_toon_matches_formatter() { let baml_value = input.to_baml_value(); let context_baml = extract_baml_field(&baml_value, "context"); - let output_format = ::baml_output_format(); - let expected = dspy_rs::baml_bridge::internal_baml_jinja::format_baml_value( + let output_format = ::baml_output_format(); + let expected = dspy_rs::bamltype::internal_baml_jinja::format_baml_value( context_baml, output_format, "toon", @@ -182,5 +182,5 @@ fn typed_input_default_non_string_is_json() { .and_then(|items| items.first()) .and_then(|value| value.as_object()) .expect("expected array with object"); - assert_eq!(first.get("docText").and_then(|v| v.as_str()), Some("Hello")); + assert_eq!(first.get("text").and_then(|v| v.as_str()), Some("Hello")); } diff --git a/crates/dspy-rs/tests/test_typed_prompt_format.rs b/crates/dspy-rs/tests/test_typed_prompt_format.rs index bb384ed8..78307676 100644 --- a/crates/dspy-rs/tests/test_typed_prompt_format.rs +++ b/crates/dspy-rs/tests/test_typed_prompt_format.rs @@ -1,6 +1,7 @@ use dspy_rs::{BamlType, ChatAdapter, Signature}; -#[derive(Clone, Debug, BamlType)] +#[derive(Clone, Debug)] +#[BamlType] /// A citation reference. struct Citation { /// Document identifier @@ -9,7 +10,8 @@ struct Citation { quote: String, } -#[derive(Clone, Debug, BamlType)] +#[derive(Clone, Debug)] +#[BamlType] /// Sentiment classification. enum Sentiment { Positive, diff --git a/crates/dsrs-macros/Cargo.toml b/crates/dsrs-macros/Cargo.toml index b3e67628..1c5de3ea 100644 --- a/crates/dsrs-macros/Cargo.toml +++ b/crates/dsrs-macros/Cargo.toml @@ -17,13 +17,9 @@ proc-macro = true syn = { version = "2", features = ["full"] } quote = "1" proc-macro2 = "1" -serde = { version = "1.0.219", features = ["derive"] } +proc-macro-crate = "3.2" serde_json = { version = "1.0.143", features = ["preserve_order"] } -anyhow = "1.0.99" -schemars = "1.0.4" -indexmap = "2.11.0" [dev-dependencies] dspy-rs = { path = "../dspy-rs" } trybuild = "1.0.110" -baml-bridge = { path = "../baml-bridge" } diff --git a/crates/dsrs-macros/src/lib.rs b/crates/dsrs-macros/src/lib.rs index ca2fa1cb..ae9ff9d4 100644 --- a/crates/dsrs-macros/src/lib.rs +++ b/crates/dsrs-macros/src/lib.rs @@ -1,5 +1,3 @@ -extern crate self as dsrs_macros; - use proc_macro::TokenStream; use quote::{format_ident, quote}; use serde_json::{Value, json}; @@ -12,6 +10,9 @@ use syn::{ }; mod optim; +mod runtime_path; + +use runtime_path::resolve_dspy_rs_path; #[proc_macro_derive(Optimizable, attributes(parameter))] pub fn derive_optimizable(input: TokenStream) -> TokenStream { @@ -21,13 +22,21 @@ pub fn derive_optimizable(input: TokenStream) -> TokenStream { #[proc_macro_derive(Signature, attributes(input, output, check, assert, alias, format))] pub fn derive_signature(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as DeriveInput); - match expand_signature(&input) { + let runtime = match resolve_dspy_rs_path() { + Ok(path) => path, + Err(err) => return err.to_compile_error().into(), + }; + + match expand_signature(&input, &runtime) { Ok(tokens) => tokens.into(), Err(err) => err.to_compile_error().into(), } } -fn expand_signature(input: &DeriveInput) -> syn::Result { +fn expand_signature( + input: &DeriveInput, + runtime: &syn::Path, +) -> syn::Result { let data = match &input.data { Data::Struct(data) => data, _ => { @@ -49,7 +58,7 @@ fn expand_signature(input: &DeriveInput) -> syn::Result syn::Result { for attr in &field.attrs { if attr.path().is_ident("input") { is_input = true; - if let Some(desc) = parse_desc_from_attr(attr) { + if let Some(desc) = parse_desc_from_attr(attr, "input")? { desc_override = Some(desc); } } else if attr.path().is_ident("output") { is_output = true; - if let Some(desc) = parse_desc_from_attr(attr) { + if let Some(desc) = parse_desc_from_attr(attr, "output")? { desc_override = Some(desc); } } else if attr.path().is_ident("alias") { @@ -248,10 +257,37 @@ fn parse_single_field(field: &syn::Field) -> syn::Result { }) } -fn parse_desc_from_attr(attr: &Attribute) -> Option { - let list = attr.meta.require_list().ok()?; - let desc = parse_desc_from_tokens(list.tokens.clone()); - if desc.is_empty() { None } else { Some(desc) } +fn parse_desc_from_attr(attr: &Attribute, attr_name: &str) -> syn::Result> { + match &attr.meta { + Meta::Path(_) => Ok(None), + Meta::List(list) => { + let metas = list.parse_args_with( + syn::punctuated::Punctuated::::parse_terminated, + )?; + + if metas.is_empty() { + return Ok(None); + } + + if metas.len() == 1 + && let Some(Meta::NameValue(meta)) = metas.first() + && meta.path.is_ident("desc") + { + return Ok(Some(parse_string_expr(&meta.value, meta.span())?)); + } + + Err(syn::Error::new_spanned( + attr, + format!( + "unsupported arguments for #[{attr_name}(...)]; only desc = \"...\" is allowed" + ), + )) + } + _ => Err(syn::Error::new_spanned( + attr, + format!("expected #[{attr_name}] or #[{attr_name}(desc = \"...\")]"), + )), + } } fn parse_string_attr(attr: &Attribute, attr_name: &str) -> syn::Result { @@ -272,7 +308,8 @@ fn parse_constraint_attr( attr: &Attribute, kind: ParsedConstraintKind, ) -> syn::Result { - let args: ConstraintArgs = attr.parse_args()?; + let mut args: ConstraintArgs = attr.parse_args()?; + normalize_constraint_expression(&mut args.expression); if kind == ParsedConstraintKind::Check && args.label.is_none() { return Err(syn::Error::new_spanned( attr, @@ -287,6 +324,17 @@ fn parse_constraint_attr( }) } +fn normalize_constraint_expression(expression: &mut String) { + // Accept common Rust-style logical operators in docs/examples and normalize + // to the Jinja expression syntax expected by downstream evaluation. + let normalized = expression + .replace(" && ", " and ") + .replace(" || ", " or ") + .replace("&&", " and ") + .replace("||", " or "); + *expression = normalized; +} + fn collect_doc_comment(attrs: &[Attribute]) -> String { let mut docs = Vec::new(); for attr in attrs { @@ -320,15 +368,16 @@ fn parse_string_expr(expr: &Expr, span: proc_macro2::Span) -> syn::Result syn::Result { let name = &input.ident; let vis = &input.vis; - let helper_structs = generate_helper_structs(name, parsed, vis)?; - let input_fields = generate_field_specs(name, &parsed.input_fields, "INPUT")?; - let output_fields = generate_field_specs(name, &parsed.output_fields, "OUTPUT")?; - let baml_delegation = generate_baml_delegation(name, parsed); - let signature_impl = generate_signature_impl(name, parsed); + let helper_structs = generate_helper_structs(name, parsed, vis, runtime)?; + let input_fields = generate_field_specs(name, &parsed.input_fields, "INPUT", runtime)?; + let output_fields = generate_field_specs(name, &parsed.output_fields, "OUTPUT", runtime)?; + let baml_delegation = generate_baml_delegation(name, parsed, runtime); + let signature_impl = generate_signature_impl(name, parsed, runtime); Ok(quote! { #helper_structs @@ -343,6 +392,7 @@ fn generate_helper_structs( name: &Ident, parsed: &ParsedSignature, vis: &Visibility, + runtime: &syn::Path, ) -> syn::Result { let input_name = format_ident!("{}Input", name); let output_name = format_ident!("__{}Output", name); @@ -353,17 +403,18 @@ fn generate_helper_structs( let all_fields: Vec<_> = parsed.all_fields.iter().map(field_tokens).collect(); Ok(quote! { - #[derive(Debug, Clone, ::dspy_rs::BamlType)] + #[#runtime::BamlType] + #[derive(Debug, Clone)] #vis struct #input_name { #(#input_fields),* } - #[derive(Debug, Clone, ::dspy_rs::BamlType)] + #[#runtime::BamlType] pub struct #output_name { #(#output_fields),* } - #[derive(Debug, Clone, ::dspy_rs::BamlType)] + #[#runtime::BamlType] pub struct #all_name { #(#all_fields),* } @@ -380,24 +431,9 @@ fn field_tokens(field: &ParsedField) -> proc_macro2::TokenStream { attrs.push(quote! { #[doc = #doc] }); } - if let Some(alias) = &field.alias { - let alias = LitStr::new(alias, proc_macro2::Span::call_site()); - attrs.push(quote! { #[baml(alias = #alias)] }); - } - - for constraint in &field.constraints { - let expr = LitStr::new(&constraint.expression, proc_macro2::Span::call_site()); - let label = constraint.label.as_deref().unwrap_or(""); - let label = LitStr::new(label, proc_macro2::Span::call_site()); - match constraint.kind { - ParsedConstraintKind::Check => { - attrs.push(quote! { #[baml(check(label = #label, expr = #expr))] }); - } - ParsedConstraintKind::Assert => { - attrs.push(quote! { #[baml(assert(label = #label, expr = #expr))] }); - } - } - } + // Note: aliases and constraints are handled at the FieldSpec level in + // generate_field_specs, not via struct attributes. The adapter layer uses + // FieldSpec metadata for LLM name mapping and constraint enforcement. quote! { #(#attrs)* @@ -409,6 +445,7 @@ fn generate_field_specs( name: &Ident, fields: &[ParsedField], kind: &str, + runtime: &syn::Path, ) -> syn::Result { let prefix = name.to_string().to_lowercase(); let array_name = format_ident!("__{}_{}_FIELDS", name.to_string().to_uppercase(), kind); @@ -438,8 +475,8 @@ fn generate_field_specs( if field.constraints.is_empty() { type_ir_fns.push(quote! { - fn #type_ir_fn_name() -> ::dspy_rs::TypeIR { - <#ty as ::dspy_rs::baml_bridge::BamlTypeInternal>::baml_type_ir() + fn #type_ir_fn_name() -> #runtime::TypeIR { + <#ty as #runtime::bamltype::compat::BamlTypeInternal>::baml_type_ir() } }); } else { @@ -452,19 +489,19 @@ fn generate_field_specs( let label = LitStr::new(label, proc_macro2::Span::call_site()); match constraint.kind { ParsedConstraintKind::Check => { - quote! { ::dspy_rs::Constraint::new_check(#label, #expr) } + quote! { #runtime::Constraint::new_check(#label, #expr) } } ParsedConstraintKind::Assert => { - quote! { ::dspy_rs::Constraint::new_assert(#label, #expr) } + quote! { #runtime::Constraint::new_assert(#label, #expr) } } } }) .collect(); type_ir_fns.push(quote! { - fn #type_ir_fn_name() -> ::dspy_rs::TypeIR { - let base = <#ty as ::dspy_rs::baml_bridge::BamlTypeInternal>::baml_type_ir(); - ::dspy_rs::baml_bridge::with_constraints(base, vec![#(#constraint_tokens),*]) + fn #type_ir_fn_name() -> #runtime::TypeIR { + let base = <#ty as #runtime::bamltype::compat::BamlTypeInternal>::baml_type_ir(); + #runtime::bamltype::compat::with_constraints(base, vec![#(#constraint_tokens),*]) } }); } @@ -477,7 +514,7 @@ fn generate_field_specs( if field.constraints.is_empty() { constraint_arrays.push(quote! { - const #constraints_name: &[::dspy_rs::ConstraintSpec] = &[]; + const #constraints_name: &[#runtime::ConstraintSpec] = &[]; }); } else { let constraint_specs: Vec<_> = field @@ -485,16 +522,16 @@ fn generate_field_specs( .iter() .map(|constraint| { let kind = match constraint.kind { - ParsedConstraintKind::Check => quote! { ::dspy_rs::ConstraintKind::Check }, + ParsedConstraintKind::Check => quote! { #runtime::ConstraintKind::Check }, ParsedConstraintKind::Assert => { - quote! { ::dspy_rs::ConstraintKind::Assert } + quote! { #runtime::ConstraintKind::Assert } } }; let label = constraint.label.as_deref().unwrap_or(""); let label = LitStr::new(label, proc_macro2::Span::call_site()); let expr = LitStr::new(&constraint.expression, proc_macro2::Span::call_site()); quote! { - ::dspy_rs::ConstraintSpec { + #runtime::ConstraintSpec { kind: #kind, label: #label, expression: #expr, @@ -504,14 +541,14 @@ fn generate_field_specs( .collect(); constraint_arrays.push(quote! { - const #constraints_name: &[::dspy_rs::ConstraintSpec] = &[ + const #constraints_name: &[#runtime::ConstraintSpec] = &[ #(#constraint_specs),* ]; }); } field_specs.push(quote! { - ::dspy_rs::FieldSpec { + #runtime::FieldSpec { name: #llm_name, rust_name: #rust_name, description: #description, @@ -526,13 +563,17 @@ fn generate_field_specs( #(#type_ir_fns)* #(#constraint_arrays)* - static #array_name: &[::dspy_rs::FieldSpec] = &[ + static #array_name: &[#runtime::FieldSpec] = &[ #(#field_specs),* ]; }) } -fn generate_baml_delegation(name: &Ident, parsed: &ParsedSignature) -> proc_macro2::TokenStream { +fn generate_baml_delegation( + name: &Ident, + parsed: &ParsedSignature, + runtime: &syn::Path, +) -> proc_macro2::TokenStream { let all_name = format_ident!("__{}All", name); let field_names: Vec<_> = parsed.all_fields.iter().map(|field| &field.ident).collect(); @@ -543,32 +584,32 @@ fn generate_baml_delegation(name: &Ident, parsed: &ParsedSignature) -> proc_macr to_value_inserts.push(quote! { fields.insert( #field_name.to_string(), - ::dspy_rs::baml_bridge::ToBamlValue::to_baml_value(&self.#ident), + #runtime::bamltype::compat::ToBamlValue::to_baml_value(&self.#ident), ); }); } quote! { - impl ::dspy_rs::baml_bridge::BamlTypeInternal for #name { + impl #runtime::bamltype::compat::BamlTypeInternal for #name { fn baml_internal_name() -> &'static str { - <#all_name as ::dspy_rs::baml_bridge::BamlTypeInternal>::baml_internal_name() + <#all_name as #runtime::bamltype::compat::BamlTypeInternal>::baml_internal_name() } - fn baml_type_ir() -> ::dspy_rs::TypeIR { - <#all_name as ::dspy_rs::baml_bridge::BamlTypeInternal>::baml_type_ir() + fn baml_type_ir() -> #runtime::TypeIR { + <#all_name as #runtime::bamltype::compat::BamlTypeInternal>::baml_type_ir() } - fn register(reg: &mut ::dspy_rs::baml_bridge::Registry) { - <#all_name as ::dspy_rs::baml_bridge::BamlTypeInternal>::register(reg) + fn register(reg: &mut #runtime::bamltype::compat::Registry) { + <#all_name as #runtime::bamltype::compat::BamlTypeInternal>::register(reg) } } - impl ::dspy_rs::baml_bridge::BamlValueConvert for #name { + impl #runtime::bamltype::compat::BamlValueConvert for #name { fn try_from_baml_value( - value: ::dspy_rs::BamlValue, + value: #runtime::BamlValue, path: Vec, - ) -> Result { - let all = <#all_name as ::dspy_rs::baml_bridge::BamlValueConvert> + ) -> Result { + let all = <#all_name as #runtime::bamltype::compat::BamlValueConvert> ::try_from_baml_value(value, path)?; Ok(Self { #(#field_names: all.#field_names),* @@ -576,18 +617,19 @@ fn generate_baml_delegation(name: &Ident, parsed: &ParsedSignature) -> proc_macr } } - impl ::dspy_rs::baml_bridge::BamlType for #name { - fn baml_output_format() -> &'static ::dspy_rs::OutputFormatContent { - <#all_name as ::dspy_rs::baml_bridge::BamlType>::baml_output_format() + impl #runtime::bamltype::compat::BamlTypeTrait for #name { + fn baml_output_format() -> &'static #runtime::OutputFormatContent { + <#all_name as #runtime::bamltype::compat::BamlTypeTrait>::baml_output_format() } } - impl ::dspy_rs::baml_bridge::ToBamlValue for #name { - fn to_baml_value(&self) -> ::dspy_rs::BamlValue { - let mut fields = ::dspy_rs::baml_bridge::baml_types::BamlMap::new(); + impl #runtime::bamltype::compat::ToBamlValue for #name { + fn to_baml_value(&self) -> #runtime::BamlValue { + let mut fields = #runtime::bamltype::baml_types::BamlMap::new(); #(#to_value_inserts)* - ::dspy_rs::baml_bridge::baml_types::BamlValue::Class( - ::baml_internal_name().to_string(), + #runtime::bamltype::baml_types::BamlValue::Class( + ::baml_internal_name() + .to_string(), fields, ) } @@ -595,7 +637,11 @@ fn generate_baml_delegation(name: &Ident, parsed: &ParsedSignature) -> proc_macr } } -fn generate_signature_impl(name: &Ident, parsed: &ParsedSignature) -> proc_macro2::TokenStream { +fn generate_signature_impl( + name: &Ident, + parsed: &ParsedSignature, + runtime: &syn::Path, +) -> proc_macro2::TokenStream { let input_name = format_ident!("{}Input", name); let output_name = format_ident!("__{}Output", name); @@ -616,7 +662,7 @@ fn generate_signature_impl(name: &Ident, parsed: &ParsedSignature) -> proc_macro let output_fields_static = format_ident!("__{}_OUTPUT_FIELDS", name.to_string().to_uppercase()); quote! { - impl ::dspy_rs::Signature for #name { + impl #runtime::Signature for #name { type Input = #input_name; type Output = #output_name; @@ -624,16 +670,16 @@ fn generate_signature_impl(name: &Ident, parsed: &ParsedSignature) -> proc_macro #instruction } - fn input_fields() -> &'static [::dspy_rs::FieldSpec] { + fn input_fields() -> &'static [#runtime::FieldSpec] { &#input_fields_static } - fn output_fields() -> &'static [::dspy_rs::FieldSpec] { + fn output_fields() -> &'static [#runtime::FieldSpec] { &#output_fields_static } - fn output_format_content() -> &'static ::dspy_rs::OutputFormatContent { - <#output_name as ::dspy_rs::baml_bridge::BamlType>::baml_output_format() + fn output_format_content() -> &'static #runtime::OutputFormatContent { + <#output_name as #runtime::bamltype::compat::BamlTypeTrait>::baml_output_format() } fn from_parts(input: Self::Input, output: Self::Output) -> Self { @@ -661,6 +707,10 @@ fn generate_signature_impl(name: &Ident, parsed: &ParsedSignature) -> proc_macro #[proc_macro_attribute] pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { let input = parse_macro_input!(item as DeriveInput); + let runtime = match resolve_dspy_rs_path() { + Ok(path) => path, + Err(err) => return err.to_compile_error().into(), + }; // Parse the attributes (cot, hint, etc.) let attr_str = attr.to_string(); @@ -693,21 +743,41 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { let mut found_first_input = false; for field in &named.named { - let field_name = field.ident.as_ref().unwrap().clone(); + let field_name = match field.ident.as_ref() { + Some(name) => name.clone(), + None => { + return syn::Error::new_spanned( + field, + "LegacySignature requires named fields", + ) + .to_compile_error() + .into(); + } + }; let field_type = field.ty.clone(); // Check for #[input] or #[output] attributes - let (is_input, desc) = has_in_attribute(&field.attrs); - let (is_output, desc2) = has_out_attribute(&field.attrs); + let (is_input, desc) = has_io_attribute(&field.attrs, "input"); + let (is_output, desc2) = has_io_attribute(&field.attrs, "output"); if is_input && is_output { - panic!("Field {field_name} cannot be both input and output"); + return syn::Error::new_spanned( + field, + format!("Field `{field_name}` cannot be both input and output"), + ) + .to_compile_error() + .into(); } if !is_input && !is_output { - panic!( - "Field {field_name} must have either #[input] or #[output] attribute" - ); + return syn::Error::new_spanned( + field, + format!( + "Field `{field_name}` must have either #[input] or #[output] attribute" + ), + ) + .to_compile_error() + .into(); } let field_desc = if is_input { desc } else { desc2 }; @@ -751,8 +821,8 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { let field_name_str = field_name.to_string(); schema_updates.push(quote! { { - let schema = schemars::schema_for!(#field_type); - let schema_json = serde_json::to_value(schema).unwrap(); + let schema = #runtime::schemars::schema_for!(#field_type); + let schema_json = #runtime::serde_json::to_value(schema).unwrap(); // Extract just the properties if it's an object schema if let Some(obj) = schema_json.as_object() { if obj.contains_key("properties") { @@ -773,8 +843,8 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { let field_name_str = field_name.to_string(); schema_updates.push(quote! { { - let schema = schemars::schema_for!(#field_type); - let schema_json = serde_json::to_value(schema).unwrap(); + let schema = #runtime::schemars::schema_for!(#field_type); + let schema_json = #runtime::serde_json::to_value(schema).unwrap(); // Extract just the properties if it's an object schema if let Some(obj) = schema_json.as_object() { if obj.contains_key("properties") { @@ -792,7 +862,14 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { } } } - _ => panic!("Signature can only be applied to structs"), + _ => { + return syn::Error::new_spanned( + &input, + "LegacySignature can only be applied to structs with named fields", + ) + .to_compile_error() + .into(); + } } if has_hint { @@ -809,18 +886,18 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { let output_schema_str = serde_json::to_string(&output_schema).unwrap(); let generated = quote! { - #[derive(Default, Debug, Clone, serde::Serialize, serde::Deserialize)] + #[derive(Default, Debug, Clone, #runtime::serde::Serialize, #runtime::serde::Deserialize)] struct #struct_name { instruction: String, - input_fields: serde_json::Value, - output_fields: serde_json::Value, - demos: Vec, + input_fields: #runtime::serde_json::Value, + output_fields: #runtime::serde_json::Value, + demos: Vec<#runtime::Example>, } impl #struct_name { pub fn new() -> Self { - let mut input_fields: serde_json::Value = serde_json::from_str(#input_schema_str).unwrap(); - let mut output_fields: serde_json::Value = serde_json::from_str(#output_schema_str).unwrap(); + let mut input_fields: #runtime::serde_json::Value = #runtime::serde_json::from_str(#input_schema_str).unwrap(); + let mut output_fields: #runtime::serde_json::Value = #runtime::serde_json::from_str(#output_schema_str).unwrap(); // Update schemas for complex types #(#schema_updates)* @@ -842,12 +919,12 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { } } - impl dspy_rs::core::MetaSignature for #struct_name { - fn demos(&self) -> Vec { + impl #runtime::core::MetaSignature for #struct_name { + fn demos(&self) -> Vec<#runtime::Example> { self.demos.clone() } - fn set_demos(&mut self, demos: Vec) -> anyhow::Result<()> { + fn set_demos(&mut self, demos: Vec<#runtime::Example>) -> #runtime::anyhow::Result<()> { self.demos = demos; Ok(()) } @@ -856,20 +933,20 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { self.instruction.clone() } - fn input_fields(&self) -> serde_json::Value { + fn input_fields(&self) -> #runtime::serde_json::Value { self.input_fields.clone() } - fn output_fields(&self) -> serde_json::Value { + fn output_fields(&self) -> #runtime::serde_json::Value { self.output_fields.clone() } - fn update_instruction(&mut self, instruction: String) -> anyhow::Result<()> { + fn update_instruction(&mut self, instruction: String) -> #runtime::anyhow::Result<()> { self.instruction = instruction; Ok(()) } - fn append(&mut self, name: &str, field_value: serde_json::Value) -> anyhow::Result<()> { + fn append(&mut self, name: &str, field_value: #runtime::serde_json::Value) -> #runtime::anyhow::Result<()> { match field_value["__dsrs_field_type"].as_str() { Some("input") => { self.input_fields[name] = field_value; @@ -878,7 +955,7 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { self.output_fields[name] = field_value; } _ => { - return Err(anyhow::anyhow!("Invalid field type: {:?}", field_value["__dsrs_field_type"].as_str())); + return Err(#runtime::anyhow::anyhow!("Invalid field type: {:?}", field_value["__dsrs_field_type"].as_str())); } } Ok(()) @@ -889,33 +966,17 @@ pub fn LegacySignature(attr: TokenStream, item: TokenStream) -> TokenStream { generated.into() } -fn has_in_attribute(attrs: &[Attribute]) -> (bool, String) { +fn has_io_attribute(attrs: &[Attribute], attr_name: &str) -> (bool, String) { for attr in attrs { - if attr.path().is_ident("input") { + if attr.path().is_ident(attr_name) { // Try to parse desc parameter if let Ok(list) = attr.meta.require_list() { let desc = parse_desc_from_tokens(list.tokens.clone()); return (true, desc); - } else { - // Just #[input] without parameters - return (true, String::new()); } - } - } - (false, String::new()) -} -fn has_out_attribute(attrs: &[Attribute]) -> (bool, String) { - for attr in attrs { - if attr.path().is_ident("output") { - // Try to parse desc parameter - if let Ok(list) = attr.meta.require_list() { - let desc = parse_desc_from_tokens(list.tokens.clone()); - return (true, desc); - } else { - // Just #[output] without parameters - return (true, String::new()); - } + // Just #[input] or #[output] without parameters. + return (true, String::new()); } } (false, String::new()) diff --git a/crates/dsrs-macros/src/optim.rs b/crates/dsrs-macros/src/optim.rs index 8b49498c..e59d8cd6 100644 --- a/crates/dsrs-macros/src/optim.rs +++ b/crates/dsrs-macros/src/optim.rs @@ -1,37 +1,53 @@ use proc_macro::TokenStream; use quote::quote; -use syn::{Data, DeriveInput, Field, Fields, parse_macro_input, parse_str}; +use syn::{Data, DeriveInput, Field, Fields, parse_macro_input}; + +use crate::runtime_path::resolve_dspy_rs_path; pub fn optimizable_impl(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as DeriveInput); - // Define trait path as a constant - easy to change in one place - let trait_path = parse_str::("::dspy_rs::core::module::Optimizable").unwrap(); + let runtime = match resolve_dspy_rs_path() { + Ok(path) => path, + Err(err) => return err.to_compile_error().into(), + }; + let trait_path: syn::Path = syn::parse_quote!(#runtime::core::module::Optimizable); // Extract parameter field names - let parameter_fields = extract_parameter_fields(&input); + let parameter_fields = match extract_parameter_fields(&input) { + Ok(fields) => fields, + Err(err) => return err.to_compile_error().into(), + }; let name = &input.ident; let generics = &input.generics; let (impl_generics, type_generics, where_clause) = generics.split_for_impl(); - let parameter_names: Vec<_> = parameter_fields - .iter() - .map(|field| field.ident.as_ref().unwrap()) - .collect(); + let mut parameter_names = Vec::with_capacity(parameter_fields.len()); + for field in ¶meter_fields { + let Some(ident) = field.ident.as_ref() else { + return syn::Error::new_spanned( + field, + "Optimizable can only be derived for structs with named fields", + ) + .to_compile_error() + .into(); + }; + parameter_names.push(ident); + } // Generate the Optimizable implementation (flatten nested parameters with compound names) let expanded = quote! { impl #impl_generics #trait_path for #name #type_generics #where_clause { fn parameters( &mut self, - ) -> indexmap::IndexMap<::std::string::String, &mut dyn #trait_path> { - let mut params: indexmap::IndexMap<::std::string::String, &mut dyn #trait_path> = indexmap::IndexMap::new(); + ) -> #runtime::indexmap::IndexMap<::std::string::String, &mut dyn #trait_path> { + let mut params: #runtime::indexmap::IndexMap<::std::string::String, &mut dyn #trait_path> = #runtime::indexmap::IndexMap::new(); #( { let __field_name = stringify!(#parameter_names).to_string(); // SAFETY: We only create disjoint mutable borrows to distinct struct fields let __field_ptr: *mut dyn #trait_path = &mut self.#parameter_names as *mut dyn #trait_path; - let __child_params: indexmap::IndexMap<::std::string::String, &mut dyn #trait_path> = unsafe { (&mut *__field_ptr).parameters() }; + let __child_params: #runtime::indexmap::IndexMap<::std::string::String, &mut dyn #trait_path> = unsafe { (&mut *__field_ptr).parameters() }; if __child_params.is_empty() { // Leaf: insert the field itself unsafe { @@ -53,21 +69,23 @@ pub fn optimizable_impl(input: TokenStream) -> TokenStream { TokenStream::from(expanded) } -fn extract_parameter_fields(input: &DeriveInput) -> Vec<&Field> { +fn extract_parameter_fields(input: &DeriveInput) -> syn::Result> { match &input.data { Data::Struct(data_struct) => match &data_struct.fields { - Fields::Named(fields_named) => fields_named + Fields::Named(fields_named) => Ok(fields_named .named .iter() .filter(|field| has_parameter_attribute(field)) - .collect(), - _ => { - panic!("Optimizable can only be derived for structs with named fields"); - } + .collect()), + _ => Err(syn::Error::new_spanned( + input, + "Optimizable can only be derived for structs with named fields", + )), }, - _ => { - panic!("Optimizable can only be derived for structs"); - } + _ => Err(syn::Error::new_spanned( + input, + "Optimizable can only be derived for structs", + )), } } diff --git a/crates/dsrs-macros/src/runtime_path.rs b/crates/dsrs-macros/src/runtime_path.rs new file mode 100644 index 00000000..e97307d3 --- /dev/null +++ b/crates/dsrs-macros/src/runtime_path.rs @@ -0,0 +1,19 @@ +use proc_macro_crate::{FoundCrate, crate_name}; +use proc_macro2::Span; + +pub(crate) fn resolve_dspy_rs_path() -> syn::Result { + match crate_name("dspy-rs") { + // `crate` fails in examples/binaries inside the dspy-rs package because + // there it points at the example crate, not the library. Use the crate + // alias (`extern crate self as dspy_rs`) for a stable path. + Ok(FoundCrate::Itself) => Ok(syn::parse_quote!(::dspy_rs)), + Ok(FoundCrate::Name(name)) => { + let ident = syn::Ident::new(&name.replace('-', "_"), Span::call_site()); + Ok(syn::parse_quote!(::#ident)) + } + Err(_) => Err(syn::Error::new( + Span::call_site(), + "could not resolve `dspy-rs`; add it as a dependency (renamed dependencies are supported)", + )), + } +} diff --git a/crates/dsrs-macros/tests/optim/derive_optimizable.rs b/crates/dsrs-macros/tests/optim/derive_optimizable.rs new file mode 100644 index 00000000..9108ee73 --- /dev/null +++ b/crates/dsrs-macros/tests/optim/derive_optimizable.rs @@ -0,0 +1,24 @@ +use dspy_rs::{Optimizable, Predict, Signature}; + +#[derive(Signature, Clone, Debug)] +struct QA { + #[input] + question: String, + + #[output] + answer: String, +} + +#[derive(Optimizable)] +struct Pipeline { + #[parameter] + qa: Predict, +} + +fn main() { + let mut pipeline = Pipeline { + qa: Predict::::new(), + }; + let params = dspy_rs::core::module::Optimizable::parameters(&mut pipeline); + let _qa = params.get("qa").expect("qa parameter should be present"); +} diff --git a/crates/dsrs-macros/tests/signature_derive.rs b/crates/dsrs-macros/tests/signature_derive.rs index 21f52873..27b8155c 100644 --- a/crates/dsrs-macros/tests/signature_derive.rs +++ b/crates/dsrs-macros/tests/signature_derive.rs @@ -12,6 +12,17 @@ struct TestSig { answer: String, } +/// Test logical operators are normalized to Jinja syntax. +#[derive(dsrs_macros::Signature)] +struct NormalizedConstraintSig { + #[input] + question: String, + + #[output] + #[check("this >= 0.0 && this <= 1.0", label = "valid_range")] + score: f64, +} + #[test] fn test_generates_input_struct() { let input = TestSigInput { @@ -46,7 +57,7 @@ fn test_from_parts_into_parts() { answer: "a".to_string(), }; - let full = TestSig::from_parts(input.clone(), output.clone()); + let full = TestSig::from_parts(input, output); assert_eq!(full.question, "q"); assert_eq!(full.answer, "a"); @@ -57,5 +68,16 @@ fn test_from_parts_into_parts() { #[test] fn test_baml_type_impl() { - let _ = ::baml_output_format(); + let _ = ::baml_output_format(); +} + +#[test] +fn test_constraint_operator_normalization() { + let output_fields = ::output_fields(); + assert_eq!(output_fields.len(), 1); + assert_eq!(output_fields[0].constraints.len(), 1); + assert_eq!( + output_fields[0].constraints[0].expression, + "this >= 0.0 and this <= 1.0" + ); } diff --git a/crates/dsrs-macros/tests/ui/input_unknown_arg.rs b/crates/dsrs-macros/tests/ui/input_unknown_arg.rs new file mode 100644 index 00000000..912cb91c --- /dev/null +++ b/crates/dsrs-macros/tests/ui/input_unknown_arg.rs @@ -0,0 +1,12 @@ +use dsrs_macros::Signature; + +#[derive(Signature)] +struct InputUnknownArg { + #[input(foo = "bar")] + question: String, + + #[output] + answer: String, +} + +fn main() {} diff --git a/crates/dsrs-macros/tests/ui/input_unknown_arg.stderr b/crates/dsrs-macros/tests/ui/input_unknown_arg.stderr new file mode 100644 index 00000000..f99d7395 --- /dev/null +++ b/crates/dsrs-macros/tests/ui/input_unknown_arg.stderr @@ -0,0 +1,5 @@ +error: unsupported arguments for #[input(...)]; only desc = "..." is allowed + --> tests/ui/input_unknown_arg.rs:5:5 + | +5 | #[input(foo = "bar")] + | ^^^^^^^^^^^^^^^^^^^^^ diff --git a/crates/dsrs-macros/tests/ui/output_unknown_arg.rs b/crates/dsrs-macros/tests/ui/output_unknown_arg.rs new file mode 100644 index 00000000..ba594d3e --- /dev/null +++ b/crates/dsrs-macros/tests/ui/output_unknown_arg.rs @@ -0,0 +1,12 @@ +use dsrs_macros::Signature; + +#[derive(Signature)] +struct OutputUnknownArg { + #[input] + question: String, + + #[output(bad = "x")] + answer: String, +} + +fn main() {} diff --git a/crates/dsrs-macros/tests/ui/output_unknown_arg.stderr b/crates/dsrs-macros/tests/ui/output_unknown_arg.stderr new file mode 100644 index 00000000..20f201f7 --- /dev/null +++ b/crates/dsrs-macros/tests/ui/output_unknown_arg.stderr @@ -0,0 +1,5 @@ +error: unsupported arguments for #[output(...)]; only desc = "..." is allowed + --> tests/ui/output_unknown_arg.rs:8:5 + | +8 | #[output(bad = "x")] + | ^^^^^^^^^^^^^^^^^^^^ diff --git a/docs/docs/building-blocks/signature.mdx b/docs/docs/building-blocks/signature.mdx index 7c261724..f06d61e6 100644 --- a/docs/docs/building-blocks/signature.mdx +++ b/docs/docs/building-blocks/signature.mdx @@ -93,19 +93,21 @@ struct SpamAnalysis { ## Custom types -When you have a non-standard type in a field, derive `BamlType` on it: +When you have a non-standard type in a field, add `#[BamlType]` on it: ```rust use dspy_rs::{Signature, BamlType}; -#[derive(BamlType, Clone, Debug)] +#[BamlType] +#[derive(Clone, Debug)] enum Sentiment { Positive, Negative, Neutral, } -#[derive(BamlType, Clone, Debug)] +#[BamlType] +#[derive(Clone, Debug)] struct Citation { /// Document ID doc_id: String, @@ -174,11 +176,12 @@ Options: `"json"`, `"yaml"`, `"toon"` ```rust #[output] -#[check("this >= 0.0 && this <= 1.0", label = "valid_range")] +#[check("this >= 0.0 and this <= 1.0", label = "valid_range")] confidence: f64, ``` Recorded in metadata, doesn't fail parsing. Label is required. +Use Jinja boolean operators (`and`, `or`) in expressions. ### `#[assert]` - Hard constraint diff --git a/docs/docs/building-blocks/types.mdx b/docs/docs/building-blocks/types.mdx index a91e07df..1d3fd584 100644 --- a/docs/docs/building-blocks/types.mdx +++ b/docs/docs/building-blocks/types.mdx @@ -1,245 +1,220 @@ --- title: 'Custom Types' -description: 'Define your own types for use in signatures' +description: 'Define your own types and see exactly how they change the model-facing schema' icon: 'shapes' --- -When your [signature](/docs/building-blocks/signature) fields use types beyond the built-ins, derive `BamlType` to make them work. +Types are not just Rust structure. They directly change what the model sees in the system prompt and what the parser accepts in responses. -## When you need it +If docs only say "use this attribute" without showing prompt impact, they are not useful. This page is written in user-visible terms. -**Built-in types** (no derive needed): +## Mental model + +When you add `#[BamlType]` to a type used in a `#[derive(Signature)]` output field: + +1. The type is rendered into the schema/instructions sent to the model. +2. The parser uses the same shape and attributes to parse the model response. +3. Attributes like `alias`, `skip`, and `default` affect both rendering and parse behavior. + +## When you need `#[BamlType]` + +Built-in types work without extra annotations: - `String`, `bool` - `i8`, `i16`, `i32`, `i64`, `f32`, `f64` - `Option`, `Vec`, `HashMap` - `Box`, `Arc`, `Rc` -**Custom types** (need `#[derive(BamlType)]`): -- Your own structs -- Your own enums -- Nested combinations of the above +Custom types need `#[BamlType]`: +- Your structs +- Your enums +- Nested custom types + +## What the model sees -## Structs +Example: ```rust -use dspy_rs::BamlType; +use dspy_rs::{BamlType, Signature}; -/// A citation from a document. -#[derive(BamlType, Clone, Debug)] -struct Citation { - /// The document identifier - doc_id: String, +#[derive(Clone, Debug)] +#[BamlType] +#[baml(rename_all = "camelCase")] +struct UserSummary { + user_name: String, + created_at: String, +} - /// Relevant quote from the source - quote: String, +#[derive(Signature, Clone, Debug)] +/// Summarize the user. +struct SummarizeUser { + #[input] + question: String, - /// Page number if available - page: Option, + #[output] + summary: UserSummary, } ``` -Doc comments become descriptions in the schema shown to the LLM. +Prompt-visible effect: +- Model sees `userName` and `createdAt`. +- Model does not see `user_name` or `created_at`. -## Enums +## Attribute impact (visible behavior) -### Simple enums +### `#[baml(alias = "...")]` ```rust -#[derive(BamlType, Clone, Debug)] -enum Sentiment { - Positive, - Negative, - Neutral, -} -``` - -### Enums with descriptions - -```rust -#[derive(BamlType, Clone, Debug)] -enum Priority { - /// Needs immediate attention - Critical, - /// Should be addressed soon - High, - /// Normal processing time - Medium, - /// Can wait - Low, +#[BamlType] +struct User { + #[baml(alias = "fullName")] + full_name: String, } ``` -### Data enums (tagged unions) +Visible effect: +- Prompt/schema uses `fullName`. +- Parser accepts model output keyed by `fullName`. -For enums where variants have different fields: +### `#[baml(rename_all = "...")]` ```rust -#[derive(BamlType, Clone, Debug)] -#[baml(tag = "type")] -enum Response { - Success { data: String }, - Error { code: i32, message: String }, +#[BamlType] +#[baml(rename_all = "camelCase")] +struct ApiResponse { + user_name: String, + created_at: String, } ``` -The `tag` attribute specifies the discriminator field name. LLM output looks like: -```json -{ "type": "Error", "code": 404, "message": "Not found" } -``` - -## Field attributes - -### `#[baml(alias = "name")]` - -Rename a field for the LLM: +Visible effect: +- Prompt/schema keys follow that naming convention. -```rust -#[derive(BamlType, Clone, Debug)] -struct User { - #[baml(alias = "userName")] - name: String, // LLM sees "userName", Rust uses "name" -} -``` +Options: `camelCase`, `snake_case`, `PascalCase`, `kebab-case`, `SCREAMING_SNAKE_CASE`. ### `#[baml(skip)]` -Exclude a field from the schema entirely: - ```rust -#[derive(BamlType, Clone, Debug)] -struct Document { +#[BamlType] +struct Doc { content: String, - #[baml(skip)] - internal_id: u64, // not sent to LLM, uses Default + internal_id: u64, } ``` -The field must implement `Default`. It won't appear in prompts or be expected in responses. +Visible effect: +- `internal_id` is omitted from the model-facing schema. -### `#[baml(default)]` +Parse effect: +- Field is filled via `Default` on parse. -Make a field optional with a default: +### `#[baml(default)]` ```rust -#[derive(BamlType, Clone, Debug)] +#[BamlType] struct Config { - name: String, // required - + name: String, #[baml(default)] - retries: i32, // optional, defaults to 0 if missing + retries: i32, } ``` -Different from `Option`: -- `Option` - explicitly nullable, LLM can return `null` -- `#[baml(default)]` - if LLM omits it, use `Default::default()` +Visible effect: +- Field is rendered as optional in schema output. -### `#[baml(check(...))]` and `#[baml(assert(...))]` +Parse effect: +- Missing `retries` parses as `Default::default()`. -Add constraints at the type level: +### `#[baml(name = "...")]` ```rust -#[derive(BamlType, Clone, Debug)] -struct Score { - #[baml(check(label = "valid_range", expr = "this >= 0 && this <= 100"))] - value: i32, - - #[baml(assert(label = "not_empty", expr = "this.len() > 0"))] - explanation: String, +#[BamlType] +#[baml(name = "UserProfile")] +struct User { + name: String, } ``` -See [Constraints](/docs/building-blocks/constraints) for the expression language. +Visible effect: +- Typed prompt type labels use `UserProfile`. +- Hoisted schema rendering also uses `UserProfile`. -## Container attributes +Important: +- Default rendering may inline objects, so class headers can be less visible. +- If you want class headers to be explicit, render with class hoisting enabled. -### `#[baml(rename_all = "...")]` - -Apply a naming convention to all fields: +### `#[baml(int_repr = "string")]` ```rust -#[derive(BamlType, Clone, Debug)] -#[baml(rename_all = "camelCase")] -struct ApiResponse { - user_name: String, // becomes "userName" - created_at: String, // becomes "createdAt" +#[BamlType] +struct BigIds { + #[baml(int_repr = "string")] + large_id: u64, } ``` -Options: `"camelCase"`, `"snake_case"`, `"PascalCase"`, `"kebab-case"`, `"SCREAMING_SNAKE_CASE"` +Visible effect: +- Prompt/schema shows `large_id` as string-like instead of int-like. -### `#[baml(name = "...")]` +Why: +- Helps with integer widths that exceed JSON number precision. -Override the type name in the schema: +### `#[baml(map_key_repr = "string" | "pairs")]` ```rust -#[derive(BamlType, Clone, Debug)] -#[baml(name = "UserProfile")] -struct User { - name: String, +#[BamlType] +struct Scores { + #[baml(map_key_repr = "string")] + by_id: std::collections::HashMap, } ``` -## Advanced: Large integers +Visible effect: +- Controls how non-string map keys are represented to the model. -For `u64`, `i128`, `u128` (which exceed JSON number precision): +## Enums + +### Simple enums ```rust -#[derive(BamlType, Clone, Debug)] -struct BigIds { - #[baml(int_repr = "string")] - large_id: u64, // serialized as "18446744073709551615" +#[BamlType] +enum Sentiment { + Positive, + Negative, + Neutral, } ``` -## Advanced: Map key types +Visible effect: +- Prompt/schema presents the allowed literal values. -For maps with non-string keys: +### Data enums (tagged unions) ```rust -#[derive(BamlType, Clone, Debug)] -struct Scores { - #[baml(map_key_repr = "string")] - by_id: HashMap, // keys become strings: {"1": 0.5, "2": 0.8} +#[BamlType] +#[baml(tag = "type")] +enum Response { + Success { data: String }, + Error { code: i32, message: String }, } ``` -## What's NOT supported - -These will give compile errors: +Visible effect: +- Prompt/schema uses `type` discriminator and variant-specific fields. -| Pattern | Why | -|---------|-----| -| Tuple structs `struct Foo(A, B)` | Use named fields instead | -| Unit structs `struct Foo;` | Nothing to serialize | -| Tuple enum variants `Enum::Var(A)` | Use `Var { field: A }` instead | -| `serde_json::Value` | Too dynamic, use concrete types | -| Trait objects `dyn Trait` | Not serializable | +## Unsupported patterns (compile-time errors) -## Nesting +- Tuple structs: `struct Foo(A, B)` +- Unit structs: `struct Foo;` +- Tuple enum variants: `Enum::Var(A)` +- `serde_json::Value` +- Trait objects (`dyn Trait`) -Types compose naturally: - -```rust -#[derive(BamlType, Clone, Debug)] -struct Author { - name: String, - email: Option, -} +## Contract tests for these docs -#[derive(BamlType, Clone, Debug)] -struct Book { - title: String, - authors: Vec, // nested custom type in vec - metadata: HashMap, -} +The user-visible claims on this page are locked by tests: -// For enums with data, use struct variants (not tuple variants): -#[derive(BamlType, Clone, Debug)] -#[baml(tag = "item_type")] -enum LibraryItem { - Book { book: Book }, - Magazine { title: String, issue: i32 }, -} -``` +- Prompt/render effects: `crates/dspy-rs/tests/test_bamltype_docs_contract.rs` +- Basic `#[BamlType]` end-user contract: `crates/dspy-rs/tests/test_bamltype_attr_contract.rs` +- Unsupported/compile-fail type shapes: `crates/bamltype/tests/ui.rs` +- Signature macro compile-fail coverage: `crates/dsrs-macros/tests/ui.rs` diff --git a/docs/docs/getting-started/quickstart.mdx b/docs/docs/getting-started/quickstart.mdx index 1be8c228..19e21107 100644 --- a/docs/docs/getting-started/quickstart.mdx +++ b/docs/docs/getting-started/quickstart.mdx @@ -166,12 +166,13 @@ Compose multi-step pipelines ### Custom types -When you need more than primitives, derive [`BamlType`](/docs/building-blocks/types): +When you need more than primitives, add [`#[BamlType]`](/docs/building-blocks/types): ```rust use dspy_rs::{Signature, BamlType}; -#[derive(BamlType, Clone, Debug)] +#[BamlType] +#[derive(Clone, Debug)] enum Sentiment { Positive, Negative, diff --git a/vendor/baml/crates/internal-baml-jinja/src/output_format/types.rs b/vendor/baml/crates/internal-baml-jinja/src/output_format/types.rs index 0e6cff90..23cb1f0b 100644 --- a/vendor/baml/crates/internal-baml-jinja/src/output_format/types.rs +++ b/vendor/baml/crates/internal-baml-jinja/src/output_format/types.rs @@ -575,6 +575,28 @@ impl OutputFormatContent { .to_string(options) } + fn rendered_class_name( + &self, + class_name: &str, + preferred_mode: Option, + ) -> String { + let preferred = + preferred_mode.and_then(|mode| self.classes.get(&(class_name.to_string(), mode))); + preferred + .or_else(|| { + self.classes.get(&( + class_name.to_string(), + baml_types::StreamingMode::NonStreaming, + )) + }) + .or_else(|| { + self.classes + .get(&(class_name.to_string(), baml_types::StreamingMode::Streaming)) + }) + .map(|class| class.name.rendered_name().to_string()) + .unwrap_or_else(|| class_name.to_string()) + } + /// Renders either the schema or the name of a type. /// /// Prompt rendering is somewhat confusing because of hoisted types, so @@ -699,8 +721,13 @@ impl OutputFormatContent { ) -> Result { match field_type { TypeIR::Class { - name: nested_class, .. - } if render_ctx.hoisted_classes.contains(nested_class) => Ok(nested_class.to_owned()), + name: nested_class, + mode, + .. + } if render_ctx.hoisted_classes.contains(nested_class) => { + // Use rendered name (alias) for hoisted class references. + Ok(self.rendered_class_name(nested_class, Some(*mode))) + } _ => self.inner_type_render(options, field_type, render_ctx), } @@ -936,9 +963,13 @@ impl OutputFormatContent { // Top level recursive classes will just use their name instead of the // entire schema which should already be hoisted. - if let TypeIR::Class { name: class, .. } = &self.target { + if let TypeIR::Class { + name: class, mode, .. + } = &self.target + { if render_ctx.hoisted_classes.contains(class) { - message = Some(class.to_owned()); + // Use rendered name (alias) for the target message. + message = Some(self.rendered_class_name(class, Some(*mode))); } } @@ -985,11 +1016,14 @@ impl OutputFormatContent { (Vec::new(), schema) }; + // Use rendered name (alias) for hoisted class headers. + let displayed_name = self.rendered_class_name(class_name, None); + let class_def = match &options.hoisted_class_prefix { RenderSetting::Always(prefix) if !prefix.is_empty() => { - format!("{prefix} {class_name} {schema_body}") + format!("{prefix} {displayed_name} {schema_body}") } - _ => format!("{class_name} {schema_body}"), + _ => format!("{displayed_name} {schema_body}"), }; // Prepend description if present diff --git a/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_class.rs b/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_class.rs index 2cf1e931..2b736338 100644 --- a/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_class.rs +++ b/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_class.rs @@ -1,5 +1,3 @@ -use std::collections::HashSet; - use anyhow::Result; use baml_types::{BamlMap, Constraint, TypeIR}; use internal_baml_jinja::types::{Class, Name}; @@ -65,24 +63,28 @@ impl TypeCoercer for Class { #[derive(Debug)] enum Triple { Pending, - NotPresent, Present(Box), } - let mut fill_result = self + // Build field states indexed by field position + let mut field_states: Vec<(&Name, &TypeIR, Triple)> = self .fields .iter() - .map(|(name, field_type, _, streaming_needed)| { - ( - name.rendered_name(), - (name, field_type, *streaming_needed, Triple::Pending), - ) - }) - .collect::>(); + .map(|(name, field_type, _, _)| (name, field_type, Triple::Pending)) + .collect(); + + // Build key-to-index map accepting both real and rendered names + let mut key_to_idx: std::collections::HashMap<&str, usize> = + std::collections::HashMap::new(); + for (idx, (name, ..)) in self.fields.iter().enumerate() { + key_to_idx.insert(name.real_name(), idx); + key_to_idx.insert(name.rendered_name(), idx); + } let flags = DeserializerConditions::new(); for (k, v) in obj.iter() { - if let Some((_, field_type, streaming_needed, val)) = fill_result.get_mut(k.as_str()) { + if let Some(&idx) = key_to_idx.get(k.as_str()) { + let (_, field_type, ref mut val) = &mut field_states[idx]; if matches!(val, Triple::Present(_)) { continue; } @@ -98,7 +100,7 @@ impl TypeCoercer for Class { } let mut result = BamlMap::new(); - for (_, (name, field_type, streaming_needed, val)) in fill_result.into_iter() { + for (name, field_type, val) in field_states.into_iter() { if let Triple::Present(ref val_ref) = val { // Check if field is required (non-optional) and is incomplete in streaming mode if !field_type.is_optional() @@ -206,11 +208,11 @@ impl TypeCoercer for Class { let mut extra_keys = vec![]; let mut found_keys = false; obj.iter().for_each(|(key, v)| { - if let Some(field) = self - .fields - .iter() - .find(|(name, ..)| matches_string_to_string(ctx, key, name.rendered_name())) - { + // Accept both real and rendered field names + if let Some(field) = self.fields.iter().find(|(name, ..)| { + matches_string_to_string(ctx, key, name.rendered_name()) + || matches_string_to_string(ctx, key, name.real_name()) + }) { let scope = ctx.enter_scope(field.0.real_name()); let parsed = field.1.coerce(&scope, &field.1, Some(v)); update_map(&mut required_values, &mut optional_values, field, parsed); diff --git a/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_enum.rs b/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_enum.rs index bef1c5a3..1300ffb2 100644 --- a/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_enum.rs +++ b/vendor/baml/crates/jsonish/src/deserializer/coercer/ir_ref/coerce_enum.rs @@ -18,12 +18,25 @@ fn enum_match_candidates(enm: &Enum) -> Vec<(&str, Vec)> { ( name.real_name(), match desc.as_ref().map(|d| d.trim()) { - Some(d) if !d.is_empty() => vec![ - name.rendered_name().into(), - d.into(), - format!("{}: {}", name.rendered_name(), d), - ], - _ => vec![name.rendered_name().into()], + Some(d) if !d.is_empty() => { + let mut candidates = vec![ + name.rendered_name().into(), + name.real_name().into(), // Also accept real name + d.into(), + format!("{}: {}", name.rendered_name(), d), + ]; + // Dedupe if real == rendered + candidates.dedup(); + candidates + } + _ => { + let mut candidates = vec![ + name.rendered_name().into(), + name.real_name().into(), // Also accept real name + ]; + candidates.dedup(); + candidates + } }, ) }) @@ -38,14 +51,14 @@ impl TypeCoercer for Enum { value: Option<&crate::jsonish::Value>, ) -> Option { // Enums can only be cast from string values - let Some(crate::jsonish::Value::String(s, _)) = value else { + let Some(crate::jsonish::Value::String(s, completion_state)) = value else { return None; }; - // Check if the string exactly matches any enum variant + // Check if the string exactly matches any enum variant (accept both real and rendered names) let mut result = None; for (variant_name, _) in &self.values { - if variant_name.rendered_name() == s { + if variant_name.rendered_name() == s || variant_name.real_name() == s { result = Some(BamlValueWithFlags::Enum( self.name.real_name().to_string(), target.clone(), @@ -55,17 +68,15 @@ impl TypeCoercer for Enum { } } - // Check completion state - if let Some(v) = value { - if let Some(ref mut res) = result { - match v.completion_state() { - baml_types::CompletionState::Complete => {} - baml_types::CompletionState::Incomplete => { - res.add_flag(crate::deserializer::deserialize_flags::Flag::Incomplete); - } - baml_types::CompletionState::Pending => { - unreachable!("jsonish::Value may never be in a Pending state.") - } + // Check completion state. + if let Some(ref mut res) = result { + match completion_state { + baml_types::CompletionState::Complete => {} + baml_types::CompletionState::Incomplete => { + res.add_flag(crate::deserializer::deserialize_flags::Flag::Incomplete); + } + baml_types::CompletionState::Pending => { + unreachable!("jsonish::Value may never be in a Pending state.") } } }