use crate::client::config::ClientConfig;
use crate::client::connection::ConnectionStateManager;
use crate::client::heartbeat::HeartbeatManager;
use crate::client::router::MessageRouter;
use crate::common::error::{FlareError, Result};
use crate::common::protocol::Frame;
use crate::common::protocol::flare::core::commands::command::Type;
use crate::common::protocol::flare::core::commands::notification_command::Type as NotificationCommandType;
use crate::common::protocol::flare::core::commands::payload_command::Type as PayloadCommandType;
use crate::common::protocol::flare::core::commands::system_command::Type as SystemCommandType;
use crate::common::{HeartbeatAppState, HeartbeatConfig, MessageParser};
use crate::transport::connection::Connection;
use crate::transport::events::{ArcObserver, ConnectionEvent};
use std::collections::HashMap;
use std::sync::{
Arc, Mutex as StdMutex, RwLock as StdRwLock,
atomic::{AtomicBool, Ordering},
};
use tokio::sync::{Mutex, Notify, oneshot};
#[path = "client_core_connect.rs"]
mod client_core_connect;
const NEGOTIATION_TIMEOUT_HINT: &str =
"Ensure `flare_chat_server` is running, not `simple_server`.";
#[cfg(target_arch = "wasm32")]
const MAX_WASM_INBOUND_QUEUE: usize = 512;
fn negotiation_timeout_error(timeout: std::time::Duration) -> FlareError {
FlareError::connection_timeout(format!(
"Negotiation timeout after {:?} (CONNECT_ACK not received). {}",
timeout, NEGOTIATION_TIMEOUT_HINT
))
}
#[cfg(not(target_arch = "wasm32"))]
async fn wait_for_negotiation_notify(
flag: Arc<AtomicBool>,
failure_reason: Arc<StdMutex<Option<String>>>,
notify: Arc<Notify>,
timeout: std::time::Duration,
) -> Result<()> {
let wait = async move {
loop {
if flag.load(Ordering::SeqCst) {
return Ok(());
}
if let Ok(reason) = failure_reason.lock()
&& let Some(msg) = reason.as_ref()
{
return Err(FlareError::protocol_error(msg.clone()));
}
notify.notified().await;
}
};
match crate::common::platform::timeout(timeout, wait).await {
Ok(result) => result,
Err(_) => Err(negotiation_timeout_error(timeout)),
}
}
pub struct ClientCore {
pub state_manager: Arc<ConnectionStateManager>,
pub parser: Arc<tokio::sync::Mutex<MessageParser>>,
heartbeat_manager: Arc<StdMutex<Option<Arc<tokio::sync::Mutex<HeartbeatManager>>>>>,
heartbeat_config: Arc<StdRwLock<HeartbeatConfig>>,
message_router: Option<MessageRouter>,
pub observers: Arc<StdMutex<Vec<ArcObserver>>>,
pub config: ClientConfig,
event_handler: Option<Arc<dyn crate::client::events::handler::ClientEventHandler>>,
#[allow(clippy::type_complexity)]
client_connection: Arc<std::sync::Mutex<Option<Arc<Mutex<Box<dyn Connection>>>>>>,
pub(crate) pending_map: Arc<tokio::sync::Mutex<HashMap<String, oneshot::Sender<Frame>>>>,
pub(crate) negotiation_completed: Arc<AtomicBool>,
pub(crate) negotiation_notify: Arc<Notify>,
negotiation_failure_reason: Arc<StdMutex<Option<String>>>,
disconnect_requested: Arc<AtomicBool>,
#[cfg(target_arch = "wasm32")]
wasm_inbound: Arc<StdMutex<Vec<Vec<u8>>>>,
}
impl ClientCore {
pub fn new(config: &ClientConfig) -> Self {
let (format, compression) = Self::determine_initial_format(config);
let parser = MessageParser::new(
format,
compression,
crate::common::encryption::EncryptionAlgorithm::None,
);
let message_router = config.enable_router.then(MessageRouter::new);
Self {
state_manager: Arc::new(ConnectionStateManager::new()),
parser: Arc::new(tokio::sync::Mutex::new(parser)),
heartbeat_manager: Arc::new(StdMutex::new(None)),
heartbeat_config: Arc::new(StdRwLock::new(config.heartbeat.clone())),
message_router,
observers: Arc::new(StdMutex::new(Vec::new())),
config: config.clone(),
event_handler: None,
client_connection: Arc::new(std::sync::Mutex::new(None)),
pending_map: Arc::new(tokio::sync::Mutex::new(HashMap::new())),
negotiation_completed: Arc::new(AtomicBool::new(false)),
negotiation_notify: Arc::new(Notify::new()),
negotiation_failure_reason: Arc::new(StdMutex::new(None)),
disconnect_requested: Arc::new(AtomicBool::new(false)),
#[cfg(target_arch = "wasm32")]
wasm_inbound: Arc::new(StdMutex::new(Vec::new())),
}
}
#[cfg(target_arch = "wasm32")]
pub fn push_wasm_inbound(&self, data: Vec<u8>) {
if let Ok(mut queue) = self.wasm_inbound.lock() {
if queue.len() >= MAX_WASM_INBOUND_QUEUE {
queue.remove(0);
tracing::warn!(
"[ClientCore] wasm inbound queue full (max {}), dropping oldest frame",
MAX_WASM_INBOUND_QUEUE
);
}
queue.push(data);
}
self.negotiation_notify.notify_waiters();
}
#[cfg(target_arch = "wasm32")]
pub async fn drain_wasm_inbound(&self) {
let batch: Vec<Vec<u8>> = match self.wasm_inbound.lock() {
Ok(mut queue) if !queue.is_empty() => queue.drain(..).collect(),
_ => return,
};
for data in batch {
self.handle_message(data).await;
}
}
pub fn set_disconnect_requested(&self, value: bool) {
self.disconnect_requested.store(value, Ordering::SeqCst);
}
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "websocket", feature = "quic", feature = "tcp")
))]
pub(crate) fn share_race_state_from(&mut self, shared: &ClientCore) {
self.observers = Arc::clone(&shared.observers);
self.pending_map = Arc::clone(&shared.pending_map);
self.disconnect_requested = Arc::clone(&shared.disconnect_requested);
self.negotiation_completed = Arc::clone(&shared.negotiation_completed);
self.negotiation_notify = Arc::clone(&shared.negotiation_notify);
self.negotiation_failure_reason = Arc::clone(&shared.negotiation_failure_reason);
self.heartbeat_config = Arc::clone(&shared.heartbeat_config);
}
fn determine_initial_format(
config: &ClientConfig,
) -> (
crate::common::protocol::SerializationFormat,
crate::common::compression::CompressionAlgorithm,
) {
if config.is_force_format() {
(config.get_serialization_format(), config.get_compression())
} else {
(
crate::common::protocol::SerializationFormat::Json,
crate::common::compression::CompressionAlgorithm::None,
)
}
}
pub async fn update_parser(
&self,
format: crate::common::protocol::SerializationFormat,
compression: crate::common::compression::CompressionAlgorithm,
encryption: crate::common::encryption::EncryptionAlgorithm,
) {
let compression_clone = compression.clone();
let encryption_clone = encryption.clone();
let mut parser = self.parser.lock().await;
*parser = MessageParser::new(format, compression, encryption);
self.negotiation_completed.store(true, Ordering::SeqCst);
self.negotiation_notify.notify_waiters();
tracing::info!(
"[ClientCore] ✅ 协商完成,解析器已更新: 最终序列化方式={:?}, 最终压缩方式={:?}, 最终加密方式={:?}, negotiation_completed={}",
format,
compression_clone,
encryption_clone,
self.negotiation_completed.load(Ordering::SeqCst)
);
self.try_start_heartbeat().await;
}
pub fn set_client_connection(&mut self, connection: Arc<Mutex<Box<dyn Connection>>>) {
if let Ok(mut conn) = self.client_connection.lock() {
*conn = Some(connection);
}
}
pub fn clear_client_connection(&self) {
if let Ok(mut conn) = self.client_connection.lock() {
*conn = None;
}
}
pub fn take_client_connection(&self) -> Option<Arc<Mutex<Box<dyn Connection>>>> {
self.client_connection
.lock()
.ok()
.and_then(|mut conn| conn.take())
}
pub fn set_event_handler(
&mut self,
handler: Option<Arc<dyn crate::client::events::handler::ClientEventHandler>>,
) {
self.event_handler = handler;
}
pub async fn start_heartbeat(&self, connection: Arc<Mutex<Box<dyn Connection>>>) {
if let Ok(mut conn) = self.client_connection.lock() {
*conn = Some(Arc::clone(&connection));
}
self.try_start_heartbeat().await;
}
async fn try_start_heartbeat(&self) {
if !self.current_heartbeat_config().enabled {
return;
}
if !self.negotiation_completed.load(Ordering::SeqCst) {
return;
}
let Ok(mut slot) = self.heartbeat_manager.lock() else {
return;
};
if slot.is_some() {
return;
}
let Some(connection) = self
.client_connection
.lock()
.ok()
.and_then(|guard| guard.clone())
else {
tracing::debug!("[ClientCore] heartbeat deferred: no active connection");
return;
};
let mut heartbeat =
HeartbeatManager::with_shared_config(Arc::clone(&self.heartbeat_config));
let parser_ref = Arc::clone(&self.parser);
heartbeat.start(connection, parser_ref);
*slot = Some(Arc::new(tokio::sync::Mutex::new(heartbeat)));
tracing::debug!("[ClientCore] heartbeat started after negotiation");
}
pub fn stop_heartbeat(&self) {
let taken = self
.heartbeat_manager
.lock()
.ok()
.and_then(|mut slot| slot.take());
if let Some(heartbeat) = taken {
Self::stop_heartbeat_async(heartbeat);
}
}
fn stop_heartbeat_async(heartbeat: Arc<tokio::sync::Mutex<HeartbeatManager>>) {
#[cfg(not(target_arch = "wasm32"))]
{
crate::client::runtime::run_client_async(async {
let mut hb_guard = heartbeat.lock().await;
hb_guard.stop();
});
}
#[cfg(target_arch = "wasm32")]
{
crate::client::runtime::spawn_client_task(async move {
let mut hb_guard = heartbeat.lock().await;
hb_guard.stop();
});
}
}
pub async fn handle_message(&self, data: Vec<u8>) {
let negotiation_completed = self.negotiation_completed.load(Ordering::SeqCst);
let frame = if !negotiation_completed {
use crate::common::message::parser::PRE_NEGOTIATION_PARSER;
match PRE_NEGOTIATION_PARSER.parse(&data) {
Ok(frame) => frame,
Err(e) => {
#[cfg(target_arch = "wasm32")]
web_sys::console::warn_1(
&format!("[flare-core] parse failed pre-negotiation: {e}").into(),
);
tracing::warn!("Failed to parse message (pre-negotiation): {}", e);
return;
}
}
} else {
match self.parse_message(&data).await {
Ok(frame) => frame,
Err(e) => {
#[cfg(target_arch = "wasm32")]
web_sys::console::warn_1(
&format!("[flare-core] parse failed negotiated: {e}").into(),
);
tracing::warn!("Failed to parse message (negotiated): {}", e);
return;
}
}
};
let is_pending_response = {
tracing::trace!(
"[ClientCore] 尝试匹配等待的响应: message_id={}",
frame.message_id
);
if let Some(cmd) = &frame.command
&& let Some(Type::Payload(msg_cmd)) = &cmd.r#type
&& msg_cmd.message_id != frame.message_id
{
tracing::warn!(
"[ClientCore] PayloadCommand.message_id 和 Frame.message_id 不一致: cmd_id={}, frame_id={}",
msg_cmd.message_id,
frame.message_id
);
}
let pending_ids: Vec<String> = {
let pending = self.pending_map.lock().await;
pending.keys().cloned().collect()
};
if !pending_ids.is_empty() {
tracing::debug!(
"[ClientCore] handle_message: 当前等待的响应 message_id 列表: {:?}",
pending_ids
);
}
let mut pending = self.pending_map.lock().await;
if let Some(sender) = pending.remove(&frame.message_id) {
tracing::debug!(
"[ClientCore] ✅ 匹配到等待的响应: message_id={}",
frame.message_id
);
if sender.send(frame.clone()).is_err() {
tracing::warn!(
"[ClientCore] 发送响应到等待通道失败: message_id={} (接收者可能已关闭)",
frame.message_id
);
false } else {
tracing::debug!(
"[ClientCore] ✅ 响应已发送到等待通道: message_id={}",
frame.message_id
);
true }
} else {
tracing::debug!(
"[ClientCore] ❌ 未找到等待的响应: message_id={}",
frame.message_id
);
false }
};
if is_pending_response {
self.notify_observers(&ConnectionEvent::Message(data));
return;
}
let is_system_command = self.handle_system_commands(&frame).await;
self.notify_observers(&ConnectionEvent::Message(data));
if is_system_command {
return;
}
self.handle_business_commands(&frame).await;
self.handle_message_routing(&frame).await;
}
async fn parse_message(&self, data: &[u8]) -> Result<Frame> {
let parser = self.parser.lock().await;
parser.parse(data)
}
async fn handle_system_commands(&self, frame: &Frame) -> bool {
let Some(cmd) = &frame.command else {
return false;
};
let Some(Type::System(sys_cmd)) = &cmd.r#type else {
return false;
};
let cmd_type = match SystemCommandType::try_from(sys_cmd.r#type) {
Ok(t) => t,
Err(_) => return false,
};
match cmd_type {
SystemCommandType::ConnectAck => {
self.handle_connect_ack_command(frame).await;
true
}
SystemCommandType::Pong => {
self.handle_pong_command(frame).await;
true
}
SystemCommandType::Kicked => {
self.handle_kicked_command(frame).await;
true
}
_ => false,
}
}
async fn handle_connect_ack_command(&self, frame: &Frame) {
if let Some(ref handler) = self.event_handler {
let _ = handler
.handle_system_command(SystemCommandType::ConnectAck, frame)
.await;
}
match self.handle_connect_ack(frame) {
Ok((format, compression, encryption)) => {
tracing::info!(
"[ClientCore] ✅ 收到 CONNECT_ACK: 服务端确定的序列化方式={:?}, 压缩方式={:?}, 加密方式={:?}",
format,
compression,
encryption
);
if !self.config.is_force_format() {
self.update_parser(format, compression.clone(), encryption.clone())
.await;
tracing::info!(
"[ClientCore] ✅ 解析器已更新为协商后的格式: {:?}, 压缩: {:?}, 加密: {:?}",
format,
compression,
encryption
);
} else {
tracing::info!(
"[ClientCore] ℹ️ 强制模式:继续使用客户端强制指定的格式: {:?}, 压缩: {:?}",
self.config.get_serialization_format(),
self.config.get_compression()
);
self.negotiation_completed.store(true, Ordering::SeqCst);
self.negotiation_notify.notify_waiters();
}
if let Err(e) = self.send_negotiation_ready().await {
tracing::warn!("[ClientCore] 发送 NEGOTIATION_READY 失败: {}", e);
}
self.try_start_heartbeat().await;
}
Err(e) => {
self.fail_negotiation(e.to_string()).await;
}
}
}
async fn handle_pong_command(&self, frame: &Frame) {
if let Some(ref handler) = self.event_handler {
let _ = handler
.handle_system_command(SystemCommandType::Pong, frame)
.await;
}
self.record_pong();
}
async fn handle_kicked_command(&self, frame: &Frame) {
let Some(cmd) = &frame.command else {
return;
};
let Some(Type::System(sys_cmd)) = &cmd.r#type else {
return;
};
let reason = sys_cmd.message.clone();
tracing::warn!("[ClientCore] ⚠️ 收到被踢消息: {}", reason);
let kick_reason = Self::parse_kick_reason(&reason, sys_cmd);
if let Some(ref handler) = self.event_handler
&& let Err(e) = handler
.handle_system_command(SystemCommandType::Kicked, frame)
.await
{
tracing::warn!("[ClientCore] 事件处理器处理 KICKED 失败: {}", e);
}
self.state_manager.set_disconnected();
self.disconnect_on_kicked().await;
self.cancel_all_pending_responses().await;
let should_notify = self.negotiation_completed.load(Ordering::SeqCst)
&& !self.disconnect_requested.load(Ordering::SeqCst);
if should_notify {
self.notify_observers(&ConnectionEvent::Disconnected(kick_reason.clone()));
tracing::info!("[ClientCore] 连接已断开(被踢): {}", kick_reason);
} else {
tracing::debug!(
"[ClientCore] 收到 KICKED 但不向观察者通知(协商未完成或我方已请求断开)"
);
}
}
fn parse_kick_reason(
base_reason: &str,
sys_cmd: &crate::common::protocol::SystemCommand,
) -> String {
if let Some(reason_bytes) = sys_cmd.metadata.get("reason")
&& let Ok(reason_str) = String::from_utf8(reason_bytes.clone())
&& reason_str == "device_conflict"
{
return format!("设备冲突:{}", base_reason);
}
base_reason.to_string()
}
async fn disconnect_on_kicked(&self) {
let client_conn_opt = self.take_client_connection();
if let Some(client_conn) = client_conn_opt {
let mut conn = client_conn.lock().await;
if let Err(e) = conn.close().await {
tracing::error!("[ClientCore] 断开连接失败: {}", e);
} else {
tracing::info!("[ClientCore] ✅ 已主动断开连接(被踢)");
}
} else {
tracing::warn!("[ClientCore] ⚠️ 客户端连接未设置,等待底层传输层关闭连接");
}
}
async fn handle_business_commands(&self, frame: &Frame) {
let Some(ref handler) = self.event_handler else {
return;
};
let Some(cmd) = &frame.command else {
return;
};
match &cmd.r#type {
Some(Type::Payload(msg_cmd)) => {
if let Ok(cmd_type) = PayloadCommandType::try_from(msg_cmd.r#type) {
let _ = handler.handle_message_command(cmd_type, frame).await;
}
}
Some(Type::Notification(notif_cmd)) => {
if let Ok(cmd_type) = NotificationCommandType::try_from(notif_cmd.r#type) {
let _ = handler.handle_notification_command(cmd_type, frame).await;
}
}
_ => {}
}
}
async fn handle_message_routing(&self, frame: &Frame) {
let Some(ref router) = self.message_router else {
return;
};
match router.route(frame).await {
Ok(replies) => {
tracing::debug!("Router generated {} replies", replies.len());
}
Err(e) => {
tracing::warn!("Router error: {}", e);
}
}
}
async fn fail_negotiation(&self, reason: String) {
tracing::warn!("[ClientCore] 协商失败: {}", reason);
if let Ok(mut stored) = self.negotiation_failure_reason.lock() {
*stored = Some(reason);
}
self.negotiation_notify.notify_waiters();
self.stop_heartbeat();
self.cancel_all_pending_responses().await;
if let Ok(reason) = self.negotiation_failure_reason.lock()
&& let Some(msg) = reason.as_ref()
{
self.handle_connection_event(&ConnectionEvent::Error(FlareError::protocol_error(
msg.clone(),
)));
}
}
fn negotiation_failure_error(&self) -> Option<FlareError> {
self.negotiation_failure_reason
.lock()
.ok()
.and_then(|reason| {
reason
.as_ref()
.map(|msg| FlareError::protocol_error(msg.clone()))
})
}
fn reset_negotiation_state(&self) {
self.negotiation_completed.store(false, Ordering::SeqCst);
if let Ok(mut reason) = self.negotiation_failure_reason.lock() {
*reason = None;
}
}
pub fn handle_connection_event(&self, event: &ConnectionEvent) {
if let Some(ref handler) = self.event_handler {
let handler_clone = Arc::clone(handler);
let event_clone = event.clone();
crate::client::runtime::spawn_client_task(async move {
let _ = handler_clone.handle_connection_event(&event_clone).await;
});
}
match event {
ConnectionEvent::Connected => {
self.state_manager.set_connected();
self.reset_negotiation_state();
}
ConnectionEvent::Disconnected(_) => {
self.state_manager.set_disconnected();
self.reset_negotiation_state();
let pending = Arc::clone(&self.pending_map);
crate::client::runtime::spawn_client_task(async move {
let mut map = pending.lock().await;
if !map.is_empty() {
tracing::debug!(
count = map.len(),
"[ClientCore] connection disconnected: clearing pending response waiters"
);
map.clear();
}
});
}
ConnectionEvent::Error(_) => {
self.state_manager.set_failed();
self.reset_negotiation_state();
let pending = Arc::clone(&self.pending_map);
crate::client::runtime::spawn_client_task(async move {
let mut map = pending.lock().await;
if !map.is_empty() {
tracing::debug!(
count = map.len(),
"[ClientCore] connection error: clearing pending response waiters"
);
map.clear();
}
});
}
ConnectionEvent::Message(_) => {
}
}
self.notify_observers(event);
}
pub fn add_observer(&self, observer: ArcObserver) {
if let Ok(mut observers) = self.observers.lock() {
observers.push(observer);
}
}
pub fn remove_observer(&self, observer: ArcObserver) {
if let Ok(mut observers) = self.observers.lock() {
observers.retain(|o| !Arc::ptr_eq(o, &observer));
}
}
fn notify_observers(&self, event: &ConnectionEvent) {
if let Ok(observers) = self.observers.lock() {
for observer in observers.iter() {
observer.on_event(event);
}
}
}
pub fn router_mut(&mut self) -> Option<&mut MessageRouter> {
self.message_router.as_mut()
}
pub fn router(&self) -> Option<&MessageRouter> {
self.message_router.as_ref()
}
pub fn state(&self) -> crate::client::connection::ConnectionState {
self.state_manager.get_state()
}
pub fn can_send(&self) -> bool {
self.state_manager.get_state().can_send()
}
pub fn is_negotiation_completed(&self) -> bool {
self.negotiation_completed
.load(std::sync::atomic::Ordering::SeqCst)
}
pub async fn wait_for_negotiation(&self, timeout: std::time::Duration) -> Result<()> {
if self.is_negotiation_completed() {
return Ok(());
}
if let Some(err) = self.negotiation_failure_error() {
return Err(err);
}
#[cfg(target_arch = "wasm32")]
{
use crate::common::platform::monotonic_now;
let deadline = monotonic_now() + timeout;
loop {
self.drain_wasm_inbound().await;
if self.is_negotiation_completed() {
return Ok(());
}
if let Some(err) = self.negotiation_failure_error() {
return Err(err);
}
if monotonic_now() >= deadline {
return Err(negotiation_timeout_error(timeout));
}
crate::common::platform::yield_to_event_loop().await;
}
}
#[cfg(not(target_arch = "wasm32"))]
{
wait_for_negotiation_notify(
Arc::clone(&self.negotiation_completed),
Arc::clone(&self.negotiation_failure_reason),
Arc::clone(&self.negotiation_notify),
timeout,
)
.await
}
}
pub fn can_connect(&self) -> bool {
self.state_manager.get_state().can_connect()
}
pub fn current_heartbeat_config(&self) -> HeartbeatConfig {
self.heartbeat_config
.read()
.map(|guard| guard.clone())
.unwrap_or_else(|_| self.config.heartbeat.clone())
}
pub fn heartbeat_effective_interval(&self) -> std::time::Duration {
self.current_heartbeat_config().effective_interval()
}
pub fn update_heartbeat_config(&self, update: impl FnOnce(&mut HeartbeatConfig)) {
if let Ok(mut config) = self.heartbeat_config.write() {
update(&mut config);
}
}
pub fn set_heartbeat_app_state(&self, state: HeartbeatAppState) {
self.update_heartbeat_config(|config| {
config.app_state = state;
});
}
pub fn set_heartbeat_nat_timeout(&self, timeout: Option<std::time::Duration>) {
self.update_heartbeat_config(|config| {
config.nat_timeout = timeout;
});
}
pub fn record_pong(&self) {
let heartbeat = match self.heartbeat_manager.lock() {
Ok(guard) => guard.as_ref().map(Arc::clone),
Err(_) => None,
};
let Some(heartbeat) = heartbeat else {
return;
};
#[cfg(not(target_arch = "wasm32"))]
{
crate::client::runtime::run_client_async(async {
let hb_guard = heartbeat.lock().await;
hb_guard.record_pong();
});
}
#[cfg(target_arch = "wasm32")]
{
let heartbeat = Arc::clone(&heartbeat);
crate::client::runtime::spawn_client_task(async move {
let hb_guard = heartbeat.lock().await;
hb_guard.record_pong();
});
}
}
}
impl ClientCore {
pub async fn register_pending_response(&self, message_id: &str) -> oneshot::Receiver<Frame> {
let (tx, rx) = oneshot::channel();
let mut pending = self.pending_map.lock().await;
pending.insert(message_id.to_string(), tx);
rx
}
pub async fn cancel_pending_response(&self, message_id: &str) {
let mut pending = self.pending_map.lock().await;
pending.remove(message_id);
}
pub async fn cancel_all_pending_responses(&self) {
let mut pending = self.pending_map.lock().await;
if !pending.is_empty() {
tracing::debug!(
count = pending.len(),
"[ClientCore] clearing all pending response waiters"
);
pending.clear();
}
}
}
impl Clone for ClientCore {
fn clone(&self) -> Self {
Self {
state_manager: Arc::clone(&self.state_manager),
parser: Arc::clone(&self.parser),
heartbeat_manager: Arc::clone(&self.heartbeat_manager),
heartbeat_config: Arc::clone(&self.heartbeat_config),
message_router: self.message_router.as_ref().map(|_| MessageRouter::new()), observers: Arc::clone(&self.observers),
config: self.config.clone(),
event_handler: self.event_handler.clone(), client_connection: Arc::clone(&self.client_connection), pending_map: Arc::clone(&self.pending_map),
negotiation_completed: Arc::clone(&self.negotiation_completed), negotiation_notify: Arc::clone(&self.negotiation_notify),
negotiation_failure_reason: Arc::clone(&self.negotiation_failure_reason),
disconnect_requested: Arc::clone(&self.disconnect_requested),
#[cfg(target_arch = "wasm32")]
wasm_inbound: Arc::clone(&self.wasm_inbound),
}
}
}
#[cfg(all(test, not(target_arch = "wasm32")))]
mod client_core_tests {
use super::*;
use crate::common::compression::CompressionAlgorithm;
use crate::common::encryption::EncryptionAlgorithm;
use crate::common::protocol::SerializationFormat;
use std::time::Duration;
#[tokio::test]
async fn update_parser_marks_negotiation_completed() {
let core = ClientCore::new(&ClientConfig::default());
assert!(!core.is_negotiation_completed());
core.update_parser(
SerializationFormat::Protobuf,
CompressionAlgorithm::Gzip,
EncryptionAlgorithm::None,
)
.await;
assert!(core.is_negotiation_completed());
}
#[tokio::test]
async fn wait_for_negotiation_returns_after_flag_set() {
let core = ClientCore::new(&ClientConfig::default());
let core = Arc::new(core);
let waiter = {
let core = Arc::clone(&core);
tokio::spawn(async move {
core.wait_for_negotiation(Duration::from_secs(1))
.await
.expect("negotiation wait")
})
};
tokio::time::sleep(Duration::from_millis(20)).await;
core.update_parser(
SerializationFormat::Json,
CompressionAlgorithm::None,
EncryptionAlgorithm::None,
)
.await;
waiter.await.expect("wait task");
}
#[tokio::test]
async fn start_heartbeat_before_negotiation_does_not_panic() {
let core = ClientCore::new(&ClientConfig::default());
core.start_heartbeat(Arc::new(Mutex::new(
Box::new(MockConnection) as Box<dyn Connection>
)))
.await;
assert!(!core.is_negotiation_completed());
}
#[test]
fn heartbeat_runtime_policy_is_shared_across_core_clones() {
let core = ClientCore::new(&ClientConfig::default());
let cloned = core.clone();
assert_eq!(core.heartbeat_effective_interval(), Duration::from_secs(30));
cloned.set_heartbeat_app_state(HeartbeatAppState::Background);
assert_eq!(
core.heartbeat_effective_interval(),
Duration::from_secs(120)
);
core.set_heartbeat_nat_timeout(Some(Duration::from_secs(40)));
assert_eq!(
cloned.heartbeat_effective_interval(),
Duration::from_secs(28)
);
}
#[test]
fn client_connection_take_clears_shared_core_slot() {
let mut core = ClientCore::new(&ClientConfig::default());
let cloned = core.clone();
core.set_client_connection(Arc::new(Mutex::new(
Box::new(MockConnection) as Box<dyn Connection>
)));
assert!(cloned.take_client_connection().is_some());
assert!(core.take_client_connection().is_none());
}
struct MockConnection;
#[async_trait::async_trait]
impl Connection for MockConnection {
fn add_observer(&mut self, _observer: crate::transport::events::ArcObserver) {}
fn remove_observer(&mut self, _observer: crate::transport::events::ArcObserver) {}
async fn send(&mut self, _data: &[u8]) -> Result<()> {
Ok(())
}
async fn close(&mut self) -> Result<()> {
Ok(())
}
fn last_active_time(&self) -> crate::common::platform::MonotonicInstant {
crate::common::platform::monotonic_now()
}
fn update_active_time(&mut self) {}
}
}