1use std::collections::HashMap;
17use std::ffi::CStr;
18use std::sync::atomic::{AtomicU64, Ordering};
19use std::sync::{LazyLock, Mutex, OnceLock, Weak};
20use std::thread;
21use std::time::Duration;
22
23use serde::Deserialize;
24use tokio::sync::watch;
25
26use td_types::traits::Function;
27
28use crate::connection::Connection;
29use crate::error::{Error, Result};
30
31#[derive(Deserialize)]
32struct Response<'a> {
33 #[serde(rename = "@type")]
34 kind: &'a str,
35}
36
37#[derive(Deserialize)]
38struct Incoming<'a> {
39 #[serde(rename = "@client_id")]
40 client_id: i32,
41 #[serde(rename = "@extra")]
42 extra: Option<u64>,
43 #[serde(rename = "@type")]
44 kind: &'a str,
45}
46
47struct Receiver {
48 clients: Mutex<HashMap<i32, Weak<Connection>>>,
49 thread: OnceLock<thread::Thread>,
50 transition: watch::Sender<()>,
51 timeout: AtomicU64,
52}
53
54static RECEIVER: LazyLock<Receiver> = LazyLock::new(|| {
55 let (transition, _) = watch::channel(());
56 let (clients, thread) = Default::default();
57 let timeout = 1f64.to_bits().into();
58 Receiver { clients, thread, transition, timeout }
59});
60
61impl Receiver {
62 fn run(&self) {
63 loop {
64 if self.clients.lock().unwrap().is_empty() {
65 self.transition.send_replace(());
68 thread::park();
69 } else {
70 self.receive();
71 }
72 }
73 }
74
75 fn receive(&self) {
76 let timeout = f64::from_bits(self.timeout.load(Ordering::Relaxed));
77 let raw = unsafe { td_sys::td_receive(timeout) };
79 if !raw.is_null() {
80 self.route(unsafe { CStr::from_ptr(raw) }.to_bytes());
83 }
84 }
85
86 fn route(&self, raw: &[u8]) {
87 let Incoming { client_id, extra, kind } = match serde_json::from_slice(raw) {
88 Ok(incoming) => incoming,
89 Err(error) => {
90 return match serde_json::from_slice(raw) {
93 Ok(Response { kind: "error" }) => report(parse_error(raw)),
94 _ => report(error),
95 };
96 }
97 };
98 if let (None, "error") = (extra, kind) {
99 return report(parse_error(raw));
100 }
101 let connection = self.clients.lock().unwrap().get(&client_id).and_then(Weak::upgrade);
102 let Some(connection) = connection else { return };
103 match extra {
104 Some(extra) => connection.complete_request(extra, kind, raw),
105 None => connection.update(raw),
106 }
107 }
108}
109
110pub(crate) fn register(id: i32, connection: Weak<Connection>) {
111 RECEIVER.clients.lock().unwrap().insert(id, connection);
112 RECEIVER.thread.get_or_init(|| thread::spawn(|| RECEIVER.run()).thread().clone()).unpark();
113 RECEIVER.transition.send_replace(());
114}
115
116pub(crate) async fn unregister(id: i32) {
117 let mut transition = RECEIVER.transition.subscribe();
118 {
119 let mut clients = RECEIVER.clients.lock().unwrap();
120 clients.remove(&id);
121 if !clients.is_empty() {
122 return;
123 }
124 transition.borrow_and_update();
127 }
128 let _ = transition.changed().await;
129}
130
131pub(crate) fn remove(id: i32) {
132 if let Some(receiver) = LazyLock::get(&RECEIVER) {
133 receiver.clients.lock().unwrap().remove(&id);
134 }
135}
136
137pub fn execute<F: Function>(request: &F) -> Result<F::Return> {
164 let mut bytes = serde_json::to_vec(request)?;
165 bytes.push(0);
166 let raw = unsafe { td_sys::td_execute(bytes.as_ptr().cast()) };
169 if raw.is_null() {
170 return Err(Error::UnexpectedResponse("synchronous request returned null"));
171 }
172 let raw = unsafe { CStr::from_ptr(raw) }.to_bytes();
174 match serde_json::from_slice(raw)? {
175 Response { kind: "error" } => Err(parse_error(raw)),
176 _ => serde_json::from_slice(raw).map_err(Into::into),
177 }
178}
179
180pub fn set_receive_timeout(timeout: Duration) {
197 RECEIVER.timeout.store(timeout.as_secs_f64().to_bits(), Ordering::Relaxed);
198}
199
200pub fn set_log_level(level: i32) {
211 unsafe { td_sys::td_set_log_verbosity_level(level) };
213}
214
215type ErrorCallback = Box<dyn Fn(Error) + Send + Sync>;
216static ERROR_CALLBACK: Mutex<Option<ErrorCallback>> = Mutex::new(None);
217
218pub fn on_error(callback: impl Fn(Error) + Send + Sync + 'static) {
251 *ERROR_CALLBACK.lock().unwrap() = Some(Box::new(callback));
252}
253
254pub(crate) fn report(error: impl Into<Error>) {
255 if let Some(callback) = ERROR_CALLBACK.lock().unwrap().as_ref() {
256 callback(error.into());
257 }
258}
259
260pub(crate) fn parse_error(raw: &[u8]) -> Error {
261 match serde_json::from_slice(raw) {
262 Ok(error) => Error::Td(error),
263 Err(error) => error.into(),
264 }
265}
266
267#[cfg(test)]
268mod tests {
269 use std::assert_matches;
270 use std::sync::Arc;
271 use tokio::sync::mpsc;
272
273 use super::*;
274
275 #[test]
276 fn unroutable_output_reports_original_errors_without_poisoning_clients() {
277 let (connection, mut updates) = Connection::fixture();
278 let (transition, _) = watch::channel(());
279 let clients = Mutex::new(HashMap::from([(7, Arc::downgrade(&connection))]));
280 let receiver = Receiver { clients, thread: OnceLock::new(), transition, timeout: AtomicU64::new(0) };
281 let errors = Arc::new(Mutex::new(Vec::new()));
282 let observed = Arc::clone(&errors);
283 on_error(move |error| observed.lock().unwrap().push(error));
284 receiver.route(br#"{"@type":"error","code":429,"message":"limited"}"#);
285 receiver.route(br#"{"@client_id":99,"@type":"error","code":500,"message":"unroutable"}"#);
286 receiver.route(br#"{"@client_id":7,"@type":"error","code":400,"message":"unsolicited"}"#);
287 receiver.route(br#"{"@type":"ok"}"#);
288 receiver.route(br#"{"@client_id":7,"@type":"updateMessageSendSucceeded","message":null}"#);
289 on_error(drop);
290
291 let errors = errors.lock().unwrap().drain(..).collect::<Vec<_>>();
292 let [global, unknown_client, known_client, missing_client, malformed]: [Error; 5] = errors.try_into().unwrap();
293 assert_matches!(global, Error::Td(error) if error.code == 429);
294 assert_matches!(unknown_client, Error::Td(error) if error.code == 500);
295 assert_matches!(known_client, Error::Td(error) if error.code == 400);
296 assert_matches!(missing_client, Error::Json(_));
297 assert_matches!(malformed, Error::Json(_));
298 let application = updates.try_recv();
299 assert_matches!(application, Err(mpsc::error::TryRecvError::Empty));
300 }
301}