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}
23
24/// A sign-in to a server with auth enabled: the refresh token `koan auth login`
25/// stored, and the access token it last bought.
26struct Session {
27    tokens: Mutex<Tokens>,
28    /// Called with each new refresh token. The server spends the old one on
29    /// every refresh, so a copy not written back is dead.
30    on_rotate: Box<dyn Fn(&str) + Send + Sync>,
31}
32
33struct Tokens {
34    access: Option<String>,
35    refresh: String,
36    /// Why the server refused the refresh token. Set once and kept: the token
37    /// will not become valid again, and the bridge polls several times a second.
38    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    /// A client carrying the sign-in stored in `[auth]`, when it is for this
55    /// server. Rotated refresh tokens are written back through `Config::persist`.
56    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    /// Authenticate with `refresh_token`, exchanging it for an access token
72    /// when the server first refuses a request.
73    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    /// Whether requests carry a stored sign-in.
90    pub fn has_session(&self) -> bool {
91        self.session.is_some()
92    }
93
94    /// Execute a raw GraphQL query/mutation.
95    ///
96    /// Access tokens are short-lived, so a 401 is answered by one refresh and
97    /// one retry; a second refusal is `GraphQLError::Unauthorized`.
98    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    /// A new access token, replacing `stale`.
153    ///
154    /// The lock is held across the request: the server spends a refresh token
155    /// on use, so two threads refreshing at once would leave one holding a
156    /// revoked token. A thread that waited finds `stale` already replaced and
157    /// takes the new token instead.
158    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    // -----------------------------------------------------------------------
205    // Typed helpers
206    // -----------------------------------------------------------------------
207
208    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    // -- Mutations --
268
269    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    /// Server URL (without /graphql path).
323    pub fn server_url(&self) -> &str {
324        self.url.trim_end_matches("/graphql")
325    }
326}
327
328// ---------------------------------------------------------------------------
329// Result types
330// ---------------------------------------------------------------------------
331
332#[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
342/// What the server said when it refused a request: the `message` of a JSON
343/// body, or the body as text.
344fn 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    /// The track's id on the server. `None` for a queue entry the server built
370    /// from a file with no database row, which cannot be streamed.
371    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    /// A request as the test server saw it.
414    struct Seen {
415        path: String,
416        bearer: Option<String>,
417        body: String,
418    }
419
420    /// An HTTP server on a loopback port answering each request with
421    /// `respond`. Returns its base URL and every request it received.
422    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    /// A server that takes `access-2` and trades `refresh-1` for it once.
477    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        // The new access token is kept: no second refresh.
504        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        // A refused refresh token is not offered again.
525        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}