use ipnetwork::Ipv4Network;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::net::{Ipv4Addr, SocketAddr};
use std::time::Duration;
use tokio::net::UdpSocket;
use tokio::time::{Instant, timeout_at};
const PORT: u16 = 9999;
const QUERY: &str = r#"{"system":{"get_sysinfo":{}}}"#;
#[derive(Clone, Serialize, Deserialize)]
pub struct KasaInfo {
pub alias: Option<String>,
pub model: Option<String>,
#[serde(alias = "dev_name")]
pub description: Option<String>,
#[serde(rename = "type")]
pub device_type: Option<String>,
#[serde(skip_serializing)]
mic_type: Option<String>,
}
pub async fn discover(
local_ip: Ipv4Addr,
net: Ipv4Network,
wait: Duration,
) -> HashMap<Ipv4Addr, KasaInfo> {
let Ok(sock) = UdpSocket::bind((local_ip, 0)).await else {
return HashMap::new();
};
if sock.set_broadcast(true).is_err() {
return HashMap::new();
}
let query = xor_encrypt(QUERY.as_bytes());
let _ = sock.send_to(&query, (net.broadcast(), PORT)).await;
let start = Instant::now();
collect(&sock, net, start + wait, Some(start + wait / 3)).await
}
pub async fn query(
local_ip: Ipv4Addr,
net: Ipv4Network,
targets: &[Ipv4Addr],
wait: Duration,
) -> HashMap<Ipv4Addr, KasaInfo> {
if targets.is_empty() {
return HashMap::new();
}
let Ok(sock) = UdpSocket::bind((local_ip, 0)).await else {
return HashMap::new();
};
let query = xor_encrypt(QUERY.as_bytes());
for &ip in targets {
let _ = sock.send_to(&query, (ip, PORT)).await;
}
collect(&sock, net, Instant::now() + wait, None).await
}
async fn collect(
sock: &UdpSocket,
net: Ipv4Network,
deadline: Instant,
mut rebroadcast_at: Option<Instant>,
) -> HashMap<Ipv4Addr, KasaInfo> {
let mut found = HashMap::new();
let mut buf = vec![0u8; 4096];
loop {
let until = rebroadcast_at.unwrap_or(deadline);
match timeout_at(until, sock.recv_from(&mut buf)).await {
Ok(Ok((n, SocketAddr::V4(src)))) if net.contains(*src.ip()) => {
if let Some(info) = parse(&xor_decrypt(&buf[..n])) {
found.insert(*src.ip(), info);
}
}
Ok(_) => {}
Err(_) if rebroadcast_at.is_some() => {
rebroadcast_at = None;
let query = xor_encrypt(QUERY.as_bytes());
let _ = sock.send_to(&query, (net.broadcast(), PORT)).await;
}
Err(_) => break,
}
}
found
}
fn parse(json: &[u8]) -> Option<KasaInfo> {
#[derive(Deserialize)]
struct Reply {
system: System,
}
#[derive(Deserialize)]
struct System {
get_sysinfo: KasaInfo,
}
let mut info = serde_json::from_slice::<Reply>(json)
.ok()?
.system
.get_sysinfo;
info.device_type = info.device_type.or(info.mic_type.take());
Some(info)
}
fn xor_encrypt(plain: &[u8]) -> Vec<u8> {
let mut key = 171u8;
plain
.iter()
.map(|&b| {
key ^= b;
key
})
.collect()
}
fn xor_decrypt(cipher: &[u8]) -> Vec<u8> {
let mut key = 171u8;
cipher
.iter()
.map(|&b| {
let plain = key ^ b;
key = b;
plain
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip() {
assert_eq!(
xor_decrypt(&xor_encrypt(QUERY.as_bytes())),
QUERY.as_bytes()
);
assert_eq!(xor_encrypt(b"{")[0], 0xd0);
}
#[test]
fn sysinfo() {
let reply = br#"{"system":{"get_sysinfo":{"sw_ver":"1.0.6","model":"HS105(US)",
"dev_name":"Smart Wi-Fi Plug Mini","alias":"Living Room Lamp",
"mic_type":"IOT.SMARTPLUGSWITCH","relay_state":1}}}"#;
let info = parse(reply).unwrap();
assert_eq!(info.alias.as_deref(), Some("Living Room Lamp"));
assert_eq!(info.model.as_deref(), Some("HS105(US)"));
assert_eq!(info.description.as_deref(), Some("Smart Wi-Fi Plug Mini"));
assert_eq!(info.device_type.as_deref(), Some("IOT.SMARTPLUGSWITCH"));
assert!(parse(b"{}").is_none());
}
}