1#![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 pub source: String,
22 pub format: String,
24 pub include_columns: Option<bool>,
26}
27
28pub async fn export_data(
30 headers: HeaderMap,
31 Query(params): Query<ExportParams>,
32) -> Result<impl IntoResponse, (StatusCode, String)> {
33 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 let include_columns = params.include_columns.unwrap_or(true);
53 let (source_sql, table_name) = parse_source(¶ms.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
77fn parse_source(source: &str) -> Result<(String, String), (StatusCode, String)> {
81 if let Some(table) = source.strip_prefix("table:") {
82 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 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
108fn is_simple_ident(s: &str) -> bool {
110 !s.is_empty() && s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_')
111}
112
113fn is_read_only_ast(stmt: &sqlparser::ast::Statement) -> bool {
115 use sqlparser::ast::Statement;
116 matches!(stmt, Statement::Query(_) | Statement::Explain { .. })
117}
118
119async 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(©_stmt).await.map_err(|e| {
131 tracing::error!("Export COPY failed: {}", e);
132 (StatusCode::INTERNAL_SERVER_ERROR, "COPY 失败".to_string())
133 })?;
134
135 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
141async 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
187fn 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 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 #[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 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 assert!(!is_simple_ident("../etc/passwd"));
267 assert!(!is_simple_ident("..\\windows"));
268 assert!(!is_simple_ident("/etc/passwd"));
269 assert!(!is_simple_ident("`posts`"));
271 assert!(!is_simple_ident("\"posts\""));
272 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 #[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 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 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 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 "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 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 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 assert!(parse_source("posts").is_err());
378 assert!(parse_source("").is_err());
379 }
380}