use crate::WireEnum;
use rand::Rng;
use sha2::{Digest, Sha256};
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use thiserror::Error;
use wacore_binary::builder::NodeBuilder;
use wacore_binary::{Jid, JidExt, LEGACY_USER_SERVER};
use wacore_binary::{Node, NodeContent, NodeContentRef, NodeRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq, WireEnum)]
pub enum InfoQueryType {
#[wire = "set"]
Set,
#[wire = "get"]
Get,
}
#[derive(Debug, Clone)]
pub struct InfoQuery<'a> {
pub namespace: &'a str,
pub query_type: InfoQueryType,
pub to: Jid,
pub target: Option<Jid>,
pub id: Option<String>,
pub content: Option<NodeContent>,
pub timeout: Option<Duration>,
}
impl<'a> InfoQuery<'a> {
pub fn get(namespace: &'a str, to: Jid, content: Option<NodeContent>) -> Self {
Self {
namespace,
query_type: InfoQueryType::Get,
to,
target: None,
id: None,
content,
timeout: None,
}
}
pub fn set(namespace: &'a str, to: Jid, content: Option<NodeContent>) -> Self {
Self {
namespace,
query_type: InfoQueryType::Set,
to,
target: None,
id: None,
content,
timeout: None,
}
}
pub fn with_target(mut self, target: Jid) -> Self {
self.target = Some(target);
self
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub fn get_ref(namespace: &'a str, to: &Jid, content: Option<NodeContent>) -> Self {
Self::get(namespace, to.clone(), content)
}
pub fn set_ref(namespace: &'a str, to: &Jid, content: Option<NodeContent>) -> Self {
Self::set(namespace, to.clone(), content)
}
pub fn with_target_ref(self, target: &Jid) -> Self {
self.with_target(target.clone())
}
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum IqError {
#[error("IQ request timed out")]
Timeout,
#[error("client is not connected")]
NotConnected,
#[error("received disconnect node during IQ wait: {0:?}")]
Disconnected(Box<Node>),
#[error("received a server error response: code={code}, text='{text}'")]
ServerError {
code: u16,
text: String,
error_type: Option<String>,
backoff: Option<u32>,
},
#[error("received unexpected IQ response type: {got:?}")]
UnexpectedResponseType { got: Option<String> },
#[error("internal channel closed unexpectedly")]
InternalChannelClosed,
}
impl IqError {
pub fn is_transport_unavailable(&self) -> bool {
matches!(
self,
IqError::NotConnected | IqError::Disconnected(_) | IqError::InternalChannelClosed
)
}
pub fn is_timeout(&self) -> bool {
match self {
IqError::Timeout => true,
IqError::NotConnected
| IqError::Disconnected(_)
| IqError::ServerError { .. }
| IqError::UnexpectedResponseType { .. }
| IqError::InternalChannelClosed => false,
}
}
}
#[derive(Debug, Clone, Error)]
#[error("server error: code={code}, text='{text}'")]
pub struct ServerErrorCode {
pub code: u16,
pub text: String,
pub error_type: Option<String>,
pub backoff: Option<u32>,
}
impl ServerErrorCode {
pub fn from_anyhow(err: &anyhow::Error) -> Option<&Self> {
err.chain().find_map(|cause| cause.downcast_ref::<Self>())
}
}
pub struct RequestUtils {
unique_id: String,
id_counter: std::sync::Arc<portable_atomic::AtomicU64>,
}
impl RequestUtils {
pub fn new(unique_id: String) -> Self {
Self {
unique_id,
id_counter: std::sync::Arc::new(portable_atomic::AtomicU64::new(0)),
}
}
pub fn with_counter(
unique_id: String,
id_counter: std::sync::Arc<portable_atomic::AtomicU64>,
) -> Self {
Self {
unique_id,
id_counter,
}
}
pub fn generate_request_id(&self) -> String {
let count = self.id_counter.fetch_add(1, Ordering::Relaxed);
format!(
"{unique_id}-{count}",
unique_id = self.unique_id,
count = count
)
}
pub fn generate_message_id(&self, user_jid: Option<&Jid>) -> String {
self.generate_message_id_at(user_jid, crate::time::now_secs_u64())
}
pub fn generate_message_id_at(&self, user_jid: Option<&Jid>, unix_secs: u64) -> String {
Self::message_id_at(user_jid, unix_secs)
}
pub fn message_id_at(user_jid: Option<&Jid>, unix_secs: u64) -> String {
let mut hasher = Sha256::new();
hasher.update(unix_secs.to_be_bytes());
if let Some(jid) = user_jid {
hasher.update(jid.user.as_bytes());
hasher.update(b"@");
hasher.update(LEGACY_USER_SERVER.as_bytes());
}
let mut random_bytes = [0u8; 16];
rand::rng().fill_bytes(&mut random_bytes);
hasher.update(random_bytes);
const HEX_UPPER: &[u8; 16] = b"0123456789ABCDEF";
let hash = hasher.finalize();
let truncated = &hash[..9];
let mut id = String::with_capacity(22);
id.push_str("3EB0");
for &b in truncated {
id.push(HEX_UPPER[(b >> 4) as usize] as char);
id.push(HEX_UPPER[(b & 0x0F) as usize] as char);
}
id
}
pub fn build_iq_node(&self, query: InfoQuery<'_>, req_id: Option<String>) -> Node {
let id = req_id.unwrap_or_else(|| self.generate_request_id());
let mut builder = NodeBuilder::new("iq")
.attr("id", id)
.attr("xmlns", query.namespace)
.attr("type", query.query_type.as_str())
.attr("to", query.to);
if let Some(target) = query.target
&& !target.is_empty()
{
builder = builder.attr("target", target);
}
builder.apply_content(query.content).build()
}
pub fn parse_iq_response(&self, response_node: &NodeRef<'_>) -> Result<(), IqError> {
if response_node.tag == "stream:error" || response_node.tag == "xmlstreamend" {
return Err(IqError::Disconnected(Box::new(response_node.to_owned())));
}
let response_type = response_node.get_attr("type");
if response_type
.as_ref()
.is_some_and(|res_type| res_type.as_str() == "error")
{
let error_child = response_node.get_optional_child_by_tag(&["error"]);
if let Some(error_node) = error_child {
let mut parser = error_node.attrs();
let code = parser.optional_u64("code").unwrap_or(0) as u16;
let text = parser
.optional_string("text")
.as_deref()
.unwrap_or("")
.to_string();
let error_type = parser.optional_string("type").map(|s| s.into_owned());
let backoff = parser
.optional_u64("backoff")
.and_then(|b| u32::try_from(b).ok());
warn_on_dropped_error_detail(error_node);
return Err(IqError::ServerError {
code,
text,
error_type,
backoff,
});
}
return Err(IqError::ServerError {
code: 0,
text: "Malformed error response".to_string(),
error_type: None,
backoff: None,
});
}
let got = response_type.map(|res_type| res_type.to_string());
if got.as_deref() != Some("result") {
return Err(IqError::UnexpectedResponseType { got });
}
Ok(())
}
}
#[derive(Debug, PartialEq, Eq)]
struct DroppedErrorDetail<'a> {
attrs: Vec<&'a str>,
children: Vec<&'a str>,
payload: Option<&'static str>,
}
const PARSED_ERROR_ATTRS: [&str; 4] = ["code", "text", "type", "backoff"];
fn dropped_error_detail<'a>(error_node: &'a NodeRef<'_>) -> Option<DroppedErrorDetail<'a>> {
let attrs: Vec<&str> = error_node
.attrs
.iter()
.map(|(name, _)| name.as_ref())
.filter(|name| !PARSED_ERROR_ATTRS.contains(name))
.collect();
let mut children = Vec::new();
let mut payload = None;
match error_node.content.as_ref() {
Some(NodeContentRef::Nodes(nodes)) => {
children.extend(nodes.iter().map(|child| child.tag.as_ref()));
}
Some(NodeContentRef::Bytes(bytes)) if !bytes.is_empty() => payload = Some("bytes"),
Some(NodeContentRef::String(text)) if !text.is_empty() => payload = Some("text"),
_ => {}
}
if attrs.is_empty() && children.is_empty() && payload.is_none() {
return None;
}
Some(DroppedErrorDetail {
attrs,
children,
payload,
})
}
static DROPPED_ERROR_DETAIL_WARNED: AtomicBool = AtomicBool::new(false);
#[cold]
fn warn_on_dropped_error_detail(error_node: &NodeRef<'_>) {
if DROPPED_ERROR_DETAIL_WARNED.load(Ordering::Relaxed) || !log::log_enabled!(log::Level::Warn) {
return;
}
let Some(detail) = dropped_error_detail(error_node) else {
return;
};
if DROPPED_ERROR_DETAIL_WARNED.swap(true, Ordering::Relaxed) {
return;
}
log::warn!(
"IQ error carries detail this parser drops: attributes={:?} children={:?} payload={:?}. \
Names only, no values, since a value can hold a JID. Reported once per process.",
detail.attrs,
detail.children,
detail.payload,
);
}
#[cfg(test)]
mod iq_error_tests {
use super::{IqError, RequestUtils};
use wacore_binary::builder::NodeBuilder;
#[test]
fn parse_iq_response_extracts_error_type_and_backoff() {
let node = NodeBuilder::new("iq")
.attr("type", "error")
.children([NodeBuilder::new("error")
.attr("code", "429")
.attr("text", "rate-overlimit")
.attr("type", "wait")
.attr("backoff", "30")
.build()])
.build();
let err = RequestUtils::new("t".to_string())
.parse_iq_response(&node.as_node_ref())
.unwrap_err();
match err {
IqError::ServerError {
code,
text,
error_type,
backoff,
} => {
assert_eq!(code, 429);
assert_eq!(text, "rate-overlimit");
assert_eq!(error_type.as_deref(), Some("wait"));
assert_eq!(backoff, Some(30));
}
other => panic!("expected ServerError, got {other:?}"),
}
}
#[test]
fn parse_iq_response_error_without_backoff_is_none() {
let node = NodeBuilder::new("iq")
.attr("type", "error")
.children([NodeBuilder::new("error").attr("code", "404").build()])
.build();
let err = RequestUtils::new("t".to_string())
.parse_iq_response(&node.as_node_ref())
.unwrap_err();
match err {
IqError::ServerError {
code,
error_type,
backoff,
..
} => {
assert_eq!(code, 404);
assert!(error_type.is_none());
assert!(backoff.is_none());
}
other => panic!("expected ServerError, got {other:?}"),
}
}
#[test]
fn parse_iq_response_accepts_result_type() {
let node = NodeBuilder::new("iq").attr("type", "result").build();
RequestUtils::new("t".to_string())
.parse_iq_response(&node.as_node_ref())
.unwrap();
}
#[test]
fn parse_iq_response_rejects_unexpected_type() {
let node = NodeBuilder::new("iq").attr("type", "get").build();
let err = RequestUtils::new("t".to_string())
.parse_iq_response(&node.as_node_ref())
.unwrap_err();
match err {
IqError::UnexpectedResponseType { got } => assert_eq!(got.as_deref(), Some("get")),
other => panic!("expected UnexpectedResponseType, got {other:?}"),
}
}
#[test]
fn parse_iq_response_rejects_missing_type() {
let node = NodeBuilder::new("iq").build();
let err = RequestUtils::new("t".to_string())
.parse_iq_response(&node.as_node_ref())
.unwrap_err();
match err {
IqError::UnexpectedResponseType { got } => assert!(got.is_none()),
other => panic!("expected UnexpectedResponseType, got {other:?}"),
}
}
}
#[cfg(test)]
mod message_id_tests {
use super::RequestUtils;
use wacore_binary::Jid;
fn jid() -> Jid {
"13135550100@s.whatsapp.net".parse().expect("valid jid")
}
#[test]
fn message_id_keeps_the_wa_web_shape() {
let id = RequestUtils::message_id_at(Some(&jid()), 1_700_000_000);
assert_eq!(id.len(), 22, "id must be 3EB0 plus 18 hex chars: {id}");
assert!(id.starts_with("3EB0"), "id must carry the WA prefix: {id}");
assert!(
id[4..]
.chars()
.all(|c| c.is_ascii_digit() || ('A'..='F').contains(&c)),
"id must be upper-case hex after the prefix: {id}"
);
}
#[test]
fn message_id_is_unique_within_one_second() {
let jid = jid();
let ids: std::collections::HashSet<String> = (0..64)
.map(|_| RequestUtils::message_id_at(Some(&jid), 1_700_000_000))
.collect();
assert_eq!(ids.len(), 64, "ids repeated within the same second");
}
#[test]
fn message_id_without_jid_keeps_the_shape() {
let id = RequestUtils::message_id_at(None, 1_700_000_000);
assert_eq!(id.len(), 22);
assert!(id.starts_with("3EB0"));
}
}
#[cfg(test)]
mod dropped_error_detail_tests {
use super::dropped_error_detail;
use wacore_binary::builder::NodeBuilder;
use wacore_binary::node::{Node, NodeContent};
#[test]
fn a_fully_parsed_error_reports_nothing() {
let node = NodeBuilder::new("error")
.attr("code", "429")
.attr("text", "rate-overlimit")
.attr("type", "wait")
.attr("backoff", "30")
.build();
let node_ref = node.as_node_ref();
assert_eq!(dropped_error_detail(&node_ref), None);
}
#[test]
fn an_empty_error_reports_nothing() {
let node = NodeBuilder::new("error").build();
let node_ref = node.as_node_ref();
assert_eq!(dropped_error_detail(&node_ref), None);
}
#[test]
fn an_unread_attribute_is_named() {
let node = NodeBuilder::new("error")
.attr("code", "400")
.attr("xmlns", "w:profile:picture")
.build();
let node_ref = node.as_node_ref();
let detail = dropped_error_detail(&node_ref).expect("xmlns is not parsed");
assert_eq!(detail.attrs, ["xmlns"]);
assert!(detail.children.is_empty());
assert_eq!(detail.payload, None);
}
#[test]
fn child_tags_are_named() {
let node = NodeBuilder::new("error")
.attr("code", "400")
.children([
NodeBuilder::new("bad-request").build(),
NodeBuilder::new("text").build(),
])
.build();
let node_ref = node.as_node_ref();
let detail = dropped_error_detail(&node_ref).expect("children are not parsed");
assert!(detail.attrs.is_empty());
assert_eq!(detail.children, ["bad-request", "text"]);
assert_eq!(detail.payload, None);
}
#[test]
fn a_raw_payload_is_reported_by_kind_only() {
let bytes_node = Node::new(
"error",
[("code".into(), "400".into())].into_iter().collect(),
Some(NodeContent::Bytes(vec![1, 2, 3])),
);
let bytes_ref = bytes_node.as_node_ref();
let detail = dropped_error_detail(&bytes_ref).expect("a payload is not parsed");
assert_eq!(detail.payload, Some("bytes"));
let text_node = Node::new(
"error",
Default::default(),
Some(NodeContent::String("x".into())),
);
let text_ref = text_node.as_node_ref();
let detail = dropped_error_detail(&text_ref).expect("a payload is not parsed");
assert_eq!(detail.payload, Some("text"));
}
#[test]
fn an_empty_payload_reports_nothing() {
let node = Node::new(
"error",
Default::default(),
Some(NodeContent::Bytes(vec![])),
);
let node_ref = node.as_node_ref();
assert_eq!(dropped_error_detail(&node_ref), None);
}
}