use crate::accounts::{AccountManager, accounts_config_path};
use crate::config::load_daemon_config;
use crate::daemon::DaemonState;
use crate::db::read_all_sessions;
use anyhow::Context;
use choreo_proto::socket_path;
use choreo_transport::key::ensure_transport_keypair;
use clap::Parser;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::mpsc;
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,
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>,
}
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}"))
}
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")),
};
fmt().with_env_filter(env_filter).init();
info!(effective_level = ?log_level.unwrap_or("from RUST_LOG"), "logging initialized");
let max_turns = resolve_max_turns().context("failed to resolve tool-loop iteration limit")?;
info!(max_turns, "tool loop iteration limit");
info!("choreographr starting (locked)");
let db = Arc::new(crate::db::open_db().context("failed to open database")?);
crate::db::run_migrations(&db).context("failed to migrate database")?;
match crate::db::purge_tombstoned_sessions(&db) {
Ok(n) if n > 0 => warn!(
purged = n,
"purged records left behind by interrupted session deletions"
),
Ok(_) => {}
Err(e) => warn!(error = %e, "failed to purge tombstoned sessions; continuing"),
}
let (daemon_tx, _daemon_rx) = mpsc::channel::<crate::daemon::DaemonCommand>();
let mut session_metadata = std::collections::HashMap::new();
match read_all_sessions(&db) {
Ok(sessions) => {
for (id, record) in sessions {
session_metadata.insert(id, record.into());
}
}
Err(e) => {
warn!("failed to load sessions from database: {e}");
}
}
info!(
count = session_metadata.len(),
"loaded sessions from database"
);
let accounts = match accounts_config_path() {
Ok(path) => AccountManager::load(&path).unwrap_or_else(|e| {
warn!("failed to load accounts: {e}");
AccountManager::empty()
}),
Err(e) => {
warn!("no accounts config path: {e}");
AccountManager::empty()
}
};
let mut tool_registry = crate::tools::ToolRegistry::new();
let mcp_manager = crate::mcp::McpManager::from_config(&mut tool_registry);
let tool_registry = tool_registry.build();
let state = DaemonState {
daemon_tx,
next_session_id: session_metadata
.keys()
.max()
.copied()
.map(|m| m + 1)
.unwrap_or(1),
max_turns,
active_sessions: std::collections::HashMap::new(),
session_metadata,
deleted_sessions: std::collections::HashSet::new(),
children: std::collections::HashMap::new(),
accounts,
providers: HashMap::new(),
credentials: std::collections::HashMap::new(),
x_credentials: None,
db,
tool_registry,
client_streams: Vec::new(),
summary_subscribers: std::collections::HashMap::new(),
activity_subscribers: std::collections::HashMap::new(),
client_subscribed_sessions: std::collections::HashMap::new(),
model_cache: HashMap::new(),
mcp_manager,
};
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 = std::sync::Arc::new(crate::server::acl::Acl::load(&acl_path));
let socket_path = socket_path();
crate::run_server(
&socket_path,
state,
cli.metrics_addr,
cli.tcp_addr,
transport_sk,
acl,
)
.context("failed to run server")
}
#[cfg(test)]
mod tests {
use super::*;
#[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")));
}
}