1use std::future::Future;
2use std::io::{Error as IoError, ErrorKind};
3use std::time::Duration;
4
5use bytes::Bytes;
6use futures::stream::SplitSink;
7use futures::{SinkExt, StreamExt};
8use tokio::net::TcpStream;
9use tokio_util::codec::Framed;
10
11use futu_codec::FutuCodec;
12use futu_codec::frame::FutuFrame;
13use futu_core::error::FutuError;
14
15const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
21
22pub struct Connection {
27 sink: SplitSink<Framed<TcpStream, FutuCodec>, FutuFrame>,
28 stream: futures::stream::SplitStream<Framed<TcpStream, FutuCodec>>,
29}
30
31impl Connection {
32 pub async fn connect(addr: &str) -> Result<Self, FutuError> {
34 let stream = connect_stream(addr).await?;
35 tracing::info!(addr = addr, "TCP connected");
36
37 let framed = Framed::new(stream, FutuCodec);
38 let (sink, stream) = framed.split();
39
40 Ok(Self { sink, stream })
41 }
42
43 pub async fn send(&mut self, frame: FutuFrame) -> Result<(), FutuError> {
45 self.sink.send(frame).await
46 }
47
48 pub async fn recv(&mut self) -> Result<Option<FutuFrame>, FutuError> {
52 match self.stream.next().await {
53 Some(Ok(frame)) => Ok(Some(frame)),
54 Some(Err(e)) => Err(e),
55 None => Ok(None),
56 }
57 }
58
59 pub fn build_frame(proto_id: u32, serial_no: u32, body: Vec<u8>) -> FutuFrame {
61 FutuFrame::new(proto_id, serial_no, Bytes::from(body))
62 }
63}
64
65async fn connect_stream(addr: &str) -> Result<TcpStream, FutuError> {
66 let stream = connect_with_timeout(addr, CONNECT_TIMEOUT, TcpStream::connect(addr)).await?;
67 configure_connected_stream(&stream)?;
68 Ok(stream)
69}
70
71fn configure_connected_stream(stream: &TcpStream) -> Result<(), FutuError> {
72 stream.set_nodelay(true)?;
73 socket2::SockRef::from(stream).set_keepalive(true)?;
74 Ok(())
75}
76
77async fn connect_with_timeout<T, F>(
78 addr: &str,
79 timeout: Duration,
80 connect: F,
81) -> Result<T, FutuError>
82where
83 F: Future<Output = std::io::Result<T>>,
84{
85 match tokio::time::timeout(timeout, connect).await {
86 Ok(Ok(stream)) => Ok(stream),
87 Ok(Err(err)) => Err(FutuError::Network(err)),
88 Err(_) => Err(FutuError::Network(IoError::new(
89 ErrorKind::TimedOut,
90 format!("connect to {addr} timed out after {}s", timeout.as_secs()),
91 ))),
92 }
93}
94
95#[cfg(test)]
96mod tests;