use parking_lot::Mutex;
use reqwest::{
Client as ReqwestClient, RequestBuilder, Response, Url,
header::{ACCEPT, AUTHORIZATION, HeaderMap, RETRY_AFTER},
};
use serde::{Serialize, de::DeserializeOwned};
use std::{
borrow::Borrow,
fmt::Debug,
sync::{Arc, LazyLock},
time::Duration,
};
use crate::{
endpoints::{
Endpoint, EventsEndpoint, GetEndpoint, LocationEndpoint,
MysteryMasterEndpoint, NationsEndpoint, NearbyEndpoint, OnlineEndpoint,
PlayersEndpoint, QuartersEndpoint, QueryEndpoint, ServerEndpoint,
ShopEndpoint, TownsEndpoint,
events::{Event, EventKind, parse_sse_event},
location::{LocationInfo, LocationQuery},
mystery_master::MysteryMasterEntry,
nations::Nation,
nearby::NearbyQuery,
online::OnlinePlayers,
players::Player,
quarters::Quarter,
server::Server,
shop::{Shop, ShopQuery, Shops},
towns::Town,
},
error::{Error, Result, snippet_around},
model::{
named_uuid::NamedUuid,
query::{Query, SimpleQuery, UuidQuery},
},
retry_strategy::{
JitteredBackoff, RetryContext, RetryStrategy, parse_retry_after,
},
};
pub const DEFAULT_BASE_URL: &str = "https://api.earthmc.net/v4/";
const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(10);
static DEFAULT_HTTP_CLIENT: LazyLock<ReqwestClient> = LazyLock::new(|| {
reqwest::ClientBuilder::new()
.connect_timeout(Duration::from_secs(10))
.user_agent(format!("pkg:cargo/earthmc@{}", env!("CARGO_PKG_VERSION")))
.build()
.expect("Failed to initialize HTTP client")
});
pub trait AuthenticationState: Clone + Debug + Send + Sync + 'static {
fn apply_auth(&self, request: RequestBuilder) -> RequestBuilder;
}
#[derive(Clone, Copy, Debug, Default)]
pub struct Unauthenticated;
impl AuthenticationState for Unauthenticated {
fn apply_auth(&self, request: RequestBuilder) -> RequestBuilder {
request
}
}
#[derive(Clone)]
pub struct Authenticated {
api_key: Arc<str>,
}
impl Authenticated {
fn new(api_key: impl AsRef<str>) -> Self {
Self {
api_key: Arc::from(api_key.as_ref()),
}
}
fn api_key(&self) -> &str {
&self.api_key
}
}
impl Debug for Authenticated {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Authenticated")
.field("api_key", &"<redacted>")
.finish()
}
}
impl AuthenticationState for Authenticated {
fn apply_auth(&self, request: RequestBuilder) -> RequestBuilder {
request
}
}
#[derive(Clone)]
pub struct Client<A = Unauthenticated> {
reqwest_client: ReqwestClient,
pub base_url: Url,
retry_strategy: Arc<Mutex<dyn RetryStrategy>>,
auth: A,
}
impl<A: Debug> Debug for Client<A> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Client")
.field("reqwest_client", &self.reqwest_client)
.field("base_url", &self.base_url)
.field("auth", &self.auth)
.finish()
}
}
impl Default for Client<Unauthenticated> {
fn default() -> Self {
ClientBuilder::default().build()
}
}
#[derive(Clone)]
pub struct ClientBuilder<A = Unauthenticated> {
reqwest_client: Option<ReqwestClient>,
base_url: Option<Url>,
retry_strategy: Option<Arc<Mutex<dyn RetryStrategy>>>,
auth: A,
}
impl Default for ClientBuilder<Unauthenticated> {
fn default() -> Self {
Self {
reqwest_client: None,
base_url: None,
retry_strategy: None,
auth: Unauthenticated,
}
}
}
impl<A> ClientBuilder<A>
where
A: AuthenticationState,
{
pub fn reqwest_client(
&mut self,
reqwest_client: ReqwestClient,
) -> &mut Self {
self.reqwest_client = Some(reqwest_client);
self
}
pub fn base_url(&mut self, base_url: Url) -> &mut Self {
self.base_url = Some(base_url);
self
}
pub fn retry_strategy(
&mut self,
retry_strategy: Arc<Mutex<dyn RetryStrategy>>,
) -> &mut Self {
self.retry_strategy = Some(retry_strategy);
self
}
pub fn build(&self) -> Client<A> {
Client {
reqwest_client: self
.reqwest_client
.clone()
.unwrap_or_else(|| DEFAULT_HTTP_CLIENT.clone()),
base_url: self
.base_url
.clone()
.unwrap_or_else(|| DEFAULT_BASE_URL.parse().unwrap()),
retry_strategy: self.retry_strategy.clone().unwrap_or_else(|| {
Arc::new(Mutex::new(JitteredBackoff::default()))
}),
auth: self.auth.clone(),
}
}
}
impl ClientBuilder<Unauthenticated> {
pub fn api_key(
&mut self,
api_key: impl AsRef<str>,
) -> ClientBuilder<Authenticated> {
ClientBuilder {
reqwest_client: self.reqwest_client.clone(),
base_url: self.base_url.clone(),
retry_strategy: self.retry_strategy.clone(),
auth: Authenticated::new(api_key),
}
}
}
impl Client<Unauthenticated> {
pub fn with_api_key(
&self,
api_key: impl AsRef<str>,
) -> Client<Authenticated> {
Client {
reqwest_client: self.reqwest_client.clone(),
base_url: self.base_url.clone(),
retry_strategy: self.retry_strategy.clone(),
auth: Authenticated::new(api_key),
}
}
}
impl Client<Authenticated> {
pub fn without_api_key(&self) -> Client<Unauthenticated> {
Client {
reqwest_client: self.reqwest_client.clone(),
base_url: self.base_url.clone(),
retry_strategy: self.retry_strategy.clone(),
auth: Unauthenticated,
}
}
pub async fn shops<Q: Into<UuidQuery>>(
&self,
query: Q,
key: Option<String>,
) -> Result<Vec<Vec<Shop>>> {
let key = key.unwrap_or_else(|| self.auth.api_key().to_owned());
let response = self
.post::<Vec<Shops>, ShopQuery>(
ShopEndpoint::PATH,
ShopQuery::new(query.into(), key),
)
.await?;
Ok(response.into_iter().map(Shops::into_vec).collect())
}
pub async fn events<I>(&self, listen: I) -> Result<EventStream>
where
I: IntoIterator,
I::Item: Borrow<EventKind>,
{
let mut url = self
.base_url
.join(EventsEndpoint::PATH)
.expect("Failed to construct request URL");
let listen = listen
.into_iter()
.map(|kind| kind.borrow().as_str())
.collect::<Vec<_>>()
.join(",");
url.set_query(Some(&format!("listen={listen}")));
let response = self
.reqwest_client
.get(url)
.header(AUTHORIZATION, format!("Bearer {}", self.auth.api_key()))
.header(ACCEPT, "text/event-stream")
.send()
.await?
.error_for_status()?;
Ok(EventStream {
response,
buffer: String::new(),
})
}
}
pub struct EventStream {
response: Response,
buffer: String,
}
impl EventStream {
pub async fn next(&mut self) -> Result<Option<Event>> {
loop {
while let Some((start, len)) = find_sse_separator(&self.buffer) {
let frame = self.buffer[..start].to_owned();
self.buffer.drain(..start + len);
if let Some(event) = parse_sse_event(&frame)? {
return Ok(Some(event));
}
}
let Some(chunk) = self.response.chunk().await? else {
return Ok(None);
};
self.buffer.push_str(&String::from_utf8_lossy(&chunk));
}
}
}
fn find_sse_separator(buffer: &str) -> Option<(usize, usize)> {
match (buffer.find("\r\n\r\n"), buffer.find("\n\n")) {
(Some(crlf), Some(lf)) if crlf < lf => Some((crlf, 4)),
(Some(_), Some(lf)) => Some((lf, 2)),
(Some(crlf), None) => Some((crlf, 4)),
(None, Some(lf)) => Some((lf, 2)),
(None, None) => None,
}
}
impl<A> Client<A>
where
A: AuthenticationState,
{
fn retry_delay(
&self,
attempt: usize,
retry_after: Option<Duration>,
) -> Option<Duration> {
let mut strat = self.retry_strategy.lock();
strat.should_retry_after(
RetryContext::new(attempt).with_retry_after(retry_after),
)
}
fn retry_after(headers: &HeaderMap) -> Option<Duration> {
headers
.get(RETRY_AFTER)
.and_then(|value| value.to_str().ok())
.and_then(parse_retry_after)
}
async fn get<T>(&self, path: &str) -> Result<T>
where
T: DeserializeOwned,
{
let url = self
.base_url
.join(path)
.expect("Failed to construct request URL");
let mut num_retries = 0;
loop {
let request = self
.auth
.apply_auth(self.reqwest_client.get(url.clone()))
.timeout(DEFAULT_REQUEST_TIMEOUT);
let attempt = request.send().await;
match attempt {
Ok(response) => {
let retry_after = Self::retry_after(response.headers());
match response.error_for_status() {
Ok(success) => {
let parsed = success.json::<T>().await?;
return Ok(parsed);
}
Err(e) => {
let err: Error = e.into();
let delay_opt =
self.retry_delay(num_retries, retry_after);
if let Some(delay) = delay_opt {
tokio::time::sleep(delay).await;
num_retries += 1;
continue;
}
return Err(err);
}
}
}
Err(e) => {
let err: Error = e.into();
let delay_opt = self.retry_delay(num_retries, None);
if let Some(delay) = delay_opt {
tokio::time::sleep(delay).await;
num_retries += 1;
continue;
}
return Err(err);
}
}
}
}
async fn post<T, B>(&self, path: &str, body: B) -> Result<T>
where
T: DeserializeOwned,
B: Serialize + Sized,
{
let url = self
.base_url
.join(path)
.expect("Failed to construct request URL");
let mut num_retries = 0;
loop {
let request = self
.auth
.apply_auth(self.reqwest_client.post(url.clone()).json(&body))
.timeout(DEFAULT_REQUEST_TIMEOUT);
let attempt = request.send().await;
match attempt {
Ok(response) => {
let retry_after = Self::retry_after(response.headers());
match response.error_for_status() {
Ok(ok_response) => {
let text = ok_response.text().await?;
match serde_json::from_str::<T>(&text) {
Ok(parsed) => return Ok(parsed),
Err(de_err) => {
let snippet = snippet_around(
&text,
de_err.line(),
de_err.column(),
10,
);
return Err(
Error::DeserializationWithSnippet {
source: de_err,
snippet,
},
);
}
}
}
Err(e) => {
let err: Error = e.into();
let delay_opt =
self.retry_delay(num_retries, retry_after);
if let Some(delay) = delay_opt {
tokio::time::sleep(delay).await;
num_retries += 1;
continue;
}
return Err(err);
}
}
}
Err(e) => {
let err: Error = e.into();
let delay_opt = self.retry_delay(num_retries, None);
if let Some(delay) = delay_opt {
tokio::time::sleep(delay).await;
num_retries += 1;
continue;
}
return Err(err);
}
}
}
}
async fn get_endpoint<E>(&self) -> Result<E::Response>
where
E: GetEndpoint,
E::Response: DeserializeOwned,
{
self.get::<E::Response>(E::PATH).await
}
async fn query_endpoint<E, Q>(&self, query: Q) -> Result<E::Response>
where
E: QueryEndpoint,
Q: Into<E::Query>,
E::Query: Serialize,
E::Response: DeserializeOwned,
{
self.post::<E::Response, Query<E::Query>>(E::PATH, query.into().into())
.await
}
pub async fn server(&self) -> Result<Server> {
self.get_endpoint::<ServerEndpoint>().await
}
pub async fn all_towns(&self) -> Result<Vec<NamedUuid>> {
self.get_endpoint::<TownsEndpoint>().await
}
pub async fn towns<Q: Into<SimpleQuery>>(
&self,
query: Q,
) -> Result<Vec<Town>> {
self.query_endpoint::<TownsEndpoint, Q>(query).await
}
pub async fn all_nations(&self) -> Result<Vec<NamedUuid>> {
self.get_endpoint::<NationsEndpoint>().await
}
pub async fn nations<Q: Into<SimpleQuery>>(
&self,
query: Q,
) -> Result<Vec<Nation>> {
self.query_endpoint::<NationsEndpoint, Q>(query).await
}
pub async fn nearby<Q: Into<NearbyQuery>>(
&self,
query: Q,
) -> Result<Vec<Vec<NamedUuid>>> {
self.query_endpoint::<NearbyEndpoint, Q>(query).await
}
pub async fn all_players(&self) -> Result<Vec<NamedUuid>> {
self.get_endpoint::<PlayersEndpoint>().await
}
pub async fn players<Q: Into<SimpleQuery>>(
&self,
query: Q,
) -> Result<Vec<Player>> {
self.query_endpoint::<PlayersEndpoint, Q>(query).await
}
pub async fn all_quarters(&self) -> Result<Vec<NamedUuid>> {
self.get_endpoint::<QuartersEndpoint>().await
}
pub async fn quarters<Q: Into<UuidQuery>>(
&self,
query: Q,
) -> Result<Vec<Quarter>> {
self.query_endpoint::<QuartersEndpoint, Q>(query).await
}
pub async fn locations<Q: Into<LocationQuery>>(
&self,
query: Q,
) -> Result<Vec<LocationInfo>> {
self.query_endpoint::<LocationEndpoint, Q>(query).await
}
pub async fn mystery_masters(&self) -> Result<Vec<MysteryMasterEntry>> {
self.get_endpoint::<MysteryMasterEndpoint>().await
}
pub async fn online_players(&self) -> Result<OnlinePlayers> {
self.get_endpoint::<OnlineEndpoint>().await
}
}
#[cfg(test)]
mod tests {
use std::borrow::Borrow;
use super::{ClientBuilder, EventKind};
use tokio::{
io::{AsyncReadExt as _, AsyncWriteExt as _},
net::TcpListener,
};
use crate::endpoints::events::Event;
#[test]
fn authenticated_debug_redacts_api_key() {
let client = ClientBuilder::default().api_key("secret").build();
let debug = format!("{client:?}");
assert!(debug.contains("<redacted>"));
assert!(!debug.contains("secret"));
}
#[tokio::test]
async fn events_sends_listen_query_and_authorization_header() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buffer = vec![0; 4096];
let bytes_read = socket.read(&mut buffer).await.unwrap();
let request = String::from_utf8_lossy(&buffer[..bytes_read]);
let lower_request = request.to_lowercase();
assert!(request.starts_with(
"GET /v4/events?listen=NewDay,TownDeleted HTTP/1.1\r\n"
));
assert!(lower_request.contains("authorization: bearer secret"));
assert!(lower_request.contains("accept: text/event-stream"));
socket
.write_all(
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\n\r\nevent: NewDay\ndata: {\"timestamp\":1774520132}\n\n",
)
.await
.unwrap();
});
let mut stream = ClientBuilder::default()
.base_url(
format!("http://{address}/v4/")
.parse()
.expect("valid base URL"),
)
.api_key("secret")
.build()
.events([EventKind::NewDay, EventKind::TownDeleted])
.await
.unwrap();
assert!(matches!(
stream.next().await.unwrap(),
Some(Event::NewDay(_))
));
server.await.unwrap();
}
#[test]
fn events_accepts_event_kind_slice() {
fn accepts<I>(listen: I)
where
I: IntoIterator,
I::Item: std::borrow::Borrow<EventKind>,
{
let _ = listen
.into_iter()
.map(|kind| kind.borrow().as_str())
.collect::<Vec<_>>();
}
accepts(EventKind::all());
}
}