Guide AI connection setup with private node credentials and explicit providers
This commit is contained in:
@@ -1025,10 +1025,46 @@ impl RpcHandler {
|
||||
.ok_or_else(|| anyhow::anyhow!("Missing key"))?;
|
||||
|
||||
match key {
|
||||
"claude_api_key_set" => {
|
||||
let key_file = self.config.data_dir.join("secrets/claude-api-key");
|
||||
let has_key = tokio::fs::metadata(&key_file).await.is_ok();
|
||||
Ok(serde_json::json!({ "value": has_key }))
|
||||
"claude_api_key_set" | "openai_api_key_set" => {
|
||||
let provider = if key == "claude_api_key_set" {
|
||||
"claude"
|
||||
} else {
|
||||
"openai"
|
||||
};
|
||||
Ok(
|
||||
serde_json::json!({ "value": crate::settings::model_provider::has_key(&self.config.data_dir, provider).await }),
|
||||
)
|
||||
}
|
||||
"ai_provider" => {
|
||||
let settings =
|
||||
crate::settings::model_provider::ModelProvider::load(&self.config.data_dir)
|
||||
.await?;
|
||||
Ok(serde_json::json!({ "value": settings }))
|
||||
}
|
||||
"ai_provider_status" => {
|
||||
let settings =
|
||||
crate::settings::model_provider::ModelProvider::load(&self.config.data_dir)
|
||||
.await?;
|
||||
let local = tokio::time::timeout(std::time::Duration::from_secs(4), async {
|
||||
let (detected, _) = crate::api::rpc::mesh::assistant::detect_ollama().await;
|
||||
detected
|
||||
&& crate::assistant::backends::ollama::model_supports_tools(
|
||||
crate::assistant::backends::ollama::OLLAMA_BASE_URL,
|
||||
crate::assistant::backends::ollama::OLLAMA_DEFAULT_MODEL,
|
||||
)
|
||||
.await
|
||||
});
|
||||
let (claude, openai, local) = tokio::join!(
|
||||
crate::settings::model_provider::has_key(&self.config.data_dir, "claude"),
|
||||
crate::settings::model_provider::has_key(&self.config.data_dir, "openai"),
|
||||
local,
|
||||
);
|
||||
let budget = crate::assistant::AssistantBudget::load(&self.config.data_dir).await;
|
||||
Ok(serde_json::json!({ "value": {
|
||||
"schema": 1, "settings": settings, "claude_configured": claude,
|
||||
"openai_configured": openai, "local_ready": local.ok(),
|
||||
"routstr_remaining_sats": budget.remaining_sats(),
|
||||
}}))
|
||||
}
|
||||
_ => Ok(serde_json::json!({ "value": null })),
|
||||
}
|
||||
@@ -1210,38 +1246,21 @@ impl RpcHandler {
|
||||
let value = params.get("value").and_then(|v| v.as_str()).unwrap_or("");
|
||||
|
||||
match key {
|
||||
"claude_api_key" => {
|
||||
let secrets_dir = self.config.data_dir.join("secrets");
|
||||
tokio::fs::create_dir_all(&secrets_dir)
|
||||
.await
|
||||
.context("Failed to create secrets dir")?;
|
||||
let key_file = secrets_dir.join("claude-api-key");
|
||||
|
||||
if value.is_empty() {
|
||||
// Remove key
|
||||
tokio::fs::remove_file(&key_file).await.ok();
|
||||
info!("Claude API key removed");
|
||||
"claude_api_key" | "openai_api_key" => {
|
||||
let provider = if key == "claude_api_key" {
|
||||
"claude"
|
||||
} else {
|
||||
// Save key
|
||||
tokio::fs::write(&key_file, value)
|
||||
.await
|
||||
.context("Failed to write API key")?;
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(&key_file, std::fs::Permissions::from_mode(0o600))
|
||||
.ok();
|
||||
}
|
||||
info!("Claude API key saved");
|
||||
}
|
||||
|
||||
// `secrets/claude-api-key` (above) is deliberately the ONLY
|
||||
// Claude key ledger on this node (13-02-PLAN.md). A second
|
||||
// copy used to be written alongside it for a standalone,
|
||||
// unauthenticated sidecar process on port 3142 — that
|
||||
// sidecar and its key copy are retired; the session-gated
|
||||
// Rust daemon reads this one file directly.
|
||||
|
||||
"openai"
|
||||
};
|
||||
crate::settings::model_provider::save_key(&self.config.data_dir, provider, value)
|
||||
.await?;
|
||||
info!(provider, "AI provider credential updated");
|
||||
Ok(serde_json::json!({ "saved": true }))
|
||||
}
|
||||
"ai_provider" => {
|
||||
let settings: crate::settings::model_provider::ModelProvider =
|
||||
serde_json::from_str(value).context("Invalid AI provider settings")?;
|
||||
settings.save(&self.config.data_dir).await?;
|
||||
Ok(serde_json::json!({ "saved": true }))
|
||||
}
|
||||
_ => anyhow::bail!("Unknown setting: {}", key),
|
||||
|
||||
@@ -11,6 +11,7 @@ use crate::api::rpc::RpcHandler;
|
||||
|
||||
pub mod claude;
|
||||
pub mod ollama;
|
||||
pub mod openai;
|
||||
pub mod routstr;
|
||||
#[cfg(test)]
|
||||
pub mod scripted;
|
||||
@@ -39,6 +40,8 @@ pub trait Backend: Send + Sync {
|
||||
pub enum BackendId {
|
||||
Ollama,
|
||||
Claude,
|
||||
Openai,
|
||||
Unavailable,
|
||||
/// 13-13: the third D-04 leg. Not currently returned as the "primary"
|
||||
/// id by `select_backend` (mirroring the existing convention that the
|
||||
/// returned id names the primary attempt, not necessarily which leg of
|
||||
@@ -52,11 +55,23 @@ impl std::fmt::Display for BackendId {
|
||||
match self {
|
||||
BackendId::Ollama => write!(f, "ollama"),
|
||||
BackendId::Claude => write!(f, "claude"),
|
||||
BackendId::Openai => write!(f, "openai"),
|
||||
BackendId::Unavailable => write!(f, "unavailable"),
|
||||
BackendId::Routstr => write!(f, "routstr"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct InvalidProviderSettings;
|
||||
#[async_trait]
|
||||
impl Backend for InvalidProviderSettings {
|
||||
async fn send(&self, _: &str, _: &[ToolDef], _: &[ChatMessage]) -> Result<BackendTurn> {
|
||||
anyhow::bail!(
|
||||
"AI connection settings could not be loaded. Review them before sending a message."
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// D-04's per-call fallback: try `primary`'s `send()`, and on a transport
|
||||
/// error fall through to `secondary` for that SAME call rather than
|
||||
/// failing the whole turn — a local model that answers earlier turns and
|
||||
@@ -115,6 +130,54 @@ fn ollama_is_selectable(detected: bool, tool_capable: bool) -> bool {
|
||||
/// tools-free degrade.
|
||||
pub async fn select_backend(handler: &RpcHandler) -> (Box<dyn Backend>, BackendId) {
|
||||
let data_dir = handler.data_dir();
|
||||
// An explicit provider is a privacy and billing choice. Never silently
|
||||
// fall through to another provider if its credentials or network fail.
|
||||
match crate::settings::model_provider::ModelProvider::load(data_dir).await {
|
||||
Ok(settings) => match settings.provider {
|
||||
crate::settings::model_provider::Provider::Openai => {
|
||||
return (
|
||||
Box::new(openai::OpenaiBackend::new(
|
||||
data_dir.to_path_buf(),
|
||||
settings.openai_model,
|
||||
)),
|
||||
BackendId::Openai,
|
||||
)
|
||||
}
|
||||
crate::settings::model_provider::Provider::Claude => {
|
||||
return (
|
||||
Box::new(claude::ClaudeBackend::new(data_dir.to_path_buf())),
|
||||
BackendId::Claude,
|
||||
)
|
||||
}
|
||||
crate::settings::model_provider::Provider::Local => {
|
||||
return (
|
||||
Box::new(ollama::OllamaBackend::new(
|
||||
ollama::OLLAMA_BASE_URL.to_string(),
|
||||
ollama::OLLAMA_DEFAULT_MODEL.to_string(),
|
||||
)),
|
||||
BackendId::Ollama,
|
||||
)
|
||||
}
|
||||
crate::settings::model_provider::Provider::Routstr => {
|
||||
let budget = crate::assistant::AssistantBudget::load(data_dir).await;
|
||||
let mints = crate::wallet::ecash::load_accepted_mints(data_dir)
|
||||
.await
|
||||
.map(|m| m.mints)
|
||||
.unwrap_or_default();
|
||||
return (
|
||||
Box::new(routstr::RoutstrBackend::new(
|
||||
data_dir.to_path_buf(),
|
||||
budget.payment_policy(),
|
||||
mints,
|
||||
handler.nostr_tor_proxy(),
|
||||
)),
|
||||
BackendId::Routstr,
|
||||
);
|
||||
}
|
||||
crate::settings::model_provider::Provider::Auto => {}
|
||||
},
|
||||
Err(_) => return (Box::new(InvalidProviderSettings), BackendId::Unavailable),
|
||||
}
|
||||
let (detected, _models) = crate::api::rpc::mesh::assistant::detect_ollama().await;
|
||||
let model = ollama::OLLAMA_DEFAULT_MODEL;
|
||||
let tool_capable = if detected {
|
||||
|
||||
@@ -0,0 +1,279 @@
|
||||
//! Explicit OpenAI API selection using the shared tool loop and egress policy.
|
||||
//! Keys stay node-side; no redirects, automatic retries, or provider fallback.
|
||||
use super::{Backend, BackendTurn};
|
||||
use crate::assistant::{
|
||||
egress::{self, EgressVerdict},
|
||||
tools::{ChatMessage, ToolCall, ToolDef},
|
||||
};
|
||||
use anyhow::{Context, Result};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{json, Value};
|
||||
use std::{path::PathBuf, time::Duration};
|
||||
|
||||
const URL: &str = "https://api.openai.com/v1/chat/completions";
|
||||
const RESPONSE_LIMIT: usize = 2 * 1024 * 1024;
|
||||
pub struct OpenaiBackend {
|
||||
data_dir: PathBuf,
|
||||
model: String,
|
||||
}
|
||||
impl OpenaiBackend {
|
||||
pub fn new(data_dir: PathBuf, model: String) -> Self {
|
||||
Self { data_dir, model }
|
||||
}
|
||||
async fn send_at(
|
||||
&self,
|
||||
url: &str,
|
||||
system: &str,
|
||||
tools: &[ToolDef],
|
||||
history: &[ChatMessage],
|
||||
) -> Result<BackendTurn> {
|
||||
let key = tokio::fs::read_to_string(self.data_dir.join("secrets/openai-api-key"))
|
||||
.await
|
||||
.map_err(|_| {
|
||||
anyhow::anyhow!("OpenAI API key is not configured. Open AI connection settings.")
|
||||
})?;
|
||||
anyhow::ensure!(!key.trim().is_empty(), "OpenAI API key is not configured");
|
||||
anyhow::ensure!(
|
||||
!self.model.is_empty(),
|
||||
"Choose an OpenAI model in AI connection settings"
|
||||
);
|
||||
let mut messages = vec![json!({"role": "system", "content": system})];
|
||||
messages.extend(history.iter().flat_map(super::routstr::message_to_wire));
|
||||
let mut body = json!({"model": self.model, "messages": messages, "stream": false,
|
||||
"store": false, "max_completion_tokens": 2048, "n": 1});
|
||||
if !tools.is_empty() {
|
||||
body["tools"] = json!(tools.iter().map(|tool| json!({"type": "function", "function": {
|
||||
"name": tool.name, "description": tool.description, "parameters": tool.parameters,
|
||||
}})).collect::<Vec<_>>());
|
||||
body["parallel_tool_calls"] = json!(false);
|
||||
}
|
||||
let context = egress::EgressContext::from_turn(
|
||||
history,
|
||||
&tools.iter().map(|tool| tool.name).collect::<Vec<_>>(),
|
||||
&self.data_dir.join("secrets"),
|
||||
)
|
||||
.await;
|
||||
match egress::screen_outbound(&body.to_string(), &context) {
|
||||
EgressVerdict::Allow => {}
|
||||
EgressVerdict::Truncate(value) => {
|
||||
body = serde_json::from_str(&value)
|
||||
.context("Could not apply outbound privacy filter")?;
|
||||
}
|
||||
EgressVerdict::BlockFallBackLocal => {
|
||||
crate::assistant::global_counters().note_blocked_egress();
|
||||
anyhow::bail!("This message contains private key or recovery material and was not sent to OpenAI");
|
||||
}
|
||||
}
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(180))
|
||||
.connect_timeout(Duration::from_secs(15))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()?;
|
||||
let mut response = client.post(url).bearer_auth(key.trim()).json(&body).send().await
|
||||
.map_err(|_| anyhow::anyhow!("OpenAI is temporarily unreachable. Your request was not retried automatically."))?;
|
||||
if !response.status().is_success() {
|
||||
anyhow::bail!("{}", error_message(response.status().as_u16()));
|
||||
}
|
||||
let mut bytes = Vec::new();
|
||||
while let Some(chunk) = response
|
||||
.chunk()
|
||||
.await
|
||||
.context("OpenAI response interrupted")?
|
||||
{
|
||||
anyhow::ensure!(
|
||||
bytes.len().saturating_add(chunk.len()) <= RESPONSE_LIMIT,
|
||||
"OpenAI response exceeded the size limit"
|
||||
);
|
||||
bytes.extend_from_slice(&chunk);
|
||||
}
|
||||
parse_response(
|
||||
&serde_json::from_slice(&bytes).context("OpenAI returned an invalid response")?,
|
||||
)
|
||||
}
|
||||
}
|
||||
#[async_trait]
|
||||
impl Backend for OpenaiBackend {
|
||||
async fn send(
|
||||
&self,
|
||||
system: &str,
|
||||
tools: &[ToolDef],
|
||||
history: &[ChatMessage],
|
||||
) -> Result<BackendTurn> {
|
||||
self.send_at(URL, system, tools, history).await
|
||||
}
|
||||
}
|
||||
fn error_message(status: u16) -> &'static str {
|
||||
match status {
|
||||
401 | 403 => "OpenAI rejected the API key or project access. Check AI connection settings.",
|
||||
404 => "This OpenAI model is unavailable for your account. Choose another model in AI connection settings.",
|
||||
429 => "OpenAI usage or rate limit reached. Check your API billing and retry later.",
|
||||
500..=599 => "OpenAI is temporarily unavailable. Retry later.",
|
||||
_ => "OpenAI rejected the request. Check the selected model and retry.",
|
||||
}
|
||||
}
|
||||
fn parse_response(value: &Value) -> Result<BackendTurn> {
|
||||
let choice = value["choices"]
|
||||
.as_array()
|
||||
.and_then(|items| items.first())
|
||||
.context("OpenAI returned no answer")?;
|
||||
anyhow::ensure!(
|
||||
choice["finish_reason"] != "length",
|
||||
"OpenAI reached the response limit. Try a shorter request."
|
||||
);
|
||||
let message = &choice["message"];
|
||||
if let Some(calls) = message["tool_calls"]
|
||||
.as_array()
|
||||
.filter(|calls| !calls.is_empty())
|
||||
{
|
||||
let mut parsed = Vec::new();
|
||||
for call in calls {
|
||||
anyhow::ensure!(
|
||||
call["type"] == "function",
|
||||
"Unsupported OpenAI tool response"
|
||||
);
|
||||
let id = call["id"]
|
||||
.as_str()
|
||||
.filter(|id| !id.is_empty())
|
||||
.context("Missing OpenAI tool call ID")?;
|
||||
let name = call["function"]["name"]
|
||||
.as_str()
|
||||
.filter(|name| !name.is_empty())
|
||||
.context("Missing OpenAI tool name")?;
|
||||
let arguments: Value = serde_json::from_str(
|
||||
call["function"]["arguments"]
|
||||
.as_str()
|
||||
.context("Invalid OpenAI tool arguments")?,
|
||||
)
|
||||
.context("Invalid OpenAI tool arguments")?;
|
||||
anyhow::ensure!(
|
||||
arguments.is_object()
|
||||
&& !parsed.iter().any(|previous: &ToolCall| previous.id == id),
|
||||
"Invalid OpenAI tool call"
|
||||
);
|
||||
parsed.push(ToolCall {
|
||||
id: id.into(),
|
||||
name: name.into(),
|
||||
arguments,
|
||||
});
|
||||
}
|
||||
return Ok(BackendTurn::ToolCalls(parsed));
|
||||
}
|
||||
let text = message["content"]
|
||||
.as_str()
|
||||
.or_else(|| message["refusal"].as_str())
|
||||
.filter(|text| !text.trim().is_empty())
|
||||
.context("OpenAI returned no text; check model compatibility")?;
|
||||
Ok(BackendTurn::Text(text.into()))
|
||||
}
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::assistant::tools::{Role, ToolResult};
|
||||
#[test]
|
||||
fn parses_text_and_rejects_incomplete_or_malformed_tool_calls() {
|
||||
assert!(
|
||||
matches!(parse_response(&json!({"choices":[{"message":{"content":"hello"}}]})).unwrap(), BackendTurn::Text(text) if text == "hello")
|
||||
);
|
||||
let valid = json!({"choices":[{"message":{"tool_calls":[{"id":"call_1","type":"function","function":{"name":"status","arguments":"{\"count\":1}"}}]}}]});
|
||||
assert!(
|
||||
matches!(parse_response(&valid).unwrap(), BackendTurn::ToolCalls(calls) if calls[0].arguments["count"] == 1)
|
||||
);
|
||||
for args in ["{", "null", "[]"] {
|
||||
let mut invalid = valid.clone();
|
||||
invalid["choices"][0]["message"]["tool_calls"][0]["function"]["arguments"] =
|
||||
json!(args);
|
||||
assert!(parse_response(&invalid).is_err());
|
||||
}
|
||||
assert!(parse_response(
|
||||
&json!({"choices":[{"finish_reason":"length","message":{"content":"partial"}}]})
|
||||
)
|
||||
.is_err());
|
||||
assert!(parse_response(&json!({"choices":[]})).is_err());
|
||||
}
|
||||
#[test]
|
||||
fn errors_distinguish_credentials_limits_and_outages_without_raw_provider_data() {
|
||||
assert!(error_message(401).contains("API key"));
|
||||
assert!(error_message(429).contains("limit"));
|
||||
assert!(!error_message(503).contains("key"));
|
||||
let wire = super::super::routstr::message_to_wire(&ChatMessage {
|
||||
role: Role::Tool,
|
||||
text: None,
|
||||
tool_calls: vec![],
|
||||
tool_results: vec![ToolResult {
|
||||
call_id: "call_1".into(),
|
||||
content: "result".into(),
|
||||
is_error: false,
|
||||
}],
|
||||
});
|
||||
assert_eq!(wire[0]["tool_call_id"], "call_1");
|
||||
}
|
||||
#[tokio::test]
|
||||
async fn real_http_adapter_sends_private_key_only_in_header_and_never_follows_redirect() {
|
||||
use hyper::{
|
||||
service::{make_service_fn, service_fn},
|
||||
Body, Response, Server,
|
||||
};
|
||||
use std::sync::{Arc, Mutex};
|
||||
let captured = Arc::new(Mutex::new(Vec::new()));
|
||||
let capture = captured.clone();
|
||||
let server = Server::bind(&([127, 0, 0, 1], 0).into()).serve(make_service_fn(move |_| {
|
||||
let capture = capture.clone();
|
||||
async move {
|
||||
Ok::<_, hyper::Error>(service_fn(move |request: hyper::Request<Body>| {
|
||||
let capture = capture.clone();
|
||||
async move {
|
||||
let (parts, body) = request.into_parts();
|
||||
let body = hyper::body::to_bytes(body).await?;
|
||||
capture.lock().unwrap().push((
|
||||
parts.headers,
|
||||
serde_json::from_slice::<Value>(&body).unwrap(),
|
||||
));
|
||||
Ok::<_, hyper::Error>(
|
||||
Response::builder()
|
||||
.status(302)
|
||||
.header("Location", "/leak")
|
||||
.body(Body::empty())
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
}))
|
||||
}
|
||||
}));
|
||||
let url = format!("http://{}/v1/chat/completions", server.local_addr());
|
||||
let task = tokio::spawn(server);
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
crate::settings::model_provider::save_key(dir.path(), "openai", "fixture-private-key")
|
||||
.await
|
||||
.unwrap();
|
||||
let backend = OpenaiBackend::new(dir.path().into(), "test-model".into());
|
||||
let history = [ChatMessage {
|
||||
role: Role::User,
|
||||
text: Some("Hello".into()),
|
||||
tool_calls: vec![],
|
||||
tool_results: vec![],
|
||||
}];
|
||||
assert!(backend
|
||||
.send_at(&url, "Be helpful", &[], &history)
|
||||
.await
|
||||
.is_err());
|
||||
let requests = captured.lock().unwrap();
|
||||
assert_eq!(requests.len(), 1);
|
||||
assert_eq!(requests[0].0["authorization"], "Bearer fixture-private-key");
|
||||
assert_eq!(requests[0].1["store"], false);
|
||||
assert_eq!(requests[0].1["max_completion_tokens"], 2048);
|
||||
assert!(!requests[0].1.to_string().contains("fixture-private-key"));
|
||||
drop(requests);
|
||||
let private = [ChatMessage {
|
||||
role: Role::User,
|
||||
text: Some("fixture-private-key".into()),
|
||||
tool_calls: vec![],
|
||||
tool_results: vec![],
|
||||
}];
|
||||
assert!(backend
|
||||
.send_at(&url, "Be helpful", &[], &private)
|
||||
.await
|
||||
.is_err());
|
||||
assert_eq!(captured.lock().unwrap().len(), 1);
|
||||
task.abort();
|
||||
}
|
||||
}
|
||||
@@ -295,7 +295,7 @@ fn parse_openai_tool_calls(raw_calls: &[Value]) -> Vec<ToolCall> {
|
||||
/// (the wire-format inverse of `parse_openai_tool_calls`), and tool-result
|
||||
/// turns carry `tool_call_id` so each call's id is echoed back exactly —
|
||||
/// the OpenAI-shape contract this adapter's edge is responsible for.
|
||||
fn message_to_wire(msg: &ChatMessage) -> Vec<Value> {
|
||||
pub(super) fn message_to_wire(msg: &ChatMessage) -> Vec<Value> {
|
||||
match msg.role {
|
||||
Role::System => vec![],
|
||||
Role::User => vec![json!({
|
||||
|
||||
@@ -9,3 +9,5 @@ pub mod session_policy;
|
||||
pub mod transport;
|
||||
|
||||
pub mod bitcoin_storage;
|
||||
|
||||
pub mod model_provider;
|
||||
|
||||
@@ -0,0 +1,178 @@
|
||||
//! Owner-selected chat provider. API keys remain in the node's private secret
|
||||
//! ledger and are never returned by settings or included in chat context.
|
||||
use anyhow::{Context, Result};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::path::Path;
|
||||
use tokio::{fs, io::AsyncWriteExt};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum Provider {
|
||||
#[default]
|
||||
Auto,
|
||||
Claude,
|
||||
Openai,
|
||||
Local,
|
||||
Routstr,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize, Serialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct ModelProvider {
|
||||
#[serde(default)]
|
||||
pub provider: Provider,
|
||||
#[serde(default)]
|
||||
pub openai_model: String,
|
||||
}
|
||||
|
||||
impl ModelProvider {
|
||||
pub fn validate(&self) -> Result<()> {
|
||||
anyhow::ensure!(
|
||||
!self.openai_model.starts_with("sk-")
|
||||
&& self.openai_model.len() <= 128
|
||||
&& self
|
||||
.openai_model
|
||||
.bytes()
|
||||
.all(|b| b.is_ascii_alphanumeric() || b"-_.:".contains(&b)),
|
||||
"Invalid OpenAI model name"
|
||||
);
|
||||
anyhow::ensure!(
|
||||
self.provider != Provider::Openai || !self.openai_model.is_empty(),
|
||||
"Choose an OpenAI model before connecting"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
pub async fn load(data_dir: &Path) -> Result<Self> {
|
||||
match fs::read(data_dir.join("settings/model-provider.json")).await {
|
||||
Ok(bytes) => {
|
||||
let settings: Self = serde_json::from_slice(&bytes)
|
||||
.context("Invalid AI provider settings; preserved for recovery")?;
|
||||
settings.validate()?;
|
||||
Ok(settings)
|
||||
}
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(Self::default()),
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
pub async fn save(&self, data_dir: &Path) -> Result<()> {
|
||||
self.validate()?;
|
||||
write_private(
|
||||
&data_dir.join("settings/model-provider.json"),
|
||||
&serde_json::to_vec(self)?,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
pub fn key_name(provider: &str) -> Result<&'static str> {
|
||||
match provider {
|
||||
"claude" => Ok("claude-api-key"),
|
||||
"openai" => Ok("openai-api-key"),
|
||||
_ => anyhow::bail!("Unsupported AI provider"),
|
||||
}
|
||||
}
|
||||
pub async fn has_key(data_dir: &Path, provider: &str) -> bool {
|
||||
let Ok(name) = key_name(provider) else {
|
||||
return false;
|
||||
};
|
||||
fs::read_to_string(data_dir.join("secrets").join(name))
|
||||
.await
|
||||
.is_ok_and(|key| !key.trim().is_empty())
|
||||
}
|
||||
pub async fn save_key(data_dir: &Path, provider: &str, value: &str) -> Result<()> {
|
||||
let path = data_dir.join("secrets").join(key_name(provider)?);
|
||||
let value = value.trim();
|
||||
anyhow::ensure!(
|
||||
value.len() <= 4096 && value.bytes().all(|b| b.is_ascii_graphic()),
|
||||
"Invalid API key format"
|
||||
);
|
||||
if value.is_empty() {
|
||||
match fs::remove_file(path).await {
|
||||
Ok(()) => Ok(()),
|
||||
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
|
||||
Err(e) => Err(e.into()),
|
||||
}
|
||||
} else {
|
||||
write_private(&path, value.as_bytes()).await
|
||||
}
|
||||
}
|
||||
async fn write_private(path: &Path, bytes: &[u8]) -> Result<()> {
|
||||
let parent = path.parent().context("Missing settings directory")?;
|
||||
fs::create_dir_all(parent).await?;
|
||||
let temporary = parent.join(format!(".provider-{}.tmp", uuid::Uuid::new_v4()));
|
||||
let result = async {
|
||||
let mut file = fs::OpenOptions::new()
|
||||
.create_new(true)
|
||||
.write(true)
|
||||
.mode(0o600)
|
||||
.open(&temporary)
|
||||
.await?;
|
||||
file.write_all(bytes).await?;
|
||||
file.sync_all().await?;
|
||||
drop(file);
|
||||
fs::rename(&temporary, path).await?;
|
||||
fs::File::open(parent).await?.sync_all().await?;
|
||||
Ok::<_, anyhow::Error>(())
|
||||
}
|
||||
.await;
|
||||
if result.is_err() {
|
||||
let _ = fs::remove_file(temporary).await;
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[tokio::test]
|
||||
async fn private_keys_replace_atomically_and_never_enter_public_settings() {
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
assert!(!has_key(dir.path(), "openai").await);
|
||||
save_key(dir.path(), "openai", "test-key-one")
|
||||
.await
|
||||
.unwrap();
|
||||
save_key(dir.path(), "openai", "test-key-two")
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(has_key(dir.path(), "openai").await);
|
||||
let key_path = dir.path().join("secrets/openai-api-key");
|
||||
assert_eq!(
|
||||
fs::metadata(&key_path).await.unwrap().permissions().mode() & 0o777,
|
||||
0o600
|
||||
);
|
||||
assert_eq!(fs::read_to_string(&key_path).await.unwrap(), "test-key-two");
|
||||
let settings = ModelProvider {
|
||||
provider: Provider::Openai,
|
||||
openai_model: "test-model".into(),
|
||||
};
|
||||
settings.save(dir.path()).await.unwrap();
|
||||
let body = serde_json::to_string(&ModelProvider::load(dir.path()).await.unwrap()).unwrap();
|
||||
assert!(!body.contains("test-key"));
|
||||
assert!(save_key(dir.path(), "../openai", "key").await.is_err());
|
||||
assert!(save_key(dir.path(), "openai", "key\nInjected: bad")
|
||||
.await
|
||||
.is_err());
|
||||
assert_eq!(fs::read_to_string(&key_path).await.unwrap(), "test-key-two");
|
||||
save_key(dir.path(), "openai", "").await.unwrap();
|
||||
assert!(!has_key(dir.path(), "openai").await);
|
||||
}
|
||||
#[tokio::test]
|
||||
async fn invalid_settings_preserve_existing_configuration() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
ModelProvider::default().save(dir.path()).await.unwrap();
|
||||
let invalid = ModelProvider {
|
||||
provider: Provider::Openai,
|
||||
openai_model: String::new(),
|
||||
};
|
||||
assert!(invalid.save(dir.path()).await.is_err());
|
||||
assert_eq!(
|
||||
ModelProvider::load(dir.path()).await.unwrap().provider,
|
||||
Provider::Auto
|
||||
);
|
||||
let path = dir.path().join("settings/model-provider.json");
|
||||
fs::write(&path, b"broken").await.unwrap();
|
||||
assert!(ModelProvider::load(dir.path()).await.is_err());
|
||||
assert_eq!(fs::read(&path).await.unwrap(), b"broken");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user