tokio-xmpp: rewrite for futures-0.3

This commit is contained in:
Astro 2020-03-05 01:25:24 +01:00
commit 23cb34e026
18 changed files with 801 additions and 1222 deletions

View file

@ -1,6 +1,6 @@
[package] [package]
name = "tokio-xmpp" name = "tokio-xmpp"
version = "1.0.1" version = "2.0.0"
authors = ["Astro <astro@spaceboyz.net>", "Emmanuel Gil Peyrot <linkmauve@linkmauve.fr>", "pep <pep+code@bouah.net>", "O01eg <o01eg@yandex.ru>"] authors = ["Astro <astro@spaceboyz.net>", "Emmanuel Gil Peyrot <linkmauve@linkmauve.fr>", "pep <pep+code@bouah.net>", "O01eg <o01eg@yandex.ru>"]
description = "Asynchronous XMPP for Rust with tokio" description = "Asynchronous XMPP for Rust with tokio"
license = "MPL-2.0" license = "MPL-2.0"
@ -12,17 +12,16 @@ keywords = ["xmpp", "tokio"]
edition = "2018" edition = "2018"
[dependencies] [dependencies]
bytes = "0.4" bytes = "0.5"
futures = "0.1" futures = "0.3"
idna = "0.2" idna = "0.2"
log = "0.4" log = "0.4"
native-tls = "0.2" native-tls = "0.2"
sasl = "0.4" sasl = "0.4"
tokio = "0.1" tokio = { version = "0.2", features = ["net", "stream", "rt-util", "rt-threaded", "macros"] }
tokio-codec = "0.1" tokio-util = { version = "0.2", features = ["codec"] }
trust-dns-resolver = "0.12" tokio-tls = "0.3"
trust-dns-proto = "0.8" trust-dns-resolver = "0.19"
tokio-io = "0.1" trust-dns-proto = "0.19"
tokio-tls = "0.2"
xml5ever = "0.16" xml5ever = "0.16"
xmpp-parsers = "0.17" xmpp-parsers = "0.17"

View file

@ -1,9 +1,8 @@
use futures::{future, Sink, Stream}; use futures::stream::StreamExt;
use std::convert::TryFrom; use std::convert::TryFrom;
use std::env::args; use std::env::args;
use std::process::exit; use std::process::exit;
use tokio::runtime::current_thread::Runtime; use tokio_xmpp::Client;
use tokio_xmpp::{xmpp_codec::Packet, Client};
use xmpp_parsers::{ use xmpp_parsers::{
disco::{DiscoInfoQuery, DiscoInfoResult}, disco::{DiscoInfoQuery, DiscoInfoResult},
iq::{Iq, IqType}, iq::{Iq, IqType},
@ -12,32 +11,25 @@ use xmpp_parsers::{
Element, Jid, Element, Jid,
}; };
fn main() { #[tokio::main]
async fn main() {
let args: Vec<String> = args().collect(); let args: Vec<String> = args().collect();
if args.len() != 4 { if args.len() != 4 {
println!("Usage: {} <jid> <password> <target>", args[0]); println!("Usage: {} <jid> <password> <target>", args[0]);
exit(1); exit(1);
} }
let jid = &args[1]; let jid = &args[1];
let password = &args[2]; let password = args[2].clone();
let target = &args[3]; let target = &args[3];
// tokio_core context
let mut rt = Runtime::new().unwrap();
// Client instance // Client instance
let client = Client::new(jid, password).unwrap(); let mut client = Client::new(jid, password).unwrap();
// Make the two interfaces for sending and receiving independent
// of each other so we can move one into a closure.
let (mut sink, stream) = client.split();
// Wrap sink in Option so that we can take() it for the send(self)
// to consume and return it back when ready.
let mut send = move |packet| {
sink.start_send(packet).expect("start_send");
};
// Main loop, processes events // Main loop, processes events
let mut wait_for_stream_end = false; let mut wait_for_stream_end = false;
let done = stream.for_each(|event| { let mut stream_ended = false;
while !stream_ended {
if let Some(event) = client.next().await {
if wait_for_stream_end { if wait_for_stream_end {
/* Do Nothing. */ /* Do Nothing. */
} else if event.is_online() { } else if event.is_online() {
@ -47,7 +39,7 @@ fn main() {
let iq = make_disco_iq(target_jid); let iq = make_disco_iq(target_jid);
println!("Sending disco#info request to {}", target.clone()); println!("Sending disco#info request to {}", target.clone());
println!(">> {}", String::from(&iq)); println!(">> {}", String::from(&iq));
send(Packet::Stanza(iq)); client.send_stanza(iq).await.unwrap();
} else if let Some(stanza) = event.into_stanza() { } else if let Some(stanza) = event.into_stanza() {
if stanza.is("iq", "jabber:client") { if stanza.is("iq", "jabber:client") {
let iq = Iq::try_from(stanza).unwrap(); let iq = Iq::try_from(stanza).unwrap();
@ -57,25 +49,17 @@ fn main() {
for ext in disco_info.extensions { for ext in disco_info.extensions {
if let Ok(server_info) = ServerInfo::try_from(ext) { if let Ok(server_info) = ServerInfo::try_from(ext) {
print_server_info(server_info); print_server_info(server_info);
}
}
}
}
wait_for_stream_end = true; wait_for_stream_end = true;
send(Packet::StreamEnd); client.send_end().await.unwrap();
} }
} }
} }
} } else {
} stream_ended = true;
}
}
Box::new(future::ok(()))
});
// Start polling `done`
match rt.block_on(done) {
Ok(_) => (),
Err(e) => {
println!("Fatal: {}", e);
()
} }
} }
} }

View file

@ -1,12 +1,12 @@
use futures::{future, Future, Sink, Stream}; use futures::stream::StreamExt;
use std::convert::TryFrom; use std::convert::TryFrom;
use std::env::args; use std::env::args;
use std::fs::{create_dir_all, File}; use std::fs::{create_dir_all, File};
use std::io::{self, Write}; use std::io::{self, Write};
use std::process::exit; use std::process::exit;
use std::str::FromStr; use std::str::FromStr;
use tokio::runtime::current_thread::Runtime; use tokio;
use tokio_xmpp::{Client, Packet}; use tokio_xmpp::Client;
use xmpp_parsers::{ use xmpp_parsers::{
avatar::{Data as AvatarData, Metadata as AvatarMetadata}, avatar::{Data as AvatarData, Metadata as AvatarMetadata},
caps::{compute_disco, hash_caps, Caps}, caps::{compute_disco, hash_caps, Caps},
@ -22,52 +22,29 @@ use xmpp_parsers::{
NodeName, NodeName,
}, },
stanza_error::{DefinedCondition, ErrorType, StanzaError}, stanza_error::{DefinedCondition, ErrorType, StanzaError},
Jid, Element, Jid,
}; };
fn main() { #[tokio::main]
async fn main() {
let args: Vec<String> = args().collect(); let args: Vec<String> = args().collect();
if args.len() != 3 { if args.len() != 3 {
println!("Usage: {} <jid> <password>", args[0]); println!("Usage: {} <jid> <password>", args[0]);
exit(1); exit(1);
} }
let jid = &args[1]; let jid = &args[1];
let password = &args[2]; let password = args[2].clone();
// tokio_core context
let mut rt = Runtime::new().unwrap();
// Client instance // Client instance
let client = Client::new(jid, password).unwrap(); let mut client = Client::new(jid, password).unwrap();
// Make the two interfaces for sending and receiving independent
// of each other so we can move one into a closure.
let (sink, stream) = client.split();
// Create outgoing pipe
let (mut tx, rx) = futures::unsync::mpsc::unbounded();
rt.spawn(
rx.forward(sink.sink_map_err(|_| panic!("Pipe")))
.map(|(rx, mut sink)| {
drop(rx);
let _ = sink.close();
})
.map_err(|e| {
panic!("Send error: {:?}", e);
}),
);
let disco_info = make_disco(); let disco_info = make_disco();
// Main loop, processes events // Main loop, processes events
let mut wait_for_stream_end = false; let mut wait_for_stream_end = false;
let done = stream.for_each(move |event| { let mut stream_ended = false;
// Helper function to send an iq error. while !stream_ended {
let mut send_error = |to, id, type_, condition, text: &str| { if let Some(event) = client.next().await {
let error = StanzaError::new(type_, condition, "en", text);
let iq = Iq::from_error(id, error).with_to(to);
tx.start_send(Packet::Stanza(iq.into())).unwrap();
};
if wait_for_stream_end { if wait_for_stream_end {
/* Do nothing */ /* Do nothing */
} else if event.is_online() { } else if event.is_online() {
@ -75,7 +52,7 @@ fn main() {
let caps = get_disco_caps(&disco_info, "https://gitlab.com/xmpp-rs/tokio-xmpp"); let caps = get_disco_caps(&disco_info, "https://gitlab.com/xmpp-rs/tokio-xmpp");
let presence = make_presence(caps); let presence = make_presence(caps);
tx.start_send(Packet::Stanza(presence.into())).unwrap(); client.send_stanza(presence.into()).await.unwrap();
} else if let Some(stanza) = event.into_stanza() { } else if let Some(stanza) = event.into_stanza() {
if stanza.is("iq", "jabber:client") { if stanza.is("iq", "jabber:client") {
let iq = Iq::try_from(stanza).unwrap(); let iq = Iq::try_from(stanza).unwrap();
@ -86,29 +63,33 @@ fn main() {
Ok(query) => { Ok(query) => {
let mut disco = disco_info.clone(); let mut disco = disco_info.clone();
disco.node = query.node; disco.node = query.node;
let iq = let iq = Iq::from_result(iq.id, Some(disco))
Iq::from_result(iq.id, Some(disco)).with_to(iq.from.unwrap()); .with_to(iq.from.unwrap());
tx.start_send(Packet::Stanza(iq.into())).unwrap(); client.send_stanza(iq.into()).await.unwrap();
} }
Err(err) => { Err(err) => client
send_error( .send_stanza(make_error(
iq.from.unwrap(), iq.from.unwrap(),
iq.id, iq.id,
ErrorType::Modify, ErrorType::Modify,
DefinedCondition::BadRequest, DefinedCondition::BadRequest,
&format!("{}", err), &format!("{}", err),
); ))
} .await
.unwrap(),
} }
} else { } else {
// We MUST answer unhandled get iqs with a service-unavailable error. // We MUST answer unhandled get iqs with a service-unavailable error.
send_error( client
.send_stanza(make_error(
iq.from.unwrap(), iq.from.unwrap(),
iq.id, iq.id,
ErrorType::Cancel, ErrorType::Cancel,
DefinedCondition::ServiceUnavailable, DefinedCondition::ServiceUnavailable,
"No handler defined for this kind of iq.", "No handler defined for this kind of iq.",
); ))
.await
.unwrap();
} }
} else if let IqType::Result(Some(payload)) = iq.payload { } else if let IqType::Result(Some(payload)) = iq.payload {
if payload.is("pubsub", ns::PUBSUB) { if payload.is("pubsub", ns::PUBSUB) {
@ -118,22 +99,25 @@ fn main() {
} }
} else if let IqType::Set(_) = iq.payload { } else if let IqType::Set(_) = iq.payload {
// We MUST answer unhandled set iqs with a service-unavailable error. // We MUST answer unhandled set iqs with a service-unavailable error.
send_error( client
.send_stanza(make_error(
iq.from.unwrap(), iq.from.unwrap(),
iq.id, iq.id,
ErrorType::Cancel, ErrorType::Cancel,
DefinedCondition::ServiceUnavailable, DefinedCondition::ServiceUnavailable,
"No handler defined for this kind of iq.", "No handler defined for this kind of iq.",
); ))
.await
.unwrap();
} }
} else if stanza.is("message", "jabber:client") { } else if stanza.is("message", "jabber:client") {
let message = Message::try_from(stanza).unwrap(); let message = Message::try_from(stanza).unwrap();
let from = message.from.clone().unwrap(); let from = message.from.clone().unwrap();
if let Some(body) = message.get_best_body(vec!["en"]) { if let Some(body) = message.get_best_body(vec!["en"]) {
if body.1 .0 == "die" { if body.0 == "die" {
println!("Secret die command triggered by {}", from); println!("Secret die command triggered by {}", from);
wait_for_stream_end = true; wait_for_stream_end = true;
tx.start_send(Packet::StreamEnd).unwrap(); client.send_end().await.unwrap();
} }
} }
for child in message.payloads { for child in message.payloads {
@ -145,13 +129,14 @@ fn main() {
let payload = item.payload.clone().unwrap(); let payload = item.payload.clone().unwrap();
if payload.is("metadata", ns::AVATAR_METADATA) { if payload.is("metadata", ns::AVATAR_METADATA) {
// TODO: do something with these metadata. // TODO: do something with these metadata.
let _metadata = AvatarMetadata::try_from(payload).unwrap(); let _metadata =
AvatarMetadata::try_from(payload).unwrap();
println!( println!(
"{} has published an avatar, downloading...", "{} has published an avatar, downloading...",
from.clone() from.clone()
); );
let iq = download_avatar(from.clone()); let iq = download_avatar(from.clone());
tx.start_send(Packet::Stanza(iq.into())).unwrap(); client.send_stanza(iq.into()).await.unwrap();
} }
} }
} }
@ -160,24 +145,30 @@ fn main() {
} }
} else if stanza.is("presence", "jabber:client") { } else if stanza.is("presence", "jabber:client") {
// Nothing to do here. // Nothing to do here.
()
} else { } else {
panic!("Unknown stanza: {}", String::from(&stanza)); panic!("Unknown stanza: {}", String::from(&stanza));
} }
} }
} else {
future::ok(()) println!("stream_ended");
}); stream_ended = true;
// Start polling `done`
match rt.block_on(done) {
Ok(_) => (),
Err(e) => {
println!("Fatal: {}", e);
()
} }
} }
} }
fn make_error(
to: Jid,
id: String,
type_: ErrorType,
condition: DefinedCondition,
text: &str,
) -> Element {
let error = StanzaError::new(type_, condition, "en", text);
let iq = Iq::from_error(id, error).with_to(to);
iq.into()
}
fn make_disco() -> DiscoInfoResult { fn make_disco() -> DiscoInfoResult {
let identities = vec![Identity::new("client", "bot", "en", "tokio-xmpp")]; let identities = vec![Identity::new("client", "bot", "en", "tokio-xmpp")];
let features = vec![ let features = vec![
@ -235,6 +226,7 @@ fn handle_iq_result(pubsub: PubSub, from: &Jid) {
} }
} }
// TODO: may use tokio?
fn save_avatar(from: &Jid, id: String, data: &[u8]) -> io::Result<()> { fn save_avatar(from: &Jid, id: String, data: &[u8]) -> io::Result<()> {
let directory = format!("data/{}", from); let directory = format!("data/{}", from);
let filename = format!("data/{}/{}", from, id); let filename = format!("data/{}/{}", from, id);

View file

@ -1,14 +1,15 @@
use futures::{future, Future, Sink, Stream}; use futures::stream::StreamExt;
use std::convert::TryFrom; use std::convert::TryFrom;
use std::env::args; use std::env::args;
use std::process::exit; use std::process::exit;
use tokio::runtime::current_thread::Runtime; use tokio;
use tokio_xmpp::{Client, Packet}; use tokio_xmpp::Client;
use xmpp_parsers::message::{Body, Message, MessageType}; use xmpp_parsers::message::{Body, Message, MessageType};
use xmpp_parsers::presence::{Presence, Show as PresenceShow, Type as PresenceType}; use xmpp_parsers::presence::{Presence, Show as PresenceShow, Type as PresenceType};
use xmpp_parsers::{Element, Jid}; use xmpp_parsers::{Element, Jid};
fn main() { #[tokio::main]
async fn main() {
let args: Vec<String> = args().collect(); let args: Vec<String> = args().collect();
if args.len() != 3 { if args.len() != 3 {
println!("Usage: {} <jid> <password>", args[0]); println!("Usage: {} <jid> <password>", args[0]);
@ -17,31 +18,16 @@ fn main() {
let jid = &args[1]; let jid = &args[1];
let password = &args[2]; let password = &args[2];
// tokio_core context
let mut rt = Runtime::new().unwrap();
// Client instance // Client instance
let client = Client::new(jid, password).unwrap(); let mut client = Client::new(jid, password.to_owned()).unwrap();
client.set_reconnect(true);
// Make the two interfaces for sending and receiving independent
// of each other so we can move one into a closure.
let (sink, stream) = client.split();
// Create outgoing pipe
let (mut tx, rx) = futures::unsync::mpsc::unbounded();
rt.spawn(
rx.forward(sink.sink_map_err(|_| panic!("Pipe")))
.map(|(rx, mut sink)| {
drop(rx);
let _ = sink.close();
})
.map_err(|e| {
panic!("Send error: {:?}", e);
}),
);
// Main loop, processes events // Main loop, processes events
let mut wait_for_stream_end = false; let mut wait_for_stream_end = false;
let done = stream.for_each(move |event| { let mut stream_ended = false;
while !stream_ended {
if let Some(event) = client.next().await {
println!("event: {:?}", event);
if wait_for_stream_end { if wait_for_stream_end {
/* Do nothing */ /* Do nothing */
} else if event.is_online() { } else if event.is_online() {
@ -52,7 +38,7 @@ fn main() {
println!("Online at {}", jid); println!("Online at {}", jid);
let presence = make_presence(); let presence = make_presence();
tx.start_send(Packet::Stanza(presence)).unwrap(); client.send_stanza(presence).await.unwrap();
} else if let Some(message) = event } else if let Some(message) = event
.into_stanza() .into_stanza()
.and_then(|stanza| Message::try_from(stanza).ok()) .and_then(|stanza| Message::try_from(stanza).ok())
@ -61,28 +47,21 @@ fn main() {
(Some(ref from), Some(ref body)) if body.0 == "die" => { (Some(ref from), Some(ref body)) if body.0 == "die" => {
println!("Secret die command triggered by {}", from); println!("Secret die command triggered by {}", from);
wait_for_stream_end = true; wait_for_stream_end = true;
tx.start_send(Packet::StreamEnd).unwrap(); client.send_end().await.unwrap();
} }
(Some(ref from), Some(ref body)) => { (Some(ref from), Some(ref body)) => {
if message.type_ != MessageType::Error { if message.type_ != MessageType::Error {
// This is a message we'll echo // This is a message we'll echo
let reply = make_reply(from.clone(), &body.0); let reply = make_reply(from.clone(), &body.0);
tx.start_send(Packet::Stanza(reply)).unwrap(); client.send_stanza(reply).await.unwrap();
} }
} }
_ => {} _ => {}
} }
} }
} else {
future::ok(()) println!("stream_ended");
}); stream_ended = true;
// Start polling `done`
match rt.block_on(done) {
Ok(_) => (),
Err(e) => {
println!("Fatal: {}", e);
()
} }
} }
} }

View file

@ -1,15 +1,15 @@
use futures::{future, Sink, Stream}; use futures::stream::StreamExt;
use std::convert::TryFrom; use std::convert::TryFrom;
use std::env::args; use std::env::args;
use std::process::exit; use std::process::exit;
use std::str::FromStr; use std::str::FromStr;
use tokio::runtime::current_thread::Runtime;
use tokio_xmpp::Component; use tokio_xmpp::Component;
use xmpp_parsers::message::{Body, Message, MessageType}; use xmpp_parsers::message::{Body, Message, MessageType};
use xmpp_parsers::presence::{Presence, Show as PresenceShow, Type as PresenceType}; use xmpp_parsers::presence::{Presence, Show as PresenceShow, Type as PresenceType};
use xmpp_parsers::{Element, Jid}; use xmpp_parsers::{Element, Jid};
fn main() { #[tokio::main]
async fn main() {
let args: Vec<String> = args().collect(); let args: Vec<String> = args().collect();
if args.len() < 3 || args.len() > 5 { if args.len() < 3 || args.len() > 5 {
println!("Usage: {} <jid> <password> [server] [port]", args[0]); println!("Usage: {} <jid> <password> [server] [port]", args[0]);
@ -24,57 +24,38 @@ fn main() {
.unwrap_or("127.0.0.1".to_owned()); .unwrap_or("127.0.0.1".to_owned());
let port: u16 = args.get(4).unwrap().parse().unwrap_or(5347u16); let port: u16 = args.get(4).unwrap().parse().unwrap_or(5347u16);
// tokio_core context
let mut rt = Runtime::new().unwrap();
// Component instance // Component instance
println!("{} {} {} {}", jid, password, server, port); println!("{} {} {} {}", jid, password, server, port);
let component = Component::new(jid, password, server, port).unwrap(); let mut component = Component::new(jid, password, server, port).await.unwrap();
// Make the two interfaces for sending and receiving independent // Make the two interfaces for sending and receiving independent
// of each other so we can move one into a closure. // of each other so we can move one into a closure.
println!("Got it: {}", component.jid.clone()); println!("Online: {}", component.jid);
let (mut sink, stream) = component.split();
// Wrap sink in Option so that we can take() it for the send(self)
// to consume and return it back when ready.
let mut send = move |stanza| {
sink.start_send(stanza).expect("start_send");
};
// Main loop, processes events
let done = stream.for_each(|event| {
if event.is_online() {
println!("Online!");
// TODO: replace these hardcoded JIDs // TODO: replace these hardcoded JIDs
let presence = make_presence( let presence = make_presence(
Jid::from_str("test@component.linkmauve.fr/coucou").unwrap(), Jid::from_str("test@component.linkmauve.fr/coucou").unwrap(),
Jid::from_str("linkmauve@linkmauve.fr").unwrap(), Jid::from_str("linkmauve@linkmauve.fr").unwrap(),
); );
send(presence); component.send_stanza(presence).await.unwrap();
} else if let Some(message) = event
.into_stanza() // Main loop, processes events
.and_then(|stanza| Message::try_from(stanza).ok()) loop {
{ if let Some(stanza) = component.next().await {
if let Some(message) = Message::try_from(stanza).ok() {
// This is a message we'll echo // This is a message we'll echo
match (message.from, message.bodies.get("")) { match (message.from, message.bodies.get("")) {
(Some(from), Some(body)) => { (Some(from), Some(body)) => {
if message.type_ != MessageType::Error { if message.type_ != MessageType::Error {
let reply = make_reply(from, &body.0); let reply = make_reply(from, &body.0);
send(reply); component.send_stanza(reply).await.unwrap();
} }
} }
_ => (), _ => (),
} }
} }
} else {
Box::new(future::ok(())) break;
});
// Start polling `done`
match rt.block_on(done) {
Ok(_) => (),
Err(e) => {
println!("Fatal: {}", e);
()
} }
} }
} }

View file

@ -1,7 +1,4 @@
use futures::{ use futures::stream::StreamExt;
future::{err, ok, IntoFuture},
Future, Poll, Stream,
};
use sasl::client::mechanisms::{Anonymous, Plain, Scram}; use sasl::client::mechanisms::{Anonymous, Plain, Scram};
use sasl::client::Mechanism; use sasl::client::Mechanism;
use sasl::common::scram::{Sha1, Sha256}; use sasl::common::scram::{Sha1, Sha256};
@ -9,7 +6,7 @@ use sasl::common::Credentials;
use std::collections::HashSet; use std::collections::HashSet;
use std::convert::TryFrom; use std::convert::TryFrom;
use std::str::FromStr; use std::str::FromStr;
use tokio_io::{AsyncRead, AsyncWrite}; use tokio::io::{AsyncRead, AsyncWrite};
use xmpp_parsers::sasl::{Auth, Challenge, Failure, Mechanism as XMPPMechanism, Response, Success}; use xmpp_parsers::sasl::{Auth, Challenge, Failure, Mechanism as XMPPMechanism, Response, Success};
use crate::xmpp_codec::Packet; use crate::xmpp_codec::Packet;
@ -18,13 +15,11 @@ use crate::{AuthError, Error, ProtocolError};
const NS_XMPP_SASL: &str = "urn:ietf:params:xml:ns:xmpp-sasl"; const NS_XMPP_SASL: &str = "urn:ietf:params:xml:ns:xmpp-sasl";
pub struct ClientAuth<S: AsyncRead + AsyncWrite> { pub async fn auth<S: AsyncRead + AsyncWrite + Unpin>(
future: Box<dyn Future<Item = XMPPStream<S>, Error = Error>>, mut stream: XMPPStream<S>,
} creds: Credentials,
) -> Result<S, Error> {
impl<S: AsyncRead + AsyncWrite + 'static> ClientAuth<S> { let local_mechs: Vec<Box<dyn Fn() -> Box<dyn Mechanism + Send + Sync> + Send>> = vec![
pub fn new(stream: XMPPStream<S>, creds: Credentials) -> Result<Self, Error> {
let local_mechs: Vec<Box<dyn Fn() -> Box<dyn Mechanism>>> = vec![
Box::new(|| Box::new(Scram::<Sha256>::from_credentials(creds.clone()).unwrap())), Box::new(|| Box::new(Scram::<Sha256>::from_credentials(creds.clone()).unwrap())),
Box::new(|| Box::new(Scram::<Sha1>::from_credentials(creds.clone()).unwrap())), Box::new(|| Box::new(Scram::<Sha1>::from_credentials(creds.clone()).unwrap())),
Box::new(|| Box::new(Plain::from_credentials(creds.clone()).unwrap())), Box::new(|| Box::new(Plain::from_credentials(creds.clone()).unwrap())),
@ -47,80 +42,43 @@ impl<S: AsyncRead + AsyncWrite + 'static> ClientAuth<S> {
let mechanism_name = let mechanism_name =
XMPPMechanism::from_str(mechanism.name()).map_err(ProtocolError::Parsers)?; XMPPMechanism::from_str(mechanism.name()).map_err(ProtocolError::Parsers)?;
let send_initial = Box::new(stream.send_stanza(Auth { stream
.send_stanza(Auth {
mechanism: mechanism_name, mechanism: mechanism_name,
data: initial, data: initial,
}))
.map_err(Error::Io);
let future = Box::new(
send_initial
.and_then(|stream| Self::handle_challenge(stream, mechanism))
.and_then(|stream| stream.restart()),
);
return Ok(ClientAuth { future });
}
}
Err(AuthError::NoMechanism)?
}
fn handle_challenge(
stream: XMPPStream<S>,
mut mechanism: Box<dyn Mechanism>,
) -> Box<dyn Future<Item = XMPPStream<S>, Error = Error>> {
Box::new(
stream
.into_future()
.map_err(|(e, _stream)| e.into())
.and_then(|(stanza, stream)| {
match stanza {
Some(Packet::Stanza(stanza)) => {
if let Ok(challenge) = Challenge::try_from(stanza.clone()) {
let response = mechanism.response(&challenge.data);
Box::new(
response
.map_err(|e| AuthError::Sasl(e).into())
.into_future()
.and_then(|response| {
// Send response and loop
stream
.send_stanza(Response { data: response })
.map_err(Error::Io)
.and_then(|stream| {
Self::handle_challenge(stream, mechanism)
}) })
}), .await?;
)
loop {
match stream.next().await {
Some(Ok(Packet::Stanza(stanza))) => {
if let Ok(challenge) = Challenge::try_from(stanza.clone()) {
let response = mechanism
.response(&challenge.data)
.map_err(|e| AuthError::Sasl(e))?;
// Send response and loop
stream.send_stanza(Response { data: response }).await?;
} else if let Ok(_) = Success::try_from(stanza.clone()) { } else if let Ok(_) = Success::try_from(stanza.clone()) {
Box::new(ok(stream)) return Ok(stream.into_inner());
} else if let Ok(failure) = Failure::try_from(stanza.clone()) { } else if let Ok(failure) = Failure::try_from(stanza.clone()) {
Box::new(err(Error::Auth(AuthError::Fail( return Err(Error::Auth(AuthError::Fail(failure.defined_condition)));
failure.defined_condition,
))))
} else if stanza.name() == "failure" { } else if stanza.name() == "failure" {
// Workaround for https://gitlab.com/xmpp-rs/xmpp-parsers/merge_requests/1 // Workaround for https://gitlab.com/xmpp-rs/xmpp-parsers/merge_requests/1
Box::new(err(Error::Auth(AuthError::Sasl("failure".to_string())))) return Err(Error::Auth(AuthError::Sasl("failure".to_string())));
} else { } else {
// ignore and loop // ignore and loop
Self::handle_challenge(stream, mechanism)
} }
} }
Some(_) => { Some(Ok(_)) => {
// ignore and loop // ignore and loop
Self::handle_challenge(stream, mechanism)
} }
None => Box::new(err(Error::Disconnected)), Some(Err(e)) => return Err(e),
None => return Err(Error::Disconnected),
}
} }
}),
)
} }
} }
impl<S: AsyncRead + AsyncWrite> Future for ClientAuth<S> { Err(AuthError::NoMechanism.into())
type Item = XMPPStream<S>;
type Error = Error;
fn poll(&mut self) -> Poll<Self::Item, Self::Error> {
self.future.poll()
}
} }

View file

@ -1,7 +1,7 @@
use futures::{sink, Async, Future, Poll, Stream}; use futures::stream::StreamExt;
use std::convert::TryFrom; use std::convert::TryFrom;
use std::mem::replace; use std::marker::Unpin;
use tokio_io::{AsyncRead, AsyncWrite}; use tokio::io::{AsyncRead, AsyncWrite};
use xmpp_parsers::bind::{BindQuery, BindResponse}; use xmpp_parsers::bind::{BindQuery, BindResponse};
use xmpp_parsers::iq::{Iq, IqType}; use xmpp_parsers::iq::{Iq, IqType};
use xmpp_parsers::Jid; use xmpp_parsers::Jid;
@ -13,90 +13,43 @@ use crate::{Error, ProtocolError};
const NS_XMPP_BIND: &str = "urn:ietf:params:xml:ns:xmpp-bind"; const NS_XMPP_BIND: &str = "urn:ietf:params:xml:ns:xmpp-bind";
const BIND_REQ_ID: &str = "resource-bind"; const BIND_REQ_ID: &str = "resource-bind";
pub enum ClientBind<S: AsyncWrite> { pub async fn bind<S: AsyncRead + AsyncWrite + Unpin>(
Unsupported(XMPPStream<S>), mut stream: XMPPStream<S>,
WaitSend(sink::Send<XMPPStream<S>>), ) -> Result<XMPPStream<S>, Error> {
WaitRecv(XMPPStream<S>),
Invalid,
}
impl<S: AsyncWrite> ClientBind<S> {
/// Consumes and returns the stream to express that you cannot use
/// the stream for anything else until the resource binding
/// req/resp are done.
pub fn new(stream: XMPPStream<S>) -> Self {
match stream.stream_features.get_child("bind", NS_XMPP_BIND) { match stream.stream_features.get_child("bind", NS_XMPP_BIND) {
None => None => {
// No resource binding available, // No resource binding available,
// return the (probably // usable) stream immediately // return the (probably // usable) stream immediately
{ return Ok(stream);
ClientBind::Unsupported(stream)
} }
Some(_) => { Some(_) => {
let resource; let resource = if let Jid::Full(jid) = stream.jid.clone() {
if let Jid::Full(jid) = stream.jid.clone() { Some(jid.resource)
resource = Some(jid.resource);
} else { } else {
resource = None; None
} };
let iq = Iq::from_set(BIND_REQ_ID, BindQuery::new(resource)); let iq = Iq::from_set(BIND_REQ_ID, BindQuery::new(resource));
let send = stream.send_stanza(iq); stream.send_stanza(iq).await?;
ClientBind::WaitSend(send)
}
}
}
}
impl<S: AsyncRead + AsyncWrite> Future for ClientBind<S> { loop {
type Item = XMPPStream<S>; match stream.next().await {
type Error = Error; Some(Ok(Packet::Stanza(stanza))) => match Iq::try_from(stanza) {
Ok(iq) if iq.id == BIND_REQ_ID => match iq.payload {
fn poll(&mut self) -> Poll<Self::Item, Self::Error> {
let state = replace(self, ClientBind::Invalid);
match state {
ClientBind::Unsupported(stream) => Ok(Async::Ready(stream)),
ClientBind::WaitSend(mut send) => match send.poll() {
Ok(Async::Ready(stream)) => {
replace(self, ClientBind::WaitRecv(stream));
self.poll()
}
Ok(Async::NotReady) => {
replace(self, ClientBind::WaitSend(send));
Ok(Async::NotReady)
}
Err(e) => Err(e)?,
},
ClientBind::WaitRecv(mut stream) => match stream.poll() {
Ok(Async::Ready(Some(Packet::Stanza(stanza)))) => match Iq::try_from(stanza) {
Ok(iq) => {
if iq.id == BIND_REQ_ID {
match iq.payload {
IqType::Result(payload) => { IqType::Result(payload) => {
payload payload
.and_then(|payload| BindResponse::try_from(payload).ok()) .and_then(|payload| BindResponse::try_from(payload).ok())
.map(|bind| stream.jid = bind.into()); .map(|bind| stream.jid = bind.into());
Ok(Async::Ready(stream)) return Ok(stream);
} }
_ => Err(ProtocolError::InvalidBindResponse)?, _ => return Err(ProtocolError::InvalidBindResponse.into()),
}
} else {
Ok(Async::NotReady)
}
}
_ => Ok(Async::NotReady),
}, },
Ok(Async::Ready(_)) => { _ => {}
replace(self, ClientBind::WaitRecv(stream));
self.poll()
}
Ok(Async::NotReady) => {
replace(self, ClientBind::WaitRecv(stream));
Ok(Async::NotReady)
}
Err(e) => Err(e)?,
}, },
ClientBind::Invalid => unreachable!(), Some(Ok(_)) => {}
Some(Err(e)) => return Err(e),
None => return Err(Error::Disconnected),
}
}
} }
} }
} }

View file

@ -1,28 +1,33 @@
use futures::{done, Async, AsyncSink, Future, Poll, Sink, StartSend, Stream}; use futures::{sink::SinkExt, task::Poll, Future, Sink, Stream};
use idna; use idna;
use sasl::common::{ChannelBinding, Credentials}; use sasl::common::{ChannelBinding, Credentials};
use std::mem::replace; use std::mem::replace;
use std::pin::Pin;
use std::str::FromStr; use std::str::FromStr;
use std::task::Context;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::TcpStream; use tokio::net::TcpStream;
use tokio_io::{AsyncRead, AsyncWrite}; use tokio::task::JoinHandle;
use tokio::task::LocalSet;
use tokio_tls::TlsStream; use tokio_tls::TlsStream;
use xmpp_parsers::{Jid, JidParseError}; use xmpp_parsers::{Element, Jid, JidParseError};
use super::event::Event; use super::event::Event;
use super::happy_eyeballs::Connecter; use super::happy_eyeballs::connect;
use super::starttls::{StartTlsClient, NS_XMPP_TLS}; use super::starttls::{starttls, NS_XMPP_TLS};
use super::xmpp_codec::Packet; use super::xmpp_codec::Packet;
use super::xmpp_stream; use super::xmpp_stream;
use super::{Error, ProtocolError}; use super::{Error, ProtocolError};
mod auth; mod auth;
use self::auth::ClientAuth;
mod bind; mod bind;
use self::bind::ClientBind;
/// XMPP client connection and state /// XMPP client connection and state
pub struct Client { pub struct Client {
state: ClientState, state: ClientState,
jid: Jid,
password: String,
reconnect: bool,
} }
type XMPPStream = xmpp_stream::XMPPStream<TlsStream<TcpStream>>; type XMPPStream = xmpp_stream::XMPPStream<TlsStream<TcpStream>>;
@ -31,7 +36,7 @@ const NS_JABBER_CLIENT: &str = "jabber:client";
enum ClientState { enum ClientState {
Invalid, Invalid,
Disconnected, Disconnected,
Connecting(Box<dyn Future<Item = XMPPStream, Error = Error>>), Connecting(JoinHandle<Result<XMPPStream, Error>>, LocalSet),
Connected(XMPPStream), Connected(XMPPStream),
} }
@ -40,87 +45,87 @@ impl Client {
/// ///
/// Start polling the returned instance so that it will connect /// Start polling the returned instance so that it will connect
/// and yield events. /// and yield events.
pub fn new(jid: &str, password: &str) -> Result<Self, JidParseError> { pub fn new<P: Into<String>>(jid: &str, password: P) -> Result<Self, JidParseError> {
let jid = Jid::from_str(jid)?; let jid = Jid::from_str(jid)?;
let client = Self::new_with_jid(jid, password); let client = Self::new_with_jid(jid, password.into());
Ok(client) Ok(client)
} }
/// Start a new client given that the JID is already parsed. /// Start a new client given that the JID is already parsed.
pub fn new_with_jid(jid: Jid, password: &str) -> Self { pub fn new_with_jid(jid: Jid, password: String) -> Self {
let password = password.to_owned(); let local = LocalSet::new();
let connect = Self::make_connect(jid, password.clone()); let connect = local.spawn_local(Self::connect(jid.clone(), password.clone()));
let client = Client { let client = Client {
state: ClientState::Connecting(Box::new(connect)), jid,
password,
state: ClientState::Connecting(connect, local),
reconnect: false,
}; };
client client
} }
fn make_connect(jid: Jid, password: String) -> impl Future<Item = XMPPStream, Error = Error> { /// Set whether to reconnect (`true`) or end the stream (`false`)
let username = jid.clone().node().unwrap(); /// when a connection to the server has ended.
let jid1 = jid.clone(); pub fn set_reconnect(&mut self, reconnect: bool) -> &mut Self {
let jid2 = jid.clone(); self.reconnect = reconnect;
let password = password; self
done(idna::domain_to_ascii(&jid.domain()))
.map_err(|_| Error::Idna)
.and_then(|domain| {
done(Connecter::from_lookup(
&domain,
Some("_xmpp-client._tcp"),
5222,
))
})
.flatten()
.and_then(move |tcp_stream| {
xmpp_stream::XMPPStream::start(tcp_stream, jid1, NS_JABBER_CLIENT.to_owned())
})
.and_then(|xmpp_stream| {
if Self::can_starttls(&xmpp_stream) {
Ok(Self::starttls(xmpp_stream))
} else {
Err(Error::Protocol(ProtocolError::NoTls))
}
})
.flatten()
.and_then(|tls_stream| XMPPStream::start(tls_stream, jid2, NS_JABBER_CLIENT.to_owned()))
.and_then(
move |xmpp_stream| done(Self::auth(xmpp_stream, username, password)), // TODO: flatten?
)
.and_then(|auth| auth)
.and_then(|xmpp_stream| Self::bind(xmpp_stream))
.and_then(|xmpp_stream| {
// println!("Bound to {}", xmpp_stream.jid);
Ok(xmpp_stream)
})
} }
fn can_starttls<S>(stream: &xmpp_stream::XMPPStream<S>) -> bool { async fn connect(jid: Jid, password: String) -> Result<XMPPStream, Error> {
stream let username = jid.clone().node().unwrap();
let password = password;
let domain = idna::domain_to_ascii(&jid.clone().domain()).map_err(|_| Error::Idna)?;
let tcp_stream = connect(&domain, Some("_xmpp-client._tcp"), 5222).await?;
let xmpp_stream =
xmpp_stream::XMPPStream::start(tcp_stream, jid, NS_JABBER_CLIENT.to_owned()).await?;
let xmpp_stream = if Self::can_starttls(&xmpp_stream) {
Self::starttls(xmpp_stream).await?
} else {
return Err(Error::Protocol(ProtocolError::NoTls));
};
let xmpp_stream = Self::auth(xmpp_stream, username, password).await?;
let xmpp_stream = Self::bind(xmpp_stream).await?;
Ok(xmpp_stream)
}
fn can_starttls<S: AsyncRead + AsyncWrite + Unpin>(
xmpp_stream: &xmpp_stream::XMPPStream<S>,
) -> bool {
xmpp_stream
.stream_features .stream_features
.get_child("starttls", NS_XMPP_TLS) .get_child("starttls", NS_XMPP_TLS)
.is_some() .is_some()
} }
fn starttls<S: AsyncRead + AsyncWrite>( async fn starttls<S: AsyncRead + AsyncWrite + Unpin>(
stream: xmpp_stream::XMPPStream<S>, xmpp_stream: xmpp_stream::XMPPStream<S>,
) -> StartTlsClient<S> { ) -> Result<xmpp_stream::XMPPStream<TlsStream<S>>, Error> {
StartTlsClient::from_stream(stream) let jid = xmpp_stream.jid.clone();
let tls_stream = starttls(xmpp_stream).await?;
xmpp_stream::XMPPStream::start(tls_stream, jid, NS_JABBER_CLIENT.to_owned()).await
} }
fn auth<S: AsyncRead + AsyncWrite + 'static>( async fn auth<S: AsyncRead + AsyncWrite + Unpin + 'static>(
stream: xmpp_stream::XMPPStream<S>, xmpp_stream: xmpp_stream::XMPPStream<S>,
username: String, username: String,
password: String, password: String,
) -> Result<ClientAuth<S>, Error> { ) -> Result<xmpp_stream::XMPPStream<S>, Error> {
let jid = xmpp_stream.jid.clone();
let creds = Credentials::default() let creds = Credentials::default()
.with_username(username) .with_username(username)
.with_password(password) .with_password(password)
.with_channel_binding(ChannelBinding::None); .with_channel_binding(ChannelBinding::None);
ClientAuth::new(stream, creds) let stream = auth::auth(xmpp_stream, creds).await?;
xmpp_stream::XMPPStream::start(stream, jid, NS_JABBER_CLIENT.to_owned()).await
} }
fn bind<S: AsyncWrite>(stream: xmpp_stream::XMPPStream<S>) -> ClientBind<S> { async fn bind<S: Unpin + AsyncRead + AsyncWrite>(
ClientBind::new(stream) stream: xmpp_stream::XMPPStream<S>,
) -> Result<xmpp_stream::XMPPStream<S>, Error> {
bind::bind(stream).await
} }
/// Get the client's bound JID (the one reported by the XMPP /// Get the client's bound JID (the one reported by the XMPP
@ -131,102 +136,150 @@ impl Client {
_ => None, _ => None,
} }
} }
/// Send stanza
pub async fn send_stanza(&mut self, stanza: Element) -> Result<(), Error> {
self.send(Packet::Stanza(stanza)).await
}
/// End connection
pub async fn send_end(&mut self) -> Result<(), Error> {
self.send(Packet::StreamEnd).await
}
} }
impl Stream for Client { impl Stream for Client {
type Item = Event; type Item = Event;
type Error = Error;
fn poll(&mut self) -> Poll<Option<Self::Item>, Self::Error> { fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
let state = replace(&mut self.state, ClientState::Invalid); let state = replace(&mut self.state, ClientState::Invalid);
match state { match state {
ClientState::Invalid => Err(Error::InvalidState), ClientState::Invalid => panic!("Invalid client state"),
ClientState::Disconnected => Ok(Async::Ready(None)), ClientState::Disconnected if self.reconnect => {
ClientState::Connecting(mut connect) => match connect.poll() { // TODO: add timeout
Ok(Async::Ready(stream)) => { let mut local = LocalSet::new();
let jid = stream.jid.clone(); let connect =
local.spawn_local(Self::connect(self.jid.clone(), self.password.clone()));
let _ = Pin::new(&mut local).poll(cx);
self.state = ClientState::Connecting(connect, local);
self.poll_next(cx)
}
ClientState::Disconnected => Poll::Ready(None),
ClientState::Connecting(mut connect, mut local) => {
match Pin::new(&mut connect).poll(cx) {
Poll::Ready(Ok(Ok(stream))) => {
let bound_jid = stream.jid.clone();
self.state = ClientState::Connected(stream); self.state = ClientState::Connected(stream);
Ok(Async::Ready(Some(Event::Online(jid)))) Poll::Ready(Some(Event::Online(bound_jid)))
}
Poll::Ready(Ok(Err(e))) => {
self.state = ClientState::Disconnected;
return Poll::Ready(Some(Event::Disconnected(e.into())));
}
Poll::Ready(Err(e)) => {
self.state = ClientState::Disconnected;
panic!("connect task: {}", e);
}
Poll::Pending => {
let _ = Pin::new(&mut local).poll(cx);
self.state = ClientState::Connecting(connect, local);
Poll::Pending
}
} }
Ok(Async::NotReady) => {
self.state = ClientState::Connecting(connect);
Ok(Async::NotReady)
} }
Err(e) => Err(e),
},
ClientState::Connected(mut stream) => { ClientState::Connected(mut stream) => {
// Poll sink // Poll sink
match stream.poll_complete() { match Pin::new(&mut stream).poll_ready(cx) {
Ok(Async::NotReady) => (), Poll::Pending => (),
Ok(Async::Ready(())) => (), Poll::Ready(Ok(())) => (),
Err(e) => return Err(e)?, Poll::Ready(Err(e)) => {
self.state = ClientState::Disconnected;
return Poll::Ready(Some(Event::Disconnected(e.into())));
}
}; };
// Poll stream // Poll stream
match stream.poll() { match Pin::new(&mut stream).poll_next(cx) {
Ok(Async::Ready(None)) => { Poll::Ready(None) => {
// EOF // EOF
self.state = ClientState::Disconnected; self.state = ClientState::Disconnected;
Ok(Async::Ready(Some(Event::Disconnected))) Poll::Ready(Some(Event::Disconnected(Error::Disconnected)))
} }
Ok(Async::Ready(Some(Packet::Stanza(stanza)))) => { Poll::Ready(Some(Ok(Packet::Stanza(stanza)))) => {
// Receive stanza // Receive stanza
self.state = ClientState::Connected(stream); self.state = ClientState::Connected(stream);
Ok(Async::Ready(Some(Event::Stanza(stanza)))) Poll::Ready(Some(Event::Stanza(stanza)))
} }
Ok(Async::Ready(Some(Packet::Text(_)))) => { Poll::Ready(Some(Ok(Packet::Text(_)))) => {
// Ignore text between stanzas // Ignore text between stanzas
self.state = ClientState::Connected(stream); self.state = ClientState::Connected(stream);
Ok(Async::NotReady) Poll::Pending
} }
Ok(Async::Ready(Some(Packet::StreamStart(_)))) => { Poll::Ready(Some(Ok(Packet::StreamStart(_)))) => {
// <stream:stream> // <stream:stream>
Err(ProtocolError::InvalidStreamStart.into()) self.state = ClientState::Disconnected;
Poll::Ready(Some(Event::Disconnected(
ProtocolError::InvalidStreamStart.into(),
)))
} }
Ok(Async::Ready(Some(Packet::StreamEnd))) => { Poll::Ready(Some(Ok(Packet::StreamEnd))) => {
// End of stream: </stream:stream> // End of stream: </stream:stream>
Ok(Async::Ready(None)) self.state = ClientState::Disconnected;
Poll::Ready(Some(Event::Disconnected(Error::Disconnected)))
} }
Ok(Async::NotReady) => { Poll::Pending => {
// Try again later // Try again later
self.state = ClientState::Connected(stream); self.state = ClientState::Connected(stream);
Ok(Async::NotReady) Poll::Pending
}
Poll::Ready(Some(Err(e))) => {
self.state = ClientState::Disconnected;
Poll::Ready(Some(Event::Disconnected(e.into())))
} }
Err(e) => Err(e)?,
} }
} }
} }
} }
} }
impl Sink for Client { impl Sink<Packet> for Client {
type SinkItem = Packet; type Error = Error;
type SinkError = Error;
fn start_send(&mut self, item: Self::SinkItem) -> StartSend<Self::SinkItem, Self::SinkError> { fn start_send(mut self: Pin<&mut Self>, item: Packet) -> Result<(), Self::Error> {
match self.state { match self.state {
ClientState::Connected(ref mut stream) => Ok(stream.start_send(item)?), ClientState::Connected(ref mut stream) => {
_ => Ok(AsyncSink::NotReady(item)), Pin::new(stream).start_send(item).map_err(|e| e.into())
}
_ => Err(Error::InvalidState),
} }
} }
fn poll_complete(&mut self) -> Poll<(), Self::SinkError> { fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
match self.state { match self.state {
ClientState::Connected(ref mut stream) => stream.poll_complete().map_err(|e| e.into()), ClientState::Connected(ref mut stream) => {
_ => Ok(Async::Ready(())), Pin::new(stream).poll_ready(cx).map_err(|e| e.into())
}
_ => Poll::Pending,
} }
} }
/// This closes the inner TCP stream. fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
///
/// To synchronize your shutdown with the server side, you should
/// first send `Packet::StreamEnd` and wait for the end of the
/// incoming stream before closing the connection.
fn close(&mut self) -> Poll<(), Self::SinkError> {
match self.state { match self.state {
ClientState::Connected(ref mut stream) => stream.close().map_err(|e| e.into()), ClientState::Connected(ref mut stream) => {
_ => Ok(Async::Ready(())), Pin::new(stream).poll_flush(cx).map_err(|e| e.into())
}
_ => Poll::Pending,
}
}
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
match self.state {
ClientState::Connected(ref mut stream) => {
Pin::new(stream).poll_close(cx).map_err(|e| e.into())
}
_ => Poll::Pending,
} }
} }
} }

View file

@ -1,6 +1,6 @@
use futures::{sink, Async, Future, Poll, Stream}; use futures::stream::StreamExt;
use std::mem::replace; use std::marker::Unpin;
use tokio_io::{AsyncRead, AsyncWrite}; use tokio::io::{AsyncRead, AsyncWrite};
use xmpp_parsers::component::Handshake; use xmpp_parsers::component::Handshake;
use crate::xmpp_codec::Packet; use crate::xmpp_codec::Packet;
@ -9,81 +9,27 @@ use crate::{AuthError, Error};
const NS_JABBER_COMPONENT_ACCEPT: &str = "jabber:component:accept"; const NS_JABBER_COMPONENT_ACCEPT: &str = "jabber:component:accept";
pub struct ComponentAuth<S: AsyncWrite> { pub async fn auth<S: AsyncRead + AsyncWrite + Unpin>(
state: ComponentAuthState<S>, stream: &mut XMPPStream<S>,
} password: String,
) -> Result<(), Error> {
let nonza = Handshake::from_password_and_stream_id(&password, &stream.id);
stream.send_stanza(nonza).await?;
enum ComponentAuthState<S: AsyncWrite> { loop {
WaitSend(sink::Send<XMPPStream<S>>), match stream.next().await {
WaitRecv(XMPPStream<S>), Some(Ok(Packet::Stanza(ref stanza)))
Invalid,
}
impl<S: AsyncWrite> ComponentAuth<S> {
// TODO: doesn't have to be a Result<> actually
pub fn new(stream: XMPPStream<S>, password: String) -> Result<Self, Error> {
// FIXME: huge hack, shouldnt be an element!
let sid = stream.stream_features.name().to_owned();
let mut this = ComponentAuth {
state: ComponentAuthState::Invalid,
};
this.send(
stream,
Handshake::from_password_and_stream_id(&password, &sid),
);
Ok(this)
}
fn send(&mut self, stream: XMPPStream<S>, handshake: Handshake) {
let nonza = handshake;
let send = stream.send_stanza(nonza);
self.state = ComponentAuthState::WaitSend(send);
}
}
impl<S: AsyncRead + AsyncWrite> Future for ComponentAuth<S> {
type Item = XMPPStream<S>;
type Error = Error;
fn poll(&mut self) -> Poll<Self::Item, Self::Error> {
let state = replace(&mut self.state, ComponentAuthState::Invalid);
match state {
ComponentAuthState::WaitSend(mut send) => match send.poll() {
Ok(Async::Ready(stream)) => {
self.state = ComponentAuthState::WaitRecv(stream);
self.poll()
}
Ok(Async::NotReady) => {
self.state = ComponentAuthState::WaitSend(send);
Ok(Async::NotReady)
}
Err(e) => Err(e)?,
},
ComponentAuthState::WaitRecv(mut stream) => match stream.poll() {
Ok(Async::Ready(Some(Packet::Stanza(ref stanza))))
if stanza.is("handshake", NS_JABBER_COMPONENT_ACCEPT) => if stanza.is("handshake", NS_JABBER_COMPONENT_ACCEPT) =>
{ {
self.state = ComponentAuthState::Invalid; return Ok(());
Ok(Async::Ready(stream))
} }
Ok(Async::Ready(Some(Packet::Stanza(ref stanza)))) Some(Ok(Packet::Stanza(ref stanza)))
if stanza.is("error", "http://etherx.jabber.org/streams") => if stanza.is("error", "http://etherx.jabber.org/streams") =>
{ {
Err(AuthError::ComponentFail.into()) return Err(AuthError::ComponentFail.into());
} }
Ok(Async::Ready(_event)) => { Some(_) => {}
// println!("ComponentAuth ignore {:?}", _event); None => return Err(Error::Disconnected),
Ok(Async::NotReady)
}
Ok(_) => {
self.state = ComponentAuthState::WaitRecv(stream);
Ok(Async::NotReady)
}
Err(e) => Err(e)?,
},
ComponentAuthState::Invalid => unreachable!(),
} }
} }
} }

View file

@ -1,163 +1,115 @@
//! Components in XMPP are services/gateways that are logged into an //! Components in XMPP are services/gateways that are logged into an
//! XMPP server under a JID consisting of just a domain name. They are //! XMPP server under a JID consisting of just a domain name. They are
//! allowed to use any user and resource identifiers in their stanzas. //! allowed to use any user and resource identifiers in their stanzas.
use futures::{done, Async, AsyncSink, Future, Poll, Sink, StartSend, Stream}; use futures::{sink::SinkExt, task::Poll, Sink, Stream};
use std::mem::replace; use std::pin::Pin;
use std::str::FromStr; use std::str::FromStr;
use std::task::Context;
use tokio::net::TcpStream; use tokio::net::TcpStream;
use tokio_io::{AsyncRead, AsyncWrite}; use xmpp_parsers::{Element, Jid};
use xmpp_parsers::{Element, Jid, JidParseError};
use super::event::Event; use super::happy_eyeballs::connect;
use super::happy_eyeballs::Connecter;
use super::xmpp_codec::Packet; use super::xmpp_codec::Packet;
use super::xmpp_stream; use super::xmpp_stream;
use super::Error; use super::Error;
mod auth; mod auth;
use self::auth::ComponentAuth;
/// Component connection to an XMPP server /// Component connection to an XMPP server
///
/// This simplifies the `XMPPStream` to a `Stream`/`Sink` of `Element`
/// (stanzas). Connection handling however is up to the user.
pub struct Component { pub struct Component {
/// The component's Jabber-Id /// The component's Jabber-Id
pub jid: Jid, pub jid: Jid,
state: ComponentState, stream: XMPPStream,
} }
type XMPPStream = xmpp_stream::XMPPStream<TcpStream>; type XMPPStream = xmpp_stream::XMPPStream<TcpStream>;
const NS_JABBER_COMPONENT_ACCEPT: &str = "jabber:component:accept"; const NS_JABBER_COMPONENT_ACCEPT: &str = "jabber:component:accept";
enum ComponentState {
Invalid,
Disconnected,
Connecting(Box<dyn Future<Item = XMPPStream, Error = Error>>),
Connected(XMPPStream),
}
impl Component { impl Component {
/// Start a new XMPP component /// Start a new XMPP component
/// pub async fn new(jid: &str, password: &str, server: &str, port: u16) -> Result<Self, Error> {
/// Start polling the returned instance so that it will connect
/// and yield events.
pub fn new(jid: &str, password: &str, server: &str, port: u16) -> Result<Self, JidParseError> {
let jid = Jid::from_str(jid)?; let jid = Jid::from_str(jid)?;
let password = password.to_owned(); let password = password.to_owned();
let connect = Self::make_connect(jid.clone(), password, server, port); let stream = Self::connect(jid.clone(), password, server, port).await?;
Ok(Component { Ok(Component { jid, stream })
jid,
state: ComponentState::Connecting(Box::new(connect)),
})
} }
fn make_connect( async fn connect(
jid: Jid, jid: Jid,
password: String, password: String,
server: &str, server: &str,
port: u16, port: u16,
) -> impl Future<Item = XMPPStream, Error = Error> { ) -> Result<XMPPStream, Error> {
let jid1 = jid.clone();
let password = password; let password = password;
done(Connecter::from_lookup(server, None, port)) let tcp_stream = connect(server, None, port).await?;
.flatten() let mut xmpp_stream =
.and_then(move |tcp_stream| { xmpp_stream::XMPPStream::start(tcp_stream, jid, NS_JABBER_COMPONENT_ACCEPT.to_owned())
xmpp_stream::XMPPStream::start( .await?;
tcp_stream, auth::auth(&mut xmpp_stream, password).await?;
jid1, Ok(xmpp_stream)
NS_JABBER_COMPONENT_ACCEPT.to_owned(),
)
})
.and_then(move |xmpp_stream| Self::auth(xmpp_stream, password).expect("auth"))
} }
fn auth<S: AsyncRead + AsyncWrite>( /// Send stanza
stream: xmpp_stream::XMPPStream<S>, pub async fn send_stanza(&mut self, stanza: Element) -> Result<(), Error> {
password: String, self.send(stanza).await
) -> Result<ComponentAuth<S>, Error> { }
ComponentAuth::new(stream, password)
/// End connection
pub async fn send_end(&mut self) -> Result<(), Error> {
self.close().await
} }
} }
impl Stream for Component { impl Stream for Component {
type Item = Event; type Item = Element;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
loop {
match Pin::new(&mut self.stream).poll_next(cx) {
Poll::Ready(Some(Ok(Packet::Stanza(stanza)))) => return Poll::Ready(Some(stanza)),
Poll::Ready(Some(Ok(Packet::Text(_)))) => {
// retry
}
Poll::Ready(Some(Ok(_))) =>
// unexpected
{
return Poll::Ready(None)
}
Poll::Ready(Some(Err(_))) => return Poll::Ready(None),
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
}
}
}
}
impl Sink<Element> for Component {
type Error = Error; type Error = Error;
fn poll(&mut self) -> Poll<Option<Self::Item>, Self::Error> { fn start_send(mut self: Pin<&mut Self>, item: Element) -> Result<(), Self::Error> {
let state = replace(&mut self.state, ComponentState::Invalid); Pin::new(&mut self.stream)
match state {
ComponentState::Invalid => Err(Error::InvalidState),
ComponentState::Disconnected => Ok(Async::Ready(None)),
ComponentState::Connecting(mut connect) => match connect.poll() {
Ok(Async::Ready(stream)) => {
self.state = ComponentState::Connected(stream);
Ok(Async::Ready(Some(Event::Online(self.jid.clone()))))
}
Ok(Async::NotReady) => {
self.state = ComponentState::Connecting(connect);
Ok(Async::NotReady)
}
Err(e) => Err(e),
},
ComponentState::Connected(mut stream) => {
// Poll sink
match stream.poll_complete() {
Ok(Async::NotReady) => (),
Ok(Async::Ready(())) => (),
Err(e) => return Err(e)?,
};
// Poll stream
match stream.poll() {
Ok(Async::NotReady) => {
self.state = ComponentState::Connected(stream);
Ok(Async::NotReady)
}
Ok(Async::Ready(None)) => {
// EOF
self.state = ComponentState::Disconnected;
Ok(Async::Ready(Some(Event::Disconnected)))
}
Ok(Async::Ready(Some(Packet::Stanza(stanza)))) => {
self.state = ComponentState::Connected(stream);
Ok(Async::Ready(Some(Event::Stanza(stanza))))
}
Ok(Async::Ready(_)) => {
self.state = ComponentState::Connected(stream);
Ok(Async::NotReady)
}
Err(e) => Err(e)?,
}
}
}
}
}
impl Sink for Component {
type SinkItem = Element;
type SinkError = Error;
fn start_send(&mut self, item: Self::SinkItem) -> StartSend<Self::SinkItem, Self::SinkError> {
match self.state {
ComponentState::Connected(ref mut stream) => match stream
.start_send(Packet::Stanza(item)) .start_send(Packet::Stanza(item))
{ .map_err(|e| e.into())
Ok(AsyncSink::NotReady(Packet::Stanza(stanza))) => Ok(AsyncSink::NotReady(stanza)),
Ok(AsyncSink::NotReady(_)) => {
panic!("Component.start_send with stanza but got something else back")
}
Ok(AsyncSink::Ready) => Ok(AsyncSink::Ready),
Err(e) => Err(e)?,
},
_ => Ok(AsyncSink::NotReady(item)),
}
} }
fn poll_complete(&mut self) -> Poll<(), Self::SinkError> { fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
match &mut self.state { Pin::new(&mut self.stream)
&mut ComponentState::Connected(ref mut stream) => { .poll_ready(cx)
stream.poll_complete().map_err(|e| e.into()) .map_err(|e| e.into())
}
_ => Ok(Async::Ready(())),
} }
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
Pin::new(&mut self.stream)
.poll_flush(cx)
.map_err(|e| e.into())
}
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
Pin::new(&mut self.stream)
.poll_close(cx)
.map_err(|e| e.into())
} }
} }

View file

@ -8,7 +8,7 @@ use trust_dns_proto::error::ProtoError;
use trust_dns_resolver::error::ResolveError; use trust_dns_resolver::error::ResolveError;
use xmpp_parsers::sasl::DefinedCondition as SaslDefinedCondition; use xmpp_parsers::sasl::DefinedCondition as SaslDefinedCondition;
use xmpp_parsers::Error as ParsersError; use xmpp_parsers::{Error as ParsersError, JidParseError};
/// Top-level error type /// Top-level error type
#[derive(Debug)] #[derive(Debug)]
@ -20,6 +20,8 @@ pub enum Error {
/// DNS label conversion error, no details available from module /// DNS label conversion error, no details available from module
/// `idna` /// `idna`
Idna, Idna,
/// Error parsing Jabber-Id
JidParse(JidParseError),
/// Protocol-level error /// Protocol-level error
Protocol(ProtocolError), Protocol(ProtocolError),
/// Authentication error /// Authentication error
@ -38,6 +40,7 @@ impl fmt::Display for Error {
Error::Io(e) => write!(fmt, "IO error: {}", e), Error::Io(e) => write!(fmt, "IO error: {}", e),
Error::Connection(e) => write!(fmt, "connection error: {}", e), Error::Connection(e) => write!(fmt, "connection error: {}", e),
Error::Idna => write!(fmt, "IDNA error"), Error::Idna => write!(fmt, "IDNA error"),
Error::JidParse(e) => write!(fmt, "jid parse error: {}", e),
Error::Protocol(e) => write!(fmt, "protocol error: {}", e), Error::Protocol(e) => write!(fmt, "protocol error: {}", e),
Error::Auth(e) => write!(fmt, "authentication error: {}", e), Error::Auth(e) => write!(fmt, "authentication error: {}", e),
Error::Tls(e) => write!(fmt, "TLS error: {}", e), Error::Tls(e) => write!(fmt, "TLS error: {}", e),
@ -59,6 +62,12 @@ impl From<ConnecterError> for Error {
} }
} }
impl From<JidParseError> for Error {
fn from(e: JidParseError) -> Self {
Error::JidParse(e)
}
}
impl From<ProtocolError> for Error { impl From<ProtocolError> for Error {
fn from(e: ProtocolError) -> Self { fn from(e: ProtocolError) -> Self {
Error::Protocol(e) Error::Protocol(e)

View file

@ -1,3 +1,4 @@
use super::Error;
use xmpp_parsers::{Element, Jid}; use xmpp_parsers::{Element, Jid};
/// High-level event on the Stream implemented by Client and Component /// High-level event on the Stream implemented by Client and Component
@ -6,7 +7,7 @@ pub enum Event {
/// Stream is connected and initialized /// Stream is connected and initialized
Online(Jid), Online(Jid),
/// Stream end /// Stream end
Disconnected, Disconnected(Error),
/// Received stanza/nonza /// Received stanza/nonza
Stanza(Element), Stanza(Element),
} }

View file

@ -1,195 +1,63 @@
use crate::{ConnecterError, Error}; use crate::{ConnecterError, Error};
use futures::{Async, Future, Poll};
use std::cell::RefCell;
use std::collections::BTreeMap;
use std::collections::VecDeque;
use std::io::Error as IoError;
use std::mem;
use std::net::SocketAddr; use std::net::SocketAddr;
use tokio::net::tcp::ConnectFuture;
use tokio::net::TcpStream; use tokio::net::TcpStream;
use trust_dns_resolver::config::LookupIpStrategy; use trust_dns_resolver::{IntoName, TokioAsyncResolver};
use trust_dns_resolver::lookup::SrvLookupFuture;
use trust_dns_resolver::lookup_ip::LookupIpFuture;
use trust_dns_resolver::{AsyncResolver, Background, BackgroundLookup, IntoName, Name};
enum State { async fn connect_to_host(
ResolveSrv(AsyncResolver, BackgroundLookup<SrvLookupFuture>), resolver: &TokioAsyncResolver,
ResolveTarget(AsyncResolver, Background<LookupIpFuture>, u16), host: &str,
Connecting(Option<AsyncResolver>, Vec<RefCell<ConnectFuture>>), port: u16,
Invalid, ) -> Result<TcpStream, Error> {
let ips = resolver
.lookup_ip(host)
.await
.map_err(ConnecterError::Resolve)?;
for ip in ips.iter() {
match TcpStream::connect(&SocketAddr::new(ip, port)).await {
Ok(stream) => return Ok(stream),
Err(_) => {}
}
}
Err(Error::Disconnected)
} }
pub struct Connecter { pub async fn connect(
fallback_port: u16,
srv_domain: Option<Name>,
domain: Name,
state: State,
targets: VecDeque<(Name, u16)>,
error: Option<Error>,
}
fn resolver() -> Result<AsyncResolver, IoError> {
let (config, mut opts) = trust_dns_resolver::system_conf::read_system_conf()?;
opts.ip_strategy = LookupIpStrategy::Ipv4AndIpv6;
let (resolver, resolver_background) = AsyncResolver::new(config, opts);
tokio::runtime::current_thread::spawn(resolver_background);
Ok(resolver)
}
impl Connecter {
pub fn from_lookup(
domain: &str, domain: &str,
srv: Option<&str>, srv: Option<&str>,
fallback_port: u16, fallback_port: u16,
) -> Result<Connecter, Error> { ) -> Result<TcpStream, Error> {
if let Ok(ip) = domain.parse() { if let Ok(ip) = domain.parse() {
// use specified IP address, not domain name, skip the whole dns part return Ok(TcpStream::connect(&SocketAddr::new(ip, fallback_port)).await?);
let connect = RefCell::new(TcpStream::connect(&SocketAddr::new(ip, fallback_port)));
return Ok(Connecter {
fallback_port,
srv_domain: None,
domain: "nohost".into_name().map_err(ConnecterError::Dns)?,
state: State::Connecting(None, vec![connect]),
targets: VecDeque::new(),
error: None,
});
} }
let srv_domain = match srv { let resolver = TokioAsyncResolver::tokio_from_system_conf()
Some(srv) => Some( .await
format!("{}.{}.", srv, domain) .map_err(ConnecterError::Resolve)?;
let srv_records = match srv {
Some(srv) => {
let srv_domain = format!("{}.{}.", srv, domain)
.into_name() .into_name()
.map_err(ConnecterError::Dns)?, .map_err(ConnecterError::Dns)?;
), resolver.srv_lookup(srv_domain).await.ok()
}
None => None, None => None,
}; };
let mut self_ = Connecter { match srv_records {
fallback_port, Some(lookup) => {
srv_domain, // TODO: sort lookup records by priority/weight
domain: domain.into_name().map_err(ConnecterError::Dns)?, for srv in lookup.iter() {
state: State::Invalid, match connect_to_host(&resolver, &srv.target().to_ascii(), srv.port()).await {
targets: VecDeque::new(), Ok(stream) => return Ok(stream),
error: None, Err(_) => {}
}; }
}
let resolver = resolver()?; Err(Error::Disconnected)
// Initialize state
match &self_.srv_domain {
&Some(ref srv_domain) => {
let srv_lookup = resolver.lookup_srv(srv_domain.clone());
self_.state = State::ResolveSrv(resolver, srv_lookup);
} }
None => { None => {
self_.targets = [(self_.domain.clone(), self_.fallback_port)] // SRV lookup error, retry with hostname
.iter() connect_to_host(&resolver, domain, fallback_port).await
.cloned()
.collect();
self_.state = State::Connecting(Some(resolver), vec![]);
}
}
Ok(self_)
}
}
impl Future for Connecter {
type Item = TcpStream;
type Error = Error;
fn poll(&mut self) -> Poll<Self::Item, Self::Error> {
let state = mem::replace(&mut self.state, State::Invalid);
match state {
State::ResolveSrv(resolver, mut srv_lookup) => {
match srv_lookup.poll() {
Ok(Async::NotReady) => {
self.state = State::ResolveSrv(resolver, srv_lookup);
Ok(Async::NotReady)
}
Ok(Async::Ready(srv_result)) => {
let srv_map: BTreeMap<_, _> = srv_result
.iter()
.map(|srv| (srv.priority(), (srv.target().clone(), srv.port())))
.collect();
let targets = srv_map.into_iter().map(|(_, tp)| tp).collect();
self.targets = targets;
self.state = State::Connecting(Some(resolver), vec![]);
self.poll()
}
Err(_) => {
// ignore, fallback
self.targets = [(self.domain.clone(), self.fallback_port)]
.iter()
.cloned()
.collect();
self.state = State::Connecting(Some(resolver), vec![]);
self.poll()
}
}
}
State::Connecting(resolver, mut connects) => {
if resolver.is_some() && connects.len() == 0 && self.targets.len() > 0 {
let resolver = resolver.unwrap();
let (host, port) = self.targets.pop_front().unwrap();
let ip_lookup = resolver.lookup_ip(host);
self.state = State::ResolveTarget(resolver, ip_lookup, port);
self.poll()
} else if connects.len() > 0 {
let mut success = None;
connects.retain(|connect| match connect.borrow_mut().poll() {
Ok(Async::NotReady) => true,
Ok(Async::Ready(connection)) => {
success = Some(connection);
false
}
Err(e) => {
if self.error.is_none() {
self.error = Some(e.into());
}
false
}
});
match success {
Some(connection) => Ok(Async::Ready(connection)),
None => {
self.state = State::Connecting(resolver, connects);
Ok(Async::NotReady)
}
}
} else {
// All targets tried
match self.error.take() {
None => Err(ConnecterError::AllFailed.into()),
Some(e) => Err(e),
}
}
}
State::ResolveTarget(resolver, mut ip_lookup, port) => {
match ip_lookup.poll() {
Ok(Async::NotReady) => {
self.state = State::ResolveTarget(resolver, ip_lookup, port);
Ok(Async::NotReady)
}
Ok(Async::Ready(ip_result)) => {
let connects = ip_result
.iter()
.map(|ip| RefCell::new(TcpStream::connect(&SocketAddr::new(ip, port))))
.collect();
self.state = State::Connecting(Some(resolver), connects);
self.poll()
}
Err(e) => {
if self.error.is_none() {
self.error = Some(ConnecterError::Resolve(e).into());
}
// ignore, next…
self.state = State::Connecting(Some(resolver), vec![]);
self.poll()
}
}
}
_ => panic!(""),
} }
} }
} }

View file

@ -6,10 +6,9 @@ mod starttls;
mod stream_start; mod stream_start;
pub mod xmpp_codec; pub mod xmpp_codec;
pub use crate::xmpp_codec::Packet; pub use crate::xmpp_codec::Packet;
pub mod xmpp_stream;
pub use crate::starttls::StartTlsClient;
mod event; mod event;
mod happy_eyeballs; mod happy_eyeballs;
pub mod xmpp_stream;
pub use crate::event::Event; pub use crate::event::Event;
mod client; mod client;
pub use crate::client::Client; pub use crate::client::Client;

View file

@ -1,114 +1,39 @@
use futures::sink; use futures::{sink::SinkExt, stream::StreamExt};
use futures::stream::Stream;
use futures::{Async, Future, Poll, Sink};
use native_tls::TlsConnector as NativeTlsConnector; use native_tls::TlsConnector as NativeTlsConnector;
use std::mem::replace; use tokio::io::{AsyncRead, AsyncWrite};
use tokio_io::{AsyncRead, AsyncWrite}; use tokio_tls::{TlsConnector, TlsStream};
use tokio_tls::{Connect, TlsConnector, TlsStream}; use xmpp_parsers::Element;
use xmpp_parsers::{Element, Jid};
use crate::xmpp_codec::Packet; use crate::xmpp_codec::Packet;
use crate::xmpp_stream::XMPPStream; use crate::xmpp_stream::XMPPStream;
use crate::Error; use crate::{Error, ProtocolError};
/// XMPP TLS XML namespace /// XMPP TLS XML namespace
pub const NS_XMPP_TLS: &str = "urn:ietf:params:xml:ns:xmpp-tls"; pub const NS_XMPP_TLS: &str = "urn:ietf:params:xml:ns:xmpp-tls";
/// XMPP stream that switches to TLS if available in received features pub async fn starttls<S: AsyncRead + AsyncWrite + Unpin>(
pub struct StartTlsClient<S: AsyncRead + AsyncWrite> { mut xmpp_stream: XMPPStream<S>,
state: StartTlsClientState<S>, ) -> Result<TlsStream<S>, Error> {
jid: Jid,
}
enum StartTlsClientState<S: AsyncRead + AsyncWrite> {
Invalid,
SendStartTls(sink::Send<XMPPStream<S>>),
AwaitProceed(XMPPStream<S>),
StartingTls(Connect<S>),
}
impl<S: AsyncRead + AsyncWrite> StartTlsClient<S> {
/// Waits for <stream:features>
pub fn from_stream(xmpp_stream: XMPPStream<S>) -> Self {
let jid = xmpp_stream.jid.clone();
let nonza = Element::builder("starttls").ns(NS_XMPP_TLS).build(); let nonza = Element::builder("starttls").ns(NS_XMPP_TLS).build();
let packet = Packet::Stanza(nonza); let packet = Packet::Stanza(nonza);
let send = xmpp_stream.send(packet); xmpp_stream.send(packet).await?;
StartTlsClient { loop {
state: StartTlsClientState::SendStartTls(send), match xmpp_stream.next().await {
jid, Some(Ok(Packet::Stanza(ref stanza))) if stanza.name() == "proceed" => break,
Some(Ok(Packet::Text(_))) => {}
Some(Err(e)) => return Err(e.into()),
_ => {
return Err(ProtocolError::NoTls.into());
} }
} }
} }
impl<S: AsyncRead + AsyncWrite> Future for StartTlsClient<S> { let domain = xmpp_stream.jid.clone().domain();
type Item = TlsStream<S>; let stream = xmpp_stream.into_inner();
type Error = Error; let tls_stream = TlsConnector::from(NativeTlsConnector::builder().build().unwrap())
.connect(&domain, stream)
.await?;
fn poll(&mut self) -> Poll<Self::Item, Self::Error> { Ok(tls_stream)
let old_state = replace(&mut self.state, StartTlsClientState::Invalid);
let mut retry = false;
let (new_state, result) = match old_state {
StartTlsClientState::SendStartTls(mut send) => match send.poll() {
Ok(Async::Ready(xmpp_stream)) => {
let new_state = StartTlsClientState::AwaitProceed(xmpp_stream);
retry = true;
(new_state, Ok(Async::NotReady))
}
Ok(Async::NotReady) => {
(StartTlsClientState::SendStartTls(send), Ok(Async::NotReady))
}
Err(e) => (StartTlsClientState::SendStartTls(send), Err(e.into())),
},
StartTlsClientState::AwaitProceed(mut xmpp_stream) => match xmpp_stream.poll() {
Ok(Async::Ready(Some(Packet::Stanza(ref stanza))))
if stanza.name() == "proceed" =>
{
let stream = xmpp_stream.stream.into_inner();
let connect =
TlsConnector::from(NativeTlsConnector::builder().build().unwrap())
.connect(&self.jid.clone().domain(), stream);
let new_state = StartTlsClientState::StartingTls(connect);
retry = true;
(new_state, Ok(Async::NotReady))
}
Ok(Async::Ready(_value)) => {
// println!("StartTlsClient ignore {:?}", _value);
(
StartTlsClientState::AwaitProceed(xmpp_stream),
Ok(Async::NotReady),
)
}
Ok(_) => (
StartTlsClientState::AwaitProceed(xmpp_stream),
Ok(Async::NotReady),
),
Err(e) => (
StartTlsClientState::AwaitProceed(xmpp_stream),
Err(Error::Protocol(e.into())),
),
},
StartTlsClientState::StartingTls(mut connect) => match connect.poll() {
Ok(Async::Ready(tls_stream)) => {
(StartTlsClientState::Invalid, Ok(Async::Ready(tls_stream)))
}
Ok(Async::NotReady) => (
StartTlsClientState::StartingTls(connect),
Ok(Async::NotReady),
),
Err(e) => (StartTlsClientState::Invalid, Err(e.into())),
},
StartTlsClientState::Invalid => unreachable!(),
};
self.state = new_state;
if retry {
self.poll()
} else {
result
}
}
} }

View file

@ -1,7 +1,7 @@
use futures::{sink, Async, Future, Poll, Sink, Stream}; use futures::{sink::SinkExt, stream::StreamExt};
use std::mem::replace; use std::marker::Unpin;
use tokio_codec::Framed; use tokio::io::{AsyncRead, AsyncWrite};
use tokio_io::{AsyncRead, AsyncWrite}; use tokio_util::codec::Framed;
use xmpp_parsers::{Element, Jid}; use xmpp_parsers::{Element, Jid};
use crate::xmpp_codec::{Packet, XMPPCodec}; use crate::xmpp_codec::{Packet, XMPPCodec};
@ -10,21 +10,11 @@ use crate::{Error, ProtocolError};
const NS_XMPP_STREAM: &str = "http://etherx.jabber.org/streams"; const NS_XMPP_STREAM: &str = "http://etherx.jabber.org/streams";
pub struct StreamStart<S: AsyncWrite> { pub async fn start<S: AsyncRead + AsyncWrite + Unpin>(
state: StreamStartState<S>, mut stream: Framed<S, XMPPCodec>,
jid: Jid, jid: Jid,
ns: String, ns: String,
} ) -> Result<XMPPStream<S>, Error> {
enum StreamStartState<S: AsyncWrite> {
SendStart(sink::Send<Framed<S, XMPPCodec>>),
RecvStart(Framed<S, XMPPCodec>),
RecvFeatures(Framed<S, XMPPCodec>, String),
Invalid,
}
impl<S: AsyncWrite> StreamStart<S> {
pub fn from_stream(stream: Framed<S, XMPPCodec>, jid: Jid, ns: String) -> Self {
let attrs = [ let attrs = [
("to".to_owned(), jid.clone().domain()), ("to".to_owned(), jid.clone().domain()),
("version".to_owned(), "1.0".to_owned()), ("version".to_owned(), "1.0".to_owned()),
@ -34,92 +24,52 @@ impl<S: AsyncWrite> StreamStart<S> {
.iter() .iter()
.cloned() .cloned()
.collect(); .collect();
let send = stream.send(Packet::StreamStart(attrs)); stream.send(Packet::StreamStart(attrs)).await?;
StreamStart { let stream_attrs;
state: StreamStartState::SendStart(send), loop {
jid, match stream.next().await {
ns, Some(Ok(Packet::StreamStart(attrs))) => {
stream_attrs = attrs;
break;
} }
Some(Ok(_)) => {}
Some(Err(e)) => return Err(e.into()),
None => return Err(Error::Disconnected),
} }
} }
impl<S: AsyncRead + AsyncWrite> Future for StreamStart<S> {
type Item = XMPPStream<S>;
type Error = Error;
fn poll(&mut self) -> Poll<Self::Item, Self::Error> {
let old_state = replace(&mut self.state, StreamStartState::Invalid);
let mut retry = false;
let (new_state, result) = match old_state {
StreamStartState::SendStart(mut send) => match send.poll() {
Ok(Async::Ready(stream)) => {
retry = true;
(StreamStartState::RecvStart(stream), Ok(Async::NotReady))
}
Ok(Async::NotReady) => (StreamStartState::SendStart(send), Ok(Async::NotReady)),
Err(e) => (StreamStartState::Invalid, Err(e.into())),
},
StreamStartState::RecvStart(mut stream) => match stream.poll() {
Ok(Async::Ready(Some(Packet::StreamStart(stream_attrs)))) => {
let stream_ns = stream_attrs let stream_ns = stream_attrs
.get("xmlns") .get("xmlns")
.ok_or(ProtocolError::NoStreamNamespace)? .ok_or(ProtocolError::NoStreamNamespace)?
.clone(); .clone();
if self.ns == "jabber:client" { let stream_id = stream_attrs
retry = true;
// TODO: skip RecvFeatures for version < 1.0
(
StreamStartState::RecvFeatures(stream, stream_ns),
Ok(Async::NotReady),
)
} else {
let id = stream_attrs
.get("id") .get("id")
.ok_or(ProtocolError::NoStreamId)? .ok_or(ProtocolError::NoStreamId)?
.clone(); .clone();
let stream = if stream_ns == "jabber:client" && stream_attrs.get("version").is_some() {
let stream_features;
loop {
match stream.next().await {
Some(Ok(Packet::Stanza(stanza))) if stanza.is("features", NS_XMPP_STREAM) => {
stream_features = stanza;
break;
}
Some(Ok(_)) => {}
Some(Err(e)) => return Err(e.into()),
None => return Err(Error::Disconnected),
}
}
XMPPStream::new(jid, stream, ns, stream_id, stream_features)
} else {
// FIXME: huge hack, shouldnt be an element! // FIXME: huge hack, shouldnt be an element!
let stream = XMPPStream::new( XMPPStream::new(
self.jid.clone(), jid,
stream, stream,
self.ns.clone(), ns,
Element::builder(id).build(), stream_id.clone(),
); Element::builder(stream_id).build(),
(StreamStartState::Invalid, Ok(Async::Ready(stream)))
}
}
Ok(Async::Ready(_)) => return Err(ProtocolError::InvalidToken.into()),
Ok(Async::NotReady) => (StreamStartState::RecvStart(stream), Ok(Async::NotReady)),
Err(e) => return Err(ProtocolError::from(e).into()),
},
StreamStartState::RecvFeatures(mut stream, stream_ns) => match stream.poll() {
Ok(Async::Ready(Some(Packet::Stanza(stanza)))) => {
if stanza.is("features", NS_XMPP_STREAM) {
let stream =
XMPPStream::new(self.jid.clone(), stream, self.ns.clone(), stanza);
(StreamStartState::Invalid, Ok(Async::Ready(stream)))
} else {
(
StreamStartState::RecvFeatures(stream, stream_ns),
Ok(Async::NotReady),
) )
}
}
Ok(Async::Ready(_)) | Ok(Async::NotReady) => (
StreamStartState::RecvFeatures(stream, stream_ns),
Ok(Async::NotReady),
),
Err(e) => return Err(ProtocolError::from(e).into()),
},
StreamStartState::Invalid => unreachable!(),
}; };
Ok(stream)
self.state = new_state;
if retry {
self.poll()
} else {
result
}
}
} }

View file

@ -5,16 +5,16 @@ use bytes::{BufMut, BytesMut};
use log::debug; use log::debug;
use std; use std;
use std::borrow::Cow; use std::borrow::Cow;
use std::cell::RefCell;
use std::collections::vec_deque::VecDeque; use std::collections::vec_deque::VecDeque;
use std::collections::HashMap; use std::collections::HashMap;
use std::default::Default; use std::default::Default;
use std::fmt::Write; use std::fmt::Write;
use std::io; use std::io;
use std::iter::FromIterator; use std::iter::FromIterator;
use std::rc::Rc;
use std::str::from_utf8; use std::str::from_utf8;
use tokio_codec::{Decoder, Encoder}; use std::sync::Arc;
use std::sync::Mutex;
use tokio_util::codec::{Decoder, Encoder};
use xml5ever::buffer_queue::BufferQueue; use xml5ever::buffer_queue::BufferQueue;
use xml5ever::interface::Attribute; use xml5ever::interface::Attribute;
use xml5ever::tokenizer::{Tag, TagKind, Token, TokenSink, XmlTokenizer}; use xml5ever::tokenizer::{Tag, TagKind, Token, TokenSink, XmlTokenizer};
@ -38,14 +38,14 @@ type QueueItem = Result<Packet, ParserError>;
/// Parser state /// Parser state
struct ParserSink { struct ParserSink {
// Ready stanzas, shared with XMPPCodec // Ready stanzas, shared with XMPPCodec
queue: Rc<RefCell<VecDeque<QueueItem>>>, queue: Arc<Mutex<VecDeque<QueueItem>>>,
// Parsing stack // Parsing stack
stack: Vec<Element>, stack: Vec<Element>,
ns_stack: Vec<HashMap<Option<String>, String>>, ns_stack: Vec<HashMap<Option<String>, String>>,
} }
impl ParserSink { impl ParserSink {
pub fn new(queue: Rc<RefCell<VecDeque<QueueItem>>>) -> Self { pub fn new(queue: Arc<Mutex<VecDeque<QueueItem>>>) -> Self {
ParserSink { ParserSink {
queue, queue,
stack: vec![], stack: vec![],
@ -54,11 +54,11 @@ impl ParserSink {
} }
fn push_queue(&self, pkt: Packet) { fn push_queue(&self, pkt: Packet) {
self.queue.borrow_mut().push_back(Ok(pkt)); self.queue.lock().unwrap().push_back(Ok(pkt));
} }
fn push_queue_error(&self, e: ParserError) { fn push_queue_error(&self, e: ParserError) {
self.queue.borrow_mut().push_back(Err(e)); self.queue.lock().unwrap().push_back(Err(e));
} }
/// Lookup XML namespace declaration for given prefix (or no prefix) /// Lookup XML namespace declaration for given prefix (or no prefix)
@ -169,7 +169,6 @@ impl TokenSink for ParserSink {
}, },
Token::EOFToken => self.push_queue(Packet::StreamEnd), Token::EOFToken => self.push_queue(Packet::StreamEnd),
Token::ParseError(s) => { Token::ParseError(s) => {
// println!("ParseError: {:?}", s);
self.push_queue_error(ParserError::Parse(ParseError(s))); self.push_queue_error(ParserError::Parse(ParseError(s)));
} }
_ => (), _ => (),
@ -190,13 +189,13 @@ pub struct XMPPCodec {
// TODO: optimize using tendrils? // TODO: optimize using tendrils?
buf: Vec<u8>, buf: Vec<u8>,
/// Shared with ParserSink /// Shared with ParserSink
queue: Rc<RefCell<VecDeque<QueueItem>>>, queue: Arc<Mutex<VecDeque<QueueItem>>>,
} }
impl XMPPCodec { impl XMPPCodec {
/// Constructor /// Constructor
pub fn new() -> Self { pub fn new() -> Self {
let queue = Rc::new(RefCell::new(VecDeque::new())); let queue = Arc::new(Mutex::new(VecDeque::new()));
let sink = ParserSink::new(queue.clone()); let sink = ParserSink::new(queue.clone());
// TODO: configure parser? // TODO: configure parser?
let parser = XmlTokenizer::new(sink, Default::default()); let parser = XmlTokenizer::new(sink, Default::default());
@ -222,10 +221,10 @@ impl Decoder for XMPPCodec {
fn decode(&mut self, buf: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> { fn decode(&mut self, buf: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
let buf1: Box<dyn AsRef<[u8]>> = if !self.buf.is_empty() && !buf.is_empty() { let buf1: Box<dyn AsRef<[u8]>> = if !self.buf.is_empty() && !buf.is_empty() {
let mut prefix = std::mem::replace(&mut self.buf, vec![]); let mut prefix = std::mem::replace(&mut self.buf, vec![]);
prefix.extend_from_slice(buf.take().as_ref()); prefix.extend_from_slice(&buf.split_to(buf.len()));
Box::new(prefix) Box::new(prefix)
} else { } else {
Box::new(buf.take()) Box::new(buf.split_to(buf.len()))
}; };
let buf1 = buf1.as_ref().as_ref(); let buf1 = buf1.as_ref().as_ref();
match from_utf8(buf1) { match from_utf8(buf1) {
@ -258,7 +257,7 @@ impl Decoder for XMPPCodec {
} }
} }
match self.queue.borrow_mut().pop_front() { match self.queue.lock().unwrap().pop_front() {
None => Ok(None), None => Ok(None),
Some(result) => result.map(|pkt| Some(pkt)), Some(result) => result.map(|pkt| Some(pkt)),
} }
@ -372,7 +371,7 @@ mod tests {
fn test_stream_start() { fn test_stream_start() {
let mut c = XMPPCodec::new(); let mut c = XMPPCodec::new();
let mut b = BytesMut::with_capacity(1024); let mut b = BytesMut::with_capacity(1024);
b.put(r"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>"); b.put_slice(b"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::StreamStart(_))) => true, Ok(Some(Packet::StreamStart(_))) => true,
@ -384,14 +383,14 @@ mod tests {
fn test_stream_end() { fn test_stream_end() {
let mut c = XMPPCodec::new(); let mut c = XMPPCodec::new();
let mut b = BytesMut::with_capacity(1024); let mut b = BytesMut::with_capacity(1024);
b.put(r"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>"); b.put_slice(b"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::StreamStart(_))) => true, Ok(Some(Packet::StreamStart(_))) => true,
_ => false, _ => false,
}); });
b.clear(); b.clear();
b.put(r"</stream:stream>"); b.put_slice(b"</stream:stream>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::StreamEnd)) => true, Ok(Some(Packet::StreamEnd)) => true,
@ -403,7 +402,7 @@ mod tests {
fn test_truncated_stanza() { fn test_truncated_stanza() {
let mut c = XMPPCodec::new(); let mut c = XMPPCodec::new();
let mut b = BytesMut::with_capacity(1024); let mut b = BytesMut::with_capacity(1024);
b.put(r"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>"); b.put_slice(b"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::StreamStart(_))) => true, Ok(Some(Packet::StreamStart(_))) => true,
@ -411,7 +410,7 @@ mod tests {
}); });
b.clear(); b.clear();
b.put(r"<test>ß</test"); b.put_slice("<test>ß</test".as_bytes());
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(None) => true, Ok(None) => true,
@ -419,7 +418,7 @@ mod tests {
}); });
b.clear(); b.clear();
b.put(r">"); b.put_slice(b">");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::Stanza(ref el))) if el.name() == "test" && el.text() == "ß" => true, Ok(Some(Packet::Stanza(ref el))) if el.name() == "test" && el.text() == "ß" => true,
@ -431,7 +430,7 @@ mod tests {
fn test_truncated_utf8() { fn test_truncated_utf8() {
let mut c = XMPPCodec::new(); let mut c = XMPPCodec::new();
let mut b = BytesMut::with_capacity(1024); let mut b = BytesMut::with_capacity(1024);
b.put(r"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>"); b.put_slice(b"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::StreamStart(_))) => true, Ok(Some(Packet::StreamStart(_))) => true,
@ -460,7 +459,7 @@ mod tests {
fn test_atrribute_prefix() { fn test_atrribute_prefix() {
let mut c = XMPPCodec::new(); let mut c = XMPPCodec::new();
let mut b = BytesMut::with_capacity(1024); let mut b = BytesMut::with_capacity(1024);
b.put(r"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>"); b.put_slice(b"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::StreamStart(_))) => true, Ok(Some(Packet::StreamStart(_))) => true,
@ -468,7 +467,7 @@ mod tests {
}); });
b.clear(); b.clear();
b.put(r"<status xml:lang='en'>Test status</status>"); b.put_slice(b"<status xml:lang='en'>Test status</status>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::Stanza(ref el))) Ok(Some(Packet::Stanza(ref el)))
@ -483,10 +482,10 @@ mod tests {
/// By default, encode() only get's a BytesMut that has 8kb space reserved. /// By default, encode() only get's a BytesMut that has 8kb space reserved.
#[test] #[test]
fn test_large_stanza() { fn test_large_stanza() {
use futures::{Future, Sink}; use futures::{executor::block_on, sink::SinkExt};
use std::io::Cursor; use std::io::Cursor;
use tokio_codec::FramedWrite; use tokio_util::codec::FramedWrite;
let framed = FramedWrite::new(Cursor::new(vec![]), XMPPCodec::new()); let mut framed = FramedWrite::new(Cursor::new(vec![]), XMPPCodec::new());
let mut text = "".to_owned(); let mut text = "".to_owned();
for _ in 0..2usize.pow(15) { for _ in 0..2usize.pow(15) {
text = text + "A"; text = text + "A";
@ -494,7 +493,7 @@ mod tests {
let stanza = Element::builder("message") let stanza = Element::builder("message")
.append(Element::builder("body").append(text.as_ref()).build()) .append(Element::builder("body").append(text.as_ref()).build())
.build(); .build();
let framed = framed.send(Packet::Stanza(stanza)).wait().expect("send"); block_on(framed.send(Packet::Stanza(stanza))).expect("send");
assert_eq!( assert_eq!(
framed.get_ref().get_ref(), framed.get_ref().get_ref(),
&("<message><body>".to_owned() + &text + "</body></message>").as_bytes() &("<message><body>".to_owned() + &text + "</body></message>").as_bytes()
@ -505,7 +504,7 @@ mod tests {
fn test_cut_out_stanza() { fn test_cut_out_stanza() {
let mut c = XMPPCodec::new(); let mut c = XMPPCodec::new();
let mut b = BytesMut::with_capacity(1024); let mut b = BytesMut::with_capacity(1024);
b.put(r"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>"); b.put_slice(b"<?xml version='1.0'?><stream:stream xmlns:stream='http://etherx.jabber.org/streams' version='1.0' xmlns='jabber:client'>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::StreamStart(_))) => true, Ok(Some(Packet::StreamStart(_))) => true,
@ -513,8 +512,8 @@ mod tests {
}); });
b.clear(); b.clear();
b.put(r"<message "); b.put_slice(b"<message ");
b.put(r"type='chat'><body>Foo</body></message>"); b.put_slice(b"type='chat'><body>Foo</body></message>");
let r = c.decode(&mut b); let r = c.decode(&mut b);
assert!(match r { assert!(match r {
Ok(Some(Packet::Stanza(_))) => true, Ok(Some(Packet::Stanza(_))) => true,

View file

@ -1,23 +1,28 @@
//! `XMPPStream` is the common container for all XMPP network connections //! `XMPPStream` is the common container for all XMPP network connections
use futures::sink::Send; use futures::sink::Send;
use futures::{Poll, Sink, StartSend, Stream}; use futures::{sink::SinkExt, task::Poll, Sink, Stream};
use tokio_codec::Framed; use std::ops::DerefMut;
use tokio_io::{AsyncRead, AsyncWrite}; use std::pin::Pin;
use std::sync::Mutex;
use std::task::Context;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio_util::codec::Framed;
use xmpp_parsers::{Element, Jid}; use xmpp_parsers::{Element, Jid};
use crate::stream_start::StreamStart; use crate::stream_start;
use crate::xmpp_codec::{Packet, XMPPCodec}; use crate::xmpp_codec::{Packet, XMPPCodec};
use crate::Error;
/// <stream:stream> namespace /// <stream:stream> namespace
pub const NS_XMPP_STREAM: &str = "http://etherx.jabber.org/streams"; pub const NS_XMPP_STREAM: &str = "http://etherx.jabber.org/streams";
/// Wraps a `stream` /// Wraps a `stream`
pub struct XMPPStream<S> { pub struct XMPPStream<S: AsyncRead + AsyncWrite + Unpin> {
/// The local Jabber-Id /// The local Jabber-Id
pub jid: Jid, pub jid: Jid,
/// Codec instance /// Codec instance
pub stream: Framed<S, XMPPCodec>, pub stream: Mutex<Framed<S, XMPPCodec>>,
/// `<stream:features/>` for XMPP version 1.0 /// `<stream:features/>` for XMPP version 1.0
pub stream_features: Element, pub stream_features: Element,
/// Root namespace /// Root namespace
@ -25,68 +30,94 @@ pub struct XMPPStream<S> {
/// This is different for either c2s, s2s, or component /// This is different for either c2s, s2s, or component
/// connections. /// connections.
pub ns: String, pub ns: String,
/// Stream `id` attribute
pub id: String,
} }
impl<S: AsyncRead + AsyncWrite> XMPPStream<S> { // // TODO: fix this hack
// unsafe impl<S: AsyncRead + AsyncWrite + Unpin> core::marker::Send for XMPPStream<S> {}
// unsafe impl<S: AsyncRead + AsyncWrite + Unpin> Sync for XMPPStream<S> {}
impl<S: AsyncRead + AsyncWrite + Unpin> XMPPStream<S> {
/// Constructor /// Constructor
pub fn new( pub fn new(
jid: Jid, jid: Jid,
stream: Framed<S, XMPPCodec>, stream: Framed<S, XMPPCodec>,
ns: String, ns: String,
id: String,
stream_features: Element, stream_features: Element,
) -> Self { ) -> Self {
XMPPStream { XMPPStream {
jid, jid,
stream, stream: Mutex::new(stream),
stream_features, stream_features,
ns, ns,
id,
} }
} }
/// Send a `<stream:stream>` start tag /// Send a `<stream:stream>` start tag
pub fn start(stream: S, jid: Jid, ns: String) -> StreamStart<S> { pub async fn start<'a>(stream: S, jid: Jid, ns: String) -> Result<Self, Error> {
let xmpp_stream = Framed::new(stream, XMPPCodec::new()); let xmpp_stream = Framed::new(stream, XMPPCodec::new());
StreamStart::from_stream(xmpp_stream, jid, ns) stream_start::start(xmpp_stream, jid, ns).await
} }
/// Unwraps the inner stream /// Unwraps the inner stream
// TODO: use this everywhere
pub fn into_inner(self) -> S { pub fn into_inner(self) -> S {
self.stream.into_inner() self.stream.into_inner().unwrap().into_inner()
} }
/// Re-run `start()` /// Re-run `start()`
pub fn restart(self) -> StreamStart<S> { pub async fn restart<'a>(self) -> Result<Self, Error> {
Self::start(self.stream.into_inner(), self.jid, self.ns) let stream = self.stream.into_inner().unwrap().into_inner();
Self::start(stream, self.jid, self.ns).await
} }
} }
impl<S: AsyncWrite> XMPPStream<S> { impl<S: AsyncRead + AsyncWrite + Unpin> XMPPStream<S> {
/// Convenience method /// Convenience method
pub fn send_stanza<E: Into<Element>>(self, e: E) -> Send<Self> { pub fn send_stanza<E: Into<Element>>(&mut self, e: E) -> Send<Self, Packet> {
self.send(Packet::Stanza(e.into())) self.send(Packet::Stanza(e.into()))
} }
} }
/// Proxy to self.stream /// Proxy to self.stream
impl<S: AsyncWrite> Sink for XMPPStream<S> { impl<S: AsyncRead + AsyncWrite + Unpin> Sink<Packet> for XMPPStream<S> {
type SinkItem = <Framed<S, XMPPCodec> as Sink>::SinkItem; type Error = crate::Error;
type SinkError = <Framed<S, XMPPCodec> as Sink>::SinkError;
fn start_send(&mut self, item: Self::SinkItem) -> StartSend<Self::SinkItem, Self::SinkError> { fn poll_ready(self: Pin<&mut Self>, _ctx: &mut Context) -> Poll<Result<(), Self::Error>> {
self.stream.start_send(item) // Pin::new(&mut self.stream).poll_ready(ctx)
// .map_err(|e| e.into())
Poll::Ready(Ok(()))
} }
fn poll_complete(&mut self) -> Poll<(), Self::SinkError> { fn start_send(mut self: Pin<&mut Self>, item: Packet) -> Result<(), Self::Error> {
self.stream.poll_complete() Pin::new(&mut self.stream.lock().unwrap().deref_mut())
.start_send(item)
.map_err(|e| e.into())
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
Pin::new(&mut self.stream.lock().unwrap().deref_mut())
.poll_flush(cx)
.map_err(|e| e.into())
}
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
Pin::new(&mut self.stream.lock().unwrap().deref_mut())
.poll_close(cx)
.map_err(|e| e.into())
} }
} }
/// Proxy to self.stream /// Proxy to self.stream
impl<S: AsyncRead> Stream for XMPPStream<S> { impl<S: AsyncRead + AsyncWrite + Unpin> Stream for XMPPStream<S> {
type Item = <Framed<S, XMPPCodec> as Stream>::Item; type Item = Result<Packet, crate::Error>;
type Error = <Framed<S, XMPPCodec> as Stream>::Error;
fn poll(&mut self) -> Poll<Option<Self::Item>, Self::Error> { fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
self.stream.poll() Pin::new(&mut self.stream.lock().unwrap().deref_mut())
.poll_next(cx)
.map(|result| result.map(|result| result.map_err(|e| e.into())))
} }
} }