diff --git a/.claude/skills/docs-search/SKILL.md b/.claude/skills/docs-search/SKILL.md new file mode 100644 index 00000000..910f0464 --- /dev/null +++ b/.claude/skills/docs-search/SKILL.md @@ -0,0 +1,32 @@ +--- +name: docs-search +description: Agentic search over the DSRs docs via Mixedbread toast-1. Use when answering questions about DSRs concepts, components, or APIs (signatures, predict, modules, IR, holes, optimizers, .dsrs files) instead of grepping docs/ page by page. +--- + +# DSRs docs search + +The published docs are indexed in a Mixedbread Store (`dsrs-docs`) and +the Rust sources in `dsrs-code`; both are queried with toast-1 agentic +search, which decomposes the question into subqueries and returns +curated chunks. `ask` can also grep the code store for exact symbols. + +Run from the repository root (needs `mixedbread` installed and +`MXBAI_API_KEY` in the environment or `.env`): + +```bash +python3 docs/scripts/search_docs.py search "" [--top-k N] [--json] +python3 docs/scripts/search_docs.py ask "" +``` + +`search` returns raw chunks with source paths — prefer it when you plan to +read the pages yourself. `ask` has toast-1 compose a grounded answer — +prefer it for a quick factual check. + +- Phrase the query as a full question, not keywords — toast-1 plans its + own subqueries from it. +- Each result prints the source page path (e.g. + `docs/components/holes.mdx`) and the chunk text; expect ~10s latency. +- Read the source page under `docs/` when a chunk looks truncated. +- If results look stale relative to the working tree, the index only + tracks pushed docs — trust the local files and optionally run + `python3 docs/scripts/search_docs.py sync` to refresh. diff --git a/.gitignore b/.gitignore index e2255933..867f1c63 100644 --- a/.gitignore +++ b/.gitignore @@ -42,3 +42,18 @@ thoughts/ # Reference copies of upstream source /reference/ .claude/worktrees/ + +# GEPA runs +docs-chat-worker/gepa/runs/ + +# Python caches +__pycache__/ + +# Wrangler caches +.wrangler/ + +# GEPA artifacts +docs-chat-worker/gepa/*.json +docs-chat-worker/gepa/*.jsonl +docs-chat-worker/gepa/*.log +docs-chat-worker/gepa/*.txt \ No newline at end of file diff --git a/Cargo.lock b/Cargo.lock index 89fa68aa..41aac312 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1018,16 +1018,15 @@ dependencies = [ "bon", "cranelift-entity", "csv", + "dsrs-syntax", "dsrs-tools", "dsrs_macros", "enum_dispatch", "facet", - "facet-reflect", "foyer", "futures", "hf-hub", "indexmap", - "kdam", "minijinja", "parquet", "rand 0.8.5", @@ -1062,6 +1061,13 @@ dependencies = [ "tokio", ] +[[package]] +name = "dsrs-syntax" +version = "0.1.0" +dependencies = [ + "serde_json", +] + [[package]] name = "dsrs-tools" version = "0.1.0" @@ -1082,11 +1088,11 @@ name = "dsrs_macros" version = "0.7.2" dependencies = [ "dspy-rs", + "dsrs-syntax", "minijinja", "proc-macro-crate", "proc-macro2", "quote", - "serde_json", "syn 2.0.106", "trybuild", ] @@ -1187,7 +1193,8 @@ dependencies = [ [[package]] name = "facet" version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e338357cf598728b41e45744d024bdc063338214992361766928a1421bd7541d" dependencies = [ "autocfg", "facet-core", @@ -1197,7 +1204,8 @@ dependencies = [ [[package]] name = "facet-core" version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a63e0ade4c53b40220614b8fc2a0a0ce21975941b553081521a195c848b2e9c2" dependencies = [ "autocfg", "const-fnv1a-hash", @@ -1208,7 +1216,8 @@ dependencies = [ [[package]] name = "facet-macro-parse" version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83ea29147986d0e184600cec533c41d6065c3c3d4b5b5745a8403494ca216b09" dependencies = [ "facet-macro-types", "proc-macro2", @@ -1218,7 +1227,8 @@ dependencies = [ [[package]] name = "facet-macro-types" version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31b0035cf41c0d4eeee82effc9161512d216d1378dd89c4d8721258429e38597" dependencies = [ "proc-macro2", "quote", @@ -1228,7 +1238,8 @@ dependencies = [ [[package]] name = "facet-macros" version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77a784f2fa36d3165b95639af790249dee0d8efdef7d53f9417cace91697e2e3" dependencies = [ "facet-macros-impl", ] @@ -1236,7 +1247,8 @@ dependencies = [ [[package]] name = "facet-macros-impl" version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf8f45c6380398bf74e59b97a20012de571502c609e580d84579d1140e491c1c" dependencies = [ "facet-macro-parse", "facet-macro-types", @@ -1245,26 +1257,6 @@ dependencies = [ "unsynn", ] -[[package]] -name = "facet-path" -version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" -dependencies = [ - "facet-core", -] - -[[package]] -name = "facet-reflect" -version = "0.43.2" -source = "git+https://github.com/darinkishore/facet?rev=cc8613c97cd1ec03e63659db34a947989b45c8a5#cc8613c97cd1ec03e63659db34a947989b45c8a5" -dependencies = [ - "facet-core", - "facet-path", - "hashbrown 0.16.1", - "regex", - "smallvec 2.0.0-alpha.12", -] - [[package]] name = "fastant" version = "0.1.10" @@ -1661,8 +1653,6 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" dependencies = [ "allocator-api2", - "equivalent", - "foldhash 0.2.0", ] [[package]] @@ -1776,7 +1766,7 @@ dependencies = [ "itoa", "pin-project-lite", "pin-utils", - "smallvec 1.15.1", + "smallvec", "tokio", "want", ] @@ -1900,7 +1890,7 @@ dependencies = [ "icu_normalizer_data", "icu_properties", "icu_provider", - "smallvec 1.15.1", + "smallvec", "zerovec", ] @@ -1975,7 +1965,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" dependencies = [ "idna_adapter", - "smallvec 1.15.1", + "smallvec", "utf8_iter", ] @@ -2115,16 +2105,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "kdam" -version = "0.6.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5740f66a8d86a086ebcacfb937070e8be6eb2f8fb45e4ae7fa428ca2a98a7b1f" -dependencies = [ - "terminal_size", - "windows-sys 0.59.0", -] - [[package]] name = "lazy_static" version = "1.5.0" @@ -2705,7 +2685,7 @@ dependencies = [ "cfg-if", "libc", "redox_syscall", - "smallvec 1.15.1", + "smallvec", "windows-targets 0.52.6", ] @@ -3627,12 +3607,6 @@ version = "1.15.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" -[[package]] -name = "smallvec" -version = "2.0.0-alpha.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ef784004ca8777809dcdad6ac37629f0a97caee4c685fcea805278d81dd8b857" - [[package]] name = "snap" version = "1.1.1" @@ -3800,16 +3774,6 @@ dependencies = [ "winapi-util", ] -[[package]] -name = "terminal_size" -version = "0.4.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "60b8cb979cb11c32ce1603f8137b22262a9d131aaa5c37b5678025f22b8becd0" -dependencies = [ - "rustix", - "windows-sys 0.60.2", -] - [[package]] name = "thiserror" version = "1.0.69" @@ -4134,7 +4098,7 @@ dependencies = [ "once_cell", "regex-automata", "sharded-slab", - "smallvec 1.15.1", + "smallvec", "thread_local", "tracing", "tracing-core", diff --git a/Cargo.toml b/Cargo.toml index ab61bf4d..26730c38 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,7 +4,26 @@ members = [ "crates/*", ] -[patch.crates-io] -# TODO(dsrs-facet-pin): switch back to upstream main/release once #2040/#2041 are merged and released. -facet = { git = "https://github.com/darinkishore/facet", rev = "cc8613c97cd1ec03e63659db34a947989b45c8a5" } -facet-reflect = { git = "https://github.com/darinkishore/facet", rev = "cc8613c97cd1ec03e63659db34a947989b45c8a5" } +# Single source of truth for every dependency used by 2+ member crates. +# Pins match what Cargo.lock currently resolves — bump here, not in members. +[workspace.dependencies] +# In-workspace crates +dspy-rs = { version = "0.7.3", path = "crates/dspy-rs" } +dsrs-tools = { version = "0.1.0", path = "crates/dsrs-tools" } +dsrs-syntax = { version = "0.1.0", path = "crates/dsrs-syntax" } + +# Shared third-party crates +anyhow = "1.0.99" +async-trait = "0.1.83" +futures = "0.3.31" +reqwest = "0.13" +serde = { version = "1.0.219", features = ["derive"] } +serde_json = { version = "1.0.143", features = ["preserve_order"] } +tempfile = "3.23.0" +thiserror = "2.0.17" +tokio = "1.46.1" + +# Shared git pins (features are enabled per member) +rig-core = { git = "https://github.com/0xPlaygrounds/rig", rev = "aee3b8bf6576ce41c9ac1dd82520752a65fa0127" } +minijinja = { git = "https://github.com/boundaryml/minijinja.git", branch = "main", default-features = false } + diff --git a/crates/dspy-rs/Cargo.toml b/crates/dspy-rs/Cargo.toml index b117414c..e6df4ef2 100644 --- a/crates/dspy-rs/Cargo.toml +++ b/crates/dspy-rs/Cargo.toml @@ -14,52 +14,53 @@ exclude = [ ] [dependencies] -futures = "0.3.31" +futures = { workspace = true } indexmap = { version = "2.10.0", features = ["serde"] } rayon = "1.10.0" -rstest = "0.25.0" -serde = { version = "1.0.219", features = ["derive"] } -serde_json = { version = "1.0.140", features = ["preserve_order"] } -tokio = { version = "1.46.1", features = ["full"] } -async-trait = "0.1.83" -anyhow = "1.0.99" +serde = { workspace = true } +serde_json = { workspace = true } +tokio = { workspace = true, features = ["full"] } +async-trait = { workspace = true } +anyhow = { workspace = true } bon = "3.7.0" -# Keep this direct pin in sync with workspace [patch.crates-io] for self-sufficient external path consumers. -facet = { git = "https://github.com/darinkishore/facet", rev = "cc8613c97cd1ec03e63659db34a947989b45c8a5", default-features = false, features = ["std", "doc"] } -facet-reflect = { git = "https://github.com/darinkishore/facet", rev = "cc8613c97cd1ec03e63659db34a947989b45c8a5" } -thiserror = "2.0.17" +# Upstream crates.io facet (the reflection-walker fork pin was removed in phase 3). +facet = { version = "0.43", default-features = false, features = ["std", "doc"] } +thiserror = { workspace = true } dsrs_macros = { version = "0.7.2", path = "../dsrs-macros" } -csv = { version = "1.3.1" } -hf-hub = { version = "0.4.3", features = ["tokio"] } -parquet = { version = "56.1.0" } -arrow = { version = "56.1.0" } +# Shared .dsrs lexer/structural grammar (also what include_program! validates with). +dsrs-syntax = { workspace = true } +# Heavy dataset-ingestion stack, gated behind the `data` feature (default-on). +csv = { version = "1.3.1", optional = true } +hf-hub = { version = "0.4.3", features = ["tokio"], optional = true } +parquet = { version = "56.1.0", optional = true } +arrow = { version = "56.1.0", optional = true } regex = "1.11.2" -reqwest = { version = "0.13", features = ["blocking"] } -kdam = "0.6.3" +reqwest = { workspace = true, features = ["blocking"] } rand = "0.8.5" foyer = { version = "0.20.0", features = ["serde"]} -tempfile = "3.23.0" -rig-core = { git = "https://github.com/0xPlaygrounds/rig", rev="aee3b8bf6576ce41c9ac1dd82520752a65fa0127" } +# Runtime dep (not dev-only): utils/cache.rs owns a TempDir for the disk cache tier. +tempfile = { workspace = true } +rig-core = { workspace = true } enum_dispatch = "0.3.13" tracing = "0.1.44" tracing-subscriber = { version = "0.3.22", features = ["env-filter", "fmt"] } -minijinja = { git = "https://github.com/boundaryml/minijinja.git", branch = "main", default-features = false, features = ["builtins", "serde", "debug"] } +minijinja = { workspace = true, features = ["builtins", "serde", "debug"] } # IR graph core (RFC 0002 §2): entity arenas + the sandboxed hole executor. -cranelift-entity = { version = "0.134", features = ["enable-serde"], optional = true } -dsrs-tools = { version = "0.1.0", path = "../dsrs-tools", optional = true } +# Unconditional since phase 2β: Predict runs through the IR Interpreter, so +# the graph core (and Code Mode via dsrs-tools) is load-bearing, not optional. +cranelift-entity = { version = "0.134", features = ["enable-serde"] } +dsrs-tools = { version = "0.1.0", path = "../dsrs-tools" } [package.metadata.cargo-machete] ignored = ["rig-core"] [features] -# RFC 0002 IR-3: the graph core + interpreter ship behind the `ir` feature -# (default-on) until the text format stage stabilizes the surface. -default = ["ir", "code-mode"] -ir = ["dep:cranelift-entity", "dep:dsrs-tools"] -# Code Mode (vision report §5.5): collapse a tool set into one sandboxed -# `run_js` tool. Default-on costs nothing extra (`ir` already pulls -# dsrs-tools); build with --no-default-features for the dep-light library. -code-mode = ["dep:dsrs-tools"] +default = ["data"] +# CSV/Parquet/HuggingFace dataset ingestion (the arrow stack). JSON/JSONL +# loading is always available. Build with --no-default-features for a +# meaningfully lighter dependency tree. +data = ["dep:arrow", "dep:parquet", "dep:hf-hub", "dep:csv"] [dev-dependencies] +rstest = "0.25.0" temp-env = { version = "0.3.6", features = ["async_closure"] } diff --git a/crates/dspy-rs/examples/01-simple.rs b/crates/dspy-rs/examples/01-simple.rs index 6ef0b252..04fb5f37 100644 --- a/crates/dspy-rs/examples/01-simple.rs +++ b/crates/dspy-rs/examples/01-simple.rs @@ -16,9 +16,7 @@ cargo run --example 01-simple use anyhow::Result; use bon::Builder; -use dspy_rs::{ - CallMetadata, Demo, LM, Module, Predict, PredictError, Predicted, configure, init_tracing, -}; +use dspy_rs::prelude::*; const QA_INSTRUCTION: &str = "Answer the question step by step."; const RATE_INSTRUCTION: &str = "Rate the answer on a scale of 1 (very bad) to 10 (very good)."; diff --git a/crates/dspy-rs/examples/02-module-iteration-and-updation.rs b/crates/dspy-rs/examples/02-module-iteration-and-updation.rs index 213de17a..181feeff 100644 --- a/crates/dspy-rs/examples/02-module-iteration-and-updation.rs +++ b/crates/dspy-rs/examples/02-module-iteration-and-updation.rs @@ -9,10 +9,7 @@ cargo run --example 02-module-iteration-and-updation use anyhow::Result; use bon::Builder; -use dspy_rs::{ - COPRO, LM, Eval, Module, Optimizer, Predict, PredictError, - Predicted, Signature, TypedMetric, average_score, configure, evaluate_trainset, init_tracing, -}; +use dspy_rs::prelude::*; #[derive(Signature, Clone, Debug)] struct QA { @@ -30,6 +27,8 @@ struct QAModule { answerer: Predict, } +dspy_rs::predictors!(QAModule { answerer }); + impl Module for QAModule { type Input = QAInput; type Output = QAOutput; @@ -96,7 +95,7 @@ async fn main() -> Result<()> { let optimizer = COPRO::builder().breadth(4).depth(1).build(); optimizer - .compile(&mut module, trainset.clone(), &metric) + .compile_module(&mut module, &trainset, &metric) .await?; let optimized = average_score(&evaluate_trainset(&module, &trainset, &metric).await?); diff --git a/crates/dspy-rs/examples/03-evaluate-hotpotqa.rs b/crates/dspy-rs/examples/03-evaluate-hotpotqa.rs index 39a53eed..5328eb66 100644 --- a/crates/dspy-rs/examples/03-evaluate-hotpotqa.rs +++ b/crates/dspy-rs/examples/03-evaluate-hotpotqa.rs @@ -8,11 +8,7 @@ cargo run --example 03-evaluate-hotpotqa --features dataloaders */ use anyhow::Result; -use dspy_rs::{ - DataLoader, Example, LM, Eval, Predict, Predicted, Signature, - TypedLoadOptions, TypedMetric, average_score, configure, evaluate_trainset_with_concurrency, - init_tracing, -}; +use dspy_rs::prelude::*; #[derive(Signature, Clone, Debug)] struct QA { diff --git a/crates/dspy-rs/examples/04-optimize-hotpotqa.rs b/crates/dspy-rs/examples/04-optimize-hotpotqa.rs index 814add63..dcb3179a 100644 --- a/crates/dspy-rs/examples/04-optimize-hotpotqa.rs +++ b/crates/dspy-rs/examples/04-optimize-hotpotqa.rs @@ -9,11 +9,7 @@ cargo run --example 04-optimize-hotpotqa --features dataloaders use anyhow::Result; use bon::Builder; -use dspy_rs::{ - COPRO, DataLoader, Example, LM, Eval, Module, ModuleState, Optimizer, - Predict, PredictError, Predicted, Signature, TypedLoadOptions, TypedMetric, average_score, - configure, evaluate_trainset, init_tracing, -}; +use dspy_rs::prelude::*; #[derive(Signature, Clone, Debug)] struct QA { @@ -42,6 +38,8 @@ struct QAModule { answerer: Predict, } +dspy_rs::predictors!(QAModule { answerer }); + impl Module for QAModule { type Input = QAInput; type Output = QAOutput; @@ -96,14 +94,14 @@ async fn main() -> Result<()> { .eval_concurrency(16) // candidate evaluations fan out 16 LM calls at a time .build(); optimizer - .compile(&mut module, examples.clone(), &metric) + .compile_module(&mut module, &examples, &metric) .await?; let optimized = average_score(&evaluate_trainset(&module, &examples, &metric).await?); println!("optimized score: {optimized:.3}"); // Persist the tuned instructions for later `ModuleState::load(...).apply(...)`. - ModuleState::from_module(&mut module)?.save("optimized-hotpotqa.json")?; + ModuleState::from_module(&module)?.save("optimized-hotpotqa.json")?; println!("saved optimized module state to optimized-hotpotqa.json"); Ok(()) diff --git a/crates/dspy-rs/examples/05-heterogenous-examples.rs b/crates/dspy-rs/examples/05-heterogenous-examples.rs index 944dce07..c3a45e80 100644 --- a/crates/dspy-rs/examples/05-heterogenous-examples.rs +++ b/crates/dspy-rs/examples/05-heterogenous-examples.rs @@ -13,7 +13,7 @@ cargo run --example 05-heterogenous-examples */ use anyhow::Result; -use dspy_rs::{LM, Predict, Signature, configure, init_tracing}; +use dspy_rs::prelude::*; use serde_json::json; #[derive(Signature, Clone, Debug)] diff --git a/crates/dspy-rs/examples/08-optimize-mipro.rs b/crates/dspy-rs/examples/08-optimize-mipro.rs index 6ca710e0..9d55233f 100644 --- a/crates/dspy-rs/examples/08-optimize-mipro.rs +++ b/crates/dspy-rs/examples/08-optimize-mipro.rs @@ -43,6 +43,8 @@ struct SimpleQA { answerer: Predict, } +dspy_rs::predictors!(SimpleQA { answerer }); + impl Module for SimpleQA { type Input = QuestionAnsweringInput; type Output = QuestionAnsweringOutput; @@ -122,11 +124,11 @@ async fn main() -> Result<()> { println!("Starting MIPROv2 optimization..."); optimizer - .compile(&mut qa_module, train_subset.clone(), &metric) + .compile_module(&mut qa_module, &train_subset, &metric) .await?; // Inspect what the optimizer installed: instructions + bootstrapped demos. - let state = ModuleState::from_module(&mut qa_module)?; + let state = ModuleState::from_module(&qa_module)?; for (predictor, predictor_state) in &state.predictors { println!( "Predictor `{predictor}`: {} bootstrapped demos, instruction override: {}", diff --git a/crates/dspy-rs/examples/09-gepa-sentiment.rs b/crates/dspy-rs/examples/09-gepa-sentiment.rs index b4d4f624..241d4ee6 100644 --- a/crates/dspy-rs/examples/09-gepa-sentiment.rs +++ b/crates/dspy-rs/examples/09-gepa-sentiment.rs @@ -41,6 +41,8 @@ struct SentimentAnalyzer { predictor: Predict, } +dspy_rs::predictors!(SentimentAnalyzer { predictor }); + impl Module for SentimentAnalyzer { type Input = SentimentSignatureInput; type Output = SentimentSignatureOutput; @@ -130,7 +132,7 @@ async fn main() -> Result<()> { .track_stats(true) .build(); - let result = gepa.compile(&mut module, trainset.clone(), &metric).await?; + let result = gepa.compile_module(&mut module, &trainset, &metric).await?; println!( "Best average score: {:.3}", @@ -159,7 +161,7 @@ async fn main() -> Result<()> { // Persist the optimized instructions/demos so production can reload them // with `ModuleState::load(...)?.apply(&mut module)?` — no re-optimization. - ModuleState::from_module(&mut module)?.save("optimized-sentiment.json")?; + ModuleState::from_module(&module)?.save("optimized-sentiment.json")?; println!("Saved optimized module state to optimized-sentiment.json"); Ok(()) diff --git a/crates/dspy-rs/examples/10-gepa-llm-judge.rs b/crates/dspy-rs/examples/10-gepa-llm-judge.rs index 95541b3f..f1ba77a7 100644 --- a/crates/dspy-rs/examples/10-gepa-llm-judge.rs +++ b/crates/dspy-rs/examples/10-gepa-llm-judge.rs @@ -60,6 +60,8 @@ struct MathSolver { solver: Predict, } +dspy_rs::predictors!(MathSolver { solver }); + impl Module for MathSolver { type Input = MathWordProblemInput; type Output = MathWordProblemOutput; @@ -194,7 +196,7 @@ async fn main() -> Result<()> { .track_stats(true) .build(); - let result = gepa.compile(&mut module, trainset.clone(), &metric).await?; + let result = gepa.compile_module(&mut module, &trainset, &metric).await?; println!("Best score: {:.3}", result.best_candidate.average_score()); println!("Total rollouts: {}", result.total_rollouts); diff --git a/crates/dspy-rs/examples/12-tracing.rs b/crates/dspy-rs/examples/12-tracing.rs index e0e58f2b..bf0cc397 100644 --- a/crates/dspy-rs/examples/12-tracing.rs +++ b/crates/dspy-rs/examples/12-tracing.rs @@ -129,16 +129,15 @@ async fn main() -> Result<()> { } // Each Predict call records one span: component name, invocation seq, - // typed input/output as JSON, a link to the previous span, and timing. + // typed input/output as JSON, and timing. // Failed calls stay visible with the prompt recorded and output absent. println!("Trace {} spans: {}", trace.meta.trace_id, trace.spans.len()); for span in &trace.spans { println!( - "Span {}: component={:?} seq={} links={:?} events={}", + "Span {}: component={:?} seq={} events={}", span.id.0, trace.component_name(span.component), span.seq, - span.links, span.events.len(), ); if let Some(input) = &span.input { diff --git a/crates/dspy-rs/examples/13-save-load-state.rs b/crates/dspy-rs/examples/13-save-load-state.rs index fccc9d4e..7e223280 100644 --- a/crates/dspy-rs/examples/13-save-load-state.rs +++ b/crates/dspy-rs/examples/13-save-load-state.rs @@ -36,6 +36,8 @@ struct QAPipeline { answerer: Predict, } +dspy_rs::predictors!(QAPipeline { answerer }); + impl Module for QAPipeline { type Input = QAInput; type Output = QAOutput; @@ -66,7 +68,7 @@ fn main() -> Result<()> { .build(); // --- Save ------------------------------------------------------------- - let state = ModuleState::from_module(&mut tuned)?; + let state = ModuleState::from_module(&tuned)?; println!("Snapshot of {} predictor(s):", state.predictors.len()); for (path, predictor_state) in &state.predictors { println!( @@ -84,7 +86,7 @@ fn main() -> Result<()> { let mut fresh = QAPipeline::builder().build(); ModuleState::load("qa-pipeline-state.json")?.apply(&mut fresh)?; - let restored = ModuleState::from_module(&mut fresh)?; + let restored = ModuleState::from_module(&fresh)?; let answerer = &restored.predictors["answerer"]; println!( "Restored `answerer`: instruction={:?}, demos={}", diff --git a/crates/dspy-rs/examples/14-functional.rs b/crates/dspy-rs/examples/14-functional.rs index d41c3a62..0f10ac39 100644 --- a/crates/dspy-rs/examples/14-functional.rs +++ b/crates/dspy-rs/examples/14-functional.rs @@ -112,11 +112,10 @@ async fn main() -> Result<()> { // parameter is one `for_component` call. for span in &trace.spans { println!( - " span {} = {:?} seq={} (links: {:?})", + " span {} = {:?} seq={}", span.id.0, trace.component_name(span.component), span.seq, - span.links ); } let refiner_calls = trace.for_component("refiner").count(); diff --git a/crates/dspy-rs/examples/16-insurance-claim-prompt.rs b/crates/dspy-rs/examples/16-insurance-claim-prompt.rs index da712dc9..057c73d5 100644 --- a/crates/dspy-rs/examples/16-insurance-claim-prompt.rs +++ b/crates/dspy-rs/examples/16-insurance-claim-prompt.rs @@ -5,14 +5,14 @@ Run with: cargo run --example 16-insurance-claim-prompt */ -use dspy_rs::{BamlType, ChatAdapter, Signature, init_tracing}; +use dspy_rs::{Schema, ChatAdapter, Signature, init_tracing}; // Keep the example self-contained; dates are represented as YYYY-MM-DD strings. type NaiveDate = String; /// Basic claim information (metadata about the claim intake). #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub struct ClaimHeader { /// Claim ID in format `CLM-XXXXXX`, where `X` is a digit. pub claim_id: Option, @@ -32,7 +32,7 @@ pub struct ClaimHeader { /// Channel used to report a claim. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub enum ClaimChannel { Email, Phone, @@ -42,7 +42,7 @@ pub enum ClaimChannel { /// Policy information if available. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub struct PolicyDetails { /// Policy number in format `POL-XXXXXXXXX`, where `X` is a digit. pub policy_number: Option, @@ -62,7 +62,7 @@ pub struct PolicyDetails { /// Type of insurance coverage. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub enum CoverageType { Property, Auto, @@ -74,7 +74,7 @@ pub enum CoverageType { /// An insured object involved in the claim (vehicle, building, person, etc.). #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub struct InsuredObject { /// Unique identifier for insured object. /// @@ -104,7 +104,7 @@ pub struct InsuredObject { /// Type of insured object. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub enum InsuredObjectType { Vehicle, Building, @@ -114,7 +114,7 @@ pub enum InsuredObjectType { /// Structured incident details. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub struct IncidentDescription { /// Specific standardized incident type. pub incident_type: IncidentType, @@ -131,7 +131,7 @@ pub struct IncidentDescription { /// Specific standardized incident type. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub enum IncidentType { RearEndCollision, SideImpactCollision, @@ -152,7 +152,7 @@ pub enum IncidentType { /// Standardized location type where incident occurred. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub enum LocationType { Intersection, Highway, @@ -167,7 +167,7 @@ pub enum LocationType { /// Top-level insurance claim object aggregating all extracted fields. #[derive(Debug, Clone, PartialEq, Eq)] -#[BamlType] +#[Schema] pub struct InsuranceClaim { /// Basic claim information. pub header: ClaimHeader, @@ -197,12 +197,17 @@ fn main() { init_tracing().expect("failed to initialize tracing"); let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system prompt"); - let user = adapter.format_user_message_typed::(&InsuranceClaimInfoInput { + let def = dspy_rs::ir::SignatureDef::of::(); + let types = dspy_rs::ir::SignatureDef::types_of::(); + let system = adapter.build_system_def(def, types, None); + let input = InsuranceClaimInfoInput { claim_text: "A raccoon bumped a parked scooter in a driveway. Reported by Taylor P. via phone. No policy details provided.".to_string(), - }); + }; + let input_map = match serde_json::to_value(&input).expect("serializable input") { + serde_json::Value::Object(map) => map, + _ => unreachable!("signature inputs serialize as objects"), + }; + let user = adapter.format_input_def(def, &input_map); println!("=== System ===\n{system}\n"); println!("=== User ===\n{user}"); diff --git a/crates/dspy-rs/examples/19-react.rs b/crates/dspy-rs/examples/19-react.rs deleted file mode 100644 index 509a7048..00000000 --- a/crates/dspy-rs/examples/19-react.rs +++ /dev/null @@ -1,80 +0,0 @@ -/* -Example: ReAct — the think/act/observe loop as a drop-in strategy. - -`ReAct` runs a bounded loop over any signature: the model reads the input, -decides whether it needs a tool, calls it, reads the observation, and repeats -until it finishes (or hits `max_steps`). Same contract as `Predict` / -`ChainOfThought` — the signature never hears about the strategy. - -Run with: -``` -cargo run --example 19-react -``` -*/ - -use anyhow::Result; -use dspy_rs::{LM, Module, ReAct, Signature, configure, init_tracing}; - -/// Answer the customer's support question. Use the tools to look up facts; -/// report only what you actually found. -#[derive(Signature, Clone, Debug)] -struct SupportAnswer { - #[input] - question: String, - - #[output] - answer: String, -} - -#[tokio::main] -async fn main() -> Result<()> { - init_tracing()?; - - configure( - LM::builder() - .model("openai:gpt-4o-mini".to_string()) - .build() - .await?, - ); - - // A tool is a name, a description the model reads, and an async closure. - // `max_steps` is the leash: the loop always terminates. - let agent = ReAct::::builder() - .tool( - "order_lookup", - "Look up an order by its number. Returns status, carrier, and ETA.", - |args: String| async move { - let order_id: String = args - .trim() - .trim_matches(|c| c == '"' || c == '\'') - .chars() - .filter(char::is_ascii_digit) - .collect(); - match order_id.as_str() { - "4127" => "status: shipped, carrier: UPS, eta: Thursday".to_string(), - other => format!("no order with id {other}"), - } - }, - ) - .max_steps(4) - .build(); - - let predicted = agent - .call(SupportAnswerInput { - question: "Where is my order? I paid for the ceramic starter jar two weeks ago. \ - Order #4127." - .to_string(), - }) - .await?; - - println!("answer: {}", predicted.answer); - - // The trajectory (thoughts, actions, observations) rides in metadata. - let metadata = predicted.metadata(); - println!("\ntool calls: {}", metadata.tool_calls.len()); - for entry in &metadata.tool_executions { - println!("---\n{entry}"); - } - - Ok(()) -} diff --git a/crates/dspy-rs/examples/20-frontdesk-contract.rs b/crates/dspy-rs/examples/20-frontdesk-contract.rs index 7955fe45..59c08d3c 100644 --- a/crates/dspy-rs/examples/20-frontdesk-contract.rs +++ b/crates/dspy-rs/examples/20-frontdesk-contract.rs @@ -125,7 +125,14 @@ async fn main() -> Result<()> { // into the prompt — print it instead of trusting anyone's word. let adapter = ChatAdapter; println!("=== the prompt Triage becomes ===\n"); - println!("{}", adapter.format_system_message_typed::()?); + println!( + "{}", + adapter.build_system_def( + dspy_rs::ir::SignatureDef::of::(), + dspy_rs::ir::SignatureDef::types_of::(), + None, + ) + ); configure( LM::builder() diff --git a/crates/dspy-rs/examples/93-smoke-slice4-react-operational.rs b/crates/dspy-rs/examples/93-smoke-slice4-react-operational.rs deleted file mode 100644 index 2ae2570d..00000000 --- a/crates/dspy-rs/examples/93-smoke-slice4-react-operational.rs +++ /dev/null @@ -1,129 +0,0 @@ -use anyhow::{Result, bail}; -use dspy_rs::{LM, PredictError, ReAct, Signature, configure, forward_all}; -use serde_json::Value; - -#[derive(Signature, Clone, Debug)] -struct SmokeSig { - #[input] - prompt: String, - - #[output] - answer: String, -} - -fn parse_binary_args(args: &str) -> Result<(i64, i64)> { - let value: Value = serde_json::from_str(args)?; - let a = value.get("a").and_then(Value::as_i64).unwrap_or(0); - let b = value.get("b").and_then(Value::as_i64).unwrap_or(0); - Ok((a, b)) -} - -fn extract_first_integer(text: &str) -> Option { - let mut token = String::new(); - for ch in text.chars() { - if ch.is_ascii_digit() || (token.is_empty() && ch == '-') { - token.push(ch); - continue; - } - if !token.is_empty() { - break; - } - } - token.parse::().ok() -} - -#[tokio::main] -async fn main() -> Result<()> { - // Smoke Label: Slice 4 ReAct + Operational - configure(LM::builder() - .model("openai:gpt-5.2".to_string()) - .build() - .await?); - - let module = ReAct::::builder() - .max_steps(6) - .tool("add", "Add two integers. Args JSON: {\"a\":int,\"b\":int}", |args| async move { - match parse_binary_args(&args) { - Ok((a, b)) => (a + b).to_string(), - Err(err) => format!("calculator_error: {err}"), - } - }) - .tool( - "multiply", - "Multiply two integers. Args JSON: {\"a\":int,\"b\":int}", - |args| async move { - match parse_binary_args(&args) { - Ok((a, b)) => (a * b).to_string(), - Err(err) => format!("calculator_error: {err}"), - } - }, - ) - .action_instruction( - "You are a strict ReAct planner. Choose exactly one tool each step, and use tool names exactly as declared.", - ) - .extract_instruction( - "Read trajectory and return only the final integer in output.answer.", - ) - .build(); - - let input = SmokeSigInput { - prompt: "Use tools to compute ((17 + 5) * 3) + 4. You MUST call add, then multiply, then add again, then finish. Return only the final integer string." - .to_string(), - }; - - let mut outcomes = forward_all(&module, vec![input], 1).await.into_iter(); - let outcome = outcomes.next().expect("expected one batch outcome"); - let predicted = outcome.map_err(|err| { - eprintln!("smoke call failed: {err}"); - if let PredictError::Parse { raw_response, .. } = &err { - eprintln!("raw_response: {:?}", raw_response); - } - anyhow::anyhow!("slice4 smoke failed") - })?; - let (output, metadata) = predicted.into_parts(); - - println!("tool_calls: {}", metadata.tool_calls.len()); - println!("tool_executions: {}", metadata.tool_executions.len()); - println!("trajectory:"); - for entry in &metadata.tool_executions { - if entry.trim().is_empty() { - continue; - } - println!("{entry}"); - println!("---"); - } - println!("answer: {}", output.answer); - - let called_tools: Vec = metadata - .tool_calls - .iter() - .map(|call| call.function.name.to_ascii_lowercase()) - .collect(); - let add_calls = called_tools - .iter() - .filter(|name| name.as_str() == "add") - .count(); - let multiply_calls = called_tools - .iter() - .filter(|name| name.as_str() == "multiply") - .count(); - - if add_calls < 2 || multiply_calls < 1 { - bail!( - "expected multi-tool trajectory with add x2 and multiply x1, got {:?}", - called_tools - ); - } - - let answer_value = extract_first_integer(&output.answer) - .ok_or_else(|| anyhow::anyhow!("answer did not contain integer: {}", output.answer))?; - if answer_value != 70 { - bail!( - "unexpected calculator result: expected 70, got {} (raw answer: {})", - answer_value, - output.answer - ); - } - - Ok(()) -} diff --git a/crates/dspy-rs/examples/94-smoke-slice5-optimizer-interface.rs b/crates/dspy-rs/examples/94-smoke-slice5-optimizer-interface.rs index 43d71790..57675bcd 100644 --- a/crates/dspy-rs/examples/94-smoke-slice5-optimizer-interface.rs +++ b/crates/dspy-rs/examples/94-smoke-slice5-optimizer-interface.rs @@ -50,7 +50,7 @@ async fn main() -> Result<()> { let optimizer = COPRO::builder().breadth(4).depth(1).build(); optimizer - .compile(&mut module, trainset, &SmokeMetric) + .compile_module(&mut module, &trainset, &SmokeMetric) .await?; let output = module diff --git a/crates/dspy-rs/examples/97-perf-microbench.rs b/crates/dspy-rs/examples/97-perf-microbench.rs index 1f723427..557a77a3 100644 --- a/crates/dspy-rs/examples/97-perf-microbench.rs +++ b/crates/dspy-rs/examples/97-perf-microbench.rs @@ -144,7 +144,9 @@ fn demo(idx: usize) -> Demo { Demo::new( BenchQAInput { question: format!("Demo question {idx}?"), - context: format!("Demo context {idx} with enough text to look like a real retrieval chunk for the benchmark."), + context: format!( + "Demo context {idx} with enough text to look like a real retrieval chunk for the benchmark." + ), }, BenchQAOutput { answer: format!("Demo answer {idx}."), @@ -180,8 +182,7 @@ fn report(name: &str, iters: u64, s: Snapshot) { } async fn make_lm(responses: u64, cache: bool) -> LM { - let client = - TestCompletionModel::new((0..responses).map(|_| assistant_content())); + let client = TestCompletionModel::new((0..responses).map(|_| assistant_content())); temp_env::async_with_vars( [("OPENAI_API_KEY", Some("bench"))], LM::builder() @@ -218,16 +219,19 @@ async fn main() { report("schema(): global RwLock map path", iters, s); // --- 2. Prompt build (system + 2 demos + user) -------------------------- + // `build_chat` renders through the loaded interpreter (conversation + // surface), so the predictor needs a bound LM even though no call is made. let predict = Predict::::builder() .demo(demo(1)) .demo(demo(2)) + .lm(make_lm(0, false).await) .build(); let input = bench_input(); let iters = 100_000u64; let s = snap(); for _ in 0..iters { - std::hint::black_box(predict.build_chat(&input).unwrap()); + std::hint::black_box(predict.build_chat(&input).await.unwrap()); } report("build_chat (system + 2 demos + user)", iters, s); @@ -239,11 +243,15 @@ async fn main() { for _ in 0..iters { std::hint::black_box( adapter - .parse_response_typed::(&assistant) + .parse_output_def( + dspy_rs::ir::SignatureDef::of::(), + dspy_rs::ir::SignatureDef::types_of::(), + &assistant, + ) .unwrap(), ); } - report("parse_response_typed (2 fields)", iters, s); + report("parse_output_def (2 fields)", iters, s); // --- 4. Full forward with test client (no demos) ------------------------ let iters = 50_000u64; @@ -300,16 +308,23 @@ async fn main() { for _ in 0..iters { std::hint::black_box( adapter - .parse_response_typed::(&checked_assistant) + .parse_output_def( + dspy_rs::ir::SignatureDef::of::(), + dspy_rs::ir::SignatureDef::types_of::(), + &checked_assistant, + ) .unwrap(), ); } - report("parse_response_typed (2 checks)", iters, s); + report("parse_output_def (2 checks)", iters, s); // --- 9. Forward with 1 tool attached (never called) ----------------------- let iters = 50_000u64; let lm = make_lm(iters, false).await; - let tooled = Predict::::builder().lm(lm).add_tool(NoopTool).build(); + let tooled = Predict::::builder() + .lm(lm) + .add_tool(NoopTool) + .build(); let s = snap(); for _ in 0..iters { std::hint::black_box(tooled.call(bench_input()).await.unwrap()); @@ -322,7 +337,9 @@ async fn main() { let s = snap(); for _ in 0..iters { std::hint::black_box( - fx::predict::("bench_fx", bench_input()).await.unwrap(), + fx::predict::("bench_fx", bench_input()) + .await + .unwrap(), ); } report("fx::predict (0 demos, default params)", iters, s); @@ -336,7 +353,9 @@ async fn main() { fx::with_params(params, async { for _ in 0..iters { std::hint::black_box( - fx::predict::("bench_fx", bench_input()).await.unwrap(), + fx::predict::("bench_fx", bench_input()) + .await + .unwrap(), ); } }) diff --git a/crates/dspy-rs/src/adapter/chat.rs b/crates/dspy-rs/src/adapter/chat.rs index da3815ec..5a0db561 100644 --- a/crates/dspy-rs/src/adapter/chat.rs +++ b/crates/dspy-rs/src/adapter/chat.rs @@ -4,34 +4,27 @@ use minijinja::UndefinedBehavior; use minijinja::value::{Kwargs, Value as MiniJinjaValue}; use regex::Regex; use serde_json::{Map, Value, json}; -use std::collections::HashMap; -use std::sync::{Arc, LazyLock, RwLock}; +use std::sync::LazyLock; use tracing::{debug, trace}; -use crate::CallMetadata; use crate::ir::{RenderSpec, SignatureDef}; use crate::trace::JsonMap; use crate::typesys::coerce::coerce; -use crate::typesys::constraint::{evaluate_constraint_expression, evaluate_expression}; +use crate::typesys::constraint::evaluate_expression; use crate::typesys::render::{schema_block, type_name}; use crate::typesys::{FieldType, TypeTable}; use crate::{ - ConstraintKind, ConstraintResult, FieldMeta, FieldSchema, InputRenderSpec, JsonishError, - Message, ParseError, PredictError, Predicted, Schema, Signature, + ConstraintKind, ConstraintResult, FieldMeta, JsonishError, Message, ParseError, }; /// Builds prompts and parses responses using the `[[ ## field ## ]]` delimiter protocol. /// -/// The adapter is stateless — all state comes from the [`SignatureSchema`](crate::SignatureSchema) -/// passed to each method. Two usage patterns: -/// -/// - **High-level** (what [`Predict`](crate::Predict) uses): `format_system_message_typed`, -/// `format_user_message_typed`, `parse_response_typed` — all parameterized by `S: Signature`. -/// - **Building blocks** (for module authors): `build_system`, `format_input`, `format_output`, -/// `parse_output`, `parse_sections` — parameterized by `&SignatureSchema`, not a Signature type. -/// -/// The building blocks exist so module authors can compose custom prompt flows (e.g. -/// ReAct's action/extract loop) without reimplementing the delimiter protocol. +/// The adapter is stateless — all state comes from the [`SignatureDef`] passed to +/// each method: `build_system_def`, `format_input_def`, `format_output_def`, +/// `parse_output_def`, plus the free `parse_sections`. This is the single prompt +/// lane; [`Predict`](crate::Predict) reaches it through the IR interpreter, and +/// module authors can call the `*_def` building blocks directly to compose custom +/// prompt flows without reimplementing the delimiter protocol. #[derive(Default, Clone)] pub struct ChatAdapter; @@ -40,14 +33,6 @@ static FIELD_HEADER_PATTERN: LazyLock = const INPUT_RENDER_TEMPLATE_NAME: &str = "__input_field__"; -struct CachedInputRenderTemplate { - env: minijinja::Environment<'static>, -} - -static INPUT_RENDER_TEMPLATE_CACHE: LazyLock< - RwLock>>, -> = LazyLock::new(|| RwLock::new(HashMap::new())); - fn regex_match(value: String, regex: String) -> bool { match Regex::new(®ex) { Ok(re) => re.is_match(&value), @@ -131,10 +116,11 @@ fn build_input_render_environment<'source>() -> minijinja::Environment<'source> env } -/// A borrowed, lane-neutral view of one signature field: everything the prompt -/// builders need, whether the source is a `'static` [`FieldSchema`] or an owned -/// [`ir::FieldDef`](crate::ir::FieldDef). Both lanes render through the same -/// view-based functions, so their prompt sections are equal by construction. +/// A borrowed view of one signature field: everything the prompt builders +/// need, sourced from an owned [`ir::FieldDef`](crate::ir::FieldDef). The +/// historical `'static` [`FieldSchema`](crate::FieldSchema) lane rendered +/// through these same view-based functions; with `Predict` routed through the +/// interpreter, the def lane is the only prompt path. struct FieldView<'a> { lm_name: &'a str, docs: &'a str, @@ -142,14 +128,6 @@ struct FieldView<'a> { } impl<'a> FieldView<'a> { - fn of_schema(field: &'a FieldSchema) -> Self { - Self { - lm_name: field.lm_name, - docs: &field.docs, - ty: &field.type_ir, - } - } - fn of_def(field: &'a crate::ir::FieldDef) -> Self { Self { lm_name: &field.lm_name, @@ -159,10 +137,6 @@ impl<'a> FieldView<'a> { } } -fn schema_views(fields: &[FieldSchema]) -> Vec> { - fields.iter().map(FieldView::of_schema).collect() -} - fn def_views(fields: &[crate::ir::FieldDef]) -> Vec> { fields.iter().map(FieldView::of_def).collect() } @@ -305,66 +279,15 @@ fn build_system_view( } impl ChatAdapter { - fn format_response_instructions_schema(&self, schema: &crate::SignatureSchema) -> String { - format_response_instructions_view(&schema_views(schema.output_fields())) - } - - /// Builds the system message for a signature using its default instruction. - /// - /// Shorthand for `format_system_message_typed_with_instruction::(None)`. - pub fn format_system_message_typed(&self) -> Result { - self.format_system_message_typed_with_instruction::(None) - } - - #[tracing::instrument( - name = "dsrs.adapter.chat.format_system_typed", - level = "trace", - skip(self), - fields( - signature = std::any::type_name::(), - instruction_override = instruction_override.is_some() - ) - )] - /// Builds the system message for a signature with an optional instruction override. + /// Builds a system message from an owned [`SignatureDef`] — no `'static` + /// requirement anywhere. `types` resolves the def's class/enum references + /// (for derive-bridged defs, [`SignatureDef::types_of`]). /// /// The system message includes: /// 1. Field descriptions (names, types, doc comments) /// 2. Field structure template (the `[[ ## field ## ]]` layout the LM should follow) /// 3. Response instructions (which fields to produce, in what order) - /// 4. Task description (the signature's instruction or the override) - pub fn format_system_message_typed_with_instruction( - &self, - instruction_override: Option<&str>, - ) -> Result { - self.build_system(S::schema(), instruction_override) - } - - /// Builds a system message from a [`SignatureSchema`](crate::SignatureSchema) directly. - /// - /// The schema-based equivalent of [`format_system_message_typed_with_instruction`](ChatAdapter::format_system_message_typed_with_instruction). - /// Use this when you have a schema but not a concrete `S: Signature` type (e.g. - /// in dynamic or schema-transformed contexts). - pub fn build_system( - &self, - schema: &crate::SignatureSchema, - instruction_override: Option<&str>, - ) -> Result { - Ok(build_system_view( - &schema_views(schema.input_fields()), - &schema_views(schema.output_fields()), - &schema.output_schema().types, - schema.instruction(), - instruction_override, - )) - } - - /// Builds a system message from an owned [`SignatureDef`] — the dynamic-lane - /// twin of [`build_system`](ChatAdapter::build_system), with no `'static` - /// requirement anywhere. `types` resolves the def's class/enum references - /// (for derive-bridged defs, [`SignatureDef::types_of`]). - /// - /// Renders through the same view-based builders as the schema path, so a def - /// bridged via [`SignatureDef::of`] produces byte-identical prompt sections. + /// 4. Task description (the def's instruction or the override) pub fn build_system_def( &self, def: &SignatureDef, @@ -380,68 +303,16 @@ impl ChatAdapter { ) } - /// Formats a typed input value as a user message with `[[ ## field ## ]]` delimiters. - /// - /// Each input field is serialized via serde and formatted according to its field path - /// (handling flattened fields). Appends the response instructions telling the LM which - /// output fields to produce. - pub fn format_user_message_typed(&self, input: &S::Input) -> String - where - S::Input: Schema, - { - self.format_input(S::schema(), input) - } - - /// Formats an input value using a schema — the building-block version of - /// [`format_user_message_typed`](ChatAdapter::format_user_message_typed). + /// Formats a value-level input from an owned [`SignatureDef`] as a user + /// message with `[[ ## field ## ]]` delimiters, with no `'static` + /// requirement anywhere. Appends the response instructions telling the LM + /// which output fields to produce. /// - /// Navigates the serialized JSON using each field's [`FieldPath`](crate::FieldPath) to - /// handle flattened structs correctly. A field with path `["inner", "question"]` is - /// extracted from the flattened structure but rendered as a flat `[[ ## question ## ]]` - /// section in the prompt. Appends response instructions so the LM sees output-field - /// ordering guidance in the latest user turn. - pub fn format_input(&self, schema: &crate::SignatureSchema, input: &I) -> String - where - I: Schema + for<'a> facet::Facet<'a>, - { - let json = serde_json::to_value(input).unwrap_or(Value::Null); - // The aliased input-context tree is only read by `#[render(jinja = ...)]` - // fields — skip the full clone it requires when no field uses Jinja. - let has_jinja_field = schema - .input_fields() - .iter() - .any(|field| matches!(field.input_render, InputRenderSpec::Jinja(_))); - let input_json = if has_jinja_field { - build_input_context_value(schema, &json) - } else { - Value::Null - }; - let vars = Value::Object(Map::new()); - - let mut result = String::new(); - for field_spec in schema.input_fields() { - if let Some(value) = value_for_path_relaxed(&json, field_spec.path()) { - result.push_str(&format!("[[ ## {} ## ]]\n", field_spec.lm_name)); - result.push_str(&render_input_field(field_spec, value, &input_json, &vars)); - result.push_str("\n\n"); - } - } - - result.push_str( - schema - .response_instructions_cached(|| self.format_response_instructions_schema(schema)), - ); - result - } - - /// Formats a value-level input from an owned [`SignatureDef`] — the - /// dynamic-lane twin of [`format_input`](ChatAdapter::format_input), with no - /// `'static` requirement anywhere. - /// - /// Fields absent from `input` are skipped, mirroring the static lane's - /// relaxed path navigation. Jinja render templates are compiled per call — - /// dynamic defs own their template strings, and a process-global cache keyed - /// on them would reintroduce the leak-per-load RFC 0002 IR-1 removed. + /// Fields absent from `input` are skipped (the historical relaxed path + /// navigation for flattened structs). Jinja render templates are compiled + /// per call — dynamic defs own their template strings, and a + /// process-global cache keyed on them would reintroduce the leak-per-load + /// RFC 0002 IR-1 removed. pub fn format_input_def(&self, def: &SignatureDef, input: &JsonMap) -> String { let mut result = String::new(); for field in def.inputs.iter() { @@ -457,48 +328,12 @@ impl ChatAdapter { result } - /// Formats a typed output value as an assistant message for few-shot demos. - /// - /// Each output field is serialized and delimited with `[[ ## field ## ]]` markers, - /// ending with `[[ ## completed ## ]]`. Used internally by [`Predict`](crate::Predict) - /// to format demo assistant messages. - pub fn format_assistant_message_typed(&self, output: &S::Output) -> String - where - S::Output: Schema, - { - self.format_output(S::schema(), output) - } - - /// Formats an output value using a schema — the building-block version of - /// [`format_assistant_message_typed`](ChatAdapter::format_assistant_message_typed). - pub fn format_output(&self, schema: &crate::SignatureSchema, output: &O) -> String - where - O: Schema + for<'a> facet::Facet<'a>, - { - let json = serde_json::to_value(output).unwrap_or(Value::Null); - - let mut sections = Vec::new(); - for field_spec in schema.output_fields() { - if let Some(value) = value_for_path_relaxed(&json, field_spec.path()) { - sections.push(format!( - "[[ ## {} ## ]]\n{}", - field_spec.lm_name, - format_json_value_for_prompt(value) - )); - } - } - let mut result = sections.join("\n\n"); - result.push_str("\n\n[[ ## completed ## ]]\n"); - - result - } - - /// Formats a value-level output map as an assistant message — the - /// dynamic-lane twin of [`format_output`](ChatAdapter::format_output), - /// used to render demo assistant turns from [`SignatureDef`]s. + /// Formats a value-level output map as an assistant message for few-shot + /// demos: each output field delimited with `[[ ## field ## ]]` markers, + /// ending with `[[ ## completed ## ]]`. /// - /// Fields absent from `output` are skipped, mirroring the static lane's - /// relaxed path navigation. + /// Fields absent from `output` are skipped (the historical relaxed path + /// navigation for flattened structs). pub fn format_output_def(&self, def: &SignatureDef, output: &JsonMap) -> String { let mut sections = Vec::new(); for field in def.outputs.iter() { @@ -515,217 +350,9 @@ impl ChatAdapter { result } - /// Formats a demo example as a (user_message, assistant_message) pair. - /// - /// Convenience method that calls [`format_user_message_typed`](ChatAdapter::format_user_message_typed) - /// and [`format_assistant_message_typed`](ChatAdapter::format_assistant_message_typed). - pub fn format_demo_typed( - &self, - demo: &crate::predictors::Demo, - ) -> (String, String) - where - S::Input: Schema, - S::Output: Schema, - { - let user_msg = self.format_user_message_typed::(&demo.input); - let assistant_msg = self.format_assistant_message_typed::(&demo.output); - (user_msg, assistant_msg) - } - - #[allow(clippy::result_large_err)] - #[tracing::instrument( - name = "dsrs.adapter.chat.parse_typed", - level = "debug", - skip(self, response), - fields( - signature = std::any::type_name::(), - output_field_count = S::schema().output_fields().len() - ) - )] - /// Parses an LM response into a typed output with per-field metadata. - /// - /// The full parsing pipeline: - /// 1. Split the response into `[[ ## field ## ]]` sections - /// 2. For each output field in the schema, find its section by LM name - /// 3. Coerce the raw text to the field's type via the in-house coercer - /// 4. Run `#[check]` and `#[assert]` constraints - /// 5. Assemble the flat fields into the typed output via serde - /// - /// Returns the typed output and a map of [`FieldMeta`] with per-field raw text, parse - /// flags, and constraint results. - pub fn parse_response_typed( - &self, - response: &Message, - ) -> std::result::Result<(S::Output, IndexMap), ParseError> { - self.parse_output_with_meta::(S::schema(), response) - } - - #[allow(clippy::result_large_err)] - /// Parses an LM response against a schema, returning typed output and field metadata. - /// - /// Schema-based equivalent of [`parse_response_typed`](ChatAdapter::parse_response_typed). - /// Use when you have a schema but not a `S: Signature` type. - pub fn parse_output_with_meta( - &self, - schema: &crate::SignatureSchema, - response: &Message, - ) -> std::result::Result<(O, IndexMap), ParseError> - where - O: Schema + for<'a> facet::Facet<'a>, - { - let content = response.text_content_cow(); - let output_schema = schema.output_schema(); - let sections = parse_sections_cow(&content); - - let mut metas = IndexMap::new(); - let mut errors = Vec::new(); - // Coerced fields keyed by their leaf name, fed straight into serde's - // MapDeserializer below — no intermediate `Value::Object` tree. - let mut output_fields: Vec<(&'static str, Value)> = - Vec::with_capacity(schema.output_fields().len()); - let mut checks_total = 0usize; - let mut checks_failed = 0usize; - let mut asserts_failed = 0usize; - - for field in schema.output_fields() { - let rust_name = field.rust_name.as_str(); - let field_type = &field.type_ir; - - let raw_text: &str = match sections.get(field.lm_name) { - Some(text) => text.as_ref(), - None => { - debug!(field = %rust_name, "missing output field in response"); - errors.push(ParseError::MissingField { - field: rust_name.to_string(), - raw_response: content.to_string(), - }); - continue; - } - }; - - let coerced = match coerce(raw_text, field_type, &output_schema.types) { - Ok(value) => value, - Err(err) => { - let expected_type = type_name(field_type, Some(&output_schema.types)); - debug!( - field = %rust_name, - expected_type = %expected_type, - raw_text_len = raw_text.len(), - "typed coercion failed" - ); - trace!( - field = %rust_name, - raw_preview = %crate::truncate(raw_text, 160), - "typed coercion failed preview" - ); - errors.push(ParseError::CoercionFailed { - field: rust_name.to_string(), - expected_type, - raw_text: raw_text.to_string(), - source: JsonishError::from(err), - }); - continue; - } - }; - - // Constraints are evaluated straight off the `'static` specs — the - // compiled expressions are cached process-wide, so no per-call - // Environment build, recompile, or `Vec` allocation. - let mut checks = Vec::new(); - for spec in field.constraints { - let passed = evaluate_constraint_expression(spec.expression, &coerced.value); - match spec.kind { - ConstraintKind::Assert => { - if !passed { - asserts_failed += 1; - debug!(field = %rust_name, label = %spec.label, "typed assert constraint failed"); - errors.push(ParseError::AssertFailed { - field: rust_name.to_string(), - label: spec.label.to_string(), - expression: spec.expression.to_string(), - value: coerced.value.clone(), - }); - } - } - ConstraintKind::Check => { - checks_total += 1; - if !passed { - checks_failed += 1; - trace!(field = %rust_name, label = %spec.label, "typed check constraint failed"); - } - checks.push(ConstraintResult { - label: spec.label.to_string(), - expression: spec.expression.to_string(), - passed, - }); - } - } - } - - metas.insert( - field.rust_name.clone(), - FieldMeta { - raw_text: raw_text.to_string(), - flags: coerced.flags, - checks, - }, - ); - - // `#[serde(flatten)]` wrappers serialize flat, so every field keys at - // the top level by its final path segment. Duplicate leaves keep the - // last value, matching the previous `Map::insert` behavior. - if let Some(leaf) = field.path().iter().last() { - if let Some(existing) = output_fields.iter_mut().find(|(name, _)| *name == leaf) { - existing.1 = coerced.value; - } else { - output_fields.push((leaf, coerced.value)); - } - } - } - - if !errors.is_empty() { - debug!( - errors = errors.len(), - checks_total, checks_failed, asserts_failed, "typed parse returned errors" - ); - let partial = if output_fields.is_empty() { - None - } else { - Some(Value::Object( - output_fields - .into_iter() - .map(|(name, value)| (name.to_string(), value)) - .collect(), - )) - }; - return Err(ParseError::Multiple { errors, partial }); - } - - // Deserialize straight from the coerced field pairs — the historical - // `Value::Object` assembly + `from_value` re-walk is skipped entirely. - let typed_output = O::deserialize( - serde::de::value::MapDeserializer::<_, serde_json::Error>::new( - output_fields.into_iter(), - ), - ) - .map_err(|err| ParseError::ExtractionFailed { - field: "".to_string(), - raw_response: content.to_string(), - reason: err.to_string(), - })?; - debug!( - parsed_fields = metas.len(), - checks_total, checks_failed, asserts_failed, "typed parse completed" - ); - - Ok((typed_output, metas)) - } - #[allow(clippy::result_large_err)] /// Parses an LM response against an owned [`SignatureDef`] into a value-level - /// output map — the dynamic-lane twin of - /// [`parse_output_with_meta`](ChatAdapter::parse_output_with_meta), with no - /// `'static` requirement anywhere. + /// output map, with no `'static` requirement anywhere. /// /// The returned [`JsonMap`] is keyed by canonical field name /// (`FieldDef::name`); `types` resolves class/enum references during @@ -826,22 +453,6 @@ impl ChatAdapter { Ok((output, metas)) } - #[allow(clippy::result_large_err)] - /// Parses an LM response into a typed output, discarding field metadata. - /// - /// Convenience wrapper around [`parse_output_with_meta`](ChatAdapter::parse_output_with_meta). - pub fn parse_output( - &self, - schema: &crate::SignatureSchema, - response: &Message, - ) -> std::result::Result - where - O: Schema + for<'a> facet::Facet<'a>, - { - let (output, _) = self.parse_output_with_meta::(schema, response)?; - Ok(output) - } - /// Splits raw LM response text into named sections by `[[ ## field ## ]]` delimiters. /// /// Returns an ordered map of field_name → section_content. The `completed` marker @@ -850,38 +461,6 @@ impl ChatAdapter { pub fn parse_sections(content: &str) -> IndexMap { crate::adapter::chat::parse_sections(content) } - - /// Parses a raw [`Message`] into a [`Predicted`](crate::Predicted). - /// - /// Convenience wrapper that calls [`parse_response_typed`](ChatAdapter::parse_response_typed) - /// and wraps the result in [`Predicted`] with default metadata - /// (zero usage, no tool calls). Useful for testing or replaying saved responses. - #[expect( - clippy::result_large_err, - reason = "Public API returns PredictError directly for downstream matching." - )] - pub fn parse_response_with_schema( - &self, - response: Message, - ) -> std::result::Result, PredictError> { - let raw_response = response.content(); - let (output, field_meta) = self - .parse_response_typed::(&response) - .map_err(|source| PredictError::Parse { - source, - raw_response: raw_response.clone(), - lm_usage: crate::LmUsage::default(), - })?; - let metadata = CallMetadata::new( - raw_response, - crate::LmUsage::default(), - Vec::new(), - Vec::new(), - None, - field_meta, - ); - Ok(Predicted::new(output, metadata)) - } } fn parse_sections(content: &str) -> IndexMap { @@ -944,43 +523,6 @@ fn parse_sections_cow(content: &str) -> IndexMap<&str, std::borrow::Cow<'_, str> .collect() } -/// Navigates `value` by `path`, tolerating flattened wrappers whose intermediate segments -/// are absent from the serialized (flat) structure. -fn value_for_path_relaxed<'a>(value: &'a Value, path: &crate::FieldPath) -> Option<&'a Value> { - let mut current = value; - let parts: Vec<_> = path.iter().collect(); - let mut idx = 0usize; - while idx < parts.len() { - match current { - Value::Object(fields) => { - if let Some(next) = fields.get(parts[idx]) { - current = next; - idx += 1; - continue; - } - // Flattened wrappers may remove one or more intermediate path - // segments (`outer.inner.answer` serialized as `answer`), so - // probe ahead for the next segment visible at this level. - let mut matched = None; - for (look_ahead, part) in parts.iter().enumerate().skip(idx + 1) { - if let Some(next) = fields.get(*part) { - matched = Some((look_ahead, next)); - break; - } - } - if let Some((look_ahead, next)) = matched { - current = next; - idx = look_ahead + 1; - continue; - } - return None; - } - _ => return None, - } - } - Some(current) -} - fn format_json_value_for_prompt(value: &Value) -> String { match value { Value::String(s) => s.clone(), @@ -989,29 +531,10 @@ fn format_json_value_for_prompt(value: &Value) -> String { } } -fn render_input_field( - field_spec: &FieldSchema, - value: &Value, - input: &Value, - vars: &Value, -) -> String { - match field_spec.input_render { - InputRenderSpec::Default => match value { - Value::String(s) => s.clone(), - _ => serde_json::to_string(value).unwrap_or_else(|_| "".to_string()), - }, - InputRenderSpec::Format(format) => crate::typesys::format_value(value, format), - InputRenderSpec::Jinja(template) => { - render_input_field_jinja(template, field_spec, value, input, vars) - } - } -} - /// Renders one input field of a [`SignatureDef`], honoring its [`RenderSpec`]. /// -/// The Jinja arm compiles the template per call: dynamic templates are owned -/// strings, so the process-global template cache (keyed on `&'static str`) -/// deliberately stays static-lane-only. +/// The Jinja arm compiles the template per call: templates are owned strings on +/// a runtime [`SignatureDef`], so there is no `&'static str` key to cache on. fn render_input_field_def( def: &SignatureDef, field: &crate::ir::FieldDef, @@ -1076,92 +599,3 @@ fn build_input_context_def(def: &SignatureDef, input: &JsonMap) -> Value { Value::Object(root) } -fn build_input_context_value(schema: &crate::SignatureSchema, root: &Value) -> Value { - let mut input_json = root.clone(); - let Some(root_map) = input_json.as_object_mut() else { - return input_json; - }; - - // Provide alias lookups for top-level fields so templates can use either - // Rust field names (`input.question`) or prompt aliases (`input.query`). - for field_spec in schema.input_fields() { - if field_spec.rust_name.contains('.') || field_spec.lm_name == field_spec.rust_name { - continue; - } - if field_spec.path().iter().nth(1).is_some() { - continue; - } - if let Some(value) = root_map.get(field_spec.rust_name.as_str()).cloned() { - root_map - .entry(field_spec.lm_name.to_string()) - .or_insert(value); - } - } - - input_json -} - -fn render_input_field_jinja( - template: &'static str, - field_spec: &FieldSchema, - value: &Value, - input: &Value, - vars: &Value, -) -> String { - let cached = { - let cache = INPUT_RENDER_TEMPLATE_CACHE - .read() - .expect("input render template cache lock poisoned"); - cache.get(template).cloned() - }; - let cached = match cached { - Some(cached) => cached, - None => { - let mut env = build_input_render_environment(); - env.add_template(INPUT_RENDER_TEMPLATE_NAME, template) - .unwrap_or_else(|err| { - panic!( - "failed to compile cached input render template for `{}` ({}): {err}", - field_spec.lm_name, field_spec.rust_name - ) - }); - let entry = Arc::new(CachedInputRenderTemplate { env }); - let mut cache = INPUT_RENDER_TEMPLATE_CACHE - .write() - .expect("input render template cache lock poisoned"); - cache.entry(template).or_insert(entry).clone() - } - }; - - let compiled = cached - .env - .get_template(INPUT_RENDER_TEMPLATE_NAME) - .unwrap_or_else(|err| { - panic!( - "failed to fetch cached input render template for `{}` ({}): {err}", - field_spec.lm_name, field_spec.rust_name - ) - }); - - let this = value.clone(); - let field = json!({ - "name": field_spec.lm_name, - "rust_name": field_spec.rust_name, - "type": type_name(&field_spec.type_ir, None), - }); - let context = json!({ - "this": this, - "input": input, - "field": field, - "vars": vars, - }); - - compiled - .render(minijinja::Value::from_serialize(context)) - .unwrap_or_else(|err| { - panic!( - "failed to render input field `{}` (rust `{}`) with #[render(jinja = ...)] template `{}`: {err}", - field_spec.lm_name, field_spec.rust_name, template - ) - }) -} diff --git a/crates/dspy-rs/src/adapter/mod.rs b/crates/dspy-rs/src/adapter/mod.rs index e6f66ce4..33e8b347 100644 --- a/crates/dspy-rs/src/adapter/mod.rs +++ b/crates/dspy-rs/src/adapter/mod.rs @@ -1,15 +1,17 @@ //! Prompt formatting and LM response parsing. //! -//! The adapter turns a [`SignatureSchema`](crate::SignatureSchema) into prompts and parses -//! LM responses back into typed values. All prompts use the `[[ ## field_name ## ]]` +//! The adapter turns a [`SignatureDef`](crate::ir::SignatureDef) into prompts and parses +//! LM responses back into value-level maps. All prompts use the `[[ ## field_name ## ]]` //! delimiter protocol — input fields, output fields, and the `[[ ## completed ## ]]` //! marker that signals the end of the response. //! -//! Most users never touch this — [`Predict`](crate::Predict) calls the adapter internally. -//! Module authors who need fine-grained control over prompt construction use the -//! building blocks directly: [`build_system`](ChatAdapter::build_system), -//! [`format_input`](ChatAdapter::format_input), -//! [`parse_output`](ChatAdapter::parse_output). +//! Most users never touch this — [`Predict`](crate::Predict) renders and parses through +//! the adapter via the IR interpreter. Module authors who need fine-grained control over +//! prompt construction use the building blocks directly: +//! [`build_system_def`](ChatAdapter::build_system_def), +//! [`format_input_def`](ChatAdapter::format_input_def), +//! [`format_output_def`](ChatAdapter::format_output_def), +//! [`parse_output_def`](ChatAdapter::parse_output_def). pub mod chat; diff --git a/crates/dspy-rs/src/core/dyn_predictor.rs b/crates/dspy-rs/src/core/dyn_predictor.rs deleted file mode 100644 index 0faf9cff..00000000 --- a/crates/dspy-rs/src/core/dyn_predictor.rs +++ /dev/null @@ -1,593 +0,0 @@ -use std::collections::HashSet; -use std::ops::ControlFlow; - -use anyhow::Result; -use facet::{ConstTypeId, Def, Facet, KnownPointer, Shape, Type, UserType}; -use facet_reflect::Peek; - -use crate::SignatureSchema; -use crate::trace::JsonMap; - -/// Type-erased optimizer handle to a [`crate::Predict`] leaf. -/// -/// Optimizers need to inspect and mutate Predict parameters (demos, instructions) -/// without knowing the concrete signature type. Discovery uses -/// [`visit_named_predictors_mut`], which walks the module tree and passes each -/// discovered `(path, &mut dyn DynPredictor)` leaf to a selector callback. -/// -/// Normal users never touch this — you pass your module to `optimizer.compile()` -/// and it uses `DynPredictor` internally. -pub(crate) trait DynPredictor: Send + Sync { - /// Returns the [`SignatureSchema`] for this predictor's signature. - fn schema(&self) -> &SignatureSchema; - - /// Returns the current instruction (override or default from the signature). - fn instruction(&self) -> String; - - /// Returns current demos as flat JSON rows (field name → value; input and - /// output fields merged into one object). - fn demos_as_json(&self) -> Vec; - - /// Snapshots the predictor's mutable state (demos + instruction override). - fn dump_state(&self) -> PredictState; - - /// **The mutation seam.** Every write to a predictor's optimizable state — - /// instruction override and demos — flows through this one method: optimizer - /// candidate set/restore, [`ModuleState`](crate::ModuleState) restore, and the - /// [`fx::Params`](crate::fx::Params) overlay all delegate here (the builder - /// funnels into the same typed applicator at construction time). This is the - /// single place where prompt caches are invalidated; the overlay direction for - /// v1 is that a candidate is *data* applied through this seam, never ad-hoc - /// field mutation. - /// - /// # Errors - /// - /// Returns an error if updated demos can't be converted to the predictor's - /// typed `Demo` (schema mismatch). - fn apply_update(&mut self, update: StateUpdate) -> Result<()>; - - /// Restores predictor state from a snapshot. Delegates to - /// [`apply_update`](DynPredictor::apply_update). - /// - /// # Errors - /// - /// Returns an error if the demos can't be converted to the predictor's typed format. - fn load_state(&mut self, state: PredictState) -> Result<()> { - self.apply_update(StateUpdate::from(state)) - } - - /// Assigns the component name this predictor records on trace spans. - /// - /// Optimizers call this with each leaf's dotted path before running traced - /// passes, so spans join back to the same names the mutation seam addresses - /// — identity data, not optimizable state. - fn set_trace_name(&mut self, name: &str); -} - -/// A partial update to a predictor's mutable state, applied through the single -/// mutation seam [`DynPredictor::apply_update`]. -/// -/// `None` fields are left untouched; `instruction: Some(None)` clears the -/// override back to the signature default. -#[derive(Clone, Debug, Default)] -pub(crate) struct StateUpdate { - /// `Some(override)` replaces the instruction override (`Some(None)` clears it). - pub instruction: Option>, - /// `Some(demos)` replaces the demo set (flat JSON rows, see - /// [`PredictState::demos`]). - pub demos: Option>, -} - -impl From for StateUpdate { - fn from(state: PredictState) -> Self { - Self { - instruction: Some(state.instruction_override), - demos: Some(state.demos), - } - } -} - -/// Serializable snapshot of a [`crate::Predict`]'s mutable state. -/// -/// Contains demos (as flat JSON rows) and the instruction override. -/// Produced by optimizers when they tune a predictor; persist a whole module's -/// worth of these with [`ModuleState`](crate::ModuleState). -#[derive(Clone, Debug, Default, PartialEq, serde::Serialize, serde::Deserialize)] -pub struct PredictState { - /// Demo rows as flat JSON objects: field name → value, with input and - /// output fields merged into one object. This is the serde boundary for - /// demos-as-data — rows are split back into the predictor's typed - /// `Demo` via the signature schema on load. - #[serde(default)] - pub demos: Vec, - /// The instruction override, if any. - #[serde(default)] - pub instruction_override: Option, -} - -type VisitMutFn = - fn(*mut (), &mut dyn FnMut(&mut dyn DynPredictor) -> ControlFlow<()>) -> ControlFlow<()>; - -#[derive(Clone, Copy, Debug, facet::Facet)] -#[facet(opaque)] -pub(crate) struct PredictAccessorFns { - pub visit_mut: VisitMutFn, -} - -impl PartialEq for PredictAccessorFns { - fn eq(&self, other: &Self) -> bool { - std::ptr::fn_addr_eq(self.visit_mut, other.visit_mut) - } -} - -impl Eq for PredictAccessorFns {} - -facet::define_attr_grammar! { - ns "dsrs"; - crate_path $crate::core::dyn_predictor; - - pub enum Attr { - PredictAccessor(Option<&'static PredictAccessorFns>), - } -} - -/// Error from [`visit_named_predictors_mut`] when the Facet walker encounters an unsupported structure. -#[derive(Debug, thiserror::Error, PartialEq, Eq)] -pub(crate) enum NamedParametersError { - /// A `Predict` leaf was found inside an unsupported container (`Rc`, `Arc`, etc.). - #[error("container `{ty}` at `{path}` contains a parameter leaf")] - Container { path: String, ty: &'static str }, - - /// A `Predict`-like leaf was found with missing or malformed shape-local accessor payload. - #[error( - "parameter-like leaf at `{path}` is missing a valid shape-local accessor payload (`#[facet(dsrs::predict_accessor = ...)]`)" - )] - MissingAttr { path: String }, -} - -/// Visits all [`crate::Predict`] leaves in a module by walking struct fields and -/// supported containers. -/// -/// The callback acts as a selector: it receives each `(dotted_path, predictor)` pair -/// and may return `ControlFlow::Break(())` to stop traversal early. -/// -/// Safety model: -/// - discovery has exclusive `&mut` access to `module` for the full traversal; -/// - leaf access requires a valid shape-local accessor payload attached to the leaf; -/// - unsupported shared-pointer containers (`Rc`, `Arc`) are rejected explicitly. -#[tracing::instrument( - level = "debug", - name = "dsrs.visit_named_predictors_mut", - skip(module, visitor) -)] -pub(crate) fn visit_named_predictors_mut( - module: &mut M, - mut visitor: F, -) -> std::result::Result<(), NamedParametersError> -where - M: for<'a> Facet<'a>, - F: FnMut(&str, &mut dyn DynPredictor) -> ControlFlow<()>, -{ - let _ = walk_value(Peek::new(&*module), "", &mut visitor)?; - Ok(()) -} - -fn walk_value( - value: Peek<'_, '_>, - path: &str, - visitor: &mut F, -) -> std::result::Result, NamedParametersError> -where - F: FnMut(&str, &mut dyn DynPredictor) -> ControlFlow<()>, -{ - let shape = value.shape(); - match resolve_predict_leaf(shape) { - PredictLeafResolution::Accessor(accessor) => { - let raw_ptr = (value.data().as_byte_ptr() as *mut u8).cast::<()>(); - let mut forward = |predictor: &mut dyn DynPredictor| visitor(path, predictor); - return Ok((accessor.visit_mut)(raw_ptr, &mut forward)); - } - PredictLeafResolution::Missing => { - return Err(NamedParametersError::MissingAttr { - path: display_path(path), - }); - } - PredictLeafResolution::NotLeaf => {} - } - - if matches!(shape.ty, Type::User(UserType::Struct(_))) { - let struct_value = value.into_struct().expect("shape says struct"); - for idx in 0..struct_value.field_count() { - let field = struct_value.ty().fields[idx]; - if field.should_skip_deserializing() { - continue; - } - - let field_path = push_field(path, field.name); - let child = struct_value - .field(idx) - .map_err(|_| NamedParametersError::MissingAttr { - path: display_path(&field_path), - })?; - if let ControlFlow::Break(()) = walk_value(child, &field_path, visitor)? { - return Ok(ControlFlow::Break(())); - } - } - return Ok(ControlFlow::Continue(())); - } - - match shape.def { - Def::Option(_) => { - if let Some(inner) = value.into_option().expect("shape says option").value() - && let ControlFlow::Break(()) = walk_value(inner, path, visitor)? - { - return Ok(ControlFlow::Break(())); - } - Ok(ControlFlow::Continue(())) - } - Def::List(_) | Def::Array(_) | Def::Slice(_) => { - for (idx, child) in value - .into_list_like() - .expect("shape says list-like") - .iter() - .enumerate() - { - let child_path = push_index(path, idx); - if let ControlFlow::Break(()) = walk_value(child, &child_path, visitor)? { - return Ok(ControlFlow::Break(())); - } - } - Ok(ControlFlow::Continue(())) - } - Def::Map(_) => { - let mut entries = value - .into_map() - .expect("shape says map") - .iter() - .map(|(key, value)| { - key.as_str().map(|name| (name.to_string(), value)).ok_or( - NamedParametersError::Container { - path: display_path(path), - ty: "HashMap", - }, - ) - }) - .collect::, _>>()?; - - entries.sort_by(|(left, _), (right, _)| left.as_bytes().cmp(right.as_bytes())); - for (key, child) in entries { - let child_path = push_map_key(path, &key); - if let ControlFlow::Break(()) = walk_value(child, &child_path, visitor)? { - return Ok(ControlFlow::Break(())); - } - } - Ok(ControlFlow::Continue(())) - } - Def::Pointer(pointer_def) => match pointer_def.known { - Some(KnownPointer::Box) => { - if let Some(inner) = value - .into_pointer() - .expect("shape says pointer") - .borrow_inner() - && let ControlFlow::Break(()) = walk_value(inner, path, visitor)? - { - return Ok(ControlFlow::Break(())); - } - Ok(ControlFlow::Continue(())) - } - _ => { - // TODO(dsrs-shared-ptr-policy): define safe mutable-handle policy for Arc/Rc traversal. - if contains_parameter(shape, &mut HashSet::new()) { - return Err(NamedParametersError::Container { - path: display_path(path), - ty: pointer_name(pointer_def.known), - }); - } - Ok(ControlFlow::Continue(())) - } - }, - _ => Ok(ControlFlow::Continue(())), - } -} - -fn contains_parameter(shape: &'static Shape, visiting: &mut HashSet) -> bool { - if !matches!(resolve_predict_leaf(shape), PredictLeafResolution::NotLeaf) { - return true; - } - - if !visiting.insert(shape.id) { - return false; - } - - let found = match shape.ty { - Type::User(UserType::Struct(struct_def)) => struct_def - .fields - .iter() - .filter(|field| !field.should_skip_deserializing()) - .any(|field| contains_parameter(field.shape(), visiting)), - _ => match shape.def { - Def::List(def) => contains_parameter(def.t(), visiting), - Def::Option(def) => contains_parameter(def.t(), visiting), - Def::Map(def) => { - contains_parameter(def.k(), visiting) || contains_parameter(def.v(), visiting) - } - Def::Array(def) => contains_parameter(def.t(), visiting), - Def::Slice(def) => contains_parameter(def.t(), visiting), - Def::Set(def) => contains_parameter(def.t(), visiting), - Def::Result(def) => { - contains_parameter(def.t(), visiting) || contains_parameter(def.e(), visiting) - } - Def::Pointer(def) => def - .pointee() - .is_some_and(|inner| contains_parameter(inner, visiting)), - _ => false, - }, - }; - - visiting.remove(&shape.id); - found -} - -enum PredictLeafResolution { - NotLeaf, - Accessor(PredictAccessorFns), - Missing, -} - -fn resolve_predict_leaf(shape: &'static Shape) -> PredictLeafResolution { - let has_leaf_marker = is_predict_shape_identity(shape); - let mut accessor_count = 0usize; - let mut accessor = None; - let mut invalid = false; - - for attr in shape.attributes { - if attr.ns != Some("dsrs") { - continue; - } - - if attr.key == "predict_accessor" { - accessor_count += 1; - match attr.get_as::() { - Some(Attr::PredictAccessor(Some(value))) => { - if accessor.is_some() { - invalid = true; - } else { - accessor = Some(**value); - } - } - _ => invalid = true, - } - } - } - - if !has_leaf_marker { - if accessor_count > 0 { - return PredictLeafResolution::Missing; - } - return PredictLeafResolution::NotLeaf; - } - - if invalid || accessor_count != 1 { - return PredictLeafResolution::Missing; - } - - match accessor { - Some(accessor) => PredictLeafResolution::Accessor(accessor), - None => PredictLeafResolution::Missing, - } -} - -fn is_predict_shape_identity(shape: &'static Shape) -> bool { - shape.type_identifier == "Predict" && shape.module_path == Some("dspy_rs::predictors::predict") -} - -fn push_field(path: &str, field: &str) -> String { - if path.is_empty() { - field.to_string() - } else { - format!("{path}.{field}") - } -} - -fn push_index(path: &str, index: usize) -> String { - if path.is_empty() { - format!("[{index}]") - } else { - format!("{path}[{index}]") - } -} - -fn push_map_key(path: &str, key: &str) -> String { - let escaped = escape_map_key(key); - if path.is_empty() { - format!("['{escaped}']") - } else { - format!("{path}['{escaped}']") - } -} - -fn escape_map_key(key: &str) -> String { - let mut escaped = String::with_capacity(key.len()); - for ch in key.chars() { - match ch { - '\\' => escaped.push_str("\\\\"), - '\'' => escaped.push_str("\\'"), - c if c.is_control() => escaped.push_str(&format!("\\u{{{:X}}}", c as u32)), - c => escaped.push(c), - } - } - escaped -} - -fn display_path(path: &str) -> String { - if path.is_empty() { - "".to_string() - } else { - path.to_string() - } -} - -fn pointer_name(pointer: Option) -> &'static str { - match pointer { - Some(KnownPointer::Box) => "Box", - Some(KnownPointer::Rc) => "Rc", - Some(KnownPointer::Arc) => "Arc", - _ => "Pointer", - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate as dsrs; - use crate::Signature; - use crate::predictors::Predict as RealPredict; - use std::ops::ControlFlow; - use std::rc::Rc; - use std::sync::Arc; - - #[derive(Signature, Clone, Debug)] - struct DummySig { - #[input] - value: String, - - #[output] - done: bool, - } - - #[derive(facet::Facet)] - #[facet(crate = facet)] - struct SharedPointerModule { - rc_predictor: Rc>, - arc_predictor: Arc>, - } - - #[test] - fn named_parameters_rejects_shared_pointers() { - let mut module = SharedPointerModule { - rc_predictor: Rc::new(RealPredict::::new()), - arc_predictor: Arc::new(RealPredict::::new()), - }; - - match visit_named_predictors_mut(&mut module, |_path, _predictor| ControlFlow::Continue(())) - { - Err(NamedParametersError::Container { path, ty }) => { - assert_eq!(path, "rc_predictor"); - assert_eq!(ty, "Rc"); - } - Ok(_) => panic!("walk unexpectedly succeeded"), - Err(other) => panic!("unexpected error: {other:?}"), - } - } - - #[derive(facet::Facet)] - #[facet(crate = facet, dsrs::predict_accessor)] - struct MalformedAccessorLeaf; - - #[derive(facet::Facet)] - #[facet(crate = facet)] - struct MalformedAccessorModule { - malformed: MalformedAccessorLeaf, - } - - #[test] - fn named_parameters_rejects_malformed_predict_accessor_payload() { - let mut module = MalformedAccessorModule { - malformed: MalformedAccessorLeaf, - }; - - match visit_named_predictors_mut(&mut module, |_path, _predictor| ControlFlow::Continue(())) - { - Err(NamedParametersError::MissingAttr { path }) => { - assert_eq!(path, "malformed"); - } - Err(other) => panic!("unexpected error: {other:?}"), - Ok(_) => panic!("walk unexpectedly succeeded"), - } - } - - #[derive(facet::Facet)] - #[facet( - crate = facet, - dsrs::predict_accessor, - dsrs::predict_accessor - )] - struct DuplicateAccessorLeaf; - - #[derive(facet::Facet)] - #[facet(crate = facet)] - struct DuplicateAccessorModule { - duplicate: DuplicateAccessorLeaf, - } - - #[test] - fn named_parameters_rejects_duplicate_predict_accessor_attrs() { - let mut module = DuplicateAccessorModule { - duplicate: DuplicateAccessorLeaf, - }; - - match visit_named_predictors_mut(&mut module, |_path, _predictor| ControlFlow::Continue(())) - { - Err(NamedParametersError::MissingAttr { path }) => { - assert_eq!(path, "duplicate"); - } - Err(other) => panic!("unexpected error: {other:?}"), - Ok(_) => panic!("walk unexpectedly succeeded"), - } - } - - #[derive(facet::Facet)] - #[facet(crate = facet, dsrs::predict_accessor)] - struct AccessorOnlyLeaf; - - #[derive(facet::Facet)] - #[facet(crate = facet)] - struct AccessorOnlyModule { - leaf: AccessorOnlyLeaf, - } - - #[test] - fn named_parameters_rejects_accessor_without_leaf_marker() { - let mut module = AccessorOnlyModule { - leaf: AccessorOnlyLeaf, - }; - - match visit_named_predictors_mut(&mut module, |_path, _predictor| ControlFlow::Continue(())) - { - Err(NamedParametersError::MissingAttr { path }) => { - assert_eq!(path, "leaf"); - } - Err(other) => panic!("unexpected error: {other:?}"), - Ok(_) => panic!("walk unexpectedly succeeded"), - } - } - - #[test] - fn real_predict_shape_has_strict_identity_marker() { - assert!(is_predict_shape_identity(RealPredict::::SHAPE)); - } - - #[derive(facet::Facet)] - #[facet(crate = facet)] - struct Predict; - - #[derive(facet::Facet)] - #[facet(crate = facet)] - struct SameNameModule { - predictor: Predict, - } - - #[test] - fn type_name_alone_is_not_treated_as_predict_leaf() { - let mut module = SameNameModule { predictor: Predict }; - let mut paths = Vec::new(); - - visit_named_predictors_mut(&mut module, |path, _predictor| { - paths.push(path.to_string()); - ControlFlow::Continue(()) - }) - .expect("walk should succeed"); - - assert!(paths.is_empty()); - } -} diff --git a/crates/dspy-rs/src/core/errors.rs b/crates/dspy-rs/src/core/errors.rs index 3c448354..3846dcd0 100644 --- a/crates/dspy-rs/src/core/errors.rs +++ b/crates/dspy-rs/src/core/errors.rs @@ -1,4 +1,4 @@ -use std::{error::Error as StdError, time::Duration}; +use std::error::Error as StdError; use crate::LmUsage; @@ -32,10 +32,6 @@ impl From for JsonishError { pub enum ErrorClass { /// The request itself was malformed. BadRequest, - /// The requested resource doesn't exist. - NotFound, - /// Access denied by the provider. - Forbidden, /// Transient failure (network, rate limit, timeout, server 5xx) — retry may help. Temporary, /// The LM responded, but the output couldn't be parsed — prompt-engineering problem. @@ -49,7 +45,8 @@ pub enum ErrorClass { /// A call can fail at three stages, and which stage tells you what to do about it: /// /// 1. **[`Lm`](PredictError::Lm)** — couldn't reach the LM or it errored. Network, -/// rate limit, timeout. Generally retryable. +/// rate limit, timeout. Reported as a provider error; the rig client owns +/// transport-level retries. /// 2. **[`Parse`](PredictError::Parse)** — the LM responded, but we couldn't extract /// the expected fields from its output. Prompt-engineering problem. Retryable (the /// LM might produce different output). Includes the raw response for debugging. @@ -102,6 +99,20 @@ pub enum PredictError { #[source] source: crate::trace::ReplayError, }, + + /// Candidate parameters could not be bound to the predictor they address. + /// + /// Raised before any LM call when a named parameter slot (an ambient + /// [`fx::Params`](crate::fx::Params) entry or a saved state) doesn't fit + /// the predictor's signature — a harness/optimizer configuration bug, not + /// an LM failure. **Not retryable** — the same params fail the same way. + #[error("params for predictor `{name}` don't fit its signature")] + Params { + /// The predictor name the params were addressed to. + name: String, + #[source] + source: Box, + }, } impl PredictError { @@ -111,6 +122,7 @@ impl PredictError { Self::Parse { .. } => ErrorClass::BadResponse, Self::Conversion { .. } => ErrorClass::Internal, Self::Replay { .. } => ErrorClass::Internal, + Self::Params { .. } => ErrorClass::Internal, } } @@ -120,6 +132,7 @@ impl PredictError { Self::Parse { .. } => true, Self::Conversion { .. } => false, Self::Replay { .. } => false, + Self::Params { .. } => false, } } } @@ -218,31 +231,12 @@ pub enum ConversionError { /// The LM provider failed before returning a usable response. /// -/// All variants except [`Provider`](LmError::Provider) are retryable. -/// Use [`is_retryable`](LmError::is_retryable) for retry logic. +/// Everything the provider stack reports arrives as [`Provider`](LmError::Provider): +/// the provider name plus its error message and source. Not retryable — the +/// underlying rig client owns transport-level retry behavior. #[derive(Debug, thiserror::Error)] pub enum LmError { - /// Could not reach the provider endpoint (DNS, connection refused, etc.). - #[error("could not reach {endpoint}")] - Network { - endpoint: String, - #[source] - source: std::io::Error, - }, - - /// The provider returned a rate limit response (HTTP 429). - #[error("rate limited by provider")] - RateLimit { retry_after: Option }, - - /// The provider returned an unexpected HTTP status. - #[error("invalid response from provider: HTTP {status}")] - InvalidResponse { status: u16, body: String }, - - /// The request exceeded the configured timeout. - #[error("request timed out after {after:?}")] - Timeout { after: Duration }, - - /// A provider-specific error that doesn't fit the other categories. + /// A provider-reported error. #[error("provider error from {provider}: {message}")] Provider { provider: String, @@ -254,23 +248,10 @@ pub enum LmError { impl LmError { pub fn class(&self) -> ErrorClass { - match self { - Self::Network { .. } => ErrorClass::Temporary, - Self::RateLimit { .. } => ErrorClass::Temporary, - Self::InvalidResponse { status, .. } if *status >= 500 => ErrorClass::Temporary, - Self::InvalidResponse { .. } => ErrorClass::BadRequest, - Self::Timeout { .. } => ErrorClass::Temporary, - Self::Provider { .. } => ErrorClass::Internal, - } + ErrorClass::Internal } pub fn is_retryable(&self) -> bool { - match self { - Self::Network { .. } => true, - Self::RateLimit { .. } => true, - Self::Timeout { .. } => true, - Self::InvalidResponse { status, .. } => *status >= 500, - Self::Provider { .. } => false, - } + false } } diff --git a/crates/dspy-rs/src/core/lm/chat.rs b/crates/dspy-rs/src/core/lm/chat.rs index 74b77272..66485734 100644 --- a/crates/dspy-rs/src/core/lm/chat.rs +++ b/crates/dspy-rs/src/core/lm/chat.rs @@ -1,5 +1,3 @@ -use anyhow::Result; - use serde::{Deserialize, Serialize}; use serde_json::{Value, json}; @@ -343,88 +341,6 @@ impl Message { msg } - fn from_json_value(message: &Value) -> Result { - let role_str = message - .get("role") - .and_then(Value::as_str) - .ok_or_else(|| anyhow::anyhow!("chat message missing string role"))?; - - let role = match role_str { - "system" => Role::System, - "user" => Role::User, - "assistant" => Role::Assistant, - other => return Err(anyhow::anyhow!("unsupported chat message role: {other}")), - }; - - let id = message.get("id").and_then(Value::as_str).map(String::from); - - let content_val = message.get("content"); - - // Support both formats: - // New: "content": [{ "type": "text", "text": "..." }, ...] - // Legacy: "content": "plain string" - let content = match content_val { - Some(Value::Array(arr)) => arr - .iter() - .map(parse_content_block) - .collect::>>()?, - Some(Value::String(s)) => vec![ContentBlock::text(s.clone())], - _ => { - // Legacy type-tagged format: { "type": "tool_call", "tool_call": {...} } - match message.get("type").and_then(Value::as_str) { - Some("tool_call") => { - let tc: ToolCall = serde_json::from_value(message["tool_call"].clone())?; - vec![ContentBlock::tool_call(tc)] - } - Some("tool_result") => { - let tr: ToolResult = - serde_json::from_value(message["tool_result"].clone())?; - vec![ContentBlock::tool_result(tr)] - } - Some("reasoning") => { - let r: Reasoning = serde_json::from_value(message["reasoning"].clone())?; - vec![ContentBlock::reasoning(r)] - } - Some(other) => { - return Err(anyhow::anyhow!("unsupported chat message type: {other}")); - } - None => return Err(anyhow::anyhow!("chat message missing content field")), - } - } - }; - - Ok(Self { role, content, id }) - } -} - -fn parse_content_block(value: &Value) -> Result { - let block_type = value - .get("type") - .and_then(Value::as_str) - .ok_or_else(|| anyhow::anyhow!("content block missing type"))?; - - match block_type { - "text" => { - let text = value - .get("text") - .and_then(Value::as_str) - .ok_or_else(|| anyhow::anyhow!("text block missing text field"))?; - Ok(ContentBlock::text(text)) - } - "tool_call" => { - let tc: ToolCall = serde_json::from_value(value["tool_call"].clone())?; - Ok(ContentBlock::tool_call(tc)) - } - "tool_result" => { - let tr: ToolResult = serde_json::from_value(value["tool_result"].clone())?; - Ok(ContentBlock::tool_result(tr)) - } - "reasoning" => { - let r: Reasoning = serde_json::from_value(value["reasoning"].clone())?; - Ok(ContentBlock::reasoning(r)) - } - other => Err(anyhow::anyhow!("unsupported content block type: {other}")), - } } // --------------------------------------------------------------------------- @@ -518,17 +434,6 @@ impl Chat { self.messages.pop() } - pub fn from_json(&self, json_dump: Value) -> Result { - let messages = json_dump - .as_array() - .ok_or_else(|| anyhow::anyhow!("chat dump must be an array"))?; - let messages = messages - .iter() - .map(Message::from_json_value) - .collect::>>()?; - Ok(Self { messages }) - } - pub fn to_json(&self) -> Value { let messages = self .messages diff --git a/crates/dspy-rs/src/core/lm/client_registry.rs b/crates/dspy-rs/src/core/lm/client_registry.rs index 6df3c7ca..8b97bb5a 100644 --- a/crates/dspy-rs/src/core/lm/client_registry.rs +++ b/crates/dspy-rs/src/core/lm/client_registry.rs @@ -36,6 +36,7 @@ fn to_unit_completion_response(response: CompletionResponse) -> Completion pub struct TestCompletionModel { responses: Arc>>, last_request: Arc>>, + usage: Arc>, } impl TestCompletionModel { @@ -43,6 +44,7 @@ impl TestCompletionModel { Self { responses: Arc::new(Mutex::new(responses.into_iter().collect())), last_request: Arc::new(Mutex::new(None)), + usage: Arc::new(Mutex::new(Usage::new())), } } @@ -53,6 +55,12 @@ impl TestCompletionModel { pub fn last_request(&self) -> Option { self.last_request.lock().unwrap().clone() } + + /// Sets the token usage reported with every subsequent canned response + /// (defaults to zero usage). + pub fn set_usage(&self, usage: Usage) { + *self.usage.lock().unwrap() = usage; + } } impl CompletionProvider for TestCompletionModel { @@ -66,7 +74,7 @@ impl CompletionProvider for TestCompletionModel { })?; Ok(CompletionResponse { choice: OneOrMany::one(response), - usage: Usage::new(), + usage: *self.usage.lock().unwrap(), raw_response: (), message_id: None, }) diff --git a/crates/dspy-rs/src/core/lm/mod.rs b/crates/dspy-rs/src/core/lm/mod.rs index e1cc2225..dcb5c632 100644 --- a/crates/dspy-rs/src/core/lm/mod.rs +++ b/crates/dspy-rs/src/core/lm/mod.rs @@ -17,7 +17,6 @@ use rig::{ use bon::Builder; use serde::{Deserialize, Serialize}; use std::{collections::HashMap, sync::Arc, time::Duration}; -use tokio::sync::Mutex; use tracing::{debug, trace, warn}; use crate::trace::SpanEvent; @@ -99,7 +98,6 @@ impl ToolSet { /// /// Errors if two tool names mangle to the same JS identifier (see /// [`dsrs_tools::js_identifier`]). - #[cfg(feature = "code-mode")] pub async fn code_mode( tools: Vec>, config: dsrs_tools::SandboxConfig, @@ -161,7 +159,7 @@ impl Default for LMConfig { #[derive(Clone)] pub struct LM { pub config: LMConfig, - pub cache_handler: Option>>, + pub cache_handler: Option>, client: Option>, } @@ -241,7 +239,7 @@ impl LM { let cache_handler = if config.cache { debug!("initializing response cache"); - Some(Arc::new(Mutex::new(ResponseCache::new().await))) + Some(Arc::new(ResponseCache::new().await)) } else { None }; @@ -770,7 +768,7 @@ impl LM { None }; if let (Some(key), Some(cache)) = (cache_key, self.cache_handler.as_ref()) - && let Some(entry) = cache.lock().await.get_entry(key).await? + && let Some(entry) = cache.get_entry(key).await? && let Some(raw_output) = entry.raw_output { debug!("lm response served from cache"); @@ -912,7 +910,7 @@ impl LM { usage: accumulated_usage, raw_output: Some(first_choice.content()), }; - cache.lock().await.insert_entry(key, entry); + cache.insert_entry(key, entry); trace!("lm response cached"); } @@ -975,14 +973,7 @@ impl LM { fields(n) )] pub async fn inspect_history(&self, n: usize) -> Vec { - self.cache_handler - .as_ref() - .unwrap() - .lock() - .await - .get_history(n) - .await - .unwrap() + self.cache_handler.as_ref().unwrap().get_history(n) } } diff --git a/crates/dspy-rs/src/core/mod.rs b/crates/dspy-rs/src/core/mod.rs index d88cb6f0..b768cc12 100644 --- a/crates/dspy-rs/src/core/mod.rs +++ b/crates/dspy-rs/src/core/mod.rs @@ -11,35 +11,29 @@ //! [`LmError`] — distinguishes LM failures from parse failures so callers can handle //! retries differently. [`LM`] is the language model client itself. //! -//! Optimizer leaf discovery is internal (`visit_named_predictors_mut`) and currently -//! traverses struct fields plus `Option`, `Vec`, `HashMap`, and `Box`. -//! `Rc`/`Arc` wrappers that contain `Predict` leaves are rejected with explicit -//! container errors. +//! Optimizer leaf discovery is explicit: modules declare their [`Predict`](crate::Predict) +//! leaves by name through the [`Predictors`] trait (usually one `predictors!` line), +//! and optimizers read them through the object-safe [`PredictorInfo`] view. //! //! Most users import these through the crate root (`use dspy_rs::*`). Module authors //! who need fine-grained prompt control also use [`SignatureSchema`] and the adapter //! building blocks directly. -pub(crate) mod dyn_predictor; mod errors; pub mod example; pub mod lm; pub mod module; -mod module_ext; mod predicted; mod schema; pub mod settings; pub mod signature; mod state; -pub(crate) use dyn_predictor::*; -pub use dyn_predictor::PredictState; pub use errors::{ConversionError, ErrorClass, JsonishError, LmError, ParseError, PredictError}; pub use example::{ToInput, ToOutput}; -pub use state::ModuleState; +pub use state::{ModuleState, PredictState}; pub use lm::*; pub use module::*; -pub use module_ext::*; pub use predicted::{CallMetadata, ConstraintResult, FieldMeta, Predicted}; pub use schema::{FieldMetadataSpec, FieldPath, FieldSchema, InputRenderSpec, SignatureSchema}; pub use settings::*; diff --git a/crates/dspy-rs/src/core/module.rs b/crates/dspy-rs/src/core/module.rs index 290c5f34..e9821d1a 100644 --- a/crates/dspy-rs/src/core/module.rs +++ b/crates/dspy-rs/src/core/module.rs @@ -1,11 +1,145 @@ +use anyhow::Result as AnyResult; use futures::stream::{self, StreamExt}; -use kdam::{BarExt, tqdm}; use tracing::debug; -use crate::{Facet, PredictError, Predicted, Schema}; +use crate::core::PredictState; +use crate::trace::JsonMap; +use crate::{Facet, PredictError, Predicted, Schema, SignatureSchema}; type IndexedForwardResult = (usize, Result, PredictError>); +/// What optimizers read from — and, at explicit boundaries, write to — a +/// [`Predict`](crate::Predict) leaf. +/// +/// This is the typed, object-safe view a module exposes through +/// [`Predictors::predictors`]. Optimizers only ever *read* through it during a +/// run (schema for reflection prompts, current instruction, demos as JSON); +/// the two mutating methods are boundary operations: +/// +/// - [`set_trace_name`](PredictorInfo::set_trace_name) — the naming pass, run +/// once per optimization run (and by [`ModuleState`](crate::ModuleState)) +/// so trace spans record the same name the module declared for the leaf; +/// - [`load_state`](PredictorInfo::load_state) — the install seam, called by +/// `ModuleState::apply` and by the optimizer's final one-shot install of the +/// winning candidate. Candidate *evaluation* never calls it: candidates are +/// injected ambiently per call tree, not written into the module. +pub trait PredictorInfo: Send + Sync { + /// The [`SignatureSchema`] for this predictor's signature — field names + /// and docs for optimizer reflection prompts. + fn schema(&self) -> &'static SignatureSchema; + + /// The current effective instruction (override if set, else the + /// signature's default). + fn instruction(&self) -> String; + + /// The signature's default instruction (ignoring any override). + fn default_instruction(&self) -> String; + + /// Current demos as flat JSON rows (input and output fields merged into + /// one object per row). + fn demos_as_json(&self) -> Vec; + + /// Snapshot of the optimizable state (instruction override + demos). + fn dump_state(&self) -> PredictState; + + /// Restores optimizable state from a snapshot — the one mutation seam. + /// + /// This is a *full* overwrite: `instruction_override: None` clears the + /// override, `demos` replaces the demo set. + /// + /// # Errors + /// + /// Returns an error if the demos can't be converted to the predictor's + /// typed `Demo` form (schema mismatch). + fn load_state(&mut self, state: PredictState) -> AnyResult<()>; + + /// Assigns the component name this predictor records on trace spans. + /// + /// Part of the trace-name contract (see [`Predictors`]): the optimizer + /// stamps each leaf with the name the module declared before any traced + /// pass, so spans join back to the same names candidates address. + fn set_trace_name(&mut self, name: &str); +} + +/// Explicit predictor-leaf discovery: a module *names* its optimizable +/// [`Predict`](crate::Predict) leaves. +/// +/// This replaces the old reflection walker. A module that wants to be +/// optimizable (or persistable via [`ModuleState`](crate::ModuleState)) +/// declares its leaves explicitly — no derive magic, no pointer casts: +/// +/// ```ignore +/// struct TwoStepQA { +/// retrieve: Predict, +/// answer: ChainOfThought, +/// } +/// +/// dspy_rs::predictors!(TwoStepQA { retrieve, answer }); +/// // or by hand: +/// impl Predictors for TwoStepQA { +/// fn predictors(&self) -> Vec<(String, &dyn PredictorInfo)> { +/// vec![("retrieve".into(), &self.retrieve), ("answer".into(), &self.answer)] +/// } +/// fn predictors_mut(&mut self) -> Vec<(String, &mut dyn PredictorInfo)> { +/// vec![("retrieve".into(), &mut self.retrieve), ("answer".into(), &mut self.answer)] +/// } +/// } +/// ``` +/// +/// # The trace-name contract +/// +/// The names returned here are the *canonical identity* of each leaf: +/// +/// 1. they become the leaf's trace-span component name (the optimizer stamps +/// them via [`PredictorInfo::set_trace_name`] once per run); +/// 2. optimizer candidates address leaves by these names (ambient +/// [`fx::Params`](crate::fx::Params) entries bind per leaf at call time); +/// 3. [`ModuleState`](crate::ModuleState) persists per-leaf state under them. +/// +/// Names must be unique within a module and stable across +/// [`predictors`](Predictors::predictors) / [`predictors_mut`](Predictors::predictors_mut). +pub trait Predictors { + /// The module's predictor leaves, `(name, read view)` in declaration order. + fn predictors(&self) -> Vec<(String, &dyn PredictorInfo)>; + + /// The module's predictor leaves, `(name, mutable view)` in declaration + /// order. Only the boundary operations (naming pass, state install) go + /// through this. + fn predictors_mut(&mut self) -> Vec<(String, &mut dyn PredictorInfo)>; +} + +/// Implements [`Predictors`] for a module struct from a list of predictor +/// fields, using each field's identifier as its leaf name. +/// +/// ```ignore +/// struct Pipeline { draft: Predict, refine: ChainOfThought } +/// dspy_rs::predictors!(Pipeline { draft, refine }); +/// ``` +#[macro_export] +macro_rules! predictors { + ($ty:ty { $($field:ident),* $(,)? }) => { + impl $crate::Predictors for $ty { + fn predictors(&self) -> ::std::vec::Vec<(::std::string::String, &dyn $crate::PredictorInfo)> { + ::std::vec![$( + ( + ::std::string::String::from(::core::stringify!($field)), + &self.$field as &dyn $crate::PredictorInfo, + ) + ),*] + } + + fn predictors_mut(&mut self) -> ::std::vec::Vec<(::std::string::String, &mut dyn $crate::PredictorInfo)> { + ::std::vec![$( + ( + ::std::string::String::from(::core::stringify!($field)), + &mut self.$field as &mut dyn $crate::PredictorInfo, + ) + ),*] + } + } + }; +} + /// Strategy-swapping interface for prompting modules. /// /// Everything in dsrs is a Module — a bare LM call ([`crate::Predict`]), @@ -31,16 +165,18 @@ type IndexedForwardResult = (usize, Result, PredictError>); /// /// # Implementing `Module` /// -/// Implement [`forward`](Module::forward). Derive `Facet` on your struct so the -/// optimizer's walker can find your [`Predict`](crate::Predict) leaves automatically. +/// Implement [`forward`](Module::forward). To make the module optimizable and +/// persistable, also declare its [`Predict`](crate::Predict) leaves via +/// [`Predictors`] (one `predictors!` line). /// /// ```ignore -/// #[derive(Facet)] /// struct TwoStepQA { /// retrieve: Predict, /// answer: ChainOfThought, /// } /// +/// dspy_rs::predictors!(TwoStepQA { retrieve, answer }); +/// /// impl Module for TwoStepQA { /// type Input = RetrieveInput; /// type Output = WithReasoning; @@ -93,7 +229,6 @@ pub trait Module: Send + Sync { /// /// Returns `Vec>`, not `Result>` — individual failures don't /// abort the batch. Results preserve input order regardless of completion order. -/// Shows a progress bar on stderr. /// /// ```no_run /// # async fn example() -> Result<(), Box> { @@ -129,16 +264,10 @@ pub async fn forward_all( where M: Module + ?Sized, { - let total = inputs.len(); - let mut pb = tqdm!(total = total, desc = "Processing"); - let mut indexed_results: Vec> = stream::iter(inputs.into_iter().enumerate()) .map(|(idx, input)| async move { (idx, module.call(input).await) }) .buffer_unordered(max_concurrency) - .inspect(|_| { - let _ = pb.update(1); - }) .collect() .await; diff --git a/crates/dspy-rs/src/core/module_ext.rs b/crates/dspy-rs/src/core/module_ext.rs deleted file mode 100644 index 6304ea8c..00000000 --- a/crates/dspy-rs/src/core/module_ext.rs +++ /dev/null @@ -1,114 +0,0 @@ -use std::sync::Arc; - -use crate::{Facet, PredictError, Predicted, Schema}; - -use super::Module; - -/// Output transformation combinators for any [`Module`]. -/// -/// Post-process a module's output without writing a full `impl Module`. This is -/// the intermediate step between "use a library module" and "author your own" — -/// if you just need to reshape the output, a closure is enough. -/// -/// The inner module's [`crate::Predict`] leaves remain visible to the Facet walker, -/// so optimizer discovery works through these wrappers. -/// -/// ```ignore -/// // Transform output without impl Module -/// let confident = cot.map(|r| ConfidentAnswer { -/// answer: r.answer.clone(), -/// confidence: 0.9, -/// }); -/// let result = confident.call(input).await?; -/// ``` -pub trait ModuleExt: Module + Sized { - /// Transforms the output with an infallible closure. Returns a [`Map`] wrapper. - fn map(self, map: F) -> Map - where - F: Fn(Self::Output) -> T + Send + Sync + 'static, - T: Schema + for<'a> Facet<'a> + Send + Sync, - { - Map { - inner: self, - map: Arc::new(map), - } - } - - /// Transforms the output with a fallible closure. Returns an [`AndThen`] wrapper. - fn and_then(self, and_then: F) -> AndThen - where - F: Fn(Self::Output) -> Result + Send + Sync + 'static, - T: Schema + for<'a> Facet<'a> + Send + Sync, - { - AndThen { - inner: self, - and_then: Arc::new(and_then), - } - } -} - -impl ModuleExt for M {} - -/// Output transformation wrapper created by [`ModuleExt::map`]. -/// -/// Delegates to the inner module, then applies the closure to the output. -/// The inner module's [`crate::Predict`] leaves remain visible to Facet reflection -/// (the `inner` field is a real struct field), so optimizers can still discover and -/// tune parameters through this wrapper. -#[derive(facet::Facet)] -#[facet(crate = facet)] -pub struct Map -where - M: Module, -{ - pub(crate) inner: M, - #[facet(opaque, skip)] - map: Arc T + Send + Sync>, -} - -#[allow(async_fn_in_trait)] -impl Module for Map -where - M: Module, - T: Schema + for<'a> Facet<'a> + Send + Sync, -{ - type Input = M::Input; - type Output = T; - - async fn forward(&self, input: Self::Input) -> Result, PredictError> { - let predicted = self.inner.call(input).await?; - let (output, metadata) = predicted.into_parts(); - Ok(Predicted::new((self.map)(output), metadata)) - } -} - -/// Fallible output transformation wrapper created by [`ModuleExt::and_then`]. -/// -/// Like [`Map`], but the closure returns `Result`. -#[derive(facet::Facet)] -#[facet(crate = facet)] -pub struct AndThen -where - M: Module, -{ - pub(crate) inner: M, - #[facet(opaque, skip)] - and_then: Arc Result + Send + Sync>, -} - -#[allow(async_fn_in_trait)] -impl Module for AndThen -where - M: Module, - T: Schema + for<'a> Facet<'a> + Send + Sync, -{ - type Input = M::Input; - type Output = T; - - async fn forward(&self, input: Self::Input) -> Result, PredictError> { - let predicted = self.inner.call(input).await?; - let (output, metadata) = predicted.into_parts(); - let transformed = (self.and_then)(output)?; - Ok(Predicted::new(transformed, metadata)) - } -} diff --git a/crates/dspy-rs/src/core/schema.rs b/crates/dspy-rs/src/core/schema.rs index 22090a1c..9596f9af 100644 --- a/crates/dspy-rs/src/core/schema.rs +++ b/crates/dspy-rs/src/core/schema.rs @@ -1,5 +1,4 @@ use std::collections::HashMap; -use std::sync::OnceLock; use facet::{Def, Field, Shape, Type, UserType}; use serde_json::Value; @@ -122,9 +121,6 @@ pub struct SignatureSchema { input_fields: Box<[FieldSchema]>, output_fields: Box<[FieldSchema]>, output_schema: OutputSchema, - /// Memoized response-instructions prompt fragment. Schema-constant, but the - /// text is owned by the chat adapter — it hands us a builder on first use. - response_instructions: OnceLock, } impl SignatureSchema { @@ -168,18 +164,9 @@ impl SignatureSchema { input_fields: input_fields.into_boxed_slice(), output_fields: output_fields.into_boxed_slice(), output_schema: ::output_schema(), - response_instructions: OnceLock::new(), }) } - /// Returns the memoized response-instructions fragment, building it on first use. - /// - /// The fragment is a pure function of the schema, but its wording belongs to - /// the chat adapter — hence the builder closure instead of building here. - pub(crate) fn response_instructions_cached(&self, build: impl FnOnce() -> String) -> &str { - self.response_instructions.get_or_init(build) - } - pub fn instruction(&self) -> &'static str { self.instruction } @@ -236,9 +223,6 @@ impl SignatureSchema { input_fields: input_fields.into_boxed_slice(), output_fields: output_fields.into_boxed_slice(), output_schema: self.output_schema.clone(), - // Fresh memo — the derived schema has different fields, so the - // parent's cached response instructions would be wrong. - response_instructions: OnceLock::new(), } } diff --git a/crates/dspy-rs/src/core/state.rs b/crates/dspy-rs/src/core/state.rs index 32318c7f..89389517 100644 --- a/crates/dspy-rs/src/core/state.rs +++ b/crates/dspy-rs/src/core/state.rs @@ -7,78 +7,90 @@ //! //! ```ignore //! // After optimization: -//! ModuleState::from_module(&mut module)?.save("optimized.json")?; +//! ModuleState::from_module(&module)?.save("optimized.json")?; //! //! // In production: //! let mut module = MyPipeline::new(); //! ModuleState::load("optimized.json")?.apply(&mut module)?; //! ``` +//! +//! Leaves are discovered through the explicit [`Predictors`] contract — state +//! is keyed by the names the module declares for its leaves. use std::collections::BTreeMap; -use std::ops::ControlFlow; use std::path::Path; use anyhow::{Context, Result, anyhow}; use serde::{Deserialize, Serialize}; -use crate::Facet; -use crate::core::dyn_predictor::{PredictState, visit_named_predictors_mut}; +use crate::core::Predictors; +use crate::trace::JsonMap; + +/// Serializable snapshot of a [`crate::Predict`]'s optimizable state. +/// +/// Contains demos (as flat JSON rows) and the instruction override. +/// Produced by optimizers when they tune a predictor; persist a whole module's +/// worth of these with [`ModuleState`], or inject them per call tree with +/// [`fx::Params`](crate::fx::Params). +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct PredictState { + /// Demo rows as flat JSON objects: field name → value, with input and + /// output fields merged into one object. This is the serde boundary for + /// demos-as-data — rows are split back into the predictor's typed + /// `Demo` via the signature schema on load. + #[serde(default)] + pub demos: Vec, + /// The instruction override, if any. + #[serde(default)] + pub instruction_override: Option, +} /// Serializable snapshot of every [`Predict`](crate::Predict) leaf in a module, -/// keyed by the dotted path the optimizer walker discovers. +/// keyed by the leaf name the module declares via [`Predictors`]. #[derive(Clone, Debug, Default, Serialize, Deserialize)] pub struct ModuleState { - /// Per-predictor state, keyed by dotted path (`BTreeMap` keeps JSON output stable). + /// Per-predictor state, keyed by leaf name (`BTreeMap` keeps JSON output stable). pub predictors: BTreeMap, } impl ModuleState { /// Snapshots the current state (instruction overrides + demos) of every /// `Predict` leaf in `module`. - /// - /// Takes `&mut` because leaf discovery uses the exclusive Facet walker; the - /// module is not modified. - pub fn from_module(module: &mut M) -> Result + pub fn from_module(module: &M) -> Result where - M: for<'a> Facet<'a>, + M: Predictors + ?Sized, { let mut predictors = BTreeMap::new(); - visit_named_predictors_mut(module, |name, predictor| { - predictors.insert(name.to_string(), predictor.dump_state()); - ControlFlow::Continue(()) - })?; + for (name, info) in module.predictors() { + predictors.insert(name, info.dump_state()); + } Ok(Self { predictors }) } - /// Applies this state to a module in place. + /// Applies this state to a module in place, stamping each restored leaf's + /// trace name with its declared name. /// - /// Every path in the state must resolve to a `Predict` leaf in `module` — - /// unknown paths are an error, since they mean the saved state and the module + /// Every name in the state must resolve to a `Predict` leaf in `module` — + /// unknown names are an error, since they mean the saved state and the module /// structure have diverged. Predictors not named in the state are left untouched. pub fn apply(&self, module: &mut M) -> Result<()> where - M: for<'a> Facet<'a>, + M: Predictors + ?Sized, { let mut remaining: BTreeMap<&str, &PredictState> = self .predictors .iter() .map(|(name, state)| (name.as_str(), state)) .collect(); - let mut load_error: Option = None; - visit_named_predictors_mut(module, |name, predictor| { - if let Some(state) = remaining.remove(name) - && let Err(err) = predictor.load_state(state.clone()) - { - load_error = Some(anyhow!("failed to load state for `{name}`: {err}")); - return ControlFlow::Break(()); + for (name, info) in module.predictors_mut() { + if let Some(state) = remaining.remove(name.as_str()) { + info.set_trace_name(&name); + info.load_state(state.clone()) + .map_err(|err| anyhow!("failed to load state for `{name}`: {err}"))?; } - ControlFlow::Continue(()) - })?; - - if let Some(err) = load_error { - return Err(err); } + if !remaining.is_empty() { let missing = remaining.keys().copied().collect::>().join("`, `"); return Err(anyhow!( diff --git a/crates/dspy-rs/src/data/v1/dataloader.rs b/crates/dspy-rs/src/data/dataloader.rs similarity index 97% rename from crates/dspy-rs/src/data/v1/dataloader.rs rename to crates/dspy-rs/src/data/dataloader.rs index 8eb301db..b21ed7af 100644 --- a/crates/dspy-rs/src/data/v1/dataloader.rs +++ b/crates/dspy-rs/src/data/dataloader.rs @@ -1,16 +1,22 @@ use anyhow::{Context, Result, anyhow}; +#[cfg(feature = "data")] use arrow::array::{ Array, BooleanArray, Float32Array, Float64Array, Int8Array, Int16Array, Int32Array, Int64Array, StringArray, UInt8Array, UInt16Array, UInt32Array, UInt64Array, }; +#[cfg(feature = "data")] use csv::{ReaderBuilder, StringRecord}; +#[cfg(feature = "data")] use hf_hub::api::sync::Api; +#[cfg(feature = "data")] use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder; use reqwest; use std::any::TypeId; use std::collections::{HashMap, HashSet}; use std::fs; +#[cfg(feature = "data")] use std::io::Cursor; +#[cfg(feature = "data")] use std::path::{Path, PathBuf}; use tracing::debug; @@ -151,6 +157,10 @@ pub enum DataLoadError { /// Typed dataset ingress for JSON/CSV/Parquet/HuggingFace sources. /// +/// JSON/JSONL loading is always available; the CSV, Parquet, and HuggingFace +/// loaders require the (default-on) `data` feature, which carries the heavy +/// arrow/parquet/hf-hub/csv dependency stack. +/// /// Canonical public contract: /// - Returns `Vec` for any row struct `E: Deserialize + Facet` — rows are /// *your* type, not a signature-shaped pair, so they can carry gold labels @@ -227,6 +237,7 @@ impl DataLoader { Ok(examples) } + #[cfg(feature = "data")] #[tracing::instrument( name = "dsrs.data.load_csv", level = "debug", @@ -258,6 +269,7 @@ impl DataLoader { Ok(examples) } + #[cfg(feature = "data")] #[tracing::instrument( name = "dsrs.data.load_csv_with", level = "debug", @@ -294,6 +306,7 @@ impl DataLoader { Ok(examples) } + #[cfg(feature = "data")] #[tracing::instrument( name = "dsrs.data.load_parquet", level = "debug", @@ -317,6 +330,7 @@ impl DataLoader { Ok(examples) } + #[cfg(feature = "data")] #[tracing::instrument( name = "dsrs.data.load_parquet_with", level = "debug", @@ -348,6 +362,7 @@ impl DataLoader { Ok(examples) } + #[cfg(feature = "data")] #[tracing::instrument( name = "dsrs.data.load_hf", level = "debug", @@ -381,6 +396,7 @@ impl DataLoader { Ok(examples) } + #[cfg(feature = "data")] #[tracing::instrument( name = "dsrs.data.load_hf_with", level = "debug", @@ -419,6 +435,7 @@ impl DataLoader { Ok(examples) } + #[cfg(feature = "data")] #[tracing::instrument( name = "dsrs.data.load_hf_from_parquet", level = "debug", @@ -521,6 +538,7 @@ impl DataLoader { Ok(rows) } + #[cfg(feature = "data")] fn load_csv_rows( path: &str, delimiter: char, @@ -550,6 +568,7 @@ impl DataLoader { Self::collect_csv_rows(&mut reader, has_headers) } + #[cfg(feature = "data")] fn collect_csv_rows( reader: &mut csv::Reader, has_headers: bool, @@ -584,6 +603,7 @@ impl DataLoader { Ok(rows) } + #[cfg(feature = "data")] fn load_parquet_rows(path: &Path) -> std::result::Result, DataLoadError> { let file = fs::File::open(path).map_err(|err| DataLoadError::Parquet(err.into()))?; let builder = ParquetRecordBatchReaderBuilder::try_new(file) @@ -622,6 +642,7 @@ impl DataLoader { Ok(rows) } + #[cfg(feature = "data")] fn load_rows_from_parquet_files( parquet_files: &[PathBuf], ) -> std::result::Result, DataLoadError> { @@ -640,6 +661,7 @@ impl DataLoader { Ok(all_rows) } + #[cfg(feature = "data")] fn load_hf_rows( dataset_name: &str, subset: &str, @@ -838,6 +860,7 @@ fn row_from_json_value( }) } +#[cfg(feature = "data")] fn parse_csv_cell(cell: &str) -> serde_json::Value { let trimmed = cell.trim(); if trimmed.is_empty() { @@ -848,6 +871,7 @@ fn parse_csv_cell(cell: &str) -> serde_json::Value { .unwrap_or_else(|_| serde_json::Value::String(cell.to_string())) } +#[cfg(feature = "data")] fn csv_record_to_row_record( record: &StringRecord, row_index: usize, @@ -866,6 +890,7 @@ fn csv_record_to_row_record( RowRecord { row_index, values } } +#[cfg(feature = "data")] fn parquet_value_to_json(column: &dyn Array, row_idx: usize) -> Option { if let Some(values) = column.as_any().downcast_ref::() { return (!values.is_null(row_idx)).then(|| serde_json::json!(values.value(row_idx))); diff --git a/crates/dspy-rs/src/data/mod.rs b/crates/dspy-rs/src/data/mod.rs index 56214f00..f719e30b 100644 --- a/crates/dspy-rs/src/data/mod.rs +++ b/crates/dspy-rs/src/data/mod.rs @@ -1,7 +1,17 @@ -//! Data loading, versioned under [`v1`]. +//! Data loading. +//! +//! Typed ingestion is first-class: +//! +//! - [`DataLoader`] provides `load_*` methods that return +//! plain row structs directly. +//! - Typed examples flow directly into evaluation and optimizer APIs. +//! +//! There is no untyped row type: custom mappers work with [`RowRecord`] +//! (`serde_json`-valued source rows) at the load boundary, and demo rows +//! travel as flat JSON objects (see [`crate::PredictState`]). -pub mod v1; +pub mod dataloader; +pub mod utils; -// Re-export items and submodules so both `crate::data::DataLoader` and -// versioned paths like `crate::data::v1::dataloader::DataLoader` keep resolving. -pub use v1::*; +pub use dataloader::*; +pub use utils::*; diff --git a/crates/dspy-rs/src/data/v1/utils.rs b/crates/dspy-rs/src/data/utils.rs similarity index 100% rename from crates/dspy-rs/src/data/v1/utils.rs rename to crates/dspy-rs/src/data/utils.rs diff --git a/crates/dspy-rs/src/data/v1/mod.rs b/crates/dspy-rs/src/data/v1/mod.rs deleted file mode 100644 index f719e30b..00000000 --- a/crates/dspy-rs/src/data/v1/mod.rs +++ /dev/null @@ -1,17 +0,0 @@ -//! Data loading. -//! -//! Typed ingestion is first-class: -//! -//! - [`DataLoader`] provides `load_*` methods that return -//! plain row structs directly. -//! - Typed examples flow directly into evaluation and optimizer APIs. -//! -//! There is no untyped row type: custom mappers work with [`RowRecord`] -//! (`serde_json`-valued source rows) at the load boundary, and demo rows -//! travel as flat JSON objects (see [`crate::PredictState`]). - -pub mod dataloader; -pub mod utils; - -pub use dataloader::*; -pub use utils::*; diff --git a/crates/dspy-rs/src/evaluate/evaluator.rs b/crates/dspy-rs/src/evaluate/evaluator.rs index 061e41e7..f88a5753 100644 --- a/crates/dspy-rs/src/evaluate/evaluator.rs +++ b/crates/dspy-rs/src/evaluate/evaluator.rs @@ -2,7 +2,7 @@ use anyhow::{Result, anyhow}; use futures::stream::{self, StreamExt, TryStreamExt}; use crate::core::{Module, ToInput}; -use crate::trace::{Trace, TraceMeta, TraceOutcome, capture_with_meta}; +use crate::trace::{SpanId, Trace, TraceMeta, TraceOutcome, capture_with_meta}; use crate::Predicted; pub use crate::trace::Eval; @@ -58,6 +58,32 @@ where prediction: &Predicted, trace: Option<&Trace>, ) -> Result; + + /// Optional per-span credit assignment (RFC 0004 §4). + /// + /// Called once per traced rollout, after [`evaluate`](TypedMetric::evaluate), + /// with the same example and prediction. Return `(SpanId, Eval)` pairs for + /// the spans you can score in isolation ("did the retriever return the + /// gold doc?"); every span you leave out keeps whole-rollout credit. Demo + /// harvesting ([`BootstrapFewShot`](crate::BootstrapFewShot), + /// [`MIPROv2`](crate::MIPROv2), [`SIMBA`](crate::SIMBA)) prefers a span's + /// own eval over the rollout score, so a good rollout's recovered-from + /// missteps stay out of the demo pool — and a failed rollout's good steps + /// can still make it in. + /// + /// The default returns no span scores: implementing only + /// [`evaluate`](TypedMetric::evaluate) keeps whole-rollout behavior + /// exactly. Pairs whose id is not in the trace are ignored; duplicate ids + /// keep the last eval. + async fn evaluate_spans( + &self, + example: &E, + prediction: &Predicted, + trace: &Trace, + ) -> Result> { + let _ = (example, prediction, trace); + Ok(Vec::new()) + } } /// Runs a module on every example in a trainset and scores each with a metric. @@ -122,9 +148,58 @@ where /// produced it (with [`Trace::outcome`] filled in). pub type Rollout = (Eval, Trace); -/// Concurrency core shared by the public entry points and optimizers: runs each -/// example under a capture scope, scores it with the metric (outside the scope), -/// and records the eval into the trace outcome. +fn json_object(value: Result) -> Option { + match value { + Ok(serde_json::Value::Object(map)) => Some(map), + _ => None, + } +} + +/// The one traced-rollout primitive: runs `module` on one example under a +/// capture scope seeded with `meta`, scores the result with the metric +/// *outside* the scope (LM-as-judge metrics don't pollute the execution +/// trace), and records the eval into the trace outcome — plus any span-level +/// evals the metric attaches via [`TypedMetric::evaluate_spans`]. +/// +/// Both the public evaluation entry points and the optimizer engine compose +/// this — there is exactly one traced-rollout loop in the crate. +pub(crate) async fn rollout_traced( + module: &M, + example: &E, + metric: &MT, + mut meta: TraceMeta, +) -> Result +where + E: ToInput + Sync, + M: Module, + MT: TypedMetric, +{ + let input = example.to_input()?; + if meta.input.is_none() { + meta.input = json_object(serde_json::to_value(&input)); + } + let started = std::time::Instant::now(); + let (result, mut trace) = capture_with_meta(meta, || module.call(input)).await; + let predicted = result.map_err(|err| anyhow!("{err}"))?; + let eval = metric.evaluate(example, &predicted, Some(&trace)).await?; + // Per-span credit (RFC 0004 §4): stamp any span-level evals the metric + // attaches; demo harvesting prefers these over the rollout score. + for (span_id, span_eval) in metric.evaluate_spans(example, &predicted, &trace).await? { + if let Some(span) = trace.spans.get_mut(span_id.0 as usize) { + span.eval = Some(span_eval); + } + } + trace.outcome = Some(TraceOutcome { + output: json_object(serde_json::to_value(&*predicted)), + error: None, + eval: Some(eval.clone()), + duration_us: started.elapsed().as_micros() as u64, + }); + Ok((eval, trace)) +} + +/// Concurrency core shared by the public entry points and optimizers: fans +/// [`rollout_traced`] out over the examples with bounded concurrency. pub(crate) async fn evaluate_examples_traced<'a, E, M, MT, I>( module: &M, examples: I, @@ -137,34 +212,11 @@ where MT: TypedMetric, I: IntoIterator, { - stream::iter(examples.into_iter().map(|example| async move { - let input = example.to_input()?; - let meta = TraceMeta { - input: serde_json::to_value(&input) - .ok() - .and_then(|value| match value { - serde_json::Value::Object(map) => Some(map), - _ => None, - }), - ..TraceMeta::default() - }; - let started = std::time::Instant::now(); - let (result, mut trace) = capture_with_meta(meta, || module.call(input)).await; - let predicted = result.map_err(|err| anyhow!("{err}"))?; - let eval = metric.evaluate(example, &predicted, Some(&trace)).await?; - trace.outcome = Some(TraceOutcome { - output: serde_json::to_value(&*predicted) - .ok() - .and_then(|value| match value { - serde_json::Value::Object(map) => Some(map), - _ => None, - }), - error: None, - eval: Some(eval.clone()), - duration_us: started.elapsed().as_micros() as u64, - }); - Ok::<_, anyhow::Error>((eval, trace)) - })) + stream::iter( + examples + .into_iter() + .map(|example| rollout_traced(module, example, metric, TraceMeta::default())), + ) .buffered(max_concurrency.max(1)) .try_collect() .await diff --git a/crates/dspy-rs/src/evaluate/feedback_helpers.rs b/crates/dspy-rs/src/evaluate/feedback_helpers.rs deleted file mode 100644 index 7edf842b..00000000 --- a/crates/dspy-rs/src/evaluate/feedback_helpers.rs +++ /dev/null @@ -1,403 +0,0 @@ -/// Helper functions for creating rich feedback [`Eval`]s -/// -/// This module provides utilities for common feedback patterns in different domains: -/// - Document retrieval (precision, recall, F1) -/// - Code generation (compilation, execution, testing) -/// - Multi-objective evaluation -/// - String similarity and classification -use super::Eval; -use std::collections::{HashMap, HashSet}; - -// ============================================================================ -// Retrieval Feedback Helpers -// ============================================================================ - -/// Create feedback for document retrieval tasks -/// -/// # Arguments -/// * `retrieved` - Documents retrieved by the system -/// * `expected` - Expected/gold documents -/// * `context_docs` - Optional list of all available documents for context -/// -/// # Example Feedback -/// ```text -/// Retrieved 3/5 correct documents (Precision: 0.6, Recall: 0.6, F1: 0.6) -/// Correctly retrieved: doc1, doc2, doc3 -/// Missed: doc4, doc5 -/// Incorrectly retrieved: doc6, doc7 -/// ``` -pub fn retrieval_feedback( - retrieved: &[impl AsRef], - expected: &[impl AsRef], - context_docs: Option<&[impl AsRef]>, -) -> Eval { - let retrieved_set: HashSet = retrieved.iter().map(|s| s.as_ref().to_string()).collect(); - - let expected_set: HashSet = expected.iter().map(|s| s.as_ref().to_string()).collect(); - - let correct: Vec = retrieved_set.intersection(&expected_set).cloned().collect(); - - let missed: Vec = expected_set.difference(&retrieved_set).cloned().collect(); - - let incorrect: Vec = retrieved_set.difference(&expected_set).cloned().collect(); - - let precision = if retrieved.is_empty() { - 0.0 - } else { - correct.len() as f64 / retrieved.len() as f64 - }; - - let recall = if expected.is_empty() { - 1.0 - } else { - correct.len() as f64 / expected.len() as f64 - }; - - let f1 = if precision + recall > 0.0 { - 2.0 * precision * recall / (precision + recall) - } else { - 0.0 - }; - - let mut feedback = format!( - "Retrieved {}/{} correct documents (Precision: {:.3}, Recall: {:.3}, F1: {:.3})\n", - correct.len(), - expected.len(), - precision, - recall, - f1 - ); - - if !correct.is_empty() { - feedback.push_str(&format!("Correctly retrieved: {}\n", correct.join(", "))); - } - - if !missed.is_empty() { - feedback.push_str(&format!("Missed: {}\n", missed.join(", "))); - } - - if !incorrect.is_empty() { - feedback.push_str(&format!( - "Incorrectly retrieved: {}\n", - incorrect.join(", ") - )); - } - - if let Some(docs) = context_docs { - feedback.push_str(&format!("Total available documents: {}\n", docs.len())); - } - - Eval::with_feedback(f1, feedback) -} - -// ============================================================================ -// Code Generation Feedback Helpers -// ============================================================================ - -/// Stage in code execution pipeline -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum CodeStage { - Parse, - Compile, - Execute, - Test, -} - -impl std::fmt::Display for CodeStage { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - CodeStage::Parse => write!(f, "Parse"), - CodeStage::Compile => write!(f, "Compile"), - CodeStage::Execute => write!(f, "Execute"), - CodeStage::Test => write!(f, "Test"), - } - } -} - -/// Result of a code stage -#[derive(Debug, Clone)] -pub enum StageResult { - Success, - Failure { error: String }, -} - -/// Create feedback for code generation pipelines -/// -/// # Arguments -/// * `stages` - List of (stage, result) tuples showing pipeline progression -/// * `final_score` - Overall score (0.0 to 1.0) -/// -/// # Example Feedback -/// ```text -/// Parse: Success -/// Compile: Success -/// Execute: RuntimeError: division by zero on line 10 -/// ``` -pub fn code_pipeline_feedback(stages: &[(CodeStage, StageResult)], final_score: f64) -> Eval { - let mut feedback = String::new(); - - for (stage, result) in stages { - match result { - StageResult::Success => { - feedback.push_str(&format!("{}: Success\n", stage)); - } - StageResult::Failure { error } => { - feedback.push_str(&format!("{}: {}\n", stage, error)); - feedback.push_str(&format!("Failed at stage: {}\n", stage)); - break; // Stop at first failure - } - } - } - - Eval::with_feedback(final_score, feedback) -} - -// ============================================================================ -// Multi-Objective Feedback Helpers -// ============================================================================ - -/// Create feedback for multi-objective optimization -/// -/// # Arguments -/// * `objectives` - Map of objective name to (score, feedback) pairs -/// * `weights` - Optional weights for aggregating objectives -/// -/// # Example Feedback -/// ```text -/// [Correctness] Score: 0.9 - Output matches expected format -/// [Latency] Score: 0.7 - Response took 450ms (target: <300ms) -/// [Privacy] Score: 1.0 - No PII detected in output -/// Overall: 0.87 (weighted average) -/// ``` -pub fn multi_objective_feedback( - objectives: &HashMap, - weights: Option<&HashMap>, -) -> Eval { - let mut feedback = String::new(); - - let mut total_score = 0.0; - let mut total_weight = 0.0; - - let mut objective_names: Vec<_> = objectives.keys().collect(); - objective_names.sort(); - - for name in objective_names { - if let Some((score, obj_feedback)) = objectives.get(name.as_str()) { - let weight = weights - .and_then(|w| w.get(name.as_str())) - .copied() - .unwrap_or(1.0); - - feedback.push_str(&format!( - "[{}] Score: {:.3} - {}\n", - name, score, obj_feedback - )); - - total_score += score * weight; - total_weight += weight; - } - } - - let aggregate_score = if total_weight > 0.0 { - total_score / total_weight - } else { - 0.0 - }; - - feedback.push_str(&format!( - "\nOverall: {:.3} (weighted average)", - aggregate_score - )); - - Eval::with_feedback(aggregate_score, feedback) -} - -// ============================================================================ -// String Similarity Feedback -// ============================================================================ - -/// Create feedback for string similarity tasks -/// -/// Uses simple word-level comparison to provide actionable feedback -pub fn string_similarity_feedback(predicted: &str, expected: &str) -> Eval { - let exact_match = predicted.trim() == expected.trim(); - - if exact_match { - return Eval::with_feedback(1.0, "Exact match"); - } - - let pred_lower = predicted.to_lowercase(); - let exp_lower = expected.to_lowercase(); - - if pred_lower == exp_lower { - return Eval::with_feedback(0.95, "Match ignoring case (minor formatting difference)"); - } - - // Word-level comparison - let pred_words: HashSet<&str> = pred_lower.split_whitespace().collect(); - let exp_words: HashSet<&str> = exp_lower.split_whitespace().collect(); - - let common_words: HashSet<_> = pred_words.intersection(&exp_words).collect(); - let missing_words: Vec<_> = exp_words.difference(&pred_words).collect(); - let extra_words: Vec<_> = pred_words.difference(&exp_words).collect(); - - let recall = if !exp_words.is_empty() { - common_words.len() as f64 / exp_words.len() as f64 - } else { - 1.0 - }; - - let precision = if !pred_words.is_empty() { - common_words.len() as f64 / pred_words.len() as f64 - } else { - 0.0 - }; - - let f1 = if precision + recall > 0.0 { - 2.0 * precision * recall / (precision + recall) - } else { - 0.0 - }; - - let mut feedback = format!("Partial match (F1: {:.3})\n", f1); - feedback.push_str(&format!("Expected: \"{}\"\n", expected)); - feedback.push_str(&format!("Predicted: \"{}\"\n", predicted)); - - if !missing_words.is_empty() { - feedback.push_str(&format!( - "Missing words: {}\n", - missing_words - .iter() - .map(|w| format!("\"{}\"", w)) - .collect::>() - .join(", ") - )); - } - - if !extra_words.is_empty() { - feedback.push_str(&format!( - "Extra words: {}\n", - extra_words - .iter() - .map(|w| format!("\"{}\"", w)) - .collect::>() - .join(", ") - )); - } - - Eval::with_feedback(f1, feedback) -} - -// ============================================================================ -// Classification Feedback -// ============================================================================ - -/// Create feedback for classification tasks -pub fn classification_feedback( - predicted_class: &str, - expected_class: &str, - confidence: Option, -) -> Eval { - let correct = predicted_class == expected_class; - let score = if correct { 1.0 } else { 0.0 }; - - let mut feedback = if correct { - format!("Correct classification: \"{}\"", predicted_class) - } else { - format!( - "Incorrect classification\n Expected: \"{}\"\n Predicted: \"{}\"", - expected_class, predicted_class - ) - }; - - if let Some(conf) = confidence { - feedback.push_str(&format!("\n Confidence: {:.3}", conf)); - } - - Eval::with_feedback(score, feedback) -} - -#[cfg(test)] -mod tests { - use super::*; - - fn feedback_text(eval: &Eval) -> &str { - eval.feedback.as_deref().unwrap_or("") - } - - #[test] - fn test_retrieval_feedback_perfect() { - let retrieved = vec!["doc1", "doc2", "doc3"]; - let expected = vec!["doc1", "doc2", "doc3"]; - - let eval = retrieval_feedback(&retrieved, &expected, None::<&[&str]>); - assert_eq!(eval.score, 1.0); - assert!(feedback_text(&eval).contains("3/3")); - } - - #[test] - fn test_retrieval_feedback_partial() { - let retrieved = vec!["doc1", "doc2", "doc4"]; - let expected = vec!["doc1", "doc2", "doc3"]; - - let eval = retrieval_feedback(&retrieved, &expected, None::<&[&str]>); - assert!(eval.score < 1.0 && eval.score > 0.0); - assert!(feedback_text(&eval).contains("Missed: doc3")); - assert!(feedback_text(&eval).contains("Incorrectly retrieved: doc4")); - } - - #[test] - fn test_code_pipeline_feedback() { - let stages = vec![ - (CodeStage::Parse, StageResult::Success), - (CodeStage::Compile, StageResult::Success), - ( - CodeStage::Execute, - StageResult::Failure { - error: "Division by zero".to_string(), - }, - ), - ]; - - let eval = code_pipeline_feedback(&stages, 0.6); - assert!(feedback_text(&eval).contains("Parse")); - assert!(feedback_text(&eval).contains("Compile")); - assert!(feedback_text(&eval).contains("Execute")); - assert_eq!(eval.score, 0.6); - } - - #[test] - fn test_multi_objective_feedback() { - let mut objectives = HashMap::new(); - objectives.insert("accuracy".to_string(), (0.9, "Good accuracy".to_string())); - objectives.insert("latency".to_string(), (0.7, "Slow response".to_string())); - - let eval = multi_objective_feedback(&objectives, None); - assert!(feedback_text(&eval).contains("[accuracy]")); - assert!(feedback_text(&eval).contains("[latency]")); - assert!((eval.score - 0.8).abs() < 0.01); // Average of 0.9 and 0.7 - } - - #[test] - fn test_string_similarity_exact() { - let eval = string_similarity_feedback("hello world", "hello world"); - assert_eq!(eval.score, 1.0); - } - - #[test] - fn test_string_similarity_case() { - let eval = string_similarity_feedback("Hello World", "hello world"); - assert_eq!(eval.score, 0.95); - } - - #[test] - fn test_classification_feedback() { - let eval = classification_feedback("positive", "positive", Some(0.95)); - assert_eq!(eval.score, 1.0); - assert!(feedback_text(&eval).contains("Correct")); - - let eval = classification_feedback("negative", "positive", Some(0.85)); - assert_eq!(eval.score, 0.0); - assert!(feedback_text(&eval).contains("Incorrect")); - } -} diff --git a/crates/dspy-rs/src/evaluate/mod.rs b/crates/dspy-rs/src/evaluate/mod.rs index 4248f992..dd34eb9d 100644 --- a/crates/dspy-rs/src/evaluate/mod.rs +++ b/crates/dspy-rs/src/evaluate/mod.rs @@ -18,7 +18,5 @@ //! (`trace.for_component("retriever")`). pub mod evaluator; -pub mod feedback_helpers; pub use evaluator::*; -pub use feedback_helpers::*; diff --git a/crates/dspy-rs/src/fx/mod.rs b/crates/dspy-rs/src/fx/mod.rs index 7d7c0915..961692c5 100644 --- a/crates/dspy-rs/src/fx/mod.rs +++ b/crates/dspy-rs/src/fx/mod.rs @@ -35,27 +35,42 @@ //! params scope is a tokio task-local: spawned subtasks do not inherit it. use std::any::{Any, TypeId}; -use std::collections::{BTreeMap, HashMap}; +use std::collections::{BTreeMap, HashMap, VecDeque}; use std::hash::{DefaultHasher, Hasher}; use std::marker::PhantomData; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, LazyLock, RwLock}; use tokio::task_local; -use crate::core::{DynPredictor, ModuleState, PredictState}; -use crate::{Facet, LmError, Module, Predict, PredictError, Predicted, Schema, Signature}; +use crate::core::{ModuleState, PredictState, PredictorInfo}; +use crate::{Facet, Module, Predict, PredictError, Predicted, Schema, Signature}; /// Runs a future with an [`ir::Overlay`](crate::ir::Overlay) as the ambient /// candidate — the overlay is unbound against the program into [`Params`] and /// scoped exactly like [`with_params`]. See /// [`ir::bridge`](crate::ir::bridge). -#[cfg(feature = "ir")] pub use crate::ir::bridge::with_overlay; task_local! { static CURRENT_PARAMS: Arc; } +/// One named parameter slot inside [`Params`]: a [`PredictState`] plus the +/// explicit-clear markers that make "reset to the signature default" +/// expressible (plain `PredictState` semantics are "None/empty = leave the +/// incumbent alone"). +#[derive(Clone, Debug, Default, PartialEq)] +pub(crate) struct ParamsEntry { + pub state: PredictState, + /// Explicitly clear the instruction override back to the signature + /// default (wins over any instance override when injected ambiently). + pub clear_instruction: bool, + /// `state.demos` is an explicit *set* — an empty vec means "no demos", + /// overriding any instance demos — rather than "non-empty means set". + pub explicit_demos: bool, +} + /// The optimizable state of a functional harness: named [`PredictState`]s — /// instructions and demos keyed by the names passed to [`predict`]. /// @@ -63,11 +78,18 @@ task_local! { /// with respect to its `Params`: evaluating a different candidate means /// injecting a different `Params` value via [`with_params`], never mutating a /// module in place. +/// +/// Struct-held [`Predict`](crate::Predict) leaves consult the ambient +/// `Params` too: each leaf binds the entry matching its component name (the +/// name stamped by [`Predictors`](crate::Predictors) discovery or +/// [`PredictBuilder::named`](crate::predictors::PredictBuilder::named)) at +/// call time, with ambient values winning over instance state per slot. This +/// is the optimizer's candidate-injection currency. #[derive(Clone, Debug, Default)] pub struct Params { - /// (config hash, state) per predictor name. The hash keys the predictor + /// (config hash, entry) per predictor name. The hash keys the predictor /// instance cache so unchanged configs reuse fully-warmed `Predict`s. - entries: BTreeMap, + entries: BTreeMap, } impl Params { @@ -75,28 +97,67 @@ impl Params { Self::default() } + fn upsert(&mut self, name: String, mutate: impl FnOnce(&mut ParamsEntry)) { + let mut entry = self + .entries + .remove(&name) + .map(|(_, entry)| entry) + .unwrap_or_default(); + mutate(&mut entry); + let hash = hash_entry(&entry); + self.entries.insert(name, (hash, entry)); + } + /// Sets the full state (instruction + demos) for a named predictor. + /// + /// `PredictState` semantics: `instruction_override: None` and empty + /// `demos` mean "leave the incumbent alone". For explicit resets use + /// [`clear_instruction`](Params::clear_instruction) / + /// [`set_demos`](Params::set_demos). pub fn set(&mut self, name: impl Into, state: PredictState) { - let hash = hash_state(&state); - self.entries.insert(name.into(), (hash, state)); + self.upsert(name.into(), |entry| { + *entry = ParamsEntry { + state, + clear_instruction: false, + explicit_demos: false, + }; + }); } /// Convenience: overrides just the instruction for a named predictor, /// preserving any demos already set. pub fn set_instruction(&mut self, name: impl Into, instruction: impl Into) { - let name = name.into(); - let mut state = self - .entries - .remove(&name) - .map(|(_, state)| state) - .unwrap_or_default(); - state.instruction_override = Some(instruction.into()); - self.set(name, state); + let instruction = instruction.into(); + self.upsert(name.into(), |entry| { + entry.state.instruction_override = Some(instruction); + entry.clear_instruction = false; + }); + } + + /// Explicitly clears the instruction override back to the signature + /// default, preserving any demos already set. Unlike leaving the + /// instruction unset (which lets an instance override read through), this + /// wins over instance state when injected ambiently. + pub fn clear_instruction(&mut self, name: impl Into) { + self.upsert(name.into(), |entry| { + entry.state.instruction_override = None; + entry.clear_instruction = true; + }); + } + + /// Explicitly sets the demo set for a named predictor, preserving any + /// instruction already set. An empty vec means "no demos" and wins over + /// instance demos when injected ambiently. + pub fn set_demos(&mut self, name: impl Into, demos: Vec) { + self.upsert(name.into(), |entry| { + entry.state.demos = demos; + entry.explicit_demos = true; + }); } /// Returns the state configured for `name`, if any. pub fn get(&self, name: &str) -> Option<&PredictState> { - self.entries.get(name).map(|(_, state)| state) + self.entries.get(name).map(|(_, entry)| &entry.state) } pub fn is_empty(&self) -> bool { @@ -110,7 +171,7 @@ impl Params { predictors: self .entries .iter() - .map(|(name, (_, state))| (name.clone(), state.clone())) + .map(|(name, (_, entry))| (name.clone(), entry.state.clone())) .collect(), } } @@ -124,32 +185,43 @@ impl Params { params } - fn entry(&self, name: &str) -> Option<(u64, &PredictState)> { - self.entries.get(name).map(|(hash, state)| (*hash, state)) + fn entry(&self, name: &str) -> Option<(u64, &ParamsEntry)> { + self.entries.get(name).map(|(hash, entry)| (*hash, entry)) } /// All named states, for the IR bridge (`Params::bind`). - #[cfg(feature = "ir")] pub(crate) fn iter_states(&self) -> impl Iterator { self.entries .iter() - .map(|(name, (_, state))| (name.as_str(), state)) + .map(|(name, (_, entry))| (name.as_str(), &entry.state)) } } -fn hash_state(state: &PredictState) -> u64 { +fn hash_entry(entry: &ParamsEntry) -> u64 { let mut hasher = DefaultHasher::new(); - if let Some(instruction) = &state.instruction_override { + if let Some(instruction) = &entry.state.instruction_override { hasher.write(instruction.as_bytes()); } - for demo in &state.demos { + for demo in &entry.state.demos { let serialized = serde_json::to_string(demo).unwrap_or_default(); hasher.write(serialized.as_bytes()); } - hasher.write_usize(state.demos.len()); + hasher.write_usize(entry.state.demos.len()); + hasher.write_u8(entry.clear_instruction as u8); + hasher.write_u8(entry.explicit_demos as u8); hasher.finish() } +/// The ambient [`ParamsEntry`] for `name`, if a [`with_params`] scope is +/// active on this task. Read by struct-held [`Predict`](crate::Predict) +/// leaves at call time (each leaf binds only its own entry). +pub(crate) fn ambient_entry(name: &str) -> Option { + CURRENT_PARAMS + .try_with(|params| params.entry(name).map(|(_, entry)| entry.clone())) + .ok() + .flatten() +} + /// Runs a future with `params` as the ambient parameter set for every /// [`predict`] call inside it. /// @@ -160,8 +232,99 @@ pub async fn with_params(params: Params, fut: Fut) -> Fut::Output { CURRENT_PARAMS.scope(Arc::new(params), fut).await } +/// [`with_params`] without re-wrapping: scopes an already-shared `Params`. +/// Used by the optimizer engine, which evaluates many rollouts under one +/// candidate concurrently. +pub(crate) async fn with_params_shared(params: Arc, fut: Fut) -> Fut::Output { + CURRENT_PARAMS.scope(params, fut).await +} + type PredictorCacheKey = (TypeId, String, u64); -type PredictorCache = HashMap>; + +/// One cached predictor plus its CLOCK reference bit. The bit is set on every +/// hit (atomically, so the shared read lock suffices) and buys the entry a +/// second chance when the eviction hand sweeps past it. +struct CacheSlot { + predictor: Arc, + referenced: AtomicBool, +} + +/// Bounded predictor cache with second-chance (CLOCK) eviction. +/// +/// The previous design cleared the whole map at capacity, which flushed every +/// warm predictor mid-optimizer-run. CLOCK evicts one *cold* entry per insert +/// instead: recently-hit entries keep circulating, so an optimizer sweeping +/// many candidates retains its working set. +#[derive(Default)] +struct PredictorCache { + map: HashMap, + /// The clock ring: keys in sweep order. The hand is the front; entries + /// granted a second chance rotate to the back. + ring: VecDeque, +} + +impl PredictorCache { + fn get(&self, key: &PredictorCacheKey) -> Option> { + self.map.get(key).map(|slot| { + slot.referenced.store(true, Ordering::Relaxed); + slot.predictor.clone() + }) + } + + /// Inserts `predictor` under `key`, returning the cached instance (the + /// incumbent, if a concurrent writer got there first). Evicts at most one + /// cold entry when at capacity. + fn insert( + &mut self, + key: PredictorCacheKey, + predictor: Arc, + ) -> Arc { + if let Some(slot) = self.map.get(&key) { + slot.referenced.store(true, Ordering::Relaxed); + return slot.predictor.clone(); + } + if self.map.len() >= PREDICTOR_CACHE_CAP { + self.evict_one(); + } + self.ring.push_back(key.clone()); + self.map.insert( + key, + CacheSlot { + predictor: predictor.clone(), + referenced: AtomicBool::new(false), + }, + ); + predictor + } + + /// Advances the clock hand until it finds an entry whose reference bit is + /// clear, and evicts it. Bits are cleared as the hand passes, so this + /// terminates within one lap plus one step even if every entry was hot. + fn evict_one(&mut self) { + while let Some(key) = self.ring.pop_front() { + let Some(slot) = self.map.get(&key) else { + // Stale ring key with no map entry: drop it and keep sweeping. + continue; + }; + if slot.referenced.swap(false, Ordering::Relaxed) { + self.ring.push_back(key); + } else { + self.map.remove(&key); + return; + } + } + } + + #[cfg(test)] + fn len(&self) -> usize { + self.map.len() + } + + #[cfg(test)] + fn contains(&self, key: &PredictorCacheKey) -> bool { + self.map.contains_key(key) + } +} /// Predictor-instance cache: (signature type, name, config hash) → `Arc>`. /// @@ -169,9 +332,9 @@ type PredictorCache = HashMap>; /// structs — a cache hit reuses a fully-warmed `Predict` (prompt prefix, toolset) /// instead of rebuilding per call. static PREDICTOR_CACHE: LazyLock> = - LazyLock::new(|| RwLock::new(HashMap::new())); + LazyLock::new(|| RwLock::new(PredictorCache::default())); -/// Backstop against unbounded growth across many optimizer candidates. +/// Capacity bound; reached, the CLOCK sweep evicts one cold entry per insert. const PREDICTOR_CACHE_CAP: usize = 1024; #[allow(clippy::result_large_err)] @@ -185,7 +348,7 @@ where .try_with(|params| { params .entry(name) - .map(|(hash, state)| (hash, Some(state.clone()))) + .map(|(hash, entry)| (hash, Some(entry.state.clone()))) }) .ok() .flatten() @@ -196,7 +359,6 @@ where let cache = PREDICTOR_CACHE.read().expect("fx predictor cache poisoned"); if let Some(cached) = cache.get(&key) { return Ok(cached - .clone() .downcast::>() .expect("fx predictor cache entry has matching TypeId")); } @@ -204,25 +366,16 @@ where let mut predictor = Predict::::builder().named(name).build(); if let Some(state) = state { - predictor.load_state(state).map_err(|err| PredictError::Lm { - source: LmError::Provider { - provider: "fx".to_string(), - message: format!("params for `{name}` don't fit signature: {err}"), - source: None, - }, + PredictorInfo::load_state(&mut predictor, state).map_err(|err| PredictError::Params { + name: name.to_string(), + source: err.into(), })?; } let predictor = Arc::new(predictor); let mut cache = PREDICTOR_CACHE.write().expect("fx predictor cache poisoned"); - if cache.len() >= PREDICTOR_CACHE_CAP { - cache.clear(); - } - let entry = cache - .entry(key) - .or_insert_with(|| predictor.clone() as Arc); + let entry = cache.insert(key, predictor as Arc); Ok(entry - .clone() .downcast::>() .expect("fx predictor cache entry has matching TypeId")) } @@ -282,3 +435,65 @@ where _marker: PhantomData, } } + +#[cfg(test)] +mod tests { + use super::*; + + fn key(n: u64) -> PredictorCacheKey { + (TypeId::of::<()>(), format!("predictor-{n}"), n) + } + + fn slot() -> Arc { + Arc::new(()) as Arc + } + + #[test] + fn cache_stays_bounded_at_cap() { + let mut cache = PredictorCache::default(); + for n in 0..(PREDICTOR_CACHE_CAP as u64 + 200) { + cache.insert(key(n), slot()); + } + assert_eq!(cache.len(), PREDICTOR_CACHE_CAP); + } + + #[test] + fn eviction_spares_the_warm_working_set() { + let mut cache = PredictorCache::default(); + for n in 0..PREDICTOR_CACHE_CAP as u64 { + cache.insert(key(n), slot()); + } + // A warm working set: the first 16 entries keep getting hit. + let working_set: Vec<_> = (0..16).map(key).collect(); + for k in &working_set { + assert!(cache.get(k).is_some()); + } + // Sweep in twice the capacity of fresh candidates; each insert evicts + // one cold entry. The warm set must survive the first sweep wave, and + // as long as it keeps getting hit between waves, every wave. + for n in 0..PREDICTOR_CACHE_CAP as u64 { + cache.insert(key(1_000_000 + n), slot()); + if n % 64 == 0 { + for k in &working_set { + assert!(cache.get(k).is_some(), "warm entry evicted mid-sweep"); + } + } + } + for k in &working_set { + assert!(cache.contains(k), "warm entry evicted by candidate sweep"); + } + assert_eq!(cache.len(), PREDICTOR_CACHE_CAP); + } + + #[test] + fn reinserting_existing_key_returns_incumbent() { + let mut cache = PredictorCache::default(); + let first = slot(); + let incumbent = cache.insert(key(1), first.clone()); + assert!(Arc::ptr_eq(&incumbent, &first)); + let second = slot(); + let returned = cache.insert(key(1), second); + assert!(Arc::ptr_eq(&returned, &first), "incumbent must win"); + assert_eq!(cache.len(), 1); + } +} diff --git a/crates/dspy-rs/src/ir/bridge.rs b/crates/dspy-rs/src/ir/bridge.rs index 6c1d42f5..eb4c3343 100644 --- a/crates/dspy-rs/src/ir/bridge.rs +++ b/crates/dspy-rs/src/ir/bridge.rs @@ -29,10 +29,10 @@ //! [`OverlayError::DemoField`] error. //! //! In the overlay → params/state direction the projection is **restricted to -//! `Instruction` and `Demos` kinds** (RFC 0002 §2.4): `ToolDesc`, `ModelRef`, -//! `ContextPolicy`, and `Code` entries have no fx/ModuleState representation -//! and are skipped, never errors — the static lane simply has no slot for -//! them. +//! `Instruction` and `Demos` kinds** (RFC 0002 §2.4): `ToolDesc`, `ToolSet`, +//! `ModelRef`, `ContextPolicy`, and `Code` entries have no fx/ModuleState +//! representation and are skipped, never errors — the static lane simply has +//! no slot for them. use std::collections::BTreeMap; @@ -201,6 +201,56 @@ where Ok(overlay) } +/// Resolves one named [`fx::ParamsEntry`](crate::fx) against `program` into +/// `(ParamId, ParamValue)` pairs — the flag-aware sibling of +/// [`states_to_overlay`] used by `Predict`'s ambient-params composition: +/// +/// - `clear_instruction` resolves to the slot's *default* value (an explicit +/// entry, so it wins over instance state when composed); +/// - `explicit_demos` emits a `Demos` entry even when the row set is empty +/// (clearing instance demos); +/// - otherwise plain `PredictState` semantics apply (no override/no demos = +/// no entry, the incumbent reads through). +pub(crate) fn entry_slot_values( + program: &Program, + name: &str, + entry: &crate::fx::ParamsEntry, +) -> Result, OverlayError> { + let mut values = Vec::new(); + + let instruction_path = format!("{name}.instruction"); + let instruction_id = + program + .param_id(&instruction_path) + .ok_or_else(|| OverlayError::UnknownPath { + path: instruction_path, + })?; + if entry.clear_instruction { + values.push((instruction_id, program.params[instruction_id].default.clone())); + } else if let Some(text) = &entry.state.instruction_override { + values.push((instruction_id, ParamValue::Instruction { text: text.clone() })); + } + + if entry.explicit_demos || !entry.state.demos.is_empty() { + let demos_path = format!("{name}.demos"); + let demos_id = program + .param_id(&demos_path) + .ok_or_else(|| OverlayError::UnknownPath { + path: demos_path.clone(), + })?; + let def = leaf_sig_of(program, demos_id); + let rows = entry + .state + .demos + .iter() + .map(|flat| split_demo_row(def, &demos_path, flat)) + .collect::, _>>()?; + values.push((demos_id, ParamValue::Demos { rows })); + } + + Ok(values) +} + /// Projects an overlay down to `(leaf name → PredictState)`, restricted to /// `Instruction`/`Demos` kinds — the shared core of /// [`fx::Params::from_overlay`](crate::fx::Params::from_overlay) and @@ -234,8 +284,9 @@ pub(crate) fn overlay_to_states( states.entry(leaf.to_string()).or_default().demos = rows.iter().map(flatten_demo_row).collect(); } - // Restricted projection (RFC 0002 §2.4): ToolDesc / ModelRef / - // ContextPolicy / Code have no representation in the static lane. + // Restricted projection (RFC 0002 §2.4): ToolDesc / ToolSet / + // ModelRef / ContextPolicy / Code have no representation in the + // static lane. _ => {} } } diff --git a/crates/dspy-rs/src/ir/builder.rs b/crates/dspy-rs/src/ir/builder.rs index 8c131943..69319117 100644 --- a/crates/dspy-rs/src/ir/builder.rs +++ b/crates/dspy-rs/src/ir/builder.rs @@ -155,6 +155,9 @@ enum SpecKind { instruction: Option, demos: Vec, tools: Vec, + /// The `tool_set` gene's default selection; `None` = the full + /// declared `tools` list. + tool_set: Option>, stop_tools: Vec, max_turns: NonZeroU32, until_parse: bool, @@ -274,6 +277,17 @@ impl NodeSpec { self } + /// Seeds the `tool_set` gene: which declared tools the loop carries at + /// run time. Default (absent): the full declared `tools` list. Must be a + /// subset of the declared tools — validation refuses anything else. + pub fn tool_set(mut self, ids: impl IntoIterator) -> Self { + match &mut self.kind { + SpecKind::Agent { tool_set, .. } => *tool_set = Some(ids.into_iter().collect()), + _ => panic!("tool_set() applies to agent specs"), + } + self + } + /// Declares stop tools — calling one ends the loop, its args become the /// raw final output. pub fn stop_tools(mut self, ids: impl IntoIterator) -> Self { @@ -460,6 +474,7 @@ pub fn agent(name: &str, sig: SigId) -> NodeSpec { instruction: None, demos: Vec::new(), tools: Vec::new(), + tool_set: None, stop_tools: Vec::new(), max_turns: StopSpec::default().max_turns, until_parse: true, @@ -875,6 +890,7 @@ impl Lowering { instruction, demos, tools, + tool_set, stop_tools, max_turns, until_parse, @@ -919,6 +935,18 @@ impl Lowering { ParamKind::ContextPolicy, ParamValue::ContextPolicy { policy: context }, ); + // The gene defaults to the full declared table: absent + // selection = every declared tool, so pre-ToolSet programs + // print, hash, and run unchanged. + let tool_set = self.leaf_param( + &name, + "tool_set", + node_id, + ParamKind::ToolSet, + ParamValue::ToolSet { + tools: tool_set.unwrap_or_else(|| tools.clone()), + }, + ); let binding = self.lower_binds(binds)?; Node::AgentLoop(AgentLoopNode { name: name_sym, @@ -927,6 +955,7 @@ impl Lowering { demos, model, tools: tools.into_boxed_slice(), + tool_set, context_policy, stop: StopSpec { max_turns, diff --git a/crates/dspy-rs/src/ir/edit.rs b/crates/dspy-rs/src/ir/edit.rs new file mode 100644 index 00000000..cb9860cb --- /dev/null +++ b/crates/dspy-rs/src/ir/edit.rs @@ -0,0 +1,1105 @@ +//! The graph-edit calculus: the *structural* mutation half of the IR. +//! +//! [`Overlay`](crate::ir::params::Overlay) mutates parameter **values** over a +//! fixed skeleton; [`Edit`] mutates the skeleton itself. Edits are plain serde +//! values — inspectable, diffable, replayable — and are only ever applied +//! through [`Program::edited`], which is pure: clone the arenas, apply the +//! edits in order, re-run the same load-time validation the builder and loader +//! use, and seal a **new** content hash. A program value is never mutated in +//! place, so every hash-bound artifact (overlays, traces, caches) minted +//! against the parent stays coherent. +//! +//! # Design decisions +//! +//! - **NodeIds are positional handles against the parent.** Within one +//! `edited()` batch, ids stay stable (swaps happen in place, removals only +//! detach); dead nodes/sigs/params are garbage-collected once at the end. +//! Ids in the child may therefore differ from the parent — re-locate leaves +//! by name ([`Program::leaf_id`]) and params by `ParamPath`. +//! - **Signatures are copy-on-write.** [`Edit::AugmentSig`] never touches the +//! leaf's current [`SignatureDef`] (other nodes may share it); it pushes an +//! augmented copy. When the prepended field is exactly the `cot` reasoning +//! field on a `Predict`, the copy keeps the base name so the canonical +//! printer re-sugars it as `cot `; otherwise it gets a fresh unique +//! name (`_`) because two same-named `sig` blocks cannot print. +//! - **Batch validation.** Edits are validated as a *sequence*: intermediate +//! states may be inconsistent (e.g. remove a producer, then its consumer); +//! only the final program must pass `validate()`. Apply-time errors +//! ([`ApplyError`]) cover what is checkable locally (stale ids, wrong node +//! kind, unknown tools, capability ceilings); everything data-flow shaped is +//! deliberately left to `validate.rs` so the edit layer and the loader can +//! never disagree. +//! - **Garbage collection preserves identity.** `edited(&[])` returns a +//! program with the parent's hash (only lineage differs, and lineage is +//! outside the hash preimage). Signatures that were already unreferenced in +//! the parent are kept; only *newly* orphaned ones are collected. Orphaned +//! param slots never reach the canonical text, so they are always dropped. +//! - **Lineage.** `edited()` stamps `lineage.parent` with the parent's +//! `program_hash` exactly like `bake()`; the other provenance fields are +//! left empty for the optimizer to fill (an edit is not an optimization run +//! record). +//! - **[`migrate_overlay`]** carries value-level progress across a structural +//! edit. An entry survives when the child has a slot at the same `ParamPath` +//! and kind whose owning leaf/tool still has a *carrying* signature: inputs +//! identical (names + types, in order) and every parent output present in +//! the child's outputs (name + type). Outputs may widen — that is what lets +//! instruction and demos survive [`Edit::AugmentSig`] (demo rows still map +//! onto the base fields; the new field is simply absent from the row). +//! `ModelRef` entries are re-minted by model *name*, not ordinal. + +use std::collections::{HashMap, HashSet}; +use std::num::NonZeroU32; + +use cranelift_entity::{EntityRef, PrimaryMap}; +use serde::{Deserialize, Serialize}; + +use crate::ir::builder::cot_reasoning_field; +use crate::ir::graph::{ + AgentLoopNode, HoleImpl, Lineage, Node, NodeBudget, NodeId, PortRef, PredictNode, Program, + RetryNode, SigId, StopSpec, ToolId, ToolKind, +}; +use crate::ir::params::{ + ContextPolicy, Overlay, ParamId, ParamKind, ParamOwner, ParamSlot, ParamValue, +}; +use crate::ir::sig::{FieldDef, SignatureDef}; +use crate::ir::validate::ValidateError; + +// --------------------------------------------------------------------------- +// Edits +// --------------------------------------------------------------------------- + +/// One structural edit. Serde values: an optimizer's proposal is data, not +/// code — it can be logged, replayed against the same parent, and diffed. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "edit", rename_all = "snake_case")] +pub enum Edit { + /// Prepend an output field to a `Predict`/`AgentLoop` leaf's signature — + /// the CoT move (mirrors [`SignatureDef::augmented_with`]). Copy-on-write: + /// a new [`SigId`] is created; nodes sharing the old signature keep it. + AugmentSig { leaf: NodeId, prepend: FieldDef }, + /// Swap a leaf's kind: `Predict` → `AgentLoop` (with the given tools ⊆ + /// `program.tools`, stop spec and budget) or `AgentLoop` → `Predict`. + /// Name, signature, bindings, and the instruction/demos/model param slots + /// are preserved; the `AgentLoop` direction mints a `.context` + /// slot, the `Predict` direction drops it. + SwapLeaf { leaf: NodeId, to: SwapTarget }, + /// Wrap an existing node in a [`RetryNode`], rewiring the parent + /// reference and redirecting downstream `Out` ports to the wrapper (the + /// wrapper, not the child, is what later siblings can see). + WrapRetry { + node: NodeId, + max_attempts: NonZeroU32, + backoff_ms: u32, + feedback: bool, + }, + /// Remove a node from its parent `Seq` body (subtree and its params are + /// garbage-collected). If a later binding still references its outputs, + /// `validate()` rejects the batch with its own error. + Remove { node: NodeId }, + /// Declare an existing program tool on an agent leaf. + AddTool { agent: NodeId, tool: ToolId }, + /// Undeclare a tool from an agent leaf (also removed from `stop_tools`). + RemoveTool { agent: NodeId, tool: ToolId }, + /// Replace an agent leaf's [`StopSpec`]. + SetStop { agent: NodeId, stop: StopSpec }, + /// Set the leaf's instruction slot *default* (a bake-like change without + /// an overlay) — for structural optimizers that also seed text. + SetInstructionDefault { leaf: NodeId, text: String }, +} + +/// Target kind of [`Edit::SwapLeaf`]. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(tag = "to", rename_all = "snake_case")] +pub enum SwapTarget { + Agent { + tools: Vec, + #[serde(default)] + stop: StopSpec, + #[serde(default)] + budget: NodeBudget, + }, + Predict, +} + +/// A lightweight, serializable descriptor of an edit kind admissible at a +/// node — the menu [`Program::legal_edits`] returns, suitable for prompting +/// an LLM proposer. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "kind", rename_all = "snake_case")] +pub enum EditKind { + AugmentSig, + SwapToAgent, + SwapToPredict, + WrapRetry, + Remove, + AddTool { tool: ToolId }, + RemoveTool { tool: ToolId }, + SetStop, + SetInstructionDefault, +} + +// --------------------------------------------------------------------------- +// Errors +// --------------------------------------------------------------------------- + +/// Why [`Program::edited`] refused. +#[derive(Debug, Clone, PartialEq, thiserror::Error)] +pub enum EditError { + /// Edit `index` could not be applied to the (partially edited) program. + /// Carries the offending edit (boxed: errors stay word-sized on the Ok + /// path). + #[error("edit #{index} ({edit:?}) failed: {reason}")] + Apply { + index: usize, + edit: Box, + reason: ApplyError, + }, + /// Every edit applied, but the resulting program failed the load-time + /// rules — the error is `validate.rs`'s own. + #[error("edited program failed validation: {0}")] + Invalid(#[from] ValidateError), +} + +/// A locally-checkable application failure. +#[derive(Debug, Clone, PartialEq, thiserror::Error)] +pub enum ApplyError { + #[error("no node {node} in this program (stale NodeId)")] + StaleNode { node: NodeId }, + #[error("{node} is a `{got}` node, expected {expected}")] + WrongKind { + node: NodeId, + expected: &'static str, + got: &'static str, + }, + #[error("signature of {node} already has a field `{field}`")] + DuplicateField { node: NodeId, field: String }, + #[error("no tool {tool} in this program (stale ToolId)")] + UnknownTool { tool: ToolId }, + #[error("tool `{name}` caps exceed the program ceiling: missing {missing:?}")] + ToolCapsExceedProgram { name: String, missing: Vec }, + #[error("tool `{name}` is already declared on {agent}")] + ToolAlreadyDeclared { agent: NodeId, name: String }, + #[error("tool `{name}` is not declared on {agent}")] + ToolNotDeclared { agent: NodeId, name: String }, + #[error("{node} is not a step of a `Seq` (only Seq steps can be removed)")] + NotInSeq { node: NodeId }, + #[error("{node} has no parent to rewire (detached by an earlier edit?)")] + Unparented { node: NodeId }, +} + +// --------------------------------------------------------------------------- +// Program surface +// --------------------------------------------------------------------------- + +impl Program { + /// Applies `edits` in order to a clone of `self` and returns the sealed, + /// validated result. `self` is never mutated. The child gets a **new** + /// content hash and `lineage.parent` set to `self`'s hash (like + /// [`Program::bake`]); overlays minted against `self` must be re-minted + /// (see [`migrate_overlay`]). + pub fn edited(&self, edits: &[Edit]) -> Result { + let mut work = self.clone(); + for (index, edit) in edits.iter().enumerate() { + apply(&mut work, edit).map_err(|reason| EditError::Apply { + index, + edit: Box::new(edit.clone()), + reason, + })?; + } + + // `Remove` only detaches; dead subtrees are still in the arena here. + // Validate the un-collected graph first so a downstream reference to + // a removed node surfaces as validate.rs's own error (NodeNotVisible + // et al.) rather than a remap failure. If the *only* complaint is + // unreachable nodes, collection is exactly the fix. + let reachable = reachable_nodes(&work); + let node_map: HashMap = if reachable.len() == work.nodes.len() { + work.nodes.keys().map(|id| (id, id)).collect() + } else { + if let Err(err) = work.validate() + && !matches!(err, ValidateError::UnreachableNodes { .. }) + { + return Err(EditError::Invalid(err)); + } + gc_nodes(&mut work, &reachable) + }; + + gc_sigs(&mut work, self); + gc_params(&mut work, &node_map); + + work.rebuild_param_index().map_err(EditError::Invalid)?; + // Validate before sealing: the hash preimage is the canonical printed + // text, and printing assumes structurally valid arenas. + work.validate().map_err(EditError::Invalid)?; + work.meta.lineage = Some(Lineage { + parent: Some(format!("{:016x}", self.meta.program_hash).into()), + ..Lineage::default() + }); + work.seal(); + Ok(work) + } + + /// The menu of edit kinds structurally admissible at `at`: leaf-only + /// moves for leaves (split by `Predict`/`AgentLoop`), per-tool add/remove + /// entries for agents, `WrapRetry` for any non-root node that is not a + /// `Refine` judge (judges must stay bare leaves), `Remove` for `Seq` + /// steps. Purely structural — data-flow legality (e.g. whether a removal + /// orphans a downstream binding) is still `validate()`'s call. A stale id + /// yields an empty menu. + pub fn legal_edits(&self, at: NodeId) -> Vec { + let Some(node) = self.nodes.get(at) else { + return Vec::new(); + }; + let mut out = Vec::new(); + match node { + Node::Predict(_) => { + out.push(EditKind::AugmentSig); + out.push(EditKind::SetInstructionDefault); + out.push(EditKind::SwapToAgent); + } + Node::AgentLoop(n) => { + out.push(EditKind::AugmentSig); + out.push(EditKind::SetInstructionDefault); + out.push(EditKind::SwapToPredict); + out.push(EditKind::SetStop); + for (tool, _) in self.tools.iter() { + if n.tools.contains(&tool) { + out.push(EditKind::RemoveTool { tool }); + } else { + out.push(EditKind::AddTool { tool }); + } + } + } + _ => {} + } + let parent = self.parent_of(at); + let is_judge = matches!( + parent.map(|p| &self.nodes[p]), + Some(Node::Refine(r)) if r.judge == at + ); + if at != self.root && !is_judge { + out.push(EditKind::WrapRetry); + } + if matches!(parent.map(|p| &self.nodes[p]), Some(Node::Seq(_))) { + out.push(EditKind::Remove); + } + out + } + + /// The node id of the leaf named `name`, if any. Leaf names are + /// program-unique and survive edits, which makes them the stable way to + /// re-locate a node across [`Program::edited`]. + pub fn leaf_id(&self, name: &str) -> Option { + self.nodes.iter().find_map(|(id, node)| { + node.leaf_name() + .is_some_and(|sym| self.syms.get(sym) == name) + .then_some(id) + }) + } + + /// The structural parent of `at` (`None` for the root or a stale id). + fn parent_of(&self, at: NodeId) -> Option { + self.nodes + .iter() + .find_map(|(id, node)| structural_children(node).contains(&at).then_some(id)) + } +} + +// --------------------------------------------------------------------------- +// Overlay migration +// --------------------------------------------------------------------------- + +/// Carries tuned values across a structural edit: for every entry in +/// `overlay` (minted against `parent`), re-mint it against `child` when the +/// child has a slot at the same `ParamPath` and kind whose owning leaf/tool +/// signature still *carries* the parent's — inputs identical, parent outputs +/// a subset of the child's (so [`Edit::AugmentSig`] keeps instruction and +/// demos alive; see the module docs). Entries that no longer fit are dropped. +/// The result is based on `child`'s hash. A base-mismatched `overlay` yields +/// an empty result rather than indexing with foreign ids. +pub fn migrate_overlay(parent: &Program, overlay: &Overlay, child: &Program) -> Overlay { + let mut out = Overlay::new(child); + if overlay.base != parent.meta.program_hash { + return out; + } + for (id, value) in overlay.entries() { + let path = parent.param_path(id); + let Some(child_id) = child.param_id(path) else { + continue; + }; + if child.params[child_id].kind != value.kind() { + continue; + } + let (Some(psig), Some(csig)) = ( + owner_sig(parent, parent.params[id].owner), + owner_sig(child, child.params[child_id].owner), + ) else { + continue; + }; + if !sig_carries(&parent.sigs[psig], &child.sigs[csig]) { + continue; + } + let value = match value { + // Model refs are ordinals into `models`; re-mint by name. + ParamValue::ModelRef { model } => { + let Some(def) = parent.models.get(*model) else { + continue; + }; + let Some((child_model, _)) = child.models.iter().find(|(_, m)| m.name == def.name) + else { + continue; + }; + ParamValue::ModelRef { model: child_model } + } + // Tool sets carry ordinals into `tools`; re-mint each by name + // and keep the intersection with what the child's agent still + // declares — partial survival is the point of migration. A + // selection with no survivors no longer fits and is dropped. + ParamValue::ToolSet { tools } => { + let declared: &[ToolId] = match child.params[child_id].owner { + ParamOwner::Node(node) => match &child.nodes[node] { + Node::AgentLoop(n) => &n.tools, + _ => continue, + }, + ParamOwner::Tool(_) => continue, + }; + let mut migrated: Vec = Vec::new(); + for t in tools { + let Some(def) = parent.tools.get(*t) else { + continue; + }; + let name = parent.syms.get(def.name); + let Some((child_tool, _)) = child + .tools + .iter() + .find(|(_, d)| child.syms.get(d.name) == name) + else { + continue; + }; + if declared.contains(&child_tool) && !migrated.contains(&child_tool) { + migrated.push(child_tool); + } + } + if migrated.is_empty() && !tools.is_empty() { + continue; + } + ParamValue::ToolSet { tools: migrated } + } + other => other.clone(), + }; + // Kind was checked above; set cannot fail, but stay total. + let _ = out.set(child, child_id, value); + } + out +} + +/// `parent` signature values still make sense on `child`: inputs identical +/// (names + types, in order), every parent output present among the child's +/// outputs (name + type). Docs/constraints/render are shape-irrelevant. +fn sig_carries(parent: &SignatureDef, child: &SignatureDef) -> bool { + parent.inputs.len() == child.inputs.len() + && parent + .inputs + .iter() + .zip(child.inputs.iter()) + .all(|(a, b)| a.name == b.name && a.ty == b.ty) + && parent.outputs.iter().all(|f| { + child + .outputs + .iter() + .any(|g| g.name == f.name && g.ty == f.ty) + }) +} + +/// The signature of a slot's owning leaf or tool (`None` for a stale owner +/// or a non-leaf node). +fn owner_sig(p: &Program, owner: ParamOwner) -> Option { + match owner { + ParamOwner::Node(id) => match p.nodes.get(id)? { + Node::Predict(n) => Some(n.sig), + Node::AgentLoop(n) => Some(n.sig), + Node::Hole(n) => Some(n.sig), + _ => None, + }, + ParamOwner::Tool(id) => p.tools.get(id).map(|t| t.sig), + } +} + +// --------------------------------------------------------------------------- +// Edit application +// --------------------------------------------------------------------------- + +fn apply(work: &mut Program, edit: &Edit) -> Result<(), ApplyError> { + match edit { + Edit::AugmentSig { leaf, prepend } => augment_sig(work, *leaf, prepend), + Edit::SwapLeaf { leaf, to } => swap_leaf(work, *leaf, to), + Edit::WrapRetry { + node, + max_attempts, + backoff_ms, + feedback, + } => wrap_retry(work, *node, *max_attempts, *backoff_ms, *feedback), + Edit::Remove { node } => remove(work, *node), + Edit::AddTool { agent, tool } => add_tool(work, *agent, *tool), + Edit::RemoveTool { agent, tool } => remove_tool(work, *agent, *tool), + Edit::SetStop { agent, stop } => { + agent_mut(work, *agent)?.stop = stop.clone(); + Ok(()) + } + Edit::SetInstructionDefault { leaf, text } => set_instruction_default(work, *leaf, text), + } +} + +fn node_checked(work: &Program, id: NodeId) -> Result<&Node, ApplyError> { + work.nodes.get(id).ok_or(ApplyError::StaleNode { node: id }) +} + +fn agent_mut(work: &mut Program, id: NodeId) -> Result<&mut AgentLoopNode, ApplyError> { + match node_checked(work, id)? { + Node::AgentLoop(_) => {} + other => { + return Err(ApplyError::WrongKind { + node: id, + expected: "an `agent` leaf", + got: kind_label(other), + }); + } + } + match &mut work.nodes[id] { + Node::AgentLoop(n) => Ok(n), + _ => unreachable!("checked above"), + } +} + +fn kind_label(node: &Node) -> &'static str { + match node { + Node::Predict(_) => "predict", + Node::AgentLoop(_) => "agent", + Node::Hole(_) => "hole", + Node::Seq(_) => "seq", + Node::ForkJoin(_) => "fork", + Node::Route(_) => "route", + Node::Retry(_) => "retry", + Node::Refine(_) => "refine", + Node::Loop(_) => "loop", + } +} + +fn augment_sig(work: &mut Program, leaf: NodeId, prepend: &FieldDef) -> Result<(), ApplyError> { + let (sig_id, is_predict) = match node_checked(work, leaf)? { + Node::Predict(n) => (n.sig, true), + Node::AgentLoop(n) => (n.sig, false), + other => { + return Err(ApplyError::WrongKind { + node: leaf, + expected: "a `predict` or `agent` leaf", + got: kind_label(other), + }); + } + }; + let base = work.sigs[sig_id].clone(); + if base + .inputs + .iter() + .chain(base.outputs.iter()) + .any(|f| f.name == prepend.name) + { + return Err(ApplyError::DuplicateField { + node: leaf, + field: prepend.name.to_string(), + }); + } + let mut augmented = base.augmented_with(std::slice::from_ref(prepend)); + // The exact `cot` reasoning field on a Predict keeps the base name: the + // canonical printer re-sugars it as `cot ` against the base, so the + // shared name never prints twice. Any other augmentation is a new + // declaration and needs its own name. + if !(is_predict && *prepend == cot_reasoning_field()) { + augmented.name = unique_sig_name(work, &base.name, &prepend.name); + } + let new_sig = work.sigs.push(augmented); + match &mut work.nodes[leaf] { + Node::Predict(n) => n.sig = new_sig, + Node::AgentLoop(n) => n.sig = new_sig, + _ => unreachable!("leaf kind checked above"), + } + Ok(()) +} + +/// `_`, uniquified against every name the text format resolves +/// in or near the signature namespace (sigs, tools, class/enum tokens). +fn unique_sig_name(p: &Program, base: &str, field: &str) -> Box { + let mut taken: HashSet = p.sigs.values().map(|s| s.name.to_string()).collect(); + taken.extend(p.tools.values().map(|t| p.syms.get(t.name).to_string())); + taken.extend(p.types.classes.keys().cloned()); + taken.extend(p.types.enums.keys().cloned()); + let stem = format!("{base}_{field}"); + if !taken.contains(&stem) { + return stem.into(); + } + let mut i = 2usize; + loop { + let candidate = format!("{stem}{i}"); + if !taken.contains(&candidate) { + return candidate.into(); + } + i += 1; + } +} + +fn swap_leaf(work: &mut Program, leaf: NodeId, to: &SwapTarget) -> Result<(), ApplyError> { + let node = node_checked(work, leaf)?.clone(); + match (node, to) { + ( + Node::Predict(n), + SwapTarget::Agent { + tools, + stop, + budget, + }, + ) => { + let mut declared: Vec = Vec::new(); + for &tool in tools { + let def = tool_checked(work, tool)?; + let name = work.syms.get(def.name).to_string(); + if !def.caps.is_subset(&work.caps) { + return Err(ApplyError::ToolCapsExceedProgram { + name, + missing: def.caps.missing_from(&work.caps), + }); + } + if declared.contains(&tool) { + return Err(ApplyError::ToolAlreadyDeclared { agent: leaf, name }); + } + declared.push(tool); + } + let leaf_name = work.syms.get(n.name).to_string(); + // Same slot-creation convention as the builder: `.context` + // and `.tool_set` (default = the full declared table). + let context_policy = work.params.push(ParamSlot { + path: format!("{leaf_name}.context").into(), + owner: ParamOwner::Node(leaf), + kind: ParamKind::ContextPolicy, + default: ParamValue::ContextPolicy { + policy: ContextPolicy::default(), + }, + }); + let tool_set = work.params.push(ParamSlot { + path: format!("{leaf_name}.tool_set").into(), + owner: ParamOwner::Node(leaf), + kind: ParamKind::ToolSet, + default: ParamValue::ToolSet { + tools: declared.clone(), + }, + }); + work.nodes[leaf] = Node::AgentLoop(AgentLoopNode { + name: n.name, + sig: n.sig, + instruction: n.instruction, + demos: n.demos, + model: n.model, + tools: declared.into_boxed_slice(), + tool_set, + context_policy, + stop: stop.clone(), + budget: budget.clone(), + binding: n.binding, + }); + Ok(()) + } + (Node::AgentLoop(n), SwapTarget::Predict) => { + // The context and tool_set slots are orphaned here and collected + // in `edited()`. + work.nodes[leaf] = Node::Predict(PredictNode { + name: n.name, + sig: n.sig, + instruction: n.instruction, + demos: n.demos, + model: n.model, + binding: n.binding, + }); + Ok(()) + } + (other, SwapTarget::Agent { .. }) => Err(ApplyError::WrongKind { + node: leaf, + expected: "a `predict` leaf", + got: kind_label(&other), + }), + (other, SwapTarget::Predict) => Err(ApplyError::WrongKind { + node: leaf, + expected: "an `agent` leaf", + got: kind_label(&other), + }), + } +} + +fn wrap_retry( + work: &mut Program, + node: NodeId, + max_attempts: NonZeroU32, + backoff_ms: u32, + feedback: bool, +) -> Result<(), ApplyError> { + node_checked(work, node)?; + let attached = work.root == node + || work + .nodes + .values() + .any(|n| structural_children(n).contains(&node)); + if !attached { + return Err(ApplyError::Unparented { node }); + } + let retry = work.nodes.push(Node::Retry(RetryNode { + child: node, + max_attempts, + backoff_ms, + feedback, + })); + // Rewire the single structural parent (nodes form a tree) — or the root + // slot itself, in which case validate() rejects with RootNotSeq. + if work.root == node { + work.root = retry; + } else { + let ids: Vec = work.nodes.keys().collect(); + for id in ids { + if id != retry && replace_child(&mut work.nodes[id], node, retry) { + break; + } + } + } + // Downstream dataflow must reference the wrapper: scope visibility is + // sibling-level, and the retry is the sibling now. (Nothing inside the + // wrapped subtree can reference the subtree's own root, so a global + // redirect is safe.) + for (id, n) in work.nodes.iter_mut() { + if id != retry { + redirect_out_ports(n, node, retry); + } + } + Ok(()) +} + +fn remove(work: &mut Program, node: NodeId) -> Result<(), ApplyError> { + node_checked(work, node)?; + let ids: Vec = work.nodes.keys().collect(); + for id in ids { + if let Node::Seq(seq) = &mut work.nodes[id] + && seq.body.contains(&node) + { + let mut body = seq.body.to_vec(); + body.retain(|&child| child != node); + seq.body = body.into_boxed_slice(); + return Ok(()); + } + } + Err(ApplyError::NotInSeq { node }) +} + +fn add_tool(work: &mut Program, agent: NodeId, tool: ToolId) -> Result<(), ApplyError> { + let def = tool_checked(work, tool)?; + let name = work.syms.get(def.name).to_string(); + if !def.caps.is_subset(&work.caps) { + return Err(ApplyError::ToolCapsExceedProgram { + name, + missing: def.caps.missing_from(&work.caps), + }); + } + let n = agent_mut(work, agent)?; + if n.tools.contains(&tool) { + return Err(ApplyError::ToolAlreadyDeclared { agent, name }); + } + let mut tools = n.tools.to_vec(); + tools.push(tool); + n.tools = tools.into_boxed_slice(); + // Declaring a tool makes it live: the tool_set gene's default tracks the + // declaration (a baked subset grows by exactly the tool just declared). + let tool_set = n.tool_set; + if let ParamValue::ToolSet { tools } = &mut work.params[tool_set].default + && !tools.contains(&tool) + { + tools.push(tool); + } + Ok(()) +} + +fn remove_tool(work: &mut Program, agent: NodeId, tool: ToolId) -> Result<(), ApplyError> { + let def = tool_checked(work, tool)?; + let name = work.syms.get(def.name).to_string(); + let n = agent_mut(work, agent)?; + if !n.tools.contains(&tool) { + return Err(ApplyError::ToolNotDeclared { agent, name }); + } + n.tools = n.tools.iter().copied().filter(|t| *t != tool).collect(); + n.stop.stop_tools = n + .stop + .stop_tools + .iter() + .copied() + .filter(|t| *t != tool) + .collect(); + // The tool_set gene's alphabet shrank; drop the tool from the default + // selection too (validation would otherwise refuse the child). + let tool_set = n.tool_set; + if let ParamValue::ToolSet { tools } = &mut work.params[tool_set].default { + tools.retain(|t| *t != tool); + } + Ok(()) +} + +fn set_instruction_default(work: &mut Program, leaf: NodeId, text: &str) -> Result<(), ApplyError> { + let param = match node_checked(work, leaf)? { + Node::Predict(n) => n.instruction, + Node::AgentLoop(n) => n.instruction, + other => { + return Err(ApplyError::WrongKind { + node: leaf, + expected: "a `predict` or `agent` leaf", + got: kind_label(other), + }); + } + }; + work.params[param].default = ParamValue::Instruction { + text: text.to_string(), + }; + Ok(()) +} + +fn tool_checked(work: &Program, tool: ToolId) -> Result<&crate::ir::graph::ToolDef, ApplyError> { + work.tools.get(tool).ok_or(ApplyError::UnknownTool { tool }) +} + +// --------------------------------------------------------------------------- +// Tree plumbing +// --------------------------------------------------------------------------- + +/// The structural children of a node — the same set `validate()` walks. +fn structural_children(node: &Node) -> Vec { + match node { + Node::Predict(_) | Node::AgentLoop(_) | Node::Hole(_) => Vec::new(), + Node::Seq(n) => n.body.to_vec(), + Node::ForkJoin(n) => n.branches.to_vec(), + Node::Route(n) => n + .arms + .iter() + .map(|(_, arm)| *arm) + .chain(n.default) + .collect(), + Node::Retry(n) => vec![n.child], + Node::Refine(n) => vec![n.child, n.judge], + Node::Loop(n) => vec![n.body], + } +} + +/// Replaces `from` with `to` in a structural child position. Returns whether +/// a replacement happened. +fn replace_child(node: &mut Node, from: NodeId, to: NodeId) -> bool { + let slot_in = |slots: &mut [NodeId]| { + for slot in slots { + if *slot == from { + *slot = to; + return true; + } + } + false + }; + match node { + Node::Predict(_) | Node::AgentLoop(_) | Node::Hole(_) => false, + Node::Seq(n) => slot_in(&mut n.body), + Node::ForkJoin(n) => slot_in(&mut n.branches), + Node::Route(n) => { + for (_, arm) in n.arms.iter_mut() { + if *arm == from { + *arm = to; + return true; + } + } + if n.default == Some(from) { + n.default = Some(to); + return true; + } + false + } + Node::Retry(n) => { + if n.child == from { + n.child = to; + true + } else { + false + } + } + Node::Refine(n) => { + if n.child == from { + n.child = to; + true + } else if n.judge == from { + n.judge = to; + true + } else { + false + } + } + Node::Loop(n) => { + if n.body == from { + n.body = to; + true + } else { + false + } + } + } +} + +fn redirect_out_ports(node: &mut Node, from: NodeId, to: NodeId) { + let redirect = |port: &mut PortRef| { + if let PortRef::Out { node, .. } = port + && *node == from + { + *node = to; + } + }; + for_each_port(node, redirect); +} + +fn for_each_port(node: &mut Node, mut f: impl FnMut(&mut PortRef)) { + match node { + Node::Predict(n) => n.binding.iter_mut().for_each(|b| f(&mut b.src)), + Node::AgentLoop(n) => n.binding.iter_mut().for_each(|b| f(&mut b.src)), + Node::Hole(n) => n.binding.iter_mut().for_each(|b| f(&mut b.src)), + Node::Seq(n) => n.out.iter_mut().for_each(|b| f(&mut b.src)), + Node::ForkJoin(n) => n.join.iter_mut().for_each(|b| f(&mut b.src)), + Node::Route(n) => f(&mut n.on), + Node::Retry(_) | Node::Refine(_) => {} + Node::Loop(n) => { + if let Some(port) = &mut n.while_ { + f(port); + } + n.carry.iter_mut().for_each(|b| f(&mut b.src)); + n.out.iter_mut().for_each(|b| f(&mut b.src)); + } + } +} + +// --------------------------------------------------------------------------- +// Garbage collection +// --------------------------------------------------------------------------- + +fn reachable_nodes(p: &Program) -> HashSet { + let mut seen = HashSet::new(); + let mut stack = vec![p.root]; + while let Some(id) = stack.pop() { + if seen.insert(id) { + stack.extend(structural_children(&p.nodes[id])); + } + } + seen +} + +/// Drops unreachable nodes, remapping ids everywhere they occur. Only called +/// after `validate()` has confirmed the reachable subgraph is self-contained, +/// so every surviving reference remaps. +fn gc_nodes(work: &mut Program, reachable: &HashSet) -> HashMap { + let mut map = HashMap::new(); + let mut nodes: PrimaryMap = PrimaryMap::new(); + for (id, node) in work.nodes.iter() { + if reachable.contains(&id) { + map.insert(id, nodes.push(node.clone())); + } + } + for node in nodes.values_mut() { + remap_children(node, &map); + for_each_port(node, |port| { + if let PortRef::Out { node, .. } = port { + *node = map[node]; + } + }); + } + work.nodes = nodes; + work.root = map[&work.root]; + map +} + +fn remap_children(node: &mut Node, map: &HashMap) { + match node { + Node::Predict(_) | Node::AgentLoop(_) | Node::Hole(_) => {} + Node::Seq(n) => n.body.iter_mut().for_each(|c| *c = map[c]), + Node::ForkJoin(n) => n.branches.iter_mut().for_each(|c| *c = map[c]), + Node::Route(n) => { + n.arms.iter_mut().for_each(|(_, arm)| *arm = map[arm]); + if let Some(default) = &mut n.default { + *default = map[default]; + } + } + Node::Retry(n) => n.child = map[&n.child], + Node::Refine(n) => { + n.child = map[&n.child]; + n.judge = map[&n.judge]; + } + Node::Loop(n) => n.body = map[&n.body], + } +} + +/// Signatures a program actually uses: the program interface, every leaf and +/// tool signature, and the `cot` base of every sugar-detected Predict (the +/// printer needs the base in the arena to re-sugar). +fn referenced_sigs(p: &Program) -> HashSet { + let mut set = HashSet::new(); + set.insert(p.sig); + for node in p.nodes.values() { + match node { + Node::Predict(n) => { + set.insert(n.sig); + if let Some(base) = cot_base_of(p, n) { + set.insert(base); + } + } + Node::AgentLoop(n) => { + set.insert(n.sig); + } + Node::Hole(n) => { + set.insert(n.sig); + } + _ => {} + } + } + for tool in p.tools.values() { + set.insert(tool.sig); + } + set +} + +/// Mirrors the canonical printer's `cot` detection (print.rs): the node's +/// signature is `base.augmented_with([reasoning])` for some *other* arena +/// signature with identical name/instruction/inputs. +fn cot_base_of(p: &Program, n: &PredictNode) -> Option { + let sig = &p.sigs[n.sig]; + let first = sig.outputs.first()?; + if *first != cot_reasoning_field() { + return None; + } + p.sigs.iter().find_map(|(id, base)| { + (id != n.sig + && base.name == sig.name + && base.instruction == sig.instruction + && base.inputs == sig.inputs + && *base.outputs == sig.outputs[1..]) + .then_some(id) + }) +} + +/// Collects *newly* orphaned signatures. Sigs that were already unreferenced +/// in the parent stay (they print in both, keeping `edited(&[])` a hash +/// no-op); sigs orphaned by this batch (replaced by `AugmentSig`, or owned by +/// removed leaves) are dropped so same-named `sig` blocks never print twice. +fn gc_sigs(work: &mut Program, parent: &Program) { + let used = referenced_sigs(work); + if used.len() == work.sigs.len() { + return; + } + let parent_used = referenced_sigs(parent); + let parent_len = parent.sigs.len(); + let retained: Vec = work + .sigs + .keys() + .filter(|id| used.contains(id) || (id.index() < parent_len && !parent_used.contains(id))) + .collect(); + if retained.len() == work.sigs.len() { + return; + } + let mut map = HashMap::new(); + let mut sigs: PrimaryMap = PrimaryMap::new(); + for id in retained { + map.insert(id, sigs.push(work.sigs[id].clone())); + } + work.sigs = sigs; + work.sig = map[&work.sig]; + for node in work.nodes.values_mut() { + match node { + Node::Predict(n) => n.sig = map[&n.sig], + Node::AgentLoop(n) => n.sig = map[&n.sig], + Node::Hole(n) => n.sig = map[&n.sig], + _ => {} + } + } + for tool in work.tools.values_mut() { + tool.sig = map[&tool.sig]; + } +} + +/// Drops param slots no node or tool references (orphaned by `Remove` and by +/// `AgentLoop` → `Predict` swaps), remapping surviving [`ParamId`]s and their +/// owners. Orphans never reach the canonical text, so this never moves the +/// hash; it does keep `ParamPath`s collision-free across swap sequences. +fn gc_params(work: &mut Program, node_map: &HashMap) { + let mut used: HashSet = HashSet::new(); + for node in work.nodes.values() { + match node { + Node::Predict(n) => used.extend([n.instruction, n.demos, n.model]), + Node::AgentLoop(n) => { + used.extend([ + n.instruction, + n.demos, + n.model, + n.tool_set, + n.context_policy, + ]); + } + Node::Hole(n) => { + if let HoleImpl::Sandboxed { code } = n.imp { + used.insert(code); + } + } + _ => {} + } + } + for tool in work.tools.values() { + used.insert(tool.desc); + if let ToolKind::Sandboxed { code } = tool.kind { + used.insert(code); + } + } + let identity_nodes = node_map.iter().all(|(from, to)| from == to); + if used.len() == work.params.len() && identity_nodes { + return; + } + let mut map = HashMap::new(); + let mut params: PrimaryMap = PrimaryMap::new(); + for (id, slot) in work.params.iter() { + if !used.contains(&id) { + continue; + } + let mut slot = slot.clone(); + if let ParamOwner::Node(owner) = slot.owner { + slot.owner = ParamOwner::Node(node_map[&owner]); + } + map.insert(id, params.push(slot)); + } + work.params = params; + for node in work.nodes.values_mut() { + match node { + Node::Predict(n) => { + n.instruction = map[&n.instruction]; + n.demos = map[&n.demos]; + n.model = map[&n.model]; + } + Node::AgentLoop(n) => { + n.instruction = map[&n.instruction]; + n.demos = map[&n.demos]; + n.model = map[&n.model]; + n.tool_set = map[&n.tool_set]; + n.context_policy = map[&n.context_policy]; + } + Node::Hole(n) => { + if let HoleImpl::Sandboxed { code } = &mut n.imp { + *code = map[code]; + } + } + _ => {} + } + } + for tool in work.tools.values_mut() { + tool.desc = map[&tool.desc]; + if let ToolKind::Sandboxed { code } = &mut tool.kind { + *code = map[code]; + } + } +} diff --git a/crates/dspy-rs/src/ir/graph.rs b/crates/dspy-rs/src/ir/graph.rs index 1d1b777c..31d9be57 100644 --- a/crates/dspy-rs/src/ir/graph.rs +++ b/crates/dspy-rs/src/ir/graph.rs @@ -277,7 +277,13 @@ pub struct AgentLoopNode { pub instruction: ParamId, pub demos: ParamId, pub model: ParamId, + /// The declared tool table — structural: it is the loop's capability + /// footprint and the legal alphabet of the `tool_set` gene. pub tools: Box<[ToolId]>, + /// `ParamKind::ToolSet`: which declared tools the loop carries at run + /// time (defaults to the full declared table). Selection is optimizable; + /// declaration is not. + pub tool_set: ParamId, /// `ParamKind::ContextPolicy`. pub context_policy: ParamId, pub stop: StopSpec, @@ -550,8 +556,8 @@ impl Program { /// Promotion (RFC 0002 §5, overlay lifecycle): folds `overlay` into a new /// `Program` value — every overlay entry (instruction, demos, tool desc, - /// model ref, context policy, code) becomes the corresponding slot's - /// *default* — stamps [`Lineage`], and recomputes the program hash. + /// tool set, model ref, context policy, code) becomes the corresponding + /// slot's *default* — stamps [`Lineage`], and recomputes the program hash. /// /// `self` is untouched; the returned program is a first-class artifact: /// it validates, serializes, and runs identically to `self` + `overlay` diff --git a/crates/dspy-rs/src/ir/interp.rs b/crates/dspy-rs/src/ir/interp.rs index 9a3a82ce..b9d9e6f9 100644 --- a/crates/dspy-rs/src/ir/interp.rs +++ b/crates/dspy-rs/src/ir/interp.rs @@ -21,9 +21,11 @@ use std::time::Instant; use cranelift_entity::SecondaryMap; use futures::future::BoxFuture; +use indexmap::IndexMap; use serde_json::{Value, json}; use crate::adapter::chat::ChatAdapter; +use crate::core::FieldMeta; use crate::ir::graph::{ AgentLoopNode, Binding, BudgetPolicy, CapSet, HoleImpl, HoleNode, ModelId, Node, NodeId, PortRef, PredictNode, Program, ToolId, ToolKind, @@ -140,6 +142,139 @@ impl BudgetMeter { } } +// --------------------------------------------------------------------------- +// Leaf metadata +// --------------------------------------------------------------------------- + +/// Parse/coercion metadata from one successful `Predict`-leaf evaluation, +/// collected by [`Interpreter::run_collecting`] in execution order. +/// +/// This is what `Predict` reassembles into [`CallMetadata`](crate::CallMetadata): +/// the same [`FieldMeta`] type (jsonish coercion [`Flag`](crate::Flag)s + +/// `#[check]` [`ConstraintResult`](crate::ConstraintResult)s) that +/// `ChatAdapter::parse_output_def` produces, so a typed `Predict` routed +/// through the interpreter loses none of `Predicted`'s metadata contract. +/// +/// Scope and semantics: +/// - `Predict` and `AgentLoop` leaves report; `Hole` leaves make no LM call +/// and parse no delimited text. An `AgentLoop` reports once per evaluation +/// with loop-accumulated usage and its tool call/execution record; +/// `field_meta` is empty when the final output came from stop-tool args +/// (no parse metadata exists for those). +/// - Only *successful* evaluations report — a failed attempt inside `Retry` +/// propagates its error and leaves no outcome; the succeeding attempt +/// reports one. A leaf re-evaluated by `Refine`/`Loop` reports once per +/// evaluation. +/// - A replay-served leaf reports raw text and usage from the recorded span +/// with **empty** `field_meta`, matching the static lane ("served +/// predictions carry no per-field parse metadata"). +#[derive(Debug, Clone)] +pub struct LeafOutcome { + /// The program-unique leaf name (the trace span `component`). + pub name: String, + /// The full text the LM returned, before parsing. + pub raw_response: String, + /// Per-field parse details keyed by canonical field name + /// (`FieldDef::name`) — raw section text, coercion flags, check results. + pub field_meta: IndexMap, + /// Token usage for this leaf's LM call. + pub usage: LmUsage, + /// Stable hash of the redacted model config used — identical to the + /// trace's [`ModelEntry::config_hash`](crate::trace::ModelEntry) for the + /// same config. + pub model_config_hash: u64, + /// The trace span this evaluation recorded, when a capture scope was + /// active — what `Predict` surfaces as `CallMetadata::span_id`. + pub span_id: Option, + /// Tool calls the model requested during an `AgentLoop` evaluation, in + /// execution order (stop-tool calls excluded). Empty for `Predict` leaves. + pub tool_calls: Vec, + /// Results of executed tool calls, aligned with [`tool_calls`](Self::tool_calls). + pub tool_executions: Vec, +} + +/// Program output plus per-leaf metadata, returned by +/// [`Interpreter::run_collecting`]. +#[derive(Debug, Clone)] +pub struct RunOutput { + /// The program's output map — exactly what [`Interpreter::run`] returns. + pub output: JsonMap, + /// One entry per successful `Predict`-leaf evaluation, in execution + /// order. `ForkJoin` branches append in declared branch order. + pub leaves: Vec, +} + +// --------------------------------------------------------------------------- +// Conversation surface (RFC 0004 §1–2) +// --------------------------------------------------------------------------- + +/// One caller-driven conversation turn, returned by +/// [`Interpreter::run_conversation_caller_managed`] and +/// [`Interpreter::resume_conversation`]. +#[derive(Debug)] +pub enum ConversationTurn { + /// The leaf produced its final output; `chat` is the conversation + /// extended with everything the turn generated. + Complete { run: RunOutput, chat: Chat }, + /// The model requested tool calls the caller must execute. Run them and + /// feed the results back through + /// [`Interpreter::resume_conversation`]. + Suspended(ToolSuspension), +} + +/// A caller-managed agent turn suspended on pending tool calls (RFC 0004 §2). +/// +/// Everything the loop needs to continue travels inside: the open trace span, +/// the loop and run budget meters, the accumulated event/usage record, and +/// the turn cursor. Dropping a suspension without resuming closes its span as +/// `Cancelled` — the same contract as cancelling a run mid-flight. +pub struct ToolSuspension { + calls: Vec, + chat: Chat, + state: SuspendState, +} + +impl ToolSuspension { + /// The tool calls the model requested, in request order. Stop-tool calls + /// never appear here — they complete the turn instead of suspending it. + pub fn calls(&self) -> &[rig::message::ToolCall] { + &self.calls + } + + /// The conversation so far, including the assistant tool-call turn. + pub fn chat(&self) -> &Chat { + &self.chat + } +} + +impl std::fmt::Debug for ToolSuspension { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ToolSuspension") + .field("calls", &self.calls) + .field("chat_len", &self.chat.len()) + .finish_non_exhaustive() + } +} + +/// The private half of a [`ToolSuspension`]: resumption state for +/// [`Interpreter::resume_conversation`]. +struct SuspendState { + node: NodeId, + overlay: Option>, + /// The loop's child meter — chained under `run_meter` exactly as in + /// dispatching mode. + meter: Arc, + run_meter: Arc, + guard: Option, + run: AgentRun, + /// The loop turn `resume_conversation` re-enters at. + next_turn: u32, + prefix_len: usize, + /// When the loop suspended — the `ToolRun` events recorded at resume + /// meter the time the calls were outstanding. + suspended_at: Instant, +} + // --------------------------------------------------------------------------- // Errors // --------------------------------------------------------------------------- @@ -175,7 +310,17 @@ pub enum RunError { source: LmError, }, #[error("parse error at `{at}`")] - Parse { at: Box, raw: String }, + Parse { + at: Box, + raw: String, + /// The structured def-lane parse failure, when the leaf's response + /// text was parsed (`Predict`/`AgentLoop` field parsing); `None` for + /// structural coercions with no section parse (stop-tool args, hole + /// outputs). + source: Option>, + /// Usage accumulated by the failing leaf evaluation. + usage: LmUsage, + }, #[error("tool `{tool}` failed at `{at}`: {message}")] Tool { at: Box, @@ -260,7 +405,6 @@ pub struct RuntimeEnv { /// the serialized program. `ToolKind` is also per-tool, while code mode /// collapses a whole loop's tool surface; and leaving the closed enum /// untouched keeps the `.dsrs` text format stable. - #[cfg(feature = "code-mode")] pub code_mode: Option, } @@ -303,7 +447,6 @@ impl RuntimeEnv { /// Enables Code Mode for every `AgentLoop` (see /// [`code_mode`](Self::code_mode)): tools are presented to the model as a /// JS API behind one `run_js` tool executing under `config`. - #[cfg(feature = "code-mode")] pub fn with_code_mode(mut self, config: dsrs_tools::SandboxConfig) -> Self { self.code_mode = Some(config); self @@ -338,7 +481,6 @@ pub struct Interpreter { /// default code gene; overlay code variants register on first use. registered: tokio::sync::Mutex>, /// Code Mode sandbox config (see [`RuntimeEnv::code_mode`]). - #[cfg(feature = "code-mode")] code_mode: Option, } @@ -422,7 +564,6 @@ impl Interpreter { // Code Mode: JS-identifier collisions among a loop's non-stop tool // names are a load-time refusal (nothing lazy, nothing at call time). - #[cfg(feature = "code-mode")] if env.code_mode.is_some() { for (_, node) in program.nodes.iter() { let Node::AgentLoop(agent) = node else { @@ -482,7 +623,6 @@ impl Interpreter { host_holes, sandbox, registered: tokio::sync::Mutex::new(registered), - #[cfg(feature = "code-mode")] code_mode: env.code_mode, }) } @@ -499,7 +639,708 @@ impl Interpreter { overlay: Option>, budget: Budget, ) -> Result { - if let Some(overlay) = &overlay + self.run_inner(input, overlay, budget, false) + .await + .map(|out| out.output) + } + + /// Like [`run`](Interpreter::run), additionally collecting a + /// [`LeafOutcome`] per successful `Predict`-leaf evaluation (raw response, + /// per-field flags and check results, usage, model config hash). See + /// [`LeafOutcome`] for the exact scope and semantics. + pub async fn run_collecting( + &self, + input: JsonMap, + overlay: Option>, + budget: Budget, + ) -> Result { + self.run_inner(input, overlay, budget, true).await + } + + /// Conversation-in/conversation-out evaluation (RFC 0004 §1): one turn + /// with the program's single leaf over a caller-owned [`Chat`]. + /// + /// A turn is not a run: each call evaluates the leaf once against the + /// given conversation and returns the extended conversation alongside the + /// turn's [`RunOutput`] (one [`LeafOutcome`], same semantics as + /// [`run_collecting`](Self::run_collecting)). Call again with the + /// returned chat — plus your follow-up message or a fresh `input` — for + /// the next turn. Each turn records one trace span (`seq` increments per + /// turn), meters against its own `budget`, and replays like any other + /// leaf evaluation: the span keys on the full chat sent, so a recorded + /// conversation serves turn by turn. + /// + /// `chat`/`input` combinations: + /// - **empty chat + `Some(input)`** — opening turn: renders system + + /// demos + the formatted input, exactly like + /// [`conversation_opening`](Self::conversation_opening). + /// - **non-empty chat + `None`** — continuation: the chat is sent as-is + /// (append your follow-up first). + /// - **non-empty chat + `Some(input)`** — typed continuation: the + /// formatted input is appended as the next user message. + /// + /// Tools on an agent leaf dispatch through their bound executors, exactly + /// like [`run`](Self::run). For the suspend-on-tools variant see + /// [`run_conversation_caller_managed`](Self::run_conversation_caller_managed). + /// + /// Only single-leaf programs (what `Predict` compiles to) have a + /// conversation surface; a multi-node graph is refused with + /// [`RunError::Input`]. + pub async fn run_conversation( + &self, + chat: Chat, + input: Option, + overlay: Option>, + budget: Budget, + ) -> Result<(RunOutput, Chat), RunError> { + match self + .conversation_turn(chat, input, overlay, budget, false) + .await? + { + ConversationTurn::Complete { run, chat } => Ok((run, chat)), + ConversationTurn::Suspended(_) => { + unreachable!("dispatching conversations never suspend") + } + } + } + + /// [`run_conversation`](Self::run_conversation) in caller-managed mode + /// (RFC 0004 §2): when the model requests tool calls on an agent leaf, + /// the loop *suspends* instead of dispatching — you get the pending calls + /// plus a resumption token ([`ToolSuspension`]) and feed results back + /// through [`resume_conversation`](Self::resume_conversation). + /// + /// Everything else is identical to dispatching mode: one span per turn + /// with the same `Exchange`/`ToolRun` event stream, the same meters (the + /// node budget chained under the run budget), and the same stop-tool + /// semantics — a stop-tool call completes the turn, it never suspends. + /// Two deliberate differences: Code Mode does not apply (the caller + /// executes tools, so the host's execution strategy can't), and a replay + /// scope never suspends (served turns have every tool effect baked in). + pub async fn run_conversation_caller_managed( + &self, + chat: Chat, + input: Option, + overlay: Option>, + budget: Budget, + ) -> Result { + self.conversation_turn(chat, input, overlay, budget, true) + .await + } + + /// Continues a suspended caller-managed turn with the results of the + /// pending tool calls, aligned with [`ToolSuspension::calls`]. + /// + /// Results land in the conversation and the trace exactly as dispatched + /// executions would: one `ToolRun` event per call (metering the time the + /// suspension was outstanding), one batched tool-result user turn, and + /// the loop continues under the same meters and turn budget. Feed a + /// failed tool's error text as its result to keep dispatching mode's + /// LATM-style conversational repair. + pub async fn resume_conversation( + &self, + suspension: ToolSuspension, + results: Vec, + ) -> Result { + let ToolSuspension { + calls, + mut chat, + state, + } = suspension; + let SuspendState { + node, + overlay, + meter, + run_meter, + guard, + mut run, + next_turn, + prefix_len, + suspended_at, + } = state; + let p = &*self.program; + let at = p + .leaf_name(node) + .expect("conversation leaves are named") + .to_string(); + let Node::AgentLoop(n) = &p.nodes[node] else { + return Err(RunError::Internal { + at: at.into(), + message: "suspension does not reference an agent leaf".to_string(), + }); + }; + if results.len() != calls.len() { + return Err(RunError::Input { + at: at.into(), + message: format!( + "expected {} tool results, got {}", + calls.len(), + results.len() + ), + }); + } + let cx = self.conversation_cx(overlay.clone(), Arc::clone(&run_meter)); + let def = &p.sigs[n.sig]; + let lm = self.p_model(&at, &cx, n.model)?; + let policy = self.p_context(&cx, n.context_policy); + + // Feed the caller's results back: same event, record, and + // conversation shape as dispatched executions. + let duration_us = suspended_at.elapsed().as_micros() as u64; + let mut blocks = Vec::with_capacity(calls.len()); + for (call, result) in calls.iter().zip(results) { + let result = clip_tool_result(result, &policy); + run.events.push(SpanEvent::ToolRun { + id: call.id.clone(), + name: call.function.name.clone(), + args: call.function.arguments.clone(), + result: result.clone(), + duration_us, + error: None, + }); + run.tool_calls.push(call.clone()); + run.tool_executions.push(result.clone()); + blocks.push(tool_result_block(call, result)); + } + chat.push_message(Message::with_content(Role::User, blocks)); + + let surface = self.build_agent_surface(&at, n, &cx, false).await?; + let lc = AgentLoopCx { + at: &at, + n, + def, + lm: &lm, + toolset: &surface.toolset, + by_name: &surface.by_name, + sandbox_code: &surface.sandbox_code, + stop_names: &surface.stop_names, + prefix_len, + meter: &meter, + run_meter: &run_meter, + policy: &policy, + code_mode: None, + }; + let outcome = self.agent_loop(&lc, chat, &mut run, next_turn, true).await; + self.conclude_agent_turn( + &at, + &lm, + guard, + run, + outcome, + AgentTurnCtx { + node, + overlay, + meter: Arc::clone(&meter), + run_meter, + prefix_len, + }, + ) + } + + /// Renders the opening [`Chat`] of a conversation with the program's + /// leaf — system + demos (+ the agent playbook) + the formatted `input` + /// turn — reading instruction/demos through `overlay` exactly like a run + /// would. Inspect or edit the result, then hand it to + /// [`run_conversation`](Self::run_conversation); or skip this and pass + /// `run_conversation` an empty chat plus the input — same rendering. + pub fn conversation_opening( + &self, + input: &JsonMap, + overlay: Option>, + ) -> Result { + self.check_overlay(overlay.as_ref())?; + let node = self.conversation_leaf()?; + let at = self + .program + .leaf_name(node) + .expect("conversation leaves are named") + .to_string(); + self.validate_input(&at, self.leaf_sig(node), input)?; + let cx = self.conversation_cx(overlay, Arc::new(BudgetMeter::new(Budget::unlimited()))); + let (mut messages, suffix) = self.render_leaf_opening(node, input, &cx)?; + messages.extend(suffix); + Ok(Chat::new(messages)) + } + + /// One conversation turn against the program's single leaf: assembles + /// this turn's messages, opens the span, consults replay, and evaluates + /// the leaf (single exchange for `Predict`, the agent loop for + /// `AgentLoop` — suspending on tool calls when `suspend_on_tools`). + async fn conversation_turn( + &self, + mut chat: Chat, + input: Option, + overlay: Option>, + budget: Budget, + suspend_on_tools: bool, + ) -> Result { + self.check_overlay(overlay.as_ref())?; + let node = self.conversation_leaf()?; + let p = &*self.program; + let at = p + .leaf_name(node) + .expect("conversation leaves are named") + .to_string(); + let def = self.leaf_sig(node); + let cx = self.conversation_cx(overlay.clone(), Arc::new(BudgetMeter::new(budget))); + + // This turn's messages, plus the rendered prefix when we own it (an + // opening turn interns system+demos; a caller-owned continuation has + // no known prefix split — the span records the full chat as suffix). + let (prefix, suffix) = if chat.is_empty() { + let input_map = input.as_ref().ok_or_else(|| RunError::Input { + at: at.clone().into(), + message: "an empty conversation needs an opening `input`".to_string(), + })?; + self.validate_input(&at, def, input_map)?; + self.render_leaf_opening(node, input_map, &cx)? + } else { + if let Some(input_map) = input.as_ref() { + self.validate_input(&at, def, input_map)?; + chat.push_message(Message::user(ChatAdapter.format_input_def(def, input_map))); + } + (Vec::new(), chat.messages) + }; + let prefix_len = if prefix.is_empty() { + // Continuation: shield the leading system prompt from history + // truncation, the only thing `prefix_len` is used for downstream. + suffix + .iter() + .take_while(|message| message.role == Role::System) + .count() + } else { + prefix.len() + }; + + let lm = match &p.nodes[node] { + Node::Predict(n) => self.p_model(&at, &cx, n.model)?, + Node::AgentLoop(n) => self.p_model(&at, &cx, n.model)?, + _ => unreachable!("conversation_leaf returns only leaves"), + }; + + let guard = begin_span(SpanRequest { + component: &at, + prefix: (!prefix.is_empty()).then_some(prefix.as_slice()), + suffix: &suffix, + input: input.clone(), + model: &lm.config, + request_hash: None, + }); + + let mut messages = prefix; + messages.extend(suffix); + + // Replay (RFC 0001 §4d/e): a conversation turn keys on the full chat + // it would send — the same preimage the span records (`request_hash` + // streams prefix ++ suffix, so the split is irrelevant). Serving a + // turn extends the chat with the recorded completion; tool effects + // are baked in, so caller-managed turns never suspend under replay. + match crate::trace::replay::intercept(&at, &lm.config, &messages) { + Some(crate::trace::replay::ReplayDirective::Serve(span)) => { + return self.serve_conversation_turn(&at, &lm, *span, messages, guard); + } + Some(crate::trace::replay::ReplayDirective::Refuse(err)) => { + if let Some(guard) = guard { + guard.finish(span_error( + crate::trace::SpanErrorKind::Lm, + err.to_string(), + Vec::new(), + None, + LmUsage::default(), + )); + } + return Err(RunError::Replay { + at: at.into(), + source: err, + }); + } + Some(crate::trace::replay::ReplayDirective::Live) | None => {} + } + + match &p.nodes[node] { + Node::Predict(_) => { + self.predict_conversation_turn(&at, def, &lm, &cx, messages, guard) + .await + } + Node::AgentLoop(n) => { + let policy = self.p_context(&cx, n.context_policy); + let surface = self + .build_agent_surface(&at, n, &cx, !suspend_on_tools) + .await?; + let meter = Arc::new(BudgetMeter::child(&cx.meter, node_budget(&n.budget))); + let lc = AgentLoopCx { + at: &at, + n, + def, + lm: &lm, + toolset: &surface.toolset, + by_name: &surface.by_name, + sandbox_code: &surface.sandbox_code, + stop_names: &surface.stop_names, + prefix_len, + meter: &meter, + run_meter: &cx.meter, + policy: &policy, + code_mode: surface.code_mode.as_ref(), + }; + let mut run = AgentRun::default(); + let outcome = self + .agent_loop(&lc, Chat::new(messages), &mut run, 0, suspend_on_tools) + .await; + self.conclude_agent_turn( + &at, + &lm, + guard, + run, + outcome, + AgentTurnCtx { + node, + overlay, + meter: Arc::clone(&meter), + run_meter: Arc::clone(&cx.meter), + prefix_len, + }, + ) + } + _ => unreachable!("conversation_leaf returns only leaves"), + } + } + + /// A `Predict` leaf's conversation turn: one exchange over the assembled + /// chat, parsed against the leaf signature — `eval_predict` minus the + /// binding resolution and prompt rendering the conversation already owns. + async fn predict_conversation_turn( + &self, + at: &str, + def: &SignatureDef, + lm: &Arc, + cx: &Cx, + messages: Vec, + guard: Option, + ) -> Result { + if cx.meter.try_reserve_call().is_err() { + if let Some(guard) = guard { + guard.finish(span_error( + crate::trace::SpanErrorKind::Lm, + "budget exhausted".to_string(), + Vec::new(), + None, + LmUsage::default(), + )); + } + return Err(RunError::Budget { at: at.into() }); + } + + let response = match lm.call(Chat::new(messages), Vec::new()).await { + Ok(response) => response, + Err(err) => { + if let Some(guard) = guard { + guard.finish(span_error( + crate::trace::SpanErrorKind::Lm, + err.to_string(), + Vec::new(), + None, + LmUsage::default(), + )); + } + return Err(RunError::Lm { + at: at.into(), + source: LmError::Provider { + provider: lm.config.model.clone(), + message: err.to_string(), + source: None, + }, + }); + } + }; + cx.meter.record_usage(&response.usage); + + let raw = response.output.content(); + match ChatAdapter.parse_output_def(def, &self.program.types, &response.output) { + Ok((output, metas)) => { + let leaf = LeafOutcome { + name: at.to_string(), + raw_response: raw.clone(), + field_meta: metas, + usage: response.usage, + model_config_hash: crate::trace::ModelEntry::from_config(&lm.config) + .config_hash, + span_id: guard.as_ref().map(|guard| guard.id()), + tool_calls: Vec::new(), + tool_executions: Vec::new(), + }; + if let Some(guard) = guard { + guard.finish(SpanOutcome { + events: response.events, + raw_output: Some(raw), + output: Some(output.clone()), + usage: response.usage, + error: None, + }); + } + Ok(ConversationTurn::Complete { + run: RunOutput { + output, + leaves: vec![leaf], + }, + chat: response.chat, + }) + } + Err(err) => { + if let Some(guard) = guard { + guard.finish(span_error( + crate::trace::SpanErrorKind::Parse, + err.to_string(), + response.events, + Some(raw.clone()), + response.usage, + )); + } + Err(RunError::Parse { + at: at.into(), + raw, + source: Some(Box::new(err)), + usage: response.usage, + }) + } + } + } + + /// Serves one conversation turn from a recorded span: the chat extends + /// with the recorded completion (exchanges verbatim, tool results + /// batched), and the [`LeafOutcome`] carries the recorded usage with no + /// per-field parse metadata — the same contract as every served leaf. + fn serve_conversation_turn( + &self, + at: &str, + lm: &Arc, + span: crate::trace::Span, + messages: Vec, + guard: Option, + ) -> Result { + let output = span + .output + .clone() + .expect("replay serves only spans with parsed output"); + let mut completion = span.completion_messages(); + if completion.is_empty() + && let Some(raw) = span.raw_output.as_deref().filter(|raw| !raw.is_empty()) + { + completion.push(Message::assistant(raw)); + } + let tool_calls: Vec = completion + .iter() + .flat_map(|message| message.tool_calls().into_iter().cloned()) + .collect(); + let tool_executions: Vec = span + .events + .iter() + .filter_map(|event| match event { + SpanEvent::ToolRun { result, .. } => Some(result.clone()), + _ => None, + }) + .collect(); + let mut chat = Chat::new(messages); + for message in &completion { + chat.push_message(message.clone()); + } + let leaf = LeafOutcome { + name: at.to_string(), + raw_response: span.raw_output.clone().unwrap_or_default(), + field_meta: IndexMap::new(), + usage: span.usage, + model_config_hash: crate::trace::ModelEntry::from_config(&lm.config).config_hash, + span_id: guard.as_ref().map(|guard| guard.id()), + tool_calls, + tool_executions, + }; + if let Some(guard) = guard { + guard.finish(SpanOutcome { + events: span.events.clone(), + raw_output: span.raw_output.clone(), + output: Some(output.clone()), + usage: span.usage, + error: None, + }); + } + Ok(ConversationTurn::Complete { + run: RunOutput { + output, + leaves: vec![leaf], + }, + chat, + }) + } + + /// Closes out an agent-leaf conversation turn: finishes the span and + /// builds the [`ConversationTurn`] — `Complete` on `Done`, a + /// [`ToolSuspension`] carrying the open span and meters on `Suspend`. + /// Shared by [`conversation_turn`](Self::conversation_turn) and + /// [`resume_conversation`](Self::resume_conversation). + fn conclude_agent_turn( + &self, + at: &str, + lm: &Arc, + guard: Option, + run: AgentRun, + outcome: Result, + ctx: AgentTurnCtx, + ) -> Result { + match outcome { + Ok(LoopOutcome::Done { + output, + raw, + field_meta, + chat, + }) => { + let leaf = LeafOutcome { + name: at.to_string(), + raw_response: raw.clone(), + field_meta, + usage: run.usage, + model_config_hash: crate::trace::ModelEntry::from_config(&lm.config) + .config_hash, + span_id: guard.as_ref().map(|guard| guard.id()), + tool_calls: run.tool_calls, + tool_executions: run.tool_executions, + }; + if let Some(guard) = guard { + guard.finish(SpanOutcome { + events: run.events, + raw_output: Some(raw), + output: Some(output.clone()), + usage: run.usage, + error: None, + }); + } + Ok(ConversationTurn::Complete { + run: RunOutput { + output, + leaves: vec![leaf], + }, + chat, + }) + } + Ok(LoopOutcome::Suspend { + calls, + chat, + next_turn, + }) => Ok(ConversationTurn::Suspended(ToolSuspension { + calls, + chat, + state: SuspendState { + node: ctx.node, + overlay: ctx.overlay, + meter: ctx.meter, + run_meter: ctx.run_meter, + guard, + run, + next_turn, + prefix_len: ctx.prefix_len, + suspended_at: Instant::now(), + }, + })), + Err(err) => { + if let Some(guard) = guard { + let kind = match &err { + RunError::Parse { .. } => crate::trace::SpanErrorKind::Parse, + RunError::Tool { .. } => crate::trace::SpanErrorKind::Tool, + _ => crate::trace::SpanErrorKind::Lm, + }; + guard.finish(span_error( + kind, + err.to_string(), + run.events, + None, + run.usage, + )); + } + Err(err) + } + } + } + + /// The conversation leaf: the program's single `Predict` or `AgentLoop` + /// node, unwrapped through single-child `Seq`s. The conversation surface + /// deliberately covers only 1-leaf programs (what `Predict` compiles + /// to) — a multi-node graph has no single conversation to own. + fn conversation_leaf(&self) -> Result { + let mut node = self.program.root; + loop { + match &self.program.nodes[node] { + Node::Predict(_) | Node::AgentLoop(_) => return Ok(node), + Node::Seq(n) if n.body.len() == 1 => node = n.body[0], + _ => { + return Err(RunError::Input { + at: "$".into(), + message: "the conversation surface requires a single-leaf program \ + (one predict or agent node)" + .to_string(), + }); + } + } + } + } + + /// The declared signature of a conversation leaf. + fn leaf_sig(&self, node: NodeId) -> &SignatureDef { + match &self.program.nodes[node] { + Node::Predict(n) => &self.program.sigs[n.sig], + Node::AgentLoop(n) => &self.program.sigs[n.sig], + _ => unreachable!("conversation_leaf returns only leaves"), + } + } + + /// Renders a conversation's opening (prefix, suffix) for the leaf with + /// overlay-resolved instruction/demos (and the agent playbook) — exactly + /// what a map-in evaluation of the same leaf would send. + fn render_leaf_opening( + &self, + node: NodeId, + input: &JsonMap, + cx: &Cx, + ) -> Result<(Vec, Vec), RunError> { + let p = &*self.program; + Ok(match &p.nodes[node] { + Node::Predict(n) => { + let instruction = self.p_text(cx, n.instruction); + let demos = self.p_demos(cx, n.demos); + render_prompt(&p.sigs[n.sig], &p.types, &instruction, &demos, input, None) + } + Node::AgentLoop(n) => { + let instruction = self.p_text(cx, n.instruction); + let demos = self.p_demos(cx, n.demos); + let policy = self.p_context(cx, n.context_policy); + render_prompt( + &p.sigs[n.sig], + &p.types, + &instruction, + &demos, + input, + policy.playbook.as_deref(), + ) + } + _ => unreachable!("conversation_leaf returns only leaves"), + }) + } + + /// A conversation's run state: overlay + run meter, no frames, no + /// collection — conversation turns build their [`LeafOutcome`] directly. + fn conversation_cx(&self, overlay: Option>, meter: Arc) -> Cx { + Cx { + overlay, + meter, + frames: SecondaryMap::new(), + inputs: Vec::new(), + feedback: None, + refine_feedback: None, + leaves: None, + } + } + + /// Overlay/program pairing check, shared by every entry point. + fn check_overlay(&self, overlay: Option<&Arc>) -> Result<(), RunError> { + if let Some(overlay) = overlay && overlay.base != self.program.meta.program_hash { return Err(RunError::Overlay { @@ -507,15 +1348,22 @@ impl Interpreter { got: self.program.meta.program_hash, }); } + Ok(()) + } - // Input surface check against the program's external signature. - let sig = &self.program.sigs[self.program.sig]; + /// Validates an input map against a signature's declared input fields. + fn validate_input( + &self, + at: &str, + sig: &SignatureDef, + input: &JsonMap, + ) -> Result<(), RunError> { for field in sig.inputs.iter() { match input.get(&*field.name) { Some(value) => { if !json_matches_type(value, &field.ty, &self.program.types) { return Err(RunError::Input { - at: "$".into(), + at: at.into(), message: format!( "field `{}` does not match its declared type", field.name @@ -526,12 +1374,26 @@ impl Interpreter { None if field.ty.is_optional() => {} None => { return Err(RunError::Input { - at: "$".into(), + at: at.into(), message: format!("missing input field `{}`", field.name), }); } } } + Ok(()) + } + + async fn run_inner( + &self, + input: JsonMap, + overlay: Option>, + budget: Budget, + collect: bool, + ) -> Result { + self.check_overlay(overlay.as_ref())?; + + // Input surface check against the program's external signature. + self.validate_input("$", &self.program.sigs[self.program.sig], &input)?; let mut cx = Cx { overlay, @@ -540,8 +1402,13 @@ impl Interpreter { inputs: vec![input], feedback: None, refine_feedback: None, + leaves: collect.then(Vec::new), }; - self.eval(self.program.root, &mut cx).await + let output = self.eval(self.program.root, &mut cx).await?; + Ok(RunOutput { + output, + leaves: cx.leaves.unwrap_or_default(), + }) } // -- node dispatch -------------------------------------------------------- @@ -569,13 +1436,13 @@ impl Interpreter { .into_iter() .map(|(branch, mut branch_cx)| async move { self.eval(branch, &mut branch_cx).await?; - Ok::<_, RunError>(branch_cx.frames) + Ok::<_, RunError>((branch_cx.frames, branch_cx.leaves)) }); // First error aborts siblings: try_join_all drops the // remaining futures, whose open spans record Cancelled // via guard-drop (RFC 0001 §3.1). let all_frames = futures::future::try_join_all(futures).await?; - for frames in all_frames { + for (frames, leaves) in all_frames { for (node, out) in frames.iter() { if let Some(out) = out && cx.frames[node].is_none() @@ -583,6 +1450,11 @@ impl Interpreter { cx.frames[node] = Some(out.clone()); } } + // Branch outcomes merge in declared branch order. + if let (Some(collected), Some(branch_leaves)) = (cx.leaves.as_mut(), leaves) + { + collected.extend(branch_leaves); + } } self.resolve_exports(&n.join, cx)? } @@ -754,6 +1626,23 @@ impl Interpreter { .output .clone() .expect("replay serves only spans with parsed output"); + // Parity with the static lane's `serve_recorded_span`: served + // predictions carry no per-field parse metadata — the + // recording stores the parsed output, not the parser's + // field-level bookkeeping. + if let Some(leaves) = cx.leaves.as_mut() { + leaves.push(LeafOutcome { + name: at.clone(), + raw_response: span.raw_output.clone().unwrap_or_default(), + field_meta: IndexMap::new(), + usage: span.usage, + model_config_hash: crate::trace::ModelEntry::from_config(&lm.config) + .config_hash, + span_id: guard.as_ref().map(|guard| guard.id()), + tool_calls: Vec::new(), + tool_executions: Vec::new(), + }); + } if let Some(guard) = guard { guard.finish(SpanOutcome { events: span.events.clone(), @@ -824,7 +1713,20 @@ impl Interpreter { let raw = response.output.content(); match ChatAdapter.parse_output_def(def, &p.types, &response.output) { - Ok((output, _metas)) => { + Ok((output, metas)) => { + if let Some(leaves) = cx.leaves.as_mut() { + leaves.push(LeafOutcome { + name: at.clone(), + raw_response: raw.clone(), + field_meta: metas, + usage: response.usage, + model_config_hash: crate::trace::ModelEntry::from_config(&lm.config) + .config_hash, + span_id: guard.as_ref().map(|guard| guard.id()), + tool_calls: Vec::new(), + tool_executions: Vec::new(), + }); + } if let Some(guard) = guard { guard.finish(SpanOutcome { events: response.events, @@ -846,7 +1748,12 @@ impl Interpreter { response.usage, )); } - Err(RunError::Parse { at: at.into(), raw }) + Err(RunError::Parse { + at: at.into(), + raw, + source: Some(Box::new(err)), + usage: response.usage, + }) } } } @@ -1026,10 +1933,12 @@ impl Interpreter { at: at.into(), message: "host hole not bound (load should have refused)".into(), })?; - let value = host(input.clone()).await.map_err(|message| RunError::Hole { - at: at.into(), - source: dsrs_tools::ExecError::Internal { message }, - })?; + let value = host(input.clone()) + .await + .map_err(|message| RunError::Hole { + at: at.into(), + source: dsrs_tools::ExecError::Internal { message }, + })?; Ok((at.to_string(), value)) } } @@ -1062,41 +1971,7 @@ impl Interpreter { } let lm = self.p_model(&at, cx, n.model)?; - // Tool surface: definitions from declared signatures with - // overlay-resolved descriptions (ToolDesc is a first-class gene), and - // overlay-resolved code for sandboxed tools. - let mut definitions = Vec::with_capacity(n.tools.len()); - let mut by_name: HashMap = HashMap::new(); - let mut sandbox_code: HashMap = HashMap::new(); - for &tool_id in n.tools.iter() { - let tool = &p.tools[tool_id]; - let name = p.syms.get(tool.name).to_string(); - definitions.push(rig::completion::ToolDefinition { - name: name.clone(), - description: self.p_text(cx, tool.desc), - parameters: input_schema_of(&p.sigs[tool.sig], &p.types), - }); - by_name.insert(name, tool_id); - if let ToolKind::Sandboxed { code } = tool.kind - && let ParamValue::Code { source, hash, .. } = self.p_value(cx, code) - { - sandbox_code.insert(tool_id, (source.clone(), *hash)); - } - } - let stop_names: Vec = n - .stop - .stop_tools - .iter() - .map(|&t| p.syms.get(p.tools[t].name).to_string()) - .collect(); - // Code Mode: collapse the non-stop tool surface into one `run_js` - // definition (stop tools stay individual — the loop must see their - // calls by name to end). - #[cfg(feature = "code-mode")] - let code_mode = self - .build_code_mode_surface(&at, n, &mut definitions, &sandbox_code, &stop_names) - .await?; - let toolset = ToolSet::from_definitions(definitions); + let surface = self.build_agent_surface(&at, n, cx, true).await?; let meter = Arc::new(BudgetMeter::child(&cx.meter, node_budget(&n.budget))); @@ -1123,6 +1998,34 @@ impl Interpreter { .output .clone() .expect("replay serves only spans with parsed output"); + if let Some(leaves) = cx.leaves.as_mut() { + // Served evaluations carry no per-field parse metadata; + // the tool record is reconstructed from the recording. + let tool_calls = span + .completion_messages() + .iter() + .flat_map(|message| message.tool_calls().into_iter().cloned()) + .collect(); + let tool_executions = span + .events + .iter() + .filter_map(|event| match event { + SpanEvent::ToolRun { result, .. } => Some(result.clone()), + _ => None, + }) + .collect(); + leaves.push(LeafOutcome { + name: at.clone(), + raw_response: span.raw_output.clone().unwrap_or_default(), + field_meta: IndexMap::new(), + usage: span.usage, + model_config_hash: crate::trace::ModelEntry::from_config(&lm.config) + .config_hash, + span_id: guard.as_ref().map(|guard| guard.id()), + tool_calls, + tool_executions, + }); + } if let Some(guard) = guard { guard.finish(SpanOutcome { events: span.events.clone(), @@ -1157,31 +2060,51 @@ impl Interpreter { n, def, lm: &lm, - toolset: &toolset, - by_name: &by_name, - sandbox_code: &sandbox_code, - stop_names: &stop_names, + toolset: &surface.toolset, + by_name: &surface.by_name, + sandbox_code: &surface.sandbox_code, + stop_names: &surface.stop_names, prefix_len, meter: &meter, run_meter: &cx.meter, policy: &policy, - #[cfg(feature = "code-mode")] - code_mode: code_mode.as_ref(), + code_mode: surface.code_mode.as_ref(), + }; + let mut run = AgentRun::default(); + let outcome = match self.agent_loop(&loop_cx, chat, &mut run, 0, false).await { + Ok(LoopOutcome::Done { + output, + raw, + field_meta, + .. + }) => Ok((output, raw, field_meta)), + Ok(LoopOutcome::Suspend { .. }) => { + unreachable!("dispatching loops never suspend") + } + Err(err) => Err(err), }; - let mut events: Vec = Vec::new(); - let mut usage = LmUsage::default(); - let outcome = self - .agent_loop(&loop_cx, chat, &mut events, &mut usage) - .await; match outcome { - Ok((output, raw)) => { + Ok((output, raw, field_meta)) => { + if let Some(leaves) = cx.leaves.as_mut() { + leaves.push(LeafOutcome { + name: at.clone(), + raw_response: raw.clone(), + field_meta, + usage: run.usage, + model_config_hash: crate::trace::ModelEntry::from_config(&lm.config) + .config_hash, + span_id: guard.as_ref().map(|guard| guard.id()), + tool_calls: run.tool_calls, + tool_executions: run.tool_executions, + }); + } if let Some(guard) = guard { guard.finish(SpanOutcome { - events, + events: run.events, raw_output: Some(raw), output: Some(output.clone()), - usage, + usage: run.usage, error: None, }); } @@ -1194,35 +2117,49 @@ impl Interpreter { RunError::Tool { .. } => crate::trace::SpanErrorKind::Tool, _ => crate::trace::SpanErrorKind::Lm, }; - guard.finish(span_error(kind, err.to_string(), events, None, usage)); + guard.finish(span_error( + kind, + err.to_string(), + run.events, + None, + run.usage, + )); } Err(err) } } } + /// Runs the agent loop from `start_turn`. Dispatching mode + /// (`suspend_on_tools = false`) executes tool calls through their bound + /// executors; suspending mode returns [`LoopOutcome::Suspend`] instead, + /// leaving execution to the caller + /// ([`resume_conversation`](Self::resume_conversation) re-enters here at + /// `next_turn`). Stop-tool and budget semantics are identical in both + /// modes. async fn agent_loop( &self, lc: &AgentLoopCx<'_>, mut chat: Chat, - events: &mut Vec, - usage: &mut LmUsage, - ) -> Result<(JsonMap, String), RunError> { + run: &mut AgentRun, + start_turn: u32, + suspend_on_tools: bool, + ) -> Result { let types = &self.program.types; - for turn in 0..lc.n.stop.max_turns.get() { + for turn in start_turn..lc.n.stop.max_turns.get() { if lc.meter.try_reserve_call().is_err() { return match lc.n.budget.on_exhausted { BudgetPolicy::Fail => Err(RunError::Budget { at: lc.at.into() }), - BudgetPolicy::Finalize => self.finalize(lc, chat, events, usage).await, + BudgetPolicy::Finalize => self.finalize(lc, chat, run).await, }; } chat = truncate_history(chat, lc.prefix_len, lc.policy); let response = lm_call_toolset(lc.lm, chat, lc.toolset, lc.at).await?; lc.meter.record_usage(&response.usage); - *usage = *usage + response.usage; - events.extend(response.events.clone()); + run.usage = run.usage + response.usage; + run.events.extend(response.events.clone()); chat = response.chat; if !response.tool_calls.is_empty() { @@ -1234,14 +2171,27 @@ impl Interpreter { { let args = call.function.arguments.clone(); let output = coerce_outputs(lc.at, lc.def, types, &args)?; - return Ok((output, args.to_string())); + return Ok(LoopOutcome::Done { + output, + raw: args.to_string(), + field_meta: IndexMap::new(), + chat, + }); + } + + if suspend_on_tools { + return Ok(LoopOutcome::Suspend { + calls: response.tool_calls, + chat, + next_turn: turn + 1, + }); } let mut blocks = Vec::with_capacity(response.tool_calls.len()); for call in &response.tool_calls { let started = Instant::now(); let (result, error) = self.execute_agent_tool(lc, call).await; - events.push(SpanEvent::ToolRun { + run.events.push(SpanEvent::ToolRun { id: call.id.clone(), name: call.function.name.clone(), args: call.function.arguments.clone(), @@ -1249,6 +2199,8 @@ impl Interpreter { duration_us: started.elapsed().as_micros() as u64, error, }); + run.tool_calls.push(call.clone()); + run.tool_executions.push(result.clone()); blocks.push(tool_result_block(call, result)); } chat.push_message(Message::with_content(Role::User, blocks)); @@ -1259,7 +2211,14 @@ impl Interpreter { if lc.n.stop.until_parse { let raw = response.output.content(); match ChatAdapter.parse_output_def(lc.def, types, &response.output) { - Ok((output, _)) => return Ok((output, raw)), + Ok((output, metas)) => { + return Ok(LoopOutcome::Done { + output, + raw, + field_meta: metas, + chat, + }); + } Err(err) if turn + 1 < lc.n.stop.max_turns.get() => { chat.push_message(Message::user(format!( "Your response could not be parsed: {err}. Respond again, \ @@ -1267,10 +2226,12 @@ impl Interpreter { `[[ ## field ## ]]` format." ))); } - Err(_) => { + Err(err) => { return Err(RunError::Parse { at: lc.at.into(), raw, + source: Some(Box::new(err)), + usage: run.usage, }); } } @@ -1280,7 +2241,7 @@ impl Interpreter { // Turns exhausted without an accepted answer. match lc.n.budget.on_exhausted { BudgetPolicy::Fail => Err(RunError::Budget { at: lc.at.into() }), - BudgetPolicy::Finalize => self.finalize(lc, chat, events, usage).await, + BudgetPolicy::Finalize => self.finalize(lc, chat, run).await, } } @@ -1290,9 +2251,8 @@ impl Interpreter { &self, lc: &AgentLoopCx<'_>, mut chat: Chat, - events: &mut Vec, - usage: &mut LmUsage, - ) -> Result<(JsonMap, String), RunError> { + run: &mut AgentRun, + ) -> Result { lc.run_meter .try_reserve_call() .map_err(|_| RunError::Budget { at: lc.at.into() })?; @@ -1313,17 +2273,23 @@ impl Interpreter { }, })?; lc.run_meter.record_usage(&response.usage); - *usage = *usage + response.usage; - events.extend(response.events.clone()); + run.usage = run.usage + response.usage; + run.events.extend(response.events.clone()); let raw = response.output.content(); - let output = ChatAdapter + let (output, metas) = ChatAdapter .parse_output_def(lc.def, &self.program.types, &response.output) - .map_err(|_| RunError::Parse { + .map_err(|err| RunError::Parse { at: lc.at.into(), raw: raw.clone(), - })? - .0; - Ok((output, raw)) + source: Some(Box::new(err)), + usage: run.usage, + })?; + Ok(LoopOutcome::Done { + output, + raw, + field_meta: metas, + chat: response.chat, + }) } /// Executes one agent tool call. Failures are conversational: the error @@ -1333,31 +2299,17 @@ impl Interpreter { lc: &AgentLoopCx<'_>, call: &rig::message::ToolCall, ) -> (String, Option) { - #[cfg(feature = "code-mode")] let outcome: Result = match lc.code_mode { Some(surface) if call.function.name == dsrs_tools::RUN_JS_TOOL_NAME => { execute_code_mode_script(surface, &call.function.arguments).await } _ => self.dispatch_agent_tool(lc, call).await, }; - #[cfg(not(feature = "code-mode"))] - let outcome = self.dispatch_agent_tool(lc, call).await; - let (mut text, error) = match outcome { + let (text, error) = match outcome { Ok(text) => (text, None), Err(message) => (message.clone(), Some(message)), }; - if let Some(max) = lc.policy.tool_result_max_bytes { - let max = max as usize; - if text.len() > max { - let mut cut = max; - while cut > 0 && !text.is_char_boundary(cut) { - cut -= 1; - } - text.truncate(cut); - text.push_str("… [truncated]"); - } - } - (text, error) + (clip_tool_result(text, lc.policy), error) } /// Routes one tool call to its bound executor (host `ToolDyn` or the @@ -1416,6 +2368,69 @@ impl Interpreter { .map_err(|err| err.to_llm_json()) } + /// Resolves one agent loop's model-facing tool surface: definitions from + /// the declared signatures with overlay-resolved descriptions (`ToolDesc` + /// is a first-class gene) and overlay-resolved code for sandboxed tools, + /// plus the stop-tool names and (when `allow_code_mode`) the Code Mode + /// collapse. Caller-managed conversations pass `allow_code_mode = false`: + /// the caller executes tools itself, so the host's code-mode execution + /// strategy cannot apply and the declared per-tool surface is presented. + async fn build_agent_surface( + &self, + at: &str, + n: &AgentLoopNode, + cx: &Cx, + allow_code_mode: bool, + ) -> Result { + let p = &*self.program; + // Tool surface: the ToolSet gene selects which *declared* tools the + // loop carries this run (absent overlay entry = the slot default = + // the full declared table). `by_name` is built from the same + // selection, so a deselected tool cannot execute even if the model + // hallucinates its name. + let tools = self.p_tool_set(cx, n.tool_set, &n.tools); + let mut definitions = Vec::with_capacity(tools.len()); + let mut by_name: HashMap = HashMap::new(); + let mut sandbox_code: HashMap = HashMap::new(); + for &tool_id in tools.iter() { + let tool = &p.tools[tool_id]; + let name = p.syms.get(tool.name).to_string(); + definitions.push(rig::completion::ToolDefinition { + name: name.clone(), + description: self.p_text(cx, tool.desc), + parameters: input_schema_of(&p.sigs[tool.sig], &p.types), + }); + by_name.insert(name, tool_id); + if let ToolKind::Sandboxed { code } = tool.kind + && let ParamValue::Code { source, hash, .. } = self.p_value(cx, code) + { + sandbox_code.insert(tool_id, (source.clone(), *hash)); + } + } + let stop_names: Vec = n + .stop + .stop_tools + .iter() + .map(|&t| p.syms.get(p.tools[t].name).to_string()) + .collect(); + // Code Mode: collapse the non-stop tool surface into one `run_js` + // definition (stop tools stay individual — the loop must see their + // calls by name to end). + let code_mode = if allow_code_mode { + self.build_code_mode_surface(at, &tools, &mut definitions, &sandbox_code, &stop_names) + .await? + } else { + None + }; + Ok(AgentSurface { + toolset: ToolSet::from_definitions(definitions), + by_name, + sandbox_code, + stop_names, + code_mode, + }) + } + /// Builds one agent loop's Code Mode surface: wraps every non-stop tool /// as a sandbox [`Capability`](dsrs_tools::Capability) (host tools call /// straight through `ToolDyn`; sandboxed tools route through the @@ -1424,11 +2439,10 @@ impl Interpreter { /// the *overlay-resolved* tool descriptions — `ToolDesc` genes keep /// flowing into the surface the model sees. Returns `None` when code /// mode is off or the loop has no non-stop tools. - #[cfg(feature = "code-mode")] async fn build_code_mode_surface( &self, at: &str, - n: &AgentLoopNode, + tools: &[ToolId], definitions: &mut Vec, sandbox_code: &HashMap, stop_names: &[String], @@ -1444,8 +2458,9 @@ impl Interpreter { let mut kept = Vec::new(); let mut apis = Vec::new(); let mut capabilities = Vec::new(); - // `definitions[i]` was built from `n.tools[i]` — same order. - for (&tool_id, definition) in n.tools.iter().zip(definitions.iter()) { + // `definitions[i]` was built from `tools[i]` (the ToolSet-selected + // surface) — same order. + for (&tool_id, definition) in tools.iter().zip(definitions.iter()) { if stop_names.contains(&definition.name) { kept.push(definition.clone()); continue; @@ -1544,6 +2559,21 @@ impl Interpreter { } } + /// The ToolSet gene's effective selection. Membership in `declared` was + /// enforced when the value entered ([`Overlay::set`] / load-time + /// validation); the filter here only keeps downstream arena indexing + /// total against hostile ids. + fn p_tool_set(&self, cx: &Cx, id: ParamId, declared: &[ToolId]) -> Box<[ToolId]> { + match self.p_value(cx, id) { + ParamValue::ToolSet { tools } => tools + .iter() + .copied() + .filter(|tool| declared.contains(tool)) + .collect(), + other => panic!("tool_set slot resolved to {:?}", other.kind()), + } + } + fn p_context(&self, cx: &Cx, id: ParamId) -> ContextPolicy { match self.p_value(cx, id) { ParamValue::ContextPolicy { policy } => policy.clone(), @@ -1658,6 +2688,56 @@ impl Interpreter { } } +/// Mutable per-invocation loop state: the span's event stream, accumulated +/// usage, and the tool call/execution record surfaced through +/// [`LeafOutcome`]. +#[derive(Default)] +struct AgentRun { + events: Vec, + usage: LmUsage, + tool_calls: Vec, + tool_executions: Vec, +} + +/// How one `agent_loop` invocation ended (short of an error): the accepted +/// final output, or — caller-managed conversations only — a suspension on +/// pending tool calls. +enum LoopOutcome { + Done { + output: JsonMap, + raw: String, + field_meta: IndexMap, + /// The conversation as of the accepting turn — what the conversation + /// surface returns; map-in/map-out evaluation drops it. + chat: Chat, + }, + Suspend { + calls: Vec, + chat: Chat, + next_turn: u32, + }, +} + +/// One agent loop's resolved tool surface (see +/// [`Interpreter::build_agent_surface`]). +struct AgentSurface { + toolset: ToolSet, + by_name: HashMap, + sandbox_code: HashMap, + stop_names: Vec, + code_mode: Option, +} + +/// Owned context threaded into [`Interpreter::conclude_agent_turn`] — what a +/// [`ToolSuspension`] needs to carry when the loop suspends. +struct AgentTurnCtx { + node: NodeId, + overlay: Option>, + meter: Arc, + run_meter: Arc, + prefix_len: usize, +} + /// Borrowed context for one agent-loop invocation — keeps `agent_loop` and /// its helpers at a sane arity. struct AgentLoopCx<'a> { @@ -1673,13 +2753,11 @@ struct AgentLoopCx<'a> { meter: &'a Arc, run_meter: &'a Arc, policy: &'a ContextPolicy, - #[cfg(feature = "code-mode")] code_mode: Option<&'a CodeModeSurface>, } /// One agent loop's Code Mode surface: the wrapped tool capabilities and the /// sandbox config `run_js` scripts execute under. -#[cfg(feature = "code-mode")] struct CodeModeSurface { capabilities: Vec, config: dsrs_tools::SandboxConfig, @@ -1687,7 +2765,6 @@ struct CodeModeSurface { /// Executes one `run_js` call against the loop's Code Mode surface. Failures /// are conversational, like every agent tool failure. -#[cfg(feature = "code-mode")] async fn execute_code_mode_script( surface: &CodeModeSurface, args: &Value, @@ -1783,6 +2860,8 @@ fn coerce_outputs( return Err(RunError::Parse { at: at.into(), raw: value.to_string(), + source: None, + usage: LmUsage::default(), }); }; let mut output = JsonMap::new(); @@ -1798,6 +2877,8 @@ fn coerce_outputs( return Err(RunError::Parse { at: at.into(), raw: value.to_string(), + source: None, + usage: LmUsage::default(), }); } } @@ -1903,6 +2984,24 @@ fn truncate_history(chat: Chat, prefix_len: usize, policy: &ContextPolicy) -> Ch Chat::new(kept) } +/// Applies `ContextPolicy.tool_result_max_bytes` to one tool result — the +/// same clip whether the tool was dispatched or the result came back through +/// [`Interpreter::resume_conversation`]. +fn clip_tool_result(mut text: String, policy: &ContextPolicy) -> String { + if let Some(max) = policy.tool_result_max_bytes { + let max = max as usize; + if text.len() > max { + let mut cut = max; + while cut > 0 && !text.is_char_boundary(cut) { + cut -= 1; + } + text.truncate(cut); + text.push_str("… [truncated]"); + } + } + text +} + /// Builds a tool-result content block for the conversation. fn tool_result_block(call: &rig::message::ToolCall, result: String) -> crate::ContentBlock { use rig::OneOrMany; @@ -2009,11 +3108,16 @@ struct Cx { feedback: Option, /// Refine feedback injection: (child leaf, input field, value). refine_feedback: Option<(NodeId, String, Value)>, + /// `Some` when the caller asked for per-leaf metadata + /// ([`Interpreter::run_collecting`]); successful `Predict` leaves push + /// here in execution order. `None` = plain `run`, zero collection cost. + leaves: Option>, } impl Cx { /// A branch context for `ForkJoin`: shared overlay/meter, snapshotted - /// frames and inputs, branch-local feedback. + /// frames and inputs, branch-local feedback. Branch-local leaf collection + /// when the parent collects; merged back in branch order at the join. fn branch(&self) -> Cx { Cx { overlay: self.overlay.clone(), @@ -2022,6 +3126,7 @@ impl Cx { inputs: self.inputs.clone(), feedback: None, refine_feedback: None, + leaves: self.leaves.as_ref().map(|_| Vec::new()), } } } diff --git a/crates/dspy-rs/src/ir/mod.rs b/crates/dspy-rs/src/ir/mod.rs index e7ed904d..a4b4860c 100644 --- a/crates/dspy-rs/src/ir/mod.rs +++ b/crates/dspy-rs/src/ir/mod.rs @@ -6,15 +6,15 @@ //! [`SignatureDef::of`], serde-derivable for the program artifact. The type //! model is [`typesys::FieldType`](crate::typesys::FieldType) unchanged; //! class/enum definitions live in a program-owned [`TypeTable`]. -//! - **IR-2, the graph** (`ir` feature, default-on) — [`Program`]: entity +//! - **IR-2, the graph** — [`Program`]: entity //! arenas over the closed [`Node`] enum, field-level [`Binding`] dataflow, //! [`ParamSlot`] parameters addressed by `ParamPath`, [`Overlay`] //! candidates, capability ceilings, and load-time validation. -//! - **IR-3, the interpreter** (`ir` feature) — [`Interpreter`]: async +//! - **IR-3, the interpreter** — [`Interpreter`]: async //! evaluation of a loaded program with overlay read-through at render time, //! RFC 0001 trace spans (component = leaf name), budget metering, and //! sandboxed [`Hole`](Node::Hole) execution via `dsrs-tools`. -//! - **IR-5, the `.dsrs` text format** (`ir` feature) — the wire form of a +//! - **IR-5, the `.dsrs` text format** — the wire form of a //! program: [`Program::from_dsrs`]/[`Program::to_dsrs`] parse and //! canonically print RFC 0002 §4 text, and the canonical text (minus //! lineage) is the [`Program::compute_hash`] preimage. A canonical JSON @@ -28,59 +28,42 @@ pub use sig::{ ConstraintDef, FieldDef, RenderSpec, SigError, SigMismatch, SignatureBuilder, SignatureDef, }; -#[cfg(feature = "ir")] pub mod bridge; -#[cfg(feature = "ir")] pub mod builder; -#[cfg(feature = "ir")] +pub mod edit; pub mod graph; -#[cfg(feature = "ir")] pub mod interp; -#[cfg(feature = "ir")] pub mod module_build; -#[cfg(feature = "ir")] pub mod params; -#[cfg(feature = "ir")] pub mod step; -#[cfg(feature = "ir")] pub mod text; -#[cfg(feature = "ir")] pub mod validate; -#[cfg(feature = "ir")] pub use bridge::{current_overlay, with_ambient_overlay, with_overlay}; -#[cfg(feature = "ir")] -pub use module_build::{ - ModuleBuildError, ModuleSpec, ModuleStep, ModuleStepKind, PortSpec, build_module_program, - default_lm, unbound_model_config, -}; -#[cfg(feature = "ir")] -pub use step::{AgentStepOpts, HoleReport, StepDef, StepKind, ToolStepDef}; -#[cfg(feature = "ir")] pub use builder::{ AsNodeName, BuildError, NodeSpec, Port, ProgramBuilder, agent, carried, cot, extern_hole, fork, hole, input, lit, loop_, out, predict, refine, retry, route, seq, }; -#[cfg(feature = "ir")] +pub use edit::{ApplyError, Edit, EditError, EditKind, SwapTarget, migrate_overlay}; pub use graph::{ AgentLoopNode, BakeError, Binding, BudgetPolicy, CapSet, ForkJoinNode, HoleImpl, HoleNode, - Interner, - Lineage, LoopNode, ModelDef, ModelId, Node, NodeBudget, NodeId, PortRef, PredictNode, Program, - ProgramMeta, RefineNode, RetryNode, RouteNode, SeqNode, SigId, StopSpec, Sym, ToolDef, ToolId, - ToolKind, + Interner, Lineage, LoopNode, ModelDef, ModelId, Node, NodeBudget, NodeId, PortRef, PredictNode, + Program, ProgramMeta, RefineNode, RetryNode, RouteNode, SeqNode, SigId, StopSpec, Sym, ToolDef, + ToolId, ToolKind, }; -#[cfg(feature = "ir")] pub use interp::{ - Budget, BudgetMeter, Exhausted, HostHoleFn, Interpreter, LoadError, RunError, RuntimeEnv, - input_schema_of, + Budget, BudgetMeter, ConversationTurn, Exhausted, HostHoleFn, Interpreter, LeafOutcome, + LoadError, RunError, RunOutput, RuntimeEnv, ToolSuspension, input_schema_of, +}; +pub use module_build::{ + ModuleBuildError, ModuleSpec, ModuleStep, ModuleStepKind, PortSpec, build_module_program, + default_lm, unbound_model_config, }; -#[cfg(feature = "ir")] pub use params::{ CodeK, CodeLang, ContextK, ContextPolicy, DemoRow, Demos, Instruction, KindTag, ModelRefK, Overlay, OverlayError, ParamId, ParamKind, ParamOwner, ParamSlot, ParamValue, Slot, ToolDesc, - code_hash, + ToolSetK, code_hash, }; -#[cfg(feature = "ir")] +pub use step::{AgentStepOpts, HoleReport, StepDef, StepKind, ToolStepDef}; pub use text::{DsrsFileError, ParseError}; -#[cfg(feature = "ir")] pub use validate::ValidateError; diff --git a/crates/dspy-rs/src/ir/params.rs b/crates/dspy-rs/src/ir/params.rs index 3fd468d6..872b41ce 100644 --- a/crates/dspy-rs/src/ir/params.rs +++ b/crates/dspy-rs/src/ir/params.rs @@ -3,7 +3,7 @@ //! mutation. //! //! Canonical `ParamPath`s: `".instruction"`, `".demos"`, -//! `".model"`, `".context"`, `".code"`, +//! `".model"`, `".context"`, `".tool_set"`, `".code"`, //! `"tool..desc"`, `"tool..code"`. Paths are the serde boundary; //! everything after load speaks [`ParamId`]s. @@ -13,7 +13,7 @@ use std::marker::PhantomData; use cranelift_entity::{SecondaryMap, entity_impl}; use serde::{Deserialize, Serialize}; -use crate::ir::graph::{ModelId, NodeId, Program, ToolId}; +use crate::ir::graph::{ModelId, Node, NodeId, Program, ToolId}; use crate::trace::JsonMap; #[derive(Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] @@ -47,6 +47,10 @@ pub enum ParamKind { Instruction, Demos, ToolDesc, + /// Tool membership as a gene (RFC 0004 §5): which of an agent node's + /// *declared* tools the loop carries. Declaration stays structural + /// (`AgentLoopNode::tools`); selection is this slot. + ToolSet, ModelRef, ContextPolicy, Code, @@ -58,6 +62,7 @@ impl ParamKind { Self::Instruction => "instruction", Self::Demos => "demos", Self::ToolDesc => "tool_desc", + Self::ToolSet => "tool_set", Self::ModelRef => "model_ref", Self::ContextPolicy => "context_policy", Self::Code => "code", @@ -77,6 +82,13 @@ pub enum ParamValue { ToolDesc { text: String, }, + /// A subset of the owning agent node's declared tools, in surface order. + /// The legal alphabet is the declared table — [`Overlay::set`] and + /// load-time validation both refuse a value naming an undeclared tool, so + /// any accepted subset is capability-safe by construction. + ToolSet { + tools: Vec, + }, ModelRef { model: ModelId, }, @@ -104,6 +116,7 @@ impl ParamValue { Self::Instruction { .. } => ParamKind::Instruction, Self::Demos { .. } => ParamKind::Demos, Self::ToolDesc { .. } => ParamKind::ToolDesc, + Self::ToolSet { .. } => ParamKind::ToolSet, Self::ModelRef { .. } => ParamKind::ModelRef, Self::ContextPolicy { .. } => ParamKind::ContextPolicy, Self::Code { .. } => ParamKind::Code, @@ -173,6 +186,7 @@ impl Copy for Slot {} pub enum Instruction {} pub enum Demos {} pub enum ToolDesc {} +pub enum ToolSetK {} pub enum ModelRefK {} pub enum ContextK {} pub enum CodeK {} @@ -189,6 +203,9 @@ impl KindTag for Demos { impl KindTag for ToolDesc { const KIND: ParamKind = ParamKind::ToolDesc; } +impl KindTag for ToolSetK { + const KIND: ParamKind = ParamKind::ToolSet; +} impl KindTag for ModelRefK { const KIND: ParamKind = ParamKind::ModelRef; } @@ -215,6 +232,11 @@ pub enum OverlayError { }, #[error("unknown param path `{path}`")] UnknownPath { path: String }, + /// A `ToolSet` value names a tool outside the owning agent node's + /// declared table — the overlay would smuggle in a capability the + /// program's grants were never checked against. + #[error("tool set for `{path}` includes `{tool}`, which the owning agent does not declare")] + ToolSetUndeclared { path: String, tool: String }, /// A flat demo row (fx/`ModuleState` form) carries a field the owning /// leaf's signature does not declare, so it cannot be split into a /// [`DemoRow`]'s input/output maps. @@ -258,6 +280,30 @@ impl Overlay { got: v.kind(), }); } + // The ToolSet alphabet is the owning agent node's declared table + // (RFC 0004 §5): a candidate can drop declared tools, never add + // undeclared ones — subsets are capability-safe by construction. + if let ParamValue::ToolSet { tools } = &v { + let declared: &[ToolId] = match slot.owner { + ParamOwner::Node(node) => match &p.nodes[node] { + Node::AgentLoop(n) => &n.tools, + _ => &[], + }, + ParamOwner::Tool(_) => &[], + }; + for tool in tools { + if !declared.contains(tool) { + return Err(OverlayError::ToolSetUndeclared { + path: slot.path.to_string(), + tool: p + .tools + .get(*tool) + .map(|def| p.syms.get(def.name).to_string()) + .unwrap_or_else(|| format!("{tool}")), + }); + } + } + } self.values[id] = Some(v); Ok(()) } @@ -274,6 +320,19 @@ impl Overlay { self.values[s.id] = Some(ParamValue::code(CodeLang::Js, source)); } + /// Sets a tool-set gene. Unlike the other typed setters this one takes + /// the program and can fail: membership in the owning agent node's + /// declared tool table cannot be checked by the type system, so the value + /// goes through the same guard as [`Overlay::set`]. + pub fn set_tool_set( + &mut self, + p: &Program, + s: Slot, + tools: Vec, + ) -> Result<(), OverlayError> { + self.set(p, s.id, ParamValue::ToolSet { tools }) + } + pub fn get(&self, id: ParamId) -> Option<&ParamValue> { self.values[id].as_ref() } diff --git a/crates/dspy-rs/src/ir/text/mod.rs b/crates/dspy-rs/src/ir/text/mod.rs index 916e77d4..d58ce850 100644 --- a/crates/dspy-rs/src/ir/text/mod.rs +++ b/crates/dspy-rs/src/ir/text/mod.rs @@ -16,31 +16,15 @@ use std::path::Path; use crate::ir::graph::Program; -pub(crate) mod lex; pub(crate) mod parse; pub(crate) mod print; /// A parse failure with the source position and what was expected — designed /// to be actionable feedback for a model regenerating the program. -#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] -#[error("line {line}, column {col}: {message}")] -pub struct ParseError { - /// 1-based source line. - pub line: u32, - /// 1-based source column (bytes). - pub col: u32, - pub message: String, -} - -impl ParseError { - pub(crate) fn at(span: lex::Span, message: impl Into) -> Self { - Self { - line: span.line, - col: span.col, - message: message.into(), - } - } -} +/// +/// Re-exported from `dsrs-syntax`, the shared home of the `.dsrs` lexer and +/// structural grammar (also used by `include_program!` at macro expansion). +pub use dsrs_syntax::ParseError; /// Failure loading or saving a `.dsrs` artifact file. #[derive(Debug, thiserror::Error)] diff --git a/crates/dspy-rs/src/ir/text/parse.rs b/crates/dspy-rs/src/ir/text/parse.rs index faf09720..7e16b51a 100644 --- a/crates/dspy-rs/src/ir/text/parse.rs +++ b/crates/dspy-rs/src/ir/text/parse.rs @@ -21,7 +21,7 @@ use crate::ir::validate::ValidateError; use crate::typesys::{ClassDef, EnumDef, EnumValueDef, FieldType, TypeTable}; use super::ParseError; -use super::lex::{Lexed, Lexer, Span, Tok}; +use dsrs_syntax::lex::{Lexed, Lexer, Span, Tok}; /// Words that cannot be used as node/sig/tool/model/class/enum names. const RESERVED: &[&str] = &[ @@ -1208,6 +1208,7 @@ impl<'a> Parser<'a> { let (key, key_span) = self.expect_ident("as an agent option")?; match key.as_str() { "tools" => spec = spec.tools(self.tool_list("tools")?), + "tool_set" => spec = spec.tool_set(self.tool_list("tool_set")?), "stop_tools" => spec = spec.stop_tools(self.tool_list("stop_tools")?), "max_turns" => { let (turns, span) = self.expect_int::("after `max_turns`")?; @@ -1303,9 +1304,9 @@ impl<'a> Parser<'a> { return Err(ParseError::at( key_span, format!( - "unknown agent option `{other}`: expected `tools`, `stop_tools`, \ - `max_turns`, `until_parse`, `budget`, `context`, `instruction`, \ - or `demos`" + "unknown agent option `{other}`: expected `tools`, `tool_set`, \ + `stop_tools`, `max_turns`, `until_parse`, `budget`, `context`, \ + `instruction`, or `demos`" ), )); } @@ -1885,6 +1886,8 @@ fn validate_error_handles(v: &ValidateError) -> (Option<&str>, Option<(&str, &st | E::ParamKindMismatch { at, .. } | E::ParamOwnerMismatch { at, .. } | E::StopToolNotDeclared { at } + | E::ToolSetUndeclared { at, .. } + | E::ToolSetDuplicate { at, .. } | E::RefineJudgeNotLeaf { at } | E::RefineJudgeInterface { at } | E::WhileNotBool { at, .. } => (Some(at), None), diff --git a/crates/dspy-rs/src/ir/text/print.rs b/crates/dspy-rs/src/ir/text/print.rs index 4815a865..2032bf86 100644 --- a/crates/dspy-rs/src/ir/text/print.rs +++ b/crates/dspy-rs/src/ir/text/print.rs @@ -592,6 +592,19 @@ impl<'p> Printer<'p> { self.indent(level + 1); let _ = writeln!(self.out, "tools [{}]", names.join(" ")); } + // The tool_set gene prints only when its default differs from the + // full declared list — pre-ToolSet programs keep their canonical + // text (and hash) byte-for-byte. + if let ParamValue::ToolSet { tools } = &self.p.params[n.tool_set].default + && tools.as_slice() != &*n.tools + { + let names: Vec<&str> = tools + .iter() + .map(|t| self.p.syms.get(self.p.tools[*t].name)) + .collect(); + self.indent(level + 1); + let _ = writeln!(self.out, "tool_set [{}]", names.join(" ")); + } if !n.stop.stop_tools.is_empty() { let names: Vec<&str> = n .stop diff --git a/crates/dspy-rs/src/ir/validate.rs b/crates/dspy-rs/src/ir/validate.rs index b95bf53c..cbf435ba 100644 --- a/crates/dspy-rs/src/ir/validate.rs +++ b/crates/dspy-rs/src/ir/validate.rs @@ -104,6 +104,12 @@ pub enum ValidateError { }, #[error("agent loop `{at}`: stop tool is not in the node's tool list")] StopToolNotDeclared { at: String }, + #[error( + "agent loop `{at}`: tool_set includes `{tool}`, which is not in the node's declared tools" + )] + ToolSetUndeclared { at: String, tool: String }, + #[error("agent loop `{at}`: tool_set lists `{tool}` more than once")] + ToolSetDuplicate { at: String, tool: String }, #[error("refine at {at}: judge must be a Predict or Hole leaf")] RefineJudgeNotLeaf { at: String }, #[error("refine at {at}: judge outputs must include score: float and feedback: string")] @@ -291,6 +297,7 @@ impl<'p> Validator<'p> { param_ok(&at, n.instruction)?; param_ok(&at, n.demos)?; param_ok(&at, n.model)?; + param_ok(&at, n.tool_set)?; param_ok(&at, n.context_policy)?; for t in n.tools.iter().chain(n.stop.stop_tools.iter()) { if t.index() >= n_tools { @@ -356,10 +363,18 @@ impl<'p> Validator<'p> { } } } - if let crate::ir::params::ParamValue::ModelRef { model } = &slot.default - && model.index() >= n_models - { - return Err(err(&at, format!("{model}"))); + match &slot.default { + crate::ir::params::ParamValue::ModelRef { model } if model.index() >= n_models => { + return Err(err(&at, format!("{model}"))); + } + crate::ir::params::ParamValue::ToolSet { tools } => { + for t in tools { + if t.index() >= n_tools { + return Err(err(&at, format!("{t}"))); + } + } + } + _ => {} } } @@ -492,6 +507,7 @@ impl<'p> Validator<'p> { )?; self.check_param_ref(&at, n.demos, ParamKind::Demos, ParamOwner::Node(id))?; self.check_param_ref(&at, n.model, ParamKind::ModelRef, ParamOwner::Node(id))?; + self.check_param_ref(&at, n.tool_set, ParamKind::ToolSet, ParamOwner::Node(id))?; self.check_param_ref( &at, n.context_policy, @@ -503,6 +519,29 @@ impl<'p> Validator<'p> { return Err(ValidateError::StopToolNotDeclared { at: at.clone() }); } } + // The tool_set gene's alphabet is the declared table (RFC + // 0004 §5): any duplicate-free subset validates, anything + // else is refused at load — never lazily at run time. + if let crate::ir::params::ParamValue::ToolSet { tools } = + &self.p.params[n.tool_set].default + { + let mut seen: HashSet = HashSet::new(); + for t in tools { + let tool = self.p.syms.get(self.p.tools[*t].name).to_string(); + if !n.tools.contains(t) { + return Err(ValidateError::ToolSetUndeclared { + at: at.clone(), + tool, + }); + } + if !seen.insert(*t) { + return Err(ValidateError::ToolSetDuplicate { + at: at.clone(), + tool, + }); + } + } + } self.check_leaf_bindings(&at, n.sig, &n.binding, scope)?; sig_outputs(&self.p.sigs[n.sig]) } diff --git a/crates/dspy-rs/src/lib.rs b/crates/dspy-rs/src/lib.rs index e0f39330..51b5b409 100644 --- a/crates/dspy-rs/src/lib.rs +++ b/crates/dspy-rs/src/lib.rs @@ -19,13 +19,17 @@ //! //! A [`Predict`] is the leaf — the only thing that actually calls the LM. Every other //! module ([`ChainOfThought`], custom pipelines) delegates to one or more `Predict` leaves. -//! Optimizers discover these leaves automatically via Facet reflection and mutate their -//! instructions and few-shot demos. +//! Modules name their leaves explicitly via [`Predictors`] (see the `predictors!` macro); +//! optimizers tune those leaves' instructions and few-shot demos by injecting candidates +//! ambiently per call, installing only the winner. //! //! # Quick start //! +//! The recommended import is [`prelude`] — the curated core surface. (The +//! crate root also glob re-exports everything for backwards compatibility.) +//! //! ```no_run -//! use dspy_rs::*; +//! use dspy_rs::prelude::*; //! //! #[derive(Signature, Clone, Debug)] //! /// Answer questions accurately and concisely. @@ -61,19 +65,20 @@ //! //! # What doesn't work (yet) //! -//! - **No dynamic graph / structural optimization.** The type-erased `ProgramGraph`, -//! `DynModule`, `StrategyFactory` layer was prototyped and intentionally removed. -//! Everything here is statically typed, which is both the strength and the constraint. -//! - **No `ReAct`, `BestOfN`, `Refine`, or other advanced modules** beyond `ChainOfThought`. -//! The module trait and augmentation system are designed for them, but nobody's built -//! them yet. +//! - **Structural optimization is program-lane only.** [`Structural`] +//! proposes graph edits ([`ir::Edit`]) over an interpreter-loaded +//! [`ir::Program`]; typed modules have no editable skeleton, so the other +//! optimizers tune their instructions and demos only. +//! - **No `BestOfN`, `Refine`, or other advanced modules** beyond +//! [`ChainOfThought`]. Agentic tool loops live in the IR instead +//! (`AgentLoopNode` via the `#[agent]` macro); the module trait and +//! augmentation system could host the rest, but nobody's built them. //! - **`CallMetadata` is not extensible.** Modules can't attach custom metadata (e.g. //! "which attempt won in BestOfN"). This should probably be a trait with associated //! types, but it isn't. -//! - **Container traversal is partial.** The optimizer walker handles `Option`, `Vec`, -//! `HashMap`, and `Box`. `Rc`/`Arc` containing `Predict` leaves return -//! explicit container errors (not silent skips), and `Predict` discovery requires -//! a valid shape-local accessor payload (`TODO(dsrs-shared-ptr-policy)`). +//! - **Leaf discovery is explicit.** Optimizable [`Predict`] leaves are whatever a +//! module declares in its [`Predictors`] impl — there is no reflection walker. +//! A leaf you forget to declare simply isn't optimized or persisted. //! //! # Crate organization //! @@ -84,14 +89,11 @@ //! - [`modules`] — [`ChainOfThought`] and augmentation types //! - [`evaluate`] — [`TypedMetric`] trait, [`evaluate_trainset`], scoring utilities //! - [`optimizer`] — [`Optimizer`] trait, [`COPRO`], [`GEPA`], [`MIPROv2`] +//! - [`ir`] — dynamic program graph, interpreter, and the `.dsrs` text format //! - [`data`] — [`DataLoader`] for JSON/CSV/Parquet/HuggingFace datasets //! - [`trace`] — Execution trace capture (spans per `Predict` call, JSONL serialization) //! - [`utils`] — Response caching -// TODO(dsrs-facet-lint-scope): remove this crate-level allow once Facet's generated -// extension-attr dispatch no longer triggers rust-lang/rust#52234 on in-crate usage. -#![allow(macro_expanded_macro_exports_accessed_by_absolute_paths)] - extern crate self as dspy_rs; pub mod adapter; @@ -118,9 +120,8 @@ pub use optimizer::*; pub use predictors::*; // The unified trace format (RFC 0001). pub use trace::{ - CompId, Eval, JsonMap, ModelEntry, ModelId, OtelEvent, OtelKeyValue, OtelSpan, OtelStatus, - OtelValue, PrefixEntry, PrefixId, ReplayError, ReplayMode, ReplayReport, RlRollout, - RlTransition, Span, SpanError, SpanErrorKind, SpanEvent, SpanGuard, SpanId, SpanOutcome, + CompId, Eval, JsonMap, ModelEntry, ModelId, PrefixEntry, PrefixId, ReplayError, ReplayMode, + ReplayReport, Span, SpanError, SpanErrorKind, SpanEvent, SpanGuard, SpanId, SpanOutcome, SpanRequest, Trace, TraceMeta, TraceOutcome, begin_span, capture, capture_with_meta, is_capturing, is_replaying, replay, }; @@ -128,16 +129,70 @@ pub use utils::*; // Code Mode (vision report §5.5): tools as a sandboxed JS API. See // `ToolSet::code_mode` for the module lane and `RuntimeEnv` for the IR lane. -#[cfg(feature = "code-mode")] pub use dsrs_tools::{Capability, CodeModeTool, RUN_JS_TOOL_NAME, SandboxConfig}; pub mod typesys; pub use dsrs_macros::*; pub use facet::{Facet, Shape}; -pub use typesys::{ - Constraint, ConstraintLevel, ConstraintOutcome, FieldType, Flag, OutputSchema, ResponseCheck, - Schema, evaluate_constraints, -}; +pub use typesys::{Constraint, ConstraintLevel, FieldType, Flag, OutputSchema, Schema}; + +/// The curated core surface — the recommended import for DSRs programs. +/// +/// ```no_run +/// use dspy_rs::prelude::*; +/// ``` +/// +/// Covers the basic path end to end: declare a [`Signature`], pick a module +/// ([`Predict`], [`ChainOfThought`]), [`configure`] an [`LM`], load data, +/// evaluate with a [`TypedMetric`], optimize with any [`Optimizer`], and +/// capture/replay traces. The crate root's glob re-exports remain for the +/// long tail (adapters, schema internals, fx, the full [`ir`] surface). +pub mod prelude { + // Signatures: the trait and the derive share a name across namespaces. + pub use crate::core::signature::Signature; + pub use dsrs_macros::{Example, Signature}; + + // Modules and predictors. + pub use crate::core::{ + CallMetadata, Module, ModuleState, PredictError, PredictState, Predicted, Predictors, + }; + pub use crate::modules::{ChainOfThought, WithReasoning}; + pub use crate::predictors::{Demo, Predict}; + // The `predictors!` leaf-declaration macro (and, same name, the module). + pub use crate::predictors; + + // LM configuration. + pub use crate::core::lm::{LM, LMConfig}; + pub use crate::core::settings::configure; + + // Data loading. + pub use crate::data::dataloader::{DataLoader, TypedLoadOptions}; + + // Evaluation. + pub use crate::evaluate::{ + TypedMetric, average_score, evaluate_trainset, evaluate_trainset_with_concurrency, + }; + // `SpanId` keys per-span evals (`TypedMetric::evaluate_spans`). + pub use crate::trace::{Eval, SpanId}; + + // Optimization: the trait, the six strategies, and the engine surface. + pub use crate::optimizer::{ + BootstrapFewShot, COPRO, Candidate, Engine, GEPA, MIPROv2, OptimizeTarget, Optimizer, + SIMBA, Structural, + }; + + // Trace capture and replay entry points. + pub use crate::trace::{ + ReplayMode, ReplayReport, Trace, capture, capture_with_meta, is_capturing, is_replaying, + replay, + }; + + // The IR: programs as data. + pub use crate::ir::{Edit, Interpreter, Overlay, Program}; + + // Telemetry sugar every quickstart uses. + pub use crate::utils::init_tracing; +} /// Pre-built signature for use in doc examples. Not part of the public API. #[doc(hidden)] @@ -163,41 +218,4 @@ pub mod __macro_support { pub use tokio; } -#[macro_export] -macro_rules! sign { - // Example Usage: signature! { - // question: String, random: bool -> answer: String - // } - // - // Example Output: - // - // #[derive(Signature, Clone)] - // struct InlineSignature { - // #[input] - // question: String, - // #[input] - // random: bool, - // #[output] - // answer: String, - // } - // - // Predict::::new() - - // Pattern: input fields -> output fields - { ($($input_name:ident : $input_type:ty),* $(,)?) -> $($output_name:ident : $output_type:ty),* $(,)? } => {{ - #[derive($crate::Signature, Clone)] - struct __InlineSignature { - $( - #[input] - $input_name: $input_type, - )* - $( - #[output] - $output_name: $output_type, - )* - } - - $crate::Predict::<__InlineSignature>::new() - }}; -} diff --git a/crates/dspy-rs/src/modules/mod.rs b/crates/dspy-rs/src/modules/mod.rs index bb78415a..138b58d0 100644 --- a/crates/dspy-rs/src/modules/mod.rs +++ b/crates/dspy-rs/src/modules/mod.rs @@ -1,5 +1,3 @@ pub mod chain_of_thought; -pub mod react; pub use chain_of_thought::{ChainOfThought, ChainOfThoughtOutput, Reasoning, WithReasoning}; -pub use react::ReAct; diff --git a/crates/dspy-rs/src/modules/react.rs b/crates/dspy-rs/src/modules/react.rs deleted file mode 100644 index d03bcfa1..00000000 --- a/crates/dspy-rs/src/modules/react.rs +++ /dev/null @@ -1,368 +0,0 @@ -use std::future::Future; -use std::sync::Arc; - -use facet::Facet; -use rig::completion::ToolDefinition; -use rig::message::{ToolCall, ToolFunction}; -use rig::tool::{ToolDyn, ToolError}; -use rig::wasm_compat::WasmBoxedFuture; - -use crate::core::{Module, Signature}; -use crate::predictors::{Predict, PredictBuilder}; -use crate::{PredictError, Predicted, Schema}; - -/// ReAct action-step schema. -#[derive(dsrs_macros::Signature, Clone, Debug)] -struct ReActActionStep { - #[input] - input: String, - - #[input] - trajectory: String, - - #[output] - thought: String, - - #[output] - action: String, - - #[output] - action_input: String, -} - -/// ReAct extraction-step schema. -#[derive(dsrs_macros::Signature, Clone, Debug)] -struct ReActExtractStep -where - O: Schema + for<'a> Facet<'a> + Send + Sync + 'static, -{ - #[input] - input: String, - - #[input] - trajectory: String, - - #[output] - output: O, -} - -#[derive(facet::Facet)] -#[facet(crate = facet)] -pub struct ReAct -where - S: Signature, - S::Input: Schema + Clone, - S::Output: Schema, -{ - action: Predict, - extract: Predict>, - #[facet(skip, opaque)] - tools: Vec>, - #[facet(skip)] - max_steps: usize, -} - -impl ReAct -where - S: Signature, - S::Input: Schema + Clone, - S::Output: Schema, -{ - pub fn new() -> Self { - Self::builder().build() - } - - pub fn builder() -> ReActBuilder { - ReActBuilder::new() - } - - async fn render_tool_manifest(&self) -> String { - if self.tools.is_empty() { - return "Available tools: (none)".to_string(); - } - - let mut lines = vec!["Available tools:".to_string()]; - for tool in &self.tools { - let definition = tool.definition(String::new()).await; - lines.push(format!("- {}: {}", definition.name, definition.description)); - } - - lines.join("\n") - } - - async fn execute_tool(&self, name: &str, args: String) -> String { - let normalized = name.trim(); - - for tool in &self.tools { - let candidate = tool.name(); - if candidate.eq_ignore_ascii_case(normalized) - || normalized.contains(&candidate) - || candidate.contains(normalized) - { - return match tool.call(args).await { - Ok(result) => result, - Err(err) => format!("tool_error: {err}"), - }; - } - } - - // Keep unknown actions explicit in trajectory instead of silently invoking - // an arbitrary tool, which hides planner/output bugs from callers. - tracing::debug!(tool = %normalized, "react tool name not found"); - let _ = args; - - format!("tool_not_found: {name}") - } - - fn is_terminal_action(action: &str) -> bool { - action.eq_ignore_ascii_case("finish") - || action.eq_ignore_ascii_case("final") - || action.eq_ignore_ascii_case("done") - } - - fn format_trace_entry( - step: usize, - thought: &str, - action: &str, - action_input: &str, - observation: Option<&str>, - ) -> String { - let observation_text = observation.unwrap_or(""); - format!( - "Step {step}\nThought: {thought}\nAction: {action}\nAction Input: {action_input}\nObservation: {observation_text}" - ) - } - - async fn run(&self, input: S::Input) -> Result, PredictError> { - let serialized_input = serde_json::to_string(&input) - .unwrap_or_else(|_| "".to_string()); - - let tool_manifest = self.render_tool_manifest().await; - let mut trajectory_text = tool_manifest.clone(); - trajectory_text.push_str("\n\n"); - - let mut tool_calls = Vec::new(); - let mut tool_executions = Vec::new(); - tool_executions.push(tool_manifest); - - for step in 0..self.max_steps { - let action_input = - ReActActionStepInput::new(serialized_input.clone(), trajectory_text.clone()); - - let action_predicted = self.action.call(action_input).await?; - let (action_output, mut action_metadata) = action_predicted.into_parts(); - tool_calls.append(&mut action_metadata.tool_calls); - tool_executions.append(&mut action_metadata.tool_executions); - - let ReActActionStepOutput { - thought, - action, - action_input, - } = action_output; - - let action_name = action - .trim() - .trim_matches('"') - .trim_matches('\'') - .to_string(); - - if Self::is_terminal_action(&action_name) { - let trace = - Self::format_trace_entry(step + 1, &thought, &action_name, &action_input, None); - tool_executions.push(trace.clone()); - trajectory_text.push_str(&format!( - "Step {}\nThought: {}\nFinal: {}\n\n", - step + 1, - thought, - action_input - )); - break; - } - - let observation = self.execute_tool(&action_name, action_input.clone()).await; - - tool_calls.push(ToolCall { - id: format!("react-step-{}", step + 1), - call_id: None, - function: ToolFunction { - name: action_name.clone(), - arguments: serde_json::json!(action_input), - }, - signature: None, - additional_params: None, - }); - tool_executions.push(Self::format_trace_entry( - step + 1, - &thought, - &action_name, - &action_input, - Some(&observation), - )); - - trajectory_text.push_str(&format!( - "Step {}\nThought: {}\nAction: {}\nAction Input: {}\nObservation: {}\n\n", - step + 1, - thought, - action_name, - action_input, - observation - )); - } - - let extract_input = ReActExtractStepInput::new(serialized_input, trajectory_text); - - let extract_predicted = self.extract.call(extract_input).await?; - let (extract_output, mut extract_metadata) = extract_predicted.into_parts(); - extract_metadata.tool_calls.extend(tool_calls); - extract_metadata.tool_executions.extend(tool_executions); - - let output: ReActExtractStepOutput = extract_output; - Ok(Predicted::new(output.output, extract_metadata)) - } -} - -impl Default for ReAct -where - S: Signature, - S::Input: Schema + Clone, - S::Output: Schema, -{ - fn default() -> Self { - Self::new() - } -} - -impl Module for ReAct -where - S: Signature, - S::Input: Schema + Clone, - S::Output: Schema, -{ - type Input = S::Input; - type Output = S::Output; - - async fn forward(&self, input: S::Input) -> Result, PredictError> { - self.run(input).await - } -} - -pub struct ReActBuilder -where - S: Signature, - S::Input: Schema + Clone, - S::Output: Schema, -{ - action: PredictBuilder, - extract: PredictBuilder>, - tools: Vec>, - max_steps: usize, -} - -impl ReActBuilder -where - S: Signature, - S::Input: Schema + Clone, - S::Output: Schema, -{ - fn new() -> Self { - Self { - action: Predict::builder(), - extract: Predict::builder(), - tools: Vec::new(), - max_steps: 4, - } - } - - pub fn action_instruction(mut self, instruction: impl Into) -> Self { - self.action = self.action.instruction(instruction); - self - } - - pub fn extract_instruction(mut self, instruction: impl Into) -> Self { - self.extract = self.extract.instruction(instruction); - self - } - - pub fn max_steps(mut self, max_steps: usize) -> Self { - self.max_steps = max_steps.max(1); - self - } - - pub fn add_tool(mut self, tool: impl ToolDyn + 'static) -> Self { - self.tools.push(Arc::new(tool)); - self - } - - pub fn with_tools(mut self, tools: impl IntoIterator>) -> Self { - self.tools.extend(tools); - self - } - - pub fn tool( - mut self, - name: impl Into, - description: impl Into, - tool_fn: F, - ) -> Self - where - F: Fn(String) -> Fut + Send + Sync + 'static, - Fut: Future + Send + 'static, - { - self.tools.push(Arc::new(PlainAsyncTool { - name: name.into(), - description: description.into(), - handler: tool_fn, - })); - self - } - - /// Sets a per-instance LM on both the action and extract predictors, - /// bypassing the global. See [`PredictBuilder::lm`]. - pub fn lm(mut self, lm: crate::core::LM) -> Self { - self.action = self.action.lm(lm.clone()); - self.extract = self.extract.lm(lm); - self - } - - pub fn build(self) -> ReAct { - ReAct { - action: self.action.build(), - extract: self.extract.build(), - tools: self.tools, - max_steps: self.max_steps, - } - } -} - -struct PlainAsyncTool { - name: String, - description: String, - handler: F, -} - -impl ToolDyn for PlainAsyncTool -where - F: Fn(String) -> Fut + Send + Sync + 'static, - Fut: Future + Send + 'static, -{ - fn name(&self) -> String { - self.name.clone() - } - - fn definition<'a>(&'a self, _prompt: String) -> WasmBoxedFuture<'a, ToolDefinition> { - Box::pin(async move { - ToolDefinition { - name: self.name.clone(), - description: self.description.clone(), - parameters: serde_json::json!({ - "type": "object", - "additionalProperties": true - }), - } - }) - } - - fn call<'a>(&'a self, args: String) -> WasmBoxedFuture<'a, Result> { - Box::pin(async move { Ok((self.handler)(args).await) }) - } -} diff --git a/crates/dspy-rs/src/optimizer/bootstrap.rs b/crates/dspy-rs/src/optimizer/bootstrap.rs index 2118ca1c..67c34bdf 100644 --- a/crates/dspy-rs/src/optimizer/bootstrap.rs +++ b/crates/dspy-rs/src/optimizer/bootstrap.rs @@ -8,35 +8,35 @@ use anyhow::{Result, anyhow}; use bon::Builder; use serde::{Deserialize, Serialize}; +use crate::core::ToInput; use crate::evaluate::{DEFAULT_EVAL_CONCURRENCY, TypedMetric}; -use crate::optimizer::engine::{ - Budget, Candidate, EngineConfig, EvalEngine, EvalOutcome, Spend, apply_candidate, -}; +use crate::optimizer::engine::{Candidate, Engine, EngineConfig, EvalOutcome, Spend}; use crate::optimizer::harvest::{collect_demo_candidates, select_demos}; -use crate::optimizer::{Optimizer, predictor_names}; +use crate::optimizer::{OptimizeTarget, Optimizer, OptimizerCommon, Report}; use crate::trace::Trace; -use crate::core::ToInput; -use crate::{Facet, Module}; +use crate::{Module, Predictors}; /// Few-shot demo bootstrapper — the simplest complete optimizer. /// /// One teacher pass, one candidate, one comparison: /// -/// 1. **Teacher pass** — runs the module over the trainset under trace -/// capture (via the shared [`EvalEngine`]), scoring each rollout with the +/// 1. **Teacher pass** — runs the target over the trainset under trace +/// capture (via the shared [`Engine`]), scoring each rollout with the /// metric. -/// 2. **Harvest** — successful `Predict` spans from rollouts scoring at least +/// 2. **Harvest** — successful `Predict` spans scoring at least /// `min_demo_score` become few-shot demo rows, joined to their predictor -/// purely by trace component name. -/// 3. **Candidate eval** — the harvested demos form a demo-overlay -/// [`Candidate`], evaluated on the same engine (teacher rollouts already -/// sit in the rollout cache, so the baseline never re-runs). -/// 4. **Keep if better** — the demos are installed permanently only when the -/// candidate's mean score beats the baseline. +/// purely by trace component name (the [`Predictors`] contract name). A +/// span scores as the rollout does unless the metric attached a span-level +/// eval ([`TypedMetric::evaluate_spans`]), which then takes precedence. +/// 3. **Candidate eval** — the harvested demos form a demo [`Candidate`], +/// evaluated ambiently on the same engine (teacher rollouts already sit in +/// the rollout cache, so the baseline never re-runs). +/// 4. **Keep if better** — the demos are installed (once, at the end) only +/// when the candidate's mean score beats the baseline. /// /// ```ignore /// let bootstrap = BootstrapFewShot::builder().max_demos(4).build(); -/// let report = bootstrap.compile(&mut module, trainset, &metric).await?; +/// let report = bootstrap.compile_module(&mut module, &trainset, &metric).await?; /// if report.adopted { /// println!("{:.3} -> {:.3}", report.baseline_score, report.candidate_score.unwrap()); /// } @@ -47,9 +47,10 @@ pub struct BootstrapFewShot { #[builder(default = 4)] pub max_demos: usize, - /// Minimum whole-rollout metric score for a rollout's spans to qualify as - /// demos. Defaults to `1.0` — full-credit rollouts only, assuming a 0–1 - /// metric; lower it for graded metrics. + /// Minimum score for a span to qualify as a demo: the span's own eval + /// when the metric attached one ([`TypedMetric::evaluate_spans`]), the + /// whole-rollout metric score otherwise. Defaults to `1.0` — full-credit + /// only, assuming a 0–1 metric; lower it for graded metrics. #[builder(default = 1.0)] pub min_demo_score: f64, @@ -73,48 +74,62 @@ pub struct BootstrapReport { pub candidate_score: Option, /// Whether the demo candidate beat the baseline and was installed. pub adopted: bool, - /// Demos harvested per predictor (dotted path -> count). + /// Demos harvested per predictor (leaf name -> count). pub demos_per_predictor: BTreeMap, /// Engine spend for the whole run. pub spend: Spend, } -impl Optimizer for BootstrapFewShot { - type Report = BootstrapReport; +impl BootstrapFewShot { + fn common(&self) -> OptimizerCommon { + OptimizerCommon { + eval_concurrency: self.eval_concurrency, + max_metric_calls: self.max_metric_calls, + max_lm_calls: self.max_lm_calls, + ..OptimizerCommon::default() + } + } - async fn compile( + /// Convenience: bootstraps a typed module over a trainset with this + /// optimizer's default engine. + pub async fn compile_module( &self, module: &mut M, - trainset: Vec, + trainset: &[E], metric: &MT, - ) -> Result + ) -> Result where E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, + M: Module + Predictors, MT: TypedMetric, { - let names = predictor_names(module)?; - if names.is_empty() { + let mut target = OptimizeTarget::module(module, trainset, metric); + let mut engine = Engine::new(Optimizer::engine_config(self)); + let report = Optimizer::compile(self, &mut target, &mut engine).await?; + report + .into_bootstrap() + .ok_or_else(|| anyhow!("BootstrapFewShot must return a bootstrap report")) + } +} + +#[async_trait::async_trait(?Send)] +impl Optimizer for BootstrapFewShot { + fn engine_config(&self) -> EngineConfig { + self.common().engine_config() + } + + async fn compile( + &self, + target: &mut OptimizeTarget<'_>, + engine: &mut Engine, + ) -> Result { + if target.leaves().is_empty() { return Err(anyhow!("no optimizable predictors found")); } - let mut engine = EvalEngine::new( - trainset, - metric, - EngineConfig { - concurrency: self.eval_concurrency, - budget: Budget { - max_metric_calls: self.max_metric_calls, - max_lm_calls: self.max_lm_calls, - max_tokens: None, - }, - cache_salt: 0, - }, - ); - // 1. Teacher pass: baseline candidate over the full trainset, traced. let baseline = engine.register(Candidate::new()); - let baseline_eval = match engine.evaluate(module, baseline, None).await? { + let baseline_eval = match engine.evaluate(target, baseline, None).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { needed } => { return Err(anyhow!( @@ -145,47 +160,47 @@ impl Optimizer for BootstrapFewShot { .collect(); if demos.is_empty() { - return Ok(BootstrapReport { + return Ok(Report::Bootstrap(BootstrapReport { baseline_score, candidate_score: None, adopted: false, demos_per_predictor, spend: *engine.spend(), - }); + })); } - // 3. Demos become a candidate, evaluated on the same engine. + // 3. Demos become a candidate, evaluated ambiently on the same engine. let mut candidate = Candidate::new(); for (name, demo_set) in demos { candidate.set_demos(name, demo_set); } let candidate_idx = engine.register(candidate.clone()); - let candidate_eval = match engine.evaluate(module, candidate_idx, None).await? { + let candidate_eval = match engine.evaluate(target, candidate_idx, None).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { .. } => { - return Ok(BootstrapReport { + return Ok(Report::Bootstrap(BootstrapReport { baseline_score, candidate_score: None, adopted: false, demos_per_predictor, spend: *engine.spend(), - }); + })); } }; let candidate_score = candidate_eval.mean(); - // 4. Keep if better: permanent install through the one candidate seam. + // 4. Keep if better: the run's one mutation, at the end. let adopted = candidate_score > baseline_score; if adopted { - let _undo = apply_candidate(module, &candidate)?; + target.install(&candidate)?; } - Ok(BootstrapReport { + Ok(Report::Bootstrap(BootstrapReport { baseline_score, candidate_score: Some(candidate_score), adopted, demos_per_predictor, spend: *engine.spend(), - }) + })) } } diff --git a/crates/dspy-rs/src/optimizer/copro.rs b/crates/dspy-rs/src/optimizer/copro.rs index 703efc65..8078d398 100644 --- a/crates/dspy-rs/src/optimizer/copro.rs +++ b/crates/dspy-rs/src/optimizer/copro.rs @@ -1,14 +1,13 @@ use anyhow::{Result, anyhow}; use bon::Builder; +use std::collections::BTreeMap; -use crate::core::DynPredictor; -use crate::evaluate::TypedMetric; -use crate::optimizer::engine::{ - Budget, Candidate, EngineConfig, EvalEngine, EvalOutcome, apply_candidate, -}; -use crate::optimizer::{Optimizer, predictor_names, with_named_predictor}; use crate::core::ToInput; -use crate::{Facet, Module}; +use crate::evaluate::TypedMetric; +use crate::optimizer::engine::{Candidate, Engine, EngineConfig, EvalOutcome}; +use crate::optimizer::target::LeafInfo; +use crate::optimizer::{OptimizeTarget, Optimizer, OptimizerCommon, Report}; +use crate::{Module, Predictors}; /// Breadth-first instruction optimizer. /// @@ -17,12 +16,14 @@ use crate::{Facet, Module}; /// `depth` rounds. Simple and predictable — good for quick iteration when you want /// better instructions without complex search. /// -/// COPRO is a thin strategy over the shared [`EvalEngine`]: each candidate -/// instruction is an overlay [`Candidate`] evaluated through the engine's -/// cached bounded-concurrency fan-out, and each round's winner is installed -/// permanently through the one candidate seam ([`apply_candidate`]). Repeated -/// candidates within a round (the base instruction always competes) are -/// deduplicated by content hash and served from the rollout cache. +/// COPRO is a thin strategy over the shared [`Engine`]: each candidate +/// instruction is a name-keyed [`Candidate`] layered on the winners +/// accumulated so far, evaluated through the engine's cached +/// bounded-concurrency ambient-injection fan-out. Nothing mutates the module +/// during the search; the accumulated winner is installed once at the end +/// through [`OptimizeTarget::install`]. Repeated candidates (the base +/// instruction always competes) deduplicate by content hash and are served +/// from the rollout cache. /// /// Does not use feedback from the metric — only the numerical score matters. If you /// have rich textual feedback, use [`GEPA`](crate::GEPA) instead. @@ -40,12 +41,13 @@ use crate::{Facet, Module}; /// /// # Cost /// -/// Total LM calls ≈ `breadth × depth × num_predictors × trainset_size`. For a module +/// Total LM calls ≈ `breadth × depth × num_predictors × trainset_size`, minus +/// rollout-cache hits (previous winners re-compete for free). For a module /// with 2 predictors, breadth=10, depth=3, and 50 training examples: ~3000 calls. /// /// ```ignore /// let copro = COPRO::builder().breadth(10).depth(3).build(); -/// copro.compile(&mut module, trainset, &metric).await?; +/// copro.compile_module(&mut module, &trainset, &metric).await?; /// ``` #[derive(Builder)] pub struct COPRO { @@ -70,29 +72,26 @@ pub struct COPRO { } impl COPRO { - fn current_instruction(module: &mut M, predictor_name: &str) -> Result - where - M: for<'a> Facet<'a>, - { - with_named_predictor(module, predictor_name, |predictor| { - Ok(predictor.instruction()) - }) + fn common(&self) -> OptimizerCommon { + OptimizerCommon { + eval_concurrency: self.eval_concurrency, + ..OptimizerCommon::default() + } } fn candidate_instructions( &self, base_instruction: &str, - predictor: &dyn DynPredictor, + leaf: &LeafInfo, depth: usize, ) -> Vec { let mut candidates = Vec::with_capacity(self.breadth.max(1)); candidates.push(base_instruction.to_string()); - let output_hint = predictor - .schema() - .output_fields() + let output_hint = leaf + .output_fields .last() - .map(|field| field.lm_name) + .map(|(name, _)| name.as_str()) .unwrap_or("output"); for idx in 0..self.breadth.saturating_sub(1) { @@ -106,59 +105,70 @@ impl COPRO { candidates } -} -impl Optimizer for COPRO { - type Report = (); - - async fn compile( + /// Convenience: optimizes a typed module over a trainset with this + /// optimizer's default engine, installing the winning instructions. + pub async fn compile_module( &self, module: &mut M, - trainset: Vec, + trainset: &[E], metric: &MT, - ) -> Result + ) -> Result<()> where E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, + M: Module + Predictors, MT: TypedMetric, { + let mut target = OptimizeTarget::module(module, trainset, metric); + let mut engine = Engine::new(Optimizer::engine_config(self)); + Optimizer::compile(self, &mut target, &mut engine).await?; + Ok(()) + } +} + +#[async_trait::async_trait(?Send)] +impl Optimizer for COPRO { + fn engine_config(&self) -> EngineConfig { + self.common().engine_config() + } + + async fn compile( + &self, + target: &mut OptimizeTarget<'_>, + engine: &mut Engine, + ) -> Result { if self.breadth <= 1 { return Err(anyhow!("breadth must be greater than 1")); } - let predictor_names = predictor_names(module)?; - - if predictor_names.is_empty() { + let leaves = target.leaves().to_vec(); + if leaves.is_empty() { return Err(anyhow!("no optimizable predictors found")); } - let mut engine = EvalEngine::new( - trainset, - metric, - EngineConfig { - concurrency: self.eval_concurrency, - budget: Budget::unlimited(), - cache_salt: 0, - }, - ); + // Winners accumulate here; the module itself is never touched until + // the final install. + let mut current = Candidate::new(); + let mut current_instructions: BTreeMap = leaves + .iter() + .map(|leaf| (leaf.name.clone(), leaf.instruction.clone())) + .collect(); for depth in 0..self.depth { - for predictor_name in &predictor_names { - let base_instruction = Self::current_instruction(module, predictor_name)?; - - let candidates = with_named_predictor(module, predictor_name, |predictor| { - Ok(self.candidate_instructions(&base_instruction, predictor, depth)) - })?; + for leaf in &leaves { + let base_instruction = current_instructions[&leaf.name].clone(); + let instructions = self.candidate_instructions(&base_instruction, leaf, depth); let mut best: Option<(f64, String)> = None; - for instruction in candidates { - let row = - engine.register(Candidate::with_instruction(predictor_name, &instruction)); - let eval = match engine.evaluate(module, row, None).await? { + for instruction in instructions { + let mut candidate = current.clone(); + candidate.set_instruction(&leaf.name, &instruction); + let row = engine.register(candidate); + let eval = match engine.evaluate(target, row, None).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { needed } => { return Err(anyhow!( - "unexpected budget exhaustion ({needed} rollouts) with an unlimited budget" + "budget exhausted ({needed} rollouts needed) during COPRO round" )); } }; @@ -170,17 +180,15 @@ impl Optimizer for COPRO { let (_, best_instruction) = best.expect("breadth > 1 guarantees at least one candidate"); - // Permanent install through the one candidate seam. The engine's - // baseline hash changes with it, correctly invalidating cached - // rollouts recorded against the previous round's skeleton. - let _undo = apply_candidate( - module, - &Candidate::with_instruction(predictor_name, best_instruction), - )?; + current.set_instruction(&leaf.name, &best_instruction); + current_instructions.insert(leaf.name.clone(), best_instruction); } } - Ok(()) + // The one mutation of the run: install the accumulated winner. + target.install(¤t)?; + + Ok(Report::None) } } @@ -191,7 +199,7 @@ mod tests { use super::*; use crate::evaluate::{Eval, TypedMetric}; use crate::trace::Trace; - use crate::{CallMetadata, Predict, PredictError, Predicted, Signature}; + use crate::{CallMetadata, Predict, PredictError, Predicted, PredictorInfo, Signature}; #[derive(Signature, Clone, Debug)] struct CoproStateSig { @@ -202,12 +210,12 @@ mod tests { answer: String, } - #[derive(facet::Facet)] - #[facet(crate = facet)] struct CoproStateModule { predictor: Predict, } + crate::predictors!(CoproStateModule { predictor }); + impl Module for CoproStateModule { type Input = CoproStateSigInput; type Output = CoproStateSigOutput; @@ -252,7 +260,7 @@ mod tests { } #[tokio::test] - async fn compile_restores_state_when_metric_errors() { + async fn compile_leaves_state_untouched_when_metric_errors() { let optimizer = COPRO::builder().breadth(2).depth(1).build(); let mut module = CoproStateModule { predictor: Predict::::builder() @@ -261,15 +269,15 @@ mod tests { }; let err = optimizer - .compile(&mut module, trainset(), &AlwaysFailMetric) + .compile_module(&mut module, &trainset(), &AlwaysFailMetric) .await .expect_err("candidate scoring should propagate metric failure"); assert!(err.to_string().contains("metric failure")); - let instruction = with_named_predictor(&mut module, "predictor", |predictor| { - Ok(predictor.instruction()) - }) - .expect("predictor lookup should succeed"); - assert_eq!(instruction, "seed-instruction"); + // Candidates are ambient — a failed run can't have leaked state. + assert_eq!( + PredictorInfo::instruction(&module.predictor), + "seed-instruction" + ); } } diff --git a/crates/dspy-rs/src/optimizer/engine.rs b/crates/dspy-rs/src/optimizer/engine.rs index 44371a53..31000239 100644 --- a/crates/dspy-rs/src/optimizer/engine.rs +++ b/crates/dspy-rs/src/optimizer/engine.rs @@ -1,76 +1,54 @@ //! The shared evaluation engine (vision §5.4): every optimizer is a thin //! strategy over this core. //! -//! # Why this lives in `optimizer/`, not `evaluate/` +//! # One engine, two lanes //! -//! The `evaluate/` module owns the *strategy-free* primitives: the -//! [`TypedMetric`] trait and the traced rollout loop -//! (`evaluate_examples_traced`). This engine composes those primitives with -//! optimizer-side vocabulary — candidates, budgets, rollout caching, Pareto -//! bookkeeping, minibatch gating — and applies candidates through the -//! crate-internal mutation seam (`with_named_predictor` / -//! `DynPredictor::apply_update`) that already lives in `optimizer/`. Putting it -//! here keeps the dependency arrow one-way: `optimizer` → `evaluate`, never the -//! reverse. +//! [`Engine`] owns the strategy-independent bookkeeping — candidate registry, +//! rollout cache, budget metering, score matrix / Pareto views, minibatch +//! gating — and evaluates candidates against an +//! [`OptimizeTarget`](crate::optimizer::OptimizeTarget), which is one of two +//! lanes: //! -//! # The pieces +//! - **module lane** — a typed module + trainset + [`TypedMetric`]. Candidates +//! are name-keyed [`Candidate`]s injected *ambiently* per rollout via +//! [`fx::with_params`](crate::fx::with_params); each +//! [`Predict`](crate::Predict) leaf binds its own entry at call time. +//! Nothing is ever mutated during evaluation, so candidates fan out +//! concurrently exactly like the program lane. +//! - **program lane** — an interpreter-loaded [`Program`](crate::ir::Program) +//! + labeled [`DemoRow`](crate::ir::DemoRow)s + a +//! [`ProgramMetric`](crate::optimizer::ProgramMetric). Candidates are +//! [`ir::Overlay`](crate::ir::Overlay)s (or [`Candidate`]s bound through +//! [`fx::Params::bind`](crate::fx::Params::bind)) read through at render +//! time. //! -//! - **[`Candidate`]** — a named set of parameter overlays (predictor name → -//! instruction/demos) plus a stable content hash. This is the §5.4 overlay -//! contract in its pre-IR form: cheap to clone, trivially serializable, and -//! applied/restored in exactly one place ([`apply_candidate`] / -//! [`restore_candidate`]) through the single mutation seam. -//! - **[`EvalEngine`]** — bounded-concurrency async fan-out over -//! (candidate × examples) with per-rollout trace capture, a rollout cache, -//! budget metering, minibatch gating, and a per-instance score matrix. -//! - **[`ScoreMatrix`] / [`ParetoView`]** — (candidates × examples) score -//! bookkeeping generalizing what GEPA's `ParetoFrontier` does, usable by any -//! strategy. -//! -//! # Concurrency model (and the candidate-parallelism seam) -//! -//! Candidates mutate shared module state through `apply_update`, so the engine -//! **serializes candidate application** and parallelizes across *examples* -//! within one candidate (`buffer_unordered`, bounded by -//! [`EngineConfig::concurrency`]). This is correct today because a module is -//! immutable (`&M`) for the duration of one candidate's fan-out. -//! Candidate-level parallelism — evaluating many candidates over one skeleton -//! simultaneously — requires overlays applied at render time, and that is the -//! IR lane's path: -//! [`ProgramEvalEngine`](crate::optimizer::program_engine::ProgramEvalEngine) -//! evaluates N `ir::Overlay` candidates over one shared `Arc` through -//! the interpreter in a single fan-out (no apply/restore at all). The module -//! lane here keeps its serialized apply/restore model unchanged. +//! The traced-rollout loop itself is owned by `evaluate/` — +//! the engine composes `rollout_traced`, it does not reimplement it. //! //! # Cache keying //! -//! Rollouts are cached on `(baseline hash, candidate hash, example uid, +//! Rollouts are cached on `(baseline identity, candidate hash, example uid, //! cache salt)`: //! -//! - the *baseline hash* is the module's [`ModuleState`] before the overlay is -//! applied, so permanently installing a winner mid-run (COPRO between -//! rounds, MIPRO after demo bootstrap) correctly invalidates stale entries; +//! - the *baseline identity* is computed once per run when the target is +//! constructed (module lane: hash of the predictors() state snapshot; +//! program lane: the program hash). Installing a winner and building a new +//! target yields a new baseline, invalidating stale entries; //! - the *cache salt* ([`EngineConfig::cache_salt`]) is the sampling-params -//! seam: today sampling params live on LM configs outside the candidate, so -//! callers that change them must bump the salt. When model refs become -//! overlay parameters, they fold into the candidate hash and the salt can -//! retire. +//! seam: sampling params live on LM configs outside the candidate, so +//! callers that change them must bump the salt. use std::collections::{BTreeMap, HashMap}; -use std::time::Instant; use anyhow::{Result, anyhow}; use futures::stream::{self, StreamExt, TryStreamExt}; use serde::{Deserialize, Serialize}; -use crate::core::{ModuleState, PredictState, StateUpdate}; -use crate::evaluate::{DEFAULT_EVAL_CONCURRENCY, Eval, TypedMetric}; -use crate::optimizer::pareto::ParetoStatistics; -use crate::optimizer::with_named_predictor; -use crate::trace::{JsonMap, Trace, TraceMeta, TraceOutcome, capture_with_meta}; +use crate::evaluate::{DEFAULT_EVAL_CONCURRENCY, Eval}; +use crate::optimizer::target::OptimizeTarget; +use crate::trace::{JsonMap, Trace}; use crate::utils::hash::StableHasher; -use crate::core::ToInput; -use crate::{Facet, LmUsage, Module}; +use crate::LmUsage; /// Score tolerance for Pareto win/tie comparisons (matches the historical /// `ParetoFrontier` tolerance). @@ -139,54 +117,47 @@ pub(crate) fn canonical_hash(value: &T) -> u64 { hasher.finish() } -fn json_object(value: Result) -> Option { - match value { - Ok(serde_json::Value::Object(map)) => Some(map), - _ => None, - } -} - // --------------------------------------------------------------------------- // Candidate as data // --------------------------------------------------------------------------- -/// A per-predictor parameter overlay: which optimizable values to install. +/// One leaf's slice of a [`Candidate`]: which optimizable values to inject. /// -/// `None` fields leave the predictor's current value untouched — an overlay is -/// a *partial* update, resolved against whatever module state is live when the -/// candidate is applied. +/// Unset fields leave the leaf's incumbent value untouched — a candidate is a +/// *partial* configuration, resolved per slot at render time (ambient entries +/// win over instance state, exactly the precedence the old mutation seam had). #[derive(Clone, Debug, Default, Serialize, Deserialize)] -pub struct Overlay { - /// `Some` installs this instruction override. +pub struct CandidateSlot { + /// `Some` injects this instruction override. #[serde(default, skip_serializing_if = "Option::is_none")] pub instruction: Option, - /// `Some` replaces the demo set. Rows are flat JSON objects (field name → - /// value, input and output fields merged), the same shape as - /// [`PredictState::demos`]. + /// Explicitly reset the instruction to the signature default, winning + /// over any instance override. Mutually exclusive with `instruction`. + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + pub clear_instruction: bool, + /// `Some` replaces the demo set (an empty vec clears it). Rows are flat + /// JSON objects (field name → value, input and output fields merged), the + /// same shape as [`PredictState::demos`](crate::core::PredictState). #[serde(default, skip_serializing_if = "Option::is_none")] pub demos: Option>, } -impl Overlay { - fn to_update(&self) -> StateUpdate { - StateUpdate { - instruction: self.instruction.clone().map(Some), - demos: self.demos.clone(), - } - } -} - -/// A candidate is *data*: a named set of [`Overlay`]s (predictor name → -/// instruction/demos) plus a stable hash ([`Candidate::stable_hash`]). +/// A candidate is *data*: name-keyed per-leaf overlays plus a stable content +/// hash ([`Candidate::stable_hash`]). The candidate currency of the module +/// lane, and — via [`to_params`](Candidate::to_params) + +/// [`fx::Params::bind`](crate::fx::Params::bind) — of the program lane too. /// -/// This is the §5.4 overlay contract in its pre-IR form — cheap to clone, -/// serializable, applied and restored in exactly one place -/// ([`apply_candidate`] / [`restore_candidate`]). The empty candidate -/// (`Candidate::default()`) is the baseline: the module exactly as it is. +/// Cheap to clone, serializable, and **never applied by mutation**: the +/// engine scopes it ambiently around each rollout +/// ([`fx::with_params`](crate::fx::with_params)); the single mutating step is +/// the caller-driven final install +/// ([`OptimizeTarget::install`](crate::optimizer::OptimizeTarget::install)). +/// The empty candidate (`Candidate::default()`) is the baseline: the module +/// exactly as it is. #[derive(Clone, Debug, Default, Serialize, Deserialize)] pub struct Candidate { - /// Predictor name (fx slot name / facet dotted path) → overlay. - pub overlays: BTreeMap, + /// Leaf name (the [`Predictors`](crate::Predictors) contract name) → slot. + pub slots: BTreeMap, } impl Candidate { @@ -208,99 +179,65 @@ impl Candidate { name: impl Into, instruction: impl Into, ) -> &mut Self { - self.overlays.entry(name.into()).or_default().instruction = Some(instruction.into()); + let slot = self.slots.entry(name.into()).or_default(); + slot.instruction = Some(instruction.into()); + slot.clear_instruction = false; + self + } + + /// Explicitly resets a named predictor's instruction to its signature + /// default (winning over any instance override). + pub fn clear_instruction(&mut self, name: impl Into) -> &mut Self { + let slot = self.slots.entry(name.into()).or_default(); + slot.instruction = None; + slot.clear_instruction = true; self } /// Sets the demo overlay for a named predictor. Each row is a flat JSON - /// object with the signature's input and output fields merged. + /// object with the signature's input and output fields merged; an empty + /// vec clears the demo set. pub fn set_demos(&mut self, name: impl Into, demos: Vec) -> &mut Self { - self.overlays.entry(name.into()).or_default().demos = Some(demos); + self.slots.entry(name.into()).or_default().demos = Some(demos); self } + /// The instruction this candidate injects for `name`, if any. + pub fn instruction_of(&self, name: &str) -> Option<&str> { + self.slots.get(name)?.instruction.as_deref() + } + + /// The demo set this candidate injects for `name`, if any. + pub fn demos_of(&self, name: &str) -> Option<&[JsonMap]> { + self.slots.get(name)?.demos.as_deref() + } + pub fn is_empty(&self) -> bool { - self.overlays.is_empty() + self.slots.is_empty() } - /// Stable content hash: identical overlay content hashes identically across - /// processes and map orderings — the cache and checkpoint identity. + /// Stable content hash: identical content hashes identically across + /// processes and map orderings — the cache identity. pub fn stable_hash(&self) -> u64 { canonical_hash(self) } -} - -/// Saved pre-overlay state for the predictors a candidate touched. Produced by -/// [`apply_candidate`], consumed by [`restore_candidate`]. -#[derive(Clone, Debug)] -pub struct CandidateUndo { - saved: BTreeMap, -} -/// Applies a candidate's overlays to a module through the single mutation seam -/// (`DynPredictor::apply_update`), returning the undo snapshot. -/// -/// This is the **one** place candidate state is written. If any overlay fails -/// to apply (unknown predictor name, demo schema mismatch), the overlays -/// applied so far are rolled back before the error returns. -pub fn apply_candidate(module: &mut M, candidate: &Candidate) -> Result -where - M: for<'a> Facet<'a>, -{ - let mut undo = CandidateUndo { - saved: BTreeMap::new(), - }; - for (name, overlay) in &candidate.overlays { - let applied = with_named_predictor(module, name, |predictor| { - let prior = predictor.dump_state(); - predictor.apply_update(overlay.to_update())?; - Ok(prior) - }); - match applied { - Ok(prior) => { - undo.saved.insert(name.clone(), prior); + /// Converts to the ambient-injection currency: name-keyed + /// [`fx::Params`](crate::fx::Params) with explicit clears preserved. + pub fn to_params(&self) -> crate::fx::Params { + let mut params = crate::fx::Params::new(); + for (name, slot) in &self.slots { + if slot.clear_instruction { + params.clear_instruction(name.clone()); + } else if let Some(text) = &slot.instruction { + params.set_instruction(name.clone(), text.clone()); } - Err(err) => { - return match restore_candidate(module, undo) { - Ok(()) => Err(err), - Err(restore_err) => Err(anyhow!( - "failed to apply candidate: {err}; and failed to roll back partial application: {restore_err}" - )), - }; + if let Some(demos) = &slot.demos { + params.set_demos(name.clone(), demos.clone()); } } + params } - Ok(undo) -} - -/// Restores the pre-candidate state captured by [`apply_candidate`]. -/// -/// Attempts every predictor even if one fails, then reports the first error. -pub fn restore_candidate(module: &mut M, undo: CandidateUndo) -> Result<()> -where - M: for<'a> Facet<'a>, -{ - let mut first_error = None; - for (name, state) in undo.saved { - if let Err(err) = with_named_predictor(module, &name, |predictor| predictor.load_state(state.clone())) - && first_error.is_none() - { - first_error = Some(anyhow!("failed to restore `{name}`: {err}")); - } - } - match first_error { - Some(err) => Err(err), - None => Ok(()), - } -} - -/// Content hash of a module's full optimizable state ([`ModuleState`]) — -/// the "skeleton" a candidate overlays. Part of the rollout-cache key. -fn baseline_hash(module: &mut M) -> Result -where - M: for<'a> Facet<'a>, -{ - Ok(canonical_hash(&ModuleState::from_module(module)?)) } // --------------------------------------------------------------------------- @@ -311,7 +248,7 @@ where /// /// `max_metric_calls` and `max_lm_calls` are metered at rollout granularity /// (one module execution = one metric call = one LM call unit; auxiliary LM -/// spend like reflection calls is charged via [`EvalEngine::charge`]). +/// spend like reflection calls is charged via [`Engine::charge`]). /// Exact per-span counts and token totals are tracked in [`Spend`] /// (`lm_spans`, `tokens`) from the captured traces; `max_tokens` stops the /// engine once recorded token usage reaches the cap. @@ -354,14 +291,13 @@ impl Budget { } } -/// What the engine has consumed so far. Serialized into checkpoints and -/// reported to strategies. +/// What the engine has consumed so far. Reported to strategies. #[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)] pub struct Spend { /// Metric evaluations executed (cache hits don't re-run the metric). pub metric_calls: usize, /// LM call units: one per executed rollout plus auxiliary charges - /// ([`EvalEngine::charge`]). + /// ([`Engine::charge`]). pub lm_calls: usize, /// Exact `Predict` spans observed across captured rollout traces. pub lm_spans: usize, @@ -378,9 +314,8 @@ pub struct Spend { /// In-memory rollout cache: `(baseline, candidate, example, salt)` → [`Eval`]. /// /// A candidate re-evaluated on a seen example returns the cached `Eval` with -/// no LM call and no metric call. Serialized into checkpoints so a resumed run -/// skips completed rollouts. (Disk-backed storage can layer on later — the key -/// is already a stable string.) +/// no LM call and no metric call. (Disk-backed storage can layer on later — +/// the key is already a stable string.) #[derive(Clone, Debug, Default, Serialize, Deserialize)] pub struct RolloutCache { entries: BTreeMap, @@ -573,6 +508,28 @@ impl ParetoView { } } +/// Snapshot of the Pareto frontier at a point in the search. +/// +/// Useful for plotting convergence. A healthy search has `num_candidates` growing +/// slowly (diversity is maintained) while `avg_coverage` increases (candidates are +/// getting more robust). If `num_candidates` is 1, the search has collapsed. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ParetoStatistics { + /// Candidates currently on the frontier. 1 means the search has converged + /// (or collapsed) to a single instruction. + pub num_candidates: usize, + /// Examples where at least one frontier candidate is the best. Should approach + /// total eval set size as the search progresses. + pub num_examples_covered: usize, + /// Mean examples won per candidate. Higher means candidates are more robust; + /// lower means more specialization. + pub avg_coverage: f32, + /// Most examples won by any single candidate. + pub max_coverage: usize, + /// Fewest examples won by any frontier candidate (always >= 1 by construction). + pub min_coverage: usize, +} + // --------------------------------------------------------------------------- // Engine // --------------------------------------------------------------------------- @@ -580,7 +537,7 @@ impl ParetoView { /// Engine tuning knobs. #[derive(Clone, Copy, Debug, Serialize, Deserialize)] pub struct EngineConfig { - /// Rollouts in flight at once within one candidate's fan-out. + /// Rollouts in flight at once within one evaluation batch. pub concurrency: usize, /// Hard spend caps; the engine stops cleanly when a batch wouldn't fit. pub budget: Budget, @@ -602,7 +559,7 @@ impl Default for EngineConfig { /// One evaluated (or cache-served) rollout. #[derive(Clone, Debug)] pub struct RolloutOutcome { - /// Index into the engine's example set. + /// Index into the target's example set. pub example: usize, pub eval: Eval, /// The captured execution trace; `None` when served from the cache. @@ -631,7 +588,7 @@ impl CandidateEval { } } -/// Result of [`EvalEngine::evaluate`]. +/// Result of [`Engine::evaluate`]. #[derive(Clone, Debug)] pub enum EvalOutcome { Complete(CandidateEval), @@ -651,7 +608,27 @@ impl EvalOutcome { } } -/// Result of [`EvalEngine::evaluate_gated`]. +/// Result of [`Engine::evaluate_many`]. +#[derive(Clone, Debug)] +pub enum BatchEvalOutcome { + /// One [`CandidateEval`] per requested candidate, in request order. + Complete(Vec), + /// The uncached portion of the batch didn't fit the remaining budget; + /// nothing ran and spend is unchanged. + BudgetExhausted { needed: usize }, +} + +impl BatchEvalOutcome { + /// The completed evaluations, if the budget allowed the batch. + pub fn completed(self) -> Option> { + match self { + Self::Complete(evals) => Some(evals), + Self::BudgetExhausted { .. } => None, + } + } +} + +/// Result of [`Engine::evaluate_gated`]. #[derive(Clone, Debug)] pub enum GateOutcome { /// Minibatch (or promotion) evaluation didn't fit the remaining budget. @@ -665,110 +642,62 @@ pub enum GateOutcome { }, } -/// Serialized engine state: candidates, score matrix, spend, and the rollout -/// cache. A resumed run skips completed rollouts via the cache. -#[derive(Serialize, Deserialize)] -struct EngineCheckpoint { - version: u32, - example_uids: Vec, - candidates: Vec, - matrix: ScoreMatrix, - cache: RolloutCache, - spend: Spend, +/// A registered candidate: its cache-identity hash plus its payload. +pub(crate) enum CandidatePayload { + /// Module-lane (name-keyed) candidate, pre-converted to the ambient + /// injection currency. Also evaluable on a program target via + /// [`fx::Params::bind`](crate::fx::Params::bind). + Params { + candidate: Candidate, + params: std::sync::Arc, + }, + /// Program-lane native candidate (can carry non-Params kinds: model refs, + /// context policies, code). Only evaluable on a program target. + Overlay(std::sync::Arc), +} + +/// A candidate bound against a concrete target, ready to inject per rollout. +#[derive(Clone)] +pub(crate) enum BoundCandidate { + Params(std::sync::Arc), + Overlay(std::sync::Arc), } /// The shared evaluation core (vision §5.4). /// -/// Owns the example set, the candidate registry, the score matrix, the rollout -/// cache, and the budget meter. Strategies (GEPA, COPRO, MIPRO, bootstrap) -/// register [`Candidate`]s and call [`evaluate`](Self::evaluate) / -/// [`evaluate_gated`](Self::evaluate_gated); the engine handles application, -/// fan-out, caching, accounting, and bookkeeping. -pub struct EvalEngine<'m, E, MT> { - examples: Vec, - example_uids: Vec, - metric: &'m MT, +/// Owns the candidate registry, the score matrix, the rollout cache, and the +/// budget meter. Strategies (GEPA, COPRO, MIPRO, SIMBA, bootstrap, +/// Structural) register +/// [`Candidate`]s and call [`evaluate`](Self::evaluate) / +/// [`evaluate_many`](Self::evaluate_many) / +/// [`evaluate_gated`](Self::evaluate_gated) against an +/// [`OptimizeTarget`](crate::optimizer::OptimizeTarget); the engine handles +/// binding, fan-out, caching, accounting, and bookkeeping. Because candidate +/// injection is ambient (never mutation), rollouts for *different candidates* +/// share one bounded-concurrency fan-out in both lanes. +pub struct Engine { config: EngineConfig, - candidates: Vec, - candidate_hashes: Vec, + candidates: Vec<(u64, CandidatePayload)>, matrix: ScoreMatrix, cache: RolloutCache, spend: Spend, + /// High-water mark of *distinct candidates* with rollouts in flight at + /// the same instant (see [`peak_candidate_concurrency`](Self::peak_candidate_concurrency)). + peak_candidates_in_flight: usize, } -impl<'m, E, MT> EvalEngine<'m, E, MT> -where - E: Serialize, -{ - pub fn new(examples: Vec, metric: &'m MT, config: EngineConfig) -> Self { - let example_uids = examples.iter().map(canonical_hash).collect(); - let matrix = ScoreMatrix::new(examples.len()); +impl Engine { + pub fn new(config: EngineConfig) -> Self { Self { - examples, - example_uids, - metric, config, candidates: Vec::new(), - candidate_hashes: Vec::new(), - matrix, + matrix: ScoreMatrix::new(0), cache: RolloutCache::default(), spend: Spend::default(), + peak_candidates_in_flight: 0, } } - /// Rebuilds an engine from a [`checkpoint`](Self::checkpoint), validating - /// that `examples` matches the checkpointed set. Completed rollouts are - /// served from the restored cache instead of re-executing. - pub fn resume( - examples: Vec, - metric: &'m MT, - config: EngineConfig, - checkpoint: &str, - ) -> Result { - let checkpoint: EngineCheckpoint = - serde_json::from_str(checkpoint).map_err(|err| anyhow!("invalid engine checkpoint: {err}"))?; - if checkpoint.version != 1 { - return Err(anyhow!( - "unsupported engine checkpoint version {}", - checkpoint.version - )); - } - let mut engine = Self::new(examples, metric, config); - if engine.example_uids != checkpoint.example_uids { - return Err(anyhow!( - "engine checkpoint does not match the provided example set" - )); - } - engine.candidate_hashes = checkpoint.candidates.iter().map(Candidate::stable_hash).collect(); - engine.candidates = checkpoint.candidates; - engine.matrix = checkpoint.matrix; - engine.cache = checkpoint.cache; - engine.spend = checkpoint.spend; - engine.matrix.ensure_rows(engine.candidates.len()); - Ok(engine) - } - - /// Serializes engine state (candidates, matrix, spend, cache) to JSON. - pub fn checkpoint(&self) -> Result { - serde_json::to_string(&EngineCheckpoint { - version: 1, - example_uids: self.example_uids.clone(), - candidates: self.candidates.clone(), - matrix: self.matrix.clone(), - cache: self.cache.clone(), - spend: self.spend, - }) - .map_err(|err| anyhow!("failed to serialize engine checkpoint: {err}")) - } - - pub fn examples(&self) -> &[E] { - &self.examples - } - - pub fn num_examples(&self) -> usize { - self.examples.len() - } - pub fn config(&self) -> &EngineConfig { &self.config } @@ -781,6 +710,13 @@ where &self.matrix } + /// The rollout cache. Keys are `(baseline identity, candidate hash, + /// example uid, salt)` — candidate identity is in the key, so two + /// candidates on the same example occupy distinct entries. + pub fn cache(&self) -> &RolloutCache { + &self.cache + } + /// Pareto view over all example columns (see [`ScoreMatrix::pareto`]). pub fn pareto(&self) -> ParetoView { self.matrix.pareto() @@ -791,24 +727,55 @@ where self.matrix.pareto_over(columns) } - /// Registers a candidate, deduplicating by content hash. Returns its index. + /// The parallelism gauge: the maximum number of **distinct candidates** + /// that have had rollouts in flight simultaneously across all batches so + /// far. + pub fn peak_candidate_concurrency(&self) -> usize { + self.peak_candidates_in_flight + } + + /// Registers a module-lane candidate, deduplicating by content hash. + /// Returns its index. pub fn register(&mut self, candidate: Candidate) -> usize { let hash = candidate.stable_hash(); - if let Some(existing) = self.candidate_hashes.iter().position(|&h| h == hash) { + if let Some(existing) = self.candidates.iter().position( + |(h, payload)| *h == hash && matches!(payload, CandidatePayload::Params { .. }), + ) { return existing; } - self.candidates.push(candidate); - self.candidate_hashes.push(hash); + let params = std::sync::Arc::new(candidate.to_params()); + self.candidates + .push((hash, CandidatePayload::Params { candidate, params })); self.matrix.ensure_rows(self.candidates.len()); self.candidates.len() - 1 } - pub fn candidate(&self, index: usize) -> &Candidate { - &self.candidates[index] + /// Registers a program-lane [`ir::Overlay`](crate::ir::Overlay) candidate, + /// deduplicating by [`Overlay::hash`](crate::ir::Overlay::hash). Returns + /// its index. + pub fn register_overlay(&mut self, overlay: crate::ir::Overlay) -> usize { + let hash = overlay.hash(); + if let Some(existing) = self.candidates.iter().position( + |(h, payload)| *h == hash && matches!(payload, CandidatePayload::Overlay(_)), + ) { + return existing; + } + self.candidates + .push((hash, CandidatePayload::Overlay(std::sync::Arc::new(overlay)))); + self.matrix.ensure_rows(self.candidates.len()); + self.candidates.len() - 1 + } + + /// The registered module-lane candidate at `index`, if it is one. + pub fn candidate(&self, index: usize) -> Option<&Candidate> { + match &self.candidates.get(index)?.1 { + CandidatePayload::Params { candidate, .. } => Some(candidate), + CandidatePayload::Overlay(_) => None, + } } pub fn candidate_hash(&self, index: usize) -> u64 { - self.candidate_hashes[index] + self.candidates[index].0 } pub fn num_candidates(&self) -> usize { @@ -827,88 +794,81 @@ where self.spend.lm_calls = self.spend.lm_calls.saturating_add(lm_calls); } - /// Evaluates a registered candidate over `subset` example indices (`None` - /// = the full set): applies the overlay through the mutation seam, fans - /// out uncached rollouts with bounded concurrency under per-rollout trace - /// capture, restores the module, and records scores into the matrix and - /// cache. + /// Evaluates N registered candidates over `subset` example indices + /// (`None` = the target's full set) in **one** bounded-concurrency + /// fan-out — candidate-level parallelism in both lanes, since candidates + /// are injected per rollout, never applied to shared state. /// /// Cached rollouts return their `Eval` with `trace: None` and consume no /// budget. If the uncached portion doesn't fit the remaining budget the - /// engine runs nothing and returns [`EvalOutcome::BudgetExhausted`]. - pub async fn evaluate( + /// engine runs nothing and returns [`BatchEvalOutcome::BudgetExhausted`]. + pub async fn evaluate_many( &mut self, - module: &mut M, - candidate: usize, + target: &OptimizeTarget<'_>, + candidates: &[usize], subset: Option<&[usize]>, - ) -> Result - where - E: ToInput + Sync, - M: Module + for<'a> Facet<'a>, - MT: TypedMetric, - { - let candidate_hash = *self - .candidate_hashes - .get(candidate) - .ok_or_else(|| anyhow!("candidate index {candidate} is not registered"))?; + ) -> Result { + for &candidate in candidates { + if candidate >= self.candidates.len() { + return Err(anyhow!("candidate index {candidate} is not registered")); + } + } + let num_examples = target.num_examples(); let indices: Vec = match subset { Some(subset) => subset.to_vec(), - None => (0..self.examples.len()).collect(), + None => (0..num_examples).collect(), }; - if let Some(&bad) = indices.iter().find(|&&idx| idx >= self.examples.len()) { + if let Some(&bad) = indices.iter().find(|&&idx| idx >= num_examples) { return Err(anyhow!( - "example index {bad} out of range ({} examples)", - self.examples.len() + "example index {bad} out of range ({num_examples} examples)" )); } - let baseline = baseline_hash(module)?; + let baseline = target.baseline(); let salt = self.config.cache_salt; - let mut cached: HashMap = HashMap::new(); - let mut pending: Vec = Vec::new(); - for &idx in &indices { - match self.cache.get(baseline, candidate_hash, self.example_uids[idx], salt) { - Some(eval) => { - cached.insert(idx, eval.clone()); + // Partition the (candidate × example) grid into cached and pending. + let mut cached: HashMap<(usize, usize), Eval> = HashMap::new(); + let mut pending: Vec<(usize, usize)> = Vec::new(); + for &candidate in candidates { + let candidate_hash = self.candidates[candidate].0; + for &idx in &indices { + let key = (candidate, idx); + if cached.contains_key(&key) || pending.contains(&key) { + continue; } - None => { - if !cached.contains_key(&idx) && !pending.contains(&idx) { - pending.push(idx); + match self + .cache + .get(baseline, candidate_hash, target.example_uid(idx), salt) + { + Some(eval) => { + cached.insert(key, eval.clone()); } + None => pending.push(key), } } } if !self.budget_allows(pending.len()) { - return Ok(EvalOutcome::BudgetExhausted { + return Ok(BatchEvalOutcome::BudgetExhausted { needed: pending.len(), }); } - let fresh: Vec<(usize, Eval, Trace)> = if pending.is_empty() { - Vec::new() - } else { - let undo = apply_candidate(module, &self.candidates[candidate])?; - let ran = self.run_rollouts(&*module, &pending, candidate_hash).await; - let restored = restore_candidate(module, undo); - match (ran, restored) { - (Ok(fresh), Ok(())) => fresh, - (Ok(_), Err(restore_err)) => return Err(restore_err), - (Err(eval_err), Ok(())) => return Err(eval_err), - (Err(eval_err), Err(restore_err)) => { - return Err(anyhow!( - "candidate evaluation failed: {eval_err}; failed to restore module state: {restore_err}" - )); - } - } - }; + // Bind each requested candidate against the target once. + let mut bound: HashMap = HashMap::new(); + for &candidate in candidates { + bound.insert(candidate, target.bind(&self.candidates[candidate].1)?); + } + + let (fresh, batch_peak) = self.run_rollouts(target, &pending, &bound).await?; + self.peak_candidates_in_flight = self.peak_candidates_in_flight.max(batch_peak); // Accounting: fresh rollouts consume budget, cached hits are free. self.spend.metric_calls += fresh.len(); self.spend.lm_calls += fresh.len(); self.spend.cache_hits += cached.len(); - for (_, _, trace) in &fresh { + for (_, _, _, trace) in &fresh { self.spend.lm_spans += trace.spans.len(); for span in &trace.spans { self.spend.tokens = self.spend.tokens + span.usage; @@ -916,62 +876,81 @@ where } // Bookkeeping: cache inserts + matrix records. - let mut fresh_by_idx: HashMap = HashMap::with_capacity(fresh.len()); - for (idx, eval, trace) in fresh { - self.cache - .insert(baseline, candidate_hash, self.example_uids[idx], salt, eval.clone()); + let mut fresh_by_key: HashMap<(usize, usize), (Eval, Trace)> = + HashMap::with_capacity(fresh.len()); + for (candidate, idx, eval, trace) in fresh { + self.cache.insert( + baseline, + self.candidates[candidate].0, + target.example_uid(idx), + salt, + eval.clone(), + ); self.matrix.record(candidate, idx, eval.score); - fresh_by_idx.insert(idx, (eval, trace)); + fresh_by_key.insert((candidate, idx), (eval, trace)); } - for (&idx, eval) in &cached { + for (&(candidate, idx), eval) in &cached { self.matrix.record(candidate, idx, eval.score); } - let rollouts = indices + let evals = candidates .iter() - .map(|&idx| { - if let Some((eval, trace)) = fresh_by_idx.remove(&idx) { - RolloutOutcome { - example: idx, - eval, - trace: Some(trace), - } - } else { - let eval = cached - .get(&idx) - .cloned() - .expect("every requested index is either fresh or cached"); - RolloutOutcome { - example: idx, - eval, - trace: None, - } - } + .map(|&candidate| CandidateEval { + candidate, + rollouts: indices + .iter() + .map(|&idx| { + if let Some((eval, trace)) = fresh_by_key.remove(&(candidate, idx)) { + RolloutOutcome { + example: idx, + eval, + trace: Some(trace), + } + } else { + let eval = cached + .get(&(candidate, idx)) + .cloned() + .expect("every requested pair is either fresh or cached"); + RolloutOutcome { + example: idx, + eval, + trace: None, + } + } + }) + .collect(), }) .collect(); - Ok(EvalOutcome::Complete(CandidateEval { - candidate, - rollouts, - })) + Ok(BatchEvalOutcome::Complete(evals)) + } + + /// Single-candidate convenience over [`evaluate_many`](Self::evaluate_many). + pub async fn evaluate( + &mut self, + target: &OptimizeTarget<'_>, + candidate: usize, + subset: Option<&[usize]>, + ) -> Result { + match self.evaluate_many(target, &[candidate], subset).await? { + BatchEvalOutcome::Complete(mut evals) => Ok(EvalOutcome::Complete(evals.remove(0))), + BatchEvalOutcome::BudgetExhausted { needed } => { + Ok(EvalOutcome::BudgetExhausted { needed }) + } + } } /// Minibatch gating (the GEPA acceptance pattern): evaluates the candidate /// on `minibatch`; only if the minibatch mean strictly beats `threshold` /// does it promote to a full-set evaluation. - pub async fn evaluate_gated( + pub async fn evaluate_gated( &mut self, - module: &mut M, + target: &OptimizeTarget<'_>, candidate: usize, minibatch: &[usize], threshold: f64, - ) -> Result - where - E: ToInput + Sync, - M: Module + for<'a> Facet<'a>, - MT: TypedMetric, - { - let minibatch_eval = match self.evaluate(module, candidate, Some(minibatch)).await? { + ) -> Result { + let minibatch_eval = match self.evaluate(target, candidate, Some(minibatch)).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { needed } => { return Ok(GateOutcome::BudgetExhausted { needed }); @@ -984,7 +963,7 @@ where }); } - match self.evaluate(module, candidate, None).await? { + match self.evaluate(target, candidate, None).await? { EvalOutcome::Complete(full) => Ok(GateOutcome::Promoted { minibatch: minibatch_eval, full, @@ -993,48 +972,75 @@ where } } - /// Bounded-concurrency fan-out over uncached examples for one applied - /// candidate. `module` is immutable here — the overlay was installed by - /// the caller — so example-level parallelism is safe. - async fn run_rollouts( + /// The one shared fan-out: every pending `(candidate, example)` pair — + /// across all candidates, in both lanes — in a single `buffer_unordered` + /// stream. Returns the fresh rollouts and the batch's distinct-candidate + /// concurrency high-water mark. + async fn run_rollouts( &self, - module: &M, - pending: &[usize], - candidate_hash: u64, - ) -> Result> - where - E: ToInput + Sync, - M: Module, - MT: TypedMetric, - { - let metric = self.metric; - stream::iter(pending.iter().map(|&idx| { - let example = &self.examples[idx]; - async move { - let input = example.to_input()?; - let meta = TraceMeta { - candidate_hash: Some(candidate_hash), - input: json_object(serde_json::to_value(&input)), - ..TraceMeta::default() - }; - let started = Instant::now(); - let (result, mut trace) = capture_with_meta(meta, || module.call(input)).await; - let predicted = result.map_err(|err| anyhow!("{err}"))?; - // Metric runs outside the capture scope so LM-as-judge metrics - // don't pollute the execution trace. - let eval = metric.evaluate(example, &predicted, Some(&trace)).await?; - trace.outcome = Some(TraceOutcome { - output: json_object(serde_json::to_value(&*predicted)), - error: None, - eval: Some(eval.clone()), - duration_us: started.elapsed().as_micros() as u64, - }); - Ok::<_, anyhow::Error>((idx, eval, trace)) + target: &OptimizeTarget<'_>, + pending: &[(usize, usize)], + bound: &HashMap, + ) -> Result<(Vec<(usize, usize, Eval, Trace)>, usize)> { + let gauge = Gauge::default(); + + let fresh: Vec<(usize, usize, Eval, Trace)> = + stream::iter(pending.iter().map(|&(candidate, idx)| { + let bound = bound[&candidate].clone(); + let candidate_hash = self.candidates[candidate].0; + let gauge = &gauge; + async move { + let _in_flight = gauge.enter(candidate); + let (eval, trace) = target.run(idx, bound, candidate_hash).await?; + Ok::<_, anyhow::Error>((candidate, idx, eval, trace)) + } + })) + .buffer_unordered(self.config.concurrency.max(1)) + .try_collect() + .await?; + + Ok((fresh, gauge.peak())) + } +} + +/// Counts distinct candidates with rollouts in flight; records the peak. +#[derive(Default)] +struct Gauge { + in_flight: std::sync::Mutex>, + peak: std::sync::atomic::AtomicUsize, +} + +impl Gauge { + fn enter(&self, candidate: usize) -> GaugeGuard<'_> { + let mut in_flight = self.in_flight.lock().unwrap(); + *in_flight.entry(candidate).or_insert(0) += 1; + self.peak + .fetch_max(in_flight.len(), std::sync::atomic::Ordering::Relaxed); + GaugeGuard { + gauge: self, + candidate, + } + } + + fn peak(&self) -> usize { + self.peak.load(std::sync::atomic::Ordering::Relaxed) + } +} + +struct GaugeGuard<'a> { + gauge: &'a Gauge, + candidate: usize, +} + +impl Drop for GaugeGuard<'_> { + fn drop(&mut self) { + let mut in_flight = self.gauge.in_flight.lock().unwrap(); + if let Some(count) = in_flight.get_mut(&self.candidate) { + *count -= 1; + if *count == 0 { + in_flight.remove(&self.candidate); } - })) - .buffer_unordered(self.config.concurrency.max(1)) - .try_collect() - .await + } } } @@ -1070,6 +1076,24 @@ mod tests { assert_eq!(empty.stable_hash(), Candidate::default().stable_hash()); } + #[test] + fn explicit_clears_are_distinct_candidate_content() { + let unset = Candidate::with_instruction("drafter", "x"); + let mut cleared = Candidate::with_instruction("drafter", "x"); + cleared.clear_instruction("drafter"); + assert_ne!(unset.stable_hash(), cleared.stable_hash()); + + // Clearing demos (empty set) differs from leaving them unset. + let mut no_demos = Candidate::new(); + no_demos.set_demos("drafter", Vec::new()); + assert_ne!(no_demos.stable_hash(), Candidate::new().stable_hash()); + + // to_params preserves the explicit clear. + let params = cleared.to_params(); + assert!(params.get("drafter").is_some()); + assert_eq!(params.get("drafter").unwrap().instruction_override, None); + } + #[test] fn candidate_round_trips_through_json() { let mut candidate = Candidate::new(); diff --git a/crates/dspy-rs/src/optimizer/gepa.rs b/crates/dspy-rs/src/optimizer/gepa.rs index d6ecc64c..d3f7712e 100644 --- a/crates/dspy-rs/src/optimizer/gepa.rs +++ b/crates/dspy-rs/src/optimizer/gepa.rs @@ -1,17 +1,17 @@ use anyhow::{Context, Result, anyhow}; use bon::Builder; -use rand::{Rng, SeedableRng, rngs::StdRng, seq::SliceRandom}; +use rand::{Rng, rngs::StdRng, seq::SliceRandom}; use serde::{Deserialize, Serialize}; +use crate::core::ToInput; use crate::evaluate::TypedMetric; use crate::optimizer::engine::{ - Budget, Candidate, CandidateEval, EngineConfig, EvalEngine, EvalOutcome, ParetoView, - apply_candidate, + Candidate, CandidateEval, Engine, EngineConfig, EvalOutcome, ParetoStatistics, ParetoView, }; -use crate::optimizer::{Optimizer, predictor_names, with_named_predictor}; +use crate::optimizer::target::LeafInfo; +use crate::optimizer::{OptimizeTarget, Optimizer, OptimizerCommon, Report}; use crate::utils::truncate; -use crate::core::ToInput; -use crate::{Facet, Module, Predict, Schema, Signature, SignatureSchema}; +use crate::{Module, Predict, Predictors, Signature}; /// Improve an LLM-pipeline module's instruction using execution feedback. /// @@ -40,30 +40,11 @@ struct ReflectOnInstruction { improved_instruction: String, } -/// Renders a predictor's input/output contract for reflection prompts. -/// Shared with [`SIMBA`](crate::SIMBA)'s introspection call. -pub(crate) fn format_schema_for_reflection(schema: &SignatureSchema) -> String { - let mut result = String::new(); - result.push_str("Input fields:\n"); - for field in schema.input_fields() { - let docs = if field.docs.is_empty() { - "No description" - } else { - field.docs.as_str() - }; - result.push_str(&format!(" - {}: {}\n", field.lm_name, docs)); - } - result.push_str("Output fields:\n"); - for field in schema.output_fields() { - let docs = if field.docs.is_empty() { - "No description" - } else { - field.docs.as_str() - }; - result.push_str(&format!(" - {}: {}\n", field.lm_name, docs)); - } - result -} +/// Character budget for the feedback appended by the *no-prompt-model* +/// fallback mutation. Without a cap the child instruction would grow by a +/// full trace dump every generation (unbounded quadratic growth across a +/// lineage); the reflection-LM path is unaffected. +const CONCAT_FEEDBACK_BUDGET: usize = 2000; /// A single instruction candidate tracked through GEPA's evolutionary search. /// @@ -124,14 +105,13 @@ pub struct GEPAResult { pub frontier_history: Vec, } -pub use super::pareto::ParetoStatistics; - /// Genetic-Pareto instruction optimizer with feedback-driven evolution. /// -/// GEPA is a thin strategy over the shared [`EvalEngine`]: candidates are -/// instruction overlays, evaluation is the engine's cached bounded-concurrency -/// fan-out, budgets are engine budgets, and the per-instance Pareto frontier is -/// a view over the engine's (candidates × examples) score matrix. +/// GEPA is a thin strategy over the shared [`Engine`]: candidates are +/// instruction [`Candidate`]s injected ambiently, evaluation is the engine's +/// cached bounded-concurrency fan-out, budgets are engine budgets, and the +/// per-instance Pareto frontier is a view over the engine's +/// (candidates × examples) score matrix. /// /// GEPA uses an evolutionary search guided by per-example feedback from your metric. /// Unlike [`COPRO`](crate::COPRO) which only uses numerical scores, GEPA requires your @@ -141,8 +121,9 @@ pub use super::pareto::ParetoStatistics; /// feedback, and the mutated component's per-invocation execution trace /// ([`Trace::for_component`](crate::Trace::for_component)), then writes an improved /// instruction each generation; without one, the feedback is appended to the -/// instruction as a deterministic mutation. Either way the quality of your feedback -/// directly determines the quality of GEPA's search. +/// instruction as a deterministic mutation (capped at a fixed character budget +/// so lineages can't grow unboundedly). Either way the quality of your +/// feedback directly determines the quality of GEPA's search. /// /// The Pareto frontier tracks candidates that aren't dominated on any individual /// training example, not just by average score. This means GEPA finds instructions @@ -151,6 +132,15 @@ pub use super::pareto::ParetoStatistics; /// Only searches instruction space — no demo mutation, no crossover between candidates. /// Each child has exactly one parent. /// +/// # Validation sets +/// +/// Build the target with +/// [`OptimizeTarget::module_with_valset`] (or use +/// [`compile_module_with_valset`](GEPA::compile_module_with_valset)): initial +/// evaluation and child scoring use the target's validation columns, parent +/// re-evaluation samples from its trainset columns. Without a valset the +/// trainset serves both roles. +/// /// # Hyperparameters /// /// - **`num_iterations`** (default: 20) — evolutionary generations. More = deeper search. @@ -166,7 +156,8 @@ pub use super::pareto::ParetoStatistics; /// - **`track_best_outputs`** (default: false) — re-run the best instruction on the /// eval set and record outputs. /// - **`prompt_model`** — reflection LM that rewrites instructions from feedback. -/// Strongly recommended; without it mutation degrades to feedback concatenation. +/// Strongly recommended; without it mutation degrades to (budget-capped) +/// feedback concatenation. /// - **`eval_concurrency`** (default: 16) — LM calls in flight during evaluation. /// - **`seed`** — fixes minibatch sampling and parent selection for reproducible runs. /// @@ -187,7 +178,7 @@ pub use super::pareto::ParetoStatistics; /// .num_iterations(20) /// .max_lm_calls(Some(500)) /// .build(); -/// let report = gepa.compile(&mut module, trainset, &feedback_metric).await?; +/// let report = gepa.compile_module(&mut module, &trainset, &feedback_metric).await?; /// println!("Best score: {:.3}", report.best_candidate.average_score()); /// ``` #[derive(Builder)] @@ -303,6 +294,16 @@ impl Lineage { } impl GEPA { + fn common(&self) -> OptimizerCommon { + OptimizerCommon { + eval_concurrency: self.eval_concurrency, + max_metric_calls: self.max_rollouts, + max_lm_calls: self.max_lm_calls, + seed: self.seed, + ..OptimizerCommon::default() + } + } + fn require_feedback(eval: &CandidateEval, module_name: &str, generation: usize) -> Result<()> { if eval .rollouts @@ -382,7 +383,9 @@ impl GEPA { /// Deterministic fallback mutation: append the feedback to the parent /// instruction. Used when no `prompt_model` is configured or the reflection - /// call fails. + /// call fails. The appended feedback is capped at + /// [`CONCAT_FEEDBACK_BUDGET`] characters so instructions can't grow by a + /// full trace dump every generation. fn concat_child_instruction( parent_instruction: &str, feedback_summary: &str, @@ -393,7 +396,7 @@ impl GEPA { "{}\n\n[GEPA gen {}] Improve based on feedback:\n{}\n(Parent score {:.3})", parent_instruction, generation + 1, - feedback_summary, + truncate(feedback_summary, CONCAT_FEEDBACK_BUDGET), parent_score, ) } @@ -404,20 +407,16 @@ impl GEPA { /// Returns the proposed instruction and the number of reflection LM calls /// consumed (0 or 1). Reflection failures degrade to the deterministic /// concatenation mutation with a warning rather than aborting the run. - #[allow(clippy::too_many_arguments)] - async fn propose_child_instruction( + async fn propose_child_instruction( &self, - module: &mut M, + leaves: &[LeafInfo], module_name: &str, parent_instruction: &str, feedback_summary: &str, parent_score: f64, generation: usize, reflector: Option<&Predict>, - ) -> (String, usize) - where - M: for<'a> Facet<'a>, - { + ) -> (String, usize) { let Some(reflector) = reflector else { return ( Self::concat_child_instruction( @@ -430,10 +429,11 @@ impl GEPA { ); }; - let task_description = with_named_predictor(module, module_name, |predictor| { - Ok(format_schema_for_reflection(predictor.schema())) - }) - .unwrap_or_default(); + let task_description = leaves + .iter() + .find(|leaf| leaf.name == module_name) + .map(LeafInfo::schema_for_reflection) + .unwrap_or_default(); let input = ReflectOnInstructionInput { task_description, @@ -477,116 +477,91 @@ impl GEPA { ) } - async fn collect_best_outputs( - module: &M, - eval_set: &[E], - ) -> Result> + /// Convenience: optimizes a typed module over a trainset with this + /// optimizer's default engine. + pub async fn compile_module( + &self, + module: &mut M, + trainset: &[E], + metric: &MT, + ) -> Result where - E: ToInput, - M: Module, - M::Output: Schema, + E: ToInput + serde::Serialize + Send + Sync, + M: Module + Predictors, + MT: TypedMetric, { - let mut outputs = Vec::with_capacity(eval_set.len()); - for example in eval_set { - let input = example.to_input()?; - let predicted = module.call(input).await.map_err(|err| anyhow!("{err}"))?; - outputs.push(serde_json::to_value(predicted.into_inner()).unwrap_or(serde_json::Value::Null)); - } - Ok(outputs) + self.compile_module_with_valset(module, trainset, None, metric) + .await } - /// Runs GEPA with an explicit validation set separate from the trainset. - /// - /// When `valset` is `Some`, initial evaluation and child scoring use the validation - /// set, while parent re-evaluation uses the trainset minibatch. When `None`, the - /// trainset serves both roles. - /// - /// # Errors - /// - /// - No optimizable predictors found - /// - Any metric evaluation returns `feedback: None` - /// - LM call failure during evaluation - pub async fn compile_with_valset( + /// [`compile_module`](Self::compile_module) with an explicit validation + /// set: initial evaluation and child scoring use the valset, parent + /// re-evaluation samples trainset minibatches. Sugar over + /// [`OptimizeTarget::module_with_valset`] + the [`Optimizer`] trait. + pub async fn compile_module_with_valset( &self, module: &mut M, - trainset: Vec, - valset: Option>, + trainset: &[E], + valset: Option<&[E]>, metric: &MT, ) -> Result where E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, + M: Module + Predictors, MT: TypedMetric, { - let predictor_names = predictor_names(module)?; - if predictor_names.is_empty() { + let mut target = OptimizeTarget::module_with_valset(module, trainset, valset, metric); + let mut engine = Engine::new(Optimizer::engine_config(self)); + let report = Optimizer::compile(self, &mut target, &mut engine).await?; + report + .into_gepa() + .ok_or_else(|| anyhow!("GEPA must return a GEPA report")) + } +} + +#[async_trait::async_trait(?Send)] +impl Optimizer for GEPA { + fn engine_config(&self) -> EngineConfig { + self.common().engine_config() + } + + async fn compile( + &self, + target: &mut OptimizeTarget<'_>, + engine: &mut Engine, + ) -> Result { + let leaves = target.leaves().to_vec(); + if leaves.is_empty() { return Err(anyhow!("no optimizable predictors found")); } - // One engine over one example universe: the validation columns come - // first (the Pareto/score columns), the trainset minibatch pool after - // them when a separate valset is supplied. - let train_len = trainset.len(); - let (engine_examples, val_cols, train_pool): (Vec, Vec, Vec) = - match valset { - Some(mut valset) => { - let val_len = valset.len(); - valset.extend(trainset); - ( - valset, - (0..val_len).collect(), - (val_len..val_len + train_len).collect(), - ) - } - None => ( - trainset, - (0..train_len).collect(), - (0..train_len).collect(), - ), - }; - - let mut engine = EvalEngine::new( - engine_examples, - metric, - EngineConfig { - concurrency: self.eval_concurrency, - budget: Budget { - max_metric_calls: self.max_rollouts, - max_lm_calls: self.max_lm_calls, - max_tokens: None, - }, - cache_salt: 0, - }, - ); + // The validation columns are the Pareto/score columns; the trainset + // columns are the minibatch pool (identical without a valset). + let val_cols = target.val_columns(); + let train_pool = target.train_columns(); let reflector = self .prompt_model .as_ref() .map(|lm| Predict::::builder().lm(lm.clone()).build()); - let mut rng = match self.seed { - Some(seed) => StdRng::seed_from_u64(seed), - None => StdRng::from_entropy(), - }; + let mut rng = self.common().rng(); let mut lineage = Lineage::new(); // Seed the frontier: each predictor's current instruction, scored on // the validation columns. - for module_name in &predictor_names { - let instruction = with_named_predictor(module, module_name, |predictor| { - Ok(predictor.instruction()) - })?; - let row = engine.register(Candidate::with_instruction(module_name, &instruction)); - let eval = match engine.evaluate(module, row, Some(&val_cols)).await? { + for leaf in &leaves { + let row = engine.register(Candidate::with_instruction(&leaf.name, &leaf.instruction)); + let eval = match engine.evaluate(target, row, Some(&val_cols)).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { .. } => break, }; - Self::require_feedback(&eval, module_name, 0)?; + Self::require_feedback(&eval, &leaf.name, 0)?; lineage.add( GEPACandidate { id: 0, - instruction, - module_name: module_name.clone(), + instruction: leaf.instruction.clone(), + module_name: leaf.name.clone(), example_scores: Vec::new(), parent_id: None, generation: 0, @@ -623,7 +598,7 @@ impl GEPA { .collect(); let parent_eval = match engine - .evaluate(module, parent_row, Some(&minibatch)) + .evaluate(target, parent_row, Some(&minibatch)) .await? { EvalOutcome::Complete(eval) => eval, @@ -641,7 +616,7 @@ impl GEPA { let (child_instruction, reflection_calls) = self .propose_child_instruction( - module, + &leaves, &parent.module_name, &parent.instruction, &feedback_summary, @@ -657,7 +632,7 @@ impl GEPA { &child.module_name, &child.instruction, )); - let child_eval = match engine.evaluate(module, child_row, Some(&val_cols)).await? { + let child_eval = match engine.evaluate(target, child_row, Some(&val_cols)).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { .. } => break, }; @@ -683,11 +658,8 @@ impl GEPA { .cloned() .context("no candidates available on Pareto frontier")?; - // Install the winner permanently through the one candidate seam. - let _undo = apply_candidate( - module, - &Candidate::with_instruction(&best_candidate.module_name, &best_candidate.instruction), - )?; + let winner = + Candidate::with_instruction(&best_candidate.module_name, &best_candidate.instruction); let highest_score_achieved_per_val_task: Vec = if lineage.entries.is_empty() { Vec::new() @@ -699,11 +671,10 @@ impl GEPA { .collect() }; - let eval_set = &engine.examples()[..val_cols.len()]; let best_outputs_valset = if self.track_best_outputs { - if !engine.budget_allows(eval_set.len()) { + if !engine.budget_allows(val_cols.len()) { tracing::debug!( - eval_examples = eval_set.len(), + eval_examples = val_cols.len(), spend = ?engine.spend(), max_lm_calls = ?self.max_lm_calls, max_rollouts = ?self.max_rollouts, @@ -711,15 +682,18 @@ impl GEPA { ); None } else { - let outputs = Self::collect_best_outputs::(module, eval_set).await?; - engine.charge(eval_set.len(), eval_set.len()); + let outputs = target.candidate_outputs(&val_cols, &winner).await?; + engine.charge(val_cols.len(), val_cols.len()); Some(outputs) } } else { None }; - Ok(GEPAResult { + // The one mutation of the run: install the winner. + target.install(&winner)?; + + Ok(Report::Gepa(GEPAResult { best_candidate, all_candidates, total_rollouts: engine.spend().metric_calls, @@ -728,26 +702,7 @@ impl GEPA { highest_score_achieved_per_val_task, best_outputs_valset, frontier_history, - }) - } -} - -impl Optimizer for GEPA { - type Report = GEPAResult; - - async fn compile( - &self, - module: &mut M, - trainset: Vec, - metric: &MT, - ) -> Result - where - E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, - MT: TypedMetric, - { - self.compile_with_valset(module, trainset, None, metric) - .await + })) } } @@ -758,7 +713,7 @@ mod tests { use super::*; use crate::evaluate::{Eval, TypedMetric}; use crate::trace::Trace; - use crate::{CallMetadata, Predict, PredictError, Predicted, Signature}; + use crate::{CallMetadata, Predict, PredictError, Predicted, PredictorInfo, Signature}; #[derive(Signature, Clone, Debug)] struct GepaStateSig { @@ -769,12 +724,12 @@ mod tests { answer: String, } - #[derive(facet::Facet)] - #[facet(crate = facet)] struct GepaStateModule { predictor: Predict, } + crate::predictors!(GepaStateModule { predictor }); + impl Module for GepaStateModule { type Input = GepaStateSigInput; type Output = GepaStateSigOutput; @@ -819,7 +774,7 @@ mod tests { } #[tokio::test] - async fn compile_restores_state_when_metric_errors() { + async fn compile_leaves_state_untouched_when_metric_errors() { let optimizer = GEPA::builder().num_iterations(1).minibatch_size(1).build(); let mut module = GepaStateModule { predictor: Predict::::builder() @@ -828,15 +783,22 @@ mod tests { }; let err = optimizer - .compile(&mut module, eval_set(), &AlwaysFailMetric) + .compile_module(&mut module, &eval_set(), &AlwaysFailMetric) .await .expect_err("candidate evaluation should propagate metric failure"); assert!(err.to_string().contains("metric failure")); - let instruction = with_named_predictor(&mut module, "predictor", |predictor| { - Ok(predictor.instruction()) - }) - .expect("predictor lookup should succeed"); - assert_eq!(instruction, "seed-instruction"); + assert_eq!( + PredictorInfo::instruction(&module.predictor), + "seed-instruction" + ); + } + + #[test] + fn concat_mutation_caps_appended_feedback() { + let huge = "x".repeat(50_000); + let child = GEPA::concat_child_instruction("seed", &huge, 0.5, 3); + assert!(child.len() < 3_000, "feedback must be capped: {}", child.len()); + assert!(child.starts_with("seed\n\n[GEPA gen 4]")); } } diff --git a/crates/dspy-rs/src/optimizer/harvest.rs b/crates/dspy-rs/src/optimizer/harvest.rs index d31dd715..818ef8ed 100644 --- a/crates/dspy-rs/src/optimizer/harvest.rs +++ b/crates/dspy-rs/src/optimizer/harvest.rs @@ -1,12 +1,19 @@ //! Demo harvesting from rollout traces — the trace name-join. //! //! A rollout trace records one span per `Predict` invocation under the same -//! component name the mutation seam addresses (the dotted path assigned by -//! `predictor_names`). Harvesting is therefore a pure name join: successful +//! component name candidates address (the leaf name declared through the +//! [`Predictors`](crate::Predictors) contract and stamped by the target's +//! naming pass). Harvesting is therefore a pure name join: successful //! spans from well-scored rollouts become few-shot demo rows for the predictor //! that produced them — no pointer identity, works identically for fx and //! struct harnesses. Shared by [`BootstrapFewShot`](crate::BootstrapFewShot) //! and [`MIPROv2`](crate::MIPROv2). +//! +//! Credit is per-span when available (RFC 0004 §4): a span carrying its own +//! [`Eval`](crate::Eval) — attached by +//! [`TypedMetric::evaluate_spans`](crate::evaluate::TypedMetric::evaluate_spans) +//! — is gated and ranked on that score instead of the whole-rollout score, so +//! a good final answer no longer vouches for every intermediate step. use std::collections::{HashMap, HashSet}; @@ -17,7 +24,8 @@ use crate::trace::{JsonMap, Trace}; /// A scored demo candidate: the flat demo row plus the input-only fingerprint /// used for deduplication. pub(crate) struct DemoCandidate { - /// Whole-rollout metric score of the trace this row came from. + /// Effective credit score: the span's own eval when the metric attached + /// one, the whole-rollout metric score otherwise. pub score: f64, /// Canonical fingerprint of the input fields, for input-level dedup. pub input_fingerprint: String, @@ -40,36 +48,44 @@ fn input_fingerprint(input: &JsonMap) -> String { /// Collects scored demo candidates per predictor name from rollout traces. /// -/// Every successful span (parsed output present) inside a rollout whose score -/// reaches `min_score` contributes one candidate to its component's bucket, -/// where the score is the *whole-rollout* metric score. +/// Every successful span (parsed output present) whose *effective* score +/// reaches `min_score` contributes one candidate to its component's bucket. +/// The effective score is the span's own [`Eval`](crate::Eval) when the +/// metric attached one via +/// [`TypedMetric::evaluate_spans`](crate::evaluate::TypedMetric::evaluate_spans), +/// the whole-rollout metric score otherwise — so a span the metric scored +/// badly is excluded even from a winning rollout, and a span it scored well +/// qualifies even from a losing one. pub(crate) fn collect_demo_candidates<'a>( rollouts: impl IntoIterator, min_score: f64, ) -> HashMap> { let mut candidates: HashMap> = HashMap::new(); - for (score, trace) in rollouts { - if score < min_score { - continue; - } + for (rollout_score, trace) in rollouts { for span in trace.successes() { + let score = span.eval.as_ref().map_or(rollout_score, |eval| eval.score); + if score < min_score { + continue; + } let name = trace.component_name(span.component); if let (Some(input), Some(output)) = (&span.input, &span.output) { - candidates.entry(name.to_string()).or_default().push( - DemoCandidate { + candidates + .entry(name.to_string()) + .or_default() + .push(DemoCandidate { score, input_fingerprint: input_fingerprint(input), row: demo_from_json(input, output), - }, - ); + }); } } } candidates } -/// Keeps the top `max_per_predictor` demos per predictor by rollout score, -/// deduplicated on input fields so repeated inputs don't crowd the demo set. +/// Keeps the top `max_per_predictor` demos per predictor by effective score +/// (span eval when present, rollout score otherwise), deduplicated on input +/// fields so repeated inputs don't crowd the demo set. pub(crate) fn select_demos( candidates: HashMap>, max_per_predictor: usize, @@ -99,3 +115,123 @@ pub(crate) fn select_demos( } selected } + +#[cfg(test)] +mod tests { + use super::*; + use crate::LmUsage; + use crate::trace::{CompId, Eval, ModelId, Span, SpanId, Trace}; + + fn json_map(pairs: &[(&str, &str)]) -> JsonMap { + pairs + .iter() + .map(|(k, v)| (k.to_string(), Value::String(v.to_string()))) + .collect() + } + + fn span(id: u32, component: u32, input: &str, output: &str, eval: Option) -> Span { + Span { + id: SpanId(id), + component: CompId(component), + seq: 0, + parent: None, + prefix: None, + suffix: Vec::new(), + input: Some(json_map(&[("prompt", input)])), + model: ModelId(0), + request_hash: 0, + events: Vec::new(), + raw_output: None, + output: Some(json_map(&[("answer", output)])), + usage: LmUsage::default(), + error: None, + eval, + started_at_us: 0, + duration_us: 0, + complete: true, + } + } + + fn trace(components: &[&str], spans: Vec) -> Trace { + Trace { + components: components.iter().map(|name| name.to_string()).collect(), + spans, + ..Trace::default() + } + } + + fn harvested_prompts(demos: &HashMap>, name: &str) -> Vec { + demos + .get(name) + .map(|rows| { + rows.iter() + .map(|row| row["prompt"].as_str().unwrap_or_default().to_string()) + .collect() + }) + .unwrap_or_default() + } + + #[test] + fn without_span_evals_rollout_score_gates_every_span() { + // Baseline behavior: no span evals anywhere, so the whole-rollout + // score decides for every span — 0.0 rollouts contribute nothing. + let good = trace( + &["draft", "refine"], + vec![span(0, 0, "q1", "a1", None), span(1, 1, "q1'", "a1'", None)], + ); + let bad = trace( + &["draft", "refine"], + vec![span(0, 0, "q2", "a2", None), span(1, 1, "q2'", "a2'", None)], + ); + + let demos = select_demos(collect_demo_candidates([(1.0, &good), (0.0, &bad)], 1.0), 4); + assert_eq!(harvested_prompts(&demos, "draft"), vec!["q1"]); + assert_eq!(harvested_prompts(&demos, "refine"), vec!["q1'"]); + } + + #[test] + fn span_eval_overrides_rollout_score_in_both_directions() { + // Winning rollout, but the metric scored the draft span 0.0 (a + // misstep the refine step recovered from) — the draft is excluded. + let recovered = trace( + &["draft", "refine"], + vec![ + span(0, 0, "q1", "wrong", Some(Eval::score(0.0))), + span(1, 1, "q1'", "right", None), + ], + ); + // Losing rollout, but the metric scored the draft span 1.0 — the + // draft qualifies anyway. + let salvaged = trace( + &["draft", "refine"], + vec![ + span(0, 0, "q2", "right", Some(Eval::score(1.0))), + span(1, 1, "q2'", "wrong", None), + ], + ); + + let demos = select_demos( + collect_demo_candidates([(1.0, &recovered), (0.0, &salvaged)], 1.0), + 4, + ); + assert_eq!(harvested_prompts(&demos, "draft"), vec!["q2"]); + assert_eq!(harvested_prompts(&demos, "refine"), vec!["q1'"]); + } + + #[test] + fn span_eval_score_ranks_candidates() { + // Both spans qualify; the span-scored one outranks the + // rollout-scored one when only one demo slot is available. + let modest = trace(&["draft"], vec![span(0, 0, "q1", "a1", None)]); + let strong = trace( + &["draft"], + vec![span(0, 0, "q2", "a2", Some(Eval::score(0.9)))], + ); + + let demos = select_demos( + collect_demo_candidates([(0.5, &modest), (0.5, &strong)], 0.5), + 1, + ); + assert_eq!(harvested_prompts(&demos, "draft"), vec!["q2"]); + } +} diff --git a/crates/dspy-rs/src/optimizer/mipro.rs b/crates/dspy-rs/src/optimizer/mipro.rs index 1683ac3d..207ea268 100644 --- a/crates/dspy-rs/src/optimizer/mipro.rs +++ b/crates/dspy-rs/src/optimizer/mipro.rs @@ -1,17 +1,16 @@ use anyhow::{Result, anyhow}; use bon::Builder; -use rand::{SeedableRng, rngs::StdRng, seq::SliceRandom}; +use rand::seq::SliceRandom; use tracing::debug; +use crate::core::ToInput; use crate::evaluate::TypedMetric; -use crate::optimizer::engine::{ - Budget, Candidate, EngineConfig, EvalEngine, EvalOutcome, apply_candidate, -}; +use crate::optimizer::engine::{Candidate, Engine, EngineConfig, EvalOutcome}; use crate::optimizer::harvest::{collect_demo_candidates, select_demos}; -use crate::optimizer::{Optimizer, predictor_names, with_named_predictor}; +use crate::optimizer::target::LeafInfo; +use crate::optimizer::{OptimizeTarget, Optimizer, OptimizerCommon, Report}; use crate::trace::Trace; -use crate::core::ToInput; -use crate::{Facet, Module, SignatureSchema}; +use crate::{Module, Predictors}; /// The whole-program score recorded on a rollout trace, if a metric ran. fn trace_score(trace: &Trace) -> Option { @@ -22,30 +21,6 @@ fn trace_score(trace: &Trace) -> Option { .map(|eval| eval.score) } -/// An instruction candidate with its evaluated score. -/// -/// Generated by [`MIPROv2`]'s candidate generation step, then scored by -/// evaluating the module with this instruction on a minibatch. -#[derive(Clone, Debug)] -pub struct PromptCandidate { - pub instruction: String, - pub score: f64, -} - -impl PromptCandidate { - pub fn new(instruction: String) -> Self { - Self { - instruction, - score: 0.0, - } - } - - pub fn with_score(mut self, score: f64) -> Self { - self.score = score; - self - } -} - /// Library of general prompting best practices used to seed candidate generation. /// /// These tips are appended to candidate instructions during [`MIPROv2`] optimization @@ -87,23 +62,47 @@ impl PromptingTips { } } +/// Renders a leaf's field contract in MIPRO's program-description format. +fn format_leaf_fields(leaf: &LeafInfo) -> String { + let mut result = String::new(); + + result.push_str("Input Fields:\n"); + for (name, docs) in &leaf.input_fields { + let desc = if docs.is_empty() { "No description" } else { docs }; + result.push_str(&format!(" - {name}: {desc}\n")); + } + + result.push_str("\nOutput Fields:\n"); + for (name, docs) in &leaf.output_fields { + let desc = if docs.is_empty() { "No description" } else { docs }; + result.push_str(&format!(" - {name}: {desc}\n")); + } + + result +} + /// Trace-guided instruction and demo optimizer. /// /// MIPROv2 (Multi-prompt Instruction PRoposal Optimizer v2) is a thin strategy -/// over the shared [`EvalEngine`], working in four phases: +/// over the shared [`Engine`], working in four phases: /// /// 1. **Trace collection** — one traced teacher pass over the trainset (the /// engine's baseline candidate), collecting whole-program scores plus /// per-`Predict` input/output spans -/// 2. **Demo bootstrapping** — successful spans from rollouts scoring at least -/// `min_demo_score` become few-shot demos on the predictor that produced -/// them via the trace name-join (top `max_bootstrapped_demos` by score, -/// deduplicated on inputs), installed through the one candidate seam +/// 2. **Demo bootstrapping** — successful spans scoring at least +/// `min_demo_score` (their own span eval when the metric attached one, +/// the rollout score otherwise) become few-shot demos on the predictor +/// that produced them via the trace name-join (top +/// `max_bootstrapped_demos` by score, deduplicated on inputs), folded +/// into the accumulating winner candidate /// 3. **Candidate generation** — uses the traces and prompting tips to generate /// `num_candidates` instruction variants per predictor /// 4. **Trial evaluation** — evaluates up to `num_trials` candidates on a sampled /// minibatch through the engine (cached, budget-metered fan-out), keeps the best /// +/// The accumulated winner (demos + best instructions) is installed once at +/// the end through [`OptimizeTarget::install`]. +/// /// Unlike [`GEPA`](crate::GEPA), MIPROv2 does not require feedback — only numerical scores. /// Unlike [`COPRO`](crate::COPRO), it uses execution traces to inform candidate generation /// rather than blind search. @@ -115,7 +114,8 @@ impl PromptingTips { /// If `num_trials` < `num_candidates`, only the first `num_trials` are evaluated. /// - **`minibatch_size`** (default: 25) — examples per candidate evaluation. /// - **`max_bootstrapped_demos`** (default: 4) — demos installed per predictor. -/// - **`min_demo_score`** (default: 0.0) — score gate for demo-eligible traces. +/// - **`min_demo_score`** (default: 0.0) — score gate for demo-eligible spans +/// (span eval when present, rollout score otherwise). /// - **`eval_concurrency`** (default: 16) — LM calls in flight during evaluation. /// - **`seed`** — fixes minibatch sampling for reproducible runs. /// @@ -135,7 +135,7 @@ impl PromptingTips { /// .num_candidates(10) /// .num_trials(20) /// .build(); -/// mipro.compile(&mut module, trainset, &metric).await?; +/// mipro.compile_module(&mut module, &trainset, &metric).await?; /// ``` #[derive(Builder)] pub struct MIPROv2 { @@ -155,8 +155,10 @@ pub struct MIPROv2 { #[builder(default = 4)] pub max_bootstrapped_demos: usize, - /// Minimum whole-program score a trace needs for its per-predictor - /// input/output pairs to qualify as bootstrapped demos. + /// Minimum score a span needs to qualify as a bootstrapped demo: its own + /// eval when the metric attached one + /// ([`TypedMetric::evaluate_spans`](crate::evaluate::TypedMetric::evaluate_spans)), + /// the whole-program score otherwise. #[builder(default = 0.0)] pub min_demo_score: f64, @@ -169,23 +171,12 @@ pub struct MIPROv2 { } impl MIPROv2 { - /// Rollout traces with the highest recorded scores, descending. Traces - /// without a recorded eval are ignored. - pub fn select_best_traces<'a>(&self, traces: &'a [Trace], num_select: usize) -> Vec<&'a Trace> { - let mut scored_traces: Vec<_> = traces - .iter() - .filter_map(|t| trace_score(t).map(|score| (score, t))) - .collect(); - - scored_traces.sort_by(|(left, _), (right, _)| { - right.partial_cmp(left).unwrap_or(std::cmp::Ordering::Equal) - }); - - scored_traces - .into_iter() - .take(num_select) - .map(|(_, t)| t) - .collect() + fn common(&self) -> OptimizerCommon { + OptimizerCommon { + eval_concurrency: self.eval_concurrency, + seed: self.seed, + ..OptimizerCommon::default() + } } fn generate_candidate_instructions( @@ -210,82 +201,54 @@ impl MIPROv2 { .collect() } - pub fn create_prompt_candidates(&self, instructions: Vec) -> Vec { - instructions.into_iter().map(PromptCandidate::new).collect() - } - - pub fn format_schema_fields(&self, signature: &SignatureSchema) -> String { - let mut result = String::new(); - - result.push_str("Input Fields:\n"); - for field in signature.input_fields() { - let desc = if field.docs.is_empty() { - "No description" - } else { - field.docs.as_str() - }; - result.push_str(&format!(" - {}: {}\n", field.lm_name, desc)); - } - - result.push_str("\nOutput Fields:\n"); - for field in signature.output_fields() { - let desc = if field.docs.is_empty() { - "No description" - } else { - field.docs.as_str() - }; - result.push_str(&format!(" - {}: {}\n", field.lm_name, desc)); - } - - result - } -} - -impl Optimizer for MIPROv2 { - type Report = (); - - async fn compile( + /// Convenience: optimizes a typed module over a trainset with this + /// optimizer's default engine, installing demos + winning instructions. + pub async fn compile_module( &self, module: &mut M, - trainset: Vec, + trainset: &[E], metric: &MT, - ) -> Result + ) -> Result<()> where E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, + M: Module + Predictors, MT: TypedMetric, { - let predictor_names = predictor_names(module)?; + let mut target = OptimizeTarget::module(module, trainset, metric); + let mut engine = Engine::new(Optimizer::engine_config(self)); + Optimizer::compile(self, &mut target, &mut engine).await?; + Ok(()) + } +} - if predictor_names.is_empty() { +#[async_trait::async_trait(?Send)] +impl Optimizer for MIPROv2 { + fn engine_config(&self) -> EngineConfig { + self.common().engine_config() + } + + async fn compile( + &self, + target: &mut OptimizeTarget<'_>, + engine: &mut Engine, + ) -> Result { + let leaves = target.leaves().to_vec(); + if leaves.is_empty() { return Err(anyhow!("no optimizable predictors found")); } - let mut rng = match self.seed { - Some(seed) => StdRng::seed_from_u64(seed), - None => StdRng::from_entropy(), - }; - - let mut engine = EvalEngine::new( - trainset, - metric, - EngineConfig { - concurrency: self.eval_concurrency, - budget: Budget::unlimited(), - cache_salt: 0, - }, - ); + let mut rng = self.common().rng(); // Phase 1: one traced teacher pass over the trainset — the engine's // baseline candidate. Whole-program traces feed candidate generation; // per-predictor spans feed demo bootstrapping via the trace name-join - // (spans record the dotted paths assigned by the naming pass above). + // (spans record the names stamped by the target's naming pass). let baseline = engine.register(Candidate::new()); - let baseline_eval = match engine.evaluate(module, baseline, None).await? { + let baseline_eval = match engine.evaluate(target, baseline, None).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { needed } => { return Err(anyhow!( - "unexpected budget exhaustion ({needed} rollouts) with an unlimited budget" + "budget exhausted ({needed} rollouts needed) during the teacher pass" )); } }; @@ -301,23 +264,20 @@ impl Optimizer for MIPROv2 { .collect(); // Phase 2: bootstrap demos from successful spans of well-scored - // rollouts and install them permanently before instruction search, so - // candidates are scored against the module as it will actually run. + // rollouts. They fold into the accumulating winner so instruction + // candidates are scored against the demos they will ship with. + let mut current = Candidate::new(); let bootstrapped_demos = select_demos( collect_demo_candidates(scored_traces, self.min_demo_score), self.max_bootstrapped_demos, ); - if !bootstrapped_demos.is_empty() { - let mut demo_candidate = Candidate::new(); - for (predictor_name, demos) in bootstrapped_demos { - debug!( - predictor = %predictor_name, - demo_count = demos.len(), - "installing bootstrapped demos" - ); - demo_candidate.set_demos(predictor_name, demos); - } - let _undo = apply_candidate(module, &demo_candidate)?; + for (predictor_name, demos) in bootstrapped_demos { + debug!( + predictor = %predictor_name, + demo_count = demos.len(), + "bootstrapping demos into the candidate" + ); + current.set_demos(predictor_name, demos); } let traces: Vec = baseline_eval @@ -329,14 +289,9 @@ impl Optimizer for MIPROv2 { // Phase 3: per-predictor instruction search on a sampled minibatch — // one minibatch per predictor round so all candidates score on the // same examples and remain comparable. - let all_indices: Vec = (0..engine.num_examples()).collect(); - for predictor_name in predictor_names { - let signature_desc = { - with_named_predictor(module, &predictor_name, |predictor| { - Ok(self.format_schema_fields(predictor.schema())) - })? - }; - + let all_indices: Vec = (0..target.num_examples()).collect(); + for leaf in &leaves { + let signature_desc = format_leaf_fields(leaf); let instructions = self.generate_candidate_instructions(&signature_desc, &traces, self.num_candidates); @@ -348,13 +303,14 @@ impl Optimizer for MIPROv2 { let mut best: Option<(f64, String)> = None; for instruction in instructions.into_iter().take(self.num_trials.max(1)) { - let row = - engine.register(Candidate::with_instruction(&predictor_name, &instruction)); - let eval = match engine.evaluate(module, row, Some(&minibatch)).await? { + let mut candidate = current.clone(); + candidate.set_instruction(&leaf.name, &instruction); + let row = engine.register(candidate); + let eval = match engine.evaluate(target, row, Some(&minibatch)).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { needed } => { return Err(anyhow!( - "unexpected budget exhaustion ({needed} rollouts) with an unlimited budget" + "budget exhausted ({needed} rollouts needed) during trial evaluation" )); } }; @@ -366,13 +322,13 @@ impl Optimizer for MIPROv2 { let (_, best_instruction) = best.ok_or_else(|| anyhow!("no candidates to evaluate"))?; - let _undo = apply_candidate( - module, - &Candidate::with_instruction(&predictor_name, best_instruction), - )?; + current.set_instruction(&leaf.name, best_instruction); } - Ok(()) + // The one mutation of the run: install demos + winning instructions. + target.install(¤t)?; + + Ok(Report::None) } } @@ -382,7 +338,7 @@ mod tests { use super::*; use crate::evaluate::{Eval, TypedMetric}; - use crate::{CallMetadata, Predict, PredictError, Predicted, Signature}; + use crate::{CallMetadata, Predict, PredictError, Predicted, PredictorInfo, Signature}; #[derive(Signature, Clone, Debug)] struct MiproStateSig { @@ -393,12 +349,12 @@ mod tests { answer: String, } - #[derive(facet::Facet)] - #[facet(crate = facet)] struct MiproStateModule { predictor: Predict, } + crate::predictors!(MiproStateModule { predictor }); + impl Module for MiproStateModule { type Input = MiproStateSigInput; type Output = MiproStateSigOutput; @@ -443,7 +399,7 @@ mod tests { } #[tokio::test] - async fn compile_restores_state_when_metric_errors() { + async fn compile_leaves_state_untouched_when_metric_errors() { let optimizer = MIPROv2::builder() .num_candidates(2) .num_trials(1) @@ -456,15 +412,14 @@ mod tests { }; let err = optimizer - .compile(&mut module, trainset(), &AlwaysFailMetric) + .compile_module(&mut module, &trainset(), &AlwaysFailMetric) .await .expect_err("candidate evaluation should propagate metric failure"); assert!(err.to_string().contains("metric failure")); - let instruction = with_named_predictor(&mut module, "predictor", |predictor| { - Ok(predictor.instruction()) - }) - .expect("predictor lookup should succeed"); - assert_eq!(instruction, "seed-instruction"); + assert_eq!( + PredictorInfo::instruction(&module.predictor), + "seed-instruction" + ); } } diff --git a/crates/dspy-rs/src/optimizer/mod.rs b/crates/dspy-rs/src/optimizer/mod.rs index 356d256c..a5c0a587 100644 --- a/crates/dspy-rs/src/optimizer/mod.rs +++ b/crates/dspy-rs/src/optimizer/mod.rs @@ -1,28 +1,35 @@ //! Automatic prompt optimization. //! -//! An optimizer takes a module, a training set, and a metric, then searches for better -//! instructions (and in some cases, demos) for each [`Predict`](crate::Predict) leaf. -//! The module is mutated in-place — after optimization, calling it produces better results -//! without any code changes. +//! An optimizer takes an [`OptimizeTarget`] (a typed module + trainset + +//! metric, or an interpreter-loaded IR program + examples + metric), searches +//! for better instructions (and in some cases demos) for each optimizable +//! leaf, and returns a [`Report`]. Candidates are **data, never mutation**: //! -//! The [`Optimizer::compile`] method takes `&mut module` (exclusive access — no concurrent -//! `call()` during optimization) and returns a report. The specific report type depends -//! on the optimizer: [`COPRO`] returns `()`, [`GEPA`] returns [`GEPAResult`] with full -//! evolution history, [`MIPROv2`] returns `()`. +//! 1. Leaves are discovered explicitly — modules declare them via +//! [`Predictors`](crate::Predictors) (see the `predictors!` macro); the +//! target snapshots their names, schemas, and current values as +//! [`LeafInfo`]s and stamps each leaf's trace name once per run. +//! 2. Each candidate is a name-keyed [`Candidate`] evaluated on the shared +//! [`Engine`] — cached, budget-metered, bounded-concurrency traced +//! rollouts with the candidate injected *ambiently* per rollout +//! ([`fx::with_params`](crate::fx::with_params)); different candidates +//! evaluate concurrently because nothing is ever applied to shared state. +//! 3. The winner is installed exactly once at the end +//! ([`OptimizeTarget::install`]) — the module lane's one mutation; the +//! program lane's winner is an [`ir::Overlay`](crate::ir::Overlay) for +//! [`Program::bake`](crate::ir::Program::bake). //! -//! # How it works internally +//! The convenience entry point for the common case is each optimizer's +//! `compile_module` inherent method: //! -//! 1. The optimizer calls `visit_named_predictors_mut` to discover all `Predict` -//! leaves via Facet reflection -//! 2. For each leaf, it reads the current instruction and generates candidates -//! 3. Each candidate becomes an overlay [`Candidate`] evaluated on the shared -//! [`EvalEngine`] — cached, budget-metered, bounded-concurrency traced -//! rollouts through the `DynPredictor::apply_update` mutation seam -//! 4. The best candidate (per optimizer's strategy) is installed through -//! [`apply_candidate`] +//! ```ignore +//! let copro = COPRO::builder().breadth(10).depth(3).build(); +//! copro.compile_module(&mut module, &trainset, &metric).await?; +//! // module is now optimized — call it as usual +//! ``` //! -//! Users never see this machinery — they call `optimizer.compile(&mut module, trainset, &metric)` -//! and their module gets better. +//! For composition (`Box` pipelines sharing one engine budget) +//! use the object-safe [`Optimizer`] trait directly. //! //! # Choosing an optimizer //! @@ -33,6 +40,7 @@ //! | [`SIMBA`] | Minibatch introspective ascent (demos + rules) | No | Low (steps × minibatch) | //! | [`GEPA`] | Genetic-Pareto evolution with feedback | **Yes** | Medium-high (iterations × eval) | //! | [`MIPROv2`] | Trace-guided candidate generation | No | Medium (candidates × trials × trainset) | +//! | [`Structural`] | LM-guided graph edits over [`ir::Edit`](crate::ir::Edit) (program lane only) | No | Medium (examples + iterations × minibatch) | pub mod bootstrap; pub mod copro; @@ -40,105 +48,121 @@ pub mod engine; pub mod gepa; pub(crate) mod harvest; pub mod mipro; -pub mod pareto; -#[cfg(feature = "ir")] -pub mod program_engine; pub mod simba; +pub mod structural; +pub mod target; pub use bootstrap::*; pub use copro::*; pub use engine::*; pub use gepa::*; pub use mipro::*; -pub use pareto::*; -#[cfg(feature = "ir")] -pub use program_engine::*; pub use simba::*; +pub use structural::*; +pub use target::{LeafInfo, OptimizeTarget, ProgramMetric}; use anyhow::Result; -use anyhow::anyhow; -use std::ops::ControlFlow; +use rand::{SeedableRng, rngs::StdRng}; -use crate::core::{DynPredictor, ToInput, visit_named_predictors_mut}; -use crate::evaluate::TypedMetric; -use crate::{Facet, Module}; - -/// Tunes a module's [`Predict`](crate::Predict) leaves for better performance. -/// -/// Takes exclusive `&mut` access to the module during optimization — you cannot call -/// the module concurrently. After `compile` returns, the module's instructions and/or -/// demos have been mutated in-place. Just call the module as before; no code changes needed. +/// A tuning strategy over the shared [`Engine`]. /// -/// ```ignore -/// let optimizer = COPRO::builder().breadth(10).depth(3).build(); -/// optimizer.compile(&mut module, trainset, &metric).await?; -/// // module is now optimized — call it as usual -/// let result = module.call(input).await?; -/// ``` +/// Object-safe by design: optimizers compose (`Box` pipelines +/// can share one [`Engine`] — one budget, one rollout cache, one score +/// matrix — across stages). The target carries the thing under optimization +/// and its example set *by reference*; the engine carries the spend. /// /// # Errors /// /// Returns an error if: -/// - No optimizable `Predict` leaves are found in the module -/// - The metric evaluation fails on any training example +/// - The target has no optimizable leaves +/// - The metric evaluation fails on any example /// - An LM call fails during candidate evaluation -#[allow(async_fn_in_trait)] -pub trait Optimizer { - type Report; +#[async_trait::async_trait(?Send)] +pub trait Optimizer: Send + Sync { + /// The engine configuration this optimizer wants when the caller doesn't + /// supply an engine explicitly (the `compile_module` convenience path). + fn engine_config(&self) -> EngineConfig { + EngineConfig::default() + } - async fn compile( + /// Runs the strategy: proposes candidates, evaluates them through + /// `engine`, installs the winner on `target`, and reports what happened. + async fn compile( &self, - module: &mut M, - trainset: Vec, - metric: &MT, - ) -> Result - where - E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, - MT: TypedMetric; + target: &mut OptimizeTarget<'_>, + engine: &mut Engine, + ) -> Result; } -/// Returns the dotted-path names of all [`Predict`](crate::Predict) leaves in a -/// module, and assigns each leaf its path as trace-span component name. -/// -/// The naming pass is what joins traces back to predictors: spans record the -/// same string [`with_named_predictor`] addresses, so demo harvesting and -/// per-component reflection need no pointer-identity bookkeeping. -pub(crate) fn predictor_names(module: &mut M) -> Result> -where - M: for<'a> Facet<'a>, -{ - let mut names = Vec::new(); - visit_named_predictors_mut(module, |name, predictor| { - predictor.set_trace_name(name); - names.push(name.to_string()); - ControlFlow::Continue(()) - })?; - Ok(names) +/// What an optimization run produced. Strategy-specific payloads for the +/// optimizers that report more than "done". +#[derive(Clone, Debug)] +pub enum Report { + /// Nothing beyond the installed winner (COPRO, MIPROv2). + None, + Gepa(GEPAResult), + Simba(SimbaReport), + Bootstrap(BootstrapReport), + /// Extension point for third-party strategies. + Custom(serde_json::Value), } -/// Looks up a single named predictor and applies a closure to it. -/// -/// # Errors -/// -/// Returns an error if the predictor name doesn't match any discovered leaf. -pub(crate) fn with_named_predictor(module: &mut M, predictor_name: &str, f: F) -> Result -where - M: for<'a> Facet<'a>, - F: FnOnce(&mut dyn DynPredictor) -> Result, -{ - let mut apply = Some(f); - let mut result = None; +impl Report { + pub fn into_gepa(self) -> Option { + match self { + Self::Gepa(report) => Some(report), + _ => None, + } + } - visit_named_predictors_mut(module, |name, predictor| { - if name != predictor_name { - return ControlFlow::Continue(()); + pub fn into_simba(self) -> Option { + match self { + Self::Simba(report) => Some(report), + _ => None, } + } - let f = apply.take().expect("selector closure should only run once"); - result = Some(f(predictor)); - ControlFlow::Break(()) - })?; + pub fn into_bootstrap(self) -> Option { + match self { + Self::Bootstrap(report) => Some(report), + _ => None, + } + } +} + +/// The engine/RNG knobs shared by every optimizer builder: evaluation +/// concurrency, budget caps, cache salt, and the sampling seed. Each +/// optimizer assembles one from its builder fields; [`engine_config`] +/// and [`rng`] replace the per-strategy construction boilerplate. +/// +/// [`engine_config`]: OptimizerCommon::engine_config +/// [`rng`]: OptimizerCommon::rng +#[derive(Clone, Copy, Debug, Default)] +pub struct OptimizerCommon { + pub eval_concurrency: usize, + pub max_metric_calls: Option, + pub max_lm_calls: Option, + pub cache_salt: u64, + pub seed: Option, +} - result.unwrap_or_else(|| Err(anyhow!("predictor `{predictor_name}` not found"))) +impl OptimizerCommon { + pub fn engine_config(&self) -> EngineConfig { + EngineConfig { + concurrency: self.eval_concurrency.max(1), + budget: Budget { + max_metric_calls: self.max_metric_calls, + max_lm_calls: self.max_lm_calls, + max_tokens: None, + }, + cache_salt: self.cache_salt, + } + } + + pub fn rng(&self) -> StdRng { + match self.seed { + Some(seed) => StdRng::seed_from_u64(seed), + None => StdRng::from_entropy(), + } + } } diff --git a/crates/dspy-rs/src/optimizer/pareto.rs b/crates/dspy-rs/src/optimizer/pareto.rs deleted file mode 100644 index 9a31bc44..00000000 --- a/crates/dspy-rs/src/optimizer/pareto.rs +++ /dev/null @@ -1,192 +0,0 @@ -use rand::Rng; -use serde::{Deserialize, Serialize}; - -use crate::optimizer::engine::{SCORE_EPS, ScoreMatrix}; -use crate::optimizer::gepa::GEPACandidate; - -/// Per-example dominance frontier for candidate selection. -/// -/// The key insight: optimizing for average score across examples lets the optimizer -/// overfit to easy examples while ignoring hard ones. The Pareto frontier prevents -/// this by keeping every candidate that's the *best on at least one example*. A -/// candidate that scores 0.3 average but is the only one to crack example #7 stays -/// on the frontier alongside a candidate that scores 0.9 average but fails #7. -/// -/// This is a standalone convenience wrapper over the engine's -/// [`ScoreMatrix`](crate::ScoreMatrix)/[`ParetoView`](crate::ParetoView) -/// bookkeeping — one dominance implementation, two entry points. -/// [`GEPA`](crate::GEPA) itself uses the engine's matrix directly (its scores -/// already live there); use `ParetoFrontier` when you track candidate payloads -/// outside an engine. -/// -/// Parents are sampled proportional to coverage (how many examples they win -/// on), so well-rounded candidates get sampled more often but specialists -/// aren't eliminated. Candidates that become dominated on every example are -/// pruned automatically. -#[derive(Debug, Clone, Default)] -pub struct ParetoFrontier { - /// All recorded scores, including rows of since-pruned candidates (their - /// columns' maxima are always matched by a surviving candidate, so keeping - /// them never changes dominance). - matrix: ScoreMatrix, - /// Surviving (frontier) candidates, in insertion order. - candidates: Vec, - /// Matrix row per surviving candidate, parallel to `candidates`. - rows: Vec, - /// Next candidate ID to assign. - next_id: usize, -} - -impl ParetoFrontier { - pub fn new() -> Self { - Self::default() - } - - pub fn len(&self) -> usize { - self.candidates.len() - } - - pub fn is_empty(&self) -> bool { - self.candidates.is_empty() - } - - pub fn candidates(&self) -> &[GEPACandidate] { - &self.candidates - } - - /// Adds a candidate if it achieves the best score on at least one example. - /// - /// Returns `true` if the candidate made it onto the frontier (won or tied on - /// at least one example). Candidates already on the frontier that no longer - /// win on any example are pruned. - pub fn add_candidate(&mut self, mut candidate: GEPACandidate, scores: &[f32]) -> bool { - candidate.id = self.next_id; - self.next_id += 1; - candidate.example_scores = scores.to_vec(); - - // Does it win or tie anywhere against the current frontier? - let best = self.matrix.pareto(); - let wins_somewhere = scores.iter().enumerate().any(|(example, &score)| { - match best.best_scores().get(example).copied().flatten() { - Some(best) => f64::from(score) + SCORE_EPS >= best, - None => true, - } - }); - if !wins_somewhere { - return false; - } - - let row = self.matrix.candidates(); - for (example, &score) in scores.iter().enumerate() { - self.matrix.record(row, example, f64::from(score)); - } - self.candidates.push(candidate); - self.rows.push(row); - - // Prune candidates the new arrival dominated everywhere. - let view = self.matrix.pareto(); - let mut idx = 0; - while idx < self.candidates.len() { - if view.wins(self.rows[idx]) == 0 { - self.candidates.remove(idx); - self.rows.remove(idx); - } else { - idx += 1; - } - } - - true - } - - /// Samples a parent candidate, weighted by how many examples it wins on. - /// - /// Well-rounded candidates get sampled more often, but specialists that only - /// win on one hard example still get a chance. This prevents the search from - /// collapsing onto a single high-average candidate. - pub fn sample_proportional_to_coverage(&self) -> Option<&GEPACandidate> { - if self.candidates.is_empty() { - return None; - } - - let view = self.matrix.pareto(); - let coverages: Vec = self.rows.iter().map(|&row| view.wins(row)).collect(); - let total_coverage: usize = coverages.iter().sum(); - - if total_coverage == 0 { - // Fallback to uniform sampling - return self.candidates.first(); - } - - let mut rng = rand::thread_rng(); - let mut target = rng.gen_range(0..total_coverage); - - for (candidate, &coverage) in self.candidates.iter().zip(coverages.iter()) { - if target < coverage { - return Some(candidate); - } - target -= coverage; - } - - // Fallback (shouldn't happen) - self.candidates.last() - } - - /// Returns the candidate with the highest average score across all examples. - /// - /// The Pareto frontier preserves diversity during search, but the winner is - /// still picked by average. - pub fn best_by_average(&self) -> Option<&GEPACandidate> { - self.candidates.iter().max_by(|a, b| { - let avg_a = a.average_score(); - let avg_b = b.average_score(); - avg_a.partial_cmp(&avg_b).unwrap_or(std::cmp::Ordering::Equal) - }) - } - - pub fn statistics(&self) -> ParetoStatistics { - let view = self.matrix.pareto(); - let coverage_per_candidate: Vec = - self.rows.iter().map(|&row| view.wins(row)).collect(); - - let avg_coverage = if !coverage_per_candidate.is_empty() { - coverage_per_candidate.iter().sum::() as f32 - / coverage_per_candidate.len() as f32 - } else { - 0.0 - }; - - ParetoStatistics { - num_candidates: self.candidates.len(), - num_examples_covered: view - .best_scores() - .iter() - .filter(|best| best.is_some()) - .count(), - avg_coverage, - max_coverage: coverage_per_candidate.iter().copied().max().unwrap_or(0), - min_coverage: coverage_per_candidate.iter().copied().min().unwrap_or(0), - } - } -} - -/// Snapshot of the Pareto frontier at a point in the search. -/// -/// Useful for plotting convergence. A healthy search has `num_candidates` growing -/// slowly (diversity is maintained) while `avg_coverage` increases (candidates are -/// getting more robust). If `num_candidates` is 1, the search has collapsed. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ParetoStatistics { - /// Candidates currently on the frontier. 1 means the search has converged - /// (or collapsed) to a single instruction. - pub num_candidates: usize, - /// Examples where at least one frontier candidate is the best. Should approach - /// total eval set size as the search progresses. - pub num_examples_covered: usize, - /// Mean examples won per candidate. Higher means candidates are more robust; - /// lower means more specialization. - pub avg_coverage: f32, - /// Most examples won by any single candidate. - pub max_coverage: usize, - /// Fewest examples won by any frontier candidate (always >= 1 by construction). - pub min_coverage: usize, -} diff --git a/crates/dspy-rs/src/optimizer/program_engine.rs b/crates/dspy-rs/src/optimizer/program_engine.rs deleted file mode 100644 index 0dff8ee8..00000000 --- a/crates/dspy-rs/src/optimizer/program_engine.rs +++ /dev/null @@ -1,490 +0,0 @@ -//! The IR-native evaluation path (RFC 0002 IR-6): candidate [`Overlay`]s -//! evaluated over one shared `Arc` through the [`Interpreter`] with -//! **candidate-level parallelism**. -//! -//! This is the seam `engine.rs` deliberately left open: the module lane must -//! serialize candidate application because candidates mutate shared predictor -//! state through `apply_update`. The dynamic lane has no mutation at all — -//! the interpreter reads instruction/demos/model/context/code through the -//! overlay at render time — so N candidates × M examples fan out in **one** -//! bounded-concurrency stream over one program instance. Nothing is applied, -//! nothing is restored. -//! -//! [`ProgramEvalEngine`] mirrors the module-lane [`EvalEngine`] machinery -//! piece for piece — same [`EngineConfig`] (concurrency/budget/salt), same -//! [`RolloutCache`] (keys are `(program hash, overlay hash, example uid, -//! salt)`, so per-candidate hits are exact), same [`ScoreMatrix`]/Pareto -//! bookkeeping, same [`Spend`] accounting, same minibatch gate — but speaks -//! JSON at the boundary: examples are labeled [`DemoRow`]s and the metric is -//! a [`ProgramMetric`] over output `JsonMap`s. The module-lane path is -//! untouched; this is an *additional* entry point -//! ([`evaluate_program_candidates`](ProgramEvalEngine::evaluate_program_candidates)), -//! not a rewrite. -//! -//! [`EvalEngine`]: crate::optimizer::engine::EvalEngine - -use std::collections::HashMap; -use std::sync::atomic::{AtomicUsize, Ordering}; -use std::sync::{Arc, Mutex}; -use std::time::Instant; - -use anyhow::{Result, anyhow}; -use futures::stream::{self, StreamExt, TryStreamExt}; - -use crate::evaluate::Eval; -use crate::ir::interp::{Budget as RunBudget, Interpreter}; -use crate::ir::params::{DemoRow, Overlay}; -use crate::optimizer::engine::{ - CandidateEval, EngineConfig, GateOutcome, RolloutCache, RolloutOutcome, ScoreMatrix, Spend, - canonical_hash, -}; -use crate::trace::{JsonMap, Trace, TraceMeta, TraceOutcome, capture_with_meta}; - -/// How a program-lane strategy tells the engine what "good" means: score one -/// interpreter output (`JsonMap` of the program's output signature fields) -/// against a labeled example. -/// -/// The JSON-native sibling of [`TypedMetric`](crate::evaluate::TypedMetric) — -/// loaded programs have no static output type, so the metric sees the same -/// value model the interpreter produces. The rollout's captured [`Trace`] is -/// always provided; LM-as-judge metrics run outside the capture scope exactly -/// as in the module lane. -#[allow(async_fn_in_trait)] -pub trait ProgramMetric: Send + Sync { - async fn evaluate( - &self, - example: &DemoRow, - output: &JsonMap, - trace: Option<&Trace>, - ) -> Result; -} - -/// Result of [`ProgramEvalEngine::evaluate_program_candidates`]. -#[derive(Clone, Debug)] -pub enum ProgramEvalOutcome { - /// One [`CandidateEval`] per requested candidate, in request order. - Complete(Vec), - /// The uncached portion of the batch didn't fit the remaining budget; - /// nothing ran and spend is unchanged. - BudgetExhausted { needed: usize }, -} - -impl ProgramEvalOutcome { - /// The completed evaluations, if the budget allowed the batch. - pub fn completed(self) -> Option> { - match self { - Self::Complete(evals) => Some(evals), - Self::BudgetExhausted { .. } => None, - } - } -} - -/// The shared evaluation core for the dynamic lane: candidate overlays over -/// one interpreter-loaded program. -/// -/// Owns the labeled example set, the candidate registry (deduplicated by -/// [`Overlay::hash`]), the score matrix, the rollout cache, and the spend -/// meter. Strategies register overlays and call -/// [`evaluate_program_candidates`](Self::evaluate_program_candidates) / -/// [`evaluate_gated`](Self::evaluate_gated). -pub struct ProgramEvalEngine<'m, MT> { - examples: Vec, - example_uids: Vec, - metric: &'m MT, - config: EngineConfig, - candidates: Vec>, - candidate_hashes: Vec, - matrix: ScoreMatrix, - cache: RolloutCache, - spend: Spend, - /// High-water mark of *distinct candidates* with rollouts in flight at - /// the same instant — the parallelism gauge (see - /// [`peak_candidate_concurrency`](Self::peak_candidate_concurrency)). - peak_candidates_in_flight: usize, -} - -impl<'m, MT: ProgramMetric> ProgramEvalEngine<'m, MT> { - pub fn new(examples: Vec, metric: &'m MT, config: EngineConfig) -> Self { - let example_uids = examples.iter().map(canonical_hash).collect(); - let matrix = ScoreMatrix::new(examples.len()); - Self { - examples, - example_uids, - metric, - config, - candidates: Vec::new(), - candidate_hashes: Vec::new(), - matrix, - cache: RolloutCache::default(), - spend: Spend::default(), - peak_candidates_in_flight: 0, - } - } - - pub fn examples(&self) -> &[DemoRow] { - &self.examples - } - - pub fn num_examples(&self) -> usize { - self.examples.len() - } - - pub fn config(&self) -> &EngineConfig { - &self.config - } - - pub fn spend(&self) -> &Spend { - &self.spend - } - - pub fn matrix(&self) -> &ScoreMatrix { - &self.matrix - } - - /// The rollout cache. Keys are `(program hash, overlay hash, example uid, - /// salt)` — candidate identity is the overlay hash, so two candidates on - /// the same example occupy distinct entries. - pub fn cache(&self) -> &RolloutCache { - &self.cache - } - - /// Pareto view over all example columns (see [`ScoreMatrix::pareto`]). - pub fn pareto(&self) -> crate::optimizer::engine::ParetoView { - self.matrix.pareto() - } - - /// Pareto view over a column subset (see [`ScoreMatrix::pareto_over`]). - pub fn pareto_over(&self, columns: &[usize]) -> crate::optimizer::engine::ParetoView { - self.matrix.pareto_over(columns) - } - - /// The parallelism gauge: the maximum number of **distinct candidates** - /// that have had rollouts in flight simultaneously across all batches so - /// far. A value ≥ 2 is positive evidence that candidate-level parallelism - /// actually happened (the module lane is structurally pinned to 1). - pub fn peak_candidate_concurrency(&self) -> usize { - self.peak_candidates_in_flight - } - - /// Registers a candidate overlay, deduplicating by [`Overlay::hash`]. - /// Returns its index. - pub fn register(&mut self, overlay: Overlay) -> usize { - let hash = overlay.hash(); - if let Some(existing) = self.candidate_hashes.iter().position(|&h| h == hash) { - return existing; - } - self.candidates.push(Arc::new(overlay)); - self.candidate_hashes.push(hash); - self.matrix.ensure_rows(self.candidates.len()); - self.candidates.len() - 1 - } - - pub fn candidate(&self, index: usize) -> &Arc { - &self.candidates[index] - } - - pub fn candidate_hash(&self, index: usize) -> u64 { - self.candidate_hashes[index] - } - - pub fn num_candidates(&self) -> usize { - self.candidates.len() - } - - /// Whether `upcoming_rollouts` more rollouts fit the remaining budget. - pub fn budget_allows(&self, upcoming_rollouts: usize) -> bool { - self.config.budget.allows(&self.spend, upcoming_rollouts) - } - - /// Charges auxiliary spend (reflection calls, teacher passes) against the - /// budget, mirroring the module lane. - pub fn charge(&mut self, metric_calls: usize, lm_calls: usize) { - self.spend.metric_calls = self.spend.metric_calls.saturating_add(metric_calls); - self.spend.lm_calls = self.spend.lm_calls.saturating_add(lm_calls); - } - - /// **The IR-native entry point**: evaluates N registered candidates over - /// `subset` example indices (`None` = the full set) through `interp`, - /// with TRUE candidate-level parallelism. - /// - /// Every uncached `(candidate, example)` pair across *all* requested - /// candidates joins one bounded-concurrency fan-out - /// ([`EngineConfig::concurrency`]); each rollout runs - /// `interp.run(input, Some(overlay), …)` under its own capture scope with - /// `TraceMeta.candidate_hash = overlay.hash()` and - /// `tags["program"] = hex(program_hash)` (RFC 0002 §3.3). There is no - /// apply/restore: the overlays read through at render over the one shared - /// `Arc`. - /// - /// Cached rollouts return their `Eval` with `trace: None` and consume no - /// budget. If the uncached portion doesn't fit the remaining budget the - /// engine runs nothing and returns - /// [`ProgramEvalOutcome::BudgetExhausted`]. - pub async fn evaluate_program_candidates( - &mut self, - interp: &Interpreter, - candidates: &[usize], - subset: Option<&[usize]>, - ) -> Result { - for &candidate in candidates { - if candidate >= self.candidates.len() { - return Err(anyhow!("candidate index {candidate} is not registered")); - } - } - let indices: Vec = match subset { - Some(subset) => subset.to_vec(), - None => (0..self.examples.len()).collect(), - }; - if let Some(&bad) = indices.iter().find(|&&idx| idx >= self.examples.len()) { - return Err(anyhow!( - "example index {bad} out of range ({} examples)", - self.examples.len() - )); - } - - let baseline = interp.program().meta.program_hash; - let salt = self.config.cache_salt; - - // Partition the full (candidate × example) grid into cached and - // pending pairs. - let mut cached: HashMap<(usize, usize), Eval> = HashMap::new(); - let mut pending: Vec<(usize, usize)> = Vec::new(); - for &candidate in candidates { - let candidate_hash = self.candidate_hashes[candidate]; - for &idx in &indices { - let key = (candidate, idx); - if cached.contains_key(&key) || pending.contains(&key) { - continue; - } - match self - .cache - .get(baseline, candidate_hash, self.example_uids[idx], salt) - { - Some(eval) => { - cached.insert(key, eval.clone()); - } - None => pending.push(key), - } - } - } - - if !self.budget_allows(pending.len()) { - return Ok(ProgramEvalOutcome::BudgetExhausted { - needed: pending.len(), - }); - } - - let (fresh, batch_peak) = self.run_rollouts(interp, &pending).await?; - self.peak_candidates_in_flight = self.peak_candidates_in_flight.max(batch_peak); - - // Accounting: fresh rollouts consume budget, cached hits are free. - self.spend.metric_calls += fresh.len(); - self.spend.lm_calls += fresh.len(); - self.spend.cache_hits += cached.len(); - for (_, _, _, trace) in &fresh { - self.spend.lm_spans += trace.spans.len(); - for span in &trace.spans { - self.spend.tokens = self.spend.tokens + span.usage; - } - } - - // Bookkeeping: cache inserts + matrix records. - let mut fresh_by_key: HashMap<(usize, usize), (Eval, Trace)> = - HashMap::with_capacity(fresh.len()); - for (candidate, idx, eval, trace) in fresh { - self.cache.insert( - baseline, - self.candidate_hashes[candidate], - self.example_uids[idx], - salt, - eval.clone(), - ); - self.matrix.record(candidate, idx, eval.score); - fresh_by_key.insert((candidate, idx), (eval, trace)); - } - for (&(candidate, idx), eval) in &cached { - self.matrix.record(candidate, idx, eval.score); - } - - let evals = candidates - .iter() - .map(|&candidate| CandidateEval { - candidate, - rollouts: indices - .iter() - .map(|&idx| { - if let Some((eval, trace)) = fresh_by_key.remove(&(candidate, idx)) { - RolloutOutcome { - example: idx, - eval, - trace: Some(trace), - } - } else { - let eval = cached - .get(&(candidate, idx)) - .cloned() - .expect("every requested pair is either fresh or cached"); - RolloutOutcome { - example: idx, - eval, - trace: None, - } - } - }) - .collect(), - }) - .collect(); - - Ok(ProgramEvalOutcome::Complete(evals)) - } - - /// Single-candidate convenience over - /// [`evaluate_program_candidates`](Self::evaluate_program_candidates). - pub async fn evaluate( - &mut self, - interp: &Interpreter, - candidate: usize, - subset: Option<&[usize]>, - ) -> Result { - use crate::optimizer::engine::EvalOutcome; - match self - .evaluate_program_candidates(interp, &[candidate], subset) - .await? - { - ProgramEvalOutcome::Complete(mut evals) => Ok(EvalOutcome::Complete(evals.remove(0))), - ProgramEvalOutcome::BudgetExhausted { needed } => { - Ok(EvalOutcome::BudgetExhausted { needed }) - } - } - } - - /// Minibatch gating (the GEPA acceptance pattern), program-lane edition: - /// evaluates the candidate on `minibatch`; only if the minibatch mean - /// strictly beats `threshold` does it promote to a full-set evaluation. - pub async fn evaluate_gated( - &mut self, - interp: &Interpreter, - candidate: usize, - minibatch: &[usize], - threshold: f64, - ) -> Result { - use crate::optimizer::engine::EvalOutcome; - let minibatch_eval = match self.evaluate(interp, candidate, Some(minibatch)).await? { - EvalOutcome::Complete(eval) => eval, - EvalOutcome::BudgetExhausted { needed } => { - return Ok(GateOutcome::BudgetExhausted { needed }); - } - }; - - if minibatch_eval.mean() <= threshold { - return Ok(GateOutcome::Rejected { - minibatch: minibatch_eval, - }); - } - - match self.evaluate(interp, candidate, None).await? { - EvalOutcome::Complete(full) => Ok(GateOutcome::Promoted { - minibatch: minibatch_eval, - full, - }), - EvalOutcome::BudgetExhausted { needed } => Ok(GateOutcome::BudgetExhausted { needed }), - } - } - - /// The one shared fan-out: every pending `(candidate, example)` pair — - /// across all candidates — in a single `buffer_unordered` stream. Returns - /// the fresh rollouts and the batch's distinct-candidate concurrency - /// high-water mark. - async fn run_rollouts( - &self, - interp: &Interpreter, - pending: &[(usize, usize)], - ) -> Result<(Vec<(usize, usize, Eval, Trace)>, usize)> { - let metric = self.metric; - let program_tag = format!("{:016x}", interp.program().meta.program_hash); - let gauge = Gauge::default(); - - let fresh: Vec<(usize, usize, Eval, Trace)> = - stream::iter(pending.iter().map(|&(candidate, idx)| { - let overlay = Arc::clone(&self.candidates[candidate]); - let candidate_hash = self.candidate_hashes[candidate]; - let example = &self.examples[idx]; - let program_tag = &program_tag; - let gauge = &gauge; - async move { - let _in_flight = gauge.enter(candidate); - let meta = TraceMeta { - candidate_hash: Some(candidate_hash), - input: Some(example.input.clone()), - tags: [("program".to_string(), program_tag.clone())] - .into_iter() - .collect(), - ..TraceMeta::default() - }; - let started = Instant::now(); - let (result, mut trace) = capture_with_meta(meta, || { - interp.run(example.input.clone(), Some(overlay), RunBudget::unlimited()) - }) - .await; - let output = result.map_err(|err| { - anyhow!("candidate {candidate} failed on example {idx}: {err}") - })?; - // Metric runs outside the capture scope so LM-as-judge - // metrics don't pollute the execution trace. - let eval = metric.evaluate(example, &output, Some(&trace)).await?; - trace.outcome = Some(TraceOutcome { - output: Some(output), - error: None, - eval: Some(eval.clone()), - duration_us: started.elapsed().as_micros() as u64, - }); - Ok::<_, anyhow::Error>((candidate, idx, eval, trace)) - } - })) - .buffer_unordered(self.config.concurrency.max(1)) - .try_collect() - .await?; - - Ok((fresh, gauge.peak())) - } -} - -/// Counts distinct candidates with rollouts in flight; records the peak. -#[derive(Default)] -struct Gauge { - in_flight: Mutex>, - peak: AtomicUsize, -} - -impl Gauge { - fn enter(&self, candidate: usize) -> GaugeGuard<'_> { - let mut in_flight = self.in_flight.lock().unwrap(); - *in_flight.entry(candidate).or_insert(0) += 1; - self.peak.fetch_max(in_flight.len(), Ordering::Relaxed); - GaugeGuard { - gauge: self, - candidate, - } - } - - fn peak(&self) -> usize { - self.peak.load(Ordering::Relaxed) - } -} - -struct GaugeGuard<'a> { - gauge: &'a Gauge, - candidate: usize, -} - -impl Drop for GaugeGuard<'_> { - fn drop(&mut self) { - let mut in_flight = self.gauge.in_flight.lock().unwrap(); - if let Some(count) = in_flight.get_mut(&self.candidate) { - *count -= 1; - if *count == 0 { - in_flight.remove(&self.candidate); - } - } - } -} diff --git a/crates/dspy-rs/src/optimizer/simba.rs b/crates/dspy-rs/src/optimizer/simba.rs index 452ba578..07a38cae 100644 --- a/crates/dspy-rs/src/optimizer/simba.rs +++ b/crates/dspy-rs/src/optimizer/simba.rs @@ -8,21 +8,20 @@ use std::collections::HashSet; use anyhow::{Result, anyhow}; use bon::Builder; -use rand::{SeedableRng, rngs::StdRng, seq::SliceRandom}; +use rand::seq::SliceRandom; use serde::{Deserialize, Serialize}; +use crate::core::ToInput; use crate::evaluate::{Eval, TypedMetric}; use crate::optimizer::engine::{ - Budget, Candidate, EngineConfig, EvalEngine, EvalOutcome, GateOutcome, Spend, apply_candidate, - canonical_hash, + Candidate, Engine, EngineConfig, EvalOutcome, GateOutcome, Spend, canonical_hash, }; -use crate::optimizer::gepa::format_schema_for_reflection; use crate::optimizer::harvest::{collect_demo_candidates, select_demos}; -use crate::optimizer::{Optimizer, predictor_names, with_named_predictor}; +use crate::optimizer::target::LeafInfo; +use crate::optimizer::{OptimizeTarget, Optimizer, OptimizerCommon, Report}; use crate::trace::Trace; use crate::utils::truncate; -use crate::core::ToInput; -use crate::{Facet, Module, Predict, Signature}; +use crate::{Module, Predict, Predictors, Signature}; /// Distill one improvement rule from contrasting rollouts. /// @@ -100,8 +99,8 @@ pub struct SimbaReport { /// Minibatch introspective ascent — the cheap agentic default (vision §4.3). /// -/// SIMBA is a thin strategy over the shared [`EvalEngine`]. It keeps one -/// *current* program as an overlay [`Candidate`] and hill-climbs: +/// SIMBA is a thin strategy over the shared [`Engine`]. It keeps one +/// *current* program as a name-keyed [`Candidate`] and hill-climbs: /// /// 1. **Sample** a trainset minibatch (seeded RNG, indices sorted for /// deterministic evaluation order). @@ -112,7 +111,8 @@ pub struct SimbaReport { /// 3. **Move** — exactly one of: /// - **append-demo**: when the best rollout reached `min_demo_score`, its /// successful `Predict` spans become one new few-shot demo per predictor -/// via the shared harvest name-join, appended to the current demo set +/// via the shared harvest name-join (spans the metric gave their own +/// sub-threshold eval are excluded), appended to the current demo set /// (capped at `max_demos`, oldest dropped, duplicates skipped); /// - **append-rule**: otherwise, a reflection LM (`prompt_model`) reads the /// contrasting rollouts and distills one rule, appended to the target @@ -124,8 +124,8 @@ pub struct SimbaReport { /// promote to a full-trainset evaluation and become the new current /// program. /// -/// The winner is installed permanently through the one candidate seam -/// ([`apply_candidate`]) when `compile` returns. +/// The winner is installed once — through [`OptimizeTarget::install`] — when +/// `compile` returns; candidate evaluation never mutates the module. /// /// # Hyperparameters /// @@ -150,7 +150,7 @@ pub struct SimbaReport { /// /// ```ignore /// let simba = SIMBA::builder().max_steps(8).minibatch_size(8).build(); -/// let report = simba.compile(&mut module, trainset, &metric).await?; +/// let report = simba.compile_module(&mut module, &trainset, &metric).await?; /// println!("{:.3} -> {:.3}", report.baseline_score, report.final_score); /// ``` #[derive(Builder)] @@ -168,7 +168,10 @@ pub struct SIMBA { pub max_demos: usize, /// Minimum rollout score for its spans to qualify as demos; below it the - /// step proposes a rule instead. + /// step proposes a rule instead. Within a qualifying rollout, a span the + /// metric gave its own eval + /// ([`TypedMetric::evaluate_spans`](crate::evaluate::TypedMetric::evaluate_spans)) + /// is gated on that score instead. #[builder(default = 1.0)] pub min_demo_score: f64, @@ -194,6 +197,16 @@ pub struct SIMBA { type RolloutStore = Vec)>>; impl SIMBA { + fn common(&self) -> OptimizerCommon { + OptimizerCommon { + eval_concurrency: self.eval_concurrency, + max_metric_calls: self.max_metric_calls, + max_lm_calls: self.max_lm_calls, + seed: self.seed, + ..OptimizerCommon::default() + } + } + /// Formats one rollout (score, feedback, per-span I/O) for the reflection /// prompt. fn summarize_rollout(eval: &Eval, trace: Option<&Trace>) -> String { @@ -232,15 +245,15 @@ impl SIMBA { /// The predictor an append-rule move targets: the one with the most spans /// in the worst rollout's trace (first name wins ties or when no trace). - fn rule_target<'a>(names: &'a [String], worst_trace: Option<&Trace>) -> &'a str { - let mut target = names[0].as_str(); + fn rule_target<'a>(leaves: &'a [LeafInfo], worst_trace: Option<&Trace>) -> &'a LeafInfo { + let mut target = &leaves[0]; if let Some(trace) = worst_trace { let mut most = 0usize; - for name in names { - let count = trace.for_component(name).count(); + for leaf in leaves { + let count = trace.for_component(&leaf.name).count(); if count > most { most = count; - target = name; + target = leaf; } } } @@ -251,16 +264,13 @@ impl SIMBA { /// the best rollout via the shared trace name-join, appended to the /// current effective demo set. Returns `None` when nothing changes (no /// qualifying spans, or every harvested demo is already present). - fn append_demo_child( + fn append_demo_child( &self, - module: &mut M, + leaves: &[LeafInfo], current: &Candidate, score: f64, trace: &Trace, - ) -> Result> - where - M: for<'a> Facet<'a>, - { + ) -> Option { let harvested = select_demos( collect_demo_candidates(std::iter::once((score, trace)), self.min_demo_score), 1, @@ -269,13 +279,15 @@ impl SIMBA { let mut child = current.clone(); let mut changed = false; for (name, new_demos) in harvested { - // Effective demo set: the current overlay if it carries one, - // otherwise whatever is installed on the module. - let mut demos = match current.overlays.get(&name).and_then(|o| o.demos.clone()) { - Some(demos) => demos, - None => with_named_predictor(module, &name, |predictor| { - Ok(predictor.demos_as_json()) - })?, + // Effective demo set: the current candidate if it carries one, + // otherwise the module's own demos (leaf snapshot). + let mut demos = match current.demos_of(&name) { + Some(demos) => demos.to_vec(), + None => leaves + .iter() + .find(|leaf| leaf.name == name) + .map(|leaf| leaf.demos.clone()) + .unwrap_or_default(), }; let seen: HashSet = demos.iter().map(canonical_hash).collect(); @@ -295,26 +307,22 @@ impl SIMBA { changed = true; } } - Ok(changed.then_some(child)) + changed.then_some(child) } /// Proposes the rule text for an append-rule move, preferring LM /// reflection. Returns the rule and the number of reflection LM calls /// consumed (0 or 1); reflection failures degrade to metric-feedback /// concatenation with a warning rather than aborting the run. - async fn propose_rule( + async fn propose_rule( &self, - module: &mut M, - target: &str, + leaf: &LeafInfo, current_instruction: &str, better_rollout: String, worse_rollout: String, worst_eval: &Eval, reflector: Option<&Predict>, - ) -> (String, usize) - where - M: for<'a> Facet<'a>, - { + ) -> (String, usize) { let fallback = || { worst_eval.feedback.clone().unwrap_or_else(|| { format!( @@ -328,13 +336,8 @@ impl SIMBA { return (fallback(), 0); }; - let task_description = with_named_predictor(module, target, |predictor| { - Ok(format_schema_for_reflection(predictor.schema())) - }) - .unwrap_or_default(); - let input = IntrospectRolloutsInput { - task_description, + task_description: leaf.schema_for_reflection(), current_instruction: current_instruction.to_string(), better_rollout, worse_rollout, @@ -345,7 +348,7 @@ impl SIMBA { let rule = predicted.rule.trim().to_string(); if rule.is_empty() { tracing::warn!( - target, + target = leaf.name, "reflection LM returned an empty rule; using metric feedback" ); } else { @@ -354,7 +357,7 @@ impl SIMBA { } Err(err) => { tracing::warn!( - target, + target = leaf.name, error = %err, "reflection LM call failed; using metric feedback" ); @@ -363,51 +366,53 @@ impl SIMBA { (fallback(), 1) } -} -impl Optimizer for SIMBA { - type Report = SimbaReport; - - async fn compile( + /// Convenience: optimizes a typed module over a trainset with this + /// optimizer's default engine. + pub async fn compile_module( &self, module: &mut M, - trainset: Vec, + trainset: &[E], metric: &MT, - ) -> Result + ) -> Result where E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, + M: Module + Predictors, MT: TypedMetric, { - let names = predictor_names(module)?; - if names.is_empty() { + let mut target = OptimizeTarget::module(module, trainset, metric); + let mut engine = Engine::new(Optimizer::engine_config(self)); + let report = Optimizer::compile(self, &mut target, &mut engine).await?; + report + .into_simba() + .ok_or_else(|| anyhow!("SIMBA must return a SIMBA report")) + } +} + +#[async_trait::async_trait(?Send)] +impl Optimizer for SIMBA { + fn engine_config(&self) -> EngineConfig { + self.common().engine_config() + } + + async fn compile( + &self, + target: &mut OptimizeTarget<'_>, + engine: &mut Engine, + ) -> Result { + let leaves = target.leaves().to_vec(); + if leaves.is_empty() { return Err(anyhow!("no optimizable predictors found")); } - let mut engine = EvalEngine::new( - trainset, - metric, - EngineConfig { - concurrency: self.eval_concurrency, - budget: Budget { - max_metric_calls: self.max_metric_calls, - max_lm_calls: self.max_lm_calls, - max_tokens: None, - }, - cache_salt: 0, - }, - ); - let num_examples = engine.num_examples(); + let num_examples = target.num_examples(); let all_indices: Vec = (0..num_examples).collect(); let reflector = self .prompt_model .as_ref() .map(|lm| Predict::::builder().lm(lm.clone()).build()); - let mut rng = match self.seed { - Some(seed) => StdRng::seed_from_u64(seed), - None => StdRng::from_entropy(), - }; + let mut rng = self.common().rng(); // Baseline: the empty candidate over the full trainset, traced. This // seeds the rollout store; because every later *current* program is @@ -415,7 +420,7 @@ impl Optimizer for SIMBA { // needs extra rollouts. let mut current = Candidate::new(); let current_idx = engine.register(current.clone()); - let baseline_eval = match engine.evaluate(module, current_idx, None).await? { + let baseline_eval = match engine.evaluate(target, current_idx, None).await? { EvalOutcome::Complete(eval) => eval, EvalOutcome::BudgetExhausted { needed } => { return Err(anyhow!( @@ -478,7 +483,7 @@ impl Optimizer for SIMBA { let demo_child = if best_eval.score >= self.min_demo_score { match &best_trace { Some(trace) => { - self.append_demo_child(module, ¤t, best_eval.score, trace)? + self.append_demo_child(&leaves, ¤t, best_eval.score, trace) } None => None, } @@ -489,21 +494,14 @@ impl Optimizer for SIMBA { let (child, move_kind) = match demo_child { Some(child) => (child, SimbaMove::AppendDemo), None => { - let target = Self::rule_target(&names, worst_trace.as_ref()); - let base_instruction = match current - .overlays - .get(target) - .and_then(|overlay| overlay.instruction.clone()) - { - Some(instruction) => instruction, - None => with_named_predictor(module, target, |predictor| { - Ok(predictor.instruction()) - })?, + let leaf = Self::rule_target(&leaves, worst_trace.as_ref()); + let base_instruction = match current.instruction_of(&leaf.name) { + Some(instruction) => instruction.to_string(), + None => leaf.instruction.clone(), }; let (rule, reflection_calls) = self .propose_rule( - module, - target, + leaf, &base_instruction, Self::summarize_rollout(&best_eval, best_trace.as_ref()), Self::summarize_rollout(&worst_eval, worst_trace.as_ref()), @@ -514,7 +512,10 @@ impl Optimizer for SIMBA { engine.charge(0, reflection_calls); let mut child = current.clone(); - child.set_instruction(target, format!("{base_instruction}\n\n[SIMBA rule] {rule}")); + child.set_instruction( + &leaf.name, + format!("{base_instruction}\n\n[SIMBA rule] {rule}"), + ); (child, SimbaMove::AppendRule) } }; @@ -522,7 +523,7 @@ impl Optimizer for SIMBA { // 4. Accept through the engine's minibatch gate. let child_idx = engine.register(child.clone()); match engine - .evaluate_gated(module, child_idx, &minibatch, threshold) + .evaluate_gated(target, child_idx, &minibatch, threshold) .await? { GateOutcome::BudgetExhausted { .. } => break, @@ -569,18 +570,18 @@ impl Optimizer for SIMBA { } } - // Install the winner permanently through the one candidate seam. + // The one mutation of the run: install the winner. if !current.is_empty() { - let _undo = apply_candidate(module, ¤t)?; + target.install(¤t)?; } - Ok(SimbaReport { + Ok(Report::Simba(SimbaReport { baseline_score, final_score, steps, accepted, rejected, spend: *engine.spend(), - }) + })) } } diff --git a/crates/dspy-rs/src/optimizer/structural.rs b/crates/dspy-rs/src/optimizer/structural.rs new file mode 100644 index 00000000..4b271ef9 --- /dev/null +++ b/crates/dspy-rs/src/optimizer/structural.rs @@ -0,0 +1,826 @@ +//! Structural: LM-guided hill-climbing over the graph-edit calculus +//! (RFC 0004 §6) — the sixth strategy over the shared [`Engine`]. +//! +//! Where the other five strategies tune parameter *values* through overlays, +//! Structural proposes [`ir::Edit`](crate::ir::Edit)s: each generation it +//! gathers the [`Program::legal_edits`] menu, has a reflection LM choose one +//! edit from the serialized menu plus the incumbent's evaluation feedback, +//! applies it via [`Program::edited`], carries the tuned overlay across the +//! structural change with [`migrate_overlay`], and accepts the child through +//! the engine's minibatch gate — parent and child scored on the same shared +//! minibatch, winner kept. + +use std::num::NonZeroU32; +use std::sync::Arc; + +use anyhow::{Result, anyhow}; +use bon::Builder; +use rand::{Rng, rngs::StdRng, seq::SliceRandom}; +use serde::{Deserialize, Serialize}; + +use crate::evaluate::Eval; +use crate::ir::builder::cot_reasoning_field; +use crate::ir::edit::{Edit, EditKind, SwapTarget, migrate_overlay}; +use crate::ir::graph::{Node, NodeBudget, NodeId, Program, StopSpec}; +use crate::ir::interp::{Interpreter, RuntimeEnv}; +use crate::ir::params::{DemoRow, Overlay}; +use crate::optimizer::OptimizerCommon; +use crate::optimizer::engine::{Engine, EngineConfig, EvalOutcome, GateOutcome, Spend}; +use crate::optimizer::target::{OptimizeTarget, ProgramMetric}; +use crate::trace::Trace; +use crate::utils::truncate; +use crate::{Predict, Signature}; + +/// Choose one structural edit for an LLM-pipeline program. +/// +/// You are optimizing the structure of an LLM pipeline program. Study the +/// program source, the menu of legal structural edits, and the per-example +/// scores and feedback from the current program's last evaluation. Choose the +/// single edit most likely to fix the failure modes the feedback names: add +/// reasoning where answers are shallow, add or drop tools where tool use goes +/// wrong, wrap flaky steps in a retry, remove steps that only add noise. +/// Return only the `option` number of the chosen menu entry, with no preamble +/// or commentary. +#[derive(Signature, Clone, Debug)] +// The struct itself is only a schema: the derive generates the +// `ChooseEditInput`/`ChooseEditOutput` types the reflection call reads. +#[allow(dead_code)] +struct ChooseEdit { + /// The program in canonical `.dsrs` text form. + #[input] + program_source: String, + + /// The menu of legal structural edits, one JSON object per line, each + /// with an `option` number. + #[input] + edit_menu: String, + + /// Per-example scores and textual feedback from the last evaluation. + #[input] + execution_feedback: String, + + /// The `option` number of the chosen edit. + #[output] + chosen_option: String, +} + +/// One entry of the proposer menu: a leaf, one [`EditKind`] admissible there, +/// and the node id to materialize the concrete [`Edit`] against. +#[derive(Clone, Debug)] +struct MenuEntry { + leaf: String, + node: NodeId, + kind: EditKind, +} + +/// What one Structural generation did. +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct StructuralStep { + /// Generation index (0-based). + pub generation: usize, + /// The leaf the chosen edit targets. + pub leaf: String, + /// The concrete edit that was proposed (serde data — replayable against + /// the generation's parent, whose hash is `parent_hash`). + pub edit: Edit, + /// `program_hash` of the parent the edit was applied to. + pub parent_hash: u64, + /// Parent's mean score on the generation's minibatch (the gate threshold). + pub parent_minibatch_score: f64, + /// Child's mean on the same minibatch; `None` when the child never scored + /// (the edit failed to apply or the child failed to load). + pub child_minibatch_score: Option, + /// Whether the gate promoted the child to the new incumbent. + pub accepted: bool, + /// Full-set mean of the child; `Some` only when accepted. + pub full_score: Option, + /// Why the child was never scored (apply/load failure), when it wasn't. + pub rejection: Option, +} + +/// What a [`Structural`] run did. The winner is returned, not installed: +/// bake it (`report.program.bake(&report.overlay, note)`) or save it — the +/// incumbent interpreter passed in is never mutated. +#[derive(Clone, Debug)] +pub struct StructuralReport { + /// The winning program (the input program when nothing was accepted). + pub program: Arc, + /// The incumbent overlay re-minted against the winner via + /// [`migrate_overlay`] at every accepted edit (empty when no overlay was + /// supplied and nothing migrated). + pub overlay: Overlay, + /// Mean metric score of the input program (+ overlay) over the examples. + pub baseline_score: f64, + /// Full-set mean of the final program (baseline if nothing was accepted). + pub final_score: f64, + /// The accepted edits, in order. Each applies to the program whose hash + /// its step's `parent_hash` records (node ids are per-parent handles, so + /// the vec is a lineage, not a single batch). + pub edits: Vec, + /// Per-generation outcomes, in order. + pub steps: Vec, + /// Generations promoted by the gate. + pub accepted: usize, + /// Generations rejected (gate losses, apply failures, load failures). + pub rejected: usize, + /// Engine spend for the whole run (reflection calls included). + pub spend: Spend, +} + +/// Structural optimizer over the graph-edit calculus (RFC 0004 §6). +/// +/// Structural is a thin strategy over the shared [`Engine`], but it searches +/// program *structure* instead of parameter values: candidates are whole +/// programs minted by [`Program::edited`], and the tuned overlay follows the +/// incumbent across every accepted edit via [`migrate_overlay`]. It runs on +/// the program lane only — the module lane has no skeleton to edit. +/// +/// Each generation: +/// +/// 1. **Sample** a shared minibatch (seeded RNG, indices sorted). The +/// incumbent's minibatch mean is the gate threshold; it is served from the +/// rollout cache (the incumbent always has full coverage), so re-scoring +/// the parent costs nothing. +/// 2. **Menu** — gather [`Program::legal_edits`] for every leaf and keep the +/// kinds Structural can materialize without free text: `AugmentSig` (the +/// CoT move), `SwapToAgent`/`SwapToPredict`, `WrapRetry`, `Remove`, and +/// per-tool `AddTool`/`RemoveTool`. `SetStop` and `SetInstructionDefault` +/// are left to value-level optimizers. +/// 3. **Choose** — a reflection LM (`prompt_model`) reads the program's +/// canonical `.dsrs` text, the serialized menu, and the incumbent's +/// per-example feedback, and returns one option number. Without a +/// `prompt_model` (or when the reply doesn't parse) the choice degrades to +/// a seeded-uniform pick from the menu. +/// 4. **Apply** — [`Program::edited`] mints the child; [`migrate_overlay`] +/// re-mints the incumbent overlay against it; the child is loaded through +/// the caller-supplied [`RuntimeEnv`] factory. An edit that fails to +/// apply, a child that fails validation, or a child that fails to load is +/// recorded and skipped — never a panic, never an abort. +/// 5. **Gate** — the child is accepted through the engine's minibatch gate: +/// only if its mean on the shared minibatch strictly beats the parent's +/// does it promote to a full-set evaluation and become the new incumbent. +/// +/// Every child is a fresh program that must be re-scored from scratch (its +/// hash keys its own rollout-cache rows), so the budget caps are the real +/// control surface: `max_rollouts` / `max_lm_calls` stop the run cleanly +/// when the next batch wouldn't fit. +/// +/// # Hyperparameters +/// +/// - **`num_iterations`** (default: 8) — structural generations to attempt. +/// - **`minibatch_size`** (default: 8) — examples in the shared minibatch +/// parent and child are compared on. +/// - **`prompt_model`** — reflection LM that chooses edits from the menu. +/// Strongly recommended; without it the choice is a seeded-uniform pick. +/// - **`max_rollouts`** / **`max_lm_calls`** — hard budget caps. Every +/// accepted child costs a full-set evaluation on top of its minibatch. +/// - **`eval_concurrency`** (default: 16) — rollouts in flight during +/// evaluation. +/// - **`seed`** — fixes minibatch sampling and the fallback edit choice. +/// +/// # Cost +/// +/// `examples.len()` for the baseline pass, then per generation: +/// `minibatch_size` rollouts for the gate, plus the remaining +/// `examples.len() - minibatch_size` only on promotion, plus one reflection +/// call when a `prompt_model` is set. Rejected children never pay for a full +/// pass. +/// +/// ```ignore +/// let structural = Structural::builder() +/// .num_iterations(8) +/// .max_rollouts(Some(400)) +/// .prompt_model(reflection_lm) +/// .build(); +/// let report = structural +/// .compile_program(&interp, &examples, &metric, || { +/// RuntimeEnv::new().bind_model("m", lm.clone()) +/// }) +/// .await?; +/// println!("{:.3} -> {:.3}", report.baseline_score, report.final_score); +/// let baked = report.program.bake(&report.overlay, Lineage::default())?; +/// ``` +#[derive(Builder)] +pub struct Structural { + /// Structural generations to attempt (one proposed edit each). + #[builder(default = 8)] + pub num_iterations: usize, + + /// Examples in the shared minibatch parent and child are compared on. + #[builder(default = 8)] + pub minibatch_size: usize, + + /// Reflection LM that chooses an edit from the serialized menu. Without + /// it, the choice degrades to a seeded-uniform pick. + pub prompt_model: Option, + + /// Hard cap on evaluation rollouts. `None` = unlimited. + pub max_rollouts: Option, + /// Hard cap on LM call units (rollouts + reflection). `None` = unlimited. + pub max_lm_calls: Option, + + /// Concurrent rollouts in flight during evaluation. + #[builder(default = crate::evaluate::DEFAULT_EVAL_CONCURRENCY)] + pub eval_concurrency: usize, + + /// Seed for minibatch sampling and the fallback edit choice. `None` uses + /// a nondeterministic seed. + pub seed: Option, +} + +/// Per-example bookkeeping for the incumbent: metric result plus the trace +/// when the engine ran it fresh (cache-served cells carry no trace). +type RolloutStore = Vec)>>; + +impl Structural { + fn common(&self) -> OptimizerCommon { + OptimizerCommon { + eval_concurrency: self.eval_concurrency, + max_metric_calls: self.max_rollouts, + max_lm_calls: self.max_lm_calls, + seed: self.seed, + ..OptimizerCommon::default() + } + } + + /// The engine configuration `compile_program` runs with. + pub fn engine_config(&self) -> EngineConfig { + self.common().engine_config() + } + + /// The proposer menu: every leaf's [`Program::legal_edits`], filtered to + /// the kinds Structural can materialize without free text. `AugmentSig` + /// is dropped when the leaf's signature already carries the reasoning + /// field (the edit would only fail with `DuplicateField`). + fn menu(program: &Program) -> Vec { + let reasoning = cot_reasoning_field(); + let mut menu = Vec::new(); + for (id, node) in program.nodes.iter() { + let Some(leaf) = program.leaf_name(id) else { + continue; + }; + let sig = match node { + Node::Predict(n) => Some(n.sig), + Node::AgentLoop(n) => Some(n.sig), + _ => None, + }; + for kind in program.legal_edits(id) { + match kind { + EditKind::SetStop | EditKind::SetInstructionDefault => continue, + EditKind::AugmentSig => { + let taken = sig.is_some_and(|sig| { + let def = &program.sigs[sig]; + def.inputs + .iter() + .chain(def.outputs.iter()) + .any(|f| f.name == reasoning.name) + }); + if taken { + continue; + } + } + _ => {} + } + menu.push(MenuEntry { + leaf: leaf.to_string(), + node: id, + kind, + }); + } + } + menu + } + + /// One human-readable line per menu entry, for the reflection prompt. + fn render_menu(program: &Program, menu: &[MenuEntry]) -> String { + let tool_name = |tool| program.syms.get(program.tools[tool].name); + let all_tools = || { + program + .tools + .values() + .map(|t| program.syms.get(t.name)) + .collect::>() + .join(", ") + }; + menu.iter() + .enumerate() + .map(|(option, entry)| { + let note = match entry.kind { + EditKind::AugmentSig => { + "prepend a chain-of-thought `reasoning` output field".to_string() + } + EditKind::SwapToAgent => format!( + "swap this `predict` leaf into a tool-using `agent` loop (tools: [{}])", + all_tools() + ), + EditKind::SwapToPredict => { + "swap this `agent` leaf back into a plain `predict`".to_string() + } + EditKind::WrapRetry => { + "wrap this node in a retry (2 attempts, feedback on)".to_string() + } + EditKind::Remove => "remove this step from its `seq`".to_string(), + EditKind::AddTool { tool } => { + format!("declare tool `{}` on this agent", tool_name(tool)) + } + EditKind::RemoveTool { tool } => { + format!("undeclare tool `{}` from this agent", tool_name(tool)) + } + EditKind::SetStop | EditKind::SetInstructionDefault => { + unreachable!("filtered out of the menu") + } + }; + serde_json::json!({ + "option": option, + "leaf": entry.leaf, + "edit": entry.kind, + "note": note, + }) + .to_string() + }) + .collect::>() + .join("\n") + } + + /// Materializes a chosen menu entry into a concrete [`Edit`] with fixed, + /// conservative parameters (the CoT reasoning field, default stop/budget, + /// 2 retry attempts). + fn materialize(program: &Program, entry: &MenuEntry) -> Edit { + match entry.kind { + EditKind::AugmentSig => Edit::AugmentSig { + leaf: entry.node, + prepend: cot_reasoning_field(), + }, + EditKind::SwapToAgent => Edit::SwapLeaf { + leaf: entry.node, + to: SwapTarget::Agent { + tools: program.tools.keys().collect(), + stop: StopSpec::default(), + budget: NodeBudget::default(), + }, + }, + EditKind::SwapToPredict => Edit::SwapLeaf { + leaf: entry.node, + to: SwapTarget::Predict, + }, + EditKind::WrapRetry => Edit::WrapRetry { + node: entry.node, + max_attempts: NonZeroU32::new(2).expect("2 is nonzero"), + backoff_ms: 0, + feedback: true, + }, + EditKind::Remove => Edit::Remove { node: entry.node }, + EditKind::AddTool { tool } => Edit::AddTool { + agent: entry.node, + tool, + }, + EditKind::RemoveTool { tool } => Edit::RemoveTool { + agent: entry.node, + tool, + }, + EditKind::SetStop | EditKind::SetInstructionDefault => { + unreachable!("filtered out of the menu") + } + } + } + + /// Extracts an in-range option number from the reflection LM's reply: + /// the first contiguous digit run (so "Option 3" parses as 3). `None` + /// when nothing parses or the number is out of range. + fn parse_choice(reply: &str, menu_len: usize) -> Option { + let digits: String = reply + .chars() + .skip_while(|c| !c.is_ascii_digit()) + .take_while(char::is_ascii_digit) + .collect(); + let option: usize = digits.parse().ok()?; + (option < menu_len).then_some(option) + } + + /// Formats the incumbent's minibatch scores/feedback plus any errored + /// spans from stored traces — what the reflection LM sees. + fn summarize_feedback(store: &RolloutStore, minibatch: &[usize]) -> String { + use std::fmt::Write as _; + + let mut text = String::new(); + for &idx in minibatch { + let Some((eval, trace)) = &store[idx] else { + continue; + }; + let _ = writeln!( + text, + "{}: score={:.3}; {}", + idx + 1, + eval.score, + eval.feedback.as_deref().unwrap_or("-") + ); + let Some(trace) = trace else { + continue; + }; + for span in &trace.spans { + if let Some(error) = &span.error { + let _ = writeln!( + text, + " {} call {}: <{}: {}>", + trace.component_name(span.component), + span.seq, + error.kind.as_str(), + truncate(span.raw_output.as_deref().unwrap_or(&error.message), 500) + ); + } + } + } + text.trim_end().to_string() + } + + /// Chooses a menu option, preferring LM reflection when a `prompt_model` + /// is configured. Returns the option index and the number of reflection + /// LM calls consumed (0 or 1). Reflection failures degrade to the + /// seeded-uniform fallback with a warning rather than aborting the run. + async fn choose( + &self, + program: &Program, + menu: &[MenuEntry], + execution_feedback: &str, + generation: usize, + reflector: Option<&Predict>, + rng: &mut StdRng, + ) -> (usize, usize) { + let fallback = |rng: &mut StdRng| rng.gen_range(0..menu.len()); + + let Some(reflector) = reflector else { + return (fallback(rng), 0); + }; + + let input = ChooseEditInput { + program_source: program.to_dsrs(), + edit_menu: Self::render_menu(program, menu), + execution_feedback: execution_feedback.to_string(), + }; + + match reflector.call(input).await { + Ok(predicted) => match Self::parse_choice(&predicted.chosen_option, menu.len()) { + Some(option) => return (option, 1), + None => { + tracing::warn!( + generation, + reply = %truncate(&predicted.chosen_option, 200), + "reflection LM reply is not a menu option; picking uniformly" + ); + } + }, + Err(err) => { + tracing::warn!( + generation, + error = %err, + "reflection LM call failed; picking uniformly" + ); + } + } + + (fallback(rng), 1) + } + + /// Runs the structural search over a loaded program. + /// + /// `interp` is the incumbent (never mutated); `env` supplies a fresh + /// [`RuntimeEnv`] every time a child program needs loading — the same + /// model/tool/sandbox bindings the incumbent was loaded with. The winner + /// comes back in the report as a program plus a migrated overlay. + pub async fn compile_program( + &self, + interp: &Interpreter, + examples: &[DemoRow], + metric: &M, + env: F, + ) -> Result + where + M: ProgramMetric, + F: Fn() -> RuntimeEnv, + { + self.compile_program_with_overlay(interp, None, examples, metric, env) + .await + } + + /// [`compile_program`](Self::compile_program) with an incumbent overlay — + /// tuned slot values from a prior value-level optimizer, minted against + /// `interp`'s program. The overlay is applied during every incumbent + /// evaluation and carried across each accepted edit with + /// [`migrate_overlay`]. + pub async fn compile_program_with_overlay( + &self, + interp: &Interpreter, + overlay: Option, + examples: &[DemoRow], + metric: &M, + env: F, + ) -> Result + where + M: ProgramMetric, + F: Fn() -> RuntimeEnv, + { + if examples.is_empty() { + return Err(anyhow!("no examples to optimize over")); + } + let mut program = Arc::clone(interp.program()); + let mut incumbent = overlay.unwrap_or_else(|| Overlay::new(&program)); + if incumbent.base != program.meta.program_hash { + return Err(anyhow!( + "overlay minted against program {:016x}, expected {:016x}", + incumbent.base, + program.meta.program_hash + )); + } + + let mut engine = Engine::new(self.engine_config()); + let reflector = self + .prompt_model + .as_ref() + .map(|lm| Predict::::builder().lm(lm.clone()).build()); + let mut rng = self.common().rng(); + + let num_examples = examples.len(); + let all_indices: Vec = (0..num_examples).collect(); + + // Baseline: the incumbent (+ overlay) over the full example set, + // traced. Seeds the rollout store and the cache — every later parent + // minibatch read is served for free. + let mut incumbent_row = engine.register_overlay(incumbent.clone()); + let baseline_eval = { + let target = OptimizeTarget::program(interp, examples, metric); + match engine.evaluate(&target, incumbent_row, None).await? { + EvalOutcome::Complete(eval) => eval, + EvalOutcome::BudgetExhausted { needed } => { + return Err(anyhow!( + "budget too small for the baseline pass ({needed} rollouts needed)" + )); + } + } + }; + let baseline_score = baseline_eval.mean(); + let mut final_score = baseline_score; + + let mut store: RolloutStore = vec![None; num_examples]; + for rollout in &baseline_eval.rollouts { + store[rollout.example] = Some((rollout.eval.clone(), rollout.trace.clone())); + } + + // The incumbent interpreter: the caller's until an edit is accepted. + let mut owned: Option = None; + + let mut edits = Vec::new(); + let mut steps = Vec::new(); + let mut accepted = 0usize; + let mut rejected = 0usize; + + for generation in 0..self.num_iterations { + // 1. Shared minibatch; sorted for deterministic evaluation order. + let minibatch_size = num_examples.min(self.minibatch_size.max(1)); + let mut minibatch: Vec = all_indices + .choose_multiple(&mut rng, minibatch_size) + .copied() + .collect(); + minibatch.sort_unstable(); + + // Don't spend a reflection call on a child we can't afford to + // score on the minibatch. + if !engine.budget_allows(minibatch.len()) { + break; + } + + // 2. The menu against the current incumbent. + let menu = Self::menu(&program); + if menu.is_empty() { + tracing::warn!( + generation, + "no legal edits at any leaf; stopping the structural search" + ); + break; + } + + // Parent's mean on the shared minibatch — the gate threshold. + // Cache-served (the incumbent always has full coverage). + let threshold = { + let cur = owned.as_ref().unwrap_or(interp); + let target = OptimizeTarget::program(cur, examples, metric); + match engine + .evaluate(&target, incumbent_row, Some(&minibatch)) + .await? + { + EvalOutcome::Complete(eval) => eval.mean(), + EvalOutcome::BudgetExhausted { .. } => break, + } + }; + + // 3. Choose one edit from the menu. + let execution_feedback = format!( + "Incumbent minibatch mean: {threshold:.3}\n{}", + Self::summarize_feedback(&store, &minibatch) + ); + let (option, reflection_calls) = self + .choose( + &program, + &menu, + &execution_feedback, + generation, + reflector.as_ref(), + &mut rng, + ) + .await; + engine.charge(0, reflection_calls); + let entry = &menu[option]; + let edit = Self::materialize(&program, entry); + + let mut step = StructuralStep { + generation, + leaf: entry.leaf.clone(), + edit: edit.clone(), + parent_hash: program.meta.program_hash, + parent_minibatch_score: threshold, + child_minibatch_score: None, + accepted: false, + full_score: None, + rejection: None, + }; + + // 4. Apply, migrate, load — every failure skips the generation. + let child_program = match program.edited(std::slice::from_ref(&edit)) { + Ok(child) => child, + Err(err) => { + tracing::warn!(generation, error = %err, "edit failed to apply; skipping"); + step.rejection = Some(format!("edit failed: {err}")); + rejected += 1; + steps.push(step); + continue; + } + }; + let child_interp = match Interpreter::load(child_program, env()).await { + Ok(interp) => interp, + Err(err) => { + tracing::warn!(generation, error = %err, "edited child failed to load; skipping"); + step.rejection = Some(format!("load failed: {err}")); + rejected += 1; + steps.push(step); + continue; + } + }; + let migrated = migrate_overlay(&program, &incumbent, child_interp.program()); + let child_row = engine.register_overlay(migrated.clone()); + + // 5. The gate: child vs parent on the shared minibatch; only a + // strict win promotes to the full set. + let gate = { + let child_target = OptimizeTarget::program(&child_interp, examples, metric); + engine + .evaluate_gated(&child_target, child_row, &minibatch, threshold) + .await? + }; + match gate { + GateOutcome::BudgetExhausted { .. } => { + steps.push(step); + break; + } + GateOutcome::Rejected { minibatch: mb_eval } => { + step.child_minibatch_score = Some(mb_eval.mean()); + rejected += 1; + steps.push(step); + } + GateOutcome::Promoted { + minibatch: mb_eval, + full, + } => { + // The child is the new incumbent. Refresh the store from + // the full pass, preferring fresh minibatch traces over + // cache-served cells. + for rollout in &full.rollouts { + store[rollout.example] = + Some((rollout.eval.clone(), rollout.trace.clone())); + } + for rollout in &mb_eval.rollouts { + if rollout.trace.is_some() { + store[rollout.example] = + Some((rollout.eval.clone(), rollout.trace.clone())); + } + } + final_score = full.mean(); + step.child_minibatch_score = Some(mb_eval.mean()); + step.accepted = true; + step.full_score = Some(final_score); + accepted += 1; + steps.push(step); + edits.push(edit); + + program = Arc::clone(child_interp.program()); + incumbent = migrated; + incumbent_row = child_row; + owned = Some(child_interp); + } + } + } + + Ok(StructuralReport { + program, + overlay: incumbent, + baseline_score, + final_score, + edits, + steps, + accepted, + rejected, + spend: *engine.spend(), + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::LMConfig; + use crate::ir::{self, FieldType as T, ProgramBuilder, SignatureDef}; + + fn qa_program() -> Program { + let mut b = ProgramBuilder::new("structural_unit"); + b.model( + "m", + LMConfig { + model: "openai:gpt-4o-mini".to_string(), + ..LMConfig::default() + }, + ); + let qa = b.sig( + SignatureDef::build("QA") + .instruction("Answer the question.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let node = ir::predict("answerer", qa).bind("question", ir::input("question")); + b.main( + qa, + ir::seq([node]).out("answer", ir::out("answerer", "answer")), + ) + .unwrap() + } + + #[test] + fn menu_excludes_free_text_kinds_and_taken_reasoning_fields() { + let program = qa_program(); + let menu = Structural::menu(&program); + assert!(!menu.is_empty()); + assert!(menu.iter().all(|entry| !matches!( + entry.kind, + EditKind::SetStop | EditKind::SetInstructionDefault + ))); + assert!( + menu.iter() + .any(|entry| entry.kind == EditKind::AugmentSig && entry.leaf == "answerer") + ); + + // Once the reasoning field is on the leaf, AugmentSig leaves the menu. + let leaf = program.leaf_id("answerer").unwrap(); + let child = program + .edited(&[Edit::AugmentSig { + leaf, + prepend: cot_reasoning_field(), + }]) + .unwrap(); + let child_menu = Structural::menu(&child); + assert!( + child_menu + .iter() + .all(|entry| entry.kind != EditKind::AugmentSig) + ); + } + + #[test] + fn parse_choice_extracts_in_range_options() { + assert_eq!(Structural::parse_choice("3", 5), Some(3)); + assert_eq!(Structural::parse_choice(" Option 2.", 5), Some(2)); + assert_eq!(Structural::parse_choice("option 12 of 20", 20), Some(12)); + assert_eq!(Structural::parse_choice("9", 5), None); + assert_eq!(Structural::parse_choice("none of these", 5), None); + assert_eq!(Structural::parse_choice("", 5), None); + } + + #[test] + fn materialized_menu_entries_apply() { + let program = qa_program(); + for entry in Structural::menu(&program) { + let edit = Structural::materialize(&program, &entry); + // Remove orphans the out binding — validate.rs's call, surfaced + // as an EditError, which the loop records and skips. + let result = program.edited(std::slice::from_ref(&edit)); + if entry.kind == EditKind::Remove { + assert!(result.is_err(), "removing the only producer must fail"); + } else { + assert!( + result.is_ok(), + "{:?} should apply: {:?}", + entry.kind, + result.err() + ); + } + } + } +} diff --git a/crates/dspy-rs/src/optimizer/target.rs b/crates/dspy-rs/src/optimizer/target.rs new file mode 100644 index 00000000..28a917d7 --- /dev/null +++ b/crates/dspy-rs/src/optimizer/target.rs @@ -0,0 +1,645 @@ +//! What an optimizer optimizes: [`OptimizeTarget`], the lane-erased pair of +//! (thing under optimization, evaluation harness). +//! +//! Two lanes, one currency: +//! +//! - [`OptimizeTarget::module`] — a typed [`Module`] (+ [`Predictors`] +//! discovery), a trainset slice, and a [`TypedMetric`]. Candidate injection +//! is ambient ([`fx::with_params`](crate::fx::with_params)) — evaluation +//! never mutates the module; the single mutation is the caller-driven +//! [`install`](OptimizeTarget::install) of the winner at the end. +//! - [`OptimizeTarget::program`] — an interpreter-loaded IR +//! [`Program`](crate::ir::Program), labeled [`DemoRow`] examples, and a +//! [`ProgramMetric`]. Candidates read through +//! [`ir::Overlay`](crate::ir::Overlay)s at render time; the winner is +//! retrievable as an overlay ([`winner_overlay`](OptimizeTarget::winner_overlay)) +//! for [`Program::bake`](crate::ir::Program::bake). +//! +//! Construction runs the **naming pass** (module lane): every leaf the module +//! declares via [`Predictors`] is stamped with its declared name +//! ([`PredictorInfo::set_trace_name`]), so trace spans, candidate entries, and +//! persistence all address the same names. The target also snapshots +//! [`LeafInfo`] (schema text, current instruction, demos) — the read surface +//! strategies build candidates from — and computes the run's baseline +//! identity exactly once. + +use std::sync::Arc; +use std::time::Instant; + +use anyhow::{Result, anyhow}; +use futures::future::LocalBoxFuture; +use serde::Serialize; +use serde_json::Value; + +use crate::core::{PredictState, PredictorInfo, Predictors, ToInput}; +use crate::evaluate::{Eval, TypedMetric, evaluator::rollout_traced}; +use crate::ir::interp::{Budget as RunBudget, Interpreter}; +use crate::ir::params::{DemoRow, Overlay, ParamValue}; +use crate::ir::graph::Node; +use crate::optimizer::engine::{BoundCandidate, Candidate, CandidatePayload, canonical_hash}; +use crate::trace::{JsonMap, Trace, TraceMeta, TraceOutcome, capture_with_meta}; +use crate::Module; + +/// How a program-lane strategy tells the engine what "good" means: score one +/// interpreter output (`JsonMap` of the program's output signature fields) +/// against a labeled example. +/// +/// The JSON-native sibling of [`TypedMetric`](crate::evaluate::TypedMetric) — +/// loaded programs have no static output type, so the metric sees the same +/// value model the interpreter produces. The rollout's captured [`Trace`] is +/// always provided; LM-as-judge metrics run outside the capture scope exactly +/// as in the module lane. +#[allow(async_fn_in_trait)] +pub trait ProgramMetric: Send + Sync { + async fn evaluate( + &self, + example: &DemoRow, + output: &JsonMap, + trace: Option<&Trace>, + ) -> Result; +} + +/// The read surface strategies build candidates from: one optimizable leaf's +/// name, current values, and field contract. Snapshotted at target +/// construction (candidates never mutate the target mid-run, so the snapshot +/// stays valid for the whole optimization). +#[derive(Clone, Debug)] +pub struct LeafInfo { + /// The leaf's canonical name (trace component / candidate key). + pub name: String, + /// Current effective instruction (override or default). + pub instruction: String, + /// The signature/default instruction (ignoring overrides). + pub default_instruction: String, + /// Current demos as flat JSON rows. + pub demos: Vec, + /// Input fields as `(lm name, docs)` pairs. + pub input_fields: Vec<(String, String)>, + /// Output fields as `(lm name, docs)` pairs. + pub output_fields: Vec<(String, String)>, +} + +impl LeafInfo { + fn from_predictor(name: &str, info: &dyn PredictorInfo) -> Self { + let schema = info.schema(); + Self { + name: name.to_string(), + instruction: info.instruction(), + default_instruction: info.default_instruction(), + demos: info.demos_as_json(), + input_fields: schema + .input_fields() + .iter() + .map(|field| (field.lm_name.to_string(), field.docs.clone())) + .collect(), + output_fields: schema + .output_fields() + .iter() + .map(|field| (field.lm_name.to_string(), field.docs.clone())) + .collect(), + } + } + + /// Renders the leaf's input/output contract for reflection prompts + /// (GEPA/SIMBA format). + pub fn schema_for_reflection(&self) -> String { + let mut result = String::new(); + result.push_str("Input fields:\n"); + for (name, docs) in &self.input_fields { + let docs = if docs.is_empty() { "No description" } else { docs }; + result.push_str(&format!(" - {name}: {docs}\n")); + } + result.push_str("Output fields:\n"); + for (name, docs) in &self.output_fields { + let docs = if docs.is_empty() { "No description" } else { docs }; + result.push_str(&format!(" - {name}: {docs}\n")); + } + result + } +} + +/// The lane-erased evaluation surface the engine drives. Object-safe so the +/// [`Optimizer`](crate::optimizer::Optimizer) trait can be too. +pub(crate) trait Lane { + fn num_examples(&self) -> usize; + fn example_uid(&self, idx: usize) -> u64; + fn baseline(&self) -> u64; + fn bind(&self, payload: &CandidatePayload) -> Result; + fn run<'s>( + &'s self, + idx: usize, + bound: BoundCandidate, + candidate_hash: u64, + ) -> LocalBoxFuture<'s, Result<(Eval, Trace)>>; + fn output<'s>( + &'s self, + idx: usize, + bound: BoundCandidate, + ) -> LocalBoxFuture<'s, Result>; + fn install(&mut self, winner: &Candidate) -> Result<()>; + fn winner_overlay(&self) -> Option>; +} + +// --------------------------------------------------------------------------- +// Module lane +// --------------------------------------------------------------------------- + +struct ModuleLane<'a, E, M, MT> { + module: &'a mut M, + /// Validation examples first (empty when no valset), trainset after. + val: &'a [E], + train: &'a [E], + metric: &'a MT, + uids: Vec, + baseline: u64, +} + +impl<'a, E, M, MT> ModuleLane<'a, E, M, MT> +where + E: ToInput + Serialize + Sync, + M: Module + Predictors, + MT: TypedMetric, +{ + fn example(&self, idx: usize) -> &E { + if idx < self.val.len() { + &self.val[idx] + } else { + &self.train[idx - self.val.len()] + } + } +} + +/// Content identity of a module's optimizable state: the `predictors()` +/// snapshot hashed as `{name → PredictState}`. Computed once per target. +fn module_baseline(leaves: &[(String, &dyn PredictorInfo)]) -> u64 { + let states: std::collections::BTreeMap<&str, PredictState> = leaves + .iter() + .map(|(name, info)| (name.as_str(), info.dump_state())) + .collect(); + canonical_hash(&states) +} + +impl<'a, E, M, MT> Lane for ModuleLane<'a, E, M, MT> +where + E: ToInput + Serialize + Sync, + M: Module + Predictors, + MT: TypedMetric, +{ + fn num_examples(&self) -> usize { + self.val.len() + self.train.len() + } + + fn example_uid(&self, idx: usize) -> u64 { + self.uids[idx] + } + + fn baseline(&self) -> u64 { + self.baseline + } + + fn bind(&self, payload: &CandidatePayload) -> Result { + match payload { + CandidatePayload::Params { params, .. } => { + Ok(BoundCandidate::Params(Arc::clone(params))) + } + CandidatePayload::Overlay(_) => Err(anyhow!( + "an ir::Overlay candidate cannot be evaluated on a module target; \ + register a Candidate instead" + )), + } + } + + fn run<'s>( + &'s self, + idx: usize, + bound: BoundCandidate, + candidate_hash: u64, + ) -> LocalBoxFuture<'s, Result<(Eval, Trace)>> { + Box::pin(async move { + let BoundCandidate::Params(params) = bound else { + return Err(anyhow!("module lane received a non-params candidate")); + }; + let example = self.example(idx); + let meta = TraceMeta { + candidate_hash: Some(candidate_hash), + ..TraceMeta::default() + }; + let module: &M = &*self.module; + // The candidate is scoped ambiently around the whole traced + // rollout: each Predict leaf binds its own entry at call time. + crate::fx::with_params_shared( + params, + rollout_traced(module, example, self.metric, meta), + ) + .await + }) + } + + fn output<'s>( + &'s self, + idx: usize, + bound: BoundCandidate, + ) -> LocalBoxFuture<'s, Result> { + Box::pin(async move { + let BoundCandidate::Params(params) = bound else { + return Err(anyhow!("module lane received a non-params candidate")); + }; + let example = self.example(idx); + let input = example.to_input()?; + let module: &M = &*self.module; + let predicted = crate::fx::with_params_shared(params, module.call(input)) + .await + .map_err(|err| anyhow!("{err}"))?; + Ok(serde_json::to_value(predicted.into_inner()).unwrap_or(Value::Null)) + }) + } + + fn install(&mut self, winner: &Candidate) -> Result<()> { + let mut leaves = self.module.predictors_mut(); + for (name, slot) in &winner.slots { + let Some((_, info)) = leaves.iter_mut().find(|(leaf, _)| leaf == name) else { + return Err(anyhow!("predictor `{name}` not found in the module")); + }; + let mut state = info.dump_state(); + if slot.clear_instruction { + state.instruction_override = None; + } else if let Some(text) = &slot.instruction { + state.instruction_override = Some(text.clone()); + } + if let Some(demos) = &slot.demos { + state.demos = demos.clone(); + } + info.load_state(state) + .map_err(|err| anyhow!("failed to install winner on `{name}`: {err}"))?; + } + Ok(()) + } + + fn winner_overlay(&self) -> Option> { + None + } +} + +// --------------------------------------------------------------------------- +// Program lane +// --------------------------------------------------------------------------- + +struct ProgramLane<'a, MT> { + interp: &'a Interpreter, + examples: &'a [DemoRow], + metric: &'a MT, + uids: Vec, + program_tag: String, + winner: Option>, +} + +impl<'a, MT: ProgramMetric> Lane for ProgramLane<'a, MT> { + fn num_examples(&self) -> usize { + self.examples.len() + } + + fn example_uid(&self, idx: usize) -> u64 { + self.uids[idx] + } + + fn baseline(&self) -> u64 { + self.interp.program().meta.program_hash + } + + fn bind(&self, payload: &CandidatePayload) -> Result { + match payload { + CandidatePayload::Overlay(overlay) => Ok(BoundCandidate::Overlay(Arc::clone(overlay))), + CandidatePayload::Params { params, .. } => { + let overlay = params + .bind(self.interp.program()) + .map_err(|err| anyhow!("failed to bind candidate against program: {err}"))?; + Ok(BoundCandidate::Overlay(Arc::new(overlay))) + } + } + } + + fn run<'s>( + &'s self, + idx: usize, + bound: BoundCandidate, + candidate_hash: u64, + ) -> LocalBoxFuture<'s, Result<(Eval, Trace)>> { + Box::pin(async move { + let BoundCandidate::Overlay(overlay) = bound else { + return Err(anyhow!("program lane received a non-overlay candidate")); + }; + let example = &self.examples[idx]; + let meta = TraceMeta { + candidate_hash: Some(candidate_hash), + input: Some(example.input.clone()), + tags: [("program".to_string(), self.program_tag.clone())] + .into_iter() + .collect(), + ..TraceMeta::default() + }; + let started = Instant::now(); + let (result, mut trace) = capture_with_meta(meta, || { + self.interp + .run(example.input.clone(), Some(overlay), RunBudget::unlimited()) + }) + .await; + let output = + result.map_err(|err| anyhow!("candidate failed on example {idx}: {err}"))?; + // Metric runs outside the capture scope so LM-as-judge metrics + // don't pollute the execution trace. + let eval = self.metric.evaluate(example, &output, Some(&trace)).await?; + trace.outcome = Some(TraceOutcome { + output: Some(output), + error: None, + eval: Some(eval.clone()), + duration_us: started.elapsed().as_micros() as u64, + }); + Ok((eval, trace)) + }) + } + + fn output<'s>( + &'s self, + idx: usize, + bound: BoundCandidate, + ) -> LocalBoxFuture<'s, Result> { + Box::pin(async move { + let BoundCandidate::Overlay(overlay) = bound else { + return Err(anyhow!("program lane received a non-overlay candidate")); + }; + let example = &self.examples[idx]; + let output = self + .interp + .run(example.input.clone(), Some(overlay), RunBudget::unlimited()) + .await + .map_err(|err| anyhow!("{err}"))?; + Ok(Value::Object(output)) + }) + } + + fn install(&mut self, winner: &Candidate) -> Result<()> { + let overlay = winner + .to_params() + .bind(self.interp.program()) + .map_err(|err| anyhow!("failed to bind winner against program: {err}"))?; + self.winner = Some(Arc::new(overlay)); + Ok(()) + } + + fn winner_overlay(&self) -> Option> { + self.winner.clone() + } +} + +/// Leaf snapshot for a loaded program: one [`LeafInfo`] per `predict`/`agent` +/// node, current values read from the param slot defaults. +fn program_leaves(interp: &Interpreter) -> Vec { + let program = interp.program(); + let mut leaves = Vec::new(); + for (node_id, node) in program.nodes.iter() { + let sig = match node { + Node::Predict(n) => n.sig, + Node::AgentLoop(n) => n.sig, + _ => continue, + }; + let Some(name) = program.leaf_name(node_id) else { + continue; + }; + let def = &program.sigs[sig]; + let instruction = program + .param_id(&format!("{name}.instruction")) + .and_then(|id| match &program.params[id].default { + ParamValue::Instruction { text } => Some(text.clone()), + _ => None, + }) + .unwrap_or_else(|| def.instruction.to_string()); + let demos = program + .param_id(&format!("{name}.demos")) + .and_then(|id| match &program.params[id].default { + ParamValue::Demos { rows } => Some( + rows.iter() + .map(|row| { + let mut flat = row.input.clone(); + flat.extend(row.output.iter().map(|(k, v)| (k.clone(), v.clone()))); + flat + }) + .collect(), + ), + _ => None, + }) + .unwrap_or_default(); + let field_pairs = |fields: &[crate::ir::sig::FieldDef]| { + fields + .iter() + .map(|field| { + ( + field.lm_name.to_string(), + field.docs.as_deref().unwrap_or("").to_string(), + ) + }) + .collect() + }; + leaves.push(LeafInfo { + name: name.to_string(), + default_instruction: instruction.clone(), + instruction, + demos, + input_fields: field_pairs(&def.inputs), + output_fields: field_pairs(&def.outputs), + }); + } + leaves +} + +// --------------------------------------------------------------------------- +// OptimizeTarget +// --------------------------------------------------------------------------- + +/// The thing an [`Optimizer`](crate::optimizer::Optimizer) optimizes: a +/// module or a program, packaged with its example set and metric. See the +/// module docs for the two lanes. +pub struct OptimizeTarget<'a> { + lane: Box, + leaves: Vec, + /// Number of leading example columns that are the validation set, when a + /// separate valset was supplied. + val_len: Option, +} + +impl<'a> OptimizeTarget<'a> { + /// A module-lane target: typed module + trainset (by reference) + metric. + /// + /// Runs the naming pass: each leaf declared by [`Predictors`] is stamped + /// with its declared name so traces, candidates, and persistence agree. + pub fn module(module: &'a mut M, trainset: &'a [E], metric: &'a MT) -> Self + where + E: ToInput + Serialize + Sync, + M: Module + Predictors, + MT: TypedMetric, + { + Self::module_with_valset(module, trainset, None, metric) + } + + /// [`module`](Self::module) with an optional validation set. + /// + /// When `valset` is `Some`, its examples become the *leading* columns of + /// the target ([`val_columns`](Self::val_columns)) and the trainset the + /// trailing ones ([`train_columns`](Self::train_columns)) — the layout + /// GEPA's Pareto bookkeeping uses. When `None`, both views cover the + /// whole trainset. + pub fn module_with_valset( + module: &'a mut M, + trainset: &'a [E], + valset: Option<&'a [E]>, + metric: &'a MT, + ) -> Self + where + E: ToInput + Serialize + Sync, + M: Module + Predictors, + MT: TypedMetric, + { + // Naming pass: stamp each leaf with its declared name (once per run). + for (name, info) in module.predictors_mut() { + info.set_trace_name(&name); + } + let named = module.predictors(); + let leaves: Vec = named + .iter() + .map(|(name, info)| LeafInfo::from_predictor(name, *info)) + .collect(); + let baseline = module_baseline(&named); + drop(named); + + let val_len = valset.map(<[E]>::len); + let val = valset.unwrap_or(&[]); + let uids = val + .iter() + .chain(trainset.iter()) + .map(canonical_hash) + .collect(); + + Self { + lane: Box::new(ModuleLane { + module, + val, + train: trainset, + metric, + uids, + baseline, + }), + leaves, + val_len, + } + } + + /// A program-lane target: interpreter-loaded program + labeled examples + /// (by reference) + JSON metric. + pub fn program(interp: &'a Interpreter, examples: &'a [DemoRow], metric: &'a MT) -> Self + where + MT: ProgramMetric, + { + let leaves = program_leaves(interp); + let uids = examples.iter().map(canonical_hash).collect(); + let program_tag = format!("{:016x}", interp.program().meta.program_hash); + Self { + lane: Box::new(ProgramLane { + interp, + examples, + metric, + uids, + program_tag, + winner: None, + }), + leaves, + val_len: None, + } + } + + /// The optimizable leaves' read surface (snapshotted at construction). + pub fn leaves(&self) -> &[LeafInfo] { + &self.leaves + } + + pub fn num_examples(&self) -> usize { + self.lane.num_examples() + } + + /// Whether this target carries a distinct validation set. + pub fn has_valset(&self) -> bool { + self.val_len.is_some() + } + + /// The scoring columns: the validation prefix when a valset was supplied, + /// else every example. + pub fn val_columns(&self) -> Vec { + match self.val_len { + Some(len) => (0..len).collect(), + None => (0..self.num_examples()).collect(), + } + } + + /// The minibatch-sampling pool: the trainset suffix when a valset was + /// supplied, else every example. + pub fn train_columns(&self) -> Vec { + match self.val_len { + Some(len) => (len..self.num_examples()).collect(), + None => (0..self.num_examples()).collect(), + } + } + + /// Installs the winning candidate — the **one** mutation of the run. + /// + /// Module lane: merges each slot into the named leaf's state through + /// [`PredictorInfo::load_state`]. Program lane: binds the winner to an + /// overlay retrievable via [`winner_overlay`](Self::winner_overlay) + /// (bake it with [`Program::bake`](crate::ir::Program::bake)). + pub fn install(&mut self, winner: &Candidate) -> Result<()> { + self.lane.install(winner) + } + + /// The installed winner as a bound overlay (program lane only). + pub fn winner_overlay(&self) -> Option> { + self.lane.winner_overlay() + } + + pub(crate) fn baseline(&self) -> u64 { + self.lane.baseline() + } + + pub(crate) fn example_uid(&self, idx: usize) -> u64 { + self.lane.example_uid(idx) + } + + pub(crate) fn bind(&self, payload: &CandidatePayload) -> Result { + self.lane.bind(payload) + } + + pub(crate) async fn run( + &self, + idx: usize, + bound: BoundCandidate, + candidate_hash: u64, + ) -> Result<(Eval, Trace)> { + self.lane.run(idx, bound, candidate_hash).await + } + + /// Runs the given examples under `candidate` and returns the bare output + /// values (no metric, no trace capture) — GEPA's best-output collection. + /// Sequential, in index order. + pub async fn candidate_outputs( + &self, + indices: &[usize], + candidate: &Candidate, + ) -> Result> { + let payload = CandidatePayload::Params { + candidate: candidate.clone(), + params: Arc::new(candidate.to_params()), + }; + let bound = self.lane.bind(&payload)?; + let mut outputs = Vec::with_capacity(indices.len()); + for &idx in indices { + outputs.push(self.lane.output(idx, bound.clone()).await?); + } + Ok(outputs) + } +} diff --git a/crates/dspy-rs/src/predictors/predict.rs b/crates/dspy-rs/src/predictors/predict.rs index 8d4fe99a..81205a18 100644 --- a/crates/dspy-rs/src/predictors/predict.rs +++ b/crates/dspy-rs/src/predictors/predict.rs @@ -1,19 +1,46 @@ use anyhow::Result; +use indexmap::IndexMap; use rig::tool::ToolDyn; use serde_json::{Map, Value}; use std::marker::PhantomData; -use std::ops::ControlFlow; use std::sync::{Arc, OnceLock}; use tracing::{debug, trace}; -use crate as dsrs; use crate::core::lm::ToolSet; -use crate::core::{DynPredictor, Module, PredictAccessorFns, PredictState, Signature, StateUpdate}; +use crate::core::{Module, PredictState, Signature}; +use crate::ir::{ + self, Budget, Interpreter, Overlay, Program, RunError, RunOutput, RuntimeEnv, SignatureDef, +}; use crate::{ - CallMetadata, Chat, ChatAdapter, FieldSchema, GLOBAL_SETTINGS, LmError, LmUsage, Message, - PredictError, Predicted, Schema, SignatureSchema, + CallMetadata, Chat, FieldSchema, GLOBAL_SETTINGS, LmError, LmUsage, ParseError, PredictError, + Predicted, Schema, SignatureSchema, }; +/// Loop options for a tooled predictor's 1-node `agent` program. +/// +/// The typed carrier for `#[agent(...)]` attribute options on the standalone +/// call path: `stop_tools`/`max_turns`/`until_parse` land in the node's +/// [`StopSpec`](crate::ir::StopSpec), `budget` in its +/// [`NodeBudget`](crate::ir::NodeBudget), and `context` in its +/// [`ContextPolicy`](crate::ir::ContextPolicy) slot. `None`s fall back to the +/// IR defaults ([`StopSpec::default`](crate::ir::StopSpec): 8 turns, +/// `until_parse = true`, no stop tools). +/// +/// Only meaningful on a predictor with tools attached +/// ([`PredictBuilder::with_tools`]); ignored otherwise. +#[derive(Clone, Debug, Default)] +pub struct AgentLoopSpec { + /// Names (⊆ the attached tools) whose call ends the loop; the call's args + /// become the raw final output (the "submit answer" pattern). + pub stop_tools: Vec, + /// `None` = the IR default (8). + pub max_turns: Option, + /// `None` = the IR default (`true`). + pub until_parse: Option, + pub budget: crate::ir::NodeBudget, + pub context: crate::ir::ContextPolicy, +} + /// A typed input/output pair for few-shot prompting. /// /// Demos are formatted as user/assistant exchanges in the prompt, showing the LM @@ -42,54 +69,24 @@ impl Demo { } } -fn predict_dyn_visit( - value: *mut (), - visitor: &mut dyn FnMut(&mut dyn DynPredictor) -> ControlFlow<()>, -) -> ControlFlow<()> -where - S: Signature, -{ - // SAFETY: this function is only called through the shape-local - // `dsrs::predict_accessor` payload attached to a shape with strict - // `Predict` identity (`type_identifier` + `module_path`). - let typed = unsafe { &mut *(value.cast::>()) }; - visitor(typed) -} - -type VisitPredictorMutFn = - fn(*mut (), &mut dyn FnMut(&mut dyn DynPredictor) -> ControlFlow<()>) -> ControlFlow<()>; - -trait PredictAccessorProvider { - const VISIT_MUT: VisitPredictorMutFn; -} - -impl PredictAccessorProvider for S -where - S: Signature, -{ - const VISIT_MUT: VisitPredictorMutFn = predict_dyn_visit::; -} - /// The leaf module. The only thing in the system that actually calls the LM. /// /// One `Predict` = one prompt template = one LM call. It takes a [`Signature`]'s fields /// and instruction, formats them into a prompt (with any demos and tools), calls the /// configured LM, and parses the response back into `S::Output`. Every other module — -/// [`ChainOfThought`](crate::ChainOfThought), `ReAct`, custom pipelines — ultimately +/// [`ChainOfThought`](crate::ChainOfThought), custom pipelines — ultimately /// delegates to one or more `Predict` leaves. /// /// This is also the unit of optimization. When an optimizer tunes your program, it's /// adjusting `Predict` leaves: their demos (few-shot examples) and instructions. -/// The optimizer's Facet walker discovers leaves automatically from struct fields — -/// no `#[parameter]` annotations or manual traversal needed. /// /// # Optimizer discovery /// -/// `Predict` encodes shape-local discovery payloads: -/// - strict shape identity (`type_identifier` + `module_path`) identifies the leaf -/// - `dsrs::predict_accessor` stores the typed mutable accessor visitor -/// -/// The optimizer walker consumes these through `visit_named_predictors_mut`. +/// Modules declare their `Predict` leaves by name through the +/// [`Predictors`](crate::Predictors) trait (see the `predictors!` macro); +/// optimizers read each leaf through its [`PredictorInfo`](crate::PredictorInfo) +/// view and inject candidates ambiently per call +/// ([`fx::with_params`](crate::fx::with_params)) — never by mutating the leaf. /// There is no runtime registration side effect in `new()` or `build()`. /// /// ```no_run @@ -115,22 +112,19 @@ where /// ``` #[derive(facet::Facet)] #[facet(crate = facet, opaque)] -#[facet(dsrs::predict_accessor = &PredictAccessorFns { - visit_mut: ::VISIT_MUT, -})] pub struct Predict { #[facet(skip, opaque)] tools: Vec>, + /// Loop options applied to the 1-node `agent` program when tools are + /// attached (stop spec, node budget, context policy). Settable only at + /// build time, like the tools themselves. + #[facet(skip, opaque)] + agent_spec: Option, #[facet(skip, opaque)] demos: Vec>, instruction_override: Option, #[facet(skip, opaque)] lm: Option>, - /// Formatted system + demo messages, built once per (instruction, demos) - /// configuration. Reset by every mutator (`set_instruction`, - /// `set_demos_from_examples`, `load_state`). - #[facet(skip, opaque)] - prompt_prefix: OnceLock>, /// Pre-fetched tool definitions + name-indexed executors. Tools are only /// settable at build time, so this never needs invalidation. #[facet(skip, opaque)] @@ -140,6 +134,23 @@ pub struct Predict { /// optimizer naming pass. #[facet(skip, opaque)] trace_name: Option, + /// The cached 1-node IR program this predictor executes: a `predict` leaf + /// named [`component_name`](Predict::component_name) over + /// [`SignatureDef::of::()`]. Reset by [`set_trace_name`] (the leaf name + /// is part of the program). + #[facet(skip, opaque)] + program: OnceLock>, + /// Instance state (instruction override + demos) minted as an [`Overlay`] + /// against [`program`](Self::program). `None` inner = no overrides (the + /// program defaults read through). Reset by every state mutator and by + /// `set_trace_name`. + #[facet(skip, opaque)] + instance_overlay: OnceLock>>, + /// The loaded interpreter, keyed by the LM it was bound with; rebuilt when + /// the resolved LM changes (per-instance LM is fixed at build, but the + /// global [`configure`](crate::configure)d LM can change between calls). + #[facet(skip, opaque)] + engine: tokio::sync::Mutex, Arc)>>, #[facet(skip, opaque)] _marker: PhantomData, } @@ -149,12 +160,15 @@ impl Predict { pub fn new() -> Self { Self { tools: Vec::new(), + agent_spec: None, demos: Vec::new(), instruction_override: None, lm: None, - prompt_prefix: OnceLock::new(), toolset: tokio::sync::OnceCell::new(), trace_name: None, + program: OnceLock::new(), + instance_overlay: OnceLock::new(), + engine: tokio::sync::Mutex::new(None), _marker: PhantomData, } } @@ -167,32 +181,38 @@ impl Predict { /// The typed write path for optimizable state (instruction override + demos). /// /// This is the only place that assigns those fields and invalidates the - /// cached prompt prefix; the builder and the type-erased - /// `DynPredictor::apply_update` seam both funnel here. `None` leaves a - /// field untouched. - fn apply_state( - &mut self, - instruction: Option>, - demos: Option>>, - ) { + /// cached instance overlay; the builder and the [`PredictorInfo::load_state`] + /// install seam both funnel here. `None` leaves a field untouched. + /// + /// [`PredictorInfo::load_state`]: crate::core::PredictorInfo::load_state + fn apply_state(&mut self, instruction: Option>, demos: Option>>) { if let Some(instruction) = instruction { self.instruction_override = instruction; } if let Some(demos) = demos { self.demos = demos; } - self.prompt_prefix = OnceLock::new(); + self.instance_overlay = OnceLock::new(); } - /// The typed direct call: builds the prompt, calls the LM, and parses the response. + /// The typed direct call: renders the prompt, calls the LM, and parses the + /// response — all through the IR interpreter. /// - /// The full pipeline: - /// 1. Format system message from the signature's schema and instruction override - /// 2. Format demo examples as user/assistant exchanges - /// 3. Format the input as the final user message - /// 4. Call the LM (with any tools attached) - /// 5. Parse the response into `S::Output` via the `[[ ## field ## ]]` protocol - /// 6. Record a trace span if inside a [`capture()`](crate::trace::capture) scope + /// `Predict` is a thin typed handle over a 1-node IR [`Program`] (a + /// `predict` leaf over [`SignatureDef::of::()`]) executed by + /// [`Interpreter::run_collecting`]: + /// 1. Instance state (instruction override + demos) reads through an + /// [`Overlay`]; an ambient optimizer overlay + /// ([`ir::current_overlay`](crate::ir::current_overlay)) composes on + /// top (ambient entries win per slot). + /// 2. The interpreter renders system/demos/input via the def-lane + /// [`ChatAdapter`] (byte-identical prompts), calls the bound LM, and + /// parses the `[[ ## field ## ]]` response. + /// 3. A trace span is recorded under this predictor's component name when + /// inside a [`capture()`](crate::trace::capture) scope; replay scopes + /// intercept above the LM exactly as before. + /// 4. The typed `Predicted` is reassembled from the run's + /// single [`LeafOutcome`](crate::ir::LeafOutcome). /// /// [`Module::forward`] delegates here; for multi-turn conversations, build the /// chat yourself and use [`call_and_parse`](Predict::call_and_parse). @@ -218,98 +238,316 @@ impl Predict { S::Input: Schema, S::Output: Schema, { - // Serialize the input for trace recording only when a scope is active. - let capture_input = if crate::trace::is_capturing() { - json_map_from_input::(&input).ok() + let program = self.program().await?; + let overlay = self.effective_overlay(&program)?; + let lm = self.resolve_lm(); + let interpreter = self.interpreter(&program, &lm).await?; + let input_map = json_map_from_input::(&input) + .map_err(|err| internal_error(format!("failed to serialize input: {err}")))?; + let run = interpreter + .run_collecting(input_map, overlay, Budget::unlimited()) + .await + .map_err(map_run_error::)?; + predicted_from_run::(run) + } + + /// Returns the cached 1-node program, building it on first use: a + /// `predict` leaf, or an `agent` leaf when tools are attached (the IR + /// says `Predict` carries no tools — a tooled predictor *is* an agent + /// loop). + async fn program(&self) -> Result, PredictError> { + if let Some(program) = self.program.get() { + return Ok(Arc::clone(program)); + } + let built = if self.tools.is_empty() { + self.build_predict_program()? } else { - None + let toolset = self + .cached_toolset() + .await + .expect("tools are non-empty, so the toolset exists"); + self.build_agent_program(toolset.definitions())? }; - let chat = self.build_chat(&input)?; - // The chat is prefix + one live user message; everything before the - // final message is the interned span prefix. - let prefix_len = chat.len().saturating_sub(1); - let (predicted, _) = self - .call_and_parse_with_input(chat, capture_input, prefix_len) - .await?; - Ok(predicted) + // A concurrent call may have won the race — both builds are identical. + let _ = self.program.set(Arc::new(built)); + Ok(Arc::clone(self.program.get().expect("program set above"))) } - /// Builds the first-turn chat from the signature, demos, and input. - /// - /// Returns a [`Chat`] ready to pass to [`call_and_parse`](Predict::call_and_parse). - /// Useful when you need to inspect or modify the prompt before sending it to - /// the LM. - /// - /// The system message and demo turns are formatted once per (instruction, - /// demos) configuration and cached on the instance — only the live user - /// message is formatted per call. + /// Builds the 1-node `agent` program for a tooled predictor: same leaf + /// name and signature as the `predict` form, plus a host-tool declaration + /// per attached tool. Loop behavior comes from the attached + /// [`AgentLoopSpec`] when one was built in + /// ([`PredictBuilder::with_agent_spec`]); otherwise it is the IR's + /// [`StopSpec`](crate::ir::StopSpec) default — `until_parse` with + /// `max_turns = 8` — the closest IR expression of the old LM-layer auto + /// tool loop. #[allow(clippy::result_large_err)] - pub fn build_chat(&self, input: &S::Input) -> Result - where - S::Input: Schema, - { - let prefix = self.prompt_prefix()?; - let user = ChatAdapter.format_user_message_typed::(input); + fn build_agent_program( + &self, + definitions: &[rig::completion::ToolDefinition], + ) -> Result { + let leaf = self.component_name(); + let def = SignatureDef::of::(); + let mut b = ir::ProgramBuilder::new(leaf); + let model = b.model( + "default", + crate::ir::module_build::unbound_model_config("default"), + ); + let sid = b.sig_of::(); + b.add_types(&::output_schema().types); + + let mut tool_ids = Vec::with_capacity(definitions.len()); + let mut tool_names = Vec::with_capacity(definitions.len()); + let mut seen = std::collections::HashSet::new(); + for definition in definitions { + // Duplicate names keep the first tool, mirroring `ToolSet::build`. + if !seen.insert(definition.name.as_str()) { + continue; + } + let (tool_sig, tool_types) = + tool_signature_from_definition(&definition.name, &definition.parameters); + b.add_types(&tool_types); + let tool_sid = b.sig(tool_sig); + tool_ids.push(b.host_tool(&definition.name, &definition.description, tool_sid, &[])); + tool_names.push(definition.name.clone()); + } - let mut messages = Vec::with_capacity(prefix.len() + 1); - messages.extend(prefix.iter().cloned()); - messages.push(Message::user(user)); - let chat = Chat::new(messages); - trace!(message_count = chat.len(), "chat constructed"); - Ok(chat) + let mut ns = ir::agent(leaf, sid).model(model); + if let Some(spec) = &self.agent_spec { + let mut stop_ids = Vec::with_capacity(spec.stop_tools.len()); + for name in &spec.stop_tools { + let position = tool_names.iter().position(|n| n == name).ok_or_else(|| { + internal_error(format!( + "stop tool `{name}` is not among the attached tools" + )) + })?; + stop_ids.push(tool_ids[position]); + } + ns = ns.stop_tools(stop_ids); + if let Some(turns) = spec.max_turns { + if turns == 0 { + return Err(internal_error("agent `max_turns` must be > 0".to_string())); + } + ns = ns.max_turns(turns); + } + if let Some(until_parse) = spec.until_parse { + ns = ns.until_parse(until_parse); + } + ns = ns.budget(spec.budget.clone()).context(spec.context.clone()); + } + let mut ns = ns.tools(tool_ids); + for field in def.inputs.iter() { + ns = ns.bind(&field.name, ir::input(&field.name)); + } + let mut root = ir::seq([ns]); + for field in def.outputs.iter() { + root = root.out(&field.name, ir::out(leaf, &field.name)); + } + b.main(sid, root) + .map_err(|err| internal_error(format!("failed to build agent program: {err}"))) + } + + /// Builds the 1-node `predict` program: leaf name = the trace-name + /// convention (so span identity and capture/replay keying are unchanged), + /// signature = `SignatureDef::of::()`, model = the `default` ref bound + /// through [`RuntimeEnv`] at load. + #[allow(clippy::result_large_err)] + fn build_predict_program(&self) -> Result { + let leaf = self.component_name(); + let def = SignatureDef::of::(); + let mut b = ir::ProgramBuilder::new(leaf); + let model = b.model( + "default", + crate::ir::module_build::unbound_model_config("default"), + ); + let sid = b.sig_of::(); + // `sig_of` merges output-reachable class/enum defs; input-reachable + // ones are needed too (the interpreter type-checks run inputs). + b.add_types(&::output_schema().types); + let mut ns = ir::predict(leaf, sid).model(model); + for field in def.inputs.iter() { + ns = ns.bind(&field.name, ir::input(&field.name)); + } + let mut root = ir::seq([ns]); + for field in def.outputs.iter() { + root = root.out(&field.name, ir::out(leaf, &field.name)); + } + b.main(sid, root) + .map_err(|err| internal_error(format!("failed to build predict program: {err}"))) } - /// Returns the cached system + demo message prefix, building it on first use. + /// The instance overlay: instruction override + demos as slot values + /// against this predictor's program. `None` when neither is set. #[allow(clippy::result_large_err)] - fn prompt_prefix(&self) -> Result<&[Message], PredictError> + fn instance_overlay(&self, program: &Arc) -> Result>, PredictError> where S::Input: Schema, + S::Output: Schema, { - if self.prompt_prefix.get().is_none() { - let built = self.build_prompt_prefix()?; - // A concurrent forward may have won the race — that's fine, both - // builds produce identical messages. - let _ = self.prompt_prefix.set(built); + if let Some(cached) = self.instance_overlay.get() { + return Ok(cached.clone()); } + let overlay = if self.instruction_override.is_none() && self.demos.is_empty() { + None + } else { + let state = PredictState { + demos: crate::core::PredictorInfo::demos_as_json(self), + instruction_override: self.instruction_override.clone(), + }; + let minted = + crate::ir::bridge::states_to_overlay(program, [(self.component_name(), &state)]) + .map_err(|err| { + internal_error(format!("failed to mint instance overlay: {err}")) + })?; + Some(Arc::new(minted)) + }; + let _ = self.instance_overlay.set(overlay); Ok(self - .prompt_prefix + .instance_overlay .get() - .expect("prompt prefix initialized above")) + .expect("instance overlay set above") + .clone()) } + /// The overlay a run reads through: instance state composed with the + /// ambient candidate scopes, later layers winning per slot: + /// + /// 1. instance state (instruction override + demos); + /// 2. an ambient optimizer overlay + /// ([`ir::current_overlay`](crate::ir::current_overlay)) — ignored when + /// minted against a *different* program (one scope can span several + /// modules; only the matching one accepts it); + /// 3. the ambient [`fx::Params`](crate::fx::Params) entry matching this + /// predictor's component name ([`fx::with_params`](crate::fx::with_params) + /// — the optimizer's candidate-injection scope). Explicit clears + /// resolve to the program slot defaults, so a candidate can reset a + /// slot past instance state. + /// + /// Ambient entries winning over instance entries preserves exactly the + /// precedence the old apply/restore mutation seam had. #[allow(clippy::result_large_err)] - fn build_prompt_prefix(&self) -> Result, PredictError> + fn effective_overlay( + &self, + program: &Arc, + ) -> Result>, PredictError> where S::Input: Schema, + S::Output: Schema, { - let chat_adapter = ChatAdapter; - let system = match chat_adapter - .format_system_message_typed_with_instruction::(self.instruction_override.as_deref()) - { - Ok(system) => system, - Err(err) => { - return Err(PredictError::Lm { - source: LmError::Provider { - provider: "internal".to_string(), - message: err.to_string(), - source: None, - }, - }); + let instance = self.instance_overlay(program)?; + let ambient = crate::ir::bridge::current_overlay() + .filter(|overlay| overlay.base == program.meta.program_hash); + let params_values = match crate::fx::ambient_entry(self.component_name()) { + Some(entry) => { + crate::ir::bridge::entry_slot_values(program, self.component_name(), &entry) + .map_err(|err| { + internal_error(format!("failed to bind ambient params: {err}")) + })? } + None => Vec::new(), }; - trace!(system_len = system.len(), "typed system prompt formatted"); - - let mut messages = Vec::with_capacity(1 + self.demos.len() * 2); - messages.push(Message::system(system)); - for demo in &self.demos { - messages.push(Message::user( - chat_adapter.format_user_message_typed::(&demo.input), - )); - messages.push(Message::assistant( - chat_adapter.format_assistant_message_typed::(&demo.output), - )); + + let mut merged: Option = match (instance, ambient) { + (None, None) => None, + (Some(instance), None) => Some((*instance).clone()), + (None, Some(ambient)) => Some((*ambient).clone()), + (Some(instance), Some(ambient)) => { + let mut merged = (*instance).clone(); + for (id, value) in ambient.entries() { + merged.set(program, id, value.clone()).map_err(|err| { + internal_error(format!("failed to compose ambient overlay: {err}")) + })?; + } + Some(merged) + } + }; + + if !params_values.is_empty() { + let mut overlay = merged.take().unwrap_or_else(|| Overlay::new(program)); + for (id, value) in params_values { + overlay.set(program, id, value).map_err(|err| { + internal_error(format!("failed to compose ambient params: {err}")) + })?; + } + merged = Some(overlay); } - Ok(messages) + + Ok(merged.map(Arc::new)) + } + + /// Resolves the LM this call uses: instance LM > global + /// [`configure()`](crate::configure) settings. Panics exactly like the + /// pre-IR path when no global LM is configured and no instance LM is set. + fn resolve_lm(&self) -> Arc { + match &self.lm { + Some(lm) => Arc::clone(lm), + None => { + let guard = GLOBAL_SETTINGS.read().unwrap(); + let settings = guard.as_ref().unwrap(); + Arc::clone(&settings.lm) + } + } + } + + /// Returns the loaded interpreter for `lm`, reloading when the resolved + /// LM changed since the last call (the program is bound to its model at + /// load, not per call). + async fn interpreter( + &self, + program: &Arc, + lm: &Arc, + ) -> Result, PredictError> { + let mut slot = self.engine.lock().await; + if let Some((bound, interpreter)) = slot.as_ref() + && Arc::ptr_eq(bound, lm) + { + return Ok(Arc::clone(interpreter)); + } + let env = self.runtime_env(lm); + let interpreter = Interpreter::load((**program).clone(), env) + .await + .map_err(|err| internal_error(format!("failed to load predict program: {err}")))?; + let interpreter = Arc::new(interpreter); + *slot = Some((Arc::clone(lm), Arc::clone(&interpreter))); + Ok(interpreter) + } + + /// The runtime environment a load binds against: the resolved model under + /// the `default` ref, plus every attached tool bound as a host tool. + fn runtime_env(&self, lm: &Arc) -> RuntimeEnv { + let mut env = RuntimeEnv::new().bind_model("default", Arc::clone(lm)); + for tool in &self.tools { + env = env.bind_host_tool(&tool.name(), Arc::clone(tool)); + } + env + } + + /// Builds the first-turn chat from the signature, demos, and input. + /// + /// Returns a [`Chat`] ready to pass to [`call_and_parse`](Predict::call_and_parse). + /// Useful when you need to inspect or modify the prompt before sending it to + /// the LM. + /// + /// Thin wrapper over the interpreter's + /// [`conversation_opening`](crate::ir::Interpreter::conversation_opening): + /// the same overlay-resolved rendering the typed [`call`](Predict::call) + /// path does, so the opening prompt is byte-identical across both. + pub async fn build_chat(&self, input: &S::Input) -> Result + where + S::Input: Schema, + S::Output: Schema, + { + let program = self.program().await?; + let overlay = self.effective_overlay(&program)?; + let lm = self.resolve_lm(); + let interpreter = self.interpreter(&program, &lm).await?; + let input_map = json_map_from_input::(input) + .map_err(|err| internal_error(format!("failed to serialize input: {err}")))?; + let chat = interpreter + .conversation_opening(&input_map, overlay) + .map_err(map_run_error::)?; + trace!(message_count = chat.len(), "chat constructed"); + Ok(chat) } /// The component name recorded on trace spans: the assigned `trace_name` @@ -341,8 +579,9 @@ impl Predict { )) } - /// The chat-level call: sends `chat` to the LM and parses the response, - /// returning both the prediction and the updated conversation history. + /// The chat-level call: sends `chat` through the interpreter's + /// conversation surface and parses the response, returning both the + /// prediction and the updated conversation history. /// /// This is the one conversation seam. The caller owns the `Chat` between /// turns: @@ -354,6 +593,13 @@ impl Predict { /// Every turn parses with the same `[[ ## field ## ]]` protocol; the caller /// is responsible for including format instructions in follow-up messages if /// the model needs reminding of the output format. + /// + /// Thin wrapper over + /// [`Interpreter::run_conversation`](crate::ir::Interpreter::run_conversation): + /// each turn is one interpreter conversation turn — one trace span, + /// replay intercepted above the LM, and (when tools are attached) the + /// same `AgentLoop` the typed [`call`](Predict::call) path runs, with + /// tool calls dispatched through the attached executors. pub async fn call_and_parse( &self, chat: Chat, @@ -363,288 +609,16 @@ impl Predict { S::Output: Schema, { trace!(message_count = chat.len(), "chat-level call"); - self.call_and_parse_with_input(chat, None, 0).await - } - - /// [`call_and_parse`](Predict::call_and_parse) with the typed input captured - /// for trace recording. `capture_input` is only recorded when a capture - /// scope is active; pass `None` when the input is unavailable (e.g. - /// multi-turn continuations). `prefix_len` is the number of leading chat - /// messages that are the cached system+demos prefix (0 for caller-owned chats). - async fn call_and_parse_with_input( - &self, - chat: Chat, - capture_input: Option>, - prefix_len: usize, - ) -> Result<(Predicted, Chat), PredictError> - where - S::Input: Schema, - S::Output: Schema, - { - let lm = match &self.lm { - Some(lm) => Arc::clone(lm), - None => { - let guard = GLOBAL_SETTINGS.read().unwrap(); - let settings = guard.as_ref().unwrap(); - Arc::clone(&settings.lm) - } - }; - - // Open the span before the LM call: a call that dies mid-flight still - // leaves (component, seq, prompt, input) in the trace. - let guard = crate::trace::begin_span(crate::trace::SpanRequest { - component: self.component_name(), - prefix: (prefix_len > 0).then(|| &chat.messages[..prefix_len]), - suffix: &chat.messages[prefix_len..], - input: capture_input, - model: &lm.config, - request_hash: None, - }); - - // Replay scope (RFC 0001 §4d/e): intercept above the LM. A served call - // constructs no client, reaches no provider, and re-executes no tool. - match crate::trace::replay::intercept(self.component_name(), &lm.config, &chat.messages) { - Some(crate::trace::replay::ReplayDirective::Serve(span)) => { - return self.serve_recorded_span(*span, chat, guard); - } - Some(crate::trace::replay::ReplayDirective::Refuse(err)) => { - if let Some(guard) = guard { - guard.finish(crate::trace::SpanOutcome { - events: Vec::new(), - raw_output: None, - output: None, - usage: LmUsage::default(), - error: Some(crate::trace::SpanError { - kind: crate::trace::SpanErrorKind::Lm, - message: err.to_string(), - }), - }); - } - return Err(PredictError::Replay { source: err }); - } - // Live directive (post-divergence) or no replay scope: proceed. - Some(crate::trace::replay::ReplayDirective::Live) | None => {} - } - - let toolset = self.cached_toolset().await; - let empty_toolset = ToolSet::default(); - let toolset_ref = toolset.as_deref().unwrap_or(&empty_toolset); - let response = match lm - .call_with_toolset(chat, toolset_ref, crate::ToolLoopMode::Auto) + let program = self.program().await?; + let overlay = self.effective_overlay(&program)?; + let lm = self.resolve_lm(); + let interpreter = self.interpreter(&program, &lm).await?; + let (run, chat) = interpreter + .run_conversation(chat, None, overlay, Budget::unlimited()) .await - { - Ok(response) => response, - Err(err) => { - if let Some(guard) = guard { - guard.finish(crate::trace::SpanOutcome { - events: Vec::new(), - raw_output: None, - output: None, - usage: LmUsage::default(), - error: Some(crate::trace::SpanError { - kind: crate::trace::SpanErrorKind::Lm, - message: err.to_string(), - }), - }); - } - return Err(PredictError::Lm { - source: LmError::Provider { - provider: lm.config.model.clone(), - message: err.to_string(), - source: None, - }, - }); - } - }; - debug!( - prompt_tokens = response.usage.prompt_tokens, - completion_tokens = response.usage.completion_tokens, - total_tokens = response.usage.total_tokens, - tool_calls = response.tool_calls.len(), - "lm response received" - ); - - let crate::core::lm::LMResponse { - output, - usage, - chat, - tool_calls, - tool_executions, - events, - } = response; - - let chat_adapter = ChatAdapter; - let raw_response = output.content(); - let lm_usage = usage; - - let (typed_output, field_metas) = match chat_adapter.parse_response_typed::(&output) { - Ok(parsed) => parsed, - Err(err) => { - let failed_fields = err.fields(); - debug!( - failed_fields = failed_fields.len(), - fields = ?failed_fields, - raw_response_len = raw_response.len(), - "typed parse failed" - ); - // Parse failures keep raw_output in the span — the model's - // unparseable prose is prime reflection material. - if let Some(guard) = guard { - guard.finish(crate::trace::SpanOutcome { - events, - raw_output: Some(raw_response.clone()), - output: None, - usage: lm_usage, - error: Some(crate::trace::SpanError { - kind: crate::trace::SpanErrorKind::Parse, - message: err.to_string(), - }), - }); - } - return Err(PredictError::Parse { - source: err, - raw_response, - lm_usage, - }); - } - }; - - let span_id = guard.as_ref().map(|guard| guard.id()); - if let Some(guard) = guard { - guard.finish(crate::trace::SpanOutcome { - events, - raw_output: Some(raw_response.clone()), - output: json_map_from_output::(&typed_output).ok(), - usage: lm_usage, - error: None, - }); - } - - let checks_total = field_metas - .values() - .map(|meta| meta.checks.len()) - .sum::(); - let checks_failed = field_metas - .values() - .flat_map(|meta| meta.checks.iter()) - .filter(|check| !check.passed) - .count(); - let flagged_fields = field_metas - .values() - .filter(|meta| !meta.flags.is_empty()) - .count(); - debug!( - output_fields = field_metas.len(), - checks_total, checks_failed, flagged_fields, "typed parse completed" - ); - - let metadata = CallMetadata::new( - raw_response, - lm_usage, - tool_calls, - tool_executions, - span_id, - field_metas, - ); - - Ok((Predicted::new(typed_output, metadata), chat)) - } - - /// Serves one call from a recorded span (replay scope, RFC 0001 §4d): - /// deserializes the recorded parsed output into `S::Output`, extends the - /// chat with the recorded completion, and records the served span into any - /// active capture scope. Zero provider calls, zero tool executions. - #[allow(clippy::result_large_err)] - fn serve_recorded_span( - &self, - span: crate::trace::Span, - mut chat: Chat, - guard: Option, - ) -> Result<(Predicted, Chat), PredictError> - where - S::Output: Schema, - { - let output_map = span - .output - .clone() - .expect("replay serves only spans with parsed output"); - let typed_output: S::Output = match serde_json::from_value(Value::Object(output_map)) { - Ok(output) => output, - Err(err) => { - let source = crate::trace::ReplayError::OutputDecode { - component: self.component_name().to_string(), - seq: span.seq, - span: span.id, - message: err.to_string(), - }; - if let Some(guard) = guard { - guard.finish(crate::trace::SpanOutcome { - events: Vec::new(), - raw_output: span.raw_output.clone(), - output: None, - usage: span.usage, - error: Some(crate::trace::SpanError { - kind: crate::trace::SpanErrorKind::Parse, - message: source.to_string(), - }), - }); - } - return Err(PredictError::Replay { source }); - } - }; - - let raw_response = span.raw_output.clone().unwrap_or_default(); - debug!( - component = self.component_name(), - seq = span.seq, - "predict call served from recorded trace" - ); - - // Rebuild the conversation the live call would have returned: the - // recorded exchanges and tool results, in order. Tool effects are baked - // into the recording — nothing re-executes. - let mut completion = span.completion_messages(); - if completion.is_empty() && !raw_response.is_empty() { - completion.push(Message::assistant(raw_response.clone())); - } - let tool_calls = completion - .iter() - .flat_map(|message| message.tool_calls().into_iter().cloned()) - .collect(); - let tool_executions = span - .events - .iter() - .filter_map(|event| match event { - crate::trace::SpanEvent::ToolRun { result, .. } => Some(result.clone()), - _ => None, - }) - .collect(); - for message in &completion { - chat.push_message(message.clone()); - } - - let span_id = guard.as_ref().map(|guard| guard.id()); - if let Some(guard) = guard { - guard.finish(crate::trace::SpanOutcome { - events: span.events.clone(), - raw_output: span.raw_output.clone(), - output: span.output.clone(), - usage: span.usage, - error: None, - }); - } - - // Served predictions carry no per-field parse metadata: the recording - // stores the parsed output, not the parser's field-level bookkeeping. - let metadata = CallMetadata::new( - raw_response, - span.usage, - tool_calls, - tool_executions, - span_id, - Default::default(), - ); - Ok((Predicted::new(typed_output, metadata), chat)) + .map_err(map_run_error::)?; + let predicted = predicted_from_run::(run)?; + Ok((predicted, chat)) } } @@ -666,6 +640,7 @@ impl Default for Predict { /// ``` pub struct PredictBuilder { tools: Vec>, + agent_spec: Option, demos: Vec>, instruction_override: Option, lm: Option>, @@ -677,6 +652,7 @@ impl PredictBuilder { fn new() -> Self { Self { tools: Vec::new(), + agent_spec: None, demos: Vec::new(), instruction_override: None, lm: None, @@ -715,6 +691,14 @@ impl PredictBuilder { self } + /// Sets loop options ([`AgentLoopSpec`]) for the 1-node `agent` program a + /// tooled predictor executes: stop tools, max turns, until-parse, node + /// budget, and context policy. No effect without tools. + pub fn with_agent_spec(mut self, spec: AgentLoopSpec) -> Self { + self.agent_spec = Some(spec); + self + } + /// Overrides the signature's default instruction for this predictor. pub fn instruction(mut self, instruction: impl Into) -> Self { self.instruction_override = Some(instruction.into()); @@ -739,16 +723,20 @@ impl PredictBuilder { } /// Builds the [`Predict`], routing state through the same applicator the - /// mutation seam uses. + /// install seam ([`PredictorInfo::load_state`](crate::core::PredictorInfo::load_state)) + /// uses. pub fn build(self) -> Predict { let mut predict = Predict { tools: self.tools, + agent_spec: self.agent_spec, demos: Vec::new(), instruction_override: None, lm: self.lm, - prompt_prefix: OnceLock::new(), toolset: tokio::sync::OnceCell::new(), trace_name: self.trace_name, + program: OnceLock::new(), + instance_overlay: OnceLock::new(), + engine: tokio::sync::Mutex::new(None), _marker: PhantomData, }; predict.apply_state(Some(self.instruction_override), Some(self.demos)); @@ -821,6 +809,363 @@ where } } +// --------------------------------------------------------------------------- +// Tool schema projection (JSON Schema → SignatureDef) +// --------------------------------------------------------------------------- + +/// Best-effort projection of a rig tool's JSON-Schema `parameters` object +/// into a tool [`SignatureDef`] (input side) plus the class/enum definitions +/// it references. The interpreter regenerates the model-facing schema from +/// this signature via [`ir::input_schema_of`](crate::ir::input_schema_of); +/// JSON-Schema features outside [`FieldType`](crate::typesys::FieldType) +/// (open `additionalProperties`, `oneOf`, per-property formats, …) degrade to +/// their closest `FieldType` equivalent. +fn tool_signature_from_definition( + tool_name: &str, + parameters: &Value, +) -> (SignatureDef, crate::typesys::TypeTable) { + let mut types = crate::typesys::TypeTable::default(); + let empty = Map::new(); + let properties = parameters + .get("properties") + .and_then(Value::as_object) + .unwrap_or(&empty); + let required: Vec<&str> = parameters + .get("required") + .and_then(Value::as_array) + .map(|entries| entries.iter().filter_map(Value::as_str).collect()) + .unwrap_or_default(); + + let mut inputs = Vec::with_capacity(properties.len()); + for (name, schema) in properties { + let token = format!("{tool_name}_{name}"); + let mut ty = field_type_from_json_schema(&token, schema, &mut types); + if !required.contains(&name.as_str()) { + ty = crate::typesys::FieldType::optional(ty); + } + let mut field = ir::FieldDef::new(name, ty); + if let Some(docs) = schema.get("description").and_then(Value::as_str) { + field = field.with_docs(docs); + } + inputs.push(field); + } + + let def = SignatureDef { + name: format!("{tool_name}_tool").into(), + instruction: "".into(), + inputs: inputs.into_boxed_slice(), + // The output side is unused for host tools (results come back as raw + // strings); a single string field keeps the signature well-formed. + outputs: Box::new([ir::FieldDef::new( + "result", + crate::typesys::FieldType::String, + )]), + }; + (def, types) +} + +fn field_type_from_json_schema( + token: &str, + schema: &Value, + types: &mut crate::typesys::TypeTable, +) -> crate::typesys::FieldType { + use crate::typesys::FieldType; + + if let Some(any_of) = schema.get("anyOf").and_then(Value::as_array) { + let items = any_of + .iter() + .enumerate() + .map(|(i, sub)| field_type_from_json_schema(&format!("{token}_{i}"), sub, types)) + .collect(); + return FieldType::Union(items); + } + if let Some(values) = schema.get("enum").and_then(Value::as_array) { + let names: Vec = values + .iter() + .filter_map(Value::as_str) + .map(str::to_string) + .collect(); + if !names.is_empty() { + types.enums.insert( + token.to_string(), + crate::typesys::EnumDef { + internal_name: token.to_string(), + rendered_name: token.to_string(), + docs: None, + values: names + .into_iter() + .map(|name| crate::typesys::EnumValueDef { + rendered_name: name.clone(), + name, + docs: None, + }) + .collect(), + }, + ); + return FieldType::Enum(token.to_string()); + } + } + match schema.get("type").and_then(Value::as_str) { + Some("string") => FieldType::String, + Some("integer") => FieldType::Int, + Some("number") => FieldType::Float, + Some("boolean") => FieldType::Bool, + Some("array") => FieldType::List(Box::new( + schema + .get("items") + .map(|items| field_type_from_json_schema(&format!("{token}_items"), items, types)) + .unwrap_or(FieldType::String), + )), + Some("object") => { + if let Some(properties) = schema.get("properties").and_then(Value::as_object) { + let required: Vec<&str> = schema + .get("required") + .and_then(Value::as_array) + .map(|entries| entries.iter().filter_map(Value::as_str).collect()) + .unwrap_or_default(); + let fields = properties + .iter() + .map(|(name, sub)| { + let mut ty = + field_type_from_json_schema(&format!("{token}_{name}"), sub, types); + if !required.contains(&name.as_str()) { + ty = crate::typesys::FieldType::optional(ty); + } + crate::typesys::FieldDef { + name: name.clone(), + rendered_name: name.clone(), + field_type: ty, + docs: sub + .get("description") + .and_then(Value::as_str) + .map(str::to_string), + constraints: Vec::new(), + } + }) + .collect(); + types.classes.insert( + token.to_string(), + crate::typesys::ClassDef { + internal_name: token.to_string(), + rendered_name: token.to_string(), + docs: None, + fields, + constraints: Vec::new(), + }, + ); + FieldType::Class(token.to_string()) + } else { + let inner = match schema.get("additionalProperties") { + Some(additional) if additional.is_object() => { + field_type_from_json_schema(&format!("{token}_value"), additional, types) + } + _ => FieldType::String, + }; + FieldType::Map(Box::new(FieldType::String), Box::new(inner)) + } + } + _ => FieldType::String, + } +} + +// --------------------------------------------------------------------------- +// Interpreter boundary: RunError → PredictError, RunOutput → Predicted +// --------------------------------------------------------------------------- + +/// Internal (non-provider, non-parse) failure surfaced through the historical +/// `PredictError::Lm { provider: "internal" }` shape. +fn internal_error(message: String) -> PredictError { + PredictError::Lm { + source: LmError::Provider { + provider: "internal".to_string(), + message, + source: None, + }, + } +} + +/// Deserializes the canonical output map into the typed output struct — +/// failure is the historical `ParseError::ExtractionFailed` shape. +fn typed_output_from_map( + output_map: &Map, + raw_response: &str, +) -> std::result::Result { + serde_json::from_value(Value::Object(output_map.clone())).map_err(|err| { + ParseError::ExtractionFailed { + field: "".to_string(), + raw_response: raw_response.to_string(), + reason: err.to_string(), + } + }) +} + +/// Re-keys def-lane [`FieldMeta`](crate::FieldMeta) entries (canonical +/// `FieldDef::name`s) to the static lane's `rust_name` keying, preserving +/// schema field order — the user-visible `CallMetadata` contract. +fn translate_field_meta( + leaf_meta: &IndexMap, +) -> IndexMap { + let mut field_meta = IndexMap::new(); + for field in S::schema().output_fields() { + let leaf_name = field.path().iter().last().unwrap_or(field.lm_name); + if let Some(meta) = leaf_meta.get(leaf_name) { + field_meta.insert(field.rust_name.clone(), meta.clone()); + } + } + field_meta +} + +/// The canonical (leaf) name of an output field → the static lane's +/// `rust_name` (dotted flatten path). Identity for non-flattened signatures. +fn output_rust_name(canonical: &str) -> String { + for field in S::schema().output_fields() { + let leaf = field.path().iter().last().unwrap_or(field.lm_name); + if leaf == canonical { + return field.rust_name.clone(); + } + } + canonical.to_string() +} + +/// Re-keys a def-lane [`ParseError`] (canonical `FieldDef::name`s) to the +/// static lane's `rust_name` keying, preserving the user-visible error shape +/// for flattened signatures. +fn translate_parse_error(err: ParseError) -> ParseError { + match err { + ParseError::MissingField { + field, + raw_response, + } => ParseError::MissingField { + field: output_rust_name::(&field), + raw_response, + }, + ParseError::ExtractionFailed { + field, + raw_response, + reason, + } => ParseError::ExtractionFailed { + field: output_rust_name::(&field), + raw_response, + reason, + }, + ParseError::CoercionFailed { + field, + expected_type, + raw_text, + source, + } => ParseError::CoercionFailed { + field: output_rust_name::(&field), + expected_type, + raw_text, + source, + }, + ParseError::AssertFailed { + field, + label, + expression, + value, + } => ParseError::AssertFailed { + field: output_rust_name::(&field), + label, + expression, + value, + }, + ParseError::Multiple { errors, partial } => ParseError::Multiple { + errors: errors.into_iter().map(translate_parse_error::).collect(), + partial, + }, + } +} + +/// Maps an interpreter [`RunError`] onto the historical [`PredictError`] +/// variants: `Lm`/`Parse`/`Replay` map structurally; everything else (input +/// validation, internal invariants) surfaces as an internal LM error. +fn map_run_error(err: RunError) -> PredictError { + match err { + RunError::Lm { source, .. } => PredictError::Lm { source }, + RunError::Parse { + raw, source, usage, .. + } => PredictError::Parse { + source: match source { + Some(parse) => translate_parse_error::(*parse), + None => ParseError::ExtractionFailed { + field: "".to_string(), + raw_response: raw.clone(), + reason: "response did not match the output signature".to_string(), + }, + }, + raw_response: raw, + lm_usage: usage, + }, + RunError::Replay { source, .. } => PredictError::Replay { source }, + other => internal_error(other.to_string()), + } +} + +/// Reassembles the typed [`Predicted`] from a 1-node run: typed deserialize of +/// the output map + [`CallMetadata`] from the single [`LeafOutcome`](crate::ir::LeafOutcome), +/// with `FieldMeta` re-keyed from canonical names to the static lane's +/// `rust_name` keying (schema field order preserved). +#[allow(clippy::result_large_err)] +fn predicted_from_run(run: RunOutput) -> Result, PredictError> +where + S::Output: Schema, +{ + let leaf = run.leaves.into_iter().next(); + let (raw_response, lm_usage, span_id, leaf_meta, tool_calls, tool_executions) = match leaf { + Some(leaf) => ( + leaf.raw_response, + leaf.usage, + leaf.span_id, + leaf.field_meta, + leaf.tool_calls, + leaf.tool_executions, + ), + None => ( + String::new(), + LmUsage::default(), + None, + IndexMap::new(), + Vec::new(), + Vec::new(), + ), + }; + + let typed: S::Output = + typed_output_from_map::(&run.output, &raw_response).map_err(|err| { + PredictError::Parse { + source: err, + raw_response: raw_response.clone(), + lm_usage, + } + })?; + let field_meta = translate_field_meta::(&leaf_meta); + + let checks_total = field_meta + .values() + .map(|meta| meta.checks.len()) + .sum::(); + let checks_failed = field_meta + .values() + .flat_map(|meta| meta.checks.iter()) + .filter(|check| !check.passed) + .count(); + debug!( + output_fields = field_meta.len(), + checks_total, checks_failed, "typed parse completed" + ); + + let metadata = CallMetadata::new( + raw_response, + lm_usage, + tool_calls, + tool_executions, + span_id, + field_meta, + ); + Ok(Predicted::new(typed, metadata)) +} + impl Module for Predict where S: Signature + Clone, @@ -844,13 +1189,30 @@ where } } -impl DynPredictor for Predict +impl Predict { + /// Assigns the component name recorded on this predictor's trace spans. + /// + /// The leaf name is part of the 1-node program (and its hash), so the + /// cached program, the overlay minted against it, and the loaded + /// interpreter are all invalidated when the name changes. + pub(crate) fn assign_trace_name(&mut self, name: &str) { + if self.trace_name.as_deref() == Some(name) { + return; + } + self.trace_name = Some(name.to_string()); + self.program = OnceLock::new(); + self.instance_overlay = OnceLock::new(); + *self.engine.get_mut() = None; + } +} + +impl crate::core::PredictorInfo for Predict where S: Signature, S::Input: Schema, S::Output: Schema, { - fn schema(&self) -> &SignatureSchema { + fn schema(&self) -> &'static SignatureSchema { S::schema() } @@ -860,41 +1222,65 @@ where .unwrap_or_else(|| S::instruction().to_string()) } + fn default_instruction(&self) -> String { + S::instruction().to_string() + } + fn demos_as_json(&self) -> Vec> { self.demos .iter() .map(|example| { - json_from_demo::(example) - .expect("typed Predict demo conversion should succeed") + json_from_demo::(example).expect("typed Predict demo conversion should succeed") }) .collect() } fn dump_state(&self) -> PredictState { PredictState { - demos: self.demos_as_json(), + demos: crate::core::PredictorInfo::demos_as_json(self), instruction_override: self.instruction_override.clone(), } } - fn apply_update(&mut self, update: StateUpdate) -> Result<()> { + fn load_state(&mut self, state: PredictState) -> Result<()> { // Convert demos before touching any state so a schema mismatch leaves // the predictor unchanged. - let demos = update + let demos = state .demos - .map(|demos| { - demos - .iter() - .map(demo_from_json::) - .collect::>>() - }) - .transpose()?; - self.apply_state(update.instruction, demos); + .iter() + .map(demo_from_json::) + .collect::>>()?; + self.apply_state(Some(state.instruction_override), Some(demos)); Ok(()) } fn set_trace_name(&mut self, name: &str) { - self.trace_name = Some(name.to_string()); + Predict::assign_trace_name(self, name); + } +} + +impl crate::core::Predictors for Predict +where + S: Signature, + S::Input: Schema, + S::Output: Schema, +{ + /// A bare `Predict` used as a module is itself the one leaf. Its name is + /// the assigned trace name when present, else `"self"`. + fn predictors(&self) -> Vec<(String, &dyn crate::core::PredictorInfo)> { + let name = self + .trace_name + .clone() + .unwrap_or_else(|| "self".to_string()); + vec![(name, self as &dyn crate::core::PredictorInfo)] + } + + fn predictors_mut(&mut self) -> Vec<(String, &mut dyn crate::core::PredictorInfo)> { + let name = self + .trace_name + .clone() + .unwrap_or_else(|| "self".to_string()); + vec![(name, self as &mut dyn crate::core::PredictorInfo)] } } @@ -953,20 +1339,24 @@ mod tests { } #[test] - fn dyn_predictor_apply_update_round_trips_json_demo_rows() { + fn predictor_info_load_state_round_trips_json_demo_rows() { + use crate::core::PredictorInfo; + let typed = typed_row("demo-input", "demo-output"); let row = json_from_demo::(&typed) .expect("typed demo should convert to a flat row"); let mut predictor = Predict::::new(); - let update = StateUpdate { - instruction: None, - demos: Some(vec![row]), - }; - DynPredictor::apply_update(&mut predictor, update) - .expect("predictor should accept JSON demo rows"); + PredictorInfo::load_state( + &mut predictor, + PredictState { + demos: vec![row], + instruction_override: None, + }, + ) + .expect("predictor should accept JSON demo rows"); - let demos = DynPredictor::demos_as_json(&predictor); + let demos = PredictorInfo::demos_as_json(&predictor); assert_eq!(demos.len(), 1); assert_eq!(demos[0].get("prompt"), Some(&json!("demo-input"))); assert_eq!(demos[0].get("answer"), Some(&json!("demo-output"))); diff --git a/crates/dspy-rs/src/trace/capture.rs b/crates/dspy-rs/src/trace/capture.rs index 988bf5b5..349394e4 100644 --- a/crates/dspy-rs/src/trace/capture.rs +++ b/crates/dspy-rs/src/trace/capture.rs @@ -128,19 +128,12 @@ impl TraceSink { let id = SpanId(inner.trace.spans.len() as u32); let parent = inner.open.last().copied(); - let links = inner - .trace - .spans - .last() - .map(|span| vec![span.id]) - .unwrap_or_default(); inner.trace.spans.push(Span { id, component, seq, parent, - links, prefix, suffix: req.suffix.to_vec(), input: req.input, @@ -151,6 +144,7 @@ impl TraceSink { output: None, usage: LmUsage::default(), error: None, + eval: None, started_at_us: now_us(), duration_us: 0, complete: true, diff --git a/crates/dspy-rs/src/trace/export/mod.rs b/crates/dspy-rs/src/trace/export/mod.rs deleted file mode 100644 index 43f4042e..00000000 --- a/crates/dspy-rs/src/trace/export/mod.rs +++ /dev/null @@ -1,11 +0,0 @@ -//! Exports: projections of the trace format onto external training and -//! observability conventions (RFC 0001 §4f/§4g). -//! -//! Everything here is a pure serialization-side projection — no new capture -//! machinery, no external dependencies. - -pub mod otel; -pub mod rl; - -pub use otel::{OtelEvent, OtelKeyValue, OtelSpan, OtelStatus, OtelValue}; -pub use rl::{RlRollout, RlTransition}; diff --git a/crates/dspy-rs/src/trace/export/otel.rs b/crates/dspy-rs/src/trace/export/otel.rs deleted file mode 100644 index ac56d6a0..00000000 --- a/crates/dspy-rs/src/trace/export/otel.rs +++ /dev/null @@ -1,368 +0,0 @@ -//! OTel export (RFC 0001 §4g): one-way batch mapping of a finished [`Trace`] -//! onto OpenTelemetry GenAI semantic conventions — as plain serializable -//! structs in the OTLP/JSON wire shape, with **no OpenTelemetry dependency**. -//! -//! Mapping (RFC §4g): -//! -//! | trace format | OTel | -//! |-------------------------------------|-------------------------------------------------| -//! | `meta.trace_id` | trace id (used verbatim if 32-hex, else hashed) | -//! | the rollout | root span `dsrs.rollout`, kind `INTERNAL` | -//! | `Span` | span, name = component name, kind = `CLIENT` | -//! | `span.parent` (else the root span) | parent span id | -//! | `started_at_us` / `duration_us` | start/end timestamps (ns) | -//! | `models[span.model].config.model` | `gen_ai.request.model` | -//! | `temperature` / `max_tokens` | `gen_ai.request.temperature` / `.max_tokens` | -//! | `usage.prompt/completion_tokens` | `gen_ai.usage.input_tokens` / `.output_tokens` | -//! | `prefix ++ suffix` (content opt-in) | `gen_ai.prompt` events | -//! | `raw_output` (content opt-in) | `gen_ai.completion` event | -//! | `SpanEvent::ToolRun` | child span `tool:{name}`, `gen_ai.tool.*` | -//! | `span.error` | status `ERROR` + message | -//! | `seq` / `component` / hashes | `dsrs.*` attributes | -//! -//! Prompt/completion/tool content is **opt-in** (`include_content`), mirroring -//! OTel's GenAI content-capture switch — spans stay exportable to shared -//! collectors without leaking prompt text. -//! -//! # Shipping to a collector -//! -//! [`Trace::to_otlp_json`] wraps the spans in a complete -//! `resourceSpans` envelope, directly acceptable to any OTLP/HTTP collector -//! (Jaeger, Tempo, otel-collector) — no SDK required: -//! -//! ```ignore -//! let payload = trace.to_otlp_json("my-service", /* include_content */ false); -//! reqwest::Client::new() -//! .post("http://localhost:4318/v1/traces") -//! .json(&payload) -//! .send() -//! .await?; -//! ``` -//! -//! Timing note: `ToolRun` events record only their duration, so tool child -//! spans start at their parent span's start time — durations are exact, -//! offsets within the parent are not. - -use serde::Serialize; - -use crate::trace::span::{Span, SpanEvent, Trace}; -use crate::utils::hash::stable_hash_debug; - -/// OTLP `SPAN_KIND_INTERNAL`. -pub const SPAN_KIND_INTERNAL: u32 = 1; -/// OTLP `SPAN_KIND_CLIENT`. -pub const SPAN_KIND_CLIENT: u32 = 3; -/// OTLP `STATUS_CODE_ERROR`. -pub const STATUS_CODE_ERROR: u32 = 2; - -/// One span in the OTLP/JSON wire shape (proto3 JSON mapping: camelCase keys, -/// 64-bit integers as decimal strings, ids as lowercase hex). -#[derive(Clone, Debug, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct OtelSpan { - /// 128-bit trace id, 32 lowercase hex chars. - pub trace_id: String, - /// 64-bit span id, 16 lowercase hex chars. - pub span_id: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub parent_span_id: Option, - pub name: String, - pub kind: u32, - pub start_time_unix_nano: String, - pub end_time_unix_nano: String, - pub attributes: Vec, - #[serde(skip_serializing_if = "Vec::is_empty")] - pub events: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - pub status: Option, -} - -/// OTLP `KeyValue`. -#[derive(Clone, Debug, Serialize)] -pub struct OtelKeyValue { - pub key: String, - pub value: OtelValue, -} - -impl OtelKeyValue { - fn str(key: &str, value: impl Into) -> Self { - Self { - key: key.to_string(), - value: OtelValue::StringValue(value.into()), - } - } - - fn int(key: &str, value: i64) -> Self { - Self { - key: key.to_string(), - // proto3 JSON maps int64 to a decimal string. - value: OtelValue::IntValue(value.to_string()), - } - } - - fn double(key: &str, value: f64) -> Self { - Self { - key: key.to_string(), - value: OtelValue::DoubleValue(value), - } - } -} - -/// OTLP `AnyValue` (the oneof arms this export emits). -#[derive(Clone, Debug, Serialize)] -#[serde(rename_all = "camelCase")] -pub enum OtelValue { - StringValue(String), - IntValue(String), - DoubleValue(f64), -} - -/// OTLP span `Event`. -#[derive(Clone, Debug, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct OtelEvent { - pub time_unix_nano: String, - pub name: String, - pub attributes: Vec, -} - -/// OTLP span `Status`. -#[derive(Clone, Debug, Serialize)] -#[serde(rename_all = "camelCase")] -pub struct OtelStatus { - pub code: u32, - #[serde(skip_serializing_if = "String::is_empty")] - pub message: String, -} - -/// Deterministic span-id namespaces: predict spans, the root span, and tool -/// child spans never collide and stay greppable in collector UIs. -const PREDICT_SPAN_NS: u64 = 0x0100_0000_0000_0000; -const ROOT_SPAN_NS: u64 = 0x0200_0000_0000_0000; -const TOOL_SPAN_NS: u64 = 0x0300_0000_0000_0000; - -fn hex_span_id(id: u64) -> String { - format!("{id:016x}") -} - -fn ns(us: u64) -> String { - (us.saturating_mul(1_000)).to_string() -} - -/// 32-hex trace id: the recorded id verbatim when it already is one (the -/// capture scope mints exactly this shape), otherwise a stable hash of it. -fn otel_trace_id(trace_id: &str) -> String { - let is_32_hex = - trace_id.len() == 32 && trace_id.bytes().all(|b| b.is_ascii_hexdigit() && !b.is_ascii_uppercase()); - if is_32_hex { - trace_id.to_string() - } else { - format!( - "{:016x}{:016x}", - stable_hash_debug(&("dsrs.otel.hi", trace_id)), - stable_hash_debug(&("dsrs.otel.lo", trace_id)), - ) - } -} - -impl Trace { - /// Maps this trace onto OTel-shaped spans: one root span for the rollout, - /// one `CLIENT` span per `Predict` invocation (parented to its recorded - /// parent span, else the root), and one child span per `ToolRun`. - /// - /// `include_content` gates prompt/completion/tool-payload capture (the - /// GenAI content-capture switch); identity, usage, and timing attributes - /// are always emitted. - pub fn to_otel_spans(&self, include_content: bool) -> Vec { - let trace_id = otel_trace_id(&self.meta.trace_id); - let root_span_id = hex_span_id(ROOT_SPAN_NS); - - let root_end_us = match &self.outcome { - Some(outcome) => self.meta.started_at_us + outcome.duration_us, - None => self - .spans - .iter() - .map(|span| span.started_at_us + span.duration_us) - .max() - .unwrap_or(self.meta.started_at_us), - }; - let mut root_attributes = vec![OtelKeyValue::str("dsrs.trace_id", &self.meta.trace_id)]; - if let Some(candidate_hash) = self.meta.candidate_hash { - root_attributes.push(OtelKeyValue::str( - "dsrs.candidate_hash", - format!("{candidate_hash:016x}"), - )); - } - for (key, value) in &self.meta.tags { - root_attributes.push(OtelKeyValue::str(&format!("dsrs.tag.{key}"), value)); - } - let root_status = self - .outcome - .as_ref() - .and_then(|outcome| outcome.error.as_ref()) - .map(|error| OtelStatus { - code: STATUS_CODE_ERROR, - message: error.clone(), - }); - - let mut spans = vec![OtelSpan { - trace_id: trace_id.clone(), - span_id: root_span_id.clone(), - parent_span_id: None, - name: "dsrs.rollout".to_string(), - kind: SPAN_KIND_INTERNAL, - start_time_unix_nano: ns(self.meta.started_at_us), - end_time_unix_nano: ns(root_end_us), - attributes: root_attributes, - events: Vec::new(), - status: root_status, - }]; - - for span in &self.spans { - spans.push(self.otel_predict_span(span, &trace_id, &root_span_id, include_content)); - spans.extend(self.otel_tool_spans(span, &trace_id, include_content)); - } - spans - } - - fn otel_predict_span( - &self, - span: &Span, - trace_id: &str, - root_span_id: &str, - include_content: bool, - ) -> OtelSpan { - let config = self.model(span); - let mut attributes = vec![ - OtelKeyValue::str("gen_ai.request.model", &config.model), - OtelKeyValue::double("gen_ai.request.temperature", config.temperature as f64), - OtelKeyValue::int("gen_ai.request.max_tokens", config.max_tokens as i64), - OtelKeyValue::int("gen_ai.usage.input_tokens", span.usage.prompt_tokens as i64), - OtelKeyValue::int( - "gen_ai.usage.output_tokens", - span.usage.completion_tokens as i64, - ), - OtelKeyValue::str("dsrs.component", self.component_name(span.component)), - OtelKeyValue::int("dsrs.seq", span.seq as i64), - OtelKeyValue::str("dsrs.request_hash", format!("{:016x}", span.request_hash)), - ]; - if let Some(candidate_hash) = self.meta.candidate_hash { - attributes.push(OtelKeyValue::str( - "dsrs.candidate_hash", - format!("{candidate_hash:016x}"), - )); - } - - let mut events = Vec::new(); - if include_content { - for message in self.prompt(span) { - events.push(OtelEvent { - time_unix_nano: ns(span.started_at_us), - name: "gen_ai.prompt".to_string(), - attributes: vec![ - OtelKeyValue::str("gen_ai.prompt.role", message.role.as_str()), - OtelKeyValue::str("gen_ai.prompt.content", message.content()), - ], - }); - } - if let Some(raw_output) = &span.raw_output { - events.push(OtelEvent { - time_unix_nano: ns(span.started_at_us + span.duration_us), - name: "gen_ai.completion".to_string(), - attributes: vec![OtelKeyValue::str("gen_ai.completion.content", raw_output)], - }); - } - } - - OtelSpan { - trace_id: trace_id.to_string(), - span_id: hex_span_id(PREDICT_SPAN_NS | span.id.0 as u64), - parent_span_id: Some(match span.parent { - Some(parent) => hex_span_id(PREDICT_SPAN_NS | parent.0 as u64), - None => root_span_id.to_string(), - }), - name: self.component_name(span.component).to_string(), - kind: SPAN_KIND_CLIENT, - start_time_unix_nano: ns(span.started_at_us), - end_time_unix_nano: ns(span.started_at_us + span.duration_us), - attributes, - events, - status: span.error.as_ref().map(|error| OtelStatus { - code: STATUS_CODE_ERROR, - message: format!("{}: {}", error.kind.as_str(), error.message), - }), - } - } - - fn otel_tool_spans( - &self, - span: &Span, - trace_id: &str, - include_content: bool, - ) -> Vec { - span.events - .iter() - .enumerate() - .filter_map(|(index, event)| match event { - SpanEvent::ToolRun { - id, - name, - args, - result, - duration_us, - error, - } => { - let mut attributes = vec![ - OtelKeyValue::str("gen_ai.tool.name", name), - OtelKeyValue::str("gen_ai.tool.call.id", id), - ]; - if include_content { - attributes.push(OtelKeyValue::str( - "gen_ai.tool.call.arguments", - serde_json::to_string(args).unwrap_or_default(), - )); - attributes.push(OtelKeyValue::str("gen_ai.tool.call.result", result)); - } - Some(OtelSpan { - trace_id: trace_id.to_string(), - span_id: hex_span_id( - TOOL_SPAN_NS | ((span.id.0 as u64) << 16) | index as u64, - ), - parent_span_id: Some(hex_span_id(PREDICT_SPAN_NS | span.id.0 as u64)), - name: format!("tool:{name}"), - kind: SPAN_KIND_INTERNAL, - // ToolRuns record duration only; anchor at parent start. - start_time_unix_nano: ns(span.started_at_us), - end_time_unix_nano: ns(span.started_at_us + duration_us), - attributes, - events: Vec::new(), - status: error.as_ref().map(|message| OtelStatus { - code: STATUS_CODE_ERROR, - message: message.clone(), - }), - }) - } - _ => None, - }) - .collect() - } - - /// Complete OTLP/HTTP JSON payload (`resourceSpans` envelope) for this - /// trace — POST it to a collector's `/v1/traces` endpoint as-is. See the - /// module docs for an example. - pub fn to_otlp_json(&self, service_name: &str, include_content: bool) -> serde_json::Value { - serde_json::json!({ - "resourceSpans": [{ - "resource": { - "attributes": [ - { "key": "service.name", "value": { "stringValue": service_name } } - ] - }, - "scopeSpans": [{ - "scope": { "name": "dsrs.trace", "version": env!("CARGO_PKG_VERSION") }, - "spans": self.to_otel_spans(include_content) - }] - }] - }) - } -} diff --git a/crates/dspy-rs/src/trace/export/rl.rs b/crates/dspy-rs/src/trace/export/rl.rs deleted file mode 100644 index 0edb5249..00000000 --- a/crates/dspy-rs/src/trace/export/rl.rs +++ /dev/null @@ -1,103 +0,0 @@ -//! RL rollout export (RFC 0001 §4f): the Agent Lightning / verifiers span -//! convention — one rollout as message lists + reward + per-subcall -//! transitions. -//! -//! A pure projection of the trace: because spans keep full [`Message`] -//! structure (tool-call blocks, reasoning blocks), the export needs no lossy -//! text munging. Each transition's `messages` is the exact rendered prompt -//! ([`Trace::prompt`]: prefix ++ suffix) and `completion` is everything the -//! policy emitted, rebuilt from the span's events -//! ([`Span::completion_messages`]): each `Exchange`'s assistant message -//! verbatim, `ToolRun`s as intervening tool-result context. -//! -//! ```ignore -//! let (result, mut trace) = capture(|| pipeline(input)).await; -//! trace.outcome = Some(TraceOutcome { eval: Some(metric_eval), ..Default::default() }); -//! let rollout = trace.to_rl_rollout().expect("eval recorded"); -//! writeln!(dataset, "{}", rollout.to_json_line()?)?; // one JSONL record per rollout -//! ``` - -use std::collections::BTreeMap; - -use anyhow::Result; -use serde::Serialize; - -use crate::trace::span::{Span, Trace}; -use crate::{LmUsage, Message}; - -/// One rollout: reward plus per-subcall transitions. Serializes to a single -/// JSON object — the JSONL record RL trainers consume. -#[derive(Debug, Serialize)] -pub struct RlRollout<'a> { - pub trace_id: &'a str, - /// The rollout-level reward: `trace.outcome.eval.score`. - pub reward: f64, - pub transitions: Vec>, - /// Free-form run tags (`trace.meta.tags`). - pub metadata: &'a BTreeMap, -} - -/// One policy subcall — a `Predict` invocation as (prompt messages, emitted -/// completion) with its span metadata. -#[derive(Debug, Serialize)] -pub struct RlTransition<'a> { - /// The optimizable unit this subcall belongs to (`"drafter"`) — the - /// per-agent credit assignment key. - pub component: &'a str, - /// 0-based invocation index of the component within the rollout. - pub seq: u32, - /// Full prompt as messages (prefix ++ suffix), provider-agnostic roles. - pub messages: Vec, - /// Everything the policy emitted for this subcall: each `Exchange`'s - /// assistant message, with `ToolRun`s as intervening tool-result context. - pub completion: Vec, - pub usage: LmUsage, - /// Model identifier from the span's interned config. - pub model: &'a str, -} - -impl RlRollout<'_> { - /// Serializes to one JSONL record. - pub fn to_json_line(&self) -> Result { - Ok(serde_json::to_string(self)?) - } -} - -impl Trace { - /// Projects this rollout onto the RL export convention. `None` when no - /// eval was recorded — a rollout without a reward is not trainable. - /// - /// Spans whose policy emitted nothing (provider failures, cancelled spans - /// — no `Exchange` event) are omitted: they contribute no completion to - /// train on. Parse-failure spans keep their transition: the emitted text - /// exists even though it did not parse. - pub fn to_rl_rollout(&self) -> Option> { - let reward = self.outcome.as_ref()?.eval.as_ref()?.score; - let transitions = self - .spans - .iter() - .filter_map(|span| self.rl_transition(span)) - .collect(); - Some(RlRollout { - trace_id: &self.meta.trace_id, - reward, - transitions, - metadata: &self.meta.tags, - }) - } - - fn rl_transition<'a>(&'a self, span: &'a Span) -> Option> { - let completion = span.completion_messages(); - if completion.is_empty() { - return None; - } - Some(RlTransition { - component: self.component_name(span.component), - seq: span.seq, - messages: self.prompt(span), - completion, - usage: span.usage, - model: &self.models[span.model.0 as usize].config.model, - }) - } -} diff --git a/crates/dspy-rs/src/trace/mod.rs b/crates/dspy-rs/src/trace/mod.rs index d8b7b64f..1e3ab258 100644 --- a/crates/dspy-rs/src/trace/mod.rs +++ b/crates/dspy-rs/src/trace/mod.rs @@ -26,13 +26,11 @@ //! until divergence (counterfactual replay of mutated candidates). pub mod capture; -pub mod export; pub mod replay; pub mod serialize; pub mod span; pub use capture::*; -pub use export::*; pub use replay::{ReplayError, ReplayMode, ReplayReport, is_replaying, replay}; pub use serialize::TRACE_FORMAT_VERSION; pub use span::*; diff --git a/crates/dspy-rs/src/trace/serialize.rs b/crates/dspy-rs/src/trace/serialize.rs index 131380ea..be4c6658 100644 --- a/crates/dspy-rs/src/trace/serialize.rs +++ b/crates/dspy-rs/src/trace/serialize.rs @@ -31,7 +31,6 @@ struct HeaderRef<'a> { h: &'a TraceMeta, components: &'a [String], /// RFC 0001 §1's reserved join column — additive, omitted when empty. - #[cfg(feature = "ir")] #[serde(skip_serializing_if = "<[_]>::is_empty")] param_ids: &'a [Option>], models: &'a [ModelEntry], @@ -43,7 +42,6 @@ struct HeaderOwned { h: TraceMeta, #[serde(default)] components: Vec, - #[cfg(feature = "ir")] #[serde(default)] param_ids: Vec>>, #[serde(default)] @@ -65,7 +63,6 @@ impl Trace { out.push_str(&serde_json::to_string(&HeaderRef { h: &self.meta, components: &self.components, - #[cfg(feature = "ir")] param_ids: &self.param_ids, models: &self.models, prefixes: &self.prefixes, @@ -105,7 +102,6 @@ impl Trace { let mut trace = Trace { meta: header.h, components: header.components, - #[cfg(feature = "ir")] param_ids: header.param_ids, models: header.models, prefixes: header.prefixes, @@ -216,7 +212,6 @@ mod tests { component: CompId(0), seq: 0, parent: None, - links: Vec::new(), prefix: None, suffix: vec![Message::user(suffix_text)], input: None, @@ -230,6 +225,7 @@ mod tests { output: None, usage: LmUsage::default(), error: None, + eval: None, started_at_us: 1, duration_us: 2, complete: true, @@ -275,6 +271,28 @@ mod tests { assert!(err.to_string().contains("newer than supported")); } + #[test] + fn span_eval_roundtrips_and_stays_off_the_wire_when_absent() { + use crate::trace::span::Eval; + + // Absent span eval: the span line carries no `eval` key at all, so + // pre-RFC-0004 traces and eval-free traces serialize byte-identically. + let trace = test_trace(test_span("hi")); + let jsonl = trace.to_jsonl().expect("serialize"); + let span_line = jsonl.lines().nth(1).expect("span line"); + assert!(!span_line.contains("\"eval\"")); + + // Present span eval: survives a write/read cycle intact. + let mut span = test_span("hi"); + span.eval = Some(Eval::with_feedback(0.25, "partial credit")); + let trace = test_trace(span); + let parsed = Trace::from_jsonl(&trace.to_jsonl().expect("serialize")).expect("deserialize"); + assert_eq!( + parsed.spans[0].eval, + Some(Eval::with_feedback(0.25, "partial credit")) + ); + } + #[test] fn unknown_event_tags_are_skipped_on_read() { let trace = test_trace(test_span("hi")); diff --git a/crates/dspy-rs/src/trace/span.rs b/crates/dspy-rs/src/trace/span.rs index cc954bd0..20819f9f 100644 --- a/crates/dspy-rs/src/trace/span.rs +++ b/crates/dspy-rs/src/trace/span.rs @@ -53,10 +53,6 @@ pub struct Span { /// execution (Predict-in-tool). Best-effort: set from the innermost open /// span at begin time; `None` at top level. pub parent: Option, - /// Dataflow predecessors. v1 records the sequential approximation the old - /// Graph recorded: `[previous span in this scope]`, empty for the first. - /// Exact for hand-written sequential pipelines; an approximation elsewhere. - pub links: Vec, // ---- request (recorded eagerly at span open) ---- /// Interned system+demos prefix. `None` when the call had no prefix @@ -85,6 +81,13 @@ pub struct Span { /// Aggregated across all exchanges. pub usage: LmUsage, pub error: Option, + /// Span-level metric result, when the metric attached one after the + /// rollout ran (see `TypedMetric::evaluate_spans`). `None` means no + /// per-span credit was assigned — demo harvesting then falls back to the + /// whole-rollout score in [`TraceOutcome::eval`]. Additive under §5.1's + /// rules — no format version bump; absent on the wire when `None`. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub eval: Option, // ---- timing ---- /// Microseconds since UNIX epoch. @@ -122,9 +125,6 @@ pub enum SpanEvent { /// Tool-level failure that was reported back to the model as text. error: Option, }, - /// RESERVED for streaming. Not emitted in v1. A streamed predict will - /// interleave these before the closing `Exchange`. - Chunk { text: String }, /// Unknown tag from a newer writer; preserved as a placeholder on read, /// dropped from the canonical JSONL on re-serialize. #[doc(hidden)] @@ -226,7 +226,6 @@ pub struct Trace { /// [`attach_program`](Trace::attach_program) fills it; a `None` entry is /// a component the program has no leaf for (static-lane harnesses leave /// the whole column empty). Additive — no format version bump. - #[cfg(feature = "ir")] #[serde(default, skip_serializing_if = "Vec::is_empty")] pub param_ids: Vec>>, pub models: Vec, @@ -313,11 +312,6 @@ impl Trace { .filter(move |span| Some(span.component) == id) } - /// Id-form overload of [`for_component`](Trace::for_component) for hot paths. - pub fn for_component_id(&self, id: CompId) -> impl Iterator + '_ { - self.spans.iter().filter(move |span| span.component == id) - } - /// Reconstructs the exact rendered prompt of a span (prefix ++ suffix). pub fn prompt(&self, span: &Span) -> Vec { let prefix = span @@ -350,7 +344,6 @@ impl Trace { /// One addressing story: a span's `component` string == the leaf name == /// the `ParamPath` prefix, so after attaching, spans join to optimizable /// slots without string surgery. - #[cfg(feature = "ir")] pub fn attach_program(&mut self, program: &crate::ir::Program) { let mut by_leaf: std::collections::HashMap<&str, Vec> = std::collections::HashMap::new(); @@ -368,110 +361,6 @@ impl Trace { .collect(); } - /// Merges another trace's spans into this one, remapping intern tables and - /// span ids. Spans are ordered by `started_at_us` after the merge — provided - /// for harnesses that fan out with `tokio::spawn` and capture per subtask. - pub fn absorb(&mut self, other: Trace) { - let comp_map: Vec = other - .components - .iter() - .map(|name| match self.component_id(name) { - Some(id) => id, - None => { - self.components.push(name.clone()); - CompId((self.components.len() - 1) as u32) - } - }) - .collect(); - // Keep the param_ids column parallel to the merged components, - // preferring already-attached entries on either side. - #[cfg(feature = "ir")] - if !self.param_ids.is_empty() || !other.param_ids.is_empty() { - self.param_ids.resize(self.components.len(), None); - for (idx, ids) in other.param_ids.iter().enumerate() { - let mapped = comp_map[idx].0 as usize; - if self.param_ids[mapped].is_none() { - self.param_ids[mapped] = ids.clone(); - } - } - } - - let model_map: Vec = other - .models - .iter() - .map(|entry| { - match self - .models - .iter() - .position(|existing| existing.config_hash == entry.config_hash) - { - Some(idx) => ModelId(idx as u32), - None => { - self.models.push(entry.clone()); - ModelId((self.models.len() - 1) as u32) - } - } - }) - .collect(); - let prefix_base = self.prefixes.len() as u32; - self.prefixes.extend(other.prefixes); - - // Seq counters continue per component name across the merge. - let mut seqs: Vec = self - .components - .iter() - .enumerate() - .map(|(idx, _)| { - self.spans - .iter() - .filter(|span| span.component.0 as usize == idx) - .count() as u32 - }) - .collect(); - - let id_base = self.spans.len() as u32; - for mut span in other.spans { - span.id = SpanId(span.id.0 + id_base); - span.parent = span.parent.map(|id| SpanId(id.0 + id_base)); - span.links = span - .links - .into_iter() - .map(|id| SpanId(id.0 + id_base)) - .collect(); - span.component = comp_map[span.component.0 as usize]; - span.model = model_map[span.model.0 as usize]; - span.prefix = span.prefix.map(|id| PrefixId(id.0 + prefix_base)); - let seq = &mut seqs[span.component.0 as usize]; - span.seq = *seq; - *seq += 1; - self.spans.push(span); - } - self.spans.sort_by_key(|span| span.started_at_us); - // Ids stay unique but no longer dense after the sort; re-densify. - let remap: std::collections::HashMap = self - .spans - .iter() - .enumerate() - .map(|(idx, span)| (span.id.0, idx as u32)) - .collect(); - for span in &mut self.spans { - span.id = SpanId(remap[&span.id.0]); - span.parent = span.parent.map(|id| SpanId(remap[&id.0])); - for link in &mut span.links { - *link = SpanId(remap[&link.0]); - } - } - } - - /// Structural redaction hook: scrub span content before persisting. - /// Redaction invalidates `request_hash` replay, so touched spans are - /// marked `complete = false`. - pub fn redact(&mut self, mut f: impl FnMut(&mut Span)) { - for span in &mut self.spans { - f(span); - span.complete = false; - } - } } impl Span { diff --git a/crates/dspy-rs/src/typesys/constraint.rs b/crates/dspy-rs/src/typesys/constraint.rs index 1578624e..9fc1eed2 100644 --- a/crates/dspy-rs/src/typesys/constraint.rs +++ b/crates/dspy-rs/src/typesys/constraint.rs @@ -41,72 +41,11 @@ pub struct Constraint { pub expression: String, } -impl Constraint { - pub fn new_check(label: impl Into, expression: impl Into) -> Self { - Self { - level: ConstraintKind::Check, - label: Some(label.into()), - expression: expression.into(), - } - } - - pub fn new_assert(label: impl Into, expression: impl Into) -> Self { - Self { - level: ConstraintKind::Assert, - label: Some(label.into()), - expression: expression.into(), - } - } -} - -/// The outcome of evaluating a constraint against a value. -#[derive(Debug, Clone)] -pub struct ConstraintOutcome { - pub level: ConstraintKind, - pub label: String, - pub expression: String, - pub passed: bool, -} - -/// A reported check result, mirroring the old `ResponseCheck` shape used by GEPA/optimizers. -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct ResponseCheck { - pub name: String, - pub expression: String, - pub status: String, -} - -/// Evaluates every constraint against `value`, binding it as `this` in a jinja expression. -/// -/// A constraint that fails to evaluate (bad expression, wrong type) is treated as not -/// passing rather than erroring, matching the tolerant spirit of the old pipeline. -pub fn evaluate_constraints(value: &Value, constraints: &[Constraint]) -> Vec { - constraints - .iter() - .map(|constraint| { - let passed = eval_expression(&constraint.expression, value).unwrap_or(false); - let label = constraint - .label - .clone() - .unwrap_or_else(|| match constraint.level { - ConstraintKind::Assert => "assert".to_string(), - ConstraintKind::Check => "check".to_string(), - }); - ConstraintOutcome { - level: constraint.level, - label, - expression: constraint.expression.clone(), - passed, - } - }) - .collect() -} - /// Evaluates a runtime (non-`'static`) constraint expression against `value`, /// binding it as `this`. Compiles per call — dynamic-lane constraints are owned /// strings, and caching them process-wide would reintroduce the leak-per-load -/// that RFC 0002 IR-1 removed. Failed evaluations return `false`, matching -/// [`evaluate_constraints`]. +/// that RFC 0002 IR-1 removed. Failed evaluations (bad expression, wrong type) +/// return `false`. pub fn evaluate_expression(expression: &str, value: &Value) -> bool { eval_expression(expression, value).unwrap_or(false) } @@ -124,7 +63,7 @@ fn eval_expression(expression: &str, value: &Value) -> Result bool { { let cache = COMPILED_EXPRESSIONS diff --git a/crates/dspy-rs/src/typesys/mod.rs b/crates/dspy-rs/src/typesys/mod.rs index b66ffc03..d192eaf8 100644 --- a/crates/dspy-rs/src/typesys/mod.rs +++ b/crates/dspy-rs/src/typesys/mod.rs @@ -13,10 +13,7 @@ pub mod render; pub mod schema; pub use coerce::{Coerced, Flag, coerce}; -pub use constraint::{ - Constraint, ConstraintKind, ConstraintLevel, ConstraintOutcome, ResponseCheck, - evaluate_constraints, evaluate_expression, -}; +pub use constraint::{Constraint, ConstraintKind, ConstraintLevel, evaluate_expression}; pub use render::{schema_block, type_name}; pub use schema::{ ClassDef, EnumDef, EnumValueDef, FieldDef, FieldType, OutputSchema, Schema, TypeTable, diff --git a/crates/dspy-rs/src/typesys/schema.rs b/crates/dspy-rs/src/typesys/schema.rs index 5d5f5f57..042811f0 100644 --- a/crates/dspy-rs/src/typesys/schema.rs +++ b/crates/dspy-rs/src/typesys/schema.rs @@ -8,12 +8,12 @@ use std::collections::HashMap; -use facet::{Def, Facet, Field, ScalarType, Shape, Type, UserType}; +use facet::{Def, Facet, ScalarType, Shape, Type, UserType}; use indexmap::IndexMap; use serde::de::DeserializeOwned; use serde::{Deserialize, Serialize}; -use super::constraint::{Constraint, ConstraintKind}; +use super::constraint::Constraint; /// Structural type of a signature/nested field, mirroring the subset of BAML's `TypeIR` /// that DSRs actually uses. Class/enum variants carry the *internal name* used as a key @@ -331,7 +331,10 @@ impl SchemaBuilder { rendered_name: field.effective_name().to_string(), field_type, docs: doc_to_description(field.doc), - constraints: constraints_from_field(field), + // Constraints on signature fields are attached via + // `FieldMetadataSpec` in `core/schema.rs`; nested-type field + // constraints are only populated by the .dsrs text parser. + constraints: Vec::new(), }); } @@ -427,12 +430,3 @@ fn doc_to_description(doc: &'static [&'static str]) -> Option { } } -/// Reads `#[check]`/`#[assert]` constraints declared directly on a facet field via the -/// `bamltype`-namespaced attrs the derive still emits. -fn constraints_from_field(field: &Field) -> Vec { - // Constraints on signature fields are attached via `FieldMetadataSpec` in - // `core/schema.rs`; nested-type field constraints are not currently surfaced. - let _ = field; - let _ = ConstraintKind::Check; - Vec::new() -} diff --git a/crates/dspy-rs/src/utils/cache.rs b/crates/dspy-rs/src/utils/cache.rs index 6360b91a..07354f21 100644 --- a/crates/dspy-rs/src/utils/cache.rs +++ b/crates/dspy-rs/src/utils/cache.rs @@ -1,8 +1,11 @@ +use std::collections::VecDeque; +use std::sync::{Arc, Mutex, MutexGuard}; + use anyhow::Result; use foyer::{BlockEngineBuilder, DeviceBuilder, FsDeviceBuilder, HybridCache, HybridCacheBuilder}; use serde::{Deserialize, Serialize}; -use tempfile; -use tracing::{debug, trace}; +use tempfile::TempDir; +use tracing::{debug, trace, warn}; use crate::LmUsage; @@ -14,6 +17,9 @@ use crate::LmUsage; /// [`LM`](crate::LM) from the rendered chat — callers never build them by hand. pub type CacheKey = u64; +const MEMORY_CAPACITY: usize = 256 * 1024 * 1024; +const DISK_CAPACITY: usize = 1024 * 1024 * 1024; + /// A cached prompt-response pair. #[derive(Clone, Debug, Serialize, Deserialize)] pub struct CacheEntry { @@ -31,46 +37,94 @@ pub struct CacheEntry { /// Hybrid memory + disk LM response cache. /// /// Uses [foyer](https://docs.rs/foyer) with 256MB memory and 1GB disk (in a -/// temp directory). Maintains a sliding window of the 100 most recent entries -/// for [`inspect_history`](crate::LM::inspect_history). +/// temp directory owned by the cache for its whole lifetime). If the disk +/// tier cannot be initialized, the cache degrades to memory-only with a +/// warning instead of panicking. Maintains a sliding window of the 100 most +/// recent entries for [`inspect_history`](crate::LM::inspect_history). +/// +/// All methods take `&self`: the foyer cache is internally synchronized and +/// the history ring sits behind its own small mutex, so concurrent LM calls +/// never serialize on a cache-wide lock. /// /// Created automatically by [`LM`](crate::LM) — you don't construct this directly. #[derive(Clone)] pub struct ResponseCache { handler: HybridCache, window_size: usize, - history_window: Vec, + /// Debug ring buffer (newest at the back) backing `get_history`. Isolated + /// in its own mutex so history bookkeeping never blocks cache lookups. + history_window: Arc>>, + /// Keeps the disk tier's temp directory alive: `TempDir` deletes the + /// directory on drop, so it must outlive every clone of the cache. + /// `None` when running memory-only. + _disk_dir: Option>, } impl ResponseCache { #[tracing::instrument(name = "dsrs.cache.new", level = "debug")] pub async fn new() -> Self { - let dir = tempfile::tempdir().unwrap(); - - let device = FsDeviceBuilder::new(dir.path()) - .with_capacity(1024 * 1024 * 1024) - .build() - .unwrap(); + let (handler, disk_dir) = match Self::try_build_hybrid().await { + Ok((handler, dir)) => (handler, Some(Arc::new(dir))), + Err(error) => { + warn!( + error = %error, + "disk cache tier unavailable; falling back to memory-only response cache" + ); + (Self::build_memory_only().await, None) + } + }; - let hybrid: HybridCache = HybridCacheBuilder::new() - .memory(256 * 1024 * 1024) - .storage() - .with_engine_config(BlockEngineBuilder::new(device)) - .build() - .await - .unwrap(); let cache = Self { - handler: hybrid, + handler, window_size: 100, - history_window: Vec::new(), + history_window: Arc::new(Mutex::new(VecDeque::new())), + _disk_dir: disk_dir, }; debug!( window_size = cache.window_size, + disk_tier = cache._disk_dir.is_some(), "response cache initialized" ); cache } + /// Builds the full memory + disk hybrid, returning the `TempDir` guard + /// that must be held for as long as the cache lives. + async fn try_build_hybrid() -> Result<(HybridCache, TempDir)> { + let dir = tempfile::tempdir()?; + + let device = FsDeviceBuilder::new(dir.path()) + .with_capacity(DISK_CAPACITY) + .build()?; + + let hybrid = HybridCacheBuilder::new() + .memory(MEMORY_CAPACITY) + .storage() + .with_engine_config(BlockEngineBuilder::new(device)) + .build() + .await?; + Ok((hybrid, dir)) + } + + /// Memory-only fallback: foyer's storage phase defaults to a noop engine, + /// which cannot fail to build (no I/O involved). + async fn build_memory_only() -> HybridCache { + HybridCacheBuilder::new() + .memory(MEMORY_CAPACITY) + .storage() + .build() + .await + .expect("memory-only foyer cache construction cannot fail") + } + + fn lock_history(&self) -> MutexGuard<'_, VecDeque> { + // The ring is debug bookkeeping: recover from a poisoned lock rather + // than propagating a panic into unrelated LM calls. + self.history_window + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + } + /// Fetches the full cached entry (including raw output) for a key. #[tracing::instrument(name = "dsrs.cache.get_entry", level = "trace", skip(self))] pub async fn get_entry(&self, key: CacheKey) -> Result> { @@ -86,17 +140,18 @@ impl ResponseCache { skip(self, entry), fields(window_size = self.window_size) )] - pub fn insert_entry(&mut self, key: CacheKey, entry: CacheEntry) { - self.history_window.insert(0, entry.clone()); - if self.history_window.len() > self.window_size { - self.history_window.pop(); - } - self.handler.insert(key, entry.clone()); - trace!( - history_len = self.history_window.len(), - prompt_len = entry.prompt.len(), - "cache entry inserted" - ); + pub fn insert_entry(&self, key: CacheKey, entry: CacheEntry) { + let prompt_len = entry.prompt.len(); + let history_len = { + let mut history = self.lock_history(); + history.push_back(entry.clone()); + if history.len() > self.window_size { + history.pop_front(); + } + history.len() + }; + self.handler.insert(key, entry); + trace!(history_len, prompt_len, "cache entry inserted"); } /// Returns the `n` most recent cached entries (newest first). @@ -106,9 +161,10 @@ impl ResponseCache { skip(self), fields(n = n) )] - pub async fn get_history(&self, n: usize) -> Result> { - let actual_n = n.min(self.history_window.len()); - trace!(actual_n, "cache history fetched"); - Ok(self.history_window[..actual_n].to_vec()) + pub fn get_history(&self, n: usize) -> Vec { + let history = self.lock_history(); + let entries: Vec = history.iter().rev().take(n).cloned().collect(); + trace!(actual_n = entries.len(), "cache history fetched"); + entries } } diff --git a/crates/dspy-rs/src/utils/telemetry.rs b/crates/dspy-rs/src/utils/telemetry.rs index 43cb8184..a0fb7d54 100644 --- a/crates/dspy-rs/src/utils/telemetry.rs +++ b/crates/dspy-rs/src/utils/telemetry.rs @@ -21,9 +21,13 @@ pub enum TelemetryInitError { /// Behavior: /// - Uses `RUST_LOG` when present. /// - Falls back to `dspy_rs=debug` when `RUST_LOG` is unset/invalid. -/// - Is idempotent: repeated calls are no-ops after first successful init. +/// - Is idempotent: repeated calls are no-ops after the first one, including +/// concurrent calls (initialization is claimed atomically up front, so a +/// racing loser returns `Ok` instead of a spurious `SetGlobalDefault` error). pub fn init_tracing() -> Result<(), TelemetryInitError> { - if TRACING_INITIALIZED.get().is_some() { + // Claim initialization before doing any work: exactly one caller wins the + // `set`, every other (possibly concurrent) caller no-ops with Ok. + if TRACING_INITIALIZED.set(()).is_err() { return Ok(()); } @@ -39,7 +43,6 @@ pub fn init_tracing() -> Result<(), TelemetryInitError> { .finish(); tracing::subscriber::set_global_default(subscriber)?; - let _ = TRACING_INITIALIZED.set(()); Ok(()) } diff --git a/crates/dspy-rs/tests/test_adapters.rs b/crates/dspy-rs/tests/test_adapters.rs index 65ee7279..505e81c9 100644 --- a/crates/dspy-rs/tests/test_adapters.rs +++ b/crates/dspy-rs/tests/test_adapters.rs @@ -1,4 +1,6 @@ +use dspy_rs::ir::SignatureDef; use dspy_rs::{ChatAdapter, Message, Signature}; +use serde_json::Value; #[derive(Signature, Clone, Debug, PartialEq)] struct BasicSignature { @@ -49,12 +51,21 @@ struct DeepFlattenSig { answer: String, } +fn json_map(value: &T) -> serde_json::Map { + match serde_json::to_value(value).expect("serializable") { + Value::Object(map) => map, + other => panic!("expected object, got {other:?}"), + } +} + #[test] fn chat_adapter_formats_typed_system_prompt() { let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system prompt should format"); + let system = adapter.build_system_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + None, + ); assert!(system.contains("Your input fields are:")); assert!(system.contains("`problem`")); @@ -66,14 +77,20 @@ fn chat_adapter_formats_typed_system_prompt() { #[test] fn chat_adapter_formats_user_and_assistant_messages() { let adapter = ChatAdapter; + let def = SignatureDef::of::(); - let user = adapter.format_user_message_typed::(&BasicSignatureInput { - problem: "What is the capital of France?".to_string(), - }); - let assistant = - adapter.format_assistant_message_typed::(&BasicSignatureOutput { + let user = adapter.format_input_def( + def, + &json_map(&BasicSignatureInput { + problem: "What is the capital of France?".to_string(), + }), + ); + let assistant = adapter.format_output_def( + def, + &json_map(&BasicSignatureOutput { answer: "Paris".to_string(), - }); + }), + ); assert!(user.contains("[[ ## problem ## ]]")); assert!(user.contains("What is the capital of France?")); @@ -90,10 +107,16 @@ fn chat_adapter_parses_typed_response() { let adapter = ChatAdapter; let response = Message::assistant("[[ ## answer ## ]]\nParis\n\n[[ ## completed ## ]]"); - let (output, field_meta) = adapter - .parse_response_typed::(&response) + let (output_map, field_meta) = adapter + .parse_output_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + &response, + ) .expect("typed response should parse"); + let output: BasicSignatureOutput = + serde_json::from_value(Value::Object(output_map)).expect("typed assembly"); assert_eq!(output.answer, "Paris"); assert_eq!( field_meta.get("answer").map(|meta| meta.raw_text.as_str()), @@ -115,14 +138,19 @@ fn parse_sections_accepts_non_word_field_names() { #[test] fn chat_adapter_formats_user_messages_with_multi_level_flatten_paths() { let adapter = ChatAdapter; - let user = adapter.format_user_message_typed::(&DeepFlattenSigInput { - question: "What should we answer?".to_string(), - middle: FlattenMiddleSigInput { - inner: FlattenLeafSigInput { - leaf: "flattened-value".to_string(), + // Multi-level `#[flatten]` inputs serialize flat, keyed by leaf name — + // exactly the canonical `FieldDef::name` keys the def lane renders from. + let user = adapter.format_input_def( + SignatureDef::of::(), + &json_map(&DeepFlattenSigInput { + question: "What should we answer?".to_string(), + middle: FlattenMiddleSigInput { + inner: FlattenLeafSigInput { + leaf: "flattened-value".to_string(), + }, }, - }, - }); + }), + ); assert!( user.contains("[[ ## question ## ]]"), diff --git a/crates/dspy-rs/tests/test_bootstrap_fewshot.rs b/crates/dspy-rs/tests/test_bootstrap_fewshot.rs index 4684f2f5..321e2d89 100644 --- a/crates/dspy-rs/tests/test_bootstrap_fewshot.rs +++ b/crates/dspy-rs/tests/test_bootstrap_fewshot.rs @@ -4,8 +4,8 @@ use anyhow::Result; use dspy_rs::{ - BootstrapFewShot, Eval, LM, LMClient, Module, ModuleState, Optimizer, Predict, PredictError, - Predicted, Signature, TestCompletionModel, TypedMetric, + BootstrapFewShot, Eval, LM, LMClient, Module, ModuleState, Predict, PredictError, Predicted, + Signature, SpanId, TestCompletionModel, Trace, TypedMetric, }; use rig::completion::AssistantContent; use rig::message::Text; @@ -20,12 +20,12 @@ struct BootSig { answer: String, } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct BootModule { predictor: Predict, } +dspy_rs::predictors!(BootModule { predictor }); + impl Module for BootModule { type Input = BootSigInput; type Output = BootSigOutput; @@ -106,7 +106,7 @@ async fn bootstrap_harvests_demos_and_adopts_when_better() { .build(); let report = bootstrap - .compile(&mut module, trainset(), &ExactMatch) + .compile_module(&mut module, &trainset(), &ExactMatch) .await .expect("bootstrap should succeed on canned responses"); @@ -152,7 +152,7 @@ async fn bootstrap_keeps_baseline_when_candidate_is_worse() { .build(); let report = bootstrap - .compile(&mut module, trainset(), &ExactMatch) + .compile_module(&mut module, &trainset(), &ExactMatch) .await .unwrap(); @@ -184,7 +184,7 @@ async fn bootstrap_without_qualifying_rollouts_skips_candidate_eval() { .build(); let report = bootstrap - .compile(&mut module, trainset(), &ExactMatch) + .compile_module(&mut module, &trainset(), &ExactMatch) .await .unwrap(); @@ -195,6 +195,166 @@ async fn bootstrap_without_qualifying_rollouts_skips_candidate_eval() { assert_eq!(report.spend.metric_calls, 3, "only the teacher pass ran"); } +// --------------------------------------------------------------------------- +// Per-span credit assignment (RFC 0004 §4): a draft/refine pipeline where the +// refine step recovers from a bad draft. Whole-rollout credit harvests the bad +// drafts as demos; a metric that attaches span evals keeps them out. +// --------------------------------------------------------------------------- + +struct TwoStepModule { + draft: Predict, + refine: Predict, +} + +dspy_rs::predictors!(TwoStepModule { draft, refine }); + +impl Module for TwoStepModule { + type Input = BootSigInput; + type Output = BootSigOutput; + + async fn forward(&self, input: BootSigInput) -> Result, PredictError> { + let draft = self.draft.call(input).await?; + self.refine + .call(BootSigInput { + prompt: draft.answer.clone(), + }) + .await + } +} + +/// Whole-rollout exact match, no span hook — the pre-RFC-0004 behavior. +struct TwoStepExactMatch; + +impl TypedMetric<(BootSigInput, BootSigOutput), TwoStepModule> for TwoStepExactMatch { + async fn evaluate( + &self, + example: &(BootSigInput, BootSigOutput), + prediction: &Predicted, + _trace: Option<&Trace>, + ) -> Result { + let score = (prediction.answer == example.1.answer) as u8 as f64; + Ok(Eval::score(score)) + } +} + +/// Same rollout score, plus span-level credit: each draft span is scored on +/// its own answer, so a recovered-from draft gets 0.0 while the rollout +/// still gets full credit. +struct SpanAwareExactMatch; + +impl TypedMetric<(BootSigInput, BootSigOutput), TwoStepModule> for SpanAwareExactMatch { + async fn evaluate( + &self, + example: &(BootSigInput, BootSigOutput), + prediction: &Predicted, + _trace: Option<&Trace>, + ) -> Result { + let score = (prediction.answer == example.1.answer) as u8 as f64; + Ok(Eval::score(score)) + } + + async fn evaluate_spans( + &self, + example: &(BootSigInput, BootSigOutput), + _prediction: &Predicted, + trace: &Trace, + ) -> Result> { + Ok(trace + .for_component("draft") + .filter_map(|span| { + let answer = span.output.as_ref()?.get("answer")?.as_str()?; + let score = (answer == example.1.answer) as u8 as f64; + Some((span.id, Eval::score(score))) + }) + .collect()) + } +} + +async fn two_step_module(client: TestCompletionModel) -> TwoStepModule { + let lm = temp_env::async_with_vars( + [("OPENAI_API_KEY", Some("test"))], + LM::builder() + .model("openai:gpt-4o-mini".to_string()) + .build(), + ) + .await + .unwrap() + .with_client(LMClient::Test(client)) + .await + .unwrap(); + TwoStepModule { + draft: Predict::::builder().lm(lm.clone()).build(), + refine: Predict::::builder().lm(lm).build(), + } +} + +/// Every rollout: draft answers wrong (`a0`/`a1`/`a2`), refine recovers with +/// the gold answer. Teacher pass then candidate pass, sequential. +fn recovery_responses() -> TestCompletionModel { + TestCompletionModel::new([ + // Teacher pass: (draft, refine) per example. + answer_response("a0"), + answer_response("0"), + answer_response("a1"), + answer_response("1"), + answer_response("a2"), + answer_response("2"), + // Candidate pass: same shape. + answer_response("a0"), + answer_response("0"), + answer_response("a1"), + answer_response("1"), + answer_response("a2"), + answer_response("2"), + ]) +} + +#[tokio::test] +async fn whole_rollout_metric_harvests_recovered_from_drafts() { + // Control: without span evals, the winning rollouts vouch for every span + // — the wrong drafts become draft demos. + let mut module = two_step_module(recovery_responses()).await; + + let bootstrap = BootstrapFewShot::builder() + .min_demo_score(1.0) + .eval_concurrency(1) + .build(); + + let report = bootstrap + .compile_module(&mut module, &trainset(), &TwoStepExactMatch) + .await + .unwrap(); + + assert!((report.baseline_score - 1.0).abs() < 1e-9); + assert_eq!(report.demos_per_predictor.get("draft"), Some(&3)); + assert_eq!(report.demos_per_predictor.get("refine"), Some(&3)); +} + +#[tokio::test] +async fn span_evals_keep_recovered_from_drafts_out_of_the_demo_pool() { + // Same rollouts, but the metric attaches per-span credit: draft spans + // score 0.0, so only the refine step's demos are harvested. + let mut module = two_step_module(recovery_responses()).await; + + let bootstrap = BootstrapFewShot::builder() + .min_demo_score(1.0) + .eval_concurrency(1) + .build(); + + let report = bootstrap + .compile_module(&mut module, &trainset(), &SpanAwareExactMatch) + .await + .unwrap(); + + assert!((report.baseline_score - 1.0).abs() < 1e-9); + assert_eq!( + report.demos_per_predictor.get("draft"), + None, + "a span the metric scored 0.0 must not become a demo, even from a full-credit rollout" + ); + assert_eq!(report.demos_per_predictor.get("refine"), Some(&3)); +} + #[tokio::test] async fn bootstrap_respects_max_demos() { let client = TestCompletionModel::new([ @@ -214,7 +374,7 @@ async fn bootstrap_respects_max_demos() { .build(); let report = bootstrap - .compile(&mut module, trainset(), &ExactMatch) + .compile_module(&mut module, &trainset(), &ExactMatch) .await .unwrap(); diff --git a/crates/dspy-rs/tests/test_caller_managed_conversation.rs b/crates/dspy-rs/tests/test_caller_managed_conversation.rs index a9a47662..d88e3280 100644 --- a/crates/dspy-rs/tests/test_caller_managed_conversation.rs +++ b/crates/dspy-rs/tests/test_caller_managed_conversation.rs @@ -4,8 +4,7 @@ //! manages the conversation loop, not the LM layer's auto tool loop. use dspy_rs::{ - LM, LMClient, Message, Predict, Role, Signature, TestCompletionModel, - ToolLoopMode, configure, + LM, LMClient, Message, Predict, Role, Signature, TestCompletionModel, ToolLoopMode, configure, }; use rig::completion::AssistantContent; use rig::message::{Text, ToolCall, ToolFunction}; @@ -83,6 +82,7 @@ async fn caller_managed_tool_loop_with_conversation() { // Turn 1: Build chat and call LM let chat = predict .build_chat(&input) + .await .expect("build_chat should succeed"); let (first_result, mut chat) = predict .call_and_parse(chat) @@ -178,7 +178,7 @@ async fn parse_failure_on_second_turn_includes_correct_raw_response() { }; // Turn 1: succeeds - let chat = predict.build_chat(&input).expect("build_chat"); + let chat = predict.build_chat(&input).await.expect("build_chat"); let (first_result, mut chat) = predict.call_and_parse(chat).await.expect("turn 1"); assert_eq!(first_result.into_inner().result, "first answer"); diff --git a/crates/dspy-rs/tests/test_chat.rs b/crates/dspy-rs/tests/test_chat.rs index fd8e0fe9..d2a20122 100644 --- a/crates/dspy-rs/tests/test_chat.rs +++ b/crates/dspy-rs/tests/test_chat.rs @@ -44,47 +44,6 @@ fn test_chat_pop() { assert_eq!(chat.len(), 0); } -#[rstest] -fn test_chat_to_json_and_back() { - let chat = Chat::new(vec![ - Message::system("You are a helpful assistant."), - Message::user("Hello, world!"), - Message::assistant("Hello, world to you!"), - ]); - let json_dump = chat.to_json(); - let reparsed = Chat::new(vec![]).from_json(json_dump).unwrap(); - - assert_eq!(reparsed.len(), 3); - assert_eq!(reparsed.messages[0].role, Role::System); - assert_eq!( - reparsed.messages[0].content(), - "You are a helpful assistant." - ); - assert_eq!(reparsed.messages[1].role, Role::User); - assert_eq!(reparsed.messages[1].content(), "Hello, world!"); - assert_eq!(reparsed.messages[2].role, Role::Assistant); - assert_eq!(reparsed.messages[2].content(), "Hello, world to you!"); -} - -#[rstest] -fn test_chat_from_legacy_json() { - // Legacy format: "content" is a plain string - let json = json!([ - {"role":"system","content":"You are a helpful assistant."}, - {"role":"user","content":"Hello, world!"}, - {"role":"assistant","content":"Hello, world to you!"} - ]); - let chat = Chat::new(vec![]).from_json(json).unwrap(); - - assert_eq!(chat.len(), 3); - assert_eq!(chat.messages[0].role, Role::System); - assert_eq!(chat.messages[0].content(), "You are a helpful assistant."); - assert_eq!(chat.messages[1].role, Role::User); - assert_eq!(chat.messages[1].content(), "Hello, world!"); - assert_eq!(chat.messages[2].role, Role::Assistant); - assert_eq!(chat.messages[2].content(), "Hello, world to you!"); -} - #[rstest] fn test_chat_push_all() { let mut chat1 = Chat::new(vec![ @@ -125,47 +84,6 @@ fn test_chat_push_all_empty() { assert_eq!(chat1.messages[0].content(), "System message"); } -#[rstest] -fn test_new_variants_round_trip_json() { - let call = ToolCall::new( - "call-1".to_string(), - ToolFunction { - name: "lookup".to_string(), - arguments: json!({ "query": "rust" }), - }, - ); - let result = ToolResult { - id: "call-1".to_string(), - call_id: Some("provider-call-1".to_string()), - content: OneOrMany::one(ToolResultContent::text("result payload")), - }; - let reasoning = Reasoning::new("thinking..."); - - let chat = Chat::new(vec![ - Message::system("You are a tool-using assistant."), - Message::tool_call(call.clone()), - Message::tool_result(result.clone()), - Message::reasoning(reasoning.clone()), - ]); - - let json_dump = chat.to_json(); - let reparsed = Chat::new(vec![]).from_json(json_dump).unwrap(); - assert_eq!(reparsed.len(), 4); - - assert_eq!(reparsed.messages[0].role, Role::System); - - assert_eq!(reparsed.messages[1].role, Role::Assistant); - assert!(reparsed.messages[1].has_tool_calls()); - let reparsed_calls = reparsed.messages[1].tool_calls(); - assert_eq!(reparsed_calls[0].function.name, call.function.name); - - assert_eq!(reparsed.messages[2].role, Role::User); - assert!(reparsed.messages[2].has_tool_results()); - - assert_eq!(reparsed.messages[3].role, Role::Assistant); - assert!(reparsed.messages[3].has_reasoning()); -} - #[rstest] fn test_system_prompt_and_rig_chat_history() { let chat = Chat::new(vec![ diff --git a/crates/dspy-rs/tests/test_chat_adapter_schema.rs b/crates/dspy-rs/tests/test_chat_adapter_schema.rs index 388218a7..357fd4a0 100644 --- a/crates/dspy-rs/tests/test_chat_adapter_schema.rs +++ b/crates/dspy-rs/tests/test_chat_adapter_schema.rs @@ -1,4 +1,12 @@ +//! Def-lane response parsing: canonical field-name keying and dotted +//! (`alias`) markers. Since Predict routes through the IR interpreter, the +//! adapter parses against [`SignatureDef`]s; output maps and metadata key by +//! canonical `FieldDef::name` (Predict translates to `rust_name` keying at +//! its boundary). + +use dspy_rs::ir::SignatureDef; use dspy_rs::{CallMetadata, ChatAdapter, Message, Predicted, Signature}; +use serde_json::Value; #[derive(Signature, Clone, Debug)] /// Adapter schema parse fixture. @@ -22,14 +30,21 @@ struct AliasSig { } #[test] -fn parse_response_typed_uses_schema_field_names() { +fn parse_output_def_uses_canonical_field_names() { let adapter = ChatAdapter; let response = Message::assistant("[[ ## answer ## ]]\nParis\n\n[[ ## completed ## ]]\n"); - let (output, field_meta) = adapter - .parse_response_typed::(&response) - .expect("typed parse should succeed"); + let (output_map, field_meta) = adapter + .parse_output_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + &response, + ) + .expect("def parse should succeed"); + assert_eq!(output_map.get("answer"), Some(&Value::from("Paris"))); + let output: ExampleSigOutput = + serde_json::from_value(Value::Object(output_map)).expect("typed assembly"); assert_eq!(output.answer, "Paris"); let answer_meta = field_meta.get("answer").expect("answer field metadata"); assert_eq!(answer_meta.raw_text.trim(), "Paris"); @@ -50,14 +65,20 @@ fn parse_response_typed_uses_schema_field_names() { } #[test] -fn parse_response_typed_accepts_dotted_field_markers() { +fn parse_output_def_accepts_dotted_field_markers() { let adapter = ChatAdapter; let response = Message::assistant("[[ ## answer.value ## ]]\nParis\n\n[[ ## completed ## ]]\n"); - let (output, field_meta) = adapter - .parse_response_typed::(&response) - .expect("typed parse should succeed for dotted aliases"); + let (output_map, field_meta) = adapter + .parse_output_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + &response, + ) + .expect("def parse should succeed for dotted aliases"); + let output: AliasSigOutput = + serde_json::from_value(Value::Object(output_map)).expect("typed assembly"); assert_eq!(output.answer, "Paris"); assert_eq!( field_meta diff --git a/crates/dspy-rs/tests/test_chat_prompt_composition.rs b/crates/dspy-rs/tests/test_chat_prompt_composition.rs index c5c4934a..236e0489 100644 --- a/crates/dspy-rs/tests/test_chat_prompt_composition.rs +++ b/crates/dspy-rs/tests/test_chat_prompt_composition.rs @@ -1,4 +1,9 @@ +//! Prompt composition through the def lane — the one prompt path since +//! Predict routes through the IR interpreter. + +use dspy_rs::ir::SignatureDef; use dspy_rs::{ChatAdapter, Demo, Signature}; +use serde_json::Value; #[derive(Signature, Clone, Debug)] /// Answer the prompt using the provided context. @@ -25,6 +30,21 @@ struct EmptyInstructionSig { summary: String, } +fn json_map(value: &T) -> serde_json::Map { + match serde_json::to_value(value).expect("serializable") { + Value::Object(map) => map, + other => panic!("expected object, got {other:?}"), + } +} + +fn system_prompt(instruction_override: Option<&str>) -> String { + ChatAdapter.build_system_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + instruction_override, + ) +} + fn find_required(haystack: &str, needle: &str) -> usize { haystack .find(needle) @@ -40,10 +60,7 @@ fn response_instruction_line(message: &str) -> &str { #[test] fn system_prompt_includes_all_sections_in_order_with_boundaries() { - let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system prompt should format"); + let system = system_prompt::(None); let descriptions_idx = find_required(&system, "Your input fields are:"); let structure_idx = find_required( @@ -80,10 +97,7 @@ fn system_prompt_includes_all_sections_in_order_with_boundaries() { #[test] fn system_prompt_field_descriptions_and_structure_are_present() { - let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system prompt should format"); + let system = system_prompt::(None); assert!(system.contains("`question` (string): User question")); assert!(system.contains("`context` (string): Retrieved context")); @@ -101,10 +115,7 @@ fn system_prompt_field_descriptions_and_structure_are_present() { #[test] fn response_instruction_line_orders_output_fields() { - let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system prompt should format"); + let system = system_prompt::(None); let line = response_instruction_line(&system); let answer_idx = find_required(line, "[[ ## answer ## ]]"); @@ -115,11 +126,8 @@ fn response_instruction_line_orders_output_fields() { #[test] fn instruction_override_is_used_in_objective_section() { - let adapter = ChatAdapter; let override_instruction = "Follow the rubric.\nCite the context."; - let system = adapter - .format_system_message_typed_with_instruction::(Some(override_instruction)) - .expect("system prompt should format with override"); + let system = system_prompt::(Some(override_instruction)); assert!(system.contains("In adhering to this structure, your objective is:")); assert!(system.contains(" Follow the rubric.")); @@ -129,57 +137,37 @@ fn instruction_override_is_used_in_objective_section() { #[test] fn empty_instruction_uses_generated_fallback_objective() { - let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system prompt should format"); + let system = system_prompt::(None); assert!(system.contains("In adhering to this structure, your objective is:")); assert!(system.contains("Given the fields `topic`, produce the fields `summary`.")); } #[test] -fn typed_and_schema_system_builders_match() { - let adapter = ChatAdapter; - let typed = adapter - .format_system_message_typed_with_instruction::(Some("Override objective")) - .expect("typed system prompt"); - let schema = adapter - .build_system(PromptPartsSig::schema(), Some("Override objective")) - .expect("schema system prompt"); - - assert_eq!(typed, schema); -} - -#[test] -fn typed_and_schema_user_builders_match_and_append_requirements() { +fn user_builder_appends_requirements() { let adapter = ChatAdapter; let input = PromptPartsSigInput { question: "What is the capital of France?".to_string(), context: "Facts: Paris is the capital city of France.".to_string(), }; - let typed = adapter.format_user_message_typed::(&input); - let schema = adapter.format_input(PromptPartsSig::schema(), &input); - assert_eq!(typed, schema); + let user = adapter.format_input_def(SignatureDef::of::(), &json_map(&input)); - assert!(typed.contains("[[ ## question ## ]]")); - assert!(typed.contains("What is the capital of France?")); - assert!(typed.contains("[[ ## context ## ]]")); - assert!(typed.contains("Facts: Paris is the capital city of France.")); + assert!(user.contains("[[ ## question ## ]]")); + assert!(user.contains("What is the capital of France?")); + assert!(user.contains("[[ ## context ## ]]")); + assert!(user.contains("Facts: Paris is the capital city of France.")); - let context_idx = find_required(&typed, "Facts: Paris is the capital city of France."); - let instruction_idx = find_required(&typed, "Respond with the corresponding output fields"); + let context_idx = find_required(&user, "Facts: Paris is the capital city of France."); + let instruction_idx = find_required(&user, "Respond with the corresponding output fields"); assert!(context_idx < instruction_idx); assert_eq!( - typed - .matches("Respond with the corresponding output fields") + user.matches("Respond with the corresponding output fields") .count(), 1 ); assert!( - typed - .trim_end() + user.trim_end() .ends_with("and then ending with the marker for `[[ ## completed ## ]]`.") ); } @@ -198,7 +186,9 @@ fn demo_format_composes_user_and_assistant_parts() { }, ); - let (user_msg, assistant_msg) = adapter.format_demo_typed::(&demo); + let def = SignatureDef::of::(); + let user_msg = adapter.format_input_def(def, &json_map(&demo.input)); + let assistant_msg = adapter.format_output_def(def, &json_map(&demo.output)); assert!(user_msg.contains("[[ ## question ## ]]")); assert!(user_msg.contains("[[ ## context ## ]]")); @@ -212,19 +202,18 @@ fn demo_format_composes_user_and_assistant_parts() { } #[test] -fn typed_and_schema_assistant_builders_match_and_end_with_completed_marker() { +fn assistant_builder_orders_fields_and_ends_with_completed_marker() { let adapter = ChatAdapter; let output = PromptPartsSigOutput { answer: "Paris".to_string(), confidence: 0.9, }; - let typed = adapter.format_assistant_message_typed::(&output); - let schema = adapter.format_output(PromptPartsSig::schema(), &output); - assert_eq!(typed, schema); + let assistant = + adapter.format_output_def(SignatureDef::of::(), &json_map(&output)); - let answer_idx = find_required(&typed, "[[ ## answer ## ]]"); - let confidence_idx = find_required(&typed, "[[ ## confidence ## ]]"); + let answer_idx = find_required(&assistant, "[[ ## answer ## ]]"); + let confidence_idx = find_required(&assistant, "[[ ## confidence ## ]]"); assert!(answer_idx < confidence_idx); - assert!(typed.trim_end().ends_with("[[ ## completed ## ]]")); + assert!(assistant.trim_end().ends_with("[[ ## completed ## ]]")); } diff --git a/crates/dspy-rs/tests/test_chat_prompt_golden.rs b/crates/dspy-rs/tests/test_chat_prompt_golden.rs index 00eac98d..da2369fa 100644 --- a/crates/dspy-rs/tests/test_chat_prompt_golden.rs +++ b/crates/dspy-rs/tests/test_chat_prompt_golden.rs @@ -1,4 +1,13 @@ +//! Golden prompt tests: the exact bytes of the `[[ ## field ## ]]` protocol. +//! +//! These render through the def lane ([`SignatureDef::of`] + the `*_def` +//! adapter methods) — the one prompt path since Predict routes through the IR +//! interpreter. The expected strings are the historical static-lane bytes and +//! must never change silently. + +use dspy_rs::ir::SignatureDef; use dspy_rs::{ChatAdapter, Demo, Signature}; +use serde_json::Value; #[derive(Signature, Clone, Debug)] struct GoldenSig { @@ -9,12 +18,21 @@ struct GoldenSig { answer: String, } +fn json_map(value: &T) -> serde_json::Map { + match serde_json::to_value(value).expect("serializable") { + Value::Object(map) => map, + other => panic!("expected object, got {other:?}"), + } +} + #[test] fn golden_system_prompt_is_stable() { let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system prompt should format"); + let system = adapter.build_system_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + None, + ); let expected = concat!( "Your input fields are:\n", @@ -48,7 +66,7 @@ fn golden_user_prompt_is_stable() { let input = GoldenSigInput { question: "What is 2+2?".to_string(), }; - let user = adapter.format_user_message_typed::(&input); + let user = adapter.format_input_def(SignatureDef::of::(), &json_map(&input)); let expected = concat!( "[[ ## question ## ]]\n", @@ -66,7 +84,7 @@ fn golden_assistant_prompt_is_stable() { let output = GoldenSigOutput { answer: "4".to_string(), }; - let assistant = adapter.format_assistant_message_typed::(&output); + let assistant = adapter.format_output_def(SignatureDef::of::(), &json_map(&output)); let expected = concat!( "[[ ## answer ## ]]\n", @@ -89,7 +107,9 @@ fn golden_demo_messages_are_stable() { }, ); - let (user, assistant) = adapter.format_demo_typed::(&demo); + let def = SignatureDef::of::(); + let user = adapter.format_input_def(def, &json_map(&demo.input)); + let assistant = adapter.format_output_def(def, &json_map(&demo.output)); let expected_user = concat!( "[[ ## question ## ]]\n", diff --git a/crates/dspy-rs/tests/test_code_mode.rs b/crates/dspy-rs/tests/test_code_mode.rs index 2ef5ce60..679fb1ea 100644 --- a/crates/dspy-rs/tests/test_code_mode.rs +++ b/crates/dspy-rs/tests/test_code_mode.rs @@ -2,7 +2,6 @@ //! sandboxed `run_js` tool, usable in the LM tool loop today. A canned LM //! emits a `run_js` call whose script chains two tool calls in a single //! execution — the token-economy win over N JSON tool calls. -#![cfg(feature = "code-mode")] use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; diff --git a/crates/dspy-rs/tests/test_dataloader.rs b/crates/dspy-rs/tests/test_dataloader.rs index 6db91346..d7078537 100644 --- a/crates/dspy-rs/tests/test_dataloader.rs +++ b/crates/dspy-rs/tests/test_dataloader.rs @@ -1,10 +1,14 @@ +//! Exercises the CSV/Parquet/HF loaders, so the whole suite requires the +//! (default-on) `data` feature. +#![cfg(feature = "data")] + use anyhow::{Result, anyhow}; use arrow::array::{ArrayRef, Int64Array, StringArray}; use arrow::datatypes::{DataType, Field, Schema}; use arrow::record_batch::RecordBatch; use bon::Builder; use dspy_rs::{ - COPRO, CallMetadata, DataLoader, Eval, Example, Module, Optimizer, Predict, PredictError, + COPRO, CallMetadata, DataLoader, Eval, Example, Module, Predict, PredictError, Predicted, Signature, TypedLoadOptions, TypedMetric, UnknownFieldPolicy, average_score, evaluate_trainset, }; @@ -50,6 +54,8 @@ struct EchoModule { predictor: Predict, } +dspy_rs::predictors!(EchoModule { predictor }); + impl Module for EchoModule { type Input = LoaderSigInput; type Output = LoaderSigOutput; @@ -512,7 +518,7 @@ async fn typed_loader_outputs_feed_evaluator_and_optimizer_paths() -> Result<()> let optimizer = COPRO::builder().breadth(2).depth(1).build(); optimizer - .compile(&mut module, trainset, &metric) + .compile_module(&mut module, &trainset, &metric) .await?; Ok(()) diff --git a/crates/dspy-rs/tests/test_eval_engine.rs b/crates/dspy-rs/tests/test_eval_engine.rs index 3a6699af..22a90531 100644 --- a/crates/dspy-rs/tests/test_eval_engine.rs +++ b/crates/dspy-rs/tests/test_eval_engine.rs @@ -1,6 +1,6 @@ //! Shared evaluation engine (vision §5.4) coverage: bounded-concurrency //! fan-out, rollout caching, budget metering, minibatch gating, matrix/Pareto -//! bookkeeping, checkpoint/resume, and the candidate apply/restore seam. +//! bookkeeping, and the ambient candidate-injection + one-shot install model. use std::collections::HashMap; use std::sync::Arc; @@ -9,9 +9,9 @@ use std::time::Duration; use anyhow::Result; use dspy_rs::{ - Budget, CallMetadata, Candidate, EngineConfig, Eval, EvalEngine, EvalOutcome, GateOutcome, LM, - LMClient, Module, ModuleState, Predict, PredictError, Predicted, Signature, - TestCompletionModel, TypedMetric, apply_candidate, restore_candidate, + Budget, CallMetadata, Candidate, Engine, EngineConfig, Eval, EvalOutcome, GateOutcome, LM, + LMClient, Module, ModuleState, OptimizeTarget, Predict, PredictError, Predicted, Signature, + TestCompletionModel, TypedMetric, }; use rig::completion::AssistantContent; use rig::message::Text; @@ -46,12 +46,12 @@ fn trainset(n: usize) -> Vec<(EngSigInput, EngSigOutput)> { // --------------------------------------------------------------------------- /// No-LM module: echoes the prompt back as the answer. -#[derive(facet::Facet)] -#[facet(crate = facet)] struct EchoModule { predictor: Predict, } +dspy_rs::predictors!(EchoModule { predictor }); + impl EchoModule { fn new() -> Self { Self { @@ -77,14 +77,13 @@ impl Module for EchoModule { /// Echo module that blocks each rollout on a barrier: the test only completes /// if all N rollouts are genuinely in flight at once. -#[derive(facet::Facet)] -#[facet(crate = facet)] struct BarrierModule { predictor: Predict, - #[facet(opaque, skip)] barrier: Arc, } +dspy_rs::predictors!(BarrierModule { predictor }); + impl Module for BarrierModule { type Input = EngSigInput; type Output = EngSigOutput; @@ -103,18 +102,15 @@ impl Module for BarrierModule { /// Echo module that gauges how many rollouts are inside `forward` at once. /// The pair barrier forces at least two to overlap, so a sequential engine /// deadlocks; the gauge proves the bound is never exceeded. -#[derive(facet::Facet)] -#[facet(crate = facet)] struct GaugeModule { predictor: Predict, - #[facet(opaque, skip)] in_flight: Arc, - #[facet(opaque, skip)] max_in_flight: Arc, - #[facet(opaque, skip)] pair_barrier: Arc, } +dspy_rs::predictors!(GaugeModule { predictor }); + impl Module for GaugeModule { type Input = EngSigInput; type Output = EngSigOutput; @@ -135,12 +131,12 @@ impl Module for GaugeModule { /// Real `Predict` leaf backed by [`TestCompletionModel`] — every rollout is an /// actual LM call against the canned response queue. -#[derive(facet::Facet)] -#[facet(crate = facet)] struct LmModule { predictor: Predict, } +dspy_rs::predictors!(LmModule { predictor }); + impl Module for LmModule { type Input = EngSigInput; type Output = EngSigOutput; @@ -253,21 +249,19 @@ async fn fan_out_runs_examples_concurrently_with_correct_results() { barrier: Arc::new(tokio::sync::Barrier::new(N)), }; let metric = IndexMetric; - let mut engine = EvalEngine::new( - trainset(N), - &metric, - EngineConfig { - concurrency: N, - ..EngineConfig::default() - }, - ); + let examples = trainset(N); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig { + concurrency: N, + ..EngineConfig::default() + }); let candidate = engine.register(Candidate::new()); // The barrier only opens once all N rollouts are in flight simultaneously; // a sequential engine would deadlock and trip the timeout. let outcome = tokio::time::timeout( Duration::from_secs(10), - engine.evaluate(&mut module, candidate, None), + engine.evaluate(&target, candidate, None), ) .await .expect("fan-out must run all rollouts concurrently") @@ -294,11 +288,13 @@ async fn fan_out_runs_examples_concurrently_with_correct_results() { async fn subset_evaluation_respects_order_and_matrix_cells() { let mut module = EchoModule::new(); let metric = IndexMetric; - let mut engine = EvalEngine::new(trainset(4), &metric, EngineConfig::default()); + let examples = trainset(4); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig::default()); let candidate = engine.register(Candidate::new()); let eval = engine - .evaluate(&mut module, candidate, Some(&[3, 1, 2])) + .evaluate(&target, candidate, Some(&[3, 1, 2])) .await .unwrap() .completed() @@ -321,11 +317,13 @@ async fn rollout_cache_serves_repeats_without_lm_calls() { ]); let mut module = lm_module(client.clone()).await; let metric = ExactMatch; - let mut engine = EvalEngine::new(trainset(3), &metric, EngineConfig::default()); + let examples = trainset(3); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig::default()); let candidate = engine.register(Candidate::new()); let first = engine - .evaluate(&mut module, candidate, None) + .evaluate(&target, candidate, None) .await .unwrap() .completed() @@ -337,7 +335,7 @@ async fn rollout_cache_serves_repeats_without_lm_calls() { // Re-evaluation: the response queue is now EMPTY, so any LM call would // error. The cache must serve all three rollouts. let second = engine - .evaluate(&mut module, candidate, None) + .evaluate(&target, candidate, None) .await .expect("cached re-evaluation must not touch the LM") .completed() @@ -348,13 +346,14 @@ async fn rollout_cache_serves_repeats_without_lm_calls() { assert_eq!(engine.spend().metric_calls, 3, "metric must not re-run"); assert_eq!(engine.spend().cache_hits, 3); - // A *different* candidate is a cache miss: it runs fresh LM calls. + // A *different* candidate is a cache miss: it runs fresh LM calls, with + // the candidate's instruction injected ambiently. client.push_response(answer_response("0")); client.push_response(answer_response("wrong")); client.push_response(answer_response("2")); let other = engine.register(Candidate::with_instruction("predictor", "be brief")); let third = engine - .evaluate(&mut module, other, None) + .evaluate(&target, other, None) .await .unwrap() .completed() @@ -362,8 +361,9 @@ async fn rollout_cache_serves_repeats_without_lm_calls() { assert!((third.mean() - 2.0 / 3.0).abs() < 1e-9); assert_eq!(engine.spend().lm_calls, 6); - // Candidate evaluation restored the module: no instruction override left. - let state = ModuleState::from_module(&mut module).unwrap(); + // Injection was ambient: the module itself never changed. + drop(target); + let state = ModuleState::from_module(&module).unwrap(); assert_eq!(state.predictors["predictor"].instruction_override, None); } @@ -371,24 +371,22 @@ async fn rollout_cache_serves_repeats_without_lm_calls() { async fn budget_stops_cleanly_and_reports_spend() { let mut module = EchoModule::new(); let metric = IndexMetric; - let mut engine = EvalEngine::new( - trainset(3), - &metric, - EngineConfig { - budget: Budget { - max_metric_calls: Some(4), - ..Budget::unlimited() - }, - ..EngineConfig::default() + let examples = trainset(3); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig { + budget: Budget { + max_metric_calls: Some(4), + ..Budget::unlimited() }, - ); + ..EngineConfig::default() + }); let base = engine.register(Candidate::new()); let better = engine.register(Candidate::with_instruction("predictor", "improved")); assert!(engine.budget_allows(3)); engine - .evaluate(&mut module, base, None) + .evaluate(&target, base, None) .await .unwrap() .completed() @@ -396,7 +394,7 @@ async fn budget_stops_cleanly_and_reports_spend() { assert_eq!(engine.spend().metric_calls, 3); // A second full eval needs 3 more rollouts; only 1 remains. - match engine.evaluate(&mut module, better, None).await.unwrap() { + match engine.evaluate(&target, better, None).await.unwrap() { EvalOutcome::BudgetExhausted { needed } => assert_eq!(needed, 3), EvalOutcome::Complete(_) => panic!("engine must stop at the budget"), } @@ -404,7 +402,7 @@ async fn budget_stops_cleanly_and_reports_spend() { // Cache-served batches are free even at the budget edge. let replay = engine - .evaluate(&mut module, base, None) + .evaluate(&target, base, None) .await .unwrap() .completed() @@ -413,7 +411,7 @@ async fn budget_stops_cleanly_and_reports_spend() { // The final budget unit still fits a single-example batch. engine - .evaluate(&mut module, better, Some(&[0])) + .evaluate(&target, better, Some(&[0])) .await .unwrap() .completed() @@ -430,13 +428,15 @@ async fn minibatch_gate_promotes_only_above_threshold() { let metric = HashKeyedMetric { scores: HashMap::from([(strong.stable_hash(), 0.9), (weak.stable_hash(), 0.1)]), }; - let mut engine = EvalEngine::new(trainset(6), &metric, EngineConfig::default()); + let examples = trainset(6); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig::default()); let strong = engine.register(strong); let weak = engine.register(weak); match engine - .evaluate_gated(&mut module, weak, &[0, 1], 0.5) + .evaluate_gated(&target, weak, &[0, 1], 0.5) .await .unwrap() { @@ -451,7 +451,7 @@ async fn minibatch_gate_promotes_only_above_threshold() { assert!(engine.matrix().score(weak, 2).is_none()); match engine - .evaluate_gated(&mut module, strong, &[0, 1], 0.5) + .evaluate_gated(&target, strong, &[0, 1], 0.5) .await .unwrap() { @@ -476,64 +476,6 @@ async fn minibatch_gate_promotes_only_above_threshold() { assert_eq!(engine.matrix().best_by_mean(), Some(strong)); } -#[tokio::test] -async fn checkpoint_resume_skips_completed_rollouts() { - let client = TestCompletionModel::new([ - answer_response("0"), - answer_response("1"), - answer_response("2"), - ]); - let metric = ExactMatch; - let candidate = Candidate::with_instruction("predictor", "resume-me"); - - let checkpoint = { - let mut module = lm_module(client.clone()).await; - let mut engine = EvalEngine::new(trainset(3), &metric, EngineConfig::default()); - let idx = engine.register(candidate.clone()); - let eval = engine - .evaluate(&mut module, idx, None) - .await - .unwrap() - .completed() - .unwrap(); - assert!((eval.mean() - 1.0).abs() < 1e-9); - engine.checkpoint().unwrap() - }; - - // Fresh process: new module, EMPTY response queue — any LM call errors. - let mut module = lm_module(TestCompletionModel::new([])).await; - let mut engine = - EvalEngine::resume(trainset(3), &metric, EngineConfig::default(), &checkpoint).unwrap(); - assert_eq!(engine.num_candidates(), 1); - assert_eq!(engine.spend().metric_calls, 3, "spend carries over"); - - let idx = engine.register(candidate); - assert_eq!(idx, 0, "re-registering dedups by content hash"); - - let eval = engine - .evaluate(&mut module, idx, None) - .await - .expect("resumed run must serve completed rollouts from cache") - .completed() - .unwrap(); - assert!((eval.mean() - 1.0).abs() < 1e-9); - assert!(eval.rollouts.iter().all(|r| r.trace.is_none())); - assert_eq!(engine.spend().metric_calls, 3, "no new metric calls"); - assert_eq!(engine.spend().cache_hits, 3); - assert_eq!(engine.matrix().mean(idx), Some(1.0)); - - // A checkpoint against a different example set is rejected. - match EvalEngine::<(EngSigInput, EngSigOutput), ExactMatch>::resume( - trainset(4), - &metric, - EngineConfig::default(), - &checkpoint, - ) { - Err(err) => assert!(err.to_string().contains("does not match")), - Ok(_) => panic!("mismatched examples must fail resume"), - } -} - #[tokio::test] async fn fan_out_never_exceeds_the_concurrency_bound() { const N: usize = 6; @@ -548,19 +490,17 @@ async fn fan_out_never_exceeds_the_concurrency_bound() { pair_barrier: Arc::new(tokio::sync::Barrier::new(BOUND)), }; let metric = IndexMetric; - let mut engine = EvalEngine::new( - trainset(N), - &metric, - EngineConfig { - concurrency: BOUND, - ..EngineConfig::default() - }, - ); + let examples = trainset(N); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig { + concurrency: BOUND, + ..EngineConfig::default() + }); let candidate = engine.register(Candidate::new()); let eval = tokio::time::timeout( Duration::from_secs(10), - engine.evaluate(&mut module, candidate, None), + engine.evaluate(&target, candidate, None), ) .await .expect("bounded fan-out must still overlap rollouts") @@ -582,20 +522,18 @@ async fn gate_reports_budget_exhaustion_for_minibatch_and_promotion() { // Minibatch itself doesn't fit: nothing runs, spend unchanged. let mut module = EchoModule::new(); - let mut engine = EvalEngine::new( - trainset(4), - &metric, - EngineConfig { - budget: Budget { - max_metric_calls: Some(1), - ..Budget::unlimited() - }, - ..EngineConfig::default() + let examples = trainset(4); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig { + budget: Budget { + max_metric_calls: Some(1), + ..Budget::unlimited() }, - ); + ..EngineConfig::default() + }); let candidate = engine.register(Candidate::new()); match engine - .evaluate_gated(&mut module, candidate, &[0, 1], 0.0) + .evaluate_gated(&target, candidate, &[0, 1], 0.0) .await .unwrap() { @@ -605,20 +543,16 @@ async fn gate_reports_budget_exhaustion_for_minibatch_and_promotion() { assert_eq!(engine.spend().metric_calls, 0); // Minibatch fits and passes the gate, but the full-set promotion doesn't. - let mut engine = EvalEngine::new( - trainset(4), - &metric, - EngineConfig { - budget: Budget { - max_metric_calls: Some(3), - ..Budget::unlimited() - }, - ..EngineConfig::default() + let mut engine = Engine::new(EngineConfig { + budget: Budget { + max_metric_calls: Some(3), + ..Budget::unlimited() }, - ); + ..EngineConfig::default() + }); let candidate = engine.register(Candidate::new()); match engine - .evaluate_gated(&mut module, candidate, &[2, 3], 0.0) + .evaluate_gated(&target, candidate, &[2, 3], 0.0) .await .unwrap() { @@ -635,21 +569,19 @@ async fn gate_reports_budget_exhaustion_for_minibatch_and_promotion() { async fn auxiliary_charges_count_against_the_budget() { let mut module = EchoModule::new(); let metric = IndexMetric; - let mut engine = EvalEngine::new( - trainset(3), - &metric, - EngineConfig { - budget: Budget { - max_lm_calls: Some(5), - ..Budget::unlimited() - }, - ..EngineConfig::default() + let examples = trainset(3); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig { + budget: Budget { + max_lm_calls: Some(5), + ..Budget::unlimited() }, - ); + ..EngineConfig::default() + }); let candidate = engine.register(Candidate::new()); engine - .evaluate(&mut module, candidate, None) + .evaluate(&target, candidate, None) .await .unwrap() .completed() @@ -661,66 +593,39 @@ async fn auxiliary_charges_count_against_the_budget() { engine.charge(0, 2); assert_eq!(engine.spend().lm_calls, 5); assert!(!engine.budget_allows(1)); - - // Charged spend survives checkpoint/resume. - let checkpoint = engine.checkpoint().unwrap(); - let resumed = - EvalEngine::<(EngSigInput, EngSigOutput), IndexMetric>::resume(trainset(3), &metric, *engine.config(), &checkpoint) - .unwrap(); - assert_eq!(resumed.spend().lm_calls, 5); - assert_eq!(resumed.spend().metric_calls, 3); } #[tokio::test] -async fn checkpoint_with_unknown_version_is_rejected() { - let metric = IndexMetric; - let engine = EvalEngine::new(trainset(2), &metric, EngineConfig::default()); - let checkpoint = engine.checkpoint().unwrap(); - - let mut doctored: serde_json::Value = serde_json::from_str(&checkpoint).unwrap(); - doctored["version"] = serde_json::json!(99); - let doctored = serde_json::to_string(&doctored).unwrap(); - - match EvalEngine::<(EngSigInput, EngSigOutput), IndexMetric>::resume( - trainset(2), - &metric, - EngineConfig::default(), - &doctored, - ) { - Err(err) => assert!(err.to_string().contains("version")), - Ok(_) => panic!("unknown checkpoint versions must fail resume"), - } -} - -#[tokio::test] -async fn permanent_install_invalidates_the_cache_via_baseline_hash() { +async fn install_changes_baseline_identity_and_invalidates_the_cache() { let client = TestCompletionModel::new([answer_response("0"), answer_response("1")]); let mut module = lm_module(client.clone()).await; let metric = ExactMatch; - let mut engine = EvalEngine::new(trainset(2), &metric, EngineConfig::default()); + let examples = trainset(2); + let mut engine = Engine::new(EngineConfig::default()); + let mut target = OptimizeTarget::module(&mut module, &examples, &metric); let candidate = engine.register(Candidate::new()); engine - .evaluate(&mut module, candidate, None) + .evaluate(&target, candidate, None) .await .unwrap() .completed() .unwrap(); assert_eq!(engine.spend().lm_calls, 2); - // Permanently install a winner mid-run (the COPRO-between-rounds shape): - // the module skeleton changed, so cached entries for the old baseline - // must NOT be served for the same candidate on the new baseline. - apply_candidate( - &mut module, - &Candidate::with_instruction("predictor", "installed"), - ) - .unwrap(); + // Install a winner (the run's one mutation) and start a new run: the + // module skeleton changed, so a fresh target computes a new baseline + // identity and cached entries for the old baseline must NOT be served. + target + .install(&Candidate::with_instruction("predictor", "installed")) + .unwrap(); + drop(target); + let target = OptimizeTarget::module(&mut module, &examples, &metric); client.push_response(answer_response("0")); client.push_response(answer_response("1")); let eval = engine - .evaluate(&mut module, candidate, None) + .evaluate(&target, candidate, None) .await .unwrap() .completed() @@ -731,67 +636,9 @@ async fn permanent_install_invalidates_the_cache_via_baseline_hash() { } #[tokio::test] -async fn cache_salt_partitions_the_rollout_cache() { - let client = TestCompletionModel::new([answer_response("0"), answer_response("1")]); - let metric = ExactMatch; - let candidate = Candidate::with_instruction("predictor", "salted"); - - let checkpoint = { - let mut module = lm_module(client.clone()).await; - let mut engine = EvalEngine::new(trainset(2), &metric, EngineConfig::default()); - let idx = engine.register(candidate.clone()); - engine - .evaluate(&mut module, idx, None) - .await - .unwrap() - .completed() - .unwrap(); - engine.checkpoint().unwrap() - }; - - // Same checkpoint, bumped salt (the sampling-params seam): every rollout - // is a cache miss and needs fresh LM responses. - let mut module = lm_module(client.clone()).await; - let mut engine = EvalEngine::resume( - trainset(2), - &metric, - EngineConfig { - cache_salt: 1, - ..EngineConfig::default() - }, - &checkpoint, - ) - .unwrap(); - client.push_response(answer_response("0")); - client.push_response(answer_response("1")); - let idx = engine.register(candidate.clone()); - engine - .evaluate(&mut module, idx, None) - .await - .expect("salted evaluation must run fresh rollouts") - .completed() - .unwrap(); - assert_eq!(engine.spend().cache_hits, 0, "bumped salt never hits the cache"); - - // Salt 0 again: the checkpointed entries are served with no LM calls. - let mut module = lm_module(TestCompletionModel::new([])).await; - let mut engine = - EvalEngine::resume(trainset(2), &metric, EngineConfig::default(), &checkpoint).unwrap(); - let idx = engine.register(candidate); - let eval = engine - .evaluate(&mut module, idx, None) - .await - .expect("original salt must serve from the checkpointed cache") - .completed() - .unwrap(); - assert_eq!(engine.spend().cache_hits, 2); - assert!(eval.rollouts.iter().all(|r| r.trace.is_none())); -} - -#[tokio::test] -async fn apply_and_restore_are_the_single_candidate_seam() { +async fn install_is_the_single_mutation_seam() { let mut module = EchoModule::new(); - let before = ModuleState::from_module(&mut module).unwrap(); + let before = ModuleState::from_module(&module).unwrap(); assert_eq!( before.predictors["predictor"].instruction_override.as_deref(), Some("seed") @@ -799,7 +646,7 @@ async fn apply_and_restore_are_the_single_candidate_seam() { assert!(before.predictors["predictor"].demos.is_empty()); let mut candidate = Candidate::new(); - candidate.set_instruction("predictor", "overlaid"); + candidate.set_instruction("predictor", "installed"); candidate.set_demos( "predictor", vec![ @@ -810,29 +657,70 @@ async fn apply_and_restore_are_the_single_candidate_seam() { ], ); - let undo = apply_candidate(&mut module, &candidate).unwrap(); - let applied = ModuleState::from_module(&mut module).unwrap(); + // Evaluation never mutates; install does, once, at the boundary. + let metric = IndexMetric; + let examples = trainset(1); + let mut target = OptimizeTarget::module(&mut module, &examples, &metric); + target.install(&candidate).unwrap(); + + // Unknown predictor names are rejected. + let bad = Candidate::with_instruction("missing", "nope"); + let err = target.install(&bad).expect_err("unknown predictor must fail"); + assert!(err.to_string().contains("missing")); + drop(target); + + let applied = ModuleState::from_module(&module).unwrap(); assert_eq!( applied.predictors["predictor"].instruction_override.as_deref(), - Some("overlaid") + Some("installed") ); assert_eq!(applied.predictors["predictor"].demos.len(), 1); - restore_candidate(&mut module, undo).unwrap(); - let after = ModuleState::from_module(&mut module).unwrap(); - assert_eq!( - after.predictors["predictor"].instruction_override.as_deref(), - Some("seed") - ); - assert!(after.predictors["predictor"].demos.is_empty()); + // Explicit clear-to-default is expressible and installs as a clear. + let mut clear = Candidate::new(); + clear.clear_instruction("predictor"); + clear.set_demos("predictor", Vec::new()); + let mut target = OptimizeTarget::module(&mut module, &examples, &metric); + target.install(&clear).unwrap(); + drop(target); + + let cleared = ModuleState::from_module(&module).unwrap(); + assert_eq!(cleared.predictors["predictor"].instruction_override, None); + assert!(cleared.predictors["predictor"].demos.is_empty()); +} - // Unknown predictor: error, and no partial application sticks. - let bad = Candidate::with_instruction("missing", "nope"); - let err = apply_candidate(&mut module, &bad).expect_err("unknown predictor must fail"); - assert!(err.to_string().contains("missing")); - let unchanged = ModuleState::from_module(&mut module).unwrap(); - assert_eq!( - unchanged.predictors["predictor"].instruction_override.as_deref(), - Some("seed") +#[tokio::test] +async fn candidates_fan_out_concurrently_in_the_module_lane() { + // Two DISTINCT candidates rendezvous on one barrier: the test only + // passes if their rollouts are in flight simultaneously — the module + // lane's ambient injection has no serialized apply/restore step. + const CANDIDATES: usize = 2; + let mut module = BarrierModule { + predictor: Predict::::builder().instruction("seed").build(), + barrier: Arc::new(tokio::sync::Barrier::new(CANDIDATES)), + }; + let metric = IndexMetric; + let examples = trainset(1); + let target = OptimizeTarget::module(&mut module, &examples, &metric); + let mut engine = Engine::new(EngineConfig::default()); + + let a = engine.register(Candidate::with_instruction("predictor", "A")); + let b = engine.register(Candidate::with_instruction("predictor", "B")); + + let evals = tokio::time::timeout( + Duration::from_secs(10), + engine.evaluate_many(&target, &[a, b], None), + ) + .await + .expect("candidate-level parallelism must release the barrier") + .unwrap() + .completed() + .unwrap(); + + assert_eq!(evals.len(), 2); + assert!( + engine.peak_candidate_concurrency() >= 2, + "expected candidate-level concurrency, gauge read {}", + engine.peak_candidate_concurrency() ); } diff --git a/crates/dspy-rs/tests/test_flatten_roundtrip.rs b/crates/dspy-rs/tests/test_flatten_roundtrip.rs index e9857ce9..be8a6251 100644 --- a/crates/dspy-rs/tests/test_flatten_roundtrip.rs +++ b/crates/dspy-rs/tests/test_flatten_roundtrip.rs @@ -1,4 +1,6 @@ +use dspy_rs::ir::SignatureDef; use dspy_rs::{Augmented, ChatAdapter, Demo, Message, Reasoning, Signature, WithReasoning}; +use serde_json::Value; #[derive(Signature, Clone, Debug)] struct QA { @@ -9,6 +11,13 @@ struct QA { answer: String, } +fn json_map(value: &T) -> serde_json::Map { + match serde_json::to_value(value).expect("serializable") { + Value::Object(map) => map, + other => panic!("expected object, got {other:?}"), + } +} + #[test] fn augmented_demo_roundtrips_through_adapter() { let adapter = ChatAdapter; @@ -24,7 +33,12 @@ fn augmented_demo_roundtrips_through_adapter() { }, ); - let (user_msg, assistant_msg) = adapter.format_demo_typed::>(&demo); + let def = SignatureDef::of::>(); + let types = SignatureDef::types_of::>(); + // `WithReasoning` flattens: the serialized demo output keys flat by leaf + // name (`reasoning`, `answer`) — exactly the def's canonical field names. + let user_msg = adapter.format_input_def(def, &json_map(&demo.input)); + let assistant_msg = adapter.format_output_def(def, &json_map(&demo.output)); let schema = as Signature>::schema(); let output_names: Vec<&str> = schema.output_fields().iter().map(|f| f.lm_name).collect(); @@ -33,9 +47,11 @@ fn augmented_demo_roundtrips_through_adapter() { assert!(assistant_msg.contains("answer")); let response = Message::assistant(assistant_msg); - let (parsed, _meta) = adapter - .parse_response_typed::>(&response) - .expect("typed parse should succeed"); + let (output_map, _meta) = adapter + .parse_output_def(def, types, &response) + .expect("def parse should succeed"); + let parsed: WithReasoning = + serde_json::from_value(Value::Object(output_map)).expect("typed assembly"); assert_eq!(parsed.reasoning, "Add the numbers"); assert_eq!(parsed.answer, "4"); diff --git a/crates/dspy-rs/tests/test_fx.rs b/crates/dspy-rs/tests/test_fx.rs index b5794eca..7289cf4d 100644 --- a/crates/dspy-rs/tests/test_fx.rs +++ b/crates/dspy-rs/tests/test_fx.rs @@ -98,6 +98,71 @@ async fn with_params_injects_instruction_ambiently() { ); } +#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] +#[tokio::test] +async fn with_params_drives_struct_held_predict_leaves_by_component_name() { + use dspy_rs::Predict; + + let _lock = SETTINGS_LOCK.lock().await; + let (lm, client) = make_test_lm(vec![ + response_with_fields(&[("answer", "instance")]), + response_with_fields(&[("answer", "ambient")]), + response_with_fields(&[("answer", "cleared")]), + ]) + .await; + + // A struct-held leaf with instance state, named "leaf". + let predictor = Predict::::builder() + .named("leaf") + .instruction("INSTANCE-INSTRUCTION") + .lm(lm) + .build(); + + // No params scope: instance override renders. + predictor + .call(FxQAInput { + question: "q1".to_string(), + }) + .await + .expect("instance call should succeed"); + let preamble = client.last_request().unwrap().preamble.unwrap_or_default(); + assert!(preamble.contains("INSTANCE-INSTRUCTION")); + + // Ambient params keyed by the component name win over instance state. + let mut params = fx::Params::new(); + params.set_instruction("leaf", "AMBIENT-INSTRUCTION"); + fx::with_params( + params, + predictor.call(FxQAInput { + question: "q2".to_string(), + }), + ) + .await + .expect("ambient call should succeed"); + let preamble = client.last_request().unwrap().preamble.unwrap_or_default(); + assert!( + preamble.contains("AMBIENT-INSTRUCTION") && !preamble.contains("INSTANCE-INSTRUCTION"), + "ambient candidate must win over instance state: {preamble}" + ); + + // Explicit clear resets past the instance override to the signature default. + let mut params = fx::Params::new(); + params.clear_instruction("leaf"); + fx::with_params( + params, + predictor.call(FxQAInput { + question: "q3".to_string(), + }), + ) + .await + .expect("cleared call should succeed"); + let preamble = client.last_request().unwrap().preamble.unwrap_or_default(); + assert!( + preamble.contains("Answer the question.") && !preamble.contains("INSTANCE-INSTRUCTION"), + "explicit clear must reset to the signature default: {preamble}" + ); +} + #[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] #[tokio::test] async fn capture_names_spans_after_fx_slots() { diff --git a/crates/dspy-rs/tests/test_gepa_typed_metric_feedback.rs b/crates/dspy-rs/tests/test_gepa_typed_metric_feedback.rs index 23805ad7..d813cae3 100644 --- a/crates/dspy-rs/tests/test_gepa_typed_metric_feedback.rs +++ b/crates/dspy-rs/tests/test_gepa_typed_metric_feedback.rs @@ -1,6 +1,6 @@ use anyhow::Result; use dspy_rs::{ - CallMetadata, Eval, GEPA, Module, Optimizer, Predict, PredictError, Predicted, Signature, + CallMetadata, Eval, GEPA, Module, Predict, PredictError, Predicted, Signature, TypedMetric, }; use std::sync::atomic::{AtomicUsize, Ordering}; @@ -15,12 +15,12 @@ struct OptimizerSig { answer: String, } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct InstructionEchoModule { predictor: Predict, } +dspy_rs::predictors!(InstructionEchoModule { predictor }); + impl Module for InstructionEchoModule { type Input = OptimizerSigInput; type Output = OptimizerSigOutput; @@ -200,7 +200,7 @@ async fn gepa_compile_succeeds_when_feedback_present() { .build(); let result = optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect("GEPA compile should succeed when feedback is present"); @@ -220,7 +220,7 @@ async fn gepa_compile_fails_without_feedback() { let optimizer = GEPA::builder().num_iterations(1).minibatch_size(2).build(); let err = optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect_err("GEPA should reject score-only metrics"); @@ -242,7 +242,7 @@ async fn gepa_compile_fails_when_feedback_is_partial() { let optimizer = GEPA::builder().num_iterations(1).minibatch_size(2).build(); let err = optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect_err("GEPA should reject partially-populated feedback outcomes"); @@ -272,7 +272,7 @@ async fn gepa_compile_fails_when_feedback_disappears_during_generation() { .build(); let err = optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect_err("GEPA should fail once feedback becomes unavailable mid-loop"); @@ -304,12 +304,7 @@ async fn gepa_compile_with_valset_uses_valset_and_tracks_best_outputs_when_enabl .build(); let result = optimizer - .compile_with_valset( - &mut module, - trainset(), - Some(valset.clone()), - &metric, - ) + .compile_module_with_valset(&mut module, &trainset(), Some(&valset), &metric) .await .expect("GEPA compile should succeed with a dedicated valset"); @@ -356,7 +351,7 @@ async fn gepa_compile_respects_max_lm_calls_budget() { .build(); let result = optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect("GEPA compile should succeed under LM call budget"); @@ -383,7 +378,7 @@ async fn gepa_compile_respects_max_rollouts_budget() { .build(); let result = optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect("GEPA compile should succeed under rollout budget"); @@ -411,7 +406,7 @@ async fn gepa_track_best_outputs_respects_lm_call_budget() { .build(); let result = optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect("GEPA compile should respect LM call budget when tracking outputs"); diff --git a/crates/dspy-rs/tests/test_include_program.rs b/crates/dspy-rs/tests/test_include_program.rs index 0caedad7..21d40966 100644 --- a/crates/dspy-rs/tests/test_include_program.rs +++ b/crates/dspy-rs/tests/test_include_program.rs @@ -3,7 +3,6 @@ //! resolves the runtime crate via proc-macro-crate, so this exercises the //! `FoundCrate::Itself` → `::dspy_rs` alias branch that examples/tests of the //! dspy-rs package itself hit. -#![cfg(feature = "ir")] dspy_rs::include_program!("tests/fixtures/qa.dsrs"); diff --git a/crates/dspy-rs/tests/test_input_format.rs b/crates/dspy-rs/tests/test_input_format.rs index ec51faec..c63b2c2a 100644 --- a/crates/dspy-rs/tests/test_input_format.rs +++ b/crates/dspy-rs/tests/test_input_format.rs @@ -1,7 +1,20 @@ -use dspy_rs::{BamlType, ChatAdapter, Signature}; +use dspy_rs::ir::SignatureDef; +use dspy_rs::{Schema, ChatAdapter, Signature}; +use serde_json::Value; + +fn json_map(value: &T) -> serde_json::Map { + match serde_json::to_value(value).expect("serializable") { + Value::Object(map) => map, + other => panic!("expected object, got {other:?}"), + } +} + +fn user_message(input: &S::Input) -> String { + ChatAdapter.format_input_def(SignatureDef::of::(), &json_map(input)) +} #[derive(Clone, Debug)] -#[BamlType] +#[Schema] struct Document { text: String, } @@ -158,7 +171,6 @@ fn extract_field(message: &str, field_name: &str) -> String { #[test] fn typed_input_format_yaml_falls_back_to_json() { - let adapter = ChatAdapter; let input = FormatSigInput { question: "What is YAML?".to_string(), context: vec![Document { @@ -166,7 +178,7 @@ fn typed_input_format_yaml_falls_back_to_json() { }], }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let context_value = extract_field(&message, "context"); let question_value = extract_field(&message, "question"); @@ -182,7 +194,6 @@ fn typed_input_format_yaml_falls_back_to_json() { #[test] fn typed_input_format_json_is_parsable() { - let adapter = ChatAdapter; let input = FormatJsonSigInput { question: "What is JSON?".to_string(), context: vec![Document { @@ -190,7 +201,7 @@ fn typed_input_format_json_is_parsable() { }], }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let context_value = extract_field(&message, "context"); let parsed: serde_json::Value = serde_json::from_str(&context_value).expect("valid JSON"); @@ -204,7 +215,6 @@ fn typed_input_format_json_is_parsable() { #[test] fn typed_input_format_toon_falls_back_to_json() { - let adapter = ChatAdapter; let input = FormatToonSigInput { question: "What is TOON?".to_string(), context: vec![Document { @@ -212,7 +222,7 @@ fn typed_input_format_toon_falls_back_to_json() { }], }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let context_value = extract_field(&message, "context"); let parsed: serde_json::Value = serde_json::from_str(&context_value).expect("valid JSON"); @@ -226,7 +236,6 @@ fn typed_input_format_toon_falls_back_to_json() { #[test] fn typed_input_default_string_is_raw() { - let adapter = ChatAdapter; let input = DefaultFormatSigInput { question: "Raw string".to_string(), context: vec![Document { @@ -234,7 +243,7 @@ fn typed_input_default_string_is_raw() { }], }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let question_value = extract_field(&message, "question"); assert_eq!(question_value, "Raw string"); @@ -242,7 +251,6 @@ fn typed_input_default_string_is_raw() { #[test] fn typed_input_default_non_string_is_json() { - let adapter = ChatAdapter; let input = DefaultFormatSigInput { question: "Default JSON".to_string(), context: vec![Document { @@ -250,7 +258,7 @@ fn typed_input_default_non_string_is_json() { }], }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let context_value = extract_field(&message, "context"); let parsed: serde_json::Value = serde_json::from_str(&context_value).expect("valid JSON"); let first = parsed @@ -263,7 +271,6 @@ fn typed_input_default_non_string_is_json() { #[test] fn typed_input_appends_response_instruction_reminder() { - let adapter = ChatAdapter; let input = DefaultFormatSigInput { question: "Reminder check".to_string(), context: vec![Document { @@ -271,7 +278,7 @@ fn typed_input_appends_response_instruction_reminder() { }], }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); assert!(message.contains("Respond with the corresponding output fields")); assert!(message.contains("[[ ## answer ## ]]")); assert!(message.contains("[[ ## completed ## ]]")); @@ -279,7 +286,6 @@ fn typed_input_appends_response_instruction_reminder() { #[test] fn typed_input_render_jinja_uses_context_values() { - let adapter = ChatAdapter; let input = RenderJinjaSigInput { question: "Question".to_string(), context: Document { @@ -287,7 +293,7 @@ fn typed_input_render_jinja_uses_context_values() { }, }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let context_value = extract_field(&message, "ctx"); assert_eq!( @@ -298,27 +304,25 @@ fn typed_input_render_jinja_uses_context_values() { #[test] fn typed_input_render_jinja_missing_var_panics() { - let adapter = ChatAdapter; let input = RenderJinjaStrictSigInput { question: "Question".to_string(), }; let result = std::panic::catch_unwind(|| { - adapter.format_user_message_typed::(&input) + user_message::(&input) }); assert!(result.is_err(), "missing Jinja variables should panic"); } #[test] fn typed_input_render_jinja_exposes_field_metadata_and_vars() { - let adapter = ChatAdapter; let input = RenderJinjaFieldMetaSigInput { context: Document { text: "Hello".to_string(), }, }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let context_value = extract_field(&message, "ctx"); let parts: Vec<&str> = context_value.split('|').collect(); @@ -331,13 +335,12 @@ fn typed_input_render_jinja_exposes_field_metadata_and_vars() { #[test] fn typed_input_render_jinja_non_string_primitives() { - let adapter = ChatAdapter; let input = RenderPrimitiveSigInput { count: 42, is_ready: true, }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let count_value = extract_field(&message, "count"); let ready_value = extract_field(&message, "is_ready"); @@ -347,14 +350,13 @@ fn typed_input_render_jinja_non_string_primitives() { #[test] fn typed_input_render_jinja_supports_contrib_filters() { - let adapter = ChatAdapter; let input = RenderContribFilterSigInput { context: Document { text: "abcdefg".to_string(), }, }; - let message = adapter.format_user_message_typed::(&input); + let message = user_message::(&input); let context_value = extract_field(&message, "context"); assert_eq!(context_value, "abcde"); diff --git a/crates/dspy-rs/tests/test_interp_conversation.rs b/crates/dspy-rs/tests/test_interp_conversation.rs new file mode 100644 index 00000000..2c6725f6 --- /dev/null +++ b/crates/dspy-rs/tests/test_interp_conversation.rs @@ -0,0 +1,905 @@ +//! Conversation surface (RFC 0004 §1–2): conversation-in/conversation-out +//! turns through the interpreter — chat growth and per-turn spans, opening +//! rendering parity with the map-in path, the caller-managed suspend/resume +//! loop, stop-tool and budget parity across dispatching and suspending modes, +//! and replay of recorded conversation runs. + +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use dspy_rs::ir::{ + self, Budget, BudgetPolicy, ConversationTurn, FieldType as T, Interpreter, NodeBudget, Program, + ProgramBuilder, RunError, RuntimeEnv, SignatureDef, +}; +use dspy_rs::trace::{ReplayMode, SpanEvent, capture, replay}; +use dspy_rs::{LM, LMClient, LMConfig, Message, Role, TestCompletionModel}; +use rig::completion::{AssistantContent, ToolDefinition}; +use rig::message::{Text, ToolCall, ToolFunction}; +use serde_json::json; + +// --------------------------------------------------------------------------- +// Fixtures +// --------------------------------------------------------------------------- + +fn fields(pairs: &[(&str, &str)]) -> String { + let mut out = String::new(); + for (name, value) in pairs { + out.push_str(&format!("[[ ## {name} ## ]]\n{value}\n\n")); + } + out.push_str("[[ ## completed ## ]]\n"); + out +} + +fn text(content: impl Into) -> AssistantContent { + AssistantContent::Text(Text { + text: content.into(), + }) +} + +fn tool_call(name: &str, args: serde_json::Value) -> AssistantContent { + AssistantContent::ToolCall(ToolCall::new( + format!("tc-{name}"), + ToolFunction { + name: name.to_string(), + arguments: args, + }, + )) +} + +async fn canned_lm(responses: Vec) -> (Arc, TestCompletionModel) { + let client = TestCompletionModel::new(responses); + let lm = temp_env::async_with_vars( + [("OPENAI_API_KEY", Some("test"))], + LM::builder() + .model("openai:gpt-4o-mini".to_string()) + .build(), + ) + .await + .unwrap() + .with_client(LMClient::Test(client.clone())) + .await + .unwrap(); + (Arc::new(lm), client) +} + +fn config() -> LMConfig { + LMConfig { + model: "openai:gpt-4o-mini".to_string(), + ..LMConfig::default() + } +} + +fn obj(pairs: &[(&str, serde_json::Value)]) -> serde_json::Map { + pairs + .iter() + .map(|(k, v)| (k.to_string(), v.clone())) + .collect() +} + +/// A 1-leaf `predict` program — what `Predict` compiles to. +fn qa_program() -> Program { + let mut b = ProgramBuilder::new("qa"); + b.model("m", config()); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .instruction("Answer the question.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let assistant = ir::predict("assistant", qa).bind("question", ir::input("question")); + b.main( + main_sig, + ir::seq([assistant]).out("answer", ir::out("assistant", "answer")), + ) + .unwrap() +} + +/// A tool the loop can dispatch, counting executions so caller-managed runs +/// can assert it was never invoked. +#[derive(Clone)] +struct CountingSearch { + calls: Arc, +} + +#[derive(Debug)] +struct CountingSearchError; + +impl std::fmt::Display for CountingSearchError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "counting search error") + } +} + +impl std::error::Error for CountingSearchError {} + +impl rig::tool::Tool for CountingSearch { + const NAME: &'static str = "search"; + type Error = CountingSearchError; + type Args = serde_json::Value; + type Output = String; + + async fn definition(&self, _prompt: String) -> ToolDefinition { + ToolDefinition { + name: Self::NAME.to_string(), + description: "host-side definition (ignored: the IR declares the interface)" + .to_string(), + parameters: json!({"type": "object", "additionalProperties": true}), + } + } + + async fn call(&self, args: Self::Args) -> Result { + self.calls.fetch_add(1, Ordering::SeqCst); + Ok(format!( + "results for {}: dsrs is a rust dspy", + args.get("query").and_then(|v| v.as_str()).unwrap_or("?") + )) + } +} + +/// A 1-leaf `agent` program with a `search` host tool. `budget` lands on the +/// node; `with_stop` adds a `submit` stop tool. +fn agent_program(budget: Option, with_stop: bool) -> Program { + let mut b = ProgramBuilder::new("agentic"); + b.model("m", config()); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .instruction("Research and answer.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::String) + .finish() + .unwrap(), + ); + let search = b.host_tool("search", "Web search", search_sig, &[]); + let mut tools = vec![search]; + let mut stop = Vec::new(); + if with_stop { + let submit_sig = b.sig( + SignatureDef::build("Submit") + .input("answer", T::String) + .output("ok", T::String) + .finish() + .unwrap(), + ); + let submit = b.host_tool("submit", "Submit the final answer", submit_sig, &[]); + tools.push(submit); + stop.push(submit); + } + let mut researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools(tools) + .stop_tools(stop) + .max_turns(4); + if let Some(budget) = budget { + researcher = researcher.budget(budget); + } + b.main( + main_sig, + ir::seq([researcher]).out("answer", ir::out("researcher", "answer")), + ) + .unwrap() +} + +async fn load_agent(program: Program, lm: Arc, counter: &Arc) -> Interpreter { + let mut env = RuntimeEnv::new().bind_model("m", lm).bind_host_tool( + "search", + Arc::new(CountingSearch { + calls: Arc::clone(counter), + }), + ); + if program + .tools + .iter() + .any(|(_, tool)| program.syms.get(tool.name) == "submit") + { + env = env.bind_host_tool( + "submit", + Arc::new(CountingSearch { + calls: Arc::clone(counter), + }), + ); + } + Interpreter::load(program, env).await.unwrap() +} + +// --------------------------------------------------------------------------- +// Seam 1: conversation-in/conversation-out +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn multi_turn_conversation_grows_chat_and_records_one_span_per_turn() { + let (lm, _client) = canned_lm(vec![ + text(fields(&[("answer", "42")])), + text(fields(&[("answer", "yes, exactly 42")])), + ]) + .await; + let interp = Interpreter::load(qa_program(), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let (result, trace) = capture(|| async { + // Turn 1: opening — empty chat plus the typed input. + let (run1, mut chat) = interp + .run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("what is 6*7?"))])), + None, + Budget::unlimited(), + ) + .await?; + assert_eq!(run1.output["answer"], "42"); + assert_eq!(chat.len(), 3); + assert_eq!(chat.messages[0].role, Role::System); + assert_eq!(chat.messages[1].role, Role::User); + assert_eq!(chat.messages[2].role, Role::Assistant); + + // Turn 2: continuation — the caller appends the follow-up. + chat.push_message(Message::user("are you sure?")); + let (run2, chat) = interp + .run_conversation(chat, None, None, Budget::unlimited()) + .await?; + assert_eq!(run2.output["answer"], "yes, exactly 42"); + assert_eq!(chat.len(), 5); + assert_eq!(chat.messages[4].role, Role::Assistant); + Ok::<_, RunError>(()) + }) + .await; + result.unwrap(); + + // A turn is not a run: one span per turn, seq increments per component. + assert_eq!(trace.components, vec!["assistant"]); + assert_eq!(trace.spans.len(), 2); + let opening = &trace.spans[0]; + assert_eq!(opening.seq, 0); + assert_eq!(opening.input.as_ref().unwrap()["question"], "what is 6*7?"); + assert!( + opening.prefix.is_some(), + "rendered opening prefix is interned" + ); + assert_eq!(opening.output.as_ref().unwrap()["answer"], "42"); + let continuation = &trace.spans[1]; + assert_eq!(continuation.seq, 1); + assert!( + continuation.prefix.is_none(), + "caller-owned chat has no prefix split" + ); + assert_eq!(continuation.suffix.len(), 4, "full chat recorded as suffix"); + assert!(continuation.input.is_none()); +} + +#[tokio::test] +async fn typed_continuation_appends_the_formatted_input_turn() { + let (lm, client) = canned_lm(vec![ + text(fields(&[("answer", "first")])), + text(fields(&[("answer", "second")])), + ]) + .await; + let interp = Interpreter::load(qa_program(), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let (_, mut chat) = interp + .run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("first question"))])), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + let before = chat.len(); + (_, chat) = interp + .run_conversation( + chat, + Some(obj(&[("question", json!("second question"))])), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + // input rendered as the next user turn + the assistant reply. + assert_eq!(chat.len(), before + 2); + let last = client.last_request().unwrap(); + let sent = format!("{:?}", last.chat_history); + assert!(sent.contains("second question")); +} + +#[tokio::test] +async fn conversation_opening_matches_the_map_in_rendering() { + let (lm, _client) = canned_lm(vec![ + text(fields(&[("answer", "a")])), + text(fields(&[("answer", "b")])), + ]) + .await; + let interp = Interpreter::load(qa_program(), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + let input = obj(&[("question", json!("what is 6*7?"))]); + + // Map-in evaluation and the conversation opening must hash identically — + // same rendered prompt, same model config, same replay key. + let (result, map_trace) = + capture(|| interp.run_collecting(input.clone(), None, Budget::unlimited())).await; + result.unwrap(); + let (result, conv_trace) = capture(|| { + interp.run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(input.clone()), + None, + Budget::unlimited(), + ) + }) + .await; + result.unwrap(); + assert_eq!( + map_trace.spans[0].request_hash, conv_trace.spans[0].request_hash, + "conversation opening renders byte-identically to the map-in path" + ); + + // And `conversation_opening` returns exactly the prompt the turn sent. + let opening = interp.conversation_opening(&input, None).unwrap(); + let recorded = conv_trace.prompt(&conv_trace.spans[0]); + assert_eq!(format!("{:?}", opening.messages), format!("{recorded:?}")); +} + +#[tokio::test] +async fn conversation_surface_refuses_multi_node_programs_and_empty_turns() { + // question → drafter → checker: two leaves, no conversation to own. + let mut b = ProgramBuilder::new("pipeline"); + b.model("m", config()); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let check = b.sig( + SignatureDef::build("Check") + .input("answer", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let drafter = ir::predict("drafter", qa).bind("question", ir::input("question")); + let checker = ir::predict("checker", check).bind("answer", ir::out("drafter", "answer")); + let program = b + .main( + main_sig, + ir::seq([drafter, checker]).out("verdict", ir::out("checker", "verdict")), + ) + .unwrap(); + + let (lm, _client) = canned_lm(vec![]).await; + let two_leaves = Interpreter::load(program, RuntimeEnv::new().bind_model("m", Arc::clone(&lm))) + .await + .unwrap(); + let err = two_leaves + .run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("q"))])), + None, + Budget::unlimited(), + ) + .await + .unwrap_err(); + assert!(matches!(err, RunError::Input { .. })); + + // An empty chat with no input has nothing to send. + let one_leaf = Interpreter::load(qa_program(), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + let err = one_leaf + .run_conversation( + dspy_rs::Chat::new(Vec::new()), + None, + None, + Budget::unlimited(), + ) + .await + .unwrap_err(); + assert!(matches!(err, RunError::Input { .. })); +} + +// --------------------------------------------------------------------------- +// Seam 2: caller-managed suspend/resume +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn caller_managed_turn_suspends_on_tool_calls_and_resumes_with_results() { + let (lm, _client) = canned_lm(vec![ + tool_call("search", json!({"query": "dsrs"})), + text(fields(&[("answer", "dsrs is a rust dspy")])), + ]) + .await; + let executed = Arc::new(AtomicUsize::new(0)); + let interp = load_agent(agent_program(None, false), lm, &executed).await; + + let (result, trace) = capture(|| async { + let turn = interp + .run_conversation_caller_managed( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("what is dsrs?"))])), + None, + Budget::unlimited(), + ) + .await?; + let suspension = match turn { + ConversationTurn::Suspended(suspension) => suspension, + ConversationTurn::Complete { .. } => panic!("expected a suspension"), + }; + assert_eq!(suspension.calls().len(), 1); + assert_eq!(suspension.calls()[0].function.name, "search"); + assert!( + suspension.chat().messages.last().unwrap().has_tool_calls(), + "the assistant tool-call turn is already in the conversation" + ); + + // The caller executes the tool and feeds the result back. + let turn = interp + .resume_conversation( + suspension, + vec!["caller says: dsrs is a rust dspy".to_string()], + ) + .await?; + let (run, chat) = match turn { + ConversationTurn::Complete { run, chat } => (run, chat), + ConversationTurn::Suspended(_) => panic!("expected completion"), + }; + assert_eq!(run.output["answer"], "dsrs is a rust dspy"); + assert_eq!(run.leaves.len(), 1); + assert_eq!(run.leaves[0].tool_calls.len(), 1); + assert_eq!( + run.leaves[0].tool_executions, + vec!["caller says: dsrs is a rust dspy".to_string()] + ); + // ... conversation shape: tool-result turn then the final answer. + let roles: Vec = chat.messages.iter().map(|m| m.role).collect(); + assert_eq!( + roles, + vec![ + Role::System, + Role::User, + Role::Assistant, + Role::User, + Role::Assistant + ] + ); + assert!(chat.messages[3].has_tool_results()); + Ok::<_, RunError>(()) + }) + .await; + result.unwrap(); + + // The interpreter never dispatched the tool. + assert_eq!(executed.load(Ordering::SeqCst), 0); + + // One span for the whole suspended-and-resumed turn, with the same + // exchange/tool_run/exchange stream a dispatched loop records. + assert_eq!(trace.spans.len(), 1); + let span = &trace.spans[0]; + let kinds: Vec<&str> = span + .events + .iter() + .map(|event| match event { + SpanEvent::Exchange { .. } => "exchange", + SpanEvent::ToolRun { .. } => "tool_run", + _ => "other", + }) + .collect(); + assert_eq!(kinds, vec!["exchange", "tool_run", "exchange"]); + match &span.events[1] { + SpanEvent::ToolRun { + name, + result, + error, + .. + } => { + assert_eq!(name, "search"); + assert_eq!(result, "caller says: dsrs is a rust dspy"); + assert!(error.is_none()); + } + other => panic!("expected ToolRun, got {other:?}"), + } + assert_eq!( + span.output.as_ref().unwrap()["answer"], + "dsrs is a rust dspy" + ); +} + +#[tokio::test] +async fn dropping_a_suspension_cancels_its_span() { + let (lm, _client) = canned_lm(vec![tool_call("search", json!({"query": "dsrs"}))]).await; + let executed = Arc::new(AtomicUsize::new(0)); + let interp = load_agent(agent_program(None, false), lm, &executed).await; + + let (_, trace) = capture(|| async { + let turn = interp + .run_conversation_caller_managed( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("q"))])), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + assert!(matches!(turn, ConversationTurn::Suspended(_))); + drop(turn); + }) + .await; + + assert_eq!(trace.spans.len(), 1); + let error = trace.spans[0].error.as_ref().expect("span closed as error"); + assert_eq!(error.kind, dspy_rs::trace::SpanErrorKind::Cancelled); +} + +#[tokio::test] +async fn suspending_and_dispatching_modes_record_identical_spans() { + let question = obj(&[("question", json!("what is dsrs?"))]); + + // Dispatching mode: the bound tool executes. + let (lm, _client) = canned_lm(vec![ + tool_call("search", json!({"query": "dsrs"})), + text(fields(&[("answer", "dsrs is a rust dspy")])), + ]) + .await; + let executed = Arc::new(AtomicUsize::new(0)); + let interp = load_agent(agent_program(None, false), lm, &executed).await; + let (result, dispatched) = capture(|| { + interp.run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(question.clone()), + None, + Budget::unlimited(), + ) + }) + .await; + let (dispatched_run, dispatched_chat) = result.unwrap(); + assert_eq!(executed.load(Ordering::SeqCst), 1); + let dispatched_result = match &dispatched.spans[0].events[1] { + SpanEvent::ToolRun { result, .. } => result.clone(), + other => panic!("expected ToolRun, got {other:?}"), + }; + + // Suspending mode over the same exchanges, feeding the exact result the + // dispatched tool produced. + let (lm, _client) = canned_lm(vec![ + tool_call("search", json!({"query": "dsrs"})), + text(fields(&[("answer", "dsrs is a rust dspy")])), + ]) + .await; + let executed = Arc::new(AtomicUsize::new(0)); + let interp = load_agent(agent_program(None, false), lm, &executed).await; + let (result, suspended) = capture(|| async { + let turn = interp + .run_conversation_caller_managed( + dspy_rs::Chat::new(Vec::new()), + Some(question.clone()), + None, + Budget::unlimited(), + ) + .await?; + let ConversationTurn::Suspended(suspension) = turn else { + panic!("expected a suspension"); + }; + interp + .resume_conversation(suspension, vec![dispatched_result.clone()]) + .await + }) + .await; + let ConversationTurn::Complete { run, chat } = result.unwrap() else { + panic!("expected completion"); + }; + assert_eq!(executed.load(Ordering::SeqCst), 0); + + // Same output, same conversation, same span identity and event stream. + assert_eq!(run.output, dispatched_run.output); + assert_eq!( + format!("{:?}", chat.messages), + format!("{:?}", dispatched_chat.messages) + ); + let a = &dispatched.spans[0]; + let b = &suspended.spans[0]; + assert_eq!(a.request_hash, b.request_hash); + assert_eq!(a.usage.total_tokens, b.usage.total_tokens); + assert_eq!(a.usage.prompt_tokens, b.usage.prompt_tokens); + assert_eq!(a.usage.completion_tokens, b.usage.completion_tokens); + assert_eq!(a.events.len(), b.events.len()); + for (left, right) in a.events.iter().zip(b.events.iter()) { + match (left, right) { + (SpanEvent::Exchange { message: l, .. }, SpanEvent::Exchange { message: r, .. }) => { + assert_eq!(format!("{l:?}"), format!("{r:?}")) + } + ( + SpanEvent::ToolRun { + name: ln, + result: lr, + error: le, + .. + }, + SpanEvent::ToolRun { + name: rn, + result: rr, + error: re, + .. + }, + ) => { + assert_eq!(ln, rn); + assert_eq!(lr, rr); + assert_eq!(le, re); + } + (left, right) => panic!("event streams diverge: {left:?} vs {right:?}"), + } + } +} + +#[tokio::test] +async fn stop_tool_completes_the_turn_in_both_modes_without_suspending() { + let stop_call = || tool_call("submit", json!({"answer": "42"})); + + let (lm, _client) = canned_lm(vec![stop_call()]).await; + let executed = Arc::new(AtomicUsize::new(0)); + let interp = load_agent(agent_program(None, true), lm, &executed).await; + let (dispatched, _) = capture(|| { + interp.run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("q"))])), + None, + Budget::unlimited(), + ) + }) + .await; + let (run, _) = dispatched.unwrap(); + assert_eq!(run.output["answer"], "42"); + + let (lm, _client) = canned_lm(vec![stop_call()]).await; + let interp = load_agent(agent_program(None, true), lm, &executed).await; + let (turn, trace) = capture(|| { + interp.run_conversation_caller_managed( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("q"))])), + None, + Budget::unlimited(), + ) + }) + .await; + let ConversationTurn::Complete { run, .. } = turn.unwrap() else { + panic!("a stop-tool call completes the turn — it never suspends"); + }; + assert_eq!(run.output["answer"], "42"); + // Stop tools are never executed, in either mode. + assert_eq!(executed.load(Ordering::SeqCst), 0); + assert!( + trace.spans[0] + .events + .iter() + .all(|event| matches!(event, SpanEvent::Exchange { .. })), + "no ToolRun events for a stop call" + ); +} + +#[tokio::test] +async fn budget_exhaustion_is_identical_across_modes() { + let budget = NodeBudget { + max_lm_calls: Some(1), + max_tokens: None, + deadline_ms: None, + on_exhausted: BudgetPolicy::Fail, + }; + + // Dispatching: the tool executes, then the next loop turn is refused. + let (lm, _client) = canned_lm(vec![tool_call("search", json!({"query": "dsrs"}))]).await; + let executed = Arc::new(AtomicUsize::new(0)); + let interp = load_agent(agent_program(Some(budget.clone()), false), lm, &executed).await; + let (result, dispatched) = capture(|| { + interp.run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("q"))])), + None, + Budget::unlimited(), + ) + }) + .await; + let dispatched_err = result.unwrap_err(); + assert!(matches!(dispatched_err, RunError::Budget { .. })); + + // Suspending: the resume hits the same refusal at the same point. + let (lm, _client) = canned_lm(vec![tool_call("search", json!({"query": "dsrs"}))]).await; + let interp = load_agent(agent_program(Some(budget), false), lm, &executed).await; + let (result, suspended) = capture(|| async { + let turn = interp + .run_conversation_caller_managed( + dspy_rs::Chat::new(Vec::new()), + Some(obj(&[("question", json!("q"))])), + None, + Budget::unlimited(), + ) + .await?; + let ConversationTurn::Suspended(suspension) = turn else { + panic!("expected a suspension"); + }; + interp + .resume_conversation(suspension, vec!["result".to_string()]) + .await + }) + .await; + let suspended_err = result.unwrap_err(); + assert!(matches!(suspended_err, RunError::Budget { .. })); + assert_eq!(dispatched_err.to_string(), suspended_err.to_string()); + + // Both spans closed as the same error with one recorded exchange. + let a = &dispatched.spans[0]; + let b = &suspended.spans[0]; + assert_eq!( + a.error.as_ref().map(|e| (e.kind, e.message.clone())), + b.error.as_ref().map(|e| (e.kind, e.message.clone())), + ); + assert_eq!( + a.events + .iter() + .filter(|e| matches!(e, SpanEvent::Exchange { .. })) + .count(), + b.events + .iter() + .filter(|e| matches!(e, SpanEvent::Exchange { .. })) + .count(), + ); +} + +// --------------------------------------------------------------------------- +// Replay of conversation runs +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn conversation_runs_replay_turn_by_turn() { + let input = obj(&[("question", json!("what is 6*7?"))]); + + // Record a 2-turn conversation live. + let (lm, _client) = canned_lm(vec![ + text(fields(&[("answer", "42")])), + text(fields(&[("answer", "yes, exactly 42")])), + ]) + .await; + let interp = Interpreter::load(qa_program(), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + let (result, recording) = capture(|| async { + let (run1, mut chat) = interp + .run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(input.clone()), + None, + Budget::unlimited(), + ) + .await?; + chat.push_message(Message::user("are you sure?")); + let (run2, chat) = interp + .run_conversation(chat, None, None, Budget::unlimited()) + .await?; + Ok::<_, RunError>((run1.output, run2.output, chat)) + }) + .await; + let (live1, live2, live_chat) = result.unwrap(); + + // Replay strictly against an LM with no responses left: every turn must + // be served from the recording, and the chats must rebuild identically. + let (empty_lm, _client) = canned_lm(vec![]).await; + let replayed = Interpreter::load(qa_program(), RuntimeEnv::new().bind_model("m", empty_lm)) + .await + .unwrap(); + let (result, report) = replay(&recording, ReplayMode::Strict, || async { + let (run1, mut chat) = replayed + .run_conversation( + dspy_rs::Chat::new(Vec::new()), + Some(input.clone()), + None, + Budget::unlimited(), + ) + .await?; + chat.push_message(Message::user("are you sure?")); + let (run2, chat) = replayed + .run_conversation(chat, None, None, Budget::unlimited()) + .await?; + Ok::<_, RunError>((run1.output, run2.output, chat)) + }) + .await; + let (served1, served2, served_chat) = result.unwrap(); + + assert_eq!(report.served, 2); + assert_eq!(report.live, 0); + assert_eq!(served1, live1); + assert_eq!(served2, live2); + assert_eq!( + format!("{:?}", served_chat.messages), + format!("{:?}", live_chat.messages) + ); +} + +#[tokio::test] +async fn caller_managed_turns_never_suspend_under_replay() { + let input = obj(&[("question", json!("what is dsrs?"))]); + + // Record a suspended-and-resumed turn live. + let (lm, _client) = canned_lm(vec![ + tool_call("search", json!({"query": "dsrs"})), + text(fields(&[("answer", "dsrs is a rust dspy")])), + ]) + .await; + let executed = Arc::new(AtomicUsize::new(0)); + let interp = load_agent(agent_program(None, false), lm, &executed).await; + let (result, recording) = capture(|| async { + let turn = interp + .run_conversation_caller_managed( + dspy_rs::Chat::new(Vec::new()), + Some(input.clone()), + None, + Budget::unlimited(), + ) + .await?; + let ConversationTurn::Suspended(suspension) = turn else { + panic!("expected a suspension"); + }; + interp + .resume_conversation(suspension, vec!["tool says 42".to_string()]) + .await + }) + .await; + let ConversationTurn::Complete { run: live_run, .. } = result.unwrap() else { + panic!("expected completion"); + }; + + // Under replay the whole turn is served — tool effects are baked into the + // recorded span, so the caller never sees a suspension. + let (empty_lm, _client) = canned_lm(vec![]).await; + let replayed = load_agent(agent_program(None, false), empty_lm, &executed).await; + let (turn, report) = replay(&recording, ReplayMode::Strict, || { + replayed.run_conversation_caller_managed( + dspy_rs::Chat::new(Vec::new()), + Some(input.clone()), + None, + Budget::unlimited(), + ) + }) + .await; + let ConversationTurn::Complete { run, chat } = turn.unwrap() else { + panic!("served turns never suspend"); + }; + assert_eq!(report.served, 1); + assert_eq!(run.output, live_run.output); + assert_eq!( + run.leaves[0].tool_executions, + vec!["tool says 42".to_string()] + ); + // The served chat carries the recorded tool-call turn and result. + assert!(chat.messages.iter().any(|m| m.has_tool_calls())); + assert!(chat.messages.iter().any(|m| m.has_tool_results())); +} diff --git a/crates/dspy-rs/tests/test_interp_metadata.rs b/crates/dspy-rs/tests/test_interp_metadata.rs new file mode 100644 index 00000000..fab26b77 --- /dev/null +++ b/crates/dspy-rs/tests/test_interp_metadata.rs @@ -0,0 +1,376 @@ +//! Interpreter per-leaf metadata seam (`Interpreter::run_collecting`): +//! coercion flags, constraint outcomes, raw response text, usage, and model +//! config hash surface per `Predict` leaf, in execution order — the exact +//! parity data the historical static lane kept, so a typed +//! `Predict` routed through the interpreter loses none of `Predicted`'s +//! metadata contract. + +use std::sync::Arc; + +use dspy_rs::Flag; +use dspy_rs::ir::{ + self, Budget, ConstraintDef, FieldDef, FieldType as T, Interpreter, Program, ProgramBuilder, + RunError, RuntimeEnv, SignatureDef, +}; +use dspy_rs::trace::ModelEntry; +use dspy_rs::{LM, LMClient, LMConfig, TestCompletionModel}; +use rig::completion::{AssistantContent, Usage}; +use rig::message::Text; +use serde_json::json; + +// --------------------------------------------------------------------------- +// Fixtures (mirrors test_ir_interp.rs) +// --------------------------------------------------------------------------- + +fn fields(pairs: &[(&str, &str)]) -> String { + let mut out = String::new(); + for (name, value) in pairs { + out.push_str(&format!("[[ ## {name} ## ]]\n{value}\n\n")); + } + out.push_str("[[ ## completed ## ]]\n"); + out +} + +fn text(content: impl Into) -> AssistantContent { + AssistantContent::Text(Text { + text: content.into(), + }) +} + +async fn canned_lm(responses: Vec) -> (Arc, TestCompletionModel) { + let client = TestCompletionModel::new(responses); + let lm = temp_env::async_with_vars( + [("OPENAI_API_KEY", Some("test"))], + LM::builder() + .model("openai:gpt-4o-mini".to_string()) + .build(), + ) + .await + .unwrap() + .with_client(LMClient::Test(client.clone())) + .await + .unwrap(); + (Arc::new(lm), client) +} + +fn config() -> LMConfig { + LMConfig { + model: "openai:gpt-4o-mini".to_string(), + ..LMConfig::default() + } +} + +fn obj(pairs: &[(&str, serde_json::Value)]) -> serde_json::Map { + pairs + .iter() + .map(|(k, v)| (k.to_string(), v.clone())) + .collect() +} + +/// question → rater → rating (Int, with the given constraints). +fn rater_program(constraints: Vec) -> Program { + let mut b = ProgramBuilder::new("rater-pipeline"); + b.model("m", config()); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("rating", T::Int) + .finish() + .unwrap(), + ); + let mut rating_field = FieldDef::new("rating", T::Int); + for constraint in constraints { + rating_field = rating_field.with_constraint(constraint); + } + let rate = b.sig( + SignatureDef::build("Rate") + .instruction("Rate the thing.") + .input("question", T::String) + .output_full(rating_field) + .finish() + .unwrap(), + ); + let rater = ir::predict("rater", rate).bind("question", ir::input("question")); + b.main( + main_sig, + ir::seq([rater]).out("rating", ir::out("rater", "rating")), + ) + .unwrap() +} + +/// question → drafter (QA) → checker (Check) → verdict. Two Predict leaves. +fn seq_program() -> Program { + let mut b = ProgramBuilder::new("pipeline"); + b.model("m", config()); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .instruction("Answer the question.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let check = b.sig( + SignatureDef::build("Check") + .instruction("Judge the answer.") + .input("answer", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let drafter = ir::predict("drafter", qa).bind("question", ir::input("question")); + let checker = ir::predict("checker", check).bind("answer", ir::out("drafter", "answer")); + b.main( + main_sig, + ir::seq([drafter, checker]).out("verdict", ir::out("checker", "verdict")), + ) + .unwrap() +} + +// --------------------------------------------------------------------------- +// Coercion flags + raw response + model config hash +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn coerced_field_surfaces_flags_raw_text_and_model_hash() { + // "1,000" needs thousands-separator coercion into Int → CoercedFromString. + let raw_response = fields(&[("rating", "1,000")]); + let (lm, _client) = canned_lm(vec![text(raw_response.clone())]).await; + let expected_hash = ModelEntry::from_config(&lm.config).config_hash; + let interp = Interpreter::load(rater_program(vec![]), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let run = interp + .run_collecting( + obj(&[("question", json!("how many?"))]), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + + assert_eq!(run.output["rating"], 1000); + assert_eq!(run.leaves.len(), 1); + + let leaf = &run.leaves[0]; + assert_eq!(leaf.name, "rater"); + assert_eq!(leaf.raw_response, raw_response); + assert_eq!(leaf.model_config_hash, expected_hash); + assert_ne!(leaf.model_config_hash, 0); + + let meta = &leaf.field_meta["rating"]; + assert_eq!(meta.raw_text, "1,000"); + assert_eq!(meta.flags, vec![Flag::CoercedFromString]); + assert!(meta.checks.is_empty()); +} + +#[tokio::test] +async fn clean_field_reports_no_flags() { + let (lm, _client) = canned_lm(vec![text(fields(&[("rating", "7")]))]).await; + let interp = Interpreter::load(rater_program(vec![]), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let run = interp + .run_collecting( + obj(&[("question", json!("how many?"))]), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + + assert_eq!(run.output["rating"], 7); + let meta = &run.leaves[0].field_meta["rating"]; + assert_eq!(meta.raw_text, "7"); + assert!(meta.flags.is_empty()); +} + +// --------------------------------------------------------------------------- +// Constraint outcomes +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn check_outcomes_surface_pass_and_fail() { + let program = rater_program(vec![ + ConstraintDef::check("positive", "this > 0"), + ConstraintDef::check("small", "this < 5"), + ]); + let (lm, _client) = canned_lm(vec![text(fields(&[("rating", "7")]))]).await; + let interp = Interpreter::load(program, RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let run = interp + .run_collecting( + obj(&[("question", json!("how many?"))]), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + + // Failed #[check]s never abort the run — they surface as outcomes. + assert_eq!(run.output["rating"], 7); + let checks = &run.leaves[0].field_meta["rating"].checks; + assert_eq!(checks.len(), 2); + assert_eq!(checks[0].label, "positive"); + assert_eq!(checks[0].expression, "this > 0"); + assert!(checks[0].passed); + assert_eq!(checks[1].label, "small"); + assert_eq!(checks[1].expression, "this < 5"); + assert!(!checks[1].passed); +} + +#[tokio::test] +async fn assert_failure_is_a_parse_error() { + // Same semantics as the static lane: a failed #[assert] is a parse error + // (no LeafOutcome — the evaluation did not succeed), while a passing + // assert records no ConstraintResult. + let program = rater_program(vec![ConstraintDef::assert("this < 5")]); + let (lm, _client) = canned_lm(vec![text(fields(&[("rating", "7")]))]).await; + let interp = Interpreter::load(program, RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let err = interp + .run_collecting( + obj(&[("question", json!("how many?"))]), + None, + Budget::unlimited(), + ) + .await + .expect_err("failed assert must fail the run"); + match err { + RunError::Parse { at, .. } => assert_eq!(&*at, "rater"), + other => panic!("expected parse error, got {other:?}"), + } +} + +#[tokio::test] +async fn passing_assert_records_no_constraint_result() { + let program = rater_program(vec![ + ConstraintDef::assert("this > 0"), + ConstraintDef::check("small", "this < 5"), + ]); + let (lm, _client) = canned_lm(vec![text(fields(&[("rating", "3")]))]).await; + let interp = Interpreter::load(program, RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let run = interp + .run_collecting( + obj(&[("question", json!("how many?"))]), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + + // Only #[check] constraints produce ConstraintResults — asserts are + // pass-or-error, exactly like the historical typed parse path. + let checks = &run.leaves[0].field_meta["rating"].checks; + assert_eq!(checks.len(), 1); + assert_eq!(checks[0].label, "small"); + assert!(checks[0].passed); +} + +// --------------------------------------------------------------------------- +// Multi-leaf execution order + usage accumulation +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn two_leaf_seq_reports_outcomes_in_execution_order() { + let drafter_raw = fields(&[("answer", "42")]); + let checker_raw = fields(&[("verdict", "correct")]); + let (lm, client) = canned_lm(vec![text(drafter_raw.clone()), text(checker_raw.clone())]).await; + let mut usage = Usage::new(); + usage.input_tokens = 3; + usage.output_tokens = 4; + usage.total_tokens = 7; + client.set_usage(usage); + + let interp = Interpreter::load(seq_program(), RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap(); + + let run = interp + .run_collecting( + obj(&[("question", json!("what is 6*7?"))]), + None, + Budget::unlimited(), + ) + .await + .unwrap(); + + assert_eq!(run.output["verdict"], "correct"); + + // Two leaves, in execution order, each with its own raw response. + assert_eq!(run.leaves.len(), 2); + assert_eq!(run.leaves[0].name, "drafter"); + assert_eq!(run.leaves[0].raw_response, drafter_raw); + assert_eq!(run.leaves[0].field_meta["answer"].raw_text, "42"); + assert_eq!(run.leaves[1].name, "checker"); + assert_eq!(run.leaves[1].raw_response, checker_raw); + assert_eq!(run.leaves[1].field_meta["verdict"].raw_text, "correct"); + + // Both leaves used the same model config. + assert_eq!( + run.leaves[0].model_config_hash, + run.leaves[1].model_config_hash + ); + + // Per-leaf usage is each call's own; the run total accumulates across leaves. + for leaf in &run.leaves { + assert_eq!(leaf.usage.prompt_tokens, 3); + assert_eq!(leaf.usage.completion_tokens, 4); + assert_eq!(leaf.usage.total_tokens, 7); + } + let total: u64 = run.leaves.iter().map(|leaf| leaf.usage.total_tokens).sum(); + assert_eq!(total, 14); +} + +// --------------------------------------------------------------------------- +// `run` is unchanged: same output, no collection +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn run_output_matches_run_collecting_output() { + let (lm_a, _c1) = canned_lm(vec![ + text(fields(&[("answer", "42")])), + text(fields(&[("verdict", "correct")])), + ]) + .await; + let (lm_b, _c2) = canned_lm(vec![ + text(fields(&[("answer", "42")])), + text(fields(&[("verdict", "correct")])), + ]) + .await; + + let plain = Interpreter::load(seq_program(), RuntimeEnv::new().bind_model("m", lm_a)) + .await + .unwrap(); + let collecting = Interpreter::load(seq_program(), RuntimeEnv::new().bind_model("m", lm_b)) + .await + .unwrap(); + + let input = obj(&[("question", json!("what is 6*7?"))]); + let from_run = plain + .run(input.clone(), None, Budget::unlimited()) + .await + .unwrap(); + let from_collecting = collecting + .run_collecting(input, None, Budget::unlimited()) + .await + .unwrap(); + + assert_eq!(from_run, from_collecting.output); +} diff --git a/crates/dspy-rs/tests/test_ir_bake.rs b/crates/dspy-rs/tests/test_ir_bake.rs index f70665ad..6a1e55e5 100644 --- a/crates/dspy-rs/tests/test_ir_bake.rs +++ b/crates/dspy-rs/tests/test_ir_bake.rs @@ -1,7 +1,6 @@ //! IR-6 (RFC 0002 §5): `Program::bake` — folding an overlay into a new //! program value, lineage stamping, hash recompute, and behavioral equality //! of base+overlay vs. baked on canned LMs. -#![cfg(feature = "ir")] use std::sync::Arc; @@ -214,10 +213,18 @@ fn bake_folds_all_slot_kinds_into_defaults() { .unwrap(), ); let search = b.host_tool("search", "old tool desc", search_sig, &["net:search"]); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expression", T::String) + .output("value", T::Float) + .finish() + .unwrap(), + ); + let calc = b.host_tool("calc", "Evaluate arithmetic", calc_sig, &[]); let researcher = ir::agent("researcher", main_sig) .model(m1) .bind("question", ir::input("question")) - .tools([search]); + .tools([search, calc]); let shouter = ir::hole("shouter", shout_sig, "(a) => ({shout: a.question})", &[]) .bind("question", ir::input("question")); let program = b @@ -245,6 +252,13 @@ fn bake_folds_all_slot_kinds_into_defaults() { let code_id = program.param_id("shouter.code").unwrap(); let new_code = ParamValue::code(CodeLang::Js, "(a) => ({shout: a.question.toUpperCase()})"); overlay.set(&program, code_id, new_code.clone()).unwrap(); + let tool_set_id = program.param_id("researcher.tool_set").unwrap(); + let restricted = ParamValue::ToolSet { + tools: vec![search], + }; + overlay + .set(&program, tool_set_id, restricted.clone()) + .unwrap(); let baked = program.bake(&overlay, note()).unwrap(); @@ -259,6 +273,13 @@ fn bake_folds_all_slot_kinds_into_defaults() { } ); assert_eq!(baked.params[code_id].default, new_code); + assert_eq!(baked.params[tool_set_id].default, restricted); + // The restricted selection is a first-class artifact: it prints as a + // `tool_set` line and the baked program reloads identically. + let text = baked.to_dsrs(); + assert!(text.contains("tool_set [search]"), "{text}"); + let reloaded = Program::from_dsrs(&text).unwrap(); + assert_eq!(reloaded.meta.program_hash, baked.meta.program_hash); // The base program is untouched (bake is pure). assert_eq!( diff --git a/crates/dspy-rs/tests/test_ir_bridge.rs b/crates/dspy-rs/tests/test_ir_bridge.rs index 794a9525..f1da1421 100644 --- a/crates/dspy-rs/tests/test_ir_bridge.rs +++ b/crates/dspy-rs/tests/test_ir_bridge.rs @@ -1,7 +1,6 @@ //! IR-2 bridge leftovers (RFC 0002 §2.4 migration contract): `fx::Params` ↔ //! `Overlay` bind/unbind, the `with_overlay` fx-lane scope, and the //! `ModuleState` ↔ `Overlay` serde projection. -#![cfg(feature = "ir")] use dspy_rs::ir::{ self, FieldType as T, Overlay, OverlayError, ParamValue, Program, ProgramBuilder, SignatureDef, diff --git a/crates/dspy-rs/tests/test_ir_code_mode.rs b/crates/dspy-rs/tests/test_ir_code_mode.rs index 478eae6e..f3128bf3 100644 --- a/crates/dspy-rs/tests/test_ir_code_mode.rs +++ b/crates/dspy-rs/tests/test_ir_code_mode.rs @@ -2,7 +2,6 @@ //! present its non-stop tools as one sandboxed `run_js` tool. Same canned //! flow as the module lane, run through the interpreter — spans record the //! `ToolRun`, stop tools stay individual, collisions refuse the load. -#![cfg(all(feature = "ir", feature = "code-mode"))] use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; diff --git a/crates/dspy-rs/tests/test_ir_dynamic_render.rs b/crates/dspy-rs/tests/test_ir_dynamic_render.rs index 4e9fbd5e..60c6464c 100644 --- a/crates/dspy-rs/tests/test_ir_dynamic_render.rs +++ b/crates/dspy-rs/tests/test_ir_dynamic_render.rs @@ -1,8 +1,12 @@ //! IR-1 (RFC 0002 §1) dynamic rendering proof. //! //! Two claims under test: -//! 1. **Lane parity** — a `SignatureDef` bridged from a derive renders byte-identical -//! prompt sections (system, input, parse) to the static `SignatureSchema` path. +//! 1. **Derive-bridged rendering** — a `SignatureDef` bridged from a derive +//! ([`SignatureDef::of`]) renders the full prompt protocol: aliases, docs, +//! `#[format]`/`#[render]` hints, and constraint metadata. (The historical +//! static `SignatureSchema` render lane was collapsed when `Predict` moved +//! onto the interpreter; byte-stability is pinned by the golden prompt +//! tests.) //! 2. **No `'static` anywhere** — a `SignatureDef` constructed at runtime, never //! mentioned in any derive, formats a system prompt + input and round-trips a //! canned LM response through parse, with all data owned and dropped after use. @@ -10,7 +14,7 @@ use dspy_rs::ChatAdapter; use dspy_rs::ir::{ConstraintDef, FieldDef, FieldType, RenderSpec, SignatureDef, TypeTable}; use dspy_rs::typesys::{ClassDef, EnumDef, EnumValueDef, FieldDef as TypeFieldDef}; -use dspy_rs::{BamlType, Message, ParseError, Signature}; +use dspy_rs::{Schema, Message, ParseError, Signature}; #[derive(Signature, Clone, Debug)] /// Grade an answer against a question. @@ -34,14 +38,14 @@ struct Graded { } #[derive(Clone, Debug)] -#[BamlType] +#[Schema] struct Citation { url: String, title: String, } #[derive(Clone, Debug)] -#[BamlType] +#[Schema] enum Stance { Support, Refute, @@ -75,79 +79,78 @@ fn graded_response() -> Message { } #[test] -fn system_prompt_parity_between_lanes() { +fn derive_bridged_system_prompt_renders_aliases_docs_and_types() { let adapter = ChatAdapter; - let static_lane = adapter.build_system(Graded::schema(), None).unwrap(); - let value_lane = adapter.build_system_def( + let system = adapter.build_system_def( SignatureDef::of::(), SignatureDef::types_of::(), None, ); - assert_eq!(static_lane, value_lane); + assert!(system.contains("`question` (string): The question being graded.")); + assert!(system.contains("[[ ## score ## ]]")); + assert!(!system.contains("[[ ## confidence ## ]]")); + assert!(system.contains("Output field `score` should be of type: float")); + assert!(system.contains("Grade an answer against a question.")); - let static_lane = adapter.build_system(Structured::schema(), None).unwrap(); - let value_lane = adapter.build_system_def( + let system = adapter.build_system_def( SignatureDef::of::(), SignatureDef::types_of::(), None, ); - assert_eq!(static_lane, value_lane); + assert!(system.contains("[[ ## citations ## ]]")); + assert!(system.contains("url: string,")); + assert!(system.contains("- Support")); let with_override = "Grade strictly."; - let static_lane = adapter - .build_system(Graded::schema(), Some(with_override)) - .unwrap(); - let value_lane = adapter.build_system_def( + let system = adapter.build_system_def( SignatureDef::of::(), SignatureDef::types_of::(), Some(with_override), ); - assert_eq!(static_lane, value_lane); + assert!(system.contains("Grade strictly.")); + assert!(!system.contains("Grade an answer against a question.")); } #[test] -fn input_format_parity_between_lanes() { +fn derive_bridged_input_honors_format_hints() { let adapter = ChatAdapter; let typed = GradedInput::new( "Is the sky blue?".to_string(), vec!["observation log".to_string()], ); - let static_lane = adapter.format_input(Graded::schema(), &typed); - let value_input = serde_json::to_value(&typed) .unwrap() .as_object() .cloned() .unwrap(); let value_lane = adapter.format_input_def(SignatureDef::of::(), &value_input); - assert_eq!(static_lane, value_lane); + assert!(value_lane.contains("[[ ## question ## ]]\nIs the sky blue?")); + // `#[format("json")]` renders the list as JSON. + assert!(value_lane.contains(r#"["observation log"]"#)); + assert!(value_lane.contains("starting with the field `[[ ## score ## ]]`")); } #[test] -fn jinja_input_parity_between_lanes() { +fn derive_bridged_jinja_input_renders_template() { let adapter = ChatAdapter; let typed = JinjaSigInput::new("What is 2+2?".to_string()); - let static_lane = adapter.format_input(JinjaSig::schema(), &typed); - let value_input = serde_json::to_value(&typed) .unwrap() .as_object() .cloned() .unwrap(); let value_lane = adapter.format_input_def(SignatureDef::of::(), &value_input); - assert_eq!(static_lane, value_lane); assert!(value_lane.contains("Q: What is 2+2? [What is 2+2?]")); } #[test] -fn parse_parity_between_lanes() { +fn derive_bridged_parse_assembles_typed_output_and_checks() { let adapter = ChatAdapter; let response = graded_response(); - let (typed, typed_meta) = adapter.parse_response_typed::(&response).unwrap(); let (value_map, value_meta) = adapter .parse_output_def( SignatureDef::of::(), @@ -156,21 +159,18 @@ fn parse_parity_between_lanes() { ) .unwrap(); - assert_eq!( - serde_json::to_value(&typed).unwrap(), - serde_json::Value::Object(value_map) - ); - - // Same per-field metadata: raw text and check outcomes line up. - assert_eq!( - typed_meta.get("confidence").unwrap().raw_text, - value_meta.get("confidence").unwrap().raw_text - ); - let static_checks = &typed_meta.get("confidence").unwrap().checks; - let value_checks = &value_meta.get("confidence").unwrap().checks; - assert_eq!(static_checks.len(), value_checks.len()); - assert_eq!(static_checks[0].label, value_checks[0].label); - assert_eq!(static_checks[0].passed, value_checks[0].passed); + // Canonical field-name keying assembles straight into the typed output. + let typed: GradedOutput = + serde_json::from_value(serde_json::Value::Object(value_map)).unwrap(); + assert_eq!(typed.confidence, 0.9); + assert_eq!(typed.verdict, "Supported"); + + // Per-field metadata: raw text and `#[check]` outcomes. + let confidence = value_meta.get("confidence").unwrap(); + assert_eq!(confidence.raw_text, "0.9"); + assert_eq!(confidence.checks.len(), 1); + assert_eq!(confidence.checks[0].label, "range"); + assert!(confidence.checks[0].passed); } /// A signature that exists nowhere as a type: built at runtime from owned diff --git a/crates/dspy-rs/tests/test_ir_edit.rs b/crates/dspy-rs/tests/test_ir_edit.rs new file mode 100644 index 00000000..0ec84c0a --- /dev/null +++ b/crates/dspy-rs/tests/test_ir_edit.rs @@ -0,0 +1,1333 @@ +//! The graph-edit calculus (`ir::edit`): typed structural edits over the IR — +//! application, apply/validate failure modes, garbage collection, the +//! `legal_edits` menu, and overlay migration across structural change. + +use std::num::NonZeroU32; + +use dspy_rs::LMConfig; +use dspy_rs::ir::{ + self, ApplyError, CapSet, DemoRow, Edit, EditError, EditKind, FieldDef, FieldType as T, Node, + NodeBudget, NodeId, Overlay, ParamKind, ParamValue, PortRef, Program, ProgramBuilder, + SignatureDef, StopSpec, SwapTarget, ToolId, ValidateError, migrate_overlay, +}; +use serde_json::json; + +fn model_config(name: &str) -> LMConfig { + LMConfig { + model: name.to_string(), + ..LMConfig::default() + } +} + +fn obj(pairs: &[(&str, serde_json::Value)]) -> serde_json::Map { + pairs + .iter() + .map(|(k, v)| (k.to_string(), v.clone())) + .collect() +} + +/// A stale/fabricated NodeId (ids are serde-transparent u32s). +fn nid(raw: u32) -> NodeId { + serde_json::from_value(json!(raw)).unwrap() +} + +fn tid(raw: u32) -> ToolId { + serde_json::from_value(json!(raw)).unwrap() +} + +fn tool_id(p: &Program, name: &str) -> ToolId { + p.tools + .iter() + .find_map(|(id, tool)| (p.syms.get(tool.name) == name).then_some(id)) + .unwrap_or_else(|| panic!("tool `{name}` exists")) +} + +/// The exact `cot` reasoning field the builder's `cot()` sugar prepends. +fn reasoning_field() -> FieldDef { + FieldDef::new("reasoning", T::String).with_docs("Think step by step to reach the answer.") +} + +/// Full round trip: canonical text parses back to the same hash and re-prints +/// byte-identically; the JSON projection loads to the same hash. +fn assert_round_trips(p: &Program) { + let text = p.to_dsrs(); + let reparsed = Program::from_dsrs(&text).expect("canonical text parses"); + assert_eq!(reparsed.meta.program_hash, p.meta.program_hash); + assert_eq!(reparsed.to_dsrs(), text, "print(parse(t)) == t"); + + let json = serde_json::to_string(p).unwrap(); + let loaded: Program = serde_json::from_str(&json).unwrap(); + assert_eq!(loaded.meta.program_hash, p.meta.program_hash); +} + +/// question → drafter (QA) → checker (Check) → verdict. +fn pipeline() -> Program { + let mut b = ProgramBuilder::new("pipeline"); + b.model("m", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .instruction("Answer the question.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let check = b.sig( + SignatureDef::build("Check") + .instruction("Judge the answer.") + .input("answer", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let drafter = ir::predict("drafter", qa).bind("question", ir::input("question")); + let checker = ir::predict("checker", check).bind("answer", ir::out("drafter", "answer")); + b.main( + main_sig, + ir::seq([drafter, checker]).out("verdict", ir::out("checker", "verdict")), + ) + .unwrap() +} + +/// question → researcher (agent, tools [search], stop_tools [search]) → +/// summarizer (predict) → summary. Declares a second tool (calc) nothing uses. +fn agent_program() -> Program { + let mut b = ProgramBuilder::new("agents"); + b.cap("net:search"); + b.cap("math:eval"); + b.model("m", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("summary", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .instruction("Research the question.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let sum = b.sig( + SignatureDef::build("Sum") + .instruction("Summarize the answer.") + .input("answer", T::String) + .output("summary", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::List(Box::new(T::String))) + .finish() + .unwrap(), + ); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expr", T::String) + .output("value", T::String) + .finish() + .unwrap(), + ); + let search = b.host_tool("search", "Web search", search_sig, &["net:search"]); + let _calc = b.host_tool("calc", "Calculator", calc_sig, &["math:eval"]); + let researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools([search]) + .stop_tools([search]) + .max_turns(6); + let summarizer = ir::predict("summarizer", sum).bind("answer", ir::out("researcher", "answer")); + b.main( + main_sig, + ir::seq([researcher, summarizer]).out("summary", ir::out("summarizer", "summary")), + ) + .unwrap() +} + +// --------------------------------------------------------------------------- +// Edits are serde values +// --------------------------------------------------------------------------- + +#[test] +fn every_edit_variant_serde_round_trips() { + let edits = vec![ + Edit::AugmentSig { + leaf: nid(1), + prepend: reasoning_field(), + }, + Edit::SwapLeaf { + leaf: nid(1), + to: SwapTarget::Agent { + tools: vec![tid(0)], + stop: StopSpec::default(), + budget: NodeBudget::default(), + }, + }, + Edit::SwapLeaf { + leaf: nid(2), + to: SwapTarget::Predict, + }, + Edit::WrapRetry { + node: nid(3), + max_attempts: NonZeroU32::new(2).unwrap(), + backoff_ms: 250, + feedback: true, + }, + Edit::Remove { node: nid(4) }, + Edit::AddTool { + agent: nid(1), + tool: tid(1), + }, + Edit::RemoveTool { + agent: nid(1), + tool: tid(0), + }, + Edit::SetStop { + agent: nid(1), + stop: StopSpec { + max_turns: NonZeroU32::new(3).unwrap(), + stop_tools: Box::new([tid(0)]), + until_parse: false, + }, + }, + Edit::SetInstructionDefault { + leaf: nid(1), + text: "Be terse.".to_string(), + }, + ]; + for edit in &edits { + let json = serde_json::to_string(edit).unwrap(); + let back: Edit = serde_json::from_str(&json).unwrap(); + assert_eq!(&back, edit, "round trip of {json}"); + } + // The whole list round-trips as one value (edits are replayable scripts). + let json = serde_json::to_string(&edits).unwrap(); + let back: Vec = serde_json::from_str(&json).unwrap(); + assert_eq!(back, edits); + + // EditKind is serde too (the proposer menu is promptable data). + let kinds = vec![ + EditKind::AugmentSig, + EditKind::AddTool { tool: tid(1) }, + EditKind::SetStop, + ]; + let json = serde_json::to_string(&kinds).unwrap(); + let back: Vec = serde_json::from_str(&json).unwrap(); + assert_eq!(back, kinds); +} + +// --------------------------------------------------------------------------- +// AugmentSig +// --------------------------------------------------------------------------- + +#[test] +fn augment_sig_is_the_cot_move() { + let parent = pipeline(); + let drafter = parent.leaf_id("drafter").unwrap(); + let old_sig = match &parent.nodes[drafter] { + Node::Predict(n) => n.sig, + _ => unreachable!(), + }; + + let child = parent + .edited(&[Edit::AugmentSig { + leaf: drafter, + prepend: reasoning_field(), + }]) + .unwrap(); + + // New SigId, reasoning prepended, base outputs preserved. + let new_sig = match &child.nodes[child.leaf_id("drafter").unwrap()] { + Node::Predict(n) => n.sig, + _ => unreachable!(), + }; + assert_ne!(new_sig, old_sig); + let sig = &child.sigs[new_sig]; + assert_eq!(&*sig.outputs[0].name, "reasoning"); + assert_eq!(&*sig.outputs[1].name, "answer"); + assert_eq!(&*sig.name, "QA", "the cot move keeps the base name"); + + // The parent is untouched; the child re-validates under a new hash. + assert_eq!(parent.sigs[old_sig].outputs.len(), 1); + assert_ne!(child.meta.program_hash, parent.meta.program_hash); + child.validate().unwrap(); + + // The canonical text re-sugars the augmented Predict. + assert!(child.to_dsrs().contains("cot QA")); + assert_round_trips(&child); +} + +#[test] +fn augment_sig_with_other_fields_renames_the_new_signature() { + // Two leaves share one SigId: copy-on-write must not disturb the sibling. + let mut b = ProgramBuilder::new("shared"); + b.model("m", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let first = ir::predict("first", qa).bind("question", ir::input("question")); + let second = ir::predict("second", qa).bind("question", ir::out("first", "answer")); + let parent = b + .main( + main_sig, + ir::seq([first, second]).out("answer", ir::out("second", "answer")), + ) + .unwrap(); + + let first_id = parent.leaf_id("first").unwrap(); + let child = parent + .edited(&[Edit::AugmentSig { + leaf: first_id, + prepend: FieldDef::new("plan", T::String), + }]) + .unwrap(); + + let (first_sig, second_sig) = ( + match &child.nodes[child.leaf_id("first").unwrap()] { + Node::Predict(n) => n.sig, + _ => unreachable!(), + }, + match &child.nodes[child.leaf_id("second").unwrap()] { + Node::Predict(n) => n.sig, + _ => unreachable!(), + }, + ); + // The sibling keeps the shared signature untouched. + assert_ne!(first_sig, second_sig); + assert_eq!(&*child.sigs[second_sig].name, "QA"); + assert_eq!(child.sigs[second_sig].outputs.len(), 1); + // The augmented copy got a fresh name (two `sig QA` blocks cannot print). + assert_eq!(&*child.sigs[first_sig].name, "QA_plan"); + assert_eq!(&*child.sigs[first_sig].outputs[0].name, "plan"); + assert!(child.to_dsrs().contains("sig QA_plan")); + assert_round_trips(&child); +} + +#[test] +fn augment_sig_rejects_duplicate_fields_and_containers() { + let parent = pipeline(); + let drafter = parent.leaf_id("drafter").unwrap(); + + let err = parent + .edited(&[Edit::AugmentSig { + leaf: drafter, + prepend: FieldDef::new("answer", T::String), + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + index: 0, + reason: ApplyError::DuplicateField { ref field, .. }, + .. + } if field == "answer" + )); + + let err = parent + .edited(&[Edit::AugmentSig { + leaf: parent.root, + prepend: reasoning_field(), + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::WrongKind { got: "seq", .. }, + .. + } + )); +} + +// --------------------------------------------------------------------------- +// SwapLeaf +// --------------------------------------------------------------------------- + +#[test] +fn swap_predict_to_agent_and_back_is_an_involution() { + let parent = agent_program(); + let summarizer = parent.leaf_id("summarizer").unwrap(); + let calc = tool_id(&parent, "calc"); + + let agentic = parent + .edited(&[Edit::SwapLeaf { + leaf: summarizer, + to: SwapTarget::Agent { + tools: vec![calc], + stop: StopSpec::default(), + budget: NodeBudget::default(), + }, + }]) + .unwrap(); + + // Name, sig shape, bindings preserved; kind changed; context slot minted. + let node = &agentic.nodes[agentic.leaf_id("summarizer").unwrap()]; + let Node::AgentLoop(n) = node else { + panic!("summarizer should be an agent now, got {node:?}"); + }; + assert_eq!(&*n.tools, &[calc]); + assert_eq!(agentic.syms.get(n.name), "summarizer"); + assert!(matches!(n.binding[0].src, PortRef::Out { .. })); + let context = agentic.param_id("summarizer.context").unwrap(); + assert_eq!(agentic.params[context].kind, ParamKind::ContextPolicy); + // The instruction/demos/model slots carried over, values intact. + let instr = agentic.param_id("summarizer.instruction").unwrap(); + assert_eq!( + agentic.params[instr].default, + ParamValue::Instruction { + text: "Summarize the answer.".to_string() + } + ); + assert_ne!(agentic.meta.program_hash, parent.meta.program_hash); + assert_round_trips(&agentic); + + // Swap back: the context slot is collected and the content hash returns + // to the parent's (lineage differs, but lineage is outside the hash). + let back = agentic + .edited(&[Edit::SwapLeaf { + leaf: agentic.leaf_id("summarizer").unwrap(), + to: SwapTarget::Predict, + }]) + .unwrap(); + assert!(matches!( + back.nodes[back.leaf_id("summarizer").unwrap()], + Node::Predict(_) + )); + assert!(back.param_id("summarizer.context").is_none()); + assert_eq!(back.meta.program_hash, parent.meta.program_hash); + assert_round_trips(&back); +} + +#[test] +fn swap_rejects_containers_wrong_directions_and_unknown_tools() { + let parent = agent_program(); + let researcher = parent.leaf_id("researcher").unwrap(); + let summarizer = parent.leaf_id("summarizer").unwrap(); + + let to_agent = SwapTarget::Agent { + tools: vec![], + stop: StopSpec::default(), + budget: NodeBudget::default(), + }; + + // Containers cannot swap. + let err = parent + .edited(&[Edit::SwapLeaf { + leaf: parent.root, + to: to_agent.clone(), + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::WrongKind { got: "seq", .. }, + .. + } + )); + + // Predict → Predict and Agent → Agent are not swaps. + let err = parent + .edited(&[Edit::SwapLeaf { + leaf: summarizer, + to: SwapTarget::Predict, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::WrongKind { got: "predict", .. }, + .. + } + )); + let err = parent + .edited(&[Edit::SwapLeaf { + leaf: researcher, + to: to_agent, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::WrongKind { got: "agent", .. }, + .. + } + )); + + // Tools must exist in program.tools. + let err = parent + .edited(&[Edit::SwapLeaf { + leaf: summarizer, + to: SwapTarget::Agent { + tools: vec![tid(99)], + stop: StopSpec::default(), + budget: NodeBudget::default(), + }, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::UnknownTool { .. }, + .. + } + )); +} + +// --------------------------------------------------------------------------- +// WrapRetry +// --------------------------------------------------------------------------- + +#[test] +fn wrap_retry_rewires_the_parent_and_downstream_ports() { + let parent = pipeline(); + let drafter = parent.leaf_id("drafter").unwrap(); + + let child = parent + .edited(&[Edit::WrapRetry { + node: drafter, + max_attempts: NonZeroU32::new(2).unwrap(), + backoff_ms: 100, + feedback: true, + }]) + .unwrap(); + + // The retry sits where the drafter sat; the drafter is its child. + let new_drafter = child.leaf_id("drafter").unwrap(); + let (retry_id, retry) = child + .nodes + .iter() + .find_map(|(id, node)| match node { + Node::Retry(r) => Some((id, r)), + _ => None, + }) + .expect("a retry node exists"); + assert_eq!(retry.child, new_drafter); + assert_eq!(retry.max_attempts.get(), 2); + assert_eq!(retry.backoff_ms, 100); + assert!(retry.feedback); + let Node::Seq(root) = &child.nodes[child.root] else { + panic!("root is a seq"); + }; + assert!(root.body.contains(&retry_id)); + assert!(!root.body.contains(&new_drafter)); + + // The checker's binding was redirected to the wrapper (sibling-level + // visibility) and the whole thing still validates and round-trips. + let Node::Predict(checker) = &child.nodes[child.leaf_id("checker").unwrap()] else { + panic!("checker is a predict"); + }; + assert!(matches!( + checker.binding[0].src, + PortRef::Out { node, .. } if node == retry_id + )); + assert_ne!(child.meta.program_hash, parent.meta.program_hash); + assert_round_trips(&child); +} + +#[test] +fn wrap_retry_on_the_root_fails_validation() { + let parent = pipeline(); + let err = parent + .edited(&[Edit::WrapRetry { + node: parent.root, + max_attempts: NonZeroU32::new(2).unwrap(), + backoff_ms: 0, + feedback: false, + }]) + .unwrap_err(); + assert!(matches!(err, EditError::Invalid(ValidateError::RootNotSeq))); +} + +// --------------------------------------------------------------------------- +// Remove +// --------------------------------------------------------------------------- + +#[test] +fn remove_collects_the_subtree_and_its_params() { + // Variant where the checker's output is unused: main exports the draft. + let mut b = ProgramBuilder::new("removable"); + b.model("m", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let check = b.sig( + SignatureDef::build("Check") + .input("answer", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let drafter = ir::predict("drafter", qa).bind("question", ir::input("question")); + let checker = ir::predict("checker", check).bind("answer", ir::out("drafter", "answer")); + let parent = b + .main( + main_sig, + ir::seq([drafter, checker]).out("answer", ir::out("drafter", "answer")), + ) + .unwrap(); + + let child = parent + .edited(&[Edit::Remove { + node: parent.leaf_id("checker").unwrap(), + }]) + .unwrap(); + + assert!(child.leaf_id("checker").is_none()); + assert_eq!(child.nodes.len(), parent.nodes.len() - 1); + // The checker's param slots and now-orphaned signature went with it. + assert!(child.param_id("checker.instruction").is_none()); + assert!(child.param_id("drafter.instruction").is_some()); + assert!(!child.to_dsrs().contains("sig Check")); + assert_ne!(child.meta.program_hash, parent.meta.program_hash); + assert_round_trips(&child); +} + +#[test] +fn remove_with_a_downstream_reference_fails_validation() { + let parent = pipeline(); + // checker binds drafter.answer; main exports checker.verdict. + let err = parent + .edited(&[Edit::Remove { + node: parent.leaf_id("drafter").unwrap(), + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Invalid(ValidateError::NodeNotVisible { ref at, .. }) if at == "checker" + )); +} + +#[test] +fn remove_of_the_root_is_rejected() { + let parent = pipeline(); + let err = parent + .edited(&[Edit::Remove { node: parent.root }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::NotInSeq { .. }, + .. + } + )); +} + +// --------------------------------------------------------------------------- +// AddTool / RemoveTool / SetStop +// --------------------------------------------------------------------------- + +#[test] +fn add_tool_declares_and_remove_tool_clears_stop_tools() { + let parent = agent_program(); + let researcher = parent.leaf_id("researcher").unwrap(); + let search = tool_id(&parent, "search"); + let calc = tool_id(&parent, "calc"); + + let child = parent + .edited(&[Edit::AddTool { + agent: researcher, + tool: calc, + }]) + .unwrap(); + let Node::AgentLoop(n) = &child.nodes[child.leaf_id("researcher").unwrap()] else { + panic!("researcher is an agent"); + }; + assert_eq!(&*n.tools, &[search, calc]); + assert_round_trips(&child); + + // RemoveTool drops the declaration *and* the stop_tools entry. + let cleared = child + .edited(&[Edit::RemoveTool { + agent: child.leaf_id("researcher").unwrap(), + tool: search, + }]) + .unwrap(); + let Node::AgentLoop(n) = &cleared.nodes[cleared.leaf_id("researcher").unwrap()] else { + panic!("researcher is an agent"); + }; + assert_eq!(&*n.tools, &[calc]); + assert!(n.stop.stop_tools.is_empty()); + assert_round_trips(&cleared); +} + +#[test] +fn add_tool_failure_modes() { + let parent = agent_program(); + let researcher = parent.leaf_id("researcher").unwrap(); + let search = tool_id(&parent, "search"); + let calc = tool_id(&parent, "calc"); + + // Unknown ToolId. + let err = parent + .edited(&[Edit::AddTool { + agent: researcher, + tool: tid(42), + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::UnknownTool { .. }, + .. + } + )); + + // Already declared. + let err = parent + .edited(&[Edit::AddTool { + agent: researcher, + tool: search, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::ToolAlreadyDeclared { ref name, .. }, + .. + } if name == "search" + )); + + // Caps outside the program ceiling. (Builders can never produce this + // state; simulate a hostile/edited artifact by shrinking the ceiling.) + let mut stripped = parent.clone(); + stripped.caps = CapSet::new(); + let err = stripped + .edited(&[Edit::AddTool { + agent: researcher, + tool: calc, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::ToolCapsExceedProgram { ref name, ref missing }, + .. + } if name == "calc" && missing == &["math:eval".to_string()] + )); + + // Target must be an agent leaf. + let err = parent + .edited(&[Edit::AddTool { + agent: parent.leaf_id("summarizer").unwrap(), + tool: calc, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::WrongKind { got: "predict", .. }, + .. + } + )); + + // RemoveTool of an undeclared tool. + let err = parent + .edited(&[Edit::RemoveTool { + agent: researcher, + tool: calc, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Apply { + reason: ApplyError::ToolNotDeclared { ref name, .. }, + .. + } if name == "calc" + )); +} + +#[test] +fn tool_edits_keep_the_tool_set_gene_in_sync() { + let parent = agent_program(); + let researcher = parent.leaf_id("researcher").unwrap(); + let search = tool_id(&parent, "search"); + let calc = tool_id(&parent, "calc"); + + // AddTool grows the gene's default alongside the declaration: the tool + // just declared is live by default. + let child = parent + .edited(&[Edit::AddTool { + agent: researcher, + tool: calc, + }]) + .unwrap(); + let slot = child.param_id("researcher.tool_set").unwrap(); + assert!(matches!( + &child.params[slot].default, + ParamValue::ToolSet { tools } if tools == &vec![search, calc] + )); + assert_round_trips(&child); + + // RemoveTool shrinks it — the child still validates (the gene never + // outlives its alphabet). + let cleared = child + .edited(&[Edit::RemoveTool { + agent: child.leaf_id("researcher").unwrap(), + tool: search, + }]) + .unwrap(); + let slot = cleared.param_id("researcher.tool_set").unwrap(); + assert!(matches!( + &cleared.params[slot].default, + ParamValue::ToolSet { tools } if tools == &vec![tool_id(&cleared, "calc")] + )); + assert_round_trips(&cleared); + + // SwapLeaf Predict → Agent mints the gene (default = declared tools); + // swapping back collects it with the other agent-only slot. + let summarizer = parent.leaf_id("summarizer").unwrap(); + let agentic = parent + .edited(&[Edit::SwapLeaf { + leaf: summarizer, + to: SwapTarget::Agent { + tools: vec![calc], + stop: StopSpec::default(), + budget: NodeBudget::default(), + }, + }]) + .unwrap(); + let slot = agentic.param_id("summarizer.tool_set").unwrap(); + assert_eq!(agentic.params[slot].kind, ParamKind::ToolSet); + assert!(matches!( + &agentic.params[slot].default, + ParamValue::ToolSet { tools } if tools == &vec![tool_id(&agentic, "calc")] + )); + let back = agentic + .edited(&[Edit::SwapLeaf { + leaf: agentic.leaf_id("summarizer").unwrap(), + to: SwapTarget::Predict, + }]) + .unwrap(); + assert!(back.param_id("summarizer.tool_set").is_none()); + assert_eq!(back.meta.program_hash, parent.meta.program_hash); +} + +#[test] +fn migrate_overlay_tool_set_survives_by_intersection() { + // Parent: researcher declares [search, calc]. + let base = agent_program(); + let parent = base + .edited(&[Edit::AddTool { + agent: base.leaf_id("researcher").unwrap(), + tool: tool_id(&base, "calc"), + }]) + .unwrap(); + let search = tool_id(&parent, "search"); + let calc = tool_id(&parent, "calc"); + let slot = parent + .slot_of::("researcher.tool_set") + .unwrap(); + + let mut both = Overlay::new(&parent); + both.set_tool_set(&parent, slot, vec![search, calc]) + .unwrap(); + let mut calc_only = Overlay::new(&parent); + calc_only.set_tool_set(&parent, slot, vec![calc]).unwrap(); + + // Child: calc undeclared again. + let child = parent + .edited(&[Edit::RemoveTool { + agent: parent.leaf_id("researcher").unwrap(), + tool: calc, + }]) + .unwrap(); + let child_slot = child.param_id("researcher.tool_set").unwrap(); + + // The selection carries by intersection with the child's alphabet. + let migrated = migrate_overlay(&parent, &both, &child); + assert!(matches!( + migrated.get(child_slot), + Some(ParamValue::ToolSet { tools }) if tools == &vec![tool_id(&child, "search")] + )); + + // A selection with no survivors no longer fits and is dropped. + let migrated = migrate_overlay(&parent, &calc_only, &child); + assert!(migrated.get(child_slot).is_none()); +} + +#[test] +fn set_stop_replaces_the_spec_and_is_validated() { + let parent = agent_program(); + let researcher = parent.leaf_id("researcher").unwrap(); + let search = tool_id(&parent, "search"); + let calc = tool_id(&parent, "calc"); + + let stop = StopSpec { + max_turns: NonZeroU32::new(3).unwrap(), + stop_tools: Box::new([search]), + until_parse: false, + }; + let child = parent + .edited(&[Edit::SetStop { + agent: researcher, + stop: stop.clone(), + }]) + .unwrap(); + let Node::AgentLoop(n) = &child.nodes[child.leaf_id("researcher").unwrap()] else { + panic!("researcher is an agent"); + }; + assert_eq!(n.stop, stop); + assert_round_trips(&child); + + // A stop tool the node does not declare is validate.rs's call. + let err = parent + .edited(&[Edit::SetStop { + agent: researcher, + stop: StopSpec { + stop_tools: Box::new([calc]), + ..StopSpec::default() + }, + }]) + .unwrap_err(); + assert!(matches!( + err, + EditError::Invalid(ValidateError::StopToolNotDeclared { ref at }) if at == "researcher" + )); +} + +// --------------------------------------------------------------------------- +// SetInstructionDefault +// --------------------------------------------------------------------------- + +#[test] +fn set_instruction_default_is_a_bake_like_change() { + let parent = pipeline(); + let child = parent + .edited(&[Edit::SetInstructionDefault { + leaf: parent.leaf_id("drafter").unwrap(), + text: "EDITED: answer in one word.".to_string(), + }]) + .unwrap(); + + let instr = child.param_id("drafter.instruction").unwrap(); + assert_eq!( + child.params[instr].default, + ParamValue::Instruction { + text: "EDITED: answer in one word.".to_string() + } + ); + // The parent's default is untouched. + let parent_instr = parent.param_id("drafter.instruction").unwrap(); + assert_eq!( + parent.params[parent_instr].default, + ParamValue::Instruction { + text: "Answer the question.".to_string() + } + ); + assert!(child.to_dsrs().contains("EDITED: answer in one word.")); + assert_ne!(child.meta.program_hash, parent.meta.program_hash); + assert_round_trips(&child); +} + +// --------------------------------------------------------------------------- +// Stale ids, purity, lineage, identity +// --------------------------------------------------------------------------- + +#[test] +fn edits_on_stale_node_ids_fail() { + let parent = pipeline(); + let stale = nid(999); + let edits = [ + Edit::AugmentSig { + leaf: stale, + prepend: reasoning_field(), + }, + Edit::SwapLeaf { + leaf: stale, + to: SwapTarget::Predict, + }, + Edit::WrapRetry { + node: stale, + max_attempts: NonZeroU32::new(2).unwrap(), + backoff_ms: 0, + feedback: false, + }, + Edit::Remove { node: stale }, + Edit::SetStop { + agent: stale, + stop: StopSpec::default(), + }, + Edit::SetInstructionDefault { + leaf: stale, + text: "x".to_string(), + }, + ]; + for edit in edits { + let err = parent.edited(std::slice::from_ref(&edit)).unwrap_err(); + assert!( + matches!( + err, + EditError::Apply { + index: 0, + reason: ApplyError::StaleNode { .. }, + .. + } + ), + "expected StaleNode for {edit:?}, got {err:?}" + ); + } +} + +#[test] +fn edited_is_pure_and_stamps_lineage() { + let parent = pipeline(); + let parent_hash = parent.meta.program_hash; + let parent_json = serde_json::to_string(&parent).unwrap(); + + let child = parent + .edited(&[Edit::AugmentSig { + leaf: parent.leaf_id("drafter").unwrap(), + prepend: reasoning_field(), + }]) + .unwrap(); + + // Purity: the parent is bit-identical. + assert_eq!(serde_json::to_string(&parent).unwrap(), parent_json); + assert_eq!(parent.meta.program_hash, parent_hash); + assert!(parent.meta.lineage.is_none()); + + // The child records its parent the way bake() does, and its hash is the + // recomputed content hash (lineage is outside the preimage). + let lineage = child.meta.lineage.as_ref().unwrap(); + assert_eq!( + lineage.parent.as_deref(), + Some(format!("{parent_hash:016x}").as_str()) + ); + assert_eq!(child.meta.program_hash, child.compute_hash()); +} + +#[test] +fn edited_with_no_edits_is_a_hash_no_op() { + let parent = pipeline(); + let child = parent.edited(&[]).unwrap(); + assert_eq!(child.meta.program_hash, parent.meta.program_hash); + assert!(child.meta.lineage.is_some()); + assert_round_trips(&child); +} + +#[test] +fn a_batch_applies_in_order_over_intermediate_states() { + let parent = agent_program(); + let summarizer = parent.leaf_id("summarizer").unwrap(); + let researcher = parent.leaf_id("researcher").unwrap(); + let search = tool_id(&parent, "search"); + let calc = tool_id(&parent, "calc"); + + // The AddTool/SetStop target the leaf *after* it becomes an agent in the + // same batch — NodeIds are stable within a batch (swaps are in-place). + let child = parent + .edited(&[ + Edit::AugmentSig { + leaf: researcher, + prepend: FieldDef::new("plan", T::String), + }, + Edit::SwapLeaf { + leaf: summarizer, + to: SwapTarget::Agent { + tools: vec![calc], + stop: StopSpec::default(), + budget: NodeBudget::default(), + }, + }, + Edit::AddTool { + agent: summarizer, + tool: search, + }, + Edit::SetStop { + agent: summarizer, + stop: StopSpec { + max_turns: NonZeroU32::new(4).unwrap(), + stop_tools: Box::new([calc]), + until_parse: true, + }, + }, + ]) + .unwrap(); + + let Node::AgentLoop(n) = &child.nodes[child.leaf_id("summarizer").unwrap()] else { + panic!("summarizer became an agent"); + }; + assert_eq!(&*n.tools, &[calc, search]); + assert_eq!(n.stop.max_turns.get(), 4); + let Node::AgentLoop(r) = &child.nodes[child.leaf_id("researcher").unwrap()] else { + panic!("researcher is still an agent"); + }; + assert_eq!(&*child.sigs[r.sig].name, "Research_plan"); + assert_round_trips(&child); +} + +// --------------------------------------------------------------------------- +// legal_edits +// --------------------------------------------------------------------------- + +#[test] +fn legal_edits_menus_track_node_shape() { + let p = agent_program(); + let researcher = p.leaf_id("researcher").unwrap(); + let summarizer = p.leaf_id("summarizer").unwrap(); + let search = tool_id(&p, "search"); + let calc = tool_id(&p, "calc"); + + let predict_menu = p.legal_edits(summarizer); + assert!(predict_menu.contains(&EditKind::AugmentSig)); + assert!(predict_menu.contains(&EditKind::SetInstructionDefault)); + assert!(predict_menu.contains(&EditKind::SwapToAgent)); + assert!(predict_menu.contains(&EditKind::WrapRetry)); + assert!(predict_menu.contains(&EditKind::Remove)); + assert!(!predict_menu.contains(&EditKind::SwapToPredict)); + assert!(!predict_menu.contains(&EditKind::SetStop)); + assert!(!predict_menu.contains(&EditKind::AddTool { tool: search })); + + let agent_menu = p.legal_edits(researcher); + assert!(agent_menu.contains(&EditKind::SwapToPredict)); + assert!(agent_menu.contains(&EditKind::SetStop)); + // Declared tools are removable, undeclared ones addable. + assert!(agent_menu.contains(&EditKind::RemoveTool { tool: search })); + assert!(agent_menu.contains(&EditKind::AddTool { tool: calc })); + assert!(!agent_menu.contains(&EditKind::AddTool { tool: search })); + assert!(!agent_menu.contains(&EditKind::SwapToAgent)); + + // The root cannot be wrapped or removed; a stale id gets no menu. + assert!(p.legal_edits(p.root).is_empty()); + assert!(p.legal_edits(nid(999)).is_empty()); + + // Every menu entry serializes (the menu is proposer-facing data). + serde_json::to_string(&agent_menu).unwrap(); +} + +// --------------------------------------------------------------------------- +// migrate_overlay +// --------------------------------------------------------------------------- + +fn tuned_overlay(p: &Program) -> Overlay { + let mut overlay = Overlay::new(p); + let instr = p.slot_of::("drafter.instruction").unwrap(); + overlay.set_instruction(instr, "TUNED: be terse."); + let demos = p.slot_of::("drafter.demos").unwrap(); + overlay.set_demos( + demos, + vec![DemoRow { + input: obj(&[("question", json!("demo q"))]), + output: obj(&[("answer", json!("demo a"))]), + }], + ); + overlay +} + +#[test] +fn migrate_overlay_survives_an_unrelated_edit() { + let parent = pipeline(); + let overlay = tuned_overlay(&parent); + + // Edit a *different* leaf: the drafter's genes must carry over. + let child = parent + .edited(&[Edit::SetInstructionDefault { + leaf: parent.leaf_id("checker").unwrap(), + text: "Judge strictly.".to_string(), + }]) + .unwrap(); + + let migrated = migrate_overlay(&parent, &overlay, &child); + assert_eq!(migrated.base, child.meta.program_hash); + let instr = child.param_id("drafter.instruction").unwrap(); + assert_eq!( + migrated.resolve(&child, instr), + &ParamValue::Instruction { + text: "TUNED: be terse.".to_string() + } + ); + let demos = child.param_id("drafter.demos").unwrap(); + assert!(matches!( + migrated.resolve(&child, demos), + ParamValue::Demos { rows } if rows.len() == 1 + )); +} + +#[test] +fn migrate_overlay_survives_augment_sig() { + // Documented decision: AugmentSig keeps inputs identical and only widens + // the outputs, so demo rows still map onto the base fields — instruction + // AND demos survive the CoT move. + let parent = pipeline(); + let overlay = tuned_overlay(&parent); + let child = parent + .edited(&[Edit::AugmentSig { + leaf: parent.leaf_id("drafter").unwrap(), + prepend: reasoning_field(), + }]) + .unwrap(); + + let migrated = migrate_overlay(&parent, &overlay, &child); + let instr = child.param_id("drafter.instruction").unwrap(); + let demos = child.param_id("drafter.demos").unwrap(); + assert!(migrated.get(instr).is_some(), "instruction survives"); + assert!(migrated.get(demos).is_some(), "demos survive"); + assert!(matches!( + migrated.resolve(&child, demos), + ParamValue::Demos { rows } if rows[0].input.contains_key("question") + )); +} + +#[test] +fn migrate_overlay_drops_entries_when_the_shape_changed() { + let parent = pipeline(); + let overlay = tuned_overlay(&parent); + + // A child where `drafter` exists at the same ParamPaths but its signature + // has a different shape (different output field name). + let mut b = ProgramBuilder::new("pipeline"); + b.model("m", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .instruction("Answer the question.") + .input("question", T::String) + .output("reply", T::String) + .finish() + .unwrap(), + ); + let check = b.sig( + SignatureDef::build("Check") + .instruction("Judge the answer.") + .input("answer", T::String) + .output("verdict", T::String) + .finish() + .unwrap(), + ); + let drafter = ir::predict("drafter", qa).bind("question", ir::input("question")); + let checker = ir::predict("checker", check).bind("answer", ir::out("drafter", "reply")); + let child = b + .main( + main_sig, + ir::seq([drafter, checker]).out("verdict", ir::out("checker", "verdict")), + ) + .unwrap(); + + let migrated = migrate_overlay(&parent, &overlay, &child); + assert_eq!(migrated.base, child.meta.program_hash); + assert!( + migrated.is_empty(), + "shape-changed leaf carries nothing over" + ); +} + +#[test] +fn migrate_overlay_remints_model_refs_by_name() { + // Two models so the ModelRef ordinal is meaningful. + let build = |name: &str, flip: bool| { + let mut b = ProgramBuilder::new(name); + // Registration order differs between parent and child: the ordinal + // for "deep" is m1 in one and m0 in the other. + let (fast, deep) = if flip { + let deep = b.model("deep", model_config("anthropic:claude-sonnet-4-5")); + let fast = b.model("fast", model_config("openai:gpt-4o-mini")); + (fast, deep) + } else { + let fast = b.model("fast", model_config("openai:gpt-4o-mini")); + let deep = b.model("deep", model_config("anthropic:claude-sonnet-4-5")); + (fast, deep) + }; + let _ = deep; + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let node = ir::predict("drafter", main_sig) + .model(fast) + .bind("question", ir::input("question")); + b.main( + main_sig, + ir::seq([node]).out("answer", ir::out("drafter", "answer")), + ) + .unwrap() + }; + let parent = build("models", false); + let child = build("models", true); + + let mut overlay = Overlay::new(&parent); + let model_slot = parent.param_id("drafter.model").unwrap(); + let deep_in_parent = parent + .models + .iter() + .find_map(|(id, m)| (&*m.name == "deep").then_some(id)) + .unwrap(); + overlay + .set( + &parent, + model_slot, + ParamValue::ModelRef { + model: deep_in_parent, + }, + ) + .unwrap(); + + let migrated = migrate_overlay(&parent, &overlay, &child); + let child_slot = child.param_id("drafter.model").unwrap(); + let deep_in_child = child + .models + .iter() + .find_map(|(id, m)| (&*m.name == "deep").then_some(id)) + .unwrap(); + assert_ne!( + deep_in_parent, deep_in_child, + "the ordinals genuinely differ" + ); + assert_eq!( + migrated.get(child_slot), + Some(&ParamValue::ModelRef { + model: deep_in_child + }) + ); +} + +#[test] +fn migrate_overlay_with_a_mismatched_base_yields_empty() { + let parent = pipeline(); + let other = agent_program(); + let overlay = tuned_overlay(&parent); + // Overlay minted against `parent` cannot be interpreted against `other`. + let migrated = migrate_overlay(&other, &overlay, &parent); + assert!(migrated.is_empty()); + assert_eq!(migrated.base, parent.meta.program_hash); +} diff --git a/crates/dspy-rs/tests/test_ir_interp.rs b/crates/dspy-rs/tests/test_ir_interp.rs index 8d83b7e1..fd28659c 100644 --- a/crates/dspy-rs/tests/test_ir_interp.rs +++ b/crates/dspy-rs/tests/test_ir_interp.rs @@ -1,7 +1,6 @@ //! IR-3 (RFC 0002 §3): interpreter end-to-end over canned LM responses — //! every node kind, overlay read-through at render time, trace capture with //! program leaf names, budget metering, and load-time refusals. -#![cfg(feature = "ir")] use std::sync::Arc; @@ -677,6 +676,118 @@ async fn agent_loop_is_one_span_with_tool_events() { ); } +/// Like [`agent_program`] but with two declared tools (`search`, `calc`). +fn two_tool_agent_program() -> Program { + let mut b = ProgramBuilder::new("agentic2"); + b.cap("net:search"); + b.model("m", config()); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .instruction("Research and answer.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::List(Box::new(T::String))) + .finish() + .unwrap(), + ); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expression", T::String) + .output("value", T::Float) + .finish() + .unwrap(), + ); + let search = b.host_tool( + "search", + "Web search; returns result snippets with URLs", + search_sig, + &["net:search"], + ); + let calc = b.host_tool("calc", "Evaluate an arithmetic expression", calc_sig, &[]); + let researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools([search, calc]) + .max_turns(4); + b.main( + main_sig, + ir::seq([researcher]).out("answer", ir::out("researcher", "answer")), + ) + .unwrap() +} + +#[tokio::test] +async fn tool_set_overlay_selects_the_tool_surface() { + let (lm, client) = canned_lm(vec![ + text(fields(&[("answer", "one")])), + text(fields(&[("answer", "two")])), + ]) + .await; + let env = RuntimeEnv::new() + .bind_model("m", lm) + .bind_host_tool("search", Arc::new(SearchTool)) + .bind_host_tool("calc", Arc::new(SearchTool)) + .grant("net:search"); + let interp = Interpreter::load(two_tool_agent_program(), env) + .await + .unwrap(); + let program = Arc::clone(interp.program()); + + // Absent slot: the loop carries the full declared table. + interp + .run(obj(&[("question", json!("q"))]), None, Budget::unlimited()) + .await + .unwrap(); + let full = client.last_request().unwrap(); + let names: Vec<&str> = full.tools.iter().map(|t| t.name.as_str()).collect(); + assert_eq!(names, vec!["search", "calc"]); + + // A ToolSet overlay entry restricts the surface to exactly the selection + // — read through at render time, the program untouched. + let search_id = program + .tools + .iter() + .find_map(|(id, t)| (program.syms.get(t.name) == "search").then_some(id)) + .unwrap(); + let slot = program + .slot_of::("researcher.tool_set") + .unwrap(); + let mut overlay = Overlay::new(&program); + overlay + .set_tool_set(&program, slot, vec![search_id]) + .unwrap(); + interp + .run( + obj(&[("question", json!("q"))]), + Some(Arc::new(overlay)), + Budget::unlimited(), + ) + .await + .unwrap(); + let restricted = client.last_request().unwrap(); + let names: Vec<&str> = restricted.tools.iter().map(|t| t.name.as_str()).collect(); + assert_eq!(names, vec!["search"]); + + // The declared table is still structural: the program defaults are + // untouched by the overlay run. + assert!(matches!( + &program.params[slot.id].default, + ParamValue::ToolSet { tools } if tools.len() == 2 + )); +} + // --------------------------------------------------------------------------- // Hole via sandbox // --------------------------------------------------------------------------- diff --git a/crates/dspy-rs/tests/test_ir_m1_holes.rs b/crates/dspy-rs/tests/test_ir_m1_holes.rs index ebd2a93b..48d24843 100644 --- a/crates/dspy-rs/tests/test_ir_m1_holes.rs +++ b/crates/dspy-rs/tests/test_ir_m1_holes.rs @@ -2,7 +2,6 @@ //! extern/host holes), the non-degenerate hole `request_hash` preimage, and //! interpreter-lane replay — predicts, agent loops, and holes served from a //! recorded trace with divergence detection on changed hole implementations. -#![cfg(feature = "ir")] use std::sync::Arc; diff --git a/crates/dspy-rs/tests/test_ir_program.rs b/crates/dspy-rs/tests/test_ir_program.rs index 76fb8cb2..1ed35ba2 100644 --- a/crates/dspy-rs/tests/test_ir_program.rs +++ b/crates/dspy-rs/tests/test_ir_program.rs @@ -1,6 +1,5 @@ //! IR-2 (RFC 0002 §2): program construction, load-time validation, parameter //! addressing, overlays, and the canonical serde round trip. -#![cfg(feature = "ir")] use dspy_rs::LMConfig; use dspy_rs::ir::{ @@ -129,6 +128,7 @@ fn param_paths_are_addressable() { "researcher.demos", "researcher.model", "researcher.context", + "researcher.tool_set", "checker.code", "tool.search.desc", ] { @@ -613,6 +613,294 @@ fn overlay_hash_and_named_round_trip() { )); } +// --------------------------------------------------------------------------- +// The ToolSet gene (RFC 0004 §5) +// --------------------------------------------------------------------------- + +/// Two declared tools; the agent declares both. Returns the program and the +/// two tool ids in declaration order (`search`, `calc`). +fn two_tool_program() -> (Program, ir::ToolId, ir::ToolId) { + let mut b = ProgramBuilder::new("twotools"); + b.cap("net:search"); + b.model("only", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .instruction("Research and answer.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::List(Box::new(T::String))) + .finish() + .unwrap(), + ); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expression", T::String) + .output("value", T::Float) + .finish() + .unwrap(), + ); + let search = b.host_tool("search", "Web search", search_sig, &["net:search"]); + let calc = b.host_tool("calc", "Evaluate arithmetic", calc_sig, &[]); + let researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools([search, calc]) + .max_turns(4); + let program = b + .main( + main_sig, + ir::seq([researcher]).out("answer", ir::out("researcher", "answer")), + ) + .unwrap(); + (program, search, calc) +} + +#[test] +fn tool_set_defaults_to_the_full_declared_table() { + let (program, search, calc) = two_tool_program(); + let id = program.param_id("researcher.tool_set").unwrap(); + assert_eq!(program.params[id].kind, ParamKind::ToolSet); + assert!(matches!( + &program.params[id].default, + ParamValue::ToolSet { tools } if tools == &vec![search, calc] + )); + + // Enumerable like every other gene. + let paths: Vec<&str> = program + .slots(ParamKind::ToolSet) + .map(|(_, slot)| &*slot.path) + .collect(); + assert_eq!(paths, vec!["researcher.tool_set"]); + assert!( + program + .slot_of::("researcher.tool_set") + .is_some() + ); +} + +#[test] +fn tool_set_subset_builds_but_undeclared_tool_is_a_build_error() { + // A declared subset is fine. + let mut b = ProgramBuilder::new("subset"); + b.cap("net:search"); + b.model("only", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::List(Box::new(T::String))) + .finish() + .unwrap(), + ); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expression", T::String) + .output("value", T::Float) + .finish() + .unwrap(), + ); + let search = b.host_tool("search", "Web search", search_sig, &["net:search"]); + let calc = b.host_tool("calc", "Evaluate arithmetic", calc_sig, &[]); + let researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools([search, calc]) + .tool_set([search]) + .max_turns(4); + let program = b + .main( + main_sig, + ir::seq([researcher]).out("answer", ir::out("researcher", "answer")), + ) + .expect("a declared subset validates"); + let id = program.param_id("researcher.tool_set").unwrap(); + assert!(matches!( + &program.params[id].default, + ParamValue::ToolSet { tools } if tools == &vec![search] + )); + + // A tool_set naming a tool the node does not declare is refused at build. + let mut b = ProgramBuilder::new("smuggle"); + b.cap("net:search"); + b.model("only", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::List(Box::new(T::String))) + .finish() + .unwrap(), + ); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expression", T::String) + .output("value", T::Float) + .finish() + .unwrap(), + ); + let search = b.host_tool("search", "Web search", search_sig, &["net:search"]); + let calc = b.host_tool("calc", "Evaluate arithmetic", calc_sig, &[]); + let researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools([search]) + .tool_set([calc]) + .max_turns(4); + let err = b + .main( + main_sig, + ir::seq([researcher]).out("answer", ir::out("researcher", "answer")), + ) + .unwrap_err(); + assert!(matches!( + err, + BuildError::Invalid(ValidateError::ToolSetUndeclared { ref at, ref tool }) + if at == "researcher" && tool == "calc" + )); +} + +#[test] +fn tool_set_duplicates_are_rejected() { + let (program, search, _calc) = two_tool_program(); + let mut hostile = program.clone(); + let id = hostile.param_id("researcher.tool_set").unwrap(); + hostile.params[id].default = ParamValue::ToolSet { + tools: vec![search, search], + }; + let err = hostile.validate().unwrap_err(); + assert!(matches!( + err, + ValidateError::ToolSetDuplicate { ref at, ref tool } + if at == "researcher" && tool == "search" + )); +} + +#[test] +fn overlay_tool_set_accepts_subsets_and_refuses_undeclared_tools() { + let (program, search, calc) = two_tool_program(); + let slot = program + .slot_of::("researcher.tool_set") + .unwrap(); + + // A subset of the declared table is a legal gene value. + let mut overlay = Overlay::new(&program); + overlay.set_tool_set(&program, slot, vec![search]).unwrap(); + assert!(matches!( + overlay.resolve(&program, slot.id), + ParamValue::ToolSet { tools } if tools == &vec![search] + )); + // The empty subset too — an agent may run tool-less. + let mut empty = Overlay::new(&program); + empty.set_tool_set(&program, slot, vec![]).unwrap(); + + // Drop `calc` from the *node* (not the program) and the old overlay + // alphabet no longer applies: a value naming it is refused on set. + let restricted = { + let mut b = ProgramBuilder::new("restricted"); + b.cap("net:search"); + b.model("only", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::List(Box::new(T::String))) + .finish() + .unwrap(), + ); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expression", T::String) + .output("value", T::Float) + .finish() + .unwrap(), + ); + let search = b.host_tool("search", "Web search", search_sig, &["net:search"]); + let _calc = b.host_tool("calc", "Evaluate arithmetic", calc_sig, &[]); + // `calc` is in the program's tool table but NOT declared on the node. + let researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools([search]) + .max_turns(4); + b.main( + main_sig, + ir::seq([researcher]).out("answer", ir::out("researcher", "answer")), + ) + .unwrap() + }; + let slot = restricted + .slot_of::("researcher.tool_set") + .unwrap(); + let mut overlay = Overlay::new(&restricted); + let err = overlay + .set_tool_set(&restricted, slot, vec![calc]) + .unwrap_err(); + assert!(matches!( + err, + OverlayError::ToolSetUndeclared { ref path, ref tool } + if path == "researcher.tool_set" && tool == "calc" + )); + + // The serde load path (`from_named`) applies the same guard. + let mut named = std::collections::BTreeMap::new(); + named.insert( + "researcher.tool_set".to_string(), + ParamValue::ToolSet { tools: vec![calc] }, + ); + assert!(matches!( + Overlay::from_named(&restricted, named), + Err(OverlayError::ToolSetUndeclared { .. }) + )); +} + // --------------------------------------------------------------------------- // Serde round trip // --------------------------------------------------------------------------- diff --git a/crates/dspy-rs/tests/test_ir_sigdef.rs b/crates/dspy-rs/tests/test_ir_sigdef.rs index ddf62597..0356b007 100644 --- a/crates/dspy-rs/tests/test_ir_sigdef.rs +++ b/crates/dspy-rs/tests/test_ir_sigdef.rs @@ -6,7 +6,7 @@ use dspy_rs::ir::{ConstraintDef, FieldDef, FieldType, RenderSpec, SigError, SignatureDef}; use dspy_rs::modules::Reasoning; -use dspy_rs::{Augmented, BamlType, Schema, Signature}; +use dspy_rs::{Augmented, Schema, Signature}; #[derive(Signature, Clone, Debug)] /// Answer questions accurately and concisely. @@ -39,7 +39,7 @@ struct Graded { } #[derive(Clone, Debug)] -#[BamlType] +#[Schema] struct Citation { /// Source URL. url: String, @@ -47,7 +47,7 @@ struct Citation { } #[derive(Clone, Debug)] -#[BamlType] +#[Schema] enum Stance { Support, Refute, @@ -66,7 +66,7 @@ struct Structured { } #[derive(Clone, Debug)] -#[BamlType] +#[Schema] struct DetailOutput { answer: String, } diff --git a/crates/dspy-rs/tests/test_ir_text.rs b/crates/dspy-rs/tests/test_ir_text.rs index fca6a8e6..0e43b706 100644 --- a/crates/dspy-rs/tests/test_ir_text.rs +++ b/crates/dspy-rs/tests/test_ir_text.rs @@ -1,6 +1,5 @@ //! IR-5 (RFC 0002 §4/§5): the `.dsrs` text format — parse, canonical print, //! text-preimage program hash, and parse-error quality. -#![cfg(feature = "ir")] use dspy_rs::LMConfig; use dspy_rs::ir::{ @@ -385,6 +384,120 @@ fn a_json_artifact_is_not_dsrs_text() { assert!(err.message.contains("`dsrs`"), "message: {}", err.message); } +// --------------------------------------------------------------------------- +// The ToolSet gene in the text form +// --------------------------------------------------------------------------- + +/// Two declared tools, `tool_set` restricted to `search`. +fn two_tool_program(restrict: bool) -> Program { + let mut b = ProgramBuilder::new("twotools"); + b.cap("net:search"); + b.model("m", model_config("openai:gpt-4o-mini")); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let research = b.sig( + SignatureDef::build("Research") + .instruction("Research and answer.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let search_sig = b.sig( + SignatureDef::build("Search") + .input("query", T::String) + .output("results", T::List(Box::new(T::String))) + .finish() + .unwrap(), + ); + let calc_sig = b.sig( + SignatureDef::build("Calc") + .input("expression", T::String) + .output("value", T::Float) + .finish() + .unwrap(), + ); + let search = b.host_tool("search", "Web search", search_sig, &["net:search"]); + let calc = b.host_tool("calc", "Evaluate arithmetic", calc_sig, &[]); + let mut researcher = ir::agent("researcher", research) + .bind("question", ir::input("question")) + .tools([search, calc]) + .max_turns(4); + if restrict { + researcher = researcher.tool_set([search]); + } + b.main( + main_sig, + ir::seq([researcher]).out("answer", ir::out("researcher", "answer")), + ) + .unwrap() +} + +#[test] +fn tool_set_gene_round_trips_through_text() { + // A restricted default prints as a `tool_set [...]` line and survives + // print → parse → print with the same canonical text and hash. + let program = two_tool_program(true); + let text = program.to_dsrs(); + assert!(text.contains("tool_set [search]"), "{text}"); + let parsed = Program::from_dsrs(&text).expect("restricted tool_set parses"); + assert_eq!(parsed.to_dsrs(), text); + assert_eq!(parsed.meta.program_hash, program.meta.program_hash); + let id = parsed.param_id("researcher.tool_set").unwrap(); + assert!(matches!( + &parsed.params[id].default, + ParamValue::ToolSet { tools } if tools.len() == 1 + )); + + // The full-table default is elided: selection == declaration prints + // nothing, so pre-ToolSet canonical text (and hashes) are unchanged. + let full = two_tool_program(false); + assert!(!full.to_dsrs().contains("tool_set")); + assert!(!qa_program().to_dsrs().contains("tool_set")); +} + +#[test] +fn tool_set_with_undeclared_tool_fails_the_text_load() { + // `calc` is a program tool but not declared on the node: refused at load + // with the validator's own words. + let src = r#"dsrs 1 +program p + +caps { net:search } + +model m = "openai:gpt-4o-mini" + +sig Main { in question: string out answer: string } + +tool search "Web search" caps [net:search] { + in query: string + out results: string[] +} + +tool calc "Evaluate arithmetic" { + in expression: string + out value: float +} + +main: Main = seq { + researcher = agent Main @m (question = $.question) { + tools [search] + tool_set [calc] + max_turns 4 + } + out { answer = researcher.answer } +} +"#; + let err = Program::from_dsrs(src).expect_err("undeclared tool_set entry must not load"); + let message = err.to_string(); + assert!(message.contains("tool_set includes `calc`"), "{message}"); +} + // --------------------------------------------------------------------------- // Parse-error quality: line + problem, actionable for a generating model // --------------------------------------------------------------------------- diff --git a/crates/dspy-rs/tests/test_lm.rs b/crates/dspy-rs/tests/test_lm.rs index d8c617d7..888c1333 100644 --- a/crates/dspy-rs/tests/test_lm.rs +++ b/crates/dspy-rs/tests/test_lm.rs @@ -145,7 +145,7 @@ async fn test_lm_cache_direct_operations() { let key: CacheKey = 0xDEAD_BEEF; // Initially cache should be empty - let cached = cache.lock().await.get_entry(key).await.unwrap(); + let cached = cache.get_entry(key).await.unwrap(); assert!(cached.is_none()); // Insert an entry @@ -154,12 +154,10 @@ async fn test_lm_cache_direct_operations() { usage: LmUsage::default(), raw_output: Some("answer: Paris".to_string()), }; - cache.lock().await.insert_entry(key, entry.clone()); + cache.insert_entry(key, entry.clone()); // Now cache should return the entry let cached = cache - .lock() - .await .get_entry(key) .await .unwrap() @@ -168,7 +166,7 @@ async fn test_lm_cache_direct_operations() { assert_eq!(cached.raw_output, entry.raw_output); // Unknown keys still miss - let missing = cache.lock().await.get_entry(key ^ 1).await.unwrap(); + let missing = cache.get_entry(key ^ 1).await.unwrap(); assert!(missing.is_none()); } @@ -229,10 +227,10 @@ async fn test_cache_preserves_usage_and_history() { raw_output: Some("answer: A fox jumps over a dog".to_string()), }; - cache.lock().await.insert_entry(key, entry.clone()); + cache.insert_entry(key, entry.clone()); // The cache stores and retrieves the full entry including usage stats. - let cached = cache.lock().await.get_entry(key).await.unwrap().unwrap(); + let cached = cache.get_entry(key).await.unwrap().unwrap(); assert_eq!(cached.prompt, entry.prompt); assert_eq!(cached.raw_output, entry.raw_output); assert_eq!(cached.usage.prompt_tokens, 50); @@ -245,9 +243,9 @@ async fn test_cache_preserves_usage_and_history() { usage: LmUsage::default(), raw_output: None, }; - cache.lock().await.insert_entry(43, later.clone()); + cache.insert_entry(43, later.clone()); - let history = cache.lock().await.get_history(2).await.unwrap(); + let history = cache.get_history(2); assert_eq!(history.len(), 2); assert_eq!(history[0].prompt, later.prompt); assert_eq!(history[1].prompt, entry.prompt); diff --git a/crates/dspy-rs/tests/test_message_roundtrip.rs b/crates/dspy-rs/tests/test_message_roundtrip.rs index 96461483..34fe8447 100644 --- a/crates/dspy-rs/tests/test_message_roundtrip.rs +++ b/crates/dspy-rs/tests/test_message_roundtrip.rs @@ -1,8 +1,7 @@ //! Round-trip tests for the new Message model. //! //! Verifies that the grouped Role + ContentBlock representation preserves -//! all content through: DSRs Message → rig Message → DSRs Message, and -//! through JSON serialization/deserialization. +//! all content through: DSRs Message → rig Message → DSRs Message. use dspy_rs::core::lm::chat::{Chat, ContentBlock, Message, Role}; use rig::OneOrMany; @@ -172,77 +171,6 @@ fn multi_turn_conversation_preserves_earlier_reasoning() { ); } -// --------------------------------------------------------------------------- -// JSON serialization round-trip -// --------------------------------------------------------------------------- - -/// Full multi-content message survives JSON serialization. -#[test] -fn grouped_message_json_roundtrip() { - let original = Chat::new(vec![ - Message::system("Be helpful"), - Message::with_content( - Role::Assistant, - vec![ - ContentBlock::reasoning(Reasoning::new("let me think")), - ContentBlock::text("the answer is 42"), - ContentBlock::tool_call(ToolCall::new( - "tc-1".to_string(), - ToolFunction { - name: "verify".to_string(), - arguments: json!({"answer": 42}), - }, - )), - ], - ), - Message::with_content( - Role::User, - vec![ - ContentBlock::tool_result(ToolResult { - id: "tc-1".to_string(), - call_id: None, - content: OneOrMany::one(ToolResultContent::text("confirmed")), - }), - ContentBlock::text("Thanks! Can you also check 43?"), - ], - ), - ]); - - let json = original.to_json(); - let reparsed = Chat::new(vec![]).from_json(json).unwrap(); - - assert_eq!(reparsed.len(), 3); - - // Verify the assistant message preserved all 3 content blocks - let asst = &reparsed.messages[1]; - assert_eq!(asst.role, Role::Assistant); - assert_eq!(asst.content.len(), 3); - assert!(asst.has_reasoning()); - assert!(asst.has_tool_calls()); - - // Verify the user message preserved both blocks - let user = &reparsed.messages[2]; - assert_eq!(user.role, Role::User); - assert_eq!(user.content.len(), 2); - assert!(user.has_tool_results()); -} - -/// Legacy JSON format (content as plain string) still parses correctly. -#[test] -fn legacy_plain_string_json_parses_into_new_model() { - let legacy_json = json!([ - {"role": "system", "content": "Be helpful"}, - {"role": "user", "content": "Hello"}, - {"role": "assistant", "content": "Hi there!"} - ]); - - let chat = Chat::new(vec![]).from_json(legacy_json).unwrap(); - assert_eq!(chat.len(), 3); - assert_eq!(chat.messages[0].role, Role::System); - assert_eq!(chat.messages[0].content(), "Be helpful"); - assert_eq!(chat.messages[2].text_content(), "Hi there!"); -} - // --------------------------------------------------------------------------- // Accessor correctness // --------------------------------------------------------------------------- diff --git a/crates/dspy-rs/tests/test_miprov2.rs b/crates/dspy-rs/tests/test_miprov2.rs index efcdf5b3..c0f4087e 100644 --- a/crates/dspy-rs/tests/test_miprov2.rs +++ b/crates/dspy-rs/tests/test_miprov2.rs @@ -1,37 +1,6 @@ -use dspy_rs::{Eval, MIPROv2, PromptCandidate, PromptingTips, Signature, Trace, TraceOutcome}; +use dspy_rs::{MIPROv2, PromptingTips}; use rstest::*; -#[derive(Signature, Clone, Debug)] -struct TestSignature { - #[input] - question: String, - - #[output] - answer: String, -} - -/// A rollout trace carrying only a whole-program score, the projection MIPRO's -/// candidate generation reads. -fn scored_trace(score: Option) -> Trace { - Trace { - outcome: score.map(|score| TraceOutcome { - output: None, - error: None, - eval: Some(Eval::score(score)), - duration_us: 0, - }), - ..Trace::default() - } -} - -fn trace_score(trace: &Trace) -> Option { - trace - .outcome - .as_ref() - .and_then(|outcome| outcome.eval.as_ref()) - .map(|eval| eval.score) -} - #[rstest] fn test_prompting_tips_default() { let tips = PromptingTips::default_tips(); @@ -49,20 +18,6 @@ fn test_prompting_tips_formatting() { assert!(formatted.contains("\n")); } -#[rstest] -fn test_prompt_candidate_creation() { - let candidate = PromptCandidate::new("Test instruction".to_string()); - - assert_eq!(candidate.instruction, "Test instruction"); - assert_eq!(candidate.score, 0.0); -} - -#[rstest] -fn test_prompt_candidate_with_score() { - let candidate = PromptCandidate::new("test".to_string()).with_score(0.85); - assert_eq!(candidate.score, 0.85); -} - #[rstest] fn test_miprov2_default_configuration() { let optimizer = MIPROv2::builder().build(); @@ -71,54 +26,3 @@ fn test_miprov2_default_configuration() { assert_eq!(optimizer.num_trials, 20); assert_eq!(optimizer.minibatch_size, 25); } - -#[rstest] -fn test_select_best_traces_descending_order() { - let optimizer = MIPROv2::builder().build(); - - let traces = vec![ - scored_trace(Some(0.1)), - scored_trace(Some(0.5)), - scored_trace(Some(0.3)), - ]; - - let best = optimizer.select_best_traces(&traces, 2); - assert_eq!(best.len(), 2); - assert_eq!(trace_score(best[0]), Some(0.5)); - assert_eq!(trace_score(best[1]), Some(0.3)); -} - -#[rstest] -fn test_select_best_traces_ignores_unscored() { - let optimizer = MIPROv2::builder().build(); - - let traces = vec![scored_trace(None), scored_trace(Some(0.8))]; - - let best = optimizer.select_best_traces(&traces, 2); - assert_eq!(best.len(), 1); - assert_eq!(trace_score(best[0]), Some(0.8)); -} - -#[rstest] -fn test_create_prompt_candidates_uses_all_instructions() { - let optimizer = MIPROv2::builder().build(); - let candidates = optimizer.create_prompt_candidates(vec![ - "instruction-1".to_string(), - "instruction-2".to_string(), - ]); - - assert_eq!(candidates.len(), 2); - assert_eq!(candidates[0].instruction, "instruction-1"); - assert_eq!(candidates[1].instruction, "instruction-2"); -} - -#[rstest] -fn test_format_schema_fields_reads_typed_schema() { - let optimizer = MIPROv2::builder().build(); - let rendered = optimizer.format_schema_fields(TestSignature::schema()); - - assert!(rendered.contains("Input Fields:")); - assert!(rendered.contains("question")); - assert!(rendered.contains("Output Fields:")); - assert!(rendered.contains("answer")); -} diff --git a/crates/dspy-rs/tests/test_module_agent.rs b/crates/dspy-rs/tests/test_module_agent.rs index 497a3d3e..020561fc 100644 --- a/crates/dspy-rs/tests/test_module_agent.rs +++ b/crates/dspy-rs/tests/test_module_agent.rs @@ -2,7 +2,6 @@ //! `#[agent]` (StepDef + standalone static-lane execution), and a `#[module]` //! whose body lowers an agent step to a first-class `AgentLoop` node with the //! tool bound and the capability ceiling self-granted. -#![cfg(feature = "ir")] use std::sync::LazyLock; @@ -67,6 +66,20 @@ fn shout(text: String) -> String { #[agent(tools(shout), max_turns = 3, budget(tokens = 50_000, on_exhausted = finalize))] fn research(question: String) -> String; +/// Record the final report. +#[tool] +fn submit(report: String) -> String { + report +} + +/// Compile a report; call submit when done. +#[agent(tools(shout, submit), stop_tools(submit), max_turns = 4)] +fn report(question: String) -> String; + +/// Answer immediately — one turn, fail when exhausted. +#[agent(tools(shout), max_turns = 1)] +fn quick(question: String) -> String; + #[dspy_rs::Schema] #[derive(Debug)] pub struct AOut { @@ -162,7 +175,8 @@ async fn agent_runs_in_module_and_standalone() { "the host tool executed inside the loop: {events:?}" ); - // Standalone: the same fn, static-lane tool loop, same tool binding. + // Standalone: the same fn, the same 1-node AgentLoop program, same tool + // binding — and the same attribute options (see the dedicated tests below). client.push_response(tool_call("shout", json!({"text": "again"}))); client.push_response(text_fields(&[("research", "LOUDER")])); let predicted = research("standalone?".to_string()) @@ -170,3 +184,52 @@ async fn agent_runs_in_module_and_standalone() { .expect("standalone agent call succeeds"); assert_eq!(predicted.research, "LOUDER"); } + +// --------------------------------------------------------------------------- +// Standalone path honors the `#[agent(...)]` options (phase 4) +// --------------------------------------------------------------------------- + +#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] +#[tokio::test] +async fn standalone_honors_stop_tools() { + let _lock = SETTINGS_LOCK.lock().await; + let (lm, _client) = make_test_lm(vec![ + // A `submit` call must end the loop with its args as the raw final + // output — no tool execution, no second LM turn (the queue has none). + tool_call("submit", json!({"report": "FINAL"})), + ]) + .await; + configure(lm); + + let predicted = report("summarize".to_string()) + .await + .expect("stop tool ends the loop with its args as the output"); + assert_eq!(predicted.report, "FINAL"); +} + +#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] +#[tokio::test] +async fn standalone_honors_max_turns() { + let _lock = SETTINGS_LOCK.lock().await; + let (lm, _client) = make_test_lm(vec![ + // One tool turn spends the whole `max_turns = 1` allowance; the old + // ignored-options path would take a second LM turn and fail on the + // empty response queue instead. + tool_call("shout", json!({"text": "stall"})), + ]) + .await; + configure(lm); + + let err = quick("now".to_string()) + .await + .expect_err("one tool turn exhausts max_turns = 1"); + let message = format!("{err:?}"); + assert!( + message.contains("budget exhausted"), + "max_turns bounded the loop: {message}" + ); + assert!( + !message.contains("queue is empty"), + "the loop must not take a second LM turn: {message}" + ); +} diff --git a/crates/dspy-rs/tests/test_module_ext.rs b/crates/dspy-rs/tests/test_module_ext.rs deleted file mode 100644 index f6d0693e..00000000 --- a/crates/dspy-rs/tests/test_module_ext.rs +++ /dev/null @@ -1,136 +0,0 @@ -use dspy_rs::{BamlType, CallMetadata, Module, ModuleExt, ParseError, PredictError, Predicted}; - -struct MaybeFails; - -#[derive(Clone, Debug, PartialEq, Eq)] -#[BamlType] -struct IntPayload { - value: i32, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -#[BamlType] -struct TextPayload { - value: String, -} - -impl Module for MaybeFails { - type Input = IntPayload; - type Output = IntPayload; - - async fn forward(&self, input: Self::Input) -> Result, PredictError> { - let input_value = input.value; - let metadata = CallMetadata::new( - format!("raw:{input_value}"), - dspy_rs::LmUsage::default(), - Vec::new(), - Vec::new(), - Some(dspy_rs::SpanId(input_value.max(0) as u32)), - indexmap::IndexMap::new(), - ); - - if input_value < 0 { - Err(PredictError::Parse { - source: ParseError::MissingField { - field: "value".to_string(), - raw_response: format!("raw:{input_value}"), - }, - raw_response: format!("raw:{input_value}"), - lm_usage: dspy_rs::LmUsage::default(), - }) - } else { - Ok(Predicted::new( - IntPayload { - value: input_value * 2, - }, - metadata, - )) - } - } -} - -#[expect( - clippy::result_large_err, - reason = "Tests ModuleExt::and_then using the crate's public PredictError type." -)] -fn transform_int_payload(value: IntPayload) -> Result { - if value.value >= 4 { - Ok(TextPayload { - value: value.value.to_string(), - }) - } else { - Err(PredictError::Parse { - source: ParseError::MissingField { - field: "transformed".to_string(), - raw_response: "transform".to_string(), - }, - raw_response: "transform".to_string(), - lm_usage: dspy_rs::LmUsage::default(), - }) - } -} - -#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] -#[tokio::test] -async fn map_transforms_success_and_preserves_metadata() { - let mapped = MaybeFails.map(|value| TextPayload { - value: format!("v={}", value.value), - }); - - let success = mapped.call(IntPayload { value: 3 }).await.unwrap(); - assert_eq!(success.metadata().raw_response, "raw:3"); - assert_eq!( - success.into_inner(), - TextPayload { - value: "v=6".to_string() - } - ); - - let err = mapped - .call(IntPayload { value: -7 }) - .await - .expect_err("failure expected"); - match err { - PredictError::Parse { - source: ParseError::MissingField { field, .. }, - raw_response, - .. - } => { - assert_eq!(field, "value"); - assert_eq!(raw_response, "raw:-7"); - } - other => panic!("unexpected error: {other:?}"), - } -} - -#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] -#[tokio::test] -async fn and_then_applies_fallible_transform_and_keeps_metadata() { - let module = MaybeFails - .and_then(transform_int_payload as fn(IntPayload) -> Result); - - let success = module.call(IntPayload { value: 3 }).await.unwrap(); - assert_eq!(success.metadata().raw_response, "raw:3"); - assert_eq!( - success.into_inner(), - TextPayload { - value: "6".to_string() - } - ); - - let err = module - .call(IntPayload { value: 1 }) - .await - .expect_err("transform error expected"); - match err { - PredictError::Parse { - source: ParseError::MissingField { field, .. }, - raw_response, - .. - } => { - assert_eq!(field, "transformed"); - assert_eq!(raw_response, "transform"); - } - other => panic!("unexpected error: {other:?}"), - } -} diff --git a/crates/dspy-rs/tests/test_module_facet_shapes.rs b/crates/dspy-rs/tests/test_module_facet_shapes.rs index ab89fa82..faa058e2 100644 --- a/crates/dspy-rs/tests/test_module_facet_shapes.rs +++ b/crates/dspy-rs/tests/test_module_facet_shapes.rs @@ -1,5 +1,4 @@ -use dspy_rs::{ChainOfThought, Facet, ModuleExt, PredictError, ReAct, Signature}; -use facet::{self, Type, UserType}; +use dspy_rs::{ChainOfThought, Facet, Signature}; #[derive(Signature, Clone, Debug, facet::Facet)] #[facet(crate = facet)] @@ -15,46 +14,6 @@ fn shape_of Facet<'a>>(_: &T) -> &'static facet::Shape { >::SHAPE } -fn struct_fields(shape: &'static facet::Shape) -> &'static [facet::Field] { - match shape.ty { - Type::User(UserType::Struct(struct_ty)) => struct_ty.fields, - _ => panic!( - "expected struct shape for {}, got {:?}", - shape.type_identifier, shape.ty - ), - } -} - -fn find_field(shape: &'static facet::Shape, name: &str) -> &'static facet::Field { - struct_fields(shape) - .iter() - .find(|field| field.name == name) - .unwrap_or_else(|| { - let available = struct_fields(shape) - .iter() - .map(|field| field.name) - .collect::>(); - panic!( - "field `{name}` not found on shape `{}` (available: {:?})", - shape.type_identifier, available - ) - }) -} - -fn drop_reasoning(output: dspy_rs::WithReasoning) -> QAOutput { - output.inner -} - -#[expect( - clippy::result_large_err, - reason = "Test verifies ModuleExt::and_then shape with the crate's public PredictError." -)] -fn drop_reasoning_checked( - output: dspy_rs::WithReasoning, -) -> Result { - Ok(output.inner) -} - #[test] fn chain_of_thought_is_a_predict_leaf() { // `ChainOfThought` is an alias for `Predict>`, @@ -64,45 +23,3 @@ fn chain_of_thought_is_a_predict_leaf() { assert_eq!(shape.type_identifier, "Predict"); } - -#[test] -fn react_shape_exposes_action_and_extract_and_skips_non_parameters() { - let module = ReAct::::new(); - let shape = shape_of(&module); - - let action = find_field(shape, "action"); - let extract = find_field(shape, "extract"); - assert!(!action.should_skip_deserializing()); - assert!(!extract.should_skip_deserializing()); - assert_eq!(action.shape().type_identifier, "Predict"); - assert_eq!(extract.shape().type_identifier, "Predict"); - - let tools = find_field(shape, "tools"); - let max_steps = find_field(shape, "max_steps"); - assert!(tools.should_skip_deserializing()); - assert!(max_steps.should_skip_deserializing()); -} - -#[test] -fn map_shape_exposes_inner_chain_of_thought_shape() { - let mapped = ChainOfThought::::new() - .map(drop_reasoning as fn(dspy_rs::WithReasoning) -> QAOutput); - let map_shape = shape_of(&mapped); - let inner = find_field(map_shape, "inner"); - - assert!(!inner.should_skip_deserializing()); - assert_eq!(inner.shape().type_identifier, "Predict"); -} - -#[test] -fn and_then_shape_exposes_inner_chain_of_thought_shape() { - let chained = ChainOfThought::::new().and_then( - drop_reasoning_checked - as fn(dspy_rs::WithReasoning) -> Result, - ); - let and_then_shape = shape_of(&chained); - let inner = find_field(and_then_shape, "inner"); - - assert!(!inner.should_skip_deserializing()); - assert_eq!(inner.shape().type_identifier, "Predict"); -} diff --git a/crates/dspy-rs/tests/test_module_forward_all.rs b/crates/dspy-rs/tests/test_module_forward_all.rs index a2376455..973966ca 100644 --- a/crates/dspy-rs/tests/test_module_forward_all.rs +++ b/crates/dspy-rs/tests/test_module_forward_all.rs @@ -1,19 +1,19 @@ use std::time::Duration; -use dspy_rs::{BamlType, CallMetadata, Module, PredictError, Predicted, forward_all}; +use dspy_rs::{Schema, CallMetadata, Module, PredictError, Predicted, forward_all}; use tokio::time::sleep; struct DelayEcho; #[derive(Clone, Debug, PartialEq, Eq)] -#[BamlType] +#[Schema] struct DelayInput { value: i64, delay_ms: i64, } #[derive(Clone, Debug, PartialEq, Eq)] -#[BamlType] +#[Schema] struct DelayOutput { value: i64, } diff --git a/crates/dspy-rs/tests/test_module_macro.rs b/crates/dspy-rs/tests/test_module_macro.rs index ee3ea142..f4933e5c 100644 --- a/crates/dspy-rs/tests/test_module_macro.rs +++ b/crates/dspy-rs/tests/test_module_macro.rs @@ -3,7 +3,6 @@ //! parse, two projections. Covers: metadata, the lowered artifact (extern //! hole, leaf names, Main sig), OPACITY, end-to-end execution through the //! interpreter, and ambient-overlay mutation of a step instruction. -#![cfg(feature = "ir")] use std::sync::{Arc, LazyLock}; diff --git a/crates/dspy-rs/tests/test_optimizer_named_parameters_integration.rs b/crates/dspy-rs/tests/test_optimizer_named_parameters_integration.rs index 35f581d5..6c6c82a5 100644 --- a/crates/dspy-rs/tests/test_optimizer_named_parameters_integration.rs +++ b/crates/dspy-rs/tests/test_optimizer_named_parameters_integration.rs @@ -1,6 +1,6 @@ use anyhow::Result; use dspy_rs::{ - COPRO, CallMetadata, Eval, Module, Optimizer, Predict, PredictError, Predicted, Signature, + COPRO, CallMetadata, Eval, Module, Predict, PredictError, Predicted, Signature, TypedMetric, }; @@ -13,12 +13,12 @@ struct OptimizerSig { answer: String, } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct InstructionEchoModule { predictor: Predict, } +dspy_rs::predictors!(InstructionEchoModule { predictor }); + impl Module for InstructionEchoModule { type Input = OptimizerSigInput; type Output = OptimizerSigOutput; @@ -83,7 +83,7 @@ async fn optimizer_compile_succeeds_without_public_named_parameter_access() { let optimizer = COPRO::builder().breadth(4).depth(1).build(); optimizer - .compile(&mut module, trainset(), &InstructionLengthMetric) + .compile_module(&mut module, &trainset(), &InstructionLengthMetric) .await .expect("COPRO compile should succeed with internal predictor discovery"); } diff --git a/crates/dspy-rs/tests/test_optimizer_typed_metric.rs b/crates/dspy-rs/tests/test_optimizer_typed_metric.rs index ad114596..bf972178 100644 --- a/crates/dspy-rs/tests/test_optimizer_typed_metric.rs +++ b/crates/dspy-rs/tests/test_optimizer_typed_metric.rs @@ -1,6 +1,6 @@ use anyhow::{Result, anyhow}; use dspy_rs::{ - COPRO, CallMetadata, MIPROv2, Eval, Module, Optimizer, Predict, PredictError, Predicted, + COPRO, CallMetadata, MIPROv2, Eval, Module, Predict, PredictError, Predicted, Signature, TypedMetric, }; use std::collections::HashSet; @@ -15,12 +15,12 @@ struct OptimizerSig { answer: String, } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct InstructionEchoModule { predictor: Predict, } +dspy_rs::predictors!(InstructionEchoModule { predictor }); + impl Module for InstructionEchoModule { type Input = OptimizerSigInput; type Output = OptimizerSigOutput; @@ -113,7 +113,7 @@ async fn copro_compile_uses_typed_metric_predictions() { let optimizer = COPRO::builder().breadth(3).depth(1).build(); optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect("COPRO compile should succeed on typed metric"); @@ -147,7 +147,7 @@ async fn mipro_compile_uses_typed_metric_predictions() { .build(); optimizer - .compile(&mut module, trainset(), &metric) + .compile_module(&mut module, &trainset(), &metric) .await .expect("MIPRO compile should succeed on typed metric"); @@ -171,7 +171,7 @@ async fn copro_compile_propagates_metric_errors() { let optimizer = COPRO::builder().breadth(3).depth(1).build(); let err = optimizer - .compile(&mut module, trainset(), &FailingMetric) + .compile_module(&mut module, &trainset(), &FailingMetric) .await .expect_err("COPRO should propagate typed metric errors"); @@ -192,7 +192,7 @@ async fn mipro_compile_propagates_metric_errors() { .build(); let err = optimizer - .compile(&mut module, trainset(), &FailingMetric) + .compile_module(&mut module, &trainset(), &FailingMetric) .await .expect_err("MIPRO should propagate typed metric errors"); diff --git a/crates/dspy-rs/tests/test_pareto.rs b/crates/dspy-rs/tests/test_pareto.rs deleted file mode 100644 index a7c5f7bb..00000000 --- a/crates/dspy-rs/tests/test_pareto.rs +++ /dev/null @@ -1,83 +0,0 @@ -use dspy_rs::optimizer::gepa::GEPACandidate; -use dspy_rs::optimizer::pareto::ParetoFrontier; - -fn make_test_candidate(instruction: &str) -> GEPACandidate { - GEPACandidate { - id: 0, - instruction: instruction.to_string(), - module_name: "test_module".to_string(), - example_scores: Vec::new(), - parent_id: None, - generation: 0, - } -} - -#[test] -fn test_frontier_empty() { - let frontier = ParetoFrontier::new(); - assert!(frontier.is_empty()); - assert_eq!(frontier.len(), 0); - assert!(frontier.sample_proportional_to_coverage().is_none()); -} - -#[test] -fn test_add_first_candidate() { - let mut frontier = ParetoFrontier::new(); - let candidate = make_test_candidate("instruction 1"); - let scores = vec![0.8, 0.7, 0.9]; - - let added = frontier.add_candidate(candidate, &scores); - assert!(added); - assert_eq!(frontier.len(), 1); -} - -#[test] -fn test_pareto_dominance() { - let mut frontier = ParetoFrontier::new(); - - // Add first candidate - wins on example 0 - let candidate1 = make_test_candidate("instruction 1"); - frontier.add_candidate(candidate1, &[0.9, 0.5, 0.5]); - - // Add second candidate - wins on examples 1 and 2 - let candidate2 = make_test_candidate("instruction 2"); - frontier.add_candidate(candidate2, &[0.5, 0.9, 0.9]); - - // Both should be on frontier (complementary strengths) - assert_eq!(frontier.len(), 2); - - // Add dominated candidate - loses on all examples - let candidate3 = make_test_candidate("instruction 3"); - let added = frontier.add_candidate(candidate3, &[0.3, 0.3, 0.3]); - - // Should not be added - assert!(!added); - assert_eq!(frontier.len(), 2); -} - -#[test] -fn test_coverage_weighted_sampling() { - let mut frontier = ParetoFrontier::new(); - - // Add candidates with different coverage - frontier.add_candidate(make_test_candidate("wins on 1"), &[0.9, 0.3, 0.3, 0.3]); - frontier.add_candidate(make_test_candidate("wins on 3"), &[0.3, 0.9, 0.9, 0.9]); - - assert_eq!(frontier.len(), 2); - - // Sample should return one of the candidates - let sampled = frontier.sample_proportional_to_coverage(); - assert!(sampled.is_some()); -} - -#[test] -fn test_statistics() { - let mut frontier = ParetoFrontier::new(); - - frontier.add_candidate(make_test_candidate("c1"), &[0.9, 0.5, 0.5]); - frontier.add_candidate(make_test_candidate("c2"), &[0.5, 0.9, 0.9]); - - let stats = frontier.statistics(); - assert_eq!(stats.num_candidates, 2); - assert_eq!(stats.num_examples_covered, 3); -} diff --git a/crates/dspy-rs/tests/test_phase_upgrades.rs b/crates/dspy-rs/tests/test_phase_upgrades.rs index 7d34e580..3db36352 100644 --- a/crates/dspy-rs/tests/test_phase_upgrades.rs +++ b/crates/dspy-rs/tests/test_phase_upgrades.rs @@ -5,7 +5,7 @@ use anyhow::Result; use dspy_rs::{ CallMetadata, Chat, Demo, Eval, GEPA, LM, LMClient, MIPROv2, Message, Module, ModuleState, - Optimizer, Predict, PredictError, Predicted, Signature, TestCompletionModel, TypedMetric, + Predict, PredictError, Predicted, Signature, TestCompletionModel, TypedMetric, evaluate_trainset, }; use rig::completion::AssistantContent; @@ -105,10 +105,6 @@ async fn capture_records_inputs_edges_and_component_names() { "intermediate" ); - // The second span is chained to the first. - assert!(first.links.is_empty()); - assert_eq!(second.links, vec![first.id]); - // Outputs are recorded. assert_eq!( first.output.as_ref().expect("first output")["answer"], @@ -129,12 +125,12 @@ async fn capture_records_inputs_edges_and_component_names() { // --- ModuleState: save / load round trip ----------------------------------- -#[derive(facet::Facet)] -#[facet(crate = facet)] struct OneStep { predictor: Predict, } +dspy_rs::predictors!(OneStep { predictor }); + #[test] fn module_state_round_trips_through_json() { let mut tuned = OneStep { @@ -263,12 +259,12 @@ impl TypedMetric<(QAInput, QAOutput), OneStepEcho> for FeedbackEcho { } } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct OneStepEcho { predictor: Predict, } +dspy_rs::predictors!(OneStepEcho { predictor }); + impl Module for OneStepEcho { type Input = QAInput; type Output = QAOutput; @@ -322,7 +318,7 @@ async fn gepa_uses_reflection_lm_to_rewrite_instructions() { )]; let report = optimizer - .compile_with_valset(&mut module, trainset, Some(valset), &FeedbackEcho) + .compile_module_with_valset(&mut module, &trainset, Some(&valset), &FeedbackEcho) .await .expect("gepa compile should succeed"); @@ -410,7 +406,7 @@ async fn gepa_reflection_receives_component_subtrace() { )]; optimizer - .compile_with_valset(&mut module, trainset, Some(valset), &FeedbackForPredict) + .compile_module_with_valset(&mut module, &trainset, Some(&valset), &FeedbackForPredict) .await .expect("gepa compile should succeed"); @@ -454,12 +450,12 @@ impl TypedMetric<(QAInput, QAOutput), OneStepPredict> for ExactMatch { } } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct OneStepPredict { predictor: Predict, } +dspy_rs::predictors!(OneStepPredict { predictor }); + impl Module for OneStepPredict { type Input = QAInput; type Output = QAOutput; @@ -498,7 +494,7 @@ async fn mipro_bootstraps_demos_from_successful_traces() { )]; optimizer - .compile(&mut module, trainset, &ExactMatch) + .compile_module(&mut module, &trainset, &ExactMatch) .await .expect("mipro compile should succeed"); diff --git a/crates/dspy-rs/tests/test_predict_conversation.rs b/crates/dspy-rs/tests/test_predict_conversation.rs index 260ca5d0..ba4a2018 100644 --- a/crates/dspy-rs/tests/test_predict_conversation.rs +++ b/crates/dspy-rs/tests/test_predict_conversation.rs @@ -1,6 +1,4 @@ -use dspy_rs::{ - LM, LMClient, Message, Predict, Role, Signature, TestCompletionModel, configure, -}; +use dspy_rs::{LM, LMClient, Message, Predict, Role, Signature, TestCompletionModel, configure}; use rig::completion::{AssistantContent, CompletionRequest}; use rig::message::{Message as RigMessage, Text, UserContent}; use std::sync::LazyLock; @@ -103,6 +101,7 @@ async fn forward_returns_chat_and_prediction() { let chat = predict .build_chat(&input) + .await .expect("build_chat should succeed"); let (predicted, chat) = predict .call_and_parse(chat) @@ -132,6 +131,7 @@ async fn call_and_parse_supports_two_turn_roundtrip() { // First turn: build fresh chat let chat = predict .build_chat(&first_input) + .await .expect("build_chat should succeed"); let (first_predicted, mut chat) = predict .call_and_parse(chat) @@ -157,3 +157,57 @@ async fn call_and_parse_supports_two_turn_roundtrip() { .expect("test model should capture last request"); assert!(request_contains_text(&last_request, caller_follow_up)); } + +/// `build_chat`/`call_and_parse` are thin wrappers over the interpreter's +/// conversation surface: the opening turn must render byte-identically to the +/// typed `call` path — same span `request_hash`, so a trace recorded on one +/// path replays on the other. +#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] +#[tokio::test] +async fn conversation_wrappers_render_identically_to_typed_call() { + let _lock = SETTINGS_LOCK.lock().await; + let responses = vec![ + response_with_fields(&[("answer", "typed")]), + response_with_fields(&[("answer", "conversational")]), + ]; + let _client = configure_test_lm(responses).await; + + // Instruction override + a demo, so the equivalence covers the + // overlay-resolved rendering, not just the bare signature. + let predict = Predict::::builder() + .named("qa") + .instruction("Answer in one word.") + .demo(dspy_rs::Demo::new( + ConversationQAInput { + question: "What is 1+1?".to_string(), + }, + ConversationQAOutput { + answer: "2".to_string(), + }, + )) + .build(); + let input = ConversationQAInput { + question: "What is the capital of France?".to_string(), + }; + + let (typed, typed_trace) = dspy_rs::trace::capture(|| predict.call(input.clone())).await; + assert_eq!(typed.unwrap().into_inner().answer, "typed"); + + let (conversational, chat_trace) = dspy_rs::trace::capture(|| async { + let chat = predict.build_chat(&input).await?; + predict.call_and_parse(chat).await + }) + .await; + let (predicted, chat) = conversational.unwrap(); + assert_eq!(predicted.into_inner().answer, "conversational"); + assert_eq!(chat.messages.last().unwrap().role, Role::Assistant); + + // One span each, same component, byte-identical rendered prompt. + assert_eq!(typed_trace.spans.len(), 1); + assert_eq!(chat_trace.spans.len(), 1); + assert_eq!(typed_trace.components, chat_trace.components); + assert_eq!( + typed_trace.spans[0].request_hash, chat_trace.spans[0].request_hash, + "conversation wrappers must render exactly what the typed call renders" + ); +} diff --git a/crates/dspy-rs/tests/test_predict_conversation_live.rs b/crates/dspy-rs/tests/test_predict_conversation_live.rs index 2bb1b292..6b52252a 100644 --- a/crates/dspy-rs/tests/test_predict_conversation_live.rs +++ b/crates/dspy-rs/tests/test_predict_conversation_live.rs @@ -36,6 +36,7 @@ async fn live_call_and_parse_two_turn_roundtrip() { }; let chat = predict .build_chat(&first_input) + .await .expect("build_chat should succeed"); let (first, mut chat) = predict .call_and_parse(chat) diff --git a/crates/dspy-rs/tests/test_program_engine.rs b/crates/dspy-rs/tests/test_program_engine.rs index 75ca4f2a..22522dba 100644 --- a/crates/dspy-rs/tests/test_program_engine.rs +++ b/crates/dspy-rs/tests/test_program_engine.rs @@ -2,7 +2,6 @@ //! evaluated over ONE shared program through the interpreter with true //! candidate-level parallelism, per-candidate rollout caching keyed on the //! overlay hash, budget gating, and the minibatch gate. -#![cfg(feature = "ir")] use std::sync::Arc; use std::time::Duration; @@ -14,8 +13,8 @@ use dspy_rs::ir::{ }; use dspy_rs::trace::JsonMap; use dspy_rs::{ - Budget, EngineConfig, Eval, EvalOutcome, GateOutcome, LM, LMClient, LMConfig, - ProgramEvalEngine, ProgramEvalOutcome, ProgramMetric, TestCompletionModel, Trace, + BatchEvalOutcome, Budget, Engine, EngineConfig, Eval, EvalOutcome, GateOutcome, LM, LMClient, + LMConfig, OptimizeTarget, ProgramMetric, TestCompletionModel, Trace, }; use rig::completion::AssistantContent; use rig::message::Text; @@ -189,14 +188,15 @@ async fn candidates_evaluate_concurrently_with_per_candidate_outputs_and_cache() let metric = RendezvousMetric { barrier: Barrier::new(2), }; - let mut engine = - ProgramEvalEngine::new(vec![example("q", "any")], &metric, EngineConfig::default()); - let a = engine.register(cand_a); - let b = engine.register(cand_b); + let examples = vec![example("q", "any")]; + let target = OptimizeTarget::program(&interp, &examples, &metric); + let mut engine = Engine::new(EngineConfig::default()); + let a = engine.register_overlay(cand_a); + let b = engine.register_overlay(cand_b); // --- One batch, two candidates, one shared Arc. --- let evals = engine - .evaluate_program_candidates(&interp, &[a, b], None) + .evaluate_many(&target, &[a, b], None) .await .unwrap() .completed() @@ -259,7 +259,7 @@ async fn candidates_evaluate_concurrently_with_per_candidate_outputs_and_cache() // The canned queues are empty, so any live call would error, and the // metric is never re-run for cached rollouts (the barrier stays idle). let cached = engine - .evaluate_program_candidates(&interp, &[a, b], None) + .evaluate_many(&target, &[a, b], None) .await .unwrap() .completed() @@ -293,32 +293,27 @@ async fn budget_gate_runs_nothing_when_the_batch_does_not_fit() { .unwrap(); let metric = ExactMatch; - let mut engine = ProgramEvalEngine::new( - vec![example("q", "any")], - &metric, - EngineConfig { - budget: Budget { - max_metric_calls: Some(1), - ..Budget::unlimited() - }, - ..EngineConfig::default() + let examples = vec![example("q", "any")]; + let target = OptimizeTarget::program(&interp, &examples, &metric); + let mut engine = Engine::new(EngineConfig { + budget: Budget { + max_metric_calls: Some(1), + ..Budget::unlimited() }, - ); + ..EngineConfig::default() + }); let mut cand_a = Overlay::new(&program); cand_a.set_instruction(instr, "A"); let mut cand_b = Overlay::new(&program); cand_b.set_instruction(instr, "B"); - let a = engine.register(cand_a); - let b = engine.register(cand_b); + let a = engine.register_overlay(cand_a); + let b = engine.register_overlay(cand_b); // Two pending rollouts against a one-rollout budget: nothing runs. - let outcome = engine - .evaluate_program_candidates(&interp, &[a, b], None) - .await - .unwrap(); + let outcome = engine.evaluate_many(&target, &[a, b], None).await.unwrap(); assert!(matches!( outcome, - ProgramEvalOutcome::BudgetExhausted { needed: 2 } + BatchEvalOutcome::BudgetExhausted { needed: 2 } )); assert_eq!(engine.spend().metric_calls, 0); assert_eq!(engine.spend().lm_calls, 0); @@ -352,21 +347,19 @@ async fn minibatch_gate_promotes_and_rejects() { .unwrap(); let metric = ExactMatch; - let mut engine = ProgramEvalEngine::new( - vec![example("q0", "ok"), example("q1", "ok")], - &metric, - EngineConfig::default(), - ); + let examples = vec![example("q0", "ok"), example("q1", "ok")]; + let target = OptimizeTarget::program(&interp, &examples, &metric); + let mut engine = Engine::new(EngineConfig::default()); let mut cand_1 = Overlay::new(&program); cand_1.set_instruction(instr, "GATE 1"); let mut cand_2 = Overlay::new(&program); cand_2.set_instruction(instr, "GATE 2"); - let c1 = engine.register(cand_1); - let c2 = engine.register(cand_2); + let c1 = engine.register_overlay(cand_1); + let c2 = engine.register_overlay(cand_2); // Minibatch mean 1.0 > 0.5: promoted to the full set, where the // minibatch example replays from the cache. - match engine.evaluate_gated(&interp, c1, &[0], 0.5).await.unwrap() { + match engine.evaluate_gated(&target, c1, &[0], 0.5).await.unwrap() { GateOutcome::Promoted { minibatch, full } => { assert_eq!(minibatch.mean(), 1.0); assert_eq!(full.rollouts.len(), 2); @@ -377,13 +370,13 @@ async fn minibatch_gate_promotes_and_rejects() { } // Minibatch mean 1.0 <= 2.0: rejected, no full evaluation. - match engine.evaluate_gated(&interp, c2, &[0], 2.0).await.unwrap() { + match engine.evaluate_gated(&target, c2, &[0], 2.0).await.unwrap() { GateOutcome::Rejected { minibatch } => assert_eq!(minibatch.mean(), 1.0), other => panic!("expected rejection, got {other:?}"), } // Single-candidate convenience path reuses the same cache. - match engine.evaluate(&interp, c1, Some(&[0])).await.unwrap() { + match engine.evaluate(&target, c1, Some(&[0])).await.unwrap() { EvalOutcome::Complete(eval) => assert!(eval.rollouts[0].trace.is_none()), other => panic!("expected completion, got {other:?}"), } @@ -407,19 +400,17 @@ async fn stale_candidate_fails_the_batch() { .unwrap(); let metric = ExactMatch; - let mut engine = - ProgramEvalEngine::new(vec![example("q", "any")], &metric, EngineConfig::default()); + let examples = vec![example("q", "any")]; + let target = OptimizeTarget::program(&interp, &examples, &metric); + let mut engine = Engine::new(EngineConfig::default()); // Minted against nothing: the interpreter's base check refuses it. - let stale = engine.register(Overlay::default()); + let stale = engine.register_overlay(Overlay::default()); let err = engine - .evaluate_program_candidates(&interp, &[stale], None) + .evaluate_many(&target, &[stale], None) .await .unwrap_err(); assert!(err.to_string().contains("overlay minted against program")); - let unregistered = engine - .evaluate_program_candidates(&interp, &[7], None) - .await - .unwrap_err(); + let unregistered = engine.evaluate_many(&target, &[7], None).await.unwrap_err(); assert!(unregistered.to_string().contains("not registered")); } diff --git a/crates/dspy-rs/tests/test_public_api_compile_fail.rs b/crates/dspy-rs/tests/test_public_api_compile_fail.rs index e579fd69..61c78b40 100644 --- a/crates/dspy-rs/tests/test_public_api_compile_fail.rs +++ b/crates/dspy-rs/tests/test_public_api_compile_fail.rs @@ -87,7 +87,7 @@ fn optimizer_compile_rejects_wrong_signature_input_type() { "wrong_signature_case", r#" use anyhow::Result; -use dspy_rs::{COPRO, ChainOfThought, Eval, Optimizer, Predicted, Signature, TypedMetric, WithReasoning}; +use dspy_rs::{COPRO, ChainOfThought, Eval, Predicted, Signature, TypedMetric, WithReasoning}; #[derive(Signature, Clone, Debug)] struct RightSig { @@ -123,7 +123,7 @@ fn main() { // Rows for the wrong signature never project into the module's input. let trainset: Vec<(WrongSigInput, WrongSigOutput)> = Vec::new(); let optimizer = COPRO::builder().breadth(1).depth(1).build(); - let _future = optimizer.compile(&mut module, trainset, &Metric); + let _future = optimizer.compile_module(&mut module, &trainset, &Metric); } "#, ); diff --git a/crates/dspy-rs/tests/test_react_builder.rs b/crates/dspy-rs/tests/test_react_builder.rs deleted file mode 100644 index bb872dc2..00000000 --- a/crates/dspy-rs/tests/test_react_builder.rs +++ /dev/null @@ -1,226 +0,0 @@ -use std::sync::LazyLock; -use std::sync::atomic::{AtomicUsize, Ordering}; - -use dspy_rs::{LM, LMClient, Module, ReAct, Signature, TestCompletionModel, configure}; -use rig::completion::AssistantContent; -use rig::message::Text; -use serde_json::Value; -use tokio::sync::Mutex; - -static SETTINGS_LOCK: LazyLock> = LazyLock::new(|| Mutex::new(())); - -fn response_with_fields(fields: &[(&str, &str)]) -> String { - let mut response = String::new(); - for (name, value) in fields { - response.push_str(&format!("[[ ## {name} ## ]]\n{value}\n\n")); - } - response.push_str("[[ ## completed ## ]]\n"); - response -} - -fn text_response(text: impl Into) -> AssistantContent { - AssistantContent::Text(Text { text: text.into() }) -} - -fn parse_calculator_args(args: &str) -> (i64, i64) { - let value: Value = - serde_json::from_str(args).unwrap_or_else(|_| serde_json::json!({ "a": 0, "b": 0 })); - let a = value.get("a").and_then(Value::as_i64).unwrap_or(0); - let b = value.get("b").and_then(Value::as_i64).unwrap_or(0); - (a, b) -} - -async fn configure_test_lm(responses: Vec) { - let client = TestCompletionModel::new(responses.into_iter().map(text_response)); - let lm = temp_env::async_with_vars( - [("OPENAI_API_KEY", Some("test"))], - LM::builder() - .model("openai:gpt-4o-mini".to_string()) - .build(), - ) - .await - .unwrap() - .with_client(LMClient::Test(client)) - .await - .unwrap(); - - configure(lm); -} - -#[derive(Signature, Clone, Debug)] -struct QA { - #[input] - question: String, - - #[output] - answer: String, -} - -#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] -#[tokio::test] -async fn react_builder_executes_multi_tool_calculator_loop_and_extracts_output() { - let _lock = SETTINGS_LOCK.lock().await; - - let action_1 = response_with_fields(&[ - ("thought", "Need to add first"), - ("action", "add"), - ("action_input", "{\"a\":17,\"b\":5}"), - ]); - let action_2 = response_with_fields(&[ - ("thought", "Now multiply the intermediate result"), - ("action", "multiply"), - ("action_input", "{\"a\":22,\"b\":3}"), - ]); - let action_3 = response_with_fields(&[ - ("thought", "Done"), - ("action", "finish"), - ("action_input", "66"), - ]); - let extract = response_with_fields(&[("output", "{\"answer\":\"66\"}")]); - - configure_test_lm(vec![action_1, action_2, action_3, extract]).await; - - let add_calls = std::sync::Arc::new(AtomicUsize::new(0)); - let multiply_calls = std::sync::Arc::new(AtomicUsize::new(0)); - let add_calls_for_tool = add_calls.clone(); - let multiply_calls_for_tool = multiply_calls.clone(); - - let react = ReAct::::builder() - .max_steps(4) - .tool("add", "Adds two integers {a,b}", move |args| { - let add_calls = add_calls_for_tool.clone(); - async move { - add_calls.fetch_add(1, Ordering::SeqCst); - let (a, b) = parse_calculator_args(&args); - (a + b).to_string() - } - }) - .tool("multiply", "Multiplies two integers {a,b}", move |args| { - let multiply_calls = multiply_calls_for_tool.clone(); - async move { - multiply_calls.fetch_add(1, Ordering::SeqCst); - let (a, b) = parse_calculator_args(&args); - (a * b).to_string() - } - }) - .build(); - - let predicted = react - .call(QAInput { - question: "Compute (17 + 5) * 3 using tools.".to_string(), - }) - .await - .expect("react call should succeed"); - - let (result, metadata) = predicted.into_parts(); - assert_eq!( - add_calls.load(Ordering::SeqCst), - 1, - "add tool execution count mismatch; metadata raw_response: {}", - metadata.raw_response - ); - assert_eq!( - multiply_calls.load(Ordering::SeqCst), - 1, - "multiply tool execution count mismatch; metadata raw_response: {}", - metadata.raw_response - ); - let tool_names: Vec = metadata - .tool_calls - .iter() - .map(|call| call.function.name.clone()) - .collect(); - assert!( - tool_names.iter().any(|name| name == "add") - && tool_names.iter().any(|name| name == "multiply"), - "expected add and multiply in tool call trajectory; got {:?}", - tool_names - ); - assert!( - metadata - .tool_executions - .iter() - .any(|entry| entry.contains("Step 1")) - && metadata - .tool_executions - .iter() - .any(|entry| entry.contains("Step 2")) - && metadata - .tool_executions - .iter() - .any(|entry| entry.contains("Step 3")), - "expected full multi-step trajectory in metadata; got {:?}", - metadata.tool_executions - ); - assert!( - metadata - .tool_executions - .iter() - .any(|entry| entry.contains("Observation: 22")) - && metadata - .tool_executions - .iter() - .any(|entry| entry.contains("Observation: 66")), - "expected calculator observations in trajectory; got {:?}", - metadata.tool_executions - ); - - let result: QAOutput = result; - assert_eq!(result.answer, "66"); -} - -#[cfg_attr(miri, ignore = "MIRI has issues with tokio's I/O driver")] -#[tokio::test] -async fn react_unknown_tool_name_does_not_execute_first_tool() { - let _lock = SETTINGS_LOCK.lock().await; - - let action_1 = response_with_fields(&[ - ("thought", "Try a missing tool"), - ("action", "missing_tool"), - ("action_input", "{\"a\":1,\"b\":2}"), - ]); - let action_2 = response_with_fields(&[ - ("thought", "Stop after observing failure"), - ("action", "finish"), - ("action_input", "done"), - ]); - let extract = response_with_fields(&[("output", "{\"answer\":\"done\"}")]); - configure_test_lm(vec![action_1, action_2, extract]).await; - - let add_calls = std::sync::Arc::new(AtomicUsize::new(0)); - let add_calls_for_tool = add_calls.clone(); - - let react = ReAct::::builder() - .max_steps(3) - .tool("add", "Adds two integers {a,b}", move |args| { - let add_calls = add_calls_for_tool.clone(); - async move { - add_calls.fetch_add(1, Ordering::SeqCst); - let (a, b) = parse_calculator_args(&args); - (a + b).to_string() - } - }) - .build(); - - let predicted = react - .call(QAInput { - question: "Call a tool that does not exist.".to_string(), - }) - .await - .expect("react call should succeed"); - let (_, metadata) = predicted.into_parts(); - - assert_eq!( - add_calls.load(Ordering::SeqCst), - 0, - "unknown tool actions should not run arbitrary registered tools" - ); - assert!( - metadata - .tool_executions - .iter() - .any(|entry| entry.contains("tool_not_found: missing_tool")), - "trajectory should record missing-tool observation; got {:?}", - metadata.tool_executions - ); -} diff --git a/crates/dspy-rs/tests/test_signature_schema.rs b/crates/dspy-rs/tests/test_signature_schema.rs index dc15b369..d3102549 100644 --- a/crates/dspy-rs/tests/test_signature_schema.rs +++ b/crates/dspy-rs/tests/test_signature_schema.rs @@ -1,13 +1,13 @@ -use dspy_rs::{BamlType, Schema, Signature, SignatureSchema}; +use dspy_rs::{Schema, Signature, SignatureSchema}; #[derive(Clone, Debug)] -#[BamlType] +#[Schema] struct DetailInput { note: String, } #[derive(Clone, Debug)] -#[BamlType] +#[Schema] struct DetailOutput { answer: String, } diff --git a/crates/dspy-rs/tests/test_simba.rs b/crates/dspy-rs/tests/test_simba.rs index 928c713f..d1352352 100644 --- a/crates/dspy-rs/tests/test_simba.rs +++ b/crates/dspy-rs/tests/test_simba.rs @@ -4,7 +4,7 @@ use anyhow::Result; use dspy_rs::{ - Eval, LM, LMClient, Module, ModuleState, Optimizer, Predict, PredictError, Predicted, SIMBA, + Eval, LM, LMClient, Module, ModuleState, Predict, PredictError, Predicted, SIMBA, Signature, SimbaMove, TestCompletionModel, TypedMetric, }; use rig::completion::AssistantContent; @@ -20,12 +20,12 @@ struct SimbaSig { answer: String, } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct SimbaModule { predictor: Predict, } +dspy_rs::predictors!(SimbaModule { predictor }); + impl Module for SimbaModule { type Input = SimbaSigInput; type Output = SimbaSigOutput; @@ -123,7 +123,7 @@ async fn append_demo_move_is_harvested_gated_and_installed() { .build(); let report = simba - .compile(&mut module, trainset(3), &ExactMatch) + .compile_module(&mut module, &trainset(3), &ExactMatch) .await .expect("SIMBA should succeed on canned responses"); @@ -179,7 +179,7 @@ async fn append_rule_move_uses_the_reflection_lm() { .build(); let report = simba - .compile(&mut module, trainset(2), &ExactMatch) + .compile_module(&mut module, &trainset(2), &ExactMatch) .await .unwrap(); @@ -222,7 +222,7 @@ async fn append_rule_falls_back_to_metric_feedback_without_prompt_model() { .build(); let report = simba - .compile(&mut module, trainset(2), &ExactMatch) + .compile_module(&mut module, &trainset(2), &ExactMatch) .await .unwrap(); @@ -262,7 +262,7 @@ async fn gate_rejection_leaves_the_module_untouched() { .build(); let report = simba - .compile(&mut module, trainset(2), &ExactMatch) + .compile_module(&mut module, &trainset(2), &ExactMatch) .await .unwrap(); @@ -300,7 +300,7 @@ async fn budget_stops_the_ascent_cleanly() { .build(); let report = simba - .compile(&mut module, trainset(2), &ExactMatch) + .compile_module(&mut module, &trainset(2), &ExactMatch) .await .unwrap(); diff --git a/crates/dspy-rs/tests/test_structural.rs b/crates/dspy-rs/tests/test_structural.rs new file mode 100644 index 00000000..669de1b8 --- /dev/null +++ b/crates/dspy-rs/tests/test_structural.rs @@ -0,0 +1,464 @@ +//! Structural (RFC 0004 §6): the sixth strategy — LM-guided edits over the +//! graph-edit calculus. The loop applies a chosen edit via `Program::edited`, +//! migrates the incumbent overlay with `migrate_overlay`, gates the child +//! against the parent on a shared minibatch, and degrades gracefully on every +//! rejection path (edit fails to apply, reflection reply doesn't parse, +//! budget exhausted). + +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use anyhow::Result; +use dspy_rs::ir::{ + self, DemoRow, Edit, FieldType as T, Interpreter, Node, Overlay, ParamValue, Program, + ProgramBuilder, RuntimeEnv, SignatureDef, +}; +use dspy_rs::trace::JsonMap; +use dspy_rs::{ + Eval, LM, LMClient, LMConfig, ProgramMetric, Structural, TestCompletionModel, Trace, +}; +use rig::completion::AssistantContent; +use rig::message::Text; +use serde_json::json; + +fn fields(pairs: &[(&str, &str)]) -> String { + let mut out = String::new(); + for (name, value) in pairs { + out.push_str(&format!("[[ ## {name} ## ]]\n{value}\n\n")); + } + out.push_str("[[ ## completed ## ]]\n"); + out +} + +fn text(content: impl Into) -> AssistantContent { + AssistantContent::Text(Text { + text: content.into(), + }) +} + +async fn canned_lm(responses: Vec) -> Arc { + Arc::new(canned_lm_owned(responses).await) +} + +async fn canned_lm_owned(responses: Vec) -> LM { + let client = TestCompletionModel::new(responses); + temp_env::async_with_vars( + [("OPENAI_API_KEY", Some("test"))], + LM::builder() + .model("openai:gpt-4o-mini".to_string()) + .build(), + ) + .await + .unwrap() + .with_client(LMClient::Test(client)) + .await + .unwrap() +} + +fn config() -> LMConfig { + LMConfig { + model: "openai:gpt-4o-mini".to_string(), + ..LMConfig::default() + } +} + +fn obj(pairs: &[(&str, serde_json::Value)]) -> JsonMap { + pairs + .iter() + .map(|(k, v)| (k.to_string(), v.clone())) + .collect() +} + +fn example(question: &str, answer: &str) -> DemoRow { + DemoRow { + input: obj(&[("question", json!(question))]), + output: obj(&[("answer", json!(answer))]), + } +} + +/// question → answerer (QA predict) → answer. The single-leaf menu, in +/// order: AugmentSig (0), SwapToAgent (1), WrapRetry (2), Remove (3). +fn qa_program() -> Program { + let mut b = ProgramBuilder::new("structural"); + b.model("m", config()); + let qa = b.sig( + SignatureDef::build("QA") + .instruction("Answer the question.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let node = ir::predict("answerer", qa).bind("question", ir::input("question")); + b.main( + qa, + ir::seq([node]).out("answer", ir::out("answerer", "answer")), + ) + .unwrap() +} + +/// Same pipeline, but the leaf is already `cot`-augmented — `AugmentSig` +/// drops out of the menu, and every remaining option is safe for the canned +/// two-field response. +fn cot_program() -> Program { + let mut b = ProgramBuilder::new("structural_cot"); + b.model("m", config()); + let main_sig = b.sig( + SignatureDef::build("Main") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let qa = b.sig( + SignatureDef::build("QA") + .instruction("Answer the question.") + .input("question", T::String) + .output("answer", T::String) + .finish() + .unwrap(), + ); + let node = ir::cot("answerer", qa).bind("question", ir::input("question")); + b.main( + main_sig, + ir::seq([node]).out("answer", ir::out("answerer", "answer")), + ) + .unwrap() +} + +/// Plain exact-match on the labeled answer, with feedback text. +struct ExactMatch; + +impl ProgramMetric for ExactMatch { + async fn evaluate( + &self, + example: &DemoRow, + output: &JsonMap, + _trace: Option<&Trace>, + ) -> Result { + let expected = example.output.get("answer"); + let got = output.get("answer"); + if got == expected { + Ok(Eval::with_feedback(1.0, "correct")) + } else { + Ok(Eval::with_feedback( + 0.0, + format!("expected {expected:?}, got {got:?}"), + )) + } + } +} + +async fn load(program: Program, lm: Arc) -> Interpreter { + Interpreter::load(program, RuntimeEnv::new().bind_model("m", lm)) + .await + .unwrap() +} + +// --------------------------------------------------------------------------- +// The loop applies an edit and keeps the child only when it wins +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn accepts_the_child_when_it_wins_the_shared_minibatch() { + // Parent answers wrong (score 0); the CoT-augmented child answers right. + let parent_lm = canned_lm(vec![text(fields(&[("answer", "wrong")]))]).await; + let interp = load(qa_program(), parent_lm).await; + let parent_hash = interp.program().meta.program_hash; + + let child_lm = canned_lm(vec![text(fields(&[ + ("reasoning", "think"), + ("answer", "right"), + ]))]) + .await; + // Reflection chooses option 0: AugmentSig on `answerer`. + let reflection = canned_lm_owned(vec![text(fields(&[("chosen_option", "0")]))]).await; + + let examples = vec![example("q", "right")]; + let structural = Structural::builder() + .num_iterations(1) + .minibatch_size(1) + .prompt_model(reflection) + .seed(0) + .build(); + + let report = structural + .compile_program(&interp, &examples, &ExactMatch, move || { + RuntimeEnv::new().bind_model("m", child_lm.clone()) + }) + .await + .unwrap(); + + // The child won: the winner is a new program with the reasoning field. + assert_ne!(report.program.meta.program_hash, parent_hash); + assert_eq!(report.accepted, 1); + assert_eq!(report.rejected, 0); + assert_eq!(report.edits.len(), 1); + assert!(matches!(report.edits[0], Edit::AugmentSig { .. })); + assert_eq!(report.baseline_score, 0.0); + assert_eq!(report.final_score, 1.0); + + let leaf = report.program.leaf_id("answerer").expect("leaf survives"); + let Node::Predict(node) = &report.program.nodes[leaf] else { + panic!("answerer is still a predict leaf"); + }; + assert!( + report.program.sigs[node.sig] + .outputs + .iter() + .any(|f| &*f.name == "reasoning"), + "the accepted edit augmented the leaf signature" + ); + + // Lineage points back at the parent. + let lineage = report.program.meta.lineage.as_ref().unwrap(); + assert_eq!( + lineage.parent.as_deref(), + Some(format!("{parent_hash:016x}").as_str()) + ); + + let step = &report.steps[0]; + assert!(step.accepted); + assert_eq!(step.parent_hash, parent_hash); + assert_eq!(step.parent_minibatch_score, 0.0); + assert_eq!(step.child_minibatch_score, Some(1.0)); + assert_eq!(step.full_score, Some(1.0)); +} + +#[tokio::test] +async fn rejects_the_child_when_it_loses_the_shared_minibatch() { + // Parent answers right (score 1); the child answers wrong. + let parent_lm = canned_lm(vec![text(fields(&[("answer", "right")]))]).await; + let interp = load(qa_program(), parent_lm).await; + let parent_hash = interp.program().meta.program_hash; + + let child_lm = canned_lm(vec![text(fields(&[ + ("reasoning", "hmm"), + ("answer", "wrong"), + ]))]) + .await; + let reflection = canned_lm_owned(vec![text(fields(&[("chosen_option", "0")]))]).await; + + let examples = vec![example("q", "right")]; + let structural = Structural::builder() + .num_iterations(1) + .minibatch_size(1) + .prompt_model(reflection) + .seed(0) + .build(); + + let report = structural + .compile_program(&interp, &examples, &ExactMatch, move || { + RuntimeEnv::new().bind_model("m", child_lm.clone()) + }) + .await + .unwrap(); + + // The parent stays the incumbent. + assert_eq!(report.program.meta.program_hash, parent_hash); + assert_eq!(report.accepted, 0); + assert_eq!(report.rejected, 1); + assert!(report.edits.is_empty()); + assert_eq!(report.baseline_score, 1.0); + assert_eq!(report.final_score, 1.0); + assert_eq!(report.overlay.base, parent_hash); + + let step = &report.steps[0]; + assert!(!step.accepted); + assert_eq!(step.child_minibatch_score, Some(0.0)); + assert_eq!(step.full_score, None); +} + +// --------------------------------------------------------------------------- +// Overlay migration preserves tuned values for surviving nodes +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn migrates_the_tuned_overlay_onto_the_winning_child() { + let parent_lm = canned_lm(vec![text(fields(&[("answer", "wrong")]))]).await; + let interp = load(qa_program(), parent_lm).await; + let program = Arc::clone(interp.program()); + let parent_hash = program.meta.program_hash; + + // A prior value-level optimizer tuned the instruction. + let slot = program + .slot_of::("answerer.instruction") + .unwrap(); + let mut tuned = Overlay::new(&program); + tuned.set_instruction(slot, "TUNED: answer tersely."); + + let child_lm = canned_lm(vec![text(fields(&[ + ("reasoning", "think"), + ("answer", "right"), + ]))]) + .await; + let reflection = canned_lm_owned(vec![text(fields(&[("chosen_option", "0")]))]).await; + + let examples = vec![example("q", "right")]; + let structural = Structural::builder() + .num_iterations(1) + .minibatch_size(1) + .prompt_model(reflection) + .seed(0) + .build(); + + let report = structural + .compile_program_with_overlay(&interp, Some(tuned), &examples, &ExactMatch, move || { + RuntimeEnv::new().bind_model("m", child_lm.clone()) + }) + .await + .unwrap(); + + // AugmentSig widens outputs, so the tuned instruction survives, re-minted + // against the winning child. + assert_ne!(report.program.meta.program_hash, parent_hash); + assert_eq!(report.overlay.base, report.program.meta.program_hash); + let id = report.program.param_id("answerer.instruction").unwrap(); + assert_eq!( + report.overlay.get(id), + Some(&ParamValue::Instruction { + text: "TUNED: answer tersely.".to_string() + }) + ); +} + +// --------------------------------------------------------------------------- +// Rejection paths never crash the run +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn edit_that_fails_validation_is_recorded_and_skipped() { + // Option 3 is Remove{answerer}: the program's out binding still + // references the removed leaf, so `edited()` refuses (validate.rs's + // error) and the generation is skipped without loading a child. + let parent_lm = canned_lm(vec![text(fields(&[("answer", "right")]))]).await; + let interp = load(qa_program(), parent_lm).await; + let parent_hash = interp.program().meta.program_hash; + + let reflection = canned_lm_owned(vec![text(fields(&[("chosen_option", "3")]))]).await; + let loads = Arc::new(AtomicUsize::new(0)); + let loads_seen = Arc::clone(&loads); + + let examples = vec![example("q", "right")]; + let structural = Structural::builder() + .num_iterations(1) + .minibatch_size(1) + .prompt_model(reflection) + .seed(0) + .build(); + + let report = structural + .compile_program(&interp, &examples, &ExactMatch, move || { + loads.fetch_add(1, Ordering::SeqCst); + RuntimeEnv::new() + }) + .await + .unwrap(); + + assert_eq!(report.program.meta.program_hash, parent_hash); + assert_eq!(report.accepted, 0); + assert_eq!(report.rejected, 1); + assert_eq!(loads_seen.load(Ordering::SeqCst), 0, "no child was loaded"); + + let step = &report.steps[0]; + assert!(matches!(step.edit, Edit::Remove { .. })); + assert!(!step.accepted); + assert_eq!(step.child_minibatch_score, None); + let rejection = step.rejection.as_deref().expect("rejection recorded"); + assert!(rejection.contains("edit failed"), "got: {rejection}"); +} + +#[tokio::test] +async fn unparseable_reflection_reply_falls_back_without_crashing() { + // The leaf is already cot-augmented, so every menu option is safe for + // the canned two-field response — whatever the seeded fallback picks + // (SwapToAgent, WrapRetry, or a validation-rejected Remove), the run + // completes. + let parent_lm = canned_lm(vec![text(fields(&[ + ("reasoning", "base"), + ("answer", "right"), + ]))]) + .await; + let interp = load(cot_program(), parent_lm).await; + + let child_lm = canned_lm(vec![ + text(fields(&[("reasoning", "child"), ("answer", "right")])), + text(fields(&[("reasoning", "child"), ("answer", "right")])), + ]) + .await; + let reflection = canned_lm_owned(vec![text(fields(&[( + "chosen_option", + "definitely the CoT one", + )]))]) + .await; + + let examples = vec![example("q", "right")]; + let structural = Structural::builder() + .num_iterations(1) + .minibatch_size(1) + .prompt_model(reflection) + .seed(7) + .build(); + + let report = structural + .compile_program(&interp, &examples, &ExactMatch, move || { + RuntimeEnv::new().bind_model("m", child_lm.clone()) + }) + .await + .expect("fallback choice keeps the run alive"); + + assert_eq!(report.steps.len(), 1); + assert_eq!(report.baseline_score, 1.0); + // The reflection call was charged against the budget either way. + assert!(report.spend.lm_calls >= 2); +} + +// --------------------------------------------------------------------------- +// Budget bounds +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn budget_too_small_for_the_baseline_errors_cleanly() { + let parent_lm = canned_lm(vec![]).await; + let interp = load(qa_program(), parent_lm).await; + + let examples = vec![example("q0", "right"), example("q1", "right")]; + let structural = Structural::builder().max_rollouts(1).build(); + + let err = structural + .compile_program(&interp, &examples, &ExactMatch, RuntimeEnv::new) + .await + .expect_err("baseline needs 2 rollouts against a cap of 1"); + assert!(err.to_string().contains("budget too small"), "got: {err}"); +} + +#[tokio::test] +async fn exhausted_budget_stops_before_proposing() { + // The cap fits exactly the baseline pass: the loop breaks before + // spending a reflection call or scoring any child. + let parent_lm = canned_lm(vec![ + text(fields(&[("answer", "right")])), + text(fields(&[("answer", "right")])), + ]) + .await; + let interp = load(qa_program(), parent_lm).await; + let parent_hash = interp.program().meta.program_hash; + + let examples = vec![example("q0", "right"), example("q1", "right")]; + let structural = Structural::builder() + .num_iterations(4) + .minibatch_size(1) + .max_rollouts(2) + .seed(0) + .build(); + + let report = structural + .compile_program(&interp, &examples, &ExactMatch, RuntimeEnv::new) + .await + .unwrap(); + + assert_eq!(report.program.meta.program_hash, parent_hash); + assert!(report.steps.is_empty()); + assert_eq!(report.final_score, report.baseline_score); + assert_eq!(report.spend.metric_calls, 2); +} diff --git a/crates/dspy-rs/tests/test_trace_attach.rs b/crates/dspy-rs/tests/test_trace_attach.rs index 2a477c3e..2a79385e 100644 --- a/crates/dspy-rs/tests/test_trace_attach.rs +++ b/crates/dspy-rs/tests/test_trace_attach.rs @@ -1,7 +1,6 @@ //! RFC 0001 §1's reserved `param_ids` column + RFC 0002 §3.3's //! `Trace::attach_program`: joining span components to a program's global //! `ParamId`s — one addressing story across traces, overlays, and slots. -#![cfg(feature = "ir")] use std::collections::BTreeSet; use std::sync::Arc; @@ -258,6 +257,7 @@ async fn attach_program_resolves_the_join_for_every_leaf_span() { "researcher.demos", "researcher.model", "researcher.context", + "researcher.tool_set", ] .into_iter() .map(String::from) @@ -300,20 +300,3 @@ async fn param_ids_survive_the_jsonl_round_trip_and_stay_additive() { assert_eq!(restored.meta.v, trace.meta.v, "no format version bump"); assert_eq!(restored.param_ids, trace.param_ids); } - -#[tokio::test] -async fn absorb_keeps_the_column_parallel() { - let (program, mut attached) = captured_run().await; - attached.attach_program(&program); - - // Absorb an unattached trace with a foreign component. - let mut other = Trace::default(); - other.components.push("foreign".to_string()); - attached.absorb(other); - - assert_eq!(attached.param_ids.len(), attached.components.len()); - let foreign = attached.component_id("foreign").unwrap(); - assert!(attached.param_ids[foreign.0 as usize].is_none()); - let drafter = attached.component_id("drafter").unwrap(); - assert!(attached.param_ids[drafter.0 as usize].is_some()); -} diff --git a/crates/dspy-rs/tests/test_trace_capture.rs b/crates/dspy-rs/tests/test_trace_capture.rs index 546d4681..066bb96c 100644 --- a/crates/dspy-rs/tests/test_trace_capture.rs +++ b/crates/dspy-rs/tests/test_trace_capture.rs @@ -103,11 +103,6 @@ async fn capture_records_spans_with_component_names_and_seq() { assert_eq!(trace.for_component("refiner").count(), 1); assert_eq!(trace.for_component("missing").count(), 0); - // Sequential links: each span points at its predecessor. - assert!(trace.spans[0].links.is_empty()); - assert_eq!(trace.spans[1].links, vec![trace.spans[0].id]); - assert_eq!(trace.spans[2].links, vec![trace.spans[1].id]); - // Prefix interning: both drafter calls share one prefix entry (system + // demo turns); the refiner has its own. let d0 = drafter_spans[0]; diff --git a/crates/dspy-rs/tests/test_trace_export.rs b/crates/dspy-rs/tests/test_trace_export.rs deleted file mode 100644 index e25c2d86..00000000 --- a/crates/dspy-rs/tests/test_trace_export.rs +++ /dev/null @@ -1,321 +0,0 @@ -//! Golden-file coverage for the trace exports (RFC 0001 §4f/§4g): a fully -//! deterministic canned trace (fixed ids, hashes, and timestamps — nothing to -//! normalize) must project to byte-stable JSON. -//! -//! To regenerate the goldens after an intentional format change: -//! `DSRS_BLESS=1 cargo test -p dspy-rs --test test_trace_export` - -use std::collections::BTreeMap; - -use dspy_rs::{ - CompId, Eval, LMConfig, LmUsage, Message, ModelEntry, ModelId, PrefixEntry, PrefixId, Span, - SpanError, SpanErrorKind, SpanEvent, SpanId, Trace, TraceMeta, TraceOutcome, -}; -use serde_json::{Value, json}; - -fn tool_call(id: &str, name: &str, args: Value) -> rig::message::ToolCall { - match rig::completion::AssistantContent::tool_call(id, name, args) { - rig::completion::AssistantContent::ToolCall(tc) => tc, - _ => unreachable!("tool_call constructor returns a ToolCall"), - } -} - -fn usage(prompt: u64, completion: u64) -> LmUsage { - LmUsage { - prompt_tokens: prompt, - completion_tokens: completion, - total_tokens: prompt + completion, - } -} - -fn json_map(value: Value) -> serde_json::Map { - match value { - Value::Object(map) => map, - _ => unreachable!("json_map takes an object literal"), - } -} - -/// A deterministic 3-span rollout: a prefix-carrying drafter call, a -/// tool-looping call, and a failed provider call (no events). -fn canned_trace() -> Trace { - let base_span = Span { - id: SpanId(0), - component: CompId(0), - seq: 0, - parent: None, - links: Vec::new(), - prefix: None, - suffix: Vec::new(), - input: None, - model: ModelId(0), - request_hash: 0, - events: Vec::new(), - raw_output: None, - output: None, - usage: LmUsage::default(), - error: None, - started_at_us: 0, - duration_us: 0, - complete: true, - }; - - let drafter = Span { - id: SpanId(0), - prefix: Some(PrefixId(0)), - suffix: vec![Message::user("[[ ## question ## ]]\nq")], - input: Some(json_map(json!({"question": "q"}))), - request_hash: 0x1111, - events: vec![SpanEvent::Exchange { - message: Message::assistant("[[ ## answer ## ]]\ndraft"), - usage: usage(10, 5), - }], - raw_output: Some("[[ ## answer ## ]]\ndraft".to_string()), - output: Some(json_map(json!({"answer": "draft"}))), - usage: usage(10, 5), - started_at_us: 1_000_100, - duration_us: 200, - ..base_span.clone() - }; - - let tooler = Span { - id: SpanId(1), - component: CompId(1), - links: vec![SpanId(0)], - suffix: vec![Message::user("use the tool on: draft")], - input: Some(json_map(json!({"question": "use the tool on: draft"}))), - request_hash: 0x2222, - events: vec![ - SpanEvent::Exchange { - message: Message::tool_call(tool_call("call_1", "search", json!({"q": "draft"}))), - usage: usage(13, 5), - }, - SpanEvent::ToolRun { - id: "call_1".to_string(), - name: "search".to_string(), - args: json!({"q": "draft"}), - result: "found: relevant doc".to_string(), - duration_us: 300, - error: None, - }, - SpanEvent::Exchange { - message: Message::assistant("[[ ## answer ## ]]\nfinal"), - usage: usage(7, 3), - }, - ], - raw_output: Some("[[ ## answer ## ]]\nfinal".to_string()), - output: Some(json_map(json!({"answer": "final"}))), - usage: usage(20, 8), - started_at_us: 1_000_400, - duration_us: 700, - ..base_span.clone() - }; - - // A failed provider call: no events, no output — excluded from RL - // transitions, exported to OTel with ERROR status. - let failed = Span { - id: SpanId(2), - seq: 1, - links: vec![SpanId(1)], - suffix: vec![Message::user("[[ ## question ## ]]\nretry q")], - input: Some(json_map(json!({"question": "retry q"}))), - request_hash: 0x3333, - error: Some(SpanError { - kind: SpanErrorKind::Lm, - message: "provider unreachable".to_string(), - }), - started_at_us: 1_001_200, - duration_us: 90, - ..base_span - }; - - let config = LMConfig { - base_url: None, - api_key: None, - model: "openai:gpt-4o-mini".to_string(), - temperature: 0.0, - max_tokens: 128, - max_tool_iterations: 4, - max_retries: 0, - retry_base_delay_ms: 1, - cache: false, - }; - - Trace { - meta: TraceMeta { - v: 1, - trace_id: "0123456789abcdef0123456789abcdef".to_string(), - started_at_us: 1_000_000, - candidate_hash: Some(42), - input: Some(json_map(json!({"question": "q"}))), - tags: BTreeMap::from([("optimizer".to_string(), "gepa".to_string())]), - }, - components: vec!["drafter".to_string(), "tooler".to_string()], - models: vec![ModelEntry::from_config(&config)], - prefixes: vec![PrefixEntry { - messages: vec![ - Message::system("You draft answers."), - Message::user("[[ ## question ## ]]\ndemo-q"), - Message::assistant("[[ ## answer ## ]]\ndemo-a"), - ], - }], - spans: vec![drafter, tooler, failed], - outcome: Some(TraceOutcome { - output: Some(json_map(json!({"answer": "final"}))), - error: None, - eval: Some(Eval::with_feedback(0.75, "ok")), - duration_us: 5_000, - }), - ..Trace::default() - } -} - -/// Compares produced JSON against a golden file; `DSRS_BLESS=1` rewrites it. -fn assert_matches_golden(produced: &Value, file_name: &str) { - let path = format!( - "{}/tests/fixtures/{file_name}", - env!("CARGO_MANIFEST_DIR") - ); - if std::env::var("DSRS_BLESS").is_ok() { - let mut pretty = serde_json::to_string_pretty(produced).expect("serialize golden"); - pretty.push('\n'); - std::fs::write(&path, pretty).expect("write golden"); - } - let golden = std::fs::read_to_string(&path) - .unwrap_or_else(|err| panic!("missing golden {path} (run with DSRS_BLESS=1): {err}")); - let golden: Value = serde_json::from_str(&golden).expect("parse golden"); - assert_eq!( - produced, &golden, - "export drifted from {file_name}; if intentional, re-bless with DSRS_BLESS=1" - ); -} - -#[test] -fn rl_rollout_matches_golden() { - let trace = canned_trace(); - let rollout = trace.to_rl_rollout().expect("eval recorded"); - let produced = serde_json::to_value(&rollout).expect("serialize rollout"); - - // Shape sanity before the byte-level golden: reward from the outcome, - // failed span dropped, tool loop flattened into completion messages. - assert_eq!(produced["reward"], json!(0.75)); - assert_eq!(produced["trace_id"], json!("0123456789abcdef0123456789abcdef")); - let transitions = produced["transitions"].as_array().expect("transitions"); - assert_eq!(transitions.len(), 2, "the failed span emits no transition"); - assert_eq!(transitions[0]["component"], json!("drafter")); - assert_eq!( - transitions[0]["messages"].as_array().unwrap().len(), - 4, - "prefix (system + demo pair) ++ suffix" - ); - assert_eq!(transitions[1]["component"], json!("tooler")); - assert_eq!( - transitions[1]["completion"].as_array().unwrap().len(), - 3, - "tool-call turn, tool-result turn, final answer" - ); - assert_eq!(transitions[1]["usage"]["total_tokens"], json!(28)); - assert_eq!(transitions[1]["model"], json!("openai:gpt-4o-mini")); - - assert_matches_golden(&produced, "rl_rollout.golden.json"); - - // The JSONL line is the same object. - let line = rollout.to_json_line().expect("jsonl line"); - assert_eq!( - serde_json::from_str::(&line).expect("parse line"), - produced - ); -} - -#[test] -fn rl_rollout_requires_a_recorded_eval() { - let mut trace = canned_trace(); - trace.outcome.as_mut().unwrap().eval = None; - assert!(trace.to_rl_rollout().is_none()); - trace.outcome = None; - assert!(trace.to_rl_rollout().is_none()); -} - -#[test] -fn otel_export_matches_golden() { - let trace = canned_trace(); - let produced = trace.to_otlp_json("dsrs-test", true); - - // Shape sanity before the byte-level golden. - let spans = &produced["resourceSpans"][0]["scopeSpans"][0]["spans"]; - let spans = spans.as_array().expect("spans array"); - // Root + 3 predict spans + 1 tool child. - assert_eq!(spans.len(), 5); - - let root = &spans[0]; - assert_eq!(root["name"], json!("dsrs.rollout")); - assert_eq!(root["traceId"], json!("0123456789abcdef0123456789abcdef")); - assert!(root.get("parentSpanId").is_none()); - assert_eq!(root["startTimeUnixNano"], json!("1000000000")); - assert_eq!(root["endTimeUnixNano"], json!("1005000000")); - - let drafter = &spans[1]; - assert_eq!(drafter["name"], json!("drafter")); - assert_eq!(drafter["kind"], json!(3), "predict spans are CLIENT"); - assert_eq!(drafter["parentSpanId"], root["spanId"]); - let attrs: Vec<(&str, &Value)> = drafter["attributes"] - .as_array() - .unwrap() - .iter() - .map(|kv| (kv["key"].as_str().unwrap(), &kv["value"])) - .collect(); - assert!(attrs.contains(&( - "gen_ai.request.model", - &json!({"stringValue": "openai:gpt-4o-mini"}) - ))); - assert!(attrs.contains(&("gen_ai.usage.input_tokens", &json!({"intValue": "10"})))); - assert!(attrs.contains(&( - "dsrs.request_hash", - &json!({"stringValue": "0000000000001111"}) - ))); - - let tool = &spans[3]; - assert_eq!(tool["name"], json!("tool:search")); - assert_eq!(tool["parentSpanId"], spans[2]["spanId"]); - assert_eq!(tool["kind"], json!(1), "tool spans are INTERNAL"); - - let failed = &spans[4]; - assert_eq!(failed["status"]["code"], json!(2)); - assert!( - failed["status"]["message"] - .as_str() - .unwrap() - .contains("provider unreachable") - ); - - assert_matches_golden(&produced, "otel_export.golden.json"); -} - -#[test] -fn otel_content_capture_is_opt_in() { - let trace = canned_trace(); - let produced = trace.to_otlp_json("dsrs-test", false); - let spans = produced["resourceSpans"][0]["scopeSpans"][0]["spans"] - .as_array() - .expect("spans array") - .clone(); - - let all_attr_keys: Vec = spans - .iter() - .flat_map(|span| span["attributes"].as_array().unwrap().iter()) - .map(|kv| kv["key"].as_str().unwrap().to_string()) - .collect(); - assert!( - !all_attr_keys.iter().any(|key| key.contains("arguments") || key.contains("result")), - "tool payloads are content" - ); - for span in &spans { - assert!( - span.get("events").is_none(), - "prompt/completion events are content: {span}" - ); - } - // Identity, usage, and timing stay. - assert!(all_attr_keys.iter().any(|key| key == "gen_ai.usage.input_tokens")); - assert!(all_attr_keys.iter().any(|key| key == "dsrs.request_hash")); - assert!(all_attr_keys.iter().any(|key| key == "gen_ai.tool.name")); -} diff --git a/crates/dspy-rs/tests/test_typed_alias.rs b/crates/dspy-rs/tests/test_typed_alias.rs index 55118527..b2c564f3 100644 --- a/crates/dspy-rs/tests/test_typed_alias.rs +++ b/crates/dspy-rs/tests/test_typed_alias.rs @@ -1,4 +1,6 @@ +use dspy_rs::ir::SignatureDef; use dspy_rs::{ChatAdapter, Message, Signature}; +use serde_json::Value; #[derive(Signature, Clone, Debug)] /// Provide an answer using aliases. @@ -12,12 +14,21 @@ struct AliasSignature { answer: String, } +fn json_map(value: &T) -> serde_json::Map { + match serde_json::to_value(value).expect("serializable") { + Value::Object(map) => map, + other => panic!("expected object, got {other:?}"), + } +} + #[test] fn typed_alias_is_used_in_prompt_and_user_message() { let adapter = ChatAdapter; - let system = adapter - .format_system_message_typed::() - .expect("system message"); + let system = adapter.build_system_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + None, + ); assert!(system.contains("[[ ## question_text ## ]]")); assert!(system.contains("[[ ## final_answer ## ]]")); @@ -29,7 +40,7 @@ fn typed_alias_is_used_in_prompt_and_user_message() { let input = AliasSignatureInput { question: "Hello".to_string(), }; - let user = adapter.format_user_message_typed::(&input); + let user = adapter.format_input_def(SignatureDef::of::(), &json_map(&input)); assert!(user.contains("[[ ## question_text ## ]]")); assert!(user.contains("Hello")); assert!(!user.contains("[[ ## question ## ]]")); @@ -39,10 +50,16 @@ fn typed_alias_is_used_in_prompt_and_user_message() { fn typed_alias_parses_output_and_maps_to_rust_name() { let adapter = ChatAdapter; let response = Message::assistant("[[ ## final_answer ## ]]\nHi\n\n[[ ## completed ## ]]"); - let (output, metas) = adapter - .parse_response_typed::(&response) + let (output_map, metas) = adapter + .parse_output_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + &response, + ) .expect("parse response"); + let output: AliasSignatureOutput = + serde_json::from_value(Value::Object(output_map)).expect("typed assembly"); assert_eq!(output.answer, "Hi"); assert!(metas.contains_key("answer")); let meta = metas.get("answer").expect("meta for answer"); diff --git a/crates/dspy-rs/tests/test_typed_prompt_format.rs b/crates/dspy-rs/tests/test_typed_prompt_format.rs index c8e9f7dd..6a261b0a 100644 --- a/crates/dspy-rs/tests/test_typed_prompt_format.rs +++ b/crates/dspy-rs/tests/test_typed_prompt_format.rs @@ -3,10 +3,11 @@ reason = "Signature derive emits multi-field constructors for schema coverage tests." )] -use dspy_rs::{BamlType, ChatAdapter, Signature}; +use dspy_rs::ir::SignatureDef; +use dspy_rs::{Schema, ChatAdapter, Signature}; #[derive(Clone, Debug)] -#[BamlType] +#[Schema] /// A citation reference. struct Citation { /// Document identifier @@ -16,7 +17,7 @@ struct Citation { } #[derive(Clone, Debug)] -#[BamlType] +#[Schema] /// Sentiment classification. enum Sentiment { Positive, @@ -49,10 +50,11 @@ struct ComprehensiveSignature { } fn system_message() -> String { - let adapter = ChatAdapter; - adapter - .format_system_message_typed::() - .expect("system message") + ChatAdapter.build_system_def( + SignatureDef::of::(), + SignatureDef::types_of::(), + None, + ) } fn extract_field_block(message: &str, field_name: &str) -> String { diff --git a/crates/dsrs-cli/Cargo.toml b/crates/dsrs-cli/Cargo.toml index 94449bcf..65fd6b71 100644 --- a/crates/dsrs-cli/Cargo.toml +++ b/crates/dsrs-cli/Cargo.toml @@ -15,17 +15,17 @@ name = "dsrs" path = "src/main.rs" [dependencies] -dspy-rs = { version = "0.7.3", path = "../dspy-rs" } -dsrs-tools = { version = "0.1.0", path = "../dsrs-tools" } -anyhow = "1" +dspy-rs = { workspace = true } +dsrs-tools = { workspace = true } +anyhow = { workspace = true } axum = "0.8" clap = { version = "4", features = ["derive"] } -serde_json = { version = "1", features = ["preserve_order"] } -tokio = { version = "1", features = ["rt-multi-thread", "macros", "net"] } +serde_json = { workspace = true } +tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net"] } [dev-dependencies] -reqwest = { version = "0.13", features = ["json"] } +reqwest = { workspace = true, features = ["json"] } # Same pin as dspy-rs: needed to construct canned AssistantContent responses # for the TestCompletionModel in server integration tests. -rig-core = { git = "https://github.com/0xPlaygrounds/rig", rev = "aee3b8bf6576ce41c9ac1dd82520752a65fa0127" } -tempfile = "3" +rig-core = { workspace = true } +tempfile = { workspace = true } diff --git a/crates/dsrs-macros/Cargo.toml b/crates/dsrs-macros/Cargo.toml index cb216bfa..cbb4f00c 100644 --- a/crates/dsrs-macros/Cargo.toml +++ b/crates/dsrs-macros/Cargo.toml @@ -18,9 +18,10 @@ syn = { version = "2", features = ["full", "visit", "visit-mut"] } quote = "1" proc-macro2 = "1" proc-macro-crate = "3.2" -serde_json = { version = "1.0.143", features = ["preserve_order"] } -minijinja = { git = "https://github.com/boundaryml/minijinja.git", branch = "main", default-features = false, features = ["serde"] } +# Shared .dsrs lexer/structural grammar — include_program!'s build-time syntax gate. +dsrs-syntax = { workspace = true } +minijinja = { workspace = true, features = ["serde"] } [dev-dependencies] -dspy-rs = { path = "../dspy-rs" } +dspy-rs = { workspace = true } trybuild = "1.0.110" diff --git a/crates/dsrs-macros/src/include_program.rs b/crates/dsrs-macros/src/include_program.rs index 6ffe1db0..261962e1 100644 --- a/crates/dsrs-macros/src/include_program.rs +++ b/crates/dsrs-macros/src/include_program.rs @@ -2,16 +2,16 @@ //! //! See [`expand`] and the macro's rustdoc in `lib.rs` for the layering //! contract. The short version: **syntax** is validated here at macro -//! expansion (via [`crate::dsrs_syntax`]); **semantics** are validated by the -//! real parser at first use of the emitted `LazyLock`, and forced at CI time -//! by an emitted `#[cfg(test)]` test — the sqlx-offline analogue. +//! expansion (via [`dsrs_syntax::check`], the shared structural grammar over +//! the same lexer `Program::from_dsrs` uses); **semantics** are validated by +//! the full parser at first use of the emitted `LazyLock`, and forced at CI +//! time by an emitted `#[cfg(test)]` test — the sqlx-offline analogue. use std::path::{Path, PathBuf}; use proc_macro2::Span; use quote::{format_ident, quote}; -use crate::dsrs_syntax; use crate::runtime_path::resolve_dspy_rs_path; /// Parsed macro input: a single string-literal path. @@ -92,8 +92,9 @@ pub(crate) fn expand( )) })?; - // Build-time gate: syntax only. Semantic validation happens through the - // real parser in the emitted runtime/test code below. + // Build-time gate: syntax only, through the shared dsrs-syntax grammar. + // Semantic validation happens through the full parser in the emitted + // runtime/test code below. dsrs_syntax::check(&text).map_err(|e| { err_at(format!( "include_program!(\"{rel}\"): line {}, column {}: {}", diff --git a/crates/dsrs-macros/src/lib.rs b/crates/dsrs-macros/src/lib.rs index 9c2f63da..12c9642d 100644 --- a/crates/dsrs-macros/src/lib.rs +++ b/crates/dsrs-macros/src/lib.rs @@ -9,7 +9,6 @@ use syn::{ visit::Visit, }; -mod dsrs_syntax; mod example_derive; mod include_program; mod module_macro; @@ -53,8 +52,13 @@ pub fn tool(attr: TokenStream, item: TokenStream) -> TokenStream { /// `context(max_history_turns = N, tool_result_max_bytes = N, playbook = "…")`. /// /// Inside `#[module]` bodies this lowers to a first-class `AgentLoop` node. -/// Called standalone it runs the static-lane tool loop (`Predict` + -/// `ToolLoopMode::Auto`) — loop options apply to the lowered form only. +/// Called standalone it executes the same 1-node `AgentLoop` program, with +/// the loop options honored: `max_turns`/`stop_tools`/`until_parse` land in +/// the node's `StopSpec`, `budget` in its `NodeBudget`, `context` in its +/// `ContextPolicy`. The exception is `model`: model refs bind only inside a +/// `#[module]` program, so setting `model = "…"` removes the standalone fn — +/// calling it is a compile error rather than a silent fallback to the +/// globally configured LM. #[proc_macro_attribute] pub fn agent(attr: TokenStream, item: TokenStream) -> TokenStream { tool_agent::expand_agent(attr, item) @@ -103,14 +107,16 @@ pub fn module(attr: TokenStream, item: TokenStream) -> TokenStream { /// /// The full `.dsrs` parser lives in `dspy-rs`, which depends on this macro /// crate — it cannot be called from here without a dependency cycle. The -/// macro therefore validates **syntax only** at build time: the `dsrs 1` -/// pragma, the top-level keyword vocabulary and declaration shapes, balanced -/// delimiters, strings/numbers/code fences — a standalone check against the -/// same surface grammar (see `docs/dsrs-format.md`). Types, dataflow, -/// capability subsets, and everything else semantic are checked by the real -/// parser at first use of `program()`/`try_program()` — and at CI time by the -/// generated test, which is the sqlx-offline analogue: run `cargo test` and a -/// semantically invalid artifact fails the suite even if never executed. +/// macro therefore validates **syntax only** at build time, via the shared +/// `dsrs-syntax` crate: the `dsrs 1` pragma, the top-level keyword +/// vocabulary and declaration shapes, balanced delimiters, +/// strings/numbers/code fences (see `docs/dsrs-format.md`). `dsrs-syntax` +/// also supplies the lexer `Program::from_dsrs` parses with, so the two +/// layers cannot drift. Types, dataflow, capability subsets, and everything +/// else semantic are checked by the full parser at first use of +/// `program()`/`try_program()` — and at CI time by the generated test, which +/// is the sqlx-offline analogue: run `cargo test` and a semantically invalid +/// artifact fails the suite even if never executed. #[proc_macro] pub fn include_program(input: TokenStream) -> TokenStream { let source_dir = proc_macro::Span::call_site() @@ -175,25 +181,13 @@ pub fn derive_signature(input: TokenStream) -> TokenStream { /// /// Expands to `#[derive(facet::Facet, serde::Serialize, serde::Deserialize)]` (plus the /// crate-path attrs), which is all a type needs to satisfy the blanket `Schema` impl. -/// Replaces the old BAML `#[BamlType]` attribute. +/// Replaced the old BAML `#[BamlType]` attribute (the compat alias is gone). #[proc_macro_attribute] #[allow(non_snake_case)] pub fn Schema(_attr: TokenStream, item: TokenStream) -> TokenStream { expand_schema_attr(item) } -/// Backwards-compatible alias for [`macro@Schema`]. -/// -/// The old vendored BAML integration exposed a `#[BamlType]` attribute that derived the -/// type-system plumbing. BAML is gone, but this alias keeps existing signatures/tests that -/// still spell it `#[BamlType]` compiling — it expands to exactly the same facet + serde -/// derives as `#[Schema]`. -#[proc_macro_attribute] -#[allow(non_snake_case)] -pub fn BamlType(_attr: TokenStream, item: TokenStream) -> TokenStream { - expand_schema_attr(item) -} - fn expand_schema_attr(item: TokenStream) -> TokenStream { let input = parse_macro_input!(item as DeriveInput); let runtime = match resolve_dspy_rs_path() { diff --git a/crates/dsrs-macros/src/tool_agent.rs b/crates/dsrs-macros/src/tool_agent.rs index 67cc21be..784223f6 100644 --- a/crates/dsrs-macros/src/tool_agent.rs +++ b/crates/dsrs-macros/src/tool_agent.rs @@ -269,7 +269,14 @@ fn parse_agent_attr(attr: TokenStream2) -> syn::Result { out.model = Some(model_ref_value(&nv.value)?); } Meta::NameValue(nv) if nv.path.is_ident("max_turns") => { - out.max_turns = Some(int_value(&nv.value)?); + let turns: u32 = int_value(&nv.value)?; + if turns == 0 { + return Err(syn::Error::new_spanned( + &nv.value, + "`max_turns` must be > 0 (the loop is mandatory and bounded)", + )); + } + out.max_turns = Some(turns); } Meta::NameValue(nv) if nv.path.is_ident("until_parse") => { let Expr::Lit(ExprLit { @@ -492,6 +499,68 @@ fn expand_agent_inner( None => quote! { ::core::option::Option::None }, }; + // The standalone fn executes the same 1-node `AgentLoop` program the + // `#[module]` lowering produces, so the loop options are honored on both + // paths. The one exception is `model`: model refs bind only inside a + // `#[module]` program's model table, and silently falling back to the + // globally configured LM would misreport which model ran — so with + // `model = "…"` set, no standalone fn is generated at all and calling it + // is a compile error ("expected function, found module"). + let standalone = if attrs.model.is_some() { + quote! {} + } else { + quote! { + #(#doc_attrs)* + /// + /// Standalone calls execute the same 1-node `AgentLoop` program + /// the `#[module]` lowering produces, with the attribute options + /// honored: `max_turns`/`stop_tools`/`until_parse` land in the + /// node's `StopSpec`, `budget` in its `NodeBudget`, and `context` + /// in its `ContextPolicy`. The predictor (and thus the tool set) + /// is built once, on first call. + #vis async fn #fn_name( + #(#arg_names: #arg_types),* + ) -> ::core::result::Result< + #runtime::Predicted<#fn_name::SigOutput>, + #runtime::PredictError, + > { + static __PREDICT: ::std::sync::OnceLock< + ::std::sync::Arc<#runtime::Predict<#fn_name::Sig>>, + > = ::std::sync::OnceLock::new(); + let predictor = __PREDICT.get_or_init(|| { + let step = #fn_name::__dsrs_step(); + let agent = step + .agent + .expect("#[agent] steps always carry agent opts"); + let spec = #runtime::predictors::AgentLoopSpec { + stop_tools: agent + .stop_tools + .iter() + .map(|name| ::std::string::String::from(*name)) + .collect(), + max_turns: agent.max_turns, + until_parse: agent.until_parse, + budget: agent.budget.clone(), + context: agent.context.clone(), + }; + let tools: ::std::vec::Vec< + ::std::sync::Arc, + > = agent.tools.into_iter().map(|t| t.dyn_tool).collect(); + ::std::sync::Arc::new( + #runtime::Predict::<#fn_name::Sig>::builder() + .named(#fn_name_str) + .with_tools(tools) + .with_agent_spec(spec) + .build(), + ) + }); + predictor + .call(#fn_name::SigInput { #(#arg_names),* }) + .await + } + } + }; + Ok(quote! { #vis mod #fn_name { #![allow(non_camel_case_types, unused_imports)] @@ -537,40 +606,7 @@ fn expand_agent_inner( } } - #(#doc_attrs)* - /// - /// Standalone calls run the static-lane tool loop (`Predict` + - /// `ToolLoopMode::Auto`); loop options (`max_turns`, `budget`, - /// `context`) apply to the lowered `AgentLoop` node inside - /// `#[module]` programs. - #vis async fn #fn_name( - #(#arg_names: #arg_types),* - ) -> ::core::result::Result< - #runtime::Predicted<#fn_name::SigOutput>, - #runtime::PredictError, - > { - static __PREDICT: ::std::sync::OnceLock< - ::std::sync::Arc<#runtime::Predict<#fn_name::Sig>>, - > = ::std::sync::OnceLock::new(); - let predictor = __PREDICT.get_or_init(|| { - let step = #fn_name::__dsrs_step(); - let tools: ::std::vec::Vec< - ::std::sync::Arc, - > = step - .agent - .map(|agent| agent.tools.into_iter().map(|t| t.dyn_tool).collect()) - .unwrap_or_default(); - ::std::sync::Arc::new( - #runtime::Predict::<#fn_name::Sig>::builder() - .named(#fn_name_str) - .with_tools(tools) - .build(), - ) - }); - predictor - .call(#fn_name::SigInput { #(#arg_names),* }) - .await - } + #standalone }) } diff --git a/crates/dsrs-macros/tests/ui/agent_model_standalone.rs b/crates/dsrs-macros/tests/ui/agent_model_standalone.rs new file mode 100644 index 00000000..6481fe9b --- /dev/null +++ b/crates/dsrs-macros/tests/ui/agent_model_standalone.rs @@ -0,0 +1,21 @@ +//! `#[agent(model = "…")]` cannot be honored on the standalone call path — +//! model refs bind only inside a `#[module]` program — so no standalone fn is +//! generated and calling one is a compile error, not a silent fallback to the +//! globally configured LM. +use dsrs_macros::{agent, tool}; + +/// Uppercase text. +#[tool] +fn shout(text: String) -> String { + text.to_uppercase() +} + +/// Research the question with the fast model. +#[agent(model = "@fast", tools(shout))] +fn research(question: String) -> String; + +async fn call_it() { + let _ = research("hi".to_string()).await; +} + +fn main() {} diff --git a/crates/dsrs-macros/tests/ui/agent_model_standalone.stderr b/crates/dsrs-macros/tests/ui/agent_model_standalone.stderr new file mode 100644 index 00000000..7d78d996 --- /dev/null +++ b/crates/dsrs-macros/tests/ui/agent_model_standalone.stderr @@ -0,0 +1,5 @@ +error[E0423]: expected function, found module `research` + --> tests/ui/agent_model_standalone.rs:18:13 + | +18 | let _ = research("hi".to_string()).await; + | ^^^^^^^^ not a function diff --git a/crates/dsrs-macros/tests/ui/agent_zero_max_turns.rs b/crates/dsrs-macros/tests/ui/agent_zero_max_turns.rs new file mode 100644 index 00000000..8c0b266f --- /dev/null +++ b/crates/dsrs-macros/tests/ui/agent_zero_max_turns.rs @@ -0,0 +1,15 @@ +//! `max_turns = 0` contradicts the IR's mandatory-and-bounded loop contract +//! and is rejected at macro expansion. +use dsrs_macros::{agent, tool}; + +/// Uppercase text. +#[tool] +fn shout(text: String) -> String { + text.to_uppercase() +} + +/// Research the question. +#[agent(tools(shout), max_turns = 0)] +fn research(question: String) -> String; + +fn main() {} diff --git a/crates/dsrs-macros/tests/ui/agent_zero_max_turns.stderr b/crates/dsrs-macros/tests/ui/agent_zero_max_turns.stderr new file mode 100644 index 00000000..493f915a --- /dev/null +++ b/crates/dsrs-macros/tests/ui/agent_zero_max_turns.stderr @@ -0,0 +1,5 @@ +error: `max_turns` must be > 0 (the loop is mandatory and bounded) + --> tests/ui/agent_zero_max_turns.rs:12:35 + | +12 | #[agent(tools(shout), max_turns = 0)] + | ^ diff --git a/crates/dsrs-syntax/Cargo.toml b/crates/dsrs-syntax/Cargo.toml new file mode 100644 index 00000000..824cf711 --- /dev/null +++ b/crates/dsrs-syntax/Cargo.toml @@ -0,0 +1,14 @@ +[package] +name = "dsrs-syntax" +version = "0.1.0" +edition = "2024" +authors = ["Herumb Shandilya "] +description = "Lexer and structural grammar for the .dsrs text format, shared by dspy-rs and dsrs_macros" +readme = "../../README.md" +documentation = "https://dsrs.herumbshandilya.com" +homepage = "https://dsrs.herumbshandilya.com" +repository = "https://github.com/krypticmouse/DSRs" +license = "Apache-2.0" + +[dependencies] +serde_json = "1.0.140" diff --git a/crates/dspy-rs/src/ir/text/lex.rs b/crates/dsrs-syntax/src/lex.rs similarity index 95% rename from crates/dspy-rs/src/ir/text/lex.rs rename to crates/dsrs-syntax/src/lex.rs index 87c1f58a..8aa27a36 100644 --- a/crates/dspy-rs/src/ir/text/lex.rs +++ b/crates/dsrs-syntax/src/lex.rs @@ -9,17 +9,17 @@ //! parser dispatches on the string (the RFC grammar is keyword-led, so one //! token of lookahead suffices). -use super::ParseError; +use crate::ParseError; /// A source position, 1-based. #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(crate) struct Span { +pub struct Span { pub line: u32, pub col: u32, } #[derive(Clone, Debug, PartialEq)] -pub(crate) enum Tok { +pub enum Tok { /// Identifier or keyword (the parser decides). Ident(String), /// JSON string literal, unescaped. @@ -54,7 +54,7 @@ pub(crate) enum Tok { impl Tok { /// Human-readable token description for error messages. - pub(crate) fn describe(&self) -> String { + pub fn describe(&self) -> String { match self { Tok::Ident(s) => format!("`{s}`"), Tok::Str(_) => "a string".to_string(), @@ -86,7 +86,7 @@ impl Tok { /// One lexed token with its source position and byte extent. #[derive(Clone, Debug)] -pub(crate) struct Lexed { +pub struct Lexed { pub tok: Tok, pub span: Span, /// Byte offset of the first byte of the token (raw-mode scans restart @@ -94,7 +94,7 @@ pub(crate) struct Lexed { pub start: usize, } -pub(crate) struct Lexer<'a> { +pub struct Lexer<'a> { src: &'a str, bytes: &'a [u8], pos: usize, @@ -103,7 +103,7 @@ pub(crate) struct Lexer<'a> { } impl<'a> Lexer<'a> { - pub(crate) fn new(src: &'a str) -> Self { + pub fn new(src: &'a str) -> Self { Self { src, bytes: src.as_bytes(), @@ -157,7 +157,7 @@ impl<'a> Lexer<'a> { } /// Repositions the cursor (used by the parser after raw-mode scans). - pub(crate) fn seek(&mut self, pos: usize, span: Span) { + pub fn seek(&mut self, pos: usize, span: Span) { self.pos = pos; self.line = span.line; self.col = span.col; @@ -165,7 +165,7 @@ impl<'a> Lexer<'a> { /// Line/col of an arbitrary byte offset (computed by rescanning; raw-mode /// scans are rare, so this stays off every hot path). - pub(crate) fn span_at(&self, pos: usize) -> Span { + pub fn span_at(&self, pos: usize) -> Span { let mut line = 1u32; let mut col = 1u32; for &b in &self.bytes[..pos.min(self.bytes.len())] { @@ -179,7 +179,7 @@ impl<'a> Lexer<'a> { Span { line, col } } - pub(crate) fn next_token(&mut self) -> Result { + pub fn next_token(&mut self) -> Result { self.skip_trivia(); let span = self.span(); let start = self.pos; @@ -386,7 +386,7 @@ impl<'a> Lexer<'a> { /// backticks followed by a newline; the source is every byte up to (not /// including) the newline that precedes a line consisting of exactly `k` /// backticks and nothing else. - pub(crate) fn scan_code_fence(&self, from: usize) -> Result<(String, usize), ParseError> { + pub fn scan_code_fence(&self, from: usize) -> Result<(String, usize), ParseError> { let span = self.span_at(from); let bytes = self.bytes; let mut i = from; @@ -454,7 +454,7 @@ impl<'a> Lexer<'a> { /// Scans one raw JSON value starting at byte `from`. Returns the parsed /// value and the byte offset one past its end. - pub(crate) fn scan_json(&self, from: usize) -> Result<(serde_json::Value, usize), ParseError> { + pub fn scan_json(&self, from: usize) -> Result<(serde_json::Value, usize), ParseError> { let span = self.span_at(from); let rest = &self.src[from..]; let mut de = serde_json::Deserializer::from_str(rest).into_iter::(); diff --git a/crates/dsrs-syntax/src/lib.rs b/crates/dsrs-syntax/src/lib.rs new file mode 100644 index 00000000..1128a6ed --- /dev/null +++ b/crates/dsrs-syntax/src/lib.rs @@ -0,0 +1,52 @@ +//! Shared syntax layer for the `.dsrs` text format (RFC 0002 §4). +//! +//! This crate is the single home of the `.dsrs` **lexer** ([`lex`]) and the +//! **structural grammar** check ([`check`]). It exists so both consumers of +//! the grammar read from one source of truth: +//! +//! - `dspy-rs` — the full parser (`Program::from_dsrs`) pulls tokens from +//! [`lex`] and lowers them through its program builder; types, dataflow, +//! and every other semantic rule live there. +//! - `dsrs_macros` — `include_program!` validates artifacts at macro +//! expansion via [`check`], the syntax-only structural pass. The macro +//! crate cannot depend on `dspy-rs` (which depends on it), so this leaf +//! crate is what breaks the cycle. +//! +//! Deliberately a leaf: no dependency on dspy-rs, dsrs_macros, facet, or rig +//! — proc-macro crates can depend on it without cycles. Grammar changes are +//! made **here once** and both frontends pick them up. + +pub mod lex; +mod structure; + +pub use structure::check; + +/// A parse failure with the source position and what was expected — designed +/// to be actionable feedback for a model regenerating the program. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ParseError { + /// 1-based source line. + pub line: u32, + /// 1-based source column (bytes). + pub col: u32, + pub message: String, +} + +impl std::fmt::Display for ParseError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "line {}, column {}: {}", self.line, self.col, self.message) + } +} + +impl std::error::Error for ParseError {} + +impl ParseError { + /// Positions `message` at `span`. + pub fn at(span: lex::Span, message: impl Into) -> Self { + Self { + line: span.line, + col: span.col, + message: message.into(), + } + } +} diff --git a/crates/dsrs-macros/src/dsrs_syntax.rs b/crates/dsrs-syntax/src/structure.rs similarity index 52% rename from crates/dsrs-macros/src/dsrs_syntax.rs rename to crates/dsrs-syntax/src/structure.rs index f4312e19..b17f0f87 100644 --- a/crates/dsrs-macros/src/dsrs_syntax.rs +++ b/crates/dsrs-syntax/src/structure.rs @@ -1,16 +1,16 @@ //! Build-time **syntax-only** validation of `.dsrs` artifacts (RFC 0002 §6.1). //! -//! # Why this module exists (the layering, honestly) +//! # Why this layer exists (the layering, honestly) //! //! `include_program!` wants to validate the artifact at macro expansion, the //! sqlx/prost way. Full validation lives in `dspy-rs` (`Program::from_dsrs` //! lowers through `ProgramBuilder` and runs `Program::validate`), and -//! `dspy-rs` depends on this proc-macro crate — depending on it back would be -//! a dependency cycle. Splitting the parser into a shared syntax crate is the -//! long-term clean answer; until that split, this module implements a -//! **standalone structural grammar check against the same surface grammar** -//! (`docs/dsrs-format.md`, mirror of `dspy-rs/src/ir/text/{lex,parse}.rs` — -//! keep in sync when the grammar changes). +//! `dspy-rs` depends on the proc-macro crate — so the macro cannot call the +//! full parser without a dependency cycle. This module is the shared +//! **structural grammar** over the one shared lexer ([`crate::lex`]): the +//! macro gets real syntax errors with positions, and the full parser and this +//! checker can never disagree about tokens because there is only one +//! tokenizer. //! //! # What it checks (and what it deliberately does not) //! @@ -29,417 +29,16 @@ //! parser): types, signature/field validity, dataflow ordering, capability //! subsets, model references, demo row shapes — anything semantic. Inside a //! balanced block this checker is deliberately *more permissive* than the -//! real parser: it must never reject an artifact `Program::from_dsrs` +//! full parser: it must never reject an artifact `Program::from_dsrs` //! accepts. -/// A syntax failure with a 1-based source position. -#[derive(Debug, Clone, PartialEq, Eq)] -pub(crate) struct SyntaxError { - pub line: u32, - pub col: u32, - pub message: String, -} - -impl std::fmt::Display for SyntaxError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "line {}, column {}: {}", self.line, self.col, self.message) - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -struct Span { - line: u32, - col: u32, -} - -impl SyntaxError { - fn at(span: Span, message: impl Into) -> Self { - Self { - line: span.line, - col: span.col, - message: message.into(), - } - } -} - -/// Token vocabulary — mirror of `dspy-rs/src/ir/text/lex.rs`. -#[derive(Clone, Debug, PartialEq)] -enum Tok { - Ident(String), - Str, - Num(String), - LBrace, - RBrace, - LParen, - RParen, - LBracket, - RBracket, - Comma, - Eq, - Dot, - Pipe, - Question, - At, - Dollar, - Caret, - Colon, - ColonColon, - Arrow, - Lt, - Gt, - /// A run of three-or-more backticks; the checker scans the raw code - /// region from this token's start offset. - Fence, - Eof, -} - -impl Tok { - fn describe(&self) -> String { - match self { - Tok::Ident(s) => format!("`{s}`"), - Tok::Str => "a string".to_string(), - Tok::Num(n) => format!("number `{n}`"), - Tok::LBrace => "`{`".to_string(), - Tok::RBrace => "`}`".to_string(), - Tok::LParen => "`(`".to_string(), - Tok::RParen => "`)`".to_string(), - Tok::LBracket => "`[`".to_string(), - Tok::RBracket => "`]`".to_string(), - Tok::Comma => "`,`".to_string(), - Tok::Eq => "`=`".to_string(), - Tok::Dot => "`.`".to_string(), - Tok::Pipe => "`|`".to_string(), - Tok::Question => "`?`".to_string(), - Tok::At => "`@`".to_string(), - Tok::Dollar => "`$`".to_string(), - Tok::Caret => "`^`".to_string(), - Tok::Colon => "`:`".to_string(), - Tok::ColonColon => "`::`".to_string(), - Tok::Arrow => "`->`".to_string(), - Tok::Lt => "`<`".to_string(), - Tok::Gt => "`>`".to_string(), - Tok::Fence => "a ``` code fence".to_string(), - Tok::Eof => "end of file".to_string(), - } - } -} - -#[derive(Clone, Debug)] -struct Lexed { - tok: Tok, - span: Span, - /// Byte offset of the token's first byte (fence raw scans restart here). - start: usize, -} - -struct Lexer<'a> { - src: &'a str, - bytes: &'a [u8], - pos: usize, - line: u32, - col: u32, -} - -impl<'a> Lexer<'a> { - fn new(src: &'a str) -> Self { - Self { - src, - bytes: src.as_bytes(), - pos: 0, - line: 1, - col: 1, - } - } - - fn span(&self) -> Span { - Span { - line: self.line, - col: self.col, - } - } - - fn bump(&mut self) -> Option { - let b = *self.bytes.get(self.pos)?; - self.pos += 1; - if b == b'\n' { - self.line += 1; - self.col = 1; - } else { - self.col += 1; - } - Some(b) - } - - fn peek_byte(&self) -> Option { - self.bytes.get(self.pos).copied() - } - - fn skip_trivia(&mut self) { - loop { - match self.peek_byte() { - Some(b' ') | Some(b'\t') | Some(b'\r') | Some(b'\n') => { - self.bump(); - } - Some(b'/') if self.bytes.get(self.pos + 1) == Some(&b'/') => { - while let Some(b) = self.peek_byte() { - if b == b'\n' { - break; - } - self.bump(); - } - } - _ => break, - } - } - } - - /// Repositions the cursor after a raw fence scan. - fn seek(&mut self, pos: usize) { - let mut line = 1u32; - let mut col = 1u32; - for &b in &self.bytes[..pos.min(self.bytes.len())] { - if b == b'\n' { - line += 1; - col = 1; - } else { - col += 1; - } - } - self.pos = pos; - self.line = line; - self.col = col; - } - - fn next_token(&mut self) -> Result { - self.skip_trivia(); - let span = self.span(); - let start = self.pos; - let Some(b) = self.peek_byte() else { - return Ok(Lexed { - tok: Tok::Eof, - span, - start, - }); - }; - - macro_rules! single { - ($tok:expr) => {{ - self.bump(); - $tok - }}; - } - - let tok = match b { - b'{' => single!(Tok::LBrace), - b'}' => single!(Tok::RBrace), - b'(' => single!(Tok::LParen), - b')' => single!(Tok::RParen), - b'[' => single!(Tok::LBracket), - b']' => single!(Tok::RBracket), - b',' => single!(Tok::Comma), - b'=' => single!(Tok::Eq), - b'.' => single!(Tok::Dot), - b'|' => single!(Tok::Pipe), - b'?' => single!(Tok::Question), - b'@' => single!(Tok::At), - b'$' => single!(Tok::Dollar), - b'^' => single!(Tok::Caret), - b'<' => single!(Tok::Lt), - b'>' => single!(Tok::Gt), - b':' => { - self.bump(); - if self.peek_byte() == Some(b':') { - self.bump(); - Tok::ColonColon - } else { - Tok::Colon - } - } - b'-' if self.bytes.get(self.pos + 1) == Some(&b'>') => { - self.bump(); - self.bump(); - Tok::Arrow - } - b'"' => { - self.lex_string(span)?; - Tok::Str - } - b'-' | b'0'..=b'9' => { - let n = self.lex_number(span)?; - Tok::Num(n) - } - b'A'..=b'Z' | b'a'..=b'z' | b'_' => { - let ident_start = self.pos; - while let Some(b) = self.peek_byte() { - if b.is_ascii_alphanumeric() || b == b'_' { - self.bump(); - } else { - break; - } - } - Tok::Ident(self.src[ident_start..self.pos].to_string()) - } - b'`' => { - let mut ticks = 0usize; - while self.peek_byte() == Some(b'`') { - ticks += 1; - self.bump(); - } - if ticks < 3 { - return Err(SyntaxError::at( - span, - "unexpected ` — code fences are three or more backticks and only follow `js`", - )); - } - Tok::Fence - } - other => { - return Err(SyntaxError::at( - span, - format!("unexpected character `{}`", char::from(other)), - )); - } - }; - - Ok(Lexed { tok, span, start }) - } - - /// Lexes a JSON string starting at the current `"` and validates its - /// escapes with serde_json — exactly the real lexer's rule. - fn lex_string(&mut self, span: Span) -> Result<(), SyntaxError> { - let start = self.pos; - self.bump(); // opening quote - loop { - match self.peek_byte() { - None => { - return Err(SyntaxError::at(span, "unterminated string literal")); - } - Some(b'\\') => { - self.bump(); - if self.bump().is_none() { - return Err(SyntaxError::at(span, "unterminated string literal")); - } - } - Some(b'"') => { - self.bump(); - break; - } - Some(b'\n') => { - return Err(SyntaxError::at( - span, - "unterminated string literal (strings are JSON strings; escape newlines as \\n)", - )); - } - Some(_) => { - self.bump(); - } - } - } - let raw = &self.src[start..self.pos]; - serde_json::from_str::(raw) - .map(|_| ()) - .map_err(|e| SyntaxError::at(span, format!("invalid string literal: {e}"))) - } - - fn lex_number(&mut self, span: Span) -> Result { - let start = self.pos; - if self.peek_byte() == Some(b'-') { - self.bump(); - } - let mut saw_digit = false; - while let Some(b) = self.peek_byte() { - match b { - b'0'..=b'9' => { - saw_digit = true; - self.bump(); - } - b'.' | b'e' | b'E' | b'+' | b'-' => { - self.bump(); - } - _ => break, - } - } - if !saw_digit { - return Err(SyntaxError::at(span, "malformed number")); - } - let raw = &self.src[start..self.pos]; - if raw.parse::().is_err() { - return Err(SyntaxError::at(span, format!("malformed number `{raw}`"))); - } - Ok(raw.to_string()) - } - - /// Scans a raw fenced code block starting at byte `from` (the opening - /// backtick run). Returns the byte offset one past the closing fence - /// line. Mirrors `Lexer::scan_code_fence` in the real lexer. - fn scan_code_fence(&self, from: usize) -> Result { - let span = { - let mut probe = Lexer::new(self.src); - probe.seek(from); - probe.span() - }; - let bytes = self.bytes; - let mut i = from; - let mut ticks = 0usize; - while bytes.get(i) == Some(&b'`') { - ticks += 1; - i += 1; - } - if ticks < 3 { - return Err(SyntaxError::at( - span, - "expected a code fence of at least three backticks (```) after `js`", - )); - } - if bytes.get(i) == Some(&b'\r') { - i += 1; - } - if bytes.get(i) != Some(&b'\n') { - return Err(SyntaxError::at( - span, - "the opening code fence must be followed by a newline", - )); - } - i += 1; - let fence = vec![b'`'; ticks]; - let mut line_start = i; - loop { - let line_end = bytes[line_start..] - .iter() - .position(|&b| b == b'\n') - .map(|off| line_start + off); - let (content_end, next_line) = match line_end { - Some(e) => (e, e + 1), - None => (bytes.len(), bytes.len()), - }; - let mut trimmed_end = content_end; - if trimmed_end > line_start && bytes[trimmed_end - 1] == b'\r' { - trimmed_end -= 1; - } - if &bytes[line_start..trimmed_end] == fence.as_slice() { - return Ok(next_line); - } - if line_end.is_none() { - return Err(SyntaxError::at( - span, - format!( - "unterminated code block: no closing fence of {ticks} backticks on its own line" - ), - )); - } - line_start = next_line; - } - } -} - -// --------------------------------------------------------------------------- -// Structural checker -// --------------------------------------------------------------------------- +use crate::ParseError; +use crate::lex::{Lexed, Lexer, Span, Tok}; /// Checks `.dsrs` source for structural syntax validity. See the module docs /// for the exact contract: this is syntax-only, and strictly more permissive /// than `dspy_rs::ir::Program::from_dsrs`. -pub(crate) fn check(src: &str) -> Result<(), SyntaxError> { +pub fn check(src: &str) -> Result<(), ParseError> { Checker::new(src)?.file() } @@ -449,26 +48,26 @@ struct Checker<'a> { } impl<'a> Checker<'a> { - fn new(src: &'a str) -> Result { + fn new(src: &'a str) -> Result { let mut lx = Lexer::new(src); let cur = lx.next_token()?; Ok(Self { lx, cur }) } - fn bump(&mut self) -> Result { + fn bump(&mut self) -> Result { let cur = std::mem::replace(&mut self.cur, self.lx.next_token()?); Ok(cur) } - fn err(&self, message: impl Into) -> SyntaxError { - SyntaxError::at(self.cur.span, message) + fn err(&self, message: impl Into) -> ParseError { + ParseError::at(self.cur.span, message) } fn at_kw(&self, kw: &str) -> bool { matches!(&self.cur.tok, Tok::Ident(word) if word == kw) } - fn expect_kw(&mut self, kw: &str, context: &str) -> Result<(), SyntaxError> { + fn expect_kw(&mut self, kw: &str, context: &str) -> Result<(), ParseError> { if self.at_kw(kw) { self.bump()?; Ok(()) @@ -480,7 +79,7 @@ impl<'a> Checker<'a> { } } - fn expect_ident(&mut self, context: &str) -> Result<(), SyntaxError> { + fn expect_ident(&mut self, context: &str) -> Result<(), ParseError> { match &self.cur.tok { Tok::Ident(_) => { self.bump()?; @@ -493,9 +92,9 @@ impl<'a> Checker<'a> { } } - fn expect_str(&mut self, context: &str) -> Result<(), SyntaxError> { + fn expect_str(&mut self, context: &str) -> Result<(), ParseError> { match &self.cur.tok { - Tok::Str => { + Tok::Str(_) => { self.bump()?; Ok(()) } @@ -506,7 +105,7 @@ impl<'a> Checker<'a> { } } - fn expect_tok(&mut self, tok: Tok, context: &str) -> Result<(), SyntaxError> { + fn expect_tok(&mut self, tok: Tok, context: &str) -> Result<(), ParseError> { if self.cur.tok == tok { self.bump()?; Ok(()) @@ -520,10 +119,11 @@ impl<'a> Checker<'a> { } /// Skips a raw fence region; the current token must be the fence opener. - fn skip_fence(&mut self) -> Result<(), SyntaxError> { + fn skip_fence(&mut self) -> Result<(), ParseError> { debug_assert_eq!(self.cur.tok, Tok::Fence); - let end = self.lx.scan_code_fence(self.cur.start)?; - self.lx.seek(end); + let (_, end) = self.lx.scan_code_fence(self.cur.start)?; + let span = self.lx.span_at(end); + self.lx.seek(end, span); self.cur = self.lx.next_token()?; Ok(()) } @@ -531,7 +131,7 @@ impl<'a> Checker<'a> { /// Consumes a balanced `{ … }` / `[ … ]` / `( … )` region, fence-aware. /// The current token must be the opening delimiter. Content is not /// inspected — everything semantic is the full parser's job. - fn skip_balanced(&mut self, context: &str) -> Result<(), SyntaxError> { + fn skip_balanced(&mut self, context: &str) -> Result<(), ParseError> { let open = match self.cur.tok { Tok::LBrace | Tok::LBracket | Tok::LParen => self.cur.tok.clone(), _ => { @@ -591,7 +191,7 @@ impl<'a> Checker<'a> { } } - fn file(mut self) -> Result<(), SyntaxError> { + fn file(mut self) -> Result<(), ParseError> { // dsrs 1 self.expect_kw("dsrs", "at the start of the file (`dsrs 1`)")?; match &self.cur.tok { @@ -730,7 +330,8 @@ impl<'a> Checker<'a> { #[cfg(test)] mod tests { - use super::{SyntaxError, check}; + use super::check; + use crate::ParseError; const MINI: &str = r#" dsrs 1 @@ -749,7 +350,7 @@ main: Main = seq { } "#; - fn check_err(src: &str) -> SyntaxError { + fn check_err(src: &str) -> ParseError { check(src).expect_err("expected a syntax error") } @@ -758,30 +359,6 @@ main: Main = seq { check(MINI).expect("minimal program is syntactically valid"); } - // The load-bearing property: everything the full parser accepts, this - // checker accepts. Run over the golden fixtures maintained next to the - // real parser (in-repo paths; not part of the published package). - #[test] - fn parity_accepts_everything_the_full_parser_accepts() { - let fixtures = concat!( - env!("CARGO_MANIFEST_DIR"), - "/../dspy-rs/tests/fixtures" - ); - let mut seen = 0usize; - for entry in std::fs::read_dir(fixtures).expect("fixtures dir readable") { - let path = entry.expect("dir entry").path(); - if path.extension().and_then(|e| e.to_str()) != Some("dsrs") { - continue; - } - let src = std::fs::read_to_string(&path).expect("fixture readable"); - check(&src).unwrap_or_else(|e| { - panic!("syntax checker rejected {}: {e}", path.display()) - }); - seen += 1; - } - assert!(seen >= 3, "expected the golden .dsrs fixtures, found {seen}"); - } - #[test] fn rejects_missing_pragma() { let err = check_err("program x\nmain: M = seq { }"); diff --git a/crates/dsrs-syntax/tests/fixture_parity.rs b/crates/dsrs-syntax/tests/fixture_parity.rs new file mode 100644 index 00000000..95e05b34 --- /dev/null +++ b/crates/dsrs-syntax/tests/fixture_parity.rs @@ -0,0 +1,25 @@ +//! The load-bearing property of the structural checker: everything the full +//! parser accepts, [`dsrs_syntax::check`] accepts. +//! +//! Runs over local copies of the golden `.dsrs` fixtures maintained next to +//! the full parser (`crates/dspy-rs/tests/fixtures/*.dsrs`). Keep the copies +//! in sync when a fixture changes — `dspy-rs`'s `test_include_program.rs` +//! exercises the same artifacts through `include_program!`, so drift that +//! matters (a fixture the checker would reject) fails that suite too. + +#[test] +fn parity_accepts_everything_the_full_parser_accepts() { + let fixtures = concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures"); + let mut seen = 0usize; + for entry in std::fs::read_dir(fixtures).expect("fixtures dir readable") { + let path = entry.expect("dir entry").path(); + if path.extension().and_then(|e| e.to_str()) != Some("dsrs") { + continue; + } + let src = std::fs::read_to_string(&path).expect("fixture readable"); + dsrs_syntax::check(&src) + .unwrap_or_else(|e| panic!("syntax checker rejected {}: {e}", path.display())); + seen += 1; + } + assert!(seen >= 3, "expected the golden .dsrs fixtures, found {seen}"); +} diff --git a/crates/dsrs-syntax/tests/fixtures/kitchen.dsrs b/crates/dsrs-syntax/tests/fixtures/kitchen.dsrs new file mode 100644 index 00000000..670ddc15 --- /dev/null +++ b/crates/dsrs-syntax/tests/fixtures/kitchen.dsrs @@ -0,0 +1,130 @@ +// Kitchen sink: every node kind, classes/enums, exotic types, sandboxed +// tools, demos, context/budget/stop options, lineage. Not canonical — +// exercised by the parse/print round-trip property. +dsrs 1 +program kitchen + +caps { fs:read net:fetch } + +model core = "openai:gpt-4o-mini" { temperature 0.2 max_tokens 1024 cache true } + +class Profile { + "A user profile." + name: string "display name" + age: int? check("this|int >= 0", "non-negative") + tags: string[]? + meta: map + kind: "gold" | "basic" +} + +enum Severity { + Low "minor" + High +} + +sig Main { + in ticket: string + in profile: Profile + out reply: string + out audit: string +} + +sig Classify { + "Classify the ticket." + in ticket: string + out severity: Severity +} + +sig Reply { + in ticket: string + out reply: string +} + +sig Draft { + in ticket: string + in feedback: string + out reply: string +} + +sig Judge { + in reply: string + out score: float + out feedback: string +} + +sig Summarize { + in ticket: string + in profile: Profile + out summary: string alias "final_summary" +} + +sig Improve { + in ticket: string + out better: string + out keep_going: bool +} + +sig Audit { + in reply: string + out audit: string +} + +sig Redact { + in text: string + out redacted: string +} + +tool fetch "Fetch a URL" caps [net:fetch] { + in url: string + out body: string +} + +tool shout "Uppercase" { + in text: string + out loud: string +} js``` +(a) => ({ loud: a.text.toUpperCase() }) +``` + +lineage { + optimizer "gepa-0.3" + trainset "tickets@v1" + budget "100 rollouts" + parent "00000000deadbeef" + date "2026-08-14" +} + +main: Main = seq { + classifier = predict Classify (ticket = $.ticket) + router = route classifier.severity { + Low -> low = predict Reply (ticket = $.ticket) + else -> high = agent Reply (ticket = $.ticket) { + tools [fetch shout] + stop_tools [shout] + max_turns 3 + until_parse false + budget { calls 5 tokens 2000 deadline_ms 60000 on_exhausted finalize } + context { max_history_turns 4 tool_result_max_bytes 2048 playbook "Be brief." } + instruction "Escalate." + demos [{"input":{"ticket":"x"},"output":{"reply":"y"}}] + } + } + forked = fork { + summarizer = predict Summarize (ticket = $.ticket, profile = $.profile) + refined = refine (threshold 0.5 max_rounds 2 feedback_field feedback) { + body = drafter = predict Draft (ticket = $.ticket, feedback = "start") + judge = grader = predict Judge (reply = drafter.reply) + } + } join { summary = summarizer.summary, refined_reply = refined.reply } + looped = loop (max_iters 3) { + improver = predict Improve (ticket = ^ticket) + while improver.keep_going + carry { ticket = improver.better } + join { improved = improver.better } + } + audited = retry (attempts 2 backoff_ms 50 feedback true) auditor = predict Audit (reply = router.reply) + redactor = hole Redact (text = audited.audit) caps [fs:read] js``` +(a) => ({ redacted: a.text }) +``` + out { reply = router.reply, audit = redactor.redacted } +} diff --git a/crates/dsrs-syntax/tests/fixtures/qa.dsrs b/crates/dsrs-syntax/tests/fixtures/qa.dsrs new file mode 100644 index 00000000..8a19f517 --- /dev/null +++ b/crates/dsrs-syntax/tests/fixtures/qa.dsrs @@ -0,0 +1,54 @@ +dsrs 1 +program qa + +caps { net:search } + +model fast = "openai:gpt-4o-mini" +model deep = "anthropic:claude-sonnet-4-5" + +sig Main { + in question: string + out answer: string + out sources: string[] +} + +sig Draft { + "Draft a thorough, factual answer." + in question: string + out answer: string +} + +sig Research { + "Verify the draft against sources; collect URLs." + in question: string + in draft: string + out evidence: string[] +} + +sig CiteCheck { + in draft: string + in evidence: string[] + out answer: string + out sources: string[] +} + +tool search "Web search; returns result snippets with URLs" caps [net:search] { + in query: string + out results: string[] +} + +main: Main = seq { + drafter = cot Draft @deep (question = $.question) + researcher = agent Research @fast (question = $.question, draft = drafter.answer) { + tools [search] + max_turns 6 + budget { tokens 40000 on_exhausted finalize } + } + checker = hole CiteCheck (draft = drafter.answer, evidence = researcher.evidence) caps [] js``` +(a) => ({ + answer: a.draft, + sources: a.evidence.filter(e => e.startsWith("http")), +}) +``` + out { answer = checker.answer, sources = checker.sources } +} diff --git a/crates/dsrs-syntax/tests/fixtures/qa_scrambled.dsrs b/crates/dsrs-syntax/tests/fixtures/qa_scrambled.dsrs new file mode 100644 index 00000000..7b19c04c --- /dev/null +++ b/crates/dsrs-syntax/tests/fixtures/qa_scrambled.dsrs @@ -0,0 +1,48 @@ +// Same program as qa.dsrs, deliberately non-canonical: +// comments, shuffled declaration kinds, collapsed whitespace, explicit +// defaults, split `out` steps. Must canonicalize to qa.dsrs exactly. +dsrs 1 +program qa + +model fast = "openai:gpt-4o-mini" { temperature 0.7 } // default temperature, dropped on print +model deep = "anthropic:claude-sonnet-4-5" + +// tool before the sigs — arena order of tool-inline sigs is not printed +tool search "Web search; returns result snippets with URLs" caps [ net : search ] { + in query: string + out results: string[] +} + +sig Main { in question: string out answer: string out sources: string [ ] } + +sig Draft { "Draft a thorough, factual answer." in question: string out answer: string } + +sig Research { + "Verify the draft against sources; collect URLs." + in question: string + in draft: string + out evidence: string[] +} + +sig CiteCheck { in draft: string in evidence: string[] + out answer: string out sources: string[] } + +caps { net:search } // the ceiling may be declared anywhere + +main: Main = seq { + drafter = cot Draft @deep ( question = $.question , ) + researcher = agent Research @fast (question = $.question, draft = drafter.answer) { + budget { on_exhausted finalize tokens 40000 } + max_turns 6 + until_parse true // the default, dropped on print + tools [search] + } + checker = hole CiteCheck (draft = drafter.answer, evidence = researcher.evidence) caps [] js``` +(a) => ({ + answer: a.draft, + sources: a.evidence.filter(e => e.startsWith("http")), +}) +``` + out { answer = checker.answer } + out { sources = checker.sources } +} diff --git a/crates/dsrs-tools/Cargo.toml b/crates/dsrs-tools/Cargo.toml index 1b2d7473..da654cb2 100644 --- a/crates/dsrs-tools/Cargo.toml +++ b/crates/dsrs-tools/Cargo.toml @@ -9,14 +9,14 @@ license = "Apache-2.0" [dependencies] rquickjs = { version = "0.12.2", features = ["parallel"] } -rig-core = { git = "https://github.com/0xPlaygrounds/rig", rev = "aee3b8bf6576ce41c9ac1dd82520752a65fa0127" } -serde = { version = "1.0.219", features = ["derive"] } -serde_json = { version = "1.0.140", features = ["preserve_order"] } -tokio = { version = "1.46.1", features = ["rt", "rt-multi-thread", "sync", "macros", "time"] } -async-trait = "0.1.83" -thiserror = "2.0.17" +rig-core = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } +tokio = { workspace = true, features = ["rt", "rt-multi-thread", "sync", "macros", "time"] } +async-trait = { workspace = true } +thiserror = { workspace = true } blake3 = "1.8.6" -futures = "0.3.31" +futures = { workspace = true } [dev-dependencies] -tokio = { version = "1.46.1", features = ["full"] } +tokio = { workspace = true, features = ["full"] } diff --git a/crates/dsrs-tools/src/capability.rs b/crates/dsrs-tools/src/capability.rs index f624fc30..8be59271 100644 --- a/crates/dsrs-tools/src/capability.rs +++ b/crates/dsrs-tools/src/capability.rs @@ -18,6 +18,138 @@ use serde_json::Value; use crate::error::RegisterError; +/// JavaScript reserved words and ambient globals that must never be used as +/// capability names. A capability becomes a `globalThis` property: shadowing +/// `JSON` would break every capability shim (each depends on +/// `JSON.stringify`/`JSON.parse`), shadowing `globalThis`/`undefined` breaks +/// the sandbox contract outright, and a reserved word like `class` is not +/// callable from script code at all. +const RESERVED_JS_NAMES: &[&str] = &[ + // Reserved words, contextual keywords, and literals. + "arguments", + "async", + "await", + "break", + "case", + "catch", + "class", + "const", + "continue", + "debugger", + "default", + "delete", + "do", + "else", + "enum", + "eval", + "export", + "extends", + "false", + "finally", + "for", + "function", + "get", + "if", + "implements", + "import", + "in", + "instanceof", + "interface", + "let", + "new", + "null", + "of", + "package", + "private", + "protected", + "public", + "return", + "set", + "static", + "super", + "switch", + "this", + "throw", + "true", + "try", + "typeof", + "undefined", + "var", + "void", + "while", + "with", + "yield", + // Ambient globals the sandbox (and its own shims) depend on. + "globalThis", + "NaN", + "Infinity", + "JSON", + "Object", + "Function", + "Array", + "String", + "Number", + "Boolean", + "Symbol", + "BigInt", + "Math", + "Date", + "RegExp", + "Error", + "AggregateError", + "EvalError", + "RangeError", + "ReferenceError", + "SyntaxError", + "TypeError", + "URIError", + "InternalError", + "Promise", + "Proxy", + "Reflect", + "Map", + "Set", + "WeakMap", + "WeakSet", + "WeakRef", + "FinalizationRegistry", + "ArrayBuffer", + "SharedArrayBuffer", + "DataView", + "Atomics", + "Int8Array", + "Uint8Array", + "Uint8ClampedArray", + "Int16Array", + "Uint16Array", + "Int32Array", + "Uint32Array", + "Float16Array", + "Float32Array", + "Float64Array", + "BigInt64Array", + "BigUint64Array", + "decodeURI", + "decodeURIComponent", + "encodeURI", + "encodeURIComponent", + "escape", + "unescape", + "parseFloat", + "parseInt", + "isNaN", + "isFinite", + "structuredClone", + "console", + "print", +]; + +/// Whether `name` is a JavaScript reserved word or an ambient global the +/// sandbox depends on (see [`RESERVED_JS_NAMES`]). +pub(crate) fn is_reserved_js_name(name: &str) -> bool { + RESERVED_JS_NAMES.contains(&name) +} + /// Mangle an arbitrary tool name into a valid JavaScript identifier. /// /// The mangling rule, in order: @@ -28,7 +160,11 @@ use crate::error::RegisterError; /// (`2fast` → `_2fast`), /// 3. an empty name becomes `_tool`, /// 4. a result starting with the runtime-reserved `__dsrs` prefix gets one -/// more leading `_` (`__dsrs_x` → `___dsrs_x`). +/// more leading `_` (`__dsrs_x` → `___dsrs_x`), +/// 5. a result colliding with a JavaScript reserved word or ambient global +/// gets a `_tool` suffix (`JSON` → `JSON_tool`, `class` → `class_tool`) — +/// installing such a name as a global would break the sandbox bootstrap +/// or be uncallable from script code. /// /// The result always passes capability-name validation. The mapping is not /// injective — distinct tool names can mangle to the same identifier — @@ -54,6 +190,9 @@ pub fn js_identifier(name: &str) -> String { if out.starts_with("__dsrs") { out.insert(0, '_'); } + if is_reserved_js_name(&out) { + out.push_str("_tool"); + } out } @@ -67,9 +206,10 @@ pub type CapabilityHandler = /// From JavaScript the capability looks synchronous — `const rows = query({q: /// "..."})` — the executor bridges the call onto the host's Tokio runtime and /// blocks the sandbox thread until it resolves. Capability calls are host -/// code: the sandbox deadline cannot interrupt them mid-flight (it re-arms as -/// soon as control returns to JS), so handlers should enforce their own -/// timeouts. +/// code, so the engine's interrupt handler cannot fire mid-call; instead the +/// executor bounds every call with the sandbox's *remaining* wall-clock +/// deadline (`tokio::time::timeout`): a handler that runs past it is dropped +/// and the call surfaces in JS as a deadline timeout. #[derive(Clone)] pub struct Capability { name: String, @@ -191,7 +331,15 @@ impl Capability { } /// Capability names become JS globals, so they must be valid identifiers - /// and must not collide with the runtime's reserved `__dsrs_*` namespace. + /// (`[A-Za-z_$][A-Za-z0-9_$]*`), must not collide with the runtime's + /// reserved `__dsrs_*` namespace, and must not shadow a JavaScript + /// reserved word or ambient global (see [`RESERVED_JS_NAMES`]). + /// + /// Enforced on **every** install path — [`add_capability`], the builder, + /// and per-sandbox installation (which covers `run_script` and Code + /// Mode) — so a hostile name can never reach the sandbox bootstrap. + /// + /// [`add_capability`]: crate::QuickJsExecutor::add_capability pub(crate) fn validate_name(name: &str) -> Result<(), RegisterError> { let invalid = |reason: &str| RegisterError::InvalidCapability { name: name.to_string(), @@ -212,6 +360,11 @@ impl Capability { if name.starts_with("__dsrs") { return Err(invalid("the `__dsrs` prefix is reserved by the runtime")); } + if is_reserved_js_name(name) { + return Err(invalid( + "collides with a JavaScript reserved word or built-in global", + )); + } Ok(()) } } diff --git a/crates/dsrs-tools/src/quickjs.rs b/crates/dsrs-tools/src/quickjs.rs index b8e33161..8aa4e96a 100644 --- a/crates/dsrs-tools/src/quickjs.rs +++ b/crates/dsrs-tools/src/quickjs.rs @@ -13,9 +13,12 @@ //! - bytecode reuse: sources are compiled once per unique content //! (BLAKE3-keyed) and the bytecode is shared across calls. //! -//! The blocking QuickJS work runs on Tokio's blocking pool; injected -//! capabilities are async Rust and are driven to completion from the sandbox -//! thread via `Handle::block_on`. +//! The blocking QuickJS work runs on Tokio's blocking pool. Injected +//! capabilities are async Rust: each call is spawned onto the host Tokio +//! runtime (bounded by the sandbox's remaining deadline via +//! `tokio::time::timeout`) while the sandbox thread parks on a channel for +//! the result — no `Handle::block_on`, so both current-thread and +//! multi-thread runtimes work. use std::collections::HashMap; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; @@ -48,8 +51,10 @@ const MODULE_NAME: &str = "dsrs_tool"; pub struct SandboxConfig { /// Max heap for one call, in bytes. Default: 32 MiB. pub memory_limit: usize, - /// Wall-clock budget for one call. Time spent inside a capability counts - /// against the budget but cannot be interrupted mid-call. Default: 500ms. + /// Wall-clock budget for one call. JS execution is interrupted by the + /// engine's interrupt handler; a capability call is bounded by the + /// *remaining* budget via `tokio::time::timeout` (the interrupt handler + /// cannot fire while host code runs). Default: 500ms. pub deadline: Duration, /// Max JS stack, in bytes. Default: 512 KiB. pub max_stack: usize, @@ -70,8 +75,20 @@ struct ToolEntry { meta: RegisteredTool, bytecode: Arc>, required: Arc>, + /// Raw BLAKE3 source hash: the bytecode-cache key, kept so + /// [`Executor::deregister`] can evict the tool's cache entry. + hash: [u8; 32], } +/// Content-hash bytecode cache. +/// +/// **Eviction policy (documented, deliberately simple):** +/// - bounded at [`Self::MAX_ENTRIES`] entries; when an insert would grow past +/// the bound the whole map is cleared first (cap-and-clear). Registered +/// tools hold their own `Arc` to their bytecode, so eviction never breaks a +/// registered tool — it only costs a recompile on the next cache miss. +/// - `deregister` evicts the tool's entry unless another registered tool +/// shares the same source hash. #[derive(Default)] struct BytecodeCache { by_hash: Mutex>>>, @@ -79,6 +96,35 @@ struct BytecodeCache { misses: AtomicU64, } +impl BytecodeCache { + /// Upper bound on cached compilations. An optimizer generating thousands + /// of candidate tool bodies stays bounded instead of leaking them all. + const MAX_ENTRIES: usize = 128; + + fn get(&self, hash: &[u8; 32]) -> Option>> { + self.by_hash + .lock() + .expect("cache lock poisoned") + .get(hash) + .cloned() + } + + fn insert(&self, hash: [u8; 32], bytecode: Arc>) { + let mut map = self.by_hash.lock().expect("cache lock poisoned"); + if map.len() >= Self::MAX_ENTRIES && !map.contains_key(&hash) { + map.clear(); + } + map.insert(hash, bytecode); + } + + fn remove(&self, hash: &[u8; 32]) { + self.by_hash + .lock() + .expect("cache lock poisoned") + .remove(hash); + } +} + /// Counters for the content-hash bytecode cache. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct CacheStats { @@ -254,34 +300,22 @@ impl QuickJsExecutor { } /// Compile `js_source` (wrapped as an ES module) to bytecode, reusing the - /// content-hash cache. Returns `(bytecode, hex_hash, was_cache_hit)`. + /// content-hash cache. Returns `(bytecode, raw_hash)`. fn compile_or_cached( &self, js_source: &str, - ) -> Result<(Arc>, String, bool), RegisterError> { + ) -> Result<(Arc>, [u8; 32]), RegisterError> { let hash = *blake3::hash(js_source.as_bytes()).as_bytes(); - let hex = blake3::Hash::from_bytes(hash).to_hex().to_string(); - if let Some(bytecode) = self - .cache - .by_hash - .lock() - .expect("cache lock poisoned") - .get(&hash) - .cloned() - { + if let Some(bytecode) = self.cache.get(&hash) { self.cache.hits.fetch_add(1, Ordering::Relaxed); - return Ok((bytecode, hex, true)); + return Ok((bytecode, hash)); } let bytecode = Arc::new(compile_module(&wrap_source(js_source), self.config)?); - self.cache - .by_hash - .lock() - .expect("cache lock poisoned") - .insert(hash, Arc::clone(&bytecode)); + self.cache.insert(hash, Arc::clone(&bytecode)); self.cache.misses.fetch_add(1, Ordering::Relaxed); - Ok((bytecode, hex, false)) + Ok((bytecode, hash)) } } @@ -353,7 +387,8 @@ impl Executor for QuickJsExecutor { let capabilities = self.snapshot_capabilities(); // Stage 2: parse/compile (content-hash cached). - let (bytecode, source_hash, _hit) = self.compile_or_cached(&source.js_source)?; + let (bytecode, hash) = self.compile_or_cached(&source.js_source)?; + let source_hash = blake3::Hash::from_bytes(hash).to_hex().to_string(); // Stage 3: instantiate in a sandbox, check the module evaluates to a // function, and run the self-test if present. @@ -392,6 +427,7 @@ impl Executor for QuickJsExecutor { meta: meta.clone(), bytecode, required: Arc::new(source.required_params()), + hash, }; let mut tools = self.tools.write().expect("tool lock poisoned"); if tools.contains_key(&source.name) { @@ -444,11 +480,17 @@ impl Executor for QuickJsExecutor { } fn deregister(&self, name: &str) -> bool { - self.tools - .write() - .expect("tool lock poisoned") - .remove(name) - .is_some() + let mut tools = self.tools.write().expect("tool lock poisoned"); + let Some(entry) = tools.remove(name) else { + return false; + }; + // Evict the tool's bytecode unless another registered tool shares it + // (identical sources share one cache entry). + let shared = tools.values().any(|other| other.hash == entry.hash); + if !shared { + self.cache.remove(&entry.hash); + } + true } } @@ -636,7 +678,9 @@ impl Sandbox { let handle = handle.ok_or_else(|| { internal("capabilities require a Tokio runtime handle".to_string()) })?; - context.with(|ctx| install_capabilities(&ctx, capabilities, &handle))?; + context.with(|ctx| { + install_capabilities(&ctx, capabilities, &handle, deadline, &timed_out, config) + })?; } Ok(Self { @@ -648,20 +692,64 @@ impl Sandbox { } } -/// Expose each capability as a global JS function. The raw hook takes and -/// returns JSON strings; a JS shim gives callers a plain-value API. +/// Fixed shim factory: turns the raw JSON-string hook into a plain-value JS +/// API (`undefined` argument becomes `null`). Deliberately a **constant** — +/// nothing is ever interpolated into evaluated code, so a hostile capability +/// name cannot inject into the sandbox bootstrap. `JSON.parse`/`stringify` +/// are captured at bootstrap time so later clobbering of `globalThis.JSON` +/// by sandboxed code cannot corrupt capability calls. +const SHIM_FACTORY: &str = "((parse, stringify) => (hook) => (arg) => \ + parse(hook(stringify(arg ?? null))))(JSON.parse, JSON.stringify)"; + +/// Extra time the sandbox thread waits beyond the capability timeout for the +/// host runtime to deliver the (possibly already-timed-out) result. Only +/// reached when the Tokio runtime is not being driven at all. +const CAPABILITY_DELIVERY_GRACE: Duration = Duration::from_secs(1); + +/// Expose each capability as a global JS function. +/// +/// Security-critical invariants: +/// - every name is re-validated here ([`Capability::validate_name`]), so the +/// check cannot be bypassed by any install path (`run_script`, Code Mode, +/// executor registration — they all create sandboxes through this function); +/// - names are only ever used as `globals.set()` property keys; no name is +/// spliced into evaluated JS (see [`SHIM_FACTORY`]). +/// +/// The raw hook takes and returns JSON strings. It runs the handler by +/// spawning it onto the host Tokio runtime, bounded by the sandbox's +/// **remaining** deadline (`tokio::time::timeout`) — the engine interrupt +/// handler cannot fire while host code runs, so the deadline is enforced here +/// instead — and parks the sandbox thread on a channel for the result. This +/// avoids `Handle::block_on`, which deadlocks blocking-pool threads of a +/// current-thread runtime. fn install_capabilities( ctx: &Ctx<'_>, capabilities: &[Capability], handle: &Handle, + deadline: Instant, + timed_out: &Arc, + config: SandboxConfig, ) -> Result<(), ExecError> { let globals = ctx.globals(); + let internal = |message: String| ExecError::Internal { message }; + + let factory: Function<'_> = ctx + .eval(SHIM_FACTORY) + .map_err(|e| internal(format!("failed to build capability shim factory: {e}")))?; + for capability in capabilities { + // Unbypassable install-time validation: `run_script` and Code Mode + // accept caller-supplied capabilities without going through + // `add_capability`, so the name must be checked again here. + Capability::validate_name(capability.name()) + .map_err(|e| internal(format!("refusing to install capability: {e}")))?; + let cap_name = capability.name().to_string(); - let hook_name = format!("__dsrs_cap_{cap_name}"); let handler = capability.handler(); let handle = handle.clone(); let thrown_name = cap_name.clone(); + let timed_out = Arc::clone(timed_out); + let deadline_ms = config.deadline.as_millis() as u64; let hook = move |ctx: Ctx<'_>, args_json: String| -> rquickjs::Result { let throw = |message: String| { @@ -672,25 +760,45 @@ fn install_capabilities( }; let args: Value = serde_json::from_str(&args_json) .map_err(|e| throw(format!("argument round-trip failed: {e}")))?; - let result = handle.block_on(handler(args)).map_err(throw)?; + + // Bound the handler by whatever is left of the sandbox deadline. + let remaining = deadline.saturating_duration_since(Instant::now()); + let (tx, rx) = std::sync::mpsc::channel(); + let future = handler(args); + handle.spawn(async move { + let _ = tx.send(tokio::time::timeout(remaining, future).await); + }); + let result = match rx.recv_timeout(remaining + CAPABILITY_DELIVERY_GRACE) { + Ok(Ok(result)) => result, + Ok(Err(_elapsed)) => { + // Handler ran past the deadline: flag the sandbox so the + // failure classifies as a timeout, and throw a clear + // (catchable) error into JS. + timed_out.store(true, Ordering::SeqCst); + return Err(throw(format!( + "call exceeded the sandbox's {deadline_ms}ms deadline" + ))); + } + Err(_) => { + timed_out.store(true, Ordering::SeqCst); + return Err(throw( + "host runtime failed to deliver a result before the sandbox \ + deadline (is the Tokio runtime being driven?)" + .to_string(), + )); + } + }; + let result = result.map_err(throw)?; serde_json::to_string(&result) .map_err(|e| throw(format!("result serialization failed: {e}"))) }; + let shim: Function<'_> = factory + .call((Func::from(hook),)) + .map_err(|e| internal(format!("failed to build capability shim `{cap_name}`: {e}")))?; globals - .set(hook_name.as_str(), Func::from(hook)) - .map_err(|e| ExecError::Internal { - message: format!("failed to install capability `{cap_name}`: {e}"), - })?; - - // Shim: plain JS values in/out; `undefined` argument becomes `null`. - let shim = format!( - "globalThis.{cap_name} = ((raw) => (arg) => JSON.parse(raw(JSON.stringify(arg ?? null))))(globalThis.{hook_name}); delete globalThis.{hook_name};" - ); - ctx.eval::<(), _>(shim.into_bytes()) - .map_err(|e| ExecError::Internal { - message: format!("failed to install capability shim `{cap_name}`: {e}"), - })?; + .set(cap_name.as_str(), shim) + .map_err(|e| internal(format!("failed to install capability `{cap_name}`: {e}")))?; } Ok(()) } diff --git a/crates/dsrs-tools/src/rig_tool.rs b/crates/dsrs-tools/src/rig_tool.rs index 67ededa2..12d8c857 100644 --- a/crates/dsrs-tools/src/rig_tool.rs +++ b/crates/dsrs-tools/src/rig_tool.rs @@ -1,5 +1,5 @@ //! Bridge from sandboxed tools to [`rig::tool::ToolDyn`], the trait DSRs -//! already threads through `Predict`, `ChainOfThought`, and `ReAct`. A +//! already threads through `Predict` and `ChainOfThought`. A //! graduated ephemeral tool is indistinguishable from a hand-written one. use std::sync::Arc; diff --git a/crates/dsrs-tools/tests/code_mode.rs b/crates/dsrs-tools/tests/code_mode.rs index 4f3d7c58..966e731d 100644 --- a/crates/dsrs-tools/tests/code_mode.rs +++ b/crates/dsrs-tools/tests/code_mode.rs @@ -246,6 +246,13 @@ fn js_identifier_mangling_rule() { assert_eq!(js_identifier("__dsrs_evil"), "___dsrs_evil"); assert_eq!(js_identifier("emoji🔥name"), "emoji_name"); assert_eq!(js_identifier("$ok"), "$ok"); + // Reserved globals and words are suffixed rather than shadowed. + assert_eq!(js_identifier("JSON"), "JSON_tool"); + assert_eq!(js_identifier("Object"), "Object_tool"); + assert_eq!(js_identifier("class"), "class_tool"); + assert_eq!(js_identifier("undefined"), "undefined_tool"); + // Only exact collisions are mangled. + assert_eq!(js_identifier("JSONish"), "JSONish"); } #[tokio::test] @@ -278,6 +285,58 @@ async fn run_script_returns_and_chains_capabilities() { assert_eq!(calls.load(Ordering::SeqCst), 2); } +#[tokio::test] +async fn run_script_rejects_capability_name_injection() { + // `run_script` installs caller-supplied capabilities without going + // through `add_capability`; a hostile name must still be refused before + // it can reach the sandbox bootstrap. + let evil = Capability::new( + "x; globalThis.leak = 1; //", + "injection attempt", + |_| async move { Ok(json!(null)) }, + ); + let err = run_script("return typeof leak;", vec![evil], SandboxConfig::default()) + .await + .expect_err("must reject the capability"); + match err { + ExecError::Internal { message } => { + assert!(message.contains("capability"), "{message}"); + } + other => panic!("expected Internal (host misconfiguration), got {other:?}"), + } + + // Reserved names are refused on the same path. + let shadow = Capability::new("JSON", "shadows JSON", |_| async move { Ok(json!(null)) }); + let err = run_script("return 1;", vec![shadow], SandboxConfig::default()) + .await + .expect_err("must reject the capability"); + assert!(matches!(err, ExecError::Internal { .. }), "{err:?}"); +} + +#[tokio::test] +async fn json_named_tool_is_mangled_and_does_not_break_other_capabilities() { + // A tool named `JSON` must not shadow `globalThis.JSON` (every capability + // shim depends on it); it is installed as `JSON_tool` instead, and other + // capabilities keep working. + let tools: Vec> = vec![Arc::new(NamedTool("JSON")), Arc::new(NamedTool("beta"))]; + let tool = CodeModeTool::new(tools, SandboxConfig::default()) + .await + .expect("build"); + let definition = tool.definition(String::new()).await; + assert!(definition.description.contains("JSON_tool(args)")); + + let args = json!({ + "code": "const a = JSON_tool({});\n\ + const b = beta({});\n\ + return {a: a.from, b: b.from, json_intact: typeof JSON.stringify === 'function'};" + }); + let result = tool.call(args.to_string()).await.expect("call"); + assert_eq!( + serde_json::from_str::(&result).unwrap(), + json!({"a": "JSON", "b": "beta", "json_intact": true}) + ); +} + #[tokio::test] async fn run_script_syntax_error_is_repairable() { let err = run_script("return {", Vec::new(), SandboxConfig::default()) @@ -375,6 +434,51 @@ async fn code_mode_deadline_kills_runaway_script() { ); } +/// A tool whose call never finishes within any sane deadline. +struct StallTool; + +impl rig::tool::Tool for StallTool { + const NAME: &'static str = "stall"; + type Error = CannedError; + type Args = serde_json::Value; + type Output = serde_json::Value; + + async fn definition(&self, _prompt: String) -> ToolDefinition { + ToolDefinition { + name: Self::NAME.to_string(), + description: "Never returns".to_string(), + parameters: json!({"type": "object"}), + } + } + + async fn call(&self, _args: Self::Args) -> Result { + tokio::time::sleep(Duration::from_secs(3600)).await; + Ok(json!(null)) + } +} + +#[tokio::test] +async fn code_mode_deadline_bounds_capability_calls() { + // In Code Mode every tool call is a capability call; the wall-clock + // deadline must bound the host future too, not just JS execution. + let tool = CodeModeTool::new(vec![Arc::new(StallTool) as Arc], tight_config()) + .await + .expect("build"); + let started = Instant::now(); + let result = tool + .call(json!({"code": "return stall({});"}).to_string()) + .await + .expect("repairable errors come back as Ok results"); + let elapsed = started.elapsed(); + let parsed: serde_json::Value = serde_json::from_str(&result).unwrap(); + assert_eq!(parsed["kind"], "timeout", "typed timeout error: {result}"); + assert_eq!(parsed["name"], RUN_JS_TOOL_NAME); + assert!( + elapsed < Duration::from_secs(5), + "capability call was not bounded by the deadline: {elapsed:?}" + ); +} + #[tokio::test] async fn code_mode_script_error_is_repairable_result() { let tool = CodeModeTool::new(vec![Arc::new(FailTool) as Arc], tight_config()) diff --git a/crates/dsrs-tools/tests/quickjs_executor.rs b/crates/dsrs-tools/tests/quickjs_executor.rs index 66503d10..04e0834b 100644 --- a/crates/dsrs-tools/tests/quickjs_executor.rs +++ b/crates/dsrs-tools/tests/quickjs_executor.rs @@ -362,7 +362,22 @@ async fn tool_can_catch_capability_errors() { #[test] fn reserved_and_invalid_capability_names_are_rejected() { - for bad in ["__dsrs_cap_x", "has space", "1starts_with_digit", ""] { + for bad in [ + "__dsrs_cap_x", + "has space", + "1starts_with_digit", + "", + // Injection attempt: must never reach the sandbox bootstrap. + "x; globalThis.leak = 1; //", + // Reserved globals/words: shadowing them breaks the runtime's shims. + "JSON", + "Object", + "Promise", + "globalThis", + "undefined", + "class", + "eval", + ] { let err = QuickJsExecutor::builder() .capability(Capability::new( bad, @@ -378,6 +393,69 @@ fn reserved_and_invalid_capability_names_are_rejected() { } } +#[tokio::test(flavor = "multi_thread")] +async fn capability_call_is_bounded_by_the_deadline() { + // The interrupt handler cannot fire while host code runs; the executor + // must bound the capability call itself with the remaining deadline. + let executor = QuickJsExecutor::builder() + .deadline(Duration::from_millis(100)) + .capability(Capability::new( + "stall", + "never returns within the deadline", + |_| async move { + tokio::time::sleep(Duration::from_secs(3600)).await; + Ok(json!(null)) + }, + )) + .build() + .expect("build"); + executor + .register(schemaless("stuck", "(args) => stall({})")) + .await + .expect("register"); + + let start = Instant::now(); + let err = executor + .execute(ToolInvocation::new("stuck", json!({}))) + .await + .expect_err("must time out"); + let elapsed = start.elapsed(); + + match &err { + ExecError::Timeout { name, deadline_ms } => { + assert_eq!(name, "stuck"); + assert_eq!(*deadline_ms, 100); + } + other => panic!("expected Timeout, got {other:?}"), + } + assert!( + elapsed < Duration::from_secs(5), + "capability call was not bounded: {elapsed:?}" + ); +} + +#[tokio::test] +async fn capabilities_work_on_a_current_thread_runtime() { + // Capability calls are bridged via spawn + channel (not + // `Handle::block_on`), so a current-thread runtime must work too. + let executor = QuickJsExecutor::builder() + .capability(Capability::new("double", "double a number", |args| async move { + let n = args["n"].as_f64().ok_or("expected {n: number}")?; + Ok(json!(n * 2.0)) + })) + .build() + .expect("build"); + executor + .register(schemaless("via_cap", "(args) => double({n: args.n})")) + .await + .expect("register"); + let result = executor + .execute(ToolInvocation::new("via_cap", json!({"n": 21}))) + .await + .expect("execute"); + assert_eq!(result.as_f64(), Some(42.0), "{result:?}"); +} + // ----------------------------------------------------- validate-then-register #[tokio::test] @@ -576,6 +654,72 @@ async fn identical_sources_share_cached_bytecode() { assert_eq!((stats.entries, stats.misses), (2, 2)); } +#[tokio::test] +async fn deregister_evicts_bytecode_unless_shared() { + let executor = QuickJsExecutor::new(); + let js = "(args) => args.x * 3"; + executor + .register(schemaless("triple_a", js)) + .await + .expect("register a"); + executor + .register(schemaless("triple_b", js)) + .await + .expect("register b"); + // Two tools, one shared cache entry. + assert_eq!(executor.cache_stats().entries, 1); + + // Still referenced by triple_b: the entry must survive. + assert!(executor.deregister("triple_a")); + assert_eq!(executor.cache_stats().entries, 1); + assert_eq!( + executor + .execute(ToolInvocation::new("triple_b", json!({"x": 2}))) + .await + .expect("execute"), + json!(6) + ); + + // Last reference gone: the bytecode is evicted with it. + assert!(executor.deregister("triple_b")); + assert_eq!(executor.cache_stats().entries, 0); + assert!(!executor.deregister("triple_b"), "already gone"); +} + +#[tokio::test] +async fn bytecode_cache_is_bounded() { + // The cache is capped (cap-and-clear); an optimizer generating many + // candidate sources must not grow it without bound, and registered tools + // must keep working after eviction (they hold their own bytecode Arc). + let executor = QuickJsExecutor::new(); + let count = 140; // > the 128-entry cap + for i in 0..count { + executor + .register(schemaless( + &format!("cand_{i}"), + &format!("(args) => args.x + {i}"), + )) + .await + .expect("register"); + } + let stats = executor.cache_stats(); + assert!( + stats.entries <= 128, + "cache exceeded its bound: {} entries", + stats.entries + ); + // Tools registered before the clear still execute. + for i in [0, count - 1] { + assert_eq!( + executor + .execute(ToolInvocation::new(format!("cand_{i}"), json!({"x": 1}))) + .await + .expect("execute"), + json!(1 + i) + ); + } +} + #[tokio::test] async fn executing_many_times_never_recompiles() { let executor = QuickJsExecutor::new(); diff --git a/docs-chat-worker/gepa/adapter.py b/docs-chat-worker/gepa/adapter.py new file mode 100644 index 00000000..48e9438a --- /dev/null +++ b/docs-chat-worker/gepa/adapter.py @@ -0,0 +1,128 @@ +"""GEPA adapter for the docs-chat worker. + +Candidate = {"system_prompt": }: the one component GEPA evolves, +deployed by pasting the winner into SYSTEM in ../worker.js. + +evaluate() runs each eval question through toast-1 (same store tools as the +worker) and scores the answer with the judge. Trajectories keep the searches +toast-1 ran (hosted_tool_calls) so the reflection step can see retrieval +behavior, not just the final answer. +""" + +import json +import os +import time +import urllib.request +from concurrent.futures import ThreadPoolExecutor + +from gepa.core.adapter import EvaluationBatch + +from judge import judge + +MXBAI_URL = "https://api.mixedbread.com/v1/chat/completions" +STORES = ["dsrs-docs", "dsrs-code"] +MAX_TOKENS = 2048 # matches the deployed worker + + +def call_student(system_prompt, question, attempts=3): + """One toast-1 rollout with retry. Returns (answer_text, hosted_tool_calls).""" + for attempt in range(attempts): + try: + return _call_student_once(system_prompt, question) + except Exception: + if attempt == attempts - 1: + raise + time.sleep(5 * (attempt + 1)) + + +def _call_student_once(system_prompt, question): + req = urllib.request.Request( + MXBAI_URL, + data=json.dumps( + { + "model": "toast-1", + "stream": False, + "max_tokens": MAX_TOKENS, + "messages": [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": question}, + ], + "tools": [ + {"type": "store_search", "store_identifiers": STORES}, + {"type": "store_grep", "store_identifiers": STORES}, + ], + } + ).encode(), + headers={ + "Authorization": f"Bearer {os.environ['MXBAI_API_KEY']}", + "Content-Type": "application/json", + }, + ) + resp = json.load(urllib.request.urlopen(req, timeout=180)) + answer = resp["choices"][0]["message"].get("content") or "" + return answer, resp.get("hosted_tool_calls") or [] + + +def _summarize_tool_calls(tool_calls): + out = [] + for tc in tool_calls: + kind = tc.get("type", "tool_call") + detail = tc.get("queries") or tc.get("pattern") or "" + out.append(f"{kind}: {detail}") + return out + + +class DocsChatAdapter: + # gepa probes this attribute; None = use the built-in instruction proposer + propose_new_texts = None + + def __init__(self, max_workers=4): + self.max_workers = max_workers + + def evaluate(self, batch, candidate, capture_traces=False): + system_prompt = candidate["system_prompt"] + + def run_one(item): + # Per-example failures score 0.0 instead of raising, per the + # GEPAAdapter contract — one bad rollout must not kill the run. + try: + answer, tool_calls = call_student(system_prompt, item["question"]) + score, feedback, _ = judge( + item["question"], + answer, + item["key_points"], + item.get("citations", []), + ) + except Exception as e: + answer, tool_calls, score = "", [], 0.0 + feedback = f"ROLLOUT ERROR (scored 0): {e}" + return { + "question": item["question"], + "answer": answer, + "tool_calls": _summarize_tool_calls(tool_calls), + "score": score, + "feedback": feedback, + } + + with ThreadPoolExecutor(max_workers=self.max_workers) as ex: + results = list(ex.map(run_one, batch)) + + return EvaluationBatch( + outputs=[r["answer"] for r in results], + scores=[r["score"] for r in results], + trajectories=results if capture_traces else None, + ) + + def make_reflective_dataset(self, candidate, eval_batch, components_to_update): + records = [ + { + "Inputs": {"question": t["question"]}, + "Generated Outputs": { + "answer": t["answer"], + "searches_run": t["tool_calls"], + }, + "Feedback": t["feedback"], + } + for t in eval_batch.trajectories + ] + return {"system_prompt": records} diff --git a/docs-chat-worker/gepa/compare.py b/docs-chat-worker/gepa/compare.py new file mode 100644 index 00000000..6af1e9ee --- /dev/null +++ b/docs-chat-worker/gepa/compare.py @@ -0,0 +1,77 @@ +"""Compare two system prompts on a holdout set the optimizer never saw. + + python compare.py --a seed --b best_prompt.txt --holdout holdout.jsonl + +--a/--b accept 'seed' (extract from ../worker.js) or a path to a prompt file. +Each question runs once per prompt; the judge scores both. Prints per-question +scores and the mean delta. +""" + +import argparse +import json +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +from run import HERE, load_env, seed_from_worker + + +def read_prompt(spec): + return seed_from_worker() if spec == "seed" else Path(spec).read_text() + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--a", default="seed") + ap.add_argument("--b", default=HERE / "best_prompt.txt") + ap.add_argument("--c", default=None, help="optional third arm") + ap.add_argument("--holdout", default=HERE / "holdout.jsonl", type=Path) + ap.add_argument("--workers", default=4, type=int) + args = ap.parse_args() + + load_env() + from adapter import call_student + from judge import judge + + prompts = {"A": read_prompt(str(args.a)), "B": read_prompt(str(args.b))} + if args.c: + prompts["C"] = read_prompt(str(args.c)) + items = [json.loads(l) for l in args.holdout.read_text().splitlines() if l.strip()] + + def run_one(task): + label, item = task + try: + answer, _ = call_student(prompts[label], item["question"]) + score, feedback, _ = judge( + item["question"], answer, item["key_points"], item.get("citations", []) + ) + except Exception as e: + answer, score, feedback = "", 0.0, f"ERROR: {e}" + return label, item["question"], score, len(answer), feedback + + tasks = [(label, item) for item in items for label in sorted(prompts)] + with ThreadPoolExecutor(max_workers=args.workers) as ex: + results = list(ex.map(run_one, tasks)) + + by_q = {} + for label, q, score, chars, feedback in results: + by_q.setdefault(q, {})[label] = (score, chars, feedback) + + labels = sorted(prompts) + header = " ".join(f"{l:>6}" for l in labels) + print(f"\n{'Q':<64} {header}") + sums = {l: 0.0 for l in labels} + for q, r in by_q.items(): + row = " ".join(f"{r.get(l, (0, 0, ''))[0]:>6.2f}" for l in labels) + for l in labels: + sums[l] += r.get(l, (0, 0, ""))[0] + print(f"{q[:62]:<64} {row}") + means = " | ".join(f"{l}: {sums[l] / len(by_q):.3f}" for l in labels) + print(f"\nmeans -> {means}") + + detail = HERE / "compare_detail.json" + detail.write_text(json.dumps(by_q, indent=2)) + print(f"per-question feedback in {detail}") + + +if __name__ == "__main__": + main() diff --git a/docs-chat-worker/gepa/judge.py b/docs-chat-worker/gepa/judge.py new file mode 100644 index 00000000..a4377c5b --- /dev/null +++ b/docs-chat-worker/gepa/judge.py @@ -0,0 +1,108 @@ +"""LLM judge for docs-chat answers: key-point coverage minus noise. + +The teacher model grades an answer against a per-question checklist and +returns a scalar score plus textual feedback. The feedback string is what +GEPA's reflection step consumes, so it must name what was missed and quote +what was fluff — a bare number wastes the optimizer's main advantage. + +Score design (coverage-dominant, so the optimizer cannot reward-hack): + 0.70 * key-point coverage -- terse-but-lossy answers lose here + 0.10 * required citations -- file paths the answer should mention + 0.20 * (1 - noise fraction) -- verbose answers lose here + -0.15 per contradicted fact (capped at 0.30) +""" + +import json +import os + +import litellm + +JUDGE_MODEL = os.environ.get("JUDGE_MODEL", "openrouter/openai/gpt-5.6-sol") + +JUDGE_PROMPT = """\ +You are grading an answer from a documentation assistant for DSRs (dspy-rs, \ +a Rust framework for building and optimizing LM pipelines). + +## Question +{question} + +## Answer under evaluation +{answer} + +## Key points a correct answer must contain +{key_points} + +## File paths the answer should cite (empty = no citation required) +{citations} + +Grade strictly. "Covered" means the fact is stated correctly, not merely \ +alluded to. "Fluff" means spans that carry no key point: preamble, restating \ +the question, hedging, generic filler, or detail nobody asked for. + +Return ONLY a JSON object with these fields: +{{ + "covered": [one boolean per key point, in the same order as listed], + "missed": ["each absent or wrong key point, restated briefly"], + "fluff": ["verbatim spans from the answer that carry no key point"], + "citations_present": ["paths from the required list that the answer mentions"], + "wrong_claims": ["claims that contradict the key points, if any"], + "advice": "2-5 sentences of direct advice to the assistant: what to add, what to cut, how to restructure." +}}""" + + +def judge(question, answer, key_points, citations): + """Returns (score: float in [0,1], feedback: str, raw: dict).""" + prompt = JUDGE_PROMPT.format( + question=question, + answer=answer or "(empty answer)", + key_points="\n".join(f"{i+1}. {k}" for i, k in enumerate(key_points)), + citations="\n".join(citations) if citations else "(none required)", + ) + resp = litellm.completion( + model=JUDGE_MODEL, + messages=[{"role": "user", "content": prompt}], + response_format={"type": "json_object"}, + num_retries=2, + ) + raw = _parse_json(resp.choices[0].message.content) + + covered = raw.get("covered", []) + coverage = sum(bool(c) for c in covered) / max(1, len(key_points)) + cite = ( + len(raw.get("citations_present", [])) / len(citations) if citations else 1.0 + ) + fluff_chars = sum(len(s) for s in raw.get("fluff", [])) + noise = min(1.0, fluff_chars / max(1, len(answer or ""))) + wrong_penalty = min(0.30, 0.15 * len(raw.get("wrong_claims", []))) + + score = 0.70 * coverage + 0.10 * min(1.0, cite) + 0.20 * (1.0 - noise) + score = max(0.0, min(1.0, score - wrong_penalty)) + + feedback = _compose_feedback(raw, coverage, noise, score) + return score, feedback, raw + + +def _compose_feedback(raw, coverage, noise, score): + lines = [ + f"Score {score:.2f} (coverage {coverage:.0%}, noise {noise:.0%} of answer)." + ] + if raw.get("missed"): + lines.append("Missing or wrong:") + lines += [f" - {m}" for m in raw["missed"]] + if raw.get("wrong_claims"): + lines.append("Contradicts the docs:") + lines += [f" - {w}" for w in raw["wrong_claims"]] + if raw.get("fluff"): + lines.append("Fluff to cut:") + lines += [f' - "{f}"' for f in raw["fluff"][:5]] + if raw.get("advice"): + lines.append(f"Advice: {raw['advice']}") + return "\n".join(lines) + + +def _parse_json(text): + text = text.strip() + if text.startswith("```"): + text = text.split("```")[1] + text = text[4:] if text.startswith("json") else text + return json.loads(text) diff --git a/docs-chat-worker/gepa/run.py b/docs-chat-worker/gepa/run.py new file mode 100644 index 00000000..9aa11330 --- /dev/null +++ b/docs-chat-worker/gepa/run.py @@ -0,0 +1,89 @@ +"""Run GEPA over the docs-chat system prompt. + + python run.py --evalset evalset.jsonl --budget 400 + +The seed prompt is extracted from ../worker.js so the run always starts from +what is actually deployed. Deploy the result by pasting best_prompt.txt into +SYSTEM in worker.js. +""" + +import argparse +import json +import random +import re +from datetime import datetime +from pathlib import Path + +HERE = Path(__file__).resolve().parent + + +def load_env(): + env = HERE.parents[1] / ".env" + import os + + for line in env.read_text().splitlines(): + line = line.strip() + if line and not line.startswith("#") and "=" in line: + k, v = line.split("=", 1) + os.environ.setdefault(k, v) + + +def seed_from_worker(): + src = (HERE.parent / "worker.js").read_text() + m = re.search(r"const SYSTEM = `(.*?)`;", src, re.S) + if m: + # undo JS template-literal line continuations + return m.group(1).replace("\\\n", "") + m = re.search(r'const SYSTEM = ("(?:[^"\\]|\\.)*");', src) + if not m: + raise SystemExit("could not find `const SYSTEM = ...` in worker.js") + return json.loads(m.group(1)) + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--evalset", default=HERE / "evalset.jsonl", type=Path) + ap.add_argument("--budget", default=400, type=int, help="max metric calls (1 call = 1 toast-1 rollout + 1 judge call)") + ap.add_argument("--val-frac", default=0.25, type=float) + ap.add_argument("--workers", default=4, type=int) + args = ap.parse_args() + + load_env() + import gepa + + from adapter import DocsChatAdapter + + items = [json.loads(l) for l in args.evalset.read_text().splitlines() if l.strip()] + random.Random(0).shuffle(items) + n_val = max(3, int(len(items) * args.val_frac)) + valset, trainset = items[:n_val], items[n_val:] + print(f"{len(trainset)} train / {len(valset)} val") + + run_dir = HERE / "runs" / datetime.now().strftime("%Y%m%d-%H%M%S") + result = gepa.optimize( + seed_candidate={"system_prompt": seed_from_worker()}, + trainset=trainset, + valset=valset, + adapter=DocsChatAdapter(max_workers=args.workers), + reflection_lm="openrouter/openai/gpt-5.6-sol", + reflection_minibatch_size=3, + max_metric_calls=args.budget, + run_dir=str(run_dir), + track_best_outputs=True, + display_progress_bar=False, + seed=0, + ) + + best = result.best_candidate["system_prompt"] + print("\n=== val scores per candidate ===") + for i, s in enumerate(result.val_aggregate_scores): + marker = " <- best" if i == result.best_idx else "" + print(f" candidate {i}: {s:.3f}{marker}") + out = HERE / "best_prompt.txt" + out.write_text(best) + print(f"\nbest prompt written to {out}") + print(f"run artifacts in {run_dir}") + + +if __name__ == "__main__": + main() diff --git a/docs-chat-worker/gepa/test_wiring.py b/docs-chat-worker/gepa/test_wiring.py new file mode 100644 index 00000000..be6bd53d --- /dev/null +++ b/docs-chat-worker/gepa/test_wiring.py @@ -0,0 +1,59 @@ +"""Wiring test: full GEPA loop with a mocked student and judge. + +Costs a few cents (the reflection LM is real — it must be, to test that +'openrouter/openai/gpt-5.6-sol' works as a reflection_lm string) but makes +zero Mixedbread calls. The mock judge rewards prompts that mention citing, +so a working loop should discover a prompt containing 'cite' and beat the +seed's 0.35 val score. + + python test_wiring.py +""" + +import run as runmod + +runmod.load_env() + +import gepa + +import adapter +import judge as judgemod + + +def fake_student(system_prompt, question): + has_cite = "cite" in system_prompt.lower() + return f"stub answer; prompt_mentions_citing={has_cite}", [ + {"type": "store_search_call", "queries": [question]} + ] + + +def fake_judge(question, answer, key_points, citations): + if "prompt_mentions_citing=True" in answer: + return 0.85, "Good: cites sources.", {} + return 0.35, "Answer never cites file paths. Instruct the assistant to cite sources.", {} + + +adapter.call_student = fake_student +adapter.judge = fake_judge +judgemod.judge = fake_judge + +trainset = [ + {"question": f"q{i}", "key_points": ["kp"], "citations": []} for i in range(4) +] + +result = gepa.optimize( + seed_candidate={"system_prompt": "You are the DSRs docs assistant. Answer concisely."}, + trainset=trainset, + valset=trainset, + adapter=adapter.DocsChatAdapter(max_workers=2), + reflection_lm="openrouter/openai/gpt-5.6-sol", + reflection_minibatch_size=2, + max_metric_calls=30, + display_progress_bar=False, + seed=0, +) + +print("candidates explored:", result.num_candidates) +print("val scores:", [round(s, 2) for s in result.val_aggregate_scores]) +print("best prompt:", repr(result.best_candidate["system_prompt"][:200])) +assert result.val_aggregate_scores[result.best_idx] > 0.35, "loop never improved on seed" +print("WIRING OK") diff --git a/docs-chat-worker/worker.js b/docs-chat-worker/worker.js index 8ed54b56..00929d8a 100644 --- a/docs-chat-worker/worker.js +++ b/docs-chat-worker/worker.js @@ -7,13 +7,7 @@ // Optional vars (wrangler.toml): STORE (default dsrs-docs), // ALLOWED_ORIGINS (comma-separated; default *). -const SYSTEM = `You are the DSRs docs assistant. DSRs (dspy-rs) is a Rust \ -framework for building and optimizing LM pipelines. Answer only from the \ -attached stores: dsrs-docs (published documentation) and dsrs-code (the \ -Rust sources under crates/). Search before answering; use grep for exact \ -symbol or signature lookups. When an answer touches implementation, quote \ -the real code with its file path. If the stores don't cover the question, \ -say so instead of guessing. Answer in concise markdown.`; +const SYSTEM = "You are the DSRs documentation assistant. DSRs (`dspy-rs`) is a Rust framework for building and optimizing language-model pipelines.\n\n## Input\n\nThe user provides a natural-language question about DSRs APIs, behavior, persistence formats, errors, optimizers, or implementation details.\n\n## Grounding requirements\n\nAnswer only from the attached stores:\n\n- `dsrs-docs`: published DSRs documentation\n- `dsrs-code`: Rust source files under `crates/`\n\nDo not answer from general Rust knowledge, DSPy knowledge, memory, or inference when the stores do not establish the claim. If the stores do not cover the question, say so explicitly.\n\nAlways search before answering:\n\n1. Search the documentation for the concept and terminology.\n2. Use grep or exact-symbol search in `dsrs-code` for type definitions, enum variants, method signatures, serde attributes, and behavior.\n3. Inspect the implementation when the question concerns runtime behavior rather than only API shape.\n4. Reconcile similarly named entry points carefully\u2014for example, trait-level `compile`, typed `compile_module` helpers, and `compile_program`.\n5. If documentation and code differ, state the discrepancy rather than silently choosing or guessing.\n\nFor implementation claims, include the smallest relevant real code excerpt and its repository file path. Never invent signatures, fields, examples, or paths.\n\n## Answer style\n\n- Use concise Markdown.\n- Answer exactly what was asked; avoid unrelated report-field inventories, introductory filler, or broad background.\n- Prefer a compact table when comparing variants or optimizers.\n- Distinguish public API signatures from implementation behavior.\n- Include important asymmetric, compatibility, retry, or failure behavior.\n- Preserve exact Rust symbol names and types from the source.\n- Cite the relevant documentation component and/or source path.\n- Do not claim behavior merely because it seems conventional.\n\n## Repository-specific facts that must be handled correctly\n\nTreat the following as known guidance, but still verify exact spellings and signatures against the stores before quoting them.\n\n### Module state persistence\n\nWhen explaining save/load behavior:\n\n- Show the `ModuleState::from_module`, `save`, `load`, and `apply` workflow if those signatures are confirmed by source.\n- `ModuleState` stores one entry per saved predictor in:\n `predictors: BTreeMap`.\n- Explain that `BTreeMap` ordering makes serialized JSON stable across runs.\n- Predictor keys are discovered module paths such as `answerer` or `inner.drafter`.\n- `PredictState` contains:\n - `demos`: a vector of flat JSON objects in which each demo\u2019s input and output fields are merged.\n - `instruction_override: Option`; JSON `null` means to use the signature\u2019s default instruction.\n- Applying state is asymmetric:\n - A saved state entry whose predictor path does not exist in the target module is an error.\n - Predictors present in the target module but omitted from the saved state are left untouched.\n- The format has no version field.\n- Field-level backward compatibility comes from `#[serde(default)]` on both `PredictState` fields.\n- Cite or quote the actual state implementation, referred to in the docs as `components/state`; do not substitute an example-file reference for the defining implementation.\n\n### `PredictError`\n\nWhen asked what can fail or what is retryable:\n\n- The four `PredictError` variants are `Lm`, `Parse`, `Conversion`, and `Replay`.\n- At the `PredictError` level, only `Parse` is retryable.\n- Do not tell callers to retry `Lm` errors based on an underlying transport classification. The rig client owns transport-level retries.\n- `Parse` errors carry both `raw_response` and `lm_usage`; failed parses therefore still account for consumed LM usage.\n- Mention both relevant methods when applicable:\n - `PredictError::is_retryable()`\n - `PredictError::class()`\n- `PredictError::class()` uses the four `ErrorClass` buckets:\n `BadRequest`, `Temporary`, `BadResponse`, and `Internal`.\n- Verify the exact variant-to-class mapping in source before presenting it; do not omit `BadRequest` or infer mappings.\n- Cite or quote the implementation documented under `components/predict`.\n\n### Optimizer return types\n\nKeep trait-level and typed entry points separate:\n\n- The optimizer trait\u2019s `compile` returns the unified `Report` enum, not an optimizer-specific associated report type.\n- The unified variants include:\n `Report::None`, `Report::Gepa`, `Report::Simba`,\n `Report::Bootstrap`, and `Report::Custom`.\n- Typed `compile_module` helpers unwrap the corresponding unified report and return the optimizer-specific result:\n - COPRO: `()`\n - MIPROv2: `()`\n - GEPA: `GEPAResult`\n - SIMBA: `SimbaReport`\n - BootstrapFewShot: `BootstrapReport`\n- Verify exact capitalization and aliases in source before answering.\n- Structural optimization is the exception: it changes program structure, operates on an interpreter-loaded program, and uses `compile_program` rather than `compile_module`.\n- For a return-type question, provide the entry-point/optimizer-to-return-type mapping but omit detailed report field lists and unrelated mutation details.\n- Cite `components/optimizers` and quote the relevant signatures or enum definition when implementation is discussed.\n\n## Quality check before responding\n\nConfirm that:\n\n- Every factual claim is supported by one of the two stores.\n- Exact symbols were looked up rather than reconstructed from memory.\n- Retryability is not confused with lower-level provider retry behavior.\n- Serialization details include ordering, flattened demos, null semantics, apply asymmetry, and compatibility when relevant.\n- Optimizer answers distinguish `compile`, `compile_module`, and Structural\u2019s `compile_program`.\n- The response contains no unsupported embellishment or unnecessary prose."; export default { async fetch(request, env) { diff --git a/docs/README.md b/docs/README.md index cebb9893..d2ab0224 100644 --- a/docs/README.md +++ b/docs/README.md @@ -35,6 +35,7 @@ public-API change, run from the repository root: RUSTC_BOOTSTRAP=1 cargo rustdoc -p dspy-rs --lib --all-features -- -Z unstable-options --output-format json RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs-tools --lib -- -Z unstable-options --output-format json RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs_macros --lib -- -Z unstable-options --output-format json +RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs-syntax --lib -- -Z unstable-options --output-format json python3 docs/scripts/gen_api.py ``` diff --git a/CURRENT_PLAN.md b/docs/archive/CURRENT_PLAN.md similarity index 99% rename from CURRENT_PLAN.md rename to docs/archive/CURRENT_PLAN.md index 33eafe2b..3d1fa34c 100644 --- a/CURRENT_PLAN.md +++ b/docs/archive/CURRENT_PLAN.md @@ -1,3 +1,5 @@ +> Archived 2026-08-19: superseded by docs/v1-vision-report.md and docs/rfcs/. Retained for history. + > Status Update (2026-02-08): **Superseded historical plan**. > > Phase 1 (Bridge Root Excision) is now the active baseline: legacy bridge crates are removed from the workspace. diff --git a/CURRENT_SPEC.md b/docs/archive/CURRENT_SPEC.md similarity index 99% rename from CURRENT_SPEC.md rename to docs/archive/CURRENT_SPEC.md index 4f390c5c..9c661630 100644 --- a/CURRENT_SPEC.md +++ b/docs/archive/CURRENT_SPEC.md @@ -1,3 +1,5 @@ +> Archived 2026-08-19: superseded by docs/v1-vision-report.md and docs/rfcs/. Retained for history. + # DSRs + BAML Integration Specification **Version:** 0.1.0 diff --git a/docs/docs.json b/docs/docs.json index 370054a8..613e8c77 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -41,6 +41,7 @@ "docs/components/holes", "docs/components/capabilities", "docs/components/program-and-nodes", + "docs/components/edit-calculus", "docs/components/dsrs-file", "docs/components/runtime", "docs/components/cli", @@ -63,6 +64,7 @@ "docs/optimizers/copro", "docs/optimizers/miprov2", "docs/optimizers/gepa", + "docs/optimizers/structural", "docs/components/optimizer-engine" ] }, @@ -97,6 +99,7 @@ "docs/api/modules", "docs/api/optimizer", "docs/api/predictors", + "docs/api/prelude", "docs/api/trace", "docs/api/typesys", "docs/api/utils" @@ -106,7 +109,8 @@ "group": "Companion crates", "pages": [ "docs/api/dsrs-tools", - "docs/api/dsrs-macros" + "docs/api/dsrs-macros", + "docs/api/dsrs-syntax" ] } ] @@ -114,70 +118,262 @@ ] }, "redirects": [ - { "source": "/docs/building-blocks/signature", "destination": "/docs/components/signatures" }, - { "source": "/docs/building-blocks/types", "destination": "/docs/components/signatures" }, - { "source": "/docs/building-blocks/constraints", "destination": "/docs/components/signatures" }, - { "source": "/docs/building-blocks/predictors", "destination": "/docs/components/predict" }, - { "source": "/docs/building-blocks/adapter", "destination": "/docs/components/adapters" }, - { "source": "/docs/building-blocks/lm", "destination": "/docs/components/lm" }, - { "source": "/docs/building-blocks/module", "destination": "/docs/components/modules" }, - { "source": "/docs/data/dataloader", "destination": "/docs/components/data" }, - { "source": "/docs/data/examples", "destination": "/docs/components/data" }, - { "source": "/docs/data/prediction", "destination": "/docs/components/predict" }, - { "source": "/docs/getting-started/introduction", "destination": "/docs/getting-started/how-dsrs-thinks" }, - { "source": "/docs/tutorials/overview", "destination": "/docs/components/module-macro" }, - { "source": "/docs/tutorials/your-first-module", "destination": "/docs/components/module-macro" }, - { "source": "/docs/tutorials/agent-with-tools", "destination": "/docs/components/tools-and-agents" }, - { "source": "/docs/tutorials/record-and-replay", "destination": "/docs/components/traces" }, - { "source": "/docs/tutorials/tune-a-module", "destination": "/docs/components/optimizers" }, - { "source": "/docs/concepts/harness-as-data", "destination": "/docs/getting-started/how-dsrs-thinks" }, - { "source": "/docs/concepts/three-lanes", "destination": "/docs/getting-started/how-dsrs-thinks" }, - { "source": "/docs/concepts/compilation", "destination": "/docs/getting-started/how-dsrs-thinks" }, - { "source": "/docs/concepts/holes", "destination": "/docs/components/holes" }, - { "source": "/docs/concepts/traces-and-replay", "destination": "/docs/components/traces" }, - { "source": "/docs/concepts/capabilities", "destination": "/docs/components/capabilities" }, - { "source": "/docs/reference/macros", "destination": "/docs/components/module-macro" }, - { "source": "/docs/reference/signatures-and-types", "destination": "/docs/components/signatures" }, - { "source": "/docs/reference/predict", "destination": "/docs/components/predict" }, - { "source": "/docs/reference/modules", "destination": "/docs/components/modules" }, - { "source": "/docs/reference/adapters", "destination": "/docs/components/adapters" }, - { "source": "/docs/reference/lm", "destination": "/docs/components/lm" }, - { "source": "/docs/reference/data", "destination": "/docs/components/data" }, - { "source": "/docs/reference/state", "destination": "/docs/components/state" }, - { "source": "/docs/reference/fx", "destination": "/docs/components/fx" }, - { "source": "/docs/reference/evaluate", "destination": "/docs/components/evaluation" }, - { "source": "/docs/reference/optimizers", "destination": "/docs/components/optimizers" }, - { "source": "/docs/reference/optimizer-engine", "destination": "/docs/components/optimizer-engine" }, - { "source": "/docs/reference/code-mode", "destination": "/docs/components/code-mode" }, - { "source": "/docs/reference/dsrs-file", "destination": "/docs/components/dsrs-file" }, - { "source": "/docs/reference/program-and-nodes", "destination": "/docs/components/program-and-nodes" }, - { "source": "/docs/reference/runtime", "destination": "/docs/components/runtime" }, - { "source": "/docs/reference/cli", "destination": "/docs/components/cli" }, - { "source": "/docs/reference/traces", "destination": "/docs/components/traces" }, - { "source": "/docs/reference/traces-and-cli", "destination": "/docs/components/traces" }, - { "source": "/docs/reference/utils", "destination": "/docs/components/utils" }, - { "source": "/docs/learn/you-never-write-the-prompt", "destination": "/docs/getting-started/quickstart" }, - { "source": "/docs/learn/the-contract-is-the-prompt", "destination": "/docs/components/signatures" }, - { "source": "/docs/learn/who-calls-the-model", "destination": "/docs/components/predict" }, - { "source": "/docs/learn/what-the-model-sees", "destination": "/docs/components/adapters" }, - { "source": "/docs/learn/pipelines-are-just-structs", "destination": "/docs/components/modules" }, - { "source": "/docs/learn/one-parse-two-projections", "destination": "/docs/components/module-macro" }, - { "source": "/docs/learn/here-be-dragons", "destination": "/docs/components/holes" }, - { "source": "/docs/learn/give-the-model-hands", "destination": "/docs/components/tools-and-agents" }, - { "source": "/docs/learn/nothing-to-declare", "destination": "/docs/components/capabilities" }, - { "source": "/docs/learn/the-expedition-log", "destination": "/docs/components/traces" }, - { "source": "/docs/learn/what-does-better-mean", "destination": "/docs/components/evaluation" }, - { "source": "/docs/learn/tracing-paper", "destination": "/docs/components/optimizers" }, - { "source": "/docs/learn/when-the-judge-is-a-model", "destination": "/docs/optimizers/gepa" }, - { "source": "/docs/learn/the-map-leaves-home", "destination": "/docs/components/cli" }, - { "source": "/docs/learn/running-in-the-wild", "destination": "/docs/components/traces" }, - { "source": "/docs/guides/write-steps", "destination": "/docs/components/module-macro" }, - { "source": "/docs/guides/add-tools", "destination": "/docs/components/tools-and-agents" }, - { "source": "/docs/guides/use-holes", "destination": "/docs/components/holes" }, - { "source": "/docs/guides/print-and-serve", "destination": "/docs/components/cli" }, - { "source": "/docs/guides/bake-a-candidate", "destination": "/docs/components/program-and-nodes" }, - { "source": "/docs/optimizers/gepa-llm-judge", "destination": "/docs/optimizers/gepa" }, - { "source": "/community", "destination": "/" } + { + "source": "/docs/building-blocks/signature", + "destination": "/docs/components/signatures" + }, + { + "source": "/docs/building-blocks/types", + "destination": "/docs/components/signatures" + }, + { + "source": "/docs/building-blocks/constraints", + "destination": "/docs/components/signatures" + }, + { + "source": "/docs/building-blocks/predictors", + "destination": "/docs/components/predict" + }, + { + "source": "/docs/building-blocks/adapter", + "destination": "/docs/components/adapters" + }, + { + "source": "/docs/building-blocks/lm", + "destination": "/docs/components/lm" + }, + { + "source": "/docs/building-blocks/module", + "destination": "/docs/components/modules" + }, + { + "source": "/docs/data/dataloader", + "destination": "/docs/components/data" + }, + { + "source": "/docs/data/examples", + "destination": "/docs/components/data" + }, + { + "source": "/docs/data/prediction", + "destination": "/docs/components/predict" + }, + { + "source": "/docs/getting-started/introduction", + "destination": "/docs/getting-started/how-dsrs-thinks" + }, + { + "source": "/docs/tutorials/overview", + "destination": "/docs/components/module-macro" + }, + { + "source": "/docs/tutorials/your-first-module", + "destination": "/docs/components/module-macro" + }, + { + "source": "/docs/tutorials/agent-with-tools", + "destination": "/docs/components/tools-and-agents" + }, + { + "source": "/docs/tutorials/record-and-replay", + "destination": "/docs/components/traces" + }, + { + "source": "/docs/tutorials/tune-a-module", + "destination": "/docs/components/optimizers" + }, + { + "source": "/docs/concepts/harness-as-data", + "destination": "/docs/getting-started/how-dsrs-thinks" + }, + { + "source": "/docs/concepts/three-lanes", + "destination": "/docs/getting-started/how-dsrs-thinks" + }, + { + "source": "/docs/concepts/compilation", + "destination": "/docs/getting-started/how-dsrs-thinks" + }, + { + "source": "/docs/concepts/holes", + "destination": "/docs/components/holes" + }, + { + "source": "/docs/concepts/traces-and-replay", + "destination": "/docs/components/traces" + }, + { + "source": "/docs/concepts/capabilities", + "destination": "/docs/components/capabilities" + }, + { + "source": "/docs/reference/macros", + "destination": "/docs/components/module-macro" + }, + { + "source": "/docs/reference/signatures-and-types", + "destination": "/docs/components/signatures" + }, + { + "source": "/docs/reference/predict", + "destination": "/docs/components/predict" + }, + { + "source": "/docs/reference/modules", + "destination": "/docs/components/modules" + }, + { + "source": "/docs/reference/adapters", + "destination": "/docs/components/adapters" + }, + { + "source": "/docs/reference/lm", + "destination": "/docs/components/lm" + }, + { + "source": "/docs/reference/data", + "destination": "/docs/components/data" + }, + { + "source": "/docs/reference/state", + "destination": "/docs/components/state" + }, + { + "source": "/docs/reference/fx", + "destination": "/docs/components/fx" + }, + { + "source": "/docs/reference/evaluate", + "destination": "/docs/components/evaluation" + }, + { + "source": "/docs/reference/optimizers", + "destination": "/docs/components/optimizers" + }, + { + "source": "/docs/reference/optimizer-engine", + "destination": "/docs/components/optimizer-engine" + }, + { + "source": "/docs/reference/code-mode", + "destination": "/docs/components/code-mode" + }, + { + "source": "/docs/reference/dsrs-file", + "destination": "/docs/components/dsrs-file" + }, + { + "source": "/docs/reference/program-and-nodes", + "destination": "/docs/components/program-and-nodes" + }, + { + "source": "/docs/reference/runtime", + "destination": "/docs/components/runtime" + }, + { + "source": "/docs/reference/cli", + "destination": "/docs/components/cli" + }, + { + "source": "/docs/reference/traces", + "destination": "/docs/components/traces" + }, + { + "source": "/docs/reference/traces-and-cli", + "destination": "/docs/components/traces" + }, + { + "source": "/docs/reference/utils", + "destination": "/docs/components/utils" + }, + { + "source": "/docs/learn/you-never-write-the-prompt", + "destination": "/docs/getting-started/quickstart" + }, + { + "source": "/docs/learn/the-contract-is-the-prompt", + "destination": "/docs/components/signatures" + }, + { + "source": "/docs/learn/who-calls-the-model", + "destination": "/docs/components/predict" + }, + { + "source": "/docs/learn/what-the-model-sees", + "destination": "/docs/components/adapters" + }, + { + "source": "/docs/learn/pipelines-are-just-structs", + "destination": "/docs/components/modules" + }, + { + "source": "/docs/learn/one-parse-two-projections", + "destination": "/docs/components/module-macro" + }, + { + "source": "/docs/learn/here-be-dragons", + "destination": "/docs/components/holes" + }, + { + "source": "/docs/learn/give-the-model-hands", + "destination": "/docs/components/tools-and-agents" + }, + { + "source": "/docs/learn/nothing-to-declare", + "destination": "/docs/components/capabilities" + }, + { + "source": "/docs/learn/the-expedition-log", + "destination": "/docs/components/traces" + }, + { + "source": "/docs/learn/what-does-better-mean", + "destination": "/docs/components/evaluation" + }, + { + "source": "/docs/learn/tracing-paper", + "destination": "/docs/components/optimizers" + }, + { + "source": "/docs/learn/when-the-judge-is-a-model", + "destination": "/docs/optimizers/gepa" + }, + { + "source": "/docs/learn/the-map-leaves-home", + "destination": "/docs/components/cli" + }, + { + "source": "/docs/learn/running-in-the-wild", + "destination": "/docs/components/traces" + }, + { + "source": "/docs/guides/write-steps", + "destination": "/docs/components/module-macro" + }, + { + "source": "/docs/guides/add-tools", + "destination": "/docs/components/tools-and-agents" + }, + { + "source": "/docs/guides/use-holes", + "destination": "/docs/components/holes" + }, + { + "source": "/docs/guides/print-and-serve", + "destination": "/docs/components/cli" + }, + { + "source": "/docs/guides/bake-a-candidate", + "destination": "/docs/components/program-and-nodes" + }, + { + "source": "/docs/optimizers/gepa-llm-judge", + "destination": "/docs/optimizers/gepa" + }, + { + "source": "/community", + "destination": "/" + } ], "logo": { "light": "/logo/main.png", diff --git a/docs/docs/api/adapter.mdx b/docs/docs/api/adapter.mdx index f116ddbc..2406eae2 100644 --- a/docs/docs/api/adapter.mdx +++ b/docs/docs/api/adapter.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. Prompt formatting and LM response parsing. diff --git a/docs/docs/api/augmentation.mdx b/docs/docs/api/augmentation.mdx index d308f9e9..dd20ac85 100644 --- a/docs/docs/api/augmentation.mdx +++ b/docs/docs/api/augmentation.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. ## Structs diff --git a/docs/docs/api/core.mdx b/docs/docs/api/core.mdx index f2b6a3aa..0984a94c 100644 --- a/docs/docs/api/core.mdx +++ b/docs/docs/api/core.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. The foundational abstractions everything else is built on. @@ -27,12 +27,11 @@ The foundational abstractions everything else is built on. | `lm::*` | Glob re-export. | | [`LmError`](https://docs.rs/dspy-rs/latest/dspy_rs/core/errors/enum.LmError.html) | Re-export of `errors::LmError`. | | `module::*` | Glob re-export. | -| `module_ext::*` | Glob re-export. | | [`ModuleState`](https://docs.rs/dspy-rs/latest/dspy_rs/core/state/struct.ModuleState.html) | Re-export of `state::ModuleState`. | | [`ParseError`](https://docs.rs/dspy-rs/latest/dspy_rs/core/errors/enum.ParseError.html) | Re-export of `errors::ParseError`. | | [`Predicted`](https://docs.rs/dspy-rs/latest/dspy_rs/core/predicted/struct.Predicted.html) | Re-export of `predicted::Predicted`. | | [`PredictError`](https://docs.rs/dspy-rs/latest/dspy_rs/core/errors/enum.PredictError.html) | Re-export of `errors::PredictError`. | -| [`PredictState`](https://docs.rs/dspy-rs/latest/dspy_rs/core/dyn_predictor/struct.PredictState.html) | Re-export of `dyn_predictor::PredictState`. | +| [`PredictState`](https://docs.rs/dspy-rs/latest/dspy_rs/core/state/struct.PredictState.html) | Re-export of `state::PredictState`. | | `settings::*` | Glob re-export. | | `signature::*` | Glob re-export. | | [`SignatureSchema`](https://docs.rs/dspy-rs/latest/dspy_rs/core/schema/struct.SignatureSchema.html) | Re-export of `schema::SignatureSchema`. | @@ -143,6 +142,8 @@ How trainset rows connect to modules. | Item | Description | |---|---| | [`Module`](https://docs.rs/dspy-rs/latest/dspy_rs/core/module/trait.Module.html) | Strategy-swapping interface for prompting modules. | +| [`PredictorInfo`](https://docs.rs/dspy-rs/latest/dspy_rs/core/module/trait.PredictorInfo.html) | What optimizers read from — and, at explicit boundaries, write to — a `Predict` leaf. | +| [`Predictors`](https://docs.rs/dspy-rs/latest/dspy_rs/core/module/trait.Predictors.html) | Explicit predictor-leaf discovery: a module *names* its optimizable `Predict` leaves. | ### Functions diff --git a/docs/docs/api/data.mdx b/docs/docs/api/data.mdx index f1d8f410..13e27d35 100644 --- a/docs/docs/api/data.mdx +++ b/docs/docs/api/data.mdx @@ -1,59 +1,50 @@ --- title: "dspy_rs::data" -description: "Data loading, versioned under `v1`." +description: "Data loading." icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. -Data loading, versioned under [`v1`]. +Data loading. ## Re-exports | Item | Description | |---|---| -| `v1::*` | Glob re-export. | +| `dataloader::*` | Glob re-export. | +| `utils::*` | Glob re-export. | ## Modules | Item | Description | |---|---| -| [`v1`](https://docs.rs/dspy-rs/latest/dspy_rs/data/v1/index.html) | Data loading. | - -## `data::v1` - -Data loading. - -### Re-exports - -| Item | Description | -|---|---| -| `dataloader::*` | Glob re-export. | -| `utils::*` | Glob re-export. | +| [`dataloader`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/index.html) | | +| [`utils`](https://docs.rs/dspy-rs/latest/dspy_rs/data/utils/index.html) | | -## `data::v1::dataloader` +## `data::dataloader` ### Structs | Item | Description | |---|---| -| [`DataLoader`](https://docs.rs/dspy-rs/latest/dspy_rs/data/v1/dataloader/struct.DataLoader.html) | Typed dataset ingress for JSON/CSV/Parquet/HuggingFace sources. | -| [`RowRecord`](https://docs.rs/dspy-rs/latest/dspy_rs/data/v1/dataloader/struct.RowRecord.html) | Raw parsed row passed to custom mapper closures in `load_*_with` APIs. | -| [`TypedLoadOptions`](https://docs.rs/dspy-rs/latest/dspy_rs/data/v1/dataloader/struct.TypedLoadOptions.html) | Options for shape-driven typed loading. | +| [`DataLoader`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/struct.DataLoader.html) | Typed dataset ingress for JSON/CSV/Parquet/HuggingFace sources. | +| [`RowRecord`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/struct.RowRecord.html) | Raw parsed row passed to custom mapper closures in `load_*_with` APIs. | +| [`TypedLoadOptions`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/struct.TypedLoadOptions.html) | Options for shape-driven typed loading. | ### Enums | Item | Description | |---|---| -| [`DataLoadError`](https://docs.rs/dspy-rs/latest/dspy_rs/data/v1/dataloader/enum.DataLoadError.html) | Row-aware errors produced by typed data loading. | -| [`UnknownFieldPolicy`](https://docs.rs/dspy-rs/latest/dspy_rs/data/v1/dataloader/enum.UnknownFieldPolicy.html) | Controls how typed loaders handle source fields that are not part of the target row struct. | +| [`DataLoadError`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/enum.DataLoadError.html) | Row-aware errors produced by typed data loading. | +| [`UnknownFieldPolicy`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/enum.UnknownFieldPolicy.html) | Controls how typed loaders handle source fields that are not part of the target row struct. | -## `data::v1::utils` +## `data::utils` ### Functions | Item | Description | |---|---| -| [`is_url`](https://docs.rs/dspy-rs/latest/dspy_rs/data/v1/utils/fn.is_url.html) | Returns `true` if the string looks like an HTTP(S) URL. | +| [`is_url`](https://docs.rs/dspy-rs/latest/dspy_rs/data/utils/fn.is_url.html) | Returns `true` if the string looks like an HTTP(S) URL. | diff --git a/docs/docs/api/dspy-rs.mdx b/docs/docs/api/dspy-rs.mdx index 3ac9528a..f07518bf 100644 --- a/docs/docs/api/dspy-rs.mdx +++ b/docs/docs/api/dspy-rs.mdx @@ -5,7 +5,7 @@ icon: "box-open" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. The crate root re-exports the surface most programs use, so `use dspy_rs::{Predict, Signature, configure}` works without module paths. Items are listed here once with their home module linked; the module pages list everything else. @@ -25,14 +25,12 @@ The crate root re-exports the surface most programs use, so `use dspy_rs::{Predi | [`CompId`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/struct.CompId.html) | Re-export of `trace::CompId`. | | [`Constraint`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.Constraint.html) | Re-export of `typesys::Constraint`. | | [`ConstraintLevel`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/type.ConstraintLevel.html) | Re-export of `typesys::ConstraintLevel`. | -| [`ConstraintOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.ConstraintOutcome.html) | Re-export of `typesys::ConstraintOutcome`. | | `core::*` | Glob re-export. | | `data::dataloader::*` | Glob re-export. | | `data::utils::*` | Glob re-export. | | `dsrs_macros::*` | Glob re-export. | | [`Eval`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/struct.Eval.html) | Re-export of `trace::Eval`. | | `evaluate::*` | Glob re-export. | -| [`evaluate_constraints`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/fn.evaluate_constraints.html) | Re-export of `typesys::evaluate_constraints`. | | [`Facet`](https://docs.rs/dspy-rs/latest/facet_core/trait.Facet.html) | Re-export of `facet::Facet`. | | [`Facet`](https://docs.rs/dspy-rs/latest/facet_macros/derive.Facet.html) | Re-export of `facet::Facet`. | | [`FieldType`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/enum.FieldType.html) | Re-export of `typesys::FieldType`. | @@ -44,11 +42,6 @@ The crate root re-exports the surface most programs use, so `use dspy_rs::{Predi | [`ModelId`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/struct.ModelId.html) | Re-export of `trace::ModelId`. | | `modules::*` | Glob re-export. | | `optimizer::*` | Glob re-export. | -| [`OtelEvent`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelEvent.html) | Re-export of `trace::OtelEvent`. | -| [`OtelKeyValue`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelKeyValue.html) | Re-export of `trace::OtelKeyValue`. | -| [`OtelSpan`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelSpan.html) | Re-export of `trace::OtelSpan`. | -| [`OtelStatus`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelStatus.html) | Re-export of `trace::OtelStatus`. | -| [`OtelValue`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/enum.OtelValue.html) | Re-export of `trace::OtelValue`. | | [`OutputSchema`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.OutputSchema.html) | Re-export of `typesys::OutputSchema`. | | `predictors::*` | Glob re-export. | | [`PrefixEntry`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/struct.PrefixEntry.html) | Re-export of `trace::PrefixEntry`. | @@ -58,9 +51,6 @@ The crate root re-exports the surface most programs use, so `use dspy_rs::{Predi | [`ReplayError`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/enum.ReplayError.html) | Re-export of `trace::ReplayError`. | | [`ReplayMode`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/enum.ReplayMode.html) | Re-export of `trace::ReplayMode`. | | [`ReplayReport`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/struct.ReplayReport.html) | Re-export of `trace::ReplayReport`. | -| [`ResponseCheck`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.ResponseCheck.html) | Re-export of `typesys::ResponseCheck`. | -| [`RlRollout`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/rl/struct.RlRollout.html) | Re-export of `trace::RlRollout`. | -| [`RlTransition`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/rl/struct.RlTransition.html) | Re-export of `trace::RlTransition`. | | [`RUN_JS_TOOL_NAME`](https://docs.rs/dspy-rs/latest/dsrs_tools/code_mode/constant.RUN_JS_TOOL_NAME.html) | Re-export of `dsrs_tools::RUN_JS_TOOL_NAME`. | | [`SandboxConfig`](https://docs.rs/dspy-rs/latest/dsrs_tools/quickjs/struct.SandboxConfig.html) | Re-export of `dsrs_tools::SandboxConfig`. | | [`Schema`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/trait.Schema.html) | Re-export of `typesys::Schema`. | @@ -82,7 +72,7 @@ The crate root re-exports the surface most programs use, so `use dspy_rs::{Predi | Item | Description | |---|---| -| [`sign`](https://docs.rs/dspy-rs/latest/dspy_rs/macro.sign.html) | | +| [`predictors`](https://docs.rs/dspy-rs/latest/dspy_rs/macro.predictors.html) | Implements `Predictors` for a module struct from a list of predictor fields, using each field's identifier as its leaf name. | ## Modules @@ -91,13 +81,14 @@ The crate root re-exports the surface most programs use, so `use dspy_rs::{Predi | [`adapter`](/docs/api/adapter) | Prompt formatting and LM response parsing. | | [`augmentation`](/docs/api/augmentation) | | | [`core`](/docs/api/core) | The foundational abstractions everything else is built on. | -| [`data`](/docs/api/data) | Data loading, versioned under `v1`. | +| [`data`](/docs/api/data) | Data loading. | | [`evaluate`](/docs/api/evaluate) | Evaluation and metrics for measuring module performance. | | [`fx`](/docs/api/fx) | Functional DSRs (experimental): harnesses as plain async functions. | | [`ir`](/docs/api/ir) | The intermediate representation (RFC 0002). | | [`modules`](/docs/api/modules) | | | [`optimizer`](/docs/api/optimizer) | Automatic prompt optimization. | | [`predictors`](/docs/api/predictors) | | +| [`prelude`](/docs/api/prelude) | The curated core surface — the recommended import for DSRs programs. | | [`trace`](/docs/api/trace) | Execution trace capture (RFC 0001). | | [`typesys`](/docs/api/typesys) | In-house type system that replaces the vendored BAML stack (`bamltype`, `baml_types`, `internal_baml_jinja`, `jsonish`). | | [`utils`](/docs/api/utils) | LM response caching. | diff --git a/docs/docs/api/dsrs-macros.mdx b/docs/docs/api/dsrs-macros.mdx index 04b8433b..3ead01ab 100644 --- a/docs/docs/api/dsrs-macros.mdx +++ b/docs/docs/api/dsrs-macros.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dsrs-macros v0.7.2` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dsrs-macros v0.7.2` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. ## Proc macros @@ -14,7 +14,6 @@ Generated from rustdoc JSON at `dsrs-macros v0.7.2` (commit `b5da857`). Do not e |---|---| | [`agent`](https://docs.rs/dsrs-macros/latest/dsrs_macros/attr.agent.html) | (attribute) Declares an LLM+tool loop as a bodyless fn (RFC 0003 M-2) — the agent sibling of `macro@predict`. | | [`Augmentation`](https://docs.rs/dsrs-macros/latest/dsrs_macros/) | (derive) | -| [`BamlType`](https://docs.rs/dsrs-macros/latest/dsrs_macros/attr.BamlType.html) | (attribute) Backwards-compatible alias for `macro@Schema`. | | [`cot`](https://docs.rs/dsrs-macros/latest/dsrs_macros/attr.cot.html) | (attribute) `macro@predict` with chain-of-thought: the LM produces a `reasoning` field before the output, and the generated function returns `Predicted +Generated from rustdoc JSON at `dsrs-syntax v0.1.0` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. + + +Shared syntax layer for the `.dsrs` text format (RFC 0002 §4). + +## Re-exports + +| Item | Description | +|---|---| +| [`check`](https://docs.rs/dsrs-syntax/latest/dsrs_syntax/structure/fn.check.html) | Re-export of `structure::check`. | + +## Modules + +| Item | Description | +|---|---| +| [`lex`](https://docs.rs/dsrs-syntax/latest/dsrs_syntax/lex/index.html) | Lexer for the `.dsrs` text format (RFC 0002 §4). | + +## Structs + +| Item | Description | +|---|---| +| [`ParseError`](https://docs.rs/dsrs-syntax/latest/dsrs_syntax/struct.ParseError.html) | A parse failure with the source position and what was expected — designed to be actionable feedback for a model regenerating the program. | + +## `dsrs_syntax::lex` + +### Structs + +| Item | Description | +|---|---| +| [`Lexed`](https://docs.rs/dsrs-syntax/latest/dsrs_syntax/lex/struct.Lexed.html) | One lexed token with its source position and byte extent. | +| [`Lexer`](https://docs.rs/dsrs-syntax/latest/dsrs_syntax/lex/struct.Lexer.html) | | +| [`Span`](https://docs.rs/dsrs-syntax/latest/dsrs_syntax/lex/struct.Span.html) | A source position, 1-based. | + +### Enums + +| Item | Description | +|---|---| +| [`Tok`](https://docs.rs/dsrs-syntax/latest/dsrs_syntax/lex/enum.Tok.html) | | diff --git a/docs/docs/api/dsrs-tools.mdx b/docs/docs/api/dsrs-tools.mdx index 7d4f3dcd..91c77e3d 100644 --- a/docs/docs/api/dsrs-tools.mdx +++ b/docs/docs/api/dsrs-tools.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dsrs-tools v0.1.0` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dsrs-tools v0.1.0` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. # dsrs-tools: sandboxed tool execution for DSRs diff --git a/docs/docs/api/evaluate.mdx b/docs/docs/api/evaluate.mdx index 59158f50..de290023 100644 --- a/docs/docs/api/evaluate.mdx +++ b/docs/docs/api/evaluate.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. Evaluation and metrics for measuring module performance. @@ -15,14 +15,12 @@ Evaluation and metrics for measuring module performance. | Item | Description | |---|---| | `evaluator::*` | Glob re-export. | -| `feedback_helpers::*` | Glob re-export. | ## Modules | Item | Description | |---|---| | [`evaluator`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/evaluator/index.html) | | -| [`feedback_helpers`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/index.html) | | ## `evaluate::evaluator` @@ -57,22 +55,3 @@ Evaluation and metrics for measuring module performance. | Item | Description | |---|---| | [`DEFAULT_EVAL_CONCURRENCY`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/evaluator/constant.DEFAULT_EVAL_CONCURRENCY.html) | Default number of examples evaluated concurrently by `evaluate_trainset`. | - -## `evaluate::feedback_helpers` - -### Enums - -| Item | Description | -|---|---| -| [`CodeStage`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/enum.CodeStage.html) | Stage in code execution pipeline | -| [`StageResult`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/enum.StageResult.html) | Result of a code stage | - -### Functions - -| Item | Description | -|---|---| -| [`classification_feedback`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/fn.classification_feedback.html) | Create feedback for classification tasks | -| [`code_pipeline_feedback`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/fn.code_pipeline_feedback.html) | Create feedback for code generation pipelines | -| [`multi_objective_feedback`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/fn.multi_objective_feedback.html) | Create feedback for multi-objective optimization | -| [`retrieval_feedback`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/fn.retrieval_feedback.html) | Create feedback for document retrieval tasks | -| [`string_similarity_feedback`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/feedback_helpers/fn.string_similarity_feedback.html) | Create feedback for string similarity tasks | diff --git a/docs/docs/api/fx.mdx b/docs/docs/api/fx.mdx index 8fa67914..a4ca3187 100644 --- a/docs/docs/api/fx.mdx +++ b/docs/docs/api/fx.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. Functional DSRs (experimental): harnesses as plain async functions. diff --git a/docs/docs/api/index.mdx b/docs/docs/api/index.mdx index 2c816a99..78c0ca67 100644 --- a/docs/docs/api/index.mdx +++ b/docs/docs/api/index.mdx @@ -8,7 +8,7 @@ This tab is an auto-generated index of the public API: every public struct, enum Every item links to its page on docs.rs, which carries the full signatures, methods, trait implementations, and long-form documentation: - + The main crate: signatures, predictors, modules, IR, traces, optimizers @@ -18,6 +18,9 @@ Sandboxed tool execution and Code Mode The proc-macro crate behind the derive and attribute surface + +The shared .dsrs lexer and structural grammar both frontends read from + For explanations and usage, the [component pages](/docs/components/signatures) are the right place; this tab answers "what exists and where does it live." @@ -30,7 +33,8 @@ The pages are rebuilt from the working tree, not fetched from a registry, so the RUSTC_BOOTSTRAP=1 cargo rustdoc -p dspy-rs --lib --all-features -- -Z unstable-options --output-format json RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs-tools --lib -- -Z unstable-options --output-format json RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs_macros --lib -- -Z unstable-options --output-format json +RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs-syntax --lib -- -Z unstable-options --output-format json python3 docs/scripts/gen_api.py ``` -Each generated page carries the crate version and commit it was built from. Run the four commands after any public-API change; CI can run them and fail on a dirty diff to keep this tab honest. +Each generated page carries the crate version and commit it was built from. Run the five commands after any public-API change; CI can run them and fail on a dirty diff to keep this tab honest. diff --git a/docs/docs/api/ir.mdx b/docs/docs/api/ir.mdx index 7c5af4b0..845ee7ab 100644 --- a/docs/docs/api/ir.mdx +++ b/docs/docs/api/ir.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. The intermediate representation (RFC 0002). @@ -17,6 +17,7 @@ The intermediate representation (RFC 0002). | [`agent`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.agent.html) | Re-export of `builder::agent`. | | [`AgentLoopNode`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.AgentLoopNode.html) | Re-export of `graph::AgentLoopNode`. | | [`AgentStepOpts`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/step/struct.AgentStepOpts.html) | Re-export of `step::AgentStepOpts`. | +| [`ApplyError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.ApplyError.html) | Re-export of `edit::ApplyError`. | | [`AsNodeName`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/trait.AsNodeName.html) | Re-export of `builder::AsNodeName`. | | [`BakeError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/enum.BakeError.html) | Re-export of `graph::BakeError`. | | [`Binding`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.Binding.html) | Re-export of `graph::Binding`. | @@ -34,12 +35,16 @@ The intermediate representation (RFC 0002). | [`ConstraintDef`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/sig/struct.ConstraintDef.html) | Re-export of `sig::ConstraintDef`. | | [`ContextK`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ContextK.html) | Re-export of `params::ContextK`. | | [`ContextPolicy`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/struct.ContextPolicy.html) | Re-export of `params::ContextPolicy`. | +| [`ConversationTurn`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/enum.ConversationTurn.html) | Re-export of `interp::ConversationTurn`. | | [`cot`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.cot.html) | Re-export of `builder::cot`. | | [`current_overlay`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/bridge/fn.current_overlay.html) | Re-export of `bridge::current_overlay`. | | [`default_lm`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/module_build/fn.default_lm.html) | Re-export of `module_build::default_lm`. | | [`DemoRow`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/struct.DemoRow.html) | Re-export of `params::DemoRow`. | | [`Demos`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.Demos.html) | Re-export of `params::Demos`. | | [`DsrsFileError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/text/enum.DsrsFileError.html) | Re-export of `text::DsrsFileError`. | +| [`Edit`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.Edit.html) | Re-export of `edit::Edit`. | +| [`EditError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.EditError.html) | Re-export of `edit::EditError`. | +| [`EditKind`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.EditKind.html) | Re-export of `edit::EditKind`. | | [`EnumDef`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.EnumDef.html) | Re-export of `crate::typesys::EnumDef`. | | [`EnumValueDef`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.EnumValueDef.html) | Re-export of `crate::typesys::EnumValueDef`. | | [`Exhausted`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.Exhausted.html) | Re-export of `interp::Exhausted`. | @@ -59,11 +64,13 @@ The intermediate representation (RFC 0002). | [`Interner`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.Interner.html) | Re-export of `graph::Interner`. | | [`Interpreter`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.Interpreter.html) | Re-export of `interp::Interpreter`. | | [`KindTag`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/trait.KindTag.html) | Re-export of `params::KindTag`. | +| [`LeafOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.LeafOutcome.html) | Re-export of `interp::LeafOutcome`. | | [`Lineage`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.Lineage.html) | Re-export of `graph::Lineage`. | | [`lit`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.lit.html) | Re-export of `builder::lit`. | | [`LoadError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/enum.LoadError.html) | Re-export of `interp::LoadError`. | | [`loop_`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.loop_.html) | Re-export of `builder::loop_`. | | [`LoopNode`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.LoopNode.html) | Re-export of `graph::LoopNode`. | +| [`migrate_overlay`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/fn.migrate_overlay.html) | Re-export of `edit::migrate_overlay`. | | [`ModelDef`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.ModelDef.html) | Re-export of `graph::ModelDef`. | | [`ModelId`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.ModelId.html) | Re-export of `graph::ModelId`. | | [`ModelRefK`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ModelRefK.html) | Re-export of `params::ModelRefK`. | @@ -83,7 +90,7 @@ The intermediate representation (RFC 0002). | [`ParamOwner`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ParamOwner.html) | Re-export of `params::ParamOwner`. | | [`ParamSlot`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/struct.ParamSlot.html) | Re-export of `params::ParamSlot`. | | [`ParamValue`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ParamValue.html) | Re-export of `params::ParamValue`. | -| [`ParseError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/text/struct.ParseError.html) | Re-export of `text::ParseError`. | +| [`ParseError`](https://docs.rs/dspy-rs/latest/dsrs_syntax/struct.ParseError.html) | Re-export of `text::ParseError`. | | [`Port`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/enum.Port.html) | Re-export of `builder::Port`. | | [`PortRef`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/enum.PortRef.html) | Re-export of `graph::PortRef`. | | [`PortSpec`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/module_build/enum.PortSpec.html) | Re-export of `module_build::PortSpec`. | @@ -100,6 +107,7 @@ The intermediate representation (RFC 0002). | [`route`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.route.html) | Re-export of `builder::route`. | | [`RouteNode`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.RouteNode.html) | Re-export of `graph::RouteNode`. | | [`RunError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/enum.RunError.html) | Re-export of `interp::RunError`. | +| [`RunOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.RunOutput.html) | Re-export of `interp::RunOutput`. | | [`RuntimeEnv`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.RuntimeEnv.html) | Re-export of `interp::RuntimeEnv`. | | [`seq`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.seq.html) | Re-export of `builder::seq`. | | [`SeqNode`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.SeqNode.html) | Re-export of `graph::SeqNode`. | @@ -112,12 +120,15 @@ The intermediate representation (RFC 0002). | [`StepDef`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/step/struct.StepDef.html) | Re-export of `step::StepDef`. | | [`StepKind`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/step/enum.StepKind.html) | Re-export of `step::StepKind`. | | [`StopSpec`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.StopSpec.html) | Re-export of `graph::StopSpec`. | +| [`SwapTarget`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.SwapTarget.html) | Re-export of `edit::SwapTarget`. | | [`Sym`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.Sym.html) | Re-export of `graph::Sym`. | | [`ToolDef`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.ToolDef.html) | Re-export of `graph::ToolDef`. | | [`ToolDesc`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ToolDesc.html) | Re-export of `params::ToolDesc`. | | [`ToolId`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.ToolId.html) | Re-export of `graph::ToolId`. | | [`ToolKind`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/enum.ToolKind.html) | Re-export of `graph::ToolKind`. | +| [`ToolSetK`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ToolSetK.html) | Re-export of `params::ToolSetK`. | | [`ToolStepDef`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/step/struct.ToolStepDef.html) | Re-export of `step::ToolStepDef`. | +| [`ToolSuspension`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.ToolSuspension.html) | Re-export of `interp::ToolSuspension`. | | [`TypeTable`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.TypeTable.html) | Re-export of `crate::typesys::TypeTable`. | | [`unbound_model_config`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/module_build/fn.unbound_model_config.html) | Re-export of `module_build::unbound_model_config`. | | [`ValidateError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/validate/enum.ValidateError.html) | Re-export of `validate::ValidateError`. | @@ -130,6 +141,7 @@ The intermediate representation (RFC 0002). |---|---| | [`bridge`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/bridge/index.html) | The fx/ModuleState ↔ `Overlay` bridge (RFC 0002 §2.4 migration contract). | | [`builder`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/index.html) | The Rust builder frontend (RFC 0002 §4.3–4.4): constructs the same runtime `Program` value the text parser will. | +| [`edit`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/index.html) | The graph-edit calculus: the *structural* mutation half of the IR. | | [`graph`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/index.html) | The IR graph core (RFC 0002 §2): entity ids, the `Interner`, the closed `Node` enum, field-level `Binding`/`PortRef` dataflow, and `Program` — arenas over value-level signatures. | | [`interp`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/index.html) | The IR interpreter (RFC 0002 §3): async evaluation of a loaded `Program`. | | [`module_build`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/module_build/index.html) | RFC 0003 stage M-3 library support: the module "linker". | @@ -216,6 +228,26 @@ The Rust builder frontend (RFC 0002 §4.3–4.4): constructs the same runtime `P | [`route`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.route.html) | Enum-discriminated branching. | | [`seq`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/builder/fn.seq.html) | Sequential composition. | +## `ir::edit` + +The graph-edit calculus: the *structural* mutation half of the IR. + +### Enums + +| Item | Description | +|---|---| +| [`ApplyError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.ApplyError.html) | A locally-checkable application failure. | +| [`Edit`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.Edit.html) | One structural edit. Serde values: an optimizer's proposal is data, not code — it can be logged, replayed against the same parent, and diffed. | +| [`EditError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.EditError.html) | Why `Program::edited` refused. | +| [`EditKind`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.EditKind.html) | A lightweight, serializable descriptor of an edit kind admissible at a node — the menu `Program::legal_edits` returns, suitable for prompting an LLM proposer. | +| [`SwapTarget`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.SwapTarget.html) | Target kind of `Edit::SwapLeaf`. | + +### Functions + +| Item | Description | +|---|---| +| [`migrate_overlay`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/fn.migrate_overlay.html) | Carries tuned values across a structural edit: for every entry in `overlay` (minted against `parent`), re-mint it against `child` when the child has a slot at the same `ParamPath`... | + ## `ir::graph` The IR graph core (RFC 0002 §2): entity ids, the `Interner`, the closed `Node` enum, field-level `Binding`/`PortRef` dataflow, and `Program` — arenas over value-level signatures. @@ -272,12 +304,16 @@ The IR interpreter (RFC 0002 §3): async evaluation of a loaded `Program`. | [`BudgetMeter`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.BudgetMeter.html) | Check-before-call metering: calls and deadline are hard-gated pre-call; token budgets are soft (checked against accumulated usage, since usage is only known post-hoc). | | [`Exhausted`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.Exhausted.html) | Budget reservation failure. | | [`Interpreter`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.Interpreter.html) | A loaded, executable program: validated graph + bound models/tools + registered sandbox code. Cheap to share; run state never lives here. | +| [`LeafOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.LeafOutcome.html) | Parse/coercion metadata from one successful `Predict`-leaf evaluation, collected by `Interpreter::run_collecting` in execution order. | +| [`RunOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.RunOutput.html) | Program output plus per-leaf metadata, returned by `Interpreter::run_collecting`. | | [`RuntimeEnv`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.RuntimeEnv.html) | What the host supplies at load: live models, host tool bindings, the sandbox, and the capability grants. | +| [`ToolSuspension`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.ToolSuspension.html) | A caller-managed agent turn suspended on pending tool calls (RFC 0004 §2). | ### Enums | Item | Description | |---|---| +| [`ConversationTurn`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/enum.ConversationTurn.html) | One caller-driven conversation turn, returned by `Interpreter::run_conversation_caller_managed` and `Interpreter::resume_conversation`. | | [`LoadError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/enum.LoadError.html) | | | [`RunError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/enum.RunError.html) | | @@ -350,6 +386,7 @@ Parameters (RFC 0002 §2.4): every mutable thing is a named, addressable slot; a | [`ParamOwner`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ParamOwner.html) | | | [`ParamValue`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ParamValue.html) | | | [`ToolDesc`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ToolDesc.html) | | +| [`ToolSetK`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/enum.ToolSetK.html) | | ### Traits @@ -386,11 +423,11 @@ RFC 0003 stage M-2: step metadata. The `.dsrs` text format (RFC 0002 §4, stage IR-5). -### Structs +### Re-exports | Item | Description | |---|---| -| [`ParseError`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/text/struct.ParseError.html) | A parse failure with the source position and what was expected — designed to be actionable feedback for a model regenerating the program. | +| [`ParseError`](https://docs.rs/dspy-rs/latest/dsrs_syntax/struct.ParseError.html) | A parse failure with the source position and what was expected — designed to be actionable feedback for a model regenerating the program. | ### Enums diff --git a/docs/docs/api/modules.mdx b/docs/docs/api/modules.mdx index 762f44a8..2b0045d7 100644 --- a/docs/docs/api/modules.mdx +++ b/docs/docs/api/modules.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. ## Re-exports @@ -14,7 +14,6 @@ Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit |---|---| | [`ChainOfThought`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/type.ChainOfThought.html) | Re-export of `chain_of_thought::ChainOfThought`. | | [`ChainOfThoughtOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/type.ChainOfThoughtOutput.html) | Re-export of `chain_of_thought::ChainOfThoughtOutput`. | -| [`ReAct`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/react/struct.ReAct.html) | Re-export of `react::ReAct`. | | [`Reasoning`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/struct.Reasoning.html) | Re-export of `chain_of_thought::Reasoning`. | | [`WithReasoning`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/struct.WithReasoning.html) | Re-export of `chain_of_thought::WithReasoning`. | @@ -23,7 +22,6 @@ Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit | Item | Description | |---|---| | [`chain_of_thought`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/index.html) | | -| [`react`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/react/index.html) | | ## `modules::chain_of_thought` @@ -40,14 +38,3 @@ Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit |---|---| | [`ChainOfThought`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/type.ChainOfThought.html) | Asks the LM to reason step-by-step before producing the answer. | | [`ChainOfThoughtOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/type.ChainOfThoughtOutput.html) | Convenience alias for `ChainOfThought`'s output type. | - -## `modules::react` - -### Structs - -| Item | Description | -|---|---| -| [`ReAct`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/react/struct.ReAct.html) | | -| [`ReActActionStepOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/react/struct.ReActActionStepOutput.html) | | -| [`ReActBuilder`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/react/struct.ReActBuilder.html) | | -| [`ReActExtractStepOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/react/struct.ReActExtractStepOutput.html) | | diff --git a/docs/docs/api/optimizer.mdx b/docs/docs/api/optimizer.mdx index 94d6287f..ac63eb08 100644 --- a/docs/docs/api/optimizer.mdx +++ b/docs/docs/api/optimizer.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. Automatic prompt optimization. @@ -18,10 +18,12 @@ Automatic prompt optimization. | `copro::*` | Glob re-export. | | `engine::*` | Glob re-export. | | `gepa::*` | Glob re-export. | +| [`LeafInfo`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/struct.LeafInfo.html) | Re-export of `target::LeafInfo`. | | `mipro::*` | Glob re-export. | -| `pareto::*` | Glob re-export. | -| `program_engine::*` | Glob re-export. | +| [`OptimizeTarget`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/struct.OptimizeTarget.html) | Re-export of `target::OptimizeTarget`. | +| [`ProgramMetric`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/trait.ProgramMetric.html) | Re-export of `target::ProgramMetric`. | | `simba::*` | Glob re-export. | +| `structural::*` | Glob re-export. | ## Modules @@ -32,15 +34,27 @@ Automatic prompt optimization. | [`engine`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/index.html) | The shared evaluation engine (vision §5.4): every optimizer is a thin strategy over this core. | | [`gepa`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/gepa/index.html) | | | [`mipro`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/mipro/index.html) | | -| [`pareto`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/pareto/index.html) | | -| [`program_engine`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/program_engine/index.html) | The IR-native evaluation path (RFC 0002 IR-6): candidate `Overlay`s evaluated over one shared `Arc` through the `Interpreter` with **candidate-level parallelism**. | | [`simba`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/index.html) | SIMBA: Stochastic Introspective Mini-Batch Ascent (vision §4.3) — the cheap agentic default. | +| [`structural`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/structural/index.html) | Structural: LM-guided hill-climbing over the graph-edit calculus (RFC 0004 §6) — the sixth strategy over the shared `Engine`. | +| [`target`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/index.html) | What an optimizer optimizes: `OptimizeTarget`, the lane-erased pair of (thing under optimization, evaluation harness). | + +## Structs + +| Item | Description | +|---|---| +| [`OptimizerCommon`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/struct.OptimizerCommon.html) | The engine/RNG knobs shared by every optimizer builder: evaluation concurrency, budget caps, cache salt, and the sampling seed. | + +## Enums + +| Item | Description | +|---|---| +| [`Report`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/enum.Report.html) | What an optimization run produced. Strategy-specific payloads for the optimizers that report more than "done". | ## Traits | Item | Description | |---|---| -| [`Optimizer`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/trait.Optimizer.html) | Tunes a module's `Predict` leaves for better performance. | +| [`Optimizer`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/trait.Optimizer.html) | A tuning strategy over the shared `Engine`. | ## `optimizer::bootstrap` @@ -72,40 +86,28 @@ The shared evaluation engine (vision §5.4): every optimizer is a thin strategy | Item | Description | |---|---| | [`Budget`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Budget.html) | Hard caps on evaluation spend. `None` = unlimited. | -| [`Candidate`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Candidate.html) | A candidate is *data*: a named set of `Overlay`s (predictor name → instruction/demos) plus a stable hash (`Candidate::stable_hash`). | +| [`Candidate`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Candidate.html) | A candidate is *data*: name-keyed per-leaf overlays plus a stable content hash (`Candidate::stable_hash`). | | [`CandidateEval`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.CandidateEval.html) | A candidate's results over one evaluation batch, in request order. | -| [`CandidateUndo`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.CandidateUndo.html) | Saved pre-overlay state for the predictors a candidate touched. Produced by `apply_candidate`, consumed by `restore_candidate`. | +| [`CandidateSlot`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.CandidateSlot.html) | One leaf's slice of a `Candidate`: which optimizable values to inject. | +| [`Engine`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Engine.html) | The shared evaluation core (vision §5.4). | | [`EngineConfig`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.EngineConfig.html) | Engine tuning knobs. | -| [`EvalEngine`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.EvalEngine.html) | The shared evaluation core (vision §5.4). | -| [`Overlay`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Overlay.html) | A per-predictor parameter overlay: which optimizable values to install. | +| [`ParetoStatistics`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.ParetoStatistics.html) | Snapshot of the Pareto frontier at a point in the search. | | [`ParetoView`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.ParetoView.html) | Dominance snapshot computed from a `ScoreMatrix`: which candidates win (or tie, within tolerance) on at least one example. | | [`RolloutCache`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.RolloutCache.html) | In-memory rollout cache: `(baseline, candidate, example, salt)` → `Eval`. | | [`RolloutOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.RolloutOutcome.html) | One evaluated (or cache-served) rollout. | | [`ScoreMatrix`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.ScoreMatrix.html) | Per-instance score matrix: candidates × examples. | -| [`Spend`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Spend.html) | What the engine has consumed so far. Serialized into checkpoints and reported to strategies. | +| [`Spend`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Spend.html) | What the engine has consumed so far. Reported to strategies. | ### Enums | Item | Description | |---|---| -| [`EvalOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/enum.EvalOutcome.html) | Result of `EvalEngine::evaluate`. | -| [`GateOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/enum.GateOutcome.html) | Result of `EvalEngine::evaluate_gated`. | - -### Functions - -| Item | Description | -|---|---| -| [`apply_candidate`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/fn.apply_candidate.html) | Applies a candidate's overlays to a module through the single mutation seam (`DynPredictor::apply_update`), returning the undo snapshot. | -| [`restore_candidate`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/fn.restore_candidate.html) | Restores the pre-candidate state captured by `apply_candidate`. | +| [`BatchEvalOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/enum.BatchEvalOutcome.html) | Result of `Engine::evaluate_many`. | +| [`EvalOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/enum.EvalOutcome.html) | Result of `Engine::evaluate`. | +| [`GateOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/enum.GateOutcome.html) | Result of `Engine::evaluate_gated`. | ## `optimizer::gepa` -### Re-exports - -| Item | Description | -|---|---| -| [`ParetoStatistics`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/pareto/struct.ParetoStatistics.html) | Re-export of `super::pareto::ParetoStatistics`. | - ### Structs | Item | Description | @@ -124,56 +126,55 @@ The shared evaluation engine (vision §5.4): every optimizer is a thin strategy |---|---| | [`MIPROv2`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/mipro/struct.MIPROv2.html) | Trace-guided instruction and demo optimizer. | | [`MIPROv2Builder`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/mipro/struct.MIPROv2Builder.html) | Use builder syntax to set the inputs and finish with `build()`). | -| [`PromptCandidate`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/mipro/struct.PromptCandidate.html) | An instruction candidate with its evaluated score. | | [`PromptingTips`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/mipro/struct.PromptingTips.html) | Library of general prompting best practices used to seed candidate generation. | -## `optimizer::pareto` - -### Structs - -| Item | Description | -|---|---| -| [`ParetoFrontier`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/pareto/struct.ParetoFrontier.html) | Per-example dominance frontier for candidate selection. | -| [`ParetoStatistics`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/pareto/struct.ParetoStatistics.html) | Snapshot of the Pareto frontier at a point in the search. | - -## `optimizer::program_engine` +## `optimizer::simba` -The IR-native evaluation path (RFC 0002 IR-6): candidate `Overlay`s evaluated over one shared `Arc` through the `Interpreter` with **candidate-level parallelism**. +SIMBA: Stochastic Introspective Mini-Batch Ascent (vision §4.3) — the cheap agentic default. ### Structs | Item | Description | |---|---| -| [`ProgramEvalEngine`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/program_engine/struct.ProgramEvalEngine.html) | The shared evaluation core for the dynamic lane: candidate overlays over one interpreter-loaded program. | +| [`IntrospectRolloutsOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.IntrospectRolloutsOutput.html) | | +| [`SIMBA`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SIMBA.html) | Minibatch introspective ascent — the cheap agentic default (vision §4.3). | +| [`SIMBABuilder`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SIMBABuilder.html) | Use builder syntax to set the inputs and finish with `build()`). | +| [`SimbaReport`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SimbaReport.html) | What a `SIMBA` run did. | +| [`SimbaStep`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SimbaStep.html) | What one SIMBA step did. | ### Enums | Item | Description | |---|---| -| [`ProgramEvalOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/program_engine/enum.ProgramEvalOutcome.html) | Result of `ProgramEvalEngine::evaluate_program_candidates`. | +| [`SimbaMove`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/enum.SimbaMove.html) | Which of SIMBA's two moves a step proposed. | -### Traits +## `optimizer::structural` + +Structural: LM-guided hill-climbing over the graph-edit calculus (RFC 0004 §6) — the sixth strategy over the shared `Engine`. + +### Structs | Item | Description | |---|---| -| [`ProgramMetric`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/program_engine/trait.ProgramMetric.html) | How a program-lane strategy tells the engine what "good" means: score one interpreter output (`JsonMap` of the program's output signature fields) against a labeled example. | +| [`ChooseEditOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/structural/struct.ChooseEditOutput.html) | | +| [`Structural`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/structural/struct.Structural.html) | Structural optimizer over the graph-edit calculus (RFC 0004 §6). | +| [`StructuralBuilder`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/structural/struct.StructuralBuilder.html) | Use builder syntax to set the inputs and finish with `build()`). | +| [`StructuralReport`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/structural/struct.StructuralReport.html) | What a `Structural` run did. The winner is returned, not installed: bake it (`report.program.bake(&report. | +| [`StructuralStep`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/structural/struct.StructuralStep.html) | What one Structural generation did. | -## `optimizer::simba` +## `optimizer::target` -SIMBA: Stochastic Introspective Mini-Batch Ascent (vision §4.3) — the cheap agentic default. +What an optimizer optimizes: `OptimizeTarget`, the lane-erased pair of (thing under optimization, evaluation harness). ### Structs | Item | Description | |---|---| -| [`IntrospectRolloutsOutput`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.IntrospectRolloutsOutput.html) | | -| [`SIMBA`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SIMBA.html) | Minibatch introspective ascent — the cheap agentic default (vision §4.3). | -| [`SIMBABuilder`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SIMBABuilder.html) | Use builder syntax to set the inputs and finish with `build()`). | -| [`SimbaReport`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SimbaReport.html) | What a `SIMBA` run did. | -| [`SimbaStep`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SimbaStep.html) | What one SIMBA step did. | +| [`LeafInfo`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/struct.LeafInfo.html) | The read surface strategies build candidates from: one optimizable leaf's name, current values, and field contract. | +| [`OptimizeTarget`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/struct.OptimizeTarget.html) | The thing an `Optimizer` optimizes: a module or a program, packaged with its example set and metric. See the module docs for the two lanes. | -### Enums +### Traits | Item | Description | |---|---| -| [`SimbaMove`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/enum.SimbaMove.html) | Which of SIMBA's two moves a step proposed. | +| [`ProgramMetric`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/trait.ProgramMetric.html) | How a program-lane strategy tells the engine what "good" means: score one interpreter output (`JsonMap` of the program's output signature fields) against a labeled example. | diff --git a/docs/docs/api/predictors.mdx b/docs/docs/api/predictors.mdx index db9bb3e7..63cb9d72 100644 --- a/docs/docs/api/predictors.mdx +++ b/docs/docs/api/predictors.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. ## Re-exports @@ -26,6 +26,7 @@ Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit | Item | Description | |---|---| +| [`AgentLoopSpec`](https://docs.rs/dspy-rs/latest/dspy_rs/predictors/predict/struct.AgentLoopSpec.html) | Loop options for a tooled predictor's 1-node `agent` program. | | [`Demo`](https://docs.rs/dspy-rs/latest/dspy_rs/predictors/predict/struct.Demo.html) | A typed input/output pair for few-shot prompting. | | [`Predict`](https://docs.rs/dspy-rs/latest/dspy_rs/predictors/predict/struct.Predict.html) | The leaf module. The only thing in the system that actually calls the LM. | | [`PredictBuilder`](https://docs.rs/dspy-rs/latest/dspy_rs/predictors/predict/struct.PredictBuilder.html) | Builder for `Predict` with demos, tools, and instruction override. | diff --git a/docs/docs/api/prelude.mdx b/docs/docs/api/prelude.mdx new file mode 100644 index 00000000..d128f5f5 --- /dev/null +++ b/docs/docs/api/prelude.mdx @@ -0,0 +1,68 @@ +--- +title: "dspy_rs::prelude" +description: "The curated core surface — the recommended import for DSRs programs." +icon: "cube" +--- + + +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. + + +The curated core surface — the recommended import for DSRs programs. + +## Re-exports + +| Item | Description | +|---|---| +| [`average_score`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/evaluator/fn.average_score.html) | Re-export of `crate::evaluate::average_score`. | +| [`BootstrapFewShot`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/bootstrap/struct.BootstrapFewShot.html) | Re-export of `crate::optimizer::BootstrapFewShot`. | +| [`CallMetadata`](https://docs.rs/dspy-rs/latest/dspy_rs/core/predicted/struct.CallMetadata.html) | Re-export of `crate::core::CallMetadata`. | +| [`Candidate`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Candidate.html) | Re-export of `crate::optimizer::Candidate`. | +| [`capture_with_meta`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/capture/fn.capture_with_meta.html) | Re-export of `crate::trace::capture_with_meta`. | +| [`capture`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/capture/index.html) | Re-export of `crate::trace::capture`. | +| [`capture`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/capture/fn.capture.html) | Re-export of `crate::trace::capture`. | +| [`ChainOfThought`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/type.ChainOfThought.html) | Re-export of `crate::modules::ChainOfThought`. | +| [`configure`](https://docs.rs/dspy-rs/latest/dspy_rs/core/settings/fn.configure.html) | Re-export of `crate::core::settings::configure`. | +| [`COPRO`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/copro/struct.COPRO.html) | Re-export of `crate::optimizer::COPRO`. | +| [`DataLoader`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/struct.DataLoader.html) | Re-export of `crate::data::dataloader::DataLoader`. | +| [`Demo`](https://docs.rs/dspy-rs/latest/dspy_rs/predictors/predict/struct.Demo.html) | Re-export of `crate::predictors::Demo`. | +| [`Edit`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/edit/enum.Edit.html) | Re-export of `crate::ir::Edit`. | +| [`Engine`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/engine/struct.Engine.html) | Re-export of `crate::optimizer::Engine`. | +| [`Eval`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/struct.Eval.html) | Re-export of `crate::trace::Eval`. | +| [`evaluate_trainset_with_concurrency`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/evaluator/fn.evaluate_trainset_with_concurrency.html) | Re-export of `crate::evaluate::evaluate_trainset_with_concurrency`. | +| [`evaluate_trainset`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/evaluator/fn.evaluate_trainset.html) | Re-export of `crate::evaluate::evaluate_trainset`. | +| [`Example`](https://docs.rs/dspy-rs/latest/dsrs_macros/derive.Example.html) | Re-export of `dsrs_macros::Example`. | +| [`GEPA`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/gepa/struct.GEPA.html) | Re-export of `crate::optimizer::GEPA`. | +| [`init_tracing`](https://docs.rs/dspy-rs/latest/dspy_rs/utils/telemetry/fn.init_tracing.html) | Re-export of `crate::utils::init_tracing`. | +| [`Interpreter`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/interp/struct.Interpreter.html) | Re-export of `crate::ir::Interpreter`. | +| [`is_capturing`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/capture/fn.is_capturing.html) | Re-export of `crate::trace::is_capturing`. | +| [`is_replaying`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/fn.is_replaying.html) | Re-export of `crate::trace::is_replaying`. | +| [`LM`](https://docs.rs/dspy-rs/latest/dspy_rs/core/lm/struct.LM.html) | Re-export of `crate::core::lm::LM`. | +| [`LMConfig`](https://docs.rs/dspy-rs/latest/dspy_rs/core/lm/struct.LMConfig.html) | Re-export of `crate::core::lm::LMConfig`. | +| [`MIPROv2`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/mipro/struct.MIPROv2.html) | Re-export of `crate::optimizer::MIPROv2`. | +| [`Module`](https://docs.rs/dspy-rs/latest/dspy_rs/core/module/trait.Module.html) | Re-export of `crate::core::Module`. | +| [`ModuleState`](https://docs.rs/dspy-rs/latest/dspy_rs/core/state/struct.ModuleState.html) | Re-export of `crate::core::ModuleState`. | +| [`Optimizer`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/trait.Optimizer.html) | Re-export of `crate::optimizer::Optimizer`. | +| [`OptimizeTarget`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/target/struct.OptimizeTarget.html) | Re-export of `crate::optimizer::OptimizeTarget`. | +| [`Overlay`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/params/struct.Overlay.html) | Re-export of `crate::ir::Overlay`. | +| [`Predict`](https://docs.rs/dspy-rs/latest/dspy_rs/predictors/predict/struct.Predict.html) | Re-export of `crate::predictors::Predict`. | +| [`Predicted`](https://docs.rs/dspy-rs/latest/dspy_rs/core/predicted/struct.Predicted.html) | Re-export of `crate::core::Predicted`. | +| [`PredictError`](https://docs.rs/dspy-rs/latest/dspy_rs/core/errors/enum.PredictError.html) | Re-export of `crate::core::PredictError`. | +| [`Predictors`](https://docs.rs/dspy-rs/latest/dspy_rs/core/module/trait.Predictors.html) | Re-export of `crate::core::Predictors`. | +| [`predictors`](https://docs.rs/dspy-rs/latest/dspy_rs/predictors/index.html) | Re-export of `crate::predictors`. | +| [`predictors`](https://docs.rs/dspy-rs/latest/dspy_rs/macro.predictors.html) | Re-export of `crate::predictors`. | +| [`PredictState`](https://docs.rs/dspy-rs/latest/dspy_rs/core/state/struct.PredictState.html) | Re-export of `crate::core::PredictState`. | +| [`Program`](https://docs.rs/dspy-rs/latest/dspy_rs/ir/graph/struct.Program.html) | Re-export of `crate::ir::Program`. | +| [`replay`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/index.html) | Re-export of `crate::trace::replay`. | +| [`replay`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/fn.replay.html) | Re-export of `crate::trace::replay`. | +| [`ReplayMode`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/enum.ReplayMode.html) | Re-export of `crate::trace::ReplayMode`. | +| [`ReplayReport`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/struct.ReplayReport.html) | Re-export of `crate::trace::ReplayReport`. | +| [`Signature`](https://docs.rs/dspy-rs/latest/dspy_rs/core/signature/trait.Signature.html) | Re-export of `crate::core::signature::Signature`. | +| [`Signature`](https://docs.rs/dspy-rs/latest/dsrs_macros/derive.Signature.html) | Re-export of `dsrs_macros::Signature`. | +| [`SIMBA`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/simba/struct.SIMBA.html) | Re-export of `crate::optimizer::SIMBA`. | +| [`SpanId`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/struct.SpanId.html) | Re-export of `crate::trace::SpanId`. | +| [`Structural`](https://docs.rs/dspy-rs/latest/dspy_rs/optimizer/structural/struct.Structural.html) | Re-export of `crate::optimizer::Structural`. | +| [`Trace`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/struct.Trace.html) | Re-export of `crate::trace::Trace`. | +| [`TypedLoadOptions`](https://docs.rs/dspy-rs/latest/dspy_rs/data/dataloader/struct.TypedLoadOptions.html) | Re-export of `crate::data::dataloader::TypedLoadOptions`. | +| [`TypedMetric`](https://docs.rs/dspy-rs/latest/dspy_rs/evaluate/evaluator/trait.TypedMetric.html) | Re-export of `crate::evaluate::TypedMetric`. | +| [`WithReasoning`](https://docs.rs/dspy-rs/latest/dspy_rs/modules/chain_of_thought/struct.WithReasoning.html) | Re-export of `crate::modules::WithReasoning`. | diff --git a/docs/docs/api/trace.mdx b/docs/docs/api/trace.mdx index eb21ed0f..f5258026 100644 --- a/docs/docs/api/trace.mdx +++ b/docs/docs/api/trace.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. Execution trace capture (RFC 0001). @@ -15,7 +15,6 @@ Execution trace capture (RFC 0001). | Item | Description | |---|---| | `capture::*` | Glob re-export. | -| `export::*` | Glob re-export. | | [`is_replaying`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/fn.is_replaying.html) | Re-export of `replay::is_replaying`. | | [`replay`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/fn.replay.html) | Re-export of `replay::replay`. | | [`ReplayError`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/enum.ReplayError.html) | Re-export of `replay::ReplayError`. | @@ -29,7 +28,6 @@ Execution trace capture (RFC 0001). | Item | Description | |---|---| | [`capture`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/capture/index.html) | Task-local capture scope for the unified trace format. | -| [`export`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/index.html) | Exports: projections of the trace format onto external training and observability conventions (RFC 0001 §4f/§4g). | | [`replay`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/replay/index.html) | Replay scope: serve `Predict` calls from a recorded `Trace` (RFC 0001 §4d/§4e). | | [`serialize`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/serialize/index.html) | JSONL wire format for `Trace`: header line, span lines, optional footer. | | [`span`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/span/index.html) | The unified trace format (RFC 0001): one `Span` per `Predict` invocation, one `Trace` per rollout. | @@ -56,60 +54,6 @@ Task-local capture scope for the unified trace format. | [`capture`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/capture/fn.capture.html) | Runs `f` while recording every `Predict` call on this task into a `Trace`. | | [`is_capturing`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/capture/fn.is_capturing.html) | Returns `true` if the current task is inside a `capture` scope. | -## `trace::export` - -Exports: projections of the trace format onto external training and observability conventions (RFC 0001 §4f/§4g). - -### Re-exports - -| Item | Description | -|---|---| -| [`OtelEvent`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelEvent.html) | Re-export of `otel::OtelEvent`. | -| [`OtelKeyValue`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelKeyValue.html) | Re-export of `otel::OtelKeyValue`. | -| [`OtelSpan`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelSpan.html) | Re-export of `otel::OtelSpan`. | -| [`OtelStatus`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelStatus.html) | Re-export of `otel::OtelStatus`. | -| [`OtelValue`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/enum.OtelValue.html) | Re-export of `otel::OtelValue`. | -| [`RlRollout`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/rl/struct.RlRollout.html) | Re-export of `rl::RlRollout`. | -| [`RlTransition`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/rl/struct.RlTransition.html) | Re-export of `rl::RlTransition`. | - -## `trace::export::otel` - -OTel export (RFC 0001 §4g): one-way batch mapping of a finished `Trace` onto OpenTelemetry GenAI semantic conventions — as plain serializable structs in the OTLP/JSON wire shape, w... - -### Structs - -| Item | Description | -|---|---| -| [`OtelEvent`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelEvent.html) | OTLP span `Event`. | -| [`OtelKeyValue`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelKeyValue.html) | OTLP `KeyValue`. | -| [`OtelSpan`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelSpan.html) | One span in the OTLP/JSON wire shape (proto3 JSON mapping: camelCase keys, 64-bit integers as decimal strings, ids as lowercase hex). | -| [`OtelStatus`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/struct.OtelStatus.html) | OTLP span `Status`. | - -### Enums - -| Item | Description | -|---|---| -| [`OtelValue`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/enum.OtelValue.html) | OTLP `AnyValue` (the oneof arms this export emits). | - -### Constants - -| Item | Description | -|---|---| -| [`SPAN_KIND_CLIENT`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/constant.SPAN_KIND_CLIENT.html) | OTLP `SPAN_KIND_CLIENT`. | -| [`SPAN_KIND_INTERNAL`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/constant.SPAN_KIND_INTERNAL.html) | OTLP `SPAN_KIND_INTERNAL`. | -| [`STATUS_CODE_ERROR`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/otel/constant.STATUS_CODE_ERROR.html) | OTLP `STATUS_CODE_ERROR`. | - -## `trace::export::rl` - -RL rollout export (RFC 0001 §4f): the Agent Lightning / verifiers span convention — one rollout as message lists + reward + per-subcall transitions. - -### Structs - -| Item | Description | -|---|---| -| [`RlRollout`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/rl/struct.RlRollout.html) | One rollout: reward plus per-subcall transitions. Serializes to a single JSON object — the JSONL record RL trainers consume. | -| [`RlTransition`](https://docs.rs/dspy-rs/latest/dspy_rs/trace/export/rl/struct.RlTransition.html) | One policy subcall — a `Predict` invocation as (prompt messages, emitted completion) with its span metadata. | - ## `trace::replay` Replay scope: serve `Predict` calls from a recorded `Trace` (RFC 0001 §4d/§4e). diff --git a/docs/docs/api/typesys.mdx b/docs/docs/api/typesys.mdx index dfed4661..7f588c43 100644 --- a/docs/docs/api/typesys.mdx +++ b/docs/docs/api/typesys.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. In-house type system that replaces the vendored BAML stack (`bamltype`, `baml_types`, @@ -21,10 +21,8 @@ In-house type system that replaces the vendored BAML stack (`bamltype`, `baml_ty | [`Constraint`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.Constraint.html) | Re-export of `constraint::Constraint`. | | [`ConstraintKind`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/enum.ConstraintKind.html) | Re-export of `constraint::ConstraintKind`. | | [`ConstraintLevel`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/type.ConstraintLevel.html) | Re-export of `constraint::ConstraintLevel`. | -| [`ConstraintOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.ConstraintOutcome.html) | Re-export of `constraint::ConstraintOutcome`. | | [`EnumDef`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.EnumDef.html) | Re-export of `schema::EnumDef`. | | [`EnumValueDef`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.EnumValueDef.html) | Re-export of `schema::EnumValueDef`. | -| [`evaluate_constraints`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/fn.evaluate_constraints.html) | Re-export of `constraint::evaluate_constraints`. | | [`evaluate_expression`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/fn.evaluate_expression.html) | Re-export of `constraint::evaluate_expression`. | | [`field_type_from_shape`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/fn.field_type_from_shape.html) | Re-export of `schema::field_type_from_shape`. | | [`FieldDef`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.FieldDef.html) | Re-export of `schema::FieldDef`. | @@ -32,7 +30,6 @@ In-house type system that replaces the vendored BAML stack (`bamltype`, `baml_ty | [`Flag`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/coerce/enum.Flag.html) | Re-export of `coerce::Flag`. | | [`internal_name_for_shape`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/fn.internal_name_for_shape.html) | Re-export of `schema::internal_name_for_shape`. | | [`OutputSchema`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/struct.OutputSchema.html) | Re-export of `schema::OutputSchema`. | -| [`ResponseCheck`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.ResponseCheck.html) | Re-export of `constraint::ResponseCheck`. | | [`schema_block`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/render/fn.schema_block.html) | Re-export of `render::schema_block`. | | [`Schema`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/schema/trait.Schema.html) | Re-export of `schema::Schema`. | | [`type_name`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/render/fn.type_name.html) | Re-export of `render::type_name`. | @@ -84,8 +81,6 @@ In-house constraint model + evaluation, replacing BAML's `Constraint` / `run_use | Item | Description | |---|---| | [`Constraint`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.Constraint.html) | A single `#check`/`#assert` constraint attached to a field. | -| [`ConstraintOutcome`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.ConstraintOutcome.html) | The outcome of evaluating a constraint against a value. | -| [`ResponseCheck`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/struct.ResponseCheck.html) | A reported check result, mirroring the old `ResponseCheck` shape used by GEPA/optimizers. | ### Enums @@ -98,7 +93,6 @@ In-house constraint model + evaluation, replacing BAML's `Constraint` / `run_use | Item | Description | |---|---| | [`evaluate_constraint_expression`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/fn.evaluate_constraint_expression.html) | Evaluates a `'static` constraint expression against `value`, compiling it at most once per process. | -| [`evaluate_constraints`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/fn.evaluate_constraints.html) | Evaluates every constraint against `value`, binding it as `this` in a jinja expression. | | [`evaluate_expression`](https://docs.rs/dspy-rs/latest/dspy_rs/typesys/constraint/fn.evaluate_expression.html) | Evaluates a runtime (non-`'static`) constraint expression against `value`, binding it as `this`. | ### Type aliases diff --git a/docs/docs/api/utils.mdx b/docs/docs/api/utils.mdx index 01e906c5..e458215d 100644 --- a/docs/docs/api/utils.mdx +++ b/docs/docs/api/utils.mdx @@ -5,7 +5,7 @@ icon: "cube" --- -Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `b5da857`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. +Generated from rustdoc JSON at `dspy-rs v0.7.3` (commit `f7d67a6`). Do not edit by hand; regenerate with `python3 docs/scripts/gen_api.py` (see the script header for the rustdoc commands). Item links lead to full signatures and method docs on docs.rs. LM response caching. diff --git a/docs/docs/components/adapters.mdx b/docs/docs/components/adapters.mdx index 0cc8a50f..40c0aa22 100644 --- a/docs/docs/components/adapters.mdx +++ b/docs/docs/components/adapters.mdx @@ -4,16 +4,16 @@ description: 'Turn signatures into prompts and parse LM responses' icon: 'arrows-turn-to-dots' --- -An adapter sits between a [signature](/docs/components/signatures) and the LM: it formats typed inputs into the prompt the model sees and parses the LM's response back into typed outputs. The default adapter is `ChatAdapter`, and [Predict](/docs/components/predict) drives it for you, so you touch it directly only to inspect what the model sees. +An adapter sits between a signature and the LM: it formats inputs into the prompt the model sees and parses the LM's response back into typed outputs. The default adapter is `ChatAdapter`, and [Predict](/docs/components/predict) drives it for you (via the IR interpreter), so you touch it directly only to inspect what the model sees. ## What adapters do ``` -Signature + Input → Adapter.format() → Prompt for LM +SignatureDef + Input → Adapter.format() → Prompt for LM LM Response → Adapter.parse() → Typed Output ``` -`ChatAdapter` handles both formatting and parsing. +`ChatAdapter` handles both formatting and parsing. It has one lane: every method is parameterized by an owned [`SignatureDef`](/docs/components/program-and-nodes) — the IR's value-level signature — whether the def came from a derived `Signature` type (`SignatureDef::of::()`) or from a loaded `.dsrs` program. There is no separate "static" formatting path. ## ChatAdapter @@ -132,23 +132,29 @@ ParseError::Multiple { ## Using the adapter directly -Usually you do not touch the adapter - `Predict` handles it. But for debugging: +Usually you do not touch the adapter - `Predict` handles it. The public surface is five building blocks, all parameterized by `&SignatureDef`: + +| Method | Produces | +|--------|----------| +| `build_system_def(def, types, instruction_override)` | The system message (field descriptions, structure template, response instructions, task description) | +| `format_input_def(def, input)` | The user message: input fields with `[[ ## field ## ]]` markers plus response instructions. `input` is a `JsonMap`; absent fields are skipped | +| `format_output_def(def, output)` | An assistant message for few-shot demos, ending with `[[ ## completed ## ]]` | +| `parse_output_def(def, types, response)` | `(JsonMap, IndexMap)` — the parsed output map plus per-field parse metadata | +| `ChatAdapter::parse_sections(content)` | Ordered map of `field_name → section_content` split by the delimiters | + +For debugging what the LM sees from a derived signature: ```rust -use dspy_rs::ChatAdapter; +use dspy_rs::{ChatAdapter, ir::SignatureDef}; let adapter = ChatAdapter; +let def = SignatureDef::of::(); +let types = SignatureDef::types_of::(); -// See what prompt would be generated -let system = adapter.format_system_message_typed::()?; -let user = adapter.format_user_message_typed::(&input); - +let system = adapter.build_system_def(def, types, None); println!("System:\n{system}"); -println!("User:\n{user}"); ``` -This is useful for understanding what the LM sees. - ## Input formatting options Input fields support two rendering paths: @@ -189,7 +195,7 @@ Available filters include: - `truncate` for length-limited string rendering If template rendering fails at runtime (for example, missing variables), `ChatAdapter` panics. -Compiled templates are cached process-wide by template string. +Templates are compiled per call — defs own their template strings, so there is no process-global template cache to leak per loaded program. ## Real example: Insurance claim extraction diff --git a/docs/docs/components/code-mode.mdx b/docs/docs/components/code-mode.mdx index 6ce9bf55..5388682d 100644 --- a/docs/docs/components/code-mode.mdx +++ b/docs/docs/components/code-mode.mdx @@ -108,12 +108,12 @@ Resource limits applied to every sandbox instance. `Copy`, so it is passed by va | Field | Type | Default | Meaning | |---|---|---|---| | `memory_limit` | `usize` | 32 MiB | Max heap for one call, in bytes. Exceeding it kills the call as `MemoryExceeded`. | -| `deadline` | `Duration` | 500 ms | Wall-clock budget for one call, enforced by the engine's interrupt handler. Time inside a capability counts against the budget but cannot be interrupted mid-call; the deadline re-arms when control returns to JS. | +| `deadline` | `Duration` | 500 ms | Wall-clock budget for one call. JS execution is interrupted by the engine's interrupt handler; a capability call is host code the interrupt handler cannot reach, so the executor bounds it with the *remaining* budget via `tokio::time::timeout` — a handler that runs past it is dropped and the call surfaces in JS as a deadline timeout. | | `max_stack` | `usize` | 512 KiB | Max JS stack, in bytes. | ### `CacheStats` -Sources compile once per unique content (BLAKE3-keyed) and the bytecode is shared across calls and tool names. `CacheStats` carries `entries: usize`, `hits: u64`, `misses: u64`. +Sources compile once per unique content (BLAKE3-keyed) and the bytecode is shared across calls and tool names. `CacheStats` carries `entries: usize`, `hits: u64`, `misses: u64`. The cache is bounded at 128 entries with a deliberately simple cap-and-clear eviction (an optimizer generating thousands of candidate tool bodies stays bounded instead of leaking them all); registered tools hold their own reference to their bytecode, so eviction only costs a recompile on the next miss. `deregister` evicts the tool's entry unless another registered tool shares the same source hash. ### `run_script` @@ -144,12 +144,13 @@ Accessors: `name()` and `description()`. `CapabilityHandler` is the public handl For wrapped tools, the args object is serialized to JSON, handed to `ToolDyn::call`, and the result string is parsed back to JSON (or returned as a plain string if it is not valid JSON). A tool error becomes a JS exception whose message names the original tool: `` tool `` failed: ``. -Capability names become JS globals, so they must be valid identifiers; the `__dsrs` prefix is reserved by the runtime. `js_identifier(name)` mangles an arbitrary tool name into a valid identifier, in order: +Capability names become JS globals, so they must be valid identifiers; the `__dsrs` prefix is reserved by the runtime, and a name that collides with a JavaScript reserved word or ambient global (`class`, `JSON`, `Object`, ...) is refused. `js_identifier(name)` mangles an arbitrary tool name into a valid identifier, in order: 1. Every character outside `[A-Za-z0-9_$]` becomes `_` (`my-tool.v2` becomes `my_tool_v2`). 2. A leading digit gets a `_` prepended (`2fast` becomes `_2fast`). 3. An empty name becomes `_tool`. 4. A result starting with `__dsrs` gets one more leading `_`. +5. A result colliding with a JavaScript reserved word or ambient global gets a `_tool` suffix (`JSON` becomes `JSON_tool`, `class` becomes `class_tool`) — declaring a `const` named `class` would break the sandbox contract outright, and shadowing `JSON` would sabotage every script that touches it. The mapping is not injective: distinct names can mangle to the same identifier, so every batch wrapper refuses collisions at registration or load time instead of silently shadowing a tool. diff --git a/docs/docs/components/data.mdx b/docs/docs/components/data.mdx index 247500f7..976815b2 100644 --- a/docs/docs/components/data.mdx +++ b/docs/docs/components/data.mdx @@ -32,9 +32,9 @@ The same row type flows through the whole loop: | Consumer | Signature | Row bound | |---|---|---| | [`evaluate_trainset`](/docs/components/evaluation) | `evaluate_trainset(&module, &[E], &metric)` | `E: ToInput + Sync` | -| [`Optimizer::compile`](/docs/components/optimizers) | `optimizer.compile(&mut module, Vec, &metric)` | `E: ToInput + serde::Serialize + Send + Sync` | +| [`compile_module`](/docs/components/optimizers) | `optimizer.compile_module(&mut module, &[E], &metric)` | `E: ToInput + serde::Serialize + Send + Sync` | | [`TypedMetric`](/docs/components/evaluation) | `evaluate(&self, example: &E, prediction, trace)` | none — the metric receives the full row | -| [`EvalEngine::new`](/docs/components/optimizer-engine) | `EvalEngine::new(Vec, &metric, config)` | `E: Serialize` (rollout-cache uids are content hashes of the whole row) | +| [`OptimizeTarget::module`](/docs/components/optimizer-engine) | `OptimizeTarget::module(&mut module, &[E], &metric)` | `E: ToInput + Serialize + Sync` (rollout-cache uids are content hashes of the whole row) | | Demo seeding | `Demo::new(row.to_input()?, row.to_output()?)` | `E: ToInput + ToOutput` | Because the metric sees the row rather than a signature-shaped pair, ground truth does not have to fit the module's output type: a metric can score a `QAOutput` prediction against supporting facts the module never produced. @@ -197,9 +197,9 @@ Mapper closure errors are wrapped as `DataLoadError::Mapper` with the failing ro The module also exposes `is_url(path: &str) -> bool`, the helper the loaders use to decide between filesystem and HTTP(S) fetching. -## Versioning +## Module layout -Data loading is versioned under `data::v1`. The module tree is `data::v1::dataloader` and `data::v1::utils`; `data/mod.rs` re-exports `v1::*`, so unversioned paths (`dspy_rs::data::DataLoader`) and versioned paths (`dspy_rs::data::v1::dataloader::DataLoader`) both resolve to the same items. The crate root flattens further: `dspy_rs::DataLoader` is the conventional import. Unversioned paths always track the current version. +Data loading lives at `data::dataloader` and `data::utils`, re-exported from `data/mod.rs`. The crate root flattens further: `dspy_rs::DataLoader` is the conventional import. (The transitional `data::v1` path alias is gone.) ### Migration note diff --git a/docs/docs/components/dsrs-file.mdx b/docs/docs/components/dsrs-file.mdx index 95d39340..e8585d11 100644 --- a/docs/docs/components/dsrs-file.mdx +++ b/docs/docs/components/dsrs-file.mdx @@ -167,6 +167,7 @@ An LM plus tool loop. The block is required. ``` researcher = agent Research @fast (question = $.question) { tools [fetch shout] + tool_set [fetch] stop_tools [shout] max_turns 6 until_parse false @@ -177,6 +178,8 @@ researcher = agent Research @fast (question = $.question) { } ``` +`tools` declares which tools the loop *may* carry — it is the loop's capability footprint. `tool_set` is the tuned selection: the subset the loop actually presents to the model, an optimizable parameter like `instruction` or `demos`. It only prints when an optimizer has restricted it; absent means the full `tools` list. + ### `hole` Typed opaque code: the type system sees a normal node, the implementation is either sandboxed JavaScript carried in the artifact or a native function the host binds by name. Every hole declares `caps [...]` (empty when it needs none), then either a `` js``` ``` `` fence or `extern ""`. @@ -276,7 +279,7 @@ Violations of any of these are compile errors: 1. `dsrs 1` first; `main: = seq { ... }` last. 2. Node names are program-unique; only earlier nodes are referenceable. -3. Every hole and tool `caps [...]` must be a subset of the program `caps { ... }`. +3. Every hole and tool `caps [...]` must be a subset of the program `caps { ... }`; an agent's `stop_tools` must come from its `tools`, and its `tool_set` must be a duplicate-free subset of them. 4. `route` needs `else` unless its arms cover every enum variant; arms export identical fields. 5. All loops carry explicit bounds (`max_iters`, `max_turns`, `attempts`, `max_rounds`). 6. Signatures need at least one `in` and one `out` field; `check` needs a label. diff --git a/docs/docs/components/edit-calculus.mdx b/docs/docs/components/edit-calculus.mdx new file mode 100644 index 00000000..9fb6d89d --- /dev/null +++ b/docs/docs/components/edit-calculus.mdx @@ -0,0 +1,95 @@ +--- +title: "The edit calculus" +description: "Structural program mutation: the Edit enum, Program::edited, legal_edits, and carrying overlays across an edit with migrate_overlay" +icon: "scissors" +--- + +The edit calculus is the *structural* mutation half of the IR. An [`Overlay`](/docs/components/program-and-nodes) mutates parameter **values** over a fixed skeleton; an `Edit` mutates the skeleton itself — add a reasoning field, swap a `Predict` for an `AgentLoop`, wrap a flaky step in a `Retry`, remove a step. Edits are plain serde values — inspectable, diffable, replayable — and are only ever applied through `Program::edited`, which is pure: it clones the arenas, applies the edits in order, re-runs the same load-time validation the builder and loader use, and seals a **new** content hash. A program value is never mutated in place, so every hash-bound artifact (overlays, traces, caches) minted against the parent stays coherent. + +All items are exported from `dspy_rs::ir`: `Edit`, `EditKind`, `SwapTarget`, `EditError`, `ApplyError`, `migrate_overlay`. + +## The edits + +An optimizer's structural proposal is data, not code — it can be logged, replayed against the same parent, and diffed. + +| `Edit` variant | Plain words | +|---|---| +| `AugmentSig { leaf, prepend }` | Prepend an output field to a `Predict`/`AgentLoop` leaf's signature — the CoT move (mirrors `SignatureDef::augmented_with`). Copy-on-write: a new `SigId` is created; nodes sharing the old signature keep it. | +| `SwapLeaf { leaf, to }` | Swap a leaf's kind: `Predict` → `AgentLoop` (with tools ⊆ `program.tools`, a stop spec, and a budget) or `AgentLoop` → `Predict`. Name, signature, bindings, and the instruction/demos/model param slots are preserved; the agent direction mints `.context` and `.tool_set` slots, the predict direction drops them. | +| `WrapRetry { node, max_attempts, backoff_ms, feedback }` | Wrap an existing node in a `Retry`, rewiring the parent reference and redirecting downstream `Out` ports to the wrapper. | +| `Remove { node }` | Remove a node from its parent `Seq` body (subtree and its params are garbage-collected). If a later binding still references its outputs, `validate()` rejects the batch. | +| `AddTool { agent, tool }` / `RemoveTool { agent, tool }` | Declare or undeclare an existing program tool on an agent leaf. The `.tool_set` default tracks the declaration: adding a tool makes it live, removing one also drops it from `stop_tools` and the tool-set default. | +| `SetStop { agent, stop }` | Replace an agent leaf's `StopSpec`. | +| `SetInstructionDefault { leaf, text }` | Set the leaf's instruction slot *default* — a bake-like change without an overlay, for structural optimizers that also seed text. | + +`SwapTarget` is the target kind of `SwapLeaf`: `Agent { tools, stop, budget }` or `Predict`. + +## `Program::edited` + +```rust +pub fn edited(&self, edits: &[Edit]) -> Result +``` + +Applies `edits` in order to a clone of `self` and returns the sealed, validated result. The child gets a **new** content hash and `lineage.parent` set to the parent's hash — exactly like `Program::bake`; the other provenance fields are left empty for the optimizer to fill (an edit is not an optimization run record). + +Behavior worth knowing: + +- **NodeIds are positional handles against the parent.** Within one `edited()` batch, ids stay stable (swaps happen in place, removals only detach); dead nodes, signatures, and params are garbage-collected once at the end. Ids in the child may therefore differ from the parent — re-locate leaves by name (`Program::leaf_id(name)`, leaf names are program-unique and survive edits) and params by path. +- **Batch validation.** Edits are validated as a *sequence*: intermediate states may be inconsistent (remove a producer, then its consumer); only the final program must pass validation. Apply-time errors cover what is checkable locally; everything data-flow shaped is deliberately left to the load-time validator, so the edit layer and the loader can never disagree. +- **Identity is preserved.** `edited(&[])` returns a program with the parent's hash — only lineage differs, and lineage is outside the hash preimage. Signatures that were already unreferenced in the parent are kept; only *newly* orphaned ones are collected. +- **CoT re-sugars.** When the prepended field is exactly the `cot` reasoning field on a `Predict`, the augmented signature copy keeps the base name so the canonical printer re-sugars it as `cot `; otherwise it gets a fresh unique name (`_`). + +## Errors + +| Error | Meaning | +|---|---| +| `EditError::Apply { index, edit, reason }` | Edit `index` could not be applied to the (partially edited) program; carries the offending edit and an `ApplyError`. | +| `EditError::Invalid(ValidateError)` | Every edit applied, but the resulting program failed the load-time rules — the error is the validator's own. | + +`ApplyError` is the locally-checkable failure set: `StaleNode`, `WrongKind` (e.g. `SetStop` on a `Predict`), `DuplicateField`, `UnknownTool`, `ToolCapsExceedProgram` (a tool's caps exceed the program ceiling), `ToolAlreadyDeclared`, `ToolNotDeclared`, `NotInSeq` (only `Seq` steps can be removed), `Unparented`. + +## `legal_edits`: the proposer menu + +```rust +pub fn legal_edits(&self, at: NodeId) -> Vec +``` + +The menu of edit kinds structurally admissible at a node — lightweight, serializable `EditKind` descriptors suitable for prompting an LLM proposer: + +| Node | Menu | +|---|---| +| `Predict` leaf | `AugmentSig`, `SetInstructionDefault`, `SwapToAgent` | +| `AgentLoop` leaf | `AugmentSig`, `SetInstructionDefault`, `SwapToPredict`, `SetStop`, plus one `AddTool { tool }` or `RemoveTool { tool }` entry per program tool | +| Any non-root node that is not a `Refine` judge | `WrapRetry` (judges must stay bare leaves) | +| Any `Seq` step | `Remove` | + +The menu is purely structural — data-flow legality (whether a removal orphans a downstream binding) is still `validate()`'s call, surfaced by `edited`. A stale id yields an empty menu. + +This is exactly how the shipped [Structural optimizer](/docs/optimizers/structural) proposes edits: it serializes the menu, has a reflection LM choose one entry, applies the choice through `edited`, and gates the child against the parent on a shared minibatch. + +## `migrate_overlay`: carrying tuned values across an edit + +```rust +pub fn migrate_overlay(parent: &Program, overlay: &Overlay, child: &Program) -> Overlay +``` + +An edit changes the program hash, so overlays minted against the parent no longer apply to the child. `migrate_overlay` carries value-level progress across the structural change: for every entry in the overlay, it re-mints the entry against the child when the child has a slot at the same path and kind whose owning leaf/tool still has a *carrying* signature — inputs identical (names and types, in order) and every parent output present in the child's outputs. Outputs may widen: that is what lets instruction and demos survive `AugmentSig` (demo rows still map onto the base fields; the new field is simply absent from the row). `ModelRef` entries are re-minted by model *name*, not ordinal. `ToolSet` entries are re-minted by tool name and intersected with what the child's agent still declares — partial survival carries the selection forward; a selection with no survivors is dropped. Entries that no longer fit are dropped; a base-mismatched overlay yields an empty result. + +```rust +use dspy_rs::ir::{Edit, migrate_overlay}; + +let leaf = program.leaf_id("drafter").expect("leaf exists"); +let child = program.edited(&[Edit::SetInstructionDefault { + leaf, + text: "Answer in one short sentence.".into(), +}])?; +let carried = migrate_overlay(&program, &tuned_overlay, &child); +``` + +## See also + +- [Program and nodes](/docs/components/program-and-nodes): the value half — params, `Overlay`, and `bake` +- [Structural](/docs/optimizers/structural): the shipped optimizer over this calculus — LM-guided edit choice, `migrate_overlay`, minibatch gating +- [Optimizer engine](/docs/components/optimizer-engine): how candidates are evaluated; a structural optimizer proposes `Edit`s where a prompt optimizer proposes overlays +- [Runtime](/docs/components/runtime): loading and running the edited program +- [The .dsrs file](/docs/components/dsrs-file): the canonical text the child prints to diff --git a/docs/docs/components/evaluation.mdx b/docs/docs/components/evaluation.mdx index af216840..e7d97821 100644 --- a/docs/docs/components/evaluation.mdx +++ b/docs/docs/components/evaluation.mdx @@ -1,6 +1,6 @@ --- title: "Evaluation" -description: "TypedMetric, Eval, the trainset evaluation loop, and feedback helper functions" +description: "TypedMetric, Eval, and the trainset evaluation loop" icon: "gauge-high" --- @@ -60,6 +60,16 @@ where prediction: &Predicted, trace: Option<&Trace>, ) -> Result; + + // Optional; the default returns no span scores. + async fn evaluate_spans( + &self, + example: &E, + prediction: &Predicted, + trace: &Trace, + ) -> Result> { + Ok(Vec::new()) + } } ``` @@ -71,6 +81,32 @@ where Return `Eval::score(f64)` for a numerical score, `Eval::with_feedback(f64, text)` to also explain why. Scores are 0.0 to 1.0 by convention. +## Per-span credit + +`evaluate` assigns one score to the whole rollout. For a multi-step module that single score over-credits: a good final answer marks every intermediate `Predict` call as good, including a step a later call had to recover from. `evaluate_spans` is the optional hook for per-span credit. The evaluation loop calls it once per traced rollout, after `evaluate`, and stamps each returned `Eval` onto its span (`Span::eval`); pairs whose id is not in the trace are ignored. + +```rust +async fn evaluate_spans( + &self, + example: &QARow, + _prediction: &Predicted, + trace: &Trace, +) -> Result> { + // Score each draft call on its own answer; the refine step may have + // recovered from a bad one. + Ok(trace + .for_component("draft") + .filter_map(|span| { + let answer = span.output.as_ref()?.get("answer")?.as_str()?; + let score = (answer == example.answer) as u8 as f64; + Some((span.id, Eval::score(score))) + }) + .collect()) +} +``` + +Demo harvesting (`BootstrapFewShot`, `MIPROv2`, `SIMBA`) prefers a span's own eval over the rollout score when gating and ranking demo candidates, so a scored-down span stays out of the demo pool even when its rollout won, and a scored-up span qualifies even when its rollout lost. Spans you leave out keep whole-rollout credit, and a metric that implements only `evaluate` behaves exactly as before. See [Optimizers](/docs/components/optimizers) for the harvesting semantics. + ## Eval and Rollout `Eval` is defined in the trace module and re-exported by `evaluate`. @@ -93,25 +129,6 @@ The public evaluation entry points return `Vec`; the traced `Rollout` path Each example runs inside a trace capture scope and the metric receives that rollout's `Trace`. Metric evaluation itself happens outside the scope, so LM-as-judge metrics do not pollute the execution trace. -## Feedback helpers - -Helper functions in `evaluate::feedback_helpers` construct `Eval`s with structured textual feedback for common domains. All return `Eval`. - -| Function | Signature | Purpose | -|---|---|---| -| `retrieval_feedback` | `(retrieved: &[impl AsRef], expected: &[impl AsRef], context_docs: Option<&[impl AsRef]>)` | Document retrieval. Score is F1; feedback lists correctly retrieved, missed, and incorrectly retrieved documents with precision, recall, and F1. Precision is `0.0` when `retrieved` is empty; recall is `1.0` when `expected` is empty | -| `code_pipeline_feedback` | `(stages: &[(CodeStage, StageResult)], final_score: f64)` | Code generation pipelines. Feedback reports each stage in order and stops at the first failure; the score is the caller-supplied `final_score` | -| `multi_objective_feedback` | `(objectives: &HashMap, weights: Option<&HashMap>)` | Multi-objective evaluation. Score is the weighted average (default weight `1.0` per objective); feedback lists each objective, sorted by name, plus the aggregate | -| `string_similarity_feedback` | `(predicted: &str, expected: &str)` | String comparison. `1.0` for an exact trimmed match, `0.95` for a case-insensitive match, otherwise word-level F1 with missing and extra words listed | -| `classification_feedback` | `(predicted_class: &str, expected_class: &str, confidence: Option)` | Classification. `1.0` on exact class match, `0.0` otherwise; feedback names the expected and predicted classes and the confidence when given | - -Supporting enums for `code_pipeline_feedback`: - -| Enum | Variants | -|---|---| -| `CodeStage` | `Parse`, `Compile`, `Execute`, `Test` | -| `StageResult` | `Success`, `Failure { error: String }` | - ## Metrics and optimizers Optimizers call the evaluation loop internally; the metric you hand them determines what they can do with the results. diff --git a/docs/docs/components/fx.mdx b/docs/docs/components/fx.mdx index 71d6b8ec..8db3baf7 100644 --- a/docs/docs/components/fx.mdx +++ b/docs/docs/components/fx.mdx @@ -44,13 +44,19 @@ Internally, resolved predictors are cached by `(signature type, name, config has | `new()` | Empty parameter set. | | `set(name, state)` | Sets the full `PredictState` (instruction plus demos) for a named predictor. | | `set_instruction(name, instruction)` | Overrides just the instruction, preserving any demos already set for that name. | +| `clear_instruction(name)` | Explicitly resets the name's instruction to the signature default — wins over any instance override when injected ambiently. | +| `set_demos(name, rows)` | Sets the demo rows (flat JSON objects) as an explicit set: an empty vec means "no demos", overriding instance demos. | | `get(name)` | Returns `Option<&PredictState>` for the name. | | `is_empty()` | True when no entries are set. | | `to_module_state()` | Converts to `ModuleState`, the persistence format shared with struct-based modules. | | `from_module_state(state)` | Builds `Params` from a saved `ModuleState`. | +| `bind(program)` (`ir` feature) | Binds name-keyed params against a compiled `Program` into an `ir::Overlay` — how a module-lane `Candidate` becomes evaluable on a program target. | +| `from_overlay(program, overlay)` (`ir` feature) | The inverse: unbinds an `ir::Overlay` into `Params`. | Because `Params` round-trips losslessly through `ModuleState::save` and `ModuleState::load`, persistence works across both authoring styles. +`Params` is also the optimizer's candidate-injection currency for struct-held modules: each `Predict` leaf consults the ambient `Params` at call time and binds the entry matching its component name (the name stamped by `Predictors` discovery or `PredictBuilder::named`), with ambient values winning over instance state per slot. See [Optimizers](/docs/components/optimizers). + ## `with_params` ```rust diff --git a/docs/docs/components/modules.mdx b/docs/docs/components/modules.mdx index ed8a4d71..615b3c83 100644 --- a/docs/docs/components/modules.mdx +++ b/docs/docs/components/modules.mdx @@ -1,10 +1,10 @@ --- title: 'Modules' -description: 'The Module trait, batch execution, combinators, ChainOfThought, ReAct, and signature augmentation' +description: 'The Module trait, batch execution, predictor discovery via Predictors, ChainOfThought, and signature augmentation' icon: 'circle-nodes' --- -A module is a prompting strategy over a signature. Everything callable in dsrs implements `Module`: the bare LM call ([`Predict`](/docs/components/predict)), `ChainOfThought`, `ReAct`, and any struct you compose from them. Swapping `Predict` for `ChainOfThought` changes the output type, and the compiler surfaces every downstream site that must change. +A module is a prompting strategy over a signature. Everything callable in dsrs implements `Module`: the bare LM call ([`Predict`](/docs/components/predict)), `ChainOfThought`, and any struct you compose from them. Swapping `Predict` for `ChainOfThought` changes the output type, and the compiler surfaces every downstream site that must change. ## Usage @@ -54,13 +54,13 @@ struct Answer { answer: String, } -#[derive(facet::Facet)] -#[facet(crate = facet)] struct Rag { condense: Predict, answer: ChainOfThought, } +dspy_rs::predictors!(Rag { condense, answer }); + impl Module for Rag { type Input = CondenseInput; type Output = WithReasoning; @@ -78,9 +78,9 @@ impl Module for Rag { let rag = Rag { condense: Predict::new(), answer: ChainOfThought::new() }; ``` -The `facet::Facet` derive is what lets optimizer discovery find the `Predict` leaves inside the struct. `forward` is plain async Rust, so branching, loops, and early returns between the LM calls need no framework support. +The `predictors!` line is what makes the module optimizable and persistable: it names the `Predict` leaves for optimizer discovery (see [Predictor discovery](#predictor-discovery-predictors) below). `forward` is plain async Rust, so branching, loops, and early returns between the LM calls need no framework support. -Source: `crates/dspy-rs/src/core/module.rs`, `core/module_ext.rs`, `modules/chain_of_thought.rs`, `modules/react.rs`, `augmentation.rs`. All items below are re-exported at the crate root unless noted. +Source: `crates/dspy-rs/src/core/module.rs`, `modules/chain_of_thought.rs`, `augmentation.rs`. All items below are re-exported at the crate root unless noted. ## The `Module` trait @@ -108,7 +108,32 @@ Every call returns [`Predicted`](/docs/components/predict): the output s `forward` takes `input` by value. This is deliberate: pipeline authors move fields into sub-module inputs with zero clones. The cost is one input clone per example in evaluation loops that reuse a trainset. -To author a module: define a struct holding `Predict`/`ChainOfThought` fields, derive `facet::Facet` so the optimizer's walker can discover the `Predict` leaves, and implement `forward`, as in the usage example above. +To author a module: define a struct holding `Predict`/`ChainOfThought` fields, declare those fields with `predictors!` so optimizers and `ModuleState` can address them by name, and implement `forward`, as in the usage example above. + +## Predictor discovery: `Predictors` + +Optimizable leaves are declared **explicitly** — there is no reflection walker and no derive magic. A module that wants to be optimizable (or persistable via [`ModuleState`](/docs/components/state)) implements the `Predictors` trait, almost always through the `predictors!` macro: + +```rust +dspy_rs::predictors!(Rag { condense, answer }); +``` + +expands to + +```rust +impl Predictors for Rag { + fn predictors(&self) -> Vec<(String, &dyn PredictorInfo)> { /* ("condense", &self.condense), ... */ } + fn predictors_mut(&mut self) -> Vec<(String, &mut dyn PredictorInfo)> { /* ... */ } +} +``` + +Each field's identifier becomes its leaf name. The names are the *canonical identity* of each leaf — the trace-name contract: + +1. They become the leaf's trace-span component name (the optimizer stamps them via `PredictorInfo::set_trace_name` once per run). +2. Optimizer candidates address leaves by these names (ambient `fx::Params` entries bind per leaf at call time). +3. `ModuleState` persists per-leaf state under them. + +Names must be unique within a module and stable across `predictors()`/`predictors_mut()`. `PredictorInfo` is the typed, object-safe per-leaf view: read methods (`schema()`, `instruction()`, `default_instruction()`, `demos_as_json()`, `dump_state()`) plus two boundary mutations — `set_trace_name` (the naming pass) and `load_state` (the install seam, used by `ModuleState::apply` and the optimizer's one-shot install of the winning candidate; candidate *evaluation* never calls it). See [Optimizers](/docs/components/optimizers). ## Batch execution: `forward_all` @@ -127,20 +152,8 @@ pub async fn forward_all( | Concurrency | Bounded by `max_concurrency` (`buffer_unordered`). | | Failure isolation | Returns `Vec>`; one failure does not abort the batch. | | Ordering | Results preserve input order regardless of completion order. | -| Progress | Renders a progress bar on stderr. | | Tracing | Instrumented as `dsrs.forward_all` at debug level. | -## Combinators: `ModuleExt` - -`ModuleExt` is blanket-implemented for every `Module`. It post-processes output without a full `impl Module`. - -| Method | Wrapper | Closure | Semantics | -|--------|---------|---------|-----------| -| `.map(f)` | `Map` | `Fn(M::Output) -> T` | Infallible output transform. | -| `.and_then(f)` | `AndThen` | `Fn(M::Output) -> Result` | Fallible output transform; an `Err` propagates. | - -Both wrappers keep `Input = M::Input`, set `Output = T`, and pass `CallMetadata` through unchanged. The wrapper structs derive `Facet` with the inner module as a real field (the closure is `#[facet(opaque, skip)]`), so the inner `Predict` leaves remain visible to optimizer discovery through the wrapper. - ## `ChainOfThought` `ChainOfThought` is pure sugar, a type alias rather than a distinct struct: @@ -164,39 +177,9 @@ pub type ChainOfThoughtOutput = WithReasoning<::Output>; For reasoning models (o1, o3, DeepSeek-R1) prefer bare `Predict`. An explicit `reasoning` field on top of internal thinking is redundant and can hurt quality. -## `ReAct` - -`ReAct` runs a thought, action, observation loop over a set of tools, then extracts a typed answer. Bounds: `S: Signature`, `S::Input: Schema + Clone`, `S::Output: Schema`. As a `Module`: `Input = S::Input`, `Output = S::Output` (no wrapper type). - -Internally it holds two `Predict` leaves: an action step (inputs `input`, `trajectory`; outputs `thought`, `action`, `action_input`, all strings) and an extract step (inputs `input`, `trajectory`; output `S::Output`). Both are visible to optimizer discovery. - -### Loop semantics - -1. The input struct is serialized to JSON and becomes the `input` field of both internal signatures. -2. The trajectory is seeded with a tool manifest (`Available tools:` with each tool's name and description, or `(none)`). -3. Each step, up to `max_steps` (default 4): the action predictor produces `thought`, `action`, `action_input`. The action name is trimmed of whitespace and surrounding quotes. -4. If the action is `finish`, `final`, or `done` (case insensitive), the loop stops. -5. Otherwise the named tool runs with `action_input` as its argument string. Matching is case insensitive: exact name, or substring containment in either direction. An unknown name yields the observation `tool_not_found: {name}`; a tool error yields `tool_error: {err}`. The step (thought, action, input, observation) is appended to the trajectory. -6. After the loop ends (terminal action or steps exhausted), the extract predictor reads the full trajectory and returns `S::Output`. - -Metadata: each executed tool is recorded in `CallMetadata::tool_calls` with id `react-step-{n}`, and the manifest plus formatted per-step traces appear in `CallMetadata::tool_executions`, merged with the extract call's metadata. - -### Builder - -`ReAct::::new()` equals `ReAct::::builder().build()`. The builder type is `ReActBuilder`; it is not re-exported at the crate root (full path `dspy_rs::modules::react::ReActBuilder`), so obtain it through `ReAct::::builder()`. - -| Method | Signature | Effect | -|--------|-----------|--------| -| `action_instruction` | `(impl Into) -> Self` | Instruction override for the action predictor. | -| `extract_instruction` | `(impl Into) -> Self` | Instruction override for the extract predictor. | -| `max_steps` | `(usize) -> Self` | Loop bound; clamped to at least 1. Default 4. | -| `add_tool` | `(impl ToolDyn + 'static) -> Self` | Registers one rig tool. | -| `with_tools` | `(impl IntoIterator>) -> Self` | Registers many tools. | -| `tool` | `(name, description, Fn(String) -> Future) -> Self` | Registers a closure as a tool; arguments arrive as a raw string (typically JSON). | -| `lm` | `(LM) -> Self` | Per-instance LM for both predictors, bypassing the global. | -| `build` | `() -> ReAct` | Constructs the module. | +## Agent loops -Tools implement rig's `ToolDyn`. See [Tools and agents](/docs/components/tools-and-agents). +There is no `ReAct` module. The tool-loop strategy lives in the IR instead: attach tools to a `Predict` (which executes as a 1-node `agent` program, see [Predict](/docs/components/predict)), or declare the loop as a first-class `AgentLoop` node with the `#[agent]` macro inside a `#[module]`. See [Tools and agents](/docs/components/tools-and-agents). ## Augmentation @@ -231,5 +214,4 @@ The derive generates a `With{Name}` wrapper struct: the augmentation fields p - [How DSRs thinks](/docs/getting-started/how-dsrs-thinks) for the call path from module to LM - [ChainOfThought smoke example](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/91-smoke-slice2-chain-of-thought.rs) - [Module authoring smoke example](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/92-smoke-slice3-module-authoring.rs) -- [ReAct operational smoke example](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/93-smoke-slice4-react-operational.rs) - [Module iteration example](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/02-module-iteration-and-updation.rs) diff --git a/docs/docs/components/optimizer-engine.mdx b/docs/docs/components/optimizer-engine.mdx index d2d965d2..e73ead21 100644 --- a/docs/docs/components/optimizer-engine.mdx +++ b/docs/docs/components/optimizer-engine.mdx @@ -1,65 +1,94 @@ --- title: "Optimizer Engine" -description: "EvalEngine, ProgramEvalEngine, Candidate, Budget, RolloutCache, ScoreMatrix, and the Pareto and outcome types every optimizer shares" +description: "Engine, OptimizeTarget, Candidate, Budget, RolloutCache, ScoreMatrix, and the Pareto and outcome types every optimizer shares" icon: "gauge-high" --- -The optimizer engine is the shared evaluation core under every optimizer. Strategies register candidates and ask the engine to evaluate them; the engine handles overlay application, rollout fan-out, caching, budget accounting, and score bookkeeping. Every optimizer in DSRs (COPRO, GEPA, MIPROv2, bootstrap) is a thin strategy over this core, which exists in two lanes: +The optimizer engine is the shared evaluation core under every optimizer. Strategies register candidates and ask the engine to evaluate them against an `OptimizeTarget`; the engine handles candidate binding, rollout fan-out, caching, budget accounting, and score bookkeeping. Every optimizer in DSRs (COPRO, GEPA, MIPROv2, SIMBA, bootstrap, Structural) is a thin strategy over this one core. -| Lane | Engine | Candidate type | Parallelism | +There is **one** `Engine`. What varies is the target — the lane-erased pair of (thing under optimization, evaluation harness): + +| Lane | Constructor | Candidate currency | Winner | |---|---|---|---| -| Module lane | `EvalEngine` | `Candidate` (named overlays), applied and restored through the mutation seam | Examples within one candidate; candidates are serialized | -| Program lane (`ir` feature) | `ProgramEvalEngine` | `ir::Overlay`, read through at render time, never applied | Candidates and examples together in one fan-out | +| Module lane | `OptimizeTarget::module(&mut module, &trainset, &metric)` | `Candidate` (name-keyed slots), injected *ambiently* per rollout via `fx::with_params` — never applied by mutation | Installed onto the module through `PredictorInfo::load_state`, once, by `OptimizeTarget::install` | +| Program lane (`ir` feature) | `OptimizeTarget::program(&interp, &examples, &metric)` | `ir::Overlay` (or a `Candidate` bound through `fx::Params::bind`), read through at render time | Retrievable as an `Arc` via `OptimizeTarget::winner_overlay`, for `Program::bake` | + +Because candidate injection is ambient in both lanes — nothing is ever applied to shared state during evaluation — rollouts for *different candidates* share one bounded-concurrency fan-out. All items on this page are exported from the crate root except where a feature gate is noted. + +## `OptimizeTarget<'a>` + +The thing an optimizer optimizes: a module or a program, packaged with its example set (by reference) and metric. + +| Method | What it does | +|---|---| +| `module(module, trainset, metric)` | Module-lane target: a typed `Module + Predictors`, a `&[E]` trainset (`E: ToInput + Serialize`), and a `TypedMetric`. Runs the **naming pass**: every leaf declared via `Predictors` is stamped with its declared name (`PredictorInfo::set_trace_name`), so trace spans, candidate entries, and persistence all address the same names. | +| `module_with_valset(module, trainset, valset, metric)` | Same, with an optional validation set. When `Some`, the valset examples become the *leading* columns and the trainset the trailing ones — the layout GEPA's Pareto bookkeeping uses. | +| `program(interp, examples, metric)` | Program-lane target: an interpreter-loaded `Program`, labeled `DemoRow` examples, and a `ProgramMetric`. | +| `leaves() -> &[LeafInfo]` | The optimizable leaves' read surface, snapshotted at construction: per leaf, `name`, current `instruction`, `default_instruction`, `demos` as flat JSON rows, and `input_fields`/`output_fields` as `(lm name, docs)` pairs. `LeafInfo::schema_for_reflection()` renders the field contract for reflection prompts. | +| `num_examples()`, `has_valset()` | Example-set access. | +| `val_columns()`, `train_columns()` | The scoring columns (validation prefix, or every example) and the minibatch-sampling pool (trainset suffix, or every example). | +| `install(&winner)` | Installs the winning `Candidate` — the **one** mutation of the run. Module lane: merges each slot into the named leaf's state through `PredictorInfo::load_state`. Program lane: binds the winner to an overlay. | +| `winner_overlay()` | The installed winner as a bound `Arc` (program lane only). | +| `candidate_outputs(indices, &candidate)` | Runs the given examples under the candidate and returns bare output values, no metric and no trace capture — GEPA's best-output collection. | -Both lanes share the same `EngineConfig`, `Budget`, `Spend`, `RolloutCache`, `ScoreMatrix`, Pareto bookkeeping, outcome types, and minibatch gate. All items on this page are exported from the crate root except where a feature gate is noted. +`ProgramMetric` is the JSON-native sibling of `TypedMetric`: loaded programs have no static output type, so the metric scores the interpreter's output `JsonMap` against a labeled `DemoRow`. + +```rust +pub trait ProgramMetric: Send + Sync { + async fn evaluate( + &self, + example: &DemoRow, + output: &JsonMap, + trace: Option<&Trace>, + ) -> Result; +} +``` -## `EvalEngine<'m, E, MT>` +## `Engine` -Owns the example set, the candidate registry, the score matrix, the rollout cache, and the budget meter. `E` is the trainset [row type](/docs/components/data); construction requires `E: Serialize`. Strategies register candidates and call `evaluate` or `evaluate_gated`; the engine handles application, fan-out, caching, accounting, and bookkeeping. +Owns the candidate registry, the score matrix, the rollout cache, and the budget meter. Strategies register candidates and call the evaluate methods against a target; the engine handles binding, fan-out, caching, accounting, and bookkeeping. | Method | What it does | |---|---| -| `new(examples, metric, config)` | Builds an engine over `Vec`, a `&MT` metric, and an `EngineConfig`. Example uids are content hashes of the whole row. | -| `evaluate(module, candidate, subset)` | Evaluates one registered candidate over `subset` example indices (`None` means the full set). Applies the overlay, fans out uncached rollouts with bounded concurrency under per-rollout trace capture, restores the module, records scores. Returns `EvalOutcome`. | -| `evaluate_gated(module, candidate, minibatch, threshold)` | The GEPA acceptance pattern: evaluates on `minibatch`; only a minibatch mean strictly greater than `threshold` promotes to a full-set evaluation. Returns `GateOutcome`. | -| `register(candidate)` | Registers a `Candidate`, deduplicating by content hash. Returns its index; a duplicate returns the existing index. | -| `candidate(i)`, `candidate_hash(i)`, `num_candidates()` | Candidate registry access. | -| `examples()`, `num_examples()`, `config()`, `spend()`, `matrix()` | State access. | +| `new(config)` | Builds an engine from an `EngineConfig`. Examples and metric live on the target, not the engine — one engine can serve successive targets (a `Box` pipeline sharing one budget). | +| `evaluate_many(target, candidates, subset)` | Evaluates N registered candidates over `subset` example indices (`None` = the target's full set) in **one** bounded-concurrency fan-out — candidate-level parallelism in both lanes. Cached rollouts return their `Eval` with `trace: None` and consume no budget. Returns `BatchEvalOutcome`. | +| `evaluate(target, candidate, subset)` | Single-candidate convenience over `evaluate_many`. Returns `EvalOutcome`. | +| `evaluate_gated(target, candidate, minibatch, threshold)` | The minibatch gate (the GEPA/SIMBA acceptance pattern): evaluates on `minibatch`; only a minibatch mean strictly greater than `threshold` promotes to a full-set evaluation. Returns `GateOutcome`. | +| `register(candidate)` | Registers a module-lane `Candidate`, deduplicating by content hash. Returns its index; a duplicate returns the existing index. | +| `register_overlay(overlay)` | Registers a program-lane `ir::Overlay`, deduplicating by `Overlay::hash()`. Returns its index. | +| `candidate(i) -> Option<&Candidate>`, `candidate_hash(i)`, `num_candidates()` | Candidate registry access (`candidate` is `None` for an overlay entry). | +| `config()`, `spend()`, `matrix()`, `cache()` | State access. | | `pareto()`, `pareto_over(columns)` | Dominance views over the score matrix (all columns, or a subset). | | `budget_allows(n)` | Whether `n` more rollouts fit the remaining budget. | | `charge(metric_calls, lm_calls)` | Charges auxiliary spend the engine did not run itself: reflection LM calls, teacher passes. | -| `checkpoint()` | Serializes engine state (example uids, candidates, matrix, cache, spend) to a JSON string, format version 1. | -| `resume(examples, metric, config, checkpoint)` | Rebuilds an engine from a checkpoint. Fails if the version or the example set does not match. Completed rollouts are served from the restored cache instead of re-executing. | +| `peak_candidate_concurrency()` | High-water mark of *distinct candidates* with rollouts in flight simultaneously — the parallelism gauge. | -`evaluate` requires `E: ToInput + Sync`, `M: Module + Facet`, and `MT: TypedMetric`. The metric runs outside the trace capture scope, so LM-as-judge metrics do not pollute the execution trace. - -**Concurrency model.** Candidates mutate shared predictor state, so the engine serializes candidate application and parallelizes across examples within one candidate (bounded by `EngineConfig::concurrency`). Candidate-level parallelism requires overlays resolved at render time; that is the program lane below. +The metric runs outside the trace capture scope, so LM-as-judge metrics do not pollute the execution trace. If the uncached portion of a batch does not fit the remaining budget, the engine runs nothing, leaves spend unchanged, and returns `BudgetExhausted`. ## `EngineConfig` | Field | Default | Meaning | |---|---|---| -| `concurrency: usize` | `16` (`DEFAULT_EVAL_CONCURRENCY`) | Rollouts in flight at once within one fan-out. | +| `concurrency: usize` | `16` (`DEFAULT_EVAL_CONCURRENCY`) | Rollouts in flight at once within one evaluation batch. | | `budget: Budget` | `Budget::unlimited()` | Hard spend caps; the engine stops cleanly when a batch would not fit. | -| `cache_salt: u64` | `0` | Folded into every cache key. Bump it when changing LM sampling settings outside the candidate; sampling params are not yet part of candidate identity. | +| `cache_salt: u64` | `0` | Folded into every cache key. Bump it when changing LM sampling settings outside the candidate; sampling params are not part of candidate identity. | -## Candidates and the mutation seam +## Candidates -A `Candidate` is data: `overlays: BTreeMap` mapping predictor name to a partial parameter update, plus a stable content hash. The empty candidate (`Candidate::default()`) is the baseline, the module exactly as it is. +A `Candidate` is data: `slots: BTreeMap` mapping leaf name (the `Predictors` contract name) to a partial per-leaf configuration, plus a stable content hash. It is cheap to clone, serializable, and **never applied by mutation**: the engine scopes it ambiently around each rollout (`fx::with_params`); the single mutating step is the caller-driven final `OptimizeTarget::install`. The empty candidate (`Candidate::default()`) is the baseline, the module exactly as it is. | Item | Signature or fields | |---|---| -| `Overlay` | `instruction: Option`, `demos: Option>`. `None` leaves the current value untouched. Demo rows are flat JSON objects, input and output fields merged. | +| `CandidateSlot` | `instruction: Option`, `clear_instruction: bool`, `demos: Option>`. Unset fields leave the leaf's incumbent value untouched; `clear_instruction` explicitly resets to the signature default, winning over any instance override. Demo rows are flat JSON objects, input and output fields merged. | | `Candidate::new()` | The empty candidate. | | `Candidate::with_instruction(name, text)` | Single-predictor instruction candidate, the COPRO and MIPRO case. | -| `set_instruction(name, text)`, `set_demos(name, rows)` | Builder-style mutators. | -| `is_empty()`, `stable_hash()` | The hash is canonical: identical content hashes identically across processes and map orderings. It is the cache and checkpoint identity. | -| `CandidateUndo` | Opaque snapshot of pre-overlay `PredictState` for every predictor the candidate touched. | -| `apply_candidate(&mut module, &candidate) -> Result` | The one place candidate state is written, through `DynPredictor::apply_update`. If any overlay fails to apply, the overlays applied so far are rolled back before the error returns. | -| `restore_candidate(&mut module, undo) -> Result<()>` | Restores the saved state. Attempts every predictor even if one fails, then reports the first error. | +| `set_instruction(name, text)`, `clear_instruction(name)`, `set_demos(name, rows)` | Builder-style mutators. `set_demos` with an empty vec clears the demo set. | +| `instruction_of(name)`, `demos_of(name)` | Read accessors. | +| `is_empty()`, `stable_hash()` | The hash is canonical: identical content hashes identically across processes and map orderings. It is the cache identity. | +| `to_params()` | Converts to the ambient-injection currency: name-keyed [`fx::Params`](/docs/components/fx) with explicit clears preserved. `fx::Params::bind(program)` turns the same value into an `ir::Overlay` for the program lane. | -This module-lane `Overlay` (instruction plus demos per predictor) is a different type from the IR `ir::Overlay`, which maps `ParamId` to `ParamValue` over a compiled `Program`. The program lane below consumes the IR type. +`CandidateSlot` (instruction plus demos per leaf name) is a different type from the IR `ir::Overlay`, which maps `ParamId` to `ParamValue` over a compiled `Program`. The program lane consumes the IR type; `Candidate::to_params()` + `Params::bind` is the bridge between them. ## `Budget` and `Spend` @@ -74,7 +103,7 @@ This module-lane `Overlay` (instruction plus demos per predictor) is a different `Budget::allows(&spend, upcoming_rollouts)` reports whether the batch fits. Zero upcoming rollouts always fit, so cache-only batches never stall. Call and metric caps are enforced prospectively (the batch must fit under the cap); the token cap is retrospective (a batch may overshoot, and the following batch is refused). -`Spend` is what the engine has consumed so far, serialized into checkpoints: +`Spend` is what the engine has consumed so far: | `Spend` field | Meaning | |---|---| @@ -90,13 +119,13 @@ This `Budget` is not the IR runtime `Budget` documented on [Runtime](/docs/compo ## `RolloutCache` -In-memory map from a rollout key to its `Eval` (score plus optional feedback). A candidate re-evaluated on a seen example returns the cached `Eval` with no LM call and no metric call. The cache is serialized into checkpoints, so a resumed run skips completed rollouts. +In-memory map from a rollout key to its `Eval` (score plus optional feedback). A candidate re-evaluated on a seen example returns the cached `Eval` with no LM call and no metric call. The key recipe is `(baseline, candidate, example, salt)`, formatted as four 16-digit hex hashes joined by colons: | Component | Module lane | Program lane | |---|---|---| -| `baseline` | Content hash of the module's `ModuleState` before the overlay is applied. Permanently installing a winner mid-run changes it, invalidating stale entries. | `program.meta.program_hash`. | +| `baseline` | Content hash of the `predictors()` state snapshot (`{name → PredictState}`), computed once at target construction. Installing a winner and building a new target yields a new baseline, invalidating stale entries. | `program.meta.program_hash`. | | `candidate` | `Candidate::stable_hash()` | `Overlay::hash()` | | `example` | Content hash of the example | Content hash of the `DemoRow` | | `salt` | `EngineConfig::cache_salt` | `EngineConfig::cache_salt` | @@ -107,9 +136,7 @@ Public surface: `get`, `insert`, `len`, `is_empty`. **`ScoreMatrix`** is a per-instance matrix of candidates (rows, registration order) by examples (columns). Cells are `None` until scored. Methods: `new(columns)`, `candidates()`, `examples()`, `ensure_rows(n)`, `record(candidate, example, score)`, `score(candidate, example)`, `row(candidate)`, `mean(candidate)`, `best_by_mean()`, `pareto()`, `pareto_over(columns)`. The column-restricted view supports GEPA-style setups where train and validation examples share one matrix. -**`ParetoView`** is a dominance snapshot computed from the matrix. `best_scores()` gives the best score per viewed column; `wins(candidate)` counts columns the candidate wins or ties on (tolerance `1e-6`); `frontier()` lists candidates winning on at least one column; `statistics()` summarizes. A candidate with zero wins is dominated. - -**`ParetoFrontier`** is a standalone convenience wrapper over the same bookkeeping for callers that track candidate payloads outside an engine; GEPA itself uses the engine's matrix directly. It stores `GEPACandidate` payloads, prunes dominated candidates automatically, and offers `add_candidate(candidate, scores) -> bool`, `sample_proportional_to_coverage()`, `best_by_average()`, `candidates()`, `len()`, `is_empty()`, `statistics()`. +**`ParetoView`** is a dominance snapshot computed from the matrix. `best_scores()` gives the best score per viewed column; `wins(candidate)` counts columns the candidate wins or ties on (tolerance `1e-6`); `frontier()` lists candidates winning on at least one column; `statistics()` summarizes. A candidate with zero wins is dominated. GEPA samples parents proportional to their Pareto coverage directly from this view. **`ParetoStatistics`** fields: `num_candidates`, `num_examples_covered`, `avg_coverage: f32`, `max_coverage`, `min_coverage`. A healthy search grows `num_candidates` slowly while `avg_coverage` rises; `num_candidates == 1` means the search has collapsed. @@ -120,39 +147,12 @@ Public surface: `get`, `insert`, `len`, `is_empty`. | `RolloutOutcome` | `example: usize`, `eval: Eval`, `trace: Option`. The trace is `None` when the rollout was served from the cache. | | `CandidateEval` | `candidate: usize`, `rollouts: Vec` in request order. `mean()` is the arithmetic mean over the batch (`0.0` when empty); `scores()` collects the raw scores. | | `EvalOutcome` | `Complete(CandidateEval)` or `BudgetExhausted { needed }`. When exhausted, nothing ran and spend is unchanged; `needed` is the uncached rollout count. `completed()` converts to `Option`. | +| `BatchEvalOutcome` | `Complete(Vec)`, one per requested candidate in request order, or `BudgetExhausted { needed }`. `completed()` converts to `Option>`. | | `GateOutcome` | `BudgetExhausted { needed }`, `Rejected { minibatch }`, or `Promoted { minibatch, full }`. | -| `ProgramEvalOutcome` | `Complete(Vec)`, one per requested candidate in request order, or `BudgetExhausted { needed }`. `completed()` converts to `Option>`. (`ir` feature) | - -## Program lane: `ProgramEvalEngine<'m, MT: ProgramMetric>` - -Behind the `ir` feature. The IR-native evaluation path: candidate `ir::Overlay`s evaluated over one shared `Arc` through the `Interpreter`, with true candidate-level parallelism. The interpreter reads instructions, demos, and code through the overlay at render time, so there is no mutation, no apply, and no restore. Every uncached `(candidate, example)` pair across all requested candidates joins one bounded-concurrency stream. - -```rust -pub trait ProgramMetric: Send + Sync { - async fn evaluate( - &self, - example: &DemoRow, - output: &JsonMap, - trace: Option<&Trace>, - ) -> Result; -} -``` -`ProgramMetric` is the JSON-native sibling of `TypedMetric`: loaded programs have no static output type, so the metric scores the interpreter's output `JsonMap` against a labeled `DemoRow`. +## Rollout mechanics -| Method | What it does | -|---|---| -| `new(examples, metric, config)` | Builds over `Vec`, a `&MT`, and the same `EngineConfig` as the module lane. | -| `evaluate_program_candidates(interp, candidates, subset)` | The IR-native entry point: evaluates N registered candidates over `subset` in one fan-out. Each rollout runs `interp.run(input, Some(overlay), Budget::unlimited())` under its own capture scope with `TraceMeta.candidate_hash` set to the overlay hash and a `program` tag carrying the program hash. Returns `ProgramEvalOutcome`. | -| `evaluate(interp, candidate, subset)` | Single-candidate convenience; returns the module-lane `EvalOutcome`. | -| `evaluate_gated(interp, candidate, minibatch, threshold)` | The same minibatch gate as the module lane; returns `GateOutcome`. | -| `register(overlay)` | Registers an `ir::Overlay`, deduplicating by `Overlay::hash()`. Returns its index. | -| `candidate(i) -> &Arc`, `candidate_hash(i)`, `num_candidates()` | Registry access. | -| `cache()` | The rollout cache; program-lane keys are listed in the table above. | -| `peak_candidate_concurrency()` | High-water mark of distinct candidates with rollouts in flight simultaneously. A value of 2 or more is positive evidence that candidate-level parallelism happened; the module lane is structurally pinned to 1. | -| `examples()`, `num_examples()`, `config()`, `spend()`, `matrix()`, `pareto()`, `pareto_over()`, `budget_allows()`, `charge()` | Identical to the module lane. | - -The program lane shares `Spend` accounting, budget gating, and matrix bookkeeping with the module lane, but it does not offer `checkpoint` or `resume`; those exist only on `EvalEngine`. +Each program-lane rollout runs `interp.run(input, Some(overlay), Budget::unlimited())` under its own capture scope with `TraceMeta.candidate_hash` set to the candidate's hash and a `program` tag carrying the program hash. Each module-lane rollout scopes the candidate's `fx::Params` ambiently around the whole traced rollout (`rollout_traced`), so every `Predict` leaf binds its own entry at call time. After the metric scores a module-lane rollout, any span-level evals it returns from `TypedMetric::evaluate_spans` are stamped onto the trace's spans; demo-harvesting optimizers prefer these over the rollout score (see [Evaluation](/docs/components/evaluation)). Every pending `(candidate, example)` pair — across all candidates, in both lanes — joins one `buffer_unordered` stream bounded by `EngineConfig::concurrency`. ## See also diff --git a/docs/docs/components/optimizers.mdx b/docs/docs/components/optimizers.mdx index 535c21f9..7a459bb7 100644 --- a/docs/docs/components/optimizers.mdx +++ b/docs/docs/components/optimizers.mdx @@ -4,10 +4,10 @@ description: "The Optimizer trait, every optimizer configuration, and every repo icon: "sliders" --- -An optimizer proposes candidates (instruction or demo overlays), evaluates them with your metric on your trainset, and keeps the best. `compile` is the entry point: it takes a module, a training set, and a metric, then searches for better instructions, and in some cases demos, for each `Predict` leaf. The module is mutated in place: after `compile` returns, calling the module produces better results with no code changes. +An optimizer proposes candidates (instruction or demo overlays), evaluates them with your metric on your trainset, and keeps the best. The convenience entry point is each optimizer's `compile_module` method: it takes a module, a training set, and a metric, then searches for better instructions, and in some cases demos, for each `Predict` leaf. After it returns, the winner is installed and calling the module produces better results with no code changes. ```rust -use dspy_rs::{COPRO, Optimizer}; +use dspy_rs::COPRO; let optimizer = COPRO::builder() .breadth(10) @@ -15,11 +15,13 @@ let optimizer = COPRO::builder() .eval_concurrency(16) .build(); optimizer - .compile(&mut module, examples.clone(), &metric) + .compile_module(&mut module, &trainset, &metric) .await?; ``` -All five optimizers are thin strategies over the shared evaluation engine. Candidates are overlays evaluated through a cached, budget-metered, bounded-concurrency fan-out, and winners are installed through the `apply_candidate` seam. The engine types (`EvalEngine`, `Candidate`, `Budget`, `Spend`, `ParetoView`) are documented in [Optimizer engine](/docs/components/optimizer-engine). +The module must declare its optimizable leaves via [`Predictors`](/docs/components/modules#predictor-discovery-predictors) (one `predictors!` line). All six optimizers are thin strategies over the shared evaluation engine. Candidates are **data, never mutation**: each candidate is a name-keyed `Candidate` injected *ambiently* per rollout (`fx::with_params`) — nothing touches the module during evaluation, so different candidates evaluate concurrently — and the winner is installed exactly once at the end (`OptimizeTarget::install`). The engine types (`Engine`, `Candidate`, `Budget`, `Spend`, `ParetoView`) are documented in [Optimizer engine](/docs/components/optimizer-engine). + +Five strategies tune parameter values through overlays; the sixth, [`Structural`](#structural), proposes graph edits over the [edit calculus](/docs/components/edit-calculus) and runs on the program lane only. @@ -64,40 +66,44 @@ All five optimizers are thin strategies over the shared evaluation engine. Candi candidates are overlays; the base program is never mutated -Step-by-step how-to pages: [COPRO](/docs/optimizers/copro), [MIPROv2](/docs/optimizers/miprov2), and [GEPA](/docs/optimizers/gepa). +Step-by-step how-to pages: [COPRO](/docs/optimizers/copro), [MIPROv2](/docs/optimizers/miprov2), [GEPA](/docs/optimizers/gepa), and [Structural](/docs/optimizers/structural). ## The Optimizer trait ```rust -#[allow(async_fn_in_trait)] -pub trait Optimizer { - type Report; +#[async_trait::async_trait(?Send)] +pub trait Optimizer: Send + Sync { + fn engine_config(&self) -> EngineConfig { EngineConfig::default() } - async fn compile( + async fn compile( &self, - module: &mut M, - trainset: Vec, - metric: &MT, - ) -> Result - where - E: ToInput + serde::Serialize + Send + Sync, - M: Module + for<'a> Facet<'a>, - MT: TypedMetric; + target: &mut OptimizeTarget<'_>, + engine: &mut Engine, + ) -> Result; } ``` -`compile` takes exclusive `&mut` access to the module: no concurrent `call()` during optimization. The trainset is `Vec` for any [row type](/docs/components/data) that projects into the module's input via `ToInput` — a `#[derive(Example)]` struct, an `(Input, Output)` tuple, or a hand-written impl; the `Serialize` bound feeds rollout-cache uids, which content-hash the whole row. All type parameters are inferred from the arguments, so no turbofish is needed: `optimizer.compile(&mut module, trainset, &metric)`. The `Facet` bound is what lets the optimizer discover `Predict` leaves by reflection and address them by dotted path. `compile` returns an error when no optimizable predictors are found, when a metric evaluation fails, or when an LM call fails during candidate evaluation. +The trait is **object-safe** by design: optimizers compose (`Box` pipelines can share one `Engine` — one budget, one rollout cache, one score matrix — across stages). The target carries the thing under optimization and its example set *by reference*; the engine carries the spend. + +`OptimizeTarget` is the lane-erased pair of (thing under optimization, evaluation harness), one of two lanes: + +- **`OptimizeTarget::module(&mut module, &trainset, &metric)`** — a typed `Module` (+ `Predictors` discovery), a trainset slice, and a `TypedMetric`. The trainset is `&[E]` for any [row type](/docs/components/data) that projects into the module's input via `ToInput`; the `Serialize` bound feeds rollout-cache uids, which content-hash the whole row. `OptimizeTarget::module_with_valset(...)` adds an optional validation set (the layout GEPA's Pareto bookkeeping uses). Construction runs the **naming pass**: every declared leaf is stamped with its declared name, so trace spans, candidate entries, and persistence all address the same names. +- **`OptimizeTarget::program(&interp, &examples, &metric)`** — an interpreter-loaded IR [`Program`](/docs/components/program-and-nodes), labeled `DemoRow` examples, and a JSON-native `ProgramMetric`. The winner is retrievable as an `ir::Overlay` (`OptimizeTarget::winner_overlay`) for `Program::bake`. + +For the common case you never build these by hand — each optimizer's `compile_module(&mut module, &trainset, &metric)` inherent method constructs a module target and a default engine, runs `compile`, and installs the winner. `compile` returns an error when the target has no optimizable leaves, when a metric evaluation fails, or when an LM call fails during candidate evaluation. -Each optimizer declares its own `Report`: +`compile` returns the `Report` enum (`Report::None`, `Report::Gepa(GEPAResult)`, `Report::Simba(SimbaReport)`, `Report::Bootstrap(BootstrapReport)`, and `Report::Custom(serde_json::Value)` as the third-party extension point), with `into_gepa()`/`into_simba()`/`into_bootstrap()` accessors. The typed `compile_module` sugar unwraps it: -| Optimizer | `Report` type | -|-----------|---------------| +| Optimizer | `compile_module` returns | +|-----------|--------------------------| | `COPRO` | `()` | | `MIPROv2` | `()` | | `GEPA` | `GEPAResult` | | `SIMBA` | `SimbaReport` | | `BootstrapFewShot` | `BootstrapReport` | +`Structural` is the exception to this table and to the trait: it edits program structure, so its input is an interpreter-loaded program rather than a lane-erased target, and its entry point is [`compile_program`](#structural) instead of `compile_module`/`compile`. + ## Choosing an optimizer | Optimizer | Strategy | Needs feedback? | Cost | @@ -107,6 +113,7 @@ Each optimizer declares its own `Report`: | `SIMBA` | Minibatch introspective ascent (demos + rules) | No | Low (steps × minibatch) | | `GEPA` | Genetic-Pareto evolution with feedback | **Yes** | Medium-high (iterations × eval) | | `MIPROv2` | Trace-guided candidate generation | No | Medium (candidates × trials × trainset) | +| `Structural` | LM-guided graph edits over `ir::Edit` (program lane only) | No | Medium (examples + iterations × minibatch) | GEPA is the only optimizer that requires textual feedback from the metric. The others use numerical scores alone. @@ -135,11 +142,11 @@ Trace-guided instruction and demo optimizer. Four phases: one traced teacher pas | `num_trials` | `usize` | `20` | Maximum candidates evaluated per predictor. When lower than `num_candidates`, only the first `num_trials` are evaluated. | | `minibatch_size` | `usize` | `25` | Examples per candidate evaluation. | | `max_bootstrapped_demos` | `usize` | `4` | Demos installed per predictor from successful traces. | -| `min_demo_score` | `f64` | `0.0` | Minimum whole-program score for a trace to qualify as a demo source. | +| `min_demo_score` | `f64` | `0.0` | Minimum score for a span to qualify as a demo source: its own span eval when present, the whole-program score otherwise. | | `eval_concurrency` | `usize` | `16` | Concurrent LM calls during candidate evaluation. | | `seed` | `Option` | `None` | Fixes minibatch sampling. `None` is nondeterministic. | -Public helper types: `PromptCandidate` (an instruction with its evaluated `score: f64`) and `PromptingTips` (the rotation of prompting best practices appended to candidates, `default_tips()` and `format_for_prompt()`). Public methods on `MIPROv2`: `select_best_traces`, `create_prompt_candidates`, `format_schema_fields`. The report is `()`. +Public helper type: `PromptingTips` (the rotation of prompting best practices appended to candidates, `default_tips()` and `format_for_prompt()`). `compile_module` returns `()`; through the trait, the report is `Report::None`. ## GEPA @@ -163,7 +170,7 @@ GEPA errors if any `Eval` from the metric has `feedback: None`. Build metrics wi | `eval_concurrency` | `usize` | `16` | Concurrent LM calls during candidate evaluation. | | `seed` | `Option` | `None` | Fixes minibatch sampling and parent selection. | -GEPA additionally exposes `compile_with_valset(module, trainset, valset, metric)`. With `Some(valset)`, initial evaluation and child scoring use the validation set while parent re-evaluation uses trainset minibatches; with `None`, the trainset serves both roles (this is what `compile` does). +GEPA additionally exposes `compile_module_with_valset(module, trainset, valset, metric)` — sugar over `OptimizeTarget::module_with_valset` plus the `Optimizer` trait. With `Some(valset)`, initial evaluation and child scoring use the validation set while parent re-evaluation uses trainset minibatches; with `None`, the trainset serves both roles (this is what `compile_module` does). ### GEPAResult @@ -236,7 +243,7 @@ The simplest complete optimizer: one teacher pass, one candidate, one comparison | Field | Type | Default | Description | |-------|------|---------|-------------| | `max_demos` | `usize` | `4` | Maximum demos installed per predictor. | -| `min_demo_score` | `f64` | `1.0` | Minimum whole-rollout score for a rollout's spans to qualify as demos. The default assumes a 0 to 1 metric; lower it for graded metrics. | +| `min_demo_score` | `f64` | `1.0` | Minimum score for a span to qualify as a demo: its own span eval when the metric attached one, the whole-rollout score otherwise. The default assumes a 0 to 1 metric; lower it for graded metrics. | | `eval_concurrency` | `usize` | `16` | Concurrent rollouts during evaluation. | | `max_metric_calls` | `Option` | `None` | Hard cap on metric calls (rollouts). | | `max_lm_calls` | `Option` | `None` | Hard cap on LM call units. | @@ -248,23 +255,71 @@ The simplest complete optimizer: one teacher pass, one candidate, one comparison | `baseline_score` | `f64` | Mean metric score of the unmodified module over the trainset. | | `candidate_score` | `Option` | Mean score with demos attached. `None` when no demos were harvested or the budget stopped before the candidate evaluation. | | `adopted` | `bool` | Whether the demo candidate beat the baseline and was installed. | -| `demos_per_predictor` | `BTreeMap` | Demos harvested per predictor, keyed by dotted path. | +| `demos_per_predictor` | `BTreeMap` | Demos harvested per predictor, keyed by leaf name. | | `spend` | `Spend` | Engine spend for the whole run. | `adopted: false` with a populated `demos_per_predictor` means demos were harvested but did not beat the baseline; the module is left unchanged. +## Structural + +LM-guided hill-climbing over the [edit calculus](/docs/components/edit-calculus), program lane only. Each generation gathers the `legal_edits` menu for every leaf, has a reflection LM (`prompt_model`) choose one edit from the serialized menu plus the incumbent's evaluation feedback, applies it with `Program::edited`, carries the incumbent overlay across the change with `migrate_overlay`, loads the child through a caller-supplied `RuntimeEnv` factory, and accepts it through the engine's minibatch gate: only a strict win on the shared minibatch promotes the child to a full-set evaluation and makes it the new incumbent. Edits that fail to apply, children that fail to load, and reflection replies that do not parse are recorded and skipped. How-to: [Structural](/docs/optimizers/structural). + +Entry points: `compile_program(&interp, &examples, &metric, env)` and `compile_program_with_overlay(&interp, Some(overlay), &examples, &metric, env)`, where `env: Fn() -> RuntimeEnv` supplies fresh bindings for each child load. The winner is returned in the report (program plus migrated overlay), never installed; bake it with `Program::bake`. + +| Field | Type | Default | Description | +|-------|------|---------|-------------| +| `num_iterations` | `usize` | `8` | Generations to attempt; each proposes exactly one edit. | +| `minibatch_size` | `usize` | `8` | Examples in the shared minibatch parent and child are compared on. | +| `prompt_model` | `Option` | `None` | Reflection LM that chooses an edit from the menu. Without it the choice is a seeded-uniform pick. | +| `max_rollouts` | `Option` | `None` | Hard cap on evaluation rollouts. Every child is a fresh program that re-scores from scratch. | +| `max_lm_calls` | `Option` | `None` | Hard cap on LM call units (rollouts plus reflection). | +| `eval_concurrency` | `usize` | `16` | Concurrent rollouts during evaluation. | +| `seed` | `Option` | `None` | Fixes minibatch sampling and the fallback edit choice. | + +### StructuralStep and StructuralReport + +Each `StructuralStep` records one generation: + +| Field | Type | Description | +|-------|------|-------------| +| `generation` | `usize` | Generation index, 0-based. | +| `leaf` | `String` | The leaf the chosen edit targets. | +| `edit` | `Edit` | The concrete proposed edit (serde data, replayable against `parent_hash`). | +| `parent_hash` | `u64` | `program_hash` of the parent the edit was applied to. | +| `parent_minibatch_score` | `f64` | Parent's mean on the shared minibatch, the gate threshold. | +| `child_minibatch_score` | `Option` | Child's mean on the same minibatch; `None` when the child never scored. | +| `accepted` | `bool` | Whether the gate promoted the child. | +| `full_score` | `Option` | Full-set mean of the child; `Some` only when accepted. | +| `rejection` | `Option` | Why the child never scored (apply/load failure), when it didn't. | + +The `StructuralReport` summarizes the run: + +| Field | Type | Description | +|-------|------|-------------| +| `program` | `Arc` | The winning program (the input program when nothing was accepted). | +| `overlay` | `Overlay` | The incumbent overlay re-minted against the winner at every accepted edit. | +| `baseline_score` | `f64` | Mean metric score of the input program (plus overlay) over the examples. | +| `final_score` | `f64` | Full-set mean of the final program; equals the baseline when nothing was accepted. | +| `edits` | `Vec` | The accepted edits, in order (a lineage, not one batch). | +| `steps` | `Vec` | Per-generation outcomes, in order. | +| `accepted` / `rejected` | `usize` | Generations promoted / not promoted. | +| `spend` | `Spend` | Engine spend for the whole run, reflection calls included. | + ## Demo harvesting -Demo harvesting is a pure name join over captured traces. A rollout trace records one span per `Predict` invocation under the same dotted-path component name the mutation seam addresses, so successful spans (parsed output present) from rollouts scoring at least the optimizer's `min_demo_score` become flat demo rows for exactly the predictor that produced them: no pointer identity, identical behavior for fx and struct harnesses. Rows are ranked by whole-rollout score, deduplicated on input fields so repeated inputs do not crowd the demo set, and capped per predictor. `BootstrapFewShot`, `MIPROv2`, and `SIMBA` share this machinery; it is internal to the crate and not part of the public API. +Demo harvesting is a pure name join over captured traces. A rollout trace records one span per `Predict` invocation under the leaf name the module declares via `Predictors` (stamped by the target's naming pass), so successful spans (parsed output present) scoring at least the optimizer's `min_demo_score` become flat demo rows for exactly the predictor that produced them: no pointer identity, identical behavior for fx and struct harnesses. Rows are gated and ranked by their effective score, deduplicated on input fields so repeated inputs do not crowd the demo set, and capped per predictor. `BootstrapFewShot`, `MIPROv2`, and `SIMBA` share this machinery; it is internal to the crate and not part of the public API. + +A span's effective score is the whole-rollout metric score unless the metric attached a span-level eval through `TypedMetric::evaluate_spans` (see [Evaluation](/docs/components/evaluation)), which then takes precedence in both directions: a span scored down stays out of the demo pool even when its rollout won, and a span scored up qualifies even when its rollout lost. Without span evals the behavior is exactly the whole-rollout join described above. ## See also -- [Optimizer engine](/docs/components/optimizer-engine): `EvalEngine`, `Candidate`, `Budget`, `Spend`, `ParetoView`, and checkpointing +- [Optimizer engine](/docs/components/optimizer-engine): `Engine`, `OptimizeTarget`, `Candidate`, `Budget`, `Spend`, and `ParetoView` - [Evaluation](/docs/components/evaluation): `TypedMetric`, `Eval`, and `Eval::with_feedback` - [Traces](/docs/components/traces): the span capture that feeds demo harvesting and reflection -- How-to pages: [COPRO](/docs/optimizers/copro), [MIPROv2](/docs/optimizers/miprov2), [GEPA](/docs/optimizers/gepa) +- [The edit calculus](/docs/components/edit-calculus): the structural moves `Structural` proposes +- How-to pages: [COPRO](/docs/optimizers/copro), [MIPROv2](/docs/optimizers/miprov2), [GEPA](/docs/optimizers/gepa), [Structural](/docs/optimizers/structural) - Runnable examples: - COPRO: https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/04-optimize-hotpotqa.rs and https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/02-module-iteration-and-updation.rs - MIPROv2: https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/08-optimize-mipro.rs diff --git a/docs/docs/components/predict.mdx b/docs/docs/components/predict.mdx index 4f4764b4..09849786 100644 --- a/docs/docs/components/predict.mdx +++ b/docs/docs/components/predict.mdx @@ -63,29 +63,30 @@ let predict = Predict::::builder() .build(); ``` -`Demo` is the typed input/output pair for few-shot prompting: `Demo::new(input, output)` with public fields `input: S::Input` and `output: S::Output`. Demos render as user/assistant exchanges in the prompt, and the types guarantee a demo matches the signature — a `Demo` cannot be attached to a `Predict`. To seed a demo from a labeled trainset row, project the row through its `ToInput`/`ToOutput` impls: `Demo::new(row.to_input()?, row.to_output()?)`, see [Data](/docs/components/data). Tools are settable only at build time. Demos and the instruction override are also writable after construction through the optimizer seam (`DynPredictor::apply_update`, `load_state`), see [State](/docs/components/state). The formatted system message and demo turns are cached once per (instruction, demos) configuration; every state mutation invalidates the cache. +`Demo` is the typed input/output pair for few-shot prompting: `Demo::new(input, output)` with public fields `input: S::Input` and `output: S::Output`. Demos render as user/assistant exchanges in the prompt, and the types guarantee a demo matches the signature — a `Demo` cannot be attached to a `Predict`. To seed a demo from a labeled trainset row, project the row through its `ToInput`/`ToOutput` impls: `Demo::new(row.to_input()?, row.to_output()?)`, see [Data](/docs/components/data). Tools are settable only at build time. Demos and the instruction override are also writable after construction through the state-install seam (`PredictorInfo::load_state`, used by `ModuleState::apply` and the optimizer's final install of the winning candidate), see [State](/docs/components/state). Every state mutation invalidates the cached instance overlay. ## Calling | Method | Returns | Use | |--------|---------|-----| -| `.call(input)` | `Result, PredictError>` | The typed direct call | +| `.call(input)` | `Result, PredictError>` | The typed direct call — runs through the IR interpreter | | `.forward(input)` | Same as `call` | `Module` trait hook; delegates to `call`. Callers should invoke `Module::call`, which exists as the future middleware seam | -| `.build_chat(&input)` | `Result` | Inspect or modify the first-turn prompt before sending | +| `.build_chat(&input).await` | `Result` | Inspect or modify the first-turn prompt before sending | | `.call_and_parse(chat)` | `Result<(Predicted, Chat), PredictError>` | Multi-turn: the caller owns the `Chat` between turns | -Each `call` executes this pipeline: +`Predict` executes as a 1-node IR [Program](/docs/components/program-and-nodes): a `predict` leaf named after the component, over `SignatureDef::of::()` — or an `agent` leaf when tools are attached (the IR says `Predict` carries no tools; a tooled predictor *is* an agent loop). Each `call` executes this pipeline: -1. Build the chat: cached system + demo prefix, plus the input formatted as the live user message. -2. Resolve the LM: the per-instance `.lm(...)` if set, otherwise the global `configure()` LM. -3. Consult any active replay scope (see below). -4. Send via `LM::call_with_toolset` in `ToolLoopMode::Auto`, executing tool calls up to `max_tool_iterations`. -5. Parse the response into `S::Output` through the `[[ ## field ## ]]` protocol, evaluating `#[check]` and `#[assert]` constraints. -6. Record a trace span when inside a `capture()` scope. +1. Build (once, then cache) the 1-node program. The leaf name is the component name, so span identity and capture/replay keying are unchanged. +2. Resolve the LM — the per-instance `.lm(...)` if set, otherwise the global `configure()` LM — and load the [Interpreter](/docs/components/runtime) against it (cached; reloaded when the resolved LM changes). +3. Compose the effective `ir::Overlay`: instance state (instruction override + demos, minted once as an overlay against the cached program) plus any ambient optimizer candidate (`fx::with_params` / `fx::with_overlay`), ambient values winning per slot. +4. Run the interpreter (`run_collecting`), which renders the prompt, consults any active replay scope, calls the LM (for an `agent` leaf: the tool loop, with the default `StopSpec` — `until_parse`, `max_turns = 8`), and parses the response through the `[[ ## field ## ]]` protocol, evaluating `#[check]` and `#[assert]` constraints. +5. Reassemble the run's per-leaf metadata (`LeafOutcome`) into `CallMetadata`, and record a trace span when inside a `capture()` scope. + +`build_chat`/`call_and_parse` are the **conversation surface**: the caller owns the `Chat` between turns. Both are thin wrappers over the interpreter's conversation entry ([Runtime](/docs/components/runtime)): `build_chat` renders the opening turn through `Interpreter::conversation_opening`, and `call_and_parse` sends the chat through `Interpreter::run_conversation` — the same overlay-resolved rendering, span recording, and replay interception as the typed `call` path, so a conversation turn and a typed call over the same state produce byte-identical prompts. A turn is not a run: each `call_and_parse` records one trace span, and `seq` increments per turn. When tools are attached, the turn runs the same `AgentLoop` as the typed path, dispatching tool calls through the attached executors; for the "return me the tool calls, I'll execute them" pattern, use `Interpreter::run_conversation_caller_managed` directly. ### Replay interception -Before constructing any client, `Predict` consults the active replay scope. A `Serve` directive returns the recorded span with zero provider calls and zero tool re-executions; a `Refuse` directive returns `PredictError::Replay`; `Live` (or no scope) proceeds normally. Served predictions carry no per-field parse metadata. See [Traces](/docs/components/traces). +Before any provider call, the interpreter consults the active replay scope — typed `call`s and conversation turns alike. A `Serve` directive returns the recorded span with zero provider calls and zero tool re-executions; a `Refuse` directive returns `PredictError::Replay`; `Live` (or no scope) proceeds normally. Served predictions carry no per-field parse metadata. See [Traces](/docs/components/traces). ## `Predicted` @@ -114,12 +115,12 @@ Each `FieldMeta` records `raw_text` (the text the LM produced for that field), ` | Variant | Payload | Fires when | Retryable | |---------|---------|------------|-----------| -| `Lm` | `source: LmError` | The provider failed before returning a usable response: network, rate limit, timeout, bad status | Per `LmError::is_retryable()` | +| `Lm` | `source: LmError` | The provider failed before returning a usable response. Everything the provider stack reports arrives as `LmError::Provider` (provider name, message, source) | No — the underlying rig client owns transport-level retries | | `Parse` | `source: ParseError`, `raw_response`, `lm_usage` | The LM responded but the expected fields could not be extracted or coerced, or an `#[assert]` failed | Yes | | `Conversion` | `source: ConversionError`, `parsed` | The parsed JSON value does not fit the typed output struct | No | | `Replay` | `source: ReplayError` | A strict replay scope refused the call: the live request diverged from its recording, or the recorded span is unusable | No | -`PredictError::class()` buckets into `ErrorClass` (`BadRequest`, `NotFound`, `Forbidden`, `Temporary`, `BadResponse`, `Internal`); `is_retryable()` drives retry logic. `Parse` errors include the raw response and the token usage: failed parses still consume tokens. +`PredictError::class()` buckets into `ErrorClass` (`BadRequest`, `Temporary`, `BadResponse`, `Internal`); `is_retryable()` drives retry logic. `Parse` errors include the raw response and the token usage: failed parses still consume tokens. ## `ToolSet` diff --git a/docs/docs/components/program-and-nodes.mdx b/docs/docs/components/program-and-nodes.mdx index 3a4731b5..3b5f9e79 100644 --- a/docs/docs/components/program-and-nodes.mdx +++ b/docs/docs/components/program-and-nodes.mdx @@ -38,7 +38,7 @@ Nodes form a tree: one parent, one use. Fan-in happens through field references, | Node | Plain words | Main fields | |---|---|---| | `Predict` | One LM call, no tools. `cot` is sugar: a Predict over a reasoning-augmented signature. | `name`, `sig`, `instruction`, `demos`, `model`, `binding` | -| `AgentLoop` | The LM plus tool loop as a first-class unit. | `name`, `sig`, `instruction`, `demos`, `model`, `tools`, `context_policy`, `stop` (max turns, stop tools, until_parse), `budget`, `binding` | +| `AgentLoop` | The LM plus tool loop as a first-class unit. | `name`, `sig`, `instruction`, `demos`, `model`, `tools` (the declared table), `tool_set` (the selection gene), `context_policy`, `stop` (max turns, stop tools, until_parse), `budget`, `binding` | | `Seq` | Runs children in order and exports named fields. | `body`, `out` | | `ForkJoin` | Runs branches concurrently (all succeed or fail fast) and joins their outputs. | `branches`, `join` | | `Route` | Picks one arm by an enum-valued port. | `on`, `arms` (variant, node pairs), `default` | @@ -60,6 +60,7 @@ Every mutable thing in a program is a named, addressable slot; a candidate is an | `Instruction` | The prompt's task description for a leaf. | | `Demos` | Few-shot demonstration rows (input map plus output map each). | | `ToolDesc` | A tool's description text. | +| `ToolSet` | Which of an agent node's *declared* tools the loop carries. Declaration (`AgentLoopNode::tools`) is structural — it is the loop's capability footprint; selection is the gene. An optimizer can drop a distracting tool or bring a declared one back, never add an undeclared one: a value naming a tool outside the declared table is refused at load and on `Overlay::set`. Absent selection = the full declared table. | | `ModelRef` | Which declared model a leaf uses. | | `ContextPolicy` | The agent context policy: history window, tool-result byte cap, playbook text. | | `Code` | Sandboxed JS source (with a stable content hash). Hole and sandboxed-tool implementations are optimizable through this kind. | @@ -68,7 +69,7 @@ Every mutable thing in a program is a named, addressable slot; a candidate is an Slots are addressed by canonical string paths. Node-owned slots use the leaf name as prefix; tool-owned slots use a `tool.` prefix: -- `".instruction"`, `".demos"`, `".model"`, `".context"`, `".code"` +- `".instruction"`, `".demos"`, `".model"`, `".context"`, `".tool_set"`, `".code"` - `"tool..desc"`, `"tool..code"` For example `"drafter.instruction"` or `"tool.search.desc"`. `Program::param_id(path)` resolves a path to its id; after load everything speaks ids. @@ -85,6 +86,7 @@ An `Overlay` is one candidate: a dense set of parameter values layered over a fi | `set_instruction(slot, text)` | Sets an instruction through a typed `Slot` handle. | | `set_demos(slot, rows)` | Sets demo rows through a typed `Slot` handle. | | `set_code(slot, source)` | Sets JS source through a typed `Slot` handle (hash computed automatically). | +| `set_tool_set(&program, slot, tools)` | Sets a tool-set gene through a typed `Slot` handle. Unlike the other typed setters it takes the program and can fail: the value must be a subset of the owning agent node's declared tools, anything else is `OverlayError::ToolSetUndeclared`. | | `resolve(&program, id)` | The effective value: the overlay entry if set, otherwise the slot's default. | | `hash()` | Stable hash over the base plus the set entries in id order. This is the trace's `candidate_hash` and the rollout-cache key. | | `to_named(&program)` | The serde boundary: a path-keyed map (`"": ParamValue`). | @@ -174,7 +176,8 @@ The same rule has a second effect worth knowing: baking changes the program's ha ## See also +- [The edit calculus](/docs/components/edit-calculus): the structural half of program mutation — `Program::edited`, `legal_edits`, `migrate_overlay` - [The .dsrs file](/docs/components/dsrs-file): the canonical text form a program prints to - [Runtime](/docs/components/runtime): loading and running a program, and how `Interpreter::run` reads through an overlay -- [Optimizer engine](/docs/components/optimizer-engine): where candidates and overlays come from during optimization, and how checkpointing saves them +- [Optimizer engine](/docs/components/optimizer-engine): where candidates and overlays come from during optimization - [CLI](/docs/components/cli): checking and serving the baked file diff --git a/docs/docs/components/runtime.mdx b/docs/docs/components/runtime.mdx index 44a9c2bd..22440097 100644 --- a/docs/docs/components/runtime.mdx +++ b/docs/docs/components/runtime.mdx @@ -67,6 +67,48 @@ let out: JsonMap = interp.run(input, overlay, budget).await?; - `overlay` is `Option>`. When present, its `base` must equal the program's hash or the run fails with `RunError::Overlay` before anything executes. - `budget` caps spend for this run. +### `Interpreter::run_collecting` + +`run_collecting(input, overlay, budget)` is `run` with per-leaf metadata: it returns a `RunOutput` — the same output map plus one `LeafOutcome` per successful `Predict`-leaf evaluation, in execution order (`ForkJoin` branches append in declared branch order). This is the seam `Predict` uses to reassemble `CallMetadata` when it executes through the interpreter. + +Each `LeafOutcome` carries: `name` (the program-unique leaf name, the trace span component), `raw_response`, `field_meta` (per-field jsonish coercion flags and `#[check]` results, keyed by canonical field name), `usage`, `model_config_hash`, `span_id` (when a capture scope was active), and — for `AgentLoop` leaves — `tool_calls` and `tool_executions`. Scope rules: `Predict` and `AgentLoop` leaves report, `Hole` leaves do not; only *successful* evaluations report (a failed `Retry` attempt leaves no outcome, the succeeding one reports); an agent whose final output came from stop-tool args has empty `field_meta`; a replay-served leaf reports the recorded raw text and usage with empty `field_meta`. + +### `Interpreter::run_conversation` + +`run_conversation(chat, input, overlay, budget)` is the conversation-in/conversation-out entry: one turn with the program's single leaf over a caller-owned `Chat`, returning `(RunOutput, Chat)` — the turn's output and metadata plus the extended conversation. It exists only for single-leaf programs (one `predict` or `agent` node, what `Predict` compiles to); a multi-node graph is refused with `RunError::Input`. + +```rust +// Opening turn: empty chat plus the typed input. +let (out, mut chat) = interp + .run_conversation(Chat::new(vec![]), Some(input), None, Budget::unlimited()) + .await?; + +// Continuation: append a follow-up and send the chat back. +chat.push_message(Message::user("are you sure?")); +let (out, chat) = interp.run_conversation(chat, None, None, Budget::unlimited()).await?; +``` + +The `chat`/`input` combinations: an empty chat with `Some(input)` renders the opening turn (system + demos + the formatted input, identical to `conversation_opening`); a non-empty chat with `None` is sent as-is; a non-empty chat with `Some(input)` appends the formatted input as the next user turn. A turn is not a run: each call records one trace span (`seq` increments per turn) and meters against its own `budget`. On an `agent` leaf the turn runs the full tool loop, dispatching tool calls through their bound executors. Replay works turn by turn — a span keys on the full chat sent, so a recorded conversation serves each turn with tool effects baked in. + +`conversation_opening(&input, overlay)` renders the opening `Chat` without calling anything: the same overlay-resolved system + demos (+ agent playbook) + input rendering a run would send. Use it to inspect or edit the first turn before `run_conversation`. `Predict::build_chat`/`call_and_parse` are thin wrappers over these two entries. + +### Caller-managed tool loops + +`run_conversation_caller_managed(chat, input, overlay, budget)` is the same turn in suspending mode: when the model requests tool calls on an `agent` leaf, the loop suspends instead of dispatching and returns `ConversationTurn::Suspended(ToolSuspension)`. Execute the calls yourself and feed the results back: + +```rust +let mut turn = interp + .run_conversation_caller_managed(Chat::new(vec![]), Some(input), None, Budget::unlimited()) + .await?; +while let ConversationTurn::Suspended(suspension) = turn { + let results = run_my_tools(suspension.calls()).await; // Vec, one per call + turn = interp.resume_conversation(suspension, results).await?; +} +let ConversationTurn::Complete { run, chat } = turn else { unreachable!() }; +``` + +`ToolSuspension::calls()` is the pending calls in request order; `ToolSuspension::chat()` is the conversation so far, including the assistant tool-call turn. `resume_conversation` records one `ToolRun` event per result (metering the time the suspension was outstanding), pushes one batched tool-result user turn, and continues the loop under the same meters and turn cursor — trace spans, budget metering, and stop-tool semantics are identical to dispatching mode. A stop-tool call completes the turn instead of suspending; a replay scope never suspends (served turns carry every tool effect); Code Mode does not apply, since the caller executes the tools. Dropping a suspension without resuming closes its span as `Cancelled`. Feed a failed tool's error text as its result to keep the conversational-repair behavior of dispatching mode. + ## `Budget` Run-level spend limits; `None` means unlimited. `Budget::default()` and `Budget::unlimited()` are the same: no limits. @@ -126,7 +168,7 @@ The macro creates a module named after the file stem (`qa.dsrs` becomes `mod qa` - `qa::try_program()`: the same, but returns a `Result`. - A generated test, so `cargo test` fails if the file ever becomes invalid. -The file is parsed and validated while your crate compiles, so a broken artifact breaks your build, not a running process. This is the shipping path for programs with host tools or host holes, which `dsrs serve` cannot bind: embed the file, bind your implementations with `bind_host_tool` and `bind_host_hole`, and serve from your own binary. +Validation is layered. **Syntax** is checked at macro expansion through `dsrs-syntax` — the shared `.dsrs` lexer and structural grammar both the macro and the full parser read from — so a malformed file breaks your build. **Semantics** (types, dataflow, capability rules) are checked by the full parser at first use of `program()`, and forced at CI time by the generated test — the sqlx-offline analogue. This is the shipping path for programs with host tools or host holes, which `dsrs serve` cannot bind: embed the file, bind your implementations with `bind_host_tool` and `bind_host_hole`, and serve from your own binary. ## See also diff --git a/docs/docs/components/signatures.mdx b/docs/docs/components/signatures.mdx index b27fd45a..571f997f 100644 --- a/docs/docs/components/signatures.mdx +++ b/docs/docs/components/signatures.mdx @@ -122,7 +122,7 @@ Every field requires exactly one of `#[input]` or `#[output]`, and the signature `#[Schema]` marks a struct or enum as usable inside signature fields. It accepts no arguments. It expands to `#[derive(facet::Facet, serde::Serialize, serde::Deserialize)]` with crate-path attributes; enums additionally receive `#[repr(u8)]` when no explicit `repr` is present. -`#[BamlType]` survives as a backwards-compatible alias with identical expansion. The old vendored BAML stack, including its `#[baml(...)]` attribute grammar, was removed. The schema layer now reads facet metadata instead. +The old vendored BAML stack, including its `#[baml(...)]` attribute grammar and the `#[BamlType]` compat alias, was removed. The schema layer reads facet metadata instead. What the schema builder honors on `#[Schema]` types: diff --git a/docs/docs/components/state.mdx b/docs/docs/components/state.mdx index cdc470b7..601f5298 100644 --- a/docs/docs/components/state.mdx +++ b/docs/docs/components/state.mdx @@ -8,7 +8,7 @@ After an optimizer tunes a module, the improved instructions and demos live only ```rust // After optimization: -ModuleState::from_module(&mut module)?.save("optimized.json")?; +ModuleState::from_module(&module)?.save("optimized.json")?; // In production: let mut module = MyPipeline::new(); @@ -17,12 +17,12 @@ ModuleState::load("optimized.json")?.apply(&mut module)?; ## `ModuleState` -A `ModuleState` holds one `PredictState` per predictor, keyed by the dotted path the optimizer walker discovers (`predictors: BTreeMap`). The `BTreeMap` keeps JSON output stable across runs. Paths follow the module structure: struct fields join with dots (`inner.predictor`), list elements append an index (`steps[0]`), and map entries append an escaped key (`stages['draft']`). +A `ModuleState` holds one `PredictState` per predictor, keyed by the leaf name the module declares via [`Predictors`](/docs/components/modules#predictor-discovery-predictors) (`predictors: BTreeMap`). The `BTreeMap` keeps JSON output stable across runs. The names are the same ones optimizer candidates and trace spans use — one naming contract across persistence, optimization, and capture. | Method | Signature | Behavior | |---|---|---| -| `from_module` | `fn from_module(module: &mut M) -> Result` | Snapshots every `Predict` leaf. Takes `&mut` because leaf discovery uses the exclusive Facet walker; the module is not modified. | -| `apply` | `fn apply(&self, module: &mut M) -> Result<()>` | Applies the state in place. Every path in the state must resolve to a `Predict` leaf; unknown paths are an error. Predictors not named in the state are left untouched. | +| `from_module` | `fn from_module(module: &M) -> Result` | Snapshots every declared `Predict` leaf (instruction override + demos). | +| `apply` | `fn apply(&self, module: &mut M) -> Result<()>` | Applies the state in place, stamping each restored leaf's trace name with its declared name. Every name in the state must resolve to a leaf in the module; unknown names are an error. Predictors not named in the state are left untouched. | | `to_json` | `fn to_json(&self) -> Result` | Serializes to pretty-printed JSON. | | `from_json` | `fn from_json(json: &str) -> Result` | Deserializes JSON produced by `to_json`. | | `save` | `fn save(&self, path: impl AsRef) -> Result<()>` | Writes `to_json` output to a file. | @@ -52,9 +52,9 @@ A saved file therefore looks like: } ``` -## The mutation seam +## The install seam -Internally, both state loading and optimizers reach predictors through one type-erased trait, `DynPredictor` (crate-private). Its `apply_update` method is the single mutation seam: every write to a predictor's optimizable state flows through it, including optimizer candidate set and restore, `ModuleState::apply`, and the `fx::Params` overlay. An update is partial: `None` fields are left untouched, `instruction: Some(None)` clears the override back to the signature default, and `Some(demos)` replaces the demo set. `load_state` (used by `apply`) delegates to `apply_update` with both fields set. This is also the single place where prompt caches are invalidated: a candidate is data applied through the seam, never ad hoc field mutation. +State loading reaches predictors through the object-safe per-leaf view `PredictorInfo` (see [Modules](/docs/components/modules#predictor-discovery-predictors)). Its `load_state` method is the install seam: a **full** overwrite of the leaf's optimizable state (`instruction_override: None` clears the override, `demos` replaces the demo set), used by `ModuleState::apply` and by the optimizer's one-shot install of the winning candidate. Candidate *evaluation* never calls it — candidates are injected ambiently per call tree (see [Optimizers](/docs/components/optimizers)). Every `load_state` invalidates the leaf's cached instance overlay, so a candidate is data applied through the seam, never ad hoc field mutation. ## Compatibility behavior @@ -63,10 +63,9 @@ There is no version field in the format. Compatibility is field level and struct | Situation | Result | |---|---| | Missing `demos` or `instruction_override` in JSON | Defaults apply (`serde(default)` on both fields). | -| State names a path absent from the module | `apply` returns an error listing the unknown predictors. | +| State names a leaf absent from the module | `apply` returns an error listing the unknown predictors. | | Demo rows do not fit the predictor's signature schema | `apply` returns `failed to load state for `name``. | | Module has predictors the state does not name | Left untouched, no error. | -| `Predict` leaf inside `Rc` or `Arc` | Traversal error; `Box`, `Option`, `Vec`, arrays, slices, and string-keyed maps are supported. | ## See also diff --git a/docs/docs/components/tools-and-agents.mdx b/docs/docs/components/tools-and-agents.mdx index 8757d82e..8ee75916 100644 --- a/docs/docs/components/tools-and-agents.mdx +++ b/docs/docs/components/tools-and-agents.mdx @@ -52,7 +52,7 @@ The bodyless-fn rules match `#[predict]`: no `async`, no generics, no `self`, pl | Option | Form | Meaning | |---|---|---| -| `model` | `model = "@name"` | Model reference for the loop. | +| `model` | `model = "@name"` | Model reference for the loop. Binds only inside `#[module]` programs; setting it removes the standalone fn (calling one is a compile error). | | `tools` | `tools(a, b)` | The `#[tool]` functions the loop may call, in order. | | `stop_tools` | `stop_tools(a)` | Tools whose call ends the loop. Each must also appear in `tools(...)`. | | `max_turns` | `max_turns = N` | Turn bound for the lowered loop node (IR default is 8 when omitted). | @@ -80,7 +80,7 @@ When the model calls a stop tool, the loop ends right there. The arguments of th ## Standalone or inside a module -`#[agent]` generates a standalone `async fn` returning `Result, PredictError>`. Called directly, it runs the static-lane tool loop (`Predict` with `ToolLoopMode::Auto`). Inside a `#[module]` body, the same call lowers to a first-class `AgentLoop` node in the program graph, and the loop options (`max_turns`, `budget`, `context`) apply only to this lowered form. +`#[agent]` generates a standalone `async fn` returning `Result, PredictError>`. Called directly, it executes the same 1-node `AgentLoop` program the `#[module]` lowering produces, with the loop options honored on both paths: `max_turns`/`stop_tools`/`until_parse` land in the node's `StopSpec`, `budget` in its `NodeBudget`, and `context` in its `ContextPolicy`. The one exception is `model`: model refs bind only inside a `#[module]` program, so setting `model = "…"` removes the standalone fn — calling it is a compile error rather than a silent fallback to the globally configured LM. ```rust use dspy_rs::module; @@ -119,10 +119,18 @@ Because the code is inside the file, the program is fully portable: anyone who c ## Related surfaces -`ReAct` is the struct-lane equivalent of `#[agent]`: a thought, action, observation loop over a set of tools that extracts a typed answer, built and configured as a Rust struct rather than declared as a bodyless function. See [Modules](/docs/components/modules). +The struct-lane way to give a model tools is to attach them to a `Predict` (`PredictBuilder::add_tool`/`with_tools`): a tooled predictor executes as a 1-node `agent` program through the interpreter, with the default stop behavior (`until_parse`, `max_turns = 8`). See [Predict](/docs/components/predict). There is no separate `ReAct` module. Code Mode is the many-tools-to-one-script alternative. Instead of advertising N tool schemas and paying one round trip per call, the model sees a single `run_js` meta-tool whose description lists your tools as a JavaScript API; it writes one script that calls them as plain functions and returns one value. See [Code Mode](/docs/components/code-mode). +## Tool membership is optimizable + +Which tools a loop carries is a tuned value, not just structure. The agent node's `tools` list is the *declaration* — the loop's capability footprint, checked against the program ceiling at load. Which of those tools the loop actually presents to the model is the `ToolSet` parameter (`".tool_set"`), a slot like `instruction` or `demos`: an optimizer's candidate can drop a distracting tool or bring a declared one back, and the descriptions the model sees are themselves `ToolDesc` slots. The alphabet is closed — a candidate can never smuggle in a tool the declaration doesn't cover; that is refused at load, not at call time. Absent selection means the full declared list, so nothing changes until an optimizer says so. See [Program and nodes](/docs/components/program-and-nodes) for the slot machinery. + +## Executing tool calls yourself + +To execute tool calls yourself instead of letting the loop dispatch them (a REPL the agent drives, tools that need caller-side state), run the agent through the interpreter's caller-managed conversation surface: `Interpreter::run_conversation_caller_managed` suspends the loop on tool calls and `resume_conversation` feeds your results back, with the same spans, budgets, and stop-tool behavior as the dispatching loop. The suspended surface presents the ToolSet-selected tools per call, same as the dispatching loop. See [Runtime](/docs/components/runtime). + ## Common mistakes **Forgetting the body on a host tool.** `#[tool]` needs a real function with a body. A bodyless function is a step, not a tool. @@ -135,6 +143,6 @@ Code Mode is the many-tools-to-one-script alternative. Instead of advertising N - [Capabilities](/docs/components/capabilities) for tool needs, program ceilings, and host grants - [Code Mode](/docs/components/code-mode) for presenting many tools as one script surface -- [Modules](/docs/components/modules) for `ReAct` and the other struct-lane strategies +- [Modules](/docs/components/modules) for the struct-lane strategies - [The module macro](/docs/components/module-macro) for the body rules that lower agent calls - [The .dsrs file](/docs/components/dsrs-file) for sandboxed tool syntax in the artifact diff --git a/docs/docs/components/traces.mdx b/docs/docs/components/traces.mdx index aaac4780..04fabe37 100644 --- a/docs/docs/components/traces.mdx +++ b/docs/docs/components/traces.mdx @@ -1,10 +1,10 @@ --- title: "Traces" -description: "Capture a run as a trace, replay it strictly or until divergence, and export it to OpenTelemetry or RL rollout formats" +description: "Capture a run as a trace and replay it strictly or until divergence" icon: "wave-pulse" --- -A trace records what a run did: one span per `Predict` call, carrying the rendered prompt, the parsed output, and a request fingerprint. Replay serves a later run back from that recording instead of a live provider, and exports project a finished trace onto external observability and training conventions. The `dsrs` CLI that checks, formats, and serves `.dsrs` artifacts has its own page: [CLI](/docs/components/cli). +A trace records what a run did: one span per `Predict` call, carrying the rendered prompt, the parsed output, and a request fingerprint. Replay serves a later run back from that recording instead of a live provider. The `dsrs` CLI that checks, formats, and serves `.dsrs` artifacts has its own page: [CLI](/docs/components/cli). ## Recording and replaying a run @@ -65,6 +65,7 @@ One span is one `Predict` invocation. At a high level it records: - **What went in**: the rendered prompt (an interned system-and-demos prefix plus the live suffix), the typed input fields as JSON, and the redacted model config. - **What happened inside**: ordered events, one `Exchange` per provider round-trip and one `ToolRun` per tool execution. - **What came out**: the raw assistant text, the parsed output fields, aggregated token usage, and any error (kinds: `lm`, `parse`, `tool`, `cancelled`). +- **What it was worth** (optional): a span-level `Eval`, present only when a metric assigned per-span credit through `TypedMetric::evaluate_spans` (see [Evaluation](/docs/components/evaluation)). Demo harvesting prefers it over the whole-rollout score; the field is omitted from the JSONL entirely when absent, so eval-free traces serialize exactly as before. - **A request fingerprint**: `request_hash`, a stable hash over the redacted model config plus the full rendered prompt. This is the replay key and the determinism check. - **Timing and completeness**: start time, duration, and a `complete` flag (false when the span was truncated or redacted; replay refuses incomplete spans). @@ -143,81 +144,10 @@ let (out, report) = dspy_rs::trace::replay(&trace, ReplayMode::Strict, || pipeli Replay scoping mirrors capture: task-local, not inherited by spawned subtasks, innermost scope wins. Compose replay outside with capture inside to record a counterfactual rollout while serving its unchanged prefix from the base trace. -## Exports - -Exports are pure serialization-side projections of a finished `Trace`: no new capture machinery, no external dependencies. Both live under `dspy_rs::trace`. - - -Exports serialize traces after a run finishes. For live, human-readable console output while a program runs, call `init_tracing` (documented on [Utils](/docs/components/utils)): it installs a process-global pretty `tracing` subscriber (respecting `RUST_LOG`, defaulting to `dspy_rs=debug`) and is independent of trace capture. - - -### OpenTelemetry - -`Trace::to_otel_spans(include_content: bool)` maps a trace onto OpenTelemetry GenAI semantic conventions as plain serializable structs in the OTLP/JSON wire shape (proto3 JSON mapping: camelCase keys, 64-bit integers as decimal strings, ids as lowercase hex), with no OpenTelemetry dependency. `Trace::to_otlp_json(service_name, include_content)` wraps those spans in a complete `resourceSpans` envelope that any OTLP/HTTP collector (Jaeger, Tempo, otel-collector) accepts at `POST /v1/traces` as-is. - -| Trace format | OTel | -|---|---| -| `meta.trace_id` | Trace id: used verbatim when already 32 lowercase hex (the capture scope mints exactly this shape), otherwise stable-hashed into one. | -| The rollout | Root span `dsrs.rollout`, kind `INTERNAL`, carrying `dsrs.trace_id`, `dsrs.candidate_hash` (when set), and one `dsrs.tag.{key}` attribute per meta tag. | -| `Span` | One span per `Predict` invocation, name = component name, kind = `CLIENT`, parented to its recorded parent span, else the root. | -| `started_at_us` / `duration_us` | Start and end timestamps in nanoseconds. | -| Interned model config | `gen_ai.request.model`, `gen_ai.request.temperature`, `gen_ai.request.max_tokens`. | -| `usage.prompt_tokens` / `usage.completion_tokens` | `gen_ai.usage.input_tokens` / `gen_ai.usage.output_tokens`. | -| Rendered prompt (content opt-in) | One `gen_ai.prompt` event per prompt message (`gen_ai.prompt.role`, `gen_ai.prompt.content`). | -| `raw_output` (content opt-in) | One `gen_ai.completion` event (`gen_ai.completion.content`). | -| `SpanEvent::ToolRun` | Child span `tool:{name}`, kind `INTERNAL`, with `gen_ai.tool.name` and `gen_ai.tool.call.id`; arguments and result attributes are content opt-in. | -| `span.error` | Status `ERROR` with message `{kind}: {message}`. | -| `component` / `seq` / `request_hash` | `dsrs.component`, `dsrs.seq`, `dsrs.request_hash` attributes (plus `dsrs.candidate_hash` when set). | - -Prompt, completion, and tool payloads are opt-in via `include_content`, mirroring OTel's GenAI content-capture switch: with `include_content: false` the spans carry identity, usage, and timing attributes only, so they stay exportable to shared collectors without leaking prompt text. - -The emitted types, all `Serialize` structs: - -| Type | Shape | -|---|---| -| `OtelSpan` | `trace_id` (32 hex chars), `span_id` (16 hex chars), optional `parent_span_id`, `name`, `kind`, `start_time_unix_nano` and `end_time_unix_nano` (decimal strings), `attributes`, `events`, optional `status`. | -| `OtelKeyValue` | `key` plus an `OtelValue`. | -| `OtelValue` | The `AnyValue` oneof arms this export emits: `StringValue`, `IntValue` (a decimal string, per proto3 JSON), `DoubleValue`. | -| `OtelEvent` | `time_unix_nano`, `name`, `attributes`. | -| `OtelStatus` | `code`, `message` (omitted when empty). | - -The OTLP enum constants are exported as plain integers: `SPAN_KIND_INTERNAL = 1`, `SPAN_KIND_CLIENT = 3`, `STATUS_CODE_ERROR = 2`. - -`ToolRun` events record only their duration, so tool child spans start at their parent span's start time: durations are exact, offsets within the parent are not. +Traces serialize after a run finishes (`Trace::to_jsonl`). For live, human-readable console output while a program runs, call `init_tracing` (documented on [Utils](/docs/components/utils)): it installs a process-global pretty `tracing` subscriber (respecting `RUST_LOG`, defaulting to `dspy_rs=debug`) and is independent of trace capture. There are no built-in exporters — the JSONL trace format is stable and self-describing, so external projections (observability, RL datasets) are serialization-side work on top of it. -### RL dataset - -`Trace::to_rl_rollout()` projects the trace onto the Agent Lightning / verifiers rollout convention: one rollout as message lists plus a reward plus per-subcall transitions. It returns `None` when no eval was recorded (`trace.outcome.eval` unset); a rollout without a reward is not trainable. Because spans keep full `Message` structure (tool-call blocks, reasoning blocks), the projection needs no lossy text munging. - -```rust -let (result, mut trace) = capture(|| pipeline(input)).await; -trace.outcome = Some(TraceOutcome { eval: Some(metric_eval), ..Default::default() }); -let rollout = trace.to_rl_rollout().expect("eval recorded"); -writeln!(dataset, "{}", rollout.to_json_line()?)?; -``` - -`RlRollout` serializes to a single JSON object; `RlRollout::to_json_line()` produces the one-line JSONL record RL trainers consume. - -| `RlRollout` field | Meaning | -|---|---| -| `trace_id` | The trace id. | -| `reward` | The rollout-level reward: `trace.outcome.eval.score`. | -| `transitions` | One `RlTransition` per surviving span, in trace order. | -| `metadata` | The free-form run tags (`trace.meta.tags`). | - -| `RlTransition` field | Meaning | -|---|---| -| `component` | The optimizable unit this subcall belongs to; the per-agent credit assignment key. | -| `seq` | 0-based invocation index of the component within the rollout. | -| `messages` | The full rendered prompt as messages (interned prefix plus live suffix), provider-agnostic roles. | -| `completion` | Everything the policy emitted for this subcall, rebuilt from the span's events: each `Exchange`'s assistant message verbatim, with `ToolRun`s as intervening tool-result context. | -| `usage` | The span's aggregated `LmUsage` (prompt, completion, and total tokens). | -| `model` | The model identifier from the span's interned config. | - -Spans whose policy emitted nothing (provider failures, cancelled spans: no `Exchange` event) are omitted; they contribute no completion to train on. Parse-failure spans keep their transition: the emitted text exists even though it did not parse. - ## See also - [CLI](/docs/components/cli): `dsrs serve` returns the capture-scope trace artifact from `POST /run?trace=1` diff --git a/docs/docs/components/utils.mdx b/docs/docs/components/utils.mdx index df0402de..aafcbdf1 100644 --- a/docs/docs/components/utils.mdx +++ b/docs/docs/components/utils.mdx @@ -16,14 +16,14 @@ let id: u64 = stable_hash_debug(&value); ## `ResponseCache` -A hybrid memory plus disk LM response cache built on [foyer](https://docs.rs/foyer): 256MB in memory and 1GB on disk in a per-process temp directory. It also maintains a sliding window of the 100 most recent entries for `LM::inspect_history`. The cache is created automatically by `LM`; you do not construct it directly. Caching is per LM instance, and entries are not shared across instances. +A hybrid memory plus disk LM response cache built on [foyer](https://docs.rs/foyer): 256MB in memory and 1GB on disk in a per-process temp directory (if the disk tier cannot be initialized, the cache degrades to memory-only with a warning instead of panicking). It also maintains a sliding window of the 100 most recent entries for `LM::inspect_history`. The cache is created automatically by `LM`; you do not construct it directly. Caching is per LM instance, and entries are not shared across instances. All methods take `&self`: the foyer cache is internally synchronized and the history ring sits behind its own small mutex, so concurrent LM calls never serialize on a cache-wide lock. | Method | Signature | Behavior | |---|---|---| | `new` | `async fn new() -> Self` | Builds the hybrid cache and its disk tier. | | `get_entry` | `async fn get_entry(&self, key: CacheKey) -> Result>` | Fetches the full cached entry, including raw output. | -| `insert_entry` | `fn insert_entry(&mut self, key: CacheKey, entry: CacheEntry)` | Synchronous insert, the direct path used by `LM::call`. Also pushes the entry into the history window. | -| `get_history` | `async fn get_history(&self, n: usize) -> Result>` | Returns the `n` most recent entries, newest first. | +| `insert_entry` | `fn insert_entry(&self, key: CacheKey, entry: CacheEntry)` | Synchronous insert, the direct path used by `LM::call`. Also pushes the entry into the history window. | +| `get_history` | `fn get_history(&self, n: usize) -> Vec` | Returns the `n` most recent entries, newest first. | ### `CacheEntry` and `CacheKey` diff --git a/docs/docs/getting-started/quickstart.mdx b/docs/docs/getting-started/quickstart.mdx index 98bd823e..4b857684 100644 --- a/docs/docs/getting-started/quickstart.mdx +++ b/docs/docs/getting-started/quickstart.mdx @@ -156,7 +156,7 @@ Every attribute, every supported type, constraints Builder surface, demos, metadata, errors -ChainOfThought, ReAct, composition +ChainOfThought, predictor discovery, composition diff --git a/docs/docs/guides/examples.mdx b/docs/docs/guides/examples.mdx index d39d207e..9300448c 100644 --- a/docs/docs/guides/examples.mdx +++ b/docs/docs/guides/examples.mdx @@ -17,7 +17,7 @@ Some examples need a feature flag (noted in the file header), for example `--fea | Example | What it shows | Key APIs | |---------|---------------|----------| | [01-simple.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/01-simple.rs) | Typed signatures, chain-of-thought via a `reasoning` field, and module composition | `Signature`, `Predict`, `Module` | -| [02-module-iteration-and-updation.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/02-module-iteration-and-updation.rs) | Optimizing a module end to end with the typed optimizer API | `COPRO`, `Optimizer::compile`, `TypedMetric` | +| [02-module-iteration-and-updation.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/02-module-iteration-and-updation.rs) | Optimizing a module end to end with the typed optimizer API | `COPRO`, `compile_module`, `TypedMetric` | | [03-evaluate-hotpotqa.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/03-evaluate-hotpotqa.rs) | Evaluating a typed QA predictor on a HotpotQA sample | `DataLoader`, `evaluate_trainset`, `Eval` | | [04-optimize-hotpotqa.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/04-optimize-hotpotqa.rs) | COPRO optimization of a QA module on HotpotQA, then saving the result | `COPRO`, `ModuleState`, `average_score` | | [05-heterogenous-examples.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/05-heterogenous-examples.rs) | Feeding a typed predictor from messy JSON rows; extra fields ignored, missing ones loud | serde boundary, generated input structs | @@ -36,12 +36,11 @@ Some examples need a feature flag (noted in the file header), for example `--fea ## The front desk examples -These nine exercise a small support-desk pipeline end to end. Each maps to a component page and is kept compiling by CI. +These exercise a small support-desk pipeline end to end. Each maps to a component page and is kept compiling by CI. | Example | Component page | What it shows | |---------|---------------------|---------------| | [18-code-mode.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/18-code-mode.rs) | [Tools & agents](/docs/components/tools-and-agents) | Two tools collapsed into the `run_js` meta-tool via `ToolSet::code_mode` | -| [19-react.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/19-react.rs) | [Tools & agents](/docs/components/tools-and-agents) | `ReAct` with a closure tool and a step cap, trajectory read from metadata | | [20-frontdesk-contract.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/20-frontdesk-contract.rs) | [Signatures](/docs/components/signatures) | Enums, nested types, constraints, demos, per-predictor models, the printed prompt | | [21-frontdesk-pipeline.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/21-frontdesk-pipeline.rs) | [Modules](/docs/components/modules) | The pipeline as a struct with a `Module` impl, plus the `fx` variant | | [22-frontdesk-module.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/22-frontdesk-module.rs) | [Module macro](/docs/components/module-macro) | `#[module]` with a hole; prints the real `.dsrs` artifact and `OPACITY` (offline) | @@ -51,5 +50,5 @@ These nine exercise a small support-desk pipeline end to end. Each maps to a com | [26-frontdesk-tune.rs](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/26-frontdesk-tune.rs) | [Optimizers](/docs/components/optimizers) | Overlay by param path, ambient overlay run, `bake` with lineage (offline path included) | -The `9x-` files ([90](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/90-smoke-slice1-typed-predict.rs) through [94](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/94-smoke-slice5-optimizer-interface.rs), plus the `97`-`99` benches) are CI surfaces: smoke slices that exercise typed predict, chain of thought, module authoring, ReAct, and the optimizer interface, and a few performance benchmarks. They are kept deliberately boring and stable; read them as compatibility contracts rather than tutorials. +The `9x-` files ([90](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/90-smoke-slice1-typed-predict.rs) through [94](https://github.com/krypticmouse/DSRs/blob/main/crates/dspy-rs/examples/94-smoke-slice5-optimizer-interface.rs), plus the `97`-`99` benches) are CI surfaces: smoke slices that exercise typed predict, chain of thought, module authoring, and the optimizer interface, and a few performance benchmarks. They are kept deliberately boring and stable; read them as compatibility contracts rather than tutorials. diff --git a/docs/docs/optimizers/copro.mdx b/docs/docs/optimizers/copro.mdx index 52332281..6b7435c4 100644 --- a/docs/docs/optimizers/copro.mdx +++ b/docs/docs/optimizers/copro.mdx @@ -36,7 +36,7 @@ Other fields: `prompt_model` (separate LM for generating candidate instructions) use anyhow::Result; use bon::Builder; use dspy_rs::{ - COPRO, Eval, LM, Module, Optimizer, Predict, PredictError, + COPRO, Eval, LM, Module, Predict, PredictError, Predicted, Signature, Trace, TypedMetric, configure, init_tracing, }; @@ -49,13 +49,14 @@ struct QA { answer: String, } -#[derive(Builder, facet::Facet)] -#[facet(crate = facet)] +#[derive(Builder)] struct MyModule { #[builder(default = Predict::::new())] predictor: Predict, } +dspy_rs::predictors!(MyModule { predictor }); + impl Module for MyModule { type Input = QAInput; type Output = QAOutput; @@ -118,13 +119,13 @@ async fn main() -> Result<()> { .build(); let metric = ExactMatchMetric; - copro.compile(&mut module, trainset, &metric).await?; + copro.compile_module(&mut module, &trainset, &metric).await?; Ok(()) } ``` -The trainset here is a `Vec` of `(QAInput, QAOutput)` tuples, the zero-boilerplate row form. Any row type that implements `ToInput` works, including `#[derive(Example)]` structs carrying gold labels and metric-only fields; see [Data](/docs/components/data). +The trainset here is a slice of `(QAInput, QAOutput)` tuples, the zero-boilerplate row form. Any row type that implements `ToInput` works, including `#[derive(Example)]` structs carrying gold labels and metric-only fields; see [Data](/docs/components/data). The `predictors!` line names the module's optimizable leaves — without it the module does not satisfy the `Predictors` bound `compile_module` requires. ### Typed data loading @@ -163,11 +164,11 @@ Recommended: 2-5 An optional separate LM used to generate candidate instructions. Falls back to the global LM set with `configure` when unset. ### Track stats -A per-round statistics flag. COPRO's report type is `()`, so nothing is returned either way; leave it at the default. +A per-round statistics flag. COPRO's `compile_module` returns `()`, so nothing is returned either way; leave it at the default. ## Implementation notes -COPRO is a thin strategy over the shared evaluation engine: each candidate instruction is an overlay evaluated through a cached, bounded-concurrency fan-out, and each round's winner is installed through the same mutation seam every optimizer uses. Repeated candidates are deduplicated by content hash and served from the rollout cache, so re-evaluating an instruction it has already seen costs nothing. See the [optimizer engine reference](/docs/components/optimizer-engine) for the machinery. +COPRO is a thin strategy over the shared evaluation engine: each candidate instruction is a name-keyed `Candidate` injected ambiently per rollout through a cached, bounded-concurrency fan-out — the module is never mutated during evaluation — and the final winner is installed once through `OptimizeTarget::install`. Repeated candidates are deduplicated by content hash and served from the rollout cache, so re-evaluating an instruction it has already seen costs nothing. See the [optimizer engine reference](/docs/components/optimizer-engine) for the machinery. Cost is approximately `breadth × depth × num_predictors × trainset_size` LM calls, minus cache hits. diff --git a/docs/docs/optimizers/gepa.mdx b/docs/docs/optimizers/gepa.mdx index 10561ac9..80d57874 100644 --- a/docs/docs/optimizers/gepa.mdx +++ b/docs/docs/optimizers/gepa.mdx @@ -54,12 +54,13 @@ GEPA can optimize at test time, not just training time (see [Inference-time sear ```rust use dspy_rs::*; -#[derive(Builder, facet::Facet)] -#[facet(crate = facet)] +#[derive(Builder)] struct MyModule { predictor: Predict, } +dspy_rs::predictors!(MyModule { predictor }); + impl Module for MyModule { type Input = MySignatureInput; type Output = MySignatureOutput; @@ -119,71 +120,12 @@ let gepa = GEPA::builder() .max_rollouts(500) // Budget control .build(); -let result = gepa.compile(&mut module, trainset, &metric).await?; +let result = gepa.compile_module(&mut module, &trainset, &metric).await?; println!("Best score: {:.3}", result.best_candidate.average_score()); println!("Best instruction: {}", result.best_candidate.instruction); ``` -## Feedback helpers - -DSRs provides utilities for common feedback patterns. Each returns a ready-made `Eval` with the score and a formatted feedback string: - -### Document retrieval -```rust -use dspy_rs::retrieval_feedback; - -let eval = retrieval_feedback( - &retrieved_docs, - &expected_docs, - Some(&all_available_docs) -); - -// Feedback: -// Retrieved 3/5 correct documents (Precision: 0.600, Recall: 0.600, F1: 0.600) -// Correctly retrieved: doc1, doc2, doc3 -// Missed: doc4, doc5 -``` - -### Code generation -```rust -use dspy_rs::{code_pipeline_feedback, CodeStage, StageResult}; - -let stages = vec![ - (CodeStage::Parse, StageResult::Success), - (CodeStage::Compile, StageResult::Success), - (CodeStage::Execute, StageResult::Failure { - error: "Division by zero on line 42".to_string(), - }), -]; - -let eval = code_pipeline_feedback(&stages, 0.6); - -// Feedback: -// Parse: Success -// Compile: Success -// Execute: Division by zero on line 42 -``` - -### Multi-objective optimization -```rust -use dspy_rs::multi_objective_feedback; - -let mut objectives = HashMap::new(); -objectives.insert("accuracy".to_string(), (0.9, "High accuracy".to_string())); -objectives.insert("latency".to_string(), (0.7, "Acceptable latency".to_string())); -objectives.insert("cost".to_string(), (0.8, "Within budget".to_string())); - -let mut weights = HashMap::new(); -weights.insert("accuracy".to_string(), 2.0); -weights.insert("latency".to_string(), 1.0); -weights.insert("cost".to_string(), 1.0); - -let eval = multi_objective_feedback(&objectives, Some(&weights)); -``` - -`string_similarity_feedback` and `classification_feedback` are also available. - ## Configuration options ```rust @@ -206,16 +148,16 @@ Pass validation data at compile time: ```rust let result = gepa - .compile_with_valset(&mut module, trainset, Some(valset), &metric) + .compile_module_with_valset(&mut module, &trainset, Some(&valset), &metric) .await?; ``` -With `Some(valset)`, initial evaluation and child scoring use the validation set while parent re-evaluation uses trainset minibatches; with `None`, the trainset serves both roles (this is what `compile` does). +With `Some(valset)`, initial evaluation and child scoring use the validation set while parent re-evaluation uses trainset minibatches; with `None`, the trainset serves both roles (this is what `compile_module` does). ## Understanding GEPA results ```rust -let result = gepa.compile(&mut module, trainset, &metric).await?; +let result = gepa.compile_module(&mut module, &trainset, &metric).await?; // Best candidate found (already installed on the module) println!("Best instruction: {}", result.best_candidate.instruction); @@ -250,10 +192,9 @@ pub struct Eval { Every rollout runs under trace capture, recording one span per `Predict` invocation. GEPA feeds the mutated component's spans (`trace.for_component(name)`) to the reflection LM alongside the feedback. See [Traces](/docs/components/traces). -**ParetoFrontier** -- Each candidate tracks which examples it wins on -- Sampling is proportional to coverage -- Automatically prunes dominated candidates +**Pareto bookkeeping** + +GEPA uses the engine's score matrix directly: `ParetoView` (see [Optimizer engine](/docs/components/optimizer-engine)) tracks which validation columns each candidate wins on, parent sampling is proportional to that coverage, and candidates with zero wins are dominated. **GEPACandidate** ```rust @@ -414,12 +355,13 @@ struct MathJudge { ### Optimized module ```rust -#[derive(Builder, facet::Facet)] -#[facet(crate = facet)] +#[derive(Builder)] struct MathSolver { #[builder(default = Predict::::new())] solver: Predict, // This gets optimized } + +dspy_rs::predictors!(MathSolver { solver }); ``` ### TypedMetric with judge @@ -660,7 +602,7 @@ let gepa = GEPA::builder() .build(); let result = gepa - .compile_with_valset(&mut module, my_tasks.clone(), Some(my_tasks), &metric) + .compile_module_with_valset(&mut module, &my_tasks, Some(&my_tasks), &metric) .await?; // Access per-task best scores and outputs diff --git a/docs/docs/optimizers/miprov2.mdx b/docs/docs/optimizers/miprov2.mdx index 30bd8f8c..29deefa1 100644 --- a/docs/docs/optimizers/miprov2.mdx +++ b/docs/docs/optimizers/miprov2.mdx @@ -18,7 +18,7 @@ One traced teacher pass over the trainset. Every example runs through your modul ### Phase 2: Demo bootstrapping -Successful spans from rollouts scoring at least `min_demo_score` become few-shot demos on the predictor that produced them (top `max_bootstrapped_demos` by score, deduplicated on inputs). Demos are installed before instruction search, so candidates are scored against the module as it will actually run. +Successful spans scoring at least `min_demo_score` become few-shot demos on the predictor that produced them (top `max_bootstrapped_demos` by score, deduplicated on inputs). A span scores as its rollout does, unless the metric attached a span-level eval via `TypedMetric::evaluate_spans` — that score then takes precedence (see [Evaluation](/docs/components/evaluation#per-span-credit)). Demos are installed before instruction search, so candidates are scored against the module as it will actually run. ### Phase 3: Candidate generation @@ -33,10 +33,10 @@ Uses the traces and a rotation of prompting tips to generate `num_candidates` in ### Phase 4: Trial evaluation -- Evaluates up to `num_trials` candidates per predictor on one sampled minibatch +- Evaluates up to `num_trials` candidates per predictor on one sampled minibatch (candidates injected ambiently — the module is never touched during evaluation) - Computes performance scores - Selects the best performing candidate -- Applies it to the module +- Installs the accumulated winner (demos + best instructions) once at the end through `OptimizeTarget::install` ## Configuration @@ -47,7 +47,7 @@ let optimizer = MIPROv2::builder() .num_trials(20) // Max candidates evaluated per predictor .minibatch_size(25) // Examples per candidate evaluation .max_bootstrapped_demos(4) // Demos installed per predictor - .min_demo_score(0.0) // Score gate for demo-eligible traces + .min_demo_score(0.0) // Score gate for demo-eligible spans .build(); ``` @@ -56,7 +56,7 @@ You can also set `eval_concurrency` (concurrent LM calls during evaluation, defa ## Usage example ```rust -use dspy_rs::{MIPROv2, Optimizer}; +use dspy_rs::MIPROv2; // Create optimizer let optimizer = MIPROv2::builder() @@ -68,13 +68,13 @@ let optimizer = MIPROv2::builder() // Typed metric implementing TypedMetric for your trainset row type let metric = ExactMatchMetric; -// Optimize your module -optimizer.compile(&mut module, train_examples, &metric).await?; +// Optimize your module (MyModule declares its leaves via predictors!) +optimizer.compile_module(&mut module, &train_examples, &metric).await?; ``` The metric is the same `TypedMetric` used by every optimizer: `evaluate(&self, example, prediction, trace) -> Result`, where `example` is your full trainset row. MIPROv2 only reads the numerical score; feedback is ignored. See the [evaluation reference](/docs/components/evaluation) for the trait. -`train_examples` is `Vec` for any row type implementing `ToInput` toward the module's input — a `#[derive(Example)]` struct or `(Input, Output)` tuples; see [Data](/docs/components/data). +`train_examples` is a slice of any row type implementing `ToInput` toward the module's input — a `#[derive(Example)]` struct or `(Input, Output)` tuples; see [Data](/docs/components/data). The module must declare its leaves with `predictors!`; see [Modules](/docs/components/modules#predictor-discovery-predictors). ### Typed data loading @@ -109,16 +109,15 @@ Use the shared data ingress reference: [`DataLoader`](/docs/components/data). The code follows standard Rust practices: - No unsafe blocks - Results for error handling with context via anyhow -- Strong types (`PromptCandidate`, `PromptingTips`) +- Strong types (`Candidate`, `PromptingTips`) - Builder pattern for configuration - Async throughout, no blocking calls -Key public types and methods: -- `PromptCandidate` - an instruction with its evaluated `score: f64` +Key public types: - `PromptingTips` - the library of best practices (`default_tips()`, `format_for_prompt()`) -- `MIPROv2::select_best_traces`, `create_prompt_candidates`, `format_schema_fields` +- `Candidate` - the shared engine currency MIPROv2 registers its instruction variants as (see [Optimizer engine](/docs/components/optimizer-engine)) -The report type is `()`: MIPROv2 mutates the module in place and returns nothing. +`compile_module` returns `()`: MIPROv2 installs the winner on the module and reports nothing further (`Report::None` through the trait). Cost is roughly `num_predictors × (trainset_size + num_trials × minibatch_size)` LM calls, minus rollout-cache hits. diff --git a/docs/docs/optimizers/structural.mdx b/docs/docs/optimizers/structural.mdx new file mode 100644 index 00000000..ca626bbc --- /dev/null +++ b/docs/docs/optimizers/structural.mdx @@ -0,0 +1,212 @@ +--- +title: 'Structural' +description: 'LM-guided graph edits over the edit calculus: config, usage, and when to use it' +icon: 'diagram-project' +--- + +import OptimizerComparison from '/snippets/optimizer-comparison.mdx'; + +**Structural** is the structural optimizer: where the other five strategies tune parameter values (instructions, demos) through overlays, Structural rewrites the program graph itself. Each generation it gathers the [`legal_edits`](/docs/components/edit-calculus#legal_edits-the-proposer-menu) menu, has a reflection LM choose one edit from the serialized menu plus the incumbent's evaluation feedback, applies it with `Program::edited`, carries the tuned overlay across the change with `migrate_overlay`, and keeps the child only if it beats the parent on a shared minibatch. + +Structural runs on the **program lane only**: it needs an interpreter-loaded [`Program`](/docs/components/program-and-nodes) whose skeleton is data. Typed modules have no editable skeleton, so there is no `compile_module` here. + +## Overview + +The [edit calculus](/docs/components/edit-calculus) makes structural mutation safe: edits are serde values, `Program::edited` is pure and re-validates, and `migrate_overlay` re-mints tuned slot values against the child. Structural is the search loop on top: a GEPA-style reflection step chooses *which* edit to try, and the engine's minibatch gate decides whether to keep the result. The moves it can propose: + +| Move | Effect | +|------|--------| +| `AugmentSig` | Prepend the chain-of-thought `reasoning` output field to a leaf's signature (the CoT move). | +| `SwapToAgent` / `SwapToPredict` | Swap a `predict` leaf into a tool-using `agent` loop over the program's declared tools, or back. | +| `WrapRetry` | Wrap a node in a `Retry` (2 attempts, feedback on). | +| `Remove` | Remove a step from its `seq`. | +| `AddTool { tool }` / `RemoveTool { tool }` | Declare or undeclare a program tool on an agent leaf. | + +`SetStop` and `SetInstructionDefault` appear in `legal_edits` but are excluded from Structural's menu: they need free-form values, which is value-level work the other optimizers already own. + +## Quick start + +### 1. Load a program and implement a `ProgramMetric` + +```rust +use dspy_rs::ir::{DemoRow, Interpreter, Program, RuntimeEnv}; +use dspy_rs::trace::JsonMap; +use dspy_rs::{Eval, ProgramMetric, Trace}; + +let program = Program::load_dsrs("qa.dsrs")?; +let interp = Interpreter::load(program, RuntimeEnv::new()).await?; + +struct ExactMatch; + +impl ProgramMetric for ExactMatch { + async fn evaluate( + &self, + example: &DemoRow, + output: &JsonMap, + _trace: Option<&Trace>, + ) -> anyhow::Result { + let correct = output.get("answer") == example.output.get("answer"); + Ok(Eval::with_feedback( + correct as u8 as f64, + if correct { "correct".into() } else { format!("expected {:?}", example.output.get("answer")) }, + )) + } +} +``` + +Feedback is optional for Structural, but whatever the metric returns is what the reflection LM reads when choosing an edit, so specific feedback buys better proposals. + +### 2. Configure and run + +```rust +use dspy_rs::Structural; + +let reflection_lm = LM::builder().model("openai:gpt-4o".to_string()).build().await?; + +let structural = Structural::builder() + .num_iterations(8) + .minibatch_size(8) + .prompt_model(reflection_lm) + .max_rollouts(400) // every child is a fresh program: cap the spend + .seed(42) + .build(); + +let report = structural + .compile_program(&interp, &examples, &ExactMatch, || { + // A fresh RuntimeEnv per child load: the same model/tool/sandbox + // bindings the incumbent was loaded with. + RuntimeEnv::new() + }) + .await?; + +println!("{:.3} -> {:.3}", report.baseline_score, report.final_score); +``` + +The closure argument supplies a fresh [`RuntimeEnv`](/docs/components/runtime) every time an edited child needs loading. Only the host knows the live bindings (models, host tools, sandbox, capability grants), so child loading cannot be implicit; return the same bindings you loaded the incumbent with. + +### 3. Keep the winner + +The winner is returned, never installed: the interpreter you passed in is untouched. Bake the migrated overlay into the winning program to get a single self-contained artifact: + +```rust +use dspy_rs::ir::Lineage; + +let baked = report.program.bake(&report.overlay, Lineage::default())?; +baked.save_dsrs("qa-structural.dsrs")?; +``` + +If you ran a value-level optimizer first (GEPA, COPRO), pass its winning overlay in and Structural carries it across every accepted edit: + +```rust +let report = structural + .compile_program_with_overlay(&interp, Some(tuned_overlay), &examples, &ExactMatch, env) + .await?; +assert_eq!(report.overlay.base, report.program.meta.program_hash); +``` + +## Configuration options + +```rust +Structural::builder() + .num_iterations(8) // Structural generations to attempt + .minibatch_size(8) // Shared minibatch for the parent/child gate + .prompt_model(reflection_lm) // Reflection LM that chooses edits (recommended) + .max_rollouts(400) // Budget: max evaluation rollouts + .max_lm_calls(500) // Budget: max LM calls (rollouts + reflection) + .eval_concurrency(16) // Rollouts in flight during evaluation + .seed(42) // Reproducible sampling and fallback choice + .build() +``` + +| Field | Type | Default | Description | +|-------|------|---------|-------------| +| `num_iterations` | `usize` | `8` | Generations to attempt; each proposes exactly one edit. | +| `minibatch_size` | `usize` | `8` | Examples in the shared minibatch parent and child are compared on. | +| `prompt_model` | `Option` | `None` | Reflection LM that chooses an edit from the menu. Without it the choice is a seeded-uniform pick. | +| `max_rollouts` | `Option` | `None` | Hard cap on evaluation rollouts. | +| `max_lm_calls` | `Option` | `None` | Hard cap on LM call units (rollouts plus reflection). | +| `eval_concurrency` | `usize` | `16` | Concurrent rollouts during evaluation. | +| `seed` | `Option` | `None` | Fixes minibatch sampling and the fallback edit choice. | + +## Understanding Structural results + +`compile_program` returns a `StructuralReport`: + +| Field | Type | Description | +|-------|------|-------------| +| `program` | `Arc` | The winning program (the input program when nothing was accepted). | +| `overlay` | `Overlay` | The incumbent overlay re-minted against the winner at every accepted edit. | +| `baseline_score` | `f64` | Mean metric score of the input program (plus overlay) over the examples. | +| `final_score` | `f64` | Full-set mean of the final program; equals the baseline when nothing was accepted. | +| `edits` | `Vec` | The accepted edits, in order. Node ids are handles against each step's parent (`parent_hash`), so this is a lineage, not one batch. | +| `steps` | `Vec` | Per-generation outcomes, in order. | +| `accepted` / `rejected` | `usize` | Generations promoted / not promoted by the gate. | +| `spend` | `Spend` | Engine spend for the whole run, reflection calls included. | + +Each `StructuralStep` records `generation`, the targeted `leaf`, the concrete `edit` (serde data, replayable against `parent_hash`), `parent_minibatch_score`, `child_minibatch_score` (`None` when the child never scored), `accepted`, `full_score` (`Some` only when accepted), and `rejection` (why a child never scored, when it didn't). + +```rust +for step in &report.steps { + println!( + "gen {}: {:?} on `{}` — parent {:.2}, child {:?}, accepted: {}", + step.generation, step.edit, step.leaf, + step.parent_minibatch_score, step.child_minibatch_score, step.accepted, + ); +} +``` + +## The loop + +1. **Baseline** the incumbent (plus overlay) over the full example set. This seeds the engine's rollout cache, so every later parent minibatch read costs nothing. +2. Each generation: + - **Sample** a shared minibatch (seeded RNG). The incumbent's minibatch mean is the gate threshold. + - **Menu**: gather `legal_edits` for every leaf; keep the materializable kinds. + - **Choose**: the reflection LM reads the program's canonical `.dsrs` text, the menu (one JSON object per line, each with an `option` number), and the incumbent's per-example feedback, and answers with one option number. + - **Apply**: `Program::edited` mints the child; `migrate_overlay` re-mints the incumbent overlay against it; the child loads through your `RuntimeEnv` factory. + - **Gate**: the child is scored on the same minibatch. Only a strict win promotes it to a full-set evaluation and makes it the new incumbent. +3. **Return** the incumbent program and its overlay. + +Every rejection path degrades gracefully: an edit that fails to apply (`EditError`), a child that fails validation or loading, or a reflection reply that does not parse is recorded in the step and skipped. A run only errors on the engine's own failure modes (a metric error, an LM error during evaluation, a budget too small for the baseline pass). + +## Cost model + +Every child is a fresh program: its hash keys fresh rollout-cache rows, so nothing it does is served from the parent's cache. Per run: + +- `examples.len()` rollouts for the baseline pass; +- per generation, `minibatch_size` rollouts for the gate, plus the remaining `examples.len() - minibatch_size` only on promotion, plus one reflection call when a `prompt_model` is set. + +Cap the spend with `max_rollouts` / `max_lm_calls`; the run stops cleanly when the next batch would not fit. + +## When to use Structural + +- The program's *shape* is the bottleneck: a leaf that should reason step by step, a step that should be an agent with tools (or should not be), a flaky node that needs a retry. +- After a value-level pass: tune instructions first, then let Structural search structure while `migrate_overlay` preserves the tuned text. +- You have a labeled example set and budget for whole-program re-evaluation. + +Prefer the value-level optimizers when instructions and demos are the lever: they are cheaper (cache-friendly, no program reloads) and search a denser space. + +## Troubleshooting + +### The run accepts nothing + +The gate requires a strict minibatch win. Small minibatches are noisy; raise `minibatch_size` for a better signal, and set a `prompt_model` so choices are informed rather than uniform. + +### Steps show `rejection: Some("edit failed: ...")` + +Normal. The menu is structural, and data-flow legality is the validator's call: removing a step whose outputs a later binding still references, for example, is refused by `Program::edited` and skipped. See [the edit calculus](/docs/components/edit-calculus#errors). + +### The run errors with "budget too small for the baseline pass" + +The baseline needs `examples.len()` rollouts before the loop can start. Raise `max_rollouts` or shrink the example set. + +## Comparison with other optimizers + + + +## See also + +- [The edit calculus](/docs/components/edit-calculus): `Edit`, `Program::edited`, `legal_edits`, `migrate_overlay` +- [Optimizers](/docs/components/optimizers): the full configuration and report tables +- [Optimizer engine](/docs/components/optimizer-engine): the shared evaluation core Structural gates through +- [Runtime](/docs/components/runtime): `Interpreter::load` and `RuntimeEnv`, which child loading goes through +- [Program and nodes](/docs/components/program-and-nodes): `Overlay` and `Program::bake` for keeping the winner diff --git a/docs/handoff-v1-unification.md b/docs/handoff-v1-unification.md new file mode 100644 index 00000000..70ff4ef8 --- /dev/null +++ b/docs/handoff-v1-unification.md @@ -0,0 +1,315 @@ +# Handoff — `v1-program-unification` branch + +Written 2026-08-20 for whoever takes this branch to `main`. This is the working +ledger of what was done, what is *verified*, what is *claimed but unverified*, +and what the next owner must do before trusting or merging it. + +## 1. Branch state + +- Branch: `v1-program-unification`, off `main` @ `4435b76`. +- Net diff vs main: **203 files, +12,131 / −8,592** (before the seam merges; + the three seam merges add ~+3,300 more). +- Structure: nine phase merges + three seam merges. Read it top-down with + `git log --oneline --merges v1-program-unification ^main`. + +| Merge | What | Author session | +|---|---|---| +| Wave 1a–1e | Dead-code kill list, sandbox hardening, `dsrs-syntax` crate, `ir/edit.rs` calculus, workspace hygiene | session A (this ledger's author) | +| Phase 2α/2β | ReAct/ModuleExt deleted; interpreter leaf metadata; **Predict demoted onto the IR Interpreter**; adapter static lane deleted | session A | +| Phase 3 | Object-safe `Optimizer`, unified `Engine`, ambient candidate injection, `DynPredictor`/facet-walker/mutation model deleted, **facet fork unpinned** | session A | +| Phase 4a/4b | Feature collapse (+ `data` feature), fx CLOCK cache, `prelude`, `#[agent]` options honored, full docs sweep, RFC 0004 seams ledger | session A | +| Seam 6 (`ca14e08`) | `Structural` — sixth optimizer, LM-guided edit proposals over `ir::Edit` | **session B — NOT reviewed by session A** | +| Seam 4 (`109c571`) | Per-span credit assignment in demo harvesting | **session B — NOT reviewed by session A** | +| Seam 5 (`16f2b27`) | Tool membership as a `ParamSlot` (`ParamKind::ToolSet`, touches `.dsrs` parse/print) | **session B — NOT reviewed by session A** | +| Seams 1+2 (`f7d67a6`) | Interpreter conversation surface (`run_conversation`) + caller-managed suspend/resume; `Predict` compat shims deleted | **session B — NOT reviewed by session A** | + +## 2. What is verified, and how + +- **Phases 1–4**: full workspace suite green after every merge (79 suites; the + count moved 80→77→79 as test files were deleted/added — each delta was + accounted for). `cargo check -p dspy-rs --no-default-features` compiles; + 35 lib tests pass without default features. +- **Prompt stability**: golden prompt tests byte-identical through the adapter + collapse (they now exercise the single `*_def` lane). +- **`.dsrs` hash stability through the `dsrs-syntax` extraction**: fixture + programs printed + hashed before/after — byte-identical (verified in wave 1c). +- **Sandbox fixes**: 7 regression tests (injection rejected, `JSON`-named tool + safe, capability timeout fires, deregister evicts bytecode, current-thread + runtime works). +- **Soundness**: `grep -rn "unsafe" crates/dspy-rs/src` → zero. The facet + `[patch.crates-io]` fork pin is gone; upstream facet 0.43 builds and tests. +- **Seams 4/5/6 (session B)**: the suite at HEAD was re-run by session A after + discovering these merges — result recorded below in §3 item 0. Beyond that, + session A has only *skimmed* this code. It has not been reviewed. + +## 3. Validation checklist for the next owner (priority order) + +0. **Confirm HEAD is green.** `cargo test --workspace` at `16f2b27`. + Session A's result (2026-08-20): exit 0, 80/80 suites ok, zero failures — + the seam merges pass on top of the phases. Cheap to re-run; do so anyway. +1. **LIVE LM smoke test — the single biggest gap.** Every test in every phase + ran against mocked completions (`TestCompletionModel`). The demoted + `Predict` path, the AgentLoop-backed `with_tools` path, and the + `Structural` optimizer have **never talked to a real provider** on this + branch. Run with a real key: examples 01–05, 22 (frontdesk module), + 18 (code-mode), 09 (GEPA, small budget), and a small `Structural` run. +2. **Review the three seam merges** (`ca14e08`, `109c571`, `16f2b27`) — they + are unreviewed by anyone except their authoring session. Specifically: + - Seam 5 changed `ir/text/parse.rs` + `print.rs` (the ToolSet gene). The + canonical text is the **program-hash preimage**: confirm whether the + grammar change invalidates pre-existing `.dsrs` artifacts and whether + that was deliberate (pre-1.0 it's acceptable, but it must be a decision, + not an accident). Check `test_ir_text` / `test_ir_bake` diffs in those + merges for regenerated expectations — regenerated goldens are a red flag + to inspect, not a proof of correctness. + - Seam 4 changes demo quality for Bootstrap/MIPRO/SIMBA (per-span credit + instead of whole-rollout). That is a *behavioral* change to optimizer + output — eyeball a before/after demo harvest on a real trainset. + - Seam 6 (`optimizer/structural.rs`): review the accept/reject gate, the + menu serialization the reflection LM sees, and `migrate_overlay` usage + across accepted edits. This is the flagship feature; it deserves the + most careful read. +3. **Performance before/after.** Never measured. Run + `examples/97-perf-microbench` (and 98/99 orchestration benches) on `main` + vs this branch. The demoted Predict adds program-cache lookup + serde + round-trip + error translation per call; the claim that this is noise + against network latency is *plausible, not proven* — and it is NOT noise + for cache-hit or replay-served calls, which skip the network entirely. +4. **The task-local footgun.** Ambient candidates (`fx::with_params` / + `with_ambient_overlay`) do not propagate into `tokio::spawn`ed tasks. A + module whose `forward` spawns tasks silently evaluates the *baseline* + during optimization. Write the demonstrating test, then decide: document + loudly, detect-and-warn (e.g. a capture-scope generation counter), or fix + (explicit context handle instead of task-locals). Until then any + user-written module with internal spawns gets silently wrong optimization. +5. **Publishability.** `cargo publish --dry-run` for `dsrs-syntax`, + `dsrs-tools`, `dsrs_macros`, `dspy-rs` (in that order). The fork pin is + gone but rig-core + minijinja are still git dependencies — crates.io will + reject those; decide the strategy (vendor, fork-publish, or wait). +6. **Docs build.** `docs/` was fully swept and API pages regenerated + (`docs/scripts/gen_api.py` — script-generated, do not hand-edit), but + nobody ran the Mintlify build/link check. Also verify session B's + `docs/docs/optimizers/structural.mdx` against the actual implementation, + and note RFC 0004 now has status annotations added by session A. +7. **Clippy + version bumps.** `cargo clippy --workspace --all-targets`; + versions are skewed (dspy-rs 0.7.3 / macros 0.7.2 / tools & cli & syntax + 0.1.0) — pick a coherent 0.8.0 story and write a CHANGELOG before the PR. + +## 4. Known sharp edges (deliberate trade-offs, documented not fixed) + +- **Explicit leaf discovery**: a leaf omitted from `predictors!` is silently + not optimized/persisted. A `#[derive(Module)]` would close this; not built. +- **`Report::Custom(json)`** on the new Optimizer trait trades type precision + for object safety. +- **Capability timeouts cancel by drop** (dsrs-tools): a tool mid-side-effect + can be interrupted half-done. Tool authors own cancellation safety; not + loudly documented. +- **Reserved JS names mangle silently** (`JSON` → `JSON_tool`) rather than + erroring. +- **Compat shims still bypass the interpreter**: conversation seam + (`TODO(dsrs-phase4-conversation)`) and caller-managed tool loop + (`TODO(dsrs-phase4-caller-managed)`) — RFC 0004 §1–2 has suggested shapes. + *(Stale as of merge `f7d67a6`: both shims are gone — see §6. Kept for the + record of what session A observed.)* + +## 5. Findings from the original architecture review that were NEVER fixed + +These were found in the pre-refactor audit, judged out of scope, and are still +true at HEAD (except where noted, but re-verify before working on them — +session B's merges may have touched some): + +- Replay is O(n²) (linear scan + deep `Span` clone per intercepted call) and + replay/caching identity hashes the `Debug` output of rig's types — any rig + `Debug` change silently invalidates every fixture. `request_hash` also + includes operational knobs (`max_retries`), so ops changes invalidate + replays. +- Trace span `parent` uses the innermost-open-span heuristic — wrong under + `futures::join!` on one task. +- Dataloader: blocking I/O (`reqwest::blocking`, sync hf-hub) inside an async + library — calling `load_hf` inside a runtime can stall or panic; `println!` + + `verbose: bool` threaded through five signatures alongside `tracing`. +- `Message::content()` is lossy and used as canonical `raw_response`/cache + payload; rig→Message conversion silently drops image/audio/document blocks. +- `LM::default()` calls `Handle::current().block_on()` — panics outside a + runtime, deadlocks a current-thread runtime. +- COPRO and MIPROv2 still don't call an LM to propose candidates (template + strings + a hardcoded tip list). Preserved deliberately in phase 3; + making them real proposers is feature work. +- `CallMetadata` is not extensible (no place for a future BestOfN/Refine to + record its decision). +- `dsrs-cli` has no `optimize`/`bake`/`run` commands; it can only serve + hand-written `.dsrs` (host tools/holes refused). +- `#[module]` accepts only straight-line `let` bodies — RFC 0003 M-4 + (match→Route, for→Loop, join!→ForkJoin) was never built. +- The crate-root glob re-exports remain (prelude is additive); the public + surface is still large. + +## 6. Session B addendum — the seam merges, from their author + +Written 2026-08-20 by session B (the seam-closure session §1's table calls +"NOT reviewed by session A"). Everything below is author-side knowledge: it +does not upgrade the review status of these merges, but it answers questions +§3 item 2 asks and records sharp edges the diffs won't volunteer. + +### 6.1 Provenance and verification + +Each seam was implemented by a dedicated agent in an isolated worktree off +`3fcef48`, then merged by session B in landing order `ca14e08` (seam 6) → +`109c571` (seam 4) → `16f2b27` (seam 5). All three were clean auto-merges. +Verified by session B: + +- Baseline `cargo test --workspace` at `3fcef48`: green (exit 0) *before* any + seam landed. +- Each agent ran the full `cargo test -p dspy-rs` suite (63 binaries, zero + failures) plus `cargo clippy --workspace --all-targets` in its worktree + before committing; each reported zero introduced warnings (the ~173 + baseline clippy warnings predate this branch's seam work). +- `cargo test -p dspy-rs` re-run green (exit 0) after the seam-6 merge and + again at `16f2b27` with all three seams combined. +- NOT done, echoing §3: no live-LM run of any seam code, no independent + review, no workspace-wide suite by session B at `16f2b27` (session A's §3 + item 0 run covers that). + +**Seams 1+2 landed after the sections above were written**: merge `f7d67a6` +(agent commit `139fba6`, +2,064/−520). Its agent ran the full +`cargo test -p dspy-rs` suite green (427 passed, 64 binaries) and clippy +clean-for-touched-files in its worktree before committing; session B re-ran +the full workspace suite + clippy after the merge (result recorded in the +commit history — verify at HEAD yourself). Both `TODO(dsrs-phase4-*)` +markers are gone repo-wide. + +**This was the one conflicted merge, and the resolution is +orchestrator-authored code — review it specifically** (it is in merge commit +`f7d67a6` itself, not in `139fba6`): the seam-1+2 branch factored +`eval_agent`'s inline tool-surface build into `build_agent_surface(...)`, +written against the pre-ToolSet tree; seam 5 had meanwhile made that build +ToolSet-aware. The resolution keeps the factored helper and moves the +selection into it: `p_tool_set` picks the surface, `build_code_mode_surface` +receives the *selected* `&[ToolId]`, and `by_name` is built from the same +selection — so both the dispatching loop and the caller-managed suspended +surface honor the ToolSet gene. The `tool_set_overlay_selects_the_tool_surface` +test plus the 11 conversation tests all pass on the merged tree, but no test +exercises ToolSet *through* the suspending path specifically — worth adding +during review. + +### 6.2 Answers to §3 item 2's specific questions + +- **Seam 5 grammar change**: pre-ToolSet `.dsrs` artifacts still parse and + their canonical text + program hashes are **byte-identical** — the + `tool_set [a b]` agent option is printed only when the selection differs + from the full declared table, and the freshly-minted slot defaults to the + full table. What *does* break is JSON serde of previously-serialized + `Program` values (new `AgentLoopNode.tool_set` field + new slot kind); that + was deliberate — canonical text is the wire form per RFC 0002. +- **Seam 5 goldens**: the only changed pre-existing expectation is one line + in `test_trace_attach.rs` (`Trace::attach_program` joins all node-owned + slots, so the new slot appears — by design). Everything else in + `test_ir_text`/`test_ir_bake` is additive new tests, not regenerated goldens. +- **Seam 4 parity**: with no span evals present, the per-span gate reduces to + the old whole-rollout gate *exactly*; there is a test asserting parity, so + before/after demo harvests differ only for metrics that opt into + `evaluate_spans`. + +### 6.3 Author-side sharp edges per seam + +Seam 6 (`Structural`, `ca14e08`): + +- It is **not an `Optimizer` trait impl**. Entry points are + `Structural::compile_program(&Interpreter, &[DemoRow], &metric, env_factory)` + (+ `_with_overlay`): every accepted edit yields a *new program* that must be + re-loaded with host bindings, hence the `RuntimeEnv` factory argument the + trait can't express. Consequences: no `Report::Structural` variant; the + winner is returned for `bake`, never installed. +- Acceptance is a SIMBA-style strict minibatch gate (hill climb), not Pareto — + cross-program score-matrix columns would be incoherent. +- The proposal menu deliberately excludes `SetStop` and + `SetInstructionDefault` (free-form values belong to the value-level + optimizers); `AugmentSig` only materializes as the canonical CoT reasoning + field and is pre-filtered when present. +- `StructuralReport.edits` is a **lineage** (node ids are per-parent), not a + batch replayable against the root program. + +Seam 4 (per-span credit, `109c571`): + +- The RFC's premise was wrong: the trace format did *not* already carry + per-span `Eval` records (eval was rollout-level only). `Span.eval: + Option` was added as an additive field (`skip_serializing_if`), no + format version bump; eval-free traces serialize byte-identically. +- The hook is `TypedMetric::evaluate_spans(&self, example, prediction, trace) + -> Result>`, default = empty (module lane only; + `ProgramMetric` untouched — extending it later is purely additive). + Unknown ids are ignored; duplicate ids last-write-win. +- SIMBA's rollout-level pre-gate (append-demo vs append-rule choice) is + untouched — span evals only exclude bad spans *within* a harvested rollout. + +Seam 5 (ToolSet gene, `16f2b27`): + +- The selection **is** the surface: stop tools are *not* auto-unioned in. A + selection that drops a stop tool with `until_parse=false` degrades to + max_turns exhaustion — bounded and scored poorly, but worth knowing. +- Duplicates in a ToolSet value are a validation error (beyond the RFC, for + canonical-text/hash unambiguity). +- Overlay membership is enforced inside `Overlay::set` (so the serde load + path `from_named` is guarded too), via + `Overlay::set_tool_set(&mut self, program, slot, tools)` — unlike other + typed setters it takes the program and is fallible. +- `migrate_overlay` re-mints ToolSet entries by tool *name* intersected with + the child's declared table; a non-empty selection whose intersection is + empty drops the entry (falls back to full table). +- The Code Mode JS-collision check at interpreter load still runs over the + full declared table (conservative superset; the overlay isn't known at + load). + +Seams 1+2 (conversation surface + caller-managed, `f7d67a6`): + +- `run_conversation(chat, input, overlay, budget) -> (RunOutput, Chat)` plus + a separate caller-managed pair: `run_conversation_caller_managed` returns + `ConversationTurn::{Complete, Suspended(ToolSuspension)}` and + `resume_conversation(suspension, results)` feeds tool results back. The + split (vs a mode flag) makes `Suspended` unrepresentable on the + dispatching entry, keeping the RFC's exact return shape there. +- `input` is `Option`: continuations have no typed input; + empty-chat + `Some` renders the opening, non-empty + `Some` appends a + typed turn. Conversation surface is restricted to single-leaf programs + (`RunError::Input` otherwise). +- The `ToolSuspension` token holds **live state** (open `SpanGuard`, budget + meters, turn cursor) — required for span/budget parity across modes, so it + is process-local and NOT serializable; no cross-process resume. Dropping + it records the span `Cancelled`. +- LM-layer `ToolLoopMode::CallerManaged` was deliberately kept: it is the + one-exchange primitive the interpreter's own loop uses, and + `LM::call_with_tool_loop_mode` stays public. +- Code Mode is intentionally not applied in suspending mode (the caller + executes tools). Resume records tool `duration_us` as time-suspended; + caller-fed results always record `error: None` (feed error text as the + result for LATM-style repair). +- `Predict::build_chat` is now `async` and needs a resolvable LM (it loads + the interpreter); ~370 lines of compat path deleted + (`call_and_parse_with_input`, `serve_recorded_span`, the prompt-prefix + cache). + +All four agents independently hit the same wall: the repo is not +rustfmt-clean (repo-wide `cargo fmt` reformats ~80 pre-existing files). Each +scoped formatting to its own diff and reverted the churn. Add "deliberate +repo-wide fmt pass" to §3 item 7's version-bump chore. + +### 6.4 Deferred by session B on purpose + +- `docs/docs/api/` was **not** regenerated by any seam agent (avoids + cross-worktree conflicts). Someone must run `docs/scripts/gen_api.py` after + the last seam lands — the API pages currently predate all seam surfaces. +- No optimizer proposes ToolSet mutations yet (the slot exists, validates, + and is honored; proposal policy is follow-up work). +- Multi-edit batches per Structural generation; program-lane + `evaluate_spans`; counterfactual/ablation credit — all explicitly out of + scope per the RFC's own deferrals. + +## 7. Housekeeping + +- ~14 agent worktrees live under `.claude/worktrees/` with merged + `worktree-agent-*` branches. After review: + `git worktree list`, `git worktree remove ` for each, then + `git branch --merged v1-program-unification | grep worktree-agent | xargs git branch -d`. +- `docs/archive/` holds the superseded CURRENT_PLAN/CURRENT_SPEC. +- The docs site search index (Mixedbread `dsrs-docs`/`dsrs-code` stores) needs + re-indexing after the docs sweep — see `docs/README.md`. diff --git a/docs/index.mdx b/docs/index.mdx index aa6fe309..afc9907f 100644 --- a/docs/index.mdx +++ b/docs/index.mdx @@ -49,7 +49,7 @@ Each component has exactly one page covering what it is, how to use it, and its | Area | Pages | |---|---| | Building blocks | [Signatures](/docs/components/signatures), [Predict](/docs/components/predict), [Modules](/docs/components/modules), [Adapters](/docs/components/adapters), [Language models](/docs/components/lm), [Data](/docs/components/data), [fx](/docs/components/fx) | -| Programs as data | [The module macro](/docs/components/module-macro), [Holes](/docs/components/holes), [Capabilities](/docs/components/capabilities), [Program & nodes](/docs/components/program-and-nodes), [The .dsrs format](/docs/components/dsrs-file), [Runtime](/docs/components/runtime), [CLI](/docs/components/cli), [State](/docs/components/state) | +| Programs as data | [The module macro](/docs/components/module-macro), [Holes](/docs/components/holes), [Capabilities](/docs/components/capabilities), [Program & nodes](/docs/components/program-and-nodes), [The edit calculus](/docs/components/edit-calculus), [The .dsrs format](/docs/components/dsrs-file), [Runtime](/docs/components/runtime), [CLI](/docs/components/cli), [State](/docs/components/state) | | Tools & agents | [Tools & agents](/docs/components/tools-and-agents), [Code Mode](/docs/components/code-mode) | | Evaluation & optimization | [Traces](/docs/components/traces), [Evaluation](/docs/components/evaluation), [Optimizers](/docs/components/optimizers), [COPRO](/docs/optimizers/copro), [MIPROv2](/docs/optimizers/miprov2), [GEPA](/docs/optimizers/gepa), [Optimizer engine](/docs/components/optimizer-engine) | | Utilities | [Utilities](/docs/components/utils) | diff --git a/docs/module_system_overview.md b/docs/module_system_overview.md index 4cb3c820..f54f6252 100644 --- a/docs/module_system_overview.md +++ b/docs/module_system_overview.md @@ -1,6 +1,6 @@ # DSRs Module System — What Changed, What It Enables -This is a quick overview of the module system redesign. It builds on everything from the paper but adds a typed core and makes Section 1.3 (graph optimization) concrete. +A quick overview of the module system as it stands after the v1 program unification. The typed core from the earlier redesign is unchanged; the graph-optimization story ("Section 1.3") is now concrete in the IR rather than a planned `ProgramGraph` layer. --- @@ -8,15 +8,14 @@ This is a quick overview of the module system redesign. It builds on everything | Before | Now | |--------|-----| -| `Example` / `Prediction` as primary I/O | Typed `S::Input` / `Predicted` for the typed path; `Example` still used at optimizer/dynamic boundary | +| `Example` / `Prediction` as primary I/O | Typed `S::Input` / `Predicted`; trainset rows are plain structs projected via `ToInput`/`ToOutput` | | `#[Signature(cot)]` applies CoT at signature level | `ChainOfThought::::new()` — strategy is the module, not the signature | -| `predict.forward(example).await` | `module.call(input).await?` on the typed path | -| Manual `#[derive(Optimizable)]` + `#[parameter]` | Automatic discovery from struct shape | -| Static `FieldSpec` arrays from macros | `SignatureSchema` derived from types at runtime | -| `CallOutcome` with `.into_result()?` | `Result, PredictError>` — `?` works on stable | -| Section 1.3 graph optimization (future work) | `ProgramGraph` being built now (V6) — walker foundation landed in V5 | - -> **TODO:** Nail down the long-term role of `Example`. It's still load-bearing at the DynPredictor boundary (demo conversion, optimizer manipulation, DataLoader). The typed path doesn't kill it — but its scope and future API need a decision. +| Reflection-based leaf discovery (facet walker, `DynPredictor` handles) | Explicit declaration: the `Predictors` trait, one `predictors!(MyModule { field_a, field_b })` line per module | +| Optimizers mutate the module during search | Candidates are data injected ambiently per rollout (`fx::with_params`); the single mutation is the final install of the winner | +| Per-optimizer engines (`EvalEngine` / `ProgramEvalEngine`) | One shared `Engine` over an `OptimizeTarget` (typed module lane or loaded-program lane) | +| `ReAct` module, `ModuleExt::map`/`and_then` combinators | Deleted. Tool loops are IR `AgentLoop` nodes (`#[agent]`, or tools on a `Predict`); output transforms are plain Rust in `forward` | +| `Predict` renders/parses on its own LM path | `Predict` executes as a 1-node IR program through the `Interpreter`; instance state is an `ir::Overlay` | +| Graph optimization as future work | The IR edit calculus: `Program::edited(&[Edit])`, `legal_edits`, `migrate_overlay` | --- @@ -36,16 +35,8 @@ let result = module.call(QAInput { question: "2+2?".into() }).await?; result.reasoning // augmented field — direct access result.answer // original field — via Deref -// Swap to ReAct — same call site -let module = ReAct::::builder() - .tool("search", "Search the web", search_fn) - .build(); - // Batch without changing the module -let results = dsrs::forward_all(&module, inputs, 5).await; - -// Simple transform without impl Module -let confident = module.map(|r| Confident { answer: r.answer, confidence: 0.9 }); +let results = dspy_rs::forward_all(&module, inputs, 5).await; ``` --- @@ -58,97 +49,88 @@ A new augmentation (like adding confidence scoring to any output): #[augment(output, append)] struct Confidence { /// Model's self-assessed confidence - confidence: f64, + #[output] confidence: f64, } // Done — WithConfidence now exists and composes with any signature // Users write: Predict> // They get: result.answer + result.confidence ``` -A new module (like BestOfN — runs N times, picks best): +A new composite module is a struct with predictor fields, a `predictors!` line, and a `forward` body of ordinary Rust: + ```rust -#[derive(Module)] -struct BestOfN { - module: M, // walker sees through — finds all Predict leaves inside - #[skip] n: usize, - #[skip] reward_fn: Box f64 + Send + Sync>, +struct TwoStepQA { + retrieve: Predict, + answer: ChainOfThought, } -impl Module for BestOfN where M::Input: Clone { - type Input = M::Input; - type Output = M::Output; - - async fn forward(&self, input: M::Input) -> Result, PredictError> { - let mut best = None; - let mut best_score = f64::NEG_INFINITY; - for _ in 0..self.n { - let result = self.module.call(input.clone()).await?; - let score = (self.reward_fn)(&input, &result); - if score > best_score { best_score = score; best = Some(result); } - } - best.ok_or(PredictError::AllAttemptsFailed) +dspy_rs::predictors!(TwoStepQA { retrieve, answer }); + +impl Module for TwoStepQA { + type Input = RetrieveInput; + type Output = WithReasoning; + + async fn forward(&self, input: Self::Input) -> Result, PredictError> { + let ctx = self.retrieve.call(input).await?; + self.answer.call(AnswerInput { context: ctx.passages.clone() }).await } } ``` -`#[derive(Module)]` makes `module: M` discoverable — optimizers automatically find and tune the Predict leaves inside whatever `M` is. `#[skip]` fields (closures, config) are invisible to the walker. No traversal code, no schema construction. +The `predictors!` line is the whole discovery story: each field identifier becomes the leaf's canonical name — its trace-span component, its optimizer-candidate key, and its `ModuleState` persistence key. No derive magic, no traversal code, no pointer casts. --- ## What optimizers see ```rust -optimizer.compile(&mut module, trainset, metric).await; +optimizer.compile_module(&mut module, &trainset, &metric).await?; // internally: -visit_named_predictors_mut(&mut module, |path, predictor| { - // mutate demos, instructions, dump/load state — all through DynPredictor handles - ControlFlow::Continue(()) -})?; +let mut target = OptimizeTarget::module(&mut module, &trainset, &metric); +// — snapshots each declared leaf as a LeafInfo (schema, instruction, demos) +// — stamps each leaf's trace name once (the naming pass) +let mut engine = Engine::new(optimizer.engine_config()); +optimizer.compile(&mut target, &mut engine).await?; +// — candidates are name-keyed `Candidate`s, injected ambiently per rollout; +// evaluation never mutates the module, so candidates fan out concurrently +// — the winner is installed exactly once via PredictorInfo::load_state // after compile returns, module.call() uses optimized params — no code change ``` ---- +The `Optimizer` trait is object-safe: `Box` pipelines can share one `Engine` — one budget, one rollout cache, one score matrix — across stages. The same trait drives the program lane (`OptimizeTarget::program`): an interpreter-loaded `.dsrs` program, JSON examples, and an overlay winner for `Program::bake`. -## What ProgramGraph enables (Section 1.3 made concrete) +--- -This is the paper's "Dynamic Workflow Optimization" — pipelines as executable graphs that can restructure themselves. +## Structural optimization (Section 1.3 made concrete) -**Current state:** the V5 walker (`visit_named_predictors_mut`) enumerates all Predict leaves in a typed module through callback traversal. Everything else — `ProgramGraph`, `DynModule`, `StrategyFactory`, registry, type-validated edges, topological execution — is being built now in V6. +The paper's "Dynamic Workflow Optimization" landed as the IR **edit calculus**, not a separate graph layer. A `#[module]` function lowers to a `Program` — a validated node tree with named leaves and addressable parameter slots — and structural moves are plain serde values applied purely: ```rust -// Project a typed module into a mutable graph (snapshot — original untouched) -let graph = ProgramGraph::from_module(&module); - -// Or build from scratch via registry -let mut graph = ProgramGraph::new(); -let cot = registry::create("chain_of_thought", &schema, Default::default())?; -graph.add_node("cot", cot)?; -graph.connect("input", "question", "cot", "question")?; // edges type-validated -let result = graph.execute(input).await?; - -// After optimization, fit back to the typed module -graph.fit(&mut module); -``` +use dspy_rs::ir::{Edit, migrate_overlay}; -**Split** from the paper: a meta planner decides a complex signature should be two steps. It calls `graph.add_node` twice with simpler schemas from `registry::create`, rewires edges with `graph.connect`, removes the original with `graph.replace_node`. Edge type validation catches wiring errors immediately. - -**Fuse**: two adjacent nodes with compatible schemas get replaced by a single node with a merged signature. Same mutation APIs. +let leaf = program.leaf_id("drafter").unwrap(); +let menu = program.legal_edits(leaf); // the proposer menu (LLM-promptable) +let child = program.edited(&[Edit::SwapLeaf { leaf, to: swap_target }])?; +let carried = migrate_overlay(&program, &tuned, &child); // value progress survives +``` -**The key architectural property**: both the typed path and the graph path use the same `SignatureSchema` → `ChatAdapter` → prompt format pipeline. A `Predict` and a `registry::create("predict", &qa_schema, ...)` produce identical prompts. The meta planner can restructure the graph without worrying about prompt divergence. +`edited` clones, applies, re-validates with the loader's own rules, and seals a new content hash; the parent program and every hash-bound artifact minted against it stay coherent. **Split**, **fuse**, wrap-in-retry, predict↔agent swaps, and tool add/remove are all expressible as `Edit` batches; data-flow legality stays with the one validator both the builder and the loader use. -**The cycle**: project → optimize (parameter and/or structural) → fit-back → evaluate → repeat. The graph is the optimizer's scratch space; the user's typed module is the stable interface. +The key architectural property is unchanged: the typed path and the program path share one rendering pipeline (`SignatureDef` → `ChatAdapter` → prompt). A `Predict` — which itself executes as a 1-node program — and a loaded `predict` leaf over the same signature produce identical prompts, so restructuring cannot cause prompt divergence. --- ## Layer stack ``` -You're here What you touch What's invisible to you -───────────────────────────────────────────────────────────────────────── -App developer Signature, module.call() Everything below -Module author #[derive(Module)], forward() Discovery, graph -Optimizer dev Optimizer::compile internals (`visit_named_predictors_mut`, DynPredictor) Graph, registry -Meta planner ProgramGraph, registry (bottom layer — Section 1.3) +You're here What you touch What's invisible to you +──────────────────────────────────────────────────────────────────────────────────── +App developer Signature, module.call() Everything below +Module author predictors!, forward() IR lowering, interpreter +Optimizer dev Optimizer::compile, OptimizeTarget, IR internals + Engine, Candidate +Structural optimizer Program, Edit, legal_edits, validator internals + migrate_overlay ``` -Each layer only exists if you need it. Simple usage never instantiates the graph layer. +Each layer only exists if you need it. Simple usage never touches the IR directly. diff --git a/docs/rfcs/0004-remaining-seams.md b/docs/rfcs/0004-remaining-seams.md new file mode 100644 index 00000000..c463e817 --- /dev/null +++ b/docs/rfcs/0004-remaining-seams.md @@ -0,0 +1,123 @@ +# RFC 0004 — Remaining seams (post-unification ledger) + +Status: informational. The v1 unification (phases 1–4) put `Predict` on the IR +interpreter, unified the optimizers over one engine with ambient candidate +injection, and collapsed the build-time feature fiction. What follows is the +deliberate remainder: seams we know about, left open on purpose, each with a +suggested shape for whoever closes it. + +## 1. Conversation surface — `TODO(dsrs-phase4-conversation)` + +> Status update (2026-08-20): an implementation landed in merge `f7d67a6` on `v1-program-unification` (`Interpreter::run_conversation`; `build_chat`/`call_and_parse` demoted to wrappers; the prompt-prefix cache and the `TODO(dsrs-phase4-conversation)` marker are gone). Not yet reviewed — see docs/handoff-v1-unification.md §6. + +**What:** multi-turn conversations still run on the LM-layer compat path. +`Predict::build_chat` / `call_and_parse` hand a caller-owned `Chat` to the LM +client directly, bypassing the interpreter, because the interpreter's only +entry is map-in/map-out (`Interpreter::run_collecting`). + +**Why deferred:** giving the interpreter a conversation-in/conversation-out +surface touches its span model (a turn is not a run) and the replay contract, +and nothing in the optimizer stack needs it yet. + +**Suggested shape:** an interpreter-native entry that accepts a prior +conversation and returns the extended one alongside outputs — e.g. +`run_conversation(chat, input, overlay, budget) -> (RunOutput, Chat)` — with +the `AgentLoop` machinery reused for turn bookkeeping. `build_chat` / +`call_and_parse` then become thin wrappers and the static prompt-prefix cache +in `Predict` can be deleted. + +## 2. Caller-managed tool loop — `TODO(dsrs-phase4-caller-managed)` + +> Status update (2026-08-20): an implementation landed in merge `f7d67a6` (`run_conversation_caller_managed` / `resume_conversation` — the loop suspends on tool calls with a process-local resumption token; the `TODO(dsrs-phase4-caller-managed)` marker is gone; the LM-layer `ToolLoopMode::CallerManaged` survives as the one-exchange primitive the interpreter itself uses). Not yet reviewed — see docs/handoff-v1-unification.md §6. + +**What:** `ToolLoopMode::CallerManaged` (the "return me the tool calls, I'll +execute them" pattern) lives on the LM-layer path +(`core/lm/mod.rs`, used by `Predict::call_and_parse_with_input`). Typed +`call`s with tools already run through the interpreter's `AgentLoop`; the +caller-managed variant does not. + +**Why deferred:** it is built on caller-owned chats, so it is blocked on seam +1 — the interpreter cannot yield mid-loop to an external executor today. + +**Suggested shape:** once the conversation surface exists, express +caller-managed as an `AgentLoop` that suspends on tool calls instead of +dispatching them: return pending calls plus a resumption token, and let the +caller feed results back in. That keeps trace spans, budget metering, and +stop-tool semantics identical across both modes. + +## 3. Shared-pointer traversal — `TODO(dsrs-shared-ptr-policy)` (retired) + +**What:** the old facet reflection walker refused to traverse `Rc`/`Arc` +containers when discovering `Predict` leaves, with an explicit error carrying +this marker. + +**Why it's gone:** phase 3 deleted the walker entirely — leaf discovery is now +explicit via the `Predictors` trait (`predictors!` macro), so there is no +container traversal left to have a policy about. The marker survives only in +`docs/specs/modules/*` prose describing the deleted design; treat those +documents as historical. + +**Suggested shape:** none. If reflection-based discovery ever returns, the +policy question returns with it; the explicit-declaration contract made it +moot. + +## 4. Whole-rollout credit assignment (`optimizer/harvest.rs`) + +> Status update (2026-08-20): an implementation landed in merge `109c571` on `v1-program-unification` (authored by a separate session). Not yet reviewed against the suggested shape below — see docs/handoff-v1-unification.md §3.2 before marking closed. + +**What:** demo harvesting scores every span in a rollout with the rollout's +single metric score — a good final answer marks *all* intermediate predictor +calls as good demos, including any that a later step had to recover from. + +**Why deferred:** per-leaf credit needs either per-span evals in the trace or +a counterfactual scorer, and the shipped optimizers (Bootstrap, MIPROv2) work +acceptably on whole-rollout signal for the shallow programs people build +today. + +**Suggested shape:** the trace format already carries per-span `Eval` records; +let metrics optionally attach span-level scores during evaluation +(`TypedMetric` gains a per-trace hook), and have `harvest.rs` prefer a span's +own eval over the rollout score when present. Deeper counterfactual schemes +(ablate-one-span replays) can layer on the replay machinery later. + +## 5. Tool membership as a `ParamSlot` (the ToolSet gene) + +> Status update (2026-08-20): an implementation landed in merge `16f2b27` (`ParamKind::ToolSet`; touches the `.dsrs` grammar and therefore program hashes). Not yet reviewed — see docs/handoff-v1-unification.md §3.2. + +**What:** which tools an `AgentLoop` carries is structural today +(`AgentLoopNode::tools: Box<[ToolId]>`); only each tool's *description* is an +optimizable slot (`ParamKind::ToolDesc`). An optimizer can rewrite what a tool +says it does, but not drop a distracting tool or add a relevant one. + +**Why deferred:** tool membership changes the capability footprint of a node, +so a membership gene has to interact with the program's cap-ceiling validation +— an overlay must not be able to smuggle in a tool the program's grants don't +cover. + +**Suggested shape:** a `ParamKind::ToolSet` slot per agent node whose value is +a subset of the *declared* tool table (declaration stays structural, selection +becomes a gene). Validation stays load-time: the legal alphabet is the +declared tools, so any subset is capability-safe by construction. Mutation +proposals then compose with overlays like any other slot value. + +## 6. Structural optimizers over `ir::Edit` + +> Status update (2026-08-20): an implementation landed in merge `ca14e08` (`optimizer/structural.rs`, the sixth strategy). Not yet reviewed — see docs/handoff-v1-unification.md §3.2. + +**What:** the edit calculus is fully shipped — `Edit`, `Program::edited`, +`Program::legal_edits` (the menu of applicable edits per node), and +`migrate_overlay` for carrying tuned slot values across a structural change — +but no shipped optimizer proposes edits. All five strategies tune +instructions/demos through overlays only. + +**Why deferred:** structural search needs an evaluation budget model (every +candidate is a new program that must be re-scored from scratch, minus what +`migrate_overlay` preserves) and a proposal policy; both are research-shaped +rather than plumbing-shaped. + +**Suggested shape:** a GEPA-style loop where the reflection step prompts over +`legal_edits` output (the menu is already serializable data), applies the +chosen `Edit` via `Program::edited`, migrates the incumbent overlay, and +scores the child against the parent on a shared minibatch. The engine's +candidate machinery already treats programs as data, so this slots in as a +sixth strategy rather than a new framework. diff --git a/docs/scripts/gen_api.py b/docs/scripts/gen_api.py index a4a558e7..4bdd4875 100644 --- a/docs/scripts/gen_api.py +++ b/docs/scripts/gen_api.py @@ -6,9 +6,10 @@ RUSTC_BOOTSTRAP=1 cargo rustdoc -p dspy-rs --lib --all-features -- -Z unstable-options --output-format json RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs-tools --lib -- -Z unstable-options --output-format json RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs_macros --lib -- -Z unstable-options --output-format json + RUSTC_BOOTSTRAP=1 cargo rustdoc -p dsrs-syntax --lib -- -Z unstable-options --output-format json python3 docs/scripts/gen_api.py -Reads target/doc/{dspy_rs,dsrs_tools,dsrs_macros}.json and rewrites +Reads target/doc/{dspy_rs,dsrs_tools,dsrs_macros,dsrs_syntax}.json and rewrites docs/docs/api/*.mdx. Every page is fully generated; do not edit them by hand. The item inventory and doc summaries come from the compiler, so the pages cannot drift from the code. Full signatures, methods, and @@ -30,6 +31,7 @@ ("dspy_rs.json", "dspy-rs", "dspy_rs", None), ("dsrs_tools.json", "dsrs-tools", "dsrs_tools", "dsrs-tools"), ("dsrs_macros.json", "dsrs-macros", "dsrs_macros", "dsrs-macros"), + ("dsrs_syntax.json", "dsrs-syntax", "dsrs_syntax", "dsrs-syntax"), ] KIND_LABELS = [ diff --git a/docs/snippets/optimizer-comparison.mdx b/docs/snippets/optimizer-comparison.mdx index f25e9b72..f180b683 100644 --- a/docs/snippets/optimizer-comparison.mdx +++ b/docs/snippets/optimizer-comparison.mdx @@ -5,5 +5,6 @@ | `SIMBA` | Minibatch introspective ascent (demos + rules) | No | Low (steps × minibatch) | | `GEPA` | Genetic-Pareto evolution with feedback | **Yes** | Medium-high (iterations × eval) | | `MIPROv2` | Trace-guided candidate generation | No | Medium (candidates × trials × trainset) | +| `Structural` | LM-guided graph edits over `ir::Edit` (program lane only) | No | Medium (examples + iterations × minibatch) | -GEPA is the only optimizer that requires textual feedback from the metric (`Eval::with_feedback`). The others use numerical scores alone. Full configuration tables for all five live in the [optimizers reference](/docs/components/optimizers). +GEPA is the only optimizer that requires textual feedback from the metric (`Eval::with_feedback`). The others use numerical scores alone. Full configuration tables for all six live in the [optimizers reference](/docs/components/optimizers).