1use std::sync::Arc;
6
7use parking_lot::Mutex;
8use reqwest::StatusCode;
9use serde_json::Value;
10
11use crate::config::Config;
12
13#[derive(Clone)]
18pub struct GraphQLClient {
19 url: String,
20 http: reqwest::blocking::Client,
21 session: Option<Arc<Session>>,
22}
23
24struct Session {
27 tokens: Mutex<Tokens>,
28 on_rotate: Box<dyn Fn(&str) + Send + Sync>,
31}
32
33struct Tokens {
34 access: Option<String>,
35 refresh: String,
36 refused: Option<String>,
39}
40
41impl GraphQLClient {
42 pub fn new(server_url: &str) -> Self {
43 let url = format!("{}/graphql", server_url.trim_end_matches('/'));
44 Self {
45 url,
46 http: reqwest::blocking::Client::builder()
47 .timeout(std::time::Duration::from_secs(30))
48 .build()
49 .expect("failed to build HTTP client"),
50 session: None,
51 }
52 }
53
54 pub fn from_config(server_url: &str) -> Self {
57 let client = Self::new(server_url);
58 let cfg = Config::load().unwrap_or_default();
59 if cfg.auth.refresh_token.is_empty()
60 || cfg.auth.server.trim_end_matches('/') != client.server_url()
61 {
62 return client;
63 }
64 client.with_session(cfg.auth.refresh_token, |token| {
65 if let Err(e) = Config::persist(|cfg| cfg.auth.refresh_token = token.to_owned()) {
66 log::warn!("could not store the rotated refresh token: {e}");
67 }
68 })
69 }
70
71 pub fn with_session(
74 mut self,
75 refresh_token: impl Into<String>,
76 on_rotate: impl Fn(&str) + Send + Sync + 'static,
77 ) -> Self {
78 self.session = Some(Arc::new(Session {
79 tokens: Mutex::new(Tokens {
80 access: None,
81 refresh: refresh_token.into(),
82 refused: None,
83 }),
84 on_rotate: Box::new(on_rotate),
85 }));
86 self
87 }
88
89 pub fn has_session(&self) -> bool {
91 self.session.is_some()
92 }
93
94 pub fn execute(&self, query: &str, variables: Option<Value>) -> Result<Value, GraphQLError> {
99 let mut body = serde_json::json!({ "query": query });
100 if let Some(vars) = variables {
101 body["variables"] = vars;
102 }
103
104 let access = self
105 .session
106 .as_ref()
107 .and_then(|s| s.tokens.lock().access.clone());
108 let mut resp = self.post(&body, access.as_deref())?;
109 if resp.status() == StatusCode::UNAUTHORIZED
110 && let Some(session) = &self.session
111 {
112 let fresh = self.refresh(session, access.as_deref())?;
113 resp = self.post(&body, Some(&fresh))?;
114 }
115
116 let status = resp.status();
117 if status == StatusCode::UNAUTHORIZED {
118 return Err(GraphQLError::Unauthorized(reason(resp)));
119 }
120 if !status.is_success() {
121 return Err(GraphQLError::Http(format!("{status}: {}", reason(resp))));
122 }
123
124 let resp: Value = resp.json().map_err(|e| GraphQLError::Http(e.to_string()))?;
125
126 if let Some(errors) = resp.get("errors")
127 && let Some(arr) = errors.as_array()
128 && !arr.is_empty()
129 {
130 let msg = arr[0]
131 .get("message")
132 .and_then(|m| m.as_str())
133 .unwrap_or("unknown error");
134 return Err(GraphQLError::Query(msg.to_string()));
135 }
136
137 Ok(resp.get("data").cloned().unwrap_or(Value::Null))
138 }
139
140 fn post(
141 &self,
142 body: &Value,
143 access: Option<&str>,
144 ) -> Result<reqwest::blocking::Response, GraphQLError> {
145 let mut req = self.http.post(&self.url).json(body);
146 if let Some(token) = access {
147 req = req.bearer_auth(token);
148 }
149 req.send().map_err(|e| GraphQLError::Http(e.to_string()))
150 }
151
152 fn refresh(&self, session: &Session, stale: Option<&str>) -> Result<String, GraphQLError> {
159 let mut tokens = session.tokens.lock();
160 if let Some(access) = &tokens.access
161 && Some(access.as_str()) != stale
162 {
163 return Ok(access.clone());
164 }
165 if let Some(reason) = &tokens.refused {
166 return Err(GraphQLError::Unauthorized(reason.clone()));
167 }
168
169 let resp = self
170 .http
171 .post(format!("{}/auth/refresh", self.server_url()))
172 .json(&serde_json::json!({ "refresh_token": tokens.refresh }))
173 .send()
174 .map_err(|e| GraphQLError::Http(e.to_string()))?;
175 if resp.status() == StatusCode::UNAUTHORIZED {
176 let refused = format!("the stored sign-in was refused: {}", reason(resp));
177 tokens.access = None;
178 tokens.refused = Some(refused.clone());
179 return Err(GraphQLError::Unauthorized(refused));
180 }
181 if !resp.status().is_success() {
182 let status = resp.status();
183 return Err(GraphQLError::Http(format!(
184 "/auth/refresh {status}: {}",
185 reason(resp)
186 )));
187 }
188
189 let body: Value = resp.json().map_err(|e| GraphQLError::Http(e.to_string()))?;
190 let (Some(access), Some(refresh)) = (
191 body["access_token"].as_str(),
192 body["refresh_token"].as_str(),
193 ) else {
194 return Err(GraphQLError::Http(
195 "malformed /auth/refresh response".into(),
196 ));
197 };
198 tokens.access = Some(access.to_owned());
199 tokens.refresh = refresh.to_owned();
200 (session.on_rotate)(refresh);
201 Ok(access.to_owned())
202 }
203
204 pub fn now_playing(&self) -> Result<NowPlaying, GraphQLError> {
209 let data = self.execute(
210 "{ nowPlaying { state positionMs durationMs queueItemId \
211 track { trackId title artist album codec sampleRate bitDepth bitrateKbps channels durationMs } } }",
212 None,
213 )?;
214 let np = &data["nowPlaying"];
215 Ok(NowPlaying {
216 state: np["state"].as_str().unwrap_or("STOPPED").to_string(),
217 position_ms: np["positionMs"].as_u64().unwrap_or(0),
218 duration_ms: np["durationMs"].as_u64(),
219 queue_item_id: np["queueItemId"].as_str().map(String::from),
220 track: np.get("track").and_then(|t| {
221 if t.is_null() {
222 return None;
223 }
224 Some(NowPlayingTrack {
225 track_id: t["trackId"].as_str().map(String::from),
226 title: t["title"].as_str().unwrap_or("").to_string(),
227 artist: t["artist"].as_str().unwrap_or("").to_string(),
228 album: t["album"].as_str().unwrap_or("").to_string(),
229 codec: t["codec"].as_str().unwrap_or("").to_string(),
230 sample_rate: t["sampleRate"].as_u64().unwrap_or(0) as u32,
231 bit_depth: t["bitDepth"].as_u64().map(|v| v as u16),
232 bitrate_kbps: t["bitrateKbps"].as_u64().map(|v| v as u32),
233 channels: t["channels"].as_u64().unwrap_or(0) as u16,
234 duration_ms: t["durationMs"].as_u64().unwrap_or(0),
235 })
236 }),
237 })
238 }
239
240 pub fn queue(&self) -> Result<Vec<QueueEntry>, GraphQLError> {
241 let data = self.execute(
242 "{ queue { queueItemId trackId title artist album codec trackNumber disc durationMs isCurrent } }",
243 None,
244 )?;
245 let entries = data["queue"]
246 .as_array()
247 .map(|arr| {
248 arr.iter()
249 .map(|e| QueueEntry {
250 queue_item_id: e["queueItemId"].as_str().unwrap_or("").to_string(),
251 track_id: e["trackId"].as_str().map(String::from),
252 title: e["title"].as_str().unwrap_or("").to_string(),
253 artist: e["artist"].as_str().unwrap_or("").to_string(),
254 album: e["album"].as_str().unwrap_or("").to_string(),
255 codec: e["codec"].as_str().map(String::from),
256 track_number: e["trackNumber"].as_i64(),
257 disc: e["disc"].as_i64(),
258 duration_ms: e["durationMs"].as_u64(),
259 is_current: e["isCurrent"].as_bool().unwrap_or(false),
260 })
261 .collect()
262 })
263 .unwrap_or_default();
264 Ok(entries)
265 }
266
267 pub fn pause(&self) -> Result<(), GraphQLError> {
270 self.execute("mutation { pause { ok } }", None)?;
271 Ok(())
272 }
273
274 pub fn resume(&self) -> Result<(), GraphQLError> {
275 self.execute("mutation { resume { ok } }", None)?;
276 Ok(())
277 }
278
279 pub fn stop(&self) -> Result<(), GraphQLError> {
280 self.execute("mutation { stop { ok } }", None)?;
281 Ok(())
282 }
283
284 pub fn next(&self) -> Result<(), GraphQLError> {
285 self.execute("mutation { next { ok } }", None)?;
286 Ok(())
287 }
288
289 pub fn previous(&self) -> Result<(), GraphQLError> {
290 self.execute("mutation { previous { ok } }", None)?;
291 Ok(())
292 }
293
294 pub fn seek(&self, position_ms: u64) -> Result<(), GraphQLError> {
295 self.execute(
296 "mutation($positionMs: Int!) { seek(positionMs: $positionMs) { ok } }",
297 Some(serde_json::json!({ "positionMs": position_ms })),
298 )?;
299 Ok(())
300 }
301
302 pub fn play(&self, queue_item_id: &str) -> Result<(), GraphQLError> {
303 self.execute(
304 "mutation($queueItemId: String!) { play(queueItemId: $queueItemId) { ok } }",
305 Some(serde_json::json!({ "queueItemId": queue_item_id })),
306 )?;
307 Ok(())
308 }
309
310 pub fn clear_queue(&self) -> Result<(), GraphQLError> {
311 self.execute("mutation { clearQueue { ok } }", None)?;
312 Ok(())
313 }
314
315 pub fn library_stats(&self) -> Result<Value, GraphQLError> {
316 self.execute(
317 "{ libraryStats { totalTracks totalArtists totalAlbums localTracks remoteTracks cachedTracks } }",
318 None,
319 )
320 }
321
322 pub fn server_url(&self) -> &str {
324 self.url.trim_end_matches("/graphql")
325 }
326}
327
328#[derive(Debug, thiserror::Error)]
333pub enum GraphQLError {
334 #[error("http error: {0}")]
335 Http(String),
336 #[error("query error: {0}")]
337 Query(String),
338 #[error("unauthorised: {0}")]
339 Unauthorized(String),
340}
341
342fn reason(resp: reqwest::blocking::Response) -> String {
345 let status = resp.status();
346 let text = resp.text().unwrap_or_default();
347 let message = serde_json::from_str::<Value>(&text)
348 .ok()
349 .and_then(|v| v["message"].as_str().map(str::to_owned))
350 .unwrap_or(text);
351 if message.trim().is_empty() {
352 status.to_string()
353 } else {
354 message.trim().to_owned()
355 }
356}
357
358#[derive(Debug, Clone)]
359pub struct NowPlaying {
360 pub state: String,
361 pub position_ms: u64,
362 pub duration_ms: Option<u64>,
363 pub queue_item_id: Option<String>,
364 pub track: Option<NowPlayingTrack>,
365}
366
367#[derive(Debug, Clone)]
368pub struct NowPlayingTrack {
369 pub track_id: Option<String>,
372 pub title: String,
373 pub artist: String,
374 pub album: String,
375 pub codec: String,
376 pub sample_rate: u32,
377 pub bit_depth: Option<u16>,
378 pub bitrate_kbps: Option<u32>,
379 pub channels: u16,
380 pub duration_ms: u64,
381}
382
383#[derive(Debug, Clone)]
384pub struct QueueEntry {
385 pub queue_item_id: String,
386 pub track_id: Option<String>,
387 pub title: String,
388 pub artist: String,
389 pub album: String,
390 pub codec: Option<String>,
391 pub track_number: Option<i64>,
392 pub disc: Option<i64>,
393 pub duration_ms: Option<u64>,
394 pub is_current: bool,
395}
396
397#[cfg(test)]
398mod tests {
399 use super::*;
400
401 #[test]
402 fn client_constructs_url() {
403 let c = GraphQLClient::new("http://localhost:4000");
404 assert_eq!(c.url, "http://localhost:4000/graphql");
405 }
406
407 #[test]
408 fn client_trailing_slash() {
409 let c = GraphQLClient::new("http://localhost:4000/");
410 assert_eq!(c.url, "http://localhost:4000/graphql");
411 }
412
413 struct Seen {
415 path: String,
416 bearer: Option<String>,
417 body: String,
418 }
419
420 fn serve(
423 respond: impl Fn(&Seen) -> (u16, &'static str) + Send + 'static,
424 ) -> (String, Arc<Mutex<Vec<String>>>) {
425 use std::io::{BufRead, BufReader, Read, Write};
426
427 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
428 let base = format!("http://{}", listener.local_addr().unwrap());
429 let log = Arc::new(Mutex::new(Vec::new()));
430 let seen_log = log.clone();
431 std::thread::spawn(move || {
432 for stream in listener.incoming() {
433 let Ok(mut stream) = stream else { return };
434 let mut reader = BufReader::new(stream.try_clone().unwrap());
435 let mut line = String::new();
436 reader.read_line(&mut line).unwrap();
437 let path = line.split_whitespace().nth(1).unwrap_or("").to_owned();
438 let (mut bearer, mut length) = (None, 0);
439 loop {
440 line.clear();
441 reader.read_line(&mut line).unwrap();
442 let header = line.trim_end();
443 if header.is_empty() {
444 break;
445 }
446 let (name, value) = header.split_once(": ").unwrap_or((header, ""));
447 match name.to_ascii_lowercase().as_str() {
448 "authorization" => {
449 bearer = value.strip_prefix("Bearer ").map(str::to_owned)
450 }
451 "content-length" => length = value.parse().unwrap_or(0),
452 _ => {}
453 }
454 }
455 let mut body = vec![0; length];
456 reader.read_exact(&mut body).unwrap();
457 let seen = Seen {
458 path,
459 bearer,
460 body: String::from_utf8_lossy(&body).into_owned(),
461 };
462 seen_log.lock().push(seen.path.clone());
463 let (code, reply) = respond(&seen);
464 write!(
465 stream,
466 "HTTP/1.1 {code} X\r\nContent-Type: application/json\r\n\
467 Content-Length: {}\r\nConnection: close\r\n\r\n{reply}",
468 reply.len()
469 )
470 .unwrap();
471 }
472 });
473 (base, log)
474 }
475
476 fn signed_in_server(seen: &Seen) -> (u16, &'static str) {
478 match seen.path.as_str() {
479 "/graphql" if seen.bearer.as_deref() == Some("access-2") => {
480 (200, r#"{"data":{"ok":true}}"#)
481 }
482 "/auth/refresh" if seen.body.contains("refresh-1") => (
483 200,
484 r#"{"access_token":"access-2","refresh_token":"refresh-2"}"#,
485 ),
486 "/auth/refresh" => (401, r#"{"message":"invalid or expired refresh token"}"#),
487 _ => (401, "missing or invalid Authorization header"),
488 }
489 }
490
491 #[test]
492 fn refused_request_refreshes_once_and_retries() {
493 let (base, log) = serve(signed_in_server);
494 let rotated = Arc::new(Mutex::new(Vec::new()));
495 let stored = rotated.clone();
496 let client = GraphQLClient::new(&base)
497 .with_session("refresh-1", move |t| stored.lock().push(t.to_owned()));
498
499 let data = client.execute("{ ok }", None).unwrap();
500 assert_eq!(data["ok"], true);
501 assert_eq!(*rotated.lock(), ["refresh-2"]);
502
503 client.clone().execute("{ ok }", None).unwrap();
505 assert_eq!(
506 *log.lock(),
507 ["/graphql", "/auth/refresh", "/graphql", "/graphql"]
508 );
509 }
510
511 #[test]
512 fn refused_refresh_is_unauthorised() {
513 let (base, log) = serve(signed_in_server);
514 let client = GraphQLClient::new(&base).with_session("revoked", |_| {});
515
516 for _ in 0..2 {
517 match client.execute("{ ok }", None) {
518 Err(GraphQLError::Unauthorized(msg)) => {
519 assert!(msg.contains("invalid or expired refresh token"), "{msg}")
520 }
521 other => panic!("expected Unauthorized, got {other:?}"),
522 }
523 }
524 assert_eq!(*log.lock(), ["/graphql", "/auth/refresh", "/graphql"]);
526 }
527
528 #[test]
529 fn unauthorised_without_a_session_says_why() {
530 let (base, log) = serve(signed_in_server);
531
532 match GraphQLClient::new(&base).execute("{ ok }", None) {
533 Err(GraphQLError::Unauthorized(msg)) => {
534 assert_eq!(msg, "missing or invalid Authorization header")
535 }
536 other => panic!("expected Unauthorized, got {other:?}"),
537 }
538 assert_eq!(*log.lock(), ["/graphql"]);
539 }
540}