use super::*;
impl ConnectionManager {
pub fn get_user_connections(&self, user_id: &str) -> Vec<String> {
self.user_connection_shard(user_id)
.read()
.ok()
.and_then(|user_connections| user_connections.get(user_id).cloned())
.unwrap_or_default()
}
pub fn bind_user(&self, connection_id: &str, user_id: String) -> Result<()> {
let mut shard = self
.connection_shard(connection_id)
.write()
.map_err(|_| FlareError::general_error("Failed to lock connection shard"))?;
let (_, _, info) = shard.get_mut(connection_id).ok_or_else(|| {
FlareError::protocol_error(format!("Connection {} not found", connection_id))
})?;
if let Some(old_user_id) = &info.user_id {
self.remove_user_connection(old_user_id, connection_id)?;
}
info.user_id = Some(user_id.clone());
self.insert_user_connection(user_id, connection_id)?;
Ok(())
}
pub fn update_connection_active(&self, connection_id: &str) -> Result<()> {
let mut shard = self
.connection_shard(connection_id)
.write()
.map_err(|_| FlareError::general_error("Failed to lock connection shard"))?;
if let Some((_, _, info)) = shard.get_mut(connection_id) {
info.update_active();
drop(shard);
Ok(())
} else {
drop(shard);
Err(FlareError::protocol_error(format!(
"Connection {} not found",
connection_id
)))
}
}
pub(super) fn update_connections_active<I>(&self, connection_ids: I) -> usize
where
I: IntoIterator<Item = String>,
{
let now = Instant::now();
let mut ids_by_shard = vec![Vec::new(); self.connection_shards.len()];
for connection_id in connection_ids {
let shard_index = self.connection_shard_index(&connection_id);
ids_by_shard[shard_index].push(connection_id);
}
ids_by_shard
.into_iter()
.enumerate()
.map(|(shard_index, ids)| {
if ids.is_empty() {
return 0;
}
let Ok(mut shard) = self.connection_shards[shard_index].write() else {
return 0;
};
let mut updated = 0;
for connection_id in ids {
if let Some((_, _, info)) = shard.get_mut(&connection_id) {
info.last_active = now;
updated += 1;
}
}
updated
})
.sum()
}
pub(super) async fn fanout_serialized_frame_to_auth_snapshots(
&self,
snapshots: Vec<ConnectionAuthSnapshot>,
frame: &crate::common::protocol::Frame,
data: &[u8],
fanout: &'static str,
) -> Vec<String> {
stream::iter(snapshots)
.map(|snapshot| async move {
let connection_id = snapshot.0.clone();
match self
.send_serialized_frame_to_auth_snapshot_without_active(snapshot, frame, data)
.await
{
Ok(connection_id) => Some(connection_id),
Err(error) => {
tracing::warn!(
connection_id = %connection_id,
error = ?error,
fanout,
"Connection frame fanout failed"
);
None
}
}
})
.buffer_unordered(self.fanout_concurrency)
.filter_map(|connection_id| async move { connection_id })
.collect()
.await
}
pub(super) async fn fanout_frame_grouped_by_encoding(
&self,
connection_ids: Vec<String>,
frame: &crate::common::protocol::Frame,
) -> (i32, i32) {
use std::collections::HashMap;
let snapshots = self.connection_snapshots_for_ids(connection_ids);
let total = snapshots.len();
if total == 0 {
return (0, 0);
}
let mut plain_groups: HashMap<
(
crate::common::protocol::SerializationFormat,
crate::common::compression::CompressionAlgorithm,
),
Vec<ConnectionSnapshot>,
> = HashMap::new();
let mut encrypted: Vec<ConnectionSnapshot> = Vec::new();
for snapshot in snapshots {
let info = &snapshot.2;
if info.encryption == crate::common::encryption::EncryptionAlgorithm::None {
plain_groups
.entry((info.serialization_format, info.compression.clone()))
.or_default()
.push(snapshot);
} else {
encrypted.push(snapshot);
}
}
let mut successful: Vec<String> = Vec::with_capacity(total);
for ((format, compression), group) in plain_groups {
let parser = crate::common::MessageParser::new(
format,
compression,
crate::common::encryption::EncryptionAlgorithm::None,
);
let data = match parser.serialize(frame) {
Ok(data) => data,
Err(error) => {
tracing::warn!(
error = ?error,
group_size = group.len(),
"grouped frame serialization failed; group counted as failures"
);
continue;
}
};
let sent: Vec<String> = stream::iter(group)
.map(|snapshot| {
let data = &data;
async move {
let (connection_id, connection, info) = snapshot;
if let Err(error) =
Self::ensure_frame_allowed_for_connection(&connection_id, &info, frame)
{
tracing::warn!(connection_id = %connection_id, error = ?error, "grouped fanout skipped");
return None;
}
match self
.send_to_connection_handle(&connection_id, connection, data)
.await
{
Ok(()) => Some(connection_id),
Err(error) => {
tracing::warn!(connection_id = %connection_id, error = ?error, "grouped fanout write failed");
None
}
}
}
})
.buffer_unordered(self.fanout_concurrency)
.filter_map(|connection_id| async move { connection_id })
.collect()
.await;
successful.extend(sent);
}
if !encrypted.is_empty() {
let sent = self
.fanout_frame_to_snapshots(encrypted, frame, None, "grouped_fanout_encrypted")
.await;
successful.extend(sent);
}
let success = successful.len() as i32;
self.update_connections_active(successful);
(success, (total as i32).saturating_sub(success))
}
pub(super) async fn fanout_frame_to_snapshots(
&self,
snapshots: Vec<ConnectionSnapshot>,
frame: &crate::common::protocol::Frame,
parser: Option<&crate::common::MessageParser>,
fanout: &'static str,
) -> Vec<String> {
stream::iter(snapshots)
.map(|snapshot| async move {
let connection_id = snapshot.0.clone();
match self
.send_frame_to_snapshot_without_active(snapshot, frame, parser)
.await
{
Ok(connection_id) => Some(connection_id),
Err(error) => {
tracing::warn!(
connection_id = %connection_id,
error = ?error,
fanout,
"Connection frame fanout failed"
);
None
}
}
})
.buffer_unordered(self.fanout_concurrency)
.filter_map(|connection_id| async move { connection_id })
.collect()
.await
}
pub fn set_connection_authenticated(
&self,
connection_id: &str,
user_id: Option<String>,
) -> Result<()> {
let mut shard = self
.connection_shard(connection_id)
.write()
.map_err(|_| FlareError::general_error("Failed to lock connection shard"))?;
let (_, _, info) = shard.get_mut(connection_id).ok_or_else(|| {
FlareError::protocol_error(format!("Connection {} not found", connection_id))
})?;
let old_user_id = info.user_id.clone();
let final_user_id = user_id.or(old_user_id.clone());
info.set_authenticated(final_user_id.clone());
if let Some(user_id) = final_user_id {
let user_id_changed = old_user_id
.as_ref()
.map(|old| old != &user_id)
.unwrap_or(true);
if user_id_changed {
if let Some(old_user_id) = old_user_id {
self.remove_user_connection(&old_user_id, connection_id)?;
}
self.insert_user_connection(user_id, connection_id)?;
} else {
self.insert_user_connection(user_id, connection_id)?;
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn update_connection_negotiation(
&self,
connection_id: &str,
device_info: Option<crate::common::device::DeviceInfo>,
serialization_format: crate::common::protocol::SerializationFormat,
compression: crate::common::compression::CompressionAlgorithm,
encryption: crate::common::encryption::EncryptionAlgorithm,
user_id: Option<String>,
metadata: Option<HashMap<String, String>>,
) -> Result<()> {
let mut shard = self
.connection_shard(connection_id)
.write()
.map_err(|_| FlareError::general_error("Failed to lock connection shard"))?;
let (_, _, info) = shard.get_mut(connection_id).ok_or_else(|| {
FlareError::protocol_error(format!("Connection {} not found", connection_id))
})?;
info.device_info = device_info;
info.serialization_format = serialization_format;
info.compression = compression;
info.encryption = encryption;
if let Some(meta) = metadata {
for (k, v) in meta {
info.metadata.insert(k, v);
}
}
let old_user_id = info.user_id.clone();
let user_id_to_set = user_id.clone().or(old_user_id.clone());
if user_id_to_set.is_none() {
tracing::trace!(connection_id = %connection_id,incoming_user_id = ?user_id,old_user_id = ?old_user_id,"update_connection_negotiation: user_id_to_set is None, user_id will not be set");
}
if let Some(user_id_val) = user_id_to_set {
if let Some(old_user_id) = old_user_id
&& old_user_id != user_id_val
{
self.remove_user_connection(&old_user_id, connection_id)?;
}
info.user_id = Some(user_id_val.clone());
self.insert_user_connection(user_id_val, connection_id)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn update_connection_negotiation_with_pipeline(
&self,
connection_id: &str,
device_info: Option<crate::common::device::DeviceInfo>,
serialization_format: crate::common::protocol::SerializationFormat,
compression: crate::common::compression::CompressionAlgorithm,
encryption: crate::common::encryption::EncryptionAlgorithm,
user_id: Option<String>,
parser: crate::common::MessageParser,
pipeline: Option<std::sync::Arc<crate::common::message::pipeline::MessagePipeline>>,
) -> Result<()> {
let mut shard = self
.connection_shard(connection_id)
.write()
.map_err(|_| FlareError::general_error("Failed to lock connection shard"))?;
let (_, _, info) = shard.get_mut(connection_id).ok_or_else(|| {
FlareError::protocol_error(format!("Connection {} not found", connection_id))
})?;
info.device_info = device_info;
info.serialization_format = serialization_format;
info.compression = compression;
info.encryption = encryption;
info.negotiation_completed = true;
info.cached_parser = Some(std::sync::Arc::new(parser));
info.cached_pipeline = pipeline;
let old_user_id = info.user_id.clone();
let user_id_to_set = user_id.clone().or(old_user_id.clone());
if user_id_to_set.is_none() {
tracing::trace!(
connection_id = %connection_id,
incoming_user_id = ?user_id,
old_user_id = ?old_user_id,
"update_connection_negotiation_with_pipeline: user_id_to_set is None, user_id will not be set"
);
}
if let Some(user_id_val) = user_id_to_set {
if let Some(old_user_id) = old_user_id
&& old_user_id != user_id_val
{
self.remove_user_connection(&old_user_id, connection_id)?;
}
info.user_id = Some(user_id_val.clone());
self.insert_user_connection(user_id_val, connection_id)?;
}
Ok(())
}
pub fn mark_negotiation_confirmed(&self, connection_id: &str) -> Result<()> {
let mut shard = self
.connection_shard(connection_id)
.write()
.map_err(|_| FlareError::general_error("Failed to lock connection shard"))?;
let (_, _, info) = shard.get_mut(connection_id).ok_or_else(|| {
FlareError::protocol_error(format!("Connection {} not found", connection_id))
})?;
if !info.negotiation_completed {
return Err(FlareError::protocol_error(format!(
"Cannot confirm negotiation for connection {}: negotiation not completed",
connection_id
)));
}
info.negotiation_confirmed = true;
tracing::trace!(
"[ConnectionManager] 协商已确认: connection_id={}",
connection_id
);
Ok(())
}
pub fn list_connections(&self) -> Vec<String> {
self.connection_shards
.iter()
.flat_map(|shard| {
shard
.read()
.ok()
.map(|connections| connections.keys().cloned().collect::<Vec<_>>())
.unwrap_or_default()
})
.collect()
}
pub fn connection_count(&self) -> usize {
self.connection_count.load(Ordering::Relaxed)
}
pub fn user_count(&self) -> usize {
self.user_count.load(Ordering::Relaxed)
}
pub fn cleanup_timeout_connections(&self, timeout: Duration) -> Vec<String> {
let timeout_connections = self.timeout_connection_snapshots(timeout);
self.remove_connection_snapshots(
timeout_connections
.iter()
.map(|(connection_id, _, _)| connection_id.clone()),
)
}
pub fn stats(&self) -> TraitConnectionStats {
let total_connections = self.connection_count();
let total_users = self.user_count();
TraitConnectionStats {
total_connections,
total_users,
}
}
pub(super) fn frame_allowed_before_auth(frame: &crate::common::protocol::Frame) -> bool {
frame
.command
.as_ref()
.and_then(|cmd| {
if let Some(crate::common::protocol::flare::core::commands::command::Type::System(
sys_cmd,
)) = &cmd.r#type
{
Some(
sys_cmd.r#type
== crate::common::protocol::flare::core::commands::system_command::Type::ConnectAck
as i32
|| sys_cmd.r#type
== crate::common::protocol::flare::core::commands::system_command::Type::Ping
as i32
|| sys_cmd.r#type
== crate::common::protocol::flare::core::commands::system_command::Type::Pong
as i32
|| sys_cmd.r#type
== crate::common::protocol::flare::core::commands::system_command::Type::Error
as i32
|| sys_cmd.r#type
== crate::common::protocol::flare::core::commands::system_command::Type::Close
as i32,
)
} else {
None
}
})
.unwrap_or(false)
}
pub(super) fn serialize_frame_for_connection(
connection_id: &str,
info: &ConnectionInfo,
frame: &crate::common::protocol::Frame,
parser: Option<&crate::common::MessageParser>,
) -> Result<Vec<u8>> {
Self::ensure_frame_allowed_for_connection(connection_id, info, frame)?;
if let Some(parser) = parser {
return parser.serialize(frame);
}
if let Some(parser) = &info.cached_parser {
return parser.serialize(frame);
}
crate::common::MessageParser::new(
info.serialization_format,
info.compression.clone(),
info.encryption.clone(),
)
.serialize(frame)
}
pub(super) fn ensure_frame_allowed_for_connection(
connection_id: &str,
info: &ConnectionInfo,
frame: &crate::common::protocol::Frame,
) -> Result<()> {
if info.authenticated || Self::frame_allowed_before_auth(frame) {
return Ok(());
}
Err(FlareError::authentication_failed(format!(
"连接 {} 未验证,无法发送消息",
connection_id
)))
}
pub(super) async fn send_to_connection_handle(
&self,
connection_id: &str,
connection: ConnectionWriteHandle,
data: &[u8],
) -> Result<()> {
match connection.try_enqueue(data) {
Ok(()) => Ok(()),
Err(err) => {
connection.close_underlying_in_background();
let _ = ConnectionManager::remove_connection(self, connection_id);
Err(err)
}
}
}
pub(super) async fn send_frame_to_snapshot(
&self,
snapshot: ConnectionSnapshot,
frame: &crate::common::protocol::Frame,
parser: Option<&crate::common::MessageParser>,
) -> Result<()> {
let connection_id = self
.send_frame_to_snapshot_without_active(snapshot, frame, parser)
.await?;
ConnectionManager::update_connection_active(self, &connection_id)?;
Ok(())
}
pub(super) async fn send_frame_to_snapshot_without_active(
&self,
snapshot: ConnectionSnapshot,
frame: &crate::common::protocol::Frame,
parser: Option<&crate::common::MessageParser>,
) -> Result<String> {
let (connection_id, connection, info) = snapshot;
let data = Self::serialize_frame_for_connection(&connection_id, &info, frame, parser)?;
self.send_to_connection_handle(&connection_id, connection, &data)
.await?;
Ok(connection_id)
}
pub(super) async fn send_serialized_frame_to_auth_snapshot_without_active(
&self,
snapshot: ConnectionAuthSnapshot,
frame: &crate::common::protocol::Frame,
data: &[u8],
) -> Result<String> {
let (connection_id, connection, authenticated) = snapshot;
if !authenticated && !Self::frame_allowed_before_auth(frame) {
return Err(FlareError::authentication_failed(format!(
"连接 {} 未验证,无法发送消息",
connection_id
)));
}
self.send_to_connection_handle(&connection_id, connection, data)
.await?;
Ok(connection_id)
}
}