Skip to main content

koan_core/
graphql_client.rs

1//! GraphQL client for connecting to a `koan serve` instance.
2//!
3//! Uses blocking reqwest — call from a background thread when used from the TUI.
4
5use std::sync::Arc;
6
7use parking_lot::Mutex;
8use reqwest::StatusCode;
9use serde_json::Value;
10
11use crate::config::Config;
12
13/// A GraphQL client that talks to a koan server.
14///
15/// Clones share one session, so a token refreshed by one thread is the token
16/// every other thread uses next.
17#[derive(Clone)]
18pub struct GraphQLClient {
19    url: String,
20    http: reqwest::blocking::Client,
21    session: Option<Arc<Session>>,
22    /// Cleared once the server turns out to predate play modes, after which
23    /// `nowPlaying` is asked for without them. Shared by clones, like the
24    /// session.
25    play_modes: Arc<std::sync::atomic::AtomicBool>,
26}
27
28/// A sign-in to a server with auth enabled: the refresh token `koan auth login`
29/// stored, and the access token it last bought.
30struct Session {
31    tokens: Mutex<Tokens>,
32    /// Called with each new refresh token. The server spends the old one on
33    /// every refresh, so a copy not written back is dead.
34    on_rotate: Box<dyn Fn(&str) + Send + Sync>,
35}
36
37struct Tokens {
38    access: Option<String>,
39    refresh: String,
40    /// Why the server refused the refresh token. Set once and kept: the token
41    /// will not become valid again, and the bridge polls several times a second.
42    refused: Option<String>,
43}
44
45impl GraphQLClient {
46    pub fn new(server_url: &str) -> Self {
47        let url = format!("{}/graphql", server_url.trim_end_matches('/'));
48        Self {
49            url,
50            http: reqwest::blocking::Client::builder()
51                .timeout(std::time::Duration::from_secs(30))
52                .build()
53                .expect("failed to build HTTP client"),
54            session: None,
55            play_modes: Arc::new(std::sync::atomic::AtomicBool::new(true)),
56        }
57    }
58
59    /// A client carrying the sign-in stored in `[auth]`, when it is for this
60    /// server. Rotated refresh tokens are written back through `Config::persist`.
61    pub fn from_config(server_url: &str) -> Self {
62        let client = Self::new(server_url);
63        let cfg = Config::load().unwrap_or_default();
64        if cfg.auth.refresh_token.is_empty()
65            || cfg.auth.server.trim_end_matches('/') != client.server_url()
66        {
67            return client;
68        }
69        client.with_session(cfg.auth.refresh_token, |token| {
70            if let Err(e) = Config::persist(|cfg| cfg.auth.refresh_token = token.to_owned()) {
71                log::warn!("could not store the rotated refresh token: {e}");
72            }
73        })
74    }
75
76    /// Authenticate with `refresh_token`, exchanging it for an access token
77    /// when the server first refuses a request.
78    pub fn with_session(
79        mut self,
80        refresh_token: impl Into<String>,
81        on_rotate: impl Fn(&str) + Send + Sync + 'static,
82    ) -> Self {
83        self.session = Some(Arc::new(Session {
84            tokens: Mutex::new(Tokens {
85                access: None,
86                refresh: refresh_token.into(),
87                refused: None,
88            }),
89            on_rotate: Box::new(on_rotate),
90        }));
91        self
92    }
93
94    /// Whether requests carry a stored sign-in.
95    pub fn has_session(&self) -> bool {
96        self.session.is_some()
97    }
98
99    /// Execute a raw GraphQL query/mutation.
100    ///
101    /// Access tokens are short-lived, so a 401 is answered by one refresh and
102    /// one retry; a second refusal is `GraphQLError::Unauthorized`.
103    pub fn execute(&self, query: &str, variables: Option<Value>) -> Result<Value, GraphQLError> {
104        let mut body = serde_json::json!({ "query": query });
105        if let Some(vars) = variables {
106            body["variables"] = vars;
107        }
108
109        let access = self
110            .session
111            .as_ref()
112            .and_then(|s| s.tokens.lock().access.clone());
113        let mut resp = self.post(&body, access.as_deref())?;
114        if resp.status() == StatusCode::UNAUTHORIZED
115            && let Some(session) = &self.session
116        {
117            let fresh = self.refresh(session, access.as_deref())?;
118            resp = self.post(&body, Some(&fresh))?;
119        }
120
121        let status = resp.status();
122        if status == StatusCode::UNAUTHORIZED {
123            return Err(GraphQLError::Unauthorized(reason(resp)));
124        }
125        if !status.is_success() {
126            return Err(GraphQLError::Http(format!("{status}: {}", reason(resp))));
127        }
128
129        let resp: Value = resp.json().map_err(|e| GraphQLError::Http(e.to_string()))?;
130
131        if let Some(errors) = resp.get("errors")
132            && let Some(arr) = errors.as_array()
133            && !arr.is_empty()
134        {
135            let msg = arr[0]
136                .get("message")
137                .and_then(|m| m.as_str())
138                .unwrap_or("unknown error");
139            return Err(GraphQLError::Query(msg.to_string()));
140        }
141
142        Ok(resp.get("data").cloned().unwrap_or(Value::Null))
143    }
144
145    fn post(
146        &self,
147        body: &Value,
148        access: Option<&str>,
149    ) -> Result<reqwest::blocking::Response, GraphQLError> {
150        let mut req = self.http.post(&self.url).json(body);
151        if let Some(token) = access {
152            req = req.bearer_auth(token);
153        }
154        req.send().map_err(|e| GraphQLError::Http(e.to_string()))
155    }
156
157    /// A new access token, replacing `stale`.
158    ///
159    /// The lock is held across the request: the server spends a refresh token
160    /// on use, so two threads refreshing at once would leave one holding a
161    /// revoked token. A thread that waited finds `stale` already replaced and
162    /// takes the new token instead.
163    fn refresh(&self, session: &Session, stale: Option<&str>) -> Result<String, GraphQLError> {
164        let mut tokens = session.tokens.lock();
165        if let Some(access) = &tokens.access
166            && Some(access.as_str()) != stale
167        {
168            return Ok(access.clone());
169        }
170        if let Some(reason) = &tokens.refused {
171            return Err(GraphQLError::Unauthorized(reason.clone()));
172        }
173
174        let resp = self
175            .http
176            .post(format!("{}/auth/refresh", self.server_url()))
177            .json(&serde_json::json!({ "refresh_token": tokens.refresh }))
178            .send()
179            .map_err(|e| GraphQLError::Http(e.to_string()))?;
180        if resp.status() == StatusCode::UNAUTHORIZED {
181            let refused = format!("the stored sign-in was refused: {}", reason(resp));
182            tokens.access = None;
183            tokens.refused = Some(refused.clone());
184            return Err(GraphQLError::Unauthorized(refused));
185        }
186        if !resp.status().is_success() {
187            let status = resp.status();
188            return Err(GraphQLError::Http(format!(
189                "/auth/refresh {status}: {}",
190                reason(resp)
191            )));
192        }
193
194        let body: Value = resp.json().map_err(|e| GraphQLError::Http(e.to_string()))?;
195        let (Some(access), Some(refresh)) = (
196            body["access_token"].as_str(),
197            body["refresh_token"].as_str(),
198        ) else {
199            return Err(GraphQLError::Http(
200                "malformed /auth/refresh response".into(),
201            ));
202        };
203        tokens.access = Some(access.to_owned());
204        tokens.refresh = refresh.to_owned();
205        (session.on_rotate)(refresh);
206        Ok(access.to_owned())
207    }
208
209    // -----------------------------------------------------------------------
210    // Typed helpers
211    // -----------------------------------------------------------------------
212
213    /// What the server is playing. A server older than play modes rejects a
214    /// query naming them, so on that refusal the query is asked again without
215    /// them, and from then on, and the mode reads as off.
216    pub fn now_playing(&self) -> Result<NowPlaying, GraphQLError> {
217        use std::sync::atomic::Ordering;
218        const TRACK: &str = "track { trackId title artist album codec sampleRate bitDepth \
219                             bitrateKbps channels durationMs }";
220        let ask = |modes: bool| {
221            let modes = if modes { "shuffle repeat " } else { "" };
222            self.execute(
223                &format!(
224                    "{{ nowPlaying {{ state positionMs durationMs queueItemId {modes}{TRACK} }} }}"
225                ),
226                None,
227            )
228        };
229        let data = if self.play_modes.load(Ordering::Relaxed) {
230            match ask(true) {
231                Err(GraphQLError::Query(e)) if e.contains("shuffle") || e.contains("repeat") => {
232                    log::info!("server predates play modes: {e}");
233                    self.play_modes.store(false, Ordering::Relaxed);
234                    ask(false)?
235                }
236                other => other?,
237            }
238        } else {
239            ask(false)?
240        };
241        let np = &data["nowPlaying"];
242        Ok(NowPlaying {
243            state: np["state"].as_str().unwrap_or("STOPPED").to_string(),
244            position_ms: np["positionMs"].as_u64().unwrap_or(0),
245            duration_ms: np["durationMs"].as_u64(),
246            queue_item_id: np["queueItemId"].as_str().map(String::from),
247            mode: crate::player::state::PlayMode {
248                shuffle: np["shuffle"].as_bool().unwrap_or(false),
249                repeat: np["repeat"]
250                    .as_str()
251                    .and_then(|r| crate::player::state::Repeat::parse(&r.to_lowercase()))
252                    .unwrap_or_default(),
253            },
254            track: np.get("track").and_then(|t| {
255                if t.is_null() {
256                    return None;
257                }
258                Some(NowPlayingTrack {
259                    track_id: t["trackId"].as_str().map(String::from),
260                    title: t["title"].as_str().unwrap_or("").to_string(),
261                    artist: t["artist"].as_str().unwrap_or("").to_string(),
262                    album: t["album"].as_str().unwrap_or("").to_string(),
263                    codec: t["codec"].as_str().unwrap_or("").to_string(),
264                    sample_rate: t["sampleRate"].as_u64().unwrap_or(0) as u32,
265                    bit_depth: t["bitDepth"].as_u64().map(|v| v as u16),
266                    bitrate_kbps: t["bitrateKbps"].as_u64().map(|v| v as u32),
267                    channels: t["channels"].as_u64().unwrap_or(0) as u16,
268                    duration_ms: t["durationMs"].as_u64().unwrap_or(0),
269                })
270            }),
271        })
272    }
273
274    pub fn queue(&self) -> Result<Vec<QueueEntry>, GraphQLError> {
275        let data = self.execute(
276            "{ queue { queueItemId trackId title artist album codec trackNumber disc durationMs isCurrent } }",
277            None,
278        )?;
279        let entries = data["queue"]
280            .as_array()
281            .map(|arr| {
282                arr.iter()
283                    .map(|e| QueueEntry {
284                        queue_item_id: e["queueItemId"].as_str().unwrap_or("").to_string(),
285                        track_id: e["trackId"].as_str().map(String::from),
286                        title: e["title"].as_str().unwrap_or("").to_string(),
287                        artist: e["artist"].as_str().unwrap_or("").to_string(),
288                        album: e["album"].as_str().unwrap_or("").to_string(),
289                        codec: e["codec"].as_str().map(String::from),
290                        track_number: e["trackNumber"].as_i64(),
291                        disc: e["disc"].as_i64(),
292                        duration_ms: e["durationMs"].as_u64(),
293                        is_current: e["isCurrent"].as_bool().unwrap_or(false),
294                    })
295                    .collect()
296            })
297            .unwrap_or_default();
298        Ok(entries)
299    }
300
301    // -- Mutations --
302
303    pub fn pause(&self) -> Result<(), GraphQLError> {
304        self.execute("mutation { pause { ok } }", None)?;
305        Ok(())
306    }
307
308    pub fn resume(&self) -> Result<(), GraphQLError> {
309        self.execute("mutation { resume { ok } }", None)?;
310        Ok(())
311    }
312
313    pub fn stop(&self) -> Result<(), GraphQLError> {
314        self.execute("mutation { stop { ok } }", None)?;
315        Ok(())
316    }
317
318    pub fn next(&self) -> Result<(), GraphQLError> {
319        self.execute("mutation { next { ok } }", None)?;
320        Ok(())
321    }
322
323    pub fn previous(&self) -> Result<(), GraphQLError> {
324        self.execute("mutation { previous { ok } }", None)?;
325        Ok(())
326    }
327
328    pub fn set_shuffle(&self, on: bool) -> Result<(), GraphQLError> {
329        self.execute(
330            "mutation($on: Boolean!) { setPlayMode(shuffle: $on) { ok } }",
331            Some(serde_json::json!({ "on": on })),
332        )?;
333        Ok(())
334    }
335
336    pub fn set_repeat(&self, repeat: crate::player::state::Repeat) -> Result<(), GraphQLError> {
337        self.execute(
338            "mutation($repeat: Repeat!) { setPlayMode(repeat: $repeat) { ok } }",
339            Some(serde_json::json!({ "repeat": repeat.as_str().to_uppercase() })),
340        )?;
341        Ok(())
342    }
343
344    pub fn seek(&self, position_ms: u64) -> Result<(), GraphQLError> {
345        self.execute(
346            "mutation($positionMs: Int!) { seek(positionMs: $positionMs) { ok } }",
347            Some(serde_json::json!({ "positionMs": position_ms })),
348        )?;
349        Ok(())
350    }
351
352    pub fn play(&self, queue_item_id: &str) -> Result<(), GraphQLError> {
353        self.execute(
354            "mutation($queueItemId: String!) { play(queueItemId: $queueItemId) { ok } }",
355            Some(serde_json::json!({ "queueItemId": queue_item_id })),
356        )?;
357        Ok(())
358    }
359
360    pub fn clear_queue(&self) -> Result<(), GraphQLError> {
361        self.execute("mutation { clearQueue { ok } }", None)?;
362        Ok(())
363    }
364
365    pub fn library_stats(&self) -> Result<Value, GraphQLError> {
366        self.execute(
367            "{ libraryStats { totalTracks totalArtists totalAlbums localTracks remoteTracks cachedTracks } }",
368            None,
369        )
370    }
371
372    /// Server URL (without /graphql path).
373    pub fn server_url(&self) -> &str {
374        self.url.trim_end_matches("/graphql")
375    }
376}
377
378// ---------------------------------------------------------------------------
379// Result types
380// ---------------------------------------------------------------------------
381
382#[derive(Debug, thiserror::Error)]
383pub enum GraphQLError {
384    #[error("http error: {0}")]
385    Http(String),
386    #[error("query error: {0}")]
387    Query(String),
388    #[error("unauthorised: {0}")]
389    Unauthorized(String),
390}
391
392/// What the server said when it refused a request: the `message` of a JSON
393/// body, or the body as text.
394fn reason(resp: reqwest::blocking::Response) -> String {
395    let status = resp.status();
396    let text = resp.text().unwrap_or_default();
397    let message = serde_json::from_str::<Value>(&text)
398        .ok()
399        .and_then(|v| v["message"].as_str().map(str::to_owned))
400        .unwrap_or(text);
401    if message.trim().is_empty() {
402        status.to_string()
403    } else {
404        message.trim().to_owned()
405    }
406}
407
408#[derive(Debug, Clone)]
409pub struct NowPlaying {
410    pub state: String,
411    pub position_ms: u64,
412    pub duration_ms: Option<u64>,
413    pub queue_item_id: Option<String>,
414    pub mode: crate::player::state::PlayMode,
415    pub track: Option<NowPlayingTrack>,
416}
417
418#[derive(Debug, Clone)]
419pub struct NowPlayingTrack {
420    /// The track's id on the server. `None` for a queue entry the server built
421    /// from a file with no database row, which cannot be streamed.
422    pub track_id: Option<String>,
423    pub title: String,
424    pub artist: String,
425    pub album: String,
426    pub codec: String,
427    pub sample_rate: u32,
428    pub bit_depth: Option<u16>,
429    pub bitrate_kbps: Option<u32>,
430    pub channels: u16,
431    pub duration_ms: u64,
432}
433
434#[derive(Debug, Clone)]
435pub struct QueueEntry {
436    pub queue_item_id: String,
437    pub track_id: Option<String>,
438    pub title: String,
439    pub artist: String,
440    pub album: String,
441    pub codec: Option<String>,
442    pub track_number: Option<i64>,
443    pub disc: Option<i64>,
444    pub duration_ms: Option<u64>,
445    pub is_current: bool,
446}
447
448#[cfg(test)]
449mod tests {
450    use super::*;
451
452    #[test]
453    fn client_constructs_url() {
454        let c = GraphQLClient::new("http://localhost:4000");
455        assert_eq!(c.url, "http://localhost:4000/graphql");
456    }
457
458    #[test]
459    fn client_trailing_slash() {
460        let c = GraphQLClient::new("http://localhost:4000/");
461        assert_eq!(c.url, "http://localhost:4000/graphql");
462    }
463
464    /// A request as the test server saw it.
465    struct Seen {
466        path: String,
467        bearer: Option<String>,
468        body: String,
469    }
470
471    /// An HTTP server on a loopback port answering each request with
472    /// `respond`. Returns its base URL and every request it received.
473    fn serve(
474        respond: impl Fn(&Seen) -> (u16, &'static str) + Send + 'static,
475    ) -> (String, Arc<Mutex<Vec<String>>>) {
476        use std::io::{BufRead, BufReader, Read, Write};
477
478        let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
479        let base = format!("http://{}", listener.local_addr().unwrap());
480        let log = Arc::new(Mutex::new(Vec::new()));
481        let seen_log = log.clone();
482        std::thread::spawn(move || {
483            for stream in listener.incoming() {
484                let Ok(mut stream) = stream else { return };
485                let mut reader = BufReader::new(stream.try_clone().unwrap());
486                let mut line = String::new();
487                reader.read_line(&mut line).unwrap();
488                let path = line.split_whitespace().nth(1).unwrap_or("").to_owned();
489                let (mut bearer, mut length) = (None, 0);
490                loop {
491                    line.clear();
492                    reader.read_line(&mut line).unwrap();
493                    let header = line.trim_end();
494                    if header.is_empty() {
495                        break;
496                    }
497                    let (name, value) = header.split_once(": ").unwrap_or((header, ""));
498                    match name.to_ascii_lowercase().as_str() {
499                        "authorization" => {
500                            bearer = value.strip_prefix("Bearer ").map(str::to_owned)
501                        }
502                        "content-length" => length = value.parse().unwrap_or(0),
503                        _ => {}
504                    }
505                }
506                let mut body = vec![0; length];
507                reader.read_exact(&mut body).unwrap();
508                let seen = Seen {
509                    path,
510                    bearer,
511                    body: String::from_utf8_lossy(&body).into_owned(),
512                };
513                seen_log.lock().push(seen.path.clone());
514                let (code, reply) = respond(&seen);
515                write!(
516                    stream,
517                    "HTTP/1.1 {code} X\r\nContent-Type: application/json\r\n\
518                     Content-Length: {}\r\nConnection: close\r\n\r\n{reply}",
519                    reply.len()
520                )
521                .unwrap();
522            }
523        });
524        (base, log)
525    }
526
527    /// A server that takes `access-2` and trades `refresh-1` for it once.
528    fn signed_in_server(seen: &Seen) -> (u16, &'static str) {
529        match seen.path.as_str() {
530            "/graphql" if seen.bearer.as_deref() == Some("access-2") => {
531                (200, r#"{"data":{"ok":true}}"#)
532            }
533            "/auth/refresh" if seen.body.contains("refresh-1") => (
534                200,
535                r#"{"access_token":"access-2","refresh_token":"refresh-2"}"#,
536            ),
537            "/auth/refresh" => (401, r#"{"message":"invalid or expired refresh token"}"#),
538            _ => (401, "missing or invalid Authorization header"),
539        }
540    }
541
542    #[test]
543    fn refused_request_refreshes_once_and_retries() {
544        let (base, log) = serve(signed_in_server);
545        let rotated = Arc::new(Mutex::new(Vec::new()));
546        let stored = rotated.clone();
547        let client = GraphQLClient::new(&base)
548            .with_session("refresh-1", move |t| stored.lock().push(t.to_owned()));
549
550        let data = client.execute("{ ok }", None).unwrap();
551        assert_eq!(data["ok"], true);
552        assert_eq!(*rotated.lock(), ["refresh-2"]);
553
554        // The new access token is kept: no second refresh.
555        client.clone().execute("{ ok }", None).unwrap();
556        assert_eq!(
557            *log.lock(),
558            ["/graphql", "/auth/refresh", "/graphql", "/graphql"]
559        );
560    }
561
562    #[test]
563    fn refused_refresh_is_unauthorised() {
564        let (base, log) = serve(signed_in_server);
565        let client = GraphQLClient::new(&base).with_session("revoked", |_| {});
566
567        for _ in 0..2 {
568            match client.execute("{ ok }", None) {
569                Err(GraphQLError::Unauthorized(msg)) => {
570                    assert!(msg.contains("invalid or expired refresh token"), "{msg}")
571                }
572                other => panic!("expected Unauthorized, got {other:?}"),
573            }
574        }
575        // A refused refresh token is not offered again.
576        assert_eq!(*log.lock(), ["/graphql", "/auth/refresh", "/graphql"]);
577    }
578
579    #[test]
580    fn unauthorised_without_a_session_says_why() {
581        let (base, log) = serve(signed_in_server);
582
583        match GraphQLClient::new(&base).execute("{ ok }", None) {
584            Err(GraphQLError::Unauthorized(msg)) => {
585                assert_eq!(msg, "missing or invalid Authorization header")
586            }
587            other => panic!("expected Unauthorized, got {other:?}"),
588        }
589        assert_eq!(*log.lock(), ["/graphql"]);
590    }
591}