websockets.rs 3.0 KB

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