use std::marker::Unpin;
use futures_core::future::BoxFuture;
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadHalf, WriteHalf};
use crate::command;
use crate::resp3::{de, from_msg, ser_cmd, token, write_cmd, Reader, Value};
#[derive(Debug)]
pub struct Connection<T> {
transport: T,
sender: SendCtx,
receiver: ReceiveCtx,
}
#[derive(Debug)]
pub struct ConnectionSendHalf<T> {
transport: WriteHalf<T>,
sender: SendCtx,
}
#[derive(Debug)]
pub struct ConnectionReceiveHalf<T> {
transport: ReadHalf<T>,
receiver: ReceiveCtx,
}
#[derive(Debug)]
struct SendCtx {
buf: Vec<u8>,
count: u64,
}
#[derive(Debug)]
struct ReceiveCtx {
reader: Reader,
count: u64,
}
#[derive(Debug, thiserror::Error)]
#[error(transparent)]
pub struct Error(#[from] pub Box<ErrorKind>);
#[derive(Debug, thiserror::Error)]
pub enum ErrorKind {
#[error("io error")]
Io(#[from] std::io::Error),
#[error("tokenize error")]
Tokenize(#[from] token::Error),
#[error("serialize error")]
Serialize(#[from] ser_cmd::Error),
#[error("deserialize error")]
Deserialize(#[from] de::Error),
#[error("invalid message order")]
InvalidMessageOrder,
}
impl<T: AsyncRead + AsyncWrite + Unpin> Connection<T> {
pub async fn new(transport: T) -> Result<(Self, Value), Error> {
Self::with_args(transport, None, None, None).await
}
pub async fn with_args(
transport: T,
auth: Option<(&str, &str)>,
setname: Option<&str>,
select: Option<u32>,
) -> Result<(Self, Value), Error> {
let mut chan = Connection {
transport,
sender: SendCtx {
buf: Vec::new(),
count: 0,
},
receiver: ReceiveCtx {
reader: Reader::new(),
count: 0,
},
};
let auth = auth.map(|(username, password)| ("AUTH", username, password));
let setname = setname.map(|clientname| ("SETNAME", clientname));
let resp = chan.raw_command(&("HELLO", 3, auth, setname)).await?;
if let Some(db) = select {
let serde::de::IgnoredAny = chan.raw_command(&("SELECT", db)).await?;
}
Ok((chan, resp))
}
pub async fn send<Req: Serialize>(&mut self, request: Req) -> Result<u64, Error> {
self.sender.send(&mut self.transport, request).await
}
pub async fn wait_response(&mut self) -> Result<Option<u64>, Error> {
self.receiver.wait_response(&mut self.transport).await
}
pub fn response(&self) -> Option<token::Message<'_>> {
self.receiver.response()
}
pub async fn raw_command<'de, Req: Serialize, Resp: Deserialize<'de>>(
&'de mut self,
request: Req,
) -> Result<Resp, Error> {
let req_cnt = self.send(request).await?;
loop {
match self.wait_response().await? {
None => continue,
Some(cnt) if cnt < req_cnt => continue,
Some(cnt) if cnt > req_cnt => return Err(ErrorKind::InvalidMessageOrder.into()),
Some(_) => break,
}
}
Ok(from_msg(self.receiver.response().unwrap())?)
}
pub fn split(self) -> (ConnectionSendHalf<T>, ConnectionReceiveHalf<T>) {
let (read, write) = tokio::io::split(self.transport);
(
ConnectionSendHalf {
transport: write,
sender: self.sender,
},
ConnectionReceiveHalf {
transport: read,
receiver: self.receiver,
},
)
}
}
impl<T: AsyncRead + AsyncWrite + Send + Unpin> command::RawCommandMut for Connection<T> {
fn raw_command<'de, Req, Resp>(
&'de mut self,
request: Req,
) -> BoxFuture<'de, Result<Resp, Error>>
where
Req: Serialize + Send + 'de,
Resp: Deserialize<'de> + 'de,
{
Box::pin(self.raw_command(request))
}
}
impl<T: AsyncWrite + Unpin> ConnectionSendHalf<T> {
pub async fn send<Req: Serialize>(&mut self, request: Req) -> Result<u64, Error> {
self.sender.send(&mut self.transport, request).await
}
pub fn is_pair_of(&self, other: &ConnectionReceiveHalf<T>) -> bool {
other.transport.is_pair_of(&self.transport)
}
pub fn unsplit(self, other: ConnectionReceiveHalf<T>) -> Connection<T> {
let transport = other.transport.unsplit(self.transport);
Connection {
transport,
sender: self.sender,
receiver: other.receiver,
}
}
}
impl<T: AsyncRead + Unpin> ConnectionReceiveHalf<T> {
pub async fn wait_response(&mut self) -> Result<Option<u64>, Error> {
self.receiver.wait_response(&mut self.transport).await
}
pub fn response(&self) -> Option<token::Message<'_>> {
self.receiver.response()
}
}
impl SendCtx {
async fn send<T, Req>(&mut self, transport: &mut T, request: Req) -> Result<u64, Error>
where
T: AsyncWrite + Unpin,
Req: Serialize,
{
let cmd = write_cmd(&mut self.buf, &request)?;
transport.write_all(cmd).await?;
self.count += 1;
Ok(self.count)
}
}
impl ReceiveCtx {
async fn wait_response<T>(&mut self, transport: &mut T) -> Result<Option<u64>, Error>
where
T: AsyncRead + Unpin,
{
while self.reader.message_not_ready()? {
transport.read_buf(self.reader.buf()).await?;
}
let msg = self.reader.message().unwrap();
let count = match msg.head() {
token::Token::Push(_) => None,
_ => {
self.count += 1;
Some(self.count)
}
};
Ok(count)
}
fn response(&self) -> Option<token::Message<'_>> {
self.reader.message()
}
}
impl From<ErrorKind> for Error {
fn from(err: ErrorKind) -> Self {
Self(Box::new(err))
}
}
impl From<std::io::Error> for Error {
fn from(err: std::io::Error) -> Self {
Self(Box::new(err.into()))
}
}
impl From<token::Error> for Error {
fn from(err: token::Error) -> Self {
Self(Box::new(err.into()))
}
}
impl From<ser_cmd::Error> for Error {
fn from(err: ser_cmd::Error) -> Self {
Self(Box::new(err.into()))
}
}
impl From<de::Error> for Error {
fn from(err: de::Error) -> Self {
Self(Box::new(err.into()))
}
}