use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use rsurl::{WsReader, WsWriter};
use spotproto::Packet;
use crate::api;
use crate::client::Inner;
use crate::error::{Error, Result};
use crate::transport::{self, Incoming};
const DIAL_TIMEOUT: Duration = Duration::from_secs(30);
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(300);
const STEADY_READ_TIMEOUT: Duration = Duration::from_secs(120);
const WRITE_POLL: Duration = Duration::from_millis(500);
pub(crate) fn main_thread(inner: Arc<Inner>) {
inner.logf(format_args!("client entering main thread"));
if let Err(e) = run_connect(&inner) {
inner.logf(format_args!("failed to perform initial connection: {e}"));
}
let mut tick = 0u32;
while !inner.is_closed() {
std::thread::sleep(Duration::from_secs(1));
tick += 1;
if tick < 30 {
continue;
}
tick = 0;
let cnt = inner.conn_cnt.load(Ordering::Relaxed);
if cnt < inner.min_conn.load(Ordering::Relaxed) {
if let Err(e) = run_connect(&inner) {
inner.logf(format_args!("failed to perform connection: {e}"));
}
}
}
}
fn run_connect(inner: &Arc<Inner>) -> Result<()> {
let (mut hosts, mut min_conn) = api::get_hosts()?;
if min_conn == 0 {
min_conn = hosts.len() as u32;
}
hosts.truncate(10);
inner.min_conn.store(min_conn, Ordering::Relaxed);
for host in hosts {
if inner.is_closed() {
break;
}
let registered = inner.hosts.lock().unwrap().insert(host.clone());
if !registered {
continue;
}
inner.logf(format_args!("connecting to host: {host}"));
let inner = inner.clone();
std::thread::spawn(move || conn_thread(inner, host));
std::thread::sleep(Duration::from_secs(2));
}
Ok(())
}
fn conn_thread(inner: Arc<Inner>, host: String) {
inner.conn_cnt.fetch_add(1, Ordering::Relaxed);
let mut fail_giveup = 0;
while !inner.is_closed() {
match transport::connect(&host, "/_websocket", DIAL_TIMEOUT) {
Err(e) => {
inner.logf(format_args!("failed to connect to server: {e}"));
fail_giveup += 1;
if fail_giveup > 10 {
break;
}
std::thread::sleep(Duration::from_secs(2));
}
Ok((reader, writer)) => {
fail_giveup = 0;
if let Err(e) = handle(&inner, reader, writer) {
inner.logf(format_args!(
"error during communications with server: {e}"
));
}
}
}
}
inner.hosts.lock().unwrap().remove(&host);
inner.conn_cnt.fetch_sub(1, Ordering::Relaxed);
}
fn handle(inner: &Arc<Inner>, mut reader: WsReader, writer: WsWriter) -> Result<()> {
let writer = Arc::new(Mutex::new(writer));
let shutdown = reader.shutdown_handle();
reader
.set_read_timeout(Some(HANDSHAKE_TIMEOUT))
.map_err(|e| Error::Ws(e.to_string()))?;
handshake(inner, &mut reader, &writer)?;
let steady = shutdown.as_ref().map_or(Some(STEADY_READ_TIMEOUT), |_| None);
reader
.set_read_timeout(steady)
.map_err(|e| Error::Ws(e.to_string()))?;
inner.online_incr();
let _online_guard = OnlineGuard { inner };
let dead = Arc::new(AtomicBool::new(false));
let writer_dead = dead.clone();
let writer_inner = inner.clone();
let writer_w = writer.clone();
let writer_shutdown = shutdown.clone();
let writer_thread = std::thread::spawn(move || {
loop {
if writer_dead.load(Ordering::Relaxed) {
return;
}
if writer_inner.is_closed() {
transport::close(&writer_w);
if let Some(s) = &writer_shutdown {
let _ = s.shutdown();
}
return;
}
let Some(msg) = writer_inner.wrq.pop_timeout(WRITE_POLL) else {
continue;
};
let buf = match Packet::Message(msg.clone()).encode() {
Ok(buf) => buf,
Err(e) => {
writer_inner.logf(format_args!("failed to encode message: {e}"));
continue;
}
};
if transport::send(&writer_w, &buf).is_err() {
writer_inner.wrq.push_front(msg);
return;
}
}
});
let res = read_loop(inner, &mut reader, &writer);
dead.store(true, Ordering::Relaxed);
let _ = writer_thread.join();
res
}
fn handshake(inner: &Arc<Inner>, reader: &mut WsReader, writer: &Mutex<WsWriter>) -> Result<()> {
loop {
let data = match transport::recv_packet(reader)? {
Incoming::Packet(data) => data,
Incoming::Closed => {
return Err(Error::Ws("connection closed during handshake".into()))
}
Incoming::Timeout => return Err(Error::Ws("handshake timed out".into())),
};
match spotproto::parse(&data, true)? {
Packet::HandshakeRequest(req) => {
if req.ready {
inner.logf(format_args!(
"authentication done, connected as c.{}",
req.client_id
));
return Ok(());
}
respond_handshake(inner, &req, writer)?;
}
other => {
inner.logf(format_args!("unsupported handshake packet type {other:?}"));
}
}
}
}
fn respond_handshake(
inner: &Arc<Inner>,
req: &spotproto::HandshakeRequest,
writer: &Mutex<WsWriter>,
) -> Result<()> {
if let Some(groups) = &req.groups {
if let Err(e) = inner.handle_groups(groups) {
inner.logf(format_args!("failed to update groups: {e}"));
}
}
let mut res = req.respond(inner.signer())?;
res.id = inner.id_bin();
transport::send(writer, &Packet::HandshakeResponse(res).encode()?)
}
fn read_loop(inner: &Arc<Inner>, reader: &mut WsReader, writer: &Mutex<WsWriter>) -> Result<()> {
loop {
let data = match transport::recv_packet(reader)? {
Incoming::Packet(data) => data,
Incoming::Closed => return Ok(()), Incoming::Timeout => {
if inner.is_closed() {
return Ok(());
}
continue;
}
};
match spotproto::parse(&data, true)? {
Packet::HandshakeRequest(req) => {
if req.ready {
continue;
}
respond_handshake(inner, &req, writer)?;
}
Packet::Message(msg) => inner.route_message(msg),
other => {
inner.logf(format_args!("unsupported packet type {other:?}"));
}
}
}
}
struct OnlineGuard<'a> {
inner: &'a Arc<Inner>,
}
impl Drop for OnlineGuard<'_> {
fn drop(&mut self) {
self.inner.online_decr();
}
}