1use 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 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 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 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 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 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 assert_eq!(body["values"].as_array().unwrap().len(), 1);
293 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 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 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}