From bca364614f2bf20411e1c142921d3a296884e4a0 Mon Sep 17 00:00:00 2001 From: Matt Stavola Date: Fri, 17 Apr 2026 20:21:00 -0400 Subject: [PATCH] Extract mlf-atproto plumbing crate Split protocol-level primitives (DID/NSID/DNS identity resolution, generic XRPC HTTP client, session management, record CRUD wrappers, client-side DAG-CBOR CID computation) out of mlf-lexicon-fetcher into a new mlf-atproto crate. The fetcher is refactored to consume it; its public API is preserved. This is the plumbing foundation for mlf-publish, mlf-plugin-host, and the DNS provider plugin binaries. --- Cargo.lock | 524 +++++++++++- Cargo.toml | 1 + codegen-plugins/mlf-codegen-go/src/lib.rs | 47 +- codegen-plugins/mlf-codegen-rust/src/lib.rs | 105 ++- .../mlf-codegen-typescript/src/lib.rs | 26 +- mlf-atproto/Cargo.toml | 22 + mlf-atproto/src/cid.rs | 69 ++ mlf-atproto/src/identity.rs | 346 ++++++++ mlf-atproto/src/lib.rs | 45 ++ mlf-atproto/src/records.rs | 359 +++++++++ mlf-atproto/src/session.rs | 194 +++++ mlf-atproto/src/xrpc.rs | 192 +++++ mlf-cli/src/check.rs | 160 ++-- mlf-cli/src/config.rs | 38 +- mlf-cli/src/fetch.rs | 159 ++-- mlf-cli/src/generate/code.rs | 78 +- mlf-cli/src/generate/lexicon.rs | 116 ++- mlf-cli/src/generate/mlf.rs | 299 ++++--- mlf-cli/src/generate/mod.rs | 58 +- mlf-cli/src/init.rs | 2 +- mlf-cli/src/main.rs | 116 ++- mlf-cli/src/workspace_ext.rs | 12 +- mlf-codegen/examples/all_generators.rs | 7 +- mlf-codegen/examples/list_generators.rs | 9 +- mlf-codegen/examples/plugin_test.rs | 3 +- mlf-codegen/src/lib.rs | 287 +++++-- mlf-diagnostics/src/lib.rs | 160 ++-- mlf-lang/src/ast.rs | 6 +- mlf-lang/src/error.rs | 72 +- mlf-lang/src/lexer.rs | 56 +- mlf-lang/src/lib.rs | 2 +- mlf-lang/src/parser.rs | 166 ++-- mlf-lang/src/workspace.rs | 515 ++++++++---- mlf-lang/tests/integration_test.rs | 9 +- mlf-lexicon-fetcher/Cargo.toml | 7 +- mlf-lexicon-fetcher/examples/usage.rs | 21 +- mlf-lexicon-fetcher/src/lib.rs | 755 ++++-------------- mlf-lexicon-fetcher/tests/dns_scenarios.rs | 129 ++- mlf-lexicon-fetcher/tests/lexicon_fetching.rs | 133 ++- mlf-lsp/src/context.rs | 28 +- mlf-lsp/src/main.rs | 10 +- mlf-lsp/src/namespace_completion.rs | 3 +- mlf-lsp/src/server.rs | 544 +++++++++---- mlf-lsp/src/utils.rs | 19 +- mlf-validation/src/lib.rs | 126 ++- tests/codegen_integration.rs | 7 +- tests/diagnostics_integration.rs | 7 +- tests/lexicon_fetcher_integration.rs | 6 +- tests/lexicon_to_mlf_integration.rs | 4 +- tests/real_world/roundtrip.rs | 64 +- tests/test_utils.rs | 12 +- website/mlf-playground-wasm/src/lib.rs | 9 +- 52 files changed, 4439 insertions(+), 1705 deletions(-) create mode 100644 mlf-atproto/Cargo.toml create mode 100644 mlf-atproto/src/cid.rs create mode 100644 mlf-atproto/src/identity.rs create mode 100644 mlf-atproto/src/lib.rs create mode 100644 mlf-atproto/src/records.rs create mode 100644 mlf-atproto/src/session.rs create mode 100644 mlf-atproto/src/xrpc.rs diff --git a/Cargo.lock b/Cargo.lock index b8de3c2..2f1a4c1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -128,6 +128,28 @@ dependencies = [ "windows-sys 0.60.2", ] +[[package]] +name = "arrayref" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76a2e8124351fda1ef8aaaa3bbd7ebbcb486bbcd4225aca0aa0d84bb2db8fecb" + +[[package]] +name = "arrayvec" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" + +[[package]] +name = "assert-json-diff" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e4f2b81832e72834d7518d8487a0396a28cc408186a2e8854c0f98011faf12" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "async-trait" version = "0.1.89" @@ -186,6 +208,22 @@ dependencies = [ "backtrace", ] +[[package]] +name = "base-x" +version = "0.2.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cbbc9d0964165b47557570cce6c952866c2678457aca742aafc9fb771d30270" + +[[package]] +name = "base256emoji" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b5e9430d9a245a77c92176e649af6e275f20839a48389859d1661e9a128d077c" +dependencies = [ + "const-str", + "match-lookup", +] + [[package]] name = "base64" version = "0.22.1" @@ -219,6 +257,42 @@ version = "2.9.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2261d10cca569e4643e526d8dc2e62e433cc8aba21ab764233731f8d369bf394" +[[package]] +name = "blake2b_simd" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b79834656f71332577234b50bfc009996f7449e0c056884e6a02492ded0ca2f3" +dependencies = [ + "arrayref", + "arrayvec", + "constant_time_eq", +] + +[[package]] +name = "blake2s_simd" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee29928bad1e3f94c9d1528da29e07a1d3d04817ae8332de1e8b846c8439f4b3" +dependencies = [ + "arrayref", + "arrayvec", + "constant_time_eq", +] + +[[package]] +name = "blake3" +version = "1.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4d2d5991425dfd0785aed03aedcf0b321d61975c9b5b3689c774a2610ae0b51e" +dependencies = [ + "arrayref", + "arrayvec", + "cc", + "cfg-if", + "constant_time_eq", + "cpufeatures 0.3.0", +] + [[package]] name = "block-buffer" version = "0.10.4" @@ -228,6 +302,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-buffer" +version = "0.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cdd35008169921d80bc60d3d0ab416eecb028c4cd653352907921d95084790be" +dependencies = [ + "hybrid-array", +] + [[package]] name = "btree-range-map" version = "0.7.2" @@ -270,6 +353,15 @@ version = "1.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e629a66d692cb9ff1a1c664e41771b3dcaf961985a9774c0eb0bd1b51cf60a48" +[[package]] +name = "cbor4ii" +version = "0.2.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b544cf8c89359205f4f990d0e6f3828db42df85b5dac95d09157a250eb0749c4" +dependencies = [ + "serde", +] + [[package]] name = "cc" version = "1.2.60" @@ -336,6 +428,20 @@ dependencies = [ "half", ] +[[package]] +name = "cid" +version = "0.11.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cbb4913a732503de004e94ce7a4e7119ffc55d1727cc9979ac3b52f511e6578c" +dependencies = [ + "multibase", + "multihash", + "no_std_io2", + "serde", + "serde_bytes", + "unsigned-varint", +] + [[package]] name = "clap" version = "4.5.48" @@ -393,6 +499,18 @@ dependencies = [ "windows-sys 0.61.1", ] +[[package]] +name = "const-str" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2f421161cb492475f1661ddc9815a745a1c894592070661180fdec3d4872e9c3" + +[[package]] +name = "constant_time_eq" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d52eff69cd5e647efe296129160853a42795992097e8af39800e1060caeea9b" + [[package]] name = "core-foundation" version = "0.9.4" @@ -418,6 +536,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + [[package]] name = "crunchy" version = "0.2.4" @@ -434,6 +561,15 @@ dependencies = [ "typenum", ] +[[package]] +name = "crypto-common" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77727bb15fa921304124b128af125e7e3b968275d1b108b379190264f4423710" +dependencies = [ + "hybrid-array", +] + [[package]] name = "dashmap" version = "5.5.3" @@ -453,6 +589,26 @@ version = "2.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2a2330da5de22e8a3cb63252ce2abb30116bf5265e89c0e01bc17015ce30a476" +[[package]] +name = "data-encoding-macro" +version = "0.1.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47ce6c96ea0102f01122a185683611bd5ac8d99e62bc59dd12e6bda344ee673d" +dependencies = [ + "data-encoding", + "data-encoding-macro-internal", +] + +[[package]] +name = "data-encoding-macro-internal" +version = "0.1.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8d162beedaa69905488a8da94f5ac3edb4dd4788b732fadb7bd120b2625c1976" +dependencies = [ + "data-encoding", + "syn 2.0.106", +] + [[package]] name = "datatest-stable" version = "0.3.3" @@ -465,6 +621,24 @@ dependencies = [ "walkdir", ] +[[package]] +name = "deadpool" +version = "0.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0be2b1d1d6ec8d846f05e137292d0b89133caf95ef33695424c09568bdd39b1b" +dependencies = [ + "deadpool-runtime", + "lazy_static", + "num_cpus", + "tokio", +] + +[[package]] +name = "deadpool-runtime" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "092966b41edc516079bdf31ec78a2e0588d1d0c08f78b91d8307215928642b2b" + [[package]] name = "deranged" version = "0.5.4" @@ -480,8 +654,18 @@ version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer", - "crypto-common", + "block-buffer 0.10.4", + "crypto-common 0.1.6", +] + +[[package]] +name = "digest" +version = "0.11.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4850db49bf08e663084f7fb5c87d202ef91a3907271aff24a94eb97ff039153c" +dependencies = [ + "block-buffer 0.12.0", + "crypto-common 0.2.1", ] [[package]] @@ -626,6 +810,7 @@ checksum = "65bc07b1a8bc7c85c5f2e110c476c7389b4554ba72af57d8445ea63a576b0876" dependencies = [ "futures-channel", "futures-core", + "futures-executor", "futures-io", "futures-sink", "futures-task", @@ -648,6 +833,17 @@ version = "0.3.31" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "05f29059c0c2090612e8d742178b0580d2dc940c837851ad723096f87af6663e" +[[package]] +name = "futures-executor" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e28d1d997f585e54aebc3f97d39e72338912123a67330d723fdbb564d646c9f" +dependencies = [ + "futures-core", + "futures-task", + "futures-util", +] + [[package]] name = "futures-io" version = "0.3.31" @@ -777,9 +973,9 @@ checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" [[package]] name = "hashbrown" -version = "0.16.0" +version = "0.17.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5419bdc4f6a9207fbeba6d11b604d481addf78ecd10c11ad51e76c2f6482748d" +checksum = "4f467dd6dccf739c208452f8014c75c18bb8301b050ad1cfb27153803edb0f51" [[package]] name = "heck" @@ -787,6 +983,12 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + [[package]] name = "hex_fmt" version = "0.3.0" @@ -878,6 +1080,21 @@ version = "1.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" +[[package]] +name = "httpdate" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" + +[[package]] +name = "hybrid-array" +version = "0.4.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3944cf8cf766b40e2a1a333ee5e9b563f854d5fa49d6a8ca2764e97c6eddb214" +dependencies = [ + "typenum", +] + [[package]] name = "hyper" version = "1.7.0" @@ -892,6 +1109,7 @@ dependencies = [ "http", "http-body", "httparse", + "httpdate", "itoa", "pin-project-lite", "pin-utils", @@ -1110,12 +1328,12 @@ dependencies = [ [[package]] name = "indexmap" -version = "2.11.4" +version = "2.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4b0f83760fb341a774ed326568e19f5a863af4a952def8c39f9ab92fd95b88e5" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", - "hashbrown 0.16.0", + "hashbrown 0.17.0", ] [[package]] @@ -1169,6 +1387,17 @@ dependencies = [ "winreg", ] +[[package]] +name = "ipld-core" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "090f624976d72f0b0bb71b86d58dc16c15e069193067cb3a3a09d655246cbbda" +dependencies = [ + "cid", + "serde", + "serde_bytes", +] + [[package]] name = "ipnet" version = "2.11.0" @@ -1213,6 +1442,16 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "keccak" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e24a010dd405bd7ed803e5253182815b41bf2e6a80cc3bfc066658e03a198aa" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", +] + [[package]] name = "langtag" version = "0.4.0" @@ -1312,6 +1551,17 @@ dependencies = [ "url", ] +[[package]] +name = "match-lookup" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "757aee279b8bdbb9f9e676796fd459e4207a1f986e87886700abf589f5abf771" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.106", +] + [[package]] name = "matchers" version = "0.2.0" @@ -1399,6 +1649,23 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "mlf-atproto" +version = "0.1.0" +dependencies = [ + "async-trait", + "cid", + "hickory-resolver", + "multihash-codetable", + "reqwest", + "serde", + "serde_ipld_dagcbor", + "serde_json", + "thiserror 2.0.17", + "tokio", + "wiremock", +] + [[package]] name = "mlf-cli" version = "0.1.0" @@ -1418,7 +1685,7 @@ dependencies = [ "reqwest", "serde", "serde_json", - "sha2", + "sha2 0.10.9", "thiserror 2.0.17", "tokio", "toml", @@ -1504,9 +1771,8 @@ name = "mlf-lexicon-fetcher" version = "0.1.0" dependencies = [ "async-trait", - "hickory-resolver", + "mlf-atproto", "reqwest", - "serde", "serde_json", "thiserror 2.0.17", "tokio", @@ -1568,6 +1834,71 @@ dependencies = [ "web-sys", ] +[[package]] +name = "multibase" +version = "0.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8694bb4835f452b0e3bb06dbebb1d6fc5385b6ca1caf2e55fd165c042390ec77" +dependencies = [ + "base-x", + "base256emoji", + "data-encoding", + "data-encoding-macro", +] + +[[package]] +name = "multihash" +version = "0.19.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "89ace881e3f514092ce9efbcb8f413d0ad9763860b828981c2de51ddc666936c" +dependencies = [ + "no_std_io2", + "serde", + "unsigned-varint", +] + +[[package]] +name = "multihash-codetable" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "facfe64780489b29aae20d32d0245219f4a8167f91193f7061589f5dae9ba307" +dependencies = [ + "blake2b_simd", + "blake2s_simd", + "blake3", + "digest 0.11.2", + "multihash-derive", + "no_std_io2", + "ripemd", + "sha1", + "sha2 0.11.0", + "sha3", +] + +[[package]] +name = "multihash-derive" +version = "0.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0576e09c49157d1910e522e595d2b32749b029dd0bc10ff6967d588490c30348" +dependencies = [ + "multihash", + "multihash-derive-impl", + "no_std_io2", +] + +[[package]] +name = "multihash-derive-impl" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3dc7141bd06405929948754f0628d247f5ca1865be745099205e5086da957cb" +dependencies = [ + "proc-macro-crate", + "proc-macro2", + "quote", + "syn 2.0.106", + "synstructure", +] + [[package]] name = "native-tls" version = "0.2.14" @@ -1585,6 +1916,15 @@ dependencies = [ "tempfile", ] +[[package]] +name = "no_std_io2" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a3564ce7035b1e4778d8cb6cacebb5d766b5e8fe5a75b9e441e33fb61a872c6" +dependencies = [ + "memchr", +] + [[package]] name = "nom" version = "7.1.3" @@ -1628,6 +1968,16 @@ dependencies = [ "autocfg", ] +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + [[package]] name = "object" version = "0.37.3" @@ -1796,6 +2146,15 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "proc-macro-crate" +version = "3.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f" +dependencies = [ + "toml_edit 0.25.11+spec-1.1.0", +] + [[package]] name = "proc-macro-error" version = "1.0.4" @@ -1989,6 +2348,15 @@ dependencies = [ "windows-sys 0.52.0", ] +[[package]] +name = "ripemd" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4dd4211456b4172d7e44261920c25acf07367c4f04bb5f5d54fc21b090d9b159" +dependencies = [ + "digest 0.11.2", +] + [[package]] name = "rustc-demangle" version = "0.1.26" @@ -2121,6 +2489,16 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "serde_bytes" +version = "0.11.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a5d440709e79d88e51ac01c4b72fc6cb7314017bb7da9eeff678aa94c10e3ea8" +dependencies = [ + "serde", + "serde_core", +] + [[package]] name = "serde_core" version = "1.0.228" @@ -2141,6 +2519,18 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "serde_ipld_dagcbor" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "46182f4f08349a02b45c998ba3215d3f9de826246ba02bb9dddfe9a2a2100778" +dependencies = [ + "cbor4ii", + "ipld-core", + "scopeguard", + "serde", +] + [[package]] name = "serde_json" version = "1.0.145" @@ -2187,6 +2577,17 @@ dependencies = [ "serde", ] +[[package]] +name = "sha1" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "aacc4cc499359472b4abe1bf11d0b12e688af9a805fa5e3016f9a386dc2d0214" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.2", +] + [[package]] name = "sha2" version = "0.10.9" @@ -2194,8 +2595,29 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", +] + +[[package]] +name = "sha2" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.2", +] + +[[package]] +name = "sha3" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be176f1a57ce4e3d31c1a166222d9768de5954f811601fb7ca06fc8203905ce1" +dependencies = [ + "digest 0.11.2", + "keccak", ] [[package]] @@ -2281,7 +2703,7 @@ dependencies = [ "proc-macro2", "quote", "serde", - "sha2", + "sha2 0.10.9", "syn 2.0.106", "thiserror 1.0.69", ] @@ -2596,8 +3018,8 @@ checksum = "dc1beb996b9d83529a9e75c17a1686767d148d70663143c7854d8b4a09ced362" dependencies = [ "serde", "serde_spanned", - "toml_datetime", - "toml_edit", + "toml_datetime 0.6.11", + "toml_edit 0.22.27", ] [[package]] @@ -2609,6 +3031,15 @@ dependencies = [ "serde", ] +[[package]] +name = "toml_datetime" +version = "1.1.1+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3165f65f62e28e0115a00b2ebdd37eb6f3b641855f9d636d3cd4103767159ad7" +dependencies = [ + "serde_core", +] + [[package]] name = "toml_edit" version = "0.22.27" @@ -2618,9 +3049,30 @@ dependencies = [ "indexmap", "serde", "serde_spanned", - "toml_datetime", + "toml_datetime 0.6.11", "toml_write", - "winnow", + "winnow 0.7.13", +] + +[[package]] +name = "toml_edit" +version = "0.25.11+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b59c4d22ed448339746c59b905d24568fcbb3ab65a500494f7b8c3e97739f2b" +dependencies = [ + "indexmap", + "toml_datetime 1.1.1+spec-1.1.0", + "toml_parser", + "winnow 1.0.1", +] + +[[package]] +name = "toml_parser" +version = "1.1.2+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526" +dependencies = [ + "winnow 1.0.1", ] [[package]] @@ -2854,6 +3306,12 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4a1a07cc7db3810833284e8d372ccdc6da29741639ecc70c9ec107df0fa6154c" +[[package]] +name = "unsigned-varint" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eb066959b24b5196ae73cb057f45598450d2c5f71460e98c49b738086eff9c06" + [[package]] name = "untrusted" version = "0.9.0" @@ -3400,6 +3858,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "winnow" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09dac053f1cd375980747450bfc7250c264eaae0583872e845c0c7cd578872b5" +dependencies = [ + "memchr", +] + [[package]] name = "winreg" version = "0.50.0" @@ -3410,6 +3877,29 @@ dependencies = [ "windows-sys 0.48.0", ] +[[package]] +name = "wiremock" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08db1edfb05d9b3c1542e521aea074442088292f00b5f28e435c714a98f85031" +dependencies = [ + "assert-json-diff", + "base64", + "deadpool", + "futures", + "http", + "http-body-util", + "hyper", + "hyper-util", + "log", + "once_cell", + "regex", + "serde", + "serde_json", + "tokio", + "url", +] + [[package]] name = "wit-bindgen" version = "0.46.0" diff --git a/Cargo.toml b/Cargo.toml index f20a7f5..4e43828 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,6 +4,7 @@ members = [ "codegen-plugins/mlf-codegen-go", "codegen-plugins/mlf-codegen-rust", "codegen-plugins/mlf-codegen-typescript", + "mlf-atproto", "mlf-cli", "mlf-codegen", "mlf-diagnostics", diff --git a/codegen-plugins/mlf-codegen-go/src/lib.rs b/codegen-plugins/mlf-codegen-go/src/lib.rs index 4a7b9c0..d20580e 100644 --- a/codegen-plugins/mlf-codegen-go/src/lib.rs +++ b/codegen-plugins/mlf-codegen-go/src/lib.rs @@ -1,4 +1,4 @@ -use mlf_codegen::{register_generator, CodeGenerator, GeneratorContext}; +use mlf_codegen::{CodeGenerator, GeneratorContext, register_generator}; use mlf_lang::ast::*; use std::fmt::Write; @@ -7,7 +7,12 @@ pub struct GoGenerator; impl GoGenerator { pub const NAME: &'static str = "go"; - fn generate_type(&self, ty: &Type, optional: bool, ctx: &GeneratorContext) -> Result { + fn generate_type( + &self, + ty: &Type, + optional: bool, + ctx: &GeneratorContext, + ) -> Result { let base_type = match ty { Type::Primitive { kind, .. } => match kind { PrimitiveType::Null => "interface{}", @@ -16,15 +21,15 @@ impl GoGenerator { PrimitiveType::String => "string", PrimitiveType::Bytes => "[]byte", PrimitiveType::Blob => "[]byte", // Annotation idea: @goType("custom.BlobType") - }.to_string(), + } + .to_string(), Type::Reference { path, .. } => { let path_str = path.to_string(); match path_str.as_str() { // Map standard library types "Datetime" => "string".to_string(), // ISO 8601 string - "Did" | "AtUri" | "Cid" | "AtIdentifier" | "Handle" | "Nsid" | "Tid" | "RecordKey" | "Uri" | "Language" => { - "string".to_string() - } + "Did" | "AtUri" | "Cid" | "AtIdentifier" | "Handle" | "Nsid" | "Tid" + | "RecordKey" | "Uri" | "Language" => "string".to_string(), _ => { // Local reference path.segments.last().unwrap().name.clone() @@ -50,7 +55,12 @@ impl GoGenerator { if !field.docs.is_empty() { write!(obj, "\t\t// {}\n", field.docs[0].text).unwrap(); } - write!(obj, "\t\t{} {} `json:\"{}", field_name, field_type, json_name).unwrap(); + write!( + obj, + "\t\t{} {} `json:\"{}", + field_name, field_type, json_name + ) + .unwrap(); if field.optional { write!(obj, ",omitempty").unwrap(); } @@ -129,7 +139,12 @@ impl CodeGenerator for GoGenerator { match item { Item::Record(record) => { output.push_str(&self.generate_doc_comment(&record.docs)); - writeln!(output, "type {} struct {{", self.capitalize(&record.name.name)).unwrap(); + writeln!( + output, + "type {} struct {{", + self.capitalize(&record.name.name) + ) + .unwrap(); for field in &record.fields { let field_name = self.capitalize(&field.name.name); @@ -139,7 +154,12 @@ impl CodeGenerator for GoGenerator { if !field.docs.is_empty() { writeln!(output, "\t// {}", field.docs[0].text).unwrap(); } - write!(output, "\t{} {} `json:\"{}\"", field_name, field_type, json_name).unwrap(); + write!( + output, + "\t{} {} `json:\"{}\"", + field_name, field_type, json_name + ) + .unwrap(); if field.optional { write!(output, ",omitempty").unwrap(); } @@ -160,7 +180,8 @@ impl CodeGenerator for GoGenerator { "type {} {}\n", type_name, self.generate_type(&def.ty, false, ctx)? - ).unwrap(); + ) + .unwrap(); } _ => { // Other types become type aliases @@ -169,7 +190,8 @@ impl CodeGenerator for GoGenerator { "type {} {}\n", type_name, self.generate_type(&def.ty, false, ctx)? - ).unwrap(); + ) + .unwrap(); } } } @@ -180,7 +202,8 @@ impl CodeGenerator for GoGenerator { "type {} {}\n", self.capitalize(&inline.name.name), self.generate_type(&inline.ty, false, ctx)? - ).unwrap(); + ) + .unwrap(); } Item::Token(token) => { output.push_str(&self.generate_doc_comment(&token.docs)); diff --git a/codegen-plugins/mlf-codegen-rust/src/lib.rs b/codegen-plugins/mlf-codegen-rust/src/lib.rs index 4b69591..e8eb0f1 100644 --- a/codegen-plugins/mlf-codegen-rust/src/lib.rs +++ b/codegen-plugins/mlf-codegen-rust/src/lib.rs @@ -1,4 +1,4 @@ -use mlf_codegen::{register_generator, CodeGenerator, GeneratorContext}; +use mlf_codegen::{CodeGenerator, GeneratorContext, register_generator}; use mlf_lang::ast::*; use std::fmt::Write; @@ -26,7 +26,12 @@ impl RustGenerator { } } - fn generate_type(&self, ty: &Type, optional: bool, ctx: &GeneratorContext) -> Result { + fn generate_type( + &self, + ty: &Type, + optional: bool, + ctx: &GeneratorContext, + ) -> Result { let base_type = match ty { Type::Primitive { kind, .. } => match kind { PrimitiveType::Null => "()".to_string(), // Unit type for null @@ -41,9 +46,8 @@ impl RustGenerator { match path_str.as_str() { // Map standard library types "Datetime" => "String".to_string(), // ISO 8601 string, could use chrono::DateTime - "Did" | "AtUri" | "Cid" | "AtIdentifier" | "Handle" | "Nsid" | "Tid" | "RecordKey" | "Uri" | "Language" => { - "String".to_string() - } + "Did" | "AtUri" | "Cid" | "AtIdentifier" | "Handle" | "Nsid" | "Tid" + | "RecordKey" | "Uri" | "Language" => "String".to_string(), _ => { // Local reference - convert to PascalCase self.to_pascal_case(&path.segments.last().unwrap().name) @@ -58,10 +62,26 @@ impl RustGenerator { // Rust doesn't have direct union types, use an enum // For now, generate a simple representation // Annotation idea: @rustEnum to customize enum generation - if types.len() == 2 && matches!(types[0], Type::Primitive { kind: PrimitiveType::Null, .. }) { + if types.len() == 2 + && matches!( + types[0], + Type::Primitive { + kind: PrimitiveType::Null, + .. + } + ) + { // Special case: null | T becomes Option return self.generate_type(&types[1], true, ctx); - } else if types.len() == 2 && matches!(types[1], Type::Primitive { kind: PrimitiveType::Null, .. }) { + } else if types.len() == 2 + && matches!( + types[1], + Type::Primitive { + kind: PrimitiveType::Null, + .. + } + ) + { return self.generate_type(&types[0], true, ctx); } // Otherwise use serde_json::Value for flexibility @@ -135,7 +155,12 @@ impl CodeGenerator for RustGenerator { Item::Record(record) => { output.push_str(&self.generate_doc_comment(&record.docs)); writeln!(output, "#[derive(Debug, Clone, Serialize, Deserialize)]").unwrap(); - writeln!(output, "pub struct {} {{", self.to_pascal_case(&record.name.name)).unwrap(); + writeln!( + output, + "pub struct {} {{", + self.to_pascal_case(&record.name.name) + ) + .unwrap(); for field in &record.fields { if !field.docs.is_empty() { @@ -144,16 +169,27 @@ impl CodeGenerator for RustGenerator { // Use serde rename for camelCase fields if field.name.name != self.to_snake_case(&field.name.name) { - writeln!(output, " #[serde(rename = \"{}\")]", field.name.name).unwrap(); + writeln!(output, " #[serde(rename = \"{}\")]", field.name.name) + .unwrap(); } // Skip serializing None values for optional fields if field.optional { - writeln!(output, " #[serde(skip_serializing_if = \"Option::is_none\")]").unwrap(); + writeln!( + output, + " #[serde(skip_serializing_if = \"Option::is_none\")]" + ) + .unwrap(); } let field_type = self.generate_type(&field.ty, field.optional, ctx)?; - writeln!(output, " pub {}: {},", self.to_snake_case(&field.name.name), field_type).unwrap(); + writeln!( + output, + " pub {}: {},", + self.to_snake_case(&field.name.name), + field_type + ) + .unwrap(); } writeln!(output, "}}\n").unwrap(); @@ -164,8 +200,14 @@ impl CodeGenerator for RustGenerator { match &def.ty { Type::Object { fields, .. } => { // Generate a struct for object types - writeln!(output, "#[derive(Debug, Clone, Serialize, Deserialize)]").unwrap(); - writeln!(output, "pub struct {} {{", self.to_pascal_case(&def.name.name)).unwrap(); + writeln!(output, "#[derive(Debug, Clone, Serialize, Deserialize)]") + .unwrap(); + writeln!( + output, + "pub struct {} {{", + self.to_pascal_case(&def.name.name) + ) + .unwrap(); for field in fields { if !field.docs.is_empty() { @@ -173,15 +215,31 @@ impl CodeGenerator for RustGenerator { } if field.name.name != self.to_snake_case(&field.name.name) { - writeln!(output, " #[serde(rename = \"{}\")]", field.name.name).unwrap(); + writeln!( + output, + " #[serde(rename = \"{}\")]", + field.name.name + ) + .unwrap(); } if field.optional { - writeln!(output, " #[serde(skip_serializing_if = \"Option::is_none\")]").unwrap(); + writeln!( + output, + " #[serde(skip_serializing_if = \"Option::is_none\")]" + ) + .unwrap(); } - let field_type = self.generate_type(&field.ty, field.optional, ctx)?; - writeln!(output, " pub {}: {},", self.to_snake_case(&field.name.name), field_type).unwrap(); + let field_type = + self.generate_type(&field.ty, field.optional, ctx)?; + writeln!( + output, + " pub {}: {},", + self.to_snake_case(&field.name.name), + field_type + ) + .unwrap(); } writeln!(output, "}}\n").unwrap(); @@ -193,7 +251,8 @@ impl CodeGenerator for RustGenerator { "pub type {} = {};\n", self.to_pascal_case(&def.name.name), self.generate_type(&def.ty, false, ctx)? - ).unwrap(); + ) + .unwrap(); } } } @@ -204,14 +263,18 @@ impl CodeGenerator for RustGenerator { "pub type {} = {};\n", self.to_pascal_case(&inline.name.name), self.generate_type(&inline.ty, false, ctx)? - ).unwrap(); + ) + .unwrap(); } Item::Token(token) => { output.push_str(&self.generate_doc_comment(&token.docs)); - writeln!(output, "pub const {}: &str = \"{}\";\n", + writeln!( + output, + "pub const {}: &str = \"{}\";\n", token.name.name.to_uppercase(), token.name.name - ).unwrap(); + ) + .unwrap(); } Item::Query(_) | Item::Procedure(_) | Item::Subscription(_) => { // TODO: Generate client methods diff --git a/codegen-plugins/mlf-codegen-typescript/src/lib.rs b/codegen-plugins/mlf-codegen-typescript/src/lib.rs index f4441d5..674c1c5 100644 --- a/codegen-plugins/mlf-codegen-typescript/src/lib.rs +++ b/codegen-plugins/mlf-codegen-typescript/src/lib.rs @@ -1,4 +1,4 @@ -use mlf_codegen::{register_generator, CodeGenerator, GeneratorContext}; +use mlf_codegen::{CodeGenerator, GeneratorContext, register_generator}; use mlf_lang::ast::*; use std::fmt::Write; @@ -24,9 +24,8 @@ impl TypeScriptGenerator { // Map standard library types to TypeScript types Ok(match path_str.as_str() { "Datetime" => "string".to_string(), // ISO 8601 - "Did" | "AtUri" | "Cid" | "AtIdentifier" | "Handle" | "Nsid" | "Tid" | "RecordKey" | "Uri" | "Language" => { - "string".to_string() - } + "Did" | "AtUri" | "Cid" | "AtIdentifier" | "Handle" | "Nsid" | "Tid" + | "RecordKey" | "Uri" | "Language" => "string".to_string(), _ => { // Local or cross-file reference path.segments.last().unwrap().name.clone() @@ -38,10 +37,8 @@ impl TypeScriptGenerator { Ok(format!("{}[]", inner_type)) } Type::Union { types, .. } => { - let type_strings: Result, _> = types - .iter() - .map(|t| self.generate_type(t, ctx)) - .collect(); + let type_strings: Result, _> = + types.iter().map(|t| self.generate_type(t, ctx)).collect(); Ok(type_strings?.join(" | ")) } Type::Object { fields, .. } => { @@ -143,7 +140,8 @@ impl CodeGenerator for TypeScriptGenerator { "export type {} = {};\n", def.name.name, self.generate_type(&def.ty, ctx)? - ).unwrap(); + ) + .unwrap(); } Item::InlineType(inline) => { output.push_str(&self.generate_doc_comment(&inline.docs)); @@ -152,11 +150,17 @@ impl CodeGenerator for TypeScriptGenerator { "export type {} = {};\n", inline.name.name, self.generate_type(&inline.ty, ctx)? - ).unwrap(); + ) + .unwrap(); } Item::Token(token) => { output.push_str(&self.generate_doc_comment(&token.docs)); - writeln!(output, "export const {} = Symbol('{}');\n", token.name.name, token.name.name).unwrap(); + writeln!( + output, + "export const {} = Symbol('{}');\n", + token.name.name, token.name.name + ) + .unwrap(); } Item::Query(_) | Item::Procedure(_) | Item::Subscription(_) => { // TODO: Generate client methods for these diff --git a/mlf-atproto/Cargo.toml b/mlf-atproto/Cargo.toml new file mode 100644 index 0000000..c24bd97 --- /dev/null +++ b/mlf-atproto/Cargo.toml @@ -0,0 +1,22 @@ +[package] +name = "mlf-atproto" +version = "0.1.0" +edition = "2024" +license = "MIT" +description = "ATProto protocol plumbing for MLF: identity, XRPC, session, records, CID" + +[dependencies] +async-trait = "0.1" +hickory-resolver = "0.24" +reqwest = { version = "0.12", features = ["json"] } +serde = { version = "1.0", features = ["derive"] } +serde_json = "1.0" +serde_ipld_dagcbor = "0.6" +cid = "0.11" +multihash-codetable = { version = "0.2", features = ["sha2"] } +thiserror = "2.0" +tokio = { version = "1", features = ["rt"] } + +[dev-dependencies] +tokio = { version = "1", features = ["full"] } +wiremock = "0.6" diff --git a/mlf-atproto/src/cid.rs b/mlf-atproto/src/cid.rs new file mode 100644 index 0000000..5a5eb0a --- /dev/null +++ b/mlf-atproto/src/cid.rs @@ -0,0 +1,69 @@ +//! Client-side CID (Content Identifier) computation for AT Protocol records. +//! +//! Records are addressed by CIDv1 over a DAG-CBOR encoding with a +//! SHA-256 multihash — this module produces exactly that. +//! +//! Use case: computing a local record's CID so we can diff against +//! the `cid` a PDS returns from `listRecords` / `getRecord` without a +//! network round-trip. + +use ::cid::Cid as Multiformat; +use multihash_codetable::{Code, MultihashDigest}; + +#[derive(thiserror::Error, Debug)] +pub enum CidError { + #[error("DAG-CBOR encoding failed: {0}")] + CborEncode(String), +} + +/// DAG-CBOR codec code, per multicodec table. +const DAG_CBOR: u64 = 0x71; + +/// Compute the CIDv1 of a JSON value by encoding it as DAG-CBOR and +/// hashing with SHA-256. +/// +/// Produces the base32-encoded `bafy…` string that AT Protocol uses. +pub fn cid_for_json(value: &serde_json::Value) -> Result { + let bytes = + serde_ipld_dagcbor::to_vec(value).map_err(|e| CidError::CborEncode(e.to_string()))?; + Ok(cid_for_dag_cbor_bytes(&bytes)) +} + +/// Compute the CIDv1 of pre-encoded DAG-CBOR bytes. +pub fn cid_for_dag_cbor_bytes(bytes: &[u8]) -> String { + let hash = Code::Sha2_256.digest(bytes); + let cid = Multiformat::new_v1(DAG_CBOR, hash); + cid.to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn round_trip_is_deterministic() { + let v = json!({"a": 1, "b": "hello"}); + let a = cid_for_json(&v).unwrap(); + let b = cid_for_json(&v).unwrap(); + assert_eq!(a, b); + assert!(a.starts_with("bafy"), "expected bafy-prefixed CID, got {a}"); + } + + #[test] + fn different_content_has_different_cid() { + let a = cid_for_json(&json!({"k": 1})).unwrap(); + let b = cid_for_json(&json!({"k": 2})).unwrap(); + assert_ne!(a, b); + } + + #[test] + fn field_order_does_not_matter_for_cbor_canonical_form() { + // serde_ipld_dagcbor sorts object keys in canonical order, + // so different insertion orders in the JSON source produce the + // same DAG-CBOR bytes and thus the same CID. + let a = cid_for_json(&json!({"a": 1, "b": 2})).unwrap(); + let b = cid_for_json(&json!({"b": 2, "a": 1})).unwrap(); + assert_eq!(a, b); + } +} diff --git a/mlf-atproto/src/identity.rs b/mlf-atproto/src/identity.rs new file mode 100644 index 0000000..560f7f3 --- /dev/null +++ b/mlf-atproto/src/identity.rs @@ -0,0 +1,346 @@ +//! Identity resolution for the AT Protocol. +//! +//! DID parsing and document resolution (`did:plc:` via plc.directory, +//! `did:web:` via well-known fetch). NSID parsing. `_lexicon.` +//! TXT resolution for lexicon publishing authorities. +//! +//! DNS is an implementation detail — the transport for the `_lexicon` +//! convention — rather than its own module. + +use async_trait::async_trait; +use hickory_resolver::TokioAsyncResolver; +use hickory_resolver::config::{ResolverConfig, ResolverOpts}; +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; + +#[derive(thiserror::Error, Debug)] +pub enum IdentityError { + #[error("Invalid NSID format: {0}")] + InvalidNsid(String), + + #[error("Invalid DID format: {0}")] + InvalidDid(String), + + #[error("DNS lookup failed for {domain}: {error}")] + DnsLookupFailed { domain: String, error: String }, + + #[error("No `did=` entry in TXT record for {0}")] + NoDidInTxt(String), + + #[error("DID resolution failed for {did}: {error}")] + DidResolutionFailed { did: String, error: String }, + + #[error("No PDS service endpoint in DID document for {0}")] + NoPdsEndpoint(String), + + #[error("Unsupported DID method: {0}")] + UnsupportedDidMethod(String), +} + +// --------------------------------------------------------------------------- +// NSID parsing +// --------------------------------------------------------------------------- + +/// Parse an NSID into (authority, name_segments). +/// +/// - `app.bsky.actor.profile` → `("app.bsky", "actor.profile")` +/// - `place.stream.key` → `("place.stream", "key")` +/// - `place.stream` → `("place.stream", "")` +/// +/// NSIDs must have at least 2 segments. +pub fn parse_nsid(nsid: &str) -> Result<(String, String), IdentityError> { + let parts: Vec<&str> = nsid.split('.').collect(); + if parts.len() < 2 { + return Err(IdentityError::InvalidNsid(format!( + "NSID must have at least 2 segments: {nsid}" + ))); + } + let authority = format!("{}.{}", parts[0], parts[1]); + let name_segments = if parts.len() > 2 { + parts[2..].join(".") + } else { + String::new() + }; + Ok((authority, name_segments)) +} + +// --------------------------------------------------------------------------- +// DNS name for _lexicon resolution +// --------------------------------------------------------------------------- + +/// Construct the DNS name for an NSID's `_lexicon` TXT lookup. +/// +/// The NSID authority is reversed into DNS order, and any name segments +/// are prepended ahead of it, all under the `_lexicon` prefix. +/// +/// - `("app.bsky", "actor.profile")` → `"_lexicon.actor.profile.bsky.app"` +/// - `("place.stream", "")` → `"_lexicon.stream.place"` +pub fn construct_dns_name(authority: &str, name_segments: &str) -> String { + let reversed_auth: Vec<&str> = authority.split('.').rev().collect(); + if name_segments.is_empty() { + format!("_lexicon.{}", reversed_auth.join(".")) + } else { + format!("_lexicon.{}.{}", name_segments, reversed_auth.join(".")) + } +} + +// --------------------------------------------------------------------------- +// DnsResolver trait +// --------------------------------------------------------------------------- + +/// Resolver for `_lexicon.` TXT records. +/// +/// Mockable so tests can inject canned responses without network access. +#[async_trait] +pub trait DnsResolver: Send + Sync { + /// Resolve an NSID (split into authority + name segments) to the DID + /// it maps to via the `_lexicon` TXT convention. + async fn resolve_lexicon_did( + &self, + authority: &str, + name_segments: &str, + ) -> Result; +} + +/// Production DNS resolver backed by hickory-resolver. +pub struct RealDnsResolver { + resolver: TokioAsyncResolver, +} + +impl RealDnsResolver { + pub fn new() -> Result { + Ok(Self { + resolver: TokioAsyncResolver::tokio(ResolverConfig::default(), ResolverOpts::default()), + }) + } + + pub fn with_config(config: ResolverConfig, opts: ResolverOpts) -> Self { + Self { + resolver: TokioAsyncResolver::tokio(config, opts), + } + } +} + +#[async_trait] +impl DnsResolver for RealDnsResolver { + async fn resolve_lexicon_did( + &self, + authority: &str, + name_segments: &str, + ) -> Result { + let dns_name = construct_dns_name(authority, name_segments); + let response = self.resolver.txt_lookup(&dns_name).await.map_err(|e| { + IdentityError::DnsLookupFailed { + domain: dns_name.clone(), + error: e.to_string(), + } + })?; + for txt in response.iter() { + for data in txt.txt_data() { + let text = String::from_utf8_lossy(data); + if let Some(did) = text.strip_prefix("did=") { + return Ok(did.trim().to_string()); + } + } + } + Err(IdentityError::NoDidInTxt(dns_name)) + } +} + +/// In-memory DNS resolver for tests. +#[derive(Clone, Default)] +pub struct MockDnsResolver { + records: Arc>>, +} + +impl MockDnsResolver { + pub fn new() -> Self { + Self::default() + } + + /// Add a record keyed by (authority, name_segments). + pub fn add_record(&mut self, authority: &str, name_segments: &str, did: String) { + let dns_name = construct_dns_name(authority, name_segments); + self.records.lock().unwrap().insert(dns_name, did); + } + + /// Add a record from a full NSID. + pub fn add_record_from_nsid(&mut self, nsid: &str, did: String) -> Result<(), IdentityError> { + let (authority, name) = parse_nsid(nsid)?; + self.add_record(&authority, &name, did); + Ok(()) + } +} + +#[async_trait] +impl DnsResolver for MockDnsResolver { + async fn resolve_lexicon_did( + &self, + authority: &str, + name_segments: &str, + ) -> Result { + let dns_name = construct_dns_name(authority, name_segments); + self.records + .lock() + .unwrap() + .get(&dns_name) + .cloned() + .ok_or_else(|| IdentityError::DnsLookupFailed { + domain: dns_name, + error: "No mock record found".into(), + }) + } +} + +// --------------------------------------------------------------------------- +// DID → PDS endpoint resolution +// --------------------------------------------------------------------------- + +/// Resolve a DID to its PDS (Personal Data Server) HTTPS endpoint. +/// +/// Supports `did:plc:*` (via https://plc.directory) and `did:web:*` +/// (via `https:///.well-known/did.json`). +pub async fn resolve_did_to_pds( + client: &reqwest::Client, + did: &str, +) -> Result { + let did_doc = fetch_did_document(client, did).await?; + extract_pds_endpoint(&did_doc, did) +} + +/// Fetch the DID document for a DID (PLC directory or `did:web:`). +pub async fn fetch_did_document( + client: &reqwest::Client, + did: &str, +) -> Result { + let url = if let Some(domain) = did.strip_prefix("did:web:") { + format!("https://{domain}/.well-known/did.json") + } else if did.starts_with("did:plc:") { + format!("https://plc.directory/{did}") + } else { + return Err(IdentityError::UnsupportedDidMethod(did.to_string())); + }; + + let resp = client + .get(&url) + .send() + .await + .map_err(|e| IdentityError::DidResolutionFailed { + did: did.to_string(), + error: e.to_string(), + })?; + if !resp.status().is_success() { + return Err(IdentityError::DidResolutionFailed { + did: did.to_string(), + error: format!("HTTP {}", resp.status()), + }); + } + resp.json() + .await + .map_err(|e| IdentityError::DidResolutionFailed { + did: did.to_string(), + error: e.to_string(), + }) +} + +/// Extract the AtprotoPersonalDataServer endpoint from a DID document. +fn extract_pds_endpoint(did_doc: &serde_json::Value, did: &str) -> Result { + let services = did_doc + .get("service") + .and_then(|v| v.as_array()) + .ok_or_else(|| IdentityError::NoPdsEndpoint(did.to_string()))?; + for service in services { + if service.get("type").and_then(|v| v.as_str()) == Some("AtprotoPersonalDataServer") + && let Some(ep) = service.get("serviceEndpoint").and_then(|v| v.as_str()) + { + return Ok(ep.trim_end_matches('/').to_string()); + } + } + Err(IdentityError::NoPdsEndpoint(did.to_string())) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_nsid() { + assert_eq!( + parse_nsid("place.stream.key").unwrap(), + ("place.stream".into(), "key".into()) + ); + assert_eq!( + parse_nsid("app.bsky.actor.profile").unwrap(), + ("app.bsky".into(), "actor.profile".into()) + ); + assert_eq!( + parse_nsid("place.stream").unwrap(), + ("place.stream".into(), "".into()) + ); + assert!(parse_nsid("invalid").is_err()); + } + + #[test] + fn test_construct_dns_name() { + assert_eq!( + construct_dns_name("place.stream", "key"), + "_lexicon.key.stream.place" + ); + assert_eq!( + construct_dns_name("app.bsky", "actor"), + "_lexicon.actor.bsky.app" + ); + assert_eq!( + construct_dns_name("app.bsky", "actor.profile"), + "_lexicon.actor.profile.bsky.app" + ); + assert_eq!( + construct_dns_name("place.stream", ""), + "_lexicon.stream.place" + ); + } + + #[tokio::test] + async fn test_mock_dns_resolver() { + let mut r = MockDnsResolver::new(); + r.add_record("place.stream", "key", "did:plc:test".into()); + assert_eq!( + r.resolve_lexicon_did("place.stream", "key").await.unwrap(), + "did:plc:test" + ); + assert!(r.resolve_lexicon_did("place.stream", "none").await.is_err()); + } + + #[tokio::test] + async fn test_mock_dns_from_nsid() { + let mut r = MockDnsResolver::new(); + r.add_record_from_nsid("app.bsky.actor.profile", "did:plc:bsky".into()) + .unwrap(); + assert_eq!( + r.resolve_lexicon_did("app.bsky", "actor.profile") + .await + .unwrap(), + "did:plc:bsky" + ); + } + + #[test] + fn test_extract_pds_endpoint() { + let doc = serde_json::json!({ + "service": [ + {"type": "Other", "serviceEndpoint": "https://nope"}, + {"type": "AtprotoPersonalDataServer", "serviceEndpoint": "https://pds.example.com/"}, + ] + }); + assert_eq!( + extract_pds_endpoint(&doc, "did:plc:x").unwrap(), + "https://pds.example.com" + ); + } + + #[test] + fn test_extract_pds_endpoint_missing() { + let doc = serde_json::json!({"service": []}); + assert!(extract_pds_endpoint(&doc, "did:plc:x").is_err()); + } +} diff --git a/mlf-atproto/src/lib.rs b/mlf-atproto/src/lib.rs new file mode 100644 index 0000000..41614b4 --- /dev/null +++ b/mlf-atproto/src/lib.rs @@ -0,0 +1,45 @@ +//! ATProto protocol plumbing for MLF. +//! +//! Low-level, protocol-level operations against the AT Protocol ecosystem: +//! +//! - [`identity`] — DID parsing, PLC / `did:web` resolution, NSID parsing, +//! `_lexicon.` TXT resolution +//! - [`xrpc`] — generic XRPC HTTP client (query + procedure) +//! - [`session`] — `createSession` + refresh, app-password auth +//! - [`records`] — typed wrappers for `getRecord` / `putRecord` / +//! `listRecords` / `deleteRecord` +//! - [`cid`] — client-side CID computation (DAG-CBOR + SHA-256 multihash) +//! +//! This crate is plumbing only — no MLF-specific domain logic lives here. +//! Consumers (`mlf-lexicon-fetcher`, `mlf-publish`) build domain operations +//! on top of these primitives. + +pub mod cid; +pub mod identity; +pub mod records; +pub mod session; +pub mod xrpc; + +pub use identity::{DnsResolver, MockDnsResolver, RealDnsResolver}; + +/// Result alias used throughout this crate. +pub type Result = std::result::Result; + +/// Umbrella error type for every submodule. +#[derive(thiserror::Error, Debug)] +pub enum Error { + #[error(transparent)] + Identity(#[from] identity::IdentityError), + + #[error(transparent)] + Xrpc(#[from] xrpc::XrpcError), + + #[error(transparent)] + Session(#[from] session::SessionError), + + #[error(transparent)] + Records(#[from] records::RecordError), + + #[error(transparent)] + Cid(#[from] cid::CidError), +} diff --git a/mlf-atproto/src/records.rs b/mlf-atproto/src/records.rs new file mode 100644 index 0000000..ad7ec7a --- /dev/null +++ b/mlf-atproto/src/records.rs @@ -0,0 +1,359 @@ +//! Typed wrappers for the AT Protocol record CRUD XRPC verbs. +//! +//! Reads (`getRecord`, `listRecords`) are unauthed public queries. +//! Writes (`putRecord`, `deleteRecord`) require a session access token. + +use crate::xrpc::{self, XrpcError}; +use serde::{Deserialize, Serialize}; + +#[derive(thiserror::Error, Debug)] +pub enum RecordError { + #[error(transparent)] + Xrpc(#[from] XrpcError), + + #[error("Record not found: {collection}/{rkey} in {repo}")] + NotFound { + repo: String, + collection: String, + rkey: String, + }, +} + +/// A single record as returned by `getRecord` / `listRecords`. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Record { + pub uri: String, + #[serde(default)] + pub cid: Option, + pub value: serde_json::Value, +} + +// --------------------------------------------------------------------------- +// getRecord +// --------------------------------------------------------------------------- + +#[derive(Serialize)] +struct GetRecordParams<'a> { + repo: &'a str, + collection: &'a str, + rkey: &'a str, +} + +/// Fetch a single record by repo + collection + rkey. +pub async fn get_record( + client: &reqwest::Client, + pds: &str, + repo: &str, + collection: &str, + rkey: &str, +) -> Result { + let params = GetRecordParams { + repo, + collection, + rkey, + }; + match xrpc::query::<_, Record>(client, pds, "com.atproto.repo.getRecord", ¶ms, None).await { + Ok(r) => Ok(r), + Err(XrpcError::HttpStatus { status: 400, .. }) => Err(RecordError::NotFound { + repo: repo.to_string(), + collection: collection.to_string(), + rkey: rkey.to_string(), + }), + Err(e) => Err(RecordError::Xrpc(e)), + } +} + +// --------------------------------------------------------------------------- +// listRecords +// --------------------------------------------------------------------------- + +#[derive(Serialize)] +struct ListRecordsParams<'a> { + repo: &'a str, + collection: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + cursor: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + limit: Option, +} + +#[derive(Debug, Deserialize)] +pub struct ListRecordsPage { + pub records: Vec, + #[serde(default)] + pub cursor: Option, +} + +/// Fetch a single page of records for a collection. +/// +/// For full enumeration, prefer [`list_all_records`]. +pub async fn list_records_page( + client: &reqwest::Client, + pds: &str, + repo: &str, + collection: &str, + cursor: Option<&str>, + limit: Option, +) -> Result { + let params = ListRecordsParams { + repo, + collection, + cursor, + limit, + }; + xrpc::query::<_, ListRecordsPage>(client, pds, "com.atproto.repo.listRecords", ¶ms, None) + .await + .map_err(Into::into) +} + +/// Fetch every record in a collection, paginating as needed. +pub async fn list_all_records( + client: &reqwest::Client, + pds: &str, + repo: &str, + collection: &str, +) -> Result, RecordError> { + let mut out = Vec::new(); + let mut cursor: Option = None; + loop { + let page = + list_records_page(client, pds, repo, collection, cursor.as_deref(), None).await?; + out.extend(page.records); + match page.cursor { + Some(c) if !c.is_empty() => cursor = Some(c), + _ => break, + } + } + Ok(out) +} + +// --------------------------------------------------------------------------- +// putRecord (authed) +// --------------------------------------------------------------------------- + +#[derive(Serialize)] +struct PutRecordInput<'a> { + repo: &'a str, + collection: &'a str, + rkey: &'a str, + record: &'a serde_json::Value, +} + +#[derive(Debug, Deserialize)] +pub struct PutRecordOutput { + pub uri: String, + pub cid: String, +} + +/// Create or replace a record at repo/collection/rkey. Requires an +/// access JWT from an authenticated session. +pub async fn put_record( + client: &reqwest::Client, + pds: &str, + access_jwt: &str, + repo: &str, + collection: &str, + rkey: &str, + record: &serde_json::Value, +) -> Result { + let input = PutRecordInput { + repo, + collection, + rkey, + record, + }; + xrpc::procedure::<_, PutRecordOutput>( + client, + pds, + "com.atproto.repo.putRecord", + &input, + Some(access_jwt), + ) + .await + .map_err(Into::into) +} + +// --------------------------------------------------------------------------- +// deleteRecord (authed) +// --------------------------------------------------------------------------- + +#[derive(Serialize)] +struct DeleteRecordInput<'a> { + repo: &'a str, + collection: &'a str, + rkey: &'a str, +} + +/// Delete a record at repo/collection/rkey. Requires an access JWT. +/// +/// Returns `Ok(())` whether or not the record existed — the PDS returns +/// success for idempotent deletes. +pub async fn delete_record( + client: &reqwest::Client, + pds: &str, + access_jwt: &str, + repo: &str, + collection: &str, + rkey: &str, +) -> Result<(), RecordError> { + let input = DeleteRecordInput { + repo, + collection, + rkey, + }; + let _: serde_json::Value = xrpc::procedure( + client, + pds, + "com.atproto.repo.deleteRecord", + &input, + Some(access_jwt), + ) + .await?; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + use wiremock::matchers::{bearer_token, body_partial_json, method, path, query_param}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + #[tokio::test] + async fn get_record_happy_path() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/xrpc/com.atproto.repo.getRecord")) + .and(query_param("repo", "did:plc:a")) + .and(query_param("collection", "com.atproto.lexicon.schema")) + .and(query_param("rkey", "com.example.thing")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "uri": "at://did:plc:a/com.atproto.lexicon.schema/com.example.thing", + "cid": "bafyREAL", + "value": {"id": "com.example.thing"}, + }))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let r = get_record( + &client, + &server.uri(), + "did:plc:a", + "com.atproto.lexicon.schema", + "com.example.thing", + ) + .await + .unwrap(); + assert_eq!(r.cid.as_deref(), Some("bafyREAL")); + } + + #[tokio::test] + async fn get_record_400_is_not_found() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/xrpc/com.atproto.repo.getRecord")) + .respond_with( + ResponseTemplate::new(400).set_body_json(json!({"error":"RecordNotFound"})), + ) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let err = get_record(&client, &server.uri(), "did:plc:a", "c", "k") + .await + .unwrap_err(); + assert!(matches!(err, RecordError::NotFound { .. })); + } + + #[tokio::test] + async fn list_all_paginates() { + let server = MockServer::start().await; + // First page returns cursor "p2". + Mock::given(method("GET")) + .and(path("/xrpc/com.atproto.repo.listRecords")) + .and(query_param("repo", "did:plc:a")) + .and(query_param("collection", "c")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "records": [{"uri":"at://did:plc:a/c/one","cid":"b1","value":{}}], + "cursor": "p2" + }))) + .up_to_n_times(1) + .mount(&server) + .await; + // Second page, no cursor. + Mock::given(method("GET")) + .and(path("/xrpc/com.atproto.repo.listRecords")) + .and(query_param("cursor", "p2")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "records": [{"uri":"at://did:plc:a/c/two","cid":"b2","value":{}}] + }))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let all = list_all_records(&client, &server.uri(), "did:plc:a", "c") + .await + .unwrap(); + assert_eq!(all.len(), 2); + assert!(all.iter().any(|r| r.uri.ends_with("/one"))); + assert!(all.iter().any(|r| r.uri.ends_with("/two"))); + } + + #[tokio::test] + async fn put_record_sends_auth_and_body() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/xrpc/com.atproto.repo.putRecord")) + .and(bearer_token("access")) + .and(body_partial_json(json!({ + "repo": "did:plc:a", + "collection": "com.atproto.lexicon.schema", + "rkey": "com.example.thing", + "record": {"id": "com.example.thing"} + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "uri": "at://did:plc:a/com.atproto.lexicon.schema/com.example.thing", + "cid": "bafyNEW" + }))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let out = put_record( + &client, + &server.uri(), + "access", + "did:plc:a", + "com.atproto.lexicon.schema", + "com.example.thing", + &json!({"id": "com.example.thing"}), + ) + .await + .unwrap(); + assert_eq!(out.cid, "bafyNEW"); + } + + #[tokio::test] + async fn delete_record_sends_auth_and_body() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/xrpc/com.atproto.repo.deleteRecord")) + .and(bearer_token("access")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({}))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + delete_record( + &client, + &server.uri(), + "access", + "did:plc:a", + "com.atproto.lexicon.schema", + "com.example.thing", + ) + .await + .unwrap(); + } +} diff --git a/mlf-atproto/src/session.rs b/mlf-atproto/src/session.rs new file mode 100644 index 0000000..ed4b9b9 --- /dev/null +++ b/mlf-atproto/src/session.rs @@ -0,0 +1,194 @@ +//! PDS session management. +//! +//! App-password authentication via `com.atproto.server.createSession` +//! plus token refresh via `com.atproto.server.refreshSession`. A [`Session`] +//! holds the credentials needed for authed XRPC calls. + +use crate::xrpc::{self, XrpcError}; +use serde::{Deserialize, Serialize}; + +#[derive(thiserror::Error, Debug)] +pub enum SessionError { + #[error(transparent)] + Xrpc(#[from] XrpcError), + + #[error("Invalid credentials (handle or password rejected by PDS)")] + InvalidCredentials, + + #[error("Session refresh failed: {0}")] + RefreshFailed(String), +} + +/// An authenticated PDS session. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Session { + /// Publishing DID (from the session response). + pub did: String, + /// The user's handle at session creation time. + pub handle: String, + /// Short-lived bearer token for authed XRPC calls. + #[serde(rename = "accessJwt")] + pub access_jwt: String, + /// Long-lived token for refreshing expired sessions. + #[serde(rename = "refreshJwt")] + pub refresh_jwt: String, +} + +#[derive(Serialize)] +struct CreateSessionInput<'a> { + identifier: &'a str, + password: &'a str, +} + +/// Create a new session using a handle + app password. +/// +/// Hits `com.atproto.server.createSession` on the given PDS. +pub async fn create_session( + client: &reqwest::Client, + pds: &str, + identifier: &str, + password: &str, +) -> Result { + let body = CreateSessionInput { + identifier, + password, + }; + match xrpc::procedure::<_, Session>( + client, + pds, + "com.atproto.server.createSession", + &body, + None, + ) + .await + { + Ok(s) => Ok(s), + Err(XrpcError::HttpStatus { status: 401, .. }) => Err(SessionError::InvalidCredentials), + Err(e) => Err(SessionError::Xrpc(e)), + } +} + +/// Refresh an existing session using its `refreshJwt`. +/// +/// On success, returns the updated session. The original `Session` should +/// be replaced with the returned one. +pub async fn refresh_session( + client: &reqwest::Client, + pds: &str, + refresh_jwt: &str, +) -> Result { + xrpc::procedure::<_, Session>( + client, + pds, + "com.atproto.server.refreshSession", + &serde_json::json!({}), + Some(refresh_jwt), + ) + .await + .map_err(|e| match e { + XrpcError::HttpStatus { status, body, .. } => { + SessionError::RefreshFailed(format!("HTTP {status}: {body}")) + } + other => SessionError::Xrpc(other), + }) +} + +/// Fetch the current session info using an access JWT. +/// +/// Useful for verifying a stored token is still valid before use. +pub async fn get_session( + client: &reqwest::Client, + pds: &str, + access_jwt: &str, +) -> Result { + xrpc::query::<_, SessionInfo>( + client, + pds, + "com.atproto.server.getSession", + &(), + Some(access_jwt), + ) + .await + .map_err(Into::into) +} + +/// Response shape for `getSession` (a subset of the full session). +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SessionInfo { + pub did: String, + pub handle: String, +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + use wiremock::matchers::{body_partial_json, method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + #[tokio::test] + async fn create_session_parses_response() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/xrpc/com.atproto.server.createSession")) + .and(body_partial_json(json!({ + "identifier": "matt.example.com", + "password": "hunter2", + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "did": "did:plc:abc", + "handle": "matt.example.com", + "accessJwt": "access-token", + "refreshJwt": "refresh-token", + }))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let s = create_session(&client, &server.uri(), "matt.example.com", "hunter2") + .await + .unwrap(); + assert_eq!(s.did, "did:plc:abc"); + assert_eq!(s.handle, "matt.example.com"); + assert_eq!(s.access_jwt, "access-token"); + assert_eq!(s.refresh_jwt, "refresh-token"); + } + + #[tokio::test] + async fn create_session_401_is_invalid_credentials() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/xrpc/com.atproto.server.createSession")) + .respond_with(ResponseTemplate::new(401).set_body_json(json!({"error":"AuthRequired"}))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let err = create_session(&client, &server.uri(), "x", "wrong") + .await + .unwrap_err(); + assert!(matches!(err, SessionError::InvalidCredentials)); + } + + #[tokio::test] + async fn refresh_session_parses_response() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/xrpc/com.atproto.server.refreshSession")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "did": "did:plc:abc", + "handle": "matt.example.com", + "accessJwt": "new-access", + "refreshJwt": "new-refresh", + }))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let s = refresh_session(&client, &server.uri(), "old-refresh") + .await + .unwrap(); + assert_eq!(s.access_jwt, "new-access"); + assert_eq!(s.refresh_jwt, "new-refresh"); + } +} diff --git a/mlf-atproto/src/xrpc.rs b/mlf-atproto/src/xrpc.rs new file mode 100644 index 0000000..c817175 --- /dev/null +++ b/mlf-atproto/src/xrpc.rs @@ -0,0 +1,192 @@ +//! Generic XRPC HTTP client. +//! +//! XRPC is the AT Protocol's HTTP RPC convention: +//! `GET /xrpc/?param=value` — a query +//! `POST /xrpc/` with JSON body — a procedure +//! +//! This module provides `query` / `procedure` helpers that handle +//! URL construction, auth headers, and JSON (de)serialization. +//! Typed wrappers for specific NSIDs (getRecord, etc.) live in +//! [`crate::records`] and [`crate::session`]. + +use serde::{Serialize, de::DeserializeOwned}; + +#[derive(thiserror::Error, Debug)] +pub enum XrpcError { + #[error("HTTP request failed: {0}")] + Transport(String), + + #[error("XRPC {nsid} returned HTTP {status}: {body}")] + HttpStatus { + nsid: String, + status: u16, + body: String, + }, + + #[error("Failed to parse XRPC response JSON: {0}")] + ResponseParse(String), +} + +/// Perform an XRPC query (GET). +/// +/// `params` is serialized as URL query parameters. `auth` is an optional +/// bearer token (the `accessJwt` from a session). +pub async fn query( + client: &reqwest::Client, + pds: &str, + nsid: &str, + params: &P, + auth: Option<&str>, +) -> Result +where + P: Serialize, + R: DeserializeOwned, +{ + let url = format!("{}/xrpc/{}", pds.trim_end_matches('/'), nsid); + let mut req = client.get(&url).query(params); + if let Some(token) = auth { + req = req.bearer_auth(token); + } + send(req, nsid).await +} + +/// Perform an XRPC procedure (POST with JSON body). +pub async fn procedure( + client: &reqwest::Client, + pds: &str, + nsid: &str, + body: &B, + auth: Option<&str>, +) -> Result +where + B: Serialize, + R: DeserializeOwned, +{ + let url = format!("{}/xrpc/{}", pds.trim_end_matches('/'), nsid); + let mut req = client.post(&url).json(body); + if let Some(token) = auth { + req = req.bearer_auth(token); + } + send(req, nsid).await +} + +async fn send( + req: reqwest::RequestBuilder, + nsid: &str, +) -> Result { + let resp = req + .send() + .await + .map_err(|e| XrpcError::Transport(e.to_string()))?; + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + return Err(XrpcError::HttpStatus { + nsid: nsid.to_string(), + status: status.as_u16(), + body, + }); + } + resp.json::() + .await + .map_err(|e| XrpcError::ResponseParse(e.to_string())) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + use wiremock::matchers::{bearer_token, method, path, query_param}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + #[tokio::test] + async fn query_builds_url_and_parses_response() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/xrpc/com.example.thing")) + .and(query_param("who", "me")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": 42}))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let got: serde_json::Value = query( + &client, + &server.uri(), + "com.example.thing", + &[("who", "me")], + None, + ) + .await + .unwrap(); + assert_eq!(got, json!({"value": 42})); + } + + #[tokio::test] + async fn query_sends_bearer_token() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/xrpc/com.example.thing")) + .and(bearer_token("tok")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"ok": true}))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let _: serde_json::Value = query( + &client, + &server.uri(), + "com.example.thing", + &(), + Some("tok"), + ) + .await + .unwrap(); + } + + #[tokio::test] + async fn procedure_posts_json_body() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/xrpc/com.example.do")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"out": "yes"}))) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let got: serde_json::Value = procedure( + &client, + &server.uri(), + "com.example.do", + &json!({"in": "yes"}), + None, + ) + .await + .unwrap(); + assert_eq!(got, json!({"out": "yes"})); + } + + #[tokio::test] + async fn http_error_surfaces_status_and_body() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/xrpc/com.example.oops")) + .respond_with(ResponseTemplate::new(418).set_body_string("teapot")) + .mount(&server) + .await; + + let client = reqwest::Client::new(); + let err = + query::<_, serde_json::Value>(&client, &server.uri(), "com.example.oops", &(), None) + .await + .unwrap_err(); + match err { + XrpcError::HttpStatus { status, body, nsid } => { + assert_eq!(status, 418); + assert_eq!(body, "teapot"); + assert_eq!(nsid, "com.example.oops"); + } + e => panic!("unexpected error: {e:?}"), + } + } +} diff --git a/mlf-cli/src/check.rs b/mlf-cli/src/check.rs index 2ad7cce..fc9852b 100644 --- a/mlf-cli/src/check.rs +++ b/mlf-cli/src/check.rs @@ -1,4 +1,4 @@ -use crate::config::{find_project_root, get_mlf_cache_dir, ConfigError, MlfConfig}; +use crate::config::{ConfigError, MlfConfig, find_project_root, get_mlf_cache_dir}; use crate::workspace_ext::workspace_with_std_and_cache; use miette::Diagnostic; use mlf_diagnostics::{ParseDiagnostic, ValidationDiagnostic}; @@ -37,7 +37,6 @@ pub enum CheckError { help: Option, }, - #[error("Record validation failed")] #[diagnostic(code(mlf::check::record_validation))] RecordValidation { @@ -49,12 +48,14 @@ pub enum CheckError { ConfigError(#[from] ConfigError), } -pub fn run_check(input_paths: Vec, explicit_root: Option) -> Result<(), CheckError> { - let current_dir = std::env::current_dir() - .map_err(|e| CheckError::ReadFile { - path: ".".to_string(), - source: e, - })?; +pub fn run_check( + input_paths: Vec, + explicit_root: Option, +) -> Result<(), CheckError> { + let current_dir = std::env::current_dir().map_err(|e| CheckError::ReadFile { + path: ".".to_string(), + source: e, + })?; // Determine root directory and input paths let (root_dir, file_paths) = if input_paths.is_empty() { @@ -65,7 +66,10 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R let config = MlfConfig::load(&config_path)?; let source_dir = project_root.join(&config.source.directory); let root = explicit_root.unwrap_or_else(|| source_dir.clone()); - println!("Using source directory from mlf.toml: {}", config.source.directory); + println!( + "Using source directory from mlf.toml: {}", + config.source.directory + ); // Collect all .mlf files from source directory let files = collect_mlf_files(&source_dir)?; @@ -120,11 +124,10 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R }; // Try to load cached lexicons from .mlf directory - let current_dir = std::env::current_dir() - .map_err(|e| CheckError::ReadFile { - path: ".".to_string(), - source: e, - })?; + let current_dir = std::env::current_dir().map_err(|e| CheckError::ReadFile { + path: ".".to_string(), + source: e, + })?; let mlf_cache_dir = find_project_root(¤t_dir) .ok() @@ -141,11 +144,9 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R let mut had_parse_errors = false; for file_path in &file_paths { - let source = std::fs::read_to_string(file_path).map_err(|source| { - CheckError::ReadFile { - path: file_path.display().to_string(), - source, - } + let source = std::fs::read_to_string(file_path).map_err(|source| CheckError::ReadFile { + path: file_path.display().to_string(), + source, })?; let filename = file_path.display().to_string(); @@ -163,7 +164,8 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R let namespace = extract_namespace(&file_path, &root_dir)?; if let Err(e) = workspace.add_module(namespace.clone(), lexicon.clone()) { - let diagnostic = ValidationDiagnostic::new(filename.clone(), source.clone(), namespace.clone(), e); + let diagnostic = + ValidationDiagnostic::new(filename.clone(), source.clone(), namespace.clone(), e); eprintln!("{:?}", miette::Report::new(diagnostic)); had_parse_errors = true; continue; @@ -181,7 +183,8 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R if let Err(e) = workspace.resolve() { // Collect all modules that have errors - let mut modules_with_errors: std::collections::BTreeMap, String)> = std::collections::BTreeMap::new(); + let mut modules_with_errors: std::collections::BTreeMap, String)> = + std::collections::BTreeMap::new(); // First, add all explicitly checked files for (filename, namespace, source) in &source_files { @@ -198,23 +201,32 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R // Try multiple locations for the source file let mut possible_paths = vec![ // Check in lexicons/ directory (common structure) - current_dir.join("lexicons").join(format!("{}.mlf", namespace_path)), + current_dir + .join("lexicons") + .join(format!("{}.mlf", namespace_path)), // Check in source directory from config - current_dir.join("src").join(format!("{}.mlf", namespace_path)), + current_dir + .join("src") + .join(format!("{}.mlf", namespace_path)), // Check relative to current directory current_dir.join(format!("{}.mlf", namespace_path)), ]; // Add cache directory if available (lexicons are in lexicons/mlf/ subdirectory) if let Some(cache_dir) = &mlf_cache_dir { - possible_paths.push(cache_dir.join("lexicons").join("mlf").join(format!("{}.mlf", namespace_path))); + possible_paths.push( + cache_dir + .join("lexicons") + .join("mlf") + .join(format!("{}.mlf", namespace_path)), + ); } for path in possible_paths { if let Ok(source) = std::fs::read_to_string(&path) { modules_with_errors.insert( error_namespace.to_string(), - (Some(path.display().to_string()), source) + (Some(path.display().to_string()), source), ); source_loaded = true; break; @@ -223,10 +235,7 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R if !source_loaded { // Couldn't load source, add placeholder - modules_with_errors.insert( - error_namespace.to_string(), - (None, String::new()) - ); + modules_with_errors.insert(error_namespace.to_string(), (None, String::new())); } } } @@ -234,21 +243,34 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R // Show diagnostics for all modules with errors for (namespace, (filename_opt, source)) in &modules_with_errors { // Only show diagnostic if this module has errors - let has_errors = e.errors.iter().any(|error| { - mlf_diagnostics::get_error_module_namespace_str(error) == namespace - }); + let has_errors = e + .errors + .iter() + .any(|error| mlf_diagnostics::get_error_module_namespace_str(error) == namespace); if has_errors { if let Some(filename) = filename_opt { // Have source file, show full diagnostic - let diagnostic = ValidationDiagnostic::new(filename.clone(), source.clone(), namespace.clone(), e.clone()); + let diagnostic = ValidationDiagnostic::new( + filename.clone(), + source.clone(), + namespace.clone(), + e.clone(), + ); eprintln!("{:?}", miette::Report::new(diagnostic)); } else { // No source available, just list the errors - let error_count = e.errors.iter() - .filter(|err| mlf_diagnostics::get_error_module_namespace_str(err) == namespace) + let error_count = e + .errors + .iter() + .filter(|err| { + mlf_diagnostics::get_error_module_namespace_str(err) == namespace + }) .count(); - eprintln!("\n{}: {} error(s) (source not available)", namespace, error_count); + eprintln!( + "\n{}: {} error(s) (source not available)", + namespace, error_count + ); } } } @@ -263,19 +285,17 @@ pub fn run_check(input_paths: Vec, explicit_root: Option) -> R } pub fn validate(lexicon_path: PathBuf, record_path: PathBuf) -> Result<(), CheckError> { - let lexicon_source = std::fs::read_to_string(&lexicon_path).map_err(|source| { - CheckError::ReadFile { + let lexicon_source = + std::fs::read_to_string(&lexicon_path).map_err(|source| CheckError::ReadFile { path: lexicon_path.display().to_string(), source, - } - })?; + })?; - let record_source = std::fs::read_to_string(&record_path).map_err(|source| { - CheckError::ReadFile { + let record_source = + std::fs::read_to_string(&record_path).map_err(|source| CheckError::ReadFile { path: record_path.display().to_string(), source, - } - })?; + })?; let lexicon = mlf_lang::parse_lexicon(&lexicon_source).map_err(|e| { let diagnostic = ParseDiagnostic::new( @@ -290,8 +310,8 @@ pub fn validate(lexicon_path: PathBuf, record_path: PathBuf) -> Result<(), Check } })?; - let record: serde_json::Value = serde_json::from_str(&record_source) - .map_err(|source| CheckError::ParseJson { source })?; + let record: serde_json::Value = + serde_json::from_str(&record_source).map_err(|source| CheckError::ParseJson { source })?; println!("✓ Lexicon parsed successfully"); println!("✓ JSON record parsed successfully"); @@ -348,27 +368,36 @@ fn collect_mlf_files(dir: &std::path::Path) -> Result, CheckError> /// Extract namespace from file path relative to root directory /// e.g., root=/project/lexicons, file=/project/lexicons/com/example/foo.mlf -> com.example.foo -fn extract_namespace(file_path: &std::path::Path, root_dir: &std::path::Path) -> Result { +fn extract_namespace( + file_path: &std::path::Path, + root_dir: &std::path::Path, +) -> Result { // Get the canonical paths to handle . and .. correctly - let file_canonical = file_path.canonicalize().map_err(|source| CheckError::ReadFile { - path: file_path.display().to_string(), - source, - })?; + let file_canonical = file_path + .canonicalize() + .map_err(|source| CheckError::ReadFile { + path: file_path.display().to_string(), + source, + })?; - let root_canonical = root_dir.canonicalize().map_err(|source| CheckError::ReadFile { - path: root_dir.display().to_string(), - source, - })?; + let root_canonical = root_dir + .canonicalize() + .map_err(|source| CheckError::ReadFile { + path: root_dir.display().to_string(), + source, + })?; // Get relative path from root to file - let relative_path = file_canonical.strip_prefix(&root_canonical) - .map_err(|_| CheckError::ValidationErrors { - help: Some(format!( - "File {} is not within root directory {}", - file_path.display(), - root_dir.display() - )), - })?; + let relative_path = + file_canonical + .strip_prefix(&root_canonical) + .map_err(|_| CheckError::ValidationErrors { + help: Some(format!( + "File {} is not within root directory {}", + file_path.display(), + root_dir.display() + )), + })?; // Convert path to namespace let mut components = Vec::new(); @@ -389,7 +418,10 @@ fn extract_namespace(file_path: &std::path::Path, root_dir: &std::path::Path) -> if components.is_empty() { return Err(CheckError::ValidationErrors { - help: Some(format!("Could not extract namespace from path: {}", file_path.display())), + help: Some(format!( + "Could not extract namespace from path: {}", + file_path.display() + )), }); } diff --git a/mlf-cli/src/config.rs b/mlf-cli/src/config.rs index 8b3f8e5..f12d79e 100644 --- a/mlf-cli/src/config.rs +++ b/mlf-cli/src/config.rs @@ -29,7 +29,10 @@ pub struct MlfConfig { #[derive(Debug, Serialize, Deserialize)] pub struct SourceConfig { - #[serde(default = "default_source_directory", skip_serializing_if = "is_default_source_directory")] + #[serde( + default = "default_source_directory", + skip_serializing_if = "is_default_source_directory" + )] pub directory: String, } @@ -60,10 +63,16 @@ pub struct DependenciesConfig { #[serde(default)] pub dependencies: Vec, - #[serde(default = "default_allow_transitive_deps", skip_serializing_if = "is_default_allow_transitive_deps")] + #[serde( + default = "default_allow_transitive_deps", + skip_serializing_if = "is_default_allow_transitive_deps" + )] pub allow_transitive_deps: bool, - #[serde(default = "default_optimize_transitive_fetches", skip_serializing_if = "is_default_optimize_transitive_fetches")] + #[serde( + default = "default_optimize_transitive_fetches", + skip_serializing_if = "is_default_optimize_transitive_fetches" + )] pub optimize_transitive_fetches: bool, } @@ -220,13 +229,22 @@ impl LockFile { Ok(()) } - pub fn add_lexicon(&mut self, nsid: String, did: String, checksum: String, dependencies: Vec) { - self.lexicons.insert(nsid.clone(), LockedLexicon { - nsid, - did, - checksum, - dependencies, - }); + pub fn add_lexicon( + &mut self, + nsid: String, + did: String, + checksum: String, + dependencies: Vec, + ) { + self.lexicons.insert( + nsid.clone(), + LockedLexicon { + nsid, + did, + checksum, + dependencies, + }, + ); } } diff --git a/mlf-cli/src/fetch.rs b/mlf-cli/src/fetch.rs index c6da523..c06a30c 100644 --- a/mlf-cli/src/fetch.rs +++ b/mlf-cli/src/fetch.rs @@ -1,6 +1,8 @@ -use crate::config::{find_project_root, get_mlf_cache_dir, init_mlf_cache, ConfigError, MlfConfig, LockFile}; -use mlf_lexicon_fetcher::{optimize_fetch_patterns, ProductionLexiconFetcher}; +use crate::config::{ + ConfigError, LockFile, MlfConfig, find_project_root, get_mlf_cache_dir, init_mlf_cache, +}; use miette::Diagnostic; +use mlf_lexicon_fetcher::{ProductionLexiconFetcher, optimize_fetch_patterns}; use sha2::{Digest, Sha256}; use std::collections::HashSet; use thiserror::Error; @@ -36,14 +38,17 @@ pub enum FetchError { InvalidNsid(String), } - - /// Main entry point for fetch command -pub async fn run_fetch(nsid: Option, save: bool, update: bool, locked: bool) -> Result<(), FetchError> { +pub async fn run_fetch( + nsid: Option, + save: bool, + update: bool, + locked: bool, +) -> Result<(), FetchError> { // Validate flags if update && locked { return Err(FetchError::HttpError( - "Cannot use --update and --locked together".to_string() + "Cannot use --update and --locked together".to_string(), )); } @@ -69,12 +74,15 @@ pub async fn run_fetch(nsid: Option, save: bool, update: bool, locked: b fetch_transitive_dependencies( &project_root, &mut lockfile, - config.dependencies.optimize_transitive_fetches - ).await?; + config.dependencies.optimize_transitive_fetches, + ) + .await?; } // Save lockfile - lockfile.save(&lockfile_path).map_err(FetchError::NoProjectRoot)?; + lockfile + .save(&lockfile_path) + .map_err(FetchError::NoProjectRoot)?; println!("\n→ Updated mlf-lock.toml"); // Save to mlf.toml if --save flag is provided @@ -117,7 +125,11 @@ fn ensure_project_root(current_dir: &std::path::Path) -> Result Result<(), FetchError> { +async fn fetch_all_dependencies( + project_root: &std::path::Path, + update: bool, + locked: bool, +) -> Result<(), FetchError> { // Load mlf.toml let config_path = project_root.join("mlf.toml"); let config = MlfConfig::load(&config_path).map_err(FetchError::NoProjectRoot)?; @@ -138,7 +150,7 @@ async fn fetch_all_dependencies(project_root: &std::path::Path, update: bool, lo if locked { if !has_existing_lockfile { return Err(FetchError::HttpError( - "No lockfile found. Run `mlf fetch` first to create mlf-lock.toml".to_string() + "No lockfile found. Run `mlf fetch` first to create mlf-lock.toml".to_string(), )); } @@ -157,10 +169,16 @@ async fn fetch_all_dependencies(project_root: &std::path::Path, update: bool, lo "fresh" }; - println!("Fetching {} dependencies... (mode: {}, transitive deps: {})", - config.dependencies.dependencies.len(), - mode, - if allow_transitive { "enabled" } else { "disabled" }); + println!( + "Fetching {} dependencies... (mode: {}, transitive deps: {})", + config.dependencies.dependencies.len(), + mode, + if allow_transitive { + "enabled" + } else { + "disabled" + } + ); // In update mode or if no lockfile, do full fetch // In normal mode with lockfile, use lockfile for cached entries @@ -190,11 +208,18 @@ async fn fetch_all_dependencies(project_root: &std::path::Path, update: bool, lo // If transitive dependencies are enabled, fetch them if allow_transitive { - fetch_transitive_dependencies(&project_root, &mut lockfile, config.dependencies.optimize_transitive_fetches).await?; + fetch_transitive_dependencies( + &project_root, + &mut lockfile, + config.dependencies.optimize_transitive_fetches, + ) + .await?; } // Save the lockfile - lockfile.save(&lockfile_path).map_err(FetchError::NoProjectRoot)?; + lockfile + .save(&lockfile_path) + .map_err(FetchError::NoProjectRoot)?; println!("\n→ Updated mlf-lock.toml"); if !errors.is_empty() { @@ -212,7 +237,10 @@ async fn fetch_all_dependencies(project_root: &std::path::Path, update: bool, lo ))); } - println!("\n✓ Successfully fetched all {} dependencies", success_count); + println!( + "\n✓ Successfully fetched all {} dependencies", + success_count + ); Ok(()) } @@ -220,7 +248,7 @@ async fn fetch_all_dependencies(project_root: &std::path::Path, update: bool, lo async fn fetch_transitive_dependencies( project_root: &std::path::Path, lockfile: &mut LockFile, - optimize_fetches: bool + optimize_fetches: bool, ) -> Result<(), FetchError> { let mut fetched_nsids = HashSet::new(); // Track NSIDs from lockfile as already fetched @@ -261,8 +289,11 @@ async fn fetch_transitive_dependencies( // Optimize the fetch patterns to reduce number of fetches let optimized_patterns = optimize_fetch_patterns(&new_deps); - println!("\n→ Found {} unresolved reference(s), fetching {} optimized pattern(s)...", - new_deps.len(), optimized_patterns.len()); + println!( + "\n→ Found {} unresolved reference(s), fetching {} optimized pattern(s)...", + new_deps.len(), + optimized_patterns.len() + ); // Track which patterns are wildcards and their constituent NSIDs let mut wildcard_failures: Vec<(String, Vec)> = Vec::new(); @@ -280,7 +311,8 @@ async fn fetch_transitive_dependencies( // If this was a wildcard that failed, collect the individual NSIDs for retry if is_wildcard { let pattern_prefix = pattern.strip_suffix(".*").unwrap(); - let matching_nsids: Vec = new_deps.iter() + let matching_nsids: Vec = new_deps + .iter() .filter(|nsid| nsid.starts_with(pattern_prefix)) .cloned() .collect(); @@ -312,7 +344,10 @@ async fn fetch_transitive_dependencies( match fetch_lexicon_with_lock(&broader, project_root, lockfile).await { Ok(()) => continue, Err(e) => { - eprintln!(" Warning: broader pattern {} also failed: {}", broader, e); + eprintln!( + " Warning: broader pattern {} also failed: {}", + broader, e + ); } } } @@ -322,9 +357,15 @@ async fn fetch_transitive_dependencies( } if !still_failing.is_empty() { - println!("\n→ Falling back to individual NSID fetches for patterns that couldn't be broadened..."); + println!( + "\n→ Falling back to individual NSID fetches for patterns that couldn't be broadened..." + ); for (failed_pattern, nsids) in still_failing { - println!(" Retrying {} NSIDs from failed pattern: {}", nsids.len(), failed_pattern); + println!( + " Retrying {} NSIDs from failed pattern: {}", + nsids.len(), + failed_pattern + ); for nsid in nsids { if !fetched_nsids.contains(&nsid) { @@ -344,8 +385,10 @@ async fn fetch_transitive_dependencies( } } else { // Fetch individually without optimization (safer, more predictable) - println!("\n→ Found {} unresolved reference(s), fetching individually...", - new_deps.len()); + println!( + "\n→ Found {} unresolved reference(s), fetching individually...", + new_deps.len() + ); for nsid in &new_deps { println!("\nFetching transitive dependency: {}", nsid); @@ -367,13 +410,19 @@ async fn fetch_transitive_dependencies( /// Fetch dependencies using the lockfile /// This refetches each lexicon from its recorded DID and verifies the checksum -async fn fetch_from_lockfile(project_root: &std::path::Path, lockfile: &LockFile) -> Result<(), FetchError> { +async fn fetch_from_lockfile( + project_root: &std::path::Path, + lockfile: &LockFile, +) -> Result<(), FetchError> { if lockfile.lexicons.is_empty() { println!("Lockfile is empty"); return Ok(()); } - println!("Fetching {} lexicon(s) from lockfile...", lockfile.lexicons.len()); + println!( + "Fetching {} lexicon(s) from lockfile...", + lockfile.lexicons.len() + ); let mut errors = Vec::new(); let mut success_count = 0; @@ -506,13 +555,19 @@ fn save_dependency(project_root: &std::path::Path, nsid: &str) -> Result<(), Fet } config.dependencies.dependencies.push(nsid.to_string()); - config.save(&config_path).map_err(FetchError::NoProjectRoot)?; + config + .save(&config_path) + .map_err(FetchError::NoProjectRoot)?; println!("Added '{}' to dependencies in mlf.toml", nsid); Ok(()) } -async fn fetch_lexicon_with_lock(nsid: &str, project_root: &std::path::Path, lockfile: &mut LockFile) -> Result<(), FetchError> { +async fn fetch_lexicon_with_lock( + nsid: &str, + project_root: &std::path::Path, + lockfile: &mut LockFile, +) -> Result<(), FetchError> { // Initialize .mlf directory init_mlf_cache(project_root).map_err(FetchError::InitFailed)?; let mlf_dir = get_mlf_cache_dir(&project_root); @@ -585,10 +640,19 @@ async fn fetch_lexicon_with_lock(nsid: &str, project_root: &std::path::Path, loc let dependencies = extract_dependencies_from_json(&fetched.lexicon); // Update lockfile with DID from fetcher metadata - lockfile.add_lexicon(fetched.nsid.clone(), fetched.did.clone(), hash, dependencies); + lockfile.add_lexicon( + fetched.nsid.clone(), + fetched.did.clone(), + hash, + dependencies, + ); } - println!("✓ Successfully fetched {} lexicon(s) for {}", result.lexicons.len(), nsid); + println!( + "✓ Successfully fetched {} lexicon(s) for {}", + result.lexicons.len(), + nsid + ); Ok(()) } @@ -613,7 +677,6 @@ fn validate_nsid_format(nsid: &str) -> Result<(), FetchError> { Ok(()) } - /// Calculate SHA-256 hash of content fn calculate_hash(content: &str) -> String { let mut hasher = Sha256::new(); @@ -661,7 +724,9 @@ fn extract_dependencies_from_json(json: &serde_json::Value) -> Vec { /// Extract external references from MLF files that need to be resolved /// Returns a set of namespace patterns (not full NSIDs) that need to be fetched -fn collect_unresolved_references(project_root: &std::path::Path) -> Result, FetchError> { +fn collect_unresolved_references( + project_root: &std::path::Path, +) -> Result, FetchError> { use mlf_lang::{parser, workspace::Workspace}; let mlf_dir = get_mlf_cache_dir(project_root); @@ -672,15 +737,19 @@ fn collect_unresolved_references(project_root: &std::path::Path) -> Result) -> std::io::Result<()> { + fn collect_mlf_files( + dir: &std::path::Path, + files: &mut Vec, + ) -> std::io::Result<()> { if dir.is_dir() { for entry in std::fs::read_dir(dir)? { let entry = entry?; @@ -704,11 +773,12 @@ fn collect_unresolved_references(project_root: &std::path::Path) -> Result "place.stream.key" - let relative_path = mlf_file.strip_prefix(&mlf_lexicons_dir) - .map_err(|_| FetchError::IoError(std::io::Error::new( + let relative_path = mlf_file.strip_prefix(&mlf_lexicons_dir).map_err(|_| { + FetchError::IoError(std::io::Error::new( std::io::ErrorKind::Other, - "Failed to compute relative path" - )))?; + "Failed to compute relative path", + )) + })?; let namespace = relative_path .with_extension("") @@ -827,4 +897,3 @@ mod tests { assert_eq!(broaden_pattern("app.rocksky.playlist"), None); } } - diff --git a/mlf-cli/src/generate/code.rs b/mlf-cli/src/generate/code.rs index 32cffbb..bc297b3 100644 --- a/mlf-cli/src/generate/code.rs +++ b/mlf-cli/src/generate/code.rs @@ -49,12 +49,10 @@ pub fn run( // Load mlf.toml if available let project_root = crate::config::find_project_root(¤t_dir).ok(); - let config = project_root - .as_ref() - .and_then(|root| { - let config_path = root.join("mlf.toml"); - crate::config::MlfConfig::load(&config_path).ok() - }); + let config = project_root.as_ref().and_then(|root| { + let config_path = root.join("mlf.toml"); + crate::config::MlfConfig::load(&config_path).ok() + }); // Determine generator name let generator_name = if let Some(explicit) = generator_name { @@ -89,7 +87,11 @@ pub fn run( } })?; - println!("Using generator: {} ({})", generator.name(), generator.description()); + println!( + "Using generator: {} ({})", + generator.name(), + generator.description() + ); println!("Output extension: {}\n", generator.file_extension()); // Determine output directory @@ -105,7 +107,10 @@ pub fn run( path: "mlf.toml".to_string(), source: std::io::Error::new( std::io::ErrorKind::NotFound, - format!("No output configured for generator '{}' in mlf.toml", generator_name) + format!( + "No output configured for generator '{}' in mlf.toml", + generator_name + ), ), })? } else { @@ -113,7 +118,7 @@ pub fn run( path: "mlf.toml".to_string(), source: std::io::Error::new( std::io::ErrorKind::NotFound, - "No mlf.toml found and no --output flag provided" + "No mlf.toml found and no --output flag provided", ), }); }; @@ -136,7 +141,7 @@ pub fn run( path: "input".to_string(), source: std::io::Error::new( std::io::ErrorKind::NotFound, - "No input files specified and no mlf.toml found" + "No input files specified and no mlf.toml found", ), }); } @@ -198,16 +203,17 @@ pub fn run( .ok() .map(|root| crate::config::get_mlf_cache_dir(&root)); - let mut workspace = match crate::workspace_ext::workspace_with_std_and_cache(mlf_cache_dir.as_deref()) { - Ok(ws) => ws, - Err(e) => { - errors.push(( - file_path.display().to_string(), - format!("Failed to load workspace: {}", e), - )); - continue; - } - }; + let mut workspace = + match crate::workspace_ext::workspace_with_std_and_cache(mlf_cache_dir.as_deref()) { + Ok(ws) => ws, + Err(e) => { + errors.push(( + file_path.display().to_string(), + format!("Failed to load workspace: {}", e), + )); + continue; + } + }; // Add the module to the workspace if let Err(e) = workspace.add_module(namespace.clone(), lexicon.clone()) { @@ -237,7 +243,10 @@ pub fn run( let generated_code = match generator.generate(&ctx) { Ok(code) => code, Err(e) => { - errors.push((file_path.display().to_string(), format!("Generation error: {}", e))); + errors.push(( + file_path.display().to_string(), + format!("Generation error: {}", e), + )); continue; } }; @@ -326,18 +335,16 @@ fn extract_namespace(file_path: &Path, root_dir: &Path) -> Result Result, output_dir: Option, explicit_root: Option, flat: bool) -> Result<(), GenerateError> { - let current_dir = std::env::current_dir() - .map_err(|e| GenerateError::WriteOutput { - path: ".".to_string(), - source: e, - })?; +pub fn run( + input_paths: Vec, + output_dir: Option, + explicit_root: Option, + flat: bool, +) -> Result<(), GenerateError> { + let current_dir = std::env::current_dir().map_err(|e| GenerateError::WriteOutput { + path: ".".to_string(), + source: e, + })?; // Load mlf.toml if available let project_root = crate::config::find_project_root(¤t_dir).ok(); - let config = project_root - .as_ref() - .and_then(|root| { - let config_path = root.join("mlf.toml"); - crate::config::MlfConfig::load(&config_path).ok() - }); + let config = project_root.as_ref().and_then(|root| { + let config_path = root.join("mlf.toml"); + crate::config::MlfConfig::load(&config_path).ok() + }); // Determine output directory let output_dir = if let Some(explicit) = output_dir { @@ -123,7 +124,10 @@ pub fn run(input_paths: Vec, output_dir: Option, explicit_root let source = match std::fs::read_to_string(&file_path) { Ok(s) => s, Err(source) => { - errors.push((file_path.display().to_string(), format!("Failed to read file: {}", source))); + errors.push(( + file_path.display().to_string(), + format!("Failed to read file: {}", source), + )); continue; } }; @@ -143,23 +147,33 @@ pub fn run(input_paths: Vec, output_dir: Option, explicit_root .ok() .map(|root| crate::config::get_mlf_cache_dir(&root)); - let mut workspace = match crate::workspace_ext::workspace_with_std_and_cache(mlf_cache_dir.as_deref()) { - Ok(ws) => ws, - Err(e) => { - errors.push((file_path.display().to_string(), format!("Failed to load workspace: {}", e))); - continue; - } - }; + let mut workspace = + match crate::workspace_ext::workspace_with_std_and_cache(mlf_cache_dir.as_deref()) { + Ok(ws) => ws, + Err(e) => { + errors.push(( + file_path.display().to_string(), + format!("Failed to load workspace: {}", e), + )); + continue; + } + }; // Add the module to the workspace if let Err(e) = workspace.add_module(namespace.clone(), lexicon.clone()) { - errors.push((file_path.display().to_string(), format!("Failed to add module: {:?}", e))); + errors.push(( + file_path.display().to_string(), + format!("Failed to add module: {:?}", e), + )); continue; } // Resolve types if let Err(e) = workspace.resolve() { - errors.push((file_path.display().to_string(), format!("Type resolution error: {:?}", e))); + errors.push(( + file_path.display().to_string(), + format!("Type resolution error: {:?}", e), + )); continue; } @@ -177,7 +191,10 @@ pub fn run(input_paths: Vec, output_dir: Option, explicit_root path.push(segment); } if let Err(source) = std::fs::create_dir_all(&path.parent().unwrap()) { - errors.push((file_path.display().to_string(), format!("Failed to create directory: {}", source))); + errors.push(( + file_path.display().to_string(), + format!("Failed to create directory: {}", source), + )); continue; } path.set_extension("json"); @@ -186,7 +203,10 @@ pub fn run(input_paths: Vec, output_dir: Option, explicit_root let json_str = serde_json::to_string_pretty(&output.json).unwrap(); if let Err(source) = std::fs::write(&output_path, format!("{}\n", json_str)) { - errors.push((output_path.display().to_string(), format!("Failed to write file: {}", source))); + errors.push(( + output_path.display().to_string(), + format!("Failed to write file: {}", source), + )); continue; } @@ -195,7 +215,11 @@ pub fn run(input_paths: Vec, output_dir: Option, explicit_root } if !errors.is_empty() { - eprintln!("\n{} file(s) generated successfully, {} error(s) encountered:\n", success_count, errors.len()); + eprintln!( + "\n{} file(s) generated successfully, {} error(s) encountered:\n", + success_count, + errors.len() + ); for (path, error) in &errors { eprintln!(" {} - {}", path, error); } @@ -248,26 +272,32 @@ fn collect_mlf_files(dir: &Path) -> Result, GenerateError> { /// e.g., root=/project/lexicons, file=/project/lexicons/com/example/foo.mlf -> com.example.foo fn extract_namespace(file_path: &Path, root_dir: &Path) -> Result { // Get the canonical paths to handle . and .. correctly - let file_canonical = file_path.canonicalize().map_err(|source| GenerateError::ReadFile { - path: file_path.display().to_string(), - source, - })?; + let file_canonical = file_path + .canonicalize() + .map_err(|source| GenerateError::ReadFile { + path: file_path.display().to_string(), + source, + })?; - let root_canonical = root_dir.canonicalize().map_err(|source| GenerateError::ReadFile { - path: root_dir.display().to_string(), - source, - })?; + let root_canonical = root_dir + .canonicalize() + .map_err(|source| GenerateError::ReadFile { + path: root_dir.display().to_string(), + source, + })?; // Get relative path from root to file - let relative_path = file_canonical.strip_prefix(&root_canonical) - .map_err(|_| GenerateError::ParseLexicon { - path: file_path.display().to_string(), - help: Some(format!( - "File {} is not within root directory {}", - file_path.display(), - root_dir.display() - )), - })?; + let relative_path = + file_canonical + .strip_prefix(&root_canonical) + .map_err(|_| GenerateError::ParseLexicon { + path: file_path.display().to_string(), + help: Some(format!( + "File {} is not within root directory {}", + file_path.display(), + root_dir.display() + )), + })?; // Convert path to namespace let mut components = Vec::new(); diff --git a/mlf-cli/src/generate/mlf.rs b/mlf-cli/src/generate/mlf.rs index cd62158..2e1647f 100644 --- a/mlf-cli/src/generate/mlf.rs +++ b/mlf-cli/src/generate/mlf.rs @@ -72,7 +72,11 @@ pub enum MlfGenerateError { }, } -pub fn run(input_patterns: Vec, output_dir: Option, flat: bool) -> Result<(), MlfGenerateError> { +pub fn run( + input_patterns: Vec, + output_dir: Option, + flat: bool, +) -> Result<(), MlfGenerateError> { let current_dir = std::env::current_dir().map_err(|source| MlfGenerateError::WriteOutput { path: "current directory".to_string(), source, @@ -80,12 +84,10 @@ pub fn run(input_patterns: Vec, output_dir: Option, flat: bool) // Load mlf.toml if available let project_root = crate::config::find_project_root(¤t_dir).ok(); - let config = project_root - .as_ref() - .and_then(|root| { - let config_path = root.join("mlf.toml"); - crate::config::MlfConfig::load(&config_path).ok() - }); + let config = project_root.as_ref().and_then(|root| { + let config_path = root.join("mlf.toml"); + crate::config::MlfConfig::load(&config_path).ok() + }); // Determine output directory let output_dir = if let Some(explicit) = output_dir { @@ -160,20 +162,16 @@ pub fn run(input_patterns: Vec, output_dir: Option, flat: bool) } }; for warning in &output.warnings { - eprintln!( - "warning ({}): {}", - warning.namespace, warning.message - ); + eprintln!("warning ({}): {}", warning.namespace, warning.message); } let mlf_content = output.mlf; // Extract namespace from JSON "id" field - let namespace = json - .get("id") - .and_then(|v| v.as_str()) - .ok_or_else(|| MlfGenerateError::InvalidLexicon { + let namespace = json.get("id").and_then(|v| v.as_str()).ok_or_else(|| { + MlfGenerateError::InvalidLexicon { message: "Missing 'id' field in lexicon".to_string(), - })?; + } + })?; let output_path = if flat { output_dir.join(format!("{}.mlf", namespace)) @@ -185,8 +183,8 @@ pub fn run(input_patterns: Vec, output_dir: Option, flat: bool) } if let Err(source) = std::fs::create_dir_all(&path.parent().unwrap()) { errors.push(( - file_path.display().to_string(), - format!("Failed to create directory: {}", source), + file_path.display().to_string(), + format!("Failed to create directory: {}", source), )); continue; } @@ -228,20 +226,20 @@ pub fn run(input_patterns: Vec, output_dir: Option, flat: bool) pub fn generate_mlf_from_json(json: &Value) -> Result { let mut output = String::new(); - let nsid = json - .get("id") - .and_then(|v| v.as_str()) - .ok_or_else(|| MlfGenerateError::InvalidLexicon { + let nsid = json.get("id").and_then(|v| v.as_str()).ok_or_else(|| { + MlfGenerateError::InvalidLexicon { message: "Missing 'id' field in lexicon".to_string(), - })?; + } + })?; let last_segment = nsid.split('.').last().unwrap_or("main"); - let defs = json.get("defs").and_then(|v| v.as_object()).ok_or_else(|| { - MlfGenerateError::InvalidLexicon { + let defs = json + .get("defs") + .and_then(|v| v.as_object()) + .ok_or_else(|| MlfGenerateError::InvalidLexicon { message: "Missing or invalid 'defs' field".to_string(), - } - })?; + })?; let ctx = ConversionContext { current_namespace: nsid.to_string(), @@ -295,9 +293,23 @@ const TOP_LEVEL_SPEC_FIELDS: &[&str] = &["lexicon", "id", "description", "defs", /// else (e.g. `permission-set`) falls through to the unknown-def /// passthrough. const KNOWN_DEF_TYPES: &[&str] = &[ - "record", "query", "procedure", "subscription", "token", - "object", "string", "integer", "boolean", "bytes", "blob", - "null", "unknown", "array", "union", "ref", "cid-link", + "record", + "query", + "procedure", + "subscription", + "token", + "object", + "string", + "integer", + "boolean", + "bytes", + "blob", + "null", + "unknown", + "array", + "union", + "ref", + "cid-link", ]; fn is_known_def_type(type_name: &str) -> bool { @@ -309,15 +321,17 @@ fn is_known_def_type(type_name: &str) -> bool { /// vendor extensions (`revision`, `x-*` flags, etc.) roundtrip /// byte-faithfully. const RECORD_SPEC_FIELDS: &[&str] = &["type", "description", "key", "record"]; -const QUERY_SPEC_FIELDS: &[&str] = &[ - "type", "description", "parameters", "output", "errors", -]; +const QUERY_SPEC_FIELDS: &[&str] = &["type", "description", "parameters", "output", "errors"]; const PROCEDURE_SPEC_FIELDS: &[&str] = &[ - "type", "description", "parameters", "input", "output", "errors", -]; -const SUBSCRIPTION_SPEC_FIELDS: &[&str] = &[ - "type", "description", "parameters", "message", "errors", + "type", + "description", + "parameters", + "input", + "output", + "errors", ]; +const SUBSCRIPTION_SPEC_FIELDS: &[&str] = + &["type", "description", "parameters", "message", "errors"]; const TOKEN_SPEC_FIELDS: &[&str] = &["type", "description"]; /// Spec-defined fields at the top of a def-type definition. Covers @@ -325,14 +339,30 @@ const TOKEN_SPEC_FIELDS: &[&str] = &["type", "description"]; /// union, ref) and unifies them all — anything outside this list on a /// def-type JSON object is treated as an extension. const DEF_TYPE_SPEC_FIELDS: &[&str] = &[ - "type", "description", + "type", + "description", // Constraint keys (mirror CONSTRAINT_KEYS). - "minLength", "maxLength", "minGraphemes", "maxGraphemes", - "minimum", "maximum", "format", "enum", "knownValues", - "accept", "maxSize", "default", "const", + "minLength", + "maxLength", + "minGraphemes", + "maxGraphemes", + "minimum", + "maximum", + "format", + "enum", + "knownValues", + "accept", + "maxSize", + "default", + "const", // Container keys. - "items", "properties", "required", "nullable", - "refs", "closed", "ref", + "items", + "properties", + "required", + "nullable", + "refs", + "closed", + "ref", ]; /// Build a `self {}` item from the top-level JSON, or `None` when @@ -342,7 +372,10 @@ const DEF_TYPE_SPEC_FIELDS: &[&str] = &[ fn render_self_item(json: &Value, ctx: &ConversionContext) -> Option { let obj = json.as_object()?; - let description = obj.get("description").and_then(|v| v.as_str()).unwrap_or(""); + let description = obj + .get("description") + .and_then(|v| v.as_str()) + .unwrap_or(""); let has_extension = obj .keys() .any(|k| !TOP_LEVEL_SPEC_FIELDS.contains(&k.as_str())); @@ -545,9 +578,32 @@ impl ConversionContext { /// Reserved words in MLF that need to be escaped const RESERVED_WORDS: &[&str] = &[ - "main", "record", "query", "procedure", "subscription", "token", "def", "type", "use", - "pub", "alias", "namespace", "constrained", "error", "unit", "null", "boolean", - "integer", "string", "bytes", "blob", "unknown", "array", "object", "union", "ref", + "main", + "record", + "query", + "procedure", + "subscription", + "token", + "def", + "type", + "use", + "pub", + "alias", + "namespace", + "constrained", + "error", + "unit", + "null", + "boolean", + "integer", + "string", + "bytes", + "blob", + "unknown", + "array", + "object", + "union", + "ref", ]; /// Escape a name if it's a reserved word @@ -559,7 +615,11 @@ fn escape_name(name: &str) -> String { } } -fn generate_record(name: &str, def: &Value, ctx: &ConversionContext) -> Result { +fn generate_record( + name: &str, + def: &Value, + ctx: &ConversionContext, +) -> Result { let mut output = String::new(); // Add doc comment if present @@ -594,11 +654,12 @@ fn generate_record(name: &str, def: &Value, ctx: &ConversionContext) -> Result Result Result Result { +fn generate_query( + name: &str, + def: &Value, + ctx: &ConversionContext, +) -> Result { let mut output = String::new(); // Add doc comment @@ -667,11 +736,7 @@ fn generate_query(name: &str, def: &Value, ctx: &ConversionContext) -> Result>() - }) + .map(|arr| arr.iter().filter_map(|v| v.as_str()).collect::>()) .unwrap_or_default(); if let Some(props) = properties { @@ -692,7 +757,10 @@ fn generate_query(name: &str, def: &Value, ctx: &ConversionContext) -> Result Result Result { +fn generate_procedure( + name: &str, + def: &Value, + ctx: &ConversionContext, +) -> Result { let mut output = String::new(); // Add doc comment @@ -748,7 +820,11 @@ fn generate_procedure(name: &str, def: &Value, ctx: &ConversionContext) -> Resul output.push_str("@main\n"); } - output.push_str(&render_extension_annotations(def, PROCEDURE_SPEC_FIELDS, ctx)); + output.push_str(&render_extension_annotations( + def, + PROCEDURE_SPEC_FIELDS, + ctx, + )); let procedure_name = if name == "main" { escape_name(&ctx.local_main_name) @@ -765,11 +841,7 @@ fn generate_procedure(name: &str, def: &Value, ctx: &ConversionContext) -> Resul let required = schema .get("required") .and_then(|v| v.as_array()) - .map(|arr| { - arr.iter() - .filter_map(|v| v.as_str()) - .collect::>() - }) + .map(|arr| arr.iter().filter_map(|v| v.as_str()).collect::>()) .unwrap_or_default(); if let Some(props) = properties { @@ -833,7 +905,11 @@ fn generate_procedure(name: &str, def: &Value, ctx: &ConversionContext) -> Resul Ok(output) } -fn generate_subscription(name: &str, def: &Value, ctx: &ConversionContext) -> Result { +fn generate_subscription( + name: &str, + def: &Value, + ctx: &ConversionContext, +) -> Result { let mut output = String::new(); // Add doc comment @@ -850,7 +926,11 @@ fn generate_subscription(name: &str, def: &Value, ctx: &ConversionContext) -> Re output.push_str("@main\n"); } - output.push_str(&render_extension_annotations(def, SUBSCRIPTION_SPEC_FIELDS, ctx)); + output.push_str(&render_extension_annotations( + def, + SUBSCRIPTION_SPEC_FIELDS, + ctx, + )); let subscription_name = if name == "main" { escape_name(&ctx.local_main_name) @@ -866,11 +946,7 @@ fn generate_subscription(name: &str, def: &Value, ctx: &ConversionContext) -> Re let required = params .get("required") .and_then(|v| v.as_array()) - .map(|arr| { - arr.iter() - .filter_map(|v| v.as_str()) - .collect::>() - }) + .map(|arr| arr.iter().filter_map(|v| v.as_str()).collect::>()) .unwrap_or_default(); if let Some(props) = properties { @@ -930,7 +1006,11 @@ fn generate_token( Ok(output) } -fn generate_def_type(name: &str, def: &Value, ctx: &ConversionContext) -> Result { +fn generate_def_type( + name: &str, + def: &Value, + ctx: &ConversionContext, +) -> Result { let mut output = String::new(); // Add doc comment if present @@ -947,7 +1027,11 @@ fn generate_def_type(name: &str, def: &Value, ctx: &ConversionContext) -> Result output.push_str("@main\n"); } - output.push_str(&render_extension_annotations(def, DEF_TYPE_SPEC_FIELDS, ctx)); + output.push_str(&render_extension_annotations( + def, + DEF_TYPE_SPEC_FIELDS, + ctx, + )); // Use last segment of NSID for "main" definitions // Keywords are now allowed by the parser, so just escape with backticks @@ -1063,7 +1147,10 @@ struct Rendered { impl Rendered { fn atom(text: impl Into) -> Self { - Self { text: text.into(), shape: Shape::Atom } + Self { + text: text.into(), + shape: Shape::Atom, + } } /// Render `base` plus any constraints from `type_def`. If no constraints @@ -1137,8 +1224,16 @@ fn generate_type( // atomic. match type_name { Some("null") => Ok(Rendered::atom("null")), - Some("boolean") => Ok(Rendered::with_constraints("boolean", type_def, indent_level)), - Some("integer") => Ok(Rendered::with_constraints("integer", type_def, indent_level)), + Some("boolean") => Ok(Rendered::with_constraints( + "boolean", + type_def, + indent_level, + )), + Some("integer") => Ok(Rendered::with_constraints( + "integer", + type_def, + indent_level, + )), Some("string") => Ok(render_string(type_def, indent_level)), Some("bytes") => Ok(Rendered::with_constraints("bytes", type_def, indent_level)), Some("blob") => Ok(Rendered::with_constraints("blob", type_def, indent_level)), @@ -1202,9 +1297,11 @@ fn render_object_inline( ctx: &ConversionContext, indent_level: usize, ) -> Result { - let obj = type_def.as_object().ok_or_else(|| MlfGenerateError::InvalidLexicon { - message: "Object type definition is not a JSON object".to_string(), - })?; + let obj = type_def + .as_object() + .ok_or_else(|| MlfGenerateError::InvalidLexicon { + message: "Object type definition is not a JSON object".to_string(), + })?; // The spec lists `properties` as required on object types, but real- // world lexicons (e.g. blog.pckt.richtext.facet marker defs) publish @@ -1214,9 +1311,11 @@ fn render_object_inline( // lexicon isn't strictly spec-compliant. let empty_map = serde_json::Map::new(); let properties = match obj.get("properties") { - Some(v) => v.as_object().ok_or_else(|| MlfGenerateError::InvalidLexicon { - message: "`properties` in object type must be a JSON object".to_string(), - })?, + Some(v) => v + .as_object() + .ok_or_else(|| MlfGenerateError::InvalidLexicon { + message: "`properties` in object type must be a JSON object".to_string(), + })?, None => { ctx.warn( "object type is missing `properties` field; \ @@ -1243,7 +1342,11 @@ fn render_object_inline( } } } - let marker = if required.contains(&field_name.as_str()) { "!" } else { "" }; + let marker = if required.contains(&field_name.as_str()) { + "!" + } else { + "" + }; let field_type = render_field_type( field_def, ctx, @@ -1301,7 +1404,10 @@ fn render_union( message: "Missing 'refs' in union type".to_string(), })?; - let closed = type_def.get("closed").and_then(|v| v.as_bool()).unwrap_or(false); + let closed = type_def + .get("closed") + .and_then(|v| v.as_bool()) + .unwrap_or(false); // An open union with zero refs is malformed per the ATProto spec — it // names no valid types at all. Real-world lexicons (e.g. @@ -1334,18 +1440,21 @@ fn render_union( // A single-member open union renders as just that member — still an atom. // Anything with a visible `|` becomes Union for postfix purposes. - let shape = if parts.len() >= 2 || closed { Shape::Union } else { Shape::Atom }; + let shape = if parts.len() >= 2 || closed { + Shape::Union + } else { + Shape::Atom + }; Ok(Rendered { text, shape }) } fn render_ref(type_def: &Value, ctx: &ConversionContext) -> Result { - let ref_str = - type_def - .get("ref") - .and_then(|v| v.as_str()) - .ok_or_else(|| MlfGenerateError::InvalidLexicon { - message: "Missing 'ref' in ref type".to_string(), - })?; + let ref_str = type_def + .get("ref") + .and_then(|v| v.as_str()) + .ok_or_else(|| MlfGenerateError::InvalidLexicon { + message: "Missing 'ref' in ref type".to_string(), + })?; Ok(resolve_ref_string(ref_str, ctx)) } diff --git a/mlf-cli/src/generate/mod.rs b/mlf-cli/src/generate/mod.rs index 5238e54..64c3e25 100644 --- a/mlf-cli/src/generate/mod.rs +++ b/mlf-cli/src/generate/mod.rs @@ -1,4 +1,4 @@ -use crate::config::{find_project_root, ConfigError, MlfConfig}; +use crate::config::{ConfigError, MlfConfig, find_project_root}; use std::path::PathBuf; pub mod code; @@ -9,20 +9,24 @@ pub mod mlf; pub fn run_all() -> Result<(), std::io::Error> { let current_dir = std::env::current_dir()?; - let project_root = find_project_root(¤t_dir) - .map_err(|e| match e { - ConfigError::NotFound => { - std::io::Error::new( - std::io::ErrorKind::NotFound, - "No mlf.toml found. Please create a configuration file or provide explicit arguments." - ) - } - _ => std::io::Error::new(std::io::ErrorKind::Other, format!("Failed to load config: {}", e)), - })?; + let project_root = find_project_root(¤t_dir).map_err(|e| match e { + ConfigError::NotFound => std::io::Error::new( + std::io::ErrorKind::NotFound, + "No mlf.toml found. Please create a configuration file or provide explicit arguments.", + ), + _ => std::io::Error::new( + std::io::ErrorKind::Other, + format!("Failed to load config: {}", e), + ), + })?; let config_path = project_root.join("mlf.toml"); - let config = MlfConfig::load(&config_path) - .map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, format!("Failed to load config: {}", e)))?; + let config = MlfConfig::load(&config_path).map_err(|e| { + std::io::Error::new( + std::io::ErrorKind::Other, + format!("Failed to load config: {}", e), + ) + })?; if config.output.is_empty() { println!("No output configurations found in mlf.toml"); @@ -42,13 +46,19 @@ pub fn run_all() -> Result<(), std::io::Error> { let output_type = &output_config.r#type; let output_dir = PathBuf::from(&output_config.directory); - println!("\nGenerating {} output to {}...", output_type, output_config.directory); + println!( + "\nGenerating {} output to {}...", + output_type, output_config.directory + ); let result = match output_type.as_str() { - "lexicon" => { - lexicon::run(input_paths.clone(), Some(output_dir), Some(source_dir.clone()), false) - .map_err(|e| format!("{}", e)) - } + "lexicon" => lexicon::run( + input_paths.clone(), + Some(output_dir), + Some(source_dir.clone()), + false, + ) + .map_err(|e| format!("{}", e)), "mlf" => { // For MLF output, we expect JSON lexicons as input // This is a bit different - we'd need JSON input patterns @@ -57,8 +67,14 @@ pub fn run_all() -> Result<(), std::io::Error> { } generator_type => { // Assume it's a code generator (typescript, go, rust, etc.) - code::run(Some(generator_type.to_string()), input_paths.clone(), Some(output_dir), Some(source_dir.clone()), false) - .map_err(|e| format!("{}", e)) + code::run( + Some(generator_type.to_string()), + input_paths.clone(), + Some(output_dir), + Some(source_dir.clone()), + false, + ) + .map_err(|e| format!("{}", e)) } }; @@ -85,7 +101,7 @@ pub fn run_all() -> Result<(), std::io::Error> { } return Err(std::io::Error::new( std::io::ErrorKind::Other, - format!("Failed to generate {} output(s)", errors.len()) + format!("Failed to generate {} output(s)", errors.len()), )); } diff --git a/mlf-cli/src/init.rs b/mlf-cli/src/init.rs index c622c84..235119c 100644 --- a/mlf-cli/src/init.rs +++ b/mlf-cli/src/init.rs @@ -1,4 +1,4 @@ -use crate::config::{init_mlf_cache, MlfConfig}; +use crate::config::{MlfConfig, init_mlf_cache}; use std::io::Write; pub fn run_init(skip_prompts: bool) -> Result<(), std::io::Error> { diff --git a/mlf-cli/src/main.rs b/mlf-cli/src/main.rs index 0122f93..31e2eb7 100644 --- a/mlf-cli/src/main.rs +++ b/mlf-cli/src/main.rs @@ -31,10 +31,15 @@ enum Commands { }, Check { - #[arg(help = "MLF lexicon file(s) or directory to validate. If omitted, checks source directory from mlf.toml")] + #[arg( + help = "MLF lexicon file(s) or directory to validate. If omitted, checks source directory from mlf.toml" + )] input: Vec, - #[arg(long, help = "Root directory for namespace calculation (defaults to mlf.toml source directory or current directory)")] + #[arg( + long, + help = "Root directory for namespace calculation (defaults to mlf.toml source directory or current directory)" + )] root: Option, }, @@ -52,13 +57,18 @@ enum Commands { }, Fetch { - #[arg(help = "Namespace to fetch (e.g., stream.place). If omitted, fetches all dependencies from mlf.toml")] + #[arg( + help = "Namespace to fetch (e.g., stream.place). If omitted, fetches all dependencies from mlf.toml" + )] nsid: Option, #[arg(long, help = "Add namespace to dependencies in mlf.toml")] save: bool, - #[arg(long, help = "Update dependencies to latest versions (ignores lockfile)")] + #[arg( + long, + help = "Update dependencies to latest versions (ignores lockfile)" + )] update: bool, #[arg(long, help = "Require lockfile and fail if dependencies need updating")] @@ -69,39 +79,73 @@ enum Commands { #[derive(Subcommand)] enum GenerateCommands { Lexicon { - #[arg(short, long, help = "Input MLF file(s) or directory. If omitted, uses source directory from mlf.toml")] + #[arg( + short, + long, + help = "Input MLF file(s) or directory. If omitted, uses source directory from mlf.toml" + )] input: Vec, - #[arg(short, long, help = "Output directory. If omitted, uses first lexicon output from mlf.toml")] + #[arg( + short, + long, + help = "Output directory. If omitted, uses first lexicon output from mlf.toml" + )] output: Option, - #[arg(long, help = "Root directory for namespace calculation (defaults to mlf.toml source directory or current directory)")] + #[arg( + long, + help = "Root directory for namespace calculation (defaults to mlf.toml source directory or current directory)" + )] root: Option, #[arg(long, help = "Use flat file structure (e.g., app.bsky.post.json)")] flat: bool, }, Code { - #[arg(short, long, help = "Generator to use (typescript, go, rust, etc.). If omitted, uses first code output from mlf.toml")] + #[arg( + short, + long, + help = "Generator to use (typescript, go, rust, etc.). If omitted, uses first code output from mlf.toml" + )] generator: Option, - #[arg(short, long, help = "Input MLF file(s) or directory. If omitted, uses source directory from mlf.toml")] + #[arg( + short, + long, + help = "Input MLF file(s) or directory. If omitted, uses source directory from mlf.toml" + )] input: Vec, - #[arg(short, long, help = "Output directory. If omitted, uses matching output from mlf.toml")] + #[arg( + short, + long, + help = "Output directory. If omitted, uses matching output from mlf.toml" + )] output: Option, - #[arg(long, help = "Root directory for namespace calculation (defaults to mlf.toml source directory or current directory)")] + #[arg( + long, + help = "Root directory for namespace calculation (defaults to mlf.toml source directory or current directory)" + )] root: Option, #[arg(long, help = "Use flat file structure (e.g., app.bsky.post.ts)")] flat: bool, }, Mlf { - #[arg(short, long, help = "Input JSON lexicon files (glob patterns supported)")] + #[arg( + short, + long, + help = "Input JSON lexicon files (glob patterns supported)" + )] input: Vec, - #[arg(short, long, help = "Output directory. If omitted, uses first mlf output from mlf.toml")] + #[arg( + short, + long, + help = "Output directory. If omitted, uses first mlf output from mlf.toml" + )] output: Option, #[arg(long, help = "Use flat file structure (e.g., app.bsky.post.json)")] @@ -114,33 +158,43 @@ async fn main() { let cli = Cli::parse(); let result: Result<(), miette::Report> = match cli.command { - Commands::Init { yes } => { - init::run_init(yes).into_diagnostic() - } - Commands::Check { input, root } => { - check::run_check(input, root).into_diagnostic() - } + Commands::Init { yes } => init::run_init(yes).into_diagnostic(), + Commands::Check { input, root } => check::run_check(input, root).into_diagnostic(), Commands::Validate { lexicon, record } => { check::validate(lexicon, record).into_diagnostic() } Commands::Generate { command } => match command { - Some(GenerateCommands::Lexicon { input, output, root, flat }) => { - generate::lexicon::run(input, output, root, flat).into_diagnostic() - } - Some(GenerateCommands::Code { generator, input, output, root, flat }) => { - generate::code::run(generator, input, output, root, flat).into_diagnostic() - } - Some(GenerateCommands::Mlf { input, output, flat }) => { - generate::mlf::run(input, output, flat).into_diagnostic() - } + Some(GenerateCommands::Lexicon { + input, + output, + root, + flat, + }) => generate::lexicon::run(input, output, root, flat).into_diagnostic(), + Some(GenerateCommands::Code { + generator, + input, + output, + root, + flat, + }) => generate::code::run(generator, input, output, root, flat).into_diagnostic(), + Some(GenerateCommands::Mlf { + input, + output, + flat, + }) => generate::mlf::run(input, output, flat).into_diagnostic(), None => { // Run all outputs from mlf.toml generate::run_all().into_diagnostic() } }, - Commands::Fetch { nsid, save, update, locked } => { - fetch::run_fetch(nsid, save, update, locked).await.into_diagnostic() - } + Commands::Fetch { + nsid, + save, + update, + locked, + } => fetch::run_fetch(nsid, save, update, locked) + .await + .into_diagnostic(), }; if let Err(e) = result { diff --git a/mlf-cli/src/workspace_ext.rs b/mlf-cli/src/workspace_ext.rs index c74ecef..c19f88e 100644 --- a/mlf-cli/src/workspace_ext.rs +++ b/mlf-cli/src/workspace_ext.rs @@ -77,9 +77,7 @@ fn extract_namespace_from_path(path: &Path, base: &Path) -> Result "place.stream.key" @@ -89,12 +87,10 @@ fn extract_namespace_from_path(path: &Path, base: &Path) -> Result, -) -> Result { +pub fn workspace_with_std_and_cache(mlf_cache_dir: Option<&Path>) -> Result { // Start with std library - let mut workspace = Workspace::with_std() - .map_err(|e| format!("Failed to load std library: {:?}", e))?; + let mut workspace = + Workspace::with_std().map_err(|e| format!("Failed to load std library: {:?}", e))?; // Load from .mlf cache if provided if let Some(cache_dir) = mlf_cache_dir { diff --git a/mlf-codegen/examples/all_generators.rs b/mlf-codegen/examples/all_generators.rs index 9edbf1e..e345924 100644 --- a/mlf-codegen/examples/all_generators.rs +++ b/mlf-codegen/examples/all_generators.rs @@ -5,9 +5,9 @@ use mlf_codegen::plugin; // Import all the plugin crates to trigger their registration -extern crate mlf_codegen_typescript; extern crate mlf_codegen_go; extern crate mlf_codegen_rust; +extern crate mlf_codegen_typescript; fn main() { println!("MLF Code Generator Plugins (All Loaded)\n"); @@ -17,10 +17,7 @@ fn main() { println!("Found {} generator(s):\n", generators.len()); for generator in generators { - println!(" {} ({}):", - generator.name(), - generator.file_extension() - ); + println!(" {} ({}):", generator.name(), generator.file_extension()); println!(" {}\n", generator.description()); } } diff --git a/mlf-codegen/examples/list_generators.rs b/mlf-codegen/examples/list_generators.rs index 6b63d0b..59cd01a 100644 --- a/mlf-codegen/examples/list_generators.rs +++ b/mlf-codegen/examples/list_generators.rs @@ -14,17 +14,16 @@ fn main() { if generators.is_empty() { println!("No generators registered!"); println!("\nTo use plugins, depend on them in your Cargo.toml:"); - println!(" mlf-codegen-typescript = {{ path = \"../codegen-plugins/mlf-codegen-typescript\" }}"); + println!( + " mlf-codegen-typescript = {{ path = \"../codegen-plugins/mlf-codegen-typescript\" }}" + ); println!(" mlf-codegen-go = {{ path = \"../codegen-plugins/mlf-codegen-go\" }}"); println!(" mlf-codegen-rust = {{ path = \"../codegen-plugins/mlf-codegen-rust\" }}"); } else { println!("Found {} generator(s):\n", generators.len()); for generator in generators { - println!(" {} ({}):", - generator.name(), - generator.file_extension() - ); + println!(" {} ({}):", generator.name(), generator.file_extension()); println!(" {}\n", generator.description()); } } diff --git a/mlf-codegen/examples/plugin_test.rs b/mlf-codegen/examples/plugin_test.rs index de1669f..6597a25 100644 --- a/mlf-codegen/examples/plugin_test.rs +++ b/mlf-codegen/examples/plugin_test.rs @@ -6,7 +6,8 @@ fn main() { // List all registered generators println!("Registered generators:"); for generator in plugin::generators() { - println!(" - {} ({}): {}", + println!( + " - {} ({}): {}", generator.name(), generator.file_extension(), generator.description() diff --git a/mlf-codegen/src/lib.rs b/mlf-codegen/src/lib.rs index f967bda..77d2c5a 100644 --- a/mlf-codegen/src/lib.rs +++ b/mlf-codegen/src/lib.rs @@ -1,6 +1,6 @@ use mlf_lang::ast::*; use mlf_lang::{ResolvedRef, Workspace}; -use serde_json::{json, Map, Value}; +use serde_json::{Map, Value, json}; use std::collections::HashMap; /// A non-fatal advisory emitted during codegen. Mirrors the shape of @@ -92,21 +92,21 @@ fn has_main_annotation(annotations: &[Annotation]) -> bool { } fn get_annotation_string_value(annotations: &[Annotation], name: &str) -> Option { - annotations.iter() + annotations + .iter() .find(|ann| ann.name.name == name) .and_then(|ann| { // Get first positional argument if it exists - ann.args.first().and_then(|arg| { - match arg { - AnnotationArg::Positional(AnnotationValue::String(s)) => Some(s.clone()), - _ => None, - } + ann.args.first().and_then(|arg| match arg { + AnnotationArg::Positional(AnnotationValue::String(s)) => Some(s.clone()), + _ => None, }) }) } fn get_encoding_annotation(annotations: &[Annotation], param_name: &str) -> Option { - annotations.iter() + annotations + .iter() .find(|ann| ann.name.name == "encoding") .and_then(|ann| { // First check for named argument matching param_name @@ -131,7 +131,11 @@ fn get_encoding_annotation(annotations: &[Annotation], param_name: &str) -> Opti }) } -pub fn generate_lexicon(namespace: &str, lexicon: &Lexicon, workspace: &Workspace) -> CodegenOutput { +pub fn generate_lexicon( + namespace: &str, + lexicon: &Lexicon, + workspace: &Workspace, +) -> CodegenOutput { let usage_counts = analyze_type_usage(lexicon); let eligibility = MainEligibility::for_lexicon(namespace, lexicon); @@ -144,35 +148,99 @@ pub fn generate_lexicon(namespace: &str, lexicon: &Lexicon, workspace: &Workspac match item { Item::Record(record) => { let mut value = generate_record_json(record, &usage_counts, workspace, namespace); - apply_extension_annotations(&mut value, &record.annotations, workspace, namespace, &mut warnings); - insert_def(&mut defs, &record.name.name, eligibility.is_main(&record.name.name, &record.annotations), value); + apply_extension_annotations( + &mut value, + &record.annotations, + workspace, + namespace, + &mut warnings, + ); + insert_def( + &mut defs, + &record.name.name, + eligibility.is_main(&record.name.name, &record.annotations), + value, + ); } Item::Query(query) => { let mut value = generate_query_json(query, &usage_counts, workspace, namespace); - apply_extension_annotations(&mut value, &query.annotations, workspace, namespace, &mut warnings); - insert_def(&mut defs, &query.name.name, eligibility.is_main(&query.name.name, &query.annotations), value); + apply_extension_annotations( + &mut value, + &query.annotations, + workspace, + namespace, + &mut warnings, + ); + insert_def( + &mut defs, + &query.name.name, + eligibility.is_main(&query.name.name, &query.annotations), + value, + ); } Item::Procedure(procedure) => { - let mut value = generate_procedure_json(procedure, &usage_counts, workspace, namespace); - apply_extension_annotations(&mut value, &procedure.annotations, workspace, namespace, &mut warnings); - insert_def(&mut defs, &procedure.name.name, eligibility.is_main(&procedure.name.name, &procedure.annotations), value); + let mut value = + generate_procedure_json(procedure, &usage_counts, workspace, namespace); + apply_extension_annotations( + &mut value, + &procedure.annotations, + workspace, + namespace, + &mut warnings, + ); + insert_def( + &mut defs, + &procedure.name.name, + eligibility.is_main(&procedure.name.name, &procedure.annotations), + value, + ); } Item::Subscription(subscription) => { - let mut value = generate_subscription_json(subscription, &usage_counts, workspace, namespace); - apply_extension_annotations(&mut value, &subscription.annotations, workspace, namespace, &mut warnings); - insert_def(&mut defs, &subscription.name.name, eligibility.is_main(&subscription.name.name, &subscription.annotations), value); + let mut value = + generate_subscription_json(subscription, &usage_counts, workspace, namespace); + apply_extension_annotations( + &mut value, + &subscription.annotations, + workspace, + namespace, + &mut warnings, + ); + insert_def( + &mut defs, + &subscription.name.name, + eligibility.is_main(&subscription.name.name, &subscription.annotations), + value, + ); } Item::DefType(def_type) => { - let mut value = generate_def_type_json(def_type, &usage_counts, workspace, namespace); - apply_extension_annotations(&mut value, &def_type.annotations, workspace, namespace, &mut warnings); - insert_def(&mut defs, &def_type.name.name, eligibility.is_main(&def_type.name.name, &def_type.annotations), value); + let mut value = + generate_def_type_json(def_type, &usage_counts, workspace, namespace); + apply_extension_annotations( + &mut value, + &def_type.annotations, + workspace, + namespace, + &mut warnings, + ); + insert_def( + &mut defs, + &def_type.name.name, + eligibility.is_main(&def_type.name.name, &def_type.annotations), + value, + ); } Item::Token(token) => { let mut token_obj = Map::new(); token_obj.insert("type".to_string(), json!("token")); insert_opt_str(&mut token_obj, "description", &extract_docs(&token.docs)); let mut value = Value::Object(token_obj); - apply_extension_annotations(&mut value, &token.annotations, workspace, namespace, &mut warnings); + apply_extension_annotations( + &mut value, + &token.annotations, + workspace, + namespace, + &mut warnings, + ); defs.insert(token.name.name.clone(), value); } Item::SelfItem(self_item) => { @@ -181,7 +249,12 @@ pub fn generate_lexicon(namespace: &str, lexicon: &Lexicon, workspace: &Workspac // annotations become top-level JSON fields alongside // `lexicon`, `id`, `defs`. self_description = extract_docs(&self_item.docs); - self_extensions = collect_extension_fields(&self_item.annotations, workspace, namespace, &mut warnings); + self_extensions = collect_extension_fields( + &self_item.annotations, + workspace, + namespace, + &mut warnings, + ); } // Inline types never appear in `defs` — they expand at their point // of use. Use statements are structural and not emitted. @@ -233,7 +306,9 @@ fn apply_extension_annotations( let Some(obj) = value.as_object_mut() else { return; }; - for (key, field_value) in collect_extension_fields(annotations, workspace, current_namespace, warnings) { + for (key, field_value) in + collect_extension_fields(annotations, workspace, current_namespace, warnings) + { obj.insert(key, field_value); } } @@ -308,7 +383,9 @@ fn annotation_value_to_json( message: format!( "@const({:?}, {}): ATProto's data model has no floats; \ emitting as string {:?} to stay spec-compliant", - key, n, n.to_string() + key, + n, + n.to_string() ), }); Value::String(n.to_string()) @@ -326,14 +403,12 @@ fn annotation_value_to_json( } AnnotationValue::Boolean(b) => Value::Bool(*b), AnnotationValue::Null => Value::Null, - AnnotationValue::Array(items) => { - Value::Array( - items - .iter() - .map(|item| annotation_value_to_json(item, key, namespace, warnings)) - .collect(), - ) - } + AnnotationValue::Array(items) => Value::Array( + items + .iter() + .map(|item| annotation_value_to_json(item, key, namespace, warnings)) + .collect(), + ), AnnotationValue::Object(entries) => { let mut obj = Map::new(); for (entry_key, entry_value) in entries { @@ -424,7 +499,11 @@ fn item_header(item: &Item) -> Option<(&str, &[Annotation])> { /// Insert a def into the lexicon's `defs` map under the canonical key — /// `"main"` for the main def, otherwise the def's own name. fn insert_def(defs: &mut Map, name: &str, is_main: bool, value: Value) { - let key = if is_main { "main".to_string() } else { name.to_string() }; + let key = if is_main { + "main".to_string() + } else { + name.to_string() + }; defs.insert(key, value); } @@ -563,7 +642,8 @@ fn build_param_properties( if !param.optional { required.push(param.name.name.clone()); } - let mut param_json = generate_type_json(¶m.ty, usage_counts, workspace, current_namespace); + let mut param_json = + generate_type_json(¶m.ty, usage_counts, workspace, current_namespace); add_description_from_docs(&mut param_json, ¶m.docs); properties.insert(param.name.name.clone(), param_json); } @@ -598,7 +678,8 @@ fn collect_object_fields( if is_nullable { nullable.push(field.name.name.clone()); } - let mut field_json = generate_type_json(&effective_ty, usage_counts, workspace, current_namespace); + let mut field_json = + generate_type_json(&effective_ty, usage_counts, workspace, current_namespace); add_description_from_docs(&mut field_json, &field.docs); properties.insert(field.name.name.clone(), field_json); } @@ -617,7 +698,12 @@ fn strip_nullable(ty: &Type) -> Option { if let Type::Parenthesized { inner, .. } = ty { return strip_nullable(inner); } - let Type::Union { types, closed, span } = ty else { + let Type::Union { + types, + closed, + span, + } = ty + else { return None; }; let has_null = types.iter().any(is_null_primitive); @@ -643,7 +729,10 @@ fn strip_nullable(ty: &Type) -> Option { fn is_null_primitive(ty: &Type) -> bool { match ty { - Type::Primitive { kind: PrimitiveType::Null, .. } => true, + Type::Primitive { + kind: PrimitiveType::Null, + .. + } => true, Type::Parenthesized { inner, .. } => is_null_primitive(inner), _ => false, } @@ -684,7 +773,10 @@ fn resolve_ref_nsid(path: &Path, workspace: &Workspace, current_namespace: &str) match workspace.resolve_ref(path, current_namespace) { Some(ResolvedRef::Local { def_name }) => format!("#{}", def_name), Some(ResolvedRef::ImplicitMain { namespace, .. }) => namespace, - Some(ResolvedRef::External { namespace, def_name }) => format!("{}#{}", namespace, def_name), + Some(ResolvedRef::External { + namespace, + def_name, + }) => format!("{}#{}", namespace, def_name), None => unresolved_ref_fallback(path), } } @@ -706,7 +798,12 @@ fn unresolved_ref_fallback(path: &Path) -> String { format!("{}#{}", namespace, def_name) } -fn generate_record_json(record: &Record, usage_counts: &HashMap, workspace: &Workspace, current_namespace: &str) -> Value { +fn generate_record_json( + record: &Record, + usage_counts: &HashMap, + workspace: &Workspace, + current_namespace: &str, +) -> Value { let (properties, required, nullable) = collect_object_fields(&record.fields, usage_counts, workspace, current_namespace); @@ -717,7 +814,8 @@ fn generate_record_json(record: &Record, usage_counts: &HashMap, record_obj.insert("properties".to_string(), Value::Object(properties)); // Check for @key annotation, default to "tid" - let key = get_annotation_string_value(&record.annotations, "key").unwrap_or_else(|| "tid".to_string()); + let key = get_annotation_string_value(&record.annotations, "key") + .unwrap_or_else(|| "tid".to_string()); let mut record_top = Map::new(); record_top.insert("type".to_string(), json!("record")); @@ -727,7 +825,12 @@ fn generate_record_json(record: &Record, usage_counts: &HashMap, Value::Object(record_top) } -fn generate_query_json(query: &Query, usage_counts: &HashMap, workspace: &Workspace, current_namespace: &str) -> Value { +fn generate_query_json( + query: &Query, + usage_counts: &HashMap, + workspace: &Workspace, + current_namespace: &str, +) -> Value { let (params_properties, params_required) = build_param_properties(&query.params, usage_counts, workspace, current_namespace); let params = build_params_object(params_properties, ¶ms_required); @@ -741,10 +844,15 @@ fn generate_query_json(query: &Query, usage_counts: &HashMap, wor ReturnType::Type(ty) => { let mut output_obj = Map::new(); output_obj.insert("encoding".to_string(), json!(output_encoding)); - output_obj.insert("schema".to_string(), generate_type_json(ty, usage_counts, workspace, current_namespace)); + output_obj.insert( + "schema".to_string(), + generate_type_json(ty, usage_counts, workspace, current_namespace), + ); (Some(Value::Object(output_obj)), None) } - ReturnType::TypeWithErrors { success, errors, .. } => { + ReturnType::TypeWithErrors { + success, errors, .. + } => { let mut error_array = Vec::new(); for error in errors { let error_docs = extract_docs(&error.docs); @@ -761,8 +869,14 @@ fn generate_query_json(query: &Query, usage_counts: &HashMap, wor let mut output_obj = Map::new(); output_obj.insert("encoding".to_string(), json!(output_encoding)); - output_obj.insert("schema".to_string(), generate_type_json(success, usage_counts, workspace, current_namespace)); - (Some(Value::Object(output_obj)), Some(Value::Array(error_array))) + output_obj.insert( + "schema".to_string(), + generate_type_json(success, usage_counts, workspace, current_namespace), + ); + ( + Some(Value::Object(output_obj)), + Some(Value::Array(error_array)), + ) } }; @@ -779,9 +893,18 @@ fn generate_query_json(query: &Query, usage_counts: &HashMap, wor Value::Object(query_obj) } -fn generate_procedure_json(procedure: &Procedure, usage_counts: &HashMap, workspace: &Workspace, current_namespace: &str) -> Value { - let (params_properties, params_required) = - build_param_properties(&procedure.params, usage_counts, workspace, current_namespace); +fn generate_procedure_json( + procedure: &Procedure, + usage_counts: &HashMap, + workspace: &Workspace, + current_namespace: &str, +) -> Value { + let (params_properties, params_required) = build_param_properties( + &procedure.params, + usage_counts, + workspace, + current_namespace, + ); // Check for @encoding annotation with "input" parameter, default to "application/json" let input_encoding = get_encoding_annotation(&procedure.annotations, "input") @@ -810,10 +933,15 @@ fn generate_procedure_json(procedure: &Procedure, usage_counts: &HashMap { let mut output_obj = Map::new(); output_obj.insert("encoding".to_string(), json!(output_encoding)); - output_obj.insert("schema".to_string(), generate_type_json(ty, usage_counts, workspace, current_namespace)); + output_obj.insert( + "schema".to_string(), + generate_type_json(ty, usage_counts, workspace, current_namespace), + ); (Some(Value::Object(output_obj)), None) } - ReturnType::TypeWithErrors { success, errors, .. } => { + ReturnType::TypeWithErrors { + success, errors, .. + } => { let mut error_array = Vec::new(); for error in errors { let error_docs = extract_docs(&error.docs); @@ -830,8 +958,14 @@ fn generate_procedure_json(procedure: &Procedure, usage_counts: &HashMap Value { - let (params_properties, params_required) = - build_param_properties(&subscription.params, usage_counts, workspace, current_namespace); + let (params_properties, params_required) = build_param_properties( + &subscription.params, + usage_counts, + workspace, + current_namespace, + ); let mut result = Map::new(); result.insert("type".to_string(), json!("subscription")); - insert_opt_str(&mut result, "description", &extract_docs(&subscription.docs)); + insert_opt_str( + &mut result, + "description", + &extract_docs(&subscription.docs), + ); let mut result = Value::Object(result); if let Some(messages) = &subscription.messages { @@ -878,8 +1020,14 @@ fn generate_subscription_json( result } -fn generate_def_type_json(def_type: &DefType, usage_counts: &HashMap, workspace: &Workspace, current_namespace: &str) -> Value { - let mut field_json = generate_type_json(&def_type.ty, usage_counts, workspace, current_namespace); +fn generate_def_type_json( + def_type: &DefType, + usage_counts: &HashMap, + workspace: &Workspace, + current_namespace: &str, +) -> Value { + let mut field_json = + generate_type_json(&def_type.ty, usage_counts, workspace, current_namespace); // def types insert `description` at position 1 (right after `type`) to // match the canonical field order in published lexicons. let text = extract_docs(&def_type.docs); @@ -891,14 +1039,24 @@ fn generate_def_type_json(def_type: &DefType, usage_counts: &HashMap, workspace: &Workspace, current_namespace: &str) -> Value { +fn generate_type_json( + ty: &Type, + usage_counts: &HashMap, + workspace: &Workspace, + current_namespace: &str, +) -> Value { match ty { Type::Primitive { kind, .. } => generate_primitive_json(*kind), Type::Reference { path, .. } => { // Inline types expand in place; other references resolve to an NSID. if workspace.is_inline_type(path) { if let Some(resolved_ty) = workspace.resolve_type_reference(path) { - return generate_type_json(&resolved_ty, usage_counts, workspace, current_namespace); + return generate_type_json( + &resolved_ty, + usage_counts, + workspace, + current_namespace, + ); } } let nsid = resolve_ref_nsid(path, workspace, current_namespace); @@ -951,8 +1109,11 @@ fn generate_type_json(ty: &Type, usage_counts: &HashMap, workspac // Parentheses are just for grouping - unwrap and process inner type generate_type_json(inner, usage_counts, workspace, current_namespace) } - Type::Constrained { base, constraints, .. } => { - let mut base_json = generate_type_json(base, usage_counts, workspace, current_namespace); + Type::Constrained { + base, constraints, .. + } => { + let mut base_json = + generate_type_json(base, usage_counts, workspace, current_namespace); if let Some(obj) = base_json.as_object_mut() { for constraint in constraints { diff --git a/mlf-diagnostics/src/lib.rs b/mlf-diagnostics/src/lib.rs index d22d6b6..a60f4c8 100644 --- a/mlf-diagnostics/src/lib.rs +++ b/mlf-diagnostics/src/lib.rs @@ -49,9 +49,10 @@ impl Diagnostic for ParseDiagnostic { ParseError::InvalidIdentifier { span, .. } => *span, }; - Some(Box::new(std::iter::once( - LabeledSpan::at(span.start..span.end, "here"), - ))) + Some(Box::new(std::iter::once(LabeledSpan::at( + span.start..span.end, + "here", + )))) } fn help<'a>(&'a self) -> Option> { @@ -75,7 +76,12 @@ pub struct ValidationDiagnostic { } impl ValidationDiagnostic { - pub fn new(filename: String, source: String, module_namespace: String, errors: ValidationErrors) -> Self { + pub fn new( + filename: String, + source: String, + module_namespace: String, + errors: ValidationErrors, + ) -> Self { Self { source_code: NamedSource::new(filename, source), module_namespace, @@ -199,7 +205,11 @@ fn format_validation_error(f: &mut fmt::Formatter<'_>, error: &ValidationError) ValidationError::ReservedName { name, .. } => { write!(f, "Reserved name '{}' cannot be used as an item name", name) } - ValidationError::AmbiguousMain { name, namespace_suffix, .. } => { + ValidationError::AmbiguousMain { + name, + namespace_suffix, + .. + } => { write!( f, "Ambiguous main definition for '{}' in namespace ending with '{}'. Use @main to disambiguate", @@ -207,9 +217,17 @@ fn format_validation_error(f: &mut fmt::Formatter<'_>, error: &ValidationError) ) } ValidationError::MultipleMain { name, .. } => { - write!(f, "Multiple items named '{}' marked with @main. Only one can be @main", name) + write!( + f, + "Multiple items named '{}' marked with @main. Only one can be @main", + name + ) } - ValidationError::ConflictNotAllowed { name, namespace_suffix, .. } => { + ValidationError::ConflictNotAllowed { + name, + namespace_suffix, + .. + } => { write!( f, "Name conflict for '{}' is not allowed. Conflicts are only allowed when the name matches the namespace suffix ('{}')", @@ -368,12 +386,16 @@ fn get_error_help(error: &ValidationError) -> Option> { "Union types must contain at least one type. Add a type to the union or use a different type.", )) } - ValidationError::ConstraintTooPermissive { message, .. } if message.contains("maxLength") => { + ValidationError::ConstraintTooPermissive { message, .. } + if message.contains("maxLength") => + { Some(Box::new( "When refining a constrained type, maxLength can only decrease (become more restrictive).", )) } - ValidationError::ConstraintTooPermissive { message, .. } if message.contains("minLength") => { + ValidationError::ConstraintTooPermissive { message, .. } + if message.contains("minLength") => + { Some(Box::new( "When refining a constrained type, minLength can only increase (become more restrictive).", )) @@ -388,12 +410,16 @@ fn get_error_help(error: &ValidationError) -> Option> { "When refining a constrained type, minimum can only increase (become more restrictive).", )) } - ValidationError::ConstraintTooPermissive { message, .. } if message.contains("maxGraphemes") => { + ValidationError::ConstraintTooPermissive { message, .. } + if message.contains("maxGraphemes") => + { Some(Box::new( "When refining a constrained type, maxGraphemes can only decrease (become more restrictive).", )) } - ValidationError::ConstraintTooPermissive { message, .. } if message.contains("minGraphemes") => { + ValidationError::ConstraintTooPermissive { message, .. } + if message.contains("minGraphemes") => + { Some(Box::new( "When refining a constrained type, minGraphemes can only increase (become more restrictive).", )) @@ -408,64 +434,72 @@ fn get_error_help(error: &ValidationError) -> Option> { "When refining a constrained type, enum values must be a subset of the base enum.", )) } - ValidationError::DuplicateDefinition { .. } => { - Some(Box::new( - "Each name can only be defined once in a module. Consider renaming one of the items or using @main annotation if they match the namespace suffix.", - )) - } - ValidationError::ReservedName { name, .. } if name == "main" => { - Some(Box::new( - "The name 'main' is reserved and cannot be used as an item name. Use @main annotation on an item instead.", - )) - } - ValidationError::ReservedName { name, .. } if name == "defs" => { - Some(Box::new( - "The name 'defs' is reserved for future use and cannot be used as an item name.", - )) - } - ValidationError::AmbiguousMain { .. } => { - Some(Box::new( - "When multiple items have the same name as the namespace suffix, use @main to mark which one is the primary definition.", - )) - } - ValidationError::MultipleMain { .. } => { - Some(Box::new( - "Only one item can be marked with @main. Remove the @main annotation from all but one item.", - )) - } - ValidationError::ConflictNotAllowed { namespace_suffix, .. } => { - Some(Box::new(format!( - "Name conflicts are only allowed when the item name matches the namespace suffix ('{}').", - namespace_suffix - ))) - } - ValidationError::CircularImport { .. } => { - Some(Box::new( - "Circular imports are not allowed. Reorganize your modules to break the cycle.", - )) - } - ValidationError::UnusedImport { .. } => { - Some(Box::new( - "This import is never used. Consider removing it to keep the code clean.", - )) - } + ValidationError::DuplicateDefinition { .. } => Some(Box::new( + "Each name can only be defined once in a module. Consider renaming one of the items or using @main annotation if they match the namespace suffix.", + )), + ValidationError::ReservedName { name, .. } if name == "main" => Some(Box::new( + "The name 'main' is reserved and cannot be used as an item name. Use @main annotation on an item instead.", + )), + ValidationError::ReservedName { name, .. } if name == "defs" => Some(Box::new( + "The name 'defs' is reserved for future use and cannot be used as an item name.", + )), + ValidationError::AmbiguousMain { .. } => Some(Box::new( + "When multiple items have the same name as the namespace suffix, use @main to mark which one is the primary definition.", + )), + ValidationError::MultipleMain { .. } => Some(Box::new( + "Only one item can be marked with @main. Remove the @main annotation from all but one item.", + )), + ValidationError::ConflictNotAllowed { + namespace_suffix, .. + } => Some(Box::new(format!( + "Name conflicts are only allowed when the item name matches the namespace suffix ('{}').", + namespace_suffix + ))), + ValidationError::CircularImport { .. } => Some(Box::new( + "Circular imports are not allowed. Reorganize your modules to break the cycle.", + )), + ValidationError::UnusedImport { .. } => Some(Box::new( + "This import is never used. Consider removing it to keep the code clean.", + )), _ => None, } } pub fn get_error_module_namespace_str(error: &ValidationError) -> &str { match error { - ValidationError::DuplicateDefinition { module_namespace, .. } => module_namespace, - ValidationError::UndefinedReference { module_namespace, .. } => module_namespace, - ValidationError::InvalidConstraint { module_namespace, .. } => module_namespace, - ValidationError::TypeMismatch { module_namespace, .. } => module_namespace, - ValidationError::ConstraintTooPermissive { module_namespace, .. } => module_namespace, - ValidationError::ReservedName { module_namespace, .. } => module_namespace, - ValidationError::AmbiguousMain { module_namespace, .. } => module_namespace, - ValidationError::MultipleMain { module_namespace, .. } => module_namespace, - ValidationError::ConflictNotAllowed { module_namespace, .. } => module_namespace, - ValidationError::CircularImport { module_namespace, .. } => module_namespace, - ValidationError::UnusedImport { module_namespace, .. } => module_namespace, + ValidationError::DuplicateDefinition { + module_namespace, .. + } => module_namespace, + ValidationError::UndefinedReference { + module_namespace, .. + } => module_namespace, + ValidationError::InvalidConstraint { + module_namespace, .. + } => module_namespace, + ValidationError::TypeMismatch { + module_namespace, .. + } => module_namespace, + ValidationError::ConstraintTooPermissive { + module_namespace, .. + } => module_namespace, + ValidationError::ReservedName { + module_namespace, .. + } => module_namespace, + ValidationError::AmbiguousMain { + module_namespace, .. + } => module_namespace, + ValidationError::MultipleMain { + module_namespace, .. + } => module_namespace, + ValidationError::ConflictNotAllowed { + module_namespace, .. + } => module_namespace, + ValidationError::CircularImport { + module_namespace, .. + } => module_namespace, + ValidationError::UnusedImport { + module_namespace, .. + } => module_namespace, } } diff --git a/mlf-lang/src/ast.rs b/mlf-lang/src/ast.rs index 3e6e2bb..2989f9b 100644 --- a/mlf-lang/src/ast.rs +++ b/mlf-lang/src/ast.rs @@ -281,7 +281,11 @@ pub enum Type { /// Array type Array { inner: Box, span: Span }, /// Union type - Union { types: Vec, closed: bool, span: Span }, + Union { + types: Vec, + closed: bool, + span: Span, + }, /// Object type (inline) Object { fields: Vec, span: Span }, /// Parenthesized type (for grouping, e.g., (A | B)[]) diff --git a/mlf-lang/src/error.rs b/mlf-lang/src/error.rs index 37c676e..ed6fbc8 100644 --- a/mlf-lang/src/error.rs +++ b/mlf-lang/src/error.rs @@ -34,17 +34,67 @@ impl ParseError { #[derive(Debug, Clone, PartialEq)] pub enum ValidationError { - DuplicateDefinition { name: String, first_span: Span, second_span: Span, module_namespace: String }, - UndefinedReference { name: String, span: Span, module_namespace: String }, - InvalidConstraint { message: String, span: Span, module_namespace: String }, - TypeMismatch { expected: String, found: String, span: Span, module_namespace: String }, - ConstraintTooPermissive { message: String, span: Span, module_namespace: String }, - ReservedName { name: String, span: Span, module_namespace: String }, - AmbiguousMain { name: String, namespace_suffix: String, first_span: Span, second_span: Span, module_namespace: String }, - MultipleMain { name: String, first_span: Span, second_span: Span, module_namespace: String }, - ConflictNotAllowed { name: String, namespace_suffix: String, span: Span, module_namespace: String }, - CircularImport { cycle: Vec, span: Span, module_namespace: String }, - UnusedImport { name: String, span: Span, module_namespace: String }, + DuplicateDefinition { + name: String, + first_span: Span, + second_span: Span, + module_namespace: String, + }, + UndefinedReference { + name: String, + span: Span, + module_namespace: String, + }, + InvalidConstraint { + message: String, + span: Span, + module_namespace: String, + }, + TypeMismatch { + expected: String, + found: String, + span: Span, + module_namespace: String, + }, + ConstraintTooPermissive { + message: String, + span: Span, + module_namespace: String, + }, + ReservedName { + name: String, + span: Span, + module_namespace: String, + }, + AmbiguousMain { + name: String, + namespace_suffix: String, + first_span: Span, + second_span: Span, + module_namespace: String, + }, + MultipleMain { + name: String, + first_span: Span, + second_span: Span, + module_namespace: String, + }, + ConflictNotAllowed { + name: String, + namespace_suffix: String, + span: Span, + module_namespace: String, + }, + CircularImport { + cycle: Vec, + span: Span, + module_namespace: String, + }, + UnusedImport { + name: String, + span: Span, + module_namespace: String, + }, } #[derive(Debug, Clone, Default)] diff --git a/mlf-lang/src/lexer.rs b/mlf-lang/src/lexer.rs index 61e09c9..1575d99 100644 --- a/mlf-lang/src/lexer.rs +++ b/mlf-lang/src/lexer.rs @@ -1,13 +1,13 @@ use alloc::string::String; use alloc::vec::Vec; use nom::{ + IResult, Parser, branch::alt, bytes::complete::{tag, take_while, take_while1}, character::complete::{char, multispace0, one_of}, combinator::{map, opt, recognize}, multi::many0, sequence::{delimited, pair, preceded}, - IResult, Parser, }; use crate::span::Span; @@ -142,7 +142,8 @@ fn identifier(input: &str) -> IResult<&str, Token> { let (rest, name) = recognize(pair( take_while1(is_ident_start), take_while(is_ident_continue), - )).parse(input)?; + )) + .parse(input)?; let token = match name { "as" => Token::As, @@ -178,17 +179,15 @@ fn type_ident(input: &str) -> IResult<&str, Token> { let (rest, name) = recognize(pair( take_while1(is_type_start), take_while(is_ident_continue), - )).parse(input)?; + )) + .parse(input)?; Ok((rest, Token::Ident(name.into()))) } fn raw_identifier(input: &str) -> IResult<&str, Token> { - let (rest, name) = delimited( - char('`'), - take_while1(|c: char| c != '`'), - char('`'), - ).parse(input)?; + let (rest, name) = + delimited(char('`'), take_while1(|c: char| c != '`'), char('`')).parse(input)?; Ok((rest, Token::Ident(name.into()))) } @@ -201,7 +200,8 @@ fn string_literal(input: &str) -> IResult<&str, Token> { recognize(pair(char('\\'), one_of(r#""\/bfnrt"#))), )))), char('"'), - ).parse(input)?; + ) + .parse(input)?; Ok((rest, Token::StringLit(s.into()))) } @@ -210,7 +210,8 @@ fn integer_literal(input: &str) -> IResult<&str, Token> { let (rest, num_str) = recognize(pair( opt(char('-')), take_while1(|c: char| c.is_ascii_digit()), - )).parse(input)?; + )) + .parse(input)?; let num = num_str.parse::().unwrap(); Ok((rest, Token::IntLit(num))) @@ -222,7 +223,8 @@ fn float_literal(input: &str) -> IResult<&str, Token> { take_while1(|c: char| c.is_ascii_digit()), char('.'), take_while1(|c: char| c.is_ascii_digit()), - )).parse(input)?; + )) + .parse(input)?; let num = num_str.parse::().unwrap(); Ok((rest, Token::FloatLit(num))) @@ -231,29 +233,23 @@ fn float_literal(input: &str) -> IResult<&str, Token> { fn doc_comment(input: &str) -> IResult<&str, Token> { let (rest, comment) = preceded( tag("///"), - map( - take_while(|c| c != '\n'), - |s: &str| Token::DocComment(s.trim().into()), - ), - ).parse(input)?; + map(take_while(|c| c != '\n'), |s: &str| { + Token::DocComment(s.trim().into()) + }), + ) + .parse(input)?; Ok((rest, comment)) } fn line_comment(input: &str) -> IResult<&str, ()> { - let (rest, _) = preceded( - tag("//"), - take_while(|c| c != '\n'), - ).parse(input)?; + let (rest, _) = preceded(tag("//"), take_while(|c| c != '\n')).parse(input)?; Ok((rest, ())) } fn hash_comment(input: &str) -> IResult<&str, ()> { - let (rest, _) = preceded( - char('#'), - take_while(|c| c != '\n'), - ).parse(input)?; + let (rest, _) = preceded(char('#'), take_while(|c| c != '\n')).parse(input)?; Ok((rest, ())) } @@ -276,7 +272,8 @@ fn symbol(input: &str) -> IResult<&str, Token> { map(char('('), |_| Token::LeftParen), map(char(')'), |_| Token::RightParen), map(char('_'), |_| Token::Underscore), - )).parse(input) + )) + .parse(input) } fn single_token(input: &str) -> IResult<&str, Option> { @@ -291,7 +288,8 @@ fn single_token(input: &str) -> IResult<&str, Option> { map(type_ident, Some), map(symbol, Some), // Parse symbols before identifiers so _ is caught map(identifier, Some), - )).parse(input) + )) + .parse(input) } pub fn tokenize(input: &str) -> Result, crate::error::ParseError> { @@ -300,7 +298,6 @@ pub fn tokenize(input: &str) -> Result, crate::error::ParseErr let mut pos = 0; while !remaining.is_empty() { - let ws_result: IResult<&str, &str> = multispace0(remaining); if let Ok((rest, ws)) = ws_result { pos += ws.len(); @@ -381,7 +378,10 @@ mod tests { fn test_doc_comment() { let input = "/// This is a doc comment"; let tokens = tokenize(input).unwrap(); - assert_eq!(tokens[0].token, Token::DocComment("This is a doc comment".into())); + assert_eq!( + tokens[0].token, + Token::DocComment("This is a doc comment".into()) + ); } #[test] diff --git a/mlf-lang/src/lib.rs b/mlf-lang/src/lib.rs index cb72860..393ddb1 100644 --- a/mlf-lang/src/lib.rs +++ b/mlf-lang/src/lib.rs @@ -17,7 +17,7 @@ pub use validate::validate_lexicon; pub use workspace::{ResolvedRef, Workspace}; // Standard library directory -use include_dir::{include_dir, Dir}; +use include_dir::{Dir, include_dir}; pub static STD_DIR: Dir<'static> = include_dir!("$CARGO_MANIFEST_DIR/../std"); diff --git a/mlf-lang/src/parser.rs b/mlf-lang/src/parser.rs index f9dafe9..18345e8 100644 --- a/mlf-lang/src/parser.rs +++ b/mlf-lang/src/parser.rs @@ -1,6 +1,11 @@ -use alloc::vec::Vec; -use crate::{Lexicon, ParseError, ast::*, lexer::{tokenize, SpannedToken}, span::{Span, Spanned}}; use crate::lexer::Token as LexToken; +use crate::{ + Lexicon, ParseError, + ast::*, + lexer::{SpannedToken, tokenize}, + span::{Span, Spanned}, +}; +use alloc::vec::Vec; struct Parser { tokens: Vec, @@ -264,7 +269,9 @@ impl Parser { let first_ident = self.parse_ident()?; // Check if this is a generator selector or bare annotation - let (selectors, name) = if matches!(self.current().token, LexToken::Comma) || matches!(self.current().token, LexToken::Colon) { + let (selectors, name) = if matches!(self.current().token, LexToken::Comma) + || matches!(self.current().token, LexToken::Colon) + { // Generator selector syntax: @rust:deprecated or @rust,typescript:deprecated let mut selectors = alloc::vec![first_ident]; @@ -427,7 +434,11 @@ impl Parser { Ok(AnnotationValue::Object(entries)) } - fn parse_record(&mut self, docs: Vec, annotations: Vec) -> Result { + fn parse_record( + &mut self, + docs: Vec, + annotations: Vec, + ) -> Result { let start = self.expect(LexToken::Record)?; let name = self.parse_ident()?; self.expect(LexToken::LeftBrace)?; @@ -467,9 +478,9 @@ impl Parser { // Check for ! to mark as required (default is optional) let optional = if matches!(self.current().token, LexToken::Exclamation) { self.advance(); - false // ! means required + false // ! means required } else { - true // default is optional + true // default is optional }; self.expect(LexToken::Colon)?; @@ -495,7 +506,11 @@ impl Parser { }) } - fn parse_inline_type(&mut self, docs: Vec, annotations: Vec) -> Result { + fn parse_inline_type( + &mut self, + docs: Vec, + annotations: Vec, + ) -> Result { let start = self.expect(LexToken::Inline)?; self.expect(LexToken::Type)?; let name = self.parse_ident()?; @@ -512,10 +527,14 @@ impl Parser { })) } - fn parse_def_type(&mut self, docs: Vec, annotations: Vec) -> Result { + fn parse_def_type( + &mut self, + docs: Vec, + annotations: Vec, + ) -> Result { let start = self.expect(LexToken::Def)?; self.expect(LexToken::Type)?; - let name = self.parse_ident()?; // Backticked keywords are already converted to Ident by lexer + let name = self.parse_ident()?; // Backticked keywords are already converted to Ident by lexer self.expect(LexToken::Equals)?; let ty = self.parse_type()?; let end = self.expect(LexToken::Semicolon)?; @@ -529,7 +548,11 @@ impl Parser { })) } - fn parse_token(&mut self, docs: Vec, annotations: Vec) -> Result { + fn parse_token( + &mut self, + docs: Vec, + annotations: Vec, + ) -> Result { let start = self.expect(LexToken::Token)?; let name = self.parse_ident()?; let end = self.expect(LexToken::Semicolon)?; @@ -542,7 +565,11 @@ impl Parser { })) } - fn parse_query(&mut self, docs: Vec, annotations: Vec) -> Result { + fn parse_query( + &mut self, + docs: Vec, + annotations: Vec, + ) -> Result { let start = self.expect(LexToken::Query)?; let name = self.parse_ident()?; self.expect(LexToken::LeftParen)?; @@ -579,14 +606,20 @@ impl Parser { let span = Span::new(types[0].span().start, types.last().unwrap().span().end); // Return type unions are open by default (no ! support in return types yet) - ReturnType::Type(Type::Union { types, closed: false, span }) + ReturnType::Type(Type::Union { + types, + closed: false, + span, + }) } } else { ReturnType::Type(output) } } else { // No return type specified - ReturnType::None { span: right_paren_span } + ReturnType::None { + span: right_paren_span, + } }; let end = self.expect(LexToken::Semicolon)?; @@ -601,7 +634,11 @@ impl Parser { })) } - fn parse_procedure(&mut self, docs: Vec, annotations: Vec) -> Result { + fn parse_procedure( + &mut self, + docs: Vec, + annotations: Vec, + ) -> Result { let start = self.expect(LexToken::Procedure)?; let name = self.parse_ident()?; self.expect(LexToken::LeftParen)?; @@ -638,14 +675,20 @@ impl Parser { let span = Span::new(types[0].span().start, types.last().unwrap().span().end); // Return type unions are open by default (no ! support in return types yet) - ReturnType::Type(Type::Union { types, closed: false, span }) + ReturnType::Type(Type::Union { + types, + closed: false, + span, + }) } } else { ReturnType::Type(output) } } else { // No return type specified - ReturnType::None { span: right_paren_span } + ReturnType::None { + span: right_paren_span, + } }; let end = self.expect(LexToken::Semicolon)?; @@ -660,7 +703,11 @@ impl Parser { })) } - fn parse_subscription(&mut self, docs: Vec, annotations: Vec) -> Result { + fn parse_subscription( + &mut self, + docs: Vec, + annotations: Vec, + ) -> Result { let start = self.expect(LexToken::Subscription)?; let name = self.parse_ident()?; self.expect(LexToken::LeftParen)?; @@ -736,7 +783,10 @@ impl Parser { self.advance(); } else if !matches!(self.current().token, LexToken::RightBrace) { return Err(ParseError::Syntax { - message: alloc::format!("Expected comma or closing brace, found {}", self.current().token), + message: alloc::format!( + "Expected comma or closing brace, found {}", + self.current().token + ), span: self.current().span, }); } @@ -785,9 +835,9 @@ impl Parser { // Check for ! to mark as required (default is optional) let optional = if matches!(self.current().token, LexToken::Exclamation) { self.advance(); - false // ! means required + false // ! means required } else { - true // default is optional + true // default is optional }; self.expect(LexToken::Colon)?; @@ -880,7 +930,11 @@ impl Parser { let span = Span::new(types[0].span().start, types.last().unwrap().span().end); // Unions are open by default, closed if ! is present let closed = has_exclamation; - return Ok(Type::Union { types, closed, span }); + return Ok(Type::Union { + types, + closed, + span, + }); } Ok(base) @@ -1113,7 +1167,9 @@ impl Parser { } _ => { return Err(ParseError::Syntax { - message: alloc::format!("Expected string literal or identifier in enum"), + message: alloc::format!( + "Expected string literal or identifier in enum" + ), span: current.span, }); } @@ -1258,7 +1314,9 @@ impl Parser { } _ => { return Err(ParseError::Syntax { - message: alloc::format!("Expected string literal or identifier in knownValues"), + message: alloc::format!( + "Expected string literal or identifier in knownValues" + ), span: current.span, }); } @@ -1314,7 +1372,9 @@ impl Parser { } _ => { return Err(ParseError::Syntax { - message: alloc::format!("Expected string, integer, boolean, or identifier for default"), + message: alloc::format!( + "Expected string, integer, boolean, or identifier for default" + ), span: current_span, }); } @@ -1354,7 +1414,9 @@ impl Parser { } _ => { return Err(ParseError::Syntax { - message: alloc::format!("Expected string, integer, boolean, or identifier for const"), + message: alloc::format!( + "Expected string, integer, boolean, or identifier for const" + ), span: current_span, }); } @@ -1500,14 +1562,12 @@ mod tests { let lexicon = result.unwrap(); assert_eq!(lexicon.items.len(), 1); match &lexicon.items[0] { - Item::InlineType(a) => { - match &a.ty { - Type::Constrained { constraints, .. } => { - assert_eq!(constraints.len(), 1); - } - _ => panic!("Expected constrained type"), + Item::InlineType(a) => match &a.ty { + Type::Constrained { constraints, .. } => { + assert_eq!(constraints.len(), 1); } - } + _ => panic!("Expected constrained type"), + }, _ => panic!("Expected inline type"), } } @@ -1520,14 +1580,12 @@ mod tests { let lexicon = result.unwrap(); assert_eq!(lexicon.items.len(), 1); match &lexicon.items[0] { - Item::InlineType(a) => { - match &a.ty { - Type::Union { types, .. } => { - assert_eq!(types.len(), 2); - } - _ => panic!("Expected union type"), + Item::InlineType(a) => match &a.ty { + Type::Union { types, .. } => { + assert_eq!(types.len(), 2); } - } + _ => panic!("Expected union type"), + }, _ => panic!("Expected inline type"), } } @@ -1540,12 +1598,10 @@ mod tests { let lexicon = result.unwrap(); assert_eq!(lexicon.items.len(), 1); match &lexicon.items[0] { - Item::InlineType(a) => { - match &a.ty { - Type::Array { .. } => {} - _ => panic!("Expected array type"), - } - } + Item::InlineType(a) => match &a.ty { + Type::Array { .. } => {} + _ => panic!("Expected array type"), + }, _ => panic!("Expected inline type"), } } @@ -1756,19 +1812,15 @@ mod tests { assert!(result.is_ok()); let lexicon = result.unwrap(); match &lexicon.items[0] { - Item::InlineType(a) => { - match &a.ty { - Type::Constrained { constraints, .. } => { - match &constraints[0] { - Constraint::Enum { values, .. } => { - assert_eq!(values.len(), 2); - } - _ => panic!("Expected enum constraint"), - } + Item::InlineType(a) => match &a.ty { + Type::Constrained { constraints, .. } => match &constraints[0] { + Constraint::Enum { values, .. } => { + assert_eq!(values.len(), 2); } - _ => panic!("Expected constrained type"), - } - } + _ => panic!("Expected enum constraint"), + }, + _ => panic!("Expected constrained type"), + }, _ => panic!("Expected inline type"), } } diff --git a/mlf-lang/src/workspace.rs b/mlf-lang/src/workspace.rs index dc4f12c..8d49b59 100644 --- a/mlf-lang/src/workspace.rs +++ b/mlf-lang/src/workspace.rs @@ -1,7 +1,11 @@ +use crate::{ + ast::*, + error::{ValidationError, ValidationErrors}, + span::Span, +}; use alloc::collections::{BTreeMap, BTreeSet}; use alloc::string::String; use alloc::vec::Vec; -use crate::{ast::*, error::{ValidationError, ValidationErrors}, span::Span}; #[derive(Debug, Clone, PartialEq)] pub struct Workspace { @@ -77,16 +81,15 @@ impl Workspace { pub fn with_prelude() -> Result { let mut ws = Self::new(); - let prelude_lexicon = crate::parser::parse_lexicon(crate::PRELUDE) - .map_err(|e| { - let mut errors = ValidationErrors::new(); - errors.push(ValidationError::InvalidConstraint { - message: alloc::format!("Failed to parse prelude: {:?}", e), - span: crate::span::Span::new(0, 0), - module_namespace: "prelude".to_string(), - }); - errors - })?; + let prelude_lexicon = crate::parser::parse_lexicon(crate::PRELUDE).map_err(|e| { + let mut errors = ValidationErrors::new(); + errors.push(ValidationError::InvalidConstraint { + message: alloc::format!("Failed to parse prelude: {:?}", e), + span: crate::span::Span::new(0, 0), + module_namespace: "prelude".to_string(), + }); + errors + })?; ws.add_module("prelude".into(), prelude_lexicon)?; Ok(ws) } @@ -95,16 +98,15 @@ impl Workspace { let mut ws = Self::new(); // Add prelude - let prelude_lexicon = crate::parser::parse_lexicon(crate::PRELUDE) - .map_err(|e| { - let mut errors = ValidationErrors::new(); - errors.push(ValidationError::InvalidConstraint { - message: alloc::format!("Failed to parse prelude: {:?}", e), - span: crate::span::Span::new(0, 0), - module_namespace: "prelude".to_string(), - }); - errors - })?; + let prelude_lexicon = crate::parser::parse_lexicon(crate::PRELUDE).map_err(|e| { + let mut errors = ValidationErrors::new(); + errors.push(ValidationError::InvalidConstraint { + message: alloc::format!("Failed to parse prelude: {:?}", e), + span: crate::span::Span::new(0, 0), + module_namespace: "prelude".to_string(), + }); + errors + })?; ws.add_module("prelude".into(), prelude_lexicon)?; // Recursively load all .mlf files from the std directory @@ -124,23 +126,28 @@ impl Workspace { .replace('/', "."); if let Some(contents) = file.contents_utf8() { - let lexicon = crate::parser::parse_lexicon(contents) - .map_err(|e| { - let mut errors = ValidationErrors::new(); - errors.push(ValidationError::InvalidConstraint { - message: alloc::format!("Failed to parse {} (file: {}): {:?}", namespace, path_str, e), - span: crate::span::Span::new(0, 0), - module_namespace: namespace.clone(), - }); - errors - })?; + let lexicon = crate::parser::parse_lexicon(contents).map_err(|e| { + let mut errors = ValidationErrors::new(); + errors.push(ValidationError::InvalidConstraint { + message: alloc::format!( + "Failed to parse {} (file: {}): {:?}", + namespace, + path_str, + e + ), + span: crate::span::Span::new(0, 0), + module_namespace: namespace.clone(), + }); + errors + })?; // If the module already exists, merge the items if ws.modules.contains_key(&namespace) { let module = ws.modules.get_mut(&namespace).unwrap(); module.lexicon.items.extend(lexicon.items); // Rebuild symbol table - module.symbols = Workspace::build_symbol_table(&namespace, &module.lexicon)?; + module.symbols = + Workspace::build_symbol_table(&namespace, &module.lexicon)?; } else { ws.add_module(namespace, lexicon)?; } @@ -161,7 +168,11 @@ impl Workspace { Ok(ws) } - pub fn add_module(&mut self, namespace: String, lexicon: Lexicon) -> Result<(), ValidationErrors> { + pub fn add_module( + &mut self, + namespace: String, + lexicon: Lexicon, + ) -> Result<(), ValidationErrors> { let symbols = Self::build_symbol_table(&namespace, &lexicon)?; let module = Module { @@ -205,14 +216,18 @@ impl Workspace { } // Filter out warnings (UnusedImport) from blocking errors - let blocking_errors: Vec = errors.errors.into_iter() + let blocking_errors: Vec = errors + .errors + .into_iter() .filter(|e| !matches!(e, ValidationError::UnusedImport { .. })) .collect(); if blocking_errors.is_empty() { Ok(()) } else { - Err(ValidationErrors { errors: blocking_errors }) + Err(ValidationErrors { + errors: blocking_errors, + }) } } @@ -284,11 +299,19 @@ impl Workspace { } } - fn typecheck_inline_type(&self, namespace: &str, inline_type: &InlineType) -> Result<(), ValidationErrors> { + fn typecheck_inline_type( + &self, + namespace: &str, + inline_type: &InlineType, + ) -> Result<(), ValidationErrors> { self.typecheck_type(namespace, &inline_type.ty) } - fn typecheck_def_type(&self, namespace: &str, def_type: &DefType) -> Result<(), ValidationErrors> { + fn typecheck_def_type( + &self, + namespace: &str, + def_type: &DefType, + ) -> Result<(), ValidationErrors> { self.typecheck_type(namespace, &def_type.ty) } @@ -352,18 +375,26 @@ impl Workspace { } } Type::Parenthesized { inner, .. } => self.typecheck_type(namespace, inner), - Type::Constrained { base, constraints, span } => { + Type::Constrained { + base, + constraints, + span, + } => { let mut errors = ValidationErrors::new(); if let Err(mut base_errors) = self.typecheck_type(namespace, base) { errors.append(&mut base_errors); } - if let Err(mut constraint_errors) = self.typecheck_constraints(namespace, base, constraints, *span) { + if let Err(mut constraint_errors) = + self.typecheck_constraints(namespace, base, constraints, *span) + { errors.append(&mut constraint_errors); } - if let Err(mut refinement_errors) = self.check_constraint_refinement(namespace, base, constraints) { + if let Err(mut refinement_errors) = + self.check_constraint_refinement(namespace, base, constraints) + { errors.append(&mut refinement_errors); } @@ -376,21 +407,28 @@ impl Workspace { } } - fn typecheck_constraints(&self, namespace: &str, base: &Type, constraints: &[Constraint], _span: Span) -> Result<(), ValidationErrors> { + fn typecheck_constraints( + &self, + namespace: &str, + base: &Type, + constraints: &[Constraint], + _span: Span, + ) -> Result<(), ValidationErrors> { let mut errors = ValidationErrors::new(); let base_kind = self.get_base_primitive(base); for constraint in constraints { match constraint { - Constraint::MinLength { span, .. } - | Constraint::MaxLength { span, .. } => { + Constraint::MinLength { span, .. } | Constraint::MaxLength { span, .. } => { // MinLength/MaxLength can be applied to strings or arrays let is_string = matches!(base_kind, Some(PrimitiveType::String)); let is_array = matches!(base, Type::Array { .. }); if !is_string && !is_array { errors.push(ValidationError::InvalidConstraint { - message: alloc::format!("Length constraint can only be applied to string or array types"), + message: alloc::format!( + "Length constraint can only be applied to string or array types" + ), span: *span, module_namespace: namespace.to_string(), }); @@ -408,8 +446,7 @@ impl Workspace { }); } } - Constraint::Minimum { span, .. } - | Constraint::Maximum { span, .. } => { + Constraint::Minimum { span, .. } | Constraint::Maximum { span, .. } => { if !matches!(base_kind, Some(PrimitiveType::Integer)) { errors.push(ValidationError::InvalidConstraint { message: alloc::format!("Numeric constraint on non-numeric type"), @@ -418,8 +455,7 @@ impl Workspace { }); } } - Constraint::Accept { span, .. } - | Constraint::MaxSize { span, .. } => { + Constraint::Accept { span, .. } | Constraint::MaxSize { span, .. } => { if !matches!(base_kind, Some(PrimitiveType::Blob)) { errors.push(ValidationError::InvalidConstraint { message: alloc::format!("Blob constraint on non-blob type"), @@ -441,7 +477,12 @@ impl Workspace { } } - fn check_constraint_refinement(&self, namespace: &str, base: &Type, new_constraints: &[Constraint]) -> Result<(), ValidationErrors> { + fn check_constraint_refinement( + &self, + namespace: &str, + base: &Type, + new_constraints: &[Constraint], + ) -> Result<(), ValidationErrors> { let base_constraints = self.get_base_constraints(base); if base_constraints.is_empty() { @@ -452,14 +493,21 @@ impl Workspace { for new_constraint in new_constraints { match new_constraint { - Constraint::MaxLength { value: new_max, span } => { + Constraint::MaxLength { + value: new_max, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::MaxLength { value: base_max, .. } = base_constraint { + if let Constraint::MaxLength { + value: base_max, .. + } = base_constraint + { if new_max > base_max { errors.push(ValidationError::ConstraintTooPermissive { message: alloc::format!( "maxLength {} is greater than base maxLength {}", - new_max, base_max + new_max, + base_max ), span: *span, module_namespace: namespace.to_string(), @@ -468,14 +516,21 @@ impl Workspace { } } } - Constraint::MinLength { value: new_min, span } => { + Constraint::MinLength { + value: new_min, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::MinLength { value: base_min, .. } = base_constraint { + if let Constraint::MinLength { + value: base_min, .. + } = base_constraint + { if new_min < base_min { errors.push(ValidationError::ConstraintTooPermissive { message: alloc::format!( "minLength {} is less than base minLength {}", - new_min, base_min + new_min, + base_min ), span: *span, module_namespace: namespace.to_string(), @@ -484,14 +539,21 @@ impl Workspace { } } } - Constraint::Maximum { value: new_max, span } => { + Constraint::Maximum { + value: new_max, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::Maximum { value: base_max, .. } = base_constraint { + if let Constraint::Maximum { + value: base_max, .. + } = base_constraint + { if new_max > base_max { errors.push(ValidationError::ConstraintTooPermissive { message: alloc::format!( "maximum {} is greater than base maximum {}", - new_max, base_max + new_max, + base_max ), span: *span, module_namespace: namespace.to_string(), @@ -500,14 +562,21 @@ impl Workspace { } } } - Constraint::Minimum { value: new_min, span } => { + Constraint::Minimum { + value: new_min, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::Minimum { value: base_min, .. } = base_constraint { + if let Constraint::Minimum { + value: base_min, .. + } = base_constraint + { if new_min < base_min { errors.push(ValidationError::ConstraintTooPermissive { message: alloc::format!( "minimum {} is less than base minimum {}", - new_min, base_min + new_min, + base_min ), span: *span, module_namespace: namespace.to_string(), @@ -516,14 +585,21 @@ impl Workspace { } } } - Constraint::MaxGraphemes { value: new_max, span } => { + Constraint::MaxGraphemes { + value: new_max, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::MaxGraphemes { value: base_max, .. } = base_constraint { + if let Constraint::MaxGraphemes { + value: base_max, .. + } = base_constraint + { if new_max > base_max { errors.push(ValidationError::ConstraintTooPermissive { message: alloc::format!( "maxGraphemes {} is greater than base maxGraphemes {}", - new_max, base_max + new_max, + base_max ), span: *span, module_namespace: namespace.to_string(), @@ -532,14 +608,21 @@ impl Workspace { } } } - Constraint::MinGraphemes { value: new_min, span } => { + Constraint::MinGraphemes { + value: new_min, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::MinGraphemes { value: base_min, .. } = base_constraint { + if let Constraint::MinGraphemes { + value: base_min, .. + } = base_constraint + { if new_min < base_min { errors.push(ValidationError::ConstraintTooPermissive { message: alloc::format!( "minGraphemes {} is less than base minGraphemes {}", - new_min, base_min + new_min, + base_min ), span: *span, module_namespace: namespace.to_string(), @@ -548,14 +631,21 @@ impl Workspace { } } } - Constraint::MaxSize { value: new_max, span } => { + Constraint::MaxSize { + value: new_max, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::MaxSize { value: base_max, .. } = base_constraint { + if let Constraint::MaxSize { + value: base_max, .. + } = base_constraint + { if new_max > base_max { errors.push(ValidationError::ConstraintTooPermissive { message: alloc::format!( "maxSize {} is greater than base maxSize {}", - new_max, base_max + new_max, + base_max ), span: *span, module_namespace: namespace.to_string(), @@ -564,9 +654,16 @@ impl Workspace { } } } - Constraint::Enum { values: new_values, span } => { + Constraint::Enum { + values: new_values, + span, + } => { for base_constraint in &base_constraints { - if let Constraint::Enum { values: base_values, .. } = base_constraint { + if let Constraint::Enum { + values: base_values, + .. + } = base_constraint + { for new_val in new_values { if !base_values.contains(new_val) { errors.push(ValidationError::ConstraintTooPermissive { @@ -595,7 +692,9 @@ impl Workspace { fn get_base_constraints(&self, ty: &Type) -> Vec { match ty { - Type::Constrained { base, constraints, .. } => { + Type::Constrained { + base, constraints, .. + } => { let mut all_constraints = constraints.clone(); all_constraints.extend(self.get_base_constraints(base)); all_constraints @@ -734,10 +833,11 @@ impl Workspace { /// Returns a vector of (local_name, original_path) tuples pub fn get_imports(&self, namespace: &str) -> Vec<(String, Vec)> { if let Some(module) = self.modules.get(namespace) { - module.imports.mappings.iter() - .map(|(local_name, imported)| { - (local_name.clone(), imported.original_path.clone()) - }) + module + .imports + .mappings + .iter() + .map(|(local_name, imported)| (local_name.clone(), imported.original_path.clone())) .collect() } else { Vec::new() @@ -802,7 +902,9 @@ impl Workspace { if let Some(module) = self.modules.get(current_namespace) { if module.symbols.types.contains_key(name) { - return Some(ResolvedRef::Local { def_name: name.clone() }); + return Some(ResolvedRef::Local { + def_name: name.clone(), + }); } if let Some(imported) = module.imports.mappings.get(name) { @@ -810,12 +912,18 @@ impl Workspace { // (e.g. `["com", "atproto", "label", "defs", "label"]`). // All but the last segment form the namespace. if imported.original_path.len() > 1 { - let namespace = imported.original_path[..imported.original_path.len() - 1].join("."); + let namespace = + imported.original_path[..imported.original_path.len() - 1].join("."); let def_name = imported.original_path.last().unwrap().clone(); - return Some(ResolvedRef::External { namespace, def_name }); + return Some(ResolvedRef::External { + namespace, + def_name, + }); } else { // Pathological single-segment import: treat as local. - return Some(ResolvedRef::Local { def_name: name.clone() }); + return Some(ResolvedRef::Local { + def_name: name.clone(), + }); } } } @@ -844,9 +952,14 @@ impl Workspace { if module.symbols.types.contains_key(type_name) { let namespace = target_namespace; if namespace == current_namespace { - return Some(ResolvedRef::Local { def_name: type_name.clone() }); + return Some(ResolvedRef::Local { + def_name: type_name.clone(), + }); } - return Some(ResolvedRef::External { namespace, def_name: type_name.clone() }); + return Some(ResolvedRef::External { + namespace, + def_name: type_name.clone(), + }); } } @@ -871,7 +984,11 @@ impl Workspace { /// Thin shim over [`Self::resolve_ref`]; prefer that for new call /// sites where the variant (local vs external vs implicit-main) /// matters. - pub fn resolve_reference_namespace(&self, path: &Path, current_namespace: &str) -> Option { + pub fn resolve_reference_namespace( + &self, + path: &Path, + current_namespace: &str, + ) -> Option { match self.resolve_ref(path, current_namespace)? { ResolvedRef::Local { .. } => Some(current_namespace.to_string()), ResolvedRef::External { namespace, .. } @@ -903,7 +1020,11 @@ impl Workspace { } } - fn resolve_use_statement(&mut self, current_namespace: &str, use_stmt: &Use) -> Result<(), ValidationErrors> { + fn resolve_use_statement( + &mut self, + current_namespace: &str, + use_stmt: &Use, + ) -> Result<(), ValidationErrors> { let mut errors = ValidationErrors::new(); // Determine if this is: @@ -921,7 +1042,8 @@ impl Workspace { // In this case, path has >=2 segments and items has 1 item whose name matches the last segment if items.len() == 1 && use_stmt.path.segments.len() >= 2 - && items[0].name.name == use_stmt.path.segments.last().unwrap().name { + && items[0].name.name == use_stmt.path.segments.last().unwrap().name + { // Old syntax: use namespace.typename as alias let namespace = use_stmt.path.segments[..use_stmt.path.segments.len() - 1] .iter() @@ -938,7 +1060,10 @@ impl Workspace { // Check if the target namespace exists or if there are modules with that prefix let namespace_exists = self.modules.contains_key(&target_namespace); - let has_children = self.modules.keys().any(|ns| ns.starts_with(&alloc::format!("{}.", target_namespace))); + let has_children = self + .modules + .keys() + .any(|ns| ns.starts_with(&alloc::format!("{}.", target_namespace))); if !namespace_exists && !has_children { errors.push(ValidationError::UndefinedReference { @@ -954,7 +1079,10 @@ impl Workspace { // Check for implicit main resolution // If namespace suffix matches a type name, import only that type // Otherwise, this is a namespace alias for path shortening - let namespace_suffix = target_namespace.split('.').last().unwrap_or(&target_namespace); + let namespace_suffix = target_namespace + .split('.') + .last() + .unwrap_or(&target_namespace); if let Some(target_module) = self.modules.get(&target_namespace) { // Module exists - check if there's a type matching the namespace suffix @@ -983,14 +1111,20 @@ impl Workspace { UseImports::All => { // Check for implicit main resolution // If namespace suffix matches a type name, import only that type - let namespace_suffix = target_namespace.split('.').last().unwrap_or(&target_namespace); + let namespace_suffix = target_namespace + .split('.') + .last() + .unwrap_or(&target_namespace); if let Some(target_module) = self.modules.get(&target_namespace) { // Check if there's a type matching the namespace suffix if target_module.symbols.types.contains_key(namespace_suffix) { // Implicit main resolution: import only the type matching the namespace suffix let imported = ImportedSymbol { - original_path: use_stmt.path.segments.iter() + original_path: use_stmt + .path + .segments + .iter() .map(|s| s.name.clone()) .collect(), local_name: namespace_suffix.to_string(), @@ -1012,7 +1146,10 @@ impl Workspace { let mut imports = Vec::new(); // Get the namespace suffix for main resolution - let namespace_suffix = target_namespace.split('.').last().unwrap_or(&target_namespace); + let namespace_suffix = target_namespace + .split('.') + .last() + .unwrap_or(&target_namespace); for item in items { // Determine the actual type name to look up @@ -1047,13 +1184,19 @@ impl Workspace { let original_path = if is_single_type_import { // Old syntax: use namespace.typename as alias // Path segments already include the type name - use_stmt.path.segments.iter() + use_stmt + .path + .segments + .iter() .map(|s| s.name.clone()) .collect() } else { // New syntax: use namespace { items }; // Need to append the type name to the namespace - use_stmt.path.segments.iter() + use_stmt + .path + .segments + .iter() .map(|s| s.name.clone()) .chain(core::iter::once(type_name.clone())) .collect() @@ -1080,7 +1223,10 @@ impl Workspace { // Add namespace alias if one was created if let Some((alias, full_namespace)) = namespace_alias_to_add { - module.imports.namespace_aliases.insert(alias, full_namespace); + module + .imports + .namespace_aliases + .insert(alias, full_namespace); } if errors.is_empty() { @@ -1094,7 +1240,10 @@ impl Workspace { annotations.iter().any(|ann| ann.name.name == "main") } - fn build_symbol_table(namespace: &str, lexicon: &Lexicon) -> Result { + fn build_symbol_table( + namespace: &str, + lexicon: &Lexicon, + ) -> Result { let mut symbols = SymbolTable { types: BTreeMap::new(), }; @@ -1119,7 +1268,10 @@ impl Workspace { }; if let Some(name) = name { - items_by_name.entry(name.to_string()).or_insert_with(Vec::new).push(item); + items_by_name + .entry(name.to_string()) + .or_insert_with(Vec::new) + .push(item); } } @@ -1253,7 +1405,8 @@ impl Workspace { } // Check @main annotations - let items_with_main: Vec<&Item> = items.iter() + let items_with_main: Vec<&Item> = items + .iter() .filter(|item| { let annotations = match item { Item::Record(r) => &r.annotations, @@ -1415,11 +1568,19 @@ impl Workspace { } } - fn resolve_inline_type(&mut self, namespace: &str, inline_type: &InlineType) -> Result<(), ValidationErrors> { + fn resolve_inline_type( + &mut self, + namespace: &str, + inline_type: &InlineType, + ) -> Result<(), ValidationErrors> { self.resolve_type(namespace, &inline_type.ty) } - fn resolve_def_type(&mut self, namespace: &str, def_type: &DefType) -> Result<(), ValidationErrors> { + fn resolve_def_type( + &mut self, + namespace: &str, + def_type: &DefType, + ) -> Result<(), ValidationErrors> { self.resolve_type(namespace, &def_type.ty) } @@ -1455,7 +1616,11 @@ impl Workspace { } } - fn resolve_procedure(&mut self, namespace: &str, procedure: &Procedure) -> Result<(), ValidationErrors> { + fn resolve_procedure( + &mut self, + namespace: &str, + procedure: &Procedure, + ) -> Result<(), ValidationErrors> { let mut errors = ValidationErrors::new(); for param in &procedure.params { @@ -1487,7 +1652,11 @@ impl Workspace { } } - fn resolve_subscription(&mut self, namespace: &str, subscription: &Subscription) -> Result<(), ValidationErrors> { + fn resolve_subscription( + &mut self, + namespace: &str, + subscription: &Subscription, + ) -> Result<(), ValidationErrors> { let mut errors = ValidationErrors::new(); for param in &subscription.params { @@ -1512,12 +1681,8 @@ impl Workspace { fn resolve_type(&mut self, namespace: &str, ty: &Type) -> Result<(), ValidationErrors> { match ty { Type::Primitive { .. } | Type::Unknown { .. } => Ok(()), - Type::Reference { path, span } => { - self.resolve_reference(namespace, path, *span) - } - Type::Array { inner, .. } => { - self.resolve_type(namespace, inner) - } + Type::Reference { path, span } => self.resolve_reference(namespace, path, *span), + Type::Array { inner, .. } => self.resolve_type(namespace, inner), Type::Union { types, .. } => { let mut errors = ValidationErrors::new(); for ty in types { @@ -1544,16 +1709,17 @@ impl Workspace { Err(errors) } } - Type::Parenthesized { inner, .. } => { - self.resolve_type(namespace, inner) - } - Type::Constrained { base, .. } => { - self.resolve_type(namespace, base) - } + Type::Parenthesized { inner, .. } => self.resolve_type(namespace, inner), + Type::Constrained { base, .. } => self.resolve_type(namespace, base), } } - fn resolve_reference(&mut self, current_namespace: &str, path: &Path, span: Span) -> Result<(), ValidationErrors> { + fn resolve_reference( + &mut self, + current_namespace: &str, + path: &Path, + span: Span, + ) -> Result<(), ValidationErrors> { let full_path = path.to_string(); if path.segments.len() == 1 { @@ -1643,7 +1809,10 @@ impl Workspace { alloc::format!("{}.{}", target_namespace, type_name) }; if let Some(module) = self.modules.get(&expanded_full_path) { - let namespace_suffix = expanded_full_path.split('.').last().unwrap_or(&expanded_full_path); + let namespace_suffix = expanded_full_path + .split('.') + .last() + .unwrap_or(&expanded_full_path); if namespace_suffix == type_name && module.symbols.types.contains_key(type_name) { return Ok(()); } @@ -1734,7 +1903,9 @@ mod tests { let a = parse_lexicon("record post {}").unwrap(); ws.add_module("app.bsky.feed".into(), a).unwrap(); - let b = parse_lexicon("use app.bsky.feed.post as FeedPost; record like { subject: FeedPost, }").unwrap(); + let b = + parse_lexicon("use app.bsky.feed.post as FeedPost; record like { subject: FeedPost, }") + .unwrap(); ws.add_module("app.bsky.feed.like".into(), b).unwrap(); let result = ws.resolve(); @@ -1880,7 +2051,12 @@ mod tests { let result = ws.resolve(); assert!(result.is_err()); let errors = result.unwrap_err(); - assert!(errors.errors.iter().any(|e| matches!(e, ValidationError::ConstraintTooPermissive { .. }))); + assert!( + errors + .errors + .iter() + .any(|e| matches!(e, ValidationError::ConstraintTooPermissive { .. })) + ); } #[test] @@ -2142,7 +2318,12 @@ mod tests { let result = ws.add_module("com.example.thread".into(), lexicon); assert!(result.is_err()); let errors = result.unwrap_err(); - assert!(errors.errors.iter().any(|e| matches!(e, ValidationError::AmbiguousMain { .. }))); + assert!( + errors + .errors + .iter() + .any(|e| matches!(e, ValidationError::AmbiguousMain { .. })) + ); } #[test] @@ -2164,7 +2345,12 @@ mod tests { let result = ws.add_module("com.example.thread".into(), lexicon); assert!(result.is_err()); let errors = result.unwrap_err(); - assert!(errors.errors.iter().any(|e| matches!(e, ValidationError::MultipleMain { .. }))); + assert!( + errors + .errors + .iter() + .any(|e| matches!(e, ValidationError::MultipleMain { .. })) + ); } #[test] @@ -2185,7 +2371,12 @@ mod tests { let result = ws.add_module("com.example.thread".into(), lexicon); assert!(result.is_err()); let errors = result.unwrap_err(); - assert!(errors.errors.iter().any(|e| matches!(e, ValidationError::ConflictNotAllowed { .. }))); + assert!( + errors + .errors + .iter() + .any(|e| matches!(e, ValidationError::ConflictNotAllowed { .. })) + ); } #[test] @@ -2203,7 +2394,12 @@ mod tests { let result = ws.add_module("com.example.thread".into(), lexicon); assert!(result.is_err()); let errors = result.unwrap_err(); - assert!(errors.errors.iter().any(|e| matches!(e, ValidationError::DuplicateDefinition { .. }))); + assert!( + errors + .errors + .iter() + .any(|e| matches!(e, ValidationError::DuplicateDefinition { .. })) + ); } #[test] @@ -2278,7 +2474,8 @@ mod tests { // Create a namespace where the namespace suffix matches a type name // No @main annotation needed for implicit main resolution - let profile_ns = parse_lexicon(r#" + let profile_ns = parse_lexicon( + r#" def type color = { red!: integer, green!: integer, @@ -2288,18 +2485,25 @@ mod tests { record profile { color: color, } - "#).unwrap(); - ws.add_module("place.stream.chat.profile".into(), profile_ns).unwrap(); + "#, + ) + .unwrap(); + ws.add_module("place.stream.chat.profile".into(), profile_ns) + .unwrap(); // Import using implicit main resolution - should only import profile, not color - let bookmark = parse_lexicon(r#" + let bookmark = parse_lexicon( + r#" use place.stream.chat.profile; record bookmark { owner!: profile, } - "#).unwrap(); - ws.add_module("com.example.bookmark".into(), bookmark).unwrap(); + "#, + ) + .unwrap(); + ws.add_module("com.example.bookmark".into(), bookmark) + .unwrap(); // Should resolve without unused import warning for color let result = ws.resolve(); @@ -2316,22 +2520,28 @@ mod tests { let mut ws = Workspace::new(); // Create a namespace where the suffix doesn't match a type name - let defs = parse_lexicon(r#" + let defs = parse_lexicon( + r#" def type foo = string; def type bar = integer; - "#).unwrap(); + "#, + ) + .unwrap(); ws.add_module("com.example.defs".into(), defs).unwrap(); // Using "use com.example;" creates a namespace alias "example" -> "com.example" // This allows referencing types via the shortened path - let app = parse_lexicon(r#" + let app = parse_lexicon( + r#" use com.example; record thing { x!: example.defs.foo, y!: example.defs.bar, } - "#).unwrap(); + "#, + ) + .unwrap(); ws.add_module("com.example.app".into(), app).unwrap(); // Should resolve successfully using the namespace alias @@ -2351,30 +2561,41 @@ mod tests { let mut ws = Workspace::new(); // Create nested namespaces - let actor_defs = parse_lexicon(r#" + let actor_defs = parse_lexicon( + r#" def type profileView = { did!: string, handle!: string, }; - "#).unwrap(); - ws.add_module("app.bsky.actor.defs".into(), actor_defs).unwrap(); - - let feed_post = parse_lexicon(r#" + "#, + ) + .unwrap(); + ws.add_module("app.bsky.actor.defs".into(), actor_defs) + .unwrap(); + + let feed_post = parse_lexicon( + r#" def type post = { text!: string, }; - "#).unwrap(); - ws.add_module("app.bsky.feed.post".into(), feed_post).unwrap(); + "#, + ) + .unwrap(); + ws.add_module("app.bsky.feed.post".into(), feed_post) + .unwrap(); // Use namespace alias to shorten references - let like = parse_lexicon(r#" + let like = parse_lexicon( + r#" use app.bsky; record like { subject!: bsky.feed.post, actor!: bsky.actor.defs.profileView, } - "#).unwrap(); + "#, + ) + .unwrap(); ws.add_module("app.bsky.feed.like".into(), like).unwrap(); let result = ws.resolve(); @@ -2389,22 +2610,28 @@ mod tests { let mut ws = Workspace::new(); // Create a namespace where the suffix doesn't match a type name - let defs = parse_lexicon(r#" + let defs = parse_lexicon( + r#" def type foo = string; def type bar = integer; - "#).unwrap(); + "#, + ) + .unwrap(); ws.add_module("com.example.defs".into(), defs).unwrap(); // Using "use com.example.defs;" where "defs" is not a type name // creates a namespace alias "defs" -> "com.example.defs" - let app = parse_lexicon(r#" + let app = parse_lexicon( + r#" use com.example.defs; record thing { x!: defs.foo, y!: defs.bar, } - "#).unwrap(); + "#, + ) + .unwrap(); ws.add_module("com.example.app".into(), app).unwrap(); // Should resolve successfully using the namespace alias diff --git a/mlf-lang/tests/integration_test.rs b/mlf-lang/tests/integration_test.rs index 92630a5..f393cfb 100644 --- a/mlf-lang/tests/integration_test.rs +++ b/mlf-lang/tests/integration_test.rs @@ -3,7 +3,7 @@ // `test.mlf` (plus optional support files) and an `expected.json`. use mlf_integration_tests::test_utils; -use mlf_lang::{parser::parse_lexicon, Workspace}; +use mlf_lang::{Workspace, parser::parse_lexicon}; use serde::Deserialize; use std::collections::HashMap; use std::fs; @@ -35,7 +35,9 @@ struct ExpectedError { } fn run_lang_test(test_mlf: &Path) -> datatest_stable::Result<()> { - let test_dir = test_mlf.parent().ok_or("test.mlf has no parent directory")?; + let test_dir = test_mlf + .parent() + .ok_or("test.mlf has no parent directory")?; let test_name = test_dir .file_name() .and_then(|s| s.to_str()) @@ -52,7 +54,8 @@ fn run_lang_test(test_mlf: &Path) -> datatest_stable::Result<()> { let expected_json = fs::read_to_string(&expected_path)?; let expected: ExpectedResult = serde_json::from_str(&expected_json)?; - let mut ws = Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; + let mut ws = + Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; // Support files load before test.mlf so cross-module refs resolve. let mut module_files: Vec<(String, PathBuf)> = config diff --git a/mlf-lexicon-fetcher/Cargo.toml b/mlf-lexicon-fetcher/Cargo.toml index daf0bb7..5e018d9 100644 --- a/mlf-lexicon-fetcher/Cargo.toml +++ b/mlf-lexicon-fetcher/Cargo.toml @@ -5,12 +5,11 @@ edition = "2024" license = "MIT" [dependencies] -hickory-resolver = "0.24" -thiserror = "2.0" -serde = { version = "1.0", features = ["derive"] } +mlf-atproto = { path = "../mlf-atproto" } +async-trait = "0.1" serde_json = "1.0" reqwest = { version = "0.12", features = ["json"] } -async-trait = "0.1" +thiserror = "2.0" tokio = { version = "1", features = ["rt"] } [dev-dependencies] diff --git a/mlf-lexicon-fetcher/examples/usage.rs b/mlf-lexicon-fetcher/examples/usage.rs index e86b9bd..c768e8a 100644 --- a/mlf-lexicon-fetcher/examples/usage.rs +++ b/mlf-lexicon-fetcher/examples/usage.rs @@ -9,7 +9,11 @@ async fn main() -> Result<(), Box> { println!("=== Example 1: Fetch with Metadata ==="); let mut dns_resolver = MockDnsResolver::new(); - dns_resolver.add_record("place.stream", "chat.profile", "did:plc:test123".to_string()); + dns_resolver.add_record( + "place.stream", + "chat.profile", + "did:plc:test123".to_string(), + ); let mut http_client = MockHttpClient::new(); http_client.add_lexicon( @@ -29,13 +33,19 @@ async fn main() -> Result<(), Box> { let fetcher = LexiconFetcher::new(dns_resolver, http_client); // New API returns metadata (DID, NSID, lexicon) - match fetcher.fetch_with_metadata("place.stream.chat.profile").await { + match fetcher + .fetch_with_metadata("place.stream.chat.profile") + .await + { Ok(result) => { println!("Successfully fetched {} lexicon(s):", result.lexicons.len()); for fetched in result.lexicons { println!(" NSID: {}", fetched.nsid); println!(" DID: {}", fetched.did); - println!(" Lexicon: {}", serde_json::to_string_pretty(&fetched.lexicon)?); + println!( + " Lexicon: {}", + serde_json::to_string_pretty(&fetched.lexicon)? + ); } } Err(e) => eprintln!("Error: {}", e), @@ -120,7 +130,10 @@ async fn main() -> Result<(), Box> { let fetcher4 = LexiconFetcher::new(dns_resolver4, http_client4); // Skip DNS resolution when DID is already known (e.g., from lockfile) - match fetcher4.fetch_from_did_with_metadata("did:plc:bsky123", "app.bsky.feed.post").await { + match fetcher4 + .fetch_from_did_with_metadata("did:plc:bsky123", "app.bsky.feed.post") + .await + { Ok(result) => { println!("Fetched from known DID:"); for fetched in result.lexicons { diff --git a/mlf-lexicon-fetcher/src/lib.rs b/mlf-lexicon-fetcher/src/lib.rs index 89ebc72..9dcdbdc 100644 --- a/mlf-lexicon-fetcher/src/lib.rs +++ b/mlf-lexicon-fetcher/src/lib.rs @@ -1,206 +1,56 @@ -// MLF Lexicon Fetcher -// Resolves ATProto lexicon NSIDs to DIDs via DNS TXT records -// and fetches lexicon JSON via HTTP +//! MLF lexicon fetcher. +//! +//! Read-side domain: given an NSID (or wildcard pattern), resolve it to +//! the publishing DID via `_lexicon.` TXT and fetch the +//! corresponding `com.atproto.lexicon.schema` record(s) from the PDS. +//! +//! Low-level primitives (DID parsing, DNS resolution, XRPC calls) live +//! in the `mlf-atproto` crate; this crate wraps them in NSID-aware +//! pattern-matching and lockfile-friendly result types. use async_trait::async_trait; -use hickory_resolver::config::{ResolverConfig, ResolverOpts}; -use hickory_resolver::TokioAsyncResolver; -use serde::Deserialize; +use mlf_atproto::identity::{self, IdentityError}; +use mlf_atproto::records::{self, RecordError}; +use mlf_atproto::xrpc::XrpcError; use std::collections::HashMap; use std::sync::{Arc, Mutex}; use thiserror::Error; -#[derive(Debug, Deserialize)] -struct AtProtoRecord { - uri: String, - value: serde_json::Value, -} +// Re-export identity primitives that existing callers depend on. +pub use mlf_atproto::identity::{ + DnsResolver, MockDnsResolver, RealDnsResolver, construct_dns_name, parse_nsid, +}; #[derive(Error, Debug)] pub enum LexiconFetcherError { - #[error("Failed to create DNS resolver: {0}")] - ResolverCreationFailed(String), + #[error("Identity resolution failed: {0}")] + Identity(#[from] IdentityError), - #[error("DNS lookup failed for {domain}: {error}")] - LookupFailed { domain: String, error: String }, + #[error("Record fetch failed: {0}")] + Records(#[from] RecordError), - #[error("No DID found in TXT record for {0}")] - NoDid(String), + #[error("XRPC call failed: {0}")] + Xrpc(#[from] XrpcError), #[error("Invalid NSID format: {0}")] InvalidNsid(String), - #[error("HTTP request failed: {0}")] - HttpRequestFailed(String), - - #[error("Failed to parse JSON response: {0}")] - JsonParseFailed(String), - #[error("Lexicon not found: {0}")] LexiconNotFound(String), - #[error("Invalid URL: {0}")] - InvalidUrl(String), + #[error("Could not extract NSID from record: {0}")] + MalformedRecord(String), } pub type Result = std::result::Result; -/// Trait for DNS resolution - allows mocking in tests -#[async_trait] -pub trait DnsResolver: Send + Sync { - /// Resolve an NSID to a DID via DNS TXT lookup - async fn resolve_lexicon_did(&self, authority: &str, name_segments: &str) -> Result; -} - -/// Real DNS resolver using hickory_resolver's async resolver -pub struct RealDnsResolver { - resolver: TokioAsyncResolver, -} - -impl RealDnsResolver { - pub async fn new() -> Result { - let resolver = TokioAsyncResolver::tokio(ResolverConfig::default(), ResolverOpts::default()); - Ok(Self { resolver }) - } - - pub async fn with_config(config: ResolverConfig, opts: ResolverOpts) -> Result { - let resolver = TokioAsyncResolver::tokio(config, opts); - Ok(Self { resolver }) - } -} - -#[async_trait] -impl DnsResolver for RealDnsResolver { - async fn resolve_lexicon_did(&self, authority: &str, name_segments: &str) -> Result { - let dns_name = construct_dns_name(authority, name_segments); - - // Lookup TXT records (async) - let response = self - .resolver - .txt_lookup(&dns_name) - .await - .map_err(|e| LexiconFetcherError::LookupFailed { - domain: dns_name.clone(), - error: e.to_string(), - })?; - - // Parse TXT records to find DID - for txt_record in response.iter() { - for txt_data in txt_record.txt_data() { - let text = String::from_utf8_lossy(txt_data); - // Look for "did=did:plc:..." or "did=did:web:..." - if let Some(did_value) = text.strip_prefix("did=") { - return Ok(did_value.trim().to_string()); - } - } - } - - Err(LexiconFetcherError::NoDid(dns_name)) - } -} - -/// Mock DNS resolver for testing -#[derive(Clone)] -pub struct MockDnsResolver { - records: Arc>>, -} - -impl MockDnsResolver { - pub fn new() -> Self { - Self { - records: Arc::new(Mutex::new(HashMap::new())), - } - } - - /// Add a mock DNS record (maps NSID authority+name to DID) - pub fn add_record(&mut self, authority: &str, name_segments: &str, did: String) { - let dns_name = construct_dns_name(authority, name_segments); - self.records.lock().unwrap().insert(dns_name, did); - } - - /// Add a mock record using full NSID - pub fn add_record_from_nsid(&mut self, nsid: &str, did: String) -> Result<()> { - let (authority, name_segments) = parse_nsid(nsid)?; - self.add_record(&authority, &name_segments, did); - Ok(()) - } -} - -impl Default for MockDnsResolver { - fn default() -> Self { - Self::new() - } -} - -#[async_trait] -impl DnsResolver for MockDnsResolver { - async fn resolve_lexicon_did(&self, authority: &str, name_segments: &str) -> Result { - let dns_name = construct_dns_name(authority, name_segments); - self.records - .lock() - .unwrap() - .get(&dns_name) - .cloned() - .ok_or_else(|| LexiconFetcherError::LookupFailed { - domain: dns_name.clone(), - error: "No mock record found".to_string(), - }) - } -} - -/// Construct DNS name from authority and name segments -/// For "app.bsky" + "actor": "_lexicon.actor.bsky.app" -/// For "place.stream" + "key": "_lexicon.key.stream.place" -pub fn construct_dns_name(authority: &str, name_segments: &str) -> String { - let auth_parts: Vec<&str> = authority.split('.').collect(); - let reversed_auth: Vec<&str> = auth_parts.iter().rev().copied().collect(); - - if name_segments.is_empty() { - // No name segments, just use reversed authority - // For "place.stream": "_lexicon.stream.place" - format!("_lexicon.{}", reversed_auth.join(".")) - } else { - // Prepend name segments before reversed authority - // For "app.bsky" + "actor": "_lexicon.actor.bsky.app" - format!("_lexicon.{}.{}", name_segments, reversed_auth.join(".")) - } -} - -/// Parse NSID into authority and name segments -/// For "place.stream.key", returns ("place.stream", "key") -/// For "app.bsky.actor.profile", returns ("app.bsky", "actor.profile") -pub fn parse_nsid(nsid: &str) -> Result<(String, String)> { - // NSID format: authority.name(.name)* - // Authority is first 2 segments (reversed domain) - let parts: Vec<&str> = nsid.split('.').collect(); - - if parts.len() < 2 { - return Err(LexiconFetcherError::InvalidNsid(format!( - "NSID must have at least 2 segments: {}", - nsid - ))); - } - - // Authority is first 2 segments - let authority = format!("{}.{}", parts[0], parts[1]); - - // Name segments are everything after the authority - let name_segments = if parts.len() > 2 { - parts[2..].join(".") - } else { - String::new() - }; - - Ok((authority, name_segments)) -} - -/// Trait for HTTP client - allows mocking in tests +/// Trait for the HTTP side of lexicon fetching. #[async_trait] pub trait HttpClient: Send + Sync { - /// Fetch a single lexicon by NSID from a DID's server + /// Fetch a single lexicon by NSID from a DID's repo. async fn fetch_lexicon(&self, did: &str, nsid: &str) -> Result; - /// Fetch all lexicons matching a pattern (e.g., "place.stream.*") + /// Fetch all lexicons under a wildcard pattern. async fn fetch_lexicons_pattern( &self, did: &str, @@ -208,7 +58,7 @@ pub trait HttpClient: Send + Sync { ) -> Result>; } -/// Real HTTP client using reqwest +/// Production HTTP client for fetching `com.atproto.lexicon.schema` records. pub struct RealHttpClient { client: reqwest::Client, } @@ -223,6 +73,13 @@ impl RealHttpClient { pub fn with_client(client: reqwest::Client) -> Self { Self { client } } + + async fn fetch_all_schema_records(&self, did: &str) -> Result> { + let pds = identity::resolve_did_to_pds(&self.client, did).await?; + records::list_all_records(&self.client, &pds, did, "com.atproto.lexicon.schema") + .await + .map_err(Into::into) + } } impl Default for RealHttpClient { @@ -234,20 +91,15 @@ impl Default for RealHttpClient { #[async_trait] impl HttpClient for RealHttpClient { async fn fetch_lexicon(&self, did: &str, nsid: &str) -> Result { - // Fetch all records from the DID's repo - let records = self.fetch_records_from_did(did).await?; - - // Find the specific NSID + let records = self.fetch_all_schema_records(did).await?; for record in records { let record_nsid = extract_nsid_from_record(&record)?; if record_nsid == nsid { return Ok(record.value); } } - Err(LexiconFetcherError::LexiconNotFound(format!( - "Lexicon {} not found in repo {}", - nsid, did + "Lexicon {nsid} not found in repo {did}" ))) } @@ -256,205 +108,34 @@ impl HttpClient for RealHttpClient { did: &str, pattern: &str, ) -> Result> { - // Fetch all records from the DID's repo - let records = self.fetch_records_from_did(did).await?; - + let records = self.fetch_all_schema_records(did).await?; let mut results = Vec::new(); - - // Handle exact match (no wildcard) - if !pattern.contains('*') && !pattern.contains('_') { - for record in records { - let record_nsid = extract_nsid_from_record(&record)?; - if record_nsid == pattern { - results.push((record_nsid, record.value)); - } - } - return Ok(results); - } - - // Handle wildcard patterns - if pattern.ends_with(".*") { - // "*" matches EVERYTHING - // For "place.stream.*", match all: place.stream.chat, place.stream.chat.profile, etc. - let base = pattern.strip_suffix(".*").unwrap(); - let prefix_with_dot = format!("{}.", base); - - for record in records { - let record_nsid = extract_nsid_from_record(&record)?; - if record_nsid.starts_with(&prefix_with_dot) { - results.push((record_nsid, record.value)); - } - } - } else if pattern.ends_with("._") { - // "_" matches only DIRECT CHILDREN - // For "place.stream._", match place.stream.chat but NOT place.stream.chat.profile - let base = pattern.strip_suffix("._").unwrap(); - let prefix_with_dot = format!("{}.", base); - - for record in records { - let record_nsid = extract_nsid_from_record(&record)?; - - if let Some(suffix) = record_nsid.strip_prefix(&prefix_with_dot) { - // Check if it's a direct child (no more dots in the suffix) - if !suffix.contains('.') && !suffix.is_empty() { - results.push((record_nsid, record.value)); - } - } + for record in records { + let record_nsid = extract_nsid_from_record(&record)?; + if nsid_matches_pattern(&record_nsid, pattern) { + results.push((record_nsid, record.value)); } } - Ok(results) } } -impl RealHttpClient { - /// Fetch all lexicon records from a DID's ATProto repository - async fn fetch_records_from_did(&self, did: &str) -> Result> { - // Resolve DID to PDS URL - let pds_url = self.resolve_did_to_pds(did).await?; - - let mut all_records = Vec::new(); - let mut cursor: Option = None; - - // Paginate through all records - loop { - let url = if let Some(ref c) = cursor { - format!( - "{}/xrpc/com.atproto.repo.listRecords?repo={}&collection=com.atproto.lexicon.schema&cursor={}", - pds_url, did, c - ) - } else { - format!( - "{}/xrpc/com.atproto.repo.listRecords?repo={}&collection=com.atproto.lexicon.schema", - pds_url, did - ) - }; - - let response = self - .client - .get(&url) - .send() - .await - .map_err(|e| LexiconFetcherError::HttpRequestFailed(e.to_string()))?; - - if !response.status().is_success() { - return Err(LexiconFetcherError::HttpRequestFailed(format!( - "HTTP {} when fetching records from {}", - response.status(), - did - ))); - } - - let mut list_response: serde_json::Value = response - .json() - .await - .map_err(|e| LexiconFetcherError::JsonParseFailed(e.to_string()))?; - - // Extract records - if let Some(records_array) = list_response.get_mut("records") { - if let Some(records) = records_array.as_array_mut() { - for record_value in records.drain(..) { - let record: AtProtoRecord = serde_json::from_value(record_value) - .map_err(|e| LexiconFetcherError::JsonParseFailed(format!("Failed to parse record: {}", e)))?; - all_records.push(record); - } - } - } - - // Check for pagination cursor - cursor = list_response - .get("cursor") - .and_then(|c| c.as_str()) - .map(|s| s.to_string()); - - if cursor.is_none() { - break; - } - } - - Ok(all_records) - } - - /// Resolve a DID to its PDS URL - async fn resolve_did_to_pds(&self, did: &str) -> Result { - // For did:web:, extract the domain - if let Some(domain) = did.strip_prefix("did:web:") { - return Ok(format!("https://{}", domain)); - } - - // For did:plc:, query the PLC directory - if did.starts_with("did:plc:") { - let url = format!("https://plc.directory/{}", did); - - let response = self - .client - .get(&url) - .send() - .await - .map_err(|e| LexiconFetcherError::HttpRequestFailed(format!("Failed to resolve DID: {}", e)))?; - - if !response.status().is_success() { - return Err(LexiconFetcherError::HttpRequestFailed(format!( - "Failed to resolve DID {}: HTTP {}", - did, - response.status() - ))); - } - - let did_doc: serde_json::Value = response - .json() - .await - .map_err(|e| LexiconFetcherError::JsonParseFailed(format!("Failed to parse DID document: {}", e)))?; - - // Extract PDS endpoint from service array - if let Some(services) = did_doc.get("service").and_then(|v| v.as_array()) { - for service in services { - if service.get("type").and_then(|v| v.as_str()) == Some("AtprotoPersonalDataServer") { - if let Some(endpoint) = service.get("serviceEndpoint").and_then(|v| v.as_str()) { - return Ok(endpoint.trim_end_matches('/').to_string()); - } - } - } - } - - return Err(LexiconFetcherError::HttpRequestFailed(format!( - "No PDS endpoint found in DID document for {}", - did - ))); - } - - Err(LexiconFetcherError::InvalidUrl(format!( - "Unsupported DID format: {}", - did - ))) - } -} - -/// Mock HTTP client for testing -#[derive(Clone)] +/// Mock HTTP client for testing. +#[derive(Clone, Default)] pub struct MockHttpClient { lexicons: Arc>>, } impl MockHttpClient { pub fn new() -> Self { - Self { - lexicons: Arc::new(Mutex::new(HashMap::new())), - } + Self::default() } - /// Add a mock lexicon response for a specific NSID pub fn add_lexicon(&mut self, nsid: String, lexicon: serde_json::Value) { self.lexicons.lock().unwrap().insert(nsid, lexicon); } } -impl Default for MockHttpClient { - fn default() -> Self { - Self::new() - } -} - #[async_trait] impl HttpClient for MockHttpClient { async fn fetch_lexicon(&self, _did: &str, nsid: &str) -> Result { @@ -473,48 +154,55 @@ impl HttpClient for MockHttpClient { ) -> Result> { let lexicons = self.lexicons.lock().unwrap(); let mut results = Vec::new(); - - // Handle exact match (no wildcard) - if !pattern.contains('*') && !pattern.contains('_') { - if let Some(lexicon) = lexicons.get(pattern) { - results.push((pattern.to_string(), lexicon.clone())); + for (nsid, lexicon) in lexicons.iter() { + if nsid_matches_pattern(nsid, pattern) { + results.push((nsid.clone(), lexicon.clone())); } - return Ok(results); } + Ok(results) + } +} - // Handle wildcard patterns - if pattern.ends_with(".*") { - // "*" matches EVERYTHING - // For "place.stream.*", match all: place.stream.chat, place.stream.chat.profile, etc. - let base = pattern.strip_suffix(".*").unwrap(); - let prefix_with_dot = format!("{}.", base); - - for (nsid, lexicon) in lexicons.iter() { - if nsid.starts_with(&prefix_with_dot) { - results.push((nsid.clone(), lexicon.clone())); - } - } - } else if pattern.ends_with("._") { - // "_" matches only DIRECT CHILDREN - // For "place.stream._", match place.stream.chat but NOT place.stream.chat.profile - let base = pattern.strip_suffix("._").unwrap(); - let prefix_with_dot = format!("{}.", base); - - for (nsid, lexicon) in lexicons.iter() { - if let Some(suffix) = nsid.strip_prefix(&prefix_with_dot) { - // Check if it's a direct child (no more dots in the suffix) - if !suffix.contains('.') && !suffix.is_empty() { - results.push((nsid.clone(), lexicon.clone())); - } - } - } +/// Check whether a concrete NSID matches an optional wildcard pattern. +/// +/// - Exact (no wildcard): `nsid == pattern` +/// - `"foo.bar.*"` matches any NSID beginning with `"foo.bar."` +/// - `"foo.bar._"` matches *direct children only* of `foo.bar` — +/// the suffix after `foo.bar.` must contain no further dots. +fn nsid_matches_pattern(nsid: &str, pattern: &str) -> bool { + if !pattern.contains('*') && !pattern.contains('_') { + return nsid == pattern; + } + if let Some(base) = pattern.strip_suffix(".*") { + let prefix = format!("{base}."); + return nsid.starts_with(&prefix); + } + if let Some(base) = pattern.strip_suffix("._") { + let prefix = format!("{base}."); + if let Some(suffix) = nsid.strip_prefix(&prefix) { + return !suffix.is_empty() && !suffix.contains('.'); } + } + false +} - Ok(results) +/// Extract the NSID from a fetched `com.atproto.lexicon.schema` record. +/// +/// Prefers the record's `id` field; falls back to the rkey in the URI. +fn extract_nsid_from_record(record: &records::Record) -> Result { + if let Some(id) = record.value.get("id").and_then(|v| v.as_str()) { + return Ok(id.to_string()); + } + if let Some(rkey) = record.uri.split('/').next_back() { + return Ok(rkey.to_string()); } + Err(LexiconFetcherError::MalformedRecord(record.uri.clone())) } -/// Main lexicon fetcher that combines DNS resolution and HTTP fetching +// --------------------------------------------------------------------------- +// Fetcher combining DNS + HTTP +// --------------------------------------------------------------------------- + pub struct LexiconFetcher { dns_resolver: D, http_client: H, @@ -528,57 +216,49 @@ impl LexiconFetcher { } } - /// Fetch a single lexicon by NSID - /// Example: "place.stream.chat.profile" -> single lexicon JSON + /// Fetch a single lexicon by exact NSID. pub async fn fetch(&self, nsid: &str) -> Result { - // Check if this is a wildcard pattern if nsid.contains('*') || nsid.contains('_') { return Err(LexiconFetcherError::InvalidNsid(format!( - "Use fetch_pattern() for wildcard patterns (* or _): {}", - nsid + "Use fetch_pattern() for wildcard patterns (* or _): {nsid}" ))); } - - // Parse NSID into authority and name segments let (authority, name_segments) = parse_nsid(nsid)?; - - // Resolve DID via DNS (async) - let did = self.dns_resolver.resolve_lexicon_did(&authority, &name_segments).await?; - - // Fetch lexicon via HTTP + let did = self + .dns_resolver + .resolve_lexicon_did(&authority, &name_segments) + .await?; self.http_client.fetch_lexicon(&did, nsid).await } - /// Fetch all lexicons matching a pattern - /// Examples: - /// - "place.stream.*" -> matches everything (place.stream.chat, place.stream.chat.profile, etc.) - /// - "place.stream._" -> matches direct children only (place.stream.chat, place.stream.key, but not place.stream.chat.profile) + /// Fetch all lexicons matching a wildcard pattern (`…*` or `…_`). pub async fn fetch_pattern(&self, pattern: &str) -> Result> { - // Parse pattern to extract authority let (authority, name_pattern) = parse_nsid(pattern)?; - - // For DNS lookup, remove wildcard suffix (both .* and ._) - // For "place.stream.*" or "place.stream._", name_pattern is "*" or "_", so we use empty string - // For "place.stream.chat.*", name_pattern is "chat.*", so we use "chat" - let dns_name_segments = if name_pattern == "*" || name_pattern == "_" { - "" - } else if let Some(pos) = name_pattern.rfind(".*") { - &name_pattern[..pos] - } else if let Some(pos) = name_pattern.rfind("._") { - &name_pattern[..pos] - } else { - &name_pattern - }; - - // Resolve DID via DNS (async) - let did = self.dns_resolver.resolve_lexicon_did(&authority, dns_name_segments).await?; - - // Fetch lexicons matching pattern via HTTP + let dns_name_segments = strip_wildcard(&name_pattern); + let did = self + .dns_resolver + .resolve_lexicon_did(&authority, dns_name_segments) + .await?; self.http_client.fetch_lexicons_pattern(&did, pattern).await } } -/// Metadata about a fetched lexicon +/// For a name-pattern like `chat.*` or `chat._`, return `"chat"`. For +/// the bare `*` / `_`, return `""`. Otherwise return the pattern itself. +fn strip_wildcard(name_pattern: &str) -> &str { + if name_pattern == "*" || name_pattern == "_" { + return ""; + } + if let Some(pos) = name_pattern.rfind(".*") { + return &name_pattern[..pos]; + } + if let Some(pos) = name_pattern.rfind("._") { + return &name_pattern[..pos]; + } + name_pattern +} + +/// Metadata about a fetched lexicon. #[derive(Debug, Clone)] pub struct FetchedLexicon { pub nsid: String, @@ -586,44 +266,32 @@ pub struct FetchedLexicon { pub did: String, } -/// Result of fetching one or more lexicons +/// Result of fetching one or more lexicons. #[derive(Debug)] pub struct FetchResult { pub lexicons: Vec, } -/// Convenience type for production use with real DNS and HTTP pub type ProductionLexiconFetcher = LexiconFetcher; impl ProductionLexiconFetcher { - /// Create a new production fetcher with default configuration + /// Create a fetcher with default DNS + HTTP configuration. pub async fn production() -> Result { - Ok(Self::new(RealDnsResolver::new().await?, RealHttpClient::new())) + Ok(Self::new(RealDnsResolver::new()?, RealHttpClient::new())) } } impl LexiconFetcher { - /// Fetch one or more lexicons and return metadata for lockfile tracking - /// Handles both exact NSIDs and patterns (* or _) + /// Fetch with metadata (NSID + DID) for lockfile tracking. + /// Handles both exact NSIDs and wildcard patterns. pub async fn fetch_with_metadata(&self, nsid: &str) -> Result { - // Parse pattern to extract authority and name segments let (authority, name_segments) = parse_nsid(nsid)?; + let dns_name_segments = strip_wildcard(&name_segments); + let did = self + .dns_resolver + .resolve_lexicon_did(&authority, dns_name_segments) + .await?; - // For DNS lookup, remove wildcard suffix (both .* and ._) - let dns_name_segments = if name_segments == "*" || name_segments == "_" { - "" - } else if let Some(pos) = name_segments.rfind(".*") { - &name_segments[..pos] - } else if let Some(pos) = name_segments.rfind("._") { - &name_segments[..pos] - } else { - &name_segments - }; - - // Resolve DID via DNS (async) - let did = self.dns_resolver.resolve_lexicon_did(&authority, dns_name_segments).await?; - - // Fetch lexicons let lexicons = if nsid.contains('*') || nsid.contains('_') { self.http_client.fetch_lexicons_pattern(&did, nsid).await? } else { @@ -631,7 +299,6 @@ impl LexiconFetcher { vec![(nsid.to_string(), lexicon)] }; - // Package results with metadata let fetched_lexicons = lexicons .into_iter() .map(|(nsid, lexicon)| FetchedLexicon { @@ -646,7 +313,7 @@ impl LexiconFetcher { }) } - /// Fetch multiple NSIDs in sequence (no optimization) + /// Fetch multiple NSIDs sequentially. pub async fn fetch_many(&self, nsids: &[String]) -> Result> { let mut results = Vec::new(); for nsid in nsids { @@ -655,44 +322,30 @@ impl LexiconFetcher { Ok(results) } - /// Fetch multiple NSIDs with optimization to reduce network requests - /// Groups similar NSIDs into wildcard patterns when beneficial - /// Example: ["app.bsky.actor.foo", "app.bsky.actor.bar"] -> fetches "app.bsky.actor.*" + /// Fetch multiple NSIDs, collapsing them into wildcard patterns where + /// beneficial to reduce network round-trips. pub async fn fetch_many_optimized(&self, nsids: &[String]) -> Result> { use std::collections::HashSet; - if nsids.is_empty() { return Ok(Vec::new()); } - - // Convert to HashSet for optimization let nsids_set: HashSet = nsids.iter().cloned().collect(); - - // Optimize into minimal set of patterns let optimized_patterns = optimize_fetch_patterns(&nsids_set); - - // Fetch each optimized pattern let mut results = Vec::new(); for pattern in optimized_patterns { results.push(self.fetch_with_metadata(&pattern).await?); } - Ok(results) } - /// Fetch lexicon(s) from a known DID, bypassing DNS resolution - /// Useful when fetching from lockfile where DID is already known - /// Handles both exact NSIDs and patterns (* or _) + /// Fetch from a known DID, skipping DNS resolution (for lockfile replay). pub async fn fetch_from_did_with_metadata(&self, did: &str, nsid: &str) -> Result { - // Fetch lexicons directly from the DID let lexicons = if nsid.contains('*') || nsid.contains('_') { self.http_client.fetch_lexicons_pattern(did, nsid).await? } else { let lexicon = self.http_client.fetch_lexicon(did, nsid).await?; vec![(nsid.to_string(), lexicon)] }; - - // Package results with metadata let fetched_lexicons = lexicons .into_iter() .map(|(nsid, lexicon)| FetchedLexicon { @@ -701,112 +354,81 @@ impl LexiconFetcher { did: did.to_string(), }) .collect(); - Ok(FetchResult { lexicons: fetched_lexicons, }) } } -/// Convenience type for testing with mocks pub type MockLexiconFetcher = LexiconFetcher; impl MockLexiconFetcher { - /// Create a new mock fetcher for testing pub fn mock() -> Self { Self::new(MockDnsResolver::new(), MockHttpClient::new()) } } -/// Extract NSID from an ATProto record -fn extract_nsid_from_record(record: &AtProtoRecord) -> Result { - // The record value should have an "id" field with the NSID - if let Some(id) = record.value.get("id").and_then(|v| v.as_str()) { - return Ok(id.to_string()); - } - - // Fallback: try to extract from URI - // URI format: at://did:plc:xxx/com.atproto.lexicon.schema/nsid - if let Some(rkey) = record.uri.split('/').last() { - return Ok(rkey.to_string()); - } - - Err(LexiconFetcherError::HttpRequestFailed(format!( - "Could not extract NSID from record: {}", - record.uri - ))) -} - -/// Optimize a set of NSIDs by collapsing them into the minimal set of fetch patterns -/// For example: ["app.bsky.actor.foo", "app.bsky.actor.bar"] -> ["app.bsky.actor.*"] -/// This function tries multiple grouping strategies to find the most efficient pattern +/// Collapse a set of NSIDs into the minimum number of fetch patterns. +/// +/// - Two or more NSIDs that share a complete prefix (all but the final +/// segment) become `.*`. +/// - Three or more NSIDs under the same authority that weren't otherwise +/// grouped become `.*`. +/// - Anything left over is emitted exact. pub fn optimize_fetch_patterns(nsids: &std::collections::HashSet) -> Vec { use std::collections::{BTreeMap, HashSet}; - if nsids.is_empty() { return Vec::new(); } - // Strategy 1: Try grouping by authority (first 2 segments) - // e.g., ["app.bsky.actor.foo", "app.bsky.feed.bar"] -> ["app.bsky.*"] let mut authority_groups: BTreeMap> = BTreeMap::new(); - + let mut prefix_groups: BTreeMap> = BTreeMap::new(); for nsid in nsids { let parts: Vec<&str> = nsid.split('.').collect(); if parts.len() >= 2 { let authority = format!("{}.{}", parts[0], parts[1]); - authority_groups.entry(authority).or_insert_with(Vec::new).push(nsid.clone()); + authority_groups + .entry(authority) + .or_default() + .push(nsid.clone()); } - } - - // Strategy 2: Try grouping by namespace prefix (all but last segment) - // e.g., ["app.bsky.actor.foo", "app.bsky.actor.bar"] -> ["app.bsky.actor.*"] - let mut prefix_groups: BTreeMap> = BTreeMap::new(); - - for nsid in nsids { - let parts: Vec<&str> = nsid.split('.').collect(); if parts.len() >= 3 { let prefix = parts[..parts.len() - 1].join("."); - prefix_groups.entry(prefix).or_insert_with(Vec::new).push(nsid.clone()); + prefix_groups.entry(prefix).or_default().push(nsid.clone()); } } let mut result = Vec::new(); - let mut handled_nsids = HashSet::new(); + let mut handled: HashSet = HashSet::new(); - // First pass: Apply namespace-level grouping (more specific) + // Specific-prefix wildcards first (stricter match, fewer false positives). for (prefix, group) in &prefix_groups { - if group.len() >= 2 && !handled_nsids.contains(&group[0]) { - result.push(format!("{}.*", prefix)); + if group.len() >= 2 && !handled.contains(&group[0]) { + result.push(format!("{prefix}.*")); for nsid in group { - handled_nsids.insert(nsid.clone()); + handled.insert(nsid.clone()); } } } - // Second pass: For remaining NSIDs, consider authority-level grouping - // Only use authority wildcard if we have 3+ different namespaces under same authority + // Authority-level wildcards only if we'd save 3+ calls. for (authority, group) in &authority_groups { - let unhandled: Vec<&String> = group.iter() - .filter(|nsid| !handled_nsids.contains(*nsid)) - .collect(); - + let unhandled: Vec<&String> = group.iter().filter(|n| !handled.contains(*n)).collect(); if unhandled.len() >= 3 { - result.push(format!("{}.*", authority)); - for nsid in &unhandled { - handled_nsids.insert((*nsid).clone()); + result.push(format!("{authority}.*")); + for nsid in unhandled { + handled.insert(nsid.clone()); } } } - // Third pass: Add remaining individual NSIDs + // Emit remaining NSIDs exact. for nsid in nsids { - if !handled_nsids.contains(nsid) { + if !handled.contains(nsid) { result.push(nsid.clone()); } } - // Sort for consistent output result.sort(); result } @@ -816,65 +438,30 @@ mod tests { use super::*; #[test] - fn test_construct_dns_name() { - assert_eq!( - construct_dns_name("place.stream", "key"), - "_lexicon.key.stream.place" - ); - assert_eq!( - construct_dns_name("app.bsky", "actor"), - "_lexicon.actor.bsky.app" - ); - assert_eq!( - construct_dns_name("app.bsky", "actor.profile"), - "_lexicon.actor.profile.bsky.app" - ); - assert_eq!( - construct_dns_name("place.stream", ""), - "_lexicon.stream.place" - ); + fn nsid_matches_pattern_exact() { + assert!(nsid_matches_pattern("foo.bar.baz", "foo.bar.baz")); + assert!(!nsid_matches_pattern("foo.bar.baz", "foo.bar.qux")); } #[test] - fn test_parse_nsid() { - let (auth, name) = parse_nsid("place.stream.key").unwrap(); - assert_eq!(auth, "place.stream"); - assert_eq!(name, "key"); - - let (auth, name) = parse_nsid("app.bsky.actor.profile").unwrap(); - assert_eq!(auth, "app.bsky"); - assert_eq!(name, "actor.profile"); - - let (auth, name) = parse_nsid("place.stream").unwrap(); - assert_eq!(auth, "place.stream"); - assert_eq!(name, ""); - - assert!(parse_nsid("invalid").is_err()); + fn nsid_matches_pattern_wildcard_star() { + assert!(nsid_matches_pattern("foo.bar.baz", "foo.bar.*")); + assert!(nsid_matches_pattern("foo.bar.baz.deep", "foo.bar.*")); + assert!(!nsid_matches_pattern("foo.qux.baz", "foo.bar.*")); } - #[tokio::test] - async fn test_mock_dns_resolver() { - let mut resolver = MockDnsResolver::new(); - resolver.add_record("place.stream", "key", "did:plc:test123".to_string()); - - let did = resolver.resolve_lexicon_did("place.stream", "key").await.unwrap(); - assert_eq!(did, "did:plc:test123"); - - let result = resolver.resolve_lexicon_did("place.stream", "notfound").await; - assert!(result.is_err()); + #[test] + fn nsid_matches_pattern_direct_child_only() { + assert!(nsid_matches_pattern("foo.bar.baz", "foo.bar._")); + assert!(!nsid_matches_pattern("foo.bar.baz.deep", "foo.bar._")); } - #[tokio::test] - async fn test_mock_dns_resolver_from_nsid() { - let mut resolver = MockDnsResolver::new(); - resolver - .add_record_from_nsid("app.bsky.actor.profile", "did:plc:bsky123".to_string()) - .unwrap(); - - let did = resolver - .resolve_lexicon_did("app.bsky", "actor.profile") - .await - .unwrap(); - assert_eq!(did, "did:plc:bsky123"); + #[test] + fn strip_wildcard_handles_all_forms() { + assert_eq!(strip_wildcard("*"), ""); + assert_eq!(strip_wildcard("_"), ""); + assert_eq!(strip_wildcard("chat.*"), "chat"); + assert_eq!(strip_wildcard("chat._"), "chat"); + assert_eq!(strip_wildcard("chat"), "chat"); } } diff --git a/mlf-lexicon-fetcher/tests/dns_scenarios.rs b/mlf-lexicon-fetcher/tests/dns_scenarios.rs index 0166a98..742d0b7 100644 --- a/mlf-lexicon-fetcher/tests/dns_scenarios.rs +++ b/mlf-lexicon-fetcher/tests/dns_scenarios.rs @@ -15,9 +15,16 @@ async fn test_successful_lookup() { #[tokio::test] async fn test_dns_record_not_found() { let resolver = MockDnsResolver::new(); - let result = resolver.resolve_lexicon_did("nonexistent.domain", "test").await; + let result = resolver + .resolve_lexicon_did("nonexistent.domain", "test") + .await; assert!(result.is_err()); - assert!(result.unwrap_err().to_string().contains("No mock record found")); + assert!( + result + .unwrap_err() + .to_string() + .contains("No mock record found") + ); } #[tokio::test] @@ -34,15 +41,24 @@ async fn test_multiple_authority_types() { resolver.add_record("com.atproto", "repo", "did:plc:atproto001".to_string()); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "actor").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "actor") + .await + .unwrap(), "did:plc:bsky001" ); assert_eq!( - resolver.resolve_lexicon_did("place.stream", "chat").await.unwrap(), + resolver + .resolve_lexicon_did("place.stream", "chat") + .await + .unwrap(), "did:plc:stream001" ); assert_eq!( - resolver.resolve_lexicon_did("com.atproto", "repo").await.unwrap(), + resolver + .resolve_lexicon_did("com.atproto", "repo") + .await + .unwrap(), "did:plc:atproto001" ); } @@ -58,18 +74,31 @@ async fn test_nested_name_segments() { resolver.add_record("app.bsky", "actor.profile", "did:plc:double".to_string()); // Three segments - resolver.add_record("app.bsky", "actor.profile.detailed", "did:plc:triple".to_string()); + resolver.add_record( + "app.bsky", + "actor.profile.detailed", + "did:plc:triple".to_string(), + ); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "actor").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "actor") + .await + .unwrap(), "did:plc:single" ); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "actor.profile").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "actor.profile") + .await + .unwrap(), "did:plc:double" ); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "actor.profile.detailed").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "actor.profile.detailed") + .await + .unwrap(), "did:plc:triple" ); } @@ -82,7 +111,10 @@ async fn test_empty_name_segments() { resolver.add_record("place.stream", "", "did:plc:root".to_string()); assert_eq!( - resolver.resolve_lexicon_did("place.stream", "").await.unwrap(), + resolver + .resolve_lexicon_did("place.stream", "") + .await + .unwrap(), "did:plc:root" ); } @@ -94,7 +126,10 @@ async fn test_did_web_format() { resolver.add_record("example.com", "api", "did:web:example.com".to_string()); assert_eq!( - resolver.resolve_lexicon_did("example.com", "api").await.unwrap(), + resolver + .resolve_lexicon_did("example.com", "api") + .await + .unwrap(), "did:web:example.com" ); } @@ -107,10 +142,13 @@ async fn test_did_plc_format() { resolver.add_record( "app.bsky", "feed", - "did:plc:z72i7hdynmk6r22z27h6tvur".to_string() + "did:plc:z72i7hdynmk6r22z27h6tvur".to_string(), ); - let did = resolver.resolve_lexicon_did("app.bsky", "feed").await.unwrap(); + let did = resolver + .resolve_lexicon_did("app.bsky", "feed") + .await + .unwrap(); assert!(did.starts_with("did:plc:")); assert_eq!(did.len(), 32); // "did:plc:" (8) + 24 chars } @@ -120,15 +158,25 @@ async fn test_add_record_from_nsid() { let mut resolver = MockDnsResolver::new(); // Add using full NSID - resolver.add_record_from_nsid("place.stream.key", "did:plc:test".to_string()).unwrap(); - resolver.add_record_from_nsid("app.bsky.actor.profile", "did:plc:bsky".to_string()).unwrap(); + resolver + .add_record_from_nsid("place.stream.key", "did:plc:test".to_string()) + .unwrap(); + resolver + .add_record_from_nsid("app.bsky.actor.profile", "did:plc:bsky".to_string()) + .unwrap(); assert_eq!( - resolver.resolve_lexicon_did("place.stream", "key").await.unwrap(), + resolver + .resolve_lexicon_did("place.stream", "key") + .await + .unwrap(), "did:plc:test" ); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "actor.profile").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "actor.profile") + .await + .unwrap(), "did:plc:bsky" ); } @@ -140,7 +188,12 @@ async fn test_invalid_nsid_format() { // NSID with only one segment let result = resolver.add_record_from_nsid("invalid", "did:plc:test".to_string()); assert!(result.is_err()); - assert!(result.unwrap_err().to_string().contains("at least 2 segments")); + assert!( + result + .unwrap_err() + .to_string() + .contains("at least 2 segments") + ); } #[tokio::test] @@ -153,11 +206,17 @@ async fn test_case_sensitivity() { // These should be treated as different domains assert_eq!( - resolver.resolve_lexicon_did("App.Bsky", "Actor").await.unwrap(), + resolver + .resolve_lexicon_did("App.Bsky", "Actor") + .await + .unwrap(), "did:plc:uppercase" ); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "actor").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "actor") + .await + .unwrap(), "did:plc:lowercase" ); } @@ -176,7 +235,10 @@ async fn test_concurrent_lookups() { for _ in 0..10 { let resolver_clone = Arc::clone(&resolver); let handle = tokio::spawn(async move { - resolver_clone.resolve_lexicon_did("app.bsky", "feed").await.unwrap() + resolver_clone + .resolve_lexicon_did("app.bsky", "feed") + .await + .unwrap() }); handles.push(handle); } @@ -199,11 +261,17 @@ async fn test_wildcard_namespace_scenarios() { // All should resolve to the same DID assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "actor.defs").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "actor.defs") + .await + .unwrap(), "did:plc:bsky" ); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", "feed.post").await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", "feed.post") + .await + .unwrap(), "did:plc:bsky" ); } @@ -253,7 +321,10 @@ async fn test_edge_case_empty_did() { resolver.add_record("test.com", "api", "".to_string()); assert_eq!( - resolver.resolve_lexicon_did("test.com", "api").await.unwrap(), + resolver + .resolve_lexicon_did("test.com", "api") + .await + .unwrap(), "" ); } @@ -267,7 +338,10 @@ async fn test_very_long_nsid() { resolver.add_record("app.bsky", long_name, "did:plc:deep".to_string()); assert_eq!( - resolver.resolve_lexicon_did("app.bsky", long_name).await.unwrap(), + resolver + .resolve_lexicon_did("app.bsky", long_name) + .await + .unwrap(), "did:plc:deep" ); } @@ -280,7 +354,10 @@ async fn test_special_characters_in_nsid() { resolver.add_record("app-test.bsky-123", "actor", "did:plc:special".to_string()); assert_eq!( - resolver.resolve_lexicon_did("app-test.bsky-123", "actor").await.unwrap(), + resolver + .resolve_lexicon_did("app-test.bsky-123", "actor") + .await + .unwrap(), "did:plc:special" ); } diff --git a/mlf-lexicon-fetcher/tests/lexicon_fetching.rs b/mlf-lexicon-fetcher/tests/lexicon_fetching.rs index 82110c6..f164901 100644 --- a/mlf-lexicon-fetcher/tests/lexicon_fetching.rs +++ b/mlf-lexicon-fetcher/tests/lexicon_fetching.rs @@ -1,15 +1,17 @@ // Integration tests for full lexicon fetching flow (DNS + HTTP) -use mlf_lexicon_fetcher::{ - LexiconFetcher, MockDnsResolver, MockHttpClient, LexiconFetcherError, -}; +use mlf_lexicon_fetcher::{LexiconFetcher, LexiconFetcherError, MockDnsResolver, MockHttpClient}; use serde_json::json; #[tokio::test] async fn test_fetch_single_lexicon() { // Setup mock DNS resolver let mut dns_resolver = MockDnsResolver::new(); - dns_resolver.add_record("place.stream", "chat.profile", "did:plc:test123".to_string()); + dns_resolver.add_record( + "place.stream", + "chat.profile", + "did:plc:test123".to_string(), + ); // Setup mock HTTP client let mut http_client = MockHttpClient::new(); @@ -23,7 +25,10 @@ async fn test_fetch_single_lexicon() { } } }); - http_client.add_lexicon("place.stream.chat.profile".to_string(), lexicon_json.clone()); + http_client.add_lexicon( + "place.stream.chat.profile".to_string(), + lexicon_json.clone(), + ); // Create fetcher let fetcher = LexiconFetcher::new(dns_resolver, http_client); @@ -32,7 +37,10 @@ async fn test_fetch_single_lexicon() { let result = fetcher.fetch("place.stream.chat.profile").await; assert!(result.is_ok()); let fetched = result.unwrap(); - assert_eq!(fetched.get("id").unwrap().as_str().unwrap(), "place.stream.chat.profile"); + assert_eq!( + fetched.get("id").unwrap().as_str().unwrap(), + "place.stream.chat.profile" + ); } #[tokio::test] @@ -86,7 +94,11 @@ async fn test_fetch_pattern_multiple_lexicons() { async fn test_fetch_lexicon_not_found() { // Setup mock DNS resolver let mut dns_resolver = MockDnsResolver::new(); - dns_resolver.add_record("place.stream", "chat.profile", "did:plc:test123".to_string()); + dns_resolver.add_record( + "place.stream", + "chat.profile", + "did:plc:test123".to_string(), + ); // Setup mock HTTP client (but don't add the lexicon) let http_client = MockHttpClient::new(); @@ -122,10 +134,12 @@ async fn test_fetch_dns_lookup_failed() { assert!(result.is_err()); match result { - Err(LexiconFetcherError::LookupFailed { domain, .. }) => { + Err(LexiconFetcherError::Identity( + mlf_atproto::identity::IdentityError::DnsLookupFailed { domain, .. }, + )) => { assert_eq!(domain, "_lexicon.chat.profile.stream.place"); } - _ => panic!("Expected LookupFailed error"), + other => panic!("Expected DnsLookupFailed, got {other:?}"), } } @@ -171,7 +185,11 @@ async fn test_fetch_pattern_empty_results() { async fn test_multiple_authorities() { // Setup mock DNS resolver with multiple authorities let mut dns_resolver = MockDnsResolver::new(); - dns_resolver.add_record("place.stream", "chat.profile", "did:plc:stream123".to_string()); + dns_resolver.add_record( + "place.stream", + "chat.profile", + "did:plc:stream123".to_string(), + ); dns_resolver.add_record("app.bsky", "actor.profile", "did:plc:bsky456".to_string()); // Setup mock HTTP client @@ -217,7 +235,9 @@ async fn test_nested_name_segments() { let fetcher = LexiconFetcher::new(dns_resolver, http_client); // Fetch deeply nested lexicon - let result = fetcher.fetch("place.stream.chat.message.attachments.image").await; + let result = fetcher + .fetch("place.stream.chat.message.attachments.image") + .await; assert!(result.is_ok()); } @@ -228,7 +248,11 @@ async fn test_concurrent_fetches() { // Setup mock DNS resolver let mut dns_resolver = MockDnsResolver::new(); - dns_resolver.add_record("place.stream", "chat.profile", "did:plc:test123".to_string()); + dns_resolver.add_record( + "place.stream", + "chat.profile", + "did:plc:test123".to_string(), + ); // Setup mock HTTP client let mut http_client = MockHttpClient::new(); @@ -244,9 +268,8 @@ async fn test_concurrent_fetches() { let mut handles = vec![]; for _ in 0..10 { let fetcher_clone = Arc::clone(&fetcher); - let handle = task::spawn(async move { - fetcher_clone.fetch("place.stream.chat.profile").await - }); + let handle = + task::spawn(async move { fetcher_clone.fetch("place.stream.chat.profile").await }); handles.push(handle); } @@ -265,10 +288,22 @@ async fn test_pattern_prefix_matching() { // Setup mock HTTP client let mut http_client = MockHttpClient::new(); - http_client.add_lexicon("app.bsky.feed.post".to_string(), json!({"id": "app.bsky.feed.post"})); - http_client.add_lexicon("app.bsky.feed.like".to_string(), json!({"id": "app.bsky.feed.like"})); - http_client.add_lexicon("app.bsky.feed.repost".to_string(), json!({"id": "app.bsky.feed.repost"})); - http_client.add_lexicon("app.bsky.actor.profile".to_string(), json!({"id": "app.bsky.actor.profile"})); + http_client.add_lexicon( + "app.bsky.feed.post".to_string(), + json!({"id": "app.bsky.feed.post"}), + ); + http_client.add_lexicon( + "app.bsky.feed.like".to_string(), + json!({"id": "app.bsky.feed.like"}), + ); + http_client.add_lexicon( + "app.bsky.feed.repost".to_string(), + json!({"id": "app.bsky.feed.repost"}), + ); + http_client.add_lexicon( + "app.bsky.actor.profile".to_string(), + json!({"id": "app.bsky.actor.profile"}), + ); // Create fetcher let fetcher = LexiconFetcher::new(dns_resolver, http_client); @@ -293,16 +328,34 @@ async fn test_underscore_wildcard_direct_children() { // Setup mock HTTP client with nested lexicons let mut http_client = MockHttpClient::new(); - + // Direct children of place.stream - http_client.add_lexicon("place.stream.chat".to_string(), json!({"id": "place.stream.chat"})); - http_client.add_lexicon("place.stream.key".to_string(), json!({"id": "place.stream.key"})); - http_client.add_lexicon("place.stream.livestream".to_string(), json!({"id": "place.stream.livestream"})); - + http_client.add_lexicon( + "place.stream.chat".to_string(), + json!({"id": "place.stream.chat"}), + ); + http_client.add_lexicon( + "place.stream.key".to_string(), + json!({"id": "place.stream.key"}), + ); + http_client.add_lexicon( + "place.stream.livestream".to_string(), + json!({"id": "place.stream.livestream"}), + ); + // Nested children (should NOT match with _) - http_client.add_lexicon("place.stream.chat.profile".to_string(), json!({"id": "place.stream.chat.profile"})); - http_client.add_lexicon("place.stream.chat.message".to_string(), json!({"id": "place.stream.chat.message"})); - http_client.add_lexicon("place.stream.key.defs".to_string(), json!({"id": "place.stream.key.defs"})); + http_client.add_lexicon( + "place.stream.chat.profile".to_string(), + json!({"id": "place.stream.chat.profile"}), + ); + http_client.add_lexicon( + "place.stream.chat.message".to_string(), + json!({"id": "place.stream.chat.message"}), + ); + http_client.add_lexicon( + "place.stream.key.defs".to_string(), + json!({"id": "place.stream.key.defs"}), + ); // Create fetcher let fetcher = LexiconFetcher::new(dns_resolver, http_client); @@ -314,12 +367,12 @@ async fn test_underscore_wildcard_direct_children() { // Should only get 3 direct children assert_eq!(lexicons.len(), 3); - + let nsids: Vec<&str> = lexicons.iter().map(|(nsid, _)| nsid.as_str()).collect(); assert!(nsids.contains(&"place.stream.chat")); assert!(nsids.contains(&"place.stream.key")); assert!(nsids.contains(&"place.stream.livestream")); - + // Should NOT contain nested children assert!(!nsids.contains(&"place.stream.chat.profile")); assert!(!nsids.contains(&"place.stream.chat.message")); @@ -334,9 +387,18 @@ async fn test_star_vs_underscore_wildcard() { // Setup mock HTTP client with nested lexicons let mut http_client = MockHttpClient::new(); - http_client.add_lexicon("place.stream.chat".to_string(), json!({"id": "place.stream.chat"})); - http_client.add_lexicon("place.stream.chat.profile".to_string(), json!({"id": "place.stream.chat.profile"})); - http_client.add_lexicon("place.stream.key".to_string(), json!({"id": "place.stream.key"})); + http_client.add_lexicon( + "place.stream.chat".to_string(), + json!({"id": "place.stream.chat"}), + ); + http_client.add_lexicon( + "place.stream.chat.profile".to_string(), + json!({"id": "place.stream.chat.profile"}), + ); + http_client.add_lexicon( + "place.stream.key".to_string(), + json!({"id": "place.stream.key"}), + ); // Create fetcher let fetcher = LexiconFetcher::new(dns_resolver, http_client); @@ -352,8 +414,11 @@ async fn test_star_vs_underscore_wildcard() { assert!(underscore_result.is_ok()); let underscore_lexicons = underscore_result.unwrap(); assert_eq!(underscore_lexicons.len(), 2); // Only direct children (chat, key) - - let nsids: Vec<&str> = underscore_lexicons.iter().map(|(nsid, _)| nsid.as_str()).collect(); + + let nsids: Vec<&str> = underscore_lexicons + .iter() + .map(|(nsid, _)| nsid.as_str()) + .collect(); assert!(nsids.contains(&"place.stream.chat")); assert!(nsids.contains(&"place.stream.key")); assert!(!nsids.contains(&"place.stream.chat.profile")); diff --git a/mlf-lsp/src/context.rs b/mlf-lsp/src/context.rs index deece8f..d7e9de4 100644 --- a/mlf-lsp/src/context.rs +++ b/mlf-lsp/src/context.rs @@ -35,7 +35,8 @@ pub fn detect_context(text_before_cursor: &str) -> CompletionContext { if parts.len() >= 2 && !parts[1].contains('(') { return CompletionContext::Annotation; } - } else if after_at.ends_with(',') || (after_at.contains(',') && !after_at.contains(':')) { + } else if after_at.ends_with(',') || (after_at.contains(',') && !after_at.contains(':')) + { // We're completing another selector: @rust, or @rust,typescript, return CompletionContext::AnnotationSelector; } else if !after_at.is_empty() && !after_at.contains(':') && !after_at.contains('(') { @@ -117,14 +118,8 @@ mod tests { #[test] fn test_use_statement() { - assert_eq!( - detect_context("use "), - CompletionContext::UseStatement - ); - assert_eq!( - detect_context("use com."), - CompletionContext::UseStatement - ); + assert_eq!(detect_context("use "), CompletionContext::UseStatement); + assert_eq!(detect_context("use com."), CompletionContext::UseStatement); assert_eq!( detect_context("use com.atproto."), CompletionContext::UseStatement @@ -170,13 +165,22 @@ mod tests { #[test] fn test_annotation_selector() { - assert_eq!(detect_context("@rust,"), CompletionContext::AnnotationSelector); - assert_eq!(detect_context("@rust,typescript,"), CompletionContext::AnnotationSelector); + assert_eq!( + detect_context("@rust,"), + CompletionContext::AnnotationSelector + ); + assert_eq!( + detect_context("@rust,typescript,"), + CompletionContext::AnnotationSelector + ); } #[test] fn test_annotation_after_selector() { assert_eq!(detect_context("@rust:"), CompletionContext::Annotation); - assert_eq!(detect_context("@rust,typescript:"), CompletionContext::Annotation); + assert_eq!( + detect_context("@rust,typescript:"), + CompletionContext::Annotation + ); } } diff --git a/mlf-lsp/src/main.rs b/mlf-lsp/src/main.rs index c4591e8..1796fd5 100644 --- a/mlf-lsp/src/main.rs +++ b/mlf-lsp/src/main.rs @@ -8,7 +8,12 @@ async fn main() { std::panic::set_hook(Box::new(|panic_info| { eprintln!("LSP PANIC: {:?}", panic_info); if let Some(location) = panic_info.location() { - eprintln!(" at {}:{}:{}", location.file(), location.line(), location.column()); + eprintln!( + " at {}:{}:{}", + location.file(), + location.line(), + location.column() + ); } if let Some(message) = panic_info.payload().downcast_ref::<&str>() { eprintln!(" message: {}", message); @@ -20,8 +25,7 @@ async fn main() { // Initialize logging with debug level tracing_subscriber::fmt() .with_env_filter( - EnvFilter::try_from_default_env() - .unwrap_or_else(|_| EnvFilter::new("debug")) + EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("debug")), ) .with_writer(std::io::stderr) .init(); diff --git a/mlf-lsp/src/namespace_completion.rs b/mlf-lsp/src/namespace_completion.rs index 9b23c47..c2d9659 100644 --- a/mlf-lsp/src/namespace_completion.rs +++ b/mlf-lsp/src/namespace_completion.rs @@ -1,11 +1,10 @@ +use mlf_lang::Workspace; /// Shared namespace path completion logic /// /// This module provides utilities for completing namespace paths in both /// `use` statements and type positions (e.g., `com.atproto.repo.strongRef`). - use std::collections::{HashMap, HashSet}; use tower_lsp::lsp_types::*; -use mlf_lang::Workspace; use crate::server::DocumentState; diff --git a/mlf-lsp/src/server.rs b/mlf-lsp/src/server.rs index 0fd5084..1c7836b 100644 --- a/mlf-lsp/src/server.rs +++ b/mlf-lsp/src/server.rs @@ -1,12 +1,12 @@ +use mlf_lang::Workspace; +use mlf_lang::ast::*; use std::collections::HashMap; use std::path::PathBuf; -use mlf_lang::ast::*; -use mlf_lang::Workspace; use tower_lsp::jsonrpc::Result; use tower_lsp::lsp_types::*; use tower_lsp::{Client, LanguageServer}; -use crate::context::{detect_context, CompletionContext as MlfCompletionContext}; +use crate::context::{CompletionContext as MlfCompletionContext, detect_context}; use crate::namespace_completion; use crate::utils::*; @@ -91,7 +91,8 @@ impl MlfLanguageServer { // Try to update workspace with partial lexicon if available let mut diagnostics = if let Some(ref lex) = partial_lexicon { - self.update_workspace(uri, lex.clone(), namespace, text).await + self.update_workspace(uri, lex.clone(), namespace, text) + .await } else { vec![] }; @@ -101,8 +102,14 @@ impl MlfLanguageServer { span_to_range(text, span) } else { Range { - start: Position { line: 0, character: 0 }, - end: Position { line: 0, character: 0 }, + start: Position { + line: 0, + character: 0, + }, + end: Position { + line: 0, + character: 0, + }, } }; @@ -135,7 +142,13 @@ impl MlfLanguageServer { None } - async fn update_workspace(&self, _uri: &Url, lexicon: Lexicon, namespace: Option, text: &str) -> Vec { + async fn update_workspace( + &self, + _uri: &Url, + lexicon: Lexicon, + namespace: Option, + text: &str, + ) -> Vec { let mut diagnostics = vec![]; if let Some(ns) = namespace { @@ -175,36 +188,60 @@ impl MlfLanguageServer { if mlf_diagnostics::get_error_module_namespace_str(&error) == ns { // Extract span and message from error variant let (span, message, is_unused) = match &error { - ValidationError::DuplicateDefinition { name, second_span, .. } => { - (*second_span, format!("Duplicate definition: {}", name), false) - } + ValidationError::DuplicateDefinition { + name, second_span, .. + } => ( + *second_span, + format!("Duplicate definition: {}", name), + false, + ), ValidationError::UndefinedReference { name, span, .. } => { (*span, format!("Undefined reference: {}", name), false) } ValidationError::InvalidConstraint { message, span, .. } => { (*span, format!("Invalid constraint: {}", message), false) } - ValidationError::TypeMismatch { expected, found, span, .. } => { - (*span, format!("Type mismatch: expected {}, found {}", expected, found), false) - } - ValidationError::ConstraintTooPermissive { message, span, .. } => { - (*span, format!("Constraint too permissive: {}", message), false) - } + ValidationError::TypeMismatch { + expected, + found, + span, + .. + } => ( + *span, + format!( + "Type mismatch: expected {}, found {}", + expected, found + ), + false, + ), + ValidationError::ConstraintTooPermissive { + message, span, .. + } => ( + *span, + format!("Constraint too permissive: {}", message), + false, + ), ValidationError::ReservedName { name, span, .. } => { (*span, format!("Reserved name: {}", name), false) } - ValidationError::AmbiguousMain { name, first_span, .. } => { - (*first_span, format!("Ambiguous main: {}", name), false) - } - ValidationError::MultipleMain { name, first_span, .. } => { - (*first_span, format!("Multiple @main annotations: {}", name), false) - } + ValidationError::AmbiguousMain { + name, first_span, .. + } => (*first_span, format!("Ambiguous main: {}", name), false), + ValidationError::MultipleMain { + name, first_span, .. + } => ( + *first_span, + format!("Multiple @main annotations: {}", name), + false, + ), ValidationError::ConflictNotAllowed { name, span, .. } => { (*span, format!("Conflict not allowed: {}", name), false) } - ValidationError::CircularImport { cycle, span, .. } => { - (*span, format!("Circular import: {}", cycle.join(" -> ")), false) - } + ValidationError::CircularImport { cycle, span, .. } => ( + *span, + format!("Circular import: {}", cycle.join(" -> ")), + false, + ), ValidationError::UnusedImport { name, span, .. } => { (*span, format!("Unused import: {}", name), true) } @@ -214,7 +251,11 @@ impl MlfLanguageServer { diagnostics.push(Diagnostic { range, - severity: Some(if is_unused { DiagnosticSeverity::HINT } else { DiagnosticSeverity::ERROR }), + severity: Some(if is_unused { + DiagnosticSeverity::HINT + } else { + DiagnosticSeverity::ERROR + }), code: None, code_description: None, source: Some("mlf".to_string()), @@ -249,19 +290,28 @@ impl MlfLanguageServer { let std_dir = global_mlf_dir.join("lexicons").join("mlf"); self.client - .log_message(MessageType::INFO, format!("Global MLF directory: {}", global_mlf_dir.display())) + .log_message( + MessageType::INFO, + format!("Global MLF directory: {}", global_mlf_dir.display()), + ) .await; // Ensure std directory exists and has files if !std_dir.join("prelude.mlf").exists() { self.client - .log_message(MessageType::INFO, "Std library not found in ~/.mlf/lexicons/mlf/, extracting embedded files...") + .log_message( + MessageType::INFO, + "Std library not found in ~/.mlf/lexicons/mlf/, extracting embedded files...", + ) .await; // Create directory if let Err(e) = std::fs::create_dir_all(&std_dir) { self.client - .log_message(MessageType::ERROR, format!("Failed to create ~/.mlf/lexicons/mlf/: {}", e)) + .log_message( + MessageType::ERROR, + format!("Failed to create ~/.mlf/lexicons/mlf/: {}", e), + ) .await; return; } @@ -270,7 +320,10 @@ impl MlfLanguageServer { self.extract_embedded_std_to_directory(&std_dir).await; } else { self.client - .log_message(MessageType::INFO, "Using existing std library from ~/.mlf/lexicons/mlf/") + .log_message( + MessageType::INFO, + "Using existing std library from ~/.mlf/lexicons/mlf/", + ) .await; } @@ -283,7 +336,10 @@ impl MlfLanguageServer { } /// Load fetched lexicons from project's .mlf cache directory - async fn load_project_lexicons(&self, workspace: &mut Workspace) -> std::result::Result<(), String> { + async fn load_project_lexicons( + &self, + workspace: &mut Workspace, + ) -> std::result::Result<(), String> { // Try to find project root by looking for mlf.toml // We'll check a few common locations let possible_roots = vec![ @@ -299,10 +355,14 @@ impl MlfLanguageServer { if lexicons_dir.exists() { self.client - .log_message(MessageType::INFO, format!("Loading project lexicons from {}", lexicons_dir.display())) + .log_message( + MessageType::INFO, + format!("Loading project lexicons from {}", lexicons_dir.display()), + ) .await; - self.load_lexicons_from_directory(workspace, &lexicons_dir, &lexicons_dir).await?; + self.load_lexicons_from_directory(workspace, &lexicons_dir, &lexicons_dir) + .await?; self.client .log_message(MessageType::INFO, "Finished loading project lexicons") @@ -352,7 +412,7 @@ impl MlfLanguageServer { self.client .log_message( MessageType::WARNING, - format!("Failed to add module {}: {:?}", namespace, e) + format!("Failed to add module {}: {:?}", namespace, e), ) .await; } @@ -363,7 +423,9 @@ impl MlfLanguageServer { uri.clone(), DocumentState { text: contents.clone(), - lexicon: Some(mlf_lang::parser::parse_lexicon(&contents).unwrap()), + lexicon: Some( + mlf_lang::parser::parse_lexicon(&contents).unwrap(), + ), namespace: Some(namespace.clone()), }, ); @@ -371,7 +433,10 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Loaded project lexicon: {} (namespace: {})", uri, namespace) + format!( + "Loaded project lexicon: {} (namespace: {})", + uri, namespace + ), ) .await; } @@ -464,19 +529,25 @@ impl MlfLanguageServer { // Write file if let Err(e) = std::fs::write(&file_path, contents_str) { - server.client + server + .client .log_message( MessageType::ERROR, - format!("Failed to write std file {}: {}", file_path.display(), e) + format!( + "Failed to write std file {}: {}", + file_path.display(), + e + ), ) .await; continue; } - server.client + server + .client .log_message( MessageType::INFO, - format!("Extracted: {}", file_path.display()) + format!("Extracted: {}", file_path.display()), ) .await; } @@ -508,8 +579,12 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("find_definition_in_workspace: target_name={}, current_namespace={}, path={}", - target_name, current_namespace, path.to_string()) + format!( + "find_definition_in_workspace: target_name={}, current_namespace={}, path={}", + target_name, + current_namespace, + path.to_string() + ), ) .await; @@ -522,7 +597,7 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Resolved target namespace: {}", target_namespace) + format!("Resolved target namespace: {}", target_namespace), ) .await; @@ -532,7 +607,11 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Searching {} documents for namespace '{}'", documents.len(), target_namespace) + format!( + "Searching {} documents for namespace '{}'", + documents.len(), + target_namespace + ), ) .await; @@ -541,7 +620,7 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Checking document {} with namespace '{}'", doc_uri, doc_ns) + format!("Checking document {} with namespace '{}'", doc_uri, doc_ns), ) .await; @@ -549,7 +628,10 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Found matching namespace! Searching for item '{}'", target_name) + format!( + "Found matching namespace! Searching for item '{}'", + target_name + ), ) .await; @@ -572,7 +654,10 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Found definition of '{}' in {}", target_name, doc_uri) + format!( + "Found definition of '{}' in {}", + target_name, doc_uri + ), ) .await; @@ -587,7 +672,7 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Definition not found for '{}'", target_name) + format!("Definition not found for '{}'", target_name), ) .await; @@ -610,8 +695,10 @@ impl MlfLanguageServer { // Check if this looks like "use namespace.typename" (old syntax) // or "use namespace" (new syntax) let (target_namespace, target_type) = if let UseImports::Items(items) = &use_stmt.imports { - if items.len() == 1 && path.segments.len() >= 2 - && items[0].name.name == path.segments.last().unwrap().name { + if items.len() == 1 + && path.segments.len() >= 2 + && items[0].name.name == path.segments.last().unwrap().name + { // Old syntax: use a.b.c as Foo // Navigate to type "c" in namespace "a.b" let ns = path.segments[..path.segments.len() - 1] @@ -631,7 +718,8 @@ impl MlfLanguageServer { (path.to_string(), None) }; - self.find_definition_in_namespace(&target_namespace, target_type.as_deref().unwrap_or("")).await + self.find_definition_in_namespace(&target_namespace, target_type.as_deref().unwrap_or("")) + .await } /// Find a definition in a specific namespace @@ -645,7 +733,10 @@ impl MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("find_definition_in_namespace: namespace='{}', name='{}'", target_namespace, target_name) + format!( + "find_definition_in_namespace: namespace='{}', name='{}'", + target_namespace, target_name + ), ) .await; @@ -829,7 +920,10 @@ impl LanguageServer for MlfLanguageServer { if let Err(e) = self.load_project_lexicons(&mut ws).await { tracing::warn!("Failed to load project lexicons: {}", e); self.client - .log_message(MessageType::WARNING, format!("Failed to load project lexicons: {}", e)) + .log_message( + MessageType::WARNING, + format!("Failed to load project lexicons: {}", e), + ) .await; } @@ -884,7 +978,10 @@ impl LanguageServer for MlfLanguageServer { async fn did_close(&self, params: DidCloseTextDocumentParams) { // Remove document from storage - self.documents.write().await.remove(¶ms.text_document.uri); + self.documents + .write() + .await + .remove(¶ms.text_document.uri); // TODO: Could rebuild workspace without this module } @@ -914,7 +1011,8 @@ impl LanguageServer for MlfLanguageServer { if let Some(annotations) = annotations_to_check { for annotation in annotations { - if annotation.span.start <= offset && offset <= annotation.span.end { + if annotation.span.start <= offset && offset <= annotation.span.end + { // Build hover content for the annotation let mut contents = vec![]; @@ -922,10 +1020,16 @@ impl LanguageServer for MlfLanguageServer { let annotation_display = if annotation.selectors.is_empty() { format!("@{}", annotation.name.name) } else { - let selector_names: Vec<_> = annotation.selectors.iter() + let selector_names: Vec<_> = annotation + .selectors + .iter() .map(|s| s.name.as_str()) .collect(); - format!("@{}:{}", selector_names.join(","), annotation.name.name) + format!( + "@{}:{}", + selector_names.join(","), + annotation.name.name + ) }; contents.push(MarkedString::LanguageString(LanguageString { @@ -936,17 +1040,29 @@ impl LanguageServer for MlfLanguageServer { // Add description based on annotation name let description = match annotation.name.name.as_str() { "deprecated" => "Marks this definition as deprecated", - "main" => "Designates this as the main definition for conflict resolution", - "key" => "Specifies the record key type (e.g., 'tid', 'literal:self')", - "encoding" => "Specifies MIME type encoding for XRPC (e.g., 'application/json', 'application/cbor')", + "main" => { + "Designates this as the main definition for conflict resolution" + } + "key" => { + "Specifies the record key type (e.g., 'tid', 'literal:self')" + } + "encoding" => { + "Specifies MIME type encoding for XRPC (e.g., 'application/json', 'application/cbor')" + } "since" => "Indicates the version when this was added", "doc" => "Provides a documentation URL", "validate" => "Specifies validation rules", "cache" => "Defines caching strategy", "indexed" => "Marks this field as indexed", - "sensitive" => "Marks this field as containing sensitive data (e.g., PII)", - "const" => "Extension field (literal). `@const(key, value)` emits `key: value` verbatim in the JSON Lexicon. Use on `self {}` for top-level fields, or any item for per-item fields.", - "reference" => "Extension field (named-type reference). `@reference(key, path)` resolves `path` through the workspace and emits the resulting NSID string under `key`.", + "sensitive" => { + "Marks this field as containing sensitive data (e.g., PII)" + } + "const" => { + "Extension field (literal). `@const(key, value)` emits `key: value` verbatim in the JSON Lexicon. Use on `self {}` for top-level fields, or any item for per-item fields." + } + "reference" => { + "Extension field (named-type reference). `@reference(key, path)` resolves `path` through the workspace and emits the resulting NSID string under `key`." + } _ => "Custom annotation", }; @@ -956,7 +1072,9 @@ impl LanguageServer for MlfLanguageServer { if !annotation.selectors.is_empty() { let selector_info = format!( "This annotation applies to: {}", - annotation.selectors.iter() + annotation + .selectors + .iter() .map(|s| s.name.as_str()) .collect::>() .join(", ") @@ -964,7 +1082,8 @@ impl LanguageServer for MlfLanguageServer { contents.push(MarkedString::String(selector_info)); } else { contents.push(MarkedString::String( - "This annotation is visible to all generators".to_string() + "This annotation is visible to all generators" + .to_string(), )); } @@ -981,22 +1100,34 @@ impl LanguageServer for MlfLanguageServer { Item::Record(r) => { for field in &r.fields { for annotation in &field.annotations { - if annotation.span.start <= offset && offset <= annotation.span.end { - let annotation_display = if annotation.selectors.is_empty() { - format!("@{}", annotation.name.name) - } else { - let selector_names: Vec<_> = annotation.selectors.iter() - .map(|s| s.name.as_str()) - .collect(); - format!("@{}:{}", selector_names.join(","), annotation.name.name) - }; + if annotation.span.start <= offset + && offset <= annotation.span.end + { + let annotation_display = + if annotation.selectors.is_empty() { + format!("@{}", annotation.name.name) + } else { + let selector_names: Vec<_> = annotation + .selectors + .iter() + .map(|s| s.name.as_str()) + .collect(); + format!( + "@{}:{}", + selector_names.join(","), + annotation.name.name + ) + }; return Ok(Some(Hover { contents: HoverContents::Scalar( MarkedString::LanguageString(LanguageString { language: "mlf".to_string(), - value: format!("Field annotation: {}", annotation_display), - }) + value: format!( + "Field annotation: {}", + annotation_display + ), + }), ), range: None, })); @@ -1049,38 +1180,55 @@ impl LanguageServer for MlfLanguageServer { Item::Record(r) => { if let Some(field) = find_field_at_offset(&r.fields, offset) { let field_type = format_type(&field.ty); - let opt = if field.optional { "optional" } else { "required" }; - contents.push(MarkedString::String( - format!("Field: {} ({})", field_type, opt) - )); + let opt = if field.optional { + "optional" + } else { + "required" + }; + contents.push(MarkedString::String(format!( + "Field: {} ({})", + field_type, opt + ))); } } Item::InlineType(i) => { - contents.push(MarkedString::String( - format!("Type: {}", format_type(&i.ty)) - )); + contents.push(MarkedString::String(format!( + "Type: {}", + format_type(&i.ty) + ))); } Item::DefType(d) => { - contents.push(MarkedString::String( - format!("Type: {}", format_type(&d.ty)) - )); + contents.push(MarkedString::String(format!( + "Type: {}", + format_type(&d.ty) + ))); } Item::Query(q) => { if let Some(field) = find_field_at_offset(&q.params, offset) { let field_type = format_type(&field.ty); - let opt = if field.optional { "optional" } else { "required" }; - contents.push(MarkedString::String( - format!("Parameter: {} ({})", field_type, opt) - )); + let opt = if field.optional { + "optional" + } else { + "required" + }; + contents.push(MarkedString::String(format!( + "Parameter: {} ({})", + field_type, opt + ))); } } Item::Procedure(p) => { if let Some(field) = find_field_at_offset(&p.params, offset) { let field_type = format_type(&field.ty); - let opt = if field.optional { "optional" } else { "required" }; - contents.push(MarkedString::String( - format!("Parameter: {} ({})", field_type, opt) - )); + let opt = if field.optional { + "optional" + } else { + "required" + }; + contents.push(MarkedString::String(format!( + "Parameter: {} ({})", + field_type, opt + ))); } } _ => {} @@ -1121,7 +1269,7 @@ impl LanguageServer for MlfLanguageServer { MarkedString::LanguageString(LanguageString { language: "mlf".to_string(), value: format!("type {}", path.to_string()), - }) + }), ), range: None, })); @@ -1149,11 +1297,12 @@ impl LanguageServer for MlfLanguageServer { let mut completions = vec![]; // Detect context - let text_before_cursor = if let Some(offset) = position_to_offset(&doc_state.text, position) { - &doc_state.text[..offset] - } else { - "" - }; + let text_before_cursor = + if let Some(offset) = position_to_offset(&doc_state.text, position) { + &doc_state.text[..offset] + } else { + "" + }; let context = detect_context(text_before_cursor); tracing::debug!("Context detected: {:?}", context); @@ -1174,10 +1323,15 @@ impl LanguageServer for MlfLanguageServer { let has_trailing_dot = after_use.ends_with('.'); let partial_path = after_use.trim_end_matches('.'); // Remove trailing dot for matching - tracing::debug!("Partial path: '{}', has_trailing_dot: {}", partial_path, has_trailing_dot); + tracing::debug!( + "Partial path: '{}', has_trailing_dot: {}", + partial_path, + has_trailing_dot + ); // Calculate the range to replace - let use_start_offset = text_before_cursor.rfind("use ").map(|i| i + 4).unwrap_or(0); + let use_start_offset = + text_before_cursor.rfind("use ").map(|i| i + 4).unwrap_or(0); let use_start_pos = offset_to_position(&doc_state.text, use_start_offset); // Use shared namespace completion logic @@ -1229,7 +1383,8 @@ impl LanguageServer for MlfLanguageServer { MlfCompletionContext::TypePosition => { // Extract text after the last ':' to detect if user is typing a namespace path let last_line = text_before_cursor.lines().last().unwrap_or(""); - let after_colon = last_line.rfind(':') + let after_colon = last_line + .rfind(':') .map(|idx| &last_line[idx + 1..]) .unwrap_or(""); @@ -1239,7 +1394,11 @@ impl LanguageServer for MlfLanguageServer { // Check if user is typing a path (contains a dot) let is_typing_path = after_colon_trimmed.contains('.'); - tracing::debug!("Type position: after_colon_trimmed='{}', is_typing_path={}", after_colon_trimmed, is_typing_path); + tracing::debug!( + "Type position: after_colon_trimmed='{}', is_typing_path={}", + after_colon_trimmed, + is_typing_path + ); if is_typing_path { // User is typing a namespace path like "com.atproto." @@ -1247,13 +1406,19 @@ impl LanguageServer for MlfLanguageServer { let has_trailing_dot = after_colon_trimmed.ends_with('.'); let partial_path = after_colon_trimmed.trim_end_matches('.'); - tracing::debug!("Namespace path mode: partial_path='{}', has_trailing_dot={}", partial_path, has_trailing_dot); + tracing::debug!( + "Namespace path mode: partial_path='{}', has_trailing_dot={}", + partial_path, + has_trailing_dot + ); // Calculate the range to replace (from after ':' and any spaces to cursor) - let colon_offset = text_before_cursor.rfind(':').map(|i| i + 1).unwrap_or(0); + let colon_offset = + text_before_cursor.rfind(':').map(|i| i + 1).unwrap_or(0); // Skip leading spaces to get to where the actual path starts let space_count = after_colon.len() - after_colon_trimmed.len(); - let replace_start_pos = offset_to_position(&doc_state.text, colon_offset + space_count); + let replace_start_pos = + offset_to_position(&doc_state.text, colon_offset + space_count); let workspace_guard = self.workspace.read().await; if let Some(workspace) = workspace_guard.as_ref() { @@ -1274,9 +1439,8 @@ impl LanguageServer for MlfLanguageServer { tracing::debug!("Local type mode"); // Primitive types - let primitives = vec![ - "null", "boolean", "integer", "string", "bytes", "blob", - ]; + let primitives = + vec!["null", "boolean", "integer", "string", "bytes", "blob"]; for prim in primitives { completions.push(CompletionItem { @@ -1334,7 +1498,11 @@ impl LanguageServer for MlfLanguageServer { let workspace_guard = self.workspace.read().await; if let Some(workspace) = workspace_guard.as_ref() { let imports = workspace.get_imports(current_namespace); - tracing::debug!("Found {} imports for namespace '{}'", imports.len(), current_namespace); + tracing::debug!( + "Found {} imports for namespace '{}'", + imports.len(), + current_namespace + ); for (local_name, original_path) in imports { // Format the original path for display @@ -1386,11 +1554,23 @@ impl LanguageServer for MlfLanguageServer { let annotations = vec![ ("deprecated", "Mark as deprecated"), ("main", "Main definition for conflict resolution"), - ("key", "Specify record key type (e.g., @key(\"literal:self\"))"), - ("encoding", "Specify MIME type encoding (e.g., @encoding(\"application/cbor\"))"), + ( + "key", + "Specify record key type (e.g., @key(\"literal:self\"))", + ), + ( + "encoding", + "Specify MIME type encoding (e.g., @encoding(\"application/cbor\"))", + ), ("since", "Version when added (e.g., @since(1, 2, 0))"), - ("doc", "Documentation URL (e.g., @doc(\"https://example.com\"))"), - ("validate", "Validation rules (e.g., @validate(min: 0, max: 100))"), + ( + "doc", + "Documentation URL (e.g., @doc(\"https://example.com\"))", + ), + ( + "validate", + "Validation rules (e.g., @validate(min: 0, max: 100))", + ), ("cache", "Caching strategy (e.g., @cache(ttl: 3600))"), ("indexed", "Mark field as indexed"), ("sensitive", "Mark field as containing sensitive data"), @@ -1429,15 +1609,39 @@ impl LanguageServer for MlfLanguageServer { MlfCompletionContext::TopLevel => { // Suggest keywords for top-level declarations let keywords = vec![ - ("record", CompletionItemKind::KEYWORD, "Define a record type"), - ("inline type", CompletionItemKind::KEYWORD, "Define an inline type"), + ( + "record", + CompletionItemKind::KEYWORD, + "Define a record type", + ), + ( + "inline type", + CompletionItemKind::KEYWORD, + "Define an inline type", + ), ("def type", CompletionItemKind::KEYWORD, "Define a def type"), ("token", CompletionItemKind::KEYWORD, "Define a token"), ("query", CompletionItemKind::KEYWORD, "Define a query"), - ("procedure", CompletionItemKind::KEYWORD, "Define a procedure"), - ("subscription", CompletionItemKind::KEYWORD, "Define a subscription"), - ("use", CompletionItemKind::KEYWORD, "Import types from another module"), - ("self", CompletionItemKind::KEYWORD, "Lexicon-as-item — attach top-level docs and @const / @reference extensions"), + ( + "procedure", + CompletionItemKind::KEYWORD, + "Define a procedure", + ), + ( + "subscription", + CompletionItemKind::KEYWORD, + "Define a subscription", + ), + ( + "use", + CompletionItemKind::KEYWORD, + "Import types from another module", + ), + ( + "self", + CompletionItemKind::KEYWORD, + "Lexicon-as-item — attach top-level docs and @const / @reference extensions", + ), ]; for (label, kind, detail) in keywords { @@ -1473,7 +1677,10 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Go to definition request at {}:{}:{}", uri, position.line, position.character), + format!( + "Go to definition request at {}:{}:{}", + uri, position.line, position.character + ), ) .await; @@ -1501,7 +1708,10 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Found use statement at cursor: {}", use_stmt.path.to_string()) + format!( + "Found use statement at cursor: {}", + use_stmt.path.to_string() + ), ) .await; @@ -1511,9 +1721,10 @@ impl LanguageServer for MlfLanguageServer { // Handle go-to-definition for use statements if let Some(ref current_ns) = current_namespace { - if let Some((def_uri, def_span)) = - self.find_definition_in_workspace_for_use(use_stmt, current_ns).await { - + if let Some((def_uri, def_span)) = self + .find_definition_in_workspace_for_use(use_stmt, current_ns) + .await + { let documents = self.documents.read().await; if let Some(target_doc) = documents.get(&def_uri) { let range = span_to_range(&target_doc.text, def_span); @@ -1521,7 +1732,10 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Returning use statement definition from: {}", def_uri) + format!( + "Returning use statement definition from: {}", + def_uri + ), ) .await; @@ -1539,11 +1753,16 @@ impl LanguageServer for MlfLanguageServer { // Check if cursor is on an imported item name if let UseImports::Items(items) = &use_stmt.imports { for import_item in items { - if import_item.name.span.start <= offset && offset <= import_item.name.span.end { + if import_item.name.span.start <= offset + && offset <= import_item.name.span.end + { self.client .log_message( MessageType::INFO, - format!("Found use item at cursor: {}", import_item.name.name) + format!( + "Found use item at cursor: {}", + import_item.name.name + ), ) .await; @@ -1553,17 +1772,25 @@ impl LanguageServer for MlfLanguageServer { let target_namespace = use_stmt.path.to_string(); let item_name = if import_item.name.name == "main" { // Special case: "main" resolves to namespace suffix - target_namespace.split('.').last().unwrap_or(&import_item.name.name) + target_namespace + .split('.') + .last() + .unwrap_or(&import_item.name.name) } else { &import_item.name.name }; - if let Some((def_uri, def_span)) = - self.find_definition_in_namespace(&target_namespace, item_name).await { - + if let Some((def_uri, def_span)) = self + .find_definition_in_namespace( + &target_namespace, + item_name, + ) + .await + { let documents = self.documents.read().await; if let Some(target_doc) = documents.get(&def_uri) { - let range = span_to_range(&target_doc.text, def_span); + let range = + span_to_range(&target_doc.text, def_span); return Ok(Some(GotoDefinitionResponse::Scalar( Location { @@ -1583,14 +1810,10 @@ impl LanguageServer for MlfLanguageServer { // Find type reference at this position for item in &lexicon.items { let type_to_check = match item { - Item::Record(r) => { - find_field_at_offset(&r.fields, offset).map(|f| &f.ty) - } + Item::Record(r) => find_field_at_offset(&r.fields, offset).map(|f| &f.ty), Item::InlineType(i) => Some(&i.ty), Item::DefType(d) => Some(&d.ty), - Item::Query(q) => { - find_field_at_offset(&q.params, offset).map(|f| &f.ty) - } + Item::Query(q) => find_field_at_offset(&q.params, offset).map(|f| &f.ty), Item::Procedure(p) => { find_field_at_offset(&p.params, offset).map(|f| &f.ty) } @@ -1598,11 +1821,12 @@ impl LanguageServer for MlfLanguageServer { }; if let Some(ty) = type_to_check { - if let Some(Type::Reference { path, .. }) = find_type_at_offset(ty, offset) { + if let Some(Type::Reference { path, .. }) = find_type_at_offset(ty, offset) + { self.client .log_message( MessageType::INFO, - format!("Found reference at cursor: {}", path.to_string()) + format!("Found reference at cursor: {}", path.to_string()), ) .await; @@ -1616,7 +1840,11 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Target name: {}, path segments: {}", target_name, path.segments.len()) + format!( + "Target name: {}, path segments: {}", + target_name, + path.segments.len() + ), ) .await; @@ -1639,23 +1867,21 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Found definition in current file") + format!("Found definition in current file"), ) .await; - return Ok(Some(GotoDefinitionResponse::Scalar( - Location { - uri: uri.clone(), - range, - }, - ))); + return Ok(Some(GotoDefinitionResponse::Scalar(Location { + uri: uri.clone(), + range, + }))); } } self.client .log_message( MessageType::INFO, - format!("Not found in current file, searching workspace...") + format!("Not found in current file, searching workspace..."), ) .await; @@ -1664,13 +1890,14 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Current namespace: {}", current_ns) + format!("Current namespace: {}", current_ns), ) .await; - if let Some((def_uri, def_span)) = - self.find_definition_in_workspace(target_name, current_ns, path).await { - + if let Some((def_uri, def_span)) = self + .find_definition_in_workspace(target_name, current_ns, path) + .await + { // Get text for span conversion let documents = self.documents.read().await; if let Some(target_doc) = documents.get(&def_uri) { @@ -1679,7 +1906,10 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::INFO, - format!("Returning definition from workspace: {}", def_uri) + format!( + "Returning definition from workspace: {}", + def_uri + ), ) .await; @@ -1695,7 +1925,7 @@ impl LanguageServer for MlfLanguageServer { self.client .log_message( MessageType::WARNING, - format!("No current namespace available") + format!("No current namespace available"), ) .await; } diff --git a/mlf-lsp/src/utils.rs b/mlf-lsp/src/utils.rs index 9bb5931..05049eb 100644 --- a/mlf-lsp/src/utils.rs +++ b/mlf-lsp/src/utils.rs @@ -52,7 +52,10 @@ pub fn offset_to_position(text: &str, offset: usize) -> Position { current_offset += ch.len_utf8(); } - Position { line: line as u32, character: character as u32 } + Position { + line: line as u32, + character: character as u32, + } } /// Convert MLF Span to LSP Range @@ -121,7 +124,9 @@ pub fn find_type_at_offset(ty: &Type, offset: usize) -> Option<&Type> { /// Find field at offset within a record/query/procedure pub fn find_field_at_offset(fields: &[Field], offset: usize) -> Option<&Field> { - fields.iter().find(|field| offset_in_span(offset, field.span)) + fields + .iter() + .find(|field| offset_in_span(offset, field.span)) } /// Get the item name @@ -185,8 +190,14 @@ pub fn format_type(ty: &Type) -> String { format!("{{ {} }}", fields_str) } Type::Parenthesized { inner, .. } => format!("({})", format_type(inner)), - Type::Constrained { base, constraints, .. } => { - format!("{} constrained {{ {} constraints }}", format_type(base), constraints.len()) + Type::Constrained { + base, constraints, .. + } => { + format!( + "{} constrained {{ {} constraints }}", + format_type(base), + constraints.len() + ) } Type::Unknown { .. } => "unknown".to_string(), } diff --git a/mlf-validation/src/lib.rs b/mlf-validation/src/lib.rs index 9263fbc..7aa3b90 100644 --- a/mlf-validation/src/lib.rs +++ b/mlf-validation/src/lib.rs @@ -1,12 +1,12 @@ +use langtag::LangTag; use mlf_lang::ast::*; +use regex::Regex; use serde_json::Value as JsonValue; use std::fmt; +use time::OffsetDateTime; +use time::format_description::well_known::Rfc3339; use unicode_segmentation::UnicodeSegmentation; -use regex::Regex; use url::Url; -use time::format_description::well_known::Rfc3339; -use time::OffsetDateTime; -use langtag::LangTag; #[derive(Debug, Clone)] pub struct ValidationError { @@ -45,7 +45,8 @@ impl<'a> RecordValidator<'a> { Item::Query(_) | Item::Procedure(_) => { errors.push(ValidationError { path: "$".to_string(), - message: "Cannot validate records against query/procedure definitions".to_string(), + message: "Cannot validate records against query/procedure definitions" + .to_string(), }); } _ => { @@ -90,7 +91,9 @@ impl<'a> RecordValidator<'a> { Type::Primitive { kind, .. } => { self.validate_primitive(value, *kind, path, errors); } - Type::Constrained { base, constraints, .. } => { + Type::Constrained { + base, constraints, .. + } => { self.validate_against_type(value, base, path, errors); self.validate_constraints(value, constraints, path, errors); } @@ -246,7 +249,11 @@ impl<'a> RecordValidator<'a> { if s.len() < *min { errors.push(ValidationError { path: path.to_string(), - message: format!("String too short: {} bytes (min: {})", s.len(), min), + message: format!( + "String too short: {} bytes (min: {})", + s.len(), + min + ), }); } } else if let Some(arr) = value.as_array() { @@ -254,7 +261,11 @@ impl<'a> RecordValidator<'a> { if arr.len() < *min { errors.push(ValidationError { path: path.to_string(), - message: format!("Array too short: {} elements (min: {})", arr.len(), min), + message: format!( + "Array too short: {} elements (min: {})", + arr.len(), + min + ), }); } } @@ -264,7 +275,11 @@ impl<'a> RecordValidator<'a> { if s.len() > *max { errors.push(ValidationError { path: path.to_string(), - message: format!("String too long: {} bytes (max: {})", s.len(), max), + message: format!( + "String too long: {} bytes (max: {})", + s.len(), + max + ), }); } } else if let Some(arr) = value.as_array() { @@ -272,7 +287,11 @@ impl<'a> RecordValidator<'a> { if arr.len() > *max { errors.push(ValidationError { path: path.to_string(), - message: format!("Array too long: {} elements (max: {})", arr.len(), max), + message: format!( + "Array too long: {} elements (max: {})", + arr.len(), + max + ), }); } } @@ -284,7 +303,10 @@ impl<'a> RecordValidator<'a> { if count < *min { errors.push(ValidationError { path: path.to_string(), - message: format!("String has too few graphemes: {} (min: {})", count, min), + message: format!( + "String has too few graphemes: {} (min: {})", + count, min + ), }); } } @@ -296,7 +318,10 @@ impl<'a> RecordValidator<'a> { if count > *max { errors.push(ValidationError { path: path.to_string(), - message: format!("String has too many graphemes: {} (max: {})", count, max), + message: format!( + "String has too many graphemes: {} (max: {})", + count, max + ), }); } } @@ -323,10 +348,13 @@ impl<'a> RecordValidator<'a> { } Constraint::Enum { values, .. } => { if let Some(s) = value.as_str() { - let enum_strings: Vec = values.iter().map(|v| match v { - mlf_lang::ast::ValueRef::Literal(lit) => lit.clone(), - mlf_lang::ast::ValueRef::Reference(path) => path.to_string(), - }).collect(); + let enum_strings: Vec = values + .iter() + .map(|v| match v { + mlf_lang::ast::ValueRef::Literal(lit) => lit.clone(), + mlf_lang::ast::ValueRef::Reference(path) => path.to_string(), + }) + .collect(); if !enum_strings.contains(&s.to_string()) { errors.push(ValidationError { path: path.to_string(), @@ -347,7 +375,10 @@ impl<'a> RecordValidator<'a> { if !mimes.iter().any(|m| m == mime) { errors.push(ValidationError { path: path.to_string(), - message: format!("MIME type '{}' not accepted (allowed: {:?})", mime, mimes), + message: format!( + "MIME type '{}' not accepted (allowed: {:?})", + mime, mimes + ), }); } } @@ -488,7 +519,10 @@ impl<'a> RecordValidator<'a> { if !matched { errors.push(ValidationError { path: path.to_string(), - message: format!("Value does not match any type in union ({} variants tried)", types.len()), + message: format!( + "Value does not match any type in union ({} variants tried)", + types.len() + ), }); } } @@ -571,9 +605,7 @@ fn validate_did(value: &str) -> bool { // DID format: did:method:method-specific-id // method: lowercase letters, numbers // method-specific-id: alphanumeric plus . - _ : - let re = Regex::new( - r"^did:[a-z0-9]+:[a-zA-Z0-9._:%-]*[a-zA-Z0-9._-]$" - ).unwrap(); + let re = Regex::new(r"^did:[a-z0-9]+:[a-zA-Z0-9._:%-]*[a-zA-Z0-9._-]$").unwrap(); re.is_match(value) } @@ -591,10 +623,14 @@ fn validate_handle(value: &str) -> bool { if segment.is_empty() || segment.starts_with('-') || segment.ends_with('-') - || segment.len() > 63 { + || segment.len() > 63 + { return false; } - if !segment.chars().all(|c| c.is_ascii_alphanumeric() || c == '-') { + if !segment + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '-') + { return false; } } @@ -624,11 +660,19 @@ fn validate_nsid(value: &str) -> bool { return false; } // NSID segments must be lowercase alphanumeric (and hyphen for domain parts) - if !part.chars().all(|c| c.is_ascii_lowercase() || c.is_ascii_digit() || c == '-') { + if !part + .chars() + .all(|c| c.is_ascii_lowercase() || c.is_ascii_digit() || c == '-') + { return false; } // Can't start with digit - if part.chars().next().map(|c| c.is_ascii_digit()).unwrap_or(false) { + if part + .chars() + .next() + .map(|c| c.is_ascii_digit()) + .unwrap_or(false) + { return false; } } @@ -649,9 +693,9 @@ fn validate_cid(value: &str) -> bool { // CIDv1: starts with 'b' (base32) or 'z' (base58btc) followed by version if value.starts_with("Qm") && value.len() == 46 { // CIDv0 - all base58btc chars - return value.chars().all(|c| { - c.is_ascii_alphanumeric() && c != '0' && c != 'O' && c != 'I' && c != 'l' - }); + return value + .chars() + .all(|c| c.is_ascii_alphanumeric() && c != '0' && c != 'O' && c != 'I' && c != 'l'); } if (value.starts_with('b') || value.starts_with('z')) && value.len() > 10 { @@ -681,9 +725,7 @@ fn validate_tid(value: &str) -> bool { return false; } - value.chars().all(|c| { - matches!(c, 'a'..='z' | '2'..='7') - }) + value.chars().all(|c| matches!(c, 'a'..='z' | '2'..='7')) } /// Validate record-key format @@ -701,9 +743,9 @@ fn validate_record_key(value: &str) -> bool { } // Otherwise, general record key validation - value.chars().all(|c| { - c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '~' || c == '-' - }) + value + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '_' || c == '~' || c == '-') } #[cfg(test)] @@ -743,8 +785,12 @@ mod tests { // Valid AT-URIs assert!(validate_at_uri("at://did:plc:abc123")); assert!(validate_at_uri("at://did:plc:abc123/com.example.foo")); - assert!(validate_at_uri("at://did:plc:abc123/com.example.foo/abc123")); - assert!(validate_at_uri("at://alice.example.com/com.example.post/abc")); + assert!(validate_at_uri( + "at://did:plc:abc123/com.example.foo/abc123" + )); + assert!(validate_at_uri( + "at://alice.example.com/com.example.post/abc" + )); // Invalid AT-URIs assert!(!validate_at_uri("https://example.com")); @@ -797,8 +843,12 @@ mod tests { #[test] fn test_validate_cid() { // Valid CIDs (examples) - assert!(validate_cid("QmYwAPJzv5CZsnA625s3Xf2nemtYgPpHdWEz79ojWnPbdG")); // CIDv0 - assert!(validate_cid("bafybeihdwdcefgh4dqkjv67uzcmw7ojee6xedzdetojuzjevtenxquvyku")); // CIDv1 + assert!(validate_cid( + "QmYwAPJzv5CZsnA625s3Xf2nemtYgPpHdWEz79ojWnPbdG" + )); // CIDv0 + assert!(validate_cid( + "bafybeihdwdcefgh4dqkjv67uzcmw7ojee6xedzdetojuzjevtenxquvyku" + )); // CIDv1 // Invalid CIDs assert!(!validate_cid("")); diff --git a/tests/codegen_integration.rs b/tests/codegen_integration.rs index 9a78f5d..e7c0a29 100644 --- a/tests/codegen_integration.rs +++ b/tests/codegen_integration.rs @@ -14,7 +14,9 @@ use std::fs; use std::path::Path; fn run_lexicon_test(input_path: &Path) -> datatest_stable::Result<()> { - let test_dir = input_path.parent().ok_or("input.mlf has no parent directory")?; + let test_dir = input_path + .parent() + .ok_or("input.mlf has no parent directory")?; let test_name = test_dir .file_name() .and_then(|s| s.to_str()) @@ -29,7 +31,8 @@ fn run_lexicon_test(input_path: &Path) -> datatest_stable::Result<()> { let input = fs::read_to_string(input_path)?; let lexicon = parse_lexicon(&input).map_err(|e| format!("Failed to parse: {:?}", e))?; - let mut ws = Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; + let mut ws = + Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; ws.add_module(namespace.clone(), lexicon) .map_err(|e| format!("Failed to add module: {:?}", e))?; ws.resolve() diff --git a/tests/diagnostics_integration.rs b/tests/diagnostics_integration.rs index edbf3f2..a3ad420 100644 --- a/tests/diagnostics_integration.rs +++ b/tests/diagnostics_integration.rs @@ -15,7 +15,9 @@ use std::fs; use std::path::Path; fn run_diagnostics_test(input_path: &Path) -> datatest_stable::Result<()> { - let test_dir = input_path.parent().ok_or("input.mlf has no parent directory")?; + let test_dir = input_path + .parent() + .ok_or("input.mlf has no parent directory")?; let test_name = test_dir .file_name() .and_then(|s| s.to_str()) @@ -33,7 +35,8 @@ fn run_diagnostics_test(input_path: &Path) -> datatest_stable::Result<()> { let input = fs::read_to_string(input_path)?; let lexicon = parse_lexicon(&input).map_err(|e| format!("Failed to parse: {:?}", e))?; - let mut ws = Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; + let mut ws = + Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; ws.add_module(namespace.clone(), lexicon) .map_err(|e| format!("Failed to add module: {:?}", e))?; diff --git a/tests/lexicon_fetcher_integration.rs b/tests/lexicon_fetcher_integration.rs index f6758cd..44dd0a3 100644 --- a/tests/lexicon_fetcher_integration.rs +++ b/tests/lexicon_fetcher_integration.rs @@ -85,7 +85,11 @@ async fn run_dns_test_async(config_path: &Path) -> datatest_stable::Result<()> { .into()); } } - } else if test_dns.resolve_lexicon_did(&authority, &dns_name).await.is_ok() { + } else if test_dns + .resolve_lexicon_did(&authority, &dns_name) + .await + .is_ok() + { return Err("Expected DNS lookup to fail, but it succeeded".into()); } diff --git a/tests/lexicon_to_mlf_integration.rs b/tests/lexicon_to_mlf_integration.rs index 4d4ebe3..c90c697 100644 --- a/tests/lexicon_to_mlf_integration.rs +++ b/tests/lexicon_to_mlf_integration.rs @@ -14,7 +14,9 @@ use std::fs; use std::path::Path; fn run_case(input_path: &Path) -> datatest_stable::Result<()> { - let test_dir = input_path.parent().ok_or("input.json has no parent directory")?; + let test_dir = input_path + .parent() + .ok_or("input.json has no parent directory")?; let test_name = test_dir .file_name() .and_then(|s| s.to_str()) diff --git a/tests/real_world/roundtrip.rs b/tests/real_world/roundtrip.rs index cd97401..4b47134 100644 --- a/tests/real_world/roundtrip.rs +++ b/tests/real_world/roundtrip.rs @@ -80,7 +80,8 @@ fn run_source_roundtrip(source: &str) { } assert_eq!( - stats.failures, 0, + stats.failures, + 0, "Round-trip failed for {} lexicon(s) under {}. See {}", stats.failures, source, @@ -143,8 +144,11 @@ optimize_transitive_fetches = false /// workspace, resolve it, and regenerate Lexicon JSON for each module. /// Returns a map keyed by the on-disk path relative to `mlf_dir` (so /// callers can line up against the original JSON tree). -fn regenerate_lexicons_from_mlf(mlf_dir: &Path) -> Result, String> { - let mut ws = Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; +fn regenerate_lexicons_from_mlf( + mlf_dir: &Path, +) -> Result, String> { + let mut ws = + Workspace::with_std().map_err(|e| format!("Failed to create workspace: {:?}", e))?; load_mlf_directory(&mut ws, mlf_dir)?; ws.resolve() @@ -231,9 +235,10 @@ fn compare_regenerated( let original_path = original_dir.join(relative_path); if !original_path.exists() { stats.failures += 1; - stats - .failed_lexicons - .push((nsid, format!("Original file not found: {}", original_path.display()))); + stats.failed_lexicons.push(( + nsid, + format!("Original file not found: {}", original_path.display()), + )); continue; } @@ -328,13 +333,11 @@ fn canonicalize_lexicon_value(value: &serde_json::Value, lexicon_id: &str) -> se } serde_json::Value::Object(sorted) } - serde_json::Value::Array(arr) => { - serde_json::Value::Array( - arr.iter() - .map(|v| canonicalize_lexicon_value(v, lexicon_id)) - .collect(), - ) - } + serde_json::Value::Array(arr) => serde_json::Value::Array( + arr.iter() + .map(|v| canonicalize_lexicon_value(v, lexicon_id)) + .collect(), + ), _ => value.clone(), } } @@ -391,13 +394,18 @@ fn canonicalize_object_in_place( // Strip empty `required: []` — ATProto treats absent and empty // equivalently, and our converter drops them. - if obj.get("required").and_then(|v| v.as_array()).map_or(false, |a| a.is_empty()) { + if obj + .get("required") + .and_then(|v| v.as_array()) + .map_or(false, |a| a.is_empty()) + { obj.remove("required"); } // Normalize missing `properties` on object types — our converter // always emits `properties: {}` even when the original omits it. - if obj.get("type").and_then(|v| v.as_str()) == Some("object") && !obj.contains_key("properties") { + if obj.get("type").and_then(|v| v.as_str()) == Some("object") && !obj.contains_key("properties") + { obj.insert("properties".to_string(), serde_json::json!({})); } @@ -439,9 +447,16 @@ fn canonicalize_object_in_place( // description is style guidance, not structural. if obj.get("type").and_then(|v| v.as_str()) == Some("array") { if let Some(serde_json::Value::Object(items)) = obj.get_mut("items") { - if items.get("type").and_then(|v| v.as_str()).map_or(false, |t| { - matches!(t, "string" | "integer" | "boolean" | "bytes" | "blob" | "unknown") - }) { + if items + .get("type") + .and_then(|v| v.as_str()) + .map_or(false, |t| { + matches!( + t, + "string" | "integer" | "boolean" | "bytes" | "blob" | "unknown" + ) + }) + { items.remove("description"); } } @@ -465,13 +480,21 @@ fn canonicalize_object_in_place( fn canonicalize_ref_string(ref_str: &str, lexicon_id: &str) -> String { let last_segment = lexicon_id.rsplit('.').next().unwrap_or(""); if let Some(fragment) = ref_str.strip_prefix('#') { - let canonical_fragment = if fragment == last_segment { "main" } else { fragment }; + let canonical_fragment = if fragment == last_segment { + "main" + } else { + fragment + }; format!("{}#{}", lexicon_id, canonical_fragment) } else if let Some(pos) = ref_str.find('#') { let ns = &ref_str[..pos]; let fragment = &ref_str[pos + 1..]; let ns_last = ns.rsplit('.').next().unwrap_or(""); - let canonical_fragment = if fragment == ns_last { "main" } else { fragment }; + let canonical_fragment = if fragment == ns_last { + "main" + } else { + fragment + }; format!("{}#{}", ns, canonical_fragment) } else { format!("{}#main", ref_str) @@ -520,4 +543,3 @@ fn write_diff_file( Ok(()) } - diff --git a/tests/test_utils.rs b/tests/test_utils.rs index 306898c..09f3b37 100644 --- a/tests/test_utils.rs +++ b/tests/test_utils.rs @@ -30,12 +30,7 @@ pub fn load_test_config( if !config_path.exists() { // Fallback: create default config using provided function - let test_name = test_dir - .file_name() - .unwrap() - .to_str() - .unwrap() - .to_string(); + let test_name = test_dir.file_name().unwrap().to_str().unwrap().to_string(); let namespace = default_namespace_fn(&test_name); let mut modules = HashMap::new(); @@ -51,9 +46,8 @@ pub fn load_test_config( }); } - let config_str = fs::read_to_string(&config_path) - .map_err(|e| format!("Failed to read test.toml: {}", e))?; + let config_str = + fs::read_to_string(&config_path).map_err(|e| format!("Failed to read test.toml: {}", e))?; toml::from_str(&config_str).map_err(|e| format!("Failed to parse test.toml: {}", e)) } - diff --git a/website/mlf-playground-wasm/src/lib.rs b/website/mlf-playground-wasm/src/lib.rs index c687c96..3c98d2d 100644 --- a/website/mlf-playground-wasm/src/lib.rs +++ b/website/mlf-playground-wasm/src/lib.rs @@ -3,15 +3,12 @@ pub use mlf_wasm::*; // Import the plugin crates and reference their static generators // This forces the linker to include them in the binary -use mlf_codegen_typescript::TYPESCRIPT_GENERATOR; use mlf_codegen_go::GO_GENERATOR; use mlf_codegen_rust::RUST_GENERATOR; +use mlf_codegen_typescript::TYPESCRIPT_GENERATOR; // Force the linker to keep the generator statics by referencing them // This function must never be optimized away #[used] -static _KEEP_GENERATORS: &[&dyn mlf_codegen::plugin::CodeGenerator] = &[ - &TYPESCRIPT_GENERATOR, - &GO_GENERATOR, - &RUST_GENERATOR, -]; +static _KEEP_GENERATORS: &[&dyn mlf_codegen::plugin::CodeGenerator] = + &[&TYPESCRIPT_GENERATOR, &GO_GENERATOR, &RUST_GENERATOR]; -- 2.51.2