1use crate::backup::extractor::{BackupSpec, extract_backup_spec};
6use crate::config::MySqlConnection;
7use crate::policy::classifier::ClassifiedStatement;
8use crate::policy::model::{Policy, SqlCategory};
9use futures_util::StreamExt;
10use mysql_async::prelude::Queryable;
11use mysql_async::{OptsBuilder, PoolConstraints, PoolOpts, Row, SslOpts, Value};
12use std::sync::Arc;
13use std::time::{Duration, Instant};
14use thiserror::Error;
15use zeroize::Zeroizing;
16
17use crate::audit::AuditDb;
18
19#[derive(Debug, Error)]
20pub enum MySqlError {
21 #[error("mysql error: {0}")]
22 Driver(#[from] mysql_async::Error),
23 #[error("read-only transaction could not be established: {0}")]
24 ReadOnlyTx(String),
25 #[error("statement timed out after {0}ms")]
26 Timeout(u64),
27 #[error("backup overflow: {0}")]
28 BackupOverflow(String),
29 #[error("backup capture failed: {0} — mutation denied")]
30 BackupFailed(String),
31 #[error("pool error: {0}")]
32 Pool(String),
33 #[error("{0}")]
34 Uncertain(String),
35 #[error("DDL target not found: {0}")]
36 DdlNotFound(String),
37 #[error("[ddl_precondition_changed] {0}")]
38 DdlPreconditionChanged(String),
39}
40
41#[derive(Debug)]
42pub struct ExecuteResult {
43 pub journal_id: Option<i64>,
45 pub ddl_absent_targets: Vec<String>,
49 pub ddl_executed_targets: Vec<String>,
53 pub ddl_no_op: bool,
56 pub warnings: Vec<&'static str>,
59 pub rows: Vec<serde_json::Value>,
60 pub fields: Vec<String>,
61 pub affected_rows: u64,
62 pub truncated: bool,
63 pub duration_ms: u64,
64 pub backup_id: Option<i64>,
65 pub backup_row_count: u64,
66}
67
68pub const CONNECT_TIMEOUT: Duration = Duration::from_secs(15);
74
75pub fn build_opts(
76 conn: &MySqlConnection,
77 password: &str,
78 database: Option<&str>,
79 host_override: Option<&str>,
80 port_override: Option<u16>,
81) -> OptsBuilder {
82 let ssl: Option<SslOpts> = if conn.ssl {
83 let mut ssl = SslOpts::default();
84 if let Some(name) = &conn.ssl_server_name {
85 ssl = ssl.with_danger_tls_hostname_override(Some(name.clone()));
88 }
89 if let Some(ca_path) = &conn.ssl_ca_path {
90 let expanded = crate::app::paths::expand_tilde(ca_path);
91 ssl = ssl.with_root_certs(vec![expanded.to_path_buf().into()]);
92 }
93 Some(ssl)
94 } else {
95 None
96 };
97 OptsBuilder::default()
98 .ip_or_hostname(
99 host_override
100 .map(str::to_string)
101 .unwrap_or_else(|| conn.host.clone()),
102 )
103 .tcp_port(port_override.unwrap_or(conn.port))
104 .user(Some(conn.user.clone()))
105 .pass(Some(password.to_string()))
106 .db_name(
107 database
108 .map(str::to_string)
109 .or_else(|| conn.database.clone()),
110 )
111 .secure_auth(true)
112 .ssl_opts(ssl)
113 .conn_ttl(Duration::from_secs(600))
114 .tcp_keepalive(Some(Duration::from_secs(30)))
115 .pool_opts(
116 PoolOpts::default()
117 .with_constraints(PoolConstraints::new(1, 4).expect("1<=4"))
118 .with_inactive_connection_ttl(Duration::from_secs(300)),
122 )
123}
124
125pub fn value_to_json(v: &Value) -> serde_json::Value {
129 match v {
130 Value::NULL => serde_json::Value::Null,
131 Value::Bytes(b) => match std::str::from_utf8(b) {
132 Ok(s) => serde_json::json!(s),
133 Err(_) => {
134 use base64::Engine;
135 serde_json::json!(base64::engine::general_purpose::STANDARD.encode(b))
136 }
137 },
138 Value::Int(i) => serde_json::json!(i),
139 Value::UInt(u) => serde_json::json!(u),
140 Value::Float(f) => serde_json::json!(f),
141 Value::Double(d) => serde_json::json!(d),
142 Value::Date(y, m, d, hh, mm, ss, us) => {
143 if *us == 0 {
144 serde_json::json!(format!("{y:04}-{m:02}-{d:02} {hh:02}:{mm:02}:{ss:02}"))
145 } else {
146 serde_json::json!(format!(
147 "{y:04}-{m:02}-{d:02} {hh:02}:{mm:02}:{ss:02}.{us:06}"
148 ))
149 }
150 }
151 Value::Time(neg, d, h, m, s, us) => {
152 let base = format!("{d:02}:{h:02}:{m:02}:{s:02}");
153 let with_us = if *us == 0 {
154 base
155 } else {
156 format!("{base}.{us:06}")
157 };
158 serde_json::json!(if *neg { format!("-{with_us}") } else { with_us })
159 }
160 }
161}
162
163static MANAGER: std::sync::OnceLock<super::pool::PoolManager> = std::sync::OnceLock::new();
164
165pub fn pool_manager() -> &'static super::pool::PoolManager {
166 MANAGER.get_or_init(super::pool::PoolManager::new)
167}
168
169pub struct MySqlExecuteParams<'a> {
170 pub connection: &'a MySqlConnection,
171 pub request_id: String,
173 pub databases_for_log: Vec<String>,
175 pub password: Zeroizing<String>,
176 pub sql: &'a str,
177 pub classified: &'a ClassifiedStatement,
178 pub policy: &'a Policy,
179 pub database: Option<&'a str>,
180 pub audit: Option<Arc<AuditDb>>,
181 pub revision: u64,
182 pub tunnel_endpoint: Option<crate::sql::ssh::TunnelLease>,
185 pub expected_ddl_targets: Option<Vec<(String, String)>>,
190}
191
192pub fn binary_json(bytes: &[u8]) -> serde_json::Value {
196 use base64::Engine;
197 serde_json::json!({
198 "type": "binary",
199 "encoding": "base64",
200 "data": base64::engine::general_purpose::STANDARD.encode(bytes),
201 })
202}
203
204pub fn value_with_column_type(
208 v: &Value,
209 ct: mysql_async::consts::ColumnType,
210 charset: u16,
211) -> serde_json::Value {
212 use mysql_async::consts::ColumnType;
213 if let Value::Bytes(b) = v
219 && (matches!(
220 ct,
221 ColumnType::MYSQL_TYPE_TINY_BLOB
222 | ColumnType::MYSQL_TYPE_MEDIUM_BLOB
223 | ColumnType::MYSQL_TYPE_LONG_BLOB
224 | ColumnType::MYSQL_TYPE_BLOB
225 | ColumnType::MYSQL_TYPE_BIT
226 | ColumnType::MYSQL_TYPE_GEOMETRY
227 ) || (charset == 63
228 && matches!(
229 ct,
230 ColumnType::MYSQL_TYPE_STRING | ColumnType::MYSQL_TYPE_VAR_STRING
231 )))
232 {
233 return binary_json(b);
234 }
235 let bytes = match v {
236 Value::Bytes(b) => b,
237 other => return value_to_json(other),
238 };
239 let text = match std::str::from_utf8(bytes) {
240 Ok(s) => s,
241 Err(_) => return value_to_json(v),
242 };
243 let numeric = matches!(
244 ct,
245 ColumnType::MYSQL_TYPE_TINY
246 | ColumnType::MYSQL_TYPE_SHORT
247 | ColumnType::MYSQL_TYPE_LONG
248 | ColumnType::MYSQL_TYPE_INT24
249 | ColumnType::MYSQL_TYPE_FLOAT
250 | ColumnType::MYSQL_TYPE_DOUBLE
251 | ColumnType::MYSQL_TYPE_YEAR
252 );
253 if numeric {
254 if let Ok(i) = text.parse::<i64>() {
255 return serde_json::json!(i);
256 }
257 if let Ok(u) = text.parse::<u64>() {
258 return serde_json::json!(u);
259 }
260 if let Ok(f) = text.parse::<f64>() {
261 return serde_json::json!(f);
262 }
263 }
264 serde_json::json!(text)
265}
266
267const READ_CATEGORIES: [SqlCategory; 1] = [SqlCategory::Read];
268const RESULT_BYTE_CAP: u64 = 4 * 1024 * 1024;
269
270fn is_mariadb(server_version: &str) -> bool {
271 server_version.to_ascii_lowercase().contains("mariadb")
272}
273
274pub async fn execute_mysql_statement(
278 params: MySqlExecuteParams<'_>,
279) -> Result<ExecuteResult, MySqlError> {
280 let start = Instant::now();
281 if crate::policy::classifier::looks_like_multiple_statements(params.sql) {
286 return Err(MySqlError::Uncertain(
287 "multi-statement input is not allowed (single statement per call)".into(),
288 ));
289 }
290 let is_read = READ_CATEGORIES.contains(¶ms.classified.category);
291 let (host, port, tunnel_generation) = match ¶ms.tunnel_endpoint {
292 Some(lease) => (lease.host.clone(), lease.port, Some(lease.generation)),
293 None => (params.connection.host.clone(), params.connection.port, None),
294 };
295 crate::app::test_mode::check_mysql_endpoint(&host, port).map_err(MySqlError::Pool)?;
298 let pool = pool_manager()
299 .verified_pool(
300 params.connection,
301 ¶ms.password,
302 params.database,
303 params.revision,
304 Some(&host),
305 Some(port),
306 tunnel_generation,
307 )
308 .await
309 .map_err(|e| MySqlError::Pool(e.to_string()))?;
310 let setup_budget = Duration::from_millis(params.policy.stmt_timeout_ms.max(1) as u64);
316 let setup = async {
317 let mut conn = pool.get_conn().await?;
318
319 let executing_id: Option<u64> = conn
321 .exec_first::<(u64,), _, _>("SELECT CONNECTION_ID()", ())
322 .await
323 .ok()
324 .flatten()
325 .map(|(id,)| id);
326
327 let (major, minor, _patch) = conn.server_version();
328 Ok::<_, MySqlError>((conn, executing_id, format!("{major}.{minor}")))
329 };
330 let (mut conn, executing_id, server_version) = tokio::time::timeout(setup_budget, setup)
331 .await
332 .map_err(|_| MySqlError::Timeout(params.policy.stmt_timeout_ms as u64))??;
333 let mut in_tx = false;
334 if params.classified.category != SqlCategory::TxCtrl {
335 let stmt = if is_read {
336 "START TRANSACTION READ ONLY"
337 } else {
338 "START TRANSACTION READ WRITE"
339 };
340 if let Err(e) = conn.query_drop(stmt).await {
341 if is_read {
342 return Err(MySqlError::ReadOnlyTx(e.to_string()));
343 }
344 if crate::backup::extractor::is_backup_required(params.classified.ast_type) {
351 return Err(MySqlError::Uncertain(format!(
352 "START TRANSACTION READ WRITE failed ({e}); refusing to run a backed-up \
353 write in autocommit — the pre-image would not be atomic with the mutation"
354 )));
355 }
356 eprintln!(
357 "[sequel-mcp] {} START TRANSACTION failed: {e}; continuing without explicit tx",
358 params.connection.name
359 );
360 } else {
361 in_tx = true;
362 }
363 if params.policy.stmt_timeout_ms > 0 {
366 let timeout_sql = if is_mariadb(&server_version) {
367 format!(
368 "SET SESSION max_statement_time = {}",
369 params.policy.stmt_timeout_ms.div_ceil(1000)
370 )
371 } else {
372 format!(
373 "SET SESSION max_execution_time = {}",
374 params.policy.stmt_timeout_ms
375 )
376 };
377 let _ = conn.query_drop(timeout_sql).await;
378 }
379 }
380
381 let journal = if matches!(
389 params.classified.category,
390 SqlCategory::Write | SqlCategory::Ddl | SqlCategory::Admin
391 ) && params.audit.is_some()
392 {
393 let Some(audit) = params.audit.as_ref() else {
394 unreachable!("guarded by is_some above")
395 };
396 let _ = crate::backup::journal::ensure_table(audit);
397 crate::backup::journal::Journal::create(
398 audit,
399 ¶ms.request_id,
400 ¶ms.connection.name,
401 ¶ms.databases_for_log,
402 params.classified.category.as_str(),
403 )
404 .ok()
405 } else {
406 None
407 };
408
409 use crate::backup::journal::JournalState;
410 let mut ddl_absent: Vec<(String, String)> = Vec::new();
414 let mut ddl_executed: Vec<String> = Vec::new();
415 let mut ddl_effective_sql: Option<String> = None;
416 let mut ddl_detail: Option<String> = None;
417
418 if params.classified.category == SqlCategory::Ddl {
427 match super::ddl::preflight_ddl(
428 &mut conn,
429 params.classified,
430 params.database.or(params.connection.database.as_deref()),
431 )
432 .await
433 {
434 Ok(super::ddl::DdlPreflight::Present) => {
435 if let (Some(expected), "drop") = (
440 params.expected_ddl_targets.as_ref(),
441 params.classified.ast_type,
442 ) {
443 let fallback = params.database.or(params.connection.database.as_deref());
444 for target in ¶ms.classified.mutated_tables {
445 let Some(schema) = target
446 .database
447 .clone()
448 .or_else(|| fallback.map(str::to_string))
449 else {
450 continue;
451 };
452 if !expected.contains(&(schema.clone(), target.table.clone())) {
453 let detail = format!(
454 "table {}.{} changed existence after the approved plan; nothing executed; re-run for a fresh plan",
455 schema, target.table
456 );
457 if let Some(j) = &journal {
458 let _ = j.transition(JournalState::Failed, Some(&detail));
459 }
460 return Err(MySqlError::DdlPreconditionChanged(detail));
461 }
462 }
463 }
464 }
465 Ok(super::ddl::DdlPreflight::Mixed { existing, missing }) => {
466 if let Some(expected) = params.expected_ddl_targets.as_ref() {
469 for (schema, table) in &existing {
470 if !expected.contains(&(schema.clone(), table.clone())) {
471 let detail = format!(
472 "table {schema}.{table} changed existence after the approved plan; nothing executed; re-run for a fresh plan"
473 );
474 if let Some(j) = &journal {
475 let _ = j.transition(JournalState::Failed, Some(&detail));
476 }
477 return Err(MySqlError::DdlPreconditionChanged(detail));
478 }
479 }
480 }
481 let absent = missing
482 .iter()
483 .map(|(s, t)| format!("{s}.{t}"))
484 .collect::<Vec<_>>()
485 .join(", ");
486 let kept = existing
487 .iter()
488 .map(|(s, t)| format!("{s}.{t}"))
489 .collect::<Vec<_>>()
490 .join(", ");
491 match super::ddl::rewrite_drop_subset(
492 params.classified.drop_object_type.unwrap_or("other"),
493 &existing,
494 ) {
495 Some(rewritten) => {
496 ddl_effective_sql = Some(rewritten);
497 ddl_executed = existing.iter().map(|(s, t)| format!("{s}.{t}")).collect();
498 ddl_absent = missing.clone();
499 ddl_detail = Some(format!(
500 "IF EXISTS: executing approved subset [{kept}]; absent targets: [{absent}]"
501 ));
502 }
503 None => {
504 let detail = format!(
505 "mixed multi-target drop of this object type cannot be safely rewritten; nothing executed (targets: [{kept}], absent: [{absent}])"
506 );
507 if let Some(j) = &journal {
508 let _ = j.transition(JournalState::Failed, Some(&detail));
509 }
510 return Err(MySqlError::DdlPreconditionChanged(detail));
511 }
512 }
513 }
514 Ok(super::ddl::DdlPreflight::MissingNoOp(missing)) => {
515 let detail = missing
516 .iter()
517 .map(|(s, t)| format!("{s}.{t}"))
518 .collect::<Vec<_>>()
519 .join(", ");
520 if let Some(j) = &journal {
521 let _ = j.transition(
522 crate::backup::journal::JournalState::Failed,
523 Some(&format!("ddl no-op: absent targets {detail}")),
524 );
525 }
526 ddl_absent = missing.clone();
527 return Ok(ExecuteResult {
528 journal_id: journal.as_ref().map(|j| j.id()),
529 ddl_no_op: true,
530 ddl_absent_targets: ddl_absent
531 .iter()
532 .map(|(s, t)| format!("{s}.{t}"))
533 .collect(),
534 ddl_executed_targets: Vec::new(),
535 warnings: vec![],
536 rows: Vec::new(),
537 fields: Vec::new(),
538 affected_rows: 0,
539 truncated: false,
540 duration_ms: start.elapsed().as_millis() as u64,
541 backup_id: None,
542 backup_row_count: 0,
543 });
544 }
545 Err(e) => {
546 if let Some(j) = &journal {
547 let _ = j.transition(
548 crate::backup::journal::JournalState::Failed,
549 Some(&e.to_string()),
550 );
551 }
552 return Err(MySqlError::DdlNotFound(e.to_string()));
553 }
554 #[allow(unreachable_patterns)]
555 Ok(_) => {}
556 }
557 }
558 let warnings: Vec<&'static str> =
559 super::ddl::protection_model_for(params.classified.category, params.classified.ast_type)
560 .warnings()
561 .to_vec();
562
563 let timeout_ms = params.policy.stmt_timeout_ms.max(1) as u64;
564 let mut conn_slot = Some(conn);
565 let work = async {
566 let mut conn = conn_slot.take().expect("conn returned by prior poll");
567 let result = run_in_transaction(
568 &mut conn,
569 ¶ms,
570 is_read,
571 start,
572 journal.as_ref(),
573 warnings.clone(),
574 ddl_absent.clone(),
575 ddl_executed.clone(),
576 ddl_effective_sql.clone(),
577 ddl_detail.clone(),
578 )
579 .await;
580 match result {
581 Ok(value) => {
582 if in_tx && let Err(e) = conn.query_drop("COMMIT").await {
583 let _ = conn.query_drop("ROLLBACK").await;
584 if let Some(j) = journal {
585 let _ = j.transition(JournalState::Uncertain, Some("COMMIT failed"));
586 }
587 (Err(MySqlError::Driver(e)), Some(conn))
588 } else {
589 if let Some(j) = journal {
590 let _ = j.transition(JournalState::MutationCommitted, None);
599 }
600 (Ok(value), Some(conn))
601 }
602 }
603 Err(e) => {
604 if in_tx {
605 let _ = conn.query_drop("ROLLBACK").await;
606 }
607 if let Some(j) = journal {
608 let _ = j.transition(JournalState::Failed, Some(&e.to_string()));
609 }
610 (Err(e), Some(conn))
611 }
612 }
613 };
614 tokio::pin!(work);
615 match tokio::time::timeout(Duration::from_millis(timeout_ms), &mut work).await {
616 Ok((result, _conn)) => result,
617 Err(_elapsed) => {
618 let kill_result = match executing_id {
627 Some(id) => tokio::time::timeout(
628 Duration::from_secs(5),
629 super::cancel::kill_query(&pool, id),
630 )
631 .await
632 .unwrap_or_else(|_| Err("kill control connection timed out".into())),
633 None => Err("no executing connection id recorded".into()),
634 };
635 let grace = tokio::time::timeout(Duration::from_secs(5), &mut work).await;
636 match (kill_result, grace) {
637 (Ok(()), Ok((Err(_interrupted), Some(mut conn)))) => {
638 if in_tx {
639 let _ = conn.query_drop("ROLLBACK").await;
640 }
641 let clean = super::cancel::verify_connection_clean(&mut conn).await;
642 match &clean {
643 Ok(true) => Err(MySqlError::Timeout(timeout_ms)),
644 Ok(false) => {
645 conn.disconnect().await.ok();
646 Err(MySqlError::Uncertain(
647 "deadline exceeded; transaction still open after rollback - connection discarded"
648 .to_string(),
649 ))
650 }
651 Err(detail) => {
652 conn.disconnect().await.ok();
653 Err(MySqlError::Uncertain(format!(
654 "deadline exceeded; verification failed ({detail}) - connection discarded"
655 )))
656 }
657 }
658 }
659 (Ok(_), Ok((Ok(_value), _conn))) => {
660 Err(MySqlError::Timeout(timeout_ms))
666 }
667 _ => {
668 Err(MySqlError::Uncertain(
672 "deadline exceeded; cancellation inconclusive - connection discarded"
673 .to_string(),
674 ))
675 }
676 }
677 }
678 }
679}
680
681#[allow(clippy::too_many_arguments)]
682async fn run_in_transaction(
683 conn: &mut mysql_async::Conn,
684 params: &MySqlExecuteParams<'_>,
685 is_read: bool,
686 start: Instant,
687 journal: Option<&crate::backup::journal::Journal<'_>>,
688 warnings_out: Vec<&'static str>,
689 ddl_absent: Vec<(String, String)>,
690 ddl_executed: Vec<String>,
691 ddl_effective_sql: Option<String>,
692 ddl_detail: Option<String>,
693) -> Result<ExecuteResult, MySqlError> {
694 use crate::backup::journal::JournalState;
695 if let Some(j) = journal {
696 let _ = j.transition(JournalState::BackupCapturing, ddl_detail.as_deref());
697 }
698 let effective_sql: &str = ddl_effective_sql.as_deref().unwrap_or(params.sql);
701 let mut backup_id: Option<i64> = None;
703 let mut backup_row_count: u64 = 0;
704 let mut pending_insert: Option<BackupSpec> = None;
705 if crate::backup::extractor::is_backup_required(params.classified.ast_type) {
706 let spec = extract_backup_spec(
707 effective_sql,
708 params.classified.ast_type,
709 crate::policy::classifier::Dialect::MySql,
710 )
711 .map_err(|e| MySqlError::BackupFailed(e.to_string()))?;
712 match &spec {
713 BackupSpec::InsertHint { .. } => pending_insert = Some(spec),
714 BackupSpec::None { .. } => {}
715 _ => {
716 if let Some(c) = capture_backup_mysql(
717 conn,
718 &spec,
719 ¶ms.connection.name,
720 params.database.or(params.connection.database.as_deref()),
721 params.policy,
722 params.audit.clone(),
723 )
724 .await?
725 {
726 backup_id = Some(c.backup_id);
727 backup_row_count = c.total_rows;
728 if let Some(j) = journal {
729 let _ = j.link_backup(c.backup_id);
730 let _ = j.transition(JournalState::BackupDurable, None);
731 }
732 }
733 }
734 }
735 }
736
737 if let Some(j) = journal {
738 let _ = j.transition(JournalState::MutationExecuting, None);
739 }
740
741 let sql_owned;
743 let sql: &str = if is_read && params.policy.stmt_timeout_ms > 0 {
744 sql_owned = crate::sql::hints::inject_max_execution_time(
745 effective_sql,
746 params.policy.stmt_timeout_ms,
747 );
748 &sql_owned
749 } else {
750 effective_sql
751 };
752
753 let cap = params.policy.row_cap;
754 let timeout = Duration::from_millis(params.policy.stmt_timeout_ms.max(1) as u64);
755 let outcome = tokio::time::timeout(timeout, collect_result(conn, sql, cap)).await;
756
757 let (rows, fields, affected, truncated, insert_id) = match outcome {
758 Ok(inner) => inner?,
759 Err(_) => return Err(MySqlError::Timeout(params.policy.stmt_timeout_ms as u64)),
760 };
761
762 if let Some(spec) = pending_insert
763 && let Some(id) = crate::backup::capture_insert_hint(
764 &spec,
765 ¶ms.connection.name,
766 params.database.or(params.connection.database.as_deref()),
767 insert_id,
768 affected,
769 params.audit.as_ref(),
770 )
771 {
772 backup_id = Some(id);
773 backup_row_count = affected;
774 }
775
776 let mut value = ExecuteResult {
777 journal_id: journal.as_ref().map(|j| j.id()),
781 ddl_no_op: false,
782 ddl_absent_targets: ddl_absent.iter().map(|(s, t)| format!("{s}.{t}")).collect(),
783 ddl_executed_targets: ddl_executed,
784 warnings: Vec::new(),
785 rows,
786 fields,
787 affected_rows: affected,
788 truncated,
789 duration_ms: start.elapsed().as_millis() as u64,
790 backup_id,
791 backup_row_count,
792 };
793 value.warnings = warnings_out;
794 Ok(value)
795}
796
797async fn collect_result(
800 conn: &mut mysql_async::Conn,
801 sql: &str,
802 row_cap: u32,
803) -> Result<(Vec<serde_json::Value>, Vec<String>, u64, bool, Option<i64>), MySqlError> {
804 use mysql_async::prelude::*;
805 let mut result = conn.query_iter(sql).await?;
806 let columns: Vec<String> = result
807 .columns()
808 .as_ref()
809 .map(|cols| cols.iter().map(|c| c.name_str().to_string()).collect())
810 .unwrap_or_default();
811 let column_types: Vec<mysql_async::consts::ColumnType> = result
812 .columns()
813 .as_ref()
814 .map(|cols| cols.iter().map(|c| c.column_type()).collect())
815 .unwrap_or_default();
816 let column_charsets: Vec<u16> = result
817 .columns()
818 .as_ref()
819 .map(|cols| cols.iter().map(|c| c.character_set()).collect())
820 .unwrap_or_default();
821 let mut rows_out: Vec<serde_json::Value> = Vec::new();
822 let mut bytes: u64 = 0;
823 let mut truncated = false;
824 let stream = result.stream::<Row>().await?;
825 let is_reader = stream.is_some();
826 if !is_reader {
827 drop(stream);
828 let affected = conn.affected_rows();
829 let insert_id = conn.last_insert_id().map(|id| id as i64);
830 return Ok((Vec::new(), Vec::new(), affected, false, insert_id));
831 }
832 let mut stream = stream.expect("checked Some above");
833 while let Some(row) = stream.next().await {
834 let row = row?;
835 if rows_out.len() >= row_cap as usize {
836 truncated = true;
837 break;
838 }
839 let mut obj = serde_json::Map::with_capacity(columns.len());
840 for (i, name) in columns.iter().enumerate() {
841 let v = row.as_ref(i).cloned().unwrap_or(Value::NULL);
842 let ct = column_types
843 .get(i)
844 .copied()
845 .unwrap_or(mysql_async::consts::ColumnType::MYSQL_TYPE_VAR_STRING);
846 let cs = column_charsets.get(i).copied().unwrap_or(255);
847 let j = value_with_column_type(&v, ct, cs);
848 bytes += name.len() as u64 + j.to_string().len() as u64;
849 obj.insert(name.clone(), j);
850 }
851 rows_out.push(serde_json::Value::Object(obj));
852 if bytes > RESULT_BYTE_CAP {
853 truncated = true;
854 break;
855 }
856 }
857 drop(stream);
858 let affected = conn.affected_rows();
859 let insert_id = conn.last_insert_id().map(|id| id as i64);
860 Ok((rows_out, columns, affected, truncated, insert_id))
861}
862
863pub async fn capture_backup_mysql(
865 conn: &mut mysql_async::Conn,
866 spec: &BackupSpec,
867 connection_name: &str,
868 database: Option<&str>,
869 policy: &Policy,
870 audit: Option<Arc<AuditDb>>,
871) -> Result<Option<crate::backup::CapturedBackup>, MySqlError> {
872 use mysql_async::prelude::*;
873 let ts = time_iso();
874 let tables = match spec {
875 BackupSpec::None { .. } | BackupSpec::InsertHint { .. } => return Ok(None),
876 BackupSpec::Rows { tables } | BackupSpec::Combined { tables } => tables,
877 BackupSpec::Schema { tables } => {
878 let mut first = None;
879 for t in tables {
880 let schema = show_create_table(conn, &t.db, &t.table).await;
881 let id = crate::backup::insert_schema_backup_row(
882 &audit,
883 &ts,
884 connection_name,
885 database,
886 &t.table,
887 schema.as_deref(),
888 )
889 .map_err(|e| MySqlError::BackupFailed(e.to_string()))?;
890 if first.is_none() {
891 first = Some(id);
892 }
893 }
894 return Ok(first.map(|id| crate::backup::CapturedBackup {
895 backup_id: id,
896 total_rows: 0,
897 truncated: false,
898 total_bytes: 0,
899 }));
900 }
901 };
902
903 let row_cap = policy.max_backup_rows as usize;
904 let byte_cap = policy.max_backup_bytes;
905 let mut first_id: Option<i64> = None;
906 let mut total_rows = 0u64;
907 let mut total_bytes = 0u64;
908 let mut truncated_any = false;
909
910 for t in tables {
911 let capped = crate::backup::extractor::with_limit(&t.select_sql, (row_cap + 1) as u64)
912 .ok_or_else(|| {
913 MySqlError::BackupOverflow(format!(
914 "backup query not safely limitable: {}",
915 t.select_sql.chars().take(80).collect::<String>()
916 ))
917 })?;
918 let mut stream = conn
922 .query_iter(capped.as_str())
923 .await
924 .map_err(|e| MySqlError::BackupFailed(e.to_string()))?;
925 let mut rows: Vec<serde_json::Value> = Vec::new();
926 let cols: Vec<String> = stream
927 .columns()
928 .as_ref()
929 .map(|cs| cs.iter().map(|c| c.name_str().to_string()).collect())
930 .unwrap_or_default();
931 let mut srows = match stream.stream::<Row>().await {
932 Ok(Some(s)) => s,
933 Ok(None) => continue,
934 Err(e) => return Err(MySqlError::BackupFailed(e.to_string())),
935 };
936 let mut truncated = false;
937 while let Some(row) = srows.next().await {
938 let row = row.map_err(|e| MySqlError::BackupFailed(e.to_string()))?;
939 if rows.len() >= row_cap {
940 truncated = true;
941 break;
942 }
943 let mut obj = serde_json::Map::with_capacity(cols.len());
944 for (i, name) in cols.iter().enumerate() {
945 let v = row.as_ref(i).cloned().unwrap_or(Value::NULL);
946 obj.insert(name.clone(), value_to_json(&v));
947 }
948 rows.push(serde_json::Value::Object(obj));
949 }
950 drop(srows);
951 let json = if rows.is_empty() {
952 None
953 } else {
954 Some(serde_json::to_string(&rows).unwrap_or_default())
955 };
956 let bytes = json.as_ref().map(|j| j.len() as u64).unwrap_or(0);
957 if truncated
958 && matches!(
959 policy.on_backup_overflow,
960 crate::policy::model::BackupOverflow::Abort
961 )
962 {
963 return Err(MySqlError::BackupOverflow(format!(
964 "row cap exceeded ({})",
965 row_cap
966 )));
967 }
968 if bytes > byte_cap
969 && matches!(
970 policy.on_backup_overflow,
971 crate::policy::model::BackupOverflow::Abort
972 )
973 {
974 return Err(MySqlError::BackupOverflow(format!(
975 "byte cap exceeded ({bytes} > {byte_cap})"
976 )));
977 }
978 let schema = if matches!(spec, BackupSpec::Combined { .. }) {
979 show_create_table(conn, &t.db, &t.table).await
980 } else {
981 None
982 };
983 let kind = match spec {
984 BackupSpec::Combined { .. } => "combined",
985 _ => "rows",
986 };
987 let count = rows.len() as u64;
988 let id = crate::backup::insert_rows_backup_row(
989 &audit,
990 &ts,
991 connection_name,
992 database.or(t.db.as_deref()),
993 &t.table,
994 kind,
995 json.as_deref(),
996 schema.as_deref(),
997 count,
998 truncated,
999 bytes,
1000 )
1001 .map_err(|e| MySqlError::BackupFailed(e.to_string()))?;
1002 if first_id.is_none() {
1003 first_id = Some(id);
1004 }
1005 total_rows += count;
1006 total_bytes += bytes;
1007 truncated_any |= truncated;
1008 }
1009
1010 Ok(first_id.map(|id| crate::backup::CapturedBackup {
1011 backup_id: id,
1012 total_rows,
1013 truncated: truncated_any,
1014 total_bytes,
1015 }))
1016}
1017
1018async fn show_create_table(
1019 conn: &mut mysql_async::Conn,
1020 db: &Option<String>,
1021 table: &str,
1022) -> Option<String> {
1023 use mysql_async::prelude::*;
1024 let target = match db {
1025 Some(db) => format!("`{}`.`{}`", db.replace('`', "``"), table.replace('`', "``")),
1026 None => format!("`{}`", table.replace('`', "``")),
1027 };
1028 let sql = format!("SHOW CREATE TABLE {target}");
1029 let row: Option<Row> = conn.exec_first(sql, ()).await.ok()?;
1030 let row = row?;
1031 let create: Option<String> = row.get(1);
1032 create
1033}
1034
1035fn time_iso() -> String {
1036 time::OffsetDateTime::now_utc()
1037 .format(&time::format_description::well_known::Rfc3339)
1038 .unwrap_or_else(|_| "1970-01-01T00:00:00Z".into())
1039}
1040
1041#[cfg(test)]
1042mod tests {
1043 use super::*;
1044
1045 fn conn() -> MySqlConnection {
1046 MySqlConnection {
1047 name: "t".into(),
1048 host: "db.example.invalid".into(),
1049 port: 3306,
1050 user: "u".into(),
1051 ..MySqlConnection::default()
1052 }
1053 }
1054
1055 #[test]
1056 fn opts_carry_tls_override_and_bounds() {
1057 let mut c = conn();
1058 c.ssl = true;
1059 c.ssl_server_name = Some("db.prod.example.invalid".into());
1060 let opts = build_opts(&c, "pw", Some("app"), None, None);
1061 let _ = opts;
1065 }
1066
1067 #[test]
1068 fn numeric_fidelity_mapping() {
1069 assert_eq!(
1070 value_to_json(&Value::Bytes(b"12345678901234567890".to_vec())),
1071 serde_json::json!("12345678901234567890")
1072 );
1073 assert_eq!(value_to_json(&Value::Int(-5)), serde_json::json!(-5));
1074 assert_eq!(
1075 value_to_json(&Value::UInt(u64::MAX)),
1076 serde_json::json!(u64::MAX)
1077 );
1078 assert_eq!(
1079 value_to_json(&Value::Date(2026, 8, 21, 15, 4, 5, 0)),
1080 serde_json::json!("2026-08-21 15:04:05")
1081 );
1082 assert_eq!(value_to_json(&Value::NULL), serde_json::Value::Null);
1083 }
1084
1085 #[tokio::test]
1086 async fn pool_init_failure_is_not_cached() {
1087 use zeroize::Zeroizing;
1090 let mgr = pool_manager();
1091 mgr.invalidate_all();
1092 let c = conn();
1093 let pw = Zeroizing::new("pw".to_string());
1094 assert!(
1095 mgr.verified_pool(&c, &pw, None, 1, None, Some(1), None)
1096 .await
1097 .is_err()
1098 );
1099 assert_eq!(mgr.pool_count(), 0);
1100 mgr.invalidate_all();
1101 }
1102}