Skip to main content

wowlab_cli/commands/snapshot/db/
copy.rs

1use sqlx::PgPool;
2use wowlab_common::output;
3use wowlab_types::copy::{CopyInsert, CopyRow};
4
5pub(super) const COPY_FLUSH_BYTES: usize = 4 * 1024 * 1024;
6
7pub(crate) async fn bulk_copy<T>(
8    pool: &PgPool,
9    table: &str,
10    label: &str,
11    rows: &[T],
12    patch: &str,
13) -> Result<(), sqlx::Error>
14where
15    T: CopyInsert,
16{
17    let pb = progress_bar(rows.len(), label);
18    let mut tx = pool.begin().await?;
19
20    sqlx::query(&format!("TRUNCATE {table} RESTART IDENTITY"))
21        .execute(&mut *tx)
22        .await?;
23
24    let stmt = format!("COPY {} ({}) FROM STDIN", table, T::COLUMNS.join(", "));
25    let mut sink = tx.copy_in_raw(&stmt).await?;
26    let mut w = CopyRow::new();
27
28    for (i, row) in rows.iter().enumerate() {
29        row.copy_row(&mut w, patch);
30        w.finish_row();
31
32        if w.len() >= COPY_FLUSH_BYTES {
33            sink.send(w.take()).await?;
34            pb.tick(i as u64 + 1);
35        }
36    }
37
38    sink.send(w.take()).await?;
39    sink.finish().await?;
40    tx.commit().await?;
41    pb.finish("");
42
43    Ok(())
44}
45
46pub(super) fn upsert_all_columns(columns: &[&str]) -> String {
47    columns[1..]
48        .iter()
49        .map(|c| format!("{c} = EXCLUDED.{c}"))
50        .chain(std::iter::once("updated_at = NOW()".to_string()))
51        .collect::<Vec<_>>()
52        .join(", ")
53}
54
55pub(super) fn progress_bar(total: usize, label: &str) -> output::ProgressBar {
56    output::ProgressBar::new(total as u64, label)
57}