Add model registry config and autocomplete

This commit is contained in:
2026-06-21 11:55:27 -05:00
parent 8dca417cd3
commit 225bda5631
12 changed files with 1607 additions and 74 deletions
+271
View File
@@ -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 <chat-id>
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<String>`
- `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`.
+16 -9
View File
@@ -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 <chat-id>
cass --resume
cass check
```
`cass --resume` without an ID lists chats for the current directory.
@@ -69,7 +75,8 @@ cass --resume
## Commands
- `/model <model>`: switch the model for future turns
- `cass check`: validate Cass config files
- `/model <model>`: switch the model for future turns; model autocomplete lists entries from `~/.cass/models.json`
- `/resume <chat>`: resume a saved chat; chat autocomplete lists chats for the current directory
- `/status`: show current chat status
+2
View File
@@ -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`.
+120
View File
@@ -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`.
+8 -7
View File
@@ -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 {
+163 -2
View File
@@ -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<AutoFillMenu> {
}
}
fn model_autofill(input: &str, selected: usize, config: &Config) -> Result<Option<AutoFillMenu>> {
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<TranscriptBlock>
}
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());
}
}
+213
View File
@@ -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<String>,
pub warnings: Vec<String>,
pub errors: Vec<String>,
}
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<CheckReport> {
run_with_root(config::cass_root(), cli)
}
pub fn run_with_root(root: PathBuf, cli: &Cli) -> Result<CheckReport> {
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()
}
+10 -1
View File
@@ -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<Command>,
/// 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<Option<String>>,
@@ -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()
}
+600 -46
View File
@@ -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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub default_model: Option<String>,
// Deprecated compatibility fields accepted from older config.json files.
#[serde(skip_serializing_if = "Option::is_none")]
pub provider: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub base_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key_env: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub default_access_mode: Option<AccessMode>,
#[serde(skip_serializing_if = "Option::is_none")]
pub context_message_limit: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model_tool_result_limit: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ui_tool_result_limit: Option<usize>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConfigFile {
pub provider: Option<String>,
pub model: Option<String>,
pub base_url: Option<String>,
pub api_key_env: Option<String>,
pub default_access_mode: Option<AccessMode>,
pub context_message_limit: Option<usize>,
pub model_tool_result_limit: Option<usize>,
pub ui_tool_result_limit: Option<usize>,
#[serde(deny_unknown_fields)]
pub struct ProvidersFile {
pub providers: Vec<ProviderDefinition>,
}
#[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<String>,
pub kind: String,
pub base_url: String,
pub api_key: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub default_model: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub models: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ModelsFile {
pub models: Vec<ModelDefinition>,
}
#[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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub context_length: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<u64>,
#[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<String>,
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<String>,
pub models: Vec<String>,
}
#[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<ModelDefinition>,
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<Self> {
let root = cass_root();
Self::load_from_root(cass_root(), cli)
}
pub fn load_from_root(root: PathBuf, cli: &Cli) -> Result<Self> {
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<Self> {
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<String> {
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<Option<ConfigFile>> {
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<ProvidersFile> {
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<ModelsFile> {
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<ApiKeyReference> {
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<String> {
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) => "<literal API key>".to_string(),
Err(_) => "<invalid API key reference>".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::<String, usize>::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<String>,
pub errors: Vec<String>,
}
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<String> {
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<String> {
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<String>,
api_key: Option<String>,
}
fn legacy_provider_override(
file: Option<&ConfigFile>,
cli: &Cli,
) -> Option<LegacyProviderOverride> {
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<ResolvedProviderConfig> {
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<Option<String>> {
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<T: Serialize>(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
}
+1
View File
@@ -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;
+14 -9
View File
@@ -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<String>,
@@ -25,15 +32,13 @@ struct PartialToolCall {
}
impl OpenAiCompatibleProvider {
pub fn new(model: String, base_url: String, api_key_env: String) -> Result<Self> {
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(
+189
View File
@@ -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"));
}