use std::io::prelude::*; use std::io; use std::net::{SocketAddr, TcpStream}; use xml::reader::{EventReader, XmlEvent as XmlReaderEvent}; use xml::writer::{EventWriter, XmlEvent as XmlWriterEvent}; use std::sync::{Arc, Mutex}; use ns; use error::Error; use openssl::ssl::{SslMethod, SslConnectorBuilder, SslStream}; pub trait Transport { fn write_event<'a, E: Into>>(&mut self, event: E) -> Result<(), Error>; fn read_event(&mut self) -> Result; } struct LockedWrite(Arc>); impl io::Write for LockedWrite { fn write(&mut self, buf: &[u8]) -> io::Result { let mut inner = self.0.lock().unwrap(); // TODO: make safer inner.write(buf) } fn flush(&mut self) -> io::Result<()> { let mut inner = self.0.lock().unwrap(); // TODO: make safer inner.flush() } } struct LockedRead(Arc>); impl io::Read for LockedRead { fn read(&mut self, buf: &mut [u8]) -> io::Result { let mut inner = self.0.lock().unwrap(); // TODO: make safer inner.read(buf) } } pub struct SslTransport { inner: Arc>>, // TODO: this feels rather ugly reader: EventReader>>, // TODO: especially feels ugly because // this read would keep the lock // held very long (potentially) writer: EventWriter>>, } impl Transport for SslTransport { fn write_event<'a, E: Into>>(&mut self, event: E) -> Result<(), Error> { Ok(self.writer.write(event)?) } fn read_event(&mut self) -> Result { Ok(self.reader.next()?) } } impl SslTransport { pub fn connect(host: &str, port: u16) -> Result { // TODO: very quick and dirty, blame starttls let mut stream = TcpStream::connect((host, port))?; write!(stream, "" , ns::CLIENT, ns::STREAM, host)?; write!(stream, "" , ns::TLS)?; let mut parser = EventReader::new(stream); loop { // TODO: possibly a timeout? match parser.next()? { XmlReaderEvent::StartElement { name, namespace, .. } => { if let Some(ns) = name.namespace { if ns == ns::TLS && name.local_name == "proceed" { break; } else if ns == ns::STREAM && name.local_name == "error" { return Err(Error::StreamError); } } }, _ => {}, } } let stream = parser.into_inner(); let ssl_connector = SslConnectorBuilder::new(SslMethod::tls())?.build(); let ssl_stream = Arc::new(Mutex::new(ssl_connector.connect(host, stream)?)); let reader = EventReader::new(LockedRead(ssl_stream.clone())); let writer = EventWriter::new(LockedWrite(ssl_stream.clone())); Ok(SslTransport { inner: ssl_stream, reader: reader, writer: writer, }) } }