use std::process::ExitCode;
use crate::cli::ProviderOp;
use crate::config::Config;
use crate::provider_acceptance::{ProviderProvisionFailureKind, ProviderProvisionResponse};
use crate::providers::{ProviderStore, ProviderUpsert};
#[derive(Debug, PartialEq, serde::Serialize)]
struct ProviderImportReport {
complete: bool,
results: Vec<serde_json::Value>,
}
fn import_failure(name: &str, outcome: ProviderProvisionFailureKind) -> serde_json::Value {
serde_json::json!({"name": name, "outcome": outcome})
}
fn remote_failure(body: &serde_json::Value) -> ProviderProvisionFailureKind {
let outcome = body.pointer("/error/outcome").cloned();
outcome
.and_then(|value| serde_json::from_value(value).ok())
.unwrap_or(ProviderProvisionFailureKind::Unverified)
}
fn print_import_report(report: &ProviderImportReport) {
println!(
"{}",
serde_json::to_string_pretty(report)
.unwrap_or_else(|_| { r#"{"complete":false,"results":[]}"#.to_string() })
);
}
async fn local_import_report(
client: &reqwest::Client,
store: &ProviderStore,
inputs: Vec<ProviderUpsert>,
) -> ProviderImportReport {
let mut report = ProviderImportReport {
complete: true,
results: Vec::with_capacity(inputs.len()),
};
for input in inputs {
let name = input.name.clone();
match crate::provider_acceptance::provision(client, store, input).await {
Ok(result) => {
report
.results
.push(serde_json::to_value(result.response()).unwrap_or_else(|_| {
import_failure(&name, ProviderProvisionFailureKind::PersistenceUncertain)
}));
}
Err(error) => {
report.complete = false;
report.results.push(import_failure(&name, error.kind()));
break;
}
}
}
report
}
async fn remote_import_report(
server: &crate::managed_server::ResolvedServer,
imported: &[ProviderUpsert],
) -> Result<ProviderImportReport, String> {
let mut report = ProviderImportReport {
complete: true,
results: Vec::with_capacity(imported.len()),
};
for record in imported {
let name = record.name.clone();
let response = crate::auth_remote::post_response(
server,
crate::route_contract::route_template(crate::route_contract::RouteId::Providers),
upsert_body(record)?,
)
.await;
match response {
Ok((status, body)) if status.is_success() => {
if let Ok(safe) = serde_json::from_value::<ProviderProvisionResponse>(body) {
report
.results
.push(serde_json::to_value(safe).unwrap_or_else(|_| {
import_failure(&name, ProviderProvisionFailureKind::Unverified)
}));
} else {
report.complete = false;
report.results.push(import_failure(
&name,
ProviderProvisionFailureKind::Unverified,
));
break;
}
}
Ok((_, body)) => {
report.complete = false;
report
.results
.push(import_failure(&name, remote_failure(&body)));
break;
}
Err(_) => {
report.complete = false;
report.results.push(import_failure(
&name,
ProviderProvisionFailureKind::Unverified,
));
break;
}
}
}
Ok(report)
}
#[must_use]
pub async fn run(config: &Config, op: &ProviderOp) -> ExitCode {
let store = match ProviderStore::open(&config.data_dir, &config.token_secret) {
Ok(store) => store,
Err(e) => {
eprintln!("error: {e}");
return ExitCode::from(1);
}
};
run_with(&store, op).await
}
pub async fn run_remote(
server: &crate::managed_server::ResolvedServer,
op: &ProviderOp,
) -> ExitCode {
warn_zai_policy(op);
match remote_result(server, op).await {
Ok(code) => code,
Err(error) => {
eprintln!("error: {error}");
ExitCode::from(1)
}
}
}
fn warn_zai_policy(op: &ProviderOp) {
if let ProviderOp::Add {
kind,
enabled,
acknowledge_intermediary_risk,
acknowledge_unsupported_client,
..
} = op
&& crate::providers::ProviderKind::from_str_opt(kind)
== Some(crate::providers::ProviderKind::ZaiCodingPlan)
&& *enabled
{
eprintln!(
"WARNING: z.ai Coding Plan is a personal subscription. Intermediary proxying is not explicitly approved; policy violations may restrict or ban the subscriber account."
);
if *acknowledge_intermediary_risk {
eprintln!("risk accepted: intermediary z.ai Coding Plan proxying");
}
for client in acknowledge_unsupported_client {
eprintln!(
"WARNING: unsupported z.ai tool risk accepted only for client '{client}'; this may cause account restrictions or a ban"
);
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Call {
pub method: &'static str,
pub path: String,
pub body: Option<serde_json::Value>,
}
pub fn call_for(op: &ProviderOp) -> Result<Option<Call>, String> {
Ok(match op {
ProviderOp::List { .. } => Some(Call {
method: "GET",
path: crate::route_contract::route_template(crate::route_contract::RouteId::Providers)
.to_string(),
body: None,
}),
ProviderOp::Show { name, .. } => Some(Call {
method: "GET",
path: format!("/api/management/providers/{name}"),
body: None,
}),
ProviderOp::Remove { name, .. } => Some(Call {
method: "DELETE",
path: format!("/api/management/providers/{name}"),
body: None,
}),
ProviderOp::Add {
name,
kind,
base_url,
model,
models,
supported_clients,
api_key,
api_key_stdin,
api_key_env,
subscriber_id,
acknowledge_intermediary_risk,
acknowledge_unsupported_client,
enabled,
if_absent,
..
} => {
reject_lefine_argv_key(kind, api_key.as_ref())?;
Some(Call {
method: "POST",
path: crate::route_contract::route_template(
crate::route_contract::RouteId::Providers,
)
.to_string(),
body: Some(upsert_body(&ProviderUpsert {
name: name.clone(),
kind: Some(kind.clone()),
base_url: base_url.clone(),
default_model: model.clone(),
models: Some(models.clone()),
supported_clients: Some(supported_clients.clone()),
api_key: supplied_api_key(api_key.as_ref(), *api_key_stdin)
.map_err(|error| error.to_string())?,
api_key_env: api_key_env.clone(),
encrypted_api_key: None,
enabled: Some(*enabled),
subscriber_id: subscriber_id.clone(),
acknowledge_intermediary_risk: Some(*acknowledge_intermediary_risk),
acknowledge_unsupported_clients: Some(acknowledge_unsupported_client.clone()),
if_absent: *if_absent,
})?),
})
}
ProviderOp::Import { .. } => None,
})
}
fn supplied_api_key(
api_key: Option<&String>,
api_key_stdin: bool,
) -> Result<Option<String>, Box<dyn std::error::Error + Send + Sync>> {
if api_key_stdin {
return crate::server_command::read_token().map(Some);
}
Ok(api_key.cloned())
}
fn reject_lefine_argv_key(kind: &str, api_key: Option<&String>) -> Result<(), String> {
if crate::providers::ProviderKind::from_str_opt(kind)
== Some(crate::providers::ProviderKind::Lefine)
&& api_key.is_some()
{
return Err("Lefine API keys must use --api-key-stdin or --api-key-env".into());
}
Ok(())
}
pub fn upsert_body(upsert: &ProviderUpsert) -> Result<serde_json::Value, String> {
serde_json::to_value(upsert).map_err(|error| error.to_string())
}
#[must_use]
pub fn records_in(answer: &serde_json::Value) -> Vec<serde_json::Value> {
answer
.get("data")
.and_then(serde_json::Value::as_array)
.cloned()
.unwrap_or_default()
}
async fn remote_result(
server: &crate::managed_server::ResolvedServer,
op: &ProviderOp,
) -> Result<ExitCode, String> {
if let ProviderOp::Import { path, .. } = op {
let text = std::fs::read_to_string(path)
.map_err(|error| format!("could not read {}: {error}", path.display()))?;
let imported = crate::providers::parse_provider_import(&text)
.map_err(|error| format!("could not parse {}: {error}", path.display()))?;
let report = remote_import_report(server, &imported).await?;
print_import_report(&report);
return Ok(if report.complete {
ExitCode::SUCCESS
} else {
ExitCode::from(1)
});
}
let Some(call) = call_for(op)? else {
return Ok(ExitCode::from(1));
};
let answer = match (call.method, call.body) {
("POST", Some(body)) => crate::auth_remote::post(server, &call.path, body).await,
("DELETE", _) => crate::auth_remote::delete(server, &call.path).await,
_ => crate::auth_remote::get(server, &call.path).await,
};
match op {
ProviderOp::List { json, .. } => {
let records = records_in(&answer?);
if *json {
println!(
"{}",
serde_json::to_string_pretty(&records).unwrap_or_else(|_| "[]".to_string())
);
} else {
print_remote_table(&records);
}
Ok(ExitCode::SUCCESS)
}
ProviderOp::Show { name, .. } => match answer {
Ok(record) => {
println!(
"{}",
serde_json::to_string_pretty(&record).unwrap_or_default()
);
Ok(ExitCode::SUCCESS)
}
Err(error) if error.contains("404") => {
eprintln!("not found: {name}");
Ok(ExitCode::from(2))
}
Err(error) => Err(error),
},
ProviderOp::Remove { name, .. } => {
answer?;
println!("removed {name}");
Ok(ExitCode::SUCCESS)
}
ProviderOp::Add { .. } => {
let answer = answer?;
println!(
"{}",
serde_json::to_string_pretty(&answer)
.map_err(|error| format!("could not encode provider outcome: {error}"))?
);
Ok(ExitCode::SUCCESS)
}
ProviderOp::Import { .. } => Ok(ExitCode::from(1)),
}
}
fn print_remote_table(records: &[serde_json::Value]) {
println!(
"{:<20} {:<18} {:<32} {:<10} default_model",
"name", "kind", "base_url", "enabled"
);
for record in records {
let text = |key: &str| {
record
.get(key)
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
};
let kind = crate::providers::ProviderKind::from_str_opt(text("kind"))
.unwrap_or_default()
.as_str();
println!(
"{:<20} {:<18} {:<32} {:<10} {}",
text("name"),
kind,
text("base_url"),
record
.get("enabled")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false),
text("default_model"),
);
}
}
#[must_use]
pub async fn run_with(store: &ProviderStore, op: &ProviderOp) -> ExitCode {
warn_zai_policy(op);
let client = match crate::upstream_client::build_upstream_client() {
Ok(client) => client,
Err(error) => {
eprintln!("error: could not build provider validation client: {error}");
return ExitCode::from(1);
}
};
match op {
ProviderOp::List { json, .. } => match store.list_redacted() {
Ok(records) if *json => {
println!(
"{}",
serde_json::to_string_pretty(&records).unwrap_or_else(|_| "[]".to_string())
);
ExitCode::SUCCESS
}
Ok(records) => {
println!(
"{:<20} {:<18} {:<32} {:<10} default_model",
"name", "kind", "base_url", "enabled"
);
for record in records {
println!(
"{:<20} {:<18} {:<32} {:<10} {}",
record.name,
record.kind.as_str(),
record.base_url,
record.enabled,
record.default_model.unwrap_or_default()
);
}
ExitCode::SUCCESS
}
Err(e) => {
eprintln!("error: {e}");
ExitCode::from(1)
}
},
ProviderOp::Add {
name,
kind,
base_url,
model,
models,
supported_clients,
api_key,
api_key_stdin,
api_key_env,
subscriber_id,
acknowledge_intermediary_risk,
acknowledge_unsupported_client,
enabled,
if_absent,
..
} => {
if let Err(error) = reject_lefine_argv_key(kind, api_key.as_ref()) {
eprintln!("error: {error}");
return ExitCode::from(2);
}
let api_key = match supplied_api_key(api_key.as_ref(), *api_key_stdin) {
Ok(api_key) => api_key,
Err(error) => {
eprintln!("error: {error}");
return ExitCode::from(2);
}
};
let input = ProviderUpsert {
name: name.clone(),
kind: Some(kind.clone()),
base_url: base_url.clone(),
default_model: model.clone(),
models: Some(models.clone()),
supported_clients: Some(supported_clients.clone()),
api_key,
api_key_env: api_key_env.clone(),
encrypted_api_key: None,
enabled: Some(*enabled),
subscriber_id: subscriber_id.clone(),
acknowledge_intermediary_risk: Some(*acknowledge_intermediary_risk),
acknowledge_unsupported_clients: Some(acknowledge_unsupported_client.clone()),
if_absent: *if_absent,
};
match crate::provider_acceptance::provision(&client, store, input).await {
Ok(result) => {
println!(
"{}",
serde_json::to_string_pretty(&result.response()).unwrap_or_default()
);
ExitCode::SUCCESS
}
Err(e) => {
eprintln!("error: {e}");
ExitCode::from(1)
}
}
}
ProviderOp::Show { name, .. } => match store.get(name) {
Ok(Some(record)) => {
println!(
"{}",
serde_json::to_string_pretty(&record.redacted()).unwrap_or_default()
);
ExitCode::SUCCESS
}
Ok(None) => {
eprintln!("not found: {name}");
ExitCode::from(2)
}
Err(e) => {
eprintln!("error: {e}");
ExitCode::from(1)
}
},
ProviderOp::Remove { name, .. } => match store.delete(name) {
Ok(true) => {
println!("removed {name}");
ExitCode::SUCCESS
}
Ok(false) => {
eprintln!("not found: {name}");
ExitCode::from(2)
}
Err(e) => {
eprintln!("error: {e}");
ExitCode::from(1)
}
},
ProviderOp::Import { path, .. } => match std::fs::read_to_string(path)
.map_err(crate::providers::ProviderError::from)
.and_then(|text| crate::providers::parse_provider_import(&text))
{
Ok(inputs) => {
let report = local_import_report(&client, store, inputs).await;
print_import_report(&report);
if report.complete {
ExitCode::SUCCESS
} else {
ExitCode::from(1)
}
}
Err(e) => {
eprintln!("error: {e}");
ExitCode::from(1)
}
},
}
}
#[cfg(test)]
#[path = "providers_cli_tests.rs"]
mod tests;