use std::{cmp, error, fmt, io};
use std::future::Future;
use tokio::time::{timeout, timeout_at, Duration, Instant};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use super::payload::{Action, Payload, Timing};
use super::pdu;
use super::server::MAX_VERSION;
use super::state::State;
const IO_TIMEOUT: Duration = Duration::from_secs(10);
const INITIAL_VERSION: u8 = 2;
pub trait PayloadTarget {
type Update: PayloadUpdate;
fn start(&mut self, reset: bool) -> Self::Update;
fn apply(
&mut self, update: Self::Update, timing: Timing
) -> Result<(), PayloadError>;
}
pub trait PayloadUpdate {
fn push_update(
&mut self, action: Action, payload: Payload
) -> Result<(), PayloadError>;
}
impl PayloadUpdate for Vec<(Action, Payload)> {
fn push_update(
&mut self, action: Action, payload: Payload
) -> Result<(), PayloadError> {
self.push((action, payload));
Ok(())
}
}
pub struct Client<Sock, Target> {
sock: Sock,
target: Target,
state: Option<State>,
version: Option<u8>,
initial_version: u8,
timing: Timing,
next_update: Option<Instant>,
}
impl<Sock, Target> Client<Sock, Target> {
pub fn new(
sock: Sock,
target: Target,
state: Option<State>
) -> Self {
Self::with_initial_version(INITIAL_VERSION, sock, target, state)
}
pub fn with_initial_version(
initial_version: u8,
sock: Sock,
target: Target,
state: Option<State>
) -> Self {
Client {
sock, target, state,
version: None,
initial_version: cmp::min(initial_version, MAX_VERSION),
timing: Timing::default(),
next_update: None,
}
}
pub fn target(&self) -> &Target {
&self.target
}
pub fn target_mut(&mut self) -> &mut Target {
&mut self.target
}
pub fn into_target(self) -> Target {
self.target
}
pub fn state(&self) -> Option<State> {
self.state
}
fn version(&self) -> u8 {
self.version.unwrap_or(self.initial_version)
}
}
impl<Sock, Target> Client<Sock, Target>
where
Sock: AsyncRead + AsyncWrite + Unpin,
Target: PayloadTarget
{
pub async fn run(&mut self) -> Result<(), io::Error> {
loop {
if let Err(err) = self.step().await {
if err.kind() == io::ErrorKind::UnexpectedEof {
return Ok(())
}
else {
return Err(err)
}
}
}
}
pub async fn step(
&mut self
) -> Result<(), io::Error> {
let update = self.update().await?;
self.apply(update).await
}
pub async fn update(
&mut self
) -> Result<Target::Update, io::Error> {
if let Some(instant) = self.next_update.take() {
if let Ok(Err(err)) = timeout_at(
instant, pdu::SerialNotify::read(&mut self.sock)
).await {
return Err(err)
}
}
if let Some(state) = self.state {
if let Some(update) = self.serial(state).await? {
self.next_update = Some(
Instant::now() + self.timing.refresh_duration()
);
return Ok(update)
}
}
let res = self.reset().await;
self.next_update = Some(
Instant::now() + self.timing.refresh_duration()
);
res
}
async fn serial(
&mut self, state: State
) -> Result<Option<Target::Update>, io::Error> {
let start = loop {
pdu::SerialQuery::new(
self.version(), state,
).write(&mut self.sock).await?;
self.sock.flush().await?;
match self.try_io(FirstSerialReply::read).await? {
FirstSerialReply::Response(start) => break start,
FirstSerialReply::Reset => {
self.state = None;
return Ok(None)
}
FirstSerialReply::VersionError(version) => {
if self.version.is_some() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"version error after successful version \
negotiation"
));
}
if version >= INITIAL_VERSION {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"version error with larger version"
));
}
self.version = Some(version);
}
}
};
self.check_version(start.version())?;
let mut target = self.target.start(false);
loop {
match pdu::Payload::read(&mut self.sock).await? {
Ok(Some(pdu)) => {
self.check_version(pdu.version())?;
let (action, payload) = match pdu.to_payload() {
Ok(some) => some,
Err(err) => {
err.write(&mut self.sock).await?;
return Err(io::Error::other(""));
}
};
if let Err(err) = target.push_update(action, payload) {
err.send(
self.version(), Some(pdu), &mut self.sock
).await?;
return Err(io::Error::other(""));
}
}
Ok(None) => {
}
Err(end) => {
self.check_version(end.version())?;
self.state = Some(end.state());
if let Some(timing) = end.timing() {
self.timing = timing
}
break;
}
}
}
Ok(Some(target))
}
pub async fn reset(&mut self) -> Result<Target::Update, io::Error> {
let start = loop {
pdu::ResetQuery::new(
self.version()
).write(&mut self.sock).await?;
self.sock.flush().await?;
match self.try_io(FirstResetReply::read).await? {
FirstResetReply::Response(start) => break start,
FirstResetReply::VersionError(version) => {
if self.version.is_some() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"version error after successful version \
negotiation"
));
}
if version >= INITIAL_VERSION {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"version error with larger version"
));
}
self.version = Some(version);
}
}
};
self.check_version(start.version())?;
let mut target = self.target.start(true);
loop {
match pdu::Payload::read(&mut self.sock).await? {
Ok(Some(pdu)) => {
self.check_version(pdu.version())?;
let (action, payload) = match pdu.to_payload() {
Ok(some) => some,
Err(err) => {
err.write(&mut self.sock).await?;
return Err(io::Error::other(""))
}
};
if let Err(err) = target.push_update(action, payload) {
err.send(
self.version(), Some(pdu), &mut self.sock
).await?;
return Err(io::Error::other(""));
}
}
Ok(None) => {
}
Err(end) => {
self.check_version(end.version())?;
self.state = Some(end.state());
if let Some(timing) = end.timing() {
self.timing = timing
}
break;
}
}
}
Ok(target)
}
pub async fn apply(
&mut self, update: Target::Update
) -> Result<(), io::Error> {
if let Err(err) = self.target.apply(update, self.timing) {
self.send_error(err).await?;
Err(io::Error::other(""))
}
else {
Ok(())
}
}
pub async fn send_error(
&mut self, err: PayloadError
) -> Result<(), io::Error> {
err.send(self.version(), None, &mut self.sock).await
}
async fn try_io<'a, F, Fut, T>(
&'a mut self, op: F
) -> Result<T, io::Error>
where
F: FnOnce(&'a mut Sock) -> Fut,
Fut: Future<Output = Result<T, io::Error>> + 'a
{
match timeout(IO_TIMEOUT, op(&mut self.sock)).await {
Ok(res) => res,
Err(_) => {
Err(io::Error::new(
io::ErrorKind::TimedOut,
"server response timed out"
))
}
}
}
fn check_version(&mut self, version: u8) -> Result<(), io::Error> {
if let Some(stored_version) = self.version {
if version != stored_version {
Err(io::Error::new(
io::ErrorKind::InvalidData,
"version has changed"
))
}
else {
Ok(())
}
}
else if version > INITIAL_VERSION {
Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"server requested unsupported protocol version {version}"
)
))
}
else {
self.version = Some(version);
Ok(())
}
}
}
enum FirstSerialReply {
Response(pdu::CacheResponse),
Reset,
VersionError(u8)
}
impl FirstSerialReply {
async fn read<Sock: AsyncRead + Unpin>(
sock: &mut Sock
) -> Result<Self, io::Error> {
let header = pdu::Header::read(sock).await?;
match header.pdu() {
pdu::CacheResponse::PDU => {
pdu::CacheResponse::read_payload(
header, sock
).await.map(FirstSerialReply::Response)
}
pdu::CacheReset::PDU => {
pdu::CacheReset::read_payload(
header, sock
).await.map(|_| FirstSerialReply::Reset)
}
pdu::Error::PDU
if header.session()
== pdu::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION
=> {
pdu::Error::skip_payload(header, sock).await?;
Ok(Self::VersionError(header.version()))
}
pdu::Error::PDU => {
Err(io::Error::other(
format!("server reported error {}", header.session())
))
}
pdu => {
Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("unexpected PDU {pdu}")
))
}
}
}
}
enum FirstResetReply {
Response(pdu::CacheResponse),
VersionError(u8)
}
impl FirstResetReply {
async fn read<Sock: AsyncRead + Unpin>(
sock: &mut Sock
) -> Result<Self, io::Error> {
let header = pdu::Header::read(sock).await?;
match header.pdu() {
pdu::CacheResponse::PDU => {
pdu::CacheResponse::read_payload(
header, sock
).await.map(Self::Response)
}
pdu::Error::PDU
if header.session()
== pdu::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION
=> {
pdu::Error::skip_payload(header, sock).await?;
Ok(Self::VersionError(header.version()))
}
pdu::Error::PDU => {
Err(io::Error::other(
format!("server reported error {}", header.session())
))
}
pdu => {
Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("unexpected PDU {pdu}")
))
}
}
}
}
#[derive(Clone, Copy, Debug)]
pub enum PayloadError {
UnknownWithdraw,
DuplicateAnnounce,
Corrupt,
Internal,
}
impl PayloadError {
fn error_code(self) -> u16 {
match self {
PayloadError::UnknownWithdraw => 6,
PayloadError::DuplicateAnnounce => 7,
PayloadError::Corrupt => 0,
PayloadError::Internal => 1
}
}
async fn send(
self, version: u8, pdu: Option<pdu::Payload>,
sock: &mut (impl AsyncWrite + Unpin)
) -> Result<(), io::Error> {
match pdu {
Some(pdu) => {
pdu::Error::new(
version, self.error_code(), pdu.as_partial_slice(), ""
).write(sock).await
}
None => {
pdu::Error::new(
version, self.error_code(), "", ""
).write(sock).await
}
}
}
}
impl fmt::Display for PayloadError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str(match *self {
PayloadError::UnknownWithdraw => "withdrawal of non-existing item",
PayloadError::DuplicateAnnounce => "duplicate announcement",
PayloadError::Corrupt => "corrup data set",
PayloadError::Internal => "internal error",
})
}
}
impl error::Error for PayloadError { }