websockets.rs 3.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798
  1. use std::{
  2. net::{TcpStream, ToSocketAddrs},
  3. pin::Pin,
  4. task::{Context, Poll},
  5. };
  6. use async_tungstenite::WebSocketStream;
  7. use futures::sink::Sink;
  8. use futures_rustls::{client::TlsStream, rustls::ServerName, TlsConnector};
  9. use smol::{prelude::*, Async};
  10. use tungstenite::{handshake::client::Response, Message};
  11. use url::Url;
  12. use crate::{Error, Result as DrkResult};
  13. #[allow(clippy::large_enum_variant)]
  14. pub enum WsStream {
  15. Tcp(WebSocketStream<Async<TcpStream>>),
  16. Tls(WebSocketStream<TlsStream<Async<TcpStream>>>),
  17. }
  18. impl Sink<Message> for WsStream {
  19. type Error = tungstenite::Error;
  20. fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
  21. match &mut *self {
  22. WsStream::Tcp(s) => Pin::new(s).poll_ready(cx),
  23. WsStream::Tls(s) => Pin::new(s).poll_ready(cx),
  24. }
  25. }
  26. fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
  27. match &mut *self {
  28. WsStream::Tcp(s) => Pin::new(s).start_send(item),
  29. WsStream::Tls(s) => Pin::new(s).start_send(item),
  30. }
  31. }
  32. fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
  33. match &mut *self {
  34. WsStream::Tcp(s) => Pin::new(s).poll_flush(cx),
  35. WsStream::Tls(s) => Pin::new(s).poll_flush(cx),
  36. }
  37. }
  38. fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
  39. match &mut *self {
  40. WsStream::Tcp(s) => Pin::new(s).poll_close(cx),
  41. WsStream::Tls(s) => Pin::new(s).poll_close(cx),
  42. }
  43. }
  44. }
  45. impl Stream for WsStream {
  46. type Item = tungstenite::Result<Message>;
  47. fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
  48. match &mut *self {
  49. WsStream::Tcp(s) => Pin::new(s).poll_next(cx),
  50. WsStream::Tls(s) => Pin::new(s).poll_next(cx),
  51. }
  52. }
  53. }
  54. /// Connects to a WebSocket address (optionally secured by TLS).
  55. pub async fn connect(addr: &str, tls: TlsConnector) -> DrkResult<(WsStream, Response)> {
  56. let url = Url::parse(addr)?;
  57. let host = url
  58. .host_str()
  59. .ok_or_else(|| Error::UrlParse(format!("Missing host in {}", url)))?
  60. .to_string();
  61. let port = url
  62. .port_or_known_default()
  63. .ok_or_else(|| Error::UrlParse(format!("Missing port in {}", url)))?;
  64. let socket_addr = {
  65. let host = host.clone();
  66. smol::unblock(move || (host.as_str(), port).to_socket_addrs())
  67. .await?
  68. .next()
  69. .ok_or(Error::NoUrlFound)?
  70. };
  71. match url.scheme() {
  72. "ws" => {
  73. let stream = Async::<TcpStream>::connect(socket_addr).await?;
  74. let (stream, resp) = async_tungstenite::client_async(addr, stream).await?;
  75. Ok((WsStream::Tcp(stream), resp))
  76. }
  77. "wss" => {
  78. let stream = Async::<TcpStream>::connect(socket_addr).await?;
  79. let stream = tls.connect(ServerName::try_from(host.as_str())?, stream).await?;
  80. let (stream, resp) = async_tungstenite::client_async(addr, stream).await?;
  81. Ok((WsStream::Tls(stream), resp))
  82. }
  83. scheme => Err(Error::UrlParse(format!("Invalid url scheme `{}`, in `{}`", scheme, url))),
  84. }
  85. }