1use std::any::Any;
11use std::collections::BTreeMap;
12use std::io::{BufRead, Write};
13
14use lora_executor::{ExecuteOptions, LoraValue, QueryResult, ResultFormat};
15use lora_io::{
16 CsvDecoder, CsvEncoder, Format, JsonArrayDecoder, JsonArrayEncoder, JsonlDecoder, JsonlEncoder,
17 RowDecoder, RowEncoder, RowMapping,
18};
19use lora_store::{GraphStorage, GraphStorageMut, InMemoryGraph};
20
21use crate::error::{LoraError, LoraErrorCode};
22use crate::Database;
23
24#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
26pub struct ExportStats {
27 pub rows: u64,
28}
29
30#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
32pub struct ImportStats {
33 pub rows: u64,
34 pub batches: u64,
35}
36
37pub const DEFAULT_IMPORT_BATCH_SIZE: usize = 1_000;
41
42impl<S> Database<S>
43where
44 S: GraphStorage + GraphStorageMut + Any + Clone + Send + Sync + 'static,
45{
46 pub fn export_query<W: Write>(
55 &self,
56 query: &str,
57 params: BTreeMap<String, LoraValue>,
58 format: Format,
59 writer: W,
60 ) -> Result<ExportStats, LoraError> {
61 match format {
62 Format::Jsonl => self.drive_export(query, params, JsonlEncoder::new(writer)),
63 Format::Json => self.drive_export(query, params, JsonArrayEncoder::new(writer)),
64 Format::Csv => self.drive_export(query, params, CsvEncoder::new(writer)),
65 }
66 }
67
68 fn drive_export<E: RowEncoder>(
69 &self,
70 query: &str,
71 params: BTreeMap<String, LoraValue>,
72 mut encoder: E,
73 ) -> Result<ExportStats, LoraError> {
74 let result = self.execute_with_params(
80 query,
81 Some(ExecuteOptions {
82 format: ResultFormat::RowArrays,
83 }),
84 params,
85 )?;
86 let QueryResult::RowArrays(arrays) = result else {
87 return Err(LoraError::new(
88 LoraErrorCode::Internal,
89 "expected RowArrays result for export".to_string(),
90 ));
91 };
92 encoder.begin(&arrays.columns).map_err(io_err)?;
93 let mut rows = 0u64;
94 for row_vals in arrays.rows {
95 let pairs: Vec<(String, LoraValue)> =
96 arrays.columns.iter().cloned().zip(row_vals).collect();
97 encoder.write_named_row(&pairs).map_err(io_err)?;
98 rows += 1;
99 }
100 encoder.finish().map_err(io_err)?;
101 Ok(ExportStats { rows })
102 }
103
104 pub fn import_rows<R: BufRead>(
109 &self,
110 reader: R,
111 format: Format,
112 mapping: &RowMapping,
113 batch_size: Option<usize>,
114 ) -> Result<ImportStats, LoraError> {
115 let template = mapping.to_cypher().map_err(|e| {
116 LoraError::new(
117 LoraErrorCode::InvalidParams,
118 format!("invalid row mapping: {e}"),
119 )
120 })?;
121 self.import_with_template(reader, format, &template, batch_size)
122 }
123
124 pub fn import_with_template<R: BufRead>(
129 &self,
130 reader: R,
131 format: Format,
132 template: &str,
133 batch_size: Option<usize>,
134 ) -> Result<ImportStats, LoraError> {
135 let batch = batch_size.unwrap_or(DEFAULT_IMPORT_BATCH_SIZE).max(1);
136 match format {
137 Format::Jsonl => self.drive_import(template, JsonlDecoder::new(reader), batch),
138 Format::Json => self.drive_import(template, JsonArrayDecoder::new(reader), batch),
139 Format::Csv => self.drive_import(template, CsvDecoder::new(reader), batch),
140 }
141 }
142
143 fn drive_import<D: RowDecoder>(
144 &self,
145 template: &str,
146 mut decoder: D,
147 batch_size: usize,
148 ) -> Result<ImportStats, LoraError> {
149 decoder.header().map_err(io_err)?;
152
153 let mut stats = ImportStats::default();
154 let mut buf: Vec<LoraValue> = Vec::with_capacity(batch_size);
155
156 while let Some(cells) = decoder.next_row().map_err(io_err)? {
157 let row_map: BTreeMap<String, LoraValue> = cells.into_iter().collect();
158 buf.push(LoraValue::Map(row_map));
159 if buf.len() >= batch_size {
160 self.flush_batch(template, &mut buf, &mut stats)?;
161 }
162 }
163 if !buf.is_empty() {
164 self.flush_batch(template, &mut buf, &mut stats)?;
165 }
166 Ok(stats)
167 }
168
169 fn flush_batch(
170 &self,
171 template: &str,
172 buf: &mut Vec<LoraValue>,
173 stats: &mut ImportStats,
174 ) -> Result<(), LoraError> {
175 let batch = std::mem::take(buf);
176 let batch_len = batch.len() as u64;
177 let mut params = BTreeMap::new();
178 params.insert("rows".to_string(), LoraValue::List(batch));
179 self.execute_with_params(template, None, params)?;
180 stats.rows += batch_len;
181 stats.batches += 1;
182 Ok(())
183 }
184}
185
186fn io_err(err: std::io::Error) -> LoraError {
187 LoraError::with_source(LoraErrorCode::Io, format!("{err}"), err)
188}
189
190impl Database<InMemoryGraph> {
191 pub fn export_query_streaming<W: Write>(
204 &self,
205 query: &str,
206 params: BTreeMap<String, LoraValue>,
207 format: Format,
208 writer: W,
209 ) -> Result<ExportStats, LoraError> {
210 match format {
211 Format::Jsonl => self.drive_streaming_export(query, params, JsonlEncoder::new(writer)),
212 Format::Json => {
213 self.drive_streaming_export(query, params, JsonArrayEncoder::new(writer))
214 }
215 Format::Csv => self.drive_streaming_export(query, params, CsvEncoder::new(writer)),
216 }
217 }
218
219 fn drive_streaming_export<E: RowEncoder>(
220 &self,
221 query: &str,
222 params: BTreeMap<String, LoraValue>,
223 mut encoder: E,
224 ) -> Result<ExportStats, LoraError> {
225 let mut stream = self.stream_with_params(query, params)?;
226 let columns = stream.columns().to_vec();
227 encoder.begin(&columns).map_err(io_err)?;
228 let mut rows = 0u64;
229 while let Some(row) = stream.next_row().map_err(LoraError::from_anyhow)? {
230 encoder.write_row(&row).map_err(io_err)?;
231 rows += 1;
232 }
233 encoder.finish().map_err(io_err)?;
234 Ok(ExportStats { rows })
235 }
236}
237
238#[cfg(test)]
239mod tests {
240 use super::*;
241
242 #[test]
243 fn round_trip_node_import_then_export() {
244 let db = Database::in_memory();
245
246 let csv = "name:string,age:int\nalice,30\nbob,25\n";
247 let mapping = RowMapping::Node {
248 label: "User".into(),
249 id_column: None,
250 id_property: None,
251 properties: vec![
252 lora_io::ColumnSpec::identity("name"),
253 lora_io::ColumnSpec::identity("age"),
254 ],
255 };
256 let stats = db
257 .import_rows(std::io::Cursor::new(csv), Format::Csv, &mapping, Some(10))
258 .unwrap();
259 assert_eq!(stats.rows, 2);
260 assert_eq!(stats.batches, 1);
261
262 let mut out = Vec::new();
263 let export = db
264 .export_query(
265 "MATCH (u:User) RETURN u.name AS name, u.age AS age ORDER BY name",
266 BTreeMap::new(),
267 Format::Jsonl,
268 &mut out,
269 )
270 .unwrap();
271 assert_eq!(export.rows, 2);
272 let text = std::str::from_utf8(&out).unwrap();
273 let lines: Vec<_> = text.lines().collect();
274 assert_eq!(lines.len(), 2);
275 assert!(lines[0].contains("\"alice\""));
276 assert!(lines[1].contains("\"bob\""));
277 }
278
279 #[test]
285 fn import_handles_property_names_with_spaces_in_non_empty_graph() {
286 let db = Database::in_memory();
287 db.execute("CREATE (:Other {tag: 1})", None).unwrap();
288
289 let csv = "User Id,First Name\nu1,Alice\nu2,Bob\nu3,Carol\n";
290 let mapping = RowMapping::Node {
291 label: "People".into(),
292 id_column: None,
293 id_property: None,
294 properties: vec![
295 lora_io::ColumnSpec::identity("User Id"),
296 lora_io::ColumnSpec::identity("First Name"),
297 ],
298 };
299 let stats = db
300 .import_rows(std::io::Cursor::new(csv), Format::Csv, &mapping, Some(2))
301 .expect("import with quoted property names should succeed");
302 assert_eq!(stats.rows, 3);
303 assert!(stats.batches >= 2, "expected at least 2 batches");
304 }
305
306 #[test]
307 fn cypher_template_import() {
308 let db = Database::in_memory();
309 let jsonl = "{\"name\":\"alice\",\"age\":30}\n{\"name\":\"bob\",\"age\":25}\n";
310 let stats = db
311 .import_with_template(
312 std::io::Cursor::new(jsonl),
313 Format::Jsonl,
314 "UNWIND $rows AS r CREATE (:Person {name: r.name, age: r.age})",
315 Some(50),
316 )
317 .unwrap();
318 assert_eq!(stats.rows, 2);
319
320 let res = db
321 .execute(
322 "MATCH (p:Person) RETURN count(p) AS c",
323 Some(ExecuteOptions {
324 format: ResultFormat::RowArrays,
325 }),
326 )
327 .unwrap();
328 match res {
329 QueryResult::RowArrays(r) => {
330 let LoraValue::Int(c) = &r.rows[0][0] else {
331 panic!("expected Int");
332 };
333 assert_eq!(*c, 2);
334 }
335 other => panic!("unexpected result {other:?}"),
336 }
337 }
338}