use std::io::{BufRead, Write};
use crate::inference::configurator::Configurator;
use crate::inference::credentials::{KeyStore, env_local_value, redact_secret, resolve_key_with};
use crate::inference::error::InferenceError;
use crate::inference::registry::{ProviderCapabilities, all, capabilities_for};
use crate::inference::types::{ChatMessage, ChatRequest};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum KeyTier {
Env,
EnvLocal,
EnvOrEnvLocal,
Store,
}
impl KeyTier {
pub fn label(self) -> &'static str {
match self {
Self::Env => "environment variable",
Self::EnvLocal => ".env.local",
Self::EnvOrEnvLocal => "environment variable or .env.local (ambiguous)",
Self::Store => "secure store",
}
}
}
pub fn classify_tier(env: bool, env_local: bool, store: bool) -> Option<KeyTier> {
if env && env_local {
Some(KeyTier::EnvOrEnvLocal)
} else if env {
Some(KeyTier::Env)
} else if env_local {
Some(KeyTier::EnvLocal)
} else if store {
Some(KeyTier::Store)
} else {
None
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ProbeOutcome {
Ok,
Unauthorized,
ModelNotFound(String),
Unconfigured,
Unsupported(String),
Failed(String),
}
impl ProbeOutcome {
pub fn label(&self) -> String {
match self {
Self::Ok => "OK — credentials accepted".to_string(),
Self::Unauthorized => "UNAUTHORIZED — provider rejected the key (401/403)".to_string(),
Self::ModelNotFound(reason) => format!(
"MODEL NOT FOUND — the provider's default model is not found/deployed on this \
account ({reason}); set an explicit, account-available model instead of relying \
on the built-in default"
),
Self::Unconfigured => {
"UNCONFIGURED — no key resolved; set one with `config keys set <provider>`"
.to_string()
}
Self::Unsupported(reason) => format!("SKIPPED — {reason}"),
Self::Failed(reason) => format!("ERROR — {reason}"),
}
}
pub fn into_result(self) -> anyhow::Result<()> {
match self {
Self::Ok | Self::Unconfigured | Self::Unsupported(_) => Ok(()),
other => Err(anyhow::anyhow!("{}", other.label())),
}
}
}
pub fn set(
store: &dyn KeyStore,
provider: &str,
value: &str,
out: &mut dyn Write,
) -> anyhow::Result<()> {
let caps = known_keyed_provider(provider)?;
let name = caps.id.as_str();
if value.is_empty() {
anyhow::bail!("refusing to store an empty key for {name}");
}
store
.set(name, value)
.map_err(|e| anyhow::anyhow!("failed to store key for {name}: {e}"))?;
writeln!(
out,
"Stored key for {name} in the secure store [{}].",
redact_secret(value)
)?;
Ok(())
}
pub fn unset(store: &dyn KeyStore, provider: &str, out: &mut dyn Write) -> anyhow::Result<()> {
let caps = known_keyed_provider(provider)?;
let name = caps.id.as_str();
let was_present = store.get(name).is_some();
store
.unset(name)
.map_err(|e| anyhow::anyhow!("failed to remove key for {name}: {e}"))?;
if was_present {
writeln!(out, "Removed key for {name} from the secure store.")?;
} else {
writeln!(
out,
"No key for {name} was set in the secure store; nothing to remove."
)?;
}
Ok(())
}
pub fn list(store: &dyn KeyStore, out: &mut dyn Write) -> anyhow::Result<()> {
writeln!(
out,
"Provider key status (names and tiers only — values are never shown):"
)?;
for caps in all() {
let name = caps.id.as_str();
match caps.credential_env {
None => writeln!(out, " {name:<12} AWS credential chain (no API key)")?,
Some(env_var) => match detect_tier(env_var, name, store) {
Some(tier) => writeln!(out, " {name:<12} configured via {}", tier.label())?,
None => writeln!(out, " {name:<12} not configured")?,
},
}
}
Ok(())
}
pub async fn probe(
store: &dyn KeyStore,
cfg: &Configurator,
provider: &str,
) -> anyhow::Result<ProbeOutcome> {
let caps = capabilities_for(provider).ok_or_else(|| {
anyhow::anyhow!("unknown provider {provider:?}; known: {}", known_names())
})?;
let name = caps.id.as_str();
if caps.credential_env.is_none() {
return Ok(ProbeOutcome::Unsupported(format!(
"{name} authenticates via the AWS credential chain, not an API key"
)));
}
let Some(resolved_key) = resolve_key_with(name, store) else {
return Ok(ProbeOutcome::Unconfigured);
};
let slug = format!("{name}/{}", caps.default_model);
let adapter = match cfg.build(&slug, store) {
Ok(adapter) => adapter,
Err(InferenceError::NoAdapterRegistered { .. }) => {
return Ok(ProbeOutcome::Unsupported(format!(
"no inference adapter is wired for {name} yet"
)));
}
Err(InferenceError::MissingCredential { .. }) => return Ok(ProbeOutcome::Unconfigured),
Err(err) => {
return Ok(ProbeOutcome::Failed(scrub_key(
&err.to_string(),
&resolved_key,
)));
}
};
let mut req = ChatRequest::new(caps.default_model, vec![ChatMessage::user("ping")]);
req.max_tokens = Some(1);
req.temperature = Some(0.0);
match adapter.chat(&req).await {
Ok(_) => Ok(ProbeOutcome::Ok),
Err(InferenceError::Api {
status: 401 | 403, ..
}) => Ok(ProbeOutcome::Unauthorized),
Err(err @ InferenceError::Api { status: 404, .. }) => Ok(ProbeOutcome::ModelNotFound(
scrub_key(&err.to_string(), &resolved_key),
)),
Err(err) => Ok(ProbeOutcome::Failed(scrub_key(
&err.to_string(),
&resolved_key,
))),
}
}
fn scrub_key(message: &str, key: &str) -> String {
if key.is_empty() {
return message.to_string();
}
message.replace(key, "[REDACTED]")
}
pub fn report_probe(
provider: &str,
outcome: &ProbeOutcome,
out: &mut dyn Write,
) -> anyhow::Result<()> {
writeln!(out, "test {provider}: {}", outcome.label())?;
Ok(())
}
pub fn read_key_line(reader: &mut dyn BufRead) -> anyhow::Result<String> {
let mut buf = String::new();
reader.read_line(&mut buf)?;
Ok(buf.trim_end_matches(['\n', '\r']).to_string())
}
fn known_keyed_provider(provider: &str) -> anyhow::Result<&'static ProviderCapabilities> {
let caps = capabilities_for(provider).ok_or_else(|| {
anyhow::anyhow!("unknown provider {provider:?}; known: {}", known_names())
})?;
if caps.credential_env.is_none() {
anyhow::bail!(
"{} authenticates via the AWS credential chain, not an API key",
caps.id.as_str()
);
}
Ok(caps)
}
fn detect_tier(env_var: &str, provider: &str, store: &dyn KeyStore) -> Option<KeyTier> {
let env = std::env::var(env_var)
.ok()
.filter(|v| !v.is_empty())
.is_some();
let env_local = env_local_value(env_var).is_some();
let stored = store.get(provider).is_some();
classify_tier(env, env_local, stored)
}
fn known_names() -> String {
all()
.iter()
.filter(|c| c.credential_env.is_some())
.map(|c| c.id.as_str())
.collect::<Vec<_>>()
.join(", ")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classify_tier_follows_precedence() {
assert_eq!(
classify_tier(true, true, true),
Some(KeyTier::EnvOrEnvLocal)
);
assert_eq!(classify_tier(true, false, true), Some(KeyTier::Env));
assert_eq!(classify_tier(false, true, true), Some(KeyTier::EnvLocal));
assert_eq!(classify_tier(false, false, true), Some(KeyTier::Store));
assert_eq!(classify_tier(false, false, false), None);
}
#[test]
fn read_key_line_strips_newline() {
let mut reader = std::io::Cursor::new(b"sk-piped-value\n".to_vec()); assert_eq!(read_key_line(&mut reader).unwrap(), "sk-piped-value");
}
#[test]
fn scrub_key_removes_every_occurrence() {
let key = "sk-or-verysecret1234"; let msg = format!("inference API error 400: bad key {key}, retry without {key}");
let scrubbed = scrub_key(&msg, key);
assert!(!scrubbed.contains(key), "leaked: {scrubbed}");
assert_eq!(scrubbed.matches("[REDACTED]").count(), 2);
}
#[test]
fn scrub_key_is_noop_for_empty_key() {
assert_eq!(scrub_key("some message", ""), "some message");
}
}