use crate::crypto::TuyaCipher;
use crate::error::{Result, TuyaError};
use crate::protocol::{self, CommandType, PREFIX_6699, TuyaMessage, Version};
use log::{debug, info, trace, warn};
use parking_lot::RwLock;
use serde_json::Value;
use socket2::{Domain, Protocol, SockAddr, Socket, Type};
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::str::FromStr;
use std::sync::Arc;
use std::sync::OnceLock;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use tokio::net::UdpSocket;
use tokio::sync::{mpsc, watch};
use tokio::time::{interval, sleep, timeout};
use serde::Serialize;
#[derive(Debug, Clone, PartialEq, Serialize)]
pub struct DiscoveryResult {
pub id: String,
pub ip: String,
pub version: Option<Version>,
pub product_key: Option<String>,
#[serde(skip)]
pub discovered_at: Instant,
}
impl DiscoveryResult {
#[must_use]
pub fn is_same_device(&self, other: &Self) -> bool {
self.id == other.id
&& self.ip == other.ip
&& self.version == other.version
&& self.product_key == other.product_key
}
}
const UDP_KEY_34: &[u8] = &[
0x6c, 0x1e, 0xc8, 0xe2, 0xbb, 0x9b, 0xb5, 0x9a, 0xb5, 0x0b, 0x0d, 0xaf, 0x64, 0x9b, 0x41, 0x0a,
];
const UDP_KEY_35: &[u8] = UDP_KEY_34;
const UDP_KEY_33: &[u8] = b"yG9shRKIBrIBUjc3";
const BROADCAST_INTERVAL: Duration = Duration::from_secs(6);
const MAX_BROADCASTS: u32 = 3;
const RECEIVE_MARGIN: Duration = BROADCAST_INTERVAL;
const DEFAULT_SCAN_TIMEOUT: Duration = Duration::from_secs(
BROADCAST_INTERVAL.as_secs() * (MAX_BROADCASTS as u64 - 1) + RECEIVE_MARGIN.as_secs(),
);
const GLOBAL_SCAN_COOLDOWN: Duration = Duration::from_secs(1800); const SCAN_THROTTLE_INTERVAL: Duration = Duration::from_secs(60); const CACHE_TTL: Duration = Duration::from_secs(24 * 60 * 60);
const PACKET_CHANNEL_CAPACITY: usize = 1024;
type PacketSender = mpsc::Sender<(Vec<u8>, SocketAddr)>;
type PacketReceiver = mpsc::Receiver<(Vec<u8>, SocketAddr)>;
#[derive(Debug)]
struct ScannerState {
cache: RwLock<HashMap<String, DiscoveryResult>>,
discovery_version: (watch::Sender<u64>, watch::Receiver<u64>),
active_scanning: AtomicBool,
last_scan_time: RwLock<Option<Instant>>,
listener_started: AtomicBool,
cancel_token: RwLock<tokio_util::sync::CancellationToken>,
sockets: RwLock<HashMap<u16, Arc<UdpSocket>>>,
packet_tx: RwLock<Option<PacketSender>>,
receiver_tasks: RwLock<HashMap<u16, tokio::task::JoinHandle<()>>>,
dispatcher_task: RwLock<Option<tokio::task::JoinHandle<()>>>,
startup_guard: parking_lot::Mutex<()>,
timeout: RwLock<Duration>,
bind_addr: RwLock<String>,
ports: RwLock<Vec<u16>>,
discovery_sources: RwLock<Vec<IpAddr>>,
}
impl ScannerState {
fn new() -> Self {
Self {
cache: RwLock::new(HashMap::new()),
discovery_version: watch::channel(0u64),
active_scanning: AtomicBool::new(false),
last_scan_time: RwLock::new(None),
listener_started: AtomicBool::new(false),
cancel_token: RwLock::new(tokio_util::sync::CancellationToken::new()),
sockets: RwLock::new(HashMap::new()),
packet_tx: RwLock::new(None),
receiver_tasks: RwLock::new(HashMap::new()),
dispatcher_task: RwLock::new(None),
startup_guard: parking_lot::Mutex::new(()),
timeout: RwLock::new(DEFAULT_SCAN_TIMEOUT),
bind_addr: RwLock::new("0.0.0.0".to_string()),
ports: RwLock::new(vec![6666, 6667, 7000]),
discovery_sources: RwLock::new(Vec::new()),
}
}
fn current_cancel_token(&self) -> tokio_util::sync::CancellationToken {
self.cancel_token.read().clone()
}
fn reset_cancel_token(&self) {
let mut guard = self.cancel_token.write();
guard.cancel();
*guard = tokio_util::sync::CancellationToken::new();
}
fn current_packet_tx(&self) -> Option<PacketSender> {
self.packet_tx.read().clone()
}
fn publish_discovery(&self) {
self.discovery_version
.0
.send_modify(|v| *v = v.wrapping_add(1));
}
fn subscribe_discoveries(&self) -> watch::Receiver<u64> {
let mut rx = self.discovery_version.1.clone();
rx.mark_unchanged();
rx
}
}
impl Drop for ScannerState {
fn drop(&mut self) {
self.cancel_token.write().cancel();
for (_, task) in self.receiver_tasks.write().drain() {
task.abort();
}
if let Some(t) = self.dispatcher_task.write().take() {
t.abort();
}
}
}
fn effective_bind_ip(configured: &str) -> Option<(IpAddr, bool)> {
let ip: IpAddr = configured.parse().ok()?;
let keep_exact = match ip {
IpAddr::V4(v4) => v4.is_unspecified() || v4.is_loopback(),
IpAddr::V6(v6) => v6.is_unspecified() || v6.is_loopback(),
};
if keep_exact {
return Some((ip, false));
}
let wildcard = match ip {
IpAddr::V4(_) => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
IpAddr::V6(_) => IpAddr::V6(Ipv6Addr::UNSPECIFIED),
};
Some((wildcard, true))
}
fn compute_port_diff(desired: &[u16], current: &[u16]) -> (Vec<u16>, Vec<u16>) {
let to_add: Vec<u16> = desired
.iter()
.filter(|p| !current.contains(p))
.copied()
.collect();
let to_remove: Vec<u16> = current
.iter()
.filter(|p| !desired.contains(p))
.copied()
.collect();
(to_add, to_remove)
}
#[derive(Debug, Clone)]
pub struct Scanner {
inner: Arc<ScannerState>,
}
static GLOBAL_SCANNER: OnceLock<Scanner> = OnceLock::new();
pub fn get() -> &'static Scanner {
GLOBAL_SCANNER.get_or_init(Scanner::new)
}
impl Scanner {
pub fn get() -> &'static Self {
get()
}
#[must_use]
pub(crate) fn new() -> Self {
let scanner = Self {
inner: Arc::new(ScannerState::new()),
};
scanner.ensure_passive_listener();
scanner
}
fn ensure_passive_listener(&self) {
let _guard = self.inner.startup_guard.lock();
self.reconcile_listener_locked();
}
fn reconcile_listener_locked(&self) {
let state = &self.inner;
let ports_snapshot: Vec<u16> = state.ports.read().clone();
let current_ports: Vec<u16> = state.sockets.read().keys().copied().collect();
let (ports_to_add, ports_to_remove) = compute_port_diff(&ports_snapshot, ¤t_ports);
if !ports_to_remove.is_empty() {
let mut tasks = state.receiver_tasks.write();
let mut sockets = state.sockets.write();
for port in &ports_to_remove {
if let Some(task) = tasks.remove(port) {
task.abort();
}
sockets.remove(port);
}
debug!(
"Passive listener: removed receivers for ports {:?}",
ports_to_remove
);
}
if ports_to_add.is_empty() && state.listener_started.load(Ordering::SeqCst) {
return;
}
let bind_addr = state.bind_addr.read().clone();
let Some((bind_ip, widened)) = effective_bind_ip(&bind_addr) else {
warn!("Invalid scanner bind address {bind_addr:?}; cannot bind listener sockets");
return;
};
if widened {
warn!(
"Scanner bind address {bind_addr} is a specific unicast IP; a socket bound to it \
does not receive limited-broadcast (255.255.255.255) discovery packets. \
Listening on the {} wildcard instead so passive discovery keeps working. To pin \
the source of active-scan broadcasts, use set_discovery_sources().",
if bind_ip.is_ipv4() { "0.0.0.0" } else { "::" }
);
}
let mut new_sockets: Vec<(u16, Arc<UdpSocket>)> = Vec::new();
{
let mut guard = state.sockets.write();
for port in ports_to_add {
if let Ok(socket) = Self::create_udp_socket(bind_ip, port) {
let arc_socket = Arc::new(socket);
guard.insert(port, arc_socket.clone());
new_sockets.push((port, arc_socket));
}
}
}
if new_sockets.is_empty() && state.listener_started.load(Ordering::SeqCst) {
return;
}
if new_sockets.is_empty() {
warn!(
"Passive listener failed to bind to any ports: {:?}",
ports_snapshot
);
return;
}
if !state.listener_started.swap(true, Ordering::SeqCst) {
let cancel_token = state.current_cancel_token();
let (tx, mut rx): (PacketSender, PacketReceiver) =
mpsc::channel(PACKET_CHANNEL_CAPACITY);
*state.packet_tx.write() = Some(tx.clone());
let recv_tasks = Self::spawn_receiver_tasks(new_sockets, tx, cancel_token.clone());
state.receiver_tasks.write().extend(recv_tasks);
let state_weak = Arc::downgrade(&self.inner);
let dispatcher_ct = cancel_token.clone();
let dispatcher = crate::runtime::spawn(async move {
debug!("Starting background passive listener task...");
loop {
tokio::select! {
() = dispatcher_ct.cancelled() => break,
msg = rx.recv() => {
let Some((data, _addr)) = msg else { break };
let Some(state) = state_weak.upgrade() else { break };
Self::dispatch_packet(&state, &data);
}
}
}
debug!("Background passive listener task stopped");
});
*state.dispatcher_task.write() = Some(dispatcher);
} else if let Some(tx) = state.current_packet_tx() {
let recv_tasks =
Self::spawn_receiver_tasks(new_sockets, tx, state.current_cancel_token());
state.receiver_tasks.write().extend(recv_tasks);
} else {
debug_assert!(
false,
"ScannerState invariant violated: listener_started=true but packet_tx=None"
);
warn!(
"Passive listener state inconsistency (listener_started=true, packet_tx=None); \
forcing rebuild"
);
state.listener_started.store(false, Ordering::SeqCst);
drop(new_sockets);
self.reconcile_listener_locked();
}
}
fn dispatch_packet(state: &Arc<ScannerState>, data: &[u8]) {
let Some(res) = parse_packet(data) else {
return;
};
let mut guard = state.cache.write();
guard.retain(|_, v| v.discovered_at.elapsed() < CACHE_TTL);
let should_log = match guard.get(&res.id) {
Some(existing) => !res.is_same_device(existing),
None => true,
};
if should_log {
let mode = if state.active_scanning.load(Ordering::SeqCst) {
"A"
} else {
"P"
};
let version = res
.version
.map_or_else(|| "unknown".to_string(), |v| v.to_string());
info!(
"Discovered device {}(v{}) at {} - {}",
res.id, version, res.ip, mode
);
}
guard.insert(res.id.clone(), res.clone());
drop(guard);
state.publish_discovery();
}
fn spawn_receiver_tasks(
sockets: Vec<(u16, Arc<UdpSocket>)>,
tx: PacketSender,
cancel_token: tokio_util::sync::CancellationToken,
) -> Vec<(u16, tokio::task::JoinHandle<()>)> {
let mut tasks = Vec::new();
for (port, socket) in sockets {
let tx = tx.clone();
let socket = socket.clone();
let ct = cancel_token.clone();
let task = crate::runtime::spawn(async move {
let mut buf = vec![0u8; 4096];
let local_addr = socket.local_addr().ok();
loop {
tokio::select! {
() = ct.cancelled() => break,
res = socket.recv_from(&mut buf) => {
match res {
Ok((len, addr)) => {
if tx.send((buf[..len].to_vec(), addr)).await.is_err() {
break;
}
}
Err(e) => {
warn!(
"Scanner UDP recv error on {:?}: {} (continuing)",
local_addr, e
);
tokio::select! {
() = ct.cancelled() => break,
() = sleep(Duration::from_millis(500)) => {}
}
}
}
}
}
}
});
tasks.push((port, task));
}
tasks
}
fn create_udp_socket(bind_ip: IpAddr, port: u16) -> Result<UdpSocket> {
let addr = SocketAddr::new(bind_ip, port);
let socket = Socket::new(Domain::for_address(addr), Type::DGRAM, Some(Protocol::UDP))?;
let _ = socket.set_reuse_address(true);
let _ = socket.set_broadcast(true);
socket.bind(&SockAddr::from(addr))?;
socket.set_nonblocking(true)?;
let std_socket: std::net::UdpSocket = socket.into();
let _guard = crate::runtime::get_runtime().enter();
Ok(UdpSocket::from_std(std_socket)?)
}
fn create_send_socket(source_ip: IpAddr) -> Result<UdpSocket> {
let addr = SocketAddr::new(source_ip, 0);
let socket = Socket::new(Domain::for_address(addr), Type::DGRAM, Some(Protocol::UDP))?;
let _ = socket.set_reuse_address(true);
let _ = socket.set_broadcast(true);
socket.bind(&SockAddr::from(addr))?;
socket.set_nonblocking(true)?;
let std_socket: std::net::UdpSocket = socket.into();
let _guard = crate::runtime::get_runtime().enter();
Ok(UdpSocket::from_std(std_socket)?)
}
pub fn stop_passive_listener(&self) {
let _guard = self.inner.startup_guard.lock();
self.inner.reset_cancel_token();
self.inner.listener_started.store(false, Ordering::SeqCst);
self.inner.sockets.write().clear();
*self.inner.packet_tx.write() = None;
for (_, task) in self.inner.receiver_tasks.write().drain() {
task.abort();
}
if let Some(t) = self.inner.dispatcher_task.write().take() {
t.abort();
}
}
#[must_use]
pub fn timeout(&self) -> Duration {
*self.inner.timeout.read()
}
#[must_use]
pub fn bind_addr(&self) -> String {
self.inner.bind_addr.read().clone()
}
#[must_use]
pub fn ports(&self) -> Vec<u16> {
self.inner.ports.read().clone()
}
pub fn set_timeout(&self, timeout: Duration) {
*self.inner.timeout.write() = timeout;
}
pub fn set_ports(&self, ports: Vec<u16>) {
*self.inner.ports.write() = ports;
self.ensure_passive_listener();
}
#[must_use]
pub fn discovery_sources(&self) -> Vec<IpAddr> {
self.inner.discovery_sources.read().clone()
}
pub fn set_discovery_sources(&self, sources: Vec<IpAddr>) {
*self.inner.discovery_sources.write() = sources;
}
pub fn set_bind_address(&self, addr: &str) -> Result<()> {
let new_addr = addr.to_string();
let _guard = self.inner.startup_guard.lock();
let old_addr = {
let mut g = self.inner.bind_addr.write();
let prev = g.clone();
*g = new_addr.clone();
prev
};
let listener_up = self.inner.listener_started.load(Ordering::SeqCst);
if old_addr != new_addr && listener_up {
debug!(
"Rebinding passive listener sockets from {} to {}",
old_addr, new_addr
);
{
let mut tasks = self.inner.receiver_tasks.write();
for (_, task) in tasks.drain() {
task.abort();
}
}
self.inner.sockets.write().clear();
}
self.reconcile_listener_locked();
Ok(())
}
pub(crate) fn subscribe_discoveries(&self) -> watch::Receiver<u64> {
self.inner.subscribe_discoveries()
}
#[must_use]
pub fn get_cached_result(&self, device_id: &str) -> Option<DiscoveryResult> {
let guard = self.inner.cache.read();
guard.get(device_id).cloned()
}
#[must_use]
pub fn is_recently_discovered(&self, device_id: &str, within: Duration) -> bool {
let guard = self.inner.cache.read();
if let Some(res) = guard.get(device_id) {
return res.discovered_at.elapsed() < within;
}
false
}
fn discover_local_ip_blocking() -> Option<String> {
const CANDIDATES: &[&str] = &["8.8.8.8:80", "255.255.255.255:80", "203.0.113.1:80"];
for dst in CANDIDATES {
let Ok(socket) = std::net::UdpSocket::bind("0.0.0.0:0") else {
continue;
};
let _ = socket.set_broadcast(true); if socket.connect(dst).is_ok()
&& let Ok(addr) = socket.local_addr()
&& !addr.ip().is_unspecified()
{
return Some(addr.ip().to_string());
}
}
None
}
async fn send_discovery_broadcast(
&self,
socket: &UdpSocket,
port: u16,
source_override: Option<IpAddr>,
) -> Result<()> {
let local_ip = match source_override {
Some(ip) => ip.to_string(),
None => Self::discover_local_ip_blocking().unwrap_or_else(|| {
warn!(
"Local IP detection failed for discovery broadcast on port {port}; \
falling back to 0.0.0.0"
);
"0.0.0.0".to_string()
}),
};
debug!("Sending discovery broadcast on port {port} (local IP: {local_ip})");
let (payload, prefix) = if port == 7000 {
(
serde_json::json!({
"from": "app",
"ip": local_ip,
}),
PREFIX_6699,
)
} else {
(
serde_json::json!({
"gwId": "",
"devId": "",
}),
protocol::PREFIX_55AA,
)
};
let msg = TuyaMessage {
seqno: 0,
cmd: if port == 7000 {
CommandType::ReqDevInfo as u32
} else {
CommandType::UdpNew as u32
},
retcode: None,
payload: serde_json::to_vec(&payload)?,
prefix,
iv: None,
};
let packed =
protocol::pack_message(&msg, if port == 7000 { Some(UDP_KEY_35) } else { None })?;
let broadcast_addr: SocketAddr = format!("255.255.255.255:{port}")
.parse()
.map_err(|_| TuyaError::Offline)?;
match socket.send_to(&packed, broadcast_addr).await {
Ok(len) => debug!("Sent discovery broadcast to {broadcast_addr}: {len} bytes"),
Err(e) => warn!("Failed to send discovery broadcast to {broadcast_addr}: {e}"),
}
Ok(())
}
pub fn scan_stream() -> impl futures_util::Stream<Item = DiscoveryResult> + Send + 'static {
Self::get().scan_stream_instance()
}
pub fn scan_stream_instance(
&self,
) -> impl futures_util::Stream<Item = DiscoveryResult> + Send + 'static {
let state = self.inner.clone();
let timeout_dur = *state.timeout.read();
let start_time = Instant::now();
let scanner = self.clone();
let cooldown_ok = {
let last_scan = state.last_scan_time.read();
last_scan.is_none_or(|t| t.elapsed() >= GLOBAL_SCAN_COOLDOWN)
};
let should_start = cooldown_ok
&& state
.active_scanning
.compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
.is_ok();
let mut discovery_rx = state.subscribe_discoveries();
if should_start {
*state.last_scan_time.write() = Some(Instant::now());
let state_clone = state.clone();
crate::runtime::spawn(async move {
let _ = scanner.perform_discovery_loop().await;
state_clone.active_scanning.store(false, Ordering::SeqCst);
state_clone.publish_discovery();
});
}
async_stream::stream! {
let mut yielded_ids = std::collections::HashSet::new();
let initial_items: Vec<_> = {
let guard = state.cache.read();
guard.values().cloned().collect()
};
for item in initial_items {
yielded_ids.insert(item.id.clone());
yield item;
}
loop {
let elapsed = start_time.elapsed();
if elapsed >= timeout_dur {
break;
}
let remaining = timeout_dur.saturating_sub(elapsed);
tokio::select! {
() = sleep(remaining) => break,
_ = discovery_rx.changed() => {
let new_items: Vec<_> = {
let guard = state.cache.read();
guard.values()
.filter(|v| !yielded_ids.contains(&v.id))
.cloned()
.collect()
};
for item in new_items {
yielded_ids.insert(item.id.clone());
yield item;
}
if !state.active_scanning.load(Ordering::SeqCst) {
let final_items: Vec<_> = {
let guard = state.cache.read();
guard.values()
.filter(|v| !yielded_ids.contains(&v.id))
.cloned()
.collect()
};
for item in final_items {
yield item;
}
break;
}
}
}
}
}
}
pub async fn scan() -> Result<Vec<DiscoveryResult>> {
Self::get().scan_instance().await
}
pub async fn scan_instance(&self) -> Result<Vec<DiscoveryResult>> {
use futures_util::StreamExt;
info!(
"Starting Tuya device scan (addr: {}, ports: {:?})...",
self.inner.bind_addr.read(),
self.inner.ports.read()
);
let results: Vec<_> = self.scan_stream_instance().collect().await;
info!("Scan finished. Found {} devices.", results.len());
Ok(results)
}
pub async fn discover_device(device_id: &str) -> Result<Option<DiscoveryResult>> {
Self::get().discover_device_instance(device_id).await
}
pub async fn discover_device_instance(
&self,
device_id: &str,
) -> Result<Option<DiscoveryResult>> {
self.discover_device_internal(device_id, false, None).await
}
pub(crate) async fn discover_device_internal(
&self,
device_id: &str,
force_scan: bool,
cancel: Option<&tokio_util::sync::CancellationToken>,
) -> Result<Option<DiscoveryResult>> {
if let Some(res) = self.check_cache_and_cooldown(device_id, force_scan) {
return Ok(Some(res));
}
self.ensure_scan_started(device_id, force_scan).await;
Ok(self.wait_for_cache_result(device_id, cancel).await)
}
fn check_cache_and_cooldown(
&self,
device_id: &str,
force_scan: bool,
) -> Option<DiscoveryResult> {
let state = &self.inner;
let guard = state.cache.read();
if let Some(res) = guard.get(device_id).cloned()
&& !force_scan
&& res.discovered_at.elapsed() < GLOBAL_SCAN_COOLDOWN
{
debug!("Found device {device_id} in discovery cache");
return Some(res);
}
if !force_scan
&& let Some(last) = *state.last_scan_time.read()
&& last.elapsed() < GLOBAL_SCAN_COOLDOWN
&& let Some(res) = guard.get(device_id).cloned()
{
debug!("Global scan cooldown active (30m). Returning cached result for {device_id}.");
return Some(res);
}
None
}
async fn ensure_scan_started(&self, device_id: &str, force_scan: bool) {
let state = self.inner.clone();
let can_scan = {
let last_scan = *state.last_scan_time.read();
match last_scan {
Some(last) if !force_scan && last.elapsed() < SCAN_THROTTLE_INTERVAL => false,
_ => !state.active_scanning.swap(true, Ordering::SeqCst),
}
};
if can_scan {
info!("Initiating background scan for device ID: {device_id}...");
*state.last_scan_time.write() = Some(Instant::now());
let scanner = self.clone();
crate::runtime::spawn(async move {
let _ = scanner.perform_discovery_loop().await;
state.active_scanning.store(false, Ordering::SeqCst);
state.publish_discovery();
});
}
}
async fn wait_for_cache_result(
&self,
device_id: &str,
cancel: Option<&tokio_util::sync::CancellationToken>,
) -> Option<DiscoveryResult> {
let state = &self.inner;
let start_wait = Instant::now();
let mut discovery_rx = state.subscribe_discoveries();
let timeout_dur = *state.timeout.read();
loop {
if let Some(res) = state.cache.read().get(device_id).cloned() {
return Some(res);
}
let elapsed = start_wait.elapsed();
if elapsed >= timeout_dur || !state.active_scanning.load(Ordering::SeqCst) {
return state.cache.read().get(device_id).cloned();
}
let remaining = timeout_dur.saturating_sub(elapsed);
if let Some(ct) = cancel {
tokio::select! {
_ = ct.cancelled() => return None,
_ = sleep(remaining) => {}
_ = discovery_rx.changed() => {}
}
} else {
let _ = timeout(remaining, discovery_rx.changed()).await;
}
}
}
async fn broadcast_once(&self, target_sockets: &[(Arc<UdpSocket>, u16)]) {
let sources = self.inner.discovery_sources.read().clone();
if sources.is_empty() {
for (socket, port) in target_sockets {
let _ = self.send_discovery_broadcast(socket, *port, None).await;
}
return;
}
for source in sources {
if !source.is_ipv4() {
debug!("Skipping non-IPv4 discovery source {source}");
continue;
}
let send_socket = match Self::create_send_socket(source) {
Ok(s) => s,
Err(e) => {
warn!("Failed to bind discovery source {source}: {e} (skipping)");
continue;
}
};
for (_, port) in target_sockets {
let _ = self
.send_discovery_broadcast(&send_socket, *port, Some(source))
.await;
}
}
}
async fn perform_discovery_loop(self) -> Result<()> {
let state = &self.inner;
let ports_snapshot: Vec<u16> = state.ports.read().clone();
let mut target_sockets = Vec::new();
{
let guard = state.sockets.read();
for &port in &ports_snapshot {
if let Some(socket) = guard.get(&port) {
target_sockets.push((socket.clone(), port));
}
}
}
if target_sockets.is_empty() {
self.ensure_passive_listener();
let guard = state.sockets.read();
for &port in &ports_snapshot {
if let Some(socket) = guard.get(&port) {
target_sockets.push((socket.clone(), port));
}
}
}
if target_sockets.is_empty() {
return Err(std::io::Error::other("No available ports for scanning").into());
}
let start = Instant::now();
let mut broadcast_interval = interval(BROADCAST_INTERVAL);
let mut broadcast_count = 0;
let timeout_dur = *state.timeout.read();
while start.elapsed() < timeout_dur {
let remaining = timeout_dur.saturating_sub(start.elapsed());
if remaining.is_zero() {
break;
}
if broadcast_count >= MAX_BROADCASTS {
sleep(remaining).await;
break;
}
tokio::select! {
() = sleep(remaining) => break,
_ = broadcast_interval.tick() => {
broadcast_count += 1;
debug!("Sent broadcast {broadcast_count}/{MAX_BROADCASTS}");
self.broadcast_once(&target_sockets).await;
}
}
}
Ok(())
}
#[must_use]
pub fn invalidate_cache(&self, id: &str) -> bool {
let mut guard = self.inner.cache.write();
guard.remove(id).is_some()
}
}
const UDP_TRY_KEYS: [Option<&[u8]>; 4] =
[Some(UDP_KEY_35), Some(UDP_KEY_34), Some(UDP_KEY_33), None];
const UDP_TRY_RETCODES: [Option<bool>; 3] = [Some(true), Some(false), None];
fn parse_packet(data: &[u8]) -> Option<DiscoveryResult> {
trace!("Parsing UDP packet of {} bytes...", data.len());
if let Ok(val) = serde_json::from_slice::<Value>(data) {
trace!("Successfully parsed raw JSON packet");
return parse_json(&val);
}
for key in UDP_TRY_KEYS {
for no_retcode in UDP_TRY_RETCODES {
match protocol::unpack_message(data, key, None, no_retcode) {
Ok(msg) => {
if msg.payload.is_empty() {
continue;
}
if let Ok(val) = serde_json::from_slice::<Value>(&msg.payload) {
trace!("Successfully parsed JSON from Tuya message payload");
return parse_json(&val);
}
let keys_to_try: Vec<&[u8]> = match key {
Some(k) => vec![k],
None => vec![UDP_KEY_33, UDP_KEY_34, UDP_KEY_35],
};
for k in keys_to_try {
if let Ok(cipher) = TuyaCipher::new(k)
&& let Ok(decrypted) =
cipher.decrypt(&msg.payload, false, None, None, None)
&& let Ok(val) = serde_json::from_slice::<Value>(&decrypted)
{
trace!(
"Successfully decrypted and parsed JSON from Tuya message payload"
);
return parse_json(&val);
}
}
}
Err(e) => {
if !matches!(
e,
crate::error::TuyaError::DecodeError(_)
| crate::error::TuyaError::HmacMismatch
| crate::error::TuyaError::CrcMismatch
| crate::error::TuyaError::InvalidHeader
) {
trace!(
"unpack_message failed with key {:?}: {e}",
key.map(crate::crypto::hex_encode),
);
}
}
}
}
}
for key in &[UDP_KEY_33, UDP_KEY_34] {
if let Ok(cipher) = TuyaCipher::new(key)
&& let Ok(decrypted) = cipher.decrypt(data, false, None, None, None)
&& let Ok(val) = serde_json::from_slice::<Value>(&decrypted)
{
trace!("Successfully decrypted and parsed JSON from entire packet");
return parse_json(&val);
}
}
if let Some(pos) = data.iter().position(|&b| b == b'{')
&& let Ok(val) = serde_json::from_slice::<Value>(&data[pos..])
{
trace!("Successfully found and parsed JSON from middle of packet");
return parse_json(&val);
}
trace!("Failed to parse UDP packet");
None
}
fn parse_json(val: &Value) -> Option<DiscoveryResult> {
let id = val
.get("gwId")
.or_else(|| val.get("devId"))
.or_else(|| val.get("id"))
.and_then(|v| v.as_str())?;
let ip = val.get("ip").and_then(|v| v.as_str())?;
let ver_s = val.get("version").and_then(|v| v.as_str());
let pk = val.get("productKey").and_then(|v| v.as_str());
Some(DiscoveryResult {
id: id.to_string(),
ip: ip.to_string(),
version: ver_s.and_then(|s| Version::from_str(s).ok()),
product_key: pk.map(std::string::ToString::to_string),
discovered_at: Instant::now(),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn scan_timeout_leaves_room_after_last_broadcast() {
let last_broadcast_at = BROADCAST_INTERVAL * (MAX_BROADCASTS - 1);
let margin = DEFAULT_SCAN_TIMEOUT.saturating_sub(last_broadcast_at);
assert!(
margin >= RECEIVE_MARGIN,
"scan timeout ({DEFAULT_SCAN_TIMEOUT:?}) leaves only {margin:?} after \
the last broadcast (at {last_broadcast_at:?}); need at least {RECEIVE_MARGIN:?}"
);
}
#[test]
fn packet_tx_lifecycle_persists_then_clears_on_stop() {
let state = Arc::new(ScannerState::new());
assert!(state.current_packet_tx().is_none());
let (tx, _rx): (PacketSender, PacketReceiver) = mpsc::channel(8);
*state.packet_tx.write() = Some(tx);
let tx_clone_a = state.current_packet_tx().expect("listener up");
let tx_clone_b = state.current_packet_tx().expect("listener up");
assert!(tx_clone_a.same_channel(&tx_clone_b));
*state.packet_tx.write() = None;
assert!(state.current_packet_tx().is_none());
}
#[test]
fn dispatch_packet_updates_cache_for_known_v31_broadcast() {
let state = Arc::new(ScannerState::new());
let payload =
br#"{"gwId":"test-device-id","ip":"10.0.0.42","version":"3.1","productKey":"pk"}"#;
Scanner::dispatch_packet(&state, payload);
let cache = state.cache.read();
let entry = cache
.get("test-device-id")
.expect("dispatch must insert into cache");
assert_eq!(entry.ip, "10.0.0.42");
assert_eq!(entry.version, Some(Version::V3_1));
}
#[test]
fn reset_cancel_token_yields_fresh_uncancelled_token() {
let state = ScannerState::new();
let first = state.current_cancel_token();
assert!(!first.is_cancelled());
state.reset_cancel_token();
assert!(first.is_cancelled());
let second = state.current_cancel_token();
assert!(!second.is_cancelled());
state.cancel_token.write().cancel();
assert!(state.current_cancel_token().is_cancelled());
state.reset_cancel_token();
let third = state.current_cancel_token();
assert!(!third.is_cancelled());
}
#[test]
fn active_scanning_compare_exchange_is_single_winner() {
use std::sync::atomic::{AtomicUsize, Ordering as O};
let state = Arc::new(ScannerState::new());
let winners = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..16 {
let s = state.clone();
let w = winners.clone();
handles.push(std::thread::spawn(move || {
if s.active_scanning
.compare_exchange(false, true, O::SeqCst, O::SeqCst)
.is_ok()
{
w.fetch_add(1, O::SeqCst);
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(
winners.load(O::SeqCst),
1,
"exactly one caller should claim the active scan slot"
);
assert!(state.active_scanning.load(O::SeqCst));
}
#[test]
fn setters_take_self_and_share_state_across_clones() {
let inner = Arc::new(ScannerState::new());
let a = Scanner {
inner: inner.clone(),
};
let b = Scanner { inner };
assert_eq!(a.timeout(), DEFAULT_SCAN_TIMEOUT);
assert_eq!(a.bind_addr(), "0.0.0.0");
assert_eq!(a.ports(), vec![6666, 6667, 7000]);
a.set_timeout(Duration::from_secs(42));
assert_eq!(b.timeout(), Duration::from_secs(42));
*a.inner.bind_addr.write() = "127.0.0.1".to_string();
*a.inner.ports.write() = vec![9999];
assert_eq!(b.bind_addr(), "127.0.0.1");
assert_eq!(b.ports(), vec![9999]);
}
#[test]
fn discovery_sources_default_empty_and_round_trip() {
let inner = Arc::new(ScannerState::new());
let a = Scanner {
inner: inner.clone(),
};
let b = Scanner { inner };
assert!(
a.discovery_sources().is_empty(),
"default must be empty (auto-detect source via route lookup)"
);
let srcs = vec![
"192.168.1.50".parse::<IpAddr>().unwrap(),
"10.0.20.5".parse::<IpAddr>().unwrap(),
];
a.set_discovery_sources(srcs.clone());
assert_eq!(b.discovery_sources(), srcs);
b.set_discovery_sources(Vec::new());
assert!(a.discovery_sources().is_empty());
}
#[test]
fn local_ip_lookup_is_not_cached() {
let src = include_str!("scanner.rs");
let forbidden_static = format!("LOCAL{}IP{}CACHE", "_", "_");
assert!(
!src.contains(&forbidden_static),
"scanner.rs must not reintroduce the local-IP cache static — see \
docs/technical-notes.md §Local-IP discovery for rationale"
);
let forbidden_ty = format!("Once{}<Option<String>>", "Lock");
assert!(
!src.contains(&forbidden_ty),
"scanner.rs must not cache an Option<String> via OnceLock — that \
would freeze the local IP for the process lifetime"
);
}
#[test]
fn effective_bind_ip_widens_only_concrete_unicast() {
assert_eq!(
effective_bind_ip("0.0.0.0"),
Some((IpAddr::V4(Ipv4Addr::UNSPECIFIED), false))
);
assert_eq!(
effective_bind_ip("::"),
Some((IpAddr::V6(Ipv6Addr::UNSPECIFIED), false))
);
assert_eq!(
effective_bind_ip("127.0.0.1"),
Some(("127.0.0.1".parse().unwrap(), false))
);
assert_eq!(
effective_bind_ip("192.168.1.50"),
Some((IpAddr::V4(Ipv4Addr::UNSPECIFIED), true))
);
assert_eq!(
effective_bind_ip("10.0.0.87"),
Some((IpAddr::V4(Ipv4Addr::UNSPECIFIED), true))
);
let (v6, widened) = effective_bind_ip("2001:db8::1").unwrap();
assert_eq!(v6, IpAddr::V6(Ipv6Addr::UNSPECIFIED));
assert!(widened);
assert_eq!(effective_bind_ip("not-an-ip"), None);
}
#[test]
fn compute_port_diff_adds_only() {
let (add, remove) = compute_port_diff(&[6666, 6667, 7000], &[]);
assert_eq!(add, vec![6666, 6667, 7000]);
assert!(remove.is_empty());
}
#[test]
fn compute_port_diff_removes_only() {
let (add, remove) = compute_port_diff(&[6666], &[6666, 6667, 7000]);
assert!(add.is_empty());
let removed: std::collections::HashSet<u16> = remove.into_iter().collect();
assert_eq!(
removed,
[6667u16, 7000]
.iter()
.copied()
.collect::<std::collections::HashSet<_>>()
);
}
#[test]
fn compute_port_diff_mixed_add_and_remove() {
let (add, remove) = compute_port_diff(&[6666, 8888], &[6666, 7000]);
assert_eq!(add, vec![8888]);
assert_eq!(remove, vec![7000]);
}
#[test]
fn compute_port_diff_no_change_returns_empty_both() {
let (add, remove) = compute_port_diff(&[6666, 7000], &[6666, 7000]);
assert!(add.is_empty());
assert!(remove.is_empty());
}
#[test]
fn set_ports_add_then_remove_reconciles_socket_map() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let scanner = Scanner {
inner: Arc::new(ScannerState::new()),
};
let port_a = 47770u16;
let port_b = 47771u16;
let port_c = 47772u16;
scanner.set_ports(vec![port_a, port_b]);
tokio::task::yield_now().await;
{
let sockets = scanner.inner.sockets.read();
assert!(sockets.contains_key(&port_a), "port_a should be bound");
assert!(sockets.contains_key(&port_b), "port_b should be bound");
assert_eq!(sockets.len(), 2);
}
{
let tasks = scanner.inner.receiver_tasks.read();
assert!(tasks.contains_key(&port_a));
assert!(tasks.contains_key(&port_b));
}
scanner.set_ports(vec![port_a, port_c]);
tokio::task::yield_now().await;
{
let sockets = scanner.inner.sockets.read();
assert!(sockets.contains_key(&port_a), "port_a kept");
assert!(
!sockets.contains_key(&port_b),
"port_b must be removed from sockets map"
);
assert!(sockets.contains_key(&port_c), "port_c added");
assert_eq!(sockets.len(), 2);
}
{
let tasks = scanner.inner.receiver_tasks.read();
assert!(tasks.contains_key(&port_a));
assert!(
!tasks.contains_key(&port_b),
"port_b receiver task must be aborted and removed"
);
assert!(tasks.contains_key(&port_c));
}
scanner.stop_passive_listener();
assert!(scanner.inner.sockets.read().is_empty());
assert!(scanner.inner.receiver_tasks.read().is_empty());
assert!(scanner.inner.dispatcher_task.read().is_none());
});
}
#[test]
fn compute_port_diff_full_swap() {
let (add, remove) = compute_port_diff(&[8888, 9999], &[6666, 7000]);
let added: std::collections::HashSet<u16> = add.into_iter().collect();
let removed: std::collections::HashSet<u16> = remove.into_iter().collect();
assert_eq!(
added,
[8888u16, 9999]
.iter()
.copied()
.collect::<std::collections::HashSet<_>>()
);
assert_eq!(
removed,
[6666u16, 7000]
.iter()
.copied()
.collect::<std::collections::HashSet<_>>()
);
}
}