use crate::action::{AmiAction, ChallengeAction, ChallengeLoginAction, LoginAction, PingAction};
use crate::codec::{AmiCodec, RawAmiMessage};
use crate::error::{AmiError, Result};
use crate::event::AmiEvent;
use crate::response::{AmiResponse, PendingActions};
use asterisk_rs_core::auth::Credentials;
use asterisk_rs_core::config::{ConnectionState, ReconnectPolicy};
use asterisk_rs_core::event::EventBus;
use futures_util::{SinkExt, StreamExt};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::net::TcpStream;
use tokio::sync::{mpsc, watch, Mutex};
use tokio_util::codec::{FramedRead, FramedWrite};
use zeroize::Zeroizing;
pub(crate) enum ConnectionCommand {
SendAction {
message: RawAmiMessage,
action_id: String,
response_tx: tokio::sync::oneshot::Sender<AmiResponse>,
},
Shutdown,
SendEventGeneratingAction {
message: RawAmiMessage,
action_id: String,
response_tx: tokio::sync::oneshot::Sender<crate::response::EventListResponse>,
},
}
pub(crate) struct ConnectionManager {
command_tx: mpsc::Sender<ConnectionCommand>,
state_rx: watch::Receiver<ConnectionState>,
}
impl ConnectionManager {
pub fn spawn(
address: String,
credentials: Credentials,
event_bus: EventBus<AmiEvent>,
reconnect_policy: ReconnectPolicy,
ping_interval: Option<Duration>,
require_challenge: bool,
) -> Self {
let (command_tx, command_rx) = mpsc::channel(256);
let (state_tx, state_rx) = watch::channel(ConnectionState::Disconnected);
tokio::spawn(connection_task(
address,
credentials,
command_rx,
event_bus,
state_tx,
reconnect_policy,
ping_interval,
require_challenge,
));
Self {
command_tx,
state_rx,
}
}
pub async fn send(&self, cmd: ConnectionCommand) -> Result<()> {
self.command_tx
.send(cmd)
.await
.map_err(|_| AmiError::Disconnected)
}
pub fn state(&self) -> ConnectionState {
*self.state_rx.borrow()
}
pub async fn wait_for_state(&self, target: ConnectionState) -> Result<()> {
let mut rx = self.state_rx.clone();
while *rx.borrow_and_update() != target {
rx.changed().await.map_err(|_| AmiError::Disconnected)?;
}
Ok(())
}
pub async fn shutdown(&self) {
let _ = self.command_tx.send(ConnectionCommand::Shutdown).await;
}
}
#[allow(clippy::too_many_arguments)]
async fn connection_task(
address: String,
credentials: Credentials,
mut command_rx: mpsc::Receiver<ConnectionCommand>,
event_bus: EventBus<AmiEvent>,
state_tx: watch::Sender<ConnectionState>,
reconnect_policy: ReconnectPolicy,
ping_interval: Option<Duration>,
require_challenge: bool,
) {
let pending = Arc::new(Mutex::new(PendingActions::new()));
let mut attempt: u32 = 0;
loop {
let _ = state_tx.send(ConnectionState::Connecting);
tracing::info!(address = %address, attempt, "connecting to AMI");
match tokio::time::timeout(Duration::from_secs(10), TcpStream::connect(&address)).await {
Ok(Ok(stream)) => {
tracing::info!(address = %address, "TCP connected to AMI");
let (read_half, write_half) = stream.into_split();
let mut reader = FramedRead::new(read_half, AmiCodec::new());
let mut writer = FramedWrite::new(write_half, AmiCodec::new());
let login_result = tokio::time::timeout(
Duration::from_secs(30),
perform_login(&credentials, &mut reader, &mut writer, require_challenge),
)
.await;
match login_result {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::error!(error = %e, "AMI login failed after connect");
continue;
}
Err(_) => {
tracing::error!("AMI login timed out after 30s");
continue;
}
}
tracing::info!("AMI login successful");
attempt = 0; let _ = state_tx.send(ConnectionState::Connected);
let mut ping_timer = ping_interval.map(tokio::time::interval);
if let Some(ref mut timer) = ping_timer {
timer.tick().await; }
let pong_received = Arc::new(AtomicBool::new(false));
let mut awaiting_pong = false;
let mut _pong_rx: Option<tokio::sync::oneshot::Receiver<AmiResponse>> = None;
loop {
tokio::select! {
biased;
frame = reader.next() => {
match frame {
Some(Ok(raw)) => {
dispatch_message(raw, &pending, &event_bus, &pong_received).await;
}
Some(Err(e)) => {
tracing::error!(error = %e, "AMI codec error");
break;
}
None => {
tracing::warn!("AMI connection closed");
break;
}
}
}
cmd = command_rx.recv() => {
match cmd {
Some(ConnectionCommand::SendAction { message, action_id, response_tx }) => {
pending.lock().await.register_with_sender(action_id, response_tx);
if let Err(e) = writer.send(message).await {
tracing::error!(error = %e, "failed to send AMI action");
break;
}
}
Some(ConnectionCommand::SendEventGeneratingAction { message, action_id, response_tx }) => {
pending.lock().await.register_event_list(action_id, response_tx);
if let Err(e) = writer.send(message).await {
tracing::error!(error = %e, "failed to send AMI action");
break;
}
}
Some(ConnectionCommand::Shutdown) => {
tracing::info!("AMI connection shutdown requested");
let _ = state_tx.send(ConnectionState::Disconnected);
return;
}
None => {
let _ = state_tx.send(ConnectionState::Disconnected);
return;
}
}
}
_ = async {
match ping_timer.as_mut() {
Some(timer) => timer.tick().await,
None => std::future::pending().await,
}
} => {
if awaiting_pong && !pong_received.load(Ordering::Acquire) {
tracing::warn!("keep-alive pong not received, treating connection as dead");
break;
}
pong_received.store(false, Ordering::Release);
let (action_id, ping_msg) = PingAction.to_message();
_pong_rx = Some(pending.lock().await.register(action_id));
awaiting_pong = true;
if let Err(e) = writer.send(ping_msg).await {
tracing::warn!(error = %e, "keep-alive ping failed, reconnecting");
break;
}
tracing::trace!("keep-alive ping sent");
}
}
}
pending.lock().await.cancel_all();
}
Ok(Err(e)) => {
tracing::error!(address = %address, error = %e, "failed to connect to AMI");
}
Err(_) => {
tracing::error!(address = %address, "AMI connection timed out");
}
}
if reconnect_policy
.max_retries
.is_some_and(|max| attempt >= max)
{
tracing::error!("max reconnection attempts reached, giving up");
let _ = state_tx.send(ConnectionState::Disconnected);
return;
}
let _ = state_tx.send(ConnectionState::Reconnecting);
let delay = reconnect_policy.delay_for_attempt(attempt);
tracing::info!(?delay, attempt, "reconnecting to AMI");
tokio::select! {
() = tokio::time::sleep(delay) => {
drain_backoff_commands(&mut command_rx, &state_tx);
}
cmd = command_rx.recv() => {
if reject_backoff_command(cmd, &state_tx) {
return; }
drain_backoff_commands(&mut command_rx, &state_tx);
}
}
attempt += 1;
}
}
async fn perform_login(
credentials: &Credentials,
reader: &mut FramedRead<tokio::net::tcp::OwnedReadHalf, AmiCodec>,
writer: &mut FramedWrite<tokio::net::tcp::OwnedWriteHalf, AmiCodec>,
require_challenge: bool,
) -> Result<()> {
let (_, challenge_msg) = ChallengeAction.to_message();
writer.send(challenge_msg).await?;
let challenge_resp = read_next_response(reader).await?;
if challenge_resp.success {
if let Some(challenge) = challenge_resp.get("Challenge") {
let key = Zeroizing::new(compute_md5_key(challenge, credentials.secret()));
let login = ChallengeLoginAction {
username: credentials.username().to_string(),
key,
};
let (_, login_msg) = login.to_message();
writer.send(login_msg).await?;
let login_resp = read_next_response(reader).await?;
if !login_resp.success {
return Err(AmiError::Auth(
asterisk_rs_core::error::AuthError::Rejected {
reason: login_resp.message.unwrap_or_default(),
},
));
}
return Ok(());
}
}
if require_challenge {
return Err(AmiError::Auth(
asterisk_rs_core::error::AuthError::Rejected {
reason: "server did not provide MD5 challenge; plaintext fallback is disabled \
(set require_challenge(false) for trusted loopback connections)"
.to_owned(),
},
));
}
tracing::warn!("MD5 challenge auth unavailable, falling back to plaintext login");
let login = LoginAction::new(credentials.username(), credentials.secret());
let (_, login_msg) = login.to_message();
writer.send(login_msg).await?;
let login_resp = read_next_response(reader).await?;
if !login_resp.success {
return Err(AmiError::Auth(
asterisk_rs_core::error::AuthError::Rejected {
reason: login_resp.message.unwrap_or_default(),
},
));
}
Ok(())
}
async fn read_next_response(
reader: &mut FramedRead<tokio::net::tcp::OwnedReadHalf, AmiCodec>,
) -> Result<AmiResponse> {
loop {
match reader.next().await {
Some(Ok(raw)) => {
if let Some(resp) = AmiResponse::from_raw(&raw) {
return Ok(resp);
}
}
Some(Err(e)) => return Err(e),
None => return Err(AmiError::Disconnected),
}
}
}
fn compute_md5_key(challenge: &str, secret: &str) -> String {
use md5::{Digest, Md5};
let mut hasher = Md5::new();
hasher.update(challenge.as_bytes());
hasher.update(secret.as_bytes());
format!("{:x}", hasher.finalize())
}
async fn dispatch_message(
raw: RawAmiMessage,
pending: &Arc<Mutex<PendingActions>>,
event_bus: &EventBus<AmiEvent>,
pong_received: &AtomicBool,
) {
if let Some(response) = AmiResponse::from_raw(&raw) {
let mut guard = pending.lock().await;
if guard.contains_event_list(&response.action_id) {
pong_received.store(true, Ordering::Release);
guard.deliver_event_list_response(response);
return;
}
let action_id = response.action_id.clone();
if guard.deliver(response) {
pong_received.store(true, Ordering::Release);
} else {
tracing::debug!(action_id, "received response for unknown action");
}
return;
}
if let Some(event) = AmiEvent::from_raw(&raw) {
if let Some(aid) = raw.get("ActionID") {
let mut guard = pending.lock().await;
if guard.deliver_event_list_event(aid, event.clone()) {
event_bus.publish(event);
return;
}
}
tracing::trace!(event = event.event_name(), "AMI event received");
event_bus.publish(event);
return;
}
tracing::debug!("received unclassifiable AMI message");
}
fn reject_backoff_command(
cmd: Option<ConnectionCommand>,
state_tx: &watch::Sender<ConnectionState>,
) -> bool {
match cmd {
None | Some(ConnectionCommand::Shutdown) => {
tracing::info!("shutdown received during reconnect backoff");
let _ = state_tx.send(ConnectionState::Disconnected);
true
}
Some(ConnectionCommand::SendAction { response_tx, .. }) => {
tracing::debug!("rejecting action received during reconnect backoff");
drop(response_tx);
false
}
Some(ConnectionCommand::SendEventGeneratingAction { response_tx, .. }) => {
tracing::debug!("rejecting event-list action received during reconnect backoff");
drop(response_tx);
false
}
}
}
fn drain_backoff_commands(
command_rx: &mut mpsc::Receiver<ConnectionCommand>,
state_tx: &watch::Sender<ConnectionState>,
) {
while let Ok(cmd) = command_rx.try_recv() {
if reject_backoff_command(Some(cmd), state_tx) {
return;
}
}
}