use crate::api::md_api::{
DepthMarketDataField, ForQuoteRspField, MdApi, MdSpiHandler, SpecificInstrumentField,
};
use crate::api::CtpApi;
use crate::error::{CtpError, CtpResult};
use crate::types::{ReqUserLoginField, RspInfoField, RspUserLoginField};
use std::sync::Arc;
use tokio::sync::{mpsc, Mutex, Notify};
use tokio::time::{timeout, Duration};
use tracing::{debug, error, warn};
#[derive(Debug, Clone)]
pub enum AsyncMdEvent {
Connected,
Disconnected(i32),
HeartBeatWarning(i32),
LoginResponse {
user_login: Option<RspUserLoginField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
LogoutResponse {
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
ErrorResponse {
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
SubMarketDataResponse {
specific_instrument: Option<SpecificInstrumentField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
UnsubMarketDataResponse {
specific_instrument: Option<SpecificInstrumentField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
DepthMarketData(DepthMarketDataField),
ForQuoteResponse(ForQuoteRspField),
}
#[derive(Debug, Clone, Default)]
pub struct AsyncMdState {
pub connected: bool,
pub logged_in: bool,
pub login_info: Option<RspUserLoginField>,
}
pub struct AsyncMdApi {
inner: Arc<Mutex<MdApi>>,
event_sender: mpsc::UnboundedSender<AsyncMdEvent>,
event_receiver: Arc<Mutex<mpsc::UnboundedReceiver<AsyncMdEvent>>>,
state: Arc<Mutex<AsyncMdState>>,
connected_notify: Arc<Notify>,
login_notify: Arc<Notify>,
}
impl AsyncMdApi {
pub async fn new(
flow_path: Option<&str>,
is_using_udp: bool,
is_multicast: bool,
is_production_mode: Option<bool>,
) -> CtpResult<Self> {
let md_api = MdApi::new(
flow_path,
is_using_udp,
is_multicast,
is_production_mode.unwrap_or(false),
)?;
let (event_sender, event_receiver) = mpsc::unbounded_channel();
Ok(Self {
inner: Arc::new(Mutex::new(md_api)),
event_sender,
event_receiver: Arc::new(Mutex::new(event_receiver)),
state: Arc::new(Mutex::new(AsyncMdState::default())),
connected_notify: Arc::new(Notify::new()),
login_notify: Arc::new(Notify::new()),
})
}
pub async fn register_front(&self, front_address: &str) -> CtpResult<()> {
let mut api = self.inner.lock().await;
api.register_front(front_address)
}
pub async fn init(&self) -> CtpResult<()> {
let mut api = self.inner.lock().await;
let handler = AsyncMdHandler::new(
self.event_sender.clone(),
self.state.clone(),
self.connected_notify.clone(),
self.login_notify.clone(),
);
api.register_spi(handler)?;
api.init()
}
pub async fn wait_connected(&self, timeout_secs: u64) -> CtpResult<()> {
let state = self.state.lock().await;
if state.connected {
return Ok(());
}
drop(state);
match timeout(
Duration::from_secs(timeout_secs),
self.connected_notify.notified(),
)
.await
{
Ok(_) => {
let state = self.state.lock().await;
if state.connected {
Ok(())
} else {
Err(CtpError::InitializationError("连接失败".to_string()))
}
}
Err(_) => Err(CtpError::InitializationError("连接超时".to_string())),
}
}
pub async fn login(
&self,
req: &ReqUserLoginField,
timeout_secs: u64,
) -> CtpResult<RspUserLoginField> {
let mut api = self.inner.lock().await;
api.req_user_login(req)?;
drop(api);
match timeout(
Duration::from_secs(timeout_secs),
self.login_notify.notified(),
)
.await
{
Ok(_) => {
let state = self.state.lock().await;
if let Some(login_info) = &state.login_info {
Ok(login_info.clone())
} else {
Err(CtpError::InitializationError("登录失败".to_string()))
}
}
Err(_) => Err(CtpError::InitializationError("登录超时".to_string())),
}
}
pub async fn subscribe_market_data(&self, instrument_ids: &[&str]) -> CtpResult<()> {
let mut api = self.inner.lock().await;
api.subscribe_market_data(instrument_ids)
}
pub async fn unsubscribe_market_data(&self, instrument_ids: &[&str]) -> CtpResult<()> {
let mut api = self.inner.lock().await;
api.unsubscribe_market_data(instrument_ids)
}
pub async fn recv_event(&self) -> Option<AsyncMdEvent> {
let mut receiver = self.event_receiver.lock().await;
receiver.recv().await
}
pub async fn try_recv_event(&self) -> Result<AsyncMdEvent, mpsc::error::TryRecvError> {
let mut receiver = self.event_receiver.lock().await;
receiver.try_recv()
}
pub async fn get_state(&self) -> AsyncMdState {
self.state.lock().await.clone()
}
}
#[derive(Clone)]
struct AsyncMdHandler {
event_sender: mpsc::UnboundedSender<AsyncMdEvent>,
state: Arc<Mutex<AsyncMdState>>,
connected_notify: Arc<Notify>,
login_notify: Arc<Notify>,
}
impl AsyncMdHandler {
fn new(
event_sender: mpsc::UnboundedSender<AsyncMdEvent>,
state: Arc<Mutex<AsyncMdState>>,
connected_notify: Arc<Notify>,
login_notify: Arc<Notify>,
) -> Self {
Self {
event_sender,
state,
connected_notify,
login_notify,
}
}
}
impl MdSpiHandler for AsyncMdHandler {
fn on_front_connected(&mut self) {
debug!("异步API: 连接成功");
if let Ok(mut state) = self.state.try_lock() {
state.connected = true;
}
self.connected_notify.notify_waiters();
let _ = self.event_sender.send(AsyncMdEvent::Connected);
}
fn on_front_disconnected(&mut self, reason: i32) {
warn!("异步API: 连接断开, 原因: {}", reason);
if let Ok(mut state) = self.state.try_lock() {
state.connected = false;
state.logged_in = false;
state.login_info = None;
}
let _ = self.event_sender.send(AsyncMdEvent::Disconnected(reason));
}
fn on_heart_beat_warning(&mut self, time_lapse: i32) {
warn!("异步API: 心跳超时警告, 时间间隔: {}秒", time_lapse);
let _ = self
.event_sender
.send(AsyncMdEvent::HeartBeatWarning(time_lapse));
}
fn on_rsp_user_login(
&mut self,
user_login: Option<RspUserLoginField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步API: 收到登录响应");
let success = rsp_info.as_ref().map_or(true, |info| info.is_success());
if success && is_last {
if let Ok(mut state) = self.state.try_lock() {
state.logged_in = true;
state.login_info = user_login.clone();
}
self.login_notify.notify_waiters();
}
let _ = self.event_sender.send(AsyncMdEvent::LoginResponse {
user_login,
rsp_info,
request_id,
is_last,
});
}
fn on_rsp_user_logout(
&mut self,
_user_logout: Option<()>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步API: 收到登出响应");
if is_last {
if let Ok(mut state) = self.state.try_lock() {
state.logged_in = false;
state.login_info = None;
}
}
let _ = self.event_sender.send(AsyncMdEvent::LogoutResponse {
rsp_info,
request_id,
is_last,
});
}
fn on_rsp_error(&mut self, rsp_info: Option<RspInfoField>, request_id: i32, is_last: bool) {
error!("异步API: 收到错误响应");
let _ = self.event_sender.send(AsyncMdEvent::ErrorResponse {
rsp_info,
request_id,
is_last,
});
}
fn on_rsp_sub_market_data(
&mut self,
specific_instrument: Option<SpecificInstrumentField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步API: 收到订阅行情响应");
let _ = self.event_sender.send(AsyncMdEvent::SubMarketDataResponse {
specific_instrument,
rsp_info,
request_id,
is_last,
});
}
fn on_rsp_unsub_market_data(
&mut self,
specific_instrument: Option<SpecificInstrumentField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步API: 收到取消订阅响应");
let _ = self
.event_sender
.send(AsyncMdEvent::UnsubMarketDataResponse {
specific_instrument,
rsp_info,
request_id,
is_last,
});
}
fn on_rtn_depth_market_data(&mut self, market_data: DepthMarketDataField) {
let _ = self
.event_sender
.send(AsyncMdEvent::DepthMarketData(market_data));
}
fn on_rtn_for_quote_rsp(&mut self, for_quote_rsp: ForQuoteRspField) {
debug!("异步API: 收到询价响应");
let _ = self
.event_sender
.send(AsyncMdEvent::ForQuoteResponse(for_quote_rsp));
}
}