use crate::api::trader_api::{
InputOrderField, InvestorPositionField, OrderField, ReqAuthenticateField, RspAuthenticateField,
TradeField, TraderApi, TraderSpiHandler, TradingAccountField,
};
use crate::api::CtpApi;
use crate::error::{CtpError, CtpResult};
use crate::types::{
InputOrderActionField, QryInvestorPositionField, QryTradingAccountField, ReqUserLoginField,
RspInfoField, RspUserLoginField,
};
use std::collections::HashMap;
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 AsyncTraderEvent {
Connected,
Disconnected(i32),
HeartBeatWarning(i32),
AuthenticateResponse {
rsp_authenticate: Option<RspAuthenticateField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
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,
},
OrderInsertResponse {
input_order: Option<InputOrderField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
OrderActionResponse {
input_order_action: Option<InputOrderActionField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
QryTradingAccountResponse {
trading_account: Option<TradingAccountField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
QryInvestorPositionResponse {
investor_position: Option<InvestorPositionField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
QryOrderResponse {
order: Option<OrderField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
QryTradeResponse {
trade: Option<TradeField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
OrderReturn(OrderField),
TradeReturn(TradeField),
ErrorResponse {
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
},
}
#[derive(Debug, Clone, Default)]
pub struct AsyncTraderState {
pub connected: bool,
pub authenticated: bool,
pub logged_in: bool,
pub auth_info: Option<RspAuthenticateField>,
pub login_info: Option<RspUserLoginField>,
}
#[derive(Debug, Clone)]
struct PendingRequest {
notify: Arc<Notify>,
response_data: Arc<Mutex<Option<AsyncTraderEvent>>>,
}
pub struct AsyncTraderApi {
inner: Arc<Mutex<TraderApi>>,
event_sender: mpsc::UnboundedSender<AsyncTraderEvent>,
event_receiver: Arc<Mutex<mpsc::UnboundedReceiver<AsyncTraderEvent>>>,
state: Arc<Mutex<AsyncTraderState>>,
connected_notify: Arc<Notify>,
auth_notify: Arc<Notify>,
login_notify: Arc<Notify>,
pending_requests: Arc<Mutex<HashMap<i32, PendingRequest>>>,
}
impl AsyncTraderApi {
pub async fn new(flow_path: Option<&str>, is_production_mode: Option<bool>) -> CtpResult<Self> {
let trader_api = TraderApi::new(flow_path, is_production_mode)?;
let (event_sender, event_receiver) = mpsc::unbounded_channel();
Ok(Self {
inner: Arc::new(Mutex::new(trader_api)),
event_sender,
event_receiver: Arc::new(Mutex::new(event_receiver)),
state: Arc::new(Mutex::new(AsyncTraderState::default())),
connected_notify: Arc::new(Notify::new()),
auth_notify: Arc::new(Notify::new()),
login_notify: Arc::new(Notify::new()),
pending_requests: Arc::new(Mutex::new(HashMap::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 = AsyncTraderHandler::new(
self.event_sender.clone(),
self.state.clone(),
self.connected_notify.clone(),
self.auth_notify.clone(),
self.login_notify.clone(),
self.pending_requests.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 authenticate(
&self,
req: &ReqAuthenticateField,
timeout_secs: u64,
) -> CtpResult<RspAuthenticateField> {
let mut api = self.inner.lock().await;
let request_id = api.req_authenticate(req)?;
drop(api);
let pending_request = PendingRequest {
notify: Arc::new(Notify::new()),
response_data: Arc::new(Mutex::new(None)),
};
{
let mut pending = self.pending_requests.lock().await;
pending.insert(request_id, pending_request.clone());
}
match timeout(
Duration::from_secs(timeout_secs),
pending_request.notify.notified(),
)
.await
{
Ok(_) => {
let state = self.state.lock().await;
if let Some(auth_info) = &state.auth_info {
Ok(auth_info.clone())
} else {
Err(CtpError::InitializationError("认证失败".to_string()))
}
}
Err(_) => {
let mut pending = self.pending_requests.lock().await;
pending.remove(&request_id);
Err(CtpError::InitializationError("认证超时".to_string()))
}
}
}
pub async fn login(
&self,
req: &ReqUserLoginField,
timeout_secs: u64,
) -> CtpResult<RspUserLoginField> {
let mut api = self.inner.lock().await;
let request_id = api.req_user_login(req)?;
drop(api);
let pending_request = PendingRequest {
notify: Arc::new(Notify::new()),
response_data: Arc::new(Mutex::new(None)),
};
{
let mut pending = self.pending_requests.lock().await;
pending.insert(request_id, pending_request.clone());
}
match timeout(
Duration::from_secs(timeout_secs),
pending_request.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(_) => {
let mut pending = self.pending_requests.lock().await;
pending.remove(&request_id);
Err(CtpError::InitializationError("登录超时".to_string()))
}
}
}
pub async fn order_insert(
&self,
req: &InputOrderField,
timeout_secs: u64,
) -> CtpResult<AsyncTraderEvent> {
let mut api = self.inner.lock().await;
let request_id = api.req_order_insert(req)?;
drop(api);
self.wait_for_response(request_id, timeout_secs).await
}
pub async fn order_action(
&self,
req: &InputOrderActionField,
timeout_secs: u64,
) -> CtpResult<AsyncTraderEvent> {
let mut api = self.inner.lock().await;
let request_id = api.req_order_action(req)?;
drop(api);
self.wait_for_response(request_id, timeout_secs).await
}
pub async fn qry_trading_account(
&self,
req: &QryTradingAccountField,
timeout_secs: u64,
) -> CtpResult<Vec<TradingAccountField>> {
let mut api = self.inner.lock().await;
let request_id = api.req_qry_trading_account(req)?;
drop(api);
let mut results = Vec::new();
let mut is_finished = false;
let start_time = std::time::Instant::now();
let timeout_duration = Duration::from_secs(timeout_secs);
while !is_finished && start_time.elapsed() < timeout_duration {
if let Some(event) = self.recv_event().await {
match event {
AsyncTraderEvent::QryTradingAccountResponse {
trading_account,
rsp_info,
request_id: resp_id,
is_last,
} if resp_id == request_id => {
if let Some(rsp) = rsp_info {
if !rsp.is_success() {
return Err(CtpError::BusinessError(
rsp.error_id,
rsp.get_error_msg().unwrap_or_default(),
));
}
}
if let Some(account) = trading_account {
results.push(account);
}
is_finished = is_last;
}
_ => continue,
}
}
}
if is_finished {
Ok(results)
} else {
Err(CtpError::InitializationError("查询超时".to_string()))
}
}
pub async fn qry_investor_position(
&self,
req: &QryInvestorPositionField,
timeout_secs: u64,
) -> CtpResult<Vec<InvestorPositionField>> {
let mut api = self.inner.lock().await;
let request_id = api.req_qry_investor_position(req)?;
drop(api);
let mut results = Vec::new();
let mut is_finished = false;
let start_time = std::time::Instant::now();
let timeout_duration = Duration::from_secs(timeout_secs);
while !is_finished && start_time.elapsed() < timeout_duration {
if let Some(event) = self.recv_event().await {
match event {
AsyncTraderEvent::QryInvestorPositionResponse {
investor_position,
rsp_info,
request_id: resp_id,
is_last,
} if resp_id == request_id => {
if let Some(rsp) = rsp_info {
if !rsp.is_success() {
return Err(CtpError::BusinessError(
rsp.error_id,
rsp.get_error_msg().unwrap_or_default(),
));
}
}
if let Some(position) = investor_position {
results.push(position);
}
is_finished = is_last;
}
_ => continue,
}
}
}
if is_finished {
Ok(results)
} else {
Err(CtpError::InitializationError("查询超时".to_string()))
}
}
async fn wait_for_response(
&self,
request_id: i32,
timeout_secs: u64,
) -> CtpResult<AsyncTraderEvent> {
let pending_request = PendingRequest {
notify: Arc::new(Notify::new()),
response_data: Arc::new(Mutex::new(None)),
};
{
let mut pending = self.pending_requests.lock().await;
pending.insert(request_id, pending_request.clone());
}
match timeout(
Duration::from_secs(timeout_secs),
pending_request.notify.notified(),
)
.await
{
Ok(_) => {
let response_data = pending_request.response_data.lock().await;
if let Some(event) = response_data.as_ref() {
Ok(event.clone())
} else {
Err(CtpError::InitializationError("响应数据为空".to_string()))
}
}
Err(_) => {
let mut pending = self.pending_requests.lock().await;
pending.remove(&request_id);
Err(CtpError::InitializationError("请求超时".to_string()))
}
}
}
pub async fn recv_event(&self) -> Option<AsyncTraderEvent> {
let mut receiver = self.event_receiver.lock().await;
receiver.recv().await
}
pub async fn try_recv_event(&self) -> Result<AsyncTraderEvent, mpsc::error::TryRecvError> {
let mut receiver = self.event_receiver.lock().await;
receiver.try_recv()
}
pub async fn get_state(&self) -> AsyncTraderState {
self.state.lock().await.clone()
}
pub async fn release(&self) -> CtpResult<()> {
let mut api = self.inner.lock().await;
api.release();
Ok(())
}
}
#[derive(Clone)]
struct AsyncTraderHandler {
event_sender: mpsc::UnboundedSender<AsyncTraderEvent>,
state: Arc<Mutex<AsyncTraderState>>,
connected_notify: Arc<Notify>,
auth_notify: Arc<Notify>,
login_notify: Arc<Notify>,
pending_requests: Arc<Mutex<HashMap<i32, PendingRequest>>>,
}
impl AsyncTraderHandler {
fn new(
event_sender: mpsc::UnboundedSender<AsyncTraderEvent>,
state: Arc<Mutex<AsyncTraderState>>,
connected_notify: Arc<Notify>,
auth_notify: Arc<Notify>,
login_notify: Arc<Notify>,
pending_requests: Arc<Mutex<HashMap<i32, PendingRequest>>>,
) -> Self {
Self {
event_sender,
state,
connected_notify,
auth_notify,
login_notify,
pending_requests,
}
}
fn notify_pending_request(&self, request_id: i32, event: AsyncTraderEvent) {
if let Ok(mut pending) = self.pending_requests.try_lock() {
if let Some(req) = pending.remove(&request_id) {
if let Ok(mut data) = req.response_data.try_lock() {
*data = Some(event);
}
req.notify.notify_waiters();
}
}
}
}
impl TraderSpiHandler for AsyncTraderHandler {
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(AsyncTraderEvent::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.authenticated = false;
state.logged_in = false;
state.auth_info = None;
state.login_info = None;
}
let _ = self
.event_sender
.send(AsyncTraderEvent::Disconnected(reason));
}
fn on_heart_beat_warning(&mut self, time_lapse: i32) {
warn!("异步交易API: 心跳超时警告, 时间间隔: {}秒", time_lapse);
let _ = self
.event_sender
.send(AsyncTraderEvent::HeartBeatWarning(time_lapse));
}
fn on_rsp_authenticate(
&mut self,
rsp_authenticate: Option<RspAuthenticateField>,
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.authenticated = true;
state.auth_info = rsp_authenticate.clone();
}
self.auth_notify.notify_waiters();
}
let event = AsyncTraderEvent::AuthenticateResponse {
rsp_authenticate,
rsp_info,
request_id,
is_last,
};
let _ = self.event_sender.send(event.clone());
self.notify_pending_request(request_id, event);
}
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 event = AsyncTraderEvent::LoginResponse {
user_login,
rsp_info,
request_id,
is_last,
};
let _ = self.event_sender.send(event.clone());
self.notify_pending_request(request_id, event);
}
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 event = AsyncTraderEvent::LogoutResponse {
rsp_info,
request_id,
is_last,
};
let _ = self.event_sender.send(event.clone());
self.notify_pending_request(request_id, event);
}
fn on_rsp_error(&mut self, rsp_info: Option<RspInfoField>, request_id: i32, is_last: bool) {
error!("异步交易API: 收到错误响应");
let event = AsyncTraderEvent::ErrorResponse {
rsp_info,
request_id,
is_last,
};
let _ = self.event_sender.send(event.clone());
self.notify_pending_request(request_id, event);
}
fn on_rsp_order_insert(
&mut self,
input_order: Option<InputOrderField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步交易API: 收到报单录入响应");
let event = AsyncTraderEvent::OrderInsertResponse {
input_order,
rsp_info,
request_id,
is_last,
};
let _ = self.event_sender.send(event.clone());
self.notify_pending_request(request_id, event);
}
fn on_rsp_order_action(
&mut self,
input_order_action: Option<InputOrderActionField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步交易API: 收到报单操作响应");
let event = AsyncTraderEvent::OrderActionResponse {
input_order_action,
rsp_info,
request_id,
is_last,
};
let _ = self.event_sender.send(event.clone());
self.notify_pending_request(request_id, event);
}
fn on_rsp_qry_trading_account(
&mut self,
trading_account: Option<TradingAccountField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步交易API: 收到查询交易账户响应");
let _ = self
.event_sender
.send(AsyncTraderEvent::QryTradingAccountResponse {
trading_account,
rsp_info,
request_id,
is_last,
});
}
fn on_rsp_qry_investor_position(
&mut self,
investor_position: Option<InvestorPositionField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步交易API: 收到查询投资者持仓响应");
let _ = self
.event_sender
.send(AsyncTraderEvent::QryInvestorPositionResponse {
investor_position,
rsp_info,
request_id,
is_last,
});
}
fn on_rsp_qry_order(
&mut self,
order: Option<OrderField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步交易API: 收到查询报单响应");
let _ = self.event_sender.send(AsyncTraderEvent::QryOrderResponse {
order,
rsp_info,
request_id,
is_last,
});
}
fn on_rsp_qry_trade(
&mut self,
trade: Option<TradeField>,
rsp_info: Option<RspInfoField>,
request_id: i32,
is_last: bool,
) {
debug!("异步交易API: 收到查询成交响应");
let _ = self.event_sender.send(AsyncTraderEvent::QryTradeResponse {
trade,
rsp_info,
request_id,
is_last,
});
}
fn on_rtn_order(&mut self, order: OrderField) {
debug!("异步交易API: 收到报单回报");
let _ = self.event_sender.send(AsyncTraderEvent::OrderReturn(order));
}
fn on_rtn_trade(&mut self, trade: TradeField) {
debug!("异步交易API: 收到成交回报");
let _ = self.event_sender.send(AsyncTraderEvent::TradeReturn(trade));
}
}