1use std::collections::HashMap;
12use std::net::IpAddr;
13use std::sync::LazyLock;
14use std::time::{Duration, Instant};
15
16use axum::extract::RawQuery;
17use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
18use axum::http::StatusCode;
19use axum::response::{IntoResponse, Response};
20use koan_core::remote::pair::PairMessage;
21use parking_lot::Mutex;
22use tokio::sync::oneshot;
23
24use crate::auth::routes::{RateLimiter, network};
25
26pub const TTL: Duration = Duration::from_secs(600);
28const MAX_PENDING: usize = 256;
30const MAX_PENDING_PER_ADDRESS: usize = 3;
34const CODE_ALPHABET: &[u8; 32] = b"0123456789ABCDEFGHJKMNPQRSTVWXYZ";
37const CODE_LEN: usize = 8;
38const KEEPALIVE: Duration = Duration::from_secs(30);
40
41struct Entry {
42 code: String,
44 device: String,
45 from: IpAddr,
47 created: Instant,
48 outcome: oneshot::Sender<PairMessage>,
49}
50
51pub struct Pairings {
52 ttl: Duration,
53 entries: Mutex<HashMap<String, Entry>>,
54}
55
56pub fn pairings() -> &'static Pairings {
59 static PAIRINGS: LazyLock<Pairings> = LazyLock::new(|| Pairings::new(ttl_from_env()));
60 &PAIRINGS
61}
62
63fn ttl_from_env() -> Duration {
66 std::env::var("KOAN_PAIR_TTL_SECS")
67 .ok()
68 .and_then(|s| s.trim().parse().ok())
69 .filter(|&secs| secs > 0)
70 .map_or(TTL, Duration::from_secs)
71}
72
73static OPENS: LazyLock<RateLimiter> = LazyLock::new(|| RateLimiter::new(60, 10));
76
77pub struct Opened<'a> {
80 pub id: String,
81 pub code: String,
83 outcome: oneshot::Receiver<PairMessage>,
84 pairings: &'a Pairings,
85}
86
87impl Drop for Opened<'_> {
88 fn drop(&mut self) {
89 self.pairings.entries.lock().remove(&self.id);
90 }
91}
92
93#[derive(Debug, Clone, PartialEq, Eq)]
95pub struct PairInfo {
96 pub device: String,
97 pub from: IpAddr,
98}
99
100impl PairInfo {
101 pub fn local(&self) -> bool {
104 is_local(self.from)
105 }
106}
107
108fn is_local(ip: IpAddr) -> bool {
111 match ip.to_canonical() {
112 IpAddr::V4(v4) => {
113 let [a, b, ..] = v4.octets();
114 v4.is_private()
115 || (a == 100 && (b & 0xc0) == 64)
116 || v4.is_link_local()
117 || v4.is_loopback()
118 }
119 IpAddr::V6(v6) => {
120 v6.is_loopback()
121 || (v6.segments()[0] & 0xfe00) == 0xfc00
122 || (v6.segments()[0] & 0xffc0) == 0xfe80
123 }
124 }
125}
126
127pub struct Taken {
129 id: String,
130 entry: Entry,
131}
132
133impl Taken {
134 pub fn device(&self) -> &str {
135 &self.entry.device
136 }
137
138 pub fn info(&self) -> PairInfo {
139 PairInfo {
140 device: self.entry.device.clone(),
141 from: self.entry.from,
142 }
143 }
144
145 pub fn settle(self, outcome: PairMessage) -> Result<(), PairMessage> {
147 self.entry.outcome.send(outcome)
148 }
149}
150
151#[derive(Debug, PartialEq, Eq)]
152pub enum OpenError {
153 Full,
154 Busy,
156 Entropy,
157}
158
159impl Pairings {
160 fn new(ttl: Duration) -> Self {
161 Self {
162 ttl,
163 entries: Mutex::default(),
164 }
165 }
166
167 pub fn open(&self, device: &str, from: IpAddr) -> Result<Opened<'_>, OpenError> {
169 let id = koan_core::auth::random_api_key().map_err(|_| OpenError::Entropy)?;
170 let (tx, rx) = oneshot::channel();
171 let mut entries = self.entries.lock();
172 self.sweep(&mut entries);
173 if entries.len() >= MAX_PENDING {
174 return Err(OpenError::Full);
175 }
176 let from = from.to_canonical();
177 let waiting = entries
178 .values()
179 .filter(|e| network(e.from) == network(from))
180 .count();
181 if waiting >= MAX_PENDING_PER_ADDRESS {
182 return Err(OpenError::Busy);
183 }
184 let code = loop {
185 let code = new_code()?;
186 if !entries.values().any(|e| e.code == code) {
187 break code;
188 }
189 };
190 entries.insert(
191 id.clone(),
192 Entry {
193 code: code.clone(),
194 device: koan_core::invite::device_name(device),
195 from,
196 created: Instant::now(),
197 outcome: tx,
198 },
199 );
200 Ok(Opened {
201 id,
202 code: format!("{}-{}", &code[..4], &code[4..]),
203 outcome: rx,
204 pairings: self,
205 })
206 }
207
208 pub fn info(&self, pair: &str) -> Option<PairInfo> {
210 let mut entries = self.entries.lock();
211 self.sweep(&mut entries);
212 let id = find(&entries, pair)?;
213 entries.get(&id).map(|e| PairInfo {
214 device: e.device.clone(),
215 from: e.from,
216 })
217 }
218
219 pub fn take(&self, pair: &str) -> Option<Taken> {
222 let mut entries = self.entries.lock();
223 self.sweep(&mut entries);
224 let id = find(&entries, pair)?;
225 let entry = entries.remove(&id)?;
226 Some(Taken { id, entry })
227 }
228
229 pub fn put_back(&self, taken: Taken) {
231 if !taken.entry.outcome.is_closed() {
232 self.entries.lock().insert(taken.id, taken.entry);
233 }
234 }
235
236 pub fn settle(
240 &self,
241 conn: &rusqlite::Connection,
242 pair: &str,
243 user_id: i64,
244 username: &str,
245 decline: bool,
246 ) -> Result<PairInfo, SettleError> {
247 use koan_core::db::queries::api_keys;
248 let taken = self.take(pair).ok_or(SettleError::NotFound)?;
249 let info = taken.info();
250 let device = info.device.clone();
251 if decline {
252 let _ = taken.settle(PairMessage::Declined);
253 log::info!("pair: {device} declined by {username}");
254 return Ok(info);
255 }
256 let api_key = match api_keys::create_api_key(conn, user_id, &device) {
257 Ok((_, key)) => key,
258 Err(e) => {
259 self.put_back(taken);
260 return Err(SettleError::Internal(e.to_string()));
261 }
262 };
263 let approved = PairMessage::Approved {
264 username: username.to_owned(),
265 api_key: api_key.clone(),
266 };
267 if taken.settle(approved).is_err() {
268 let _ = api_keys::revoke_api_key_value(conn, &api_key);
270 return Err(SettleError::NotFound);
271 }
272 log::info!("pair: {device} signed in as {username}");
273 Ok(info)
274 }
275
276 fn sweep(&self, entries: &mut HashMap<String, Entry>) {
278 let lapsed: Vec<String> = entries
279 .iter()
280 .filter(|(_, e)| e.created.elapsed() >= self.ttl)
281 .map(|(id, _)| id.clone())
282 .collect();
283 for id in lapsed {
284 if let Some(e) = entries.remove(&id) {
285 let _ = e.outcome.send(PairMessage::Expired);
286 }
287 }
288 }
289}
290
291#[derive(Debug, PartialEq, Eq)]
292pub enum SettleError {
293 NotFound,
295 Internal(String),
296}
297
298fn new_code() -> Result<String, OpenError> {
299 let mut bytes = [0u8; CODE_LEN];
300 getrandom::fill(&mut bytes).map_err(|_| OpenError::Entropy)?;
301 Ok(bytes
303 .iter()
304 .map(|b| CODE_ALPHABET[(b % 32) as usize] as char)
305 .collect())
306}
307
308fn normalise_code(typed: &str) -> Option<String> {
311 let code: String = typed
312 .chars()
313 .filter(|c| !matches!(c, '-' | ' '))
314 .map(|c| match c.to_ascii_uppercase() {
315 'I' | 'L' => '1',
316 'O' => '0',
317 c => c,
318 })
319 .collect();
320 (code.len() == CODE_LEN && code.bytes().all(|b| CODE_ALPHABET.contains(&b))).then_some(code)
321}
322
323fn find(entries: &HashMap<String, Entry>, pair: &str) -> Option<String> {
324 let pair = pair.trim();
325 if entries.contains_key(pair) {
326 return Some(pair.to_owned());
327 }
328 let code = normalise_code(pair)?;
329 entries
330 .iter()
331 .find(|(_, e)| e.code == code)
332 .map(|(id, _)| id.clone())
333}
334
335pub(crate) async fn route(
338 RawQuery(raw): RawQuery,
339 ws: WebSocketUpgrade,
340 request: axum::extract::Request,
341) -> Response {
342 if request.headers().contains_key(axum::http::header::ORIGIN) {
347 return (StatusCode::FORBIDDEN, "pairing is not open to web pages").into_response();
348 }
349 let from = crate::auth::routes::client_ip(&request);
350 if !OPENS.allow(from) {
351 return (StatusCode::TOO_MANY_REQUESTS, "too many pairings").into_response();
352 }
353 let name = form_urlencoded::parse(raw.unwrap_or_default().as_bytes())
354 .find(|(k, _)| k == "name")
355 .map(|(_, v)| v.into_owned())
356 .unwrap_or_default();
357 match pairings().open(&name, from) {
358 Ok(opened) => ws.on_upgrade(move |socket| session(socket, opened)),
359 Err(OpenError::Full) => {
360 (StatusCode::SERVICE_UNAVAILABLE, "too many pairings waiting").into_response()
361 }
362 Err(OpenError::Busy) => (
363 StatusCode::TOO_MANY_REQUESTS,
364 "too many pairings waiting from this address",
365 )
366 .into_response(),
367 Err(OpenError::Entropy) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
368 }
369}
370
371async fn session(mut socket: WebSocket, mut opened: Opened<'static>) {
372 let pending = PairMessage::Pending {
373 id: opened.id.clone(),
374 code: opened.code.clone(),
375 expires_in: opened.pairings.ttl.as_secs(),
376 };
377 if send(&mut socket, &pending).await.is_err() {
378 return;
379 }
380 log::info!("pair: {} waiting", opened.code);
381 let lapse = tokio::time::sleep(opened.pairings.ttl);
382 tokio::pin!(lapse);
383 let mut keepalive =
384 tokio::time::interval_at(tokio::time::Instant::now() + KEEPALIVE, KEEPALIVE);
385 let outcome = loop {
386 tokio::select! {
387 outcome = &mut opened.outcome => break outcome.unwrap_or(PairMessage::Expired),
389 () = &mut lapse => break PairMessage::Expired,
390 _ = keepalive.tick() => {
391 if socket.send(Message::Ping(Default::default())).await.is_err() {
392 return;
393 }
394 }
395 msg = socket.recv() => match msg {
396 Some(Ok(Message::Close(_)) | Err(_)) | None => return,
397 Some(Ok(_)) => {}
398 },
399 }
400 };
401 log::info!(
402 "pair: {} {}",
403 opened.code,
404 match outcome {
405 PairMessage::Approved { .. } => "approved",
406 PairMessage::Declined => "declined",
407 _ => "expired",
408 }
409 );
410 let _ = send(&mut socket, &outcome).await;
411 let _ = socket.send(Message::Close(None)).await;
412}
413
414async fn send(socket: &mut WebSocket, message: &PairMessage) -> Result<(), axum::Error> {
415 let text = serde_json::to_string(message).unwrap_or_default();
416 socket.send(Message::Text(text.into())).await
417}
418
419#[cfg(test)]
420mod tests {
421 use super::*;
422
423 const LAN: IpAddr = IpAddr::V4(std::net::Ipv4Addr::new(192, 168, 1, 20));
424
425 fn host(i: usize) -> IpAddr {
427 IpAddr::V4(std::net::Ipv4Addr::new(
428 10,
429 9,
430 (i / 256) as u8,
431 (i % 256) as u8,
432 ))
433 }
434
435 fn leaked(ttl: Duration) -> &'static Pairings {
436 Box::leak(Box::new(Pairings::new(ttl)))
437 }
438
439 #[test]
440 fn codes_are_crockford_and_dashed() {
441 let p = leaked(TTL);
442 for _ in 0..50 {
443 let o = p.open("tv", LAN).unwrap();
444 let (a, b) = o.code.split_once('-').unwrap();
445 assert_eq!((a.len(), b.len()), (4, 4));
446 assert!(
447 format!("{a}{b}")
448 .bytes()
449 .all(|c| CODE_ALPHABET.contains(&c) && !b"ILOU".contains(&c))
450 );
451 }
452 }
453
454 #[test]
455 fn ids_are_unguessable_and_distinct() {
456 let p = leaked(TTL);
457 let opened: Vec<_> = (0..100).map(|i| p.open("tv", host(i)).unwrap()).collect();
458 let ids: std::collections::HashSet<_> = opened.iter().map(|o| o.id.clone()).collect();
459 assert_eq!(ids.len(), 100);
460 assert!(opened.iter().all(|o| o.id.len() == 43));
461 let codes: std::collections::HashSet<_> = opened.iter().map(|o| o.code.clone()).collect();
462 assert_eq!(codes.len(), 100);
463 }
464
465 #[test]
466 fn a_code_is_found_however_it_is_typed() {
467 let p = leaked(TTL);
468 let o = p.open("Living room\u{7} TV", LAN).unwrap();
469 assert_eq!(
470 p.info(&o.id).map(|i| i.device).as_deref(),
471 Some("Living room TV")
472 );
473 let bare = o.code.replace('-', "");
474 for typed in [
475 o.code.clone(),
476 o.code.to_lowercase(),
477 bare.clone(),
478 format!(" {} {} ", &bare[..4], &bare[4..]),
479 ] {
480 assert_eq!(
481 p.info(&typed).map(|i| i.device).as_deref(),
482 Some("Living room TV"),
483 "{typed}"
484 );
485 }
486 assert_eq!(normalise_code("o1l1-IOAB").as_deref(), Some("011110AB"));
487 assert_eq!(normalise_code("ABCD-EFGU"), None);
488 assert_eq!(normalise_code("ABC"), None);
489 }
490
491 #[tokio::test]
492 async fn approving_delivers_the_key_to_the_waiting_device() {
493 let p = leaked(TTL);
494 let mut o = p.open("tv", LAN).unwrap();
495 let taken = p.take(&o.code).unwrap();
496 assert_eq!(taken.device(), "tv");
497 assert!(p.take(&o.id).is_none());
499 taken
500 .settle(PairMessage::Approved {
501 username: "alice".into(),
502 api_key: "key".into(),
503 })
504 .unwrap();
505 assert_eq!(
506 (&mut o.outcome).await.unwrap(),
507 PairMessage::Approved {
508 username: "alice".into(),
509 api_key: "key".into(),
510 }
511 );
512 }
513
514 fn database() -> (tempfile::TempDir, koan_core::db::connection::Database, i64) {
515 let dir = tempfile::tempdir().unwrap();
516 let db = koan_core::db::connection::Database::open(&dir.path().join("koan.db")).unwrap();
517 let id = koan_core::db::queries::auth::create_user(
518 &db.conn,
519 "alice",
520 "hunter2",
521 koan_core::auth::Role::Readonly,
522 )
523 .unwrap();
524 (dir, db, id)
525 }
526
527 #[tokio::test]
528 async fn approving_mints_a_key_for_the_approver() {
529 let (_dir, db, alice) = database();
530 let p = leaked(TTL);
531 let mut o = p.open("Living room TV", LAN).unwrap();
532 let settled = p.settle(&db.conn, &o.code, alice, "alice", false).unwrap();
533 assert_eq!(
534 settled,
535 PairInfo {
536 device: "Living room TV".into(),
537 from: LAN,
538 }
539 );
540 let PairMessage::Approved { username, api_key } = (&mut o.outcome).await.unwrap() else {
541 panic!("not approved");
542 };
543 assert_eq!(username, "alice");
544 let user = koan_core::db::queries::api_keys::authenticate_api_key(&db.conn, &api_key)
545 .unwrap()
546 .unwrap();
547 assert_eq!(user.id, alice);
548 let keys = koan_core::db::queries::api_keys::list_api_keys(&db.conn, Some(alice)).unwrap();
549 assert_eq!(keys[0].name, "Living room TV");
550 assert_eq!(
552 p.settle(&db.conn, &o.id, alice, "alice", false),
553 Err(SettleError::NotFound)
554 );
555 }
556
557 #[tokio::test]
558 async fn declining_mints_nothing() {
559 let (_dir, db, alice) = database();
560 let p = leaked(TTL);
561 let mut o = p.open("tv", LAN).unwrap();
562 p.settle(&db.conn, &o.id, alice, "alice", true).unwrap();
563 assert_eq!((&mut o.outcome).await.unwrap(), PairMessage::Declined);
564 let keys = koan_core::db::queries::api_keys::list_api_keys(&db.conn, Some(alice)).unwrap();
565 assert!(keys.is_empty());
566 }
567
568 #[test]
569 fn a_key_for_a_gone_device_is_revoked() {
570 let (_dir, db, alice) = database();
571 let p = leaked(TTL);
572 let o = p.open("tv", LAN).unwrap();
573 let id = o.id.clone();
574 let taken = p.take(&id).unwrap();
576 drop(o);
577 p.entries.lock().insert(taken.id, taken.entry);
578 assert_eq!(
579 p.settle(&db.conn, &id, alice, "alice", false),
580 Err(SettleError::NotFound)
581 );
582 let keys = koan_core::db::queries::api_keys::list_api_keys(&db.conn, Some(alice)).unwrap();
583 assert!(keys.is_empty());
584 }
585
586 #[test]
587 fn approving_an_unknown_pairing_is_not_found() {
588 let (_dir, db, alice) = database();
589 let p = leaked(TTL);
590 assert_eq!(
591 p.settle(&db.conn, "ABCD-EFGH", alice, "alice", false),
592 Err(SettleError::NotFound)
593 );
594 }
595
596 #[test]
597 fn a_gone_device_gives_the_outcome_back() {
598 let p = leaked(TTL);
599 let o = p.open("tv", LAN).unwrap();
600 let taken = p.take(&o.id).unwrap();
601 drop(o);
602 assert!(taken.settle(PairMessage::Declined).is_err());
603 }
604
605 #[tokio::test]
606 async fn a_lapsed_pairing_is_expired_and_unknown() {
607 let p = leaked(Duration::ZERO);
608 let mut o = p.open("tv", LAN).unwrap();
609 assert!(p.info(&o.id).is_none());
610 assert!(p.take(&o.code).is_none());
611 assert_eq!((&mut o.outcome).await.unwrap(), PairMessage::Expired);
612 }
613
614 #[test]
615 fn the_requesting_address_is_kept() {
616 let p = leaked(TTL);
617 let mapped: IpAddr = "::ffff:203.0.113.9".parse().unwrap();
618 let o = p.open("tv", mapped).unwrap();
619 let info = p.info(&o.code).unwrap();
620 assert_eq!(info.from.to_string(), "203.0.113.9");
621 assert!(!info.local());
622 assert_eq!(p.take(&o.id).unwrap().info(), info);
623 }
624
625 #[test]
626 fn private_addresses_are_local() {
627 for local in [
628 "10.1.2.3",
629 "172.16.0.1",
630 "192.168.1.20",
631 "169.254.10.1",
632 "127.0.0.1",
633 "100.64.0.1",
634 "100.101.102.103",
635 "100.127.255.254",
636 "::1",
637 "fd12:3456::1",
638 "fe80::1",
639 "::ffff:192.168.0.5",
640 ] {
641 assert!(is_local(local.parse().unwrap()), "{local}");
642 }
643 for public in [
644 "203.0.113.9",
645 "8.8.8.8",
646 "172.32.0.1",
647 "100.63.255.255",
648 "100.128.0.1",
649 "2001:db8::1",
650 "2a00:1450::1",
651 "::ffff:8.8.8.8",
652 ] {
653 assert!(!is_local(public.parse().unwrap()), "{public}");
654 }
655 }
656
657 #[test]
658 fn unknown_pairings_are_not_found() {
659 let p = leaked(TTL);
660 let o = p.open("tv", LAN).unwrap();
661 let other = if o.code == "0000-0000" {
662 "1111-1111"
663 } else {
664 "0000-0000"
665 };
666 assert!(p.info(other).is_none());
667 assert!(p.take("not-a-pairing").is_none());
668 assert!(p.info("").is_none());
669 }
670
671 #[test]
672 fn a_closed_socket_removes_its_pairing() {
673 let p = leaked(TTL);
674 let o = p.open("tv", LAN).unwrap();
675 let id = o.id.clone();
676 drop(o);
677 assert!(p.info(&id).is_none());
678 }
679
680 #[test]
681 fn pairings_are_capped() {
682 let p = leaked(TTL);
683 let held: Vec<_> = (0..MAX_PENDING)
684 .map(|i| p.open("tv", host(i)).unwrap())
685 .collect();
686 assert_eq!(p.open("tv", LAN).err(), Some(OpenError::Full));
687 drop(held);
688 assert!(p.open("tv", LAN).is_ok());
689 }
690
691 #[test]
692 fn one_address_holds_three_at_most() {
693 let p = leaked(TTL);
694 let held: Vec<_> = (0..3).map(|_| p.open("tv", LAN).unwrap()).collect();
695 assert_eq!(p.open("tv", LAN).err(), Some(OpenError::Busy));
696 let mapped: IpAddr = "::ffff:192.168.1.20".parse().unwrap();
698 assert_eq!(p.open("tv", mapped).err(), Some(OpenError::Busy));
699 assert!(p.open("tv", host(1)).is_ok());
700
701 let v6: Vec<_> = ["2001:db8:1:2::1", "2001:db8:1:2::2", "2001:db8:1:2:ffff::9"]
703 .into_iter()
704 .map(|a| p.open("tv", a.parse().unwrap()).unwrap())
705 .collect();
706 assert_eq!(
707 p.open("tv", "2001:db8:1:2:abcd::1".parse().unwrap()).err(),
708 Some(OpenError::Busy)
709 );
710 assert!(p.open("tv", "2001:db8:1:3::1".parse().unwrap()).is_ok());
711
712 drop(held);
713 assert!(p.open("tv", LAN).is_ok());
714 drop(v6);
715 }
716}