use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::LazyLock;
use std::time::{Duration, Instant};
use axum::extract::RawQuery;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use koan_core::remote::pair::PairMessage;
use parking_lot::Mutex;
use tokio::sync::oneshot;
use crate::auth::routes::RateLimiter;
pub const TTL: Duration = Duration::from_secs(600);
const MAX_PENDING: usize = 256;
const MAX_PENDING_PER_ADDRESS: usize = 3;
const CODE_ALPHABET: &[u8; 32] = b"0123456789ABCDEFGHJKMNPQRSTVWXYZ";
const CODE_LEN: usize = 8;
const KEEPALIVE: Duration = Duration::from_secs(30);
struct Entry {
code: String,
device: String,
from: IpAddr,
created: Instant,
outcome: oneshot::Sender<PairMessage>,
}
pub struct Pairings {
ttl: Duration,
entries: Mutex<HashMap<String, Entry>>,
}
pub fn pairings() -> &'static Pairings {
static PAIRINGS: LazyLock<Pairings> = LazyLock::new(|| Pairings::new(ttl_from_env()));
&PAIRINGS
}
fn ttl_from_env() -> Duration {
std::env::var("KOAN_PAIR_TTL_SECS")
.ok()
.and_then(|s| s.trim().parse().ok())
.filter(|&secs| secs > 0)
.map_or(TTL, Duration::from_secs)
}
static OPENS: LazyLock<RateLimiter> = LazyLock::new(|| RateLimiter::new(60, 10));
pub struct Opened<'a> {
pub id: String,
pub code: String,
outcome: oneshot::Receiver<PairMessage>,
pairings: &'a Pairings,
}
impl Drop for Opened<'_> {
fn drop(&mut self) {
self.pairings.entries.lock().remove(&self.id);
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PairInfo {
pub device: String,
pub from: IpAddr,
}
impl PairInfo {
pub fn local(&self) -> bool {
is_local(self.from)
}
}
fn is_local(ip: IpAddr) -> bool {
match ip.to_canonical() {
IpAddr::V4(v4) => {
let [a, b, ..] = v4.octets();
v4.is_private()
|| (a == 100 && (b & 0xc0) == 64)
|| v4.is_link_local()
|| v4.is_loopback()
}
IpAddr::V6(v6) => {
v6.is_loopback()
|| (v6.segments()[0] & 0xfe00) == 0xfc00
|| (v6.segments()[0] & 0xffc0) == 0xfe80
}
}
}
fn network(ip: IpAddr) -> IpAddr {
match ip {
IpAddr::V4(_) => ip,
IpAddr::V6(v6) => IpAddr::V6((u128::from(v6) & !((1u128 << 64) - 1)).into()),
}
}
pub struct Taken {
id: String,
entry: Entry,
}
impl Taken {
pub fn device(&self) -> &str {
&self.entry.device
}
pub fn info(&self) -> PairInfo {
PairInfo {
device: self.entry.device.clone(),
from: self.entry.from,
}
}
pub fn settle(self, outcome: PairMessage) -> Result<(), PairMessage> {
self.entry.outcome.send(outcome)
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum OpenError {
Full,
Busy,
Entropy,
}
impl Pairings {
fn new(ttl: Duration) -> Self {
Self {
ttl,
entries: Mutex::default(),
}
}
pub fn open(&self, device: &str, from: IpAddr) -> Result<Opened<'_>, OpenError> {
let id = koan_core::auth::random_api_key().map_err(|_| OpenError::Entropy)?;
let (tx, rx) = oneshot::channel();
let mut entries = self.entries.lock();
self.sweep(&mut entries);
if entries.len() >= MAX_PENDING {
return Err(OpenError::Full);
}
let from = from.to_canonical();
let waiting = entries
.values()
.filter(|e| network(e.from) == network(from))
.count();
if waiting >= MAX_PENDING_PER_ADDRESS {
return Err(OpenError::Busy);
}
let code = loop {
let code = new_code()?;
if !entries.values().any(|e| e.code == code) {
break code;
}
};
entries.insert(
id.clone(),
Entry {
code: code.clone(),
device: koan_core::invite::device_name(device),
from,
created: Instant::now(),
outcome: tx,
},
);
Ok(Opened {
id,
code: format!("{}-{}", &code[..4], &code[4..]),
outcome: rx,
pairings: self,
})
}
pub fn info(&self, pair: &str) -> Option<PairInfo> {
let mut entries = self.entries.lock();
self.sweep(&mut entries);
let id = find(&entries, pair)?;
entries.get(&id).map(|e| PairInfo {
device: e.device.clone(),
from: e.from,
})
}
pub fn take(&self, pair: &str) -> Option<Taken> {
let mut entries = self.entries.lock();
self.sweep(&mut entries);
let id = find(&entries, pair)?;
let entry = entries.remove(&id)?;
Some(Taken { id, entry })
}
pub fn put_back(&self, taken: Taken) {
if !taken.entry.outcome.is_closed() {
self.entries.lock().insert(taken.id, taken.entry);
}
}
pub fn settle(
&self,
conn: &rusqlite::Connection,
pair: &str,
user_id: i64,
username: &str,
decline: bool,
) -> Result<PairInfo, SettleError> {
use koan_core::db::queries::api_keys;
let taken = self.take(pair).ok_or(SettleError::NotFound)?;
let info = taken.info();
let device = info.device.clone();
if decline {
let _ = taken.settle(PairMessage::Declined);
log::info!("pair: {device} declined by {username}");
return Ok(info);
}
let api_key = match api_keys::create_api_key(conn, user_id, &device) {
Ok((_, key)) => key,
Err(e) => {
self.put_back(taken);
return Err(SettleError::Internal(e.to_string()));
}
};
let approved = PairMessage::Approved {
username: username.to_owned(),
api_key: api_key.clone(),
};
if taken.settle(approved).is_err() {
let _ = api_keys::revoke_api_key_value(conn, &api_key);
return Err(SettleError::NotFound);
}
log::info!("pair: {device} signed in as {username}");
Ok(info)
}
fn sweep(&self, entries: &mut HashMap<String, Entry>) {
let lapsed: Vec<String> = entries
.iter()
.filter(|(_, e)| e.created.elapsed() >= self.ttl)
.map(|(id, _)| id.clone())
.collect();
for id in lapsed {
if let Some(e) = entries.remove(&id) {
let _ = e.outcome.send(PairMessage::Expired);
}
}
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum SettleError {
NotFound,
Internal(String),
}
fn new_code() -> Result<String, OpenError> {
let mut bytes = [0u8; CODE_LEN];
getrandom::fill(&mut bytes).map_err(|_| OpenError::Entropy)?;
Ok(bytes
.iter()
.map(|b| CODE_ALPHABET[(b % 32) as usize] as char)
.collect())
}
fn normalise_code(typed: &str) -> Option<String> {
let code: String = typed
.chars()
.filter(|c| !matches!(c, '-' | ' '))
.map(|c| match c.to_ascii_uppercase() {
'I' | 'L' => '1',
'O' => '0',
c => c,
})
.collect();
(code.len() == CODE_LEN && code.bytes().all(|b| CODE_ALPHABET.contains(&b))).then_some(code)
}
fn find(entries: &HashMap<String, Entry>, pair: &str) -> Option<String> {
let pair = pair.trim();
if entries.contains_key(pair) {
return Some(pair.to_owned());
}
let code = normalise_code(pair)?;
entries
.iter()
.find(|(_, e)| e.code == code)
.map(|(id, _)| id.clone())
}
pub(crate) async fn route(
RawQuery(raw): RawQuery,
ws: WebSocketUpgrade,
request: axum::extract::Request,
) -> Response {
if request.headers().contains_key(axum::http::header::ORIGIN) {
return (StatusCode::FORBIDDEN, "pairing is not open to web pages").into_response();
}
let from = crate::auth::routes::client_ip(&request);
if !OPENS.allow(from) {
return (StatusCode::TOO_MANY_REQUESTS, "too many pairings").into_response();
}
let name = form_urlencoded::parse(raw.unwrap_or_default().as_bytes())
.find(|(k, _)| k == "name")
.map(|(_, v)| v.into_owned())
.unwrap_or_default();
match pairings().open(&name, from) {
Ok(opened) => ws.on_upgrade(move |socket| session(socket, opened)),
Err(OpenError::Full) => {
(StatusCode::SERVICE_UNAVAILABLE, "too many pairings waiting").into_response()
}
Err(OpenError::Busy) => (
StatusCode::TOO_MANY_REQUESTS,
"too many pairings waiting from this address",
)
.into_response(),
Err(OpenError::Entropy) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
}
}
async fn session(mut socket: WebSocket, mut opened: Opened<'static>) {
let pending = PairMessage::Pending {
id: opened.id.clone(),
code: opened.code.clone(),
expires_in: opened.pairings.ttl.as_secs(),
};
if send(&mut socket, &pending).await.is_err() {
return;
}
log::info!("pair: {} waiting", opened.code);
let lapse = tokio::time::sleep(opened.pairings.ttl);
tokio::pin!(lapse);
let mut keepalive =
tokio::time::interval_at(tokio::time::Instant::now() + KEEPALIVE, KEEPALIVE);
let outcome = loop {
tokio::select! {
outcome = &mut opened.outcome => break outcome.unwrap_or(PairMessage::Expired),
() = &mut lapse => break PairMessage::Expired,
_ = keepalive.tick() => {
if socket.send(Message::Ping(Default::default())).await.is_err() {
return;
}
}
msg = socket.recv() => match msg {
Some(Ok(Message::Close(_)) | Err(_)) | None => return,
Some(Ok(_)) => {}
},
}
};
log::info!(
"pair: {} {}",
opened.code,
match outcome {
PairMessage::Approved { .. } => "approved",
PairMessage::Declined => "declined",
_ => "expired",
}
);
let _ = send(&mut socket, &outcome).await;
let _ = socket.send(Message::Close(None)).await;
}
async fn send(socket: &mut WebSocket, message: &PairMessage) -> Result<(), axum::Error> {
let text = serde_json::to_string(message).unwrap_or_default();
socket.send(Message::Text(text.into())).await
}
#[cfg(test)]
mod tests {
use super::*;
const LAN: IpAddr = IpAddr::V4(std::net::Ipv4Addr::new(192, 168, 1, 20));
fn host(i: usize) -> IpAddr {
IpAddr::V4(std::net::Ipv4Addr::new(
10,
9,
(i / 256) as u8,
(i % 256) as u8,
))
}
fn leaked(ttl: Duration) -> &'static Pairings {
Box::leak(Box::new(Pairings::new(ttl)))
}
#[test]
fn codes_are_crockford_and_dashed() {
let p = leaked(TTL);
for _ in 0..50 {
let o = p.open("tv", LAN).unwrap();
let (a, b) = o.code.split_once('-').unwrap();
assert_eq!((a.len(), b.len()), (4, 4));
assert!(
format!("{a}{b}")
.bytes()
.all(|c| CODE_ALPHABET.contains(&c) && !b"ILOU".contains(&c))
);
}
}
#[test]
fn ids_are_unguessable_and_distinct() {
let p = leaked(TTL);
let opened: Vec<_> = (0..100).map(|i| p.open("tv", host(i)).unwrap()).collect();
let ids: std::collections::HashSet<_> = opened.iter().map(|o| o.id.clone()).collect();
assert_eq!(ids.len(), 100);
assert!(opened.iter().all(|o| o.id.len() == 43));
let codes: std::collections::HashSet<_> = opened.iter().map(|o| o.code.clone()).collect();
assert_eq!(codes.len(), 100);
}
#[test]
fn a_code_is_found_however_it_is_typed() {
let p = leaked(TTL);
let o = p.open("Living room\u{7} TV", LAN).unwrap();
assert_eq!(
p.info(&o.id).map(|i| i.device).as_deref(),
Some("Living room TV")
);
let bare = o.code.replace('-', "");
for typed in [
o.code.clone(),
o.code.to_lowercase(),
bare.clone(),
format!(" {} {} ", &bare[..4], &bare[4..]),
] {
assert_eq!(
p.info(&typed).map(|i| i.device).as_deref(),
Some("Living room TV"),
"{typed}"
);
}
assert_eq!(normalise_code("o1l1-IOAB").as_deref(), Some("011110AB"));
assert_eq!(normalise_code("ABCD-EFGU"), None);
assert_eq!(normalise_code("ABC"), None);
}
#[tokio::test]
async fn approving_delivers_the_key_to_the_waiting_device() {
let p = leaked(TTL);
let mut o = p.open("tv", LAN).unwrap();
let taken = p.take(&o.code).unwrap();
assert_eq!(taken.device(), "tv");
assert!(p.take(&o.id).is_none());
taken
.settle(PairMessage::Approved {
username: "alice".into(),
api_key: "key".into(),
})
.unwrap();
assert_eq!(
(&mut o.outcome).await.unwrap(),
PairMessage::Approved {
username: "alice".into(),
api_key: "key".into(),
}
);
}
fn database() -> (tempfile::TempDir, koan_core::db::connection::Database, i64) {
let dir = tempfile::tempdir().unwrap();
let db = koan_core::db::connection::Database::open(&dir.path().join("koan.db")).unwrap();
let id = koan_core::db::queries::auth::create_user(
&db.conn,
"alice",
"hunter2",
koan_core::auth::Role::Readonly,
)
.unwrap();
(dir, db, id)
}
#[tokio::test]
async fn approving_mints_a_key_for_the_approver() {
let (_dir, db, alice) = database();
let p = leaked(TTL);
let mut o = p.open("Living room TV", LAN).unwrap();
let settled = p.settle(&db.conn, &o.code, alice, "alice", false).unwrap();
assert_eq!(
settled,
PairInfo {
device: "Living room TV".into(),
from: LAN,
}
);
let PairMessage::Approved { username, api_key } = (&mut o.outcome).await.unwrap() else {
panic!("not approved");
};
assert_eq!(username, "alice");
let user = koan_core::db::queries::api_keys::authenticate_api_key(&db.conn, &api_key)
.unwrap()
.unwrap();
assert_eq!(user.id, alice);
let keys = koan_core::db::queries::api_keys::list_api_keys(&db.conn, Some(alice)).unwrap();
assert_eq!(keys[0].name, "Living room TV");
assert_eq!(
p.settle(&db.conn, &o.id, alice, "alice", false),
Err(SettleError::NotFound)
);
}
#[tokio::test]
async fn declining_mints_nothing() {
let (_dir, db, alice) = database();
let p = leaked(TTL);
let mut o = p.open("tv", LAN).unwrap();
p.settle(&db.conn, &o.id, alice, "alice", true).unwrap();
assert_eq!((&mut o.outcome).await.unwrap(), PairMessage::Declined);
let keys = koan_core::db::queries::api_keys::list_api_keys(&db.conn, Some(alice)).unwrap();
assert!(keys.is_empty());
}
#[test]
fn a_key_for_a_gone_device_is_revoked() {
let (_dir, db, alice) = database();
let p = leaked(TTL);
let o = p.open("tv", LAN).unwrap();
let id = o.id.clone();
let taken = p.take(&id).unwrap();
drop(o);
p.entries.lock().insert(taken.id, taken.entry);
assert_eq!(
p.settle(&db.conn, &id, alice, "alice", false),
Err(SettleError::NotFound)
);
let keys = koan_core::db::queries::api_keys::list_api_keys(&db.conn, Some(alice)).unwrap();
assert!(keys.is_empty());
}
#[test]
fn approving_an_unknown_pairing_is_not_found() {
let (_dir, db, alice) = database();
let p = leaked(TTL);
assert_eq!(
p.settle(&db.conn, "ABCD-EFGH", alice, "alice", false),
Err(SettleError::NotFound)
);
}
#[test]
fn a_gone_device_gives_the_outcome_back() {
let p = leaked(TTL);
let o = p.open("tv", LAN).unwrap();
let taken = p.take(&o.id).unwrap();
drop(o);
assert!(taken.settle(PairMessage::Declined).is_err());
}
#[tokio::test]
async fn a_lapsed_pairing_is_expired_and_unknown() {
let p = leaked(Duration::ZERO);
let mut o = p.open("tv", LAN).unwrap();
assert!(p.info(&o.id).is_none());
assert!(p.take(&o.code).is_none());
assert_eq!((&mut o.outcome).await.unwrap(), PairMessage::Expired);
}
#[test]
fn the_requesting_address_is_kept() {
let p = leaked(TTL);
let mapped: IpAddr = "::ffff:203.0.113.9".parse().unwrap();
let o = p.open("tv", mapped).unwrap();
let info = p.info(&o.code).unwrap();
assert_eq!(info.from.to_string(), "203.0.113.9");
assert!(!info.local());
assert_eq!(p.take(&o.id).unwrap().info(), info);
}
#[test]
fn private_addresses_are_local() {
for local in [
"10.1.2.3",
"172.16.0.1",
"192.168.1.20",
"169.254.10.1",
"127.0.0.1",
"100.64.0.1",
"100.101.102.103",
"100.127.255.254",
"::1",
"fd12:3456::1",
"fe80::1",
"::ffff:192.168.0.5",
] {
assert!(is_local(local.parse().unwrap()), "{local}");
}
for public in [
"203.0.113.9",
"8.8.8.8",
"172.32.0.1",
"100.63.255.255",
"100.128.0.1",
"2001:db8::1",
"2a00:1450::1",
"::ffff:8.8.8.8",
] {
assert!(!is_local(public.parse().unwrap()), "{public}");
}
}
#[test]
fn unknown_pairings_are_not_found() {
let p = leaked(TTL);
let o = p.open("tv", LAN).unwrap();
let other = if o.code == "0000-0000" {
"1111-1111"
} else {
"0000-0000"
};
assert!(p.info(other).is_none());
assert!(p.take("not-a-pairing").is_none());
assert!(p.info("").is_none());
}
#[test]
fn a_closed_socket_removes_its_pairing() {
let p = leaked(TTL);
let o = p.open("tv", LAN).unwrap();
let id = o.id.clone();
drop(o);
assert!(p.info(&id).is_none());
}
#[test]
fn pairings_are_capped() {
let p = leaked(TTL);
let held: Vec<_> = (0..MAX_PENDING)
.map(|i| p.open("tv", host(i)).unwrap())
.collect();
assert_eq!(p.open("tv", LAN).err(), Some(OpenError::Full));
drop(held);
assert!(p.open("tv", LAN).is_ok());
}
#[test]
fn one_address_holds_three_at_most() {
let p = leaked(TTL);
let held: Vec<_> = (0..3).map(|_| p.open("tv", LAN).unwrap()).collect();
assert_eq!(p.open("tv", LAN).err(), Some(OpenError::Busy));
let mapped: IpAddr = "::ffff:192.168.1.20".parse().unwrap();
assert_eq!(p.open("tv", mapped).err(), Some(OpenError::Busy));
assert!(p.open("tv", host(1)).is_ok());
let v6: Vec<_> = ["2001:db8:1:2::1", "2001:db8:1:2::2", "2001:db8:1:2:ffff::9"]
.into_iter()
.map(|a| p.open("tv", a.parse().unwrap()).unwrap())
.collect();
assert_eq!(
p.open("tv", "2001:db8:1:2:abcd::1".parse().unwrap()).err(),
Some(OpenError::Busy)
);
assert!(p.open("tv", "2001:db8:1:3::1".parse().unwrap()).is_ok());
drop(held);
assert!(p.open("tv", LAN).is_ok());
drop(v6);
}
}