use super::DaemonState;
use crate::accounts::{AccountConfig, AccountManager};
use crate::broadcast::SubscriberSink;
use crate::db;
use choreo_keystore::ServiceCredential;
use choreo_proto::DaemonMessage;
use std::collections::HashMap;
use std::io;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::mpsc;
use tracing::{debug, error, info, warn};
use zeroize::{Zeroize, Zeroizing};
#[derive(Debug, thiserror::Error)]
pub enum KeystoreOpError {
#[error(
"keystore not initialized — no key is bound to this daemon yet; it will be bound automatically on next client connect"
)]
Unbound,
#[error("{0}")]
Other(String),
}
impl DaemonState {
fn send_targeted(
writer: Option<&SubscriberSink>,
global_lag: &Arc<AtomicUsize>,
msg: &DaemonMessage,
) {
if let Some(w) = writer {
w.send_accounted(msg, global_lag);
} else {
warn!(
?msg,
"no client writer for targeted keystore reply; dropping reply"
);
}
}
fn send_targeted_ack(
writer: Option<&SubscriberSink>,
global_lag: &Arc<AtomicUsize>,
msg: &DaemonMessage,
reply: &mpsc::Sender<()>,
) {
Self::send_targeted(writer, global_lag, msg);
let _ = reply.send(());
}
fn send_credential_add_failed(
&self,
service: &str,
error: String,
client_writer: Option<&SubscriberSink>,
reply: &mpsc::Sender<()>,
) {
Self::send_targeted_ack(
client_writer,
&self.global_lag,
&DaemonMessage::CredentialAddFailed {
service: service.to_string(),
error,
},
reply,
);
}
pub(super) fn handle_unlock(
&mut self,
private_key: Vec<u8>,
client_writer: Option<&SubscriberSink>,
reply: &mpsc::Sender<()>,
) {
info!("Unlock attempt");
let was_locked = self.locked;
let result = handle_unlock_inner(self, private_key);
let reply_msg = match &result {
Ok(()) => DaemonMessage::Unlocked,
Err(e) => unlock_error_reply(e),
};
Self::send_targeted_ack(client_writer, &self.global_lag, &reply_msg, reply);
if result.is_ok() && was_locked {
self.broadcast_keystore_state();
}
info!("Unlock result: success={}", result.is_ok());
}
pub(super) fn handle_bind_keystore(
&mut self,
key: Vec<u8>,
client_writer: Option<&SubscriberSink>,
reply: &mpsc::Sender<()>,
) {
info!("BindKeystore attempt");
let was_locked = self.locked;
let result = handle_bind_keystore_inner(self, key);
let reply_msg = match &result {
Ok(()) => DaemonMessage::Bound,
Err(e) => unlock_error_reply(e),
};
Self::send_targeted_ack(client_writer, &self.global_lag, &reply_msg, reply);
if result.is_ok() && was_locked {
self.broadcast_keystore_state();
}
info!("BindKeystore result: success={}", result.is_ok());
}
pub(super) fn handle_lock(&mut self, reply: &mpsc::Sender<Result<(), String>>) {
let was_locked = self.locked;
let credentials_cleared = self.credentials.len();
self.credentials.clear();
self.drop_session_clients(None);
self.x_credentials = None;
self.locked = true;
info!(
credentials_cleared,
"keystore locked: in-memory credentials cleared"
);
if !was_locked && self.keystore_bound {
self.broadcast_keystore_state();
}
let _ = reply.send(Ok(()));
}
pub(super) fn handle_save_credential(
&mut self,
service: String,
encrypted_blob: &[u8],
mut unlock_key: Vec<u8>,
client_writer: Option<&SubscriberSink>,
reply: &mpsc::Sender<()>,
) {
let was_locked = self.locked;
let key = Zeroizing::new(if let Ok(k) = unlock_key.as_slice().try_into() {
k
} else {
unlock_key.zeroize();
self.send_credential_add_failed(
&service,
"invalid unlock_key: expected exactly 32 bytes".to_string(),
client_writer,
reply,
);
return;
});
unlock_key.zeroize();
if let Err(e) = verify_keystore_binding(self, &key) {
let reply_msg = match e {
KeystoreOpError::Unbound => DaemonMessage::KeystoreUnbound {
error: KeystoreOpError::Unbound.to_string(),
},
KeystoreOpError::Other(e) => DaemonMessage::CredentialAddFailed {
service: service.clone(),
error: e,
},
};
Self::send_targeted_ack(client_writer, &self.global_lag, &reply_msg, reply);
return;
}
let plaintext =
match choreo_keystore::crypto::decrypt_with_private_key(&key, encrypted_blob) {
Ok(pt) => pt,
Err(e) => {
warn!(
service = %service,
error = %e,
"AddCredential: blob failed test-decrypt with the presented unlock key; \
rejecting without persisting"
);
self.send_credential_add_failed(
&service,
format!(
"credential blob failed to decrypt with the provided unlock key: {e}"
),
client_writer,
reply,
);
return;
}
};
let cred: ServiceCredential = match postcard::from_bytes(&plaintext) {
Ok(c) => c,
Err(e) => {
self.send_credential_add_failed(
&service,
format!("credential payload is not a valid ServiceCredential: {e}"),
client_writer,
reply,
);
return;
}
};
if let Err(e) = db::set_credential_blob(&self.db, &service, encrypted_blob) {
self.send_credential_add_failed(
&service,
format!("failed to save credential: {e}"),
client_writer,
reply,
);
return;
}
if matches!(&cred, ServiceCredential::X { .. }) && service == "twitter" {
self.x_credentials = Some(cred.clone());
}
if matches!(&cred, ServiceCredential::ApiKey { .. }) {
self.drop_session_clients(Some(&service));
}
let result = unlock_tail_skip(self, &key, &service, cred.clone());
if let Err(e) = result {
error!(
service = %service,
error = %e,
"AddCredential: persisted credential but implicit unlock failed"
);
self.send_credential_add_failed(
&service,
format!("credential saved but unlock failed: {e}"),
client_writer,
reply,
);
return;
}
info!(
service = %service,
"AddCredential: persisted, tested, and implicitly unlocked the keystore"
);
Self::send_targeted_ack(
client_writer,
&self.global_lag,
&DaemonMessage::Unlocked,
reply,
);
Self::send_targeted(
client_writer,
&self.global_lag,
&DaemonMessage::CredentialAdded { service },
);
if was_locked && !self.locked {
self.broadcast_keystore_state();
}
}
}
fn handle_unlock_inner(
state: &mut DaemonState,
mut private_key: Vec<u8>,
) -> Result<(), KeystoreOpError> {
let key = zeroized_key_or_wipe(&mut private_key)?;
private_key.zeroize();
verify_keystore_binding(state, &key)?;
unlock_tail(state, &key).map_err(|e| KeystoreOpError::Other(e.to_string()))
}
fn handle_bind_keystore_inner(
state: &mut DaemonState,
mut key: Vec<u8>,
) -> Result<(), KeystoreOpError> {
let key = {
let arr = zeroized_key_or_wipe(&mut key)?;
key.zeroize();
arr
};
bind_keystore(state, &key)?;
unlock_tail(state, &key).map_err(|e| KeystoreOpError::Other(e.to_string()))
}
fn zeroized_key_or_wipe(key: &mut Vec<u8>) -> Result<Zeroizing<[u8; 32]>, KeystoreOpError> {
if let Ok(k) = key.as_slice().try_into() {
Ok(Zeroizing::new(k))
} else {
key.zeroize();
Err(KeystoreOpError::Other(
"invalid key: expected exactly 32 bytes".to_string(),
))
}
}
fn unlock_error_reply(e: &KeystoreOpError) -> DaemonMessage {
match e {
KeystoreOpError::Unbound => DaemonMessage::KeystoreUnbound {
error: KeystoreOpError::Unbound.to_string(),
},
KeystoreOpError::Other(e) => DaemonMessage::LockedError { error: e.clone() },
}
}
fn bind_keystore(state: &mut DaemonState, key: &[u8; 32]) -> Result<(), KeystoreOpError> {
let binding = db::get_keystore_binding(&state.db)
.map_err(|e| KeystoreOpError::Other(format!("failed to read keystore binding: {e}")))?;
let derived = x25519_dalek::PublicKey::from(&x25519_dalek::StaticSecret::from(*key));
match binding {
None => {
db::set_keystore_binding(&state.db, derived.as_bytes()).map_err(|e| {
KeystoreOpError::Other(format!("failed to persist keystore binding: {e}"))
})?;
info!(
"KEYSTORE BOUND: adopted unlock key via BindKeystore (TOFU); \
public key (hex) = {} — all future Unlock/AddCredential \
attempts must present this key, others are rejected",
hex::encode(derived.as_bytes())
);
state.keystore_bound = true;
Ok(())
}
Some(stored) if stored == *derived.as_bytes() => {
debug!("BindKeystore: key matches the existing binding");
Ok(())
}
Some(_) => Err(KeystoreOpError::Other(
"keystore is already bound and the presented key does not match; \
refusing to overwrite the binding"
.to_string(),
)),
}
}
pub(crate) fn verify_keystore_binding(
state: &DaemonState,
key: &[u8; 32],
) -> Result<(), KeystoreOpError> {
let binding = db::get_keystore_binding(&state.db)
.map_err(|e| KeystoreOpError::Other(format!("failed to read keystore binding: {e}")))?;
let derived = x25519_dalek::PublicKey::from(&x25519_dalek::StaticSecret::from(*key));
match binding {
None => {
debug!(
"verify_keystore_binding: keystore has no binding; refusing verify-only operation"
);
Err(KeystoreOpError::Unbound)
}
Some(stored) if stored == *derived.as_bytes() => Ok(()),
Some(_) => Err(KeystoreOpError::Other(
"unlock key does not match the daemon's keystore binding".to_string(),
)),
}
}
pub(crate) fn unlock_tail(state: &mut DaemonState, key: &[u8; 32]) -> io::Result<()> {
let credentials = decrypt_credential_blobs(state, key, None)?;
finish_unlock(state, credentials)
}
pub(crate) fn unlock_tail_skip(
state: &mut DaemonState,
key: &[u8; 32],
skip_service: &str,
seeded: ServiceCredential,
) -> io::Result<()> {
let mut credentials = decrypt_credential_blobs(state, key, Some(skip_service))?;
credentials.insert(skip_service.to_string(), seeded);
finish_unlock(state, credentials)
}
fn decrypt_credential_blobs(
state: &DaemonState,
key: &[u8; 32],
skip: Option<&str>,
) -> io::Result<HashMap<String, ServiceCredential>> {
let blobs = db::get_all_credential_blobs(&state.db)
.map_err(|e| io::Error::other(format!("failed to read credentials from database: {e}")))?;
info!("Unlock: {} credential blobs in DB", blobs.len());
let mut credentials = HashMap::new();
let mut decrypt_failures = 0usize;
for (service, blob) in &blobs {
if skip == Some(service.as_str()) {
continue;
}
match choreo_keystore::crypto::decrypt_with_private_key(key, blob) {
Ok(plaintext) => match postcard::from_bytes::<ServiceCredential>(&plaintext) {
Ok(cred) => {
credentials.insert(service.clone(), cred);
}
Err(e) => {
warn!("Unlock: failed to decode credential '{}': {e}", service);
decrypt_failures += 1;
}
},
Err(e) => {
warn!("Unlock: failed to decrypt credential '{}': {e}", service);
decrypt_failures += 1;
}
}
}
let skipped = usize::from(skip.is_some());
info!(
"Unlock: decrypted {}/{} credentials ({} failures): {:?}",
credentials.len(),
blobs.len() - skipped,
decrypt_failures,
credentials.keys().collect::<Vec<_>>()
);
Ok(credentials)
}
fn finish_unlock(
state: &mut DaemonState,
credentials: HashMap<String, ServiceCredential>,
) -> io::Result<()> {
let accounts_path = state.accounts.path().to_path_buf();
let mut accounts = AccountManager::load(&accounts_path)
.map_err(|e| io::Error::other(format!("failed to load accounts: {e}")))?;
if accounts.is_empty() && credentials.contains_key("openai") {
let default_config = AccountConfig::simple("default", "openai");
if let Err(e) = accounts.add(default_config) {
tracing::warn!("failed to create default account: {e}");
}
}
state.x_credentials = credentials
.get("twitter")
.filter(|c| matches!(c, ServiceCredential::X { .. }))
.cloned();
state.credentials = credentials;
state.accounts = accounts;
state.locked = false;
let account_names: Vec<String> = state
.accounts
.all_configs()
.iter()
.map(|c| c.name.clone())
.collect();
info!("Unlock: accounts loaded: {:?}", account_names);
for config in state.accounts.all_configs() {
info!(
"Unlock: account '{}': has_credential={}",
config.name,
state.credentials.contains_key(&config.name)
);
}
info!("Unlock: keystore decrypted; sessions will rebuild providers lazily on next use");
Ok(())
}
#[cfg(test)]
mod tests {
use super::super::tests::{make_daemon_state, test_pub};
use super::*;
#[test]
fn failed_unlock_tail_publishes_no_decrypted_state() {
let (mut state, _rx) = make_daemon_state();
let key: [u8; 32] = [7u8; 32];
let blob = choreo_keystore::crypto::encrypt_with_public_key(
&test_pub(key),
&postcard::to_allocvec(&ServiceCredential::ApiKey {
key: "sk-test".to_string(),
})
.unwrap(),
)
.unwrap();
db::set_credential_blob(&state.db, "openai", &blob).unwrap();
let accounts_path = state.accounts.path().to_path_buf();
std::fs::write(&accounts_path, "definitely not valid TOML [[[").unwrap();
let err = unlock_tail(&mut state, &key)
.expect_err("an unloadable accounts file must fail the unlock tail");
assert!(err.to_string().contains("accounts"), "got: {err}");
assert!(
state.credentials.is_empty(),
"a failed tail must not publish decrypted credentials"
);
assert!(state.x_credentials.is_none(), "X creds must stay empty");
assert!(state.locked, "a failed tail must not unlock the daemon");
assert!(state.accounts.is_empty(), "accounts must stay as they were");
}
#[test]
fn unlock_tail_skip_seeds_the_tested_credential_over_the_db_blob() {
let (mut state, _rx) = make_daemon_state();
let key: [u8; 32] = [9u8; 32];
let db_blob = choreo_keystore::crypto::encrypt_with_public_key(
&test_pub(key),
&postcard::to_allocvec(&ServiceCredential::ApiKey {
key: "stale-db-value".to_string(),
})
.unwrap(),
)
.unwrap();
db::set_credential_blob(&state.db, "openai", &db_blob).unwrap();
let seeded = ServiceCredential::ApiKey {
key: "fresh-tested-value".to_string(),
};
unlock_tail_skip(&mut state, &key, "openai", seeded).unwrap();
assert!(!state.locked, "the skip tail still unlocks");
match state.credentials.get("openai") {
Some(ServiceCredential::ApiKey { key }) => {
assert_eq!(key, "fresh-tested-value", "the seed wins");
}
other => panic!("expected the seeded openai credential, got {other:?}"),
}
}
}