use std::{collections::HashMap, sync::Arc};
use atomic_counter::{AtomicCounter, ConsistentCounter};
use dportable::time::Timeout;
use futures::{
channel::{
mpsc::{unbounded, TrySendError, UnboundedReceiver, UnboundedSender},
oneshot,
},
Sink, SinkExt,
};
use crate::{
consumer::{self, Message, Payload, RequestId},
producer::{self, StreamResponse},
ShutdownType,
};
pub trait State<Response> {
fn handle_message(&mut self, message: producer::Message<Response>) -> Result<(), ShutdownType>;
fn shutdown(&mut self, shutdown_type: ShutdownType);
fn idle(&self) -> bool;
}
#[derive(Debug, Default)]
pub struct RequestIdGenerator {
counter: ConsistentCounter,
}
impl RequestIdGenerator {
pub fn new() -> Self {
RequestIdGenerator {
counter: ConsistentCounter::new(0),
}
}
pub fn id(&self) -> RequestId {
self.counter.inc() as RequestId
}
}
pub type ValueResultSender<T> = oneshot::Sender<consumer::Result<T>>;
pub type StreamResultSender = ValueResultSender<()>;
pub type StreamValuesSender<T> = UnboundedSender<T>;
#[derive(Debug)]
pub enum ResultSender<T> {
Value(ValueResultSender<T>),
Stream {
result_sender: StreamResultSender,
values_sender: StreamValuesSender<T>,
},
Abort,
}
#[derive(Debug)]
pub struct FullRequest<Request, T> {
pub message: Message<Request>,
pub result_sender: ResultSender<T>,
}
#[derive(Debug)]
pub struct RequestSender<Request, T> {
sender: Arc<UnboundedSender<FullRequest<Request, T>>>,
}
impl<Request, T> Clone for RequestSender<Request, T> {
fn clone(&self) -> Self {
Self {
sender: self.sender.clone(),
}
}
}
impl<Request, T> RequestSender<Request, T> {
pub fn pair() -> (
RequestSender<Request, T>,
UnboundedReceiver<FullRequest<Request, T>>,
) {
let (sender, receiver) = unbounded();
let sender = Arc::new(sender);
let sender = RequestSender { sender };
(sender, receiver)
}
pub fn abort(&self, id: RequestId) {
let payload = Payload::Abort;
let message = Message { id, payload };
let full = FullRequest {
message,
result_sender: ResultSender::Abort,
};
let _ = self.sender.unbounded_send(full);
}
pub fn send(
&self,
id: RequestId,
request: Request,
result_sender: ResultSender<T>,
) -> Result<(), TrySendError<FullRequest<Request, T>>> {
let message = Message {
id,
payload: Payload::Request(request),
};
let full = FullRequest {
message,
result_sender,
};
self.sender.unbounded_send(full)
}
}
pub fn handle_response<T>(
result: T,
request_id: RequestId,
requests: &mut HashMap<RequestId, oneshot::Sender<consumer::Result<T>>>,
pending: &mut usize,
) {
if let Some(sender) = requests.remove(&request_id) {
let _ = sender.send(Ok(result));
*pending -= 1;
}
}
pub fn handle_stream_response<T>(
result: StreamResponse<T>,
request_id: RequestId,
requests: &mut HashMap<RequestId, (Option<StreamResultSender>, StreamValuesSender<T>)>,
pending: &mut usize,
) {
match result {
StreamResponse::Open => {
if let Some((sender, _)) = requests.get_mut(&request_id) {
if let Some(sender) = sender.take() {
let _ = sender.send(Ok(()));
*pending -= 1;
}
}
}
StreamResponse::Item(item) => {
if let Some((_, sender)) = requests.get_mut(&request_id) {
let _ = sender.unbounded_send(item);
}
}
StreamResponse::Closed => {
requests.remove(&request_id);
}
}
}
pub fn handle_producer_message<Response, Error, S>(
message_option: Option<Result<producer::Message<Response>, Error>>,
has_timeout: bool,
timeout_future: &mut Timeout,
state: &mut S,
should_break: &mut bool,
) where
S: State<Response>,
{
if let Some(result) = message_option {
if has_timeout {
timeout_future.reset();
}
match result {
Ok(message) => {
if let Err(shutdown_type) = state.handle_message(message) {
state.shutdown(shutdown_type);
*should_break = true;
}
}
Err(_) => {
*should_break = true;
}
}
} else {
state.shutdown(ShutdownType::Closed);
*should_break = true;
}
}
pub async fn handle_new_request<Request, T, S>(
request_option: Option<FullRequest<Request, T>>,
has_timeout: bool,
timeout_future: &mut Timeout,
requests: &mut HashMap<RequestId, oneshot::Sender<consumer::Result<T>>>,
sender: &mut S,
pending: &mut usize,
) where
S: Sink<consumer::Message<Request>> + Unpin,
{
if let Some(FullRequest {
message,
result_sender,
}) = request_option
{
if has_timeout {
timeout_future.reset();
}
let id = message.id;
let result = sender.send(message).await;
let result = result.map_err(|_| crate::Error::Closed);
match result_sender {
ResultSender::Value(sender) => {
if let Err(error) = result {
let _ = sender.send(Err(error));
} else {
requests.insert(id, sender);
*pending += 1;
}
}
ResultSender::Abort => {
if requests.remove(&id).is_some() {
*pending -= 1;
}
}
_ => unreachable!("stream sender got when value sender expected"),
}
}
}
pub async fn handle_new_no_ack_request<Request, S>(
request_option: Option<FullRequest<Request, ()>>,
has_timeout: bool,
timeout_future: &mut Timeout,
sender: &mut S,
) where
S: Sink<consumer::Message<Request>> + Unpin,
{
if let Some(FullRequest {
message,
result_sender,
}) = request_option
{
if has_timeout {
timeout_future.reset();
}
let result = sender.send(message).await;
let result = result.map_err(|_| crate::Error::Closed);
match result_sender {
ResultSender::Value(sender) => {
if let Err(error) = result {
let _ = sender.send(Err(error));
} else {
let _ = sender.send(Ok(()));
}
}
ResultSender::Abort => {}
_ => unreachable!("stream sender got when value sender expected"),
}
}
}
pub async fn handle_new_stream_request<Request, T, S>(
request_option: Option<FullRequest<Request, T>>,
has_timeout: bool,
timeout_future: &mut Timeout,
requests: &mut HashMap<RequestId, (Option<StreamResultSender>, StreamValuesSender<T>)>,
sender: &mut S,
pending: &mut usize,
) where
S: Sink<consumer::Message<Request>> + Unpin,
{
if let Some(FullRequest {
message,
result_sender,
}) = request_option
{
if has_timeout {
timeout_future.reset();
}
let id = message.id;
let result = sender.send(message).await;
let result = result.map_err(|_| crate::Error::Closed);
match result_sender {
ResultSender::Stream {
result_sender,
values_sender,
} => {
if let Err(error) = result {
let _ = result_sender.send(Err(error));
} else {
requests.insert(id, (Some(result_sender), values_sender));
*pending += 1;
}
}
ResultSender::Abort => {
if let Some((sender, _)) = requests.remove(&id) {
if sender.is_some() {
*pending -= 1;
}
}
}
_ => unreachable!("value sender got when stream sender expected"),
}
}
}