use std::fmt;
use std::sync::{Arc, Mutex, PoisonError};
use std::time::Duration;
use net_backend_protocol::auth::{AuthSession, LoginRequest, LogoutRequest, RefreshRequest, RegisterRequest, SteamLoginRequest, TokenPair};
use net_backend_protocol::{codes, GetServerInfo, HttpCall, ServerInfo};
use tokio::sync::watch;
use tokio::time::Instant;
use crate::http::{Answer, BaseUrl, Http, Outgoing};
use crate::session::{Session, TokenUpdates, UNCERTAIN_RETRY_WINDOW};
use crate::Error;
pub const DEFAULT_TIMEOUT: Duration = Duration::from_secs(15);
pub const DEFAULT_MAX_RESPONSE_BYTES: usize = 10 * 1024 * 1024;
pub const MAX_TIMEOUT: Duration = Duration::from_secs(3600);
#[derive(Clone)]
pub struct ClientBuilder {
url: String,
timeout: Duration,
max_response_bytes: usize,
refresh_margin: Duration,
allow_insecure_http: bool,
tokens: Option<TokenPair>,
}
impl fmt::Debug for ClientBuilder {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ClientBuilder")
.field("url", &self.url)
.field("timeout", &self.timeout)
.field("max_response_bytes", &self.max_response_bytes)
.field("refresh_margin", &self.refresh_margin)
.field("allow_insecure_http", &self.allow_insecure_http)
.field("tokens", &self.tokens.is_some())
.finish()
}
}
impl ClientBuilder {
fn new(url: &str) -> Self {
Self {
url: url.to_string(),
timeout: DEFAULT_TIMEOUT,
max_response_bytes: DEFAULT_MAX_RESPONSE_BYTES,
refresh_margin: Duration::from_secs(net_backend_protocol::auth::ACCESS_TOKEN_REFRESH_MARGIN_SECS),
allow_insecure_http: false,
tokens: None,
}
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout.clamp(Duration::from_millis(1), MAX_TIMEOUT);
self
}
pub fn max_response_bytes(mut self, bytes: usize) -> Self {
self.max_response_bytes = bytes.max(1024);
self
}
pub fn refresh_margin(mut self, margin: Duration) -> Self {
self.refresh_margin = margin.min(MAX_TIMEOUT);
self
}
pub fn allow_insecure_http(mut self, allow: bool) -> Self {
self.allow_insecure_http = allow;
self
}
pub fn tokens(mut self, tokens: TokenPair) -> Self {
self.tokens = Some(tokens);
self
}
pub fn build(self) -> Result<Client, Error> {
let base = BaseUrl::parse(&self.url)?;
if base.scheme == crate::http::Scheme::Http && !base.is_loopback() && !self.allow_insecure_http {
return Err(Error::invalid(format!(
"plain http:// to `{}` is refused: use https://, a loopback host, or ClientBuilder::allow_insecure_http(true)",
base.host
)));
}
let http = Http::new(base, self.max_response_bytes)?;
let session = Session::new(self.refresh_margin);
if let Some(tokens) = self.tokens {
session.set(tokens, None);
}
Ok(Client { inner: Arc::new(Inner { http, session, timeout: self.timeout, refreshing: Mutex::new(None) }) })
}
}
type RefreshOutcome = Option<Result<TokenPair, Error>>;
pub(crate) struct Inner {
pub(crate) http: Http,
pub(crate) session: Session,
pub(crate) timeout: Duration,
refreshing: Mutex<Option<watch::Receiver<RefreshOutcome>>>,
}
#[derive(Clone)]
pub struct Client {
pub(crate) inner: Arc<Inner>,
}
impl fmt::Debug for Client {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Client").field("server", &self.inner.http.base.url("")).field("logged_in", &self.inner.session.tokens().is_some()).finish()
}
}
impl Client {
pub fn new(url: &str) -> Result<Self, Error> {
Self::builder(url).build()
}
pub fn builder(url: &str) -> ClientBuilder {
ClientBuilder::new(url)
}
fn deadline(&self) -> Instant {
let now = Instant::now();
now.checked_add(self.inner.timeout).unwrap_or(now)
}
pub async fn call<C: HttpCall>(&self, call: &C) -> Result<C::Response, Error> {
crate::runtime::current()?;
let deadline = self.deadline();
let out = Outgoing::for_call(call)?;
if !C::ROUTE.auth {
return self.inner.http.send(&out, None, deadline).await?.decode();
}
self.send_authed(&out, deadline).await?.decode()
}
async fn send_authed(&self, out: &Outgoing, deadline: Instant) -> Result<Answer, Error> {
let (token, generation) = self.access_token(deadline).await?;
let answer = self.inner.http.send(out, Some(token.expose()), deadline).await?;
if answer.status != 401 {
return Ok(answer);
}
let current = self.inner.session.access().map(|(_, g, _)| g);
if current == Some(generation) {
self.refresh_shared(deadline).await?;
}
let (token, _) = self.access_token(deadline).await?;
self.inner.http.send(out, Some(token.expose()), deadline).await
}
pub(crate) async fn access_token(&self, deadline: Instant) -> Result<(net_backend_protocol::AccessToken, u64), Error> {
let (token, generation, wants_refresh) = self.inner.session.access().ok_or(Error::NotLoggedIn)?;
if !wants_refresh {
return Ok((token, generation));
}
match self.refresh_shared(deadline).await {
Ok(pair) => {
let generation = self.inner.session.access().map_or(generation, |(_, g, _)| g);
Ok((pair.access_token, generation))
}
Err(error @ (Error::SessionEnded { .. } | Error::NotLoggedIn)) => Err(error),
Err(error) if self.inner.session.expired() => Err(error),
Err(error) => {
tracing::debug!("net_backend_client: refresh failed ({error}); the current access token is still valid");
Ok((token, generation))
}
}
}
pub(crate) async fn refresh_shared(&self, deadline: Instant) -> Result<TokenPair, Error> {
let handle = crate::runtime::current()?;
let mut receiver = {
let mut slot = self.inner.refreshing.lock().unwrap_or_else(PoisonError::into_inner);
match slot.as_ref() {
Some(receiver) if receiver.has_changed().is_ok() => receiver.clone(),
_ => {
let (sender, receiver) = watch::channel::<RefreshOutcome>(None);
*slot = Some(receiver.clone());
let inner = Arc::clone(&self.inner);
handle.spawn(async move {
let outcome = refresh_task(&inner).await;
*inner.refreshing.lock().unwrap_or_else(PoisonError::into_inner) = None;
let _ = sender.send(Some(outcome));
});
receiver
}
}
};
loop {
if let Some(outcome) = receiver.borrow_and_update().clone() {
return outcome;
}
match tokio::time::timeout_at(deadline, receiver.changed()).await {
Ok(Ok(())) => {}
Ok(Err(_)) => return Err(Error::Shutdown),
Err(_) => return Err(Error::timeout("not sent: the token refresh did not finish before the deadline", Some(false))),
}
}
}
pub async fn info(&self) -> Result<ServerInfo, Error> {
self.call(&GetServerInfo::new()).await
}
pub async fn register(&self, request: RegisterRequest) -> Result<AuthSession, Error> {
self.session_call(&request, false).await
}
pub async fn login(&self, request: LoginRequest) -> Result<AuthSession, Error> {
self.session_call(&request, false).await
}
pub async fn login_steam(&self, request: SteamLoginRequest) -> Result<AuthSession, Error> {
self.session_call(&request, false).await
}
pub async fn link_steam(&self, request: SteamLoginRequest) -> Result<AuthSession, Error> {
self.session_call(&request, true).await
}
async fn session_call<C: HttpCall<Response = AuthSession>>(&self, call: &C, authed: bool) -> Result<AuthSession, Error> {
crate::runtime::current()?;
let deadline = self.deadline();
let out = Outgoing::for_call(call)?;
let answer = if authed { self.send_authed(&out, deadline).await? } else { self.inner.http.send(&out, None, deadline).await? };
let session: AuthSession = answer.decode()?;
self.inner.session.set(session.tokens.clone(), answer.server_now);
Ok(session)
}
pub async fn refresh(&self) -> Result<TokenPair, Error> {
crate::runtime::current()?;
if self.inner.session.tokens().is_none() {
return Err(Error::NotLoggedIn);
}
self.refresh_shared(self.deadline()).await
}
pub async fn logout(&self) -> Result<(), Error> {
self.logout_with(LogoutRequest::this_session()).await
}
pub async fn logout_everywhere(&self) -> Result<(), Error> {
self.logout_with(LogoutRequest::everywhere()).await
}
async fn logout_with(&self, request: LogoutRequest) -> Result<(), Error> {
crate::runtime::current()?;
let deadline = self.deadline();
let (tokens, generation) = {
let state = self.inner.session.lock();
(state.tokens.clone().ok_or(Error::NotLoggedIn)?, state.generation)
};
let request = request.with_refresh_token(tokens.refresh_token.clone());
let out = Outgoing::for_call(&request)?;
let bearer = (!self.inner.session.expired()).then(|| tokens.access_token.expose().to_string());
let answer = self.inner.http.send(&out, bearer.as_deref(), deadline).await?;
match answer.decode::<net_backend_protocol::Ack>() {
Ok(_) => {
self.inner.session.clear_if(generation);
Ok(())
}
Err(error) if error.status() == Some(401) => {
self.inner.session.clear_if(generation);
Ok(())
}
Err(error) => Err(error),
}
}
pub fn resume(&self, tokens: TokenPair) {
self.inner.session.set(tokens, None);
}
pub fn forget_session(&self) {
self.inner.session.clear();
}
pub fn tokens(&self) -> Option<TokenPair> {
self.inner.session.tokens()
}
pub fn is_logged_in(&self) -> bool {
self.inner.session.tokens().is_some()
}
pub fn token_updates(&self) -> TokenUpdates {
self.inner.session.subscribe()
}
pub fn server_url(&self) -> String {
self.inner.http.base.url("")
}
}
async fn refresh_task(inner: &Arc<Inner>) -> Result<TokenPair, Error> {
let mut attempt: u32 = 0;
loop {
let (refresh_token, generation) = {
let state = inner.session.lock();
match state.tokens.as_ref() {
Some(tokens) => (tokens.refresh_token.clone(), state.generation),
None => return Err(Error::NotLoggedIn),
}
};
let now = Instant::now();
let deadline = now.checked_add(inner.timeout).unwrap_or(now);
let out = Outgoing::for_call(&RefreshRequest::new(refresh_token))?;
let result = match inner.http.send(&out, None, deadline).await {
Ok(answer) => answer.decode::<TokenPair>().map(|pair| (pair, answer.server_now)),
Err(error) => Err(error),
};
match result {
Ok((pair, server_now)) => {
if inner.session.lock().generation != generation {
return inner.session.tokens().ok_or(Error::NotLoggedIn);
}
inner.session.set(pair.clone(), server_now);
return Ok(pair);
}
Err(error) if error.ends_session() => {
let code = error.code().unwrap_or(codes::UNAUTHORIZED).to_string();
tracing::info!("net_backend_client: the session ended: the server refused the refresh ({code})");
inner.session.clear_if(generation);
return Err(Error::SessionEnded { code });
}
Err(error) => {
let maybe_sent = matches!(error, Error::Network { sent: None, .. } | Error::Timeout { sent: None, .. });
if !maybe_sent {
return Err(error);
}
let since = {
let mut state = inner.session.lock();
if state.generation != generation {
return Err(error);
}
*state.uncertain_since.get_or_insert(now)
};
attempt = attempt.saturating_add(1);
let pause = Duration::from_millis(250u64.saturating_mul(1 << attempt.min(4)));
if attempt > 4 || Instant::now().saturating_duration_since(since).saturating_add(pause) >= UNCERTAIN_RETRY_WINDOW {
return Err(error);
}
tokio::time::sleep(pause).await;
}
}
}
}