use crate::config::load_daemon_config;
use crate::daemon::DaemonState;
use anyhow::Context;
use choreo_proto::socket_path;
use choreo_transport::key::ensure_transport_keypair;
use clap::Parser;
use tracing::{info, warn};
use tracing_subscriber::{EnvFilter, fmt};
fn clap_styles() -> clap::builder::Styles {
use clap::builder::styling::{AnsiColor, Effects, Styles};
Styles::styled()
.header(AnsiColor::Green.on_default() | Effects::BOLD)
.usage(AnsiColor::Green.on_default() | Effects::BOLD)
.literal(AnsiColor::Cyan.on_default() | Effects::BOLD)
.placeholder(AnsiColor::Cyan.on_default())
}
#[derive(Parser)]
#[command(
name = "choreographr",
version = choreo_proto::release_name::version_string(env!("CARGO_PKG_VERSION")),
about = "Choreographr AI daemon",
color = clap::ColorChoice::Auto,
styles = clap_styles()
)]
struct Cli {
#[arg(short = 'v', long = "verbose", action = clap::ArgAction::Count)]
verbose: u8,
#[arg(short = 'q', long = "quiet", action = clap::ArgAction::Count)]
quiet: u8,
#[arg(long = "metrics-addr")]
metrics_addr: Option<String>,
#[arg(long = "tcp-addr")]
tcp_addr: Option<String>,
#[arg(long = "log-file")]
log_file: Option<String>,
#[arg(long = "auto-exit")]
auto_exit: bool,
#[command(subcommand)]
command: Option<Command>,
}
#[derive(clap::Subcommand)]
enum Command {
AclAdd {
pubkey: String,
},
Fingerprint {
#[arg(default_value = None)]
path: Option<String>,
},
}
fn acl_add_to(path: &std::path::Path, pubkey_b64: &str) -> anyhow::Result<usize> {
use base64::Engine as _;
let key: [u8; 32] = base64::engine::general_purpose::STANDARD
.decode(pubkey_b64.trim())
.map_err(|e| anyhow::anyhow!("invalid pubkey: not valid base64: {e}"))?
.try_into()
.map_err(|_| anyhow::anyhow!("invalid pubkey: must decode to exactly 32 bytes"))?;
let existing = crate::server::acl::Acl::load(path);
if existing.contains(&key) {
info!("pubkey is already authorized; nothing to do");
return Ok(existing.len());
}
crate::server::acl::append_key_locked(path, &key).map_err(|e| anyhow::anyhow!(e))?;
let count = crate::server::acl::Acl::load(path).len();
info!(clients = count, "ACL: client key enrolled");
Ok(count)
}
fn fingerprint_cli(path: Option<&str>) -> anyhow::Result<()> {
use choreo_transport::key::{fingerprint, fingerprint_of_file, read_server_pk};
let fp = match path {
Some(p) => fingerprint_of_file(std::path::Path::new(p))
.with_context(|| format!("failed to fingerprint key file {p}"))?,
None => {
let pk = read_server_pk(None).context(
"failed to read this machine's transport public key (has the daemon ever run here?)",
)?;
fingerprint(&pk)
}
};
println!("{fp}");
Ok(())
}
const DEFAULT_MAX_TURNS: u32 = 0;
fn resolve_max_turns() -> anyhow::Result<u32> {
match std::env::var("CHOREOGRAPHR_MAX_TURNS") {
Ok(val) => return parse_max_turns_env(&val),
Err(std::env::VarError::NotPresent) => {}
Err(e) => {
return Err(anyhow::anyhow!(
"failed to read CHOREOGRAPHR_MAX_TURNS: {e}"
));
}
}
if let Ok(config) = load_daemon_config()
&& let Some(n) = config.max_turns
{
return Ok(n);
}
Ok(DEFAULT_MAX_TURNS)
}
fn parse_max_turns_env(val: &str) -> anyhow::Result<u32> {
val.parse::<u32>()
.map_err(|e| anyhow::anyhow!("CHOREOGRAPHR_MAX_TURNS={val:?} is not a valid u32: {e}"))
}
#[cfg(unix)]
fn open_log_file(path: &str) -> anyhow::Result<std::fs::File> {
use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt};
let file = std::fs::OpenOptions::new()
.create(true)
.append(true)
.mode(0o600)
.custom_flags(rustix::fs::OFlags::NOFOLLOW.bits() as i32)
.open(path)
.with_context(|| {
format!(
"failed to open --log-file {path} for writing; check that the \
directory exists and is writable"
)
})?;
let meta = file
.metadata()
.with_context(|| format!("failed to stat --log-file {path}"))?;
let euid = rustix::process::geteuid().as_raw();
if !meta.is_file() || meta.uid() != euid {
anyhow::bail!(
"refusing to write --log-file {path}: it is not a regular file owned by the \
current user"
);
}
file.set_permissions(std::fs::Permissions::from_mode(0o600))
.with_context(|| format!("failed to set 0600 on --log-file {path}"))?;
Ok(file)
}
#[cfg(not(unix))]
fn open_log_file(path: &str) -> anyhow::Result<std::fs::File> {
std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(path)
.with_context(|| {
format!(
"failed to open --log-file {path} for writing; check that the \
directory exists and is writable"
)
})
}
pub fn main() -> anyhow::Result<()> {
let cli = Cli::parse();
let log_level = if std::env::var("RUST_LOG").is_ok() {
if cli.verbose > 0 || cli.quiet > 0 {
warn!("RUST_LOG is set; -v/-q CLI flags are ignored");
}
None } else {
let level = match (cli.verbose, cli.quiet) {
(0, 0) => "info",
(_, q) if q > 0 => "warn",
(1, 0) => "debug",
_ => "trace",
};
Some(level)
};
let env_filter = match log_level {
Some(level) => EnvFilter::new(level),
None => EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info")),
};
if let Some(path) = &cli.log_file {
let file = open_log_file(path)?;
fmt()
.with_env_filter(env_filter)
.with_ansi(false)
.with_writer(std::sync::Mutex::new(file))
.init();
} else {
fmt().with_env_filter(env_filter).init();
}
info!(effective_level = ?log_level.unwrap_or("from RUST_LOG"), "logging initialized");
match &cli.command {
Some(Command::AclAdd { pubkey }) => {
let path = choreo_keystore::paths::authorized_clients_path()
.context("failed to resolve authorized_clients path")?;
let count = acl_add_to(&path, pubkey)?;
println!("client key authorized ({count} client(s) now trusted)");
return Ok(());
}
Some(Command::Fingerprint { path }) => return fingerprint_cli(path.as_deref()),
None => {}
}
#[cfg(feature = "blockchain")]
choreo_blockchain::runtime::init()
.map_err(|e| anyhow::anyhow!("failed to initialize blockchain tokio runtime: {e}"))?;
#[cfg(feature = "content")]
match choreo_content::init() {
Ok(()) => {}
Err(e) => {
warn!(
error = %e,
"failed to initialize the coordination platform tokio runtime; \
content write tools will be unavailable"
);
}
}
let max_turns = resolve_max_turns().context("failed to resolve tool-loop iteration limit")?;
info!(max_turns, "tool loop iteration limit");
info!(
version = %choreo_proto::release_name::version_string(env!("CARGO_PKG_VERSION")),
"choreographr starting (locked)"
);
let state = DaemonState::open(crate::daemon::OpenOptions {
db_path: crate::db::db_path().context("failed to resolve database path")?,
accounts_path: crate::accounts::accounts_config_path()
.context("failed to resolve accounts config path")?,
catalog_paths: crate::catalog::CatalogPaths::from_dirs(),
tool_policy: crate::tools::ToolPolicy::Full,
max_turns,
platform_tool_bridge: None,
})
.context("failed to open daemon state")?;
let (transport_sk, _transport_pk) =
ensure_transport_keypair().context("failed to load/generate transport keypair")?;
let acl_path = choreo_keystore::paths::authorized_clients_path()
.context("failed to resolve authorized_clients path")?;
let acl = crate::server::acl::SharedAcl::load(&acl_path);
let socket_path = socket_path();
crate::run_server(
&socket_path,
state,
cli.metrics_addr,
cli.tcp_addr,
transport_sk,
acl,
cli.auto_exit,
)
.context("failed to run server")
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(unix)]
#[test]
fn open_log_file_creates_with_0600() {
use std::os::unix::fs::PermissionsExt;
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("daemon.log");
let file = open_log_file(path.to_str().unwrap()).unwrap();
drop(file);
let mode = std::fs::metadata(&path).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o600, "fresh daemon logs must be owner-only");
}
#[cfg(unix)]
#[test]
fn open_log_file_tightens_a_preexisting_loose_file() {
use std::os::unix::fs::PermissionsExt;
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("existing.log");
std::fs::write(&path, b"old log").unwrap();
std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o644)).unwrap();
open_log_file(path.to_str().unwrap()).unwrap();
let mode = std::fs::metadata(&path).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o600, "an existing loose log must be tightened");
}
#[cfg(unix)]
#[test]
fn open_log_file_refuses_a_symlink() {
let dir = tempfile::tempdir().unwrap();
let target = dir.path().join("target.log");
std::fs::write(&target, b"secret").unwrap();
let link = dir.path().join("link.log");
std::os::unix::fs::symlink(&target, &link).unwrap();
assert!(
open_log_file(link.to_str().unwrap()).is_err(),
"a symlink at the log path must be refused"
);
assert_eq!(std::fs::read(&target).unwrap(), b"secret");
}
const CLI_KEY_A: [u8; 32] = [1u8; 32];
const CLI_KEY_B: [u8; 32] = [2u8; 32];
fn cli_b64(key: &[u8; 32]) -> String {
use base64::Engine as _;
base64::engine::general_purpose::STANDARD.encode(key)
}
#[test]
fn acl_add_to_appends_and_is_idempotent() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("authorized_clients.toml");
let count = acl_add_to(&path, &cli_b64(&CLI_KEY_A)).unwrap();
assert_eq!(count, 1);
assert!(
crate::server::acl::Acl::load(&path).contains(&CLI_KEY_A),
"the enrolled key must authorize"
);
assert_eq!(acl_add_to(&path, &cli_b64(&CLI_KEY_A)).unwrap(), 1);
let file = std::fs::read_to_string(&path).unwrap();
assert_eq!(
file.matches("pubkey").count(),
1,
"re-adding must not duplicate the entry"
);
assert_eq!(acl_add_to(&path, &cli_b64(&CLI_KEY_B)).unwrap(), 2);
assert!(crate::server::acl::Acl::load(&path).contains(&CLI_KEY_B));
}
#[test]
fn acl_add_to_rejects_bad_keys() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("authorized_clients.toml");
assert!(acl_add_to(&path, "not-base64!!!").is_err());
use base64::Engine as _;
let short = base64::engine::general_purpose::STANDARD.encode([9u8; 16]);
assert!(acl_add_to(&path, &short).is_err());
assert!(!path.exists(), "a rejected add must not create the file");
}
#[test]
fn parse_max_turns_env_accepts_zero() {
assert_eq!(parse_max_turns_env("0").unwrap(), 0);
}
#[test]
fn parse_max_turns_env_accepts_positive() {
assert_eq!(parse_max_turns_env("42").unwrap(), 42);
}
#[test]
fn parse_max_turns_env_rejects_non_numeric() {
assert!(parse_max_turns_env("abc").is_err());
}
#[test]
fn parse_max_turns_env_rejects_negative() {
assert!(parse_max_turns_env("-5").is_err());
}
#[test]
fn parse_max_turns_env_rejects_empty() {
assert!(parse_max_turns_env("").is_err());
}
#[test]
fn version_flag_displays_package_version() {
let err = match Cli::try_parse_from(["choreographr", "--version"]) {
Err(e) => e,
Ok(_) => panic!("--version should short-circuit before arg validation"),
};
assert_eq!(err.kind(), clap::error::ErrorKind::DisplayVersion);
assert!(err.to_string().contains(env!("CARGO_PKG_VERSION")));
let expected = choreo_proto::release_name::version_string(env!("CARGO_PKG_VERSION"));
assert!(err.to_string().contains(&expected));
}
}