Add model registry config and autocomplete
This commit is contained in:
@@ -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`.
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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,6 +1,7 @@
|
||||
pub mod access;
|
||||
pub mod agent;
|
||||
pub mod app;
|
||||
pub mod check;
|
||||
pub mod cli;
|
||||
pub mod config;
|
||||
pub mod conversation;
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
Reference in New Issue
Block a user