From 225bda563196cdbbef247c5f7ef301f3b52e1c2c Mon Sep 17 00:00:00 2001 From: Owen Qwen Date: Sun, 21 Jun 2026 11:55:27 -0500 Subject: [PATCH] Add model registry config and autocomplete --- PROVIDER_MODEL_CONFIG_PLAN.md | 271 ++++++++++++ README.md | 25 +- docs/README.md | 2 + docs/configuration.md | 120 ++++++ src/agent.rs | 15 +- src/app.rs | 165 +++++++- src/check.rs | 213 ++++++++++ src/cli.rs | 11 +- src/config.rs | 646 +++++++++++++++++++++++++++-- src/lib.rs | 1 + src/providers/openai_compatible.rs | 23 +- tests/config_tests.rs | 189 +++++++++ 12 files changed, 1607 insertions(+), 74 deletions(-) create mode 100644 PROVIDER_MODEL_CONFIG_PLAN.md create mode 100644 docs/configuration.md create mode 100644 src/check.rs create mode 100644 tests/config_tests.rs diff --git a/PROVIDER_MODEL_CONFIG_PLAN.md b/PROVIDER_MODEL_CONFIG_PLAN.md new file mode 100644 index 0000000..2198317 --- /dev/null +++ b/PROVIDER_MODEL_CONFIG_PLAN.md @@ -0,0 +1,271 @@ +# Provider and Model Configuration Implementation Plan + +## Goals + +- Add `~/.cass/providers.json` as the source of truth for provider definitions. +- Add `~/.cass/models.json` for optional model metadata such as context length and max output tokens. +- Support API keys as either literal strings or environment-variable references in the form `"$PROVIDER_API_KEY"`. +- Document the config files well enough that a user can manually edit them or ask Cass, in full-access mode, to add/update providers and models by reading the bundled docs. +- Add `cass check` to validate config JSON syntax, schema, references, and basic operational readiness. + +## Proposed file layout + +Cass-managed/user-editable files under `~/.cass`: + +- `config.json`: user preferences and default provider/model references only (no provider connection details or model metadata). +- `providers.json`: provider registry. +- `models.json`: model metadata registry. +- `docs/`: bundled read-only docs installed on startup. + +Keep `config.json` for user preferences, but move provider connection details out of it. Continue to accept the existing `provider`, `model`, `base_url`, and `api_key_env` fields as a backward-compatible legacy path, with docs steering users to `default_provider`/`default_model` plus the new registry files. + +## Proposed schemas + +### `config.json` + +```json +{ + "default_model": "accounts/fireworks/models/qwen3p7-plus", + "default_access_mode": "read-only", + "context_message_limit": 80, + "model_tool_result_limit": 24000, + "ui_tool_result_limit": 4000 +} +``` + +Backward-compatible deprecated fields to continue accepting for now: + +```json +{ + "base_url": "https://api.fireworks.ai/inference/v1", + "api_key_env": "FIREWORKS_API_KEY" +} +``` + +### `providers.json` + +Use an array to make manual edits straightforward and preserve room for provider-specific fields. + +```json +{ + "providers": [ + { + "id": "fireworks", + "name": "Fireworks", + "kind": "openai-compatible", + "base_url": "https://api.fireworks.ai/inference/v1", + "api_key": "$FIREWORKS_API_KEY", + "default_model": "accounts/fireworks/models/qwen3p7-plus", + "models": [ + "accounts/fireworks/models/qwen3p7-plus" + ] + } + ] +} +``` + +Initial provider fields: + +- `id` required, unique stable identifier referenced by config defaults and models. +- `name` optional display name. +- `kind` required; initially only `"openai-compatible"` is supported. +- `base_url` required for `openai-compatible`. +- `api_key` required string. If it starts with `$`, resolve the remaining text as an environment variable name. Otherwise use it as a literal API key. +- `default_model` optional model to use when `config.json` omits `model`. +- `models` optional list of model ids associated with the provider. + +### `models.json` + +```json +{ + "models": [ + { + "id": "accounts/fireworks/models/qwen3p7-plus", + "provider": "fireworks", + "display_name": "Qwen 3p7 Plus", + "context_length": 262144, + "max_output_tokens": 32768, + "supports_tools": true, + "supports_streaming": true + } + ] +} +``` + +Initial model fields: + +- `id` required model identifier sent to the provider. +- `provider` required provider id. Validate that it exists in `providers.json`. +- `display_name` optional human-friendly name. +- `context_length` optional positive integer. +- `max_output_tokens` optional positive integer. +- `supports_tools` optional boolean, defaults to `true`. +- `supports_streaming` optional boolean, defaults to `true`. + +Deduplicate models by `(provider, id)`. + +## Runtime behavior + +1. On startup, ensure `~/.cass` exists as today. +2. If `providers.json` is missing, create a default file containing the current Fireworks provider with `"api_key": "$FIREWORKS_API_KEY"`. +3. If `models.json` is missing, create a default file containing metadata for the current default Fireworks model. +4. Load `config.json`, `providers.json`, and `models.json`. +5. Resolve the active model: + - CLI `--model` overrides everything. + - Else `config.default_model` (or legacy `config.model`). + - Else selected provider `default_model`. + - Else current built-in default. +6. Resolve the active provider: + - Prefer `config.default_provider` when present. + - Else infer from the selected model when `models.json` or provider model lists identify exactly one provider. + - Else use the Fireworks default provider when available. + - If legacy `base_url`/`api_key_env` are present and no registry provider is selected, synthesize a legacy `openai-compatible` provider for backward compatibility. + - Otherwise fail with a clear config error that suggests running `cass check`. +7. Resolve API key: + - `"$NAME"` means read env var `NAME`. + - Empty env var names are invalid. + - Literal strings are passed through unchanged. + - Never print literal API key values in errors or check output. +8. Construct `OpenAiCompatibleProvider` from the resolved provider settings instead of raw `model/base_url/api_key_env` fields. + +## CLI changes + +Refactor `src/cli.rs` to support subcommands while preserving existing invocation forms: + +```text +cass [--model MODEL] [--base-url URL] [--api-key-env ENV] [--cwd PATH] +cass --resume +cass --resume +cass check +``` + +Implementation sketch: + +- Add `Command::Check` as an optional subcommand. +- Keep existing top-level flags for chat mode. +- In `app::run`, parse config, then if command is `Check`, run config checks and exit without entering the TUI. +- Exit code `0` when checks pass, non-zero when any error exists. + +## `cass check` behavior + +Initial check scope: config files only. + +Checks: + +- `config.json`, `providers.json`, and `models.json` parse as valid JSON when present. +- Files match expected schema/types and reject unknown required shapes. +- Required provider fields are present. +- Provider ids are unique. +- Provider `kind` is supported. +- `openai-compatible` providers have valid `base_url` and `api_key` strings. +- `$ENV_VAR` API key references have non-empty names. +- Active provider can be resolved from `default_provider`, selected model metadata, provider model lists, or a valid legacy fallback. +- Provider `default_model` and `models` entries can be matched against `models.json` when metadata exists. +- Model ids are unique within their provider scope and every model has a provider. +- Model numeric metadata is positive. +- Model `provider` references exist. +- Active provider API key resolves. For inactive providers, missing env vars should be warnings, not errors. + +Output format example: + +```text +Cass config check +✓ ~/.cass/config.json: valid +✓ ~/.cass/providers.json: valid (1 provider) +✓ ~/.cass/models.json: valid (1 model) +✓ active provider: fireworks +✓ active model: accounts/fireworks/models/qwen3p7-plus +✓ api key: FIREWORKS_API_KEY is set + +All checks passed. +``` + +On errors, print each error with the file path and JSON path when possible. + +## Code changes + +### Config loading + +- Extend `src/config.rs` or split into `src/config/` modules if it becomes large. +- Add structs: + - `ProvidersFile` + - `ProviderDefinition` + - `ModelsFile` + - `ModelDefinition` + - `ResolvedProviderConfig` + - `ResolvedModelMetadata` +- Add helper functions: + - `load_or_create_default_provider_registry(root)` + - `load_or_create_default_model_registry(root)` + - `resolve_api_key(spec: &str) -> Result` + - `redact_api_key_for_display(spec: &str) -> String` + - `validate_config_files(root, cli_overrides) -> CheckReport` +- Update `Config` to include resolved provider details, while keeping old fields during migration if needed. + +### Provider construction + +- Change `OpenAiCompatibleProvider::new` to accept a resolved settings struct: + +```rust +pub struct OpenAiCompatibleSettings { + pub model: String, + pub base_url: String, + pub api_key: String, +} +``` + +- Update `agent::run_turn` to create the provider from `settings.config.resolved_provider`. +- Keep current request body behavior initially. Store model metadata for future context-management and output-token use. + +### Check command + +- Add a `src/check.rs` module for report types and rendering. +- `CheckReport` should hold errors and warnings separately. +- `cass check` should not start the TUI or create a conversation. +- Prefer deterministic output for tests. + +### Documentation + +Add bundled docs: + +- `docs/configuration.md`: full schema examples, env-var API key behavior, manual editing steps, and examples prompts for asking Cass to add a provider/model. +- Update `docs/README.md` to link to `configuration.md`. +- Update top-level `README.md` Configure and Usage sections. + +Include a user-facing example: + +```text +Run cass in full-access mode and ask: +"Read ~/.cass/docs/configuration.md, then add an OpenAI-compatible provider named Together using TOGETHER_API_KEY and add model metadata for meta-llama/Llama-3.1-70B-Instruct-Turbo." +``` + +## Tests + +Add tests for: + +- Default `providers.json` and `models.json` creation in a temp Cass root. +- Loading current legacy `config.json` with `base_url` and `api_key_env` still works. +- `$ENV_VAR` API key resolution succeeds when set and fails clearly when missing. +- Literal API key strings are accepted and never appear in check output. +- Duplicate provider ids fail validation. +- Invalid JSON syntax fails validation with file path. +- Invalid provider/model references fail validation. +- `cass check` returns success for defaults and failure for invalid files. +- Docs install still includes the new configuration document. + +## Migration/backward compatibility + +- Do not break existing users with only `~/.cass/config.json`. +- Continue honoring `--base-url` and `--api-key-env`; internally convert them to provider overrides for the session. +- Mark `base_url` and `api_key_env` as deprecated in docs, but do not remove them yet. +- If both `providers.json` and legacy connection fields are present, `providers.json` wins unless CLI overrides are used. + +## Suggested implementation order + +1. Add registry structs, defaults, loading, API-key resolution, and validation helpers. +2. Add `cass check` CLI plumbing and report output. +3. Wire resolved provider settings into `OpenAiCompatibleProvider` and `agent::run_turn`. +4. Bootstrap missing `providers.json` and `models.json` with defaults. +5. Add docs and README updates. +6. Add/adjust tests. +7. Run `cargo fmt`, `cargo test`, and `cass check`. diff --git a/README.md b/README.md index 56fd0c9..e0f8c5d 100644 --- a/README.md +++ b/README.md @@ -17,11 +17,11 @@ cassady ## Configure -By default Cass uses: +By default Cass creates `~/.cass/providers.json` and `~/.cass/models.json` with Fireworks configured: - base URL: `https://api.fireworks.ai/inference/v1` - model: `accounts/fireworks/models/qwen3p7-plus` -- API key env var: `FIREWORKS_API_KEY` +- API key: `"$FIREWORKS_API_KEY"` Set your key: @@ -29,21 +29,26 @@ Set your key: export FIREWORKS_API_KEY=... ``` -Optional config lives at `~/.cass/config.json`: +User preferences live at `~/.cass/config.json`: ```json { - "provider": "openai-compatible", - "model": "accounts/fireworks/models/qwen3p7-plus", - "base_url": "https://api.fireworks.ai/inference/v1", - "api_key_env": "FIREWORKS_API_KEY", + "default_model": "accounts/fireworks/models/qwen3p7-plus", "default_access_mode": "read-only" } ``` +Provider connection details belong in `~/.cass/providers.json`. Model metadata belongs in `~/.cass/models.json`. API keys may be literal strings or environment-variable references like `"$FIREWORKS_API_KEY"`. + +Validate config with: + +```sh +cass check +``` + Extra global instructions can be placed in `~/.cass/global.md`. -Bundled documentation from this build is embedded into the binary and installed to `~/.cass/docs` on startup. +Bundled documentation from this build is embedded into the binary and installed to `~/.cass/docs` on startup. See `~/.cass/docs/configuration.md` for full configuration docs. ## Usage @@ -51,6 +56,7 @@ Bundled documentation from this build is embedded into the binary and installed cass [--model MODEL] [--base-url URL] [--api-key-env ENV] [--cwd PATH] cass --resume cass --resume +cass check ``` `cass --resume` without an ID lists chats for the current directory. @@ -69,7 +75,8 @@ cass --resume ## Commands -- `/model `: switch the model for future turns +- `cass check`: validate Cass config files +- `/model `: switch the model for future turns; model autocomplete lists entries from `~/.cass/models.json` - `/resume `: resume a saved chat; chat autocomplete lists chats for the current directory - `/status`: show current chat status diff --git a/docs/README.md b/docs/README.md index e77d297..05d1092 100644 --- a/docs/README.md +++ b/docs/README.md @@ -3,3 +3,5 @@ These docs are embedded into the `cass` binary at build time and installed to `~/.cass/docs` when Cass starts. Cass tools may list, search, and read this directory. Mutating tools are blocked from writing here, even in full-access mode. + +- [Configuration](configuration.md): `config.json`, `providers.json`, `models.json`, and `cass check`. diff --git a/docs/configuration.md b/docs/configuration.md new file mode 100644 index 0000000..e7fd988 --- /dev/null +++ b/docs/configuration.md @@ -0,0 +1,120 @@ +# Configuration + +Cass reads user-editable config files from `~/.cass`. + +- `config.json`: user preferences, such as the default model and access mode. +- `providers.json`: provider connection definitions. +- `models.json`: model metadata. + +Cass creates `providers.json` and `models.json` automatically if they are missing. The default provider is Fireworks. + +## `config.json` + +`config.json` should contain preferences only. Provider connection details belong in `providers.json`; model metadata belongs in `models.json`. + +Example: + +```json +{ + "default_model": "accounts/fireworks/models/qwen3p7-plus", + "default_access_mode": "read-only", + "context_message_limit": 80, + "model_tool_result_limit": 24000, + "ui_tool_result_limit": 4000 +} +``` + +Fields: + +- `default_provider`: optional provider id from `providers.json`. If omitted, Cass infers the provider from `default_model` when possible. +- `default_model`: optional model id to use by default. +- `default_access_mode`: `"read-only"` or `"full-access"`. +- `context_message_limit`: optional number of recent non-system messages sent to the model. +- `model_tool_result_limit`: optional max bytes of tool output sent back to the model. +- `ui_tool_result_limit`: optional max bytes of tool output shown in the UI unless full output is toggled. + +Deprecated compatibility fields from older Cass versions are still accepted: `provider`, `model`, `base_url`, and `api_key_env`. Prefer moving provider connection details to `providers.json`. + +## `providers.json` + +Example: + +```json +{ + "providers": [ + { + "id": "fireworks", + "name": "Fireworks", + "kind": "openai-compatible", + "base_url": "https://api.fireworks.ai/inference/v1", + "api_key": "$FIREWORKS_API_KEY", + "default_model": "accounts/fireworks/models/qwen3p7-plus", + "models": [ + "accounts/fireworks/models/qwen3p7-plus" + ] + } + ] +} +``` + +Fields: + +- `id`: required unique provider id. +- `name`: optional display name. +- `kind`: required provider kind. Currently only `"openai-compatible"` is supported. +- `base_url`: required OpenAI-compatible API base URL. +- `api_key`: required string. Use either a literal key or an environment-variable reference like `"$FIREWORKS_API_KEY"`. +- `default_model`: optional model id to use when no default model is configured. +- `models`: optional list of model ids associated with this provider. + +Only strings that start with `$` are resolved as environment variables. Cass does not expand partial strings or `${NAME}` syntax. + +## `models.json` + +Example: + +```json +{ + "models": [ + { + "id": "accounts/fireworks/models/qwen3p7-plus", + "provider": "fireworks", + "display_name": "Qwen 3p7 Plus", + "context_length": 262144, + "max_output_tokens": 32768, + "supports_tools": true, + "supports_streaming": true + } + ] +} +``` + +Fields: + +- `id`: required model id sent to the provider. +- `provider`: required provider id from `providers.json`. +- `display_name`: optional human-friendly name. +- `context_length`: optional positive integer. +- `max_output_tokens`: optional positive integer. +- `supports_tools`: optional boolean, defaults to `true`. +- `supports_streaming`: optional boolean, defaults to `true`. + +## Check configuration + +Run: + +```sh +cass check +``` + +This validates JSON syntax, expected schema, duplicate provider/model ids, model/provider references, active provider/model resolution, and API key environment-variable availability. Missing API keys for inactive providers are warnings; a missing active provider API key is an error. + +## Ask Cass to edit config + +Run Cass in full-access mode and ask it to read these docs before editing: + +```text +Read ~/.cass/docs/configuration.md, then add an OpenAI-compatible provider named Together using TOGETHER_API_KEY and add model metadata for meta-llama/Llama-3.1-70B-Instruct-Turbo. +``` + +After Cass edits the files, run `cass check`. diff --git a/src/agent.rs b/src/agent.rs index f2e1f40..306f562 100644 --- a/src/agent.rs +++ b/src/agent.rs @@ -2,7 +2,7 @@ use crate::access::AccessMode; use crate::config::Config; use crate::conversation::{now_ts, Conversation, Record, StoredToolCall}; use crate::prompt; -use crate::providers::openai_compatible::OpenAiCompatibleProvider; +use crate::providers::openai_compatible::{OpenAiCompatibleProvider, OpenAiCompatibleSettings}; use crate::providers::types::ModelMessage; use crate::tools::{self, ToolContext}; use anyhow::Result; @@ -46,18 +46,19 @@ pub async fn run_turn( ts: now_ts(), })?; - let provider = match OpenAiCompatibleProvider::new( - settings.config.model.clone(), - settings.config.base_url.clone(), - settings.config.api_key_env.clone(), - ) { - Ok(p) => p, + let api_key = match settings.config.resolved_api_key() { + Ok(api_key) => api_key, Err(err) => { let _ = tx.send(AgentEvent::Status(err.to_string())); let _ = tx.send(AgentEvent::TurnFinished); return Ok(conversation); } }; + let provider = OpenAiCompatibleProvider::new(OpenAiCompatibleSettings { + model: settings.config.model.clone(), + base_url: settings.config.active_provider.base_url.clone(), + api_key, + }); let docs_dir = settings.config.docs_dir(); let tool_ctx = ToolContext { diff --git a/src/app.rs b/src/app.rs index dddad2b..656b6f8 100644 --- a/src/app.rs +++ b/src/app.rs @@ -1,6 +1,6 @@ use crate::agent::{self, AgentEvent, AgentSettings}; -use crate::cli; -use crate::config::Config; +use crate::cli::{self, Command}; +use crate::config::{Config, ModelDefinition}; use crate::conversation::{self, Conversation}; use crate::prompt; use crate::ui::autofill::{AutoFillItem, AutoFillMenu}; @@ -17,6 +17,15 @@ use tokio::task::JoinHandle; pub async fn run() -> Result<()> { let cli = cli::parse(); + if matches!(cli.command, Some(Command::Check)) { + let report = crate::check::run(&cli)?; + print!("{}", report.render()); + if report.has_errors() { + std::process::exit(1); + } + return Ok(()); + } + let config = Config::load(&cli)?; let cwd = resolve_cwd(cli.cwd.clone())?; @@ -533,6 +542,10 @@ fn build_autofill( return Ok(Some(menu)); } + if let Some(menu) = model_autofill(input, selected, config)? { + return Ok(Some(menu)); + } + resume_chat_autofill(input, selected, config, cwd) } @@ -566,6 +579,83 @@ fn command_autofill(input: &str, selected: usize) -> Option { } } +fn model_autofill(input: &str, selected: usize, config: &Config) -> Result> { + let Some(rest) = input.strip_prefix("/model") else { + return Ok(None); + }; + if rest.is_empty() || !rest.chars().next().is_some_and(|c| c.is_whitespace()) { + return Ok(None); + } + + let arg = rest.trim_start_matches(char::is_whitespace); + if arg.chars().any(char::is_whitespace) { + return Ok(None); + } + let replacement_start = input.len() - arg.len(); + let query = arg.to_ascii_lowercase(); + + let models = crate::config::load_or_create_default_model_registry(&config.root)?; + if !arg.is_empty() && models.models.iter().any(|model| model.id == arg) { + return Ok(None); + } + + let mut items = Vec::new(); + for model in models.models { + if model_matches(&model, &query) { + let id = model.id.clone(); + let detail = model_detail(&model, &config.model); + items.push(AutoFillItem::new(id.clone(), id).with_detail(detail)); + } + } + + if items.is_empty() { + Ok(None) + } else { + Ok(Some( + AutoFillMenu::new("Models", replacement_start, input.len(), items) + .with_selected(selected), + )) + } +} + +fn model_matches(model: &ModelDefinition, query: &str) -> bool { + if query.is_empty() { + return true; + } + model.id.to_ascii_lowercase().contains(query) + || model.provider.to_ascii_lowercase().contains(query) + || model + .display_name + .as_ref() + .is_some_and(|name| name.to_ascii_lowercase().contains(query)) +} + +fn model_detail(model: &ModelDefinition, current_model: &str) -> String { + let mut parts = Vec::new(); + if model.id == current_model { + parts.push("current".to_string()); + } + if let Some(name) = &model.display_name { + if !name.trim().is_empty() && name != &model.id { + parts.push(name.clone()); + } + } + parts.push(format!("provider {}", model.provider)); + if let Some(context_length) = model.context_length { + parts.push(format!("ctx {context_length}")); + } + if let Some(max_output_tokens) = model.max_output_tokens { + parts.push(format!("max {max_output_tokens}")); + } + if !model.supports_tools { + parts.push("no tools".to_string()); + } + if !model.supports_streaming { + parts.push("no streaming".to_string()); + } + parts.join(" · ") +} + fn resume_chat_autofill( input: &str, selected: usize, @@ -798,3 +888,74 @@ fn blocks_from_conversation(conversation: &Conversation) -> Vec } blocks } + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::tempdir; + + fn config_with_models(models_json: &str) -> (tempfile::TempDir, Config) { + let root = tempdir().unwrap(); + std::fs::write(root.path().join("models.json"), models_json).unwrap(); + let mut config = Config::default(); + config.root = root.path().to_path_buf(); + config.model = "alpha-model".to_string(); + (root, config) + } + + #[test] + fn model_autofill_lists_models_from_models_json() { + let (_root, config) = config_with_models( + r#"{ + "models": [ + { + "id": "alpha-model", + "provider": "fireworks", + "display_name": "Alpha Model", + "context_length": 1000, + "max_output_tokens": 200 + }, + { + "id": "beta-model", + "provider": "other", + "display_name": "Beta Model" + } + ] +} +"#, + ); + + let menu = model_autofill("/model ", 0, &config).unwrap().unwrap(); + + assert_eq!(menu.items.len(), 2); + assert_eq!(menu.items[0].label, "alpha-model"); + assert_eq!(menu.items[0].insert, "alpha-model"); + assert_eq!(menu.apply("/model ").unwrap(), "/model alpha-model"); + let detail = menu.items[0].detail.as_deref().unwrap(); + assert!(detail.contains("current")); + assert!(detail.contains("Alpha Model")); + assert!(detail.contains("provider fireworks")); + } + + #[test] + fn model_autofill_filters_and_hides_exact_matches() { + let (_root, config) = config_with_models( + r#"{ + "models": [ + { "id": "alpha-model", "provider": "fireworks" }, + { "id": "beta-model", "provider": "other", "display_name": "Beta Model" } + ] +} +"#, + ); + + let menu = model_autofill("/model beta", 0, &config).unwrap().unwrap(); + assert_eq!(menu.items.len(), 1); + assert_eq!(menu.items[0].insert, "beta-model"); + assert_eq!(menu.apply("/model beta").unwrap(), "/model beta-model"); + + assert!(model_autofill("/model alpha-model", 0, &config) + .unwrap() + .is_none()); + } +} diff --git a/src/check.rs b/src/check.rs new file mode 100644 index 0000000..99f7854 --- /dev/null +++ b/src/check.rs @@ -0,0 +1,213 @@ +use crate::cli::Cli; +use crate::config::{self, ApiKeyReference, Config, ModelsFile, ProvidersFile}; +use anyhow::{Context, Result}; +use std::fs; +use std::path::{Path, PathBuf}; + +#[derive(Debug, Default, Clone)] +pub struct CheckReport { + pub successes: Vec, + pub warnings: Vec, + pub errors: Vec, +} + +impl CheckReport { + pub fn has_errors(&self) -> bool { + !self.errors.is_empty() + } + + pub fn render(&self) -> String { + let mut out = String::from("Cass config check\n"); + for success in &self.successes { + out.push_str("✓ "); + out.push_str(success); + out.push('\n'); + } + if !self.warnings.is_empty() { + out.push_str("\nWarnings:\n"); + for warning in &self.warnings { + out.push_str("! "); + out.push_str(warning); + out.push('\n'); + } + } + if !self.errors.is_empty() { + out.push_str("\nErrors:\n"); + for error in &self.errors { + out.push_str("✗ "); + out.push_str(error); + out.push('\n'); + } + out.push_str("\nConfig check failed.\n"); + } else { + out.push_str("\nAll checks passed.\n"); + } + out + } +} + +pub fn run(cli: &Cli) -> Result { + run_with_root(config::cass_root(), cli) +} + +pub fn run_with_root(root: PathBuf, cli: &Cli) -> Result { + fs::create_dir_all(&root).with_context(|| format!("creating {}", root.display()))?; + + let mut report = CheckReport::default(); + let config_path = config::config_path(&root); + let providers_path = config::providers_path(&root); + let models_path = config::models_path(&root); + + let config_file = match config::load_config_file(&root) { + Ok(Some(file)) => { + report + .successes + .push(format!("{}: valid", pretty_path(&config_path))); + Some(file) + } + Ok(None) => { + report.successes.push(format!( + "{}: not present (using defaults)", + pretty_path(&config_path) + )); + None + } + Err(err) => { + report + .errors + .push(format!("{}: {err:#}", pretty_path(&config_path))); + None + } + }; + + let providers = match config::load_or_create_default_provider_registry(&root) { + Ok(file) => { + report.successes.push(format!( + "{}: valid ({} provider{})", + pretty_path(&providers_path), + file.providers.len(), + plural(file.providers.len()) + )); + Some(file) + } + Err(err) => { + report + .errors + .push(format!("{}: {err:#}", pretty_path(&providers_path))); + None + } + }; + + let models = match config::load_or_create_default_model_registry(&root) { + Ok(file) => { + report.successes.push(format!( + "{}: valid ({} model{})", + pretty_path(&models_path), + file.models.len(), + plural(file.models.len()) + )); + Some(file) + } + Err(err) => { + report + .errors + .push(format!("{}: {err:#}", pretty_path(&models_path))); + None + } + }; + + let (Some(providers), Some(models)) = (providers, models) else { + return Ok(report); + }; + + if report.errors.is_empty() { + validate(&mut report, config_file.as_ref(), &providers, &models); + } + + if report.errors.is_empty() { + match Config::load_from_root_with_docs(root.clone(), root.join("docs"), cli) { + Ok(cfg) => check_active_config(&mut report, &cfg, &providers), + Err(err) => report.errors.push(format!("active config: {err:#}")), + } + } + + Ok(report) +} + +fn validate( + report: &mut CheckReport, + config_file: Option<&config::ConfigFile>, + providers: &ProvidersFile, + models: &ModelsFile, +) { + let summary = config::validate_registries(config_file, providers, models); + report.warnings.extend(summary.warnings); + report.errors.extend(summary.errors); +} + +fn check_active_config(report: &mut CheckReport, cfg: &Config, providers: &ProvidersFile) { + report + .successes + .push(format!("active provider: {}", cfg.provider_id)); + report + .successes + .push(format!("active model: {}", cfg.model)); + + check_api_key(report, "api key", &cfg.active_provider.api_key, true); + + for provider in &providers.providers { + if provider.id == cfg.provider_id { + continue; + } + check_api_key( + report, + &format!("provider `{}` api key", provider.id), + &provider.api_key, + false, + ); + } +} + +fn check_api_key(report: &mut CheckReport, label: &str, spec: &str, active: bool) { + match config::api_key_reference(spec) { + Ok(ApiKeyReference::Env(name)) => match std::env::var(&name) { + Ok(value) if !value.is_empty() => { + if active { + report.successes.push(format!("{label}: {name} is set")); + } + } + _ if active => report + .errors + .push(format!("{label}: environment variable `{name}` is not set")), + _ => report + .warnings + .push(format!("{label}: environment variable `{name}` is not set")), + }, + Ok(ApiKeyReference::Literal) => { + if active { + report + .successes + .push(format!("{label}: literal value configured")); + } + } + Err(err) if active => report.errors.push(format!("{label}: {err}")), + Err(err) => report.warnings.push(format!("{label}: {err}")), + } +} + +fn plural(count: usize) -> &'static str { + if count == 1 { + "" + } else { + "s" + } +} + +fn pretty_path(path: &Path) -> String { + if let Some(home) = dirs::home_dir() { + if let Ok(rest) = path.strip_prefix(&home) { + return format!("~/{}", rest.display()); + } + } + path.display().to_string() +} diff --git a/src/cli.rs b/src/cli.rs index 0baa699..5400218 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -1,9 +1,12 @@ -use clap::Parser; +use clap::{Parser, Subcommand}; use std::path::PathBuf; #[derive(Debug, Parser, Clone)] #[command(name = "cass", version, about = "Cassady/Cass terminal coding agent")] pub struct Cli { + #[command(subcommand)] + pub command: Option, + /// Resume a chat. Without a chat id, list chats for the current cwd. #[arg(long, num_args = 0..=1, value_name = "CHAT_ID")] pub resume: Option>, @@ -33,6 +36,12 @@ pub struct Cli { pub full_access: bool, } +#[derive(Debug, Subcommand, Clone, PartialEq, Eq)] +pub enum Command { + /// Validate Cass config files. + Check, +} + pub fn parse() -> Cli { Cli::parse() } diff --git a/src/config.rs b/src/config.rs index 2b5212c..e7519c7 100644 --- a/src/config.rs +++ b/src/config.rs @@ -1,28 +1,108 @@ use crate::access::AccessMode; use crate::cli::Cli; -use anyhow::{Context, Result}; +use anyhow::{bail, Context, Result}; use serde::{Deserialize, Serialize}; +use std::collections::{BTreeMap, BTreeSet}; use std::fs; -use std::path::PathBuf; +use std::path::{Path, PathBuf}; + +pub const DEFAULT_PROVIDER_ID: &str = "fireworks"; +pub const DEFAULT_PROVIDER_NAME: &str = "Fireworks"; +pub const DEFAULT_PROVIDER_KIND: &str = "openai-compatible"; +pub const DEFAULT_MODEL: &str = "accounts/fireworks/models/qwen3p7-plus"; +pub const DEFAULT_BASE_URL: &str = "https://api.fireworks.ai/inference/v1"; +pub const DEFAULT_API_KEY_ENV: &str = "FIREWORKS_API_KEY"; + +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ConfigFile { + #[serde(skip_serializing_if = "Option::is_none")] + pub default_provider: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub default_model: Option, + + // Deprecated compatibility fields accepted from older config.json files. + #[serde(skip_serializing_if = "Option::is_none")] + pub provider: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub model: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub base_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub api_key_env: Option, + + #[serde(skip_serializing_if = "Option::is_none")] + pub default_access_mode: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub context_message_limit: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub model_tool_result_limit: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub ui_tool_result_limit: Option, +} #[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ConfigFile { - pub provider: Option, - pub model: Option, - pub base_url: Option, - pub api_key_env: Option, - pub default_access_mode: Option, - pub context_message_limit: Option, - pub model_tool_result_limit: Option, - pub ui_tool_result_limit: Option, +#[serde(deny_unknown_fields)] +pub struct ProvidersFile { + pub providers: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ProviderDefinition { + pub id: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub name: Option, + pub kind: String, + pub base_url: String, + pub api_key: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub default_model: Option, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub models: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ModelsFile { + pub models: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ModelDefinition { + pub id: String, + pub provider: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub display_name: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub context_length: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_output_tokens: Option, + #[serde(default = "default_true")] + pub supports_tools: bool, + #[serde(default = "default_true")] + pub supports_streaming: bool, +} + +#[derive(Debug, Clone)] +pub struct ResolvedProviderConfig { + pub id: String, + pub name: Option, + pub kind: String, + pub base_url: String, + /// Either a literal API key or an env-var reference like "$FIREWORKS_API_KEY". + pub api_key: String, + pub default_model: Option, + pub models: Vec, } #[derive(Debug, Clone)] pub struct Config { - pub provider: String, + pub provider_id: String, pub model: String, - pub base_url: String, - pub api_key_env: String, + pub active_provider: ResolvedProviderConfig, + pub model_metadata: Option, pub default_access_mode: AccessMode, pub context_message_limit: usize, pub model_tool_result_limit: usize, @@ -31,15 +111,22 @@ pub struct Config { pub docs_dir: PathBuf, } +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ApiKeyReference { + Env(String), + Literal, +} + impl Default for Config { fn default() -> Self { let root = cass_root(); let docs_dir = root.join("docs"); + let active_provider = default_provider_definition().into_resolved(); Self { - provider: "openai-compatible".to_string(), - model: "accounts/fireworks/models/qwen3p7-plus".to_string(), - base_url: "https://api.fireworks.ai/inference/v1".to_string(), - api_key_env: "FIREWORKS_API_KEY".to_string(), + provider_id: DEFAULT_PROVIDER_ID.to_string(), + model: DEFAULT_MODEL.to_string(), + active_provider, + model_metadata: Some(default_model_definition()), default_access_mode: AccessMode::ReadOnly, context_message_limit: 80, model_tool_result_limit: 24_000, @@ -56,34 +143,41 @@ pub fn cass_root() -> PathBuf { .join(".cass") } +pub fn config_path(root: &Path) -> PathBuf { + root.join("config.json") +} + +pub fn providers_path(root: &Path) -> PathBuf { + root.join("providers.json") +} + +pub fn models_path(root: &Path) -> PathBuf { + root.join("models.json") +} + impl Config { pub fn load(cli: &Cli) -> Result { - let root = cass_root(); + Self::load_from_root(cass_root(), cli) + } + + pub fn load_from_root(root: PathBuf, cli: &Cli) -> Result { fs::create_dir_all(root.join("conversations")) .with_context(|| format!("creating {}", root.join("conversations").display()))?; let docs_dir = crate::docs::install(&root)?; + Self::load_from_root_with_docs(root, docs_dir, cli) + } + + pub fn load_from_root_with_docs(root: PathBuf, docs_dir: PathBuf, cli: &Cli) -> Result { + fs::create_dir_all(&root).with_context(|| format!("creating {}", root.display()))?; + let providers = load_or_create_default_provider_registry(&root)?; + let models = load_or_create_default_model_registry(&root)?; + let file = load_config_file(&root)?; + let mut cfg = Config::default(); cfg.root = root.clone(); cfg.docs_dir = docs_dir; - let path = root.join("config.json"); - if path.exists() { - let text = - fs::read_to_string(&path).with_context(|| format!("reading {}", path.display()))?; - let file: ConfigFile = serde_json::from_str(&text) - .with_context(|| format!("parsing {}", path.display()))?; - if let Some(v) = file.provider { - cfg.provider = v; - } - if let Some(v) = file.model { - cfg.model = v; - } - if let Some(v) = file.base_url { - cfg.base_url = v; - } - if let Some(v) = file.api_key_env { - cfg.api_key_env = v; - } + if let Some(file) = &file { if let Some(v) = file.default_access_mode { cfg.default_access_mode = v; } @@ -98,15 +192,6 @@ impl Config { } } - if let Some(v) = &cli.model { - cfg.model = v.clone(); - } - if let Some(v) = &cli.base_url { - cfg.base_url = v.clone(); - } - if let Some(v) = &cli.api_key_env { - cfg.api_key_env = v.clone(); - } if cli.readonly { cfg.default_access_mode = AccessMode::ReadOnly; } @@ -114,6 +199,35 @@ impl Config { cfg.default_access_mode = AccessMode::FullAccess; } + let requested_model = requested_model(file.as_ref(), cli); + let provider_id_from_config = requested_provider_id(file.as_ref(), &providers); + let legacy = legacy_provider_override(file.as_ref(), cli); + + let mut provider = resolve_provider( + requested_model.as_deref().unwrap_or(DEFAULT_MODEL), + provider_id_from_config.as_deref(), + file.as_ref().and_then(|f| f.provider.as_deref()), + legacy.as_ref(), + &providers, + &models, + )?; + + if let Some(base_url) = &cli.base_url { + provider.base_url = base_url.clone(); + } + if let Some(api_key_env) = &cli.api_key_env { + provider.api_key = format!("${api_key_env}"); + } + + let model = requested_model + .or_else(|| provider.default_model.clone()) + .unwrap_or_else(|| DEFAULT_MODEL.to_string()); + let metadata = find_model_for_provider(&models, &provider.id, &model).cloned(); + + cfg.provider_id = provider.id.clone(); + cfg.model = model; + cfg.active_provider = provider; + cfg.model_metadata = metadata; Ok(cfg) } @@ -128,4 +242,444 @@ impl Config { pub fn docs_dir(&self) -> PathBuf { self.docs_dir.clone() } + + pub fn resolved_api_key(&self) -> Result { + resolve_api_key(&self.active_provider.api_key) + } +} + +impl ProviderDefinition { + pub fn into_resolved(self) -> ResolvedProviderConfig { + ResolvedProviderConfig { + id: self.id, + name: self.name, + kind: self.kind, + base_url: self.base_url, + api_key: self.api_key, + default_model: self.default_model, + models: self.models, + } + } + + pub fn to_resolved(&self) -> ResolvedProviderConfig { + self.clone().into_resolved() + } +} + +pub fn load_config_file(root: &Path) -> Result> { + let path = config_path(root); + if !path.exists() { + return Ok(None); + } + let text = fs::read_to_string(&path).with_context(|| format!("reading {}", path.display()))?; + let file: ConfigFile = + serde_json::from_str(&text).with_context(|| format!("parsing {}", path.display()))?; + Ok(Some(file)) +} + +pub fn load_or_create_default_provider_registry(root: &Path) -> Result { + fs::create_dir_all(root).with_context(|| format!("creating {}", root.display()))?; + let path = providers_path(root); + if !path.exists() { + let file = ProvidersFile { + providers: vec![default_provider_definition()], + }; + write_json_pretty(&path, &file)?; + return Ok(file); + } + let text = fs::read_to_string(&path).with_context(|| format!("reading {}", path.display()))?; + serde_json::from_str(&text).with_context(|| format!("parsing {}", path.display())) +} + +pub fn load_or_create_default_model_registry(root: &Path) -> Result { + fs::create_dir_all(root).with_context(|| format!("creating {}", root.display()))?; + let path = models_path(root); + if !path.exists() { + let file = ModelsFile { + models: vec![default_model_definition()], + }; + write_json_pretty(&path, &file)?; + return Ok(file); + } + let text = fs::read_to_string(&path).with_context(|| format!("reading {}", path.display()))?; + serde_json::from_str(&text).with_context(|| format!("parsing {}", path.display())) +} + +pub fn default_provider_definition() -> ProviderDefinition { + ProviderDefinition { + id: DEFAULT_PROVIDER_ID.to_string(), + name: Some(DEFAULT_PROVIDER_NAME.to_string()), + kind: DEFAULT_PROVIDER_KIND.to_string(), + base_url: DEFAULT_BASE_URL.to_string(), + api_key: format!("${DEFAULT_API_KEY_ENV}"), + default_model: Some(DEFAULT_MODEL.to_string()), + models: vec![DEFAULT_MODEL.to_string()], + } +} + +pub fn default_model_definition() -> ModelDefinition { + ModelDefinition { + id: DEFAULT_MODEL.to_string(), + provider: DEFAULT_PROVIDER_ID.to_string(), + display_name: Some("Qwen 3p7 Plus".to_string()), + context_length: Some(262_144), + max_output_tokens: Some(32_768), + supports_tools: true, + supports_streaming: true, + } +} + +pub fn api_key_reference(spec: &str) -> Result { + if let Some(name) = spec.strip_prefix('$') { + if name.is_empty() { + bail!("API key env-var reference must include a variable name, e.g. \"$FIREWORKS_API_KEY\""); + } + return Ok(ApiKeyReference::Env(name.to_string())); + } + if spec.is_empty() { + bail!("literal API key must not be empty"); + } + Ok(ApiKeyReference::Literal) +} + +pub fn resolve_api_key(spec: &str) -> Result { + match api_key_reference(spec)? { + ApiKeyReference::Env(name) => { + let value = std::env::var(&name) + .with_context(|| format!("missing API key environment variable `{name}`"))?; + if value.is_empty() { + bail!("API key environment variable `{name}` is empty"); + } + Ok(value) + } + ApiKeyReference::Literal => Ok(spec.to_string()), + } +} + +pub fn redact_api_key_for_display(spec: &str) -> String { + match api_key_reference(spec) { + Ok(ApiKeyReference::Env(name)) => format!("${name}"), + Ok(ApiKeyReference::Literal) => "".to_string(), + Err(_) => "".to_string(), + } +} + +pub fn validate_registries( + config_file: Option<&ConfigFile>, + providers: &ProvidersFile, + models: &ModelsFile, +) -> ValidationSummary { + let mut out = ValidationSummary::default(); + let mut provider_ids = BTreeSet::new(); + let mut provider_counts = BTreeMap::::new(); + for provider in &providers.providers { + *provider_counts.entry(provider.id.clone()).or_default() += 1; + if provider.id.trim().is_empty() { + out.errors + .push("providers.json: provider id must not be empty".into()); + } + if provider.kind.trim().is_empty() { + out.errors.push(format!( + "providers.json: provider `{}` kind must not be empty", + provider.id + )); + } else if provider.kind != DEFAULT_PROVIDER_KIND { + out.errors.push(format!( + "providers.json: provider `{}` uses unsupported kind `{}`", + provider.id, provider.kind + )); + } + if provider.base_url.trim().is_empty() { + out.errors.push(format!( + "providers.json: provider `{}` base_url must not be empty", + provider.id + )); + } else if reqwest::Url::parse(&provider.base_url).is_err() { + out.errors.push(format!( + "providers.json: provider `{}` base_url is not a valid URL", + provider.id + )); + } + if let Err(err) = api_key_reference(&provider.api_key) { + out.errors.push(format!( + "providers.json: provider `{}` has invalid api_key: {err}", + provider.id + )); + } + if let Some(model) = &provider.default_model { + if model.trim().is_empty() { + out.errors.push(format!( + "providers.json: provider `{}` default_model must not be empty", + provider.id + )); + } + } + for model in &provider.models { + if model.trim().is_empty() { + out.errors.push(format!( + "providers.json: provider `{}` models entries must not be empty", + provider.id + )); + } + } + provider_ids.insert(provider.id.clone()); + } + for (id, count) in provider_counts { + if count > 1 { + out.errors + .push(format!("providers.json: duplicate provider id `{id}`")); + } + } + + let mut model_counts = BTreeMap::<(String, String), usize>::new(); + for model in &models.models { + *model_counts + .entry((model.provider.clone(), model.id.clone())) + .or_default() += 1; + if model.id.trim().is_empty() { + out.errors + .push("models.json: model id must not be empty".into()); + } + if model.provider.trim().is_empty() { + out.errors.push(format!( + "models.json: model `{}` provider must not be empty", + model.id + )); + } else if !provider_ids.contains(&model.provider) { + out.errors.push(format!( + "models.json: model `{}` references unknown provider `{}`", + model.id, model.provider + )); + } + if matches!(model.context_length, Some(0)) { + out.errors.push(format!( + "models.json: model `{}` context_length must be positive", + model.id + )); + } + if matches!(model.max_output_tokens, Some(0)) { + out.errors.push(format!( + "models.json: model `{}` max_output_tokens must be positive", + model.id + )); + } + } + for ((provider, id), count) in model_counts { + if count > 1 { + out.errors.push(format!( + "models.json: duplicate model `{id}` for provider `{provider}`" + )); + } + } + + for provider in &providers.providers { + if let Some(default_model) = &provider.default_model { + if find_model_for_provider(models, &provider.id, default_model).is_none() { + out.warnings.push(format!( + "providers.json: provider `{}` default_model `{}` has no matching models.json entry", + provider.id, default_model + )); + } + } + for model in &provider.models { + if find_model_for_provider(models, &provider.id, model).is_none() { + out.warnings.push(format!( + "providers.json: provider `{}` model `{}` has no matching models.json entry", + provider.id, model + )); + } + } + } + + if let Some(file) = config_file { + if file.provider.is_some() { + out.warnings.push( + "config.json: `provider` is deprecated; use `default_provider` or infer provider from `default_model`".into(), + ); + } + if file.model.is_some() { + out.warnings + .push("config.json: `model` is deprecated; use `default_model`".into()); + } + if file.base_url.is_some() || file.api_key_env.is_some() { + out.warnings.push( + "config.json: `base_url` and `api_key_env` are deprecated; move provider connection details to providers.json".into(), + ); + } + if let Some(default_provider) = &file.default_provider { + if !provider_ids.contains(default_provider) { + out.errors.push(format!( + "config.json: default_provider `{default_provider}` does not exist in providers.json" + )); + } + } + } + + out +} + +#[derive(Debug, Default, Clone)] +pub struct ValidationSummary { + pub warnings: Vec, + pub errors: Vec, +} + +pub fn find_model_for_provider<'a>( + models: &'a ModelsFile, + provider_id: &str, + model_id: &str, +) -> Option<&'a ModelDefinition> { + models + .models + .iter() + .find(|m| m.provider == provider_id && m.id == model_id) +} + +fn requested_model(file: Option<&ConfigFile>, cli: &Cli) -> Option { + cli.model.clone().or_else(|| { + file.and_then(|f| { + f.default_model + .clone() + .or_else(|| f.model.clone()) + .filter(|m| !m.trim().is_empty()) + }) + }) +} + +fn requested_provider_id(file: Option<&ConfigFile>, providers: &ProvidersFile) -> Option { + let file = file?; + if let Some(default_provider) = &file.default_provider { + return Some(default_provider.clone()); + } + let legacy_provider = file.provider.as_ref()?; + if providers.providers.iter().any(|p| p.id == *legacy_provider) { + Some(legacy_provider.clone()) + } else { + None + } +} + +#[derive(Debug, Clone)] +struct LegacyProviderOverride { + base_url: Option, + api_key: Option, +} + +fn legacy_provider_override( + file: Option<&ConfigFile>, + cli: &Cli, +) -> Option { + let base_url = cli + .base_url + .clone() + .or_else(|| file.and_then(|f| f.base_url.clone())); + let api_key = cli + .api_key_env + .as_ref() + .map(|env| format!("${env}")) + .or_else(|| file.and_then(|f| f.api_key_env.as_ref().map(|env| format!("${env}")))); + if base_url.is_some() || api_key.is_some() { + Some(LegacyProviderOverride { base_url, api_key }) + } else { + None + } +} + +fn resolve_provider( + model: &str, + explicit_provider_id: Option<&str>, + legacy_provider_name: Option<&str>, + legacy: Option<&LegacyProviderOverride>, + providers: &ProvidersFile, + models: &ModelsFile, +) -> Result { + if let Some(id) = explicit_provider_id { + if let Some(provider) = providers.providers.iter().find(|p| p.id == id) { + return Ok(provider.to_resolved()); + } + if let Some(legacy) = legacy { + return Ok(legacy_resolved_provider(id, legacy)); + } + bail!("configured provider `{id}` does not exist in providers.json"); + } + + if let Some(legacy) = legacy { + let id = legacy_provider_name.unwrap_or("legacy-openai-compatible"); + return Ok(legacy_resolved_provider(id, legacy)); + } + + if let Some(provider_id) = unique_provider_for_model(model, providers, models)? { + if let Some(provider) = providers.providers.iter().find(|p| p.id == provider_id) { + return Ok(provider.to_resolved()); + } + bail!("model `{model}` references unknown provider `{provider_id}` in models.json"); + } + + if let Some(provider) = providers + .providers + .iter() + .find(|p| p.id == DEFAULT_PROVIDER_ID) + { + return Ok(provider.to_resolved()); + } + + if providers.providers.len() == 1 { + return Ok(providers.providers[0].to_resolved()); + } + + bail!("could not resolve active provider; set config.json `default_provider` or choose a model with a unique provider in models.json") +} + +fn unique_provider_for_model( + model: &str, + providers: &ProvidersFile, + models: &ModelsFile, +) -> Result> { + let mut ids = BTreeSet::new(); + for model_def in &models.models { + if model_def.id == model { + ids.insert(model_def.provider.clone()); + } + } + for provider in &providers.providers { + if provider.default_model.as_deref() == Some(model) + || provider.models.iter().any(|candidate| candidate == model) + { + ids.insert(provider.id.clone()); + } + } + match ids.len() { + 0 => Ok(None), + 1 => Ok(ids.into_iter().next()), + _ => bail!( + "model `{model}` is configured for multiple providers; set config.json `default_provider`" + ), + } +} + +fn legacy_resolved_provider(id: &str, legacy: &LegacyProviderOverride) -> ResolvedProviderConfig { + ResolvedProviderConfig { + id: id.to_string(), + name: Some("Legacy OpenAI-compatible provider".to_string()), + kind: DEFAULT_PROVIDER_KIND.to_string(), + base_url: legacy + .base_url + .clone() + .unwrap_or_else(|| DEFAULT_BASE_URL.to_string()), + api_key: legacy + .api_key + .clone() + .unwrap_or_else(|| format!("${DEFAULT_API_KEY_ENV}")), + default_model: Some(DEFAULT_MODEL.to_string()), + models: vec![DEFAULT_MODEL.to_string()], + } +} + +fn write_json_pretty(path: &Path, value: &T) -> Result<()> { + let text = serde_json::to_string_pretty(value)?; + fs::write(path, format!("{text}\n")).with_context(|| format!("writing {}", path.display())) +} + +fn default_true() -> bool { + true } diff --git a/src/lib.rs b/src/lib.rs index e068c4e..980903e 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,6 +1,7 @@ pub mod access; pub mod agent; pub mod app; +pub mod check; pub mod cli; pub mod config; pub mod conversation; diff --git a/src/providers/openai_compatible.rs b/src/providers/openai_compatible.rs index a4f3592..a90da47 100644 --- a/src/providers/openai_compatible.rs +++ b/src/providers/openai_compatible.rs @@ -2,7 +2,7 @@ use super::types::{CompletionResult, ModelMessage}; use crate::agent::AgentEvent; use crate::conversation::StoredToolCall; use crate::tools::ToolSpec; -use anyhow::{bail, Context, Result}; +use anyhow::{bail, Result}; use futures_util::StreamExt; use reqwest::Client; use serde_json::{json, Value}; @@ -17,6 +17,13 @@ pub struct OpenAiCompatibleProvider { api_key: String, } +#[derive(Debug, Clone)] +pub struct OpenAiCompatibleSettings { + pub model: String, + pub base_url: String, + pub api_key: String, +} + #[derive(Debug, Default)] struct PartialToolCall { id: Option, @@ -25,15 +32,13 @@ struct PartialToolCall { } impl OpenAiCompatibleProvider { - pub fn new(model: String, base_url: String, api_key_env: String) -> Result { - let api_key = std::env::var(&api_key_env) - .with_context(|| format!("missing API key environment variable `{api_key_env}`"))?; - Ok(Self { + pub fn new(settings: OpenAiCompatibleSettings) -> Self { + Self { client: Client::new(), - model, - base_url: normalize_base_url(&base_url), - api_key, - }) + model: settings.model, + base_url: normalize_base_url(&settings.base_url), + api_key: settings.api_key, + } } pub async fn complete( diff --git a/tests/config_tests.rs b/tests/config_tests.rs new file mode 100644 index 0000000..eab4882 --- /dev/null +++ b/tests/config_tests.rs @@ -0,0 +1,189 @@ +use cassady::check; +use cassady::cli::Cli; +use cassady::config::{self, Config, ModelsFile, ProviderDefinition, ProvidersFile}; +use tempfile::tempdir; + +fn cli() -> Cli { + Cli { + command: None, + resume: None, + model: None, + base_url: None, + api_key_env: None, + cwd: None, + readonly: false, + full_access: false, + } +} + +#[test] +fn default_provider_and_model_files_are_created() { + let root = tempdir().unwrap(); + + let cfg = Config::load_from_root_with_docs( + root.path().to_path_buf(), + root.path().join("docs"), + &cli(), + ) + .unwrap(); + + assert_eq!(cfg.provider_id, config::DEFAULT_PROVIDER_ID); + assert_eq!(cfg.model, config::DEFAULT_MODEL); + assert!(root.path().join("providers.json").is_file()); + assert!(root.path().join("models.json").is_file()); + + let providers: ProvidersFile = + serde_json::from_str(&std::fs::read_to_string(root.path().join("providers.json")).unwrap()) + .unwrap(); + assert_eq!(providers.providers[0].id, "fireworks"); + assert_eq!(providers.providers[0].api_key, "$FIREWORKS_API_KEY"); + + let models: ModelsFile = + serde_json::from_str(&std::fs::read_to_string(root.path().join("models.json")).unwrap()) + .unwrap(); + assert_eq!(models.models[0].provider, "fireworks"); +} + +#[test] +fn legacy_config_connection_fields_still_work() { + let root = tempdir().unwrap(); + std::fs::write( + root.path().join("config.json"), + r#"{ + "provider": "openai-compatible", + "model": "legacy-model", + "base_url": "https://example.com/v1", + "api_key_env": "LEGACY_API_KEY" +} +"#, + ) + .unwrap(); + + let cfg = Config::load_from_root_with_docs( + root.path().to_path_buf(), + root.path().join("docs"), + &cli(), + ) + .unwrap(); + + assert_eq!(cfg.provider_id, "openai-compatible"); + assert_eq!(cfg.model, "legacy-model"); + assert_eq!(cfg.active_provider.base_url, "https://example.com/v1"); + assert_eq!(cfg.active_provider.api_key, "$LEGACY_API_KEY"); +} + +#[test] +fn api_key_resolution_supports_env_refs_and_literals() { + let key = "CASS_TEST_PROVIDER_KEY"; + let old = std::env::var(key).ok(); + std::env::set_var(key, "secret-value"); + assert_eq!( + config::resolve_api_key(&format!("${key}")).unwrap(), + "secret-value" + ); + assert_eq!( + config::resolve_api_key("literal-key").unwrap(), + "literal-key" + ); + + std::env::remove_var(key); + assert!(config::resolve_api_key(&format!("${key}")).is_err()); + + if let Some(old) = old { + std::env::set_var(key, old); + } +} + +#[test] +fn validation_rejects_duplicate_provider_ids() { + let providers = ProvidersFile { + providers: vec![ + config::default_provider_definition(), + ProviderDefinition { + id: "fireworks".into(), + name: None, + kind: "openai-compatible".into(), + base_url: "https://example.com/v1".into(), + api_key: "$OTHER_KEY".into(), + default_model: None, + models: Vec::new(), + }, + ], + }; + let models = ModelsFile { + models: vec![config::default_model_definition()], + }; + + let summary = config::validate_registries(None, &providers, &models); + assert!(summary + .errors + .iter() + .any(|error| error.contains("duplicate provider id `fireworks`"))); +} + +#[test] +fn check_reports_invalid_json() { + let root = tempdir().unwrap(); + std::fs::write(root.path().join("providers.json"), "{ invalid json").unwrap(); + + let report = check::run_with_root(root.path().to_path_buf(), &cli()).unwrap(); + + assert!(report.has_errors()); + assert!(report + .errors + .iter() + .any(|error| error.contains("providers.json"))); +} + +#[test] +fn check_passes_for_valid_literal_key_without_leaking_it() { + let root = tempdir().unwrap(); + std::fs::write( + root.path().join("providers.json"), + r#"{ + "providers": [ + { + "id": "test-provider", + "kind": "openai-compatible", + "base_url": "https://example.com/v1", + "api_key": "super-secret-literal", + "default_model": "test-model", + "models": ["test-model"] + } + ] +} +"#, + ) + .unwrap(); + std::fs::write( + root.path().join("models.json"), + r#"{ + "models": [ + { + "id": "test-model", + "provider": "test-provider", + "context_length": 128, + "max_output_tokens": 64 + } + ] +} +"#, + ) + .unwrap(); + std::fs::write( + root.path().join("config.json"), + r#"{ + "default_provider": "test-provider", + "default_model": "test-model" +} +"#, + ) + .unwrap(); + + let report = check::run_with_root(root.path().to_path_buf(), &cli()).unwrap(); + let rendered = report.render(); + + assert!(!report.has_errors(), "{rendered}"); + assert!(!rendered.contains("super-secret-literal")); + assert!(rendered.contains("api key: literal value configured")); +}