From 7486e2c7b3a74c2cdd1944dc6079dc58820228b1 Mon Sep 17 00:00:00 2001 From: Illuminated2020 <2357303264@qq.com> Date: Sun, 10 May 2026 18:59:00 +0800 Subject: [PATCH] =?UTF-8?q?feat(dbx-cli):=20=E5=AE=9E=E7=8E=B0=20headless?= =?UTF-8?q?=20schema=20=E4=B8=8E=20safe-query?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 调用 dbx_core::schema_snapshot 生成 headless schema snapshot - safe-query 从 storage 加载连接并按风险阻断非 READ SQL - 补充 SQLite fixture 覆盖 READ、limit 与阻断错误 --- crates/dbx-cli/src/commands.rs | 250 +++++++++++++++++++++++++++++++-- 1 file changed, 239 insertions(+), 11 deletions(-) diff --git a/crates/dbx-cli/src/commands.rs b/crates/dbx-cli/src/commands.rs index dbf54b30a..f44b5a8f9 100644 --- a/crates/dbx-cli/src/commands.rs +++ b/crates/dbx-cli/src/commands.rs @@ -103,7 +103,7 @@ fn flag_spec(command: Option<&str>, subcommand: Option<&str>, flag: &str) -> Opt match (command, subcommand, flag) { (Some("conn"), Some("show"), "--redacted") => Some(boolean_flag), (Some("schema"), Some("snapshot"), "--conn" | "--db") => Some(normal_value), - (Some("safe-query"), _, "--conn" | "--db") => Some(normal_value), + (Some("safe-query"), _, "--conn" | "--db" | "--limit") => Some(normal_value), (Some("safe-query"), _, "--sql") => Some(free_text_value), (Some("handoff"), _, "--conn" | "--sql-file") => Some(normal_value), (Some("handoff"), _, "--title" | "--sql" | "--description") => Some(free_text_value), @@ -223,6 +223,29 @@ async fn find_connection( } } +async fn state_with_connection( + config: dbx_core::models::connection::ConnectionConfig, +) -> Result { + let state = open_state().await?; + state.configs.lock().await.insert(config.id.clone(), config.clone()); + + match config.db_type { + dbx_core::models::connection::DatabaseType::Sqlite => { + let pool = dbx_core::db::sqlite::connect_path(&config.host).await?; + state.connections.lock().await.insert(config.id.clone(), dbx_core::connection::PoolKind::Sqlite(pool)); + } + dbx_core::models::connection::DatabaseType::DuckDb => { + let pool = dbx_core::db::duckdb_driver::connect_path(&config.host)?; + state.connections.lock().await.insert(config.id.clone(), dbx_core::connection::PoolKind::DuckDb(pool)); + } + _ => { + state.get_or_create_pool(&config.id, config.database.as_deref()).await?; + } + } + + Ok(state) +} + async fn context(args: &[String]) -> CliEnvelope { if let Err(err) = reject_unexpected_positionals(args, "context", None) { return err; @@ -279,16 +302,22 @@ async fn schema_snapshot(args: &[String]) -> CliEnvelope { return err; } - let Some(_conn) = option_value(args, "--conn") else { + let Some(conn) = option_value(args, "--conn") else { return fail(CliSource::Headless, CliErrorCode::ConnectionNotFound, "--conn is required", true); }; + let config = match find_connection(conn).await { + Ok(config) => config, + Err(err) => return err, + }; + let state = match state_with_connection(config.clone()).await { + Ok(state) => state, + Err(err) => return fail(CliSource::Headless, CliErrorCode::InternalError, err, false), + }; - fail( - CliSource::Headless, - CliErrorCode::InternalError, - "Schema snapshot headless execution is not implemented", - true, - ) + match dbx_core::schema_snapshot::snapshot(&state, &config.id, option_value(args, "--db"), None).await { + Ok(snapshot) => ok(CliSource::Headless, serde_json::to_value(snapshot).unwrap()), + Err(err) => fail(CliSource::Headless, CliErrorCode::InternalError, err, false), + } } async fn safe_query(args: &[String]) -> CliEnvelope { @@ -296,14 +325,64 @@ async fn safe_query(args: &[String]) -> CliEnvelope { return err; } - let Some(_conn) = option_value(args, "--conn") else { + let Some(conn) = option_value(args, "--conn") else { return fail(CliSource::Headless, CliErrorCode::ConnectionNotFound, "--conn is required", true); }; - let Some(_sql) = option_value(args, "--sql") else { + let Some(sql) = option_value(args, "--sql") else { return fail(CliSource::Headless, CliErrorCode::QueryClassificationFailed, "--sql is required", true); }; + let limit = match parse_result_limit(args) { + Ok(limit) => limit as usize, + Err(err) => return err, + }; + let config = match find_connection(conn).await { + Ok(config) => config, + Err(err) => return err, + }; + let risk = dbx_core::sql_safety::risk_for_connection(sql, &config.name, config.color.as_deref()); + match risk.operation_class { + dbx_core::sql_safety::OperationClass::Read => {} + dbx_core::sql_safety::OperationClass::Write if risk.is_production => { + return blocked_query(CliErrorCode::ProductionWriteBlocked, &risk); + } + dbx_core::sql_safety::OperationClass::Write => { + return blocked_query(CliErrorCode::HandoffRequired, &risk); + } + dbx_core::sql_safety::OperationClass::Ddl => { + return blocked_query(CliErrorCode::DdlBlocked, &risk); + } + dbx_core::sql_safety::OperationClass::Unknown => { + return blocked_query(CliErrorCode::QueryClassificationFailed, &risk); + } + } + let state = match state_with_connection(config.clone()).await { + Ok(state) => state, + Err(err) => return fail(CliSource::Headless, CliErrorCode::InternalError, err, false), + }; + let database = option_value(args, "--db").or(config.database.as_deref()).unwrap_or(""); - fail(CliSource::Headless, CliErrorCode::InternalError, "Safe query headless execution is not implemented", true) + match dbx_core::query::execute_sql_statement(&state, &config.id, database, sql, None, None).await { + Ok(result) => ok( + CliSource::Headless, + serde_json::json!({ + "risk": risk, + "result": limit_query_result(result, limit), + }), + ), + Err(err) => fail(CliSource::Headless, CliErrorCode::InternalError, err, false), + } +} + +fn blocked_query(code: CliErrorCode, risk: &dbx_core::sql_safety::RiskMetadata) -> CliEnvelope { + fail(CliSource::Headless, code, serde_json::to_string(risk).unwrap(), true) +} + +fn limit_query_result(mut result: dbx_core::db::QueryResult, limit: usize) -> dbx_core::db::QueryResult { + if result.rows.len() > limit { + result.rows.truncate(limit); + result.truncated = true; + } + result } async fn handoff(args: &[String]) -> CliEnvelope { @@ -492,6 +571,30 @@ mod tests { storage.save_connections(configs).await.unwrap(); } + async fn create_sqlite_fixture(path: &std::path::Path) { + std::fs::File::create(path).unwrap(); + let pool = dbx_core::db::sqlite::connect_path(&path.display().to_string()).await.unwrap(); + dbx_core::db::sqlite::execute_query(&pool, "CREATE TABLE teams (id INTEGER PRIMARY KEY, name TEXT NOT NULL)") + .await + .unwrap(); + dbx_core::db::sqlite::execute_query( + &pool, + "CREATE TABLE users (id INTEGER PRIMARY KEY, team_id INTEGER NOT NULL, email TEXT NOT NULL, FOREIGN KEY(team_id) REFERENCES teams(id))", + ) + .await + .unwrap(); + dbx_core::db::sqlite::execute_query(&pool, "INSERT INTO teams (id, name) VALUES (1, 'Core'), (2, 'Data')") + .await + .unwrap(); + dbx_core::db::sqlite::execute_query( + &pool, + "INSERT INTO users (id, team_id, email) VALUES (1, 1, 'ada@example.com'), (2, 2, 'grace@example.com')", + ) + .await + .unwrap(); + pool.close().await; + } + fn success_data(env: CliEnvelope) -> serde_json::Value { match env { CliEnvelope::Success { source, data, .. } => { @@ -782,6 +885,131 @@ mod tests { std::env::remove_var("DBX_APP_DATA_DIR"); } + #[tokio::test] + async fn schema_snapshot_executes_headless_sqlite_snapshot_from_storage() { + let _guard = ENV_LOCK.lock().unwrap(); + let dir = tempfile::tempdir().unwrap(); + let data_path = dir.path().join("fixture.sqlite"); + create_sqlite_fixture(&data_path).await; + std::env::set_var("DBX_APP_DATA_DIR", dir.path()); + seed_connections( + dir.path(), + &[embedded_fixture("sqlite-id", "local-sqlite", DatabaseType::Sqlite, &data_path)], + ) + .await; + + let data = success_data( + dispatch(vec![ + "schema".into(), + "snapshot".into(), + "--conn".into(), + "local-sqlite".into(), + "--format".into(), + "json".into(), + ]) + .await, + ); + + assert_eq!(data["connectionId"], "sqlite-id"); + assert_eq!(data["database"], "main"); + let tables = data["tables"].as_array().expect("tables should be an array"); + assert!(tables.iter().any(|table| table["name"] == "users")); + + std::env::remove_var("DBX_APP_DATA_DIR"); + } + + #[tokio::test] + async fn safe_query_executes_read_sqlite_query_with_limit_and_risk() { + let _guard = ENV_LOCK.lock().unwrap(); + let dir = tempfile::tempdir().unwrap(); + let data_path = dir.path().join("fixture.sqlite"); + create_sqlite_fixture(&data_path).await; + std::env::set_var("DBX_APP_DATA_DIR", dir.path()); + seed_connections( + dir.path(), + &[embedded_fixture("sqlite-id", "local-sqlite", DatabaseType::Sqlite, &data_path)], + ) + .await; + + let data = success_data( + dispatch(vec![ + "safe-query".into(), + "--conn".into(), + "sqlite-id".into(), + "--sql".into(), + "SELECT email FROM users ORDER BY id".into(), + "--limit".into(), + "1".into(), + "--format".into(), + "json".into(), + ]) + .await, + ); + + assert_eq!(data["risk"]["operationClass"], "read"); + assert_eq!(data["result"]["columns"], serde_json::json!(["email"])); + assert_eq!(data["result"]["rows"], serde_json::json!([["ada@example.com"]])); + assert_eq!(data["result"]["truncated"], true); + + std::env::remove_var("DBX_APP_DATA_DIR"); + } + + #[tokio::test] + async fn safe_query_blocks_non_read_sql_with_structured_risk_errors() { + let _guard = ENV_LOCK.lock().unwrap(); + let dir = tempfile::tempdir().unwrap(); + let data_path = dir.path().join("fixture.sqlite"); + create_sqlite_fixture(&data_path).await; + std::env::set_var("DBX_APP_DATA_DIR", dir.path()); + let mut prod = embedded_fixture("prod-sqlite", "prod-sqlite", DatabaseType::Sqlite, &data_path); + prod.color = Some("#ef4444".to_string()); + let dev = embedded_fixture("dev-sqlite", "dev-sqlite", DatabaseType::Sqlite, &data_path); + seed_connections(dir.path(), &[prod, dev]).await; + + let cases = [ + ( + "prod-sqlite", + "UPDATE users SET email = 'x@example.com' WHERE id = 1", + CliErrorCode::ProductionWriteBlocked, + "write", + true, + ), + ( + "dev-sqlite", + "UPDATE users SET email = 'x@example.com' WHERE id = 1", + CliErrorCode::HandoffRequired, + "write", + false, + ), + ("prod-sqlite", "DROP TABLE users", CliErrorCode::DdlBlocked, "ddl", true), + ("prod-sqlite", "VACUUM", CliErrorCode::QueryClassificationFailed, "unknown", true), + ]; + + for (conn, sql, expected_code, expected_class, expected_production) in cases { + match dispatch(vec![ + "safe-query".into(), + "--conn".into(), + conn.into(), + "--sql".into(), + sql.into(), + "--format".into(), + "json".into(), + ]) + .await + { + CliEnvelope::Failure { error, .. } => { + assert_eq!(error.code, expected_code); + let risk: serde_json::Value = serde_json::from_str(&error.message).unwrap(); + assert_eq!(risk["operationClass"], expected_class); + assert_eq!(risk["isProduction"], expected_production); + } + CliEnvelope::Success { .. } => panic!("expected safe-query to block {sql}"), + } + } + + std::env::remove_var("DBX_APP_DATA_DIR"); + } + #[tokio::test] async fn unknown_command_returns_internal_error_envelope() { assert_failure_code(