use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use dashmap::DashMap;
use tokio::sync::{mpsc, Mutex as AsyncMutex};
use tokio_stream::wrappers::ReceiverStream;
use tokio_stream::StreamExt;
use tonic::service::interceptor::InterceptedService;
use tonic::transport::Channel;
use tonic::Streaming;
use tracing::{debug, instrument, warn};
use arc_swap::ArcSwap;
use crate::auth::{ChannelAuthenticator, ChannelIdInterceptor, SaslStreamGuard};
use crate::config::GoosefsConfig;
use crate::error::{Error, Result};
use crate::proto::grpc::block::{
block_worker_client::BlockWorkerClient, write_request, OpenLocalBlockRequest,
OpenLocalBlockResponse, ReadRequest, ReadResponse, RequestType, WriteRequest,
WriteRequestCommand, WriteResponse,
};
use crate::proto::proto::dataserver::{CreateUfsFileOptions, OpenUfsBlockOptions};
use crate::proto::proto::security::Capability;
#[derive(Clone, Debug)]
pub struct WriteBlockOptions {
pub request_type: RequestType,
pub create_ufs_file_options: Option<CreateUfsFileOptions>,
pub async_write: bool,
}
impl Default for WriteBlockOptions {
fn default() -> Self {
Self {
request_type: RequestType::GoosefsBlock,
create_ufs_file_options: None,
async_write: false,
}
}
}
pub struct WriteBlockHandle {
block_id: i64,
pub request_tx: Option<mpsc::Sender<WriteRequest>>,
response_rx: mpsc::Receiver<std::result::Result<WriteResponse, tonic::Status>>,
task_handle: Option<tokio::task::JoinHandle<()>>,
}
impl WriteBlockHandle {
pub async fn recv_response(&mut self) -> Result<Option<WriteResponse>> {
match self.response_rx.recv().await {
Some(Ok(resp)) => Ok(Some(resp)),
Some(Err(status)) => Err(Error::GrpcError {
message: format!(
"WriteBlock server error for block_id={}: {}",
self.block_id, status
),
source: Box::new(status),
}),
None => Ok(None),
}
}
pub async fn close(mut self) -> Result<()> {
drop(self.request_tx.take());
debug!(
block_id = self.block_id,
"closed write stream, waiting for server finalize"
);
let mut last_err: Option<Error> = None;
while let Some(result) = self.response_rx.recv().await {
match result {
Ok(_resp) => {
debug!(
block_id = self.block_id,
"received final response from server"
);
}
Err(status) => {
last_err = Some(Error::GrpcError {
message: format!(
"WriteBlock server error for block_id={}: {}",
self.block_id, status
),
source: Box::new(status),
});
break;
}
}
}
if let Some(handle) = self.task_handle.take() {
match tokio::time::timeout(std::time::Duration::from_secs(5), handle).await {
Ok(Ok(())) => {}
Ok(Err(join_err)) => {
if join_err.is_panic() {
warn!(
block_id = self.block_id,
"WriteBlock background task panicked"
);
}
if last_err.is_none() {
last_err = Some(Error::Internal {
message: format!(
"WriteBlock background task ended abnormally for block_id={}: {}",
self.block_id, join_err
),
source: None,
});
}
}
Err(_) => {
warn!(
block_id = self.block_id,
"WriteBlock background task did not finish within 5s after stream drain; aborting"
);
}
}
}
if let Some(e) = last_err {
return Err(e);
}
Ok(())
}
pub async fn cancel(mut self) {
if let Some(handle) = self.task_handle.take() {
handle.abort();
}
debug!(block_id = self.block_id, "cancelled write stream");
}
}
impl Drop for WriteBlockHandle {
fn drop(&mut self) {
if let Some(handle) = self.task_handle.take() {
debug!(
block_id = self.block_id,
"WriteBlockHandle dropped without close()/cancel(); aborting background task"
);
handle.abort();
}
}
}
pub struct OpenLocalBlockGuard {
block_id: i64,
_request_tx: Option<mpsc::Sender<OpenLocalBlockRequest>>,
_response_stream: std::sync::Mutex<Streaming<OpenLocalBlockResponse>>,
}
impl Drop for OpenLocalBlockGuard {
fn drop(&mut self) {
debug!(
block_id = self.block_id,
"OpenLocalBlockGuard dropped — closing bidi stream, Worker will unlock block"
);
let _ = self._request_tx.take();
}
}
type AuthenticatedBlockWorkerClient =
BlockWorkerClient<InterceptedService<Channel, ChannelIdInterceptor>>;
#[derive(Clone)]
pub struct WorkerClient {
inner: AuthenticatedBlockWorkerClient,
addr: String,
generation: u64,
_sasl_guard: std::sync::Arc<Option<SaslStreamGuard>>,
}
impl WorkerClient {
pub async fn connect(addr: &str, config: &GoosefsConfig) -> Result<Self> {
let endpoint = Channel::from_shared(format!("http://{}", addr))
.map_err(|e| Error::ConfigError {
message: format!("invalid worker endpoint: {}", e),
})?
.connect_timeout(config.connect_timeout)
.timeout(config.request_timeout);
let channel = endpoint.connect().await?;
let authenticator =
ChannelAuthenticator::new(config.auth_type, config.auth_username.clone(), None)
.with_auth_timeout(config.auth_timeout);
let mut auth_channel = authenticator.authenticate(channel).await?;
let sasl_guard = auth_channel.take_sasl_guard();
debug!(addr = %addr, auth_type = %config.auth_type, "connected to Goosefs Worker");
Ok(Self {
inner: BlockWorkerClient::new(auth_channel.channel),
addr: addr.to_string(),
generation: 0,
_sasl_guard: std::sync::Arc::new(sasl_guard),
})
}
pub async fn connect_simple(addr: &str, connect_timeout: Duration) -> Result<Self> {
let endpoint = Channel::from_shared(format!("http://{}", addr))
.map_err(|e| Error::ConfigError {
message: format!("invalid worker endpoint: {}", e),
})?
.connect_timeout(connect_timeout);
let channel = endpoint.connect().await?;
let interceptor = ChannelIdInterceptor::new(uuid::Uuid::new_v4().to_string());
let intercepted = InterceptedService::new(channel, interceptor);
debug!(addr = %addr, "connected to Goosefs Worker (no auth)");
Ok(Self {
inner: BlockWorkerClient::new(intercepted),
addr: addr.to_string(),
generation: 0,
_sasl_guard: std::sync::Arc::new(None),
})
}
pub fn from_channel(channel: Channel, addr: String) -> Self {
let interceptor = ChannelIdInterceptor::new("test-no-auth".to_string());
let intercepted = InterceptedService::new(channel, interceptor);
Self {
inner: BlockWorkerClient::new(intercepted),
addr,
generation: 0,
_sasl_guard: std::sync::Arc::new(None),
}
}
#[instrument(skip(self, open_ufs_block_options), fields(block_id = %block_id, offset = %offset, length = %length))]
pub async fn read_block(
&self,
block_id: i64,
offset: i64,
length: i64,
chunk_size: i64,
prefetch_window: Option<i32>,
open_ufs_block_options: Option<OpenUfsBlockOptions>,
) -> Result<(mpsc::Sender<ReadRequest>, Streaming<ReadResponse>)> {
let (tx, rx) = mpsc::channel::<ReadRequest>(32);
let initial_request = ReadRequest {
block_id: Some(block_id),
offset: Some(offset),
length: Some(length),
chunk_size: Some(chunk_size),
open_ufs_block_options,
offset_received: None,
position_short: None,
request_id: None,
capability: None,
block_size: None,
prefetch_window,
};
tx.send(initial_request)
.await
.map_err(|_| Error::BlockIoError {
message: "failed to send initial ReadRequest".to_string(),
})?;
let stream = ReceiverStream::new(rx);
let response = self.inner.clone().read_block(stream).await?;
Ok((tx, response.into_inner()))
}
pub async fn read_block_positioned(
&self,
block_id: i64,
offset: i64,
length: i64,
chunk_size: i64,
open_ufs_block_options: Option<OpenUfsBlockOptions>,
) -> Result<(mpsc::Sender<ReadRequest>, Streaming<ReadResponse>)> {
let (tx, rx) = mpsc::channel::<ReadRequest>(32);
let initial_request = ReadRequest {
block_id: Some(block_id),
offset: Some(offset),
length: Some(length),
chunk_size: Some(chunk_size),
open_ufs_block_options,
offset_received: None,
position_short: Some(true), request_id: None,
capability: None,
block_size: None,
prefetch_window: None,
};
tx.send(initial_request)
.await
.map_err(|_| Error::BlockIoError {
message: "failed to send initial positioned ReadRequest".to_string(),
})?;
let stream = ReceiverStream::new(rx);
let response = self.inner.clone().read_block(stream).await?;
Ok((tx, response.into_inner()))
}
#[instrument(skip(self, capability), fields(block_id = %block_id, block_size = %block_size))]
pub async fn open_local_block(
&self,
block_id: i64,
block_size: i64,
capability: Option<Capability>,
) -> Result<(OpenLocalBlockResponse, OpenLocalBlockGuard)> {
let (tx, rx) = mpsc::channel::<OpenLocalBlockRequest>(4);
let request = OpenLocalBlockRequest {
block_id: Some(block_id),
capability,
block_size: Some(block_size),
};
tx.send(request).await.map_err(|_| Error::BlockIoError {
message: format!("failed to send OpenLocalBlockRequest for block_id={block_id}"),
})?;
let stream = ReceiverStream::new(rx);
let response = self.inner.clone().open_local_block(stream).await?;
let mut response_stream = response.into_inner();
let first = response_stream
.message()
.await?
.ok_or_else(|| Error::BlockIoError {
message: format!(
"OpenLocalBlock stream for block_id={block_id} closed before any response"
),
})?;
debug!(
block_id = block_id,
path = first.path.as_deref().unwrap_or(""),
block_size = ?first.block_size,
"OpenLocalBlock granted local path"
);
let guard = OpenLocalBlockGuard {
block_id,
_request_tx: Some(tx),
_response_stream: std::sync::Mutex::new(response_stream),
};
Ok((first, guard))
}
#[instrument(skip(self, options), fields(block_id = %block_id))]
pub async fn write_block(
&self,
block_id: i64,
space_to_reserve: i64,
options: WriteBlockOptions,
) -> Result<WriteBlockHandle> {
let (tx, rx) = mpsc::channel::<WriteRequest>(32);
let initial_command = WriteRequest {
value: Some(write_request::Value::Command(WriteRequestCommand {
r#type: Some(options.request_type as i32),
id: Some(block_id),
offset: Some(0),
flush: None,
create_ufs_file_options: options.create_ufs_file_options,
space_to_reserve: Some(space_to_reserve),
capability: None,
medium_type: None,
async_write: Some(options.async_write),
})),
};
let initial_stream = tokio_stream::once(initial_command);
let subsequent_stream = ReceiverStream::new(rx);
let combined_stream = initial_stream.chain(subsequent_stream);
let (resp_tx, resp_rx) =
mpsc::channel::<std::result::Result<WriteResponse, tonic::Status>>(8);
let mut client = self.inner.clone();
let addr = self.addr.clone();
let task_handle = tokio::spawn(async move {
debug!(block_id = block_id, addr = %addr, "WriteBlock gRPC task started");
let call_result = client.write_block(combined_stream).await;
match call_result {
Ok(response) => {
let mut stream = response.into_inner();
loop {
match stream.message().await {
Ok(Some(msg)) => {
if resp_tx.send(Ok(msg)).await.is_err() {
debug!(block_id = block_id, "response receiver dropped");
break;
}
}
Ok(None) => {
debug!(block_id = block_id, "server closed response stream");
break;
}
Err(status) => {
warn!(block_id = block_id, %status, "server response error");
let _ = resp_tx.send(Err(status)).await;
break;
}
}
}
}
Err(status) => {
warn!(block_id = block_id, %status, "WriteBlock RPC failed");
let _ = resp_tx.send(Err(status)).await;
}
}
debug!(block_id = block_id, "WriteBlock gRPC task finished");
});
debug!(block_id = block_id, "WriteBlock handle created");
Ok(WriteBlockHandle {
block_id,
request_tx: Some(tx),
response_rx: resp_rx,
task_handle: Some(task_handle),
})
}
pub fn addr(&self) -> &str {
&self.addr
}
pub fn generation(&self) -> u64 {
self.generation
}
}
pub struct WorkerClientPool {
clients: ArcSwap<HashMap<String, WorkerClient>>,
reconnect_locks: DashMap<String, Arc<AsyncMutex<()>>>,
addr_rr: DashMap<String, Arc<AtomicU64>>,
pool_size: usize,
next_generation: AtomicU64,
config: GoosefsConfig,
}
impl WorkerClientPool {
pub fn new(config: GoosefsConfig) -> Self {
let pool_size = config.worker_connection_pool_size.max(1);
Self {
clients: ArcSwap::from_pointee(HashMap::new()),
reconnect_locks: DashMap::new(),
addr_rr: DashMap::new(),
pool_size,
next_generation: AtomicU64::new(1),
config,
}
}
pub fn pool_size(&self) -> usize {
self.pool_size
}
fn slot_key(&self, addr: &str, slot: usize) -> String {
if self.pool_size <= 1 {
return addr.to_string();
}
let mut s = String::with_capacity(addr.len() + 21);
s.push_str(addr);
s.push('#');
let mut buf = itoa::Buffer::new();
s.push_str(buf.format(slot));
s
}
async fn next_slot(&self, addr: &str) -> usize {
if self.pool_size <= 1 {
return 0;
}
if let Some(c) = self.addr_rr.get(addr) {
return (c.fetch_add(1, Ordering::Relaxed) % self.pool_size as u64) as usize;
}
let c = self
.addr_rr
.entry(addr.to_string())
.or_insert_with(|| Arc::new(AtomicU64::new(0)));
(c.fetch_add(1, Ordering::Relaxed) % self.pool_size as u64) as usize
}
pub async fn acquire(&self, addr: &str) -> Result<WorkerClient> {
let slot = self.next_slot(addr).await;
let key = self.slot_key(addr, slot);
self.acquire_by_key(&key, addr).await
}
async fn acquire_by_key(&self, key: &str, connect_addr: &str) -> Result<WorkerClient> {
if let Some(client) = self.clients.load().get(key).cloned() {
debug!(key = %key, generation = client.generation, "reusing cached WorkerClient");
return Ok(client);
}
let lock = self.reconnect_lock_for(key).await;
let _guard = lock.lock().await;
if let Some(client) = self.clients.load().get(key).cloned() {
return Ok(client);
}
debug!(key = %key, addr = %connect_addr, "creating new WorkerClient for pool");
let mut client = WorkerClient::connect(connect_addr, &self.config).await?;
client.generation = self.next_generation.fetch_add(1, Ordering::Relaxed);
self.insert_client(key, client.clone());
Ok(client)
}
fn insert_client(&self, key: &str, client: WorkerClient) {
self.clients.rcu(|cur| {
let mut next = (**cur).clone();
next.insert(key.to_string(), client.clone());
next
});
}
fn remove_client(&self, key: &str) {
self.clients.rcu(|cur| {
let mut next = (**cur).clone();
next.remove(key);
next
});
}
pub async fn invalidate(&self, addr: &str) {
let keys: Vec<String> = (0..self.pool_size)
.map(|s| self.slot_key(addr, s))
.collect();
self.clients.rcu(|cur| {
let mut next = (**cur).clone();
for k in &keys {
if next.remove(k).is_some() {
debug!(key = %k, "invalidated WorkerClient from pool");
}
}
next
});
for k in &keys {
if self.reconnect_locks.remove(k).is_some() {
debug!(key = %k, "removed reconnect lock for invalidated worker");
}
}
if self.pool_size > 1 {
self.addr_rr.remove(addr);
}
}
async fn reconnect_lock_for(&self, addr: &str) -> Arc<AsyncMutex<()>> {
self.reconnect_locks
.entry(addr.to_string())
.or_insert_with(|| Arc::new(AsyncMutex::new(())))
.clone()
}
pub async fn reconnect_if_stale(
&self,
addr: &str,
stale_generation: u64,
) -> Result<WorkerClient> {
if self.pool_size <= 1 {
return self.reconnect_by_key(addr, addr, stale_generation).await;
}
let mut target_key: Option<String> = None;
{
let cache = self.clients.load();
for s in 0..self.pool_size {
let k = self.slot_key(addr, s);
if let Some(c) = cache.get(&k) {
if c.generation == stale_generation {
target_key = Some(k);
break;
}
}
}
}
match target_key {
Some(key) => self.reconnect_by_key(&key, addr, stale_generation).await,
None => {
debug!(
addr = %addr,
observed = stale_generation,
"reconnect coalesced — no slot matches stale generation (already refreshed)"
);
crate::metrics::counter(crate::metrics::name::CLIENT_WORKER_RECONNECTS_COALESCED)
.inc(1);
self.acquire(addr).await
}
}
}
async fn reconnect_by_key(
&self,
map_key: &str,
connect_addr: &str,
stale_generation: u64,
) -> Result<WorkerClient> {
let lock = self.reconnect_lock_for(map_key).await;
let _guard = lock.lock().await;
{
let cache = self.clients.load();
if let Some(client) = cache.get(map_key) {
if client.generation > stale_generation {
debug!(
key = %map_key,
observed = stale_generation,
current = client.generation,
"reconnect coalesced — another task already refreshed this channel"
);
crate::metrics::counter(
crate::metrics::name::CLIENT_WORKER_RECONNECTS_COALESCED,
)
.inc(1);
return Ok(client.clone());
}
}
}
debug!(
key = %map_key,
stale_generation = stale_generation,
"performing single-flight reconnect"
);
crate::metrics::counter(crate::metrics::name::CLIENT_WORKER_RECONNECTS_TOTAL).inc(1);
self.remove_client(map_key);
let mut fresh = WorkerClient::connect(connect_addr, &self.config).await?;
fresh.generation = self.next_generation.fetch_add(1, Ordering::Relaxed);
self.insert_client(map_key, fresh.clone());
debug!(
key = %map_key,
new_generation = fresh.generation,
"single-flight reconnect installed fresh WorkerClient"
);
Ok(fresh)
}
pub async fn reconnect(&self, addr: &str) -> Result<WorkerClient> {
let slot = self.next_slot(addr).await;
let key = self.slot_key(addr, slot);
self.reconnect_by_key(&key, addr, u64::MAX).await
}
pub fn new_shared(config: GoosefsConfig) -> Arc<Self> {
Arc::new(Self::new(config))
}
#[cfg(test)]
async fn test_install(&self, addr: &str, mut client: WorkerClient) -> Option<WorkerClient> {
client.generation = self.next_generation.fetch_add(1, Ordering::Relaxed);
let prev = self.clients.load().get(addr).cloned();
self.insert_client(addr, client);
prev
}
#[cfg(test)]
async fn test_current_generation(&self, addr: &str) -> Option<u64> {
self.clients.load().get(addr).map(|c| c.generation)
}
#[cfg(test)]
async fn test_reconnect_locks_len(&self) -> usize {
self.reconnect_locks.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use tonic::transport::Channel;
fn fake_client(addr: &str) -> WorkerClient {
let channel = Channel::from_static("http://127.0.0.1:1").connect_lazy();
WorkerClient::from_channel(channel, addr.to_string())
}
#[tokio::test]
async fn worker_pool_round_robins_slots() {
let config = GoosefsConfig::new("127.0.0.1:9200").with_worker_connection_pool_size(4);
let pool = WorkerClientPool::new(config);
assert_eq!(pool.pool_size(), 4);
assert_eq!(pool.slot_key("h:1", 0), "h:1#0");
assert_eq!(pool.slot_key("h:1", 3), "h:1#3");
let mut seen = Vec::new();
for _ in 0..8 {
seen.push(pool.next_slot("h:1").await);
}
assert_eq!(seen, vec![0, 1, 2, 3, 0, 1, 2, 3]);
assert_eq!(pool.next_slot("h:2").await, 0);
}
#[tokio::test]
async fn worker_pool_single_channel_keys_by_addr() {
let pool = WorkerClientPool::new(
GoosefsConfig::new("127.0.0.1:9200").with_worker_connection_pool_size(1),
);
assert_eq!(pool.pool_size(), 1);
assert_eq!(pool.slot_key("h:1", 0), "h:1");
assert_eq!(pool.next_slot("h:1").await, 0);
assert_eq!(pool.next_slot("h:1").await, 0);
}
#[tokio::test]
async fn test_reconnect_if_stale_coalesces_when_generation_advanced() {
let pool = WorkerClientPool::new(
GoosefsConfig::new("127.0.0.1:9200").with_worker_connection_pool_size(1),
);
let addr = "test-worker:9203";
pool.test_install(addr, fake_client(addr)).await;
let gen_before = pool.test_current_generation(addr).await.unwrap();
pool.test_install(addr, fake_client(addr)).await;
let gen_after = pool.test_current_generation(addr).await.unwrap();
assert!(gen_after > gen_before);
let result = pool.reconnect_if_stale(addr, gen_before).await;
assert!(
result.is_ok(),
"coalesced reconnect must short-circuit without network I/O, got {:?}",
result.err()
);
let returned = result.unwrap();
assert_eq!(
returned.generation(),
gen_after,
"caller must receive the already-replaced generation"
);
assert_eq!(
pool.test_current_generation(addr).await,
Some(gen_after),
"cached generation must not advance for a coalesced caller"
);
}
#[tokio::test]
async fn test_reconnect_locks_are_per_address() {
let pool = WorkerClientPool::new(GoosefsConfig::new("127.0.0.1:9200"));
let lock_a = pool.reconnect_lock_for("worker-a:9203").await;
let lock_b = pool.reconnect_lock_for("worker-b:9203").await;
let guard_a = lock_a.lock().await;
let guard_b = tokio::time::timeout(std::time::Duration::from_millis(50), lock_b.lock())
.await
.expect("lock for different address must not be blocked");
drop(guard_b);
drop(guard_a);
}
#[tokio::test]
async fn test_invalidate_clears_reconnect_lock_to_prevent_leak() {
let pool = WorkerClientPool::new(
GoosefsConfig::new("127.0.0.1:9200").with_worker_connection_pool_size(1),
);
for i in 0..10 {
let addr = format!("ephemeral-worker-{}:9203", i);
pool.test_install(&addr, fake_client(&addr)).await;
let _lock = pool.reconnect_lock_for(&addr).await;
}
assert_eq!(
pool.test_reconnect_locks_len().await,
10,
"reconnect_locks must be populated by reconnect_lock_for()"
);
for i in 0..10 {
pool.invalidate(&format!("ephemeral-worker-{}:9203", i))
.await;
}
assert_eq!(
pool.test_reconnect_locks_len().await,
0,
"invalidate() must remove the per-address reconnect lock so the \
map does not leak across worker churn"
);
}
#[tokio::test]
async fn test_generation_is_monotonic_across_installs() {
let pool = WorkerClientPool::new(GoosefsConfig::new("127.0.0.1:9200"));
let addr = "w:9203";
pool.test_install(addr, fake_client(addr)).await;
let g1 = pool.test_current_generation(addr).await.unwrap();
pool.test_install(addr, fake_client(addr)).await;
let g2 = pool.test_current_generation(addr).await.unwrap();
pool.test_install(addr, fake_client(addr)).await;
let g3 = pool.test_current_generation(addr).await.unwrap();
assert!(g1 < g2, "gen {} not less than {}", g1, g2);
assert!(g2 < g3, "gen {} not less than {}", g2, g3);
}
#[tokio::test]
async fn test_auth_retry_reconnect_if_stale_returns_fresh_after_rpc_failure() {
let pool = WorkerClientPool::new(
GoosefsConfig::new("127.0.0.1:9200").with_worker_connection_pool_size(1),
);
let addr = "test-worker:9203";
pool.test_install(addr, fake_client(addr)).await;
let stale_client = pool.acquire(addr).await.unwrap();
let stale_gen = stale_client.generation();
pool.test_install(addr, fake_client(addr)).await;
let expected_fresh_gen = pool.test_current_generation(addr).await.unwrap();
assert!(
expected_fresh_gen > stale_gen,
"fresh gen must exceed stale gen"
);
let fresh_client = pool
.reconnect_if_stale(addr, stale_gen)
.await
.expect("reconnect_if_stale must return Ok when generation advanced");
assert_eq!(
fresh_client.generation(),
expected_fresh_gen,
"must return the already-installed fresh client (coalesced reconnect)"
);
assert!(
fresh_client.generation() > stale_gen,
"fresh client generation ({}) must be > stale generation ({})",
fresh_client.generation(),
stale_gen
);
assert_eq!(
pool.test_current_generation(addr).await,
Some(expected_fresh_gen),
"pool generation must not advance for a coalesced reconnect"
);
}
#[tokio::test]
async fn test_auth_retry_unconditional_reconnect_never_short_circuits() {
let pool = WorkerClientPool::new(GoosefsConfig::new("127.0.0.1:9200"));
let addr = "test-worker:9203";
pool.test_install(addr, fake_client(addr)).await;
let current_gen = pool.test_current_generation(addr).await.unwrap();
let result = pool.reconnect_if_stale(addr, u64::MAX).await;
assert!(
result.is_err(),
"reconnect_if_stale(addr, u64::MAX) must NOT short-circuit \
when generation ({}) < u64::MAX — expected real connect attempt",
current_gen
);
}
#[tokio::test]
async fn test_auth_retry_multiple_observers_collapse_to_one_reconnect() {
let pool = WorkerClientPool::new(
GoosefsConfig::new("127.0.0.1:9200").with_worker_connection_pool_size(1),
);
let addr = "test-worker:9203";
pool.test_install(addr, fake_client(addr)).await;
let stale_gen = pool.test_current_generation(addr).await.unwrap();
pool.test_install(addr, fake_client(addr)).await;
let fresh_gen = pool.test_current_generation(addr).await.unwrap();
assert!(fresh_gen > stale_gen);
let client = pool
.reconnect_if_stale(addr, stale_gen)
.await
.expect("coalesced reconnect must succeed");
assert_eq!(
client.generation(),
fresh_gen,
"second observer must get the already-installed fresh client"
);
assert_eq!(
pool.test_current_generation(addr).await,
Some(fresh_gen),
"generation must not advance for a coalesced observer"
);
let client_old = pool
.reconnect_if_stale(addr, 0)
.await
.expect("observer with gen=0 must also get coalesced client");
assert_eq!(
client_old.generation(),
fresh_gen,
"observer with stale gen=0 must get the same fresh client"
);
}
#[tokio::test]
async fn write_block_handle_drop_aborts_background_task() {
let (tx, _rx) = mpsc::channel::<WriteRequest>(8);
let (_resp_tx, resp_rx) = mpsc::channel(8);
let task = tokio::spawn(async {
std::future::pending::<()>().await;
});
let abort_handle = task.abort_handle();
let handle = WriteBlockHandle {
block_id: 42,
request_tx: Some(tx),
response_rx: resp_rx,
task_handle: Some(task),
};
assert!(
!abort_handle.is_finished(),
"task should still be running before Drop"
);
drop(handle);
for _ in 0..50 {
if abort_handle.is_finished() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(
abort_handle.is_finished(),
"Drop did not abort background task — pre-fix regression"
);
}
#[tokio::test]
async fn write_block_handle_drop_after_close_is_noop() {
let (tx, _rx) = mpsc::channel::<WriteRequest>(8);
let (_, resp_rx) = mpsc::channel(8);
let task = tokio::spawn(async {});
let handle = WriteBlockHandle {
block_id: 7,
request_tx: Some(tx),
response_rx: resp_rx,
task_handle: Some(task),
};
let close_fut = handle.close();
let res = tokio::time::timeout(Duration::from_millis(500), close_fut).await;
assert!(
res.is_ok(),
"close() must complete promptly when response stream is closed"
);
assert!(
res.unwrap().is_ok(),
"close() should succeed on a graceful task"
);
}
}