From c7301faf6b014776bf0e2b6dc07e648629720468 Mon Sep 17 00:00:00 2001 From: stephen Date: Tue, 1 Sep 2026 19:20:51 -0700 Subject: [PATCH] chore: derive cassette suites from the tree and sort-check manifest lists --- .github/workflows/ci.yaml | 8 + AGENTS.md | 2 + CONTRIBUTING.md | 2 +- Cargo.toml | 107 ++++--- README.md | 2 + crates/rig-core/src/providers/mod.rs | 2 + tests/common/cassette_safety.rs | 415 ++++++++++----------------- xtask/src/main.rs | 5 + xtask/src/sorted_blocks.rs | 185 ++++++++++++ xtask/src/sorted_blocks/tests.rs | 52 ++++ 10 files changed, 465 insertions(+), 315 deletions(-) create mode 100644 xtask/src/sorted_blocks.rs create mode 100644 xtask/src/sorted_blocks/tests.rs diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index bfaabc619b..de91fd74c0 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -272,6 +272,14 @@ jobs: - name: Check test modules are sibling files run: cargo xtask check-test-layout + # Lists that grow with every provider, companion crate or dependency + # (marked `sorted: start` / `sorted: end` in Cargo.toml, rig-core's + # providers module list and the README integration table) stay in + # order, so two PRs that each add one entry touch different lines + # instead of the same hunk. + - name: Check marked lists are sorted + run: cargo xtask check-sorted-blocks + clippy: name: stable / clippy runs-on: ubuntu-latest diff --git a/AGENTS.md b/AGENTS.md index 2b5cf6f1b2..41875c3cda 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -195,6 +195,8 @@ surfaces (`AgentRunner::stream` and `AgentRunner::run` share `drive_agent`). or `tests.rs` beside `mod.rs`/`lib.rs`. `cargo xtask check-test-layout` enforces this in CI. - Avoid unrelated refactors. +- Lists between `sorted: start` / `sorted: end` markers (Cargo.toml dependency and feature blocks, rig-core's providers module list, the README integration table) stay in case-insensitive byte order; `cargo xtask check-sorted-blocks` enforces it. Insert new entries in place, never at the end. +- Cassette wrappers are discovered by convention (a function named `with_*cassette*` whose first argument is the scenario, called only from its own `tests/providers//` directory); there is no registry to update. ## Cassette Regression Tests diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 73b7af48a9..6a3ec23587 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -13,7 +13,7 @@ Additionally, please ensure that if you are submitting a bug ticket (ie, somethi Contributions are always encouraged and welcome. Before creating a pull request, create a new issue that tracks that pull request describing the problem in more detail. Pull request descriptions should include information about its implementation, especially if it makes changes to existing abstractions. -PRs should be small and focused and should avoid interacting with multiple facets of the library. This may result in a larger PR being split into two or more smaller PRs. Commit messages should follow the [Conventional Commit](https://conventionalcommits.org/en/v1.0.0) format (prefixing with `feat`, `fix`, etc.) as this integrates into our auto-releases via a [release-plz](https://github.com/MarcoIeni/release-plz) Github action. +PRs should be small and focused and should avoid interacting with multiple facets of the library. Lists marked `sorted: start` / `sorted: end` (in `Cargo.toml`, the providers module list and the README integration table) are kept sorted so concurrent PRs do not collide on the same hunk; `cargo xtask check-sorted-blocks` runs in CI. This may result in a larger PR being split into two or more smaller PRs. Commit messages should follow the [Conventional Commit](https://conventionalcommits.org/en/v1.0.0) format (prefixing with `feat`, `fix`, etc.) as this integrates into our auto-releases via a [release-plz](https://github.com/MarcoIeni/release-plz) Github action. Do not edit `CHANGELOG.md`, `crates/*/CHANGELOG.md` or `MIGRATING.md` in a pull request; CI fails the PR if you do. Put changelog bullets and migration notes in the PR description under `## Changelog` and `## Migration` (the PR template has both). The repository squash-merges, so those sections become the merge commit body, and the release PR regenerates both files from them via `scripts/release-notes.sh`. diff --git a/Cargo.toml b/Cargo.toml index 5d172d8f26..2e2844b00a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -110,24 +110,30 @@ edition = "2024" # # The `rig` facade is the workspace-root package; examples depend on it via # `rig.workspace = true` (add per-example features at the use site). +# sorted: start rig = { path = ".", version = "0.42.0" } rig-agent = { path = "crates/rig-agent", version = "0.42.0" } +rig-core = { path = "crates/rig-core", version = "0.42.0" } rig-reqwest = { path = "crates/rig-reqwest", version = "0.42.0" } rig-rmcp = { path = "crates/rig-rmcp", version = "0.42.0" } -rig-core = { path = "crates/rig-core", version = "0.42.0" } +# sorted: end + +# sorted: start anyhow = "1" arrow-array = "58" as-any = "0.3" assert_fs = "1" async-stream = "0.3" -axum = "0.8" aws-config = { version = "1", default-features = false } # 1.124: `OutputConfig`/`ServiceTier`/citations types rig-bedrock maps. aws-sdk-bedrockruntime = { version = "1.124", default-features = false } aws-smithy-eventstream = "0.60" aws-smithy-runtime-api = "1" aws-smithy-types = "1" +axum = "0.8" base64 = "0.22" +# Runtime-agnostic executor used by the `agent_no_tokio` example. +bevy_tasks = { version = "0.19", features = ["multi_threaded"] } bytes = "1" candle-core = "0.11" candle-nn = "0.11" @@ -143,15 +149,13 @@ eventsource-stream = "0.2.3" fastembed = { version = "4.5", default-features = false } fastrand = "2" futures = "0.3" -# Runtime-agnostic executor used by the `agent_no_tokio` example. -bevy_tasks = { version = "0.19", features = ["multi_threaded"] } futures-timer = "3" glob = "0.3" -google-cloud-auth = "1" # 1.7: `ThinkingLevel`, `ImageConfig`, `FunctionResponse{Blob,FileData}`. google-cloud-aiplatform-v1 = { version = "1.7", default-features = false, features = [ "prediction-service", ] } +google-cloud-auth = "1" http = "1" # 0.8.2: `Recording::export_async` (cassette harness). httpmock = "0.8.2" @@ -175,6 +179,7 @@ ordered-float = "5" pgvector = "0.4.2" pin-project-lite = "0.2" proc-macro2 = "1" +prost = "0.14" # 1.10: the `Qdrant` client, `Payload`, and the query/upsert builders. qdrant-client = { version = "1.10", default-features = false, features = [ "serde", @@ -187,17 +192,16 @@ reqwest = { version = "0.13", default-features = false } reqwest-middleware = { version = "0.5", default-features = false } reqwest-retry = "0.9" rmcp = "2" -url = "2" rusqlite = "0.32" -scylla = "1" +safetensors = "0.8" schemars = "1" +scylla = "1" serde = "1" -serde_yaml = "0.9" # 1.0.54: the `float_roundtrip` feature. serde_json = "1.0.54" -safetensors = "0.8" -sha2 = "0.10" serde_path_to_error = "0.1" +serde_yaml = "0.9" +sha2 = "0.10" sqlite-vec = "0.1" sqlx = "0.9" # 3.2: 3.0/3.1 no longer build against `surrealdb-core ^3` as published. @@ -207,27 +211,28 @@ term_size = "0.3" testcontainers = "0.27" textwrap = "0.16" thiserror = "2" +tokenizers = { version = "0.22", default-features = false } tokio = "1" tokio-rusqlite = { version = "0.6", default-features = false } -tokio-test = "0.4" tokio-stream = "0.1" +tokio-test = "0.4" tokio-tungstenite = { version = "0.29", default-features = false } -tokenizers = { version = "0.22", default-features = false } tonic = "0.14" tonic-build = "0.14" tonic-prost = "0.14" tonic-prost-build = "0.14" -prost = "0.14" tracing = "0.1" # 0.2.5: `Instrumented: Stream` under `futures-03` for futures 0.3 stable. tracing-futures = "0.2.5" tracing-opentelemetry = "0.33" tracing-subscriber = "0.3" +url = "2" uuid = "1" -wasm-bindgen-futures = "0.4" wasm-bindgen = "0.2" +wasm-bindgen-futures = "0.4" web-time = "1" zerocopy = "0.8" +# sorted: end [workspace.metadata.cargo-autoinherit] # Skip cargo-autoinherit for these packages @@ -241,15 +246,16 @@ rig-core = { path = "crates/rig-core", version = "0.42.0", default-features = fa rig-reqwest = { path = "crates/rig-reqwest", version = "0.42.0", optional = true, default-features = false } rig-agent = { path = "crates/rig-agent", version = "0.42.0", optional = true, default-features = false } rig-rmcp = { path = "crates/rig-rmcp", version = "0.42.0", optional = true } +rig-derive = { path = "crates/rig-derive", version = "0.42.0", optional = true } +# Companion crates, one optional dependency per `crates/rig-*` integration. +# sorted: start rig-bedrock = { path = "crates/rig-bedrock", version = "0.42.0", optional = true, default-features = false } rig-candle = { path = "crates/rig-candle", version = "0.42.0", optional = true, default-features = false } rig-fastembed = { path = "crates/rig-fastembed", version = "0.42.0", optional = true, default-features = false } rig-gemini-grpc = { path = "crates/rig-gemini-grpc", version = "0.42.0", optional = true, default-features = false } rig-helixdb = { path = "crates/rig-helixdb", version = "0.42.0", optional = true, default-features = false } rig-lancedb = { path = "crates/rig-lancedb", version = "0.42.0", optional = true, default-features = false } -lancedb = { workspace = true, optional = true } rig-memory = { path = "crates/rig-memory", version = "0.42.0", optional = true, default-features = false } -rig-derive = { path = "crates/rig-derive", version = "0.42.0", optional = true } rig-milvus = { path = "crates/rig-milvus", version = "0.42.0", optional = true, default-features = false } rig-mongodb = { path = "crates/rig-mongodb", version = "0.42.0", optional = true, default-features = false } rig-neo4j = { path = "crates/rig-neo4j", version = "0.42.0", optional = true, default-features = false } @@ -261,6 +267,10 @@ rig-sqlite = { path = "crates/rig-sqlite", version = "0.42.0", optional = true, rig-surrealdb = { path = "crates/rig-surrealdb", version = "0.42.0", optional = true, default-features = false } rig-vectorize = { path = "crates/rig-vectorize", version = "0.42.0", optional = true, default-features = false } rig-vertexai = { path = "crates/rig-vertexai", version = "0.42.0", optional = true, default-features = false } +# sorted: end +# Re-exported beside `rig-lancedb` for the `lancedb` feature; not a companion +# crate itself, so it sits outside the sorted block. +lancedb = { workspace = true, optional = true } # The bundled websocket backend is native-only (tungstenite needs a tokio # reactor and raises a `compile_error!` on wasm), so it is scoped to non-wasm @@ -280,6 +290,7 @@ rig-vertexai = { path = "crates/rig-vertexai", version = "0.42.0", optional = tr rig-tungstenite = { path = "crates/rig-tungstenite", version = "0.42.0", optional = true } [dev-dependencies] +# sorted: start anyhow = { workspace = true } arrow-array = { workspace = true } assert_fs = { workspace = true } @@ -289,25 +300,19 @@ aws-sdk-bedrockruntime = { workspace = true } aws-smithy-eventstream = { workspace = true } aws-smithy-runtime-api = { workspace = true } aws-smithy-types = { workspace = true } -rig-core = { path = "crates/rig-core", version = "0.42.0", default-features = false, features = [ - "test-utils", -] } -rig-agent = { path = "crates/rig-agent", version = "0.42.0", default-features = false, features = [ - "test-utils", -] } -tokio = { workspace = true, features = ["full"] } -tracing-subscriber = { workspace = true, features = ["env-filter"] } -tokio-test = { workspace = true } -redis = { workspace = true, features = ["tokio-comp", "aio", "vector-sets"] } -serde_path_to_error = { workspace = true } +axum = { workspace = true } base64 = { workspace = true } bytes = { workspace = true } futures = { workspace = true } httpmock = { workspace = true, features = ["record", "proxy"] } +hyper-util = { workspace = true, features = ["service", "server"] } mongodb = { workspace = true } neo4rs = { workspace = true } +opentelemetry = { workspace = true } +opentelemetry-otlp = { workspace = true } +opentelemetry_sdk = { workspace = true, features = ["rt-tokio"] } qdrant-client = { workspace = true } -reqwest-retry = { workspace = true } +redis = { workspace = true, features = ["tokio-comp", "aio", "vector-sets"] } reqwest = { workspace = true, features = ["json", "stream"] } reqwest-middleware = { workspace = true, features = [ "json", @@ -316,39 +321,46 @@ reqwest-middleware = { workspace = true, features = [ "http2", "rustls", ] } +reqwest-retry = { workspace = true } +rig-agent = { path = "crates/rig-agent", version = "0.42.0", default-features = false, features = [ + "test-utils", +] } +rig-core = { path = "crates/rig-core", version = "0.42.0", default-features = false, features = [ + "test-utils", +] } +rmcp = { workspace = true, features = [ + "client", + "macros", + "reqwest", + "transport-streamable-http-client", + "transport-streamable-http-client-reqwest", + "transport-streamable-http-server-session", + "transport-streamable-http-server", + "transport-worker", +] } +rusqlite = { workspace = true, features = ["bundled"] } schemars = { workspace = true } serde = { workspace = true, features = ["derive"] } -sha2 = { workspace = true } # float_roundtrip keeps cassette scrubbing idempotent on float-heavy bodies # (e.g. embedding vectors): without it, parse→print can oscillate between two # representations of the same decimal and the recorded YAML never stabilizes. serde_json = { workspace = true, features = ["float_roundtrip"] } +serde_path_to_error = { workspace = true } serde_yaml = { workspace = true } -syn = { workspace = true, features = ["full", "visit"] } -rusqlite = { workspace = true, features = ["bundled"] } +sha2 = { workspace = true } sqlite-vec = { workspace = true } sqlx = { workspace = true, features = ["runtime-tokio", "postgres", "uuid", "json"] } +syn = { workspace = true, features = ["full", "visit"] } testcontainers = { workspace = true } thiserror = { workspace = true } +tokio = { workspace = true, features = ["full"] } tokio-rusqlite = { workspace = true, features = ["bundled"] } -hyper-util = { workspace = true, features = ["service", "server"] } -rmcp = { workspace = true, features = [ - "client", - "macros", - "reqwest", - "transport-streamable-http-client", - "transport-streamable-http-client-reqwest", - "transport-streamable-http-server-session", - "transport-streamable-http-server", - "transport-worker", -] } -axum = { workspace = true } -opentelemetry = { workspace = true } -opentelemetry_sdk = { workspace = true, features = ["rt-tokio"] } -opentelemetry-otlp = { workspace = true } +tokio-test = { workspace = true } tracing = { workspace = true } tracing-opentelemetry = { workspace = true } +tracing-subscriber = { workspace = true, features = ["env-filter"] } url = { workspace = true } +# sorted: end [features] default = ["rig-core/default", "reqwest", "agent", "derive", "rustls"] @@ -358,6 +370,8 @@ test-utils = ["rig-core/test-utils", "rig-agent?/test-utils"] # (`tests/tool_facade_features.rs`). Not in `default`; `--all-features` (CI's # lane) enables it so the guard still runs, while ad-hoc `cargo test` skips it. facade-build-tests = [] +# Companion-crate features, one per optional dependency above. +# sorted: start bedrock = ["dep:rig-bedrock"] candle = ["dep:rig-candle"] fastembed = [ @@ -382,6 +396,7 @@ sqlite = ["dep:rig-sqlite"] surrealdb = ["dep:rig-surrealdb"] vectorize = ["dep:rig-vectorize"] vertexai = ["dep:rig-vertexai"] +# sorted: end audio = ["rig-core/audio", "rig-agent?/audio", "rig-reqwest?/audio"] image = ["rig-core/image", "rig-agent?/image", "rig-reqwest?/image"] derive = ["dep:rig-derive", "rig-core/derive", "rig-agent?/derive"] diff --git a/README.md b/README.md index d5dd84c38a..2b6d4e0c8c 100644 --- a/README.md +++ b/README.md @@ -155,6 +155,7 @@ The root `rig` facade exposes companion crates behind one feature per integratio rig = { version = "0.36.0", features = ["lancedb", "fastembed"] } ``` + | Integration | Crate | Feature | Module path | | --- | --- | --- | --- | | AWS Bedrock | [`rig-bedrock`](https://github.com/0xPlaygrounds/rig/tree/main/crates/rig-bedrock) | `bedrock` | `rig::bedrock` | @@ -175,6 +176,7 @@ rig = { version = "0.36.0", features = ["lancedb", "fastembed"] } | ScyllaDB | [`rig-scylladb`](https://github.com/0xPlaygrounds/rig/tree/main/crates/rig-scylladb) | `scylladb` | `rig::scylladb` | | SQLite | [`rig-sqlite`](https://github.com/0xPlaygrounds/rig/tree/main/crates/rig-sqlite) | `sqlite` | `rig::sqlite` | | SurrealDB | [`rig-surrealdb`](https://github.com/0xPlaygrounds/rig/tree/main/crates/rig-surrealdb) | `surrealdb` | `rig::surrealdb` | + `rig::memory` is available without the `memory` feature; it contains the core conversation memory traits and in-memory backend re-exported from `rig-core`. diff --git a/crates/rig-core/src/providers/mod.rs b/crates/rig-core/src/providers/mod.rs index 2bcede35bf..8968fb1b62 100644 --- a/crates/rig-core/src/providers/mod.rs +++ b/crates/rig-core/src/providers/mod.rs @@ -107,6 +107,7 @@ //! # Ok(()) //! # } //! ``` +// sorted: start pub mod anthropic; pub mod azure; pub mod chatgpt; @@ -134,3 +135,4 @@ pub mod voyageai; pub mod xai; pub mod xiaomimimo; pub mod zai; +// sorted: end diff --git a/tests/common/cassette_safety.rs b/tests/common/cassette_safety.rs index 6c7224ba0c..808338c9fd 100644 --- a/tests/common/cassette_safety.rs +++ b/tests/common/cassette_safety.rs @@ -1,4 +1,21 @@ //! Safety checks for committed cassette fixtures. +//! +//! There is no registry of providers or wrappers. Suites are discovered from +//! the tree, and the conventions the discovery relies on are the invariant: +//! +//! * a provider is a directory `tests/providers//` that has a test +//! binary `tests/.rs`, and its cassettes live under +//! `tests/cassettes//`; +//! * a cassette wrapper is any function whose name starts with `with_` and +//! contains `cassette`; its first argument is the scenario (a string +//! literal or `CassetteSpec::new("...")`); +//! * a wrapper is called only from inside its own provider's directory, so +//! the directory a call sits in is the provider its scenario belongs to. +//! +//! A wrapper that names another provider in its identifier but is called +//! from a different provider's directory is rejected, because the scenario +//! would otherwise be registered under the wrong provider and surface as a +//! confusing missing/orphaned pair. use std::collections::BTreeSet; use std::fs; @@ -9,245 +26,48 @@ use syn::visit::{self, Visit}; use syn::{Expr, ExprCall, ExprLit, ItemFn, Lit}; const CASSETTE_ROOT: &str = concat!(env!("CARGO_MANIFEST_DIR"), "/tests/cassettes"); +const PROVIDER_SOURCE_ROOT: &str = "tests/providers"; struct ProviderCassetteSuite { - provider: &'static str, - source_dir: &'static str, - wrapper_names: &'static [&'static str], + provider: String, + source_dir: PathBuf, } -const PROVIDER_CASSETTE_SUITES: &[ProviderCassetteSuite] = &[ - ProviderCassetteSuite { - provider: "openai", - source_dir: "tests/providers/openai/cassette", - wrapper_names: &[ - "with_openai_cassette", - "with_openai_lifecycle_cassette", - "with_openai_prompt_caching_cassette", - "with_openai_completions_prompt_caching_cassette", - "with_openai_turn_metadata_cassette", - "with_openai_cassette_bogus_key", - "with_openai_completions_cassette", - "with_openai_cassette_result", - "with_openai_completions_cassette_result", - "with_openai_vllm_cassette", - "with_local_reasoning_content_cassette", - "with_openai_refusal_cassette", - "with_openai_max_tokens_cassette", - "with_openai_image_params_cassette", - "with_openai_truncation_cassette", - "with_openai_chat_stream_logprobs_cassette_result", - "with_openai_tool_truncation_cassette_result", - "with_openai_tool_lifecycle_cassette_result", - "with_openai_terminal_metadata_cassette_result", - "with_openai_history_roundtrip_cassette_result", - "with_openai_transcription_cassette", - "with_openai_audio_cassette", - "with_openai_websocket_cassette", - ], - }, - ProviderCassetteSuite { - provider: "chatgpt", - source_dir: "tests/providers/chatgpt/cassette", - wrapper_names: &[ - "with_chatgpt_cassette", - "with_chatgpt_cassette_default_instructions", - "with_chatgpt_noninteractive_oauth_cassette", - ], - }, - ProviderCassetteSuite { - provider: "copilot", - source_dir: "tests/providers/copilot", - wrapper_names: &[ - "with_copilot_cassette", - "with_copilot_cassette_result", - "with_copilot_noninteractive_oauth_cassette", - ], - }, - ProviderCassetteSuite { - provider: "anthropic", - source_dir: "tests/providers/anthropic/cassette", - wrapper_names: &[ - "with_anthropic_cassette", - "with_anthropic_lifecycle_cassette", - "with_anthropic_turn_metadata_cassette", - "with_anthropic_cassette_result", - "with_anthropic_cassette_bogus_key", - "with_anthropic_files_cassette", - "with_anthropic_gateway_cassette", - "with_anthropic_stop_sequence_cassette", - "with_anthropic_empty_stop_cassette", - "with_anthropic_reasoning_usage_cassette", - ], - }, - ProviderCassetteSuite { - provider: "bedrock", - source_dir: "tests/providers/bedrock/cassette", - wrapper_names: &["with_bedrock_cassette"], - }, - ProviderCassetteSuite { - provider: "doubleword", - source_dir: "tests/providers/doubleword/cassette", - wrapper_names: &[ - "with_doubleword_prompt_caching_cassette", - "with_doubleword_cassette", - "with_doubleword_bogus_key_cassette", - "with_doubleword_cassette_result", - "with_doubleword_embedding_cassette", - ], - }, - ProviderCassetteSuite { - provider: "cohere", - source_dir: "tests/providers/cohere/cassette", - wrapper_names: &[ - "with_cohere_cassette", - "with_cohere_prompt_caching_cassette", - ], - }, - ProviderCassetteSuite { - provider: "venice", - source_dir: "tests/providers/venice/cassette", - wrapper_names: &[ - "with_venice_prompt_caching_cassette", - "with_venice_cassette", - "with_venice_cassette_result", - "with_venice_direct_cassette", - ], - }, - ProviderCassetteSuite { - provider: "gemini", - source_dir: "tests/providers/gemini/cassette", - wrapper_names: &[ - "with_gemini_prompt_caching_cassette", - "with_gemini_cassette", - "with_gemini_lifecycle_cassette", - "with_gemini_turn_metadata_cassette", - "with_gemini_cassette_bogus_key", - "with_gemini_code_execution_cassette", - "with_gemini_interactions_cassette", - "with_gemini_stream_terminal_cassette", - "with_gemini_thought_text_cassette", - ], - }, - ProviderCassetteSuite { - provider: "ollama", - source_dir: "tests/providers/ollama/cassette", - wrapper_names: &["with_ollama_cassette"], - }, - ProviderCassetteSuite { - provider: "llamacpp", - source_dir: "tests/providers/llamacpp/cassette", - wrapper_names: &[ - "with_llamacpp_cassette", - "with_llamacpp_cassette_result", - "with_llamacpp_bare_openai_cassette", - "with_llamacpp_embeddings_cassette", - "with_llamacpp_vision_cassette", - "with_llamacpp_small_context_cassette", - "with_llamacpp_no_jinja_cassette", - "with_llamacpp_rerank_cassette", - "with_llamacpp_pooling_none_cassette", - "with_llamacpp_causal_embeddings_cassette", - "with_llamacpp_competent_cassette", - "with_llamacpp_llama_family_cassette", - "with_llamacpp_mistral_family_cassette", - "with_llamacpp_gemma_family_cassette", - "with_llamacpp_prompt_caching_cassette", - "with_llamacpp_large_vision_cassette", - "with_llamacpp_raw_http_cassette", - "with_llamacpp_api_key_cassette", - "with_llamacpp_missing_api_key_cassette", - ], - }, - ProviderCassetteSuite { - provider: "xai", - source_dir: "tests/providers/xai", - wrapper_names: &[ - "with_xai_prompt_caching_cassette", - "with_xai_cassette", - "with_xai_cassette_bogus_key", - "with_xai_cassette_result", - ], - }, - ProviderCassetteSuite { - provider: "openrouter", - source_dir: "tests/providers/openrouter/cassette", - wrapper_names: &[ - "with_openrouter_prompt_caching_cassette", - "with_openrouter_cassette", - "with_openrouter_cassette_result", - "with_openrouter_cassette_bogus_key_result", - "with_openrouter_openai_cassette", - "with_openrouter_refusal_cassette", - "with_openrouter_usage_cassette", - "with_openrouter_stream_logprobs_cassette_result", - "with_openrouter_tool_truncation_cassette_result", - "with_openrouter_tool_lifecycle_cassette_result", - "with_openrouter_terminal_metadata_cassette_result", - "with_openrouter_history_roundtrip_cassette_result", - "with_openrouter_reasoning_tool_order_cassette_result", - ], - }, - ProviderCassetteSuite { - provider: "deepseek", - source_dir: "tests/providers/deepseek", - wrapper_names: &[ - "with_deepseek_prompt_caching_cassette", - "with_deepseek_cassette", - "with_deepseek_cassette_result", - "with_deepseek_cassette_bogus_key_result", - "with_deepseek_truncation_cassette_result", - "with_deepseek_block_order_cassette_result", - "with_deepseek_wire_shape_cassette_result", - "with_deepseek_followup_hunt_cassette_result", - "with_deepseek_stream_logprobs_cassette_result", - ], - }, - ProviderCassetteSuite { - provider: "groq", - source_dir: "tests/providers/groq", - wrapper_names: &[ - "with_groq_prompt_caching_cassette", - "with_groq_cassette_result", - "with_groq_cassette_bogus_key_result", - ], - }, - ProviderCassetteSuite { - provider: "mistral", - source_dir: "tests/providers/mistral", - wrapper_names: &[ - "with_mistral_embedding_cassette", - "with_mistral_prompt_caching_cassette", - "with_mistral_cassette_result", - "with_mistral_multimodal_cassette", - "with_mistral_cassette_bogus_key_result", - "with_mistral_capability_cassette", - "with_mistral_terminal_metadata_cassette_result", - "with_mistral_tool_truncation_cassette_result", - "with_mistral_tool_lifecycle_cassette_result", - "with_mistral_history_roundtrip_cassette_result", - "with_mistral_request_shape_cassette_result", - "with_mistral_logprobs_rejection_cassette_result", - ], - }, - ProviderCassetteSuite { - provider: "perplexity", - source_dir: "tests/providers/perplexity/cassette", - wrapper_names: &[ - "with_perplexity_cassette", - "with_perplexity_prompt_caching_cassette", - ], - }, - ProviderCassetteSuite { - provider: "mistralrs", - source_dir: "tests/providers/mistralrs/cassette", - wrapper_names: &[ - "with_mistralrs_cassette", - "with_mistralrs_completions_cassette", - "with_mistralrs_raw_cassette", - ], - }, -]; +/// Every provider directory under `tests/providers/` that has a `tests/.rs` +/// binary. Directories without a binary are reported by the caller if they +/// contain cassette wrapper calls (see `collect_expected_cassette_paths`). +fn discovered_suites() -> Vec { + provider_source_dirs() + .into_iter() + .filter(|(provider, _)| repo_path(&format!("tests/{provider}.rs")).is_file()) + .map(|(provider, source_dir)| ProviderCassetteSuite { + provider, + source_dir, + }) + .collect() +} + +/// `(name, path)` for every directory directly under `tests/providers/`, sorted. +fn provider_source_dirs() -> Vec<(String, PathBuf)> { + let root = repo_path(PROVIDER_SOURCE_ROOT); + let mut dirs = Vec::new(); + for entry in fs::read_dir(&root).expect("tests/providers should be readable") { + let entry = entry.expect("tests/providers entry should be readable"); + let path = entry.path(); + if path.is_dir() { + dirs.push((entry.file_name().to_string_lossy().into_owned(), path)); + } + } + dirs.sort(); + dirs +} + +/// The naming convention a cassette wrapper follows. Discovery is by call +/// site, not by definition, because some providers define their wrappers +/// through macros. +fn is_cassette_wrapper_name(name: &str) -> bool { + name.starts_with("with_") && name.contains("cassette") +} #[test] fn cassettes_do_not_contain_obvious_secrets() { @@ -268,23 +88,22 @@ fn cassettes_do_not_contain_obvious_secrets() { // binary, before anything is skipped: // // * every top-level entry under `tests/cassettes` must be a directory - // named after a suite registered in `PROVIDER_CASSETTE_SUITES` — a - // stray file or an unregistered provider directory fails everywhere - // rather than silently escaping the scan; - // * every registered suite's `tests/.rs` must include this - // module — so each registered directory is provably scanned by - // exactly the binary that owns it, and adding a suite without wiring - // the scan into its binary fails everywhere too; - // * every registered provider name must be a valid crate identifier — + // named after a discovered suite (a `tests/providers//` + // directory with a `tests/.rs` binary) — a stray file or a + // cassette directory no binary owns fails everywhere rather than + // silently escaping the scan; + // * every discovered suite's `tests/.rs` must include this + // module — so each cassette directory is provably scanned by exactly + // the binary that owns it, and adding a provider without wiring the + // scan into its binary fails everywhere too; + // * every provider name must be a valid crate identifier — // `env!("CARGO_CRATE_NAME")` mangles hyphens to underscores, so a // hyphenated provider would resolve `own_dir` to a path that never // exists and skip its own scan without a single failure. let mut failures = Vec::new(); - let registered: BTreeSet<&str> = PROVIDER_CASSETTE_SUITES - .iter() - .map(|suite| suite.provider) - .collect(); + let suites = discovered_suites(); + let discovered: BTreeSet<&str> = suites.iter().map(|suite| suite.provider.as_str()).collect(); for entry in fs::read_dir(root).expect("cassette root should be readable") { let entry = entry.expect("cassette root entry should be readable"); let name = entry.file_name(); @@ -294,15 +113,15 @@ fn cassettes_do_not_contain_obvious_secrets() { "tests/cassettes/{name} is not a provider directory; loose files under the \ cassette root are scanned by no binary" )); - } else if !registered.contains(name.as_str()) { + } else if !discovered.contains(name.as_str()) { failures.push(format!( - "tests/cassettes/{name} has no PROVIDER_CASSETTE_SUITES entry, so no test \ - binary scans it for secrets — register it in \ - tests/common/cassette_safety.rs" + "tests/cassettes/{name} has no tests/providers/{name}/ directory with a \ + tests/{name}.rs binary, so no test binary scans it for secrets — add both, \ + or delete the cassette directory" )); } } - for suite in PROVIDER_CASSETTE_SUITES { + for suite in &suites { if !suite .provider .chars() @@ -315,6 +134,12 @@ fn cassettes_do_not_contain_obvious_secrets() { suite.provider )); } + // A provider with no cassette directory has nothing to scan; the + // moment `tests/cassettes/` appears, its binary must compile + // this module or every binary fails. + if !root.join(&suite.provider).is_dir() { + continue; + } let binary_source = repo_path(&format!("tests/{}.rs", suite.provider)); if !binary_compiles_cassette_scan(&binary_source) { failures.push(format!( @@ -420,21 +245,16 @@ fn collect_expected_cassette_paths() -> (BTreeSet, Vec) { let mut expected = BTreeSet::new(); let mut failures = Vec::new(); - for suite in PROVIDER_CASSETTE_SUITES { - let source_dir = repo_path(suite.source_dir); - if !source_dir.exists() { - failures.push(format!( - "cassette source directory does not exist: {}", - display_repo_path(&source_dir) - )); - continue; - } + let suites = discovered_suites(); + let providers: BTreeSet<&str> = suites.iter().map(|suite| suite.provider.as_str()).collect(); - for source_file in collect_rust_files(&source_dir) { - match cassette_scenarios_in_file(&source_file, suite.wrapper_names) { + for suite in &suites { + for source_file in collect_rust_files(&suite.source_dir) { + match cassette_scenarios_in_file(&source_file, &suite.provider, &providers) { Ok(scenarios) => { for scenario in scenarios { - expected.insert(crate::cassettes::cassette_path(suite.provider, &scenario)); + expected + .insert(crate::cassettes::cassette_path(&suite.provider, &scenario)); } } Err(error) => failures.push(error), @@ -442,6 +262,25 @@ fn collect_expected_cassette_paths() -> (BTreeSet, Vec) { } } + // A provider directory with no binary is scanned by nothing: any cassette + // wrapper call in it would register scenarios nowhere. + for (provider, source_dir) in provider_source_dirs() { + if providers.contains(provider.as_str()) { + continue; + } + for source_file in collect_rust_files(&source_dir) { + if let Ok(scenarios) = cassette_scenarios_in_file(&source_file, &provider, &providers) + && !scenarios.is_empty() + { + failures.push(format!( + "{} calls cassette wrappers but tests/providers/{provider}/ has no \ + tests/{provider}.rs binary, so its scenarios are checked by nothing", + display_repo_path(&source_file) + )); + } + } + } + (expected, failures) } @@ -507,7 +346,8 @@ fn binary_compiles_cassette_scan(source: &Path) -> bool { fn cassette_scenarios_in_file( path: &Path, - wrapper_names: &[&'static str], + provider: &str, + providers: &BTreeSet<&str>, ) -> Result, String> { let contents = fs::read_to_string(path) .map_err(|error| format!("{} should be readable: {error}", display_repo_path(path)))?; @@ -515,7 +355,8 @@ fn cassette_scenarios_in_file( .map_err(|error| format!("{} should parse as Rust: {error}", display_repo_path(path)))?; let mut visitor = CassetteScenarioVisitor { path, - wrapper_names, + provider, + providers, scenarios: Vec::new(), failures: Vec::new(), }; @@ -530,7 +371,8 @@ fn cassette_scenarios_in_file( struct CassetteScenarioVisitor<'a> { path: &'a Path, - wrapper_names: &'a [&'static str], + provider: &'a str, + providers: &'a BTreeSet<&'a str>, scenarios: Vec, failures: Vec, } @@ -543,14 +385,30 @@ impl<'ast, 'a> Visit<'ast> for CassetteScenarioVisitor<'a> { if node.attrs.iter().any(|attr| attr.path().is_ident("ignore")) { return; } + // A wrapper's own body forwards its scenario to a base wrapper; that is + // a definition, not a call site, and its argument is a variable. + if is_cassette_wrapper_name(&node.sig.ident.to_string()) { + return; + } visit::visit_item_fn(self, node); } fn visit_expr_call(&mut self, node: &'ast ExprCall) { if let Some(wrapper_name) = cassette_wrapper_name(node) - && self.wrapper_names.contains(&wrapper_name.as_str()) + && is_cassette_wrapper_name(&wrapper_name) { + if let Some(other) = + foreign_provider_in_wrapper(&wrapper_name, self.provider, self.providers) + { + self.failures.push(format!( + "{} calls {wrapper_name}, which names provider {other:?}, from provider \ + {:?}'s directory; a wrapper is called only from its own provider's \ + tests/providers// directory", + display_repo_path(self.path), + self.provider + )); + } match node.args.first() { Some(expr) => match cassette_scenario_value(expr) { Some(scenario) => self.scenarios.push(scenario), @@ -570,6 +428,27 @@ impl<'ast, 'a> Visit<'ast> for CassetteScenarioVisitor<'a> { } } +/// The provider a wrapper name points at, when it is not the directory's own. +/// Segments of the identifier are compared against the discovered provider +/// names: `with_openrouter_openai_cassette` in openrouter's directory names its +/// own provider and passes; `with_openai_cassette` in openrouter's directory +/// does not. Names that mention no provider (`with_local_reasoning_content_cassette`) +/// pass anywhere. +fn foreign_provider_in_wrapper<'p>( + wrapper_name: &str, + provider: &str, + providers: &BTreeSet<&'p str>, +) -> Option<&'p str> { + let segments: Vec<&str> = wrapper_name.split('_').collect(); + if segments.contains(&provider) { + return None; + } + providers + .iter() + .copied() + .find(|candidate| segments.contains(candidate)) +} + fn cassette_scenario_value(expr: &Expr) -> Option { match expr { Expr::Lit(ExprLit { diff --git a/xtask/src/main.rs b/xtask/src/main.rs index 73ec610795..19dca5914f 100644 --- a/xtask/src/main.rs +++ b/xtask/src/main.rs @@ -18,11 +18,13 @@ //! cargo xtask generate-provider-aliases # rewrite the file //! cargo xtask generate-provider-aliases --check # fail if it would change //! cargo xtask check-test-layout # fail on inline `mod tests { }` +//! cargo xtask check-sorted-blocks # fail on an out-of-order marked list //! ``` mod aliases; mod reachable; mod rustdoc; +mod sorted_blocks; mod test_layout; use std::path::{Path, PathBuf}; @@ -39,6 +41,7 @@ fn main() -> ExitCode { let result = match task.as_deref() { Some("generate-provider-aliases") => generate_provider_aliases(check), Some("check-test-layout") => test_layout::check(&workspace_root()), + Some("check-sorted-blocks") => sorted_blocks::check(&workspace_root()), Some(other) => Err(format!("unknown task {other:?}\n{USAGE}")), None => Err(format!("no task given\n{USAGE}")), }; @@ -60,6 +63,8 @@ tasks: from rig-core's rustdoc output check-test-layout fail if any crates/*/src file has an inline test-gated `mod x { }` instead of `mod x;` + check-sorted-blocks fail if a list between `sorted: start` and + `sorted: end` markers is not in byte order "; fn generate_provider_aliases(check: bool) -> Result<(), String> { diff --git a/xtask/src/sorted_blocks.rs b/xtask/src/sorted_blocks.rs new file mode 100644 index 0000000000..aeaacbd41a --- /dev/null +++ b/xtask/src/sorted_blocks.rs @@ -0,0 +1,185 @@ +//! `check-sorted-blocks`: lists that grow with every provider, crate or +//! dependency stay in byte order, so two PRs that each add one entry touch +//! different lines instead of the same hunk. +//! +//! A block is the text between a line containing `sorted: start` and one +//! containing `sorted: end` (inside whatever comment syntax the file uses: +//! `#`, `//`, or ``). Inside a block every entry's key must be +//! greater than the previous entry's, comparing ASCII-lowercased bytes (so a +//! Markdown table reads `ScyllaDB, SQLite`, not `SQLite, ScyllaDB`) with the +//! original bytes as the tie-break. An entry is one of: +//! +//! * a TOML key (`name = ...`), with a multi-line inline table or array +//! counted as part of the entry until its brackets balance; +//! * a Rust `mod name;` / `pub mod name;` declaration; +//! * a Markdown table row, keyed on its first cell. +//! +//! Blank lines, comment lines and the `| --- |` table separator are skipped. +//! The check is textual on purpose: it is about line order, which a TOML or +//! Markdown parser would discard. Every file listed in `CHECKED_FILES` must +//! contain at least one block, so removing the markers cannot pass silently. + +use std::path::Path; + +/// Files that carry `sorted: start` / `sorted: end` blocks, relative to the +/// workspace root. +const CHECKED_FILES: &[&str] = &[ + "Cargo.toml", + "crates/rig-core/src/providers/mod.rs", + "README.md", +]; + +pub(crate) fn check(workspace: &Path) -> Result<(), String> { + let mut failures = Vec::new(); + let mut blocks = 0; + for relative in CHECKED_FILES { + let path = workspace.join(relative); + let source = std::fs::read_to_string(&path) + .map_err(|error| format!("could not read {}: {error}", path.display()))?; + let found = check_source(relative, &source, &mut failures); + if found == 0 { + failures.push(format!( + "{relative}: no `sorted: start` / `sorted: end` block found; the markers were \ + removed or the file no longer carries a sorted list" + )); + } + blocks += found; + } + if failures.is_empty() { + println!( + "ok: {blocks} sorted block(s) in {} file(s) are in byte order", + CHECKED_FILES.len() + ); + return Ok(()); + } + let mut message = String::from( + "sorted blocks out of order; reorder the entries so each key is greater than the one \ + before it:\n", + ); + for failure in &failures { + message.push_str(" "); + message.push_str(failure); + message.push('\n'); + } + Err(message) +} + +/// Check every block in `source`, appending failures; returns how many blocks +/// were found (a block that never closes counts and is reported). +pub(crate) fn check_source(name: &str, source: &str, failures: &mut Vec) -> usize { + let mut blocks = 0; + let mut open: Option<(usize, Option<(usize, String)>)> = None; + let mut depth: i32 = 0; + + let lines: Vec<&str> = source.lines().collect(); + for (index, line) in lines.iter().enumerate() { + let number = index + 1; + let trimmed = line.trim(); + + if trimmed.contains("sorted: start") { + if let Some((start, _)) = open { + failures.push(format!( + "{name}:{number}: `sorted: start` while the block from line {start} is still open" + )); + } + open = Some((number, None)); + depth = 0; + blocks += 1; + continue; + } + if trimmed.contains("sorted: end") { + if open.take().is_none() { + failures.push(format!( + "{name}:{number}: `sorted: end` without a matching start" + )); + } + continue; + } + let Some((_, previous)) = open.as_mut() else { + continue; + }; + + if depth > 0 { + depth += bracket_delta(line); + continue; + } + let next = lines.get(index + 1).map(|next| next.trim()).unwrap_or(""); + let Some(key) = entry_key(trimmed, next) else { + continue; + }; + depth = bracket_delta(line).max(0); + + if let Some((previous_line, previous_key)) = previous + && sort_key(&key) <= sort_key(previous_key) + { + failures.push(format!( + "{name}:{number}: `{key}` sorts before `{previous_key}` (line {previous_line})" + )); + } + *previous = Some((number, key)); + } + + if let Some((start, _)) = open { + failures.push(format!("{name}:{start}: `sorted: start` is never closed")); + } + blocks +} + +/// The key an entry line sorts on, or `None` for a line that is not an entry. +/// `next` is the following line: a table row directly above the `| --- |` +/// separator is the header, not an entry. +fn entry_key(trimmed: &str, next: &str) -> Option { + if trimmed.is_empty() + || trimmed.starts_with('#') + || trimmed.starts_with("//") + || trimmed.starts_with("\n| Name | X |\n| --- | --- |\n| B | 1 |\n| A | 2 |\n\n"; + assert_eq!(failures_for(md), vec!["f:5: `A` sorts before `B` (line 4)"]); +} + +#[test] +fn unclosed_block_is_reported_and_counted() { + let mut failures = Vec::new(); + let blocks = check_source("f", "# sorted: start\na = 1\n", &mut failures); + assert_eq!(blocks, 1); + assert_eq!(failures, vec!["f:1: `sorted: start` is never closed"]); +} + +#[test] +fn duplicate_keys_fail() { + assert_eq!( + failures_for("# sorted: start\na = 1\na = 2\n# sorted: end\n").len(), + 1 + ); +}