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