Skip to main content

yggdrasil/api/database/
export.rs

1//! 数据导出:Axum 流式路由(大文件不走 JSON 序列化)。
2//!
3//! 鉴权镜像 [`crate::api::upload`]:从 cookie 取 session,校验管理员。
4//! 两种来源:按表导出(表名白名单)/ 按查询导出(强制只读,AST 校验)。
5//! 两种格式:CSV(COPY TO STDOUT 流式)/ SQL(逐行拼 INSERT)。
6
7#![allow(clippy::unused_unit)]
8
9use axum::body::Body;
10use axum::extract::Query;
11use axum::http::{header, HeaderMap, StatusCode};
12use axum::response::IntoResponse;
13use futures::StreamExt;
14use serde::Deserialize;
15
16use crate::auth::session::parse_session_token;
17
18#[derive(Deserialize)]
19pub struct ExportParams {
20    /// `table:<name>` 或 `query:<SELECT ...>`。
21    pub source: String,
22    /// `sql` 或 `csv`。
23    pub format: String,
24    /// 是否带表头(CSV)/ 列名(SQL INSERT)。默认 true。
25    pub include_columns: Option<bool>,
26}
27
28/// POST/GET /api/database/export —— 流式导出。
29pub async fn export_data(
30    headers: HeaderMap,
31    Query(params): Query<ExportParams>,
32) -> Result<impl IntoResponse, (StatusCode, String)> {
33    // 1. 鉴权:cookie → session → admin
34    let cookie_header = headers
35        .get("cookie")
36        .and_then(|h| h.to_str().ok())
37        .unwrap_or("");
38    let token = parse_session_token(cookie_header).map(str::to_string);
39    let token = match token {
40        Some(t) => t,
41        None => return Err((StatusCode::UNAUTHORIZED, "未登录".to_string())),
42    };
43    let user = match crate::api::auth::get_user_by_token(&token).await {
44        Ok(Some(u)) => u,
45        _ => return Err((StatusCode::UNAUTHORIZED, "会话已过期".to_string())),
46    };
47    if user.role != crate::models::user::UserRole::Admin {
48        return Err((StatusCode::FORBIDDEN, "权限不足".to_string()));
49    }
50
51    // 2. 解析来源 + 白名单/只读校验
52    let include_columns = params.include_columns.unwrap_or(true);
53    let (source_sql, table_name) = parse_source(&params.source)?;
54
55    match params.format.as_str() {
56        "csv" => export_csv(source_sql, include_columns).await,
57        "sql" => export_sql(source_sql, include_columns, table_name).await,
58        _ => Err((StatusCode::BAD_REQUEST, "不支持的格式".to_string())),
59    }
60    .map(|(body, content_type, filename)| {
61        let disposition = format!("attachment; filename=\"{}\"", filename);
62        let disposition_value = axum::http::HeaderValue::from_str(&disposition)
63            .unwrap_or_else(|_| axum::http::HeaderValue::from_static("attachment"));
64        let content_type_value = axum::http::HeaderValue::from_str(content_type)
65            .unwrap_or_else(|_| axum::http::HeaderValue::from_static("application/octet-stream"));
66        (
67            StatusCode::OK,
68            [
69                (header::CONTENT_TYPE, content_type_value),
70                (header::CONTENT_DISPOSITION, disposition_value),
71            ],
72            body,
73        )
74    })
75}
76
77/// 解析导出来源,返回(内部 SQL,表名)。
78/// - `table:posts` → 校验表名合法后,`SELECT * FROM "posts"`(只读)+ 表名 "posts"。
79/// - `query:SELECT ...` → 校验为只读语句后原样返回,表名为 "export"。
80fn parse_source(source: &str) -> Result<(String, String), (StatusCode, String)> {
81    if let Some(table) = source.strip_prefix("table:") {
82        // 表名白名单:仅允许标识符字符,防注入
83        let t = table.trim();
84        if t.is_empty() || !is_simple_ident(t) {
85            return Err((StatusCode::BAD_REQUEST, "无效的表名".to_string()));
86        }
87        Ok((format!("SELECT * FROM \"{}\"", t), t.to_string()))
88    } else if let Some(query) = source.strip_prefix("query:") {
89        // 只读校验:sqlparser 解析后所有语句均为 Query/Explain
90        let dialect = sqlparser::dialect::PostgreSqlDialect {};
91        let parsed = sqlparser::parser::Parser::parse_sql(&dialect, query)
92            .map_err(|_| (StatusCode::BAD_REQUEST, "SQL 解析失败".to_string()))?;
93        if parsed.iter().any(|s| !is_read_only_ast(s)) {
94            return Err((
95                StatusCode::BAD_REQUEST,
96                "导出查询必须是只读(SELECT/EXPLAIN)".to_string(),
97            ));
98        }
99        Ok((query.to_string(), "export".to_string()))
100    } else {
101        Err((
102            StatusCode::BAD_REQUEST,
103            "source 必须是 table:<name> 或 query:<sql>".to_string(),
104        ))
105    }
106}
107
108/// 简单标识符校验(字母数字下划线,防 SQL 注入与路径穿越)。
109fn is_simple_ident(s: &str) -> bool {
110    !s.is_empty() && s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_')
111}
112
113/// 判断 AST 是否只读(SELECT/EXPLAIN)。
114fn is_read_only_ast(stmt: &sqlparser::ast::Statement) -> bool {
115    use sqlparser::ast::Statement;
116    matches!(stmt, Statement::Query(_) | Statement::Explain { .. })
117}
118
119/// CSV 导出:用 COPY ... TO STDOUT WITH CSV 流式(大表不 OOM)。
120async fn export_csv(
121    source_sql: String,
122    include_columns: bool,
123) -> Result<(Body, &'static str, String), (StatusCode, String)> {
124    let client = crate::db::pool::get_conn()
125        .await
126        .map_err(|_| (StatusCode::SERVICE_UNAVAILABLE, "数据库不可用".to_string()))?;
127
128    let header_clause = if include_columns { "HEADER" } else { "" };
129    let copy_stmt = format!("COPY ({}) TO STDOUT WITH CSV {}", source_sql, header_clause);
130    let stream = client.copy_out(&copy_stmt).await.map_err(|e| {
131        tracing::error!("Export COPY failed: {}", e);
132        (StatusCode::INTERNAL_SERVER_ERROR, "COPY 失败".to_string())
133    })?;
134
135    // tokio_postgres 的 copy_out 流产出 Bytes,直接转 axum Body。
136    let mapped = stream.map(|res| res.map_err(std::io::Error::other));
137    let body = Body::from_stream(mapped);
138    Ok((body, "text/csv; charset=utf-8", "export.csv".to_string()))
139}
140
141/// SQL 导出:逐行拼 INSERT(纯数据,不含 DDL)。
142async fn export_sql(
143    source_sql: String,
144    include_columns: bool,
145    table_name: String,
146) -> Result<(Body, &'static str, String), (StatusCode, String)> {
147    let client = crate::db::pool::get_conn()
148        .await
149        .map_err(|_| (StatusCode::SERVICE_UNAVAILABLE, "数据库不可用".to_string()))?;
150
151    let rows = client.query(&source_sql, &[]).await.map_err(|e| {
152        tracing::error!("Export query failed: {}", e);
153        (StatusCode::INTERNAL_SERVER_ERROR, "查询失败".to_string())
154    })?;
155
156    let columns: Vec<String> = rows
157        .first()
158        .map(|r| r.columns().iter().map(|c| c.name().to_string()).collect())
159        .unwrap_or_default();
160
161    let mut out = String::new();
162    out.push_str("-- Yggdrasil 数据导出(仅数据,不含 schema)\n");
163    let col_clause = if include_columns && !columns.is_empty() {
164        format!("({})", columns.join(", "))
165    } else {
166        String::new()
167    };
168
169    for r in &rows {
170        let vals: Vec<String> = (0..r.len()).map(|i| sql_quote_cell(r, i)).collect();
171        out.push_str(&format!(
172            "INSERT INTO {} {} VALUES ({});\n",
173            table_name,
174            col_clause,
175            vals.join(", ")
176        ));
177    }
178
179    let body = Body::from(out);
180    Ok((
181        body,
182        "application/sql; charset=utf-8",
183        "export.sql".to_string(),
184    ))
185}
186
187/// 把一个单元格值转成 SQL 字面量(字符串单引号转义,数字原样,NULL)。
188fn sql_quote_cell(row: &tokio_postgres::Row, idx: usize) -> String {
189    let ty = row
190        .columns()
191        .get(idx)
192        .map(|c| c.type_().name())
193        .unwrap_or("");
194    match ty {
195        "int2" => row
196            .try_get::<_, Option<i16>>(idx)
197            .ok()
198            .flatten()
199            .map(|v| v.to_string())
200            .unwrap_or_else(|| "NULL".into()),
201        "int4" => row
202            .try_get::<_, Option<i32>>(idx)
203            .ok()
204            .flatten()
205            .map(|v| v.to_string())
206            .unwrap_or_else(|| "NULL".into()),
207        "int8" => row
208            .try_get::<_, Option<i64>>(idx)
209            .ok()
210            .flatten()
211            .map(|v| v.to_string())
212            .unwrap_or_else(|| "NULL".into()),
213        "float4" | "float8" => row
214            .try_get::<_, Option<f64>>(idx)
215            .ok()
216            .flatten()
217            .map(|v| v.to_string())
218            .unwrap_or_else(|| "NULL".into()),
219        "bool" => row
220            .try_get::<_, Option<bool>>(idx)
221            .ok()
222            .flatten()
223            .map(|v| if v { "TRUE" } else { "FALSE" }.to_string())
224            .unwrap_or_else(|| "NULL".into()),
225        _ => {
226            // 字符串/时间戳等:单引号转义
227            match row.try_get::<_, Option<String>>(idx) {
228                Ok(Some(s)) => format!("'{}'", s.replace('\'', "''")),
229                _ => "NULL".into(),
230            }
231        }
232    }
233}
234
235#[cfg(all(test, feature = "server"))]
236mod tests {
237    use super::*;
238
239    // ---- is_simple_ident:表名白名单(防 SQL 注入与路径穿越)----
240
241    #[test]
242    fn simple_ident_accepts_legal_names() {
243        assert!(is_simple_ident("posts"));
244        assert!(is_simple_ident("user_posts"));
245        assert!(is_simple_ident("table1"));
246        assert!(is_simple_ident("a"));
247        assert!(is_simple_ident("_private"));
248        assert!(is_simple_ident("ABC123xyz"));
249    }
250
251    #[test]
252    fn simple_ident_rejects_empty_and_whitespace() {
253        assert!(!is_simple_ident(""));
254        assert!(!is_simple_ident("   "));
255    }
256
257    #[test]
258    fn simple_ident_rejects_injection_vectors() {
259        // SQL 注入:引号、分号、注释、空格分隔
260        assert!(!is_simple_ident("posts; DROP TABLE users"));
261        assert!(!is_simple_ident("posts' OR '1'='1"));
262        assert!(!is_simple_ident("posts--"));
263        assert!(!is_simple_ident("public.posts"));
264        assert!(!is_simple_ident("posts where 1=1"));
265        // 路径穿越
266        assert!(!is_simple_ident("../etc/passwd"));
267        assert!(!is_simple_ident("..\\windows"));
268        assert!(!is_simple_ident("/etc/passwd"));
269        // 反引号 / 双引号标识符语法
270        assert!(!is_simple_ident("`posts`"));
271        assert!(!is_simple_ident("\"posts\""));
272        // 连字符(常见合法表名片段,但本白名单禁止,防 `a-b` 被 SQL 解析为减法)
273        assert!(!is_simple_ident("my-posts"));
274    }
275
276    #[test]
277    fn simple_ident_rejects_unicode_and_non_ascii() {
278        assert!(!is_simple_ident("文章"));
279        assert!(!is_simple_ident("posts²"));
280        assert!(!is_simple_ident("café"));
281    }
282
283    // ---- parse_source:来源解析 + 只读校验(安全入口)----
284
285    #[test]
286    fn parse_table_source_builds_select_star() {
287        let (sql, name) = parse_source("table:posts").unwrap();
288        assert_eq!(sql, "SELECT * FROM \"posts\"");
289        assert_eq!(name, "posts");
290    }
291
292    #[test]
293    fn parse_table_source_trims_whitespace() {
294        let (sql, name) = parse_source("table:  posts  ").unwrap();
295        assert_eq!(sql, "SELECT * FROM \"posts\"");
296        assert_eq!(name, "posts");
297    }
298
299    #[test]
300    fn parse_table_source_rejects_bad_name() {
301        // 非法表名(注入/穿越/空)必须在拼进 SQL 前被拒,防 SQL 注入。
302        for bad in [
303            "table:",
304            "table:   ",
305            "table:posts; DROP TABLE users",
306            "table:../secret",
307            "table:a b",
308            "table:\"x\"",
309        ] {
310            let (code, msg) = parse_source(bad).expect_err(bad);
311            assert_eq!(code, StatusCode::BAD_REQUEST);
312            assert!(msg.contains("表名"), "{bad}: {msg}");
313        }
314    }
315
316    #[test]
317    fn parse_query_source_accepts_select_and_explain() {
318        assert!(parse_source("query:SELECT * FROM posts").is_ok());
319        assert!(parse_source("query:SELECT id, title FROM posts WHERE id > 0").is_ok());
320        assert!(parse_source("query:EXPLAIN SELECT * FROM posts").is_ok());
321        assert!(parse_source("query:SELECT 1").is_ok());
322        // 复杂但只读的查询
323        assert!(parse_source("query:WITH t AS (SELECT 1) SELECT * FROM t").is_ok());
324    }
325
326    #[test]
327    fn parse_query_source_rejects_all_write_operations() {
328        // 每一种写/破坏操作都必须被 AST 校验拦截——这是数据导出的核心安全不变量。
329        for evil in [
330            "query:INSERT INTO posts VALUES (1)",
331            "query:UPDATE posts SET title='x'",
332            "query:DELETE FROM posts",
333            "query:DROP TABLE posts",
334            "query:TRUNCATE posts",
335            "query:CREATE TABLE evil (id int)",
336            "query:ALTER TABLE posts ADD COLUMN x int",
337            "query:GRANT SELECT ON posts TO public",
338            // 即便嵌在 SELECT 里,子查询的写操作也要拒
339            "query:INSERT INTO logs VALUES (1) RETURNING 1",
340        ] {
341            let (code, msg) = parse_source(evil).expect_err(evil);
342            assert_eq!(code, StatusCode::BAD_REQUEST, "{evil}");
343            // 写操作要么被 AST 校验拒("只读"),要么 sqlparser 解析阶段就失败
344            assert!(
345                msg.contains("只读") || msg.contains("解析失败"),
346                "{evil}: {msg}"
347            );
348        }
349    }
350
351    #[test]
352    fn parse_query_source_rejects_unparseable_sql() {
353        for bad in [
354            "query:这不是SQL",
355            "query:SELECT FROM",
356            "query:@#$%",
357            "query:1+1",
358        ] {
359            assert!(parse_source(bad).is_err(), "{bad} 应解析失败或被拒");
360        }
361    }
362
363    #[test]
364    fn parse_query_source_allows_empty_result_set() {
365        // sqlparser 容忍纯分号(解析为零语句),`any()` 在空集上为 false → 放行。
366        // 锁定该契约:这只产生空导出,无数据泄露,属可接受行为;
367        // 若未来想收紧为"空查询拒绝",此测试会提醒你更新。
368        assert!(parse_source("query:;;;").is_ok());
369    }
370
371    #[test]
372    fn parse_source_rejects_unknown_prefix() {
373        let (code, msg) = parse_source("unknown:posts").expect_err("未知前缀");
374        assert_eq!(code, StatusCode::BAD_REQUEST);
375        assert!(msg.contains("source"));
376        // 完全无前缀的裸字符串也应被拒
377        assert!(parse_source("posts").is_err());
378        assert!(parse_source("").is_err());
379    }
380}