//! TCP+TLS connection to the cTrader Open API server. //! //! Provides a framed, heartbeat-aware transport layer built on //! `tokio-native-tls` and `tokio-util::codec`. use std::sync::Arc; use std::time::Duration; use futures::stream::{SplitSink, SplitStream}; use futures::{SinkExt, StreamExt}; use tokio::net::TcpStream; use tokio::sync::Mutex; use tokio_native_tls::TlsStream; use tokio_util::codec::Framed; use tracing::{debug, info, trace}; use crate::codec::CTraderCodec; use crate::config::CTraderConfig; use crate::error::{CTraderError, Result}; use crate::proto::ProtoMessage; type FramedTls = Framed, CTraderCodec>; type Sender = SplitSink; type Receiver = SplitStream; /// A TCP+TLS connection to a cTrader API server. /// /// Splits the framed stream into an `Arc>` (shared between /// heartbeat task and caller) and a `Receiver` (consumed by the read loop). pub struct CTraderConnection { sender: Arc>, receiver: Option, heartbeat_handle: Option>, } impl CTraderConnection { /// Establish a TCP+TLS connection (does NOT start heartbeat). /// /// Call `start_heartbeat()` after authentication completes. pub async fn connect(config: &CTraderConfig) -> Result { let host = config.environment.host(); let port = config.environment.port(); let addr = format!("{host}:{port}"); info!(host, port, "connecting to cTrader API"); // TCP connect let tcp_stream = TcpStream::connect(&addr).await.map_err(|e| { CTraderError::ConnectionFailed(format!("TCP connect to {addr} failed: {e}")) })?; debug!("TCP connected to {addr}"); // TLS handshake let native_connector = native_tls::TlsConnector::new().map_err(|e| { CTraderError::TlsError(format!("failed to create TLS connector: {e}")) })?; let tls_connector = tokio_native_tls::TlsConnector::from(native_connector); let tls_stream = tls_connector.connect(host, tcp_stream).await.map_err(|e| { CTraderError::TlsError(format!("TLS handshake with {host} failed: {e}")) })?; info!("TLS connected to {addr}"); // Wrap in codec framing let framed = Framed::new(tls_stream, CTraderCodec::new()); let (raw_sender, receiver) = framed.split(); let sender = Arc::new(Mutex::new(raw_sender)); Ok(Self { sender, receiver: Some(receiver), heartbeat_handle: None, }) } /// Start the heartbeat keepalive task. /// /// Should be called after authentication succeeds. pub fn start_heartbeat(&mut self, interval: Duration) { let heartbeat_sender = Arc::clone(&self.sender); let handle = tokio::spawn(async move { Self::heartbeat_loop(heartbeat_sender, interval).await; }); self.heartbeat_handle = Some(handle); } /// Send a `ProtoMessage` through the TLS connection. pub async fn send(&self, msg: ProtoMessage) -> Result<()> { let mut sender = self.sender.lock().await; sender.send(msg).await.map_err(|e| { CTraderError::ConnectionFailed(format!("failed to send message: {e}")) })?; Ok(()) } /// Receive the next message from the connection. /// /// This only works before the dispatcher takes over the receiver. pub async fn recv(&mut self) -> Result { let receiver = self .receiver .as_mut() .ok_or(CTraderError::NotConnected)?; match receiver.next().await { Some(Ok(msg)) => Ok(msg), Some(Err(e)) => Err(CTraderError::ConnectionFailed(format!( "receive error: {e}" ))), None => Err(CTraderError::NotConnected), } } /// Take ownership of the receive half (can only be called once). pub fn take_receiver(&mut self) -> Option { self.receiver.take() } /// Internal heartbeat loop — sends `ProtoHeartbeatEvent` on interval. async fn heartbeat_loop(sender: Arc>, interval: Duration) { let mut ticker = tokio::time::interval(interval); // Skip the first immediate tick. ticker.tick().await; loop { ticker.tick().await; let msg = ProtoMessage { payload_type: crate::proto::PT_HEARTBEAT_EVENT, payload: None, client_msg_id: None, }; trace!("sending heartbeat"); let mut s = sender.lock().await; if let Err(e) = s.send(msg).await { debug!("heartbeat send failed: {e}"); break; } } } } impl Drop for CTraderConnection { fn drop(&mut self) { if let Some(handle) = self.heartbeat_handle.take() { handle.abort(); } } } #[cfg(test)] mod tests { use super::*; #[test] fn heartbeat_msg_has_correct_payload_type() { let msg = ProtoMessage { payload_type: crate::proto::PT_HEARTBEAT_EVENT, payload: None, client_msg_id: None, }; assert_eq!(msg.payload_type, 51); } }