Skip to main content

ecat_data_iotdb/
lib.rs

1// Copyright (c) 2026 erik <erik@erik.xyz> — https://erik.xyz
2use async_trait::async_trait;
3use ecat_data::{DataPoint, FieldValue, TsdbClient};
4use ecat_errors::{Error, ErrorCode};
5use ecat_tls::TlsClientConfig;
6use serde::Deserialize;
7
8#[derive(Debug, Clone, Deserialize)]
9pub struct IotdbConfig {
10    pub base_url: String,
11    pub username: String,
12    pub password: String,
13    #[serde(default)]
14    pub tls: Option<TlsClientConfig>,
15}
16
17pub struct IotdbClient {
18    client: reqwest::Client,
19    base_url: String,
20    username: String,
21    password: String,
22}
23
24impl IotdbClient {
25    pub fn new(
26        base_url: impl Into<String>,
27        username: impl Into<String>,
28        password: impl Into<String>,
29    ) -> Self {
30        Self {
31            client: reqwest::Client::new(),
32            base_url: base_url.into(),
33            username: username.into(),
34            password: password.into(),
35        }
36    }
37
38    pub fn from_config(cfg: IotdbConfig) -> Result<Self, Error> {
39        let client = ecat_tls::build_reqwest_client(&cfg.tls)
40            .map_err(|e| Error::new(ErrorCode::Internal, "iotdb", format!("TLS: {e}")))?;
41        Ok(Self {
42            client,
43            base_url: cfg.base_url,
44            username: cfg.username,
45            password: cfg.password,
46        })
47    }
48}
49
50#[async_trait]
51impl TsdbClient for IotdbClient {
52    async fn write(&self, points: &[DataPoint]) -> Result<(), Error> {
53        for p in points {
54            // Apache IoTDB REST v2 insertTablet body:
55            // {"device": "...", "is_aligned": false, "timestamps": [...],
56            //  "measurements": [...], "data_types": [...], "values": [[...]]}
57            // `device` = measurement; tags are not representable in this API.
58            let mut measurements = Vec::with_capacity(p.fields.len());
59            let mut data_types = Vec::with_capacity(p.fields.len());
60            let mut values: Vec<serde_json::Value> = Vec::with_capacity(p.fields.len());
61            for (k, v) in &p.fields {
62                measurements.push(k.clone());
63                let (dt, val) = match v {
64                    FieldValue::Float(f) => (
65                        "DOUBLE",
66                        serde_json::Value::Number(
67                            serde_json::Number::from_f64(*f).unwrap_or(0.into()),
68                        ),
69                    ),
70                    FieldValue::Int(i) => ("INT64", serde_json::Value::Number((*i).into())),
71                    FieldValue::String(s) => ("TEXT", serde_json::Value::String(s.clone())),
72                    FieldValue::Bool(b) => ("BOOLEAN", serde_json::Value::Bool(*b)),
73                };
74                data_types.push(dt);
75                values.push(val);
76            }
77            let body = serde_json::json!({
78                "device": p.measurement,
79                "is_aligned": false,
80                "timestamps": [p.timestamp.unwrap_or(0)],
81                "measurements": measurements,
82                "data_types": data_types,
83                "values": [values],
84            });
85            let resp = self
86                .client
87                .post(format!("{}/rest/v2/insertTablet", self.base_url))
88                .basic_auth(&self.username, Some(&self.password))
89                .header("Content-Type", "application/json")
90                .json(&body)
91                .send()
92                .await
93                .map_err(|e| {
94                    Error::new(ErrorCode::Internal, "iotdb", format!("iotdb write: {e}"))
95                })?;
96            if !resp.status().is_success() {
97                return Err(Error::new(
98                    ErrorCode::Internal,
99                    "iotdb",
100                    resp.text().await.unwrap_or_default(),
101                ));
102            }
103            // IoTDB REST v2 may return HTTP 200 with a body `code` != 200 on
104            // some failures; surface those too.
105            if let Ok(v) = resp.json::<serde_json::Value>().await
106                && let Some(code) = v.get("code").and_then(|c| c.as_i64())
107                && code != 200
108            {
109                return Err(Error::new(
110                    ErrorCode::Internal,
111                    "iotdb",
112                    format!(
113                        "iotdb write failed: code {code}: {}",
114                        v.get("message")
115                            .and_then(|m| m.as_str())
116                            .unwrap_or("no message")
117                    ),
118                ));
119            }
120        }
121        Ok(())
122    }
123
124    async fn query(&self, sql: &str) -> Result<serde_json::Value, Error> {
125        let resp = self
126            .client
127            .post(format!("{}/rest/v2/query", self.base_url))
128            .basic_auth(&self.username, Some(&self.password))
129            .header("Content-Type", "text/plain; charset=utf-8")
130            .body(sql.to_string())
131            .send()
132            .await
133            .map_err(|e| Error::new(ErrorCode::Internal, "iotdb", format!("iotdb query: {e}")))?;
134        if !resp.status().is_success() {
135            return Err(Error::new(
136                ErrorCode::Internal,
137                "iotdb",
138                resp.text().await.unwrap_or_default(),
139            ));
140        }
141        resp.json()
142            .await
143            .map_err(|e| Error::new(ErrorCode::Internal, "iotdb", format!("iotdb parse: {e}")))
144    }
145}
146
147#[cfg(test)]
148mod tests {
149    use super::*;
150    use axum::extract::{Request, State};
151    use axum::response::{IntoResponse, Response};
152    use std::sync::{Arc, Mutex};
153
154    #[test]
155    fn client_constructs() {
156        let _client = IotdbClient::new("http://localhost:18080", "root", "root");
157    }
158
159    #[derive(Debug)]
160    struct CapturedRequest {
161        path: String,
162        headers: Vec<(String, String)>,
163        body: Vec<u8>,
164    }
165
166    impl CapturedRequest {
167        fn header(&self, name: &str) -> Option<&str> {
168            self.headers
169                .iter()
170                .find(|(k, _)| k.eq_ignore_ascii_case(name))
171                .map(|(_, v)| v.as_str())
172        }
173    }
174
175    /// mock IoTDB /rest/v2/insertTablet 端点:捕获请求路径/头/体,按给定
176    /// 状态码与响应体应答,返回 mock 的 base_url。
177    async fn spawn_mock_insert(
178        captured: Arc<Mutex<Vec<CapturedRequest>>>,
179        status: u16,
180        body: &'static str,
181    ) -> String {
182        let config = Arc::new(MockConfig {
183            captured,
184            status,
185            body,
186        });
187        let app = axum::Router::new()
188            .route("/rest/v2/insertTablet", axum::routing::post(handle_insert))
189            .with_state(config);
190
191        async fn handle_insert(State(config): State<Arc<MockConfig>>, req: Request) -> Response {
192            let path = req.uri().path().to_string();
193            let (parts, req_body) = req.into_parts();
194            let headers = parts
195                .headers
196                .iter()
197                .map(|(k, v)| {
198                    (
199                        k.as_str().to_string(),
200                        v.to_str().unwrap_or_default().to_string(),
201                    )
202                })
203                .collect();
204            let req_body = axum::body::to_bytes(req_body, usize::MAX)
205                .await
206                .unwrap_or_default();
207            config
208                .captured
209                .lock()
210                .unwrap_or_else(|e| e.into_inner())
211                .push(CapturedRequest {
212                    path,
213                    headers,
214                    body: req_body.to_vec(),
215                });
216            if config.body.is_empty() {
217                axum::Json(serde_json::json!({"code": 200})).into_response()
218            } else {
219                (
220                    axum::http::StatusCode::from_u16(config.status).unwrap(),
221                    axum::response::Response::new(axum::body::Body::from(config.body)),
222                )
223                    .into_response()
224            }
225        }
226
227        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
228        let addr = listener.local_addr().unwrap();
229        tokio::spawn(async move {
230            axum::serve(listener, app).await.unwrap();
231        });
232        format!("http://{addr}")
233    }
234
235    struct MockConfig {
236        captured: Arc<Mutex<Vec<CapturedRequest>>>,
237        status: u16,
238        body: &'static str,
239    }
240
241    /// 按 measurement 索引对齐断言 fields 构造的三元组
242    /// (measurements/data_types/values[0] 来自同一循环,顺序一致但
243    /// HashMap 迭代顺序不定,故按名称索引断言)。
244    fn assert_field(
245        body: &serde_json::Value,
246        field: &str,
247        expected_type: &str,
248        expected_value: serde_json::Value,
249    ) {
250        let measurements = body["measurements"].as_array().unwrap();
251        let idx = measurements
252            .iter()
253            .position(|m| m.as_str() == Some(field))
254            .unwrap_or_else(|| panic!("field {field} missing from {measurements:?}"));
255        assert_eq!(
256            body["data_types"][idx].as_str(),
257            Some(expected_type),
258            "type for {field}"
259        );
260        assert_eq!(body["values"][0][idx], expected_value, "value for {field}");
261    }
262
263    #[tokio::test]
264    async fn insert_tablet_sends_full_protocol_body() {
265        let captured = Arc::new(Mutex::new(Vec::new()));
266        let base_url = spawn_mock_insert(captured.clone(), 200, "").await;
267        let client = IotdbClient::new(base_url, "root", "root");
268
269        let point = DataPoint::new("cpu")
270            .with_field("usage", FieldValue::Float(0.85))
271            .with_field("count", FieldValue::Int(3))
272            .with_field("active", FieldValue::Bool(true))
273            .with_field("name", FieldValue::String("web".into()))
274            .with_timestamp(1_700_000_000_000);
275        client.write(&[point]).await.unwrap();
276
277        let reqs = captured.lock().unwrap_or_else(|e| e.into_inner());
278        assert_eq!(reqs.len(), 1);
279        assert_eq!(reqs[0].path, "/rest/v2/insertTablet");
280        assert_eq!(reqs[0].header("content-type"), Some("application/json"));
281        // reqwest basic_auth("root", "root") → base64("root:root")
282        assert_eq!(reqs[0].header("authorization"), Some("Basic cm9vdDpyb290"));
283
284        let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
285        assert_eq!(body["device"], "cpu");
286        assert_eq!(body["is_aligned"], false);
287        assert_eq!(
288            body["timestamps"],
289            serde_json::json!([1_700_000_000_000_i64])
290        );
291        // values 为 [时间戳] × [字段] 的二维数组,单点单时间戳
292        assert_eq!(body["values"].as_array().unwrap().len(), 1);
293        // 字段类型编码与取值逐一对齐(HashMap 顺序不定,按名断言)
294        assert_field(&body, "usage", "DOUBLE", serde_json::json!(0.85));
295        assert_field(&body, "count", "INT64", serde_json::json!(3));
296        assert_field(&body, "active", "BOOLEAN", serde_json::json!(true));
297        assert_field(&body, "name", "TEXT", serde_json::json!("web"));
298    }
299
300    #[tokio::test]
301    async fn insert_tablet_defaults_timestamp_to_zero() {
302        let captured = Arc::new(Mutex::new(Vec::new()));
303        let base_url = spawn_mock_insert(captured.clone(), 200, "").await;
304        let client = IotdbClient::new(base_url, "root", "root");
305
306        client
307            .write(&[DataPoint::new("mem").with_field("used", FieldValue::Int(7))])
308            .await
309            .unwrap();
310
311        let reqs = captured.lock().unwrap_or_else(|e| e.into_inner());
312        let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
313        assert_eq!(body["timestamps"], serde_json::json!([0]));
314    }
315
316    #[tokio::test]
317    async fn insert_tablet_sends_one_request_per_point() {
318        let captured = Arc::new(Mutex::new(Vec::new()));
319        let base_url = spawn_mock_insert(captured.clone(), 200, "").await;
320        let client = IotdbClient::new(base_url, "root", "root");
321
322        let p1 = DataPoint::new("cpu").with_field("usage", FieldValue::Float(0.5));
323        let p2 = DataPoint::new("mem").with_field("used", FieldValue::Int(1));
324        client.write(&[p1, p2]).await.unwrap();
325
326        let reqs = captured.lock().unwrap_or_else(|e| e.into_inner());
327        assert_eq!(reqs.len(), 2, "每点独立一次 insertTablet 请求");
328        let devices: Vec<String> = reqs
329            .iter()
330            .map(|r| {
331                let body: serde_json::Value = serde_json::from_slice(&r.body).unwrap();
332                body["device"].as_str().unwrap().to_string()
333            })
334            .collect();
335        assert!(devices.iter().any(|d| d == "cpu"));
336        assert!(devices.iter().any(|d| d == "mem"));
337    }
338
339    #[tokio::test]
340    async fn write_returns_err_on_http_error() {
341        let captured = Arc::new(Mutex::new(Vec::new()));
342        let base_url = spawn_mock_insert(captured.clone(), 500, "boom").await;
343        let client = IotdbClient::new(base_url, "root", "root");
344        let err = client
345            .write(&[DataPoint::new("cpu").with_field("x", FieldValue::Int(1))])
346            .await
347            .unwrap_err();
348        assert!(err.to_string().contains("boom"), "got: {err}");
349    }
350
351    #[tokio::test]
352    async fn write_returns_err_on_2xx_with_failure_code() {
353        let captured = Arc::new(Mutex::new(Vec::new()));
354        // IoTDB REST v2 部分失败返回 HTTP 200 + body code != 200
355        let base_url = spawn_mock_insert(
356            captured.clone(),
357            200,
358            r#"{"code":501,"message":"table not exists"}"#,
359        )
360        .await;
361        let client = IotdbClient::new(base_url, "root", "root");
362        let err = client
363            .write(&[DataPoint::new("cpu").with_field("x", FieldValue::Int(1))])
364            .await
365            .unwrap_err();
366        assert!(err.to_string().contains("table not exists"), "got: {err}");
367    }
368
369    #[tokio::test]
370    async fn write_converts_non_finite_float_to_zero() {
371        let captured = Arc::new(Mutex::new(Vec::new()));
372        let base_url = spawn_mock_insert(captured.clone(), 200, "").await;
373        let client = IotdbClient::new(base_url, "root", "root");
374        client
375            .write(&[DataPoint::new("cpu").with_field("x", FieldValue::Float(f64::NAN))])
376            .await
377            .unwrap();
378        let reqs = captured.lock().unwrap_or_else(|e| e.into_inner());
379        let body: serde_json::Value = serde_json::from_slice(&reqs[0].body).unwrap();
380        assert_field(&body, "x", "DOUBLE", serde_json::json!(0));
381    }
382
383    /// mock IoTDB /rest/v2/query 端点:按给定状态码与响应体应答。
384    async fn spawn_mock_query(status: u16, body: &'static str) -> String {
385        let app = axum::Router::new().route(
386            "/rest/v2/query",
387            axum::routing::post(move || async move {
388                (
389                    axum::http::StatusCode::from_u16(status).unwrap(),
390                    axum::response::Response::new(axum::body::Body::from(body)),
391                )
392            }),
393        );
394        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
395        let addr = listener.local_addr().unwrap();
396        tokio::spawn(async move {
397            axum::serve(listener, app).await.unwrap();
398        });
399        format!("http://{addr}")
400    }
401
402    #[tokio::test]
403    async fn query_parses_successful_json_response() {
404        let body = r#"{"code":200,"expression":[{"alias":"x"}],"timestamp":[],"values":[]}"#;
405        let base_url = spawn_mock_query(200, body).await;
406        let client = IotdbClient::new(base_url, "root", "root");
407        let v = client.query("select x from root.s").await.unwrap();
408        assert_eq!(v["code"], 200);
409        assert_eq!(v["expression"][0]["alias"], "x");
410    }
411
412    #[tokio::test]
413    async fn query_returns_err_on_http_error() {
414        let base_url = spawn_mock_query(500, "query failed").await;
415        let client = IotdbClient::new(base_url, "root", "root");
416        let err = client.query("select 1").await.unwrap_err();
417        assert!(err.to_string().contains("query failed"), "got: {err}");
418    }
419
420    #[tokio::test]
421    async fn query_non_json_body_returns_parse_error() {
422        let base_url = spawn_mock_query(200, "not json").await;
423        let client = IotdbClient::new(base_url, "root", "root");
424        let err = client.query("select 1").await.unwrap_err();
425        assert!(err.to_string().contains("iotdb parse"), "got: {err}");
426    }
427}