use std::{
sync::{
Arc,
atomic::{
AtomicBool, AtomicI32, AtomicI64, AtomicU64,
Ordering::{AcqRel, Acquire, Release},
},
},
thread::yield_now,
};
use gxhash::HashMap as GxHashMap;
use parking_lot::RwLock;
use super::i_garnet_server::{
ClusterSessionFace, MessageConsumerFace, ServerEnumerate, ServerError, SessionProviderFace,
WireFormat,
};
pub const DEFAULT_NETWORK_BUFFER_SIZE: usize = 1 << 16;
pub const DISPOSED_HANDLER_COUNT: i32 = i32::MIN;
pub struct GarnetServerBase {
active_handlers: RwLock<GxHashMap<u64, Arc<dyn MessageConsumerFace>>>,
next_handler_id: AtomicU64,
active_handler_count: AtomicI32,
session_providers: RwLock<GxHashMap<WireFormat, Arc<dyn SessionProviderFace>>>,
network_buffer_size: usize,
endpoint: String,
disposed: AtomicBool,
total_connections_received: AtomicI64,
total_connections_disposed: AtomicI64,
}
impl GarnetServerBase {
pub fn new(endpoint: &str, network_buffer_size: usize) -> Self {
Self {
active_handlers: RwLock::new(GxHashMap::default()),
next_handler_id: AtomicU64::new(0),
active_handler_count: AtomicI32::new(0),
session_providers: RwLock::new(GxHashMap::default()),
network_buffer_size: if network_buffer_size == 0 {
DEFAULT_NETWORK_BUFFER_SIZE
} else {
network_buffer_size
},
endpoint: endpoint.to_string(),
disposed: AtomicBool::new(false),
total_connections_received: AtomicI64::new(0),
total_connections_disposed: AtomicI64::new(0),
}
}
pub fn endpoint(&self) -> &str {
&self.endpoint
}
pub fn network_buffer_size(&self) -> usize {
self.network_buffer_size
}
pub fn disposed(&self) -> bool {
self.disposed.load(Acquire)
}
#[inline]
pub fn increment_connections_received(&self) {
self.total_connections_received.fetch_add(1, AcqRel);
}
#[inline]
pub fn increment_connections_disposed(&self) {
self.total_connections_disposed.fetch_add(1, AcqRel);
}
pub fn total_connections_received(&self) -> i64 {
self.total_connections_received.load(Acquire)
}
pub fn total_connections_disposed(&self) -> i64 {
self.total_connections_disposed.load(Acquire)
}
pub fn get_conn_active(&self) -> i64 {
self.active_handlers.read().len() as i64
}
pub fn reset_connections_received(&self) {
let active = self.active_handlers.read().len() as i64;
self.total_connections_received.store(active, Release);
}
pub fn reset_connections_disposed(&self) {
self.total_connections_disposed.store(0, Release);
}
pub fn register_handler(&self, consumer: Arc<dyn MessageConsumerFace>) -> Option<u64> {
self.admit_handler(consumer, -1)
}
pub fn admit_handler(
&self,
consumer: Arc<dyn MessageConsumerFace>,
connection_limit: i32,
) -> Option<u64> {
let count = self.active_handler_count.fetch_add(1, AcqRel) + 1;
if count < 0 || (connection_limit != -1 && count > connection_limit) {
self.active_handler_count.fetch_sub(1, AcqRel);
return None;
}
let handler_id = self.next_handler_id.fetch_add(1, AcqRel);
self.active_handlers.write().insert(handler_id, consumer);
Some(handler_id)
}
pub fn remove_handler(&self, handler_id: u64) -> bool {
self.active_handlers.write().remove(&handler_id).is_some()
}
pub fn active_handler_ids(&self) -> Vec<u64> {
self.active_handlers.read().keys().copied().collect()
}
pub fn register(
&self,
wire_format: WireFormat,
backend_provider: Arc<dyn SessionProviderFace>,
) -> Result<(), ServerError> {
let mut providers = self.session_providers.write();
if providers.contains_key(&wire_format) {
return Err(ServerError::WireFormatAlreadyRegistered(wire_format));
}
providers.insert(wire_format, backend_provider);
Ok(())
}
pub fn unregister(&self, wire_format: WireFormat) -> Option<Arc<dyn SessionProviderFace>> {
self.session_providers.write().remove(&wire_format)
}
pub fn get_session_providers(&self) -> Vec<(WireFormat, Arc<dyn SessionProviderFace>)> {
self
.session_providers
.read()
.iter()
.map(|(k, v)| (*k, v.clone()))
.collect()
}
pub fn find_session_provider(
&self,
wire_format: WireFormat,
) -> Option<Arc<dyn SessionProviderFace>> {
self.session_providers.read().get(&wire_format).cloned()
}
pub fn add_session(
self: &Arc<Self>,
wire_format: WireFormat,
backend_provider: &dyn SessionProviderFace,
network_sender_id: u64,
) -> Option<Arc<dyn MessageConsumerFace>> {
let session = backend_provider.get_session(wire_format, network_sender_id)?;
session.attach_server(self.clone());
Some(session)
}
pub fn dispose_active_handlers(&self) {
log::trace!("Begin disposing active handlers");
while self.active_handler_count.load(Acquire) >= 0 {
while self.active_handler_count.load(Acquire) > 0 {
for handler_id in self.active_handler_ids() {
if let Some(consumer) = self.active_handlers.write().remove(&handler_id) {
consumer.dispose();
self.active_handler_count.fetch_sub(1, AcqRel);
self.increment_connections_disposed();
}
}
yield_now();
}
if self
.active_handler_count
.compare_exchange(0, DISPOSED_HANDLER_COUNT, AcqRel, Acquire)
.is_ok()
{
break;
}
}
log::trace!("End disposing active handlers");
}
pub fn dispose(&self) {
self.disposed.store(true, Release);
self.dispose_active_handlers();
self.session_providers.write().clear();
}
pub fn dispose_message_consumer(&self, handler_id: u64) -> bool {
let consumer = self.active_handlers.write().remove(&handler_id);
if let Some(consumer) = consumer {
self.active_handler_count.fetch_sub(1, AcqRel);
self.increment_connections_disposed();
consumer.dispose();
true
} else {
false
}
}
}
impl ServerEnumerate for GarnetServerBase {
fn active_consumers(&self) -> Vec<Arc<dyn MessageConsumerFace>> {
self.active_handlers.read().values().cloned().collect()
}
fn active_cluster_sessions(&self) -> Vec<Arc<dyn ClusterSessionFace>> {
self
.active_handlers
.read()
.values()
.filter_map(|consumer| consumer.cluster_session())
.collect()
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
struct TestConsumer {
disposed: Mutex<usize>,
cluster: Mutex<Option<Arc<dyn ClusterSessionFace>>>,
}
impl TestConsumer {
fn new() -> Arc<Self> {
Arc::new(Self {
disposed: Mutex::new(0),
cluster: Mutex::new(None),
})
}
}
impl MessageConsumerFace for TestConsumer {
fn dispose(&self) {
*self.disposed.lock().expect("无锁中毒") += 1;
}
fn cluster_session(&self) -> Option<Arc<dyn ClusterSessionFace>> {
self.cluster.lock().expect("无锁中毒").clone()
}
fn attach_server(&self, _server: Arc<dyn ServerEnumerate>) {}
}
struct TestClusterSession;
impl ClusterSessionFace for TestClusterSession {
fn session_id(&self) -> i64 {
7
}
}
#[test]
fn connection_counters_track_receive_and_dispose() {
let base = GarnetServerBase::new("127.0.0.1:6379", 0);
assert_eq!(base.network_buffer_size(), DEFAULT_NETWORK_BUFFER_SIZE);
assert_eq!(base.endpoint(), "127.0.0.1:6379");
let consumer = TestConsumer::new();
let handler_id = base.register_handler(consumer).expect("未释放可注册");
base.increment_connections_received();
assert_eq!(base.get_conn_active(), 1);
assert_eq!(base.total_connections_received(), 1);
base.reset_connections_received();
assert_eq!(base.total_connections_received(), 1);
assert!(base.dispose_message_consumer(handler_id));
assert_eq!(base.total_connections_disposed(), 1);
assert_eq!(base.get_conn_active(), 0);
base.reset_connections_disposed();
assert_eq!(base.total_connections_disposed(), 0);
}
#[test]
fn register_rejects_duplicate_wire_format() {
let base = GarnetServerBase::new("ep", 4096);
struct Provider;
impl SessionProviderFace for Provider {
fn get_session(
&self,
_wire_format: WireFormat,
_network_sender_id: u64,
) -> Option<Arc<dyn MessageConsumerFace>> {
None
}
}
base
.register(WireFormat::Ascii, Arc::new(Provider))
.expect("首次注册成功");
assert!(matches!(
base.register(WireFormat::Ascii, Arc::new(Provider)),
Err(ServerError::WireFormatAlreadyRegistered(WireFormat::Ascii))
));
assert_eq!(base.get_session_providers().len(), 1);
assert!(base.unregister(WireFormat::Ascii).is_some());
assert!(base.unregister(WireFormat::Ascii).is_none());
}
#[test]
fn add_session_attaches_server_backref() {
let base = Arc::new(GarnetServerBase::new("ep", 0));
struct Provider;
impl SessionProviderFace for Provider {
fn get_session(
&self,
_wire_format: WireFormat,
_network_sender_id: u64,
) -> Option<Arc<dyn MessageConsumerFace>> {
Some(TestConsumer::new())
}
}
let session = base
.add_session(WireFormat::Ascii, &Provider, 42)
.expect("会话可创建");
assert_eq!(base.active_consumers().len(), 0);
let _ = base.register_handler(session.clone()).expect("登记成功");
assert_eq!(base.active_consumers().len(), 1);
drop(session);
}
#[test]
fn active_cluster_sessions_enumerates_consumers() {
let base = GarnetServerBase::new("ep", 0);
let consumer = TestConsumer::new();
*consumer.cluster.lock().expect("无锁中毒") = Some(Arc::new(TestClusterSession));
let _ = base.register_handler(consumer);
let sessions = base.active_cluster_sessions();
assert_eq!(sessions.len(), 1);
assert_eq!(sessions[0].session_id(), 7);
}
#[test]
fn dispose_drains_handlers_and_blocks_new_ones() {
let base = GarnetServerBase::new("ep", 0);
let consumer = TestConsumer::new();
let _ = base.register_handler(consumer.clone());
assert_eq!(base.get_conn_active(), 1);
base.dispose();
assert!(base.disposed());
assert_eq!(*consumer.disposed.lock().expect("无锁中毒"), 1);
assert_eq!(base.get_conn_active(), 0);
assert!(base.register_handler(TestConsumer::new()).is_none());
}
}