chore: add rustfmt.toml and apply unified code formatting

This commit is contained in:
t8y2
2026-05-05 14:16:18 +08:00
parent edcdd652ae
commit 41445ff474
53 changed files with 1177 additions and 2592 deletions
+28 -94
View File
@@ -12,15 +12,11 @@ use tokio::sync::RwLock;
// Stream cancel registry // Stream cancel registry
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
static AI_STREAMS: LazyLock<RwLock<HashMap<String, Arc<AtomicBool>>>> = static AI_STREAMS: LazyLock<RwLock<HashMap<String, Arc<AtomicBool>>>> = LazyLock::new(|| RwLock::new(HashMap::new()));
LazyLock::new(|| RwLock::new(HashMap::new()));
pub async fn register_stream(session_id: &str) -> Arc<AtomicBool> { pub async fn register_stream(session_id: &str) -> Arc<AtomicBool> {
let cancelled = Arc::new(AtomicBool::new(false)); let cancelled = Arc::new(AtomicBool::new(false));
AI_STREAMS AI_STREAMS.write().await.insert(session_id.to_string(), cancelled.clone());
.write()
.await
.insert(session_id.to_string(), cancelled.clone());
cancelled cancelled
} }
@@ -129,10 +125,7 @@ pub struct AiConversation {
pub fn resolve_endpoint(config: &AiConfig) -> String { pub fn resolve_endpoint(config: &AiConfig) -> String {
let ep = config.endpoint.trim().trim_end_matches('/'); let ep = config.endpoint.trim().trim_end_matches('/');
if ep.ends_with("/chat/completions") if ep.ends_with("/chat/completions") || ep.ends_with("/responses") || ep.ends_with("/messages") {
|| ep.ends_with("/responses")
|| ep.ends_with("/messages")
{
return ep.to_string(); return ep.to_string();
} }
match config.provider { match config.provider {
@@ -149,11 +142,7 @@ pub fn resolve_endpoint(config: &AiConfig) -> String {
pub fn stream_data_payload(line: &str) -> Option<&str> { pub fn stream_data_payload(line: &str) -> Option<&str> {
let line = line.trim(); let line = line.trim();
if line.is_empty() if line.is_empty() || line.starts_with(':') || line.starts_with("event:") || line.starts_with("id:") {
|| line.starts_with(':')
|| line.starts_with("event:")
|| line.starts_with("id:")
{
return None; return None;
} }
if let Some(data) = line.strip_prefix("data:") { if let Some(data) = line.strip_prefix("data:") {
@@ -175,11 +164,7 @@ pub fn claude_stream_text(event: &serde_json::Value) -> Option<&str> {
pub fn openai_stream_text(event: &serde_json::Value) -> Option<&str> { pub fn openai_stream_text(event: &serde_json::Value) -> Option<&str> {
event["choices"] event["choices"]
.get(0) .get(0)
.and_then(|choice| { .and_then(|choice| choice["delta"]["content"].as_str().or_else(|| choice["message"]["content"].as_str()))
choice["delta"]["content"]
.as_str()
.or_else(|| choice["message"]["content"].as_str())
})
.or_else(|| event["content"].as_str()) .or_else(|| event["content"].as_str())
.filter(|text| !text.is_empty()) .filter(|text| !text.is_empty())
} }
@@ -196,10 +181,7 @@ pub fn responses_stream_text(event: &serde_json::Value) -> Option<&str> {
} }
pub fn extract_error(data: &serde_json::Value) -> Option<String> { pub fn extract_error(data: &serde_json::Value) -> Option<String> {
data["error"]["message"] data["error"]["message"].as_str().or_else(|| data["error"].as_str()).map(ToString::to_string)
.as_str()
.or_else(|| data["error"].as_str())
.map(ToString::to_string)
} }
pub fn build_responses_input(system_prompt: &str, messages: &[AiMessage]) -> serde_json::Value { pub fn build_responses_input(system_prompt: &str, messages: &[AiMessage]) -> serde_json::Value {
@@ -240,16 +222,10 @@ fn validate_config(config: &AiConfig) -> Result<(), String> {
// Non-streaming calls // Non-streaming calls
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
pub async fn call_claude( pub async fn call_claude(client: &reqwest::Client, request: AiCompletionRequest) -> Result<String, String> {
client: &reqwest::Client,
request: AiCompletionRequest,
) -> Result<String, String> {
let mut headers = HeaderMap::new(); let mut headers = HeaderMap::new();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert( headers.insert("x-api-key", HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?);
"x-api-key",
HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?,
);
headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01")); headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
let body = json!({ let body = json!({
@@ -281,25 +257,16 @@ pub async fn call_claude(
.to_string()) .to_string())
} }
pub async fn call_openai_compatible( pub async fn call_openai_compatible(client: &reqwest::Client, request: AiCompletionRequest) -> Result<String, String> {
client: &reqwest::Client,
request: AiCompletionRequest,
) -> Result<String, String> {
let mut headers = HeaderMap::new(); let mut headers = HeaderMap::new();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert( headers.insert(
AUTHORIZATION, AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)) HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
.map_err(|e| e.to_string())?,
); );
let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })]; let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })];
messages.extend( messages.extend(request.messages.iter().map(|message| json!({ "role": message.role, "content": message.content })));
request
.messages
.iter()
.map(|message| json!({ "role": message.role, "content": message.content })),
);
let body = json!({ let body = json!({
"model": request.config.model, "model": request.config.model,
@@ -322,22 +289,15 @@ pub async fn call_openai_compatible(
return Err(extract_error(&data).unwrap_or_else(|| format!("API error: {status}"))); return Err(extract_error(&data).unwrap_or_else(|| format!("API error: {status}")));
} }
Ok(data["choices"][0]["message"]["content"] Ok(data["choices"][0]["message"]["content"].as_str().unwrap_or_default().to_string())
.as_str()
.unwrap_or_default()
.to_string())
} }
pub async fn call_responses_api( pub async fn call_responses_api(client: &reqwest::Client, request: AiCompletionRequest) -> Result<String, String> {
client: &reqwest::Client,
request: AiCompletionRequest,
) -> Result<String, String> {
let mut headers = HeaderMap::new(); let mut headers = HeaderMap::new();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert( headers.insert(
AUTHORIZATION, AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)) HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
.map_err(|e| e.to_string())?,
); );
let body = json!({ let body = json!({
@@ -365,9 +325,7 @@ pub async fn call_responses_api(
.as_array() .as_array()
.and_then(|items| { .and_then(|items| {
items.iter().find_map(|item| { items.iter().find_map(|item| {
item["content"] item["content"].as_array().and_then(|parts| parts.iter().find_map(|p| p["text"].as_str()))
.as_array()
.and_then(|parts| parts.iter().find_map(|p| p["text"].as_str()))
}) })
}) })
.unwrap_or_default() .unwrap_or_default()
@@ -381,18 +339,13 @@ pub async fn call_responses_api(
pub async fn test_connection_core(config: &AiConfig) -> Result<String, String> { pub async fn test_connection_core(config: &AiConfig) -> Result<String, String> {
validate_config(config)?; validate_config(config)?;
let client = reqwest::Client::builder() let client =
.timeout(std::time::Duration::from_secs(15)) reqwest::Client::builder().timeout(std::time::Duration::from_secs(15)).build().map_err(|e| e.to_string())?;
.build()
.map_err(|e| e.to_string())?;
let request = AiCompletionRequest { let request = AiCompletionRequest {
config: config.clone(), config: config.clone(),
system_prompt: String::new(), system_prompt: String::new(),
messages: vec![AiMessage { messages: vec![AiMessage { role: "user".into(), content: "hi".into() }],
role: "user".into(),
content: "hi".into(),
}],
max_tokens: Some(1), max_tokens: Some(1),
temperature: Some(0.0), temperature: Some(0.0),
}; };
@@ -413,10 +366,8 @@ pub async fn test_connection_core(config: &AiConfig) -> Result<String, String> {
pub async fn complete(request: &AiCompletionRequest) -> Result<String, String> { pub async fn complete(request: &AiCompletionRequest) -> Result<String, String> {
validate_config(&request.config)?; validate_config(&request.config)?;
let client = reqwest::Client::builder() let client =
.timeout(std::time::Duration::from_secs(60)) reqwest::Client::builder().timeout(std::time::Duration::from_secs(60)).build().map_err(|e| e.to_string())?;
.build()
.map_err(|e| e.to_string())?;
match request.config.provider { match request.config.provider {
AiProvider::Claude => call_claude(&client, request.clone()).await, AiProvider::Claude => call_claude(&client, request.clone()).await,
@@ -442,15 +393,11 @@ pub async fn stream(
) -> Result<(), String> { ) -> Result<(), String> {
validate_config(&request.config)?; validate_config(&request.config)?;
let client = reqwest::Client::builder() let client =
.timeout(std::time::Duration::from_secs(120)) reqwest::Client::builder().timeout(std::time::Duration::from_secs(120)).build().map_err(|e| e.to_string())?;
.build()
.map_err(|e| e.to_string())?;
match request.config.provider { match request.config.provider {
AiProvider::Claude => { AiProvider::Claude => stream_claude(&client, session_id, request, cancelled, &on_chunk).await,
stream_claude(&client, session_id, request, cancelled, &on_chunk).await
}
AiProvider::Openai | AiProvider::Custom => { AiProvider::Openai | AiProvider::Custom => {
if request.config.api_style == AiApiStyle::Responses { if request.config.api_style == AiApiStyle::Responses {
stream_responses_api(&client, session_id, request, cancelled, &on_chunk).await stream_responses_api(&client, session_id, request, cancelled, &on_chunk).await
@@ -470,10 +417,7 @@ async fn stream_claude(
) -> Result<(), String> { ) -> Result<(), String> {
let mut headers = HeaderMap::new(); let mut headers = HeaderMap::new();
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert( headers.insert("x-api-key", HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?);
"x-api-key",
HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?,
);
headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01")); headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
let body = json!({ let body = json!({
@@ -559,17 +503,11 @@ async fn stream_openai(
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert( headers.insert(
AUTHORIZATION, AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)) HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
.map_err(|e| e.to_string())?,
); );
let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })]; let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })];
messages.extend( messages.extend(request.messages.iter().map(|m| json!({ "role": m.role, "content": m.content })));
request
.messages
.iter()
.map(|m| json!({ "role": m.role, "content": m.content })),
);
let body = json!({ let body = json!({
"model": request.config.model, "model": request.config.model,
@@ -661,8 +599,7 @@ async fn stream_responses_api(
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
headers.insert( headers.insert(
AUTHORIZATION, AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)) HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
.map_err(|e| e.to_string())?,
); );
let body = json!({ let body = json!({
@@ -771,10 +708,7 @@ pub fn load_conversations(path: &Path) -> Result<Vec<AiConversation>, String> {
} }
pub fn delete_conversation(path: &Path, id: &str) -> Result<(), String> { pub fn delete_conversation(path: &Path, id: &str) -> Result<(), String> {
let conversations: Vec<AiConversation> = read_conversations(path)? let conversations: Vec<AiConversation> = read_conversations(path)?.into_iter().filter(|c| c.id != id).collect();
.into_iter()
.filter(|c| c.id != id)
.collect();
write_conversations(path, &conversations) write_conversations(path, &conversations)
} }
+15 -62
View File
@@ -49,20 +49,13 @@ impl AppState {
} }
} }
pub async fn get_or_create_pool( pub async fn get_or_create_pool(&self, connection_id: &str, database: Option<&str>) -> Result<String, String> {
&self,
connection_id: &str,
database: Option<&str>,
) -> Result<String, String> {
let db_type = { let db_type = {
let configs = self.configs.lock().await; let configs = self.configs.lock().await;
configs.get(connection_id).map(|c| c.db_type.clone()) configs.get(connection_id).map(|c| c.db_type.clone())
}; };
let is_embedded = matches!( let is_embedded = matches!(db_type, Some(DatabaseType::Sqlite) | Some(DatabaseType::DuckDb));
db_type,
Some(DatabaseType::Sqlite) | Some(DatabaseType::DuckDb)
);
if is_embedded { if is_embedded {
return Ok(connection_id.to_string()); return Ok(connection_id.to_string());
} }
@@ -84,10 +77,7 @@ impl AppState {
drop(conns); drop(conns);
let configs = self.configs.lock().await; let configs = self.configs.lock().await;
let config = configs let config = configs.get(connection_id).ok_or("Connection config not found")?.clone();
.get(connection_id)
.ok_or("Connection config not found")?
.clone();
drop(configs); drop(configs);
let mut db_config = config.clone(); let mut db_config = config.clone();
@@ -108,19 +98,14 @@ impl AppState {
DatabaseType::Doris | DatabaseType::StarRocks => { DatabaseType::Doris | DatabaseType::StarRocks => {
PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true) PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true)
} }
DatabaseType::Postgres | DatabaseType::Redshift => { DatabaseType::Postgres | DatabaseType::Redshift => PoolKind::Postgres(db::postgres::connect(&url).await?),
PoolKind::Postgres(db::postgres::connect(&url).await?) DatabaseType::Sqlite => PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&db_config.host)).await?),
}
DatabaseType::Sqlite => {
PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&db_config.host)).await?)
}
DatabaseType::Redis => { DatabaseType::Redis => {
let con = db::redis_driver::connect(&url).await?; let con = db::redis_driver::connect(&url).await?;
PoolKind::Redis(tokio::sync::Mutex::new(con)) PoolKind::Redis(tokio::sync::Mutex::new(con))
} }
DatabaseType::DuckDb => { DatabaseType::DuckDb => {
let con = duckdb::Connection::open(&expand_tilde(&db_config.host)) let con = duckdb::Connection::open(&expand_tilde(&db_config.host)).map_err(|e| e.to_string())?;
.map_err(|e| e.to_string())?;
PoolKind::DuckDb(Arc::new(std::sync::Mutex::new(con))) PoolKind::DuckDb(Arc::new(std::sync::Mutex::new(con)))
} }
DatabaseType::MongoDb => { DatabaseType::MongoDb => {
@@ -129,16 +114,8 @@ impl AppState {
PoolKind::MongoDb(client) PoolKind::MongoDb(client)
} }
DatabaseType::ClickHouse => { DatabaseType::ClickHouse => {
let username = if db_config.username.is_empty() { let username = if db_config.username.is_empty() { None } else { Some(db_config.username.clone()) };
None let password = if db_config.password.is_empty() { None } else { Some(db_config.password.clone()) };
} else {
Some(db_config.username.clone())
};
let password = if db_config.password.is_empty() {
None
} else {
Some(db_config.password.clone())
};
let client = db::clickhouse_driver::ChClient::new(&url, username, password); let client = db::clickhouse_driver::ChClient::new(&url, username, password);
db::clickhouse_driver::test_connection(&client).await?; db::clickhouse_driver::test_connection(&client).await?;
PoolKind::ClickHouse(client) PoolKind::ClickHouse(client)
@@ -166,11 +143,8 @@ impl AppState {
PoolKind::Oracle(Arc::new(tokio::sync::Mutex::new(client))) PoolKind::Oracle(Arc::new(tokio::sync::Mutex::new(client)))
} }
DatabaseType::Elasticsearch => { DatabaseType::Elasticsearch => {
let client = db::elasticsearch_driver::EsClient::new( let client =
&url, db::elasticsearch_driver::EsClient::new(&url, Some(&db_config.username), Some(&db_config.password));
Some(&db_config.username),
Some(&db_config.password),
);
db::elasticsearch_driver::test_connection(&client).await?; db::elasticsearch_driver::test_connection(&client).await?;
PoolKind::Elasticsearch(client) PoolKind::Elasticsearch(client)
} }
@@ -212,18 +186,12 @@ impl AppState {
Ok(("127.0.0.1".to_string(), local_port)) Ok(("127.0.0.1".to_string(), local_port))
} }
pub async fn reconnect_pool( pub async fn reconnect_pool(&self, connection_id: &str, database: Option<&str>) -> Result<String, String> {
&self,
connection_id: &str,
database: Option<&str>,
) -> Result<String, String> {
let is_single_conn = { let is_single_conn = {
let configs = self.configs.lock().await; let configs = self.configs.lock().await;
configs configs
.get(connection_id) .get(connection_id)
.map(|c| { .map(|c| c.db_type == DatabaseType::Oracle || c.db_type == DatabaseType::Elasticsearch)
c.db_type == DatabaseType::Oracle || c.db_type == DatabaseType::Elasticsearch
})
.unwrap_or(false) .unwrap_or(false)
}; };
let pool_key = if is_single_conn { let pool_key = if is_single_conn {
@@ -247,11 +215,7 @@ pub fn connection_url_for_endpoint(config: &ConnectionConfig, host: &str, port:
} }
} }
pub fn redacted_connection_url_for_endpoint( pub fn redacted_connection_url_for_endpoint(config: &ConnectionConfig, host: &str, port: u16) -> String {
config: &ConnectionConfig,
host: &str,
port: u16,
) -> String {
if host == config.host && port == config.port { if host == config.host && port == config.port {
config.redacted_connection_url() config.redacted_connection_url()
} else { } else {
@@ -259,21 +223,10 @@ pub fn redacted_connection_url_for_endpoint(
} }
} }
pub async fn probe_connection_endpoint( pub async fn probe_connection_endpoint(config: &ConnectionConfig, host: &str, port: u16) -> Result<(), String> {
config: &ConnectionConfig,
host: &str,
port: u16,
) -> Result<(), String> {
match config.db_type { match config.db_type {
DatabaseType::Sqlite | DatabaseType::DuckDb => Ok(()), DatabaseType::Sqlite | DatabaseType::DuckDb => Ok(()),
DatabaseType::MongoDb DatabaseType::MongoDb if config.connection_string.as_deref().is_some_and(|value| !value.is_empty()) => Ok(()),
if config
.connection_string
.as_deref()
.is_some_and(|value| !value.is_empty()) =>
{
Ok(())
}
_ => db::probe_tcp_endpoint(&format!("{:?}", config.db_type), host, port).await, _ => db::probe_tcp_endpoint(&format!("{:?}", config.db_type), host, port).await,
} }
} }
+20 -72
View File
@@ -24,10 +24,7 @@ impl FileSecretStore {
} }
fn read_store(&self) -> HashMap<String, String> { fn read_store(&self) -> HashMap<String, String> {
std::fs::read_to_string(&self.path) std::fs::read_to_string(&self.path).ok().and_then(|json| serde_json::from_str(&json).ok()).unwrap_or_default()
.ok()
.and_then(|json| serde_json::from_str(&json).ok())
.unwrap_or_default()
} }
fn write_store(&self, map: &HashMap<String, String>) -> Result<(), String> { fn write_store(&self, map: &HashMap<String, String>) -> Result<(), String> {
@@ -64,12 +61,7 @@ pub fn save_connections_to_file(
persist_secret(store, &config.id, MAIN_PASSWORD_KEY, &config.password)?; persist_secret(store, &config.id, MAIN_PASSWORD_KEY, &config.password)?;
persist_secret(store, &config.id, SSH_PASSWORD_KEY, &config.ssh_password)?; persist_secret(store, &config.id, SSH_PASSWORD_KEY, &config.ssh_password)?;
persist_secret(store, &config.id, SSH_KEY_PASSPHRASE_KEY, &config.ssh_key_passphrase)?; persist_secret(store, &config.id, SSH_KEY_PASSPHRASE_KEY, &config.ssh_key_passphrase)?;
persist_optional_secret( persist_optional_secret(store, &config.id, CONNECTION_STRING_KEY, config.connection_string.as_deref())?;
store,
&config.id,
CONNECTION_STRING_KEY,
config.connection_string.as_deref(),
)?;
} }
write_sanitized_connections(path, configs) write_sanitized_connections(path, configs)
@@ -113,11 +105,7 @@ pub fn load_connections_from_file(
needs_rewrite = true; needs_rewrite = true;
} }
match config match config.connection_string.as_deref().filter(|secret| !secret.is_empty()) {
.connection_string
.as_deref()
.filter(|secret| !secret.is_empty())
{
Some(secret) => { Some(secret) => {
store.set_secret(&config.id, CONNECTION_STRING_KEY, secret)?; store.set_secret(&config.id, CONNECTION_STRING_KEY, secret)?;
needs_rewrite = true; needs_rewrite = true;
@@ -220,8 +208,8 @@ pub fn secret_account(connection_id: &str, key: &str) -> String {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{ use super::{
load_connections_from_file, save_connections_to_file, ConnectionSecretStore, load_connections_from_file, save_connections_to_file, ConnectionSecretStore, CONNECTION_STRING_KEY,
CONNECTION_STRING_KEY, MAIN_PASSWORD_KEY, SSH_PASSWORD_KEY, MAIN_PASSWORD_KEY, SSH_PASSWORD_KEY,
}; };
use crate::models::connection::{ConnectionConfig, DatabaseType}; use crate::models::connection::{ConnectionConfig, DatabaseType};
use std::cell::RefCell; use std::cell::RefCell;
@@ -236,48 +224,31 @@ mod tests {
impl MemorySecretStore { impl MemorySecretStore {
fn set_existing(&self, connection_id: &str, key: &str, value: &str) { fn set_existing(&self, connection_id: &str, key: &str, value: &str) {
self.values self.values.borrow_mut().insert(secret_key(connection_id, key), value.to_string());
.borrow_mut()
.insert(secret_key(connection_id, key), value.to_string());
} }
fn get_existing(&self, connection_id: &str, key: &str) -> Option<String> { fn get_existing(&self, connection_id: &str, key: &str) -> Option<String> {
self.values self.values.borrow().get(&secret_key(connection_id, key)).cloned()
.borrow()
.get(&secret_key(connection_id, key))
.cloned()
} }
fn was_deleted(&self, connection_id: &str, key: &str) -> bool { fn was_deleted(&self, connection_id: &str, key: &str) -> bool {
self.deleted self.deleted.borrow().contains(&secret_key(connection_id, key))
.borrow()
.contains(&secret_key(connection_id, key))
} }
} }
impl ConnectionSecretStore for MemorySecretStore { impl ConnectionSecretStore for MemorySecretStore {
fn set_secret(&self, connection_id: &str, key: &str, secret: &str) -> Result<(), String> { fn set_secret(&self, connection_id: &str, key: &str, secret: &str) -> Result<(), String> {
self.values self.values.borrow_mut().insert(secret_key(connection_id, key), secret.to_string());
.borrow_mut()
.insert(secret_key(connection_id, key), secret.to_string());
Ok(()) Ok(())
} }
fn get_secret(&self, connection_id: &str, key: &str) -> Result<Option<String>, String> { fn get_secret(&self, connection_id: &str, key: &str) -> Result<Option<String>, String> {
Ok(self Ok(self.values.borrow().get(&secret_key(connection_id, key)).cloned())
.values
.borrow()
.get(&secret_key(connection_id, key))
.cloned())
} }
fn delete_secret(&self, connection_id: &str, key: &str) -> Result<(), String> { fn delete_secret(&self, connection_id: &str, key: &str) -> Result<(), String> {
self.values self.values.borrow_mut().remove(&secret_key(connection_id, key));
.borrow_mut() self.deleted.borrow_mut().push(secret_key(connection_id, key));
.remove(&secret_key(connection_id, key));
self.deleted
.borrow_mut()
.push(secret_key(connection_id, key));
Ok(()) Ok(())
} }
} }
@@ -287,10 +258,7 @@ mod tests {
} }
fn temp_connections_file(name: &str) -> std::path::PathBuf { fn temp_connections_file(name: &str) -> std::path::PathBuf {
let dir = std::env::temp_dir().join(format!( let dir = std::env::temp_dir().join(format!("dbx-connection-secrets-test-{}-{name}", std::process::id()));
"dbx-connection-secrets-test-{}-{name}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap(); std::fs::create_dir_all(&dir).unwrap();
dir.join("connections.json") dir.join("connections.json")
} }
@@ -335,14 +303,8 @@ mod tests {
save_connections_to_file(&path, &configs, &store).unwrap(); save_connections_to_file(&path, &configs, &store).unwrap();
assert_eq!( assert_eq!(store.get_existing("main", MAIN_PASSWORD_KEY).as_deref(), Some("db-secret"));
store.get_existing("main", MAIN_PASSWORD_KEY).as_deref(), assert_eq!(store.get_existing("main", SSH_PASSWORD_KEY).as_deref(), Some("ssh-secret"));
Some("db-secret")
);
assert_eq!(
store.get_existing("main", SSH_PASSWORD_KEY).as_deref(),
Some("ssh-secret")
);
let persisted = read_configs(&path); let persisted = read_configs(&path);
assert_eq!(persisted[0].password, ""); assert_eq!(persisted[0].password, "");
assert_eq!(persisted[0].ssh_password, ""); assert_eq!(persisted[0].ssh_password, "");
@@ -374,14 +336,8 @@ mod tests {
assert_eq!(loaded[0].password, "plain-db"); assert_eq!(loaded[0].password, "plain-db");
assert_eq!(loaded[0].ssh_password, "plain-ssh"); assert_eq!(loaded[0].ssh_password, "plain-ssh");
assert_eq!( assert_eq!(store.get_existing("legacy", MAIN_PASSWORD_KEY).as_deref(), Some("plain-db"));
store.get_existing("legacy", MAIN_PASSWORD_KEY).as_deref(), assert_eq!(store.get_existing("legacy", SSH_PASSWORD_KEY).as_deref(), Some("plain-ssh"));
Some("plain-db")
);
assert_eq!(
store.get_existing("legacy", SSH_PASSWORD_KEY).as_deref(),
Some("plain-ssh")
);
let persisted = read_configs(&path); let persisted = read_configs(&path);
assert_eq!(persisted[0].password, ""); assert_eq!(persisted[0].password, "");
assert_eq!(persisted[0].ssh_password, ""); assert_eq!(persisted[0].ssh_password, "");
@@ -401,10 +357,7 @@ mod tests {
assert!(store.was_deleted("old", MAIN_PASSWORD_KEY)); assert!(store.was_deleted("old", MAIN_PASSWORD_KEY));
assert!(store.was_deleted("old", SSH_PASSWORD_KEY)); assert!(store.was_deleted("old", SSH_PASSWORD_KEY));
assert_eq!( assert_eq!(store.get_existing("kept", MAIN_PASSWORD_KEY).as_deref(), Some("new-db"));
store.get_existing("kept", MAIN_PASSWORD_KEY).as_deref(),
Some("new-db")
);
} }
#[test] #[test]
@@ -418,18 +371,13 @@ mod tests {
save_connections_to_file(&path, &[config], &store).unwrap(); save_connections_to_file(&path, &[config], &store).unwrap();
assert_eq!( assert_eq!(
store store.get_existing("mongo", CONNECTION_STRING_KEY).as_deref(),
.get_existing("mongo", CONNECTION_STRING_KEY)
.as_deref(),
Some("mongodb://user:secret@localhost/app") Some("mongodb://user:secret@localhost/app")
); );
let persisted = read_configs(&path); let persisted = read_configs(&path);
assert_eq!(persisted[0].connection_string, None); assert_eq!(persisted[0].connection_string, None);
let loaded = load_connections_from_file(&path, &store).unwrap(); let loaded = load_connections_from_file(&path, &store).unwrap();
assert_eq!( assert_eq!(loaded[0].connection_string.as_deref(), Some("mongodb://user:secret@localhost/app"));
loaded[0].connection_string.as_deref(),
Some("mongodb://user:secret@localhost/app")
);
} }
} }
+17 -75
View File
@@ -14,16 +14,9 @@ pub struct ChClient {
impl ChClient { impl ChClient {
pub fn new(url: &str, username: Option<String>, password: Option<String>) -> Self { pub fn new(url: &str, username: Option<String>, password: Option<String>) -> Self {
let http = HttpClient::builder() let http =
.connect_timeout(connection_timeout()) HttpClient::builder().connect_timeout(connection_timeout()).build().unwrap_or_else(|_| HttpClient::new());
.build() Self { http, base_url: url.trim_end_matches('/').to_string(), username, password }
.unwrap_or_else(|_| HttpClient::new());
Self {
http,
base_url: url.trim_end_matches('/').to_string(),
username,
password,
}
} }
} }
@@ -62,55 +55,33 @@ fn build_request(client: &ChClient, req: reqwest::RequestBuilder) -> reqwest::Re
} }
} }
async fn ch_query( async fn ch_query(client: &ChClient, sql: &str, database: Option<&str>) -> Result<ChJsonResult, String> {
client: &ChClient,
sql: &str,
database: Option<&str>,
) -> Result<ChJsonResult, String> {
let mut url = format!("{}/?default_format=JSONCompact", client.base_url); let mut url = format!("{}/?default_format=JSONCompact", client.base_url);
if let Some(db) = database { if let Some(db) = database {
url.push_str(&format!("&database={}", db)); url.push_str(&format!("&database={}", db));
} }
let req = build_request(client, client.http.post(&url).body(sql.to_string())); let req = build_request(client, client.http.post(&url).body(sql.to_string()));
let resp = req let resp = req.send().await.map_err(|e| format!("ClickHouse request failed: {e}"))?;
.send()
.await
.map_err(|e| format!("ClickHouse request failed: {e}"))?;
if !resp.status().is_success() { if !resp.status().is_success() {
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
return Err(format!("ClickHouse error: {body}")); return Err(format!("ClickHouse error: {body}"));
} }
resp.json::<ChJsonResult>() resp.json::<ChJsonResult>().await.map_err(|e| format!("ClickHouse parse error: {e}"))
.await
.map_err(|e| format!("ClickHouse parse error: {e}"))
} }
pub async fn test_connection(client: &ChClient) -> Result<(), String> { pub async fn test_connection(client: &ChClient) -> Result<(), String> {
let url = format!("{}/ping", client.base_url); let url = format!("{}/ping", client.base_url);
let req = build_request(client, client.http.get(&url)); let req = build_request(client, client.http.get(&url));
with_connection_timeout("ClickHouse", async { with_connection_timeout("ClickHouse", async {
req.send() req.send().await.map_err(|e| format!("ClickHouse connection failed: {e}"))
.await
.map_err(|e| format!("ClickHouse connection failed: {e}"))
}) })
.await?; .await?;
Ok(()) Ok(())
} }
pub async fn list_databases(client: &ChClient) -> Result<Vec<DatabaseInfo>, String> { pub async fn list_databases(client: &ChClient) -> Result<Vec<DatabaseInfo>, String> {
let result = ch_query( let result = ch_query(client, "SELECT name FROM system.databases ORDER BY name", None).await?;
client, Ok(result.data.iter().map(|row| DatabaseInfo { name: row[0].as_str().unwrap_or("").to_string() }).collect())
"SELECT name FROM system.databases ORDER BY name",
None,
)
.await?;
Ok(result
.data
.iter()
.map(|row| DatabaseInfo {
name: row[0].as_str().unwrap_or("").to_string(),
})
.collect())
} }
pub async fn list_tables(client: &ChClient, database: &str) -> Result<Vec<TableInfo>, String> { pub async fn list_tables(client: &ChClient, database: &str) -> Result<Vec<TableInfo>, String> {
@@ -124,24 +95,13 @@ pub async fn list_tables(client: &ChClient, database: &str) -> Result<Vec<TableI
.iter() .iter()
.map(|row| { .map(|row| {
let engine = row.get(1).and_then(|v| v.as_str()).unwrap_or(""); let engine = row.get(1).and_then(|v| v.as_str()).unwrap_or("");
let table_type = if engine.contains("View") { let table_type = if engine.contains("View") { "VIEW" } else { "BASE TABLE" };
"VIEW" TableInfo { name: row[0].as_str().unwrap_or("").to_string(), table_type: table_type.to_string() }
} else {
"BASE TABLE"
};
TableInfo {
name: row[0].as_str().unwrap_or("").to_string(),
table_type: table_type.to_string(),
}
}) })
.collect()) .collect())
} }
pub async fn get_columns( pub async fn get_columns(client: &ChClient, database: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
client: &ChClient,
database: &str,
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
let sql = format!( let sql = format!(
"SELECT name, type, default_kind, default_expression, is_in_primary_key \ "SELECT name, type, default_kind, default_expression, is_in_primary_key \
FROM system.columns WHERE database = '{}' AND table = '{}' ORDER BY position", FROM system.columns WHERE database = '{}' AND table = '{}' ORDER BY position",
@@ -153,20 +113,12 @@ pub async fn get_columns(
.data .data
.iter() .iter()
.map(|row| { .map(|row| {
let data_type = row let data_type = row.get(1).and_then(|v| v.as_str()).unwrap_or("").to_string();
.get(1)
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let is_nullable = data_type.starts_with("Nullable"); let is_nullable = data_type.starts_with("Nullable");
let is_pk = row.get(4).and_then(|v| v.as_u64()).unwrap_or(0) == 1; let is_pk = row.get(4).and_then(|v| v.as_u64()).unwrap_or(0) == 1;
let default_kind = row.get(2).and_then(|v| v.as_str()).unwrap_or(""); let default_kind = row.get(2).and_then(|v| v.as_str()).unwrap_or("");
let default_expr = row.get(3).and_then(|v| v.as_str()).unwrap_or(""); let default_expr = row.get(3).and_then(|v| v.as_str()).unwrap_or("");
let column_default = if default_kind.is_empty() { let column_default = if default_kind.is_empty() { None } else { Some(default_expr.to_string()) };
None
} else {
Some(default_expr.to_string())
};
ColumnInfo { ColumnInfo {
name: row[0].as_str().unwrap_or("").to_string(), name: row[0].as_str().unwrap_or("").to_string(),
data_type, data_type,
@@ -183,11 +135,7 @@ pub async fn get_columns(
.collect()) .collect())
} }
pub async fn execute_query( pub async fn execute_query(client: &ChClient, database: &str, sql: &str) -> Result<QueryResult, String> {
client: &ChClient,
database: &str,
sql: &str,
) -> Result<QueryResult, String> {
let start = Instant::now(); let start = Instant::now();
let trimmed = sql.trim().to_uppercase(); let trimmed = sql.trim().to_uppercase();
@@ -207,15 +155,9 @@ pub async fn execute_query(
truncated: false, truncated: false,
}) })
} else { } else {
let url = format!( let url = format!("{}/?default_format=JSONCompact&database={}", client.base_url, database);
"{}/?default_format=JSONCompact&database={}",
client.base_url, database
);
let req = build_request(client, client.http.post(&url).body(sql.to_string())); let req = build_request(client, client.http.post(&url).body(sql.to_string()));
let resp = req let resp = req.send().await.map_err(|e| format!("ClickHouse request failed: {e}"))?;
.send()
.await
.map_err(|e| format!("ClickHouse request failed: {e}"))?;
if !resp.status().is_success() { if !resp.status().is_success() {
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
return Err(format!("ClickHouse error: {body}")); return Err(format!("ClickHouse error: {body}"));
+18 -78
View File
@@ -16,15 +16,9 @@ impl EsClient {
(Some(u), Some(p)) if !u.is_empty() => Some((u.to_string(), p.to_string())), (Some(u), Some(p)) if !u.is_empty() => Some((u.to_string(), p.to_string())),
_ => None, _ => None,
}; };
let http = HttpClient::builder() let http =
.connect_timeout(connection_timeout()) HttpClient::builder().connect_timeout(connection_timeout()).build().unwrap_or_else(|_| HttpClient::new());
.build() Self { http, base_url: url.trim_end_matches('/').to_string(), auth }
.unwrap_or_else(|_| HttpClient::new());
Self {
http,
base_url: url.trim_end_matches('/').to_string(),
auth,
}
} }
fn get(&self, path: &str) -> reqwest::RequestBuilder { fn get(&self, path: &str) -> reqwest::RequestBuilder {
@@ -58,21 +52,13 @@ impl EsClient {
impl Clone for EsClient { impl Clone for EsClient {
fn clone(&self) -> Self { fn clone(&self) -> Self {
Self { Self { http: self.http.clone(), base_url: self.base_url.clone(), auth: self.auth.clone() }
http: self.http.clone(),
base_url: self.base_url.clone(),
auth: self.auth.clone(),
}
} }
} }
pub async fn test_connection(client: &EsClient) -> Result<(), String> { pub async fn test_connection(client: &EsClient) -> Result<(), String> {
let resp = with_connection_timeout("Elasticsearch", async { let resp = with_connection_timeout("Elasticsearch", async {
client client.get("/").send().await.map_err(|e| format!("Elasticsearch connection failed: {e}"))
.get("/")
.send()
.await
.map_err(|e| format!("Elasticsearch connection failed: {e}"))
}) })
.await?; .await?;
if !resp.status().is_success() { if !resp.status().is_success() {
@@ -97,15 +83,8 @@ pub async fn list_indices(client: &EsClient) -> Result<Vec<String>, String> {
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
return Err(format!("Elasticsearch error: {body}")); return Err(format!("Elasticsearch error: {body}"));
} }
let indices: Vec<CatIndex> = resp let indices: Vec<CatIndex> = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
.json() let mut names: Vec<String> = indices.into_iter().filter(|i| !i.index.starts_with('.')).map(|i| i.index).collect();
.await
.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
let mut names: Vec<String> = indices
.into_iter()
.filter(|i| !i.index.starts_with('.'))
.map(|i| i.index)
.collect();
names.sort(); names.sort();
Ok(names) Ok(names)
} }
@@ -147,22 +126,14 @@ pub async fn find_documents(
}); });
let path = format!("/{}/_search", index); let path = format!("/{}/_search", index);
let resp = client let resp = client.post(&path).json(&body).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
.post(&path)
.json(&body)
.send()
.await
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
if !resp.status().is_success() { if !resp.status().is_success() {
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
return Err(format!("Elasticsearch error: {body}")); return Err(format!("Elasticsearch error: {body}"));
} }
let result: SearchResponse = resp let result: SearchResponse = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
.json()
.await
.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
let documents: Vec<serde_json::Value> = result let documents: Vec<serde_json::Value> = result
.hits .hits
@@ -178,56 +149,29 @@ pub async fn find_documents(
}) })
.collect(); .collect();
Ok(MongoDocumentResult { Ok(MongoDocumentResult { documents, total: result.hits.total.value })
documents,
total: result.hits.total.value,
})
} }
pub async fn insert_document( pub async fn insert_document(client: &EsClient, index: &str, doc_json: &str) -> Result<String, String> {
client: &EsClient, let doc: serde_json::Value = serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
index: &str,
doc_json: &str,
) -> Result<String, String> {
let doc: serde_json::Value =
serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
let path = format!("/{}/_doc?refresh=true", index); let path = format!("/{}/_doc?refresh=true", index);
let resp = client let resp = client.post(&path).json(&doc).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
.post(&path)
.json(&doc)
.send()
.await
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
if !resp.status().is_success() { if !resp.status().is_success() {
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
return Err(format!("Elasticsearch error: {body}")); return Err(format!("Elasticsearch error: {body}"));
} }
let result: serde_json::Value = resp let result: serde_json::Value = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
.json()
.await
.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
Ok(result["_id"].as_str().unwrap_or("").to_string()) Ok(result["_id"].as_str().unwrap_or("").to_string())
} }
pub async fn update_document( pub async fn update_document(client: &EsClient, index: &str, id: &str, doc_json: &str) -> Result<u64, String> {
client: &EsClient, let doc: serde_json::Value = serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
index: &str,
id: &str,
doc_json: &str,
) -> Result<u64, String> {
let doc: serde_json::Value =
serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
let path = format!("/{}/_doc/{}?refresh=true", index, id); let path = format!("/{}/_doc/{}?refresh=true", index, id);
let resp = client let resp = client.put(&path).json(&doc).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
.put(&path)
.json(&doc)
.send()
.await
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
if !resp.status().is_success() { if !resp.status().is_success() {
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
@@ -239,11 +183,7 @@ pub async fn update_document(
pub async fn delete_document(client: &EsClient, index: &str, id: &str) -> Result<u64, String> { pub async fn delete_document(client: &EsClient, index: &str, id: &str) -> Result<u64, String> {
let path = format!("/{}/_doc/{}?refresh=true", index, id); let path = format!("/{}/_doc/{}?refresh=true", index, id);
let resp = client let resp = client.delete(&path).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
.delete(&path)
.send()
.await
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
if !resp.status().is_success() { if !resp.status().is_success() {
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
+5 -8
View File
@@ -36,12 +36,9 @@ where
} }
pub async fn probe_tcp_endpoint(label: &str, host: &str, port: u16) -> Result<(), String> { pub async fn probe_tcp_endpoint(label: &str, host: &str, port: u16) -> Result<(), String> {
tokio::time::timeout( tokio::time::timeout(tcp_probe_timeout(), tokio::net::TcpStream::connect((host, port)))
tcp_probe_timeout(), .await
tokio::net::TcpStream::connect((host, port)), .map_err(|_| format!("{label} TCP connection timed out ({TCP_PROBE_TIMEOUT_SECS}s)"))?
) .map(|_| ())
.await .map_err(|e| format!("{label} TCP connection failed: {e}"))
.map_err(|_| format!("{label} TCP connection timed out ({TCP_PROBE_TIMEOUT_SECS}s)"))?
.map(|_| ())
.map_err(|e| format!("{label} TCP connection failed: {e}"))
} }
+11 -42
View File
@@ -14,9 +14,7 @@ pub struct MongoDocumentResult {
pub async fn connect(url: &str) -> Result<Client, String> { pub async fn connect(url: &str) -> Result<Client, String> {
with_connection_timeout("MongoDB", async { with_connection_timeout("MongoDB", async {
Client::with_uri_str(url) Client::with_uri_str(url).await.map_err(|e| format!("MongoDB connection failed: {e}"))
.await
.map_err(|e| format!("MongoDB connection failed: {e}"))
}) })
.await .await
} }
@@ -30,18 +28,11 @@ pub async fn test_connection(client: &Client) -> Result<(), String> {
} }
pub async fn list_databases(client: &Client) -> Result<Vec<String>, String> { pub async fn list_databases(client: &Client) -> Result<Vec<String>, String> {
client client.list_database_names().await.map_err(|e| e.to_string())
.list_database_names()
.await
.map_err(|e| e.to_string())
} }
pub async fn list_collections(client: &Client, database: &str) -> Result<Vec<String>, String> { pub async fn list_collections(client: &Client, database: &str) -> Result<Vec<String>, String> {
client client.database(database).list_collection_names().await.map_err(|e| e.to_string())
.database(database)
.list_collection_names()
.await
.map_err(|e| e.to_string())
} }
pub async fn find_documents( pub async fn find_documents(
@@ -53,17 +44,9 @@ pub async fn find_documents(
) -> Result<MongoDocumentResult, String> { ) -> Result<MongoDocumentResult, String> {
let col = client.database(database).collection::<Document>(collection); let col = client.database(database).collection::<Document>(collection);
let total = col let total = col.count_documents(doc! {}).await.map_err(|e| e.to_string())?;
.count_documents(doc! {})
.await
.map_err(|e| e.to_string())?;
let mut cursor = col let mut cursor = col.find(doc! {}).skip(skip).limit(limit).await.map_err(|e| e.to_string())?;
.find(doc! {})
.skip(skip)
.limit(limit)
.await
.map_err(|e| e.to_string())?;
let mut documents = Vec::new(); let mut documents = Vec::new();
while cursor.advance().await.map_err(|e| e.to_string())? { while cursor.advance().await.map_err(|e| e.to_string())? {
@@ -94,31 +77,17 @@ pub async fn update_document(
id: &str, id: &str,
doc_json: &str, doc_json: &str,
) -> Result<u64, String> { ) -> Result<u64, String> {
let oid = mongodb::bson::oid::ObjectId::parse_str(id) let oid = mongodb::bson::oid::ObjectId::parse_str(id).map_err(|e| format!("Invalid ObjectId: {e}"))?;
.map_err(|e| format!("Invalid ObjectId: {e}"))?; let new_doc: Document = serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
let new_doc: Document =
serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
let col = client.database(database).collection::<Document>(collection); let col = client.database(database).collection::<Document>(collection);
let result = col let result = col.replace_one(doc! { "_id": oid }, new_doc).await.map_err(|e| e.to_string())?;
.replace_one(doc! { "_id": oid }, new_doc)
.await
.map_err(|e| e.to_string())?;
Ok(result.modified_count) Ok(result.modified_count)
} }
pub async fn delete_document( pub async fn delete_document(client: &Client, database: &str, collection: &str, id: &str) -> Result<u64, String> {
client: &Client, let oid = mongodb::bson::oid::ObjectId::parse_str(id).map_err(|e| format!("Invalid ObjectId: {e}"))?;
database: &str,
collection: &str,
id: &str,
) -> Result<u64, String> {
let oid = mongodb::bson::oid::ObjectId::parse_str(id)
.map_err(|e| format!("Invalid ObjectId: {e}"))?;
let col = client.database(database).collection::<Document>(collection); let col = client.database(database).collection::<Document>(collection);
let result = col let result = col.delete_one(doc! { "_id": oid }).await.map_err(|e| e.to_string())?;
.delete_one(doc! { "_id": oid })
.await
.map_err(|e| e.to_string())?;
Ok(result.deleted_count) Ok(result.deleted_count)
} }
+42 -159
View File
@@ -5,9 +5,7 @@ use sqlx::{Column, Executor, Row, TypeInfo, ValueRef};
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use super::{connection_timeout, with_connection_timeout}; use super::{connection_timeout, with_connection_timeout};
use crate::types::{ use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
};
fn quote_value(s: &str) -> String { fn quote_value(s: &str) -> String {
format!("'{}'", s.replace('\\', "\\\\").replace('\'', "\\'")) format!("'{}'", s.replace('\\', "\\\\").replace('\'', "\\'"))
@@ -15,32 +13,20 @@ fn quote_value(s: &str) -> String {
fn get_str(row: &MySqlRow, idx: usize) -> String { fn get_str(row: &MySqlRow, idx: usize) -> String {
row.try_get::<String, _>(idx) row.try_get::<String, _>(idx)
.or_else(|_| { .or_else(|_| row.try_get::<Vec<u8>, _>(idx).map(|b| String::from_utf8_lossy(&b).to_string()))
row.try_get::<Vec<u8>, _>(idx)
.map(|b| String::from_utf8_lossy(&b).to_string())
})
.unwrap_or_default() .unwrap_or_default()
} }
fn get_str_by_name(row: &MySqlRow, name: &str) -> String { fn get_str_by_name(row: &MySqlRow, name: &str) -> String {
row.try_get::<String, _>(name) row.try_get::<String, _>(name)
.or_else(|_| { .or_else(|_| row.try_get::<Vec<u8>, _>(name).map(|b| String::from_utf8_lossy(&b).to_string()))
row.try_get::<Vec<u8>, _>(name)
.map(|b| String::from_utf8_lossy(&b).to_string())
})
.unwrap_or_default() .unwrap_or_default()
} }
fn get_opt_str(row: &MySqlRow, name: &str) -> Option<String> { fn get_opt_str(row: &MySqlRow, name: &str) -> Option<String> {
row.try_get::<Option<String>, _>(name) row.try_get::<Option<String>, _>(name).ok().flatten().or_else(|| {
.ok() row.try_get::<Option<Vec<u8>>, _>(name).ok().flatten().map(|b| String::from_utf8_lossy(&b).to_string())
.flatten() })
.or_else(|| {
row.try_get::<Option<Vec<u8>>, _>(name)
.ok()
.flatten()
.map(|b| String::from_utf8_lossy(&b).to_string())
})
} }
fn numeric_metadata_u64_to_i32(value: Option<u64>) -> Option<i32> { fn numeric_metadata_u64_to_i32(value: Option<u64>) -> Option<i32> {
@@ -52,9 +38,7 @@ fn numeric_metadata_i64_to_i32(value: Option<i64>) -> Option<i32> {
} }
fn numeric_metadata_str_to_i32(value: Option<String>) -> Option<i32> { fn numeric_metadata_str_to_i32(value: Option<String>) -> Option<i32> {
value value.and_then(|v| v.parse::<i64>().ok()).and_then(|v| i32::try_from(v).ok())
.and_then(|v| v.parse::<i64>().ok())
.and_then(|v| i32::try_from(v).ok())
} }
fn get_opt_i32(row: &MySqlRow, name: &str) -> Option<i32> { fn get_opt_i32(row: &MySqlRow, name: &str) -> Option<i32> {
@@ -67,9 +51,7 @@ fn get_opt_i32(row: &MySqlRow, name: &str) -> Option<i32> {
.flatten() .flatten()
.or_else(|| numeric_metadata_i64_to_i32(row.try_get::<Option<i64>, _>(name).ok().flatten())) .or_else(|| numeric_metadata_i64_to_i32(row.try_get::<Option<i64>, _>(name).ok().flatten()))
.or_else(|| numeric_metadata_u64_to_i32(row.try_get::<Option<u64>, _>(name).ok().flatten())) .or_else(|| numeric_metadata_u64_to_i32(row.try_get::<Option<u64>, _>(name).ok().flatten()))
.or_else(|| { .or_else(|| numeric_metadata_str_to_i32(row.try_get::<Option<String>, _>(name).ok().flatten()))
numeric_metadata_str_to_i32(row.try_get::<Option<String>, _>(name).ok().flatten())
})
.or_else(|| { .or_else(|| {
row.try_get::<Option<Vec<u8>>, _>(name) row.try_get::<Option<Vec<u8>>, _>(name)
.ok() .ok()
@@ -107,27 +89,20 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso
return v; return v;
} }
if let Ok(v) = row.try_get::<String, _>(idx) { if let Ok(v) = row.try_get::<String, _>(idx) {
return serde_json::from_str::<serde_json::Value>(&v) return serde_json::from_str::<serde_json::Value>(&v).unwrap_or(serde_json::Value::String(v));
.unwrap_or(serde_json::Value::String(v));
} }
return serde_json::Value::Null; return serde_json::Value::Null;
} }
if upper_type == "BOOLEAN" { if upper_type == "BOOLEAN" {
return row return row.try_get::<bool, _>(idx).map(serde_json::Value::Bool).unwrap_or(serde_json::Value::Null);
.try_get::<bool, _>(idx)
.map(serde_json::Value::Bool)
.unwrap_or(serde_json::Value::Null);
} }
if upper_type.contains("BIGINT") { if upper_type.contains("BIGINT") {
return row return row
.try_get::<i64, _>(idx) .try_get::<i64, _>(idx)
.map(|v| serde_json::Value::String(v.to_string())) .map(|v| serde_json::Value::String(v.to_string()))
.or_else(|_| { .or_else(|_| row.try_get::<u64, _>(idx).map(|v| serde_json::Value::String(v.to_string())))
row.try_get::<u64, _>(idx)
.map(|v| serde_json::Value::String(v.to_string()))
})
.unwrap_or(serde_json::Value::Null); .unwrap_or(serde_json::Value::Null);
} }
@@ -151,25 +126,16 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso
row.try_get::<String, _>(idx) row.try_get::<String, _>(idx)
.map(serde_json::Value::String) .map(serde_json::Value::String)
.or_else(|_| { .or_else(|_| row.try_get::<i64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
row.try_get::<i64, _>(idx) .or_else(|_| row.try_get::<u64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
.map(|v| serde_json::Value::Number(v.into()))
})
.or_else(|_| {
row.try_get::<u64, _>(idx)
.map(|v| serde_json::Value::Number(v.into()))
})
.or_else(|_| { .or_else(|_| {
row.try_get::<f64, _>(idx).map(|v| { row.try_get::<f64, _>(idx).map(|v| {
serde_json::Number::from_f64(v) serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
.map(serde_json::Value::Number)
.unwrap_or(serde_json::Value::Null)
}) })
}) })
.or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool)) .or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool))
.or_else(|_| { .or_else(|_| {
row.try_get::<Vec<u8>, _>(idx) row.try_get::<Vec<u8>, _>(idx).map(|b| serde_json::Value::String(String::from_utf8_lossy(&b).to_string()))
.map(|b| serde_json::Value::String(String::from_utf8_lossy(&b).to_string()))
}) })
.or_else(|e| mysql_temporal_to_json_value(row, idx).ok_or(e)) .or_else(|e| mysql_temporal_to_json_value(row, idx).ok_or(e))
.unwrap_or(serde_json::Value::Null) .unwrap_or(serde_json::Value::Null)
@@ -189,14 +155,9 @@ pub async fn connect(url: &str) -> Result<MySqlPool, String> {
} }
pub async fn connect_bare(url: &str) -> Result<MySqlPool, String> { pub async fn connect_bare(url: &str) -> Result<MySqlPool, String> {
let options: sqlx::mysql::MySqlConnectOptions = url let options: sqlx::mysql::MySqlConnectOptions =
.parse() url.parse().map_err(|e: sqlx::Error| format!("Invalid MySQL URL: {e}"))?;
.map_err(|e: sqlx::Error| format!("Invalid MySQL URL: {e}"))?; let options = options.no_engine_substitution(false).set_names(false).pipes_as_concat(false).timezone(None);
let options = options
.no_engine_substitution(false)
.set_names(false)
.pipes_as_concat(false)
.timezone(None);
with_connection_timeout("MySQL", async { with_connection_timeout("MySQL", async {
MySqlPoolOptions::new() MySqlPoolOptions::new()
.max_connections(5) .max_connections(5)
@@ -210,18 +171,12 @@ pub async fn connect_bare(url: &str) -> Result<MySqlPool, String> {
} }
pub async fn list_databases(pool: &MySqlPool) -> Result<Vec<DatabaseInfo>, String> { pub async fn list_databases(pool: &MySqlPool) -> Result<Vec<DatabaseInfo>, String> {
let rows: Vec<MySqlRow> = let rows: Vec<MySqlRow> = sqlx::raw_sql("SELECT SCHEMA_NAME FROM information_schema.SCHEMATA ORDER BY SCHEMA_NAME")
sqlx::raw_sql("SELECT SCHEMA_NAME FROM information_schema.SCHEMATA ORDER BY SCHEMA_NAME") .fetch_all(pool)
.fetch_all(pool) .await
.await .map_err(|e| e.to_string())?;
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows.iter().map(|row| DatabaseInfo { name: get_str(row, 0) }).collect())
.iter()
.map(|row| DatabaseInfo {
name: get_str(row, 0),
})
.collect())
} }
pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableInfo>, String> { pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableInfo>, String> {
@@ -229,10 +184,7 @@ pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableIn
"SELECT TABLE_NAME, TABLE_TYPE FROM information_schema.TABLES WHERE TABLE_SCHEMA = {} ORDER BY TABLE_NAME", "SELECT TABLE_NAME, TABLE_TYPE FROM information_schema.TABLES WHERE TABLE_SCHEMA = {} ORDER BY TABLE_NAME",
quote_value(database), quote_value(database),
); );
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql) let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
@@ -243,11 +195,7 @@ pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableIn
.collect()) .collect())
} }
pub async fn get_columns( pub async fn get_columns(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
pool: &MySqlPool,
database: &str,
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
let sql = format!( let sql = format!(
"SELECT c.COLUMN_NAME, c.COLUMN_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, c.EXTRA, c.COLUMN_COMMENT, \ "SELECT c.COLUMN_NAME, c.COLUMN_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, c.EXTRA, c.COLUMN_COMMENT, \
CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \ CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \
@@ -263,10 +211,7 @@ pub async fn get_columns(
quote_value(database), quote_value(database),
quote_value(table), quote_value(table),
); );
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql) let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
@@ -295,22 +240,11 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
|| trimmed.starts_with("EXPLAIN") || trimmed.starts_with("EXPLAIN")
{ {
if bare { if bare {
let rows: Vec<MySqlRow> = sqlx::raw_sql(sql) let rows: Vec<MySqlRow> = sqlx::raw_sql(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let (columns, column_types) = if let Some(first) = rows.first() { let (columns, column_types) = if let Some(first) = rows.first() {
let cols: Vec<String> = first let cols: Vec<String> = first.columns().iter().map(|c| c.name().to_string()).collect();
.columns() let types: Vec<String> = first.columns().iter().map(|c| c.type_info().name().to_string()).collect();
.iter()
.map(|c| c.name().to_string())
.collect();
let types: Vec<String> = first
.columns()
.iter()
.map(|c| c.type_info().name().to_string())
.collect();
(cols, types) (cols, types)
} else { } else {
(vec![], vec![]) (vec![], vec![])
@@ -320,13 +254,7 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
.iter() .iter()
.map(|row| { .map(|row| {
(0..row.len()) (0..row.len())
.map(|i| { .map(|i| mysql_value_to_json(row, i, column_types.get(i).map(String::as_str).unwrap_or("")))
mysql_value_to_json(
row,
i,
column_types.get(i).map(String::as_str).unwrap_or(""),
)
})
.collect() .collect()
}) })
.collect(); .collect();
@@ -340,33 +268,16 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
}) })
} else { } else {
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?; let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
let columns: Vec<String> = desc let columns: Vec<String> = desc.columns().iter().map(|c| c.name().to_string()).collect();
.columns() let column_types: Vec<String> = desc.columns().iter().map(|c| c.type_info().name().to_string()).collect();
.iter()
.map(|c| c.name().to_string())
.collect();
let column_types: Vec<String> = desc
.columns()
.iter()
.map(|c| c.type_info().name().to_string())
.collect();
let rows: Vec<MySqlRow> = sqlx::query(sql) let rows: Vec<MySqlRow> = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let result_rows: Vec<Vec<serde_json::Value>> = rows let result_rows: Vec<Vec<serde_json::Value>> = rows
.iter() .iter()
.map(|row| { .map(|row| {
(0..row.len()) (0..row.len())
.map(|i| { .map(|i| mysql_value_to_json(row, i, column_types.get(i).map(String::as_str).unwrap_or("")))
mysql_value_to_json(
row,
i,
column_types.get(i).map(String::as_str).unwrap_or(""),
)
})
.collect() .collect()
}) })
.collect(); .collect();
@@ -380,10 +291,7 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
}) })
} }
} else { } else {
let result = sqlx::raw_sql(sql) let result = sqlx::raw_sql(sql).execute(pool).await.map_err(|e| e.to_string())?;
.execute(pool)
.await
.map_err(|e| e.to_string())?;
Ok(QueryResult { Ok(QueryResult {
columns: vec![], columns: vec![],
@@ -395,11 +303,7 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
} }
} }
pub async fn list_indexes( pub async fn list_indexes(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
pool: &MySqlPool,
database: &str,
table: &str,
) -> Result<Vec<IndexInfo>, String> {
let sql = format!( let sql = format!(
"SELECT INDEX_NAME, GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) AS columns, \ "SELECT INDEX_NAME, GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) AS columns, \
MIN(NON_UNIQUE) = 0 AS is_unique, INDEX_NAME = 'PRIMARY' AS is_primary, \ MIN(NON_UNIQUE) = 0 AS is_unique, INDEX_NAME = 'PRIMARY' AS is_primary, \
@@ -411,10 +315,7 @@ pub async fn list_indexes(
quote_value(database), quote_value(database),
quote_value(table), quote_value(table),
); );
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql) let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
@@ -422,11 +323,7 @@ pub async fn list_indexes(
let cols_str = get_str_by_name(row, "columns"); let cols_str = get_str_by_name(row, "columns");
IndexInfo { IndexInfo {
name: get_str_by_name(row, "INDEX_NAME"), name: get_str_by_name(row, "INDEX_NAME"),
columns: cols_str columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(),
.split(',')
.filter(|s| !s.is_empty())
.map(|s| s.to_string())
.collect(),
is_unique: row.get::<bool, _>("is_unique"), is_unique: row.get::<bool, _>("is_unique"),
is_primary: row.get::<bool, _>("is_primary"), is_primary: row.get::<bool, _>("is_primary"),
filter: None, filter: None,
@@ -438,11 +335,7 @@ pub async fn list_indexes(
.collect()) .collect())
} }
pub async fn list_foreign_keys( pub async fn list_foreign_keys(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
pool: &MySqlPool,
database: &str,
table: &str,
) -> Result<Vec<ForeignKeyInfo>, String> {
let sql = format!( let sql = format!(
"SELECT kcu.CONSTRAINT_NAME, kcu.COLUMN_NAME, \ "SELECT kcu.CONSTRAINT_NAME, kcu.COLUMN_NAME, \
kcu.REFERENCED_TABLE_NAME, kcu.REFERENCED_COLUMN_NAME \ kcu.REFERENCED_TABLE_NAME, kcu.REFERENCED_COLUMN_NAME \
@@ -453,10 +346,7 @@ pub async fn list_foreign_keys(
quote_value(database), quote_value(database),
quote_value(table), quote_value(table),
); );
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql) let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
@@ -469,11 +359,7 @@ pub async fn list_foreign_keys(
.collect()) .collect())
} }
pub async fn list_triggers( pub async fn list_triggers(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
pool: &MySqlPool,
database: &str,
table: &str,
) -> Result<Vec<TriggerInfo>, String> {
let sql = format!( let sql = format!(
"SELECT TRIGGER_NAME, EVENT_MANIPULATION, ACTION_TIMING \ "SELECT TRIGGER_NAME, EVENT_MANIPULATION, ACTION_TIMING \
FROM information_schema.TRIGGERS \ FROM information_schema.TRIGGERS \
@@ -482,10 +368,7 @@ pub async fn list_triggers(
quote_value(database), quote_value(database),
quote_value(table), quote_value(table),
); );
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql) let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
+43 -88
View File
@@ -2,27 +2,16 @@ use oracle_rs::{Config, Connection};
use std::time::Instant; use std::time::Instant;
use super::{connection_timeout, CONNECTION_TIMEOUT_SECS}; use super::{connection_timeout, CONNECTION_TIMEOUT_SECS};
use crate::types::{ use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
};
pub type OracleClient = Connection; pub type OracleClient = Connection;
pub async fn connect( pub async fn connect(host: &str, port: u16, service: &str, user: &str, pass: &str) -> Result<OracleClient, String> {
host: &str,
port: u16,
service: &str,
user: &str,
pass: &str,
) -> Result<OracleClient, String> {
let config = Config::new(host, port, service, user, pass); let config = Config::new(host, port, service, user, pass);
tokio::time::timeout( tokio::time::timeout(connection_timeout(), Connection::connect_with_config(config))
connection_timeout(), .await
Connection::connect_with_config(config), .map_err(|_| format!("Oracle connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
) .map_err(|e| format!("Oracle connection failed: {e}"))
.await
.map_err(|_| format!("Oracle connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.map_err(|e| format!("Oracle connection failed: {e}"))
} }
fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value { fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value {
@@ -30,9 +19,9 @@ fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value {
oracle_rs::Value::Null => serde_json::Value::Null, oracle_rs::Value::Null => serde_json::Value::Null,
oracle_rs::Value::String(s) => serde_json::Value::String(s.clone()), oracle_rs::Value::String(s) => serde_json::Value::String(s.clone()),
oracle_rs::Value::Integer(n) => serde_json::Value::Number((*n).into()), oracle_rs::Value::Integer(n) => serde_json::Value::Number((*n).into()),
oracle_rs::Value::Float(f) => serde_json::Number::from_f64(*f) oracle_rs::Value::Float(f) => {
.map(serde_json::Value::Number) serde_json::Number::from_f64(*f).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
.unwrap_or(serde_json::Value::Null), }
oracle_rs::Value::Boolean(b) => serde_json::Value::Bool(*b), oracle_rs::Value::Boolean(b) => serde_json::Value::Bool(*b),
oracle_rs::Value::Json(v) => v.clone(), oracle_rs::Value::Json(v) => v.clone(),
_ => serde_json::Value::String(format!("{val:?}")), _ => serde_json::Value::String(format!("{val:?}")),
@@ -40,29 +29,15 @@ fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value {
} }
pub async fn list_databases(conn: &OracleClient) -> Result<Vec<DatabaseInfo>, String> { pub async fn list_databases(conn: &OracleClient) -> Result<Vec<DatabaseInfo>, String> {
let result = conn let result =
.query("SELECT username FROM all_users ORDER BY username", &[]) conn.query("SELECT username FROM all_users ORDER BY username", &[]).await.map_err(|e| e.to_string())?;
.await Ok(result.rows.iter().map(|row| DatabaseInfo { name: row.get_string(0).unwrap_or("").to_string() }).collect())
.map_err(|e| e.to_string())?;
Ok(result
.rows
.iter()
.map(|row| DatabaseInfo {
name: row.get_string(0).unwrap_or("").to_string(),
})
.collect())
} }
pub async fn list_schemas(conn: &OracleClient) -> Result<Vec<String>, String> { pub async fn list_schemas(conn: &OracleClient) -> Result<Vec<String>, String> {
let result = conn let result =
.query("SELECT username FROM all_users ORDER BY username", &[]) conn.query("SELECT username FROM all_users ORDER BY username", &[]).await.map_err(|e| e.to_string())?;
.await Ok(result.rows.iter().map(|row| row.get_string(0).unwrap_or("").to_string()).collect())
.map_err(|e| e.to_string())?;
Ok(result
.rows
.iter()
.map(|row| row.get_string(0).unwrap_or("").to_string())
.collect())
} }
pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result<Vec<TableInfo>, String> { pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result<Vec<TableInfo>, String> {
@@ -84,37 +59,36 @@ pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result<Vec<TableI
.collect()) .collect())
} }
pub async fn get_columns( pub async fn get_columns(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
conn: &OracleClient,
schema: &str,
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
let s = schema.replace('\'', "''"); let s = schema.replace('\'', "''");
let t = table.replace('\'', "''"); let t = table.replace('\'', "''");
let pk_result = conn.query( let pk_result = conn
&format!( .query(
"SELECT cols.COLUMN_NAME FROM ALL_CONS_COLUMNS cols \ &format!(
"SELECT cols.COLUMN_NAME FROM ALL_CONS_COLUMNS cols \
JOIN ALL_CONSTRAINTS cons ON cols.CONSTRAINT_NAME = cons.CONSTRAINT_NAME AND cols.OWNER = cons.OWNER \ JOIN ALL_CONSTRAINTS cons ON cols.CONSTRAINT_NAME = cons.CONSTRAINT_NAME AND cols.OWNER = cons.OWNER \
WHERE cons.CONSTRAINT_TYPE = 'P' AND cons.OWNER = '{s}' AND cons.TABLE_NAME = '{t}'" WHERE cons.CONSTRAINT_TYPE = 'P' AND cons.OWNER = '{s}' AND cons.TABLE_NAME = '{t}'"
), ),
&[], &[],
).await.map_err(|e| e.to_string())?; )
let pk_names: std::collections::HashSet<String> = pk_result .await
.rows .map_err(|e| e.to_string())?;
.iter() let pk_names: std::collections::HashSet<String> =
.filter_map(|row| row.get_string(0).map(|s| s.to_string())) pk_result.rows.iter().filter_map(|row| row.get_string(0).map(|s| s.to_string())).collect();
.collect();
let col_result = conn.query( let col_result = conn
&format!( .query(
"SELECT COLUMN_NAME, DATA_TYPE, NULLABLE, DATA_PRECISION, DATA_SCALE, DATA_LENGTH, CHAR_LENGTH \ &format!(
"SELECT COLUMN_NAME, DATA_TYPE, NULLABLE, DATA_PRECISION, DATA_SCALE, DATA_LENGTH, CHAR_LENGTH \
FROM ALL_TAB_COLUMNS \ FROM ALL_TAB_COLUMNS \
WHERE OWNER = '{s}' AND TABLE_NAME = '{t}' \ WHERE OWNER = '{s}' AND TABLE_NAME = '{t}' \
ORDER BY COLUMN_ID" ORDER BY COLUMN_ID"
), ),
&[], &[],
).await.map_err(|e| e.to_string())?; )
.await
.map_err(|e| e.to_string())?;
Ok(col_result Ok(col_result
.rows .rows
@@ -161,11 +135,7 @@ pub async fn get_columns(
.collect()) .collect())
} }
pub async fn list_indexes( pub async fn list_indexes(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
conn: &OracleClient,
schema: &str,
table: &str,
) -> Result<Vec<IndexInfo>, String> {
let sql = format!( let sql = format!(
"SELECT i.INDEX_NAME, \ "SELECT i.INDEX_NAME, \
LISTAGG(ic.COLUMN_NAME, ',') WITHIN GROUP (ORDER BY ic.COLUMN_POSITION) AS columns, \ LISTAGG(ic.COLUMN_NAME, ',') WITHIN GROUP (ORDER BY ic.COLUMN_POSITION) AS columns, \
@@ -189,11 +159,7 @@ pub async fn list_indexes(
let cols_str = row.get_string(1).unwrap_or(""); let cols_str = row.get_string(1).unwrap_or("");
IndexInfo { IndexInfo {
name: row.get_string(0).unwrap_or("").to_string(), name: row.get_string(0).unwrap_or("").to_string(),
columns: cols_str columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(),
.split(',')
.filter(|s| !s.is_empty())
.map(|s| s.to_string())
.collect(),
is_unique: row.get_string(2).unwrap_or("") == "UNIQUE", is_unique: row.get_string(2).unwrap_or("") == "UNIQUE",
is_primary: row.get_i64(3).unwrap_or(0) == 1, is_primary: row.get_i64(3).unwrap_or(0) == 1,
filter: None, filter: None,
@@ -205,11 +171,7 @@ pub async fn list_indexes(
.collect()) .collect())
} }
pub async fn list_foreign_keys( pub async fn list_foreign_keys(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
conn: &OracleClient,
schema: &str,
table: &str,
) -> Result<Vec<ForeignKeyInfo>, String> {
let sql = format!( let sql = format!(
"SELECT c.CONSTRAINT_NAME, cc.COLUMN_NAME, rc.TABLE_NAME, rcc.COLUMN_NAME \ "SELECT c.CONSTRAINT_NAME, cc.COLUMN_NAME, rc.TABLE_NAME, rcc.COLUMN_NAME \
FROM ALL_CONSTRAINTS c \ FROM ALL_CONSTRAINTS c \
@@ -218,7 +180,8 @@ pub async fn list_foreign_keys(
JOIN ALL_CONS_COLUMNS rcc ON rc.CONSTRAINT_NAME = rcc.CONSTRAINT_NAME AND rc.OWNER = rcc.OWNER \ JOIN ALL_CONS_COLUMNS rcc ON rc.CONSTRAINT_NAME = rcc.CONSTRAINT_NAME AND rc.OWNER = rcc.OWNER \
WHERE c.CONSTRAINT_TYPE = 'R' AND c.OWNER = '{s}' AND c.TABLE_NAME = '{t}' \ WHERE c.CONSTRAINT_TYPE = 'R' AND c.OWNER = '{s}' AND c.TABLE_NAME = '{t}' \
ORDER BY c.CONSTRAINT_NAME", ORDER BY c.CONSTRAINT_NAME",
s = schema.replace('\'', "''"), t = table.replace('\'', "''") s = schema.replace('\'', "''"),
t = table.replace('\'', "''")
); );
let result = conn.query(&sql, &[]).await.map_err(|e| e.to_string())?; let result = conn.query(&sql, &[]).await.map_err(|e| e.to_string())?;
Ok(result Ok(result
@@ -233,11 +196,7 @@ pub async fn list_foreign_keys(
.collect()) .collect())
} }
pub async fn list_triggers( pub async fn list_triggers(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
conn: &OracleClient,
schema: &str,
table: &str,
) -> Result<Vec<TriggerInfo>, String> {
let sql = format!( let sql = format!(
"SELECT TRIGGER_NAME, TRIGGERING_EVENT, TRIGGER_TYPE \ "SELECT TRIGGER_NAME, TRIGGERING_EVENT, TRIGGER_TYPE \
FROM ALL_TRIGGERS \ FROM ALL_TRIGGERS \
@@ -276,11 +235,7 @@ pub async fn execute_query(conn: &OracleClient, sql: &str) -> Result<QueryResult
.iter() .iter()
.map(|row| { .map(|row| {
(0..columns.len()) (0..columns.len())
.map(|i| { .map(|i| row.get(i).map(|v| value_to_json(v)).unwrap_or(serde_json::Value::Null))
row.get(i)
.map(|v| value_to_json(v))
.unwrap_or(serde_json::Value::Null)
})
.collect() .collect()
}) })
.collect(); .collect();
+27 -100
View File
@@ -5,9 +5,7 @@ use sqlx::{Column, Executor, Row, TypeInfo, ValueRef};
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use super::{connection_timeout, with_connection_timeout}; use super::{connection_timeout, with_connection_timeout};
use crate::types::{ use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
};
fn pg_temporal_to_json_value(row: &PgRow, idx: usize) -> Option<serde_json::Value> { fn pg_temporal_to_json_value(row: &PgRow, idx: usize) -> Option<serde_json::Value> {
if let Ok(v) = row.try_get::<DateTime<Utc>, _>(idx) { if let Ok(v) = row.try_get::<DateTime<Utc>, _>(idx) {
@@ -37,17 +35,13 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val
return v; return v;
} }
if let Ok(v) = row.try_get::<String, _>(idx) { if let Ok(v) = row.try_get::<String, _>(idx) {
return serde_json::from_str::<serde_json::Value>(&v) return serde_json::from_str::<serde_json::Value>(&v).unwrap_or(serde_json::Value::String(v));
.unwrap_or(serde_json::Value::String(v));
} }
return serde_json::Value::Null; return serde_json::Value::Null;
} }
if upper == "BOOL" { if upper == "BOOL" {
return row return row.try_get::<bool, _>(idx).map(serde_json::Value::Bool).unwrap_or(serde_json::Value::Null);
.try_get::<bool, _>(idx)
.map(serde_json::Value::Bool)
.unwrap_or(serde_json::Value::Null);
} }
if upper.contains("TIMESTAMP") if upper.contains("TIMESTAMP")
@@ -70,19 +64,11 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val
row.try_get::<String, _>(idx) row.try_get::<String, _>(idx)
.map(serde_json::Value::String) .map(serde_json::Value::String)
.or_else(|_| { .or_else(|_| row.try_get::<i64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
row.try_get::<i64, _>(idx) .or_else(|_| row.try_get::<i32, _>(idx).map(|v| serde_json::Value::Number(v.into())))
.map(|v| serde_json::Value::Number(v.into()))
})
.or_else(|_| {
row.try_get::<i32, _>(idx)
.map(|v| serde_json::Value::Number(v.into()))
})
.or_else(|_| { .or_else(|_| {
row.try_get::<f64, _>(idx).map(|v| { row.try_get::<f64, _>(idx).map(|v| {
serde_json::Number::from_f64(v) serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
.map(serde_json::Value::Number)
.unwrap_or(serde_json::Value::Null)
}) })
}) })
.or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool)) .or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool))
@@ -104,18 +90,12 @@ pub async fn connect(url: &str) -> Result<PgPool, String> {
} }
pub async fn list_databases(pool: &PgPool) -> Result<Vec<DatabaseInfo>, String> { pub async fn list_databases(pool: &PgPool) -> Result<Vec<DatabaseInfo>, String> {
let rows: Vec<PgRow> = let rows: Vec<PgRow> = sqlx::query("SELECT datname FROM pg_database WHERE datistemplate = false ORDER BY datname")
sqlx::query("SELECT datname FROM pg_database WHERE datistemplate = false ORDER BY datname") .fetch_all(pool)
.fetch_all(pool) .await
.await .map_err(|e| e.to_string())?;
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows.iter().map(|row| DatabaseInfo { name: row.get::<String, _>("datname") }).collect())
.iter()
.map(|row| DatabaseInfo {
name: row.get::<String, _>("datname"),
})
.collect())
} }
pub async fn list_tables(pool: &PgPool, schema: &str) -> Result<Vec<TableInfo>, String> { pub async fn list_tables(pool: &PgPool, schema: &str) -> Result<Vec<TableInfo>, String> {
@@ -149,17 +129,10 @@ pub async fn list_schemas(pool: &PgPool) -> Result<Vec<String>, String> {
.await .await
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
Ok(rows Ok(rows.iter().map(|row| row.get::<String, _>("schema_name")).collect())
.iter()
.map(|row| row.get::<String, _>("schema_name"))
.collect())
} }
pub async fn get_columns( pub async fn get_columns(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
pool: &PgPool,
schema: &str,
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
let rows: Vec<PgRow> = sqlx::query( let rows: Vec<PgRow> = sqlx::query(
"SELECT a.attname AS column_name, \ "SELECT a.attname AS column_name, \
format_type(a.atttypid, a.atttypmod) AS full_type, \ format_type(a.atttypid, a.atttypmod) AS full_type, \
@@ -194,9 +167,7 @@ pub async fn get_columns(
Ok(rows Ok(rows
.iter() .iter()
.map(|row| { .map(|row| {
let full_type = row let full_type = row.get::<Option<String>, _>("full_type").unwrap_or_default();
.get::<Option<String>, _>("full_type")
.unwrap_or_default();
ColumnInfo { ColumnInfo {
name: row.get::<String, _>("column_name"), name: row.get::<String, _>("column_name"),
data_type: full_type, data_type: full_type,
@@ -223,32 +194,19 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
|| trimmed.starts_with("WITH") || trimmed.starts_with("WITH")
|| trimmed.starts_with("TABLE") || trimmed.starts_with("TABLE")
{ {
let rows: Vec<PgRow> = sqlx::query(sql) let rows: Vec<PgRow> = sqlx::query(sql).persistent(false).fetch_all(pool).await.map_err(|e| e.to_string())?;
.persistent(false)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let (columns, column_types): (Vec<String>, Vec<String>) = if let Some(first) = rows.first() let (columns, column_types): (Vec<String>, Vec<String>) = if let Some(first) = rows.first() {
{
let cols = first.columns(); let cols = first.columns();
( (
cols.iter().map(|c| c.name().to_string()).collect(), cols.iter().map(|c| c.name().to_string()).collect(),
cols.iter() cols.iter().map(|c| c.type_info().name().to_string()).collect(),
.map(|c| c.type_info().name().to_string())
.collect(),
) )
} else { } else {
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?; let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
( (
desc.columns() desc.columns().iter().map(|c| c.name().to_string()).collect(),
.iter() desc.columns().iter().map(|c| c.type_info().name().to_string()).collect(),
.map(|c| c.name().to_string())
.collect(),
desc.columns()
.iter()
.map(|c| c.type_info().name().to_string())
.collect(),
) )
}; };
@@ -256,13 +214,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
.iter() .iter()
.map(|row| { .map(|row| {
(0..row.len()) (0..row.len())
.map(|i| { .map(|i| pg_value_to_json(row, i, column_types.get(i).map(String::as_str).unwrap_or("")))
pg_value_to_json(
row,
i,
column_types.get(i).map(String::as_str).unwrap_or(""),
)
})
.collect() .collect()
}) })
.collect(); .collect();
@@ -275,10 +227,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
truncated: false, truncated: false,
}) })
} else { } else {
let result = sqlx::query(sql) let result = sqlx::query(sql).execute(pool).await.map_err(|e| e.to_string())?;
.execute(pool)
.await
.map_err(|e| e.to_string())?;
Ok(QueryResult { Ok(QueryResult {
columns: vec![], columns: vec![],
@@ -290,11 +239,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
} }
} }
pub async fn list_indexes( pub async fn list_indexes(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
pool: &PgPool,
schema: &str,
table: &str,
) -> Result<Vec<IndexInfo>, String> {
let rows: Vec<PgRow> = sqlx::query( let rows: Vec<PgRow> = sqlx::query(
"SELECT i.relname AS index_name, \ "SELECT i.relname AS index_name, \
array_agg(COALESCE(a.attname, pg_get_indexdef(ix.indexrelid, k.n::int, true)) ORDER BY k.n) AS columns, \ array_agg(COALESCE(a.attname, pg_get_indexdef(ix.indexrelid, k.n::int, true)) ORDER BY k.n) AS columns, \
@@ -326,15 +271,9 @@ pub async fn list_indexes(
.iter() .iter()
.map(|row| { .map(|row| {
let all_cols: Vec<String> = row.get::<Vec<String>, _>("columns"); let all_cols: Vec<String> = row.get::<Vec<String>, _>("columns");
let nkeyatts = row let nkeyatts = row.get::<Option<i16>, _>("nkeyatts").unwrap_or(all_cols.len() as i16) as usize;
.get::<Option<i16>, _>("nkeyatts")
.unwrap_or(all_cols.len() as i16) as usize;
let key_cols = all_cols[..nkeyatts].to_vec(); let key_cols = all_cols[..nkeyatts].to_vec();
let included = if nkeyatts < all_cols.len() { let included = if nkeyatts < all_cols.len() { all_cols[nkeyatts..].to_vec() } else { vec![] };
all_cols[nkeyatts..].to_vec()
} else {
vec![]
};
IndexInfo { IndexInfo {
name: row.get::<String, _>("index_name"), name: row.get::<String, _>("index_name"),
columns: key_cols, columns: key_cols,
@@ -342,22 +281,14 @@ pub async fn list_indexes(
is_primary: row.get::<bool, _>("is_primary"), is_primary: row.get::<bool, _>("is_primary"),
filter: row.get::<Option<String>, _>("filter_expr"), filter: row.get::<Option<String>, _>("filter_expr"),
index_type: row.get::<Option<String>, _>("index_type"), index_type: row.get::<Option<String>, _>("index_type"),
included_columns: if included.is_empty() { included_columns: if included.is_empty() { None } else { Some(included) },
None
} else {
Some(included)
},
comment: row.get::<Option<String>, _>("index_comment"), comment: row.get::<Option<String>, _>("index_comment"),
} }
}) })
.collect()) .collect())
} }
pub async fn list_foreign_keys( pub async fn list_foreign_keys(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
pool: &PgPool,
schema: &str,
table: &str,
) -> Result<Vec<ForeignKeyInfo>, String> {
let rows: Vec<PgRow> = sqlx::query( let rows: Vec<PgRow> = sqlx::query(
"SELECT kcu.constraint_name, kcu.column_name, \ "SELECT kcu.constraint_name, kcu.column_name, \
ccu.table_name AS ref_table, ccu.column_name AS ref_column \ ccu.table_name AS ref_table, ccu.column_name AS ref_column \
@@ -388,11 +319,7 @@ pub async fn list_foreign_keys(
.collect()) .collect())
} }
pub async fn list_triggers( pub async fn list_triggers(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
pool: &PgPool,
schema: &str,
table: &str,
) -> Result<Vec<TriggerInfo>, String> {
let rows: Vec<PgRow> = sqlx::query( let rows: Vec<PgRow> = sqlx::query(
"SELECT trigger_name, event_manipulation, action_timing \ "SELECT trigger_name, event_manipulation, action_timing \
FROM information_schema.triggers \ FROM information_schema.triggers \
+42 -150
View File
@@ -29,44 +29,26 @@ pub struct RedisValue {
pub async fn connect(url: &str) -> Result<redis::aio::MultiplexedConnection, String> { pub async fn connect(url: &str) -> Result<redis::aio::MultiplexedConnection, String> {
let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?; let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?;
let mut con = tokio::time::timeout( let mut con = tokio::time::timeout(connection_timeout(), client.get_multiplexed_async_connection())
connection_timeout(), .await
client.get_multiplexed_async_connection(), .map_err(|_| format!("Redis connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
) .map_err(|e| format!("Redis connection failed: {e}"))?;
.await
.map_err(|_| format!("Redis connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.map_err(|e| format!("Redis connection failed: {e}"))?;
tokio::time::timeout( tokio::time::timeout(connection_timeout(), redis::cmd("PING").query_async::<String>(&mut con))
connection_timeout(), .await
redis::cmd("PING").query_async::<String>(&mut con), .map_err(|_| format!("Redis ping timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
) .map_err(|e| format!("Redis authentication failed or command rejected: {e}"))?;
.await
.map_err(|_| format!("Redis ping timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.map_err(|e| format!("Redis authentication failed or command rejected: {e}"))?;
Ok(con) Ok(con)
} }
pub async fn list_databases( pub async fn list_databases(con: &mut redis::aio::MultiplexedConnection) -> Result<Vec<u32>, String> {
con: &mut redis::aio::MultiplexedConnection, let configured_count =
) -> Result<Vec<u32>, String> { redis::cmd("CONFIG").arg("GET").arg("databases").query_async(con).await.ok().and_then(parse_database_count);
let configured_count = redis::cmd("CONFIG")
.arg("GET")
.arg("databases")
.query_async(con)
.await
.ok()
.and_then(parse_database_count);
let keyspace_dbs = list_keyspace_databases(con).await.unwrap_or_default(); let keyspace_dbs = list_keyspace_databases(con).await.unwrap_or_default();
let database_count = configured_count.unwrap_or(DEFAULT_REDIS_DATABASES); let database_count = configured_count.unwrap_or(DEFAULT_REDIS_DATABASES);
let max_db = keyspace_dbs let max_db = keyspace_dbs.iter().copied().max().map(|db| db + 1).unwrap_or(0);
.iter()
.copied()
.max()
.map(|db| db + 1)
.unwrap_or(0);
let visible_count = database_count.max(max_db).max(1); let visible_count = database_count.max(max_db).max(1);
Ok((0..visible_count).collect()) Ok((0..visible_count).collect())
@@ -88,14 +70,8 @@ fn parse_database_count(value: redis::Value) -> Option<u32> {
}) })
} }
async fn list_keyspace_databases( async fn list_keyspace_databases(con: &mut redis::aio::MultiplexedConnection) -> Result<Vec<u32>, String> {
con: &mut redis::aio::MultiplexedConnection, let info: String = redis::cmd("INFO").arg("keyspace").query_async(con).await.map_err(|e| e.to_string())?;
) -> Result<Vec<u32>, String> {
let info: String = redis::cmd("INFO")
.arg("keyspace")
.query_async(con)
.await
.map_err(|e| e.to_string())?;
let mut dbs = Vec::new(); let mut dbs = Vec::new();
for line in info.lines() { for line in info.lines() {
@@ -111,11 +87,7 @@ async fn list_keyspace_databases(
} }
pub async fn select_db(con: &mut redis::aio::MultiplexedConnection, db: u32) -> Result<(), String> { pub async fn select_db(con: &mut redis::aio::MultiplexedConnection, db: u32) -> Result<(), String> {
redis::cmd("SELECT") redis::cmd("SELECT").arg(db).query_async(con).await.map_err(|e| e.to_string())
.arg(db)
.query_async(con)
.await
.map_err(|e| e.to_string())
} }
pub async fn scan_keys_page( pub async fn scan_keys_page(
@@ -136,35 +108,18 @@ pub async fn scan_keys_page(
let mut result = Vec::new(); let mut result = Vec::new();
for key in &keys { for key in &keys {
let key_type: String = redis::cmd("TYPE") let key_type: String =
.arg(key.as_str()) redis::cmd("TYPE").arg(key.as_str()).query_async(con).await.unwrap_or_else(|_| "unknown".to_string());
.query_async(con)
.await
.unwrap_or_else(|_| "unknown".to_string());
let ttl: i64 = con.ttl(key.as_str()).await.unwrap_or(-1); let ttl: i64 = con.ttl(key.as_str()).await.unwrap_or(-1);
result.push(RedisKeyInfo { result.push(RedisKeyInfo { key: key.clone(), key_type, ttl });
key: key.clone(),
key_type,
ttl,
});
} }
Ok(RedisScanResult { Ok(RedisScanResult { cursor: next_cursor, keys: result })
cursor: next_cursor,
keys: result,
})
} }
pub async fn get_value( pub async fn get_value(con: &mut redis::aio::MultiplexedConnection, key: &str) -> Result<RedisValue, String> {
con: &mut redis::aio::MultiplexedConnection, let key_type: String = redis::cmd("TYPE").arg(key).query_async(con).await.map_err(|e| e.to_string())?;
key: &str,
) -> Result<RedisValue, String> {
let key_type: String = redis::cmd("TYPE")
.arg(key)
.query_async(con)
.await
.map_err(|e| e.to_string())?;
let ttl: i64 = con.ttl(key).await.unwrap_or(-1); let ttl: i64 = con.ttl(key).await.unwrap_or(-1);
@@ -182,33 +137,20 @@ pub async fn get_value(
serde_json::json!(v) serde_json::json!(v)
} }
"zset" => { "zset" => {
let v: Vec<(String, f64)> = con let v: Vec<(String, f64)> = con.zrange_withscores(key, 0, -1).await.map_err(|e| e.to_string())?;
.zrange_withscores(key, 0, -1) serde_json::json!(v.iter().map(|(m, s)| serde_json::json!({"member": m, "score": s})).collect::<Vec<_>>())
.await
.map_err(|e| e.to_string())?;
serde_json::json!(v
.iter()
.map(|(m, s)| serde_json::json!({"member": m, "score": s}))
.collect::<Vec<_>>())
} }
"hash" => { "hash" => {
let v: Vec<(String, String)> = con.hgetall(key).await.map_err(|e| e.to_string())?; let v: Vec<(String, String)> = con.hgetall(key).await.map_err(|e| e.to_string())?;
let map: serde_json::Map<String, serde_json::Value> = v let map: serde_json::Map<String, serde_json::Value> =
.into_iter() v.into_iter().map(|(k, v)| (k, serde_json::Value::String(v))).collect();
.map(|(k, v)| (k, serde_json::Value::String(v)))
.collect();
serde_json::Value::Object(map) serde_json::Value::Object(map)
} }
"stream" => get_stream_entries(con, key).await?, "stream" => get_stream_entries(con, key).await?,
_ => serde_json::Value::Null, _ => serde_json::Value::Null,
}; };
Ok(RedisValue { Ok(RedisValue { key: key.to_string(), key_type, ttl, value })
key: key.to_string(),
key_type,
ttl,
value,
})
} }
async fn get_stream_entries( async fn get_stream_entries(
@@ -286,23 +228,16 @@ pub async fn set_string(
value: &str, value: &str,
ttl: Option<i64>, ttl: Option<i64>,
) -> Result<(), String> { ) -> Result<(), String> {
con.set::<_, _, ()>(key, value) con.set::<_, _, ()>(key, value).await.map_err(|e| e.to_string())?;
.await
.map_err(|e| e.to_string())?;
if let Some(t) = ttl { if let Some(t) = ttl {
if t > 0 { if t > 0 {
con.expire::<_, ()>(key, t) con.expire::<_, ()>(key, t).await.map_err(|e| e.to_string())?;
.await
.map_err(|e| e.to_string())?;
} }
} }
Ok(()) Ok(())
} }
pub async fn delete_key( pub async fn delete_key(con: &mut redis::aio::MultiplexedConnection, key: &str) -> Result<(), String> {
con: &mut redis::aio::MultiplexedConnection,
key: &str,
) -> Result<(), String> {
con.del::<_, ()>(key).await.map_err(|e| e.to_string()) con.del::<_, ()>(key).await.map_err(|e| e.to_string())
} }
@@ -312,67 +247,29 @@ pub async fn hash_set(
field: &str, field: &str,
value: &str, value: &str,
) -> Result<(), String> { ) -> Result<(), String> {
con.hset::<_, _, _, ()>(key, field, value) con.hset::<_, _, _, ()>(key, field, value).await.map_err(|e| e.to_string())
.await
.map_err(|e| e.to_string())
} }
pub async fn hash_del( pub async fn hash_del(con: &mut redis::aio::MultiplexedConnection, key: &str, field: &str) -> Result<(), String> {
con: &mut redis::aio::MultiplexedConnection, con.hdel::<_, _, ()>(key, field).await.map_err(|e| e.to_string())
key: &str,
field: &str,
) -> Result<(), String> {
con.hdel::<_, _, ()>(key, field)
.await
.map_err(|e| e.to_string())
} }
pub async fn list_push( pub async fn list_push(con: &mut redis::aio::MultiplexedConnection, key: &str, value: &str) -> Result<(), String> {
con: &mut redis::aio::MultiplexedConnection, con.rpush::<_, _, ()>(key, value).await.map_err(|e| e.to_string())
key: &str,
value: &str,
) -> Result<(), String> {
con.rpush::<_, _, ()>(key, value)
.await
.map_err(|e| e.to_string())
} }
pub async fn list_remove( pub async fn list_remove(con: &mut redis::aio::MultiplexedConnection, key: &str, index: i64) -> Result<(), String> {
con: &mut redis::aio::MultiplexedConnection,
key: &str,
index: i64,
) -> Result<(), String> {
let placeholder = "__DELETED_PLACEHOLDER__"; let placeholder = "__DELETED_PLACEHOLDER__";
redis::cmd("LSET") redis::cmd("LSET").arg(key).arg(index).arg(placeholder).query_async::<()>(con).await.map_err(|e| e.to_string())?;
.arg(key) con.lrem::<_, _, ()>(key, 1, placeholder).await.map_err(|e| e.to_string())
.arg(index)
.arg(placeholder)
.query_async::<()>(con)
.await
.map_err(|e| e.to_string())?;
con.lrem::<_, _, ()>(key, 1, placeholder)
.await
.map_err(|e| e.to_string())
} }
pub async fn set_add( pub async fn set_add(con: &mut redis::aio::MultiplexedConnection, key: &str, member: &str) -> Result<(), String> {
con: &mut redis::aio::MultiplexedConnection, con.sadd::<_, _, ()>(key, member).await.map_err(|e| e.to_string())
key: &str,
member: &str,
) -> Result<(), String> {
con.sadd::<_, _, ()>(key, member)
.await
.map_err(|e| e.to_string())
} }
pub async fn set_remove( pub async fn set_remove(con: &mut redis::aio::MultiplexedConnection, key: &str, member: &str) -> Result<(), String> {
con: &mut redis::aio::MultiplexedConnection, con.srem::<_, _, ()>(key, member).await.map_err(|e| e.to_string())
key: &str,
member: &str,
) -> Result<(), String> {
con.srem::<_, _, ()>(key, member)
.await
.map_err(|e| e.to_string())
} }
#[cfg(test)] #[cfg(test)]
@@ -387,12 +284,7 @@ mod tests {
fn parses_stream_entries() { fn parses_stream_entries() {
let raw = RedisRawValue::Array(vec![RedisRawValue::Array(vec![ let raw = RedisRawValue::Array(vec![RedisRawValue::Array(vec![
bulk("1714470000000-0"), bulk("1714470000000-0"),
RedisRawValue::Array(vec![ RedisRawValue::Array(vec![bulk("event"), bulk("login"), bulk("user_id"), bulk("42")]),
bulk("event"),
bulk("login"),
bulk("user_id"),
bulk("42"),
]),
])]); ])]);
let parsed = parse_stream_entries(raw); let parsed = parse_stream_entries(raw);
+23 -75
View File
@@ -2,14 +2,10 @@ use sqlx::sqlite::{SqliteConnectOptions, SqlitePool, SqlitePoolOptions, SqliteRo
use sqlx::{Column, Executor, Row}; use sqlx::{Column, Executor, Row};
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use crate::types::{ use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
};
pub async fn connect_path(path: &str) -> Result<SqlitePool, String> { pub async fn connect_path(path: &str) -> Result<SqlitePool, String> {
let mut options = SqliteConnectOptions::new() let mut options = SqliteConnectOptions::new().filename(path).create_if_missing(true);
.filename(path)
.create_if_missing(true);
if is_network_path(path) { if is_network_path(path) {
options = options.vfs("unix-nolock"); options = options.vfs("unix-nolock");
@@ -25,16 +21,11 @@ pub async fn connect_path(path: &str) -> Result<SqlitePool, String> {
} }
fn is_network_path(path: &str) -> bool { fn is_network_path(path: &str) -> bool {
path.starts_with("\\\\") path.starts_with("\\\\") || path.starts_with("//") || path.contains("wsl.localhost") || path.contains("wsl$")
|| path.starts_with("//")
|| path.contains("wsl.localhost")
|| path.contains("wsl$")
} }
pub async fn list_databases(_pool: &SqlitePool) -> Result<Vec<DatabaseInfo>, String> { pub async fn list_databases(_pool: &SqlitePool) -> Result<Vec<DatabaseInfo>, String> {
Ok(vec![DatabaseInfo { Ok(vec![DatabaseInfo { name: "main".to_string() }])
name: "main".to_string(),
}])
} }
pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result<Vec<TableInfo>, String> { pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result<Vec<TableInfo>, String> {
@@ -51,25 +42,15 @@ pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result<Vec<TableIn
let t: String = row.get("type"); let t: String = row.get("type");
TableInfo { TableInfo {
name: row.get::<String, _>("name"), name: row.get::<String, _>("name"),
table_type: if t == "view" { table_type: if t == "view" { "VIEW".to_string() } else { "BASE TABLE".to_string() },
"VIEW".to_string()
} else {
"BASE TABLE".to_string()
},
} }
}) })
.collect()) .collect())
} }
pub async fn get_columns( pub async fn get_columns(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
pool: &SqlitePool, let rows: Vec<SqliteRow> =
_schema: &str, sqlx::query(&format!("PRAGMA table_info(\"{}\")", table)).fetch_all(pool).await.map_err(|e| e.to_string())?;
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
let rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA table_info(\"{}\")", table))
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
@@ -88,11 +69,7 @@ pub async fn get_columns(
.collect()) .collect())
} }
pub async fn list_indexes( pub async fn list_indexes(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
pool: &SqlitePool,
_schema: &str,
table: &str,
) -> Result<Vec<IndexInfo>, String> {
let safe_table = table.replace('"', "\"\""); let safe_table = table.replace('"', "\"\"");
let idx_rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA index_list(\"{safe_table}\")")) let idx_rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA index_list(\"{safe_table}\")"))
.fetch_all(pool) .fetch_all(pool)
@@ -112,10 +89,7 @@ pub async fn list_indexes(
.await .await
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
let columns: Vec<String> = col_rows let columns: Vec<String> = col_rows.iter().map(|r| r.get::<String, _>("name")).collect();
.iter()
.map(|r| r.get::<String, _>("name"))
.collect();
indexes.push(IndexInfo { indexes.push(IndexInfo {
name, name,
@@ -131,11 +105,7 @@ pub async fn list_indexes(
Ok(indexes) Ok(indexes)
} }
pub async fn list_foreign_keys( pub async fn list_foreign_keys(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
pool: &SqlitePool,
_schema: &str,
table: &str,
) -> Result<Vec<ForeignKeyInfo>, String> {
let rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA foreign_key_list(\"{}\")", table)) let rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA foreign_key_list(\"{}\")", table))
.fetch_all(pool) .fetch_all(pool)
.await .await
@@ -152,18 +122,13 @@ pub async fn list_foreign_keys(
.collect()) .collect())
} }
pub async fn list_triggers( pub async fn list_triggers(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
pool: &SqlitePool, let rows: Vec<SqliteRow> =
_schema: &str, sqlx::query("SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND tbl_name = ? ORDER BY name")
table: &str, .bind(table)
) -> Result<Vec<TriggerInfo>, String> { .fetch_all(pool)
let rows: Vec<SqliteRow> = sqlx::query( .await
"SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND tbl_name = ? ORDER BY name", .map_err(|e| e.to_string())?;
)
.bind(table)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
@@ -184,11 +149,7 @@ pub async fn list_triggers(
} else { } else {
"DELETE" "DELETE"
}; };
TriggerInfo { TriggerInfo { name: row.get::<String, _>("name"), event: event.to_string(), timing: timing.to_string() }
name: row.get::<String, _>("name"),
event: event.to_string(),
timing: timing.to_string(),
}
}) })
.collect()) .collect())
} }
@@ -203,16 +164,9 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
|| trimmed.starts_with("WITH") || trimmed.starts_with("WITH")
{ {
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?; let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
let columns: Vec<String> = desc let columns: Vec<String> = desc.columns().iter().map(|c| c.name().to_string()).collect();
.columns()
.iter()
.map(|c| c.name().to_string())
.collect();
let rows: Vec<SqliteRow> = sqlx::query(sql) let rows: Vec<SqliteRow> = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let result_rows: Vec<Vec<serde_json::Value>> = rows let result_rows: Vec<Vec<serde_json::Value>> = rows
.iter() .iter()
@@ -221,10 +175,7 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
.map(|i| { .map(|i| {
row.try_get::<String, _>(i) row.try_get::<String, _>(i)
.map(serde_json::Value::String) .map(serde_json::Value::String)
.or_else(|_| { .or_else(|_| row.try_get::<i64, _>(i).map(|v| serde_json::Value::Number(v.into())))
row.try_get::<i64, _>(i)
.map(|v| serde_json::Value::Number(v.into()))
})
.or_else(|_| { .or_else(|_| {
row.try_get::<f64, _>(i).map(|v| { row.try_get::<f64, _>(i).map(|v| {
serde_json::Number::from_f64(v) serde_json::Number::from_f64(v)
@@ -247,10 +198,7 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
truncated: false, truncated: false,
}) })
} else { } else {
let result = sqlx::query(sql) let result = sqlx::query(sql).execute(pool).await.map_err(|e| e.to_string())?;
.execute(pool)
.await
.map_err(|e| e.to_string())?;
Ok(QueryResult { Ok(QueryResult {
columns: vec![], columns: vec![],
+27 -87
View File
@@ -5,9 +5,7 @@ use tokio::net::TcpStream;
use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt}; use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt};
use super::{connection_timeout, CONNECTION_TIMEOUT_SECS}; use super::{connection_timeout, CONNECTION_TIMEOUT_SECS};
use crate::types::{ use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
};
pub type SqlServerClient = Client<Compat<TcpStream>>; pub type SqlServerClient = Client<Compat<TcpStream>>;
@@ -48,13 +46,10 @@ async fn try_connect(
.await .await
.map_err(|_| format!("SQL Server connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? .map_err(|_| format!("SQL Server connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.map_err(|e| format!("SQL Server connection failed: {e}"))?; .map_err(|e| format!("SQL Server connection failed: {e}"))?;
tokio::time::timeout( tokio::time::timeout(connection_timeout(), Client::connect(config, tcp.compat_write()))
connection_timeout(), .await
Client::connect(config, tcp.compat_write()), .map_err(|_| format!("SQL Server handshake timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
) .map_err(|e| format!("SQL Server connection failed: {e}"))
.await
.map_err(|_| format!("SQL Server handshake timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.map_err(|e| format!("SQL Server connection failed: {e}"))
} }
fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> { fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> {
@@ -69,9 +64,7 @@ fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> {
} else if let Some(v) = row.try_get::<i64, _>(i).ok().flatten() { } else if let Some(v) = row.try_get::<i64, _>(i).ok().flatten() {
serde_json::Value::Number(v.into()) serde_json::Value::Number(v.into())
} else if let Some(v) = row.try_get::<f64, _>(i).ok().flatten() { } else if let Some(v) = row.try_get::<f64, _>(i).ok().flatten() {
serde_json::Number::from_f64(v) serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
.map(serde_json::Value::Number)
.unwrap_or(serde_json::Value::Null)
} else if let Some(v) = row.try_get::<bool, _>(i).ok().flatten() { } else if let Some(v) = row.try_get::<bool, _>(i).ok().flatten() {
serde_json::Value::Bool(v) serde_json::Value::Bool(v)
} else { } else {
@@ -82,20 +75,9 @@ fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> {
} }
pub async fn list_databases(client: &mut SqlServerClient) -> Result<Vec<DatabaseInfo>, String> { pub async fn list_databases(client: &mut SqlServerClient) -> Result<Vec<DatabaseInfo>, String> {
let stream = client let stream = client.query("SELECT name FROM sys.databases ORDER BY name", &[]).await.map_err(|e| e.to_string())?;
.query("SELECT name FROM sys.databases ORDER BY name", &[]) let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.await Ok(rows.iter().map(|row| DatabaseInfo { name: row.get::<&str, _>(0).unwrap_or("").to_string() }).collect())
.map_err(|e| e.to_string())?;
let rows = stream
.into_first_result()
.await
.map_err(|e| e.to_string())?;
Ok(rows
.iter()
.map(|row| DatabaseInfo {
name: row.get::<&str, _>(0).unwrap_or("").to_string(),
})
.collect())
} }
pub async fn list_schemas(client: &mut SqlServerClient) -> Result<Vec<String>, String> { pub async fn list_schemas(client: &mut SqlServerClient) -> Result<Vec<String>, String> {
@@ -108,29 +90,17 @@ pub async fn list_schemas(client: &mut SqlServerClient) -> Result<Vec<String>, S
) )
.await .await
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
let rows = stream let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.into_first_result() Ok(rows.iter().map(|row| row.get::<&str, _>(0).unwrap_or("").to_string()).collect())
.await
.map_err(|e| e.to_string())?;
Ok(rows
.iter()
.map(|row| row.get::<&str, _>(0).unwrap_or("").to_string())
.collect())
} }
pub async fn list_tables( pub async fn list_tables(client: &mut SqlServerClient, schema: &str) -> Result<Vec<TableInfo>, String> {
client: &mut SqlServerClient,
schema: &str,
) -> Result<Vec<TableInfo>, String> {
let sql = format!( let sql = format!(
"SELECT TABLE_NAME, TABLE_TYPE FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA = '{}' ORDER BY TABLE_NAME", "SELECT TABLE_NAME, TABLE_TYPE FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA = '{}' ORDER BY TABLE_NAME",
schema.replace('\'', "''") schema.replace('\'', "''")
); );
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.into_first_result()
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
.map(|row| TableInfo { .map(|row| TableInfo {
@@ -140,11 +110,7 @@ pub async fn list_tables(
.collect()) .collect())
} }
pub async fn get_columns( pub async fn get_columns(client: &mut SqlServerClient, schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
client: &mut SqlServerClient,
schema: &str,
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
let sql = format!( let sql = format!(
"SELECT c.COLUMN_NAME, c.DATA_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, \ "SELECT c.COLUMN_NAME, c.DATA_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, \
CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \ CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \
@@ -158,10 +124,7 @@ pub async fn get_columns(
s = schema.replace('\'', "''"), t = table.replace('\'', "''") s = schema.replace('\'', "''"), t = table.replace('\'', "''")
); );
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.into_first_result()
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
.map(|row| { .map(|row| {
@@ -216,11 +179,7 @@ pub async fn get_columns(
.collect()) .collect())
} }
pub async fn list_indexes( pub async fn list_indexes(client: &mut SqlServerClient, schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
client: &mut SqlServerClient,
schema: &str,
table: &str,
) -> Result<Vec<IndexInfo>, String> {
let sql = format!( let sql = format!(
"SELECT i.name, \ "SELECT i.name, \
STRING_AGG(CASE WHEN ic.is_included_column = 0 THEN c.name END, ',') WITHIN GROUP (ORDER BY ic.key_ordinal) AS columns, \ STRING_AGG(CASE WHEN ic.is_included_column = 0 THEN c.name END, ',') WITHIN GROUP (ORDER BY ic.key_ordinal) AS columns, \
@@ -236,10 +195,7 @@ pub async fn list_indexes(
s = schema.replace('\'', "''"), t = table.replace('\'', "''") s = schema.replace('\'', "''"), t = table.replace('\'', "''")
); );
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.into_first_result()
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
.map(|row| { .map(|row| {
@@ -247,11 +203,7 @@ pub async fn list_indexes(
let inc_str = row.get::<&str, _>(5).unwrap_or(""); let inc_str = row.get::<&str, _>(5).unwrap_or("");
IndexInfo { IndexInfo {
name: row.get::<&str, _>(0).unwrap_or("").to_string(), name: row.get::<&str, _>(0).unwrap_or("").to_string(),
columns: cols_str columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(),
.split(',')
.filter(|s| !s.is_empty())
.map(|s| s.to_string())
.collect(),
is_unique: row.get::<bool, _>(2).unwrap_or(false), is_unique: row.get::<bool, _>(2).unwrap_or(false),
is_primary: row.get::<bool, _>(3).unwrap_or(false), is_primary: row.get::<bool, _>(3).unwrap_or(false),
filter: row.get::<&str, _>(6).map(|s| s.to_string()), filter: row.get::<&str, _>(6).map(|s| s.to_string()),
@@ -281,13 +233,11 @@ pub async fn list_foreign_keys(
JOIN sys.columns rc ON fkc.referenced_object_id = rc.object_id AND fkc.referenced_column_id = rc.column_id \ JOIN sys.columns rc ON fkc.referenced_object_id = rc.object_id AND fkc.referenced_column_id = rc.column_id \
WHERE fk.parent_object_id = OBJECT_ID('{s}.{t}') \ WHERE fk.parent_object_id = OBJECT_ID('{s}.{t}') \
ORDER BY fk.name", ORDER BY fk.name",
s = schema.replace('\'', "''"), t = table.replace('\'', "''") s = schema.replace('\'', "''"),
t = table.replace('\'', "''")
); );
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.into_first_result()
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
.map(|row| ForeignKeyInfo { .map(|row| ForeignKeyInfo {
@@ -310,13 +260,11 @@ pub async fn list_triggers(
JOIN sys.trigger_events te ON t.object_id = te.object_id \ JOIN sys.trigger_events te ON t.object_id = te.object_id \
WHERE t.parent_id = OBJECT_ID('{s}.{t}') \ WHERE t.parent_id = OBJECT_ID('{s}.{t}') \
ORDER BY t.name", ORDER BY t.name",
s = schema.replace('\'', "''"), t = table.replace('\'', "''") s = schema.replace('\'', "''"),
t = table.replace('\'', "''")
); );
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
let rows = stream let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.into_first_result()
.await
.map_err(|e| e.to_string())?;
Ok(rows Ok(rows
.iter() .iter()
.map(|row| TriggerInfo { .map(|row| TriggerInfo {
@@ -341,19 +289,11 @@ pub async fn execute_query(client: &mut SqlServerClient, sql: &str) -> Result<Qu
.columns() .columns()
.await .await
.map_err(|e| e.to_string())? .map_err(|e| e.to_string())?
.map(|cols| { .map(|cols| cols.iter().map(|c| c.name().to_string()).collect::<Vec<_>>())
cols.iter()
.map(|c| c.name().to_string())
.collect::<Vec<_>>()
})
.unwrap_or_default(); .unwrap_or_default();
let rows = stream let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
.into_first_result() let result_rows: Vec<Vec<serde_json::Value>> = rows.iter().map(|row| row_to_json(row)).collect();
.await
.map_err(|e| e.to_string())?;
let result_rows: Vec<Vec<serde_json::Value>> =
rows.iter().map(|row| row_to_json(row)).collect();
Ok(QueryResult { Ok(QueryResult {
columns: columns_meta, columns: columns_meta,
+24 -66
View File
@@ -32,39 +32,24 @@ async fn connect_and_authenticate(
ssh_key_path: &str, ssh_key_path: &str,
ssh_key_passphrase: &str, ssh_key_passphrase: &str,
) -> Result<Handle<SshClient>, String> { ) -> Result<Handle<SshClient>, String> {
let config = Arc::new(Config { let config = Arc::new(Config { nodelay: true, ..Default::default() });
nodelay: true,
..Default::default()
});
let mut session = tokio::time::timeout( let mut session =
connection_timeout(), tokio::time::timeout(connection_timeout(), client::connect(config, (ssh_host, ssh_port), SshClient {}))
client::connect(config, (ssh_host, ssh_port), SshClient {}), .await
) .map_err(|_| format!("SSH connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.await .map_err(|e| format!("SSH connection failed: {e}"))?;
.map_err(|_| format!("SSH connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.map_err(|e| format!("SSH connection failed: {e}"))?;
if !ssh_key_path.is_empty() { if !ssh_key_path.is_empty() {
let passphrase = if ssh_key_passphrase.is_empty() { let passphrase = if ssh_key_passphrase.is_empty() { None } else { Some(ssh_key_passphrase) };
None let key_pair = load_secret_key(ssh_key_path, passphrase).map_err(|e| format!("Failed to load SSH key: {e}"))?;
} else {
Some(ssh_key_passphrase)
};
let key_pair = load_secret_key(ssh_key_path, passphrase)
.map_err(|e| format!("Failed to load SSH key: {e}"))?;
let auth_res = tokio::time::timeout( let auth_res = tokio::time::timeout(
connection_timeout(), connection_timeout(),
session.authenticate_publickey( session.authenticate_publickey(
ssh_user, ssh_user,
PrivateKeyWithHashAlg::new( PrivateKeyWithHashAlg::new(
Arc::new(key_pair), Arc::new(key_pair),
session session.best_supported_rsa_hash().await.ok().flatten().flatten(),
.best_supported_rsa_hash()
.await
.ok()
.flatten()
.flatten(),
), ),
), ),
) )
@@ -75,13 +60,11 @@ async fn connect_and_authenticate(
return Err("SSH public key authentication failed".to_string()); return Err("SSH public key authentication failed".to_string());
} }
} else if !ssh_password.is_empty() { } else if !ssh_password.is_empty() {
let auth_res = tokio::time::timeout( let auth_res =
connection_timeout(), tokio::time::timeout(connection_timeout(), session.authenticate_password(ssh_user, ssh_password))
session.authenticate_password(ssh_user, ssh_password), .await
) .map_err(|_| format!("SSH password auth timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.await .map_err(|e| format!("SSH password auth failed: {e}"))?;
.map_err(|_| format!("SSH password auth timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
.map_err(|e| format!("SSH password auth failed: {e}"))?;
if !auth_res.success() { if !auth_res.success() {
return Err("SSH password authentication failed".to_string()); return Err("SSH password authentication failed".to_string());
} }
@@ -92,12 +75,7 @@ async fn connect_and_authenticate(
Ok(session) Ok(session)
} }
async fn forward_loop( async fn forward_loop(session: Handle<SshClient>, listener: TcpListener, remote_host: String, remote_port: u16) {
session: Handle<SshClient>,
listener: TcpListener,
remote_host: String,
remote_port: u16,
) {
loop { loop {
let (mut stream, peer_addr) = match listener.accept().await { let (mut stream, peer_addr) = match listener.accept().await {
Ok(v) => v, Ok(v) => v,
@@ -163,9 +141,7 @@ pub struct TunnelManager {
impl TunnelManager { impl TunnelManager {
pub fn new() -> Self { pub fn new() -> Self {
Self { Self { tunnels: Mutex::new(HashMap::new()) }
tunnels: Mutex::new(HashMap::new()),
}
} }
pub async fn start_tunnel( pub async fn start_tunnel(
@@ -183,42 +159,24 @@ impl TunnelManager {
) -> Result<u16, String> { ) -> Result<u16, String> {
let local_port = portpicker::pick_unused_port().ok_or("No available port")?; let local_port = portpicker::pick_unused_port().ok_or("No available port")?;
let session = connect_and_authenticate( let session =
ssh_host, connect_and_authenticate(ssh_host, ssh_port, ssh_user, ssh_password, ssh_key_path, ssh_key_passphrase)
ssh_port, .await?;
ssh_user,
ssh_password,
ssh_key_path,
ssh_key_passphrase,
)
.await?;
let bind_addr = if expose_to_lan { let bind_addr = if expose_to_lan { "0.0.0.0" } else { "127.0.0.1" };
"0.0.0.0" let listener =
} else { TcpListener::bind((bind_addr, local_port)).await.map_err(|e| format!("Failed to bind local port: {e}"))?;
"127.0.0.1"
};
let listener = TcpListener::bind((bind_addr, local_port))
.await
.map_err(|e| format!("Failed to bind local port: {e}"))?;
let remote_host = remote_host.to_string(); let remote_host = remote_host.to_string();
let handle = tokio::spawn(forward_loop(session, listener, remote_host, remote_port)); let handle = tokio::spawn(forward_loop(session, listener, remote_host, remote_port));
self.tunnels self.tunnels.lock().await.insert(connection_id.to_string(), (handle, local_port));
.lock()
.await
.insert(connection_id.to_string(), (handle, local_port));
Ok(local_port) Ok(local_port)
} }
pub async fn local_port(&self, connection_id: &str) -> Option<u16> { pub async fn local_port(&self, connection_id: &str) -> Option<u16> {
self.tunnels self.tunnels.lock().await.get(connection_id).map(|(_, port)| *port)
.lock()
.await
.get(connection_id)
.map(|(_, port)| *port)
} }
pub async fn stop_tunnel(&self, connection_id: &str) { pub async fn stop_tunnel(&self, connection_id: &str) {
+1 -4
View File
@@ -45,9 +45,6 @@ pub fn clear_history_entries(path: &Path) -> Result<(), String> {
} }
pub fn delete_history_entry_by_id(path: &Path, id: &str) -> Result<(), String> { pub fn delete_history_entry_by_id(path: &Path, id: &str) -> Result<(), String> {
let entries: Vec<HistoryEntry> = read_all(path)? let entries: Vec<HistoryEntry> = read_all(path)?.into_iter().filter(|e| e.id != id).collect();
.into_iter()
.filter(|e| e.id != id)
.collect();
write_all(path, &entries) write_all(path, &entries)
} }
+41 -34
View File
@@ -73,7 +73,10 @@ pub enum DatabaseType {
impl ConnectionConfig { impl ConnectionConfig {
pub fn needs_bare_mysql(&self) -> bool { pub fn needs_bare_mysql(&self) -> bool {
matches!(self.db_type, DatabaseType::Doris | DatabaseType::StarRocks) matches!(self.db_type, DatabaseType::Doris | DatabaseType::StarRocks)
|| self.driver_profile.as_deref().map(|p| p.to_lowercase()) || self
.driver_profile
.as_deref()
.map(|p| p.to_lowercase())
.is_some_and(|p| matches!(p.as_str(), "doris" | "starrocks" | "selectdb" | "tdengine")) .is_some_and(|p| matches!(p.as_str(), "doris" | "starrocks" | "selectdb" | "tdengine"))
} }
@@ -103,20 +106,17 @@ impl ConnectionConfig {
let scheme = if self.ssl { "rediss" } else { "redis" }; let scheme = if self.ssl { "rediss" } else { "redis" };
format!("{scheme}://{host}:{port}/") format!("{scheme}://{host}:{port}/")
} }
DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => format!("mysql://{host}:{port}{db_part}?{params}"), DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => {
format!("mysql://{host}:{port}{db_part}?{params}")
}
DatabaseType::Postgres | DatabaseType::Redshift => { DatabaseType::Postgres | DatabaseType::Redshift => {
let suffix = if params.is_empty() { let suffix = if params.is_empty() { String::new() } else { format!("?{params}") };
String::new()
} else {
format!("?{params}")
};
format!("postgres://{host}:{port}{db_part}{suffix}") format!("postgres://{host}:{port}{db_part}{suffix}")
} }
DatabaseType::ClickHouse => format!("http://{host}:{port}"), DatabaseType::ClickHouse => format!("http://{host}:{port}"),
DatabaseType::SqlServer => format!( DatabaseType::SqlServer => {
"server=tcp:{host},{port};database={}", format!("server=tcp:{host},{port};database={}", self.database.as_deref().unwrap_or("master"))
self.database.as_deref().unwrap_or("master") }
),
DatabaseType::MongoDb => { DatabaseType::MongoDb => {
if let Some(cs) = self.connection_string.as_deref().filter(|s| !s.is_empty()) { if let Some(cs) = self.connection_string.as_deref().filter(|s| !s.is_empty()) {
return cs.to_string(); return cs.to_string();
@@ -154,20 +154,12 @@ impl ConnectionConfig {
format!("{scheme}://{username}:{password}@{host}:{port}/") format!("{scheme}://{username}:{password}@{host}:{port}/")
} }
} }
DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => format!( DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => {
"mysql://{}:{}@{host}:{port}{db_part}?{params}", format!("mysql://{}:{}@{host}:{port}{db_part}?{params}", username, password)
username, password }
),
DatabaseType::Postgres | DatabaseType::Redshift => { DatabaseType::Postgres | DatabaseType::Redshift => {
let suffix = if params.is_empty() { let suffix = if params.is_empty() { String::new() } else { format!("?{params}") };
String::new() format!("postgres://{}:{}@{host}:{port}{db_part}{suffix}", username, password)
} else {
format!("?{params}")
};
format!(
"postgres://{}:{}@{host}:{port}{db_part}{suffix}",
username, password
)
} }
DatabaseType::ClickHouse => format!("http://{host}:{port}"), DatabaseType::ClickHouse => format!("http://{host}:{port}"),
DatabaseType::SqlServer => format!( DatabaseType::SqlServer => format!(
@@ -197,10 +189,15 @@ impl ConnectionConfig {
let value = self.url_params.as_deref().unwrap_or("").trim(); let value = self.url_params.as_deref().unwrap_or("").trim();
if self.needs_bare_mysql() { if self.needs_bare_mysql() {
let v = value.trim_start_matches('?'); let v = value.trim_start_matches('?');
let filtered: Vec<&str> = v.split('&') let filtered: Vec<&str> = v
.split('&')
.filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred")) .filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred"))
.collect(); .collect();
return if filtered.is_empty() { "ssl-mode=disabled".to_string() } else { format!("ssl-mode=disabled&{}", filtered.join("&")) }; return if filtered.is_empty() {
"ssl-mode=disabled".to_string()
} else {
format!("ssl-mode=disabled&{}", filtered.join("&"))
};
} }
match self.db_type { match self.db_type {
DatabaseType::Mysql => { DatabaseType::Mysql => {
@@ -209,18 +206,31 @@ impl ConnectionConfig {
base.to_string() base.to_string()
} else if value.contains("ssl-mode=") { } else if value.contains("ssl-mode=") {
let v = value.trim_start_matches('?'); let v = value.trim_start_matches('?');
if v.contains("charset=") { v.to_string() } else { format!("{v}&charset=utf8mb4") } if v.contains("charset=") {
v.to_string()
} else {
format!("{v}&charset=utf8mb4")
}
} else { } else {
let v = value.trim_start_matches('?'); let v = value.trim_start_matches('?');
if v.contains("charset=") { format!("ssl-mode=preferred&{v}") } else { format!("{base}&{v}") } if v.contains("charset=") {
format!("ssl-mode=preferred&{v}")
} else {
format!("{base}&{v}")
}
} }
} }
DatabaseType::Doris | DatabaseType::StarRocks => { DatabaseType::Doris | DatabaseType::StarRocks => {
let v = value.trim_start_matches('?'); let v = value.trim_start_matches('?');
let filtered: Vec<&str> = v.split('&') let filtered: Vec<&str> = v
.split('&')
.filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred")) .filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred"))
.collect(); .collect();
if filtered.is_empty() { "ssl-mode=disabled".to_string() } else { format!("ssl-mode=disabled&{}", filtered.join("&")) } if filtered.is_empty() {
"ssl-mode=disabled".to_string()
} else {
format!("ssl-mode=disabled&{}", filtered.join("&"))
}
} }
DatabaseType::Postgres | DatabaseType::Redshift => value.trim_start_matches('?').to_string(), DatabaseType::Postgres | DatabaseType::Redshift => value.trim_start_matches('?').to_string(),
_ => value.trim_start_matches('?').to_string(), _ => value.trim_start_matches('?').to_string(),
@@ -308,10 +318,7 @@ mod tests {
config.db_type = DatabaseType::Postgres; config.db_type = DatabaseType::Postgres;
config.url_params = Some("sslmode=disable".to_string()); config.url_params = Some("sslmode=disable".to_string());
assert_eq!( assert_eq!(config.connection_url(), "postgres://postgres:secret@10.1.2.3:2883/test?sslmode=disable");
config.connection_url(),
"postgres://postgres:secret@10.1.2.3:2883/test?sslmode=disable"
);
} }
#[test] #[test]
+6 -17
View File
@@ -1,11 +1,8 @@
use crate::connection::{AppState, PoolKind}; use crate::connection::{AppState, PoolKind};
use crate::db::mongo_driver::{self, MongoDocumentResult};
use crate::db::elasticsearch_driver; use crate::db::elasticsearch_driver;
use crate::db::mongo_driver::{self, MongoDocumentResult};
pub async fn mongo_list_databases_core( pub async fn mongo_list_databases_core(state: &AppState, connection_id: &str) -> Result<Vec<String>, String> {
state: &AppState,
connection_id: &str,
) -> Result<Vec<String>, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::MongoDb(client) => mongo_driver::list_databases(client).await, PoolKind::MongoDb(client) => mongo_driver::list_databases(client).await,
@@ -37,9 +34,7 @@ pub async fn mongo_find_documents_core(
) -> Result<MongoDocumentResult, String> { ) -> Result<MongoDocumentResult, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::MongoDb(client) => { PoolKind::MongoDb(client) => mongo_driver::find_documents(client, database, collection, skip, limit).await,
mongo_driver::find_documents(client, database, collection, skip, limit).await
}
PoolKind::Elasticsearch(client) => { PoolKind::Elasticsearch(client) => {
let client = client.clone(); let client = client.clone();
drop(connections); drop(connections);
@@ -58,9 +53,7 @@ pub async fn mongo_insert_document_core(
) -> Result<String, String> { ) -> Result<String, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::MongoDb(client) => { PoolKind::MongoDb(client) => mongo_driver::insert_document(client, database, collection, doc_json).await,
mongo_driver::insert_document(client, database, collection, doc_json).await
}
PoolKind::Elasticsearch(client) => { PoolKind::Elasticsearch(client) => {
let client = client.clone(); let client = client.clone();
drop(connections); drop(connections);
@@ -80,9 +73,7 @@ pub async fn mongo_update_document_core(
) -> Result<u64, String> { ) -> Result<u64, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::MongoDb(client) => { PoolKind::MongoDb(client) => mongo_driver::update_document(client, database, collection, id, doc_json).await,
mongo_driver::update_document(client, database, collection, id, doc_json).await
}
PoolKind::Elasticsearch(client) => { PoolKind::Elasticsearch(client) => {
let client = client.clone(); let client = client.clone();
drop(connections); drop(connections);
@@ -101,9 +92,7 @@ pub async fn mongo_delete_document_core(
) -> Result<u64, String> { ) -> Result<u64, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::MongoDb(client) => { PoolKind::MongoDb(client) => mongo_driver::delete_document(client, database, collection, id).await,
mongo_driver::delete_document(client, database, collection, id).await
}
PoolKind::Elasticsearch(client) => { PoolKind::Elasticsearch(client) => {
let client = client.clone(); let client = client.clone();
drop(connections); drop(connections);
+52 -50
View File
@@ -15,8 +15,12 @@ pub fn duckdb_execute(con: &duckdb::Connection, sql: &str) -> Result<db::QueryRe
let start = std::time::Instant::now(); let start = std::time::Instant::now();
let trimmed = sql.trim().to_uppercase(); let trimmed = sql.trim().to_uppercase();
if trimmed.starts_with("SELECT") || trimmed.starts_with("SHOW") || trimmed.starts_with("DESCRIBE") if trimmed.starts_with("SELECT")
|| trimmed.starts_with("EXPLAIN") || trimmed.starts_with("WITH") || trimmed.starts_with("PRAGMA") || trimmed.starts_with("SHOW")
|| trimmed.starts_with("DESCRIBE")
|| trimmed.starts_with("EXPLAIN")
|| trimmed.starts_with("WITH")
|| trimmed.starts_with("PRAGMA")
{ {
let mut stmt = con.prepare(sql).map_err(|e| e.to_string())?; let mut stmt = con.prepare(sql).map_err(|e| e.to_string())?;
let mut rows = stmt.query([]).map_err(|e| e.to_string())?; let mut rows = stmt.query([]).map_err(|e| e.to_string())?;
@@ -28,27 +32,45 @@ pub fn duckdb_execute(con: &duckdb::Connection, sql: &str) -> Result<db::QueryRe
let mut result_rows = Vec::new(); let mut result_rows = Vec::new();
while let Some(row) = rows.next().map_err(|e| e.to_string())? { while let Some(row) = rows.next().map_err(|e| e.to_string())? {
if result_rows.len() >= MAX_ROWS { break; } if result_rows.len() >= MAX_ROWS {
let vals: Vec<serde_json::Value> = (0..col_count).map(|i| { break;
row.get::<_, String>(i) }
.map(serde_json::Value::String) let vals: Vec<serde_json::Value> = (0..col_count)
.or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into()))) .map(|i| {
.or_else(|_| row.get::<_, f64>(i).map(|v| { row.get::<_, String>(i)
serde_json::Number::from_f64(v) .map(serde_json::Value::String)
.map(serde_json::Value::Number) .or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into())))
.unwrap_or(serde_json::Value::Null) .or_else(|_| {
})) row.get::<_, f64>(i).map(|v| {
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool)) serde_json::Number::from_f64(v)
.unwrap_or(serde_json::Value::Null) .map(serde_json::Value::Number)
}).collect(); .unwrap_or(serde_json::Value::Null)
})
})
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool))
.unwrap_or(serde_json::Value::Null)
})
.collect();
result_rows.push(vals); result_rows.push(vals);
} }
let truncated = result_rows.len() >= MAX_ROWS; let truncated = result_rows.len() >= MAX_ROWS;
Ok(db::QueryResult { columns, rows: result_rows, affected_rows: 0, execution_time_ms: start.elapsed().as_millis(), truncated }) Ok(db::QueryResult {
columns,
rows: result_rows,
affected_rows: 0,
execution_time_ms: start.elapsed().as_millis(),
truncated,
})
} else { } else {
let affected = con.execute(sql, []).map_err(|e| e.to_string())?; let affected = con.execute(sql, []).map_err(|e| e.to_string())?;
Ok(db::QueryResult { columns: vec![], rows: vec![], affected_rows: affected as u64, execution_time_ms: start.elapsed().as_millis(), truncated: false }) Ok(db::QueryResult {
columns: vec![],
rows: vec![],
affected_rows: affected as u64,
execution_time_ms: start.elapsed().as_millis(),
truncated: false,
})
} }
} }
@@ -79,16 +101,10 @@ pub fn canceled_error() -> String {
} }
pub fn is_canceled(cancel_token: &Option<CancellationToken>) -> bool { pub fn is_canceled(cancel_token: &Option<CancellationToken>) -> bool {
cancel_token cancel_token.as_ref().map(|token| token.is_cancelled()).unwrap_or(false)
.as_ref()
.map(|token| token.is_cancelled())
.unwrap_or(false)
} }
pub async fn wait_for_query<F>( pub async fn wait_for_query<F>(cancel_token: Option<CancellationToken>, future: F) -> Result<db::QueryResult, String>
cancel_token: Option<CancellationToken>,
future: F,
) -> Result<db::QueryResult, String>
where where
F: Future<Output = Result<db::QueryResult, String>>, F: Future<Output = Result<db::QueryResult, String>>,
{ {
@@ -110,9 +126,7 @@ where
result = timeout(timeout_duration, future) => result.map_err(|_| timeout_error())?, result = timeout(timeout_duration, future) => result.map_err(|_| timeout_error())?,
} }
} else { } else {
timeout(timeout_duration, future) timeout(timeout_duration, future).await.map_err(|_| timeout_error())?
.await
.map_err(|_| timeout_error())?
} }
} }
@@ -143,23 +157,17 @@ pub async fn do_execute(
let p = p.clone(); let p = p.clone();
let bare = *bare; let bare = *bare;
drop(connections); drop(connections);
wait_for_query(cancel_token, db::mysql::execute_query(&p, sql, bare)) wait_for_query(cancel_token, db::mysql::execute_query(&p, sql, bare)).await.map(truncate_result)
.await
.map(truncate_result)
} }
PoolKind::Postgres(p) => { PoolKind::Postgres(p) => {
let p = p.clone(); let p = p.clone();
drop(connections); drop(connections);
wait_for_query(cancel_token, db::postgres::execute_query(&p, sql)) wait_for_query(cancel_token, db::postgres::execute_query(&p, sql)).await.map(truncate_result)
.await
.map(truncate_result)
} }
PoolKind::Sqlite(p) => { PoolKind::Sqlite(p) => {
let p = p.clone(); let p = p.clone();
drop(connections); drop(connections);
wait_for_query(cancel_token, db::sqlite::execute_query(&p, sql)) wait_for_query(cancel_token, db::sqlite::execute_query(&p, sql)).await.map(truncate_result)
.await
.map(truncate_result)
} }
PoolKind::ClickHouse(client) => { PoolKind::ClickHouse(client) => {
let client = client.clone(); let client = client.clone();
@@ -180,9 +188,7 @@ pub async fn do_execute(
}, },
None => client.lock().await, None => client.lock().await,
}; };
wait_for_query(cancel_token, db::sqlserver::execute_query(&mut client, sql)) wait_for_query(cancel_token, db::sqlserver::execute_query(&mut client, sql)).await.map(truncate_result)
.await
.map(truncate_result)
} }
PoolKind::Oracle(client) => { PoolKind::Oracle(client) => {
let client = client.clone(); let client = client.clone();
@@ -195,9 +201,7 @@ pub async fn do_execute(
}, },
None => client.lock().await, None => client.lock().await,
}; };
wait_for_query(cancel_token, db::oracle_driver::execute_query(&*client, sql)) wait_for_query(cancel_token, db::oracle_driver::execute_query(&*client, sql)).await.map(truncate_result)
.await
.map(truncate_result)
} }
PoolKind::Elasticsearch(_) => Err("Use document browser for Elasticsearch".to_string()), PoolKind::Elasticsearch(_) => Err("Use document browser for Elasticsearch".to_string()),
PoolKind::Redis(_) => Err("Use Redis-specific commands".to_string()), PoolKind::Redis(_) => Err("Use Redis-specific commands".to_string()),
@@ -244,9 +248,7 @@ pub async fn execute_multi_core(
let statements = split_sql_statements(sql); let statements = split_sql_statements(sql);
if statements.len() <= 1 { if statements.len() <= 1 {
let single_sql = statements.into_iter().next().unwrap_or_default(); let single_sql = statements.into_iter().next().unwrap_or_default();
let result = execute_sql_statement( let result = execute_sql_statement(state, connection_id, database, &single_sql, cancel_token).await?;
state, connection_id, database, &single_sql, cancel_token,
).await?;
return Ok(vec![result]); return Ok(vec![result]);
} }
@@ -262,9 +264,7 @@ pub async fn execute_multi_core(
}); });
break; break;
} }
match execute_sql_statement( match execute_sql_statement(state, connection_id, database, stmt, cancel_token.clone()).await {
state, connection_id, database, stmt, cancel_token.clone(),
).await {
Ok(r) => results.push(r), Ok(r) => results.push(r),
Err(e) => { Err(e) => {
results.push(db::QueryResult { results.push(db::QueryResult {
@@ -308,7 +308,9 @@ pub async fn execute_statements(
} }
return Err(format!( return Err(format!(
"Statement {} failed: {}. Previous {} statement(s) may have been committed.", "Statement {} failed: {}. Previous {} statement(s) may have been committed.",
i + 1, e, i i + 1,
e,
i
)); ));
} }
} }
+5 -23
View File
@@ -10,25 +10,13 @@ pub struct RunningQueries {
impl RunningQueries { impl RunningQueries {
pub fn register(&self, execution_id: String) -> RegisteredQuery { pub fn register(&self, execution_id: String) -> RegisteredQuery {
let token = CancellationToken::new(); let token = CancellationToken::new();
self.inner self.inner.lock().expect("running query registry poisoned").insert(execution_id.clone(), token.clone());
.lock()
.expect("running query registry poisoned")
.insert(execution_id.clone(), token.clone());
RegisteredQuery { RegisteredQuery { execution_id, token, running_queries: self.clone() }
execution_id,
token,
running_queries: self.clone(),
}
} }
pub fn cancel(&self, execution_id: &str) -> bool { pub fn cancel(&self, execution_id: &str) -> bool {
let token = self let token = self.inner.lock().expect("running query registry poisoned").get(execution_id).cloned();
.inner
.lock()
.expect("running query registry poisoned")
.get(execution_id)
.cloned();
if let Some(token) = token { if let Some(token) = token {
token.cancel(); token.cancel();
@@ -40,17 +28,11 @@ impl RunningQueries {
#[cfg(test)] #[cfg(test)]
pub fn has(&self, execution_id: &str) -> bool { pub fn has(&self, execution_id: &str) -> bool {
self.inner self.inner.lock().expect("running query registry poisoned").contains_key(execution_id)
.lock()
.expect("running query registry poisoned")
.contains_key(execution_id)
} }
fn remove(&self, execution_id: &str) { fn remove(&self, execution_id: &str) {
self.inner self.inner.lock().expect("running query registry poisoned").remove(execution_id);
.lock()
.expect("running query registry poisoned")
.remove(execution_id);
} }
} }
+6 -32
View File
@@ -1,10 +1,7 @@
use crate::connection::{AppState, PoolKind}; use crate::connection::{AppState, PoolKind};
use crate::db::redis_driver::{self, RedisScanResult, RedisValue}; use crate::db::redis_driver::{self, RedisScanResult, RedisValue};
pub async fn redis_list_databases_core( pub async fn redis_list_databases_core(state: &AppState, connection_id: &str) -> Result<Vec<u32>, String> {
state: &AppState,
connection_id: &str,
) -> Result<Vec<u32>, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
let pool = connections.get(connection_id).ok_or("Connection not found")?; let pool = connections.get(connection_id).ok_or("Connection not found")?;
match pool { match pool {
@@ -36,11 +33,7 @@ pub async fn redis_scan_keys_core(
} }
} }
pub async fn redis_get_value_core( pub async fn redis_get_value_core(state: &AppState, connection_id: &str, key: &str) -> Result<RedisValue, String> {
state: &AppState,
connection_id: &str,
key: &str,
) -> Result<RedisValue, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
let pool = connections.get(connection_id).ok_or("Connection not found")?; let pool = connections.get(connection_id).ok_or("Connection not found")?;
match pool { match pool {
@@ -70,11 +63,7 @@ pub async fn redis_set_string_core(
} }
} }
pub async fn redis_delete_key_core( pub async fn redis_delete_key_core(state: &AppState, connection_id: &str, key: &str) -> Result<(), String> {
state: &AppState,
connection_id: &str,
key: &str,
) -> Result<(), String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
let pool = connections.get(connection_id).ok_or("Connection not found")?; let pool = connections.get(connection_id).ok_or("Connection not found")?;
match pool { match pool {
@@ -100,12 +89,7 @@ pub async fn redis_hash_set_core(
} }
} }
pub async fn redis_hash_del_core( pub async fn redis_hash_del_core(state: &AppState, connection_id: &str, key: &str, field: &str) -> Result<(), String> {
state: &AppState,
connection_id: &str,
key: &str,
field: &str,
) -> Result<(), String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::Redis(con) => redis_driver::hash_del(&mut *con.lock().await, key, field).await, PoolKind::Redis(con) => redis_driver::hash_del(&mut *con.lock().await, key, field).await,
@@ -113,12 +97,7 @@ pub async fn redis_hash_del_core(
} }
} }
pub async fn redis_list_push_core( pub async fn redis_list_push_core(state: &AppState, connection_id: &str, key: &str, value: &str) -> Result<(), String> {
state: &AppState,
connection_id: &str,
key: &str,
value: &str,
) -> Result<(), String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::Redis(con) => redis_driver::list_push(&mut *con.lock().await, key, value).await, PoolKind::Redis(con) => redis_driver::list_push(&mut *con.lock().await, key, value).await,
@@ -139,12 +118,7 @@ pub async fn redis_list_remove_core(
} }
} }
pub async fn redis_set_add_core( pub async fn redis_set_add_core(state: &AppState, connection_id: &str, key: &str, member: &str) -> Result<(), String> {
state: &AppState,
connection_id: &str,
key: &str,
member: &str,
) -> Result<(), String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
match connections.get(connection_id).ok_or("Not found")? { match connections.get(connection_id).ok_or("Not found")? {
PoolKind::Redis(con) => redis_driver::set_add(&mut *con.lock().await, key, member).await, PoolKind::Redis(con) => redis_driver::set_add(&mut *con.lock().await, key, member).await,
+165 -80
View File
@@ -8,18 +8,16 @@ pub fn duckdb_query_tables(con: &duckdb::Connection) -> Result<Vec<db::TableInfo
let mut stmt = con.prepare( let mut stmt = con.prepare(
"SELECT table_name, table_type FROM information_schema.tables WHERE table_schema = 'main' ORDER BY table_name" "SELECT table_name, table_type FROM information_schema.tables WHERE table_schema = 'main' ORDER BY table_name"
).map_err(|e| e.to_string())?; ).map_err(|e| e.to_string())?;
let rows = stmt.query_map([], |row| { let rows = stmt
Ok(db::TableInfo { .query_map([], |row| Ok(db::TableInfo { name: row.get::<_, String>(0)?, table_type: row.get::<_, String>(1)? }))
name: row.get::<_, String>(0)?, .map_err(|e| e.to_string())?;
table_type: row.get::<_, String>(1)?,
})
}).map_err(|e| e.to_string())?;
Ok(rows.filter_map(|r| r.ok()).collect()) Ok(rows.filter_map(|r| r.ok()).collect())
} }
pub fn duckdb_query_columns(con: &duckdb::Connection, table: &str) -> Result<Vec<db::ColumnInfo>, String> { pub fn duckdb_query_columns(con: &duckdb::Connection, table: &str) -> Result<Vec<db::ColumnInfo>, String> {
let mut pk_stmt = con.prepare( let mut pk_stmt = con
"SELECT kcu.column_name .prepare(
"SELECT kcu.column_name
FROM information_schema.table_constraints tc FROM information_schema.table_constraints tc
JOIN information_schema.key_column_usage kcu JOIN information_schema.key_column_usage kcu
ON tc.constraint_name = kcu.constraint_name ON tc.constraint_name = kcu.constraint_name
@@ -28,67 +26,81 @@ pub fn duckdb_query_columns(con: &duckdb::Connection, table: &str) -> Result<Vec
WHERE tc.constraint_type = 'PRIMARY KEY' WHERE tc.constraint_type = 'PRIMARY KEY'
AND tc.table_schema = 'main' AND tc.table_schema = 'main'
AND tc.table_name = ? AND tc.table_name = ?
ORDER BY kcu.ordinal_position" ORDER BY kcu.ordinal_position",
).map_err(|e| e.to_string())?; )
let pk_rows = pk_stmt.query_map([table], |row| row.get::<_, String>(0))
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
let pk_rows = pk_stmt.query_map([table], |row| row.get::<_, String>(0)).map_err(|e| e.to_string())?;
let primary_keys: std::collections::HashSet<String> = pk_rows.filter_map(|r| r.ok()).collect(); let primary_keys: std::collections::HashSet<String> = pk_rows.filter_map(|r| r.ok()).collect();
let mut stmt = con.prepare( let mut stmt = con
"SELECT column_name, data_type, is_nullable, column_default .prepare(
"SELECT column_name, data_type, is_nullable, column_default
FROM information_schema.columns FROM information_schema.columns
WHERE table_schema = 'main' AND table_name = ? WHERE table_schema = 'main' AND table_name = ?
ORDER BY ordinal_position" ORDER BY ordinal_position",
).map_err(|e| e.to_string())?; )
let rows = stmt.query_map([table], |row| { .map_err(|e| e.to_string())?;
let name = row.get::<_, String>(0)?; let rows = stmt
Ok(db::ColumnInfo { .query_map([table], |row| {
is_primary_key: primary_keys.contains(&name), let name = row.get::<_, String>(0)?;
name, Ok(db::ColumnInfo {
data_type: row.get::<_, String>(1)?, is_primary_key: primary_keys.contains(&name),
is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES", name,
column_default: row.get::<_, Option<String>>(3)?, data_type: row.get::<_, String>(1)?,
extra: None, comment: None, is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES",
numeric_precision: None, column_default: row.get::<_, Option<String>>(3)?,
numeric_scale: None, extra: None,
character_maximum_length: None, comment: None,
numeric_precision: None,
numeric_scale: None,
character_maximum_length: None,
})
}) })
}).map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
Ok(rows.filter_map(|r| r.ok()).collect()) Ok(rows.filter_map(|r| r.ok()).collect())
} }
pub fn extract_duckdb(connections: &HashMap<String, PoolKind>, key: &str) -> Option<Arc<std::sync::Mutex<duckdb::Connection>>> { pub fn extract_duckdb(
connections: &HashMap<String, PoolKind>,
key: &str,
) -> Option<Arc<std::sync::Mutex<duckdb::Connection>>> {
match connections.get(key)? { match connections.get(key)? {
PoolKind::DuckDb(con) => Some(con.clone()), PoolKind::DuckDb(con) => Some(con.clone()),
_ => None, _ => None,
} }
} }
pub fn extract_sqlserver(connections: &HashMap<String, PoolKind>, key: &str) -> Option<Arc<tokio::sync::Mutex<db::sqlserver::SqlServerClient>>> { pub fn extract_sqlserver(
connections: &HashMap<String, PoolKind>,
key: &str,
) -> Option<Arc<tokio::sync::Mutex<db::sqlserver::SqlServerClient>>> {
match connections.get(key)? { match connections.get(key)? {
PoolKind::SqlServer(client) => Some(client.clone()), PoolKind::SqlServer(client) => Some(client.clone()),
_ => None, _ => None,
} }
} }
pub fn extract_clickhouse(connections: &HashMap<String, PoolKind>, key: &str) -> Option<db::clickhouse_driver::ChClient> { pub fn extract_clickhouse(
connections: &HashMap<String, PoolKind>,
key: &str,
) -> Option<db::clickhouse_driver::ChClient> {
match connections.get(key)? { match connections.get(key)? {
PoolKind::ClickHouse(client) => Some(client.clone()), PoolKind::ClickHouse(client) => Some(client.clone()),
_ => None, _ => None,
} }
} }
pub fn extract_oracle(connections: &HashMap<String, PoolKind>, key: &str) -> Option<Arc<tokio::sync::Mutex<db::oracle_driver::OracleClient>>> { pub fn extract_oracle(
connections: &HashMap<String, PoolKind>,
key: &str,
) -> Option<Arc<tokio::sync::Mutex<db::oracle_driver::OracleClient>>> {
match connections.get(key)? { match connections.get(key)? {
PoolKind::Oracle(client) => Some(client.clone()), PoolKind::Oracle(client) => Some(client.clone()),
_ => None, _ => None,
} }
} }
pub async fn list_databases_core( pub async fn list_databases_core(state: &AppState, connection_id: &str) -> Result<Vec<db::DatabaseInfo>, String> {
state: &AppState,
connection_id: &str,
) -> Result<Vec<db::DatabaseInfo>, String> {
{ {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
if let Some(client) = extract_clickhouse(&connections, connection_id) { if let Some(client) = extract_clickhouse(&connections, connection_id) {
@@ -119,11 +131,7 @@ pub async fn list_databases_core(
} }
} }
pub async fn list_schemas_core( pub async fn list_schemas_core(state: &AppState, connection_id: &str, database: &str) -> Result<Vec<String>, String> {
state: &AppState,
connection_id: &str,
database: &str,
) -> Result<Vec<String>, String> {
let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?;
{ {
@@ -351,7 +359,8 @@ pub async fn get_table_ddl_core(
drop(connections); drop(connections);
let tbl = table.replace('\'', "''"); let tbl = table.replace('\'', "''");
let con = con.lock().map_err(|e| e.to_string())?; let con = con.lock().map_err(|e| e.to_string())?;
let mut stmt = con.prepare(&format!("SELECT sql FROM duckdb_tables() WHERE table_name = '{tbl}'")) let mut stmt = con
.prepare(&format!("SELECT sql FROM duckdb_tables() WHERE table_name = '{tbl}'"))
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
let mut rows = stmt.query([]).map_err(|e| e.to_string())?; let mut rows = stmt.query([]).map_err(|e| e.to_string())?;
if let Some(row) = rows.next().map_err(|e| e.to_string())? { if let Some(row) = rows.next().map_err(|e| e.to_string())? {
@@ -361,8 +370,12 @@ pub async fn get_table_ddl_core(
} }
if let Some(client) = extract_clickhouse(&connections, &pool_key) { if let Some(client) = extract_clickhouse(&connections, &pool_key) {
drop(connections); drop(connections);
let result = db::clickhouse_driver::execute_query(&client, database, &format!("SHOW CREATE TABLE `{table}`")).await?; let result =
return result.rows.first() db::clickhouse_driver::execute_query(&client, database, &format!("SHOW CREATE TABLE `{table}`"))
.await?;
return result
.rows
.first()
.and_then(|r| r.first()) .and_then(|r| r.first())
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.map(|s| s.to_string()) .map(|s| s.to_string())
@@ -394,8 +407,7 @@ pub async fn get_table_ddl_core(
pub async fn mysql_ddl(pool: &sqlx::mysql::MySqlPool, table: &str) -> Result<String, String> { pub async fn mysql_ddl(pool: &sqlx::mysql::MySqlPool, table: &str) -> Result<String, String> {
use sqlx::Row; use sqlx::Row;
let sql = format!("SHOW CREATE TABLE `{}`", table.replace('`', "``")); let sql = format!("SHOW CREATE TABLE `{}`", table.replace('`', "``"));
let row: sqlx::mysql::MySqlRow = sqlx::raw_sql(&sql) let row: sqlx::mysql::MySqlRow = sqlx::raw_sql(&sql).fetch_one(pool).await.map_err(|e| e.to_string())?;
.fetch_one(pool).await.map_err(|e| e.to_string())?;
row.try_get::<String, _>(1) row.try_get::<String, _>(1)
.or_else(|_| row.try_get::<Vec<u8>, _>(1).map(|b| String::from_utf8_lossy(&b).to_string())) .or_else(|_| row.try_get::<Vec<u8>, _>(1).map(|b| String::from_utf8_lossy(&b).to_string()))
.map_err(|e| e.to_string()) .map_err(|e| e.to_string())
@@ -405,7 +417,9 @@ pub async fn sqlite_ddl(pool: &sqlx::sqlite::SqlitePool, table: &str) -> Result<
use sqlx::Row; use sqlx::Row;
let row: sqlx::sqlite::SqliteRow = sqlx::query("SELECT sql FROM sqlite_master WHERE type='table' AND name=?") let row: sqlx::sqlite::SqliteRow = sqlx::query("SELECT sql FROM sqlite_master WHERE type='table' AND name=?")
.bind(table) .bind(table)
.fetch_one(pool).await.map_err(|e| e.to_string())?; .fetch_one(pool)
.await
.map_err(|e| e.to_string())?;
row.try_get::<String, _>(0).map_err(|e| e.to_string()) row.try_get::<String, _>(0).map_err(|e| e.to_string())
} }
@@ -415,31 +429,56 @@ pub async fn pg_ddl(pool: &sqlx::postgres::PgPool, schema: &str, table: &str) ->
let fkeys = db::postgres::list_foreign_keys(pool, schema, table).await?; let fkeys = db::postgres::list_foreign_keys(pool, schema, table).await?;
let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n"); let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n");
let col_lines: Vec<String> = columns.iter().map(|c| { let col_lines: Vec<String> = columns
let mut line = format!(" \"{}\" {}", c.name, c.data_type); .iter()
if !c.is_nullable { line.push_str(" NOT NULL"); } .map(|c| {
if let Some(ref def) = c.column_default { line.push_str(&format!(" DEFAULT {def}")); } let mut line = format!(" \"{}\" {}", c.name, c.data_type);
line if !c.is_nullable {
}).collect(); line.push_str(" NOT NULL");
}
if let Some(ref def) = c.column_default {
line.push_str(&format!(" DEFAULT {def}"));
}
line
})
.collect();
ddl.push_str(&col_lines.join(",\n")); ddl.push_str(&col_lines.join(",\n"));
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect(); let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
if !pks.is_empty() { if !pks.is_empty() {
ddl.push_str(&format!(",\n PRIMARY KEY ({})", pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", "))); ddl.push_str(&format!(
",\n PRIMARY KEY ({})",
pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", ")
));
} }
for fk in &fkeys { for fk in &fkeys {
ddl.push_str(&format!(",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")", fk.name, fk.column, fk.ref_table, fk.ref_column)); ddl.push_str(&format!(
",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")",
fk.name, fk.column, fk.ref_table, fk.ref_column
));
} }
ddl.push_str("\n);\n"); ddl.push_str("\n);\n");
for idx in &indexes { for idx in &indexes {
if idx.is_primary { continue; } if idx.is_primary {
continue;
}
let unique = if idx.is_unique { "UNIQUE " } else { "" }; let unique = if idx.is_unique { "UNIQUE " } else { "" };
let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", "); let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", ");
let using = idx.index_type.as_deref().map(|t| format!(" USING {t}")).unwrap_or_default(); let using = idx.index_type.as_deref().map(|t| format!(" USING {t}")).unwrap_or_default();
let include = idx.included_columns.as_deref().filter(|c| !c.is_empty()).map(|cols| format!(" INCLUDE ({})", cols.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", "))).unwrap_or_default(); let include = idx
.included_columns
.as_deref()
.filter(|c| !c.is_empty())
.map(|cols| {
format!(" INCLUDE ({})", cols.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", "))
})
.unwrap_or_default();
let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default(); let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default();
ddl.push_str(&format!("\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\"{using} ({cols}){include}{filter};", idx.name)); ddl.push_str(&format!(
"\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\"{using} ({cols}){include}{filter};",
idx.name
));
if let Some(ref c) = idx.comment { if let Some(ref c) = idx.comment {
ddl.push_str(&format!("\nCOMMENT ON INDEX \"{schema}\".\"{}\" IS '{}';", idx.name, c.replace('\'', "''"))); ddl.push_str(&format!("\nCOMMENT ON INDEX \"{schema}\".\"{}\" IS '{}';", idx.name, c.replace('\'', "''")));
} }
@@ -447,66 +486,112 @@ pub async fn pg_ddl(pool: &sqlx::postgres::PgPool, schema: &str, table: &str) ->
Ok(ddl) Ok(ddl)
} }
pub async fn build_sqlserver_ddl(client: &mut db::sqlserver::SqlServerClient, schema: &str, table: &str) -> Result<String, String> { pub async fn build_sqlserver_ddl(
client: &mut db::sqlserver::SqlServerClient,
schema: &str,
table: &str,
) -> Result<String, String> {
let columns = db::sqlserver::get_columns(client, schema, table).await?; let columns = db::sqlserver::get_columns(client, schema, table).await?;
let indexes = db::sqlserver::list_indexes(client, schema, table).await?; let indexes = db::sqlserver::list_indexes(client, schema, table).await?;
let fkeys = db::sqlserver::list_foreign_keys(client, schema, table).await?; let fkeys = db::sqlserver::list_foreign_keys(client, schema, table).await?;
let mut ddl = format!("CREATE TABLE [{schema}].[{table}] (\n"); let mut ddl = format!("CREATE TABLE [{schema}].[{table}] (\n");
let col_lines: Vec<String> = columns.iter().map(|c| { let col_lines: Vec<String> = columns
let mut line = format!(" [{}] {}", c.name, c.data_type); .iter()
if !c.is_nullable { line.push_str(" NOT NULL"); } .map(|c| {
if let Some(ref def) = c.column_default { line.push_str(&format!(" DEFAULT {def}")); } let mut line = format!(" [{}] {}", c.name, c.data_type);
line if !c.is_nullable {
}).collect(); line.push_str(" NOT NULL");
}
if let Some(ref def) = c.column_default {
line.push_str(&format!(" DEFAULT {def}"));
}
line
})
.collect();
ddl.push_str(&col_lines.join(",\n")); ddl.push_str(&col_lines.join(",\n"));
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect(); let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
if !pks.is_empty() { if !pks.is_empty() {
ddl.push_str(&format!(",\n PRIMARY KEY ({})", pks.iter().map(|k| format!("[{k}]")).collect::<Vec<_>>().join(", "))); ddl.push_str(&format!(
",\n PRIMARY KEY ({})",
pks.iter().map(|k| format!("[{k}]")).collect::<Vec<_>>().join(", ")
));
} }
for fk in &fkeys { for fk in &fkeys {
ddl.push_str(&format!(",\n CONSTRAINT [{}] FOREIGN KEY ([{}]) REFERENCES [{}]([{}])", fk.name, fk.column, fk.ref_table, fk.ref_column)); ddl.push_str(&format!(
",\n CONSTRAINT [{}] FOREIGN KEY ([{}]) REFERENCES [{}]([{}])",
fk.name, fk.column, fk.ref_table, fk.ref_column
));
} }
ddl.push_str("\n);\n"); ddl.push_str("\n);\n");
for idx in &indexes { for idx in &indexes {
if idx.is_primary { continue; } if idx.is_primary {
continue;
}
let unique = if idx.is_unique { "UNIQUE " } else { "" }; let unique = if idx.is_unique { "UNIQUE " } else { "" };
let idx_type = idx.index_type.as_deref().map(|t| format!("{t} ")).unwrap_or_default(); let idx_type = idx.index_type.as_deref().map(|t| format!("{t} ")).unwrap_or_default();
let cols = idx.columns.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", "); let cols = idx.columns.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", ");
let include = idx.included_columns.as_deref().filter(|c| !c.is_empty()).map(|cols| format!(" INCLUDE ({})", cols.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", "))).unwrap_or_default(); let include = idx
.included_columns
.as_deref()
.filter(|c| !c.is_empty())
.map(|cols| format!(" INCLUDE ({})", cols.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", ")))
.unwrap_or_default();
let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default(); let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default();
ddl.push_str(&format!("\nCREATE {unique}{idx_type}INDEX [{}] ON [{schema}].[{table}] ({cols}){include}{filter};", idx.name)); ddl.push_str(&format!(
"\nCREATE {unique}{idx_type}INDEX [{}] ON [{schema}].[{table}] ({cols}){include}{filter};",
idx.name
));
} }
Ok(ddl) Ok(ddl)
} }
pub async fn build_oracle_ddl(client: &db::oracle_driver::OracleClient, schema: &str, table: &str) -> Result<String, String> { pub async fn build_oracle_ddl(
client: &db::oracle_driver::OracleClient,
schema: &str,
table: &str,
) -> Result<String, String> {
let columns = db::oracle_driver::get_columns(client, schema, table).await?; let columns = db::oracle_driver::get_columns(client, schema, table).await?;
let indexes = db::oracle_driver::list_indexes(client, schema, table).await?; let indexes = db::oracle_driver::list_indexes(client, schema, table).await?;
let fkeys = db::oracle_driver::list_foreign_keys(client, schema, table).await?; let fkeys = db::oracle_driver::list_foreign_keys(client, schema, table).await?;
let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n"); let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n");
let col_lines: Vec<String> = columns.iter().map(|c| { let col_lines: Vec<String> = columns
let mut line = format!(" \"{}\" {}", c.name, c.data_type); .iter()
if !c.is_nullable { line.push_str(" NOT NULL"); } .map(|c| {
if let Some(ref def) = c.column_default { line.push_str(&format!(" DEFAULT {def}")); } let mut line = format!(" \"{}\" {}", c.name, c.data_type);
line if !c.is_nullable {
}).collect(); line.push_str(" NOT NULL");
}
if let Some(ref def) = c.column_default {
line.push_str(&format!(" DEFAULT {def}"));
}
line
})
.collect();
ddl.push_str(&col_lines.join(",\n")); ddl.push_str(&col_lines.join(",\n"));
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect(); let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
if !pks.is_empty() { if !pks.is_empty() {
ddl.push_str(&format!(",\n PRIMARY KEY ({})", pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", "))); ddl.push_str(&format!(
",\n PRIMARY KEY ({})",
pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", ")
));
} }
for fk in &fkeys { for fk in &fkeys {
ddl.push_str(&format!(",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")", fk.name, fk.column, fk.ref_table, fk.ref_column)); ddl.push_str(&format!(
",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")",
fk.name, fk.column, fk.ref_table, fk.ref_column
));
} }
ddl.push_str("\n);\n"); ddl.push_str("\n);\n");
for idx in &indexes { for idx in &indexes {
if idx.is_primary { continue; } if idx.is_primary {
continue;
}
let unique = if idx.is_unique { "UNIQUE " } else { "" }; let unique = if idx.is_unique { "UNIQUE " } else { "" };
let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", "); let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", ");
ddl.push_str(&format!("\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\" ({cols});", idx.name)); ddl.push_str(&format!("\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\" ({cols});", idx.name));
+6 -24
View File
@@ -147,17 +147,11 @@ impl SqlStatementSplitter {
} }
match ch { match ch {
'\'' if !self.in_double_quote '\'' if !self.in_double_quote && !self.in_backtick && self.previous != Some('\\') => {
&& !self.in_backtick
&& self.previous != Some('\\') =>
{
self.in_single_quote = !self.in_single_quote; self.in_single_quote = !self.in_single_quote;
self.buffer.push(ch); self.buffer.push(ch);
} }
'"' if !self.in_single_quote '"' if !self.in_single_quote && !self.in_backtick && self.previous != Some('\\') => {
&& !self.in_backtick
&& self.previous != Some('\\') =>
{
self.in_double_quote = !self.in_double_quote; self.in_double_quote = !self.in_double_quote;
self.buffer.push(ch); self.buffer.push(ch);
} }
@@ -351,10 +345,7 @@ mod tests {
let mut splitter = SqlStatementSplitter::default(); let mut splitter = SqlStatementSplitter::default();
assert_eq!(splitter.push_chunk("SELECT 1; -"), vec!["SELECT 1"]); assert_eq!(splitter.push_chunk("SELECT 1; -"), vec!["SELECT 1"]);
assert_eq!( assert_eq!(splitter.push_chunk("- comment ; ignored\nSELECT 2;"), vec!["-- comment ; ignored\nSELECT 2"]);
splitter.push_chunk("- comment ; ignored\nSELECT 2;"),
vec!["-- comment ; ignored\nSELECT 2"]
);
assert_eq!(splitter.finish(), Vec::<String>::new()); assert_eq!(splitter.finish(), Vec::<String>::new());
} }
@@ -363,10 +354,7 @@ mod tests {
let mut splitter = SqlStatementSplitter::default(); let mut splitter = SqlStatementSplitter::default();
assert_eq!(splitter.push_chunk("SELECT 1; /"), vec!["SELECT 1"]); assert_eq!(splitter.push_chunk("SELECT 1; /"), vec!["SELECT 1"]);
assert_eq!( assert_eq!(splitter.push_chunk("* comment ; ignored */\nSELECT 2;"), vec!["/* comment ; ignored */\nSELECT 2"]);
splitter.push_chunk("* comment ; ignored */\nSELECT 2;"),
vec!["/* comment ; ignored */\nSELECT 2"]
);
assert_eq!(splitter.finish(), Vec::<String>::new()); assert_eq!(splitter.finish(), Vec::<String>::new());
} }
@@ -402,14 +390,8 @@ mod tests {
#[test] #[test]
fn keeps_mysql_executable_comments_as_statements() { fn keeps_mysql_executable_comments_as_statements() {
assert_eq!( assert_eq!(
split_sql_script( split_sql_script("/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */;\nSELECT 1;",).unwrap(),
"/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */;\nSELECT 1;", vec!["/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */", "SELECT 1",]
)
.unwrap(),
vec![
"/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */",
"SELECT 1",
]
); );
} }
} }
+63 -149
View File
@@ -58,20 +58,12 @@ const SCHEMA_STATEMENTS: &[&str] = &[
impl Storage { impl Storage {
pub async fn open(db_path: &Path) -> Result<Self, String> { pub async fn open(db_path: &Path) -> Result<Self, String> {
let url = format!("sqlite:{}?mode=rwc", db_path.display()); let url = format!("sqlite:{}?mode=rwc", db_path.display());
let options = SqliteConnectOptions::from_str(&url) let options = SqliteConnectOptions::from_str(&url).map_err(|e| e.to_string())?.create_if_missing(true);
.map_err(|e| e.to_string())? let pool =
.create_if_missing(true); SqlitePoolOptions::new().max_connections(5).connect_with(options).await.map_err(|e| e.to_string())?;
let pool = SqlitePoolOptions::new()
.max_connections(5)
.connect_with(options)
.await
.map_err(|e| e.to_string())?;
for statement in SCHEMA_STATEMENTS { for statement in SCHEMA_STATEMENTS {
sqlx::query(statement) sqlx::query(statement).execute(&pool).await.map_err(|e| e.to_string())?;
.execute(&pool)
.await
.map_err(|e| e.to_string())?;
} }
Ok(Self { db: pool }) Ok(Self { db: pool })
@@ -125,11 +117,7 @@ impl Storage {
Ok(()) Ok(())
} }
pub async fn load_history_entries( pub async fn load_history_entries(&self, limit: usize, offset: usize) -> Result<Vec<HistoryEntry>, String> {
&self,
limit: usize,
offset: usize,
) -> Result<Vec<HistoryEntry>, String> {
let rows: Vec<HistoryRow> = sqlx::query_as( let rows: Vec<HistoryRow> = sqlx::query_as(
"SELECT id, connection_name, database, sql_text, executed_at, \ "SELECT id, connection_name, database, sql_text, executed_at, \
execution_time_ms, success, error \ execution_time_ms, success, error \
@@ -157,19 +145,12 @@ impl Storage {
} }
pub async fn clear_history(&self) -> Result<(), String> { pub async fn clear_history(&self) -> Result<(), String> {
sqlx::query("DELETE FROM history") sqlx::query("DELETE FROM history").execute(&self.db).await.map_err(|e| e.to_string())?;
.execute(&self.db)
.await
.map_err(|e| e.to_string())?;
Ok(()) Ok(())
} }
pub async fn delete_history_entry(&self, id: &str) -> Result<(), String> { pub async fn delete_history_entry(&self, id: &str) -> Result<(), String> {
sqlx::query("DELETE FROM history WHERE id = ?") sqlx::query("DELETE FROM history WHERE id = ?").bind(id).execute(&self.db).await.map_err(|e| e.to_string())?;
.bind(id)
.execute(&self.db)
.await
.map_err(|e| e.to_string())?;
Ok(()) Ok(())
} }
} }
@@ -190,15 +171,12 @@ impl Storage {
} }
pub async fn load_ai_config(&self) -> Result<Option<AiConfig>, String> { pub async fn load_ai_config(&self) -> Result<Option<AiConfig>, String> {
let row: Option<(String,)> = let row: Option<(String,)> = sqlx::query_as("SELECT config_json FROM ai_config WHERE id = 1")
sqlx::query_as("SELECT config_json FROM ai_config WHERE id = 1") .fetch_optional(&self.db)
.fetch_optional(&self.db) .await
.await .map_err(|e| e.to_string())?;
.map_err(|e| e.to_string())?;
match row { match row {
Some((json,)) => serde_json::from_str(&json) Some((json,)) => serde_json::from_str(&json).map(Some).map_err(|e| e.to_string()),
.map(Some)
.map_err(|e| e.to_string()),
None => Ok(None), None => Ok(None),
} }
} }
@@ -221,8 +199,7 @@ struct AiConversationRow {
impl Storage { impl Storage {
pub async fn save_ai_conversation(&self, conv: &AiConversation) -> Result<(), String> { pub async fn save_ai_conversation(&self, conv: &AiConversation) -> Result<(), String> {
let messages_json = let messages_json = serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
sqlx::query( sqlx::query(
"INSERT OR REPLACE INTO ai_conversations \ "INSERT OR REPLACE INTO ai_conversations \
(id, title, connection_name, database, messages_json, created_at, updated_at) \ (id, title, connection_name, database, messages_json, created_at, updated_at) \
@@ -263,8 +240,7 @@ impl Storage {
rows.into_iter() rows.into_iter()
.map(|r| { .map(|r| {
let messages: Vec<AiChatMessage> = let messages: Vec<AiChatMessage> = serde_json::from_str(&r.messages_json).map_err(|e| e.to_string())?;
serde_json::from_str(&r.messages_json).map_err(|e| e.to_string())?;
Ok(AiConversation { Ok(AiConversation {
id: r.id, id: r.id,
title: r.title, title: r.title,
@@ -296,10 +272,7 @@ impl Storage {
pub async fn save_connections(&self, configs: &[ConnectionConfig]) -> Result<(), String> { pub async fn save_connections(&self, configs: &[ConnectionConfig]) -> Result<(), String> {
let mut tx = self.db.begin().await.map_err(|e| e.to_string())?; let mut tx = self.db.begin().await.map_err(|e| e.to_string())?;
sqlx::query("DELETE FROM connections") sqlx::query("DELETE FROM connections").execute(&mut *tx).await.map_err(|e| e.to_string())?;
.execute(&mut *tx)
.await
.map_err(|e| e.to_string())?;
for config in configs { for config in configs {
// Store config without secrets // Store config without secrets
@@ -319,49 +292,31 @@ impl Storage {
// Store secrets // Store secrets
persist_secret_in_tx(&mut tx, &config.id, "password", &config.password).await?; persist_secret_in_tx(&mut tx, &config.id, "password", &config.password).await?;
persist_secret_in_tx(&mut tx, &config.id, "ssh_password", &config.ssh_password) persist_secret_in_tx(&mut tx, &config.id, "ssh_password", &config.ssh_password).await?;
.await?; persist_secret_in_tx(&mut tx, &config.id, "ssh_key_passphrase", &config.ssh_key_passphrase).await?;
persist_secret_in_tx(
&mut tx,
&config.id,
"ssh_key_passphrase",
&config.ssh_key_passphrase,
)
.await?;
if let Some(cs) = &config.connection_string { if let Some(cs) = &config.connection_string {
persist_secret_in_tx(&mut tx, &config.id, "connection_string", cs).await?; persist_secret_in_tx(&mut tx, &config.id, "connection_string", cs).await?;
} else { } else {
sqlx::query( sqlx::query("DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?")
"DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?", .bind(&config.id)
) .bind("connection_string")
.bind(&config.id) .execute(&mut *tx)
.bind("connection_string") .await
.execute(&mut *tx) .map_err(|e| e.to_string())?;
.await
.map_err(|e| e.to_string())?;
} }
} }
// Remove secrets for connections that no longer exist // Remove secrets for connections that no longer exist
if configs.is_empty() { if configs.is_empty() {
sqlx::query("DELETE FROM connection_secrets") sqlx::query("DELETE FROM connection_secrets").execute(&mut *tx).await.map_err(|e| e.to_string())?;
.execute(&mut *tx)
.await
.map_err(|e| e.to_string())?;
} else { } else {
let placeholders: Vec<&str> = configs.iter().map(|_| "?").collect(); let placeholders: Vec<&str> = configs.iter().map(|_| "?").collect();
let sql = format!( let sql = format!("DELETE FROM connection_secrets WHERE connection_id NOT IN ({})", placeholders.join(","));
"DELETE FROM connection_secrets WHERE connection_id NOT IN ({})",
placeholders.join(",")
);
let mut query = sqlx::query(&sql); let mut query = sqlx::query(&sql);
for config in configs { for config in configs {
query = query.bind(&config.id); query = query.bind(&config.id);
} }
query query.execute(&mut *tx).await.map_err(|e| e.to_string())?;
.execute(&mut *tx)
.await
.map_err(|e| e.to_string())?;
} }
tx.commit().await.map_err(|e| e.to_string())?; tx.commit().await.map_err(|e| e.to_string())?;
@@ -369,28 +324,17 @@ impl Storage {
} }
pub async fn load_connections(&self) -> Result<Vec<ConnectionConfig>, String> { pub async fn load_connections(&self) -> Result<Vec<ConnectionConfig>, String> {
let rows: Vec<(String, String)> = let rows: Vec<(String, String)> = sqlx::query_as("SELECT id, config_json FROM connections")
sqlx::query_as("SELECT id, config_json FROM connections") .fetch_all(&self.db)
.fetch_all(&self.db) .await
.await .map_err(|e| e.to_string())?;
.map_err(|e| e.to_string())?;
let mut configs = Vec::new(); let mut configs = Vec::new();
for (id, json) in rows { for (id, json) in rows {
let mut config: ConnectionConfig = let mut config: ConnectionConfig = serde_json::from_str(&json).map_err(|e| e.to_string())?;
serde_json::from_str(&json).map_err(|e| e.to_string())?; config.password = self.get_secret(&id, "password").await?.unwrap_or_default();
config.password = self config.ssh_password = self.get_secret(&id, "ssh_password").await?.unwrap_or_default();
.get_secret(&id, "password") config.ssh_key_passphrase = self.get_secret(&id, "ssh_key_passphrase").await?.unwrap_or_default();
.await?
.unwrap_or_default();
config.ssh_password = self
.get_secret(&id, "ssh_password")
.await?
.unwrap_or_default();
config.ssh_key_passphrase = self
.get_secret(&id, "ssh_key_passphrase")
.await?
.unwrap_or_default();
config.connection_string = self.get_secret(&id, "connection_string").await?; config.connection_string = self.get_secret(&id, "connection_string").await?;
configs.push(config); configs.push(config);
} }
@@ -403,28 +347,18 @@ impl Storage {
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
impl Storage { impl Storage {
pub async fn get_secret( pub async fn get_secret(&self, connection_id: &str, key: &str) -> Result<Option<String>, String> {
&self, let row: Option<(String,)> =
connection_id: &str, sqlx::query_as("SELECT secret FROM connection_secrets WHERE connection_id = ? AND key = ?")
key: &str, .bind(connection_id)
) -> Result<Option<String>, String> { .bind(key)
let row: Option<(String,)> = sqlx::query_as( .fetch_optional(&self.db)
"SELECT secret FROM connection_secrets WHERE connection_id = ? AND key = ?", .await
) .map_err(|e| e.to_string())?;
.bind(connection_id)
.bind(key)
.fetch_optional(&self.db)
.await
.map_err(|e| e.to_string())?;
Ok(row.map(|(s,)| s)) Ok(row.map(|(s,)| s))
} }
pub async fn set_secret( pub async fn set_secret(&self, connection_id: &str, key: &str, secret: &str) -> Result<(), String> {
&self,
connection_id: &str,
key: &str,
secret: &str,
) -> Result<(), String> {
sqlx::query( sqlx::query(
"INSERT OR REPLACE INTO connection_secrets (connection_id, key, secret) \ "INSERT OR REPLACE INTO connection_secrets (connection_id, key, secret) \
VALUES (?, ?, ?)", VALUES (?, ?, ?)",
@@ -438,11 +372,7 @@ impl Storage {
Ok(()) Ok(())
} }
pub async fn delete_secret( pub async fn delete_secret(&self, connection_id: &str, key: &str) -> Result<(), String> {
&self,
connection_id: &str,
key: &str,
) -> Result<(), String> {
sqlx::query("DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?") sqlx::query("DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?")
.bind(connection_id) .bind(connection_id)
.bind(key) .bind(key)
@@ -458,10 +388,7 @@ impl Storage {
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
impl Storage { impl Storage {
pub async fn save_sidebar_layout( pub async fn save_sidebar_layout(&self, layout: &serde_json::Value) -> Result<(), String> {
&self,
layout: &serde_json::Value,
) -> Result<(), String> {
let json = serde_json::to_string(layout).map_err(|e| e.to_string())?; let json = serde_json::to_string(layout).map_err(|e| e.to_string())?;
sqlx::query("INSERT OR REPLACE INTO sidebar_layout (id, layout_json) VALUES (1, ?)") sqlx::query("INSERT OR REPLACE INTO sidebar_layout (id, layout_json) VALUES (1, ?)")
.bind(&json) .bind(&json)
@@ -472,15 +399,12 @@ impl Storage {
} }
pub async fn load_sidebar_layout(&self) -> Result<Option<serde_json::Value>, String> { pub async fn load_sidebar_layout(&self) -> Result<Option<serde_json::Value>, String> {
let row: Option<(String,)> = let row: Option<(String,)> = sqlx::query_as("SELECT layout_json FROM sidebar_layout WHERE id = 1")
sqlx::query_as("SELECT layout_json FROM sidebar_layout WHERE id = 1") .fetch_optional(&self.db)
.fetch_optional(&self.db) .await
.await .map_err(|e| e.to_string())?;
.map_err(|e| e.to_string())?;
match row { match row {
Some((json,)) => serde_json::from_str(&json) Some((json,)) => serde_json::from_str(&json).map(Some).map_err(|e| e.to_string()),
.map(Some)
.map_err(|e| e.to_string()),
None => Ok(None), None => Ok(None),
} }
} }
@@ -509,8 +433,7 @@ impl Storage {
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?; let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
let configs: Vec<ConnectionConfig> = serde_json::from_str(&json).unwrap_or_default(); let configs: Vec<ConnectionConfig> = serde_json::from_str(&json).unwrap_or_default();
for config in &configs { for config in &configs {
let config_json = let config_json = serde_json::to_string(config).map_err(|e| e.to_string())?;
serde_json::to_string(config).map_err(|e| e.to_string())?;
sqlx::query("INSERT OR IGNORE INTO connections (id, config_json) VALUES (?, ?)") sqlx::query("INSERT OR IGNORE INTO connections (id, config_json) VALUES (?, ?)")
.bind(&config.id) .bind(&config.id)
.bind(&config_json) .bind(&config_json)
@@ -528,8 +451,7 @@ impl Storage {
return Ok(()); return Ok(());
} }
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?; let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
let secrets: std::collections::HashMap<String, String> = let secrets: std::collections::HashMap<String, String> = serde_json::from_str(&json).unwrap_or_default();
serde_json::from_str(&json).unwrap_or_default();
for (key, secret) in &secrets { for (key, secret) in &secrets {
// key format: "connection:{id}:{field}" // key format: "connection:{id}:{field}"
let parts: Vec<&str> = key.splitn(3, ':').collect(); let parts: Vec<&str> = key.splitn(3, ':').collect();
@@ -588,10 +510,7 @@ impl Storage {
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?; let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
// Only migrate if the table is empty // Only migrate if the table is empty
let count: (i64,) = let count: (i64,) =
sqlx::query_as("SELECT COUNT(*) FROM ai_config") sqlx::query_as("SELECT COUNT(*) FROM ai_config").fetch_one(&self.db).await.map_err(|e| e.to_string())?;
.fetch_one(&self.db)
.await
.map_err(|e| e.to_string())?;
if count.0 == 0 { if count.0 == 0 {
sqlx::query("INSERT OR IGNORE INTO ai_config (id, config_json) VALUES (1, ?)") sqlx::query("INSERT OR IGNORE INTO ai_config (id, config_json) VALUES (1, ?)")
.bind(&json) .bind(&json)
@@ -609,11 +528,9 @@ impl Storage {
return Ok(()); return Ok(());
} }
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?; let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
let conversations: Vec<AiConversation> = let conversations: Vec<AiConversation> = serde_json::from_str(&json).unwrap_or_default();
serde_json::from_str(&json).unwrap_or_default();
for conv in &conversations { for conv in &conversations {
let messages_json = let messages_json = serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
sqlx::query( sqlx::query(
"INSERT OR IGNORE INTO ai_conversations \ "INSERT OR IGNORE INTO ai_conversations \
(id, title, connection_name, database, messages_json, \ (id, title, connection_name, database, messages_json, \
@@ -641,19 +558,16 @@ impl Storage {
return Ok(()); return Ok(());
} }
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?; let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
let count: (i64,) = let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM sidebar_layout")
sqlx::query_as("SELECT COUNT(*) FROM sidebar_layout") .fetch_one(&self.db)
.fetch_one(&self.db)
.await
.map_err(|e| e.to_string())?;
if count.0 == 0 {
sqlx::query(
"INSERT OR IGNORE INTO sidebar_layout (id, layout_json) VALUES (1, ?)",
)
.bind(&json)
.execute(&self.db)
.await .await
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
if count.0 == 0 {
sqlx::query("INSERT OR IGNORE INTO sidebar_layout (id, layout_json) VALUES (1, ?)")
.bind(&json)
.execute(&self.db)
.await
.map_err(|e| e.to_string())?;
} }
std::fs::rename(&path, data_dir.join("sidebar_layout.json.bak")).ok(); std::fs::rename(&path, data_dir.join("sidebar_layout.json.bak")).ok();
Ok(()) Ok(())
+35 -150
View File
@@ -4,9 +4,9 @@ use std::path::Path;
use calamine::{open_workbook_auto, Data, Reader}; use calamine::{open_workbook_auto, Data, Reader};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use crate::connection::AppState;
use crate::models::connection::DatabaseType; use crate::models::connection::DatabaseType;
use crate::transfer::{execute_on_pool, generate_insert, qualified_table}; use crate::transfer::{execute_on_pool, generate_insert, qualified_table};
use crate::connection::AppState;
pub const DEFAULT_PREVIEW_LIMIT: usize = 50; pub const DEFAULT_PREVIEW_LIMIT: usize = 50;
pub const DEFAULT_BATCH_SIZE: usize = 500; pub const DEFAULT_BATCH_SIZE: usize = 500;
@@ -142,15 +142,8 @@ pub fn csv_value(value: &str) -> serde_json::Value {
} }
} }
pub fn parse_delimited_bytes( pub fn parse_delimited_bytes(bytes: &[u8], delimiter: u8, preview_limit: usize) -> Result<ParsedImportFile, String> {
bytes: &[u8], let mut reader = csv::ReaderBuilder::new().delimiter(delimiter).flexible(true).from_reader(bytes);
delimiter: u8,
preview_limit: usize,
) -> Result<ParsedImportFile, String> {
let mut reader = csv::ReaderBuilder::new()
.delimiter(delimiter)
.flexible(true)
.from_reader(bytes);
let columns = reader let columns = reader
.headers() .headers()
.map_err(|e| e.to_string())? .map_err(|e| e.to_string())?
@@ -172,21 +165,12 @@ pub fn parse_delimited_bytes(
} }
let mut row = Vec::with_capacity(columns.len()); let mut row = Vec::with_capacity(columns.len());
for index in 0..columns.len() { for index in 0..columns.len() {
row.push( row.push(record.get(index).map(csv_value).unwrap_or(serde_json::Value::Null));
record
.get(index)
.map(csv_value)
.unwrap_or(serde_json::Value::Null),
);
} }
rows.push(row); rows.push(row);
} }
Ok(ParsedImportFile { Ok(ParsedImportFile { columns, rows, total_rows })
columns,
rows,
total_rows,
})
} }
pub fn parse_csv_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImportFile, String> { pub fn parse_csv_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImportFile, String> {
@@ -229,25 +213,15 @@ pub fn parse_json_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImpo
.collect::<Vec<_>>() .collect::<Vec<_>>()
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
return Ok(ParsedImportFile { return Ok(ParsedImportFile { columns, rows, total_rows: items.len() });
columns,
rows,
total_rows: items.len(),
});
} }
if items.iter().all(|item| item.is_array()) { if items.iter().all(|item| item.is_array()) {
let max_cols = items let max_cols = items.iter().filter_map(|item| item.as_array().map(|row| row.len())).max().unwrap_or(0);
.iter()
.filter_map(|item| item.as_array().map(|row| row.len()))
.max()
.unwrap_or(0);
if max_cols == 0 { if max_cols == 0 {
return Err("Import file has no columns".to_string()); return Err("Import file has no columns".to_string());
} }
let columns = (0..max_cols) let columns = (0..max_cols).map(|index| format!("column_{}", index + 1)).collect::<Vec<_>>();
.map(|index| format!("column_{}", index + 1))
.collect::<Vec<_>>();
let rows = items let rows = items
.iter() .iter()
.take(preview_limit) .take(preview_limit)
@@ -258,11 +232,7 @@ pub fn parse_json_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImpo
.collect::<Vec<_>>() .collect::<Vec<_>>()
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
return Ok(ParsedImportFile { return Ok(ParsedImportFile { columns, rows, total_rows: items.len() });
columns,
rows,
total_rows: items.len(),
});
} }
Err("JSON rows must all be objects or all be arrays".to_string()) Err("JSON rows must all be objects or all be arrays".to_string())
@@ -272,9 +242,9 @@ pub fn xlsx_cell_value(cell: &Data) -> serde_json::Value {
match cell { match cell {
Data::Empty => serde_json::Value::Null, Data::Empty => serde_json::Value::Null,
Data::String(s) => csv_value(s), Data::String(s) => csv_value(s),
Data::Float(n) => serde_json::Number::from_f64(*n) Data::Float(n) => {
.map(serde_json::Value::Number) serde_json::Number::from_f64(*n).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
.unwrap_or(serde_json::Value::Null), }
Data::Int(n) => serde_json::Value::Number((*n).into()), Data::Int(n) => serde_json::Value::Number((*n).into()),
Data::Bool(v) => serde_json::Value::Bool(*v), Data::Bool(v) => serde_json::Value::Bool(*v),
Data::DateTime(v) => serde_json::Value::String(v.to_string()), Data::DateTime(v) => serde_json::Value::String(v.to_string()),
@@ -300,18 +270,10 @@ pub fn xlsx_cell_label(cell: &Data) -> String {
pub fn parse_xlsx_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> { pub fn parse_xlsx_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> {
let mut workbook = open_workbook_auto(path).map_err(|e| e.to_string())?; let mut workbook = open_workbook_auto(path).map_err(|e| e.to_string())?;
let sheet_name = workbook let sheet_name = workbook.sheet_names().first().cloned().ok_or_else(|| "Workbook has no sheets".to_string())?;
.sheet_names() let range = workbook.worksheet_range(&sheet_name).map_err(|e| e.to_string())?;
.first()
.cloned()
.ok_or_else(|| "Workbook has no sheets".to_string())?;
let range = workbook
.worksheet_range(&sheet_name)
.map_err(|e| e.to_string())?;
let mut rows_iter = range.rows(); let mut rows_iter = range.rows();
let header = rows_iter let header = rows_iter.next().ok_or_else(|| "Import file has no rows".to_string())?;
.next()
.ok_or_else(|| "Import file has no rows".to_string())?;
let columns = header let columns = header
.iter() .iter()
.enumerate() .enumerate()
@@ -330,21 +292,12 @@ pub fn parse_xlsx_file(path: &str, preview_limit: usize) -> Result<ParsedImportF
} }
let mut row = Vec::with_capacity(columns.len()); let mut row = Vec::with_capacity(columns.len());
for index in 0..columns.len() { for index in 0..columns.len() {
row.push( row.push(source_row.get(index).map(xlsx_cell_value).unwrap_or(serde_json::Value::Null));
source_row
.get(index)
.map(xlsx_cell_value)
.unwrap_or(serde_json::Value::Null),
);
} }
rows.push(row); rows.push(row);
} }
Ok(ParsedImportFile { Ok(ParsedImportFile { columns, rows, total_rows })
columns,
rows,
total_rows,
})
} }
pub fn parse_import_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> { pub fn parse_import_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> {
@@ -384,10 +337,7 @@ pub fn mapping_indexes(
return Err("Target column cannot be empty".to_string()); return Err("Target column cannot be empty".to_string());
} }
if !target_seen.insert(mapping.target_column.clone()) { if !target_seen.insert(mapping.target_column.clone()) {
return Err(format!( return Err(format!("Target column mapped more than once: {}", mapping.target_column));
"Target column mapped more than once: {}",
mapping.target_column
));
} }
mapped.push((source_index, mapping.target_column.clone())); mapped.push((source_index, mapping.target_column.clone()));
} }
@@ -403,10 +353,7 @@ pub fn build_import_insert_batches(
batch_size: usize, batch_size: usize,
) -> Result<Vec<ImportSqlBatch>, String> { ) -> Result<Vec<ImportSqlBatch>, String> {
let mapped = mapping_indexes(data, mappings)?; let mapped = mapping_indexes(data, mappings)?;
let columns = mapped let columns = mapped.iter().map(|(_, target)| target.clone()).collect::<Vec<_>>();
.iter()
.map(|(_, target)| target.clone())
.collect::<Vec<_>>();
let batch_size = batch_size.max(1); let batch_size = batch_size.max(1);
let mut batches = Vec::new(); let mut batches = Vec::new();
@@ -416,20 +363,13 @@ pub fn build_import_insert_batches(
.map(|row| { .map(|row| {
mapped mapped
.iter() .iter()
.map(|(source_index, _)| { .map(|(source_index, _)| row.get(*source_index).cloned().unwrap_or(serde_json::Value::Null))
row.get(*source_index)
.cloned()
.unwrap_or(serde_json::Value::Null)
})
.collect::<Vec<_>>() .collect::<Vec<_>>()
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let sql = generate_insert(&columns, &rows, table, schema, db_type); let sql = generate_insert(&columns, &rows, table, schema, db_type);
if !sql.trim().is_empty() { if !sql.trim().is_empty() {
batches.push(ImportSqlBatch { batches.push(ImportSqlBatch { sql, row_count: chunk.len() });
sql,
row_count: chunk.len(),
});
} }
} }
@@ -448,11 +388,7 @@ pub fn preview_table_import_file_core(file_path: &str) -> Result<TableImportPrev
let kind = import_file_kind(file_path)?; let kind = import_file_kind(file_path)?;
let parsed = parse_import_file(file_path, DEFAULT_PREVIEW_LIMIT)?; let parsed = parse_import_file(file_path, DEFAULT_PREVIEW_LIMIT)?;
let metadata = std::fs::metadata(file_path).map_err(|e| e.to_string())?; let metadata = std::fs::metadata(file_path).map_err(|e| e.to_string())?;
let file_name = Path::new(file_path) let file_name = Path::new(file_path).file_name().and_then(|name| name.to_str()).unwrap_or(file_path).to_string();
.file_name()
.and_then(|name| name.to_str())
.unwrap_or(file_path)
.to_string();
Ok(TableImportPreview { Ok(TableImportPreview {
file_name, file_name,
@@ -478,11 +414,7 @@ pub async fn import_table_file_core<F>(
where where
F: FnMut(TableImportProgress), F: FnMut(TableImportProgress),
{ {
let batch_size = if request.batch_size == 0 { let batch_size = if request.batch_size == 0 { DEFAULT_BATCH_SIZE } else { request.batch_size };
DEFAULT_BATCH_SIZE
} else {
request.batch_size
};
let parsed = match parse_import_file(&request.file_path, usize::MAX) { let parsed = match parse_import_file(&request.file_path, usize::MAX) {
Ok(parsed) => parsed, Ok(parsed) => parsed,
@@ -583,11 +515,7 @@ where
error: None, error: None,
}); });
Ok(TableImportSummary { Ok(TableImportSummary { import_id: request.import_id.clone(), rows_imported, total_rows })
import_id: request.import_id.clone(),
rows_imported,
total_rows,
})
} }
#[cfg(test)] #[cfg(test)]
@@ -627,81 +555,38 @@ mod tests {
assert_eq!(parsed.total_rows, 1); assert_eq!(parsed.total_rows, 1);
assert_eq!( assert_eq!(
parsed.rows[0], parsed.rows[0],
vec![ vec![serde_json::Value::String("1".to_string()), serde_json::Value::String("Ada".to_string()),]
serde_json::Value::String("1".to_string()),
serde_json::Value::String("Ada".to_string()),
]
); );
} }
#[test] #[test]
fn parses_json_array_objects_with_union_columns() { fn parses_json_array_objects_with_union_columns() {
let parsed = let parsed = parse_json_bytes(br#"[{"id":1,"name":"Ada"},{"id":2,"active":true}]"#, 10).unwrap();
parse_json_bytes(br#"[{"id":1,"name":"Ada"},{"id":2,"active":true}]"#, 10).unwrap();
assert_eq!(parsed.columns, vec!["id", "name", "active"]); assert_eq!(parsed.columns, vec!["id", "name", "active"]);
assert_eq!(parsed.total_rows, 2); assert_eq!(parsed.total_rows, 2);
assert_eq!( assert_eq!(parsed.rows[0], vec![serde_json::json!(1), serde_json::json!("Ada"), serde_json::Value::Null,]);
parsed.rows[0], assert_eq!(parsed.rows[1], vec![serde_json::json!(2), serde_json::Value::Null, serde_json::json!(true),]);
vec![
serde_json::json!(1),
serde_json::json!("Ada"),
serde_json::Value::Null,
]
);
assert_eq!(
parsed.rows[1],
vec![
serde_json::json!(2),
serde_json::Value::Null,
serde_json::json!(true),
]
);
} }
#[test] #[test]
fn builds_import_insert_batches_from_mapped_columns() { fn builds_import_insert_batches_from_mapped_columns() {
let mappings = vec![ let mappings = vec![
TableImportColumnMapping { TableImportColumnMapping { source_column: "id".to_string(), target_column: "user_id".to_string() },
source_column: "id".to_string(), TableImportColumnMapping { source_column: "name".to_string(), target_column: "display_name".to_string() },
target_column: "user_id".to_string(),
},
TableImportColumnMapping {
source_column: "name".to_string(),
target_column: "display_name".to_string(),
},
]; ];
let data = ParsedImportFile { let data = ParsedImportFile {
columns: vec!["id".to_string(), "name".to_string(), "ignored".to_string()], columns: vec!["id".to_string(), "name".to_string(), "ignored".to_string()],
rows: vec![ rows: vec![
vec![ vec![serde_json::json!(1), serde_json::json!("Ada"), serde_json::json!("x")],
serde_json::json!(1), vec![serde_json::json!(2), serde_json::json!("O'Hara"), serde_json::json!("y")],
serde_json::json!("Ada"), vec![serde_json::json!(3), serde_json::Value::Null, serde_json::json!("z")],
serde_json::json!("x"),
],
vec![
serde_json::json!(2),
serde_json::json!("O'Hara"),
serde_json::json!("y"),
],
vec![
serde_json::json!(3),
serde_json::Value::Null,
serde_json::json!("z"),
],
], ],
total_rows: 3, total_rows: 3,
}; };
let batches = build_import_insert_batches( let batches =
&data, build_import_insert_batches(&data, &mappings, "users", "public", &DatabaseType::Postgres, 2).unwrap();
&mappings,
"users",
"public",
&DatabaseType::Postgres,
2,
)
.unwrap();
assert_eq!(batches, vec![ assert_eq!(batches, vec![
ImportSqlBatch { ImportSqlBatch {
+110 -80
View File
@@ -1,5 +1,5 @@
use std::collections::HashSet;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use tokio::sync::RwLock; use tokio::sync::RwLock;
use crate::connection::{AppState, PoolKind}; use crate::connection::{AppState, PoolKind};
@@ -50,7 +50,9 @@ pub enum TransferStatus {
pub fn quote_identifier(name: &str, db_type: &DatabaseType) -> String { pub fn quote_identifier(name: &str, db_type: &DatabaseType) -> String {
match db_type { match db_type {
DatabaseType::Mysql | DatabaseType::ClickHouse | DatabaseType::Doris | DatabaseType::StarRocks => format!("`{}`", name.replace('`', "``")), DatabaseType::Mysql | DatabaseType::ClickHouse | DatabaseType::Doris | DatabaseType::StarRocks => {
format!("`{}`", name.replace('`', "``"))
}
DatabaseType::SqlServer => format!("[{}]", name.replace(']', "]]")), DatabaseType::SqlServer => format!("[{}]", name.replace(']', "]]")),
_ => format!("\"{}\"", name.replace('"', "\"\"")), _ => format!("\"{}\"", name.replace('"', "\"\"")),
} }
@@ -69,10 +71,24 @@ pub fn escape_value(val: &serde_json::Value, db_type: &DatabaseType) -> String {
match val { match val {
serde_json::Value::Null => "NULL".to_string(), serde_json::Value::Null => "NULL".to_string(),
serde_json::Value::Bool(b) => match db_type { serde_json::Value::Bool(b) => match db_type {
DatabaseType::Mysql | DatabaseType::Sqlite | DatabaseType::DuckDb | DatabaseType::Doris | DatabaseType::StarRocks => { DatabaseType::Mysql
if *b { "1".to_string() } else { "0".to_string() } | DatabaseType::Sqlite
| DatabaseType::DuckDb
| DatabaseType::Doris
| DatabaseType::StarRocks => {
if *b {
"1".to_string()
} else {
"0".to_string()
}
}
_ => {
if *b {
"TRUE".to_string()
} else {
"FALSE".to_string()
}
} }
_ => if *b { "TRUE".to_string() } else { "FALSE".to_string() },
}, },
serde_json::Value::Number(n) => n.to_string(), serde_json::Value::Number(n) => n.to_string(),
serde_json::Value::String(s) => { serde_json::Value::String(s) => {
@@ -160,20 +176,17 @@ pub fn map_column_type(source_type: &str, _source_db: &DatabaseType, target_db:
DatabaseType::Postgres => "TIMESTAMP".into(), DatabaseType::Postgres => "TIMESTAMP".into(),
_ => "DATETIME".into(), _ => "DATETIME".into(),
}, },
"timestamp" | "timestamptz" | "timestamp with time zone" "timestamp" | "timestamptz" | "timestamp with time zone" | "timestamp without time zone" => match target_db {
| "timestamp without time zone" => match target_db {
DatabaseType::Mysql => "DATETIME".into(), DatabaseType::Mysql => "DATETIME".into(),
DatabaseType::SqlServer => "DATETIME2".into(), DatabaseType::SqlServer => "DATETIME2".into(),
_ => "TIMESTAMP".into(), _ => "TIMESTAMP".into(),
}, },
"blob" | "longblob" | "mediumblob" | "tinyblob" | "binary" | "varbinary" | "image" => { "blob" | "longblob" | "mediumblob" | "tinyblob" | "binary" | "varbinary" | "image" => match target_db {
match target_db { DatabaseType::Postgres => "BYTEA".into(),
DatabaseType::Postgres => "BYTEA".into(), DatabaseType::Mysql => "BLOB".into(),
DatabaseType::Mysql => "BLOB".into(), DatabaseType::SqlServer => "VARBINARY(MAX)".into(),
DatabaseType::SqlServer => "VARBINARY(MAX)".into(), _ => "BLOB".into(),
_ => "BLOB".into(), },
}
}
"bytea" => match target_db { "bytea" => match target_db {
DatabaseType::Postgres => "BYTEA".into(), DatabaseType::Postgres => "BYTEA".into(),
DatabaseType::Mysql => "BLOB".into(), DatabaseType::Mysql => "BLOB".into(),
@@ -217,14 +230,13 @@ pub fn generate_create_table_ddl(
}) })
.collect(); .collect();
let pks: Vec<String> = columns let pks: Vec<String> =
.iter() columns.iter().filter(|c| c.is_primary_key).map(|c| quote_identifier(&c.name, target_db)).collect();
.filter(|c| c.is_primary_key)
.map(|c| quote_identifier(&c.name, target_db))
.collect();
let mut ddl = match target_db { let mut ddl = match target_db {
DatabaseType::SqlServer => format!("IF NOT EXISTS (SELECT * FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_NAME = '{table}')\n"), DatabaseType::SqlServer => {
format!("IF NOT EXISTS (SELECT * FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_NAME = '{table}')\n")
}
_ => String::new(), _ => String::new(),
}; };
@@ -261,11 +273,7 @@ pub fn generate_insert(
} }
let full_table = qualified_table(table, schema, db_type); let full_table = qualified_table(table, schema, db_type);
let col_list = columns let col_list = columns.iter().map(|c| quote_identifier(c, db_type)).collect::<Vec<_>>().join(", ");
.iter()
.map(|c| quote_identifier(c, db_type))
.collect::<Vec<_>>()
.join(", ");
let value_rows: Vec<String> = rows let value_rows: Vec<String> = rows
.iter() .iter()
@@ -287,11 +295,7 @@ pub fn pagination_sql(
limit: usize, limit: usize,
) -> String { ) -> String {
let full_table = qualified_table(table, schema, db_type); let full_table = qualified_table(table, schema, db_type);
let col_list = columns let col_list = columns.iter().map(|c| quote_identifier(c, db_type)).collect::<Vec<_>>().join(", ");
.iter()
.map(|c| quote_identifier(c, db_type))
.collect::<Vec<_>>()
.join(", ");
match db_type { match db_type {
DatabaseType::SqlServer | DatabaseType::Oracle => { DatabaseType::SqlServer | DatabaseType::Oracle => {
@@ -310,11 +314,7 @@ pub fn count_sql(table: &str, schema: &str, db_type: &DatabaseType) -> String {
format!("SELECT COUNT(*) FROM {full_table}") format!("SELECT COUNT(*) FROM {full_table}")
} }
pub async fn execute_on_pool( pub async fn execute_on_pool(state: &AppState, pool_key: &str, sql: &str) -> Result<db::QueryResult, String> {
state: &AppState,
pool_key: &str,
sql: &str,
) -> Result<db::QueryResult, String> {
let connections = state.connections.lock().await; let connections = state.connections.lock().await;
let pool = connections.get(pool_key).ok_or("Connection not found")?; let pool = connections.get(pool_key).ok_or("Connection not found")?;
@@ -361,8 +361,10 @@ pub async fn execute_on_pool(
let con = con.lock().map_err(|e| e.to_string())?; let con = con.lock().map_err(|e| e.to_string())?;
let start = std::time::Instant::now(); let start = std::time::Instant::now();
let trimmed = sql.trim().to_uppercase(); let trimmed = sql.trim().to_uppercase();
if trimmed.starts_with("SELECT") || trimmed.starts_with("SHOW") if trimmed.starts_with("SELECT")
|| trimmed.starts_with("DESCRIBE") || trimmed.starts_with("WITH") || trimmed.starts_with("SHOW")
|| trimmed.starts_with("DESCRIBE")
|| trimmed.starts_with("WITH")
|| trimmed.starts_with("PRAGMA") || trimmed.starts_with("PRAGMA")
{ {
let mut stmt = con.prepare(&sql).map_err(|e| e.to_string())?; let mut stmt = con.prepare(&sql).map_err(|e| e.to_string())?;
@@ -374,21 +376,40 @@ pub async fn execute_on_pool(
.collect(); .collect();
let mut result_rows = Vec::new(); let mut result_rows = Vec::new();
while let Some(row) = rows.next().map_err(|e| e.to_string())? { while let Some(row) = rows.next().map_err(|e| e.to_string())? {
let vals: Vec<serde_json::Value> = (0..col_count).map(|i| { let vals: Vec<serde_json::Value> = (0..col_count)
row.get::<_, String>(i).map(serde_json::Value::String) .map(|i| {
.or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into()))) row.get::<_, String>(i)
.or_else(|_| row.get::<_, f64>(i).map(|v| { .map(serde_json::Value::String)
serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null) .or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into())))
})) .or_else(|_| {
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool)) row.get::<_, f64>(i).map(|v| {
.unwrap_or(serde_json::Value::Null) serde_json::Number::from_f64(v)
}).collect(); .map(serde_json::Value::Number)
.unwrap_or(serde_json::Value::Null)
})
})
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool))
.unwrap_or(serde_json::Value::Null)
})
.collect();
result_rows.push(vals); result_rows.push(vals);
} }
Ok(db::QueryResult { columns, rows: result_rows, affected_rows: 0, execution_time_ms: start.elapsed().as_millis(), truncated: false }) Ok(db::QueryResult {
columns,
rows: result_rows,
affected_rows: 0,
execution_time_ms: start.elapsed().as_millis(),
truncated: false,
})
} else { } else {
let affected = con.execute(&sql, []).map_err(|e| e.to_string())?; let affected = con.execute(&sql, []).map_err(|e| e.to_string())?;
Ok(db::QueryResult { columns: vec![], rows: vec![], affected_rows: affected as u64, execution_time_ms: start.elapsed().as_millis(), truncated: false }) Ok(db::QueryResult {
columns: vec![],
rows: vec![],
affected_rows: affected as u64,
execution_time_ms: start.elapsed().as_millis(),
truncated: false,
})
} }
}) })
.await .await
@@ -422,28 +443,34 @@ pub async fn get_columns_for_transfer(
let table = table.to_string(); let table = table.to_string();
return tokio::task::spawn_blocking(move || { return tokio::task::spawn_blocking(move || {
let con = con.lock().map_err(|e| e.to_string())?; let con = con.lock().map_err(|e| e.to_string())?;
let mut stmt = con.prepare( let mut stmt = con
"SELECT column_name, data_type, is_nullable, column_default .prepare(
"SELECT column_name, data_type, is_nullable, column_default
FROM information_schema.columns FROM information_schema.columns
WHERE table_schema = 'main' AND table_name = ? WHERE table_schema = 'main' AND table_name = ?
ORDER BY ordinal_position" ORDER BY ordinal_position",
).map_err(|e| e.to_string())?; )
let rows = stmt.query_map([&table], |row| { .map_err(|e| e.to_string())?;
Ok(db::ColumnInfo { let rows = stmt
name: row.get::<_, String>(0)?, .query_map([&table], |row| {
data_type: row.get::<_, String>(1)?, Ok(db::ColumnInfo {
is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES", name: row.get::<_, String>(0)?,
column_default: row.get::<_, Option<String>>(3)?, data_type: row.get::<_, String>(1)?,
is_primary_key: false, is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES",
extra: None, column_default: row.get::<_, Option<String>>(3)?,
comment: None, is_primary_key: false,
numeric_precision: None, extra: None,
numeric_scale: None, comment: None,
character_maximum_length: None, numeric_precision: None,
numeric_scale: None,
character_maximum_length: None,
})
}) })
}).map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
Ok(rows.filter_map(|r| r.ok()).collect()) Ok(rows.filter_map(|r| r.ok()).collect())
}).await.map_err(|e| e.to_string())?; })
.await
.map_err(|e| e.to_string())?;
} }
if let Some(PoolKind::ClickHouse(client)) = connections.get(pool_key) { if let Some(PoolKind::ClickHouse(client)) = connections.get(pool_key) {
@@ -526,9 +553,14 @@ where
// Get source columns (deduplicate by name) // Get source columns (deduplicate by name)
let columns = { let columns = {
let raw = get_columns_for_transfer( let raw = get_columns_for_transfer(
state, source_pool_key, &request.source_connection_id, state,
&request.source_database, &request.source_schema, table, source_pool_key,
).await?; &request.source_connection_id,
&request.source_database,
&request.source_schema,
table,
)
.await?;
let mut seen = std::collections::HashSet::new(); let mut seen = std::collections::HashSet::new();
raw.into_iter().filter(|c| seen.insert(c.name.clone())).collect::<Vec<_>>() raw.into_iter().filter(|c| seen.insert(c.name.clone())).collect::<Vec<_>>()
}; };
@@ -544,13 +576,11 @@ where
let total_rows = { let total_rows = {
let sql = count_sql(table, &request.source_schema, source_db_type); let sql = count_sql(table, &request.source_schema, source_db_type);
match execute_on_pool(state, source_pool_key, &sql).await { match execute_on_pool(state, source_pool_key, &sql).await {
Ok(result) => result.rows.first() Ok(result) => result.rows.first().and_then(|r| r.first()).and_then(|v| match v {
.and_then(|r| r.first()) serde_json::Value::Number(n) => n.as_u64(),
.and_then(|v| match v { serde_json::Value::String(s) => s.parse::<u64>().ok(),
serde_json::Value::Number(n) => n.as_u64(), _ => None,
serde_json::Value::String(s) => s.parse::<u64>().ok(), }),
_ => None,
}),
Err(e) => { Err(e) => {
log::warn!("[transfer] count failed for {}: {}", table, e); log::warn!("[transfer] count failed for {}: {}", table, e);
None None
@@ -578,8 +608,7 @@ where
DatabaseType::Sqlite | DatabaseType::DuckDb => format!("DELETE FROM {full_table}"), DatabaseType::Sqlite | DatabaseType::DuckDb => format!("DELETE FROM {full_table}"),
_ => format!("TRUNCATE TABLE {full_table}"), _ => format!("TRUNCATE TABLE {full_table}"),
}; };
execute_on_pool(state, target_pool_key, &truncate_sql).await execute_on_pool(state, target_pool_key, &truncate_sql).await.map_err(|e| format!("Failed to truncate: {e}"))?;
.map_err(|e| format!("Failed to truncate: {e}"))?;
} }
// Transfer data in batches // Transfer data in batches
@@ -602,7 +631,8 @@ where
let insert_sql = generate_insert(&col_names, &result.rows, table, &request.target_schema, target_db_type); let insert_sql = generate_insert(&col_names, &result.rows, table, &request.target_schema, target_db_type);
if !insert_sql.is_empty() { if !insert_sql.is_empty() {
execute_on_pool(state, target_pool_key, &insert_sql).await execute_on_pool(state, target_pool_key, &insert_sql)
.await
.map_err(|e| format!("Insert failed at offset {offset}: {e}"))?; .map_err(|e| format!("Insert failed at offset {offset}: {e}"))?;
} }
+1 -4
View File
@@ -53,10 +53,7 @@ pub fn normalize_version(version: &str) -> String {
} }
pub fn parse_version(version: &str) -> Vec<u64> { pub fn parse_version(version: &str) -> Vec<u64> {
normalize_version(version) normalize_version(version).split(['.', '-', '+']).map(|part| part.parse::<u64>().unwrap_or(0)).collect()
.split(['.', '-', '+'])
.map(|part| part.parse::<u64>().unwrap_or(0))
.collect()
} }
pub fn is_newer_version(latest: &str, current: &str) -> bool { pub fn is_newer_version(latest: &str, current: &str) -> bool {
+3
View File
@@ -0,0 +1,3 @@
edition = "2021"
max_width = 120
use_small_heuristics = "Max"
+1 -1
View File
@@ -1,3 +1,3 @@
fn main() { fn main() {
tauri_build::build() tauri_build::build()
} }
+7 -24
View File
@@ -1,5 +1,5 @@
use std::sync::Arc; use std::sync::Arc;
use tauri::{Emitter, AppHandle, State}; use tauri::{AppHandle, Emitter, State};
use super::connection::AppState; use super::connection::AppState;
pub use dbx_core::ai::*; pub use dbx_core::ai::*;
@@ -10,17 +10,12 @@ pub async fn ai_test_connection(config: AiConfig) -> Result<String, String> {
} }
#[tauri::command] #[tauri::command]
pub async fn save_ai_config( pub async fn save_ai_config(state: State<'_, Arc<AppState>>, config: AiConfig) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
config: AiConfig,
) -> Result<(), String> {
state.storage.save_ai_config(&config).await state.storage.save_ai_config(&config).await
} }
#[tauri::command] #[tauri::command]
pub async fn load_ai_config( pub async fn load_ai_config(state: State<'_, Arc<AppState>>) -> Result<Option<AiConfig>, String> {
state: State<'_, Arc<AppState>>,
) -> Result<Option<AiConfig>, String> {
state.storage.load_ai_config().await state.storage.load_ai_config().await
} }
@@ -30,11 +25,7 @@ pub async fn ai_complete(request: AiCompletionRequest) -> Result<String, String>
} }
#[tauri::command] #[tauri::command]
pub async fn ai_stream( pub async fn ai_stream(app: AppHandle, session_id: String, request: AiCompletionRequest) -> Result<(), String> {
app: AppHandle,
session_id: String,
request: AiCompletionRequest,
) -> Result<(), String> {
let cancelled = dbx_core::ai::register_stream(&session_id).await; let cancelled = dbx_core::ai::register_stream(&session_id).await;
let result = dbx_core::ai::stream(&session_id, &request, &cancelled, |chunk| { let result = dbx_core::ai::stream(&session_id, &request, &cancelled, |chunk| {
@@ -52,24 +43,16 @@ pub async fn ai_cancel_stream(session_id: String) -> Result<bool, String> {
} }
#[tauri::command] #[tauri::command]
pub async fn save_ai_conversation( pub async fn save_ai_conversation(state: State<'_, Arc<AppState>>, conversation: AiConversation) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
conversation: AiConversation,
) -> Result<(), String> {
state.storage.save_ai_conversation(&conversation).await state.storage.save_ai_conversation(&conversation).await
} }
#[tauri::command] #[tauri::command]
pub async fn load_ai_conversations( pub async fn load_ai_conversations(state: State<'_, Arc<AppState>>) -> Result<Vec<AiConversation>, String> {
state: State<'_, Arc<AppState>>,
) -> Result<Vec<AiConversation>, String> {
state.storage.load_ai_conversations().await state.storage.load_ai_conversations().await
} }
#[tauri::command] #[tauri::command]
pub async fn delete_ai_conversation( pub async fn delete_ai_conversation(state: State<'_, Arc<AppState>>, id: String) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
id: String,
) -> Result<(), String> {
state.storage.delete_ai_conversation(&id).await state.storage.delete_ai_conversation(&id).await
} }
+62 -146
View File
@@ -2,74 +2,52 @@ use std::sync::Arc;
use tauri::State; use tauri::State;
pub use dbx_core::connection::{ pub use dbx_core::connection::{
connection_url_for_endpoint, expand_tilde, probe_connection_endpoint, connection_url_for_endpoint, expand_tilde, probe_connection_endpoint, redacted_connection_url_for_endpoint,
redacted_connection_url_for_endpoint, AppState, PoolKind, AppState, PoolKind,
}; };
use dbx_core::db; use dbx_core::db;
use dbx_core::models::connection::{ConnectionConfig, DatabaseType}; use dbx_core::models::connection::{ConnectionConfig, DatabaseType};
#[tauri::command] #[tauri::command]
pub async fn save_connections( pub async fn save_connections(state: State<'_, Arc<AppState>>, configs: Vec<ConnectionConfig>) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
configs: Vec<ConnectionConfig>,
) -> Result<(), String> {
state.storage.save_connections(&configs).await state.storage.save_connections(&configs).await
} }
#[tauri::command] #[tauri::command]
pub async fn load_connections( pub async fn load_connections(state: State<'_, Arc<AppState>>) -> Result<Vec<ConnectionConfig>, String> {
state: State<'_, Arc<AppState>>,
) -> Result<Vec<ConnectionConfig>, String> {
state.storage.load_connections().await state.storage.load_connections().await
} }
#[tauri::command] #[tauri::command]
pub async fn save_sidebar_layout( pub async fn save_sidebar_layout(state: State<'_, Arc<AppState>>, layout: serde_json::Value) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
layout: serde_json::Value,
) -> Result<(), String> {
state.storage.save_sidebar_layout(&layout).await state.storage.save_sidebar_layout(&layout).await
} }
#[tauri::command] #[tauri::command]
pub async fn load_sidebar_layout( pub async fn load_sidebar_layout(state: State<'_, Arc<AppState>>) -> Result<Option<serde_json::Value>, String> {
state: State<'_, Arc<AppState>>,
) -> Result<Option<serde_json::Value>, String> {
state.storage.load_sidebar_layout().await state.storage.load_sidebar_layout().await
} }
#[tauri::command] #[tauri::command]
pub async fn test_connection( pub async fn test_connection(state: State<'_, Arc<AppState>>, config: ConnectionConfig) -> Result<String, String> {
state: State<'_, Arc<AppState>>,
config: ConnectionConfig,
) -> Result<String, String> {
let tunnel_id = format!("{}:test", config.id); let tunnel_id = format!("{}:test", config.id);
let connection_id = if config.ssh_enabled && !config.ssh_host.is_empty() { let connection_id =
tunnel_id.as_str() if config.ssh_enabled && !config.ssh_host.is_empty() { tunnel_id.as_str() } else { config.id.as_str() };
} else {
config.id.as_str()
};
let (host, port) = state.connection_host_port(connection_id, &config).await?; let (host, port) = state.connection_host_port(connection_id, &config).await?;
let probe_result = probe_connection_endpoint(&config, &host, port).await; let probe_result = probe_connection_endpoint(&config, &host, port).await;
let url = connection_url_for_endpoint(&config, &host, port); let url = connection_url_for_endpoint(&config, &host, port);
let target = redacted_connection_url_for_endpoint(&config, &host, port); let target = redacted_connection_url_for_endpoint(&config, &host, port);
log::info!( log::info!("[test_connection] db_type={:?} target={}", config.db_type, target);
"[test_connection] db_type={:?} target={}",
config.db_type,
target
);
let result = match probe_result { let result = match probe_result {
Err(e) => Err(e), Err(e) => Err(e),
Ok(()) => match config.db_type { Ok(()) => match config.db_type {
DatabaseType::Mysql if config.needs_bare_mysql() => { DatabaseType::Mysql if config.needs_bare_mysql() => match db::mysql::connect_bare(&url).await {
match db::mysql::connect_bare(&url).await { Ok(pool) => {
Ok(pool) => { pool.close().await;
pool.close().await; Ok("Connection successful".to_string())
Ok("Connection successful".to_string())
}
Err(e) => Err(e),
} }
} Err(e) => Err(e),
},
DatabaseType::Mysql => match db::mysql::connect(&url).await { DatabaseType::Mysql => match db::mysql::connect(&url).await {
Ok(pool) => { Ok(pool) => {
pool.close().await; pool.close().await;
@@ -77,70 +55,48 @@ pub async fn test_connection(
} }
Err(e) => Err(e), Err(e) => Err(e),
}, },
DatabaseType::Doris | DatabaseType::StarRocks => { DatabaseType::Doris | DatabaseType::StarRocks => match db::mysql::connect_bare(&url).await {
match db::mysql::connect_bare(&url).await { Ok(pool) => {
Ok(pool) => { pool.close().await;
pool.close().await; Ok("Connection successful".to_string())
Ok("Connection successful".to_string())
}
Err(e) => Err(e),
} }
} Err(e) => Err(e),
DatabaseType::Postgres | DatabaseType::Redshift => { },
match db::postgres::connect(&url).await { DatabaseType::Postgres | DatabaseType::Redshift => match db::postgres::connect(&url).await {
Ok(pool) => { Ok(pool) => {
pool.close().await; pool.close().await;
Ok("Connection successful".to_string()) Ok("Connection successful".to_string())
}
Err(e) => Err(e),
} }
} Err(e) => Err(e),
DatabaseType::Sqlite => { },
match db::sqlite::connect_path(&expand_tilde(&config.host)).await { DatabaseType::Sqlite => match db::sqlite::connect_path(&expand_tilde(&config.host)).await {
Ok(pool) => { Ok(pool) => {
pool.close().await; pool.close().await;
Ok("Connection successful".to_string()) Ok("Connection successful".to_string())
}
Err(e) => Err(e),
} }
} Err(e) => Err(e),
DatabaseType::Redis => db::redis_driver::connect(&url) },
.await DatabaseType::Redis => db::redis_driver::connect(&url).await.map(|_| "Connection successful".to_string()),
.map(|_| "Connection successful".to_string()),
DatabaseType::DuckDb => duckdb::Connection::open(&expand_tilde(&config.host)) DatabaseType::DuckDb => duckdb::Connection::open(&expand_tilde(&config.host))
.map(|_| "Connection successful".to_string()) .map(|_| "Connection successful".to_string())
.map_err(|e| e.to_string()), .map_err(|e| e.to_string()),
DatabaseType::MongoDb => match db::mongo_driver::connect(&url).await { DatabaseType::MongoDb => match db::mongo_driver::connect(&url).await {
Ok(client) => db::mongo_driver::test_connection(&client) Ok(client) => {
.await db::mongo_driver::test_connection(&client).await.map(|_| "Connection successful".to_string())
.map(|_| "Connection successful".to_string()), }
Err(e) => Err(e.to_string()), Err(e) => Err(e.to_string()),
}, },
DatabaseType::ClickHouse => { DatabaseType::ClickHouse => {
let username = if config.username.is_empty() { let username = if config.username.is_empty() { None } else { Some(config.username.clone()) };
None let password = if config.password.is_empty() { None } else { Some(config.password.clone()) };
} else {
Some(config.username.clone())
};
let password = if config.password.is_empty() {
None
} else {
Some(config.password.clone())
};
let client = db::clickhouse_driver::ChClient::new(&url, username, password); let client = db::clickhouse_driver::ChClient::new(&url, username, password);
db::clickhouse_driver::test_connection(&client) db::clickhouse_driver::test_connection(&client).await.map(|_| "Connection successful".to_string())
}
DatabaseType::SqlServer => {
db::sqlserver::connect(&host, port, &config.username, &config.password, config.database.as_deref())
.await .await
.map(|_| "Connection successful".to_string()) .map(|_| "Connection successful".to_string())
} }
DatabaseType::SqlServer => db::sqlserver::connect(
&host,
port,
&config.username,
&config.password,
config.database.as_deref(),
)
.await
.map(|_| "Connection successful".to_string()),
DatabaseType::Oracle => db::oracle_driver::connect( DatabaseType::Oracle => db::oracle_driver::connect(
&host, &host,
port, port,
@@ -151,14 +107,9 @@ pub async fn test_connection(
.await .await
.map(|_| "Connection successful".to_string()), .map(|_| "Connection successful".to_string()),
DatabaseType::Elasticsearch => { DatabaseType::Elasticsearch => {
let client = db::elasticsearch_driver::EsClient::new( let client =
&url, db::elasticsearch_driver::EsClient::new(&url, Some(&config.username), Some(&config.password));
Some(&config.username), db::elasticsearch_driver::test_connection(&client).await.map(|_| "Connection successful".to_string())
Some(&config.password),
);
db::elasticsearch_driver::test_connection(&client)
.await
.map(|_| "Connection successful".to_string())
} }
}, },
}; };
@@ -171,10 +122,7 @@ pub async fn test_connection(
} }
#[tauri::command] #[tauri::command]
pub async fn connect_db( pub async fn connect_db(state: State<'_, Arc<AppState>>, config: ConnectionConfig) -> Result<String, String> {
state: State<'_, Arc<AppState>>,
config: ConnectionConfig,
) -> Result<String, String> {
let id = config.id.clone(); let id = config.id.clone();
let (host, port) = state.connection_host_port(&id, &config).await?; let (host, port) = state.connection_host_port(&id, &config).await?;
@@ -182,26 +130,17 @@ pub async fn connect_db(
let url = connection_url_for_endpoint(&config, &host, port); let url = connection_url_for_endpoint(&config, &host, port);
let pool = match config.db_type { let pool = match config.db_type {
DatabaseType::Mysql if config.needs_bare_mysql() => { DatabaseType::Mysql if config.needs_bare_mysql() => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true),
PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true)
}
DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?, false), DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?, false),
DatabaseType::Doris | DatabaseType::StarRocks => { DatabaseType::Doris | DatabaseType::StarRocks => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true),
PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true) DatabaseType::Postgres | DatabaseType::Redshift => PoolKind::Postgres(db::postgres::connect(&url).await?),
} DatabaseType::Sqlite => PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&config.host)).await?),
DatabaseType::Postgres | DatabaseType::Redshift => {
PoolKind::Postgres(db::postgres::connect(&url).await?)
}
DatabaseType::Sqlite => {
PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&config.host)).await?)
}
DatabaseType::Redis => { DatabaseType::Redis => {
let con = db::redis_driver::connect(&url).await?; let con = db::redis_driver::connect(&url).await?;
PoolKind::Redis(tokio::sync::Mutex::new(con)) PoolKind::Redis(tokio::sync::Mutex::new(con))
} }
DatabaseType::DuckDb => { DatabaseType::DuckDb => {
let con = let con = duckdb::Connection::open(&expand_tilde(&config.host)).map_err(|e| e.to_string())?;
duckdb::Connection::open(&expand_tilde(&config.host)).map_err(|e| e.to_string())?;
PoolKind::DuckDb(std::sync::Arc::new(std::sync::Mutex::new(con))) PoolKind::DuckDb(std::sync::Arc::new(std::sync::Mutex::new(con)))
} }
DatabaseType::MongoDb => { DatabaseType::MongoDb => {
@@ -210,29 +149,16 @@ pub async fn connect_db(
PoolKind::MongoDb(client) PoolKind::MongoDb(client)
} }
DatabaseType::ClickHouse => { DatabaseType::ClickHouse => {
let username = if config.username.is_empty() { let username = if config.username.is_empty() { None } else { Some(config.username.clone()) };
None let password = if config.password.is_empty() { None } else { Some(config.password.clone()) };
} else {
Some(config.username.clone())
};
let password = if config.password.is_empty() {
None
} else {
Some(config.password.clone())
};
let client = db::clickhouse_driver::ChClient::new(&url, username, password); let client = db::clickhouse_driver::ChClient::new(&url, username, password);
db::clickhouse_driver::test_connection(&client).await?; db::clickhouse_driver::test_connection(&client).await?;
PoolKind::ClickHouse(client) PoolKind::ClickHouse(client)
} }
DatabaseType::SqlServer => { DatabaseType::SqlServer => {
let client = db::sqlserver::connect( let client =
&host, db::sqlserver::connect(&host, port, &config.username, &config.password, config.database.as_deref())
port, .await?;
&config.username,
&config.password,
config.database.as_deref(),
)
.await?;
PoolKind::SqlServer(std::sync::Arc::new(tokio::sync::Mutex::new(client))) PoolKind::SqlServer(std::sync::Arc::new(tokio::sync::Mutex::new(client)))
} }
DatabaseType::Oracle => { DatabaseType::Oracle => {
@@ -247,11 +173,7 @@ pub async fn connect_db(
PoolKind::Oracle(std::sync::Arc::new(tokio::sync::Mutex::new(client))) PoolKind::Oracle(std::sync::Arc::new(tokio::sync::Mutex::new(client)))
} }
DatabaseType::Elasticsearch => { DatabaseType::Elasticsearch => {
let client = db::elasticsearch_driver::EsClient::new( let client = db::elasticsearch_driver::EsClient::new(&url, Some(&config.username), Some(&config.password));
&url,
Some(&config.username),
Some(&config.password),
);
db::elasticsearch_driver::test_connection(&client).await?; db::elasticsearch_driver::test_connection(&client).await?;
PoolKind::Elasticsearch(client) PoolKind::Elasticsearch(client)
} }
@@ -264,16 +186,10 @@ pub async fn connect_db(
} }
#[tauri::command] #[tauri::command]
pub async fn disconnect_db( pub async fn disconnect_db(state: State<'_, Arc<AppState>>, connection_id: String) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
connection_id: String,
) -> Result<(), String> {
let mut conns = state.connections.lock().await; let mut conns = state.connections.lock().await;
let keys_to_remove: Vec<String> = conns let keys_to_remove: Vec<String> =
.keys() conns.keys().filter(|k| *k == &connection_id || k.starts_with(&format!("{connection_id}:"))).cloned().collect();
.filter(|k| *k == &connection_id || k.starts_with(&format!("{connection_id}:")))
.cloned()
.collect();
for key in keys_to_remove { for key in keys_to_remove {
if let Some(pool) = conns.remove(&key) { if let Some(pool) = conns.remove(&key) {
match pool { match pool {
+2 -8
View File
@@ -5,10 +5,7 @@ use super::connection::AppState;
pub use dbx_core::history::HistoryEntry; pub use dbx_core::history::HistoryEntry;
#[tauri::command] #[tauri::command]
pub async fn save_history( pub async fn save_history(state: State<'_, Arc<AppState>>, entry: HistoryEntry) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
entry: HistoryEntry,
) -> Result<(), String> {
state.storage.save_history_entry(&entry).await state.storage.save_history_entry(&entry).await
} }
@@ -27,9 +24,6 @@ pub async fn clear_history(state: State<'_, Arc<AppState>>) -> Result<(), String
} }
#[tauri::command] #[tauri::command]
pub async fn delete_history_entry( pub async fn delete_history_entry(state: State<'_, Arc<AppState>>, id: String) -> Result<(), String> {
state: State<'_, Arc<AppState>>,
id: String,
) -> Result<(), String> {
state.storage.delete_history_entry(&id).await state.storage.delete_history_entry(&id).await
} }
+25 -8
View File
@@ -1,6 +1,6 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::sync::Arc; use std::sync::Arc;
use tauri::{Emitter, AppHandle, Manager}; use tauri::{AppHandle, Emitter, Manager};
use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener; use tokio::net::TcpListener;
@@ -82,7 +82,10 @@ pub fn start(app_handle: AppHandle, state: Arc<AppState>) {
}); });
} }
fn find_config_by_name<'a>(configs: &'a [crate::models::connection::ConnectionConfig], name: &str) -> Option<&'a crate::models::connection::ConnectionConfig> { fn find_config_by_name<'a>(
configs: &'a [crate::models::connection::ConnectionConfig],
name: &str,
) -> Option<&'a crate::models::connection::ConnectionConfig> {
configs.iter().find(|c| c.name.eq_ignore_ascii_case(name)) configs.iter().find(|c| c.name.eq_ignore_ascii_case(name))
} }
@@ -94,14 +97,21 @@ async fn respond(stream: &mut tokio::net::TcpStream, status: &str, body: &str) {
async fn handle_open_table(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) { async fn handle_open_table(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) {
let req: OpenTableRequest = match serde_json::from_str(body) { let req: OpenTableRequest = match serde_json::from_str(body) {
Ok(r) => r, Ok(r) => r,
Err(_) => { respond(stream, "400 Bad Request", "").await; return; } Err(_) => {
respond(stream, "400 Bad Request", "").await;
return;
}
}; };
let configs = match state.storage.load_connections().await { let configs = match state.storage.load_connections().await {
Ok(c) => c, Ok(c) => c,
Err(_) => { respond(stream, "500 Internal Server Error", "").await; return; } Err(_) => {
respond(stream, "500 Internal Server Error", "").await;
return;
}
}; };
let Some(config) = find_config_by_name(&configs, &req.connection_name) else { let Some(config) = find_config_by_name(&configs, &req.connection_name) else {
respond(stream, "404 Not Found", "Connection not found").await; return; respond(stream, "404 Not Found", "Connection not found").await;
return;
}; };
let event = McpOpenTableEvent { let event = McpOpenTableEvent {
connection_id: config.id.clone(), connection_id: config.id.clone(),
@@ -116,14 +126,21 @@ async fn handle_open_table(app: &AppHandle, state: &Arc<AppState>, body: &str, s
async fn handle_execute_query(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) { async fn handle_execute_query(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) {
let req: ExecuteQueryRequest = match serde_json::from_str(body) { let req: ExecuteQueryRequest = match serde_json::from_str(body) {
Ok(r) => r, Ok(r) => r,
Err(_) => { respond(stream, "400 Bad Request", "").await; return; } Err(_) => {
respond(stream, "400 Bad Request", "").await;
return;
}
}; };
let configs = match state.storage.load_connections().await { let configs = match state.storage.load_connections().await {
Ok(c) => c, Ok(c) => c,
Err(_) => { respond(stream, "500 Internal Server Error", "").await; return; } Err(_) => {
respond(stream, "500 Internal Server Error", "").await;
return;
}
}; };
let Some(config) = find_config_by_name(&configs, &req.connection_name) else { let Some(config) = find_config_by_name(&configs, &req.connection_name) else {
respond(stream, "404 Not Found", "Connection not found").await; return; respond(stream, "404 Not Found", "Connection not found").await;
return;
}; };
let event = McpExecuteQueryEvent { let event = McpExecuteQueryEvent {
connection_id: config.id.clone(), connection_id: config.id.clone(),
+2 -1
View File
@@ -53,7 +53,8 @@ pub async fn mongo_update_document(
id: String, id: String,
doc_json: String, doc_json: String,
) -> Result<u64, String> { ) -> Result<u64, String> {
dbx_core::mongo_ops::mongo_update_document_core(&state, &connection_id, &database, &collection, &id, &doc_json).await dbx_core::mongo_ops::mongo_update_document_core(&state, &connection_id, &database, &collection, &id, &doc_json)
.await
} }
#[tauri::command] #[tauri::command]
+7 -28
View File
@@ -16,20 +16,11 @@ pub async fn execute_query(
sql: String, sql: String,
execution_id: Option<String>, execution_id: Option<String>,
) -> Result<db::QueryResult, String> { ) -> Result<db::QueryResult, String> {
let registered_query = execution_id let registered_query =
.as_ref() execution_id.as_ref().filter(|id| !id.trim().is_empty()).map(|id| state.running_queries.register(id.clone()));
.filter(|id| !id.trim().is_empty())
.map(|id| state.running_queries.register(id.clone()));
let cancel_token = registered_query.as_ref().map(|query| query.token()); let cancel_token = registered_query.as_ref().map(|query| query.token());
dbx_core::query::execute_sql_statement( dbx_core::query::execute_sql_statement(&state, &connection_id, &database, &sql, cancel_token).await
&state,
&connection_id,
&database,
&sql,
cancel_token,
)
.await
} }
#[tauri::command] #[tauri::command]
@@ -40,27 +31,15 @@ pub async fn execute_multi(
sql: String, sql: String,
execution_id: Option<String>, execution_id: Option<String>,
) -> Result<Vec<db::QueryResult>, String> { ) -> Result<Vec<db::QueryResult>, String> {
let registered_query = execution_id let registered_query =
.as_ref() execution_id.as_ref().filter(|id| !id.trim().is_empty()).map(|id| state.running_queries.register(id.clone()));
.filter(|id| !id.trim().is_empty())
.map(|id| state.running_queries.register(id.clone()));
let cancel_token = registered_query.as_ref().map(|query| query.token()); let cancel_token = registered_query.as_ref().map(|query| query.token());
dbx_core::query::execute_multi_core( dbx_core::query::execute_multi_core(&state, &connection_id, &database, &sql, cancel_token).await
&state,
&connection_id,
&database,
&sql,
cancel_token,
)
.await
} }
#[tauri::command] #[tauri::command]
pub async fn cancel_query( pub async fn cancel_query(state: State<'_, Arc<AppState>>, execution_id: String) -> Result<bool, String> {
state: State<'_, Arc<AppState>>,
execution_id: String,
) -> Result<bool, String> {
Ok(state.running_queries.cancel(&execution_id)) Ok(state.running_queries.cancel(&execution_id))
} }
+20 -10
View File
@@ -5,10 +5,7 @@ use crate::commands::connection::AppState;
use dbx_core::db::redis_driver::{RedisScanResult, RedisValue}; use dbx_core::db::redis_driver::{RedisScanResult, RedisValue};
#[tauri::command] #[tauri::command]
pub async fn redis_list_databases( pub async fn redis_list_databases(state: State<'_, Arc<AppState>>, connection_id: String) -> Result<Vec<u32>, String> {
state: State<'_, Arc<AppState>>,
connection_id: String,
) -> Result<Vec<u32>, String> {
dbx_core::redis_ops::redis_list_databases_core(&state, &connection_id).await dbx_core::redis_ops::redis_list_databases_core(&state, &connection_id).await
} }
@@ -56,7 +53,10 @@ pub async fn redis_delete_key(
#[tauri::command] #[tauri::command]
pub async fn redis_hash_set( pub async fn redis_hash_set(
state: State<'_, Arc<AppState>>, state: State<'_, Arc<AppState>>,
connection_id: String, key: String, field: String, value: String, connection_id: String,
key: String,
field: String,
value: String,
) -> Result<(), String> { ) -> Result<(), String> {
dbx_core::redis_ops::redis_hash_set_core(&state, &connection_id, &key, &field, &value).await dbx_core::redis_ops::redis_hash_set_core(&state, &connection_id, &key, &field, &value).await
} }
@@ -64,7 +64,9 @@ pub async fn redis_hash_set(
#[tauri::command] #[tauri::command]
pub async fn redis_hash_del( pub async fn redis_hash_del(
state: State<'_, Arc<AppState>>, state: State<'_, Arc<AppState>>,
connection_id: String, key: String, field: String, connection_id: String,
key: String,
field: String,
) -> Result<(), String> { ) -> Result<(), String> {
dbx_core::redis_ops::redis_hash_del_core(&state, &connection_id, &key, &field).await dbx_core::redis_ops::redis_hash_del_core(&state, &connection_id, &key, &field).await
} }
@@ -72,7 +74,9 @@ pub async fn redis_hash_del(
#[tauri::command] #[tauri::command]
pub async fn redis_list_push( pub async fn redis_list_push(
state: State<'_, Arc<AppState>>, state: State<'_, Arc<AppState>>,
connection_id: String, key: String, value: String, connection_id: String,
key: String,
value: String,
) -> Result<(), String> { ) -> Result<(), String> {
dbx_core::redis_ops::redis_list_push_core(&state, &connection_id, &key, &value).await dbx_core::redis_ops::redis_list_push_core(&state, &connection_id, &key, &value).await
} }
@@ -80,7 +84,9 @@ pub async fn redis_list_push(
#[tauri::command] #[tauri::command]
pub async fn redis_list_remove( pub async fn redis_list_remove(
state: State<'_, Arc<AppState>>, state: State<'_, Arc<AppState>>,
connection_id: String, key: String, index: i64, connection_id: String,
key: String,
index: i64,
) -> Result<(), String> { ) -> Result<(), String> {
dbx_core::redis_ops::redis_list_remove_core(&state, &connection_id, &key, index).await dbx_core::redis_ops::redis_list_remove_core(&state, &connection_id, &key, index).await
} }
@@ -88,7 +94,9 @@ pub async fn redis_list_remove(
#[tauri::command] #[tauri::command]
pub async fn redis_set_add( pub async fn redis_set_add(
state: State<'_, Arc<AppState>>, state: State<'_, Arc<AppState>>,
connection_id: String, key: String, member: String, connection_id: String,
key: String,
member: String,
) -> Result<(), String> { ) -> Result<(), String> {
dbx_core::redis_ops::redis_set_add_core(&state, &connection_id, &key, &member).await dbx_core::redis_ops::redis_set_add_core(&state, &connection_id, &key, &member).await
} }
@@ -96,7 +104,9 @@ pub async fn redis_set_add(
#[tauri::command] #[tauri::command]
pub async fn redis_set_remove( pub async fn redis_set_remove(
state: State<'_, Arc<AppState>>, state: State<'_, Arc<AppState>>,
connection_id: String, key: String, member: String, connection_id: String,
key: String,
member: String,
) -> Result<(), String> { ) -> Result<(), String> {
dbx_core::redis_ops::redis_set_remove_core(&state, &connection_id, &key, &member).await dbx_core::redis_ops::redis_set_remove_core(&state, &connection_id, &key, &member).await
} }
+19 -94
View File
@@ -12,8 +12,7 @@ use crate::commands::connection::AppState;
use crate::commands::query::execute_sql_statement; use crate::commands::query::execute_sql_statement;
pub use dbx_core::sql::{ pub use dbx_core::sql::{
statement_summary, SqlFilePreview, SqlFileProgress, SqlFileRequest, statement_summary, SqlFilePreview, SqlFileProgress, SqlFileRequest, SqlFileStatus, SqlStatementSplitter,
SqlFileStatus, SqlStatementSplitter,
}; };
static SQL_FILE_EXECUTIONS: std::sync::LazyLock<RwLock<HashMap<String, CancellationToken>>> = static SQL_FILE_EXECUTIONS: std::sync::LazyLock<RwLock<HashMap<String, CancellationToken>>> =
@@ -38,25 +37,15 @@ struct SqlFileSummary {
#[tauri::command] #[tauri::command]
pub async fn preview_sql_file(file_path: String) -> Result<SqlFilePreview, String> { pub async fn preview_sql_file(file_path: String) -> Result<SqlFilePreview, String> {
let path = PathBuf::from(&file_path); let path = PathBuf::from(&file_path);
let metadata = tokio::fs::metadata(&path) let metadata = tokio::fs::metadata(&path).await.map_err(|e| e.to_string())?;
.await let mut file = tokio::fs::File::open(&path).await.map_err(|e| e.to_string())?;
.map_err(|e| e.to_string())?;
let mut file = tokio::fs::File::open(&path)
.await
.map_err(|e| e.to_string())?;
let mut buffer = vec![0; 4096]; let mut buffer = vec![0; 4096];
let bytes_read = tokio::io::AsyncReadExt::read(&mut file, &mut buffer) let bytes_read = tokio::io::AsyncReadExt::read(&mut file, &mut buffer).await.map_err(|e| e.to_string())?;
.await
.map_err(|e| e.to_string())?;
buffer.truncate(bytes_read); buffer.truncate(bytes_read);
let preview = String::from_utf8_lossy(&buffer).to_string(); let preview = String::from_utf8_lossy(&buffer).to_string();
Ok(SqlFilePreview { Ok(SqlFilePreview {
file_name: path file_name: path.file_name().and_then(|name| name.to_str()).unwrap_or("script.sql").to_string(),
.file_name()
.and_then(|name| name.to_str())
.unwrap_or("script.sql")
.to_string(),
file_path, file_path,
size_bytes: metadata.len(), size_bytes: metadata.len(),
preview, preview,
@@ -76,18 +65,7 @@ pub async fn execute_sql_file(
} }
let started_at = Instant::now(); let started_at = Instant::now();
emit_progress( emit_progress(&app, &request.execution_id, SqlFileStatus::Started, 0, 0, 0, 0, started_at, "", None);
&app,
&request.execution_id,
SqlFileStatus::Started,
0,
0,
0,
0,
started_at,
"",
None,
);
let result = execute_sql_file_inner(&app, &state, &request, token, started_at).await; let result = execute_sql_file_inner(&app, &state, &request, token, started_at).await;
{ {
@@ -334,19 +312,14 @@ fn register_sql_file_execution(
token: CancellationToken, token: CancellationToken,
) -> Result<(), String> { ) -> Result<(), String> {
if executions.contains_key(&execution_id) { if executions.contains_key(&execution_id) {
return Err(format!( return Err(format!("SQL file execution '{execution_id}' already exists"));
"SQL file execution '{execution_id}' already exists"
));
} }
executions.insert(execution_id, token); executions.insert(execution_id, token);
Ok(()) Ok(())
} }
fn remove_sql_file_execution( fn remove_sql_file_execution(executions: &mut HashMap<String, CancellationToken>, execution_id: &str) {
executions: &mut HashMap<String, CancellationToken>,
execution_id: &str,
) {
executions.remove(execution_id); executions.remove(execution_id);
} }
@@ -394,11 +367,7 @@ fn statement_error_decision(
); );
if continue_on_error { if continue_on_error {
return StatementErrorDecision { return StatementErrorDecision { progress: vec![statement_failed], failure_count, result: Ok(false) };
progress: vec![statement_failed],
failure_count,
result: Ok(false),
};
} }
let terminal_error = sql_file_progress( let terminal_error = sql_file_progress(
@@ -413,11 +382,7 @@ fn statement_error_decision(
Some(error.clone()), Some(error.clone()),
); );
StatementErrorDecision { StatementErrorDecision { progress: vec![statement_failed, terminal_error], failure_count, result: Err(error) }
progress: vec![statement_failed, terminal_error],
failure_count,
result: Err(error),
}
} }
fn emit_progress( fn emit_progress(
@@ -559,11 +524,7 @@ async fn run_statements_for_test(
} }
SqlFileSummary { SqlFileSummary {
status: if token.is_cancelled() { status: if token.is_cancelled() { SqlFileStatus::Cancelled } else { SqlFileStatus::Done },
SqlFileStatus::Cancelled
} else {
SqlFileStatus::Done
},
success_count, success_count,
failure_count, failure_count,
failed_statement_index, failed_statement_index,
@@ -586,12 +547,7 @@ mod execution_tests {
#[tokio::test] #[tokio::test]
async fn stops_on_first_failure_by_default() { async fn stops_on_first_failure_by_default() {
let summary = run_fake_script( let summary = run_fake_script(vec!["ok 1".into(), "fail 2".into(), "ok 3".into()], false, None).await;
vec!["ok 1".into(), "fail 2".into(), "ok 3".into()],
false,
None,
)
.await;
assert_eq!(summary.success_count, 1); assert_eq!(summary.success_count, 1);
assert_eq!(summary.failure_count, 1); assert_eq!(summary.failure_count, 1);
@@ -601,12 +557,7 @@ mod execution_tests {
#[tokio::test] #[tokio::test]
async fn continues_after_failure_when_enabled() { async fn continues_after_failure_when_enabled() {
let summary = run_fake_script( let summary = run_fake_script(vec!["ok 1".into(), "fail 2".into(), "ok 3".into()], true, None).await;
vec!["ok 1".into(), "fail 2".into(), "ok 3".into()],
true,
None,
)
.await;
assert_eq!(summary.success_count, 2); assert_eq!(summary.success_count, 2);
assert_eq!(summary.failure_count, 1); assert_eq!(summary.failure_count, 1);
@@ -615,12 +566,7 @@ mod execution_tests {
#[tokio::test] #[tokio::test]
async fn cancellation_stops_before_next_statement() { async fn cancellation_stops_before_next_statement() {
let summary = run_fake_script( let summary = run_fake_script(vec!["ok 1".into(), "ok 2".into(), "ok 3".into()], true, Some(1)).await;
vec!["ok 1".into(), "ok 2".into(), "ok 3".into()],
true,
Some(1),
)
.await;
assert_eq!(summary.success_count, 1); assert_eq!(summary.success_count, 1);
assert_eq!(summary.status, SqlFileStatus::Cancelled); assert_eq!(summary.status, SqlFileStatus::Cancelled);
@@ -628,15 +574,7 @@ mod execution_tests {
#[test] #[test]
fn file_io_errors_build_terminal_error_progress() { fn file_io_errors_build_terminal_error_progress() {
let progress = file_io_error_progress( let progress = file_io_error_progress("exec-1", 4, 2, 1, 17, Instant::now(), "read failed".to_string());
"exec-1",
4,
2,
1,
17,
Instant::now(),
"read failed".to_string(),
);
assert_eq!(progress.execution_id, "exec-1"); assert_eq!(progress.execution_id, "exec-1");
assert_eq!(progress.status, SqlFileStatus::Error); assert_eq!(progress.status, SqlFileStatus::Error);
@@ -655,13 +593,9 @@ mod execution_tests {
let replacement = CancellationToken::new(); let replacement = CancellationToken::new();
executions.insert("dup".to_string(), original.clone()); executions.insert("dup".to_string(), original.clone());
let result = let result = register_sql_file_execution(&mut executions, "dup".to_string(), replacement.clone());
register_sql_file_execution(&mut executions, "dup".to_string(), replacement.clone());
assert_eq!( assert_eq!(result.unwrap_err(), "SQL file execution 'dup' already exists");
result.unwrap_err(),
"SQL file execution 'dup' already exists"
);
assert_eq!(executions.len(), 1); assert_eq!(executions.len(), 1);
executions.get("dup").unwrap().cancel(); executions.get("dup").unwrap().cancel();
@@ -720,17 +654,8 @@ mod execution_tests {
#[test] #[test]
fn progress_payload_serializes_camel_case_status() { fn progress_payload_serializes_camel_case_status() {
let progress = sql_file_progress( let progress =
"exec-1", sql_file_progress("exec-1", SqlFileStatus::StatementDone, 1, 1, 0, 3, Instant::now(), "select 1", None);
SqlFileStatus::StatementDone,
1,
1,
0,
3,
Instant::now(),
"select 1",
None,
);
let value = serde_json::to_value(progress).unwrap(); let value = serde_json::to_value(progress).unwrap();
+2 -7
View File
@@ -8,10 +8,7 @@ use crate::commands::connection::AppState;
use crate::commands::transfer::get_db_type; use crate::commands::transfer::get_db_type;
// Re-export types for backward compatibility // Re-export types for backward compatibility
pub use dbx_core::table_import::{ pub use dbx_core::table_import::{TableImportPreview, TableImportProgress, TableImportRequest, TableImportSummary};
TableImportPreview, TableImportProgress,
TableImportRequest, TableImportSummary,
};
static CANCELLED_IMPORTS: std::sync::LazyLock<RwLock<HashSet<String>>> = static CANCELLED_IMPORTS: std::sync::LazyLock<RwLock<HashSet<String>>> =
std::sync::LazyLock::new(|| RwLock::new(HashSet::new())); std::sync::LazyLock::new(|| RwLock::new(HashSet::new()));
@@ -44,9 +41,7 @@ pub async fn import_table_file(
let pool_key = if request.database.is_empty() { let pool_key = if request.database.is_empty() {
request.connection_id.clone() request.connection_id.clone()
} else { } else {
state state.get_or_create_pool(&request.connection_id, Some(&request.database)).await?
.get_or_create_pool(&request.connection_id, Some(&request.database))
.await?
}; };
let result = dbx_core::table_import::import_table_file_core( let result = dbx_core::table_import::import_table_file_core(
+69 -51
View File
@@ -4,10 +4,7 @@ use tauri::{AppHandle, Emitter, State};
use crate::commands::connection::AppState; use crate::commands::connection::AppState;
// Re-export types and functions used by other modules // Re-export types and functions used by other modules
pub use dbx_core::transfer::{ pub use dbx_core::transfer::{get_db_type, TransferProgress, TransferRequest, TransferStatus};
get_db_type,
TransferProgress, TransferRequest, TransferStatus,
};
fn emit_progress(app: &AppHandle, progress: TransferProgress) { fn emit_progress(app: &AppHandle, progress: TransferProgress) {
let _ = app.emit("transfer-progress", progress); let _ = app.emit("transfer-progress", progress);
@@ -27,12 +24,10 @@ pub async fn start_transfer(
let target_db_type = get_db_type(&state, &request.target_connection_id).await?; let target_db_type = get_db_type(&state, &request.target_connection_id).await?;
// Ensure pools // Ensure pools
let source_pool_key = state let source_pool_key =
.get_or_create_pool(&request.source_connection_id, Some(&request.source_database)) state.get_or_create_pool(&request.source_connection_id, Some(&request.source_database)).await?;
.await?; let target_pool_key =
let target_pool_key = state state.get_or_create_pool(&request.target_connection_id, Some(&request.target_database)).await?;
.get_or_create_pool(&request.target_connection_id, Some(&request.target_database))
.await?;
tokio::spawn(async move { tokio::spawn(async move {
let total_tables = request.tables.len(); let total_tables = request.tables.len();
@@ -40,16 +35,19 @@ pub async fn start_transfer(
for (i, table) in request.tables.iter().enumerate() { for (i, table) in request.tables.iter().enumerate() {
if dbx_core::transfer::is_cancelled(&transfer_id).await { if dbx_core::transfer::is_cancelled(&transfer_id).await {
emit_progress(&app, TransferProgress { emit_progress(
transfer_id: transfer_id.clone(), &app,
table: table.clone(), TransferProgress {
table_index: i, transfer_id: transfer_id.clone(),
total_tables, table: table.clone(),
rows_transferred: 0, table_index: i,
total_rows: None, total_tables,
status: TransferStatus::Cancelled, rows_transferred: 0,
error: None, total_rows: None,
}); status: TransferStatus::Cancelled,
error: None,
},
);
dbx_core::transfer::clear_cancelled(&transfer_id).await; dbx_core::transfer::clear_cancelled(&transfer_id).await;
return; return;
} }
@@ -57,48 +55,68 @@ pub async fn start_transfer(
log::info!("[transfer] table {}/{}: {}", i + 1, total_tables, table); log::info!("[transfer] table {}/{}: {}", i + 1, total_tables, table);
match dbx_core::transfer::transfer_table( match dbx_core::transfer::transfer_table(
&state, &request, table, i, &state,
&source_db_type, &target_db_type, &request,
&source_pool_key, &target_pool_key, table,
i,
&source_db_type,
&target_db_type,
&source_pool_key,
&target_pool_key,
|progress| emit_progress(&app, progress), |progress| emit_progress(&app, progress),
).await { )
.await
{
Ok(rows) => { Ok(rows) => {
emit_progress(&app, TransferProgress { emit_progress(
transfer_id: transfer_id.clone(), &app,
table: table.clone(), TransferProgress {
table_index: i, transfer_id: transfer_id.clone(),
total_tables, table: table.clone(),
rows_transferred: rows, table_index: i,
total_rows: Some(rows), total_tables,
status: if i == total_tables - 1 { TransferStatus::Done } else { TransferStatus::TableDone }, rows_transferred: rows,
error: None, total_rows: Some(rows),
}); status: if i == total_tables - 1 {
TransferStatus::Done
} else {
TransferStatus::TableDone
},
error: None,
},
);
} }
Err(e) => { Err(e) => {
if e == "Cancelled" { if e == "Cancelled" {
emit_progress(&app, TransferProgress { emit_progress(
&app,
TransferProgress {
transfer_id: transfer_id.clone(),
table: table.clone(),
table_index: i,
total_tables,
rows_transferred: 0,
total_rows: None,
status: TransferStatus::Cancelled,
error: None,
},
);
dbx_core::transfer::clear_cancelled(&transfer_id).await;
return;
}
emit_progress(
&app,
TransferProgress {
transfer_id: transfer_id.clone(), transfer_id: transfer_id.clone(),
table: table.clone(), table: table.clone(),
table_index: i, table_index: i,
total_tables, total_tables,
rows_transferred: 0, rows_transferred: 0,
total_rows: None, total_rows: None,
status: TransferStatus::Cancelled, status: TransferStatus::Error,
error: None, error: Some(e),
}); },
dbx_core::transfer::clear_cancelled(&transfer_id).await; );
return;
}
emit_progress(&app, TransferProgress {
transfer_id: transfer_id.clone(),
table: table.clone(),
table_index: i,
total_tables,
rows_transferred: 0,
total_rows: None,
status: TransferStatus::Error,
error: Some(e),
});
dbx_core::transfer::clear_cancelled(&transfer_id).await; dbx_core::transfer::clear_cancelled(&transfer_id).await;
return; return;
} }
+5 -16
View File
@@ -9,9 +9,7 @@ use tauri::{Manager, RunEvent};
#[cfg_attr(mobile, tauri::mobile_entry_point)] #[cfg_attr(mobile, tauri::mobile_entry_point)]
pub fn run() { pub fn run() {
rustls::crypto::aws_lc_rs::default_provider() rustls::crypto::aws_lc_rs::default_provider().install_default().expect("Failed to install rustls crypto provider");
.install_default()
.expect("Failed to install rustls crypto provider");
tauri::Builder::default() tauri::Builder::default()
.plugin(tauri_plugin_dialog::init()) .plugin(tauri_plugin_dialog::init())
@@ -21,26 +19,17 @@ pub fn run() {
.plugin(tauri_plugin_process::init()) .plugin(tauri_plugin_process::init())
.setup(|app| { .setup(|app| {
if cfg!(debug_assertions) { if cfg!(debug_assertions) {
app.handle().plugin( app.handle().plugin(tauri_plugin_log::Builder::default().level(log::LevelFilter::Info).build())?;
tauri_plugin_log::Builder::default()
.level(log::LevelFilter::Info)
.build(),
)?;
} }
let data_dir = app let data_dir =
.path() app.path().app_data_dir().map_err(|e| e.to_string()).expect("Failed to resolve app data dir");
.app_data_dir()
.map_err(|e| e.to_string())
.expect("Failed to resolve app data dir");
std::fs::create_dir_all(&data_dir).expect("Failed to create data dir"); std::fs::create_dir_all(&data_dir).expect("Failed to create data dir");
let db_path = data_dir.join("dbx.db"); let db_path = data_dir.join("dbx.db");
let storage = tauri::async_runtime::block_on(async { let storage = tauri::async_runtime::block_on(async {
let s = Storage::open(&db_path).await.expect("Failed to open storage"); let s = Storage::open(&db_path).await.expect("Failed to open storage");
s.migrate_from_json(&data_dir) s.migrate_from_json(&data_dir).await.expect("Failed to migrate JSON data");
.await
.expect("Failed to migrate JSON data");
s s
}); });
+1 -1
View File
@@ -2,5 +2,5 @@
#![cfg_attr(not(debug_assertions), windows_subsystem = "windows")] #![cfg_attr(not(debug_assertions), windows_subsystem = "windows")]
fn main() { fn main() {
dbx_lib::run(); dbx_lib::run();
} }
+6 -22
View File
@@ -24,18 +24,11 @@ pub struct AuthCheckResponse {
const MAX_ATTEMPTS: u32 = 5; const MAX_ATTEMPTS: u32 = 5;
const LOCKOUT_SECS: u64 = 60; const LOCKOUT_SECS: u64 = 60;
pub async fn login( pub async fn login(State(state): State<Arc<WebState>>, Json(body): Json<LoginRequest>) -> Result<Response, StatusCode> {
State(state): State<Arc<WebState>>,
Json(body): Json<LoginRequest>,
) -> Result<Response, StatusCode> {
let password_hash = match &state.password_hash { let password_hash = match &state.password_hash {
Some(h) => h, Some(h) => h,
None => { None => {
return Ok(( return Ok((StatusCode::OK, Json(serde_json::json!({"ok": true}))).into_response());
StatusCode::OK,
Json(serde_json::json!({"ok": true})),
)
.into_response());
} }
}; };
@@ -48,7 +41,8 @@ pub async fn login(
return Ok(( return Ok((
StatusCode::TOO_MANY_REQUESTS, StatusCode::TOO_MANY_REQUESTS,
Json(serde_json::json!({"error": format!("请 {remaining} 秒后再试")})), Json(serde_json::json!({"error": format!("请 {remaining} 秒后再试")})),
).into_response()); )
.into_response());
} }
} }
} }
@@ -78,12 +72,7 @@ pub async fn login(
state.sessions.write().await.insert(token.clone()); state.sessions.write().await.insert(token.clone());
let cookie = format!("dbx_session={token}; Path=/; HttpOnly; SameSite=Lax"); let cookie = format!("dbx_session={token}; Path=/; HttpOnly; SameSite=Lax");
Ok(( Ok((StatusCode::OK, [("set-cookie", cookie.as_str())], Json(serde_json::json!({"ok": true}))).into_response())
StatusCode::OK,
[("set-cookie", cookie.as_str())],
Json(serde_json::json!({"ok": true})),
)
.into_response())
} }
pub async fn check(State(state): State<Arc<WebState>>, req: Request<axum::body::Body>) -> Json<AuthCheckResponse> { pub async fn check(State(state): State<Arc<WebState>>, req: Request<axum::body::Body>) -> Json<AuthCheckResponse> {
@@ -102,12 +91,7 @@ pub async fn logout(State(state): State<Arc<WebState>>, req: Request<axum::body:
state.sessions.write().await.remove(&token); state.sessions.write().await.remove(&token);
} }
let cookie = "dbx_session=; Path=/; HttpOnly; Max-Age=0"; let cookie = "dbx_session=; Path=/; HttpOnly; Max-Age=0";
( (StatusCode::OK, [("set-cookie", cookie)], Json(serde_json::json!({"ok": true}))).into_response()
StatusCode::OK,
[("set-cookie", cookie)],
Json(serde_json::json!({"ok": true})),
)
.into_response()
} }
fn extract_session_token<B>(req: &Request<B>) -> Option<String> { fn extract_session_token<B>(req: &Request<B>) -> Option<String> {
+14 -40
View File
@@ -28,28 +28,19 @@ async fn main() {
) )
.init(); .init();
rustls::crypto::aws_lc_rs::default_provider() rustls::crypto::aws_lc_rs::default_provider().install_default().expect("Failed to install rustls crypto provider");
.install_default()
.expect("Failed to install rustls crypto provider");
// Data directory // Data directory
let data_dir = std::env::var("DBX_DATA_DIR") let data_dir = std::env::var("DBX_DATA_DIR").map(std::path::PathBuf::from).unwrap_or_else(|_| {
.map(std::path::PathBuf::from) let home = std::env::var("HOME").unwrap_or_else(|_| ".".to_string());
.unwrap_or_else(|_| { std::path::PathBuf::from(home).join(".dbx-web")
let home = std::env::var("HOME").unwrap_or_else(|_| ".".to_string()); });
std::path::PathBuf::from(home).join(".dbx-web")
});
std::fs::create_dir_all(&data_dir).expect("Failed to create data directory"); std::fs::create_dir_all(&data_dir).expect("Failed to create data directory");
let app_state = { let app_state = {
let db_path = data_dir.join("dbx.db"); let db_path = data_dir.join("dbx.db");
let storage = Storage::open(&db_path) let storage = Storage::open(&db_path).await.expect("Failed to open storage");
.await storage.migrate_from_json(&data_dir).await.expect("Failed to migrate JSON data");
.expect("Failed to open storage");
storage
.migrate_from_json(&data_dir)
.await
.expect("Failed to migrate JSON data");
Arc::new(AppState::new(storage)) Arc::new(AppState::new(storage))
}; };
@@ -66,17 +57,11 @@ async fn main() {
password_hash, password_hash,
sessions: RwLock::new(HashSet::new()), sessions: RwLock::new(HashSet::new()),
sse_channels: RwLock::new(HashMap::new()), sse_channels: RwLock::new(HashMap::new()),
login_rate_limit: tokio::sync::Mutex::new(state::LoginRateLimit { login_rate_limit: tokio::sync::Mutex::new(state::LoginRateLimit { fail_count: 0, locked_until: None }),
fail_count: 0,
locked_until: None,
}),
}); });
// CORS // CORS
let cors = CorsLayer::new() let cors = CorsLayer::new().allow_origin(Any).allow_methods(Any).allow_headers(Any);
.allow_origin(Any)
.allow_methods(Any)
.allow_headers(Any);
// API routes // API routes
let api = Router::new() let api = Router::new()
@@ -159,25 +144,18 @@ async fn main() {
.with_state(web_state.clone()); .with_state(web_state.clone());
// Build app // Build app
let mut app = Router::new() let mut app = Router::new().nest("/api", api).layer(tower_http::trace::TraceLayer::new_for_http()).layer(cors);
.nest("/api", api)
.layer(tower_http::trace::TraceLayer::new_for_http())
.layer(cors);
// Static file serving // Static file serving
if let Ok(static_dir) = std::env::var("DBX_STATIC_DIR") { if let Ok(static_dir) = std::env::var("DBX_STATIC_DIR") {
use tower_http::services::{ServeDir, ServeFile}; use tower_http::services::{ServeDir, ServeFile};
let index_path = format!("{}/index.html", static_dir); let index_path = format!("{}/index.html", static_dir);
let serve_dir = ServeDir::new(&static_dir) let serve_dir = ServeDir::new(&static_dir).not_found_service(ServeFile::new(&index_path));
.not_found_service(ServeFile::new(&index_path));
app = app.fallback_service(serve_dir); app = app.fallback_service(serve_dir);
} }
// Bind address // Bind address
let port: u16 = std::env::var("DBX_PORT") let port: u16 = std::env::var("DBX_PORT").ok().and_then(|p| p.parse().ok()).unwrap_or(4224);
.ok()
.and_then(|p| p.parse().ok())
.unwrap_or(4224);
let addr = SocketAddr::from(([0, 0, 0, 0], port)); let addr = SocketAddr::from(([0, 0, 0, 0], port));
tracing::info!("DBX Web server starting on http://{}", addr); tracing::info!("DBX Web server starting on http://{}", addr);
@@ -185,10 +163,6 @@ async fn main() {
tracing::info!("Password protection is enabled"); tracing::info!("Password protection is enabled");
} }
let listener = tokio::net::TcpListener::bind(addr) let listener = tokio::net::TcpListener::bind(addr).await.expect("Failed to bind address");
.await axum::serve(listener, app).await.expect("Server error");
.expect("Failed to bind address");
axum::serve(listener, app)
.await
.expect("Server error");
} }
+15 -60
View File
@@ -6,9 +6,7 @@ use axum::Json;
use futures::stream::Stream; use futures::stream::Stream;
use serde::Deserialize; use serde::Deserialize;
use dbx_core::ai::{ use dbx_core::ai::{AiCompletionRequest, AiConfig, AiConversation, AiStreamChunk};
AiCompletionRequest, AiConfig, AiConversation, AiStreamChunk,
};
use crate::error::AppError; use crate::error::AppError;
use crate::state::WebState; use crate::state::WebState;
@@ -62,24 +60,12 @@ pub async fn save_ai_config(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(body): Json<SaveAiConfigRequest>, Json(body): Json<SaveAiConfigRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
state state.app.storage.save_ai_config(&body.config).await.map_err(AppError)?;
.app
.storage
.save_ai_config(&body.config)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
pub async fn load_ai_config( pub async fn load_ai_config(State(state): State<Arc<WebState>>) -> Result<Json<Option<AiConfig>>, AppError> {
State(state): State<Arc<WebState>>, let config = state.app.storage.load_ai_config().await.map_err(AppError)?;
) -> Result<Json<Option<AiConfig>>, AppError> {
let config = state
.app
.storage
.load_ai_config()
.await
.map_err(AppError)?;
Ok(Json(config)) Ok(Json(config))
} }
@@ -91,24 +77,12 @@ pub async fn save_ai_conversation(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(body): Json<SaveAiConversationRequest>, Json(body): Json<SaveAiConversationRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
state state.app.storage.save_ai_conversation(&body.conversation).await.map_err(AppError)?;
.app
.storage
.save_ai_conversation(&body.conversation)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
pub async fn load_ai_conversations( pub async fn load_ai_conversations(State(state): State<Arc<WebState>>) -> Result<Json<Vec<AiConversation>>, AppError> {
State(state): State<Arc<WebState>>, let conversations = state.app.storage.load_ai_conversations().await.map_err(AppError)?;
) -> Result<Json<Vec<AiConversation>>, AppError> {
let conversations = state
.app
.storage
.load_ai_conversations()
.await
.map_err(AppError)?;
Ok(Json(conversations)) Ok(Json(conversations))
} }
@@ -116,12 +90,7 @@ pub async fn delete_ai_conversation(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
state state.app.storage.delete_ai_conversation(&id).await.map_err(AppError)?;
.app
.storage
.delete_ai_conversation(&id)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
@@ -129,12 +98,8 @@ pub async fn delete_ai_conversation(
// AI complete (non-streaming) // AI complete (non-streaming)
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
pub async fn ai_complete( pub async fn ai_complete(Json(body): Json<AiCompleteRequest>) -> Result<Json<String>, AppError> {
Json(body): Json<AiCompleteRequest>, let result = dbx_core::ai::complete(&body.request).await.map_err(AppError)?;
) -> Result<Json<String>, AppError> {
let result = dbx_core::ai::complete(&body.request)
.await
.map_err(AppError)?;
Ok(Json(result)) Ok(Json(result))
} }
@@ -142,12 +107,8 @@ pub async fn ai_complete(
// AI test connection // AI test connection
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
pub async fn ai_test_connection( pub async fn ai_test_connection(Json(body): Json<AiTestConnectionRequest>) -> Result<Json<String>, AppError> {
Json(body): Json<AiTestConnectionRequest>, let result = dbx_core::ai::test_connection_core(&body.config).await.map_err(AppError)?;
) -> Result<Json<String>, AppError> {
let result = dbx_core::ai::test_connection_core(&body.config)
.await
.map_err(AppError)?;
Ok(Json(result)) Ok(Json(result))
} }
@@ -155,9 +116,7 @@ pub async fn ai_test_connection(
// AI cancel stream // AI cancel stream
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
pub async fn ai_cancel_stream( pub async fn ai_cancel_stream(Json(body): Json<AiCancelStreamRequest>) -> Result<Json<bool>, AppError> {
Json(body): Json<AiCancelStreamRequest>,
) -> Result<Json<bool>, AppError> {
let result = dbx_core::ai::cancel_stream(&body.session_id).await; let result = dbx_core::ai::cancel_stream(&body.session_id).await;
Ok(Json(result)) Ok(Json(result))
} }
@@ -184,12 +143,8 @@ pub async fn ai_stream(
.await; .await;
if let Err(_e) = result { if let Err(_e) = result {
let error_chunk = AiStreamChunk { let error_chunk =
session_id: sid.clone(), AiStreamChunk { session_id: sid.clone(), delta: String::new(), reasoning_delta: None, done: true };
delta: String::new(),
reasoning_delta: None,
done: true,
};
let _ = tx.send(serde_json::to_string(&error_chunk).unwrap_or_default()); let _ = tx.send(serde_json::to_string(&error_chunk).unwrap_or_default());
} }
+10 -40
View File
@@ -35,24 +35,15 @@ pub async fn test_connection(
// Store config temporarily // Store config temporarily
let temp_id = format!("__test_{}", uuid::Uuid::new_v4()); let temp_id = format!("__test_{}", uuid::Uuid::new_v4());
app.configs app.configs.lock().await.insert(temp_id.clone(), config.clone());
.lock()
.await
.insert(temp_id.clone(), config.clone());
// Try to connect // Try to connect
let result = app let result = app.get_or_create_pool(&temp_id, config.database.as_deref()).await;
.get_or_create_pool(&temp_id, config.database.as_deref())
.await;
// Clean up any pool keys created for the temporary connection, including // Clean up any pool keys created for the temporary connection, including
// database-scoped keys like "__test_uuid:database". // database-scoped keys like "__test_uuid:database".
let mut connections = app.connections.lock().await; let mut connections = app.connections.lock().await;
let temp_keys: Vec<String> = connections let temp_keys: Vec<String> = connections.keys().filter(|key| key.starts_with(&temp_id)).cloned().collect();
.keys()
.filter(|key| key.starts_with(&temp_id))
.cloned()
.collect();
for key in temp_keys { for key in temp_keys {
connections.remove(&key); connections.remove(&key);
} }
@@ -73,15 +64,9 @@ pub async fn connect_db(
let app = &state.app; let app = &state.app;
let connection_id = config.id.clone(); let connection_id = config.id.clone();
app.configs app.configs.lock().await.insert(connection_id.clone(), config.clone());
.lock()
.await
.insert(connection_id.clone(), config.clone());
let pool_key = app let pool_key = app.get_or_create_pool(&connection_id, config.database.as_deref()).await.map_err(AppError)?;
.get_or_create_pool(&connection_id, config.database.as_deref())
.await
.map_err(AppError)?;
Ok(Json(pool_key)) Ok(Json(pool_key))
} }
@@ -94,11 +79,8 @@ pub async fn disconnect_db(
let mut connections = app.connections.lock().await; let mut connections = app.connections.lock().await;
// Remove all pool keys that start with this connection_id // Remove all pool keys that start with this connection_id
let keys_to_remove: Vec<String> = connections let keys_to_remove: Vec<String> =
.keys() connections.keys().filter(|k| k.starts_with(&body.connection_id)).cloned().collect();
.filter(|k| k.starts_with(&body.connection_id))
.cloned()
.collect();
for key in keys_to_remove { for key in keys_to_remove {
connections.remove(&key); connections.remove(&key);
} }
@@ -114,23 +96,11 @@ pub async fn save_connections(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(body): Json<SaveConnectionsRequest>, Json(body): Json<SaveConnectionsRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
state state.app.storage.save_connections(&body.configs).await.map_err(AppError)?;
.app
.storage
.save_connections(&body.configs)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
pub async fn load_connections( pub async fn load_connections(State(state): State<Arc<WebState>>) -> Result<Json<Vec<ConnectionConfig>>, AppError> {
State(state): State<Arc<WebState>>, let configs = state.app.storage.load_connections().await.map_err(AppError)?;
) -> Result<Json<Vec<ConnectionConfig>>, AppError> {
let configs = state
.app
.storage
.load_connections()
.await
.map_err(AppError)?;
Ok(Json(configs)) Ok(Json(configs))
} }
+3 -18
View File
@@ -24,12 +24,7 @@ pub async fn save_history(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(body): Json<SaveHistoryRequest>, Json(body): Json<SaveHistoryRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
state state.app.storage.save_history_entry(&body.entry).await.map_err(AppError)?;
.app
.storage
.save_history_entry(&body.entry)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
@@ -39,12 +34,7 @@ pub async fn load_history(
) -> Result<Json<Vec<HistoryEntry>>, AppError> { ) -> Result<Json<Vec<HistoryEntry>>, AppError> {
let limit = q.limit.unwrap_or(100); let limit = q.limit.unwrap_or(100);
let offset = q.offset.unwrap_or(0); let offset = q.offset.unwrap_or(0);
let entries = state let entries = state.app.storage.load_history_entries(limit, offset).await.map_err(AppError)?;
.app
.storage
.load_history_entries(limit, offset)
.await
.map_err(AppError)?;
Ok(Json(entries)) Ok(Json(entries))
} }
@@ -57,11 +47,6 @@ pub async fn delete_history_entry(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
state state.app.storage.delete_history_entry(&id).await.map_err(AppError)?;
.app
.storage
.delete_history_entry(&id)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
+3 -15
View File
@@ -17,23 +17,11 @@ pub async fn save_sidebar_layout(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(body): Json<SaveLayoutRequest>, Json(body): Json<SaveLayoutRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
state state.app.storage.save_sidebar_layout(&body.layout).await.map_err(AppError)?;
.app
.storage
.save_sidebar_layout(&body.layout)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
pub async fn load_sidebar_layout( pub async fn load_sidebar_layout(State(state): State<Arc<WebState>>) -> Result<Json<serde_json::Value>, AppError> {
State(state): State<Arc<WebState>>, let layout = state.app.storage.load_sidebar_layout().await.map_err(AppError)?;
) -> Result<Json<serde_json::Value>, AppError> {
let layout = state
.app
.storage
.load_sidebar_layout()
.await
.map_err(AppError)?;
Ok(Json(layout.unwrap_or(serde_json::json!(null)))) Ok(Json(layout.unwrap_or(serde_json::json!(null))))
} }
+5 -10
View File
@@ -62,9 +62,8 @@ pub async fn list_databases(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<MongoConnectionRequest>, Json(req): Json<MongoConnectionRequest>,
) -> Result<Json<Vec<String>>, AppError> { ) -> Result<Json<Vec<String>>, AppError> {
let result = dbx_core::mongo_ops::mongo_list_databases_core(&state.app, &req.connection_id) let result =
.await dbx_core::mongo_ops::mongo_list_databases_core(&state.app, &req.connection_id).await.map_err(AppError)?;
.map_err(AppError)?;
Ok(Json(result)) Ok(Json(result))
} }
@@ -72,13 +71,9 @@ pub async fn list_collections(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<MongoCollectionRequest>, Json(req): Json<MongoCollectionRequest>,
) -> Result<Json<Vec<String>>, AppError> { ) -> Result<Json<Vec<String>>, AppError> {
let result = dbx_core::mongo_ops::mongo_list_collections_core( let result = dbx_core::mongo_ops::mongo_list_collections_core(&state.app, &req.connection_id, &req.database)
&state.app, .await
&req.connection_id, .map_err(AppError)?;
&req.database,
)
.await
.map_err(AppError)?;
Ok(Json(result)) Ok(Json(result))
} }
+8 -22
View File
@@ -34,9 +34,7 @@ pub async fn execute_query(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<ExecuteQueryRequest>, Json(req): Json<ExecuteQueryRequest>,
) -> Result<Json<serde_json::Value>, AppError> { ) -> Result<Json<serde_json::Value>, AppError> {
let execution_id = req let execution_id = req.execution_id.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
.execution_id
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
let registered = state.app.running_queries.register(execution_id); let registered = state.app.running_queries.register(execution_id);
let cancel_token = registered.token(); let cancel_token = registered.token();
@@ -59,9 +57,7 @@ pub async fn execute_multi(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<ExecuteQueryRequest>, Json(req): Json<ExecuteQueryRequest>,
) -> Result<Json<serde_json::Value>, AppError> { ) -> Result<Json<serde_json::Value>, AppError> {
let execution_id = req let execution_id = req.execution_id.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
.execution_id
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
let registered = state.app.running_queries.register(execution_id); let registered = state.app.running_queries.register(execution_id);
let cancel_token = registered.token(); let cancel_token = registered.token();
@@ -84,14 +80,9 @@ pub async fn execute_batch(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<ExecuteBatchRequest>, Json(req): Json<ExecuteBatchRequest>,
) -> Result<Json<serde_json::Value>, AppError> { ) -> Result<Json<serde_json::Value>, AppError> {
let result = dbx_core::query::execute_statements( let result = dbx_core::query::execute_statements(&state.app, &req.connection_id, &req.database, &req.statements)
&state.app, .await
&req.connection_id, .map_err(AppError)?;
&req.database,
&req.statements,
)
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -109,14 +100,9 @@ pub async fn execute_script(
Json(req): Json<ExecuteQueryRequest>, Json(req): Json<ExecuteQueryRequest>,
) -> Result<Json<serde_json::Value>, AppError> { ) -> Result<Json<serde_json::Value>, AppError> {
let statements = dbx_core::sql::split_sql_statements(&req.sql); let statements = dbx_core::sql::split_sql_statements(&req.sql);
let result = dbx_core::query::execute_statements( let result = dbx_core::query::execute_statements(&state.app, &req.connection_id, &req.database, &statements)
&state.app, .await
&req.connection_id, .map_err(AppError)?;
&req.database,
&statements,
)
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
+17 -43
View File
@@ -69,9 +69,8 @@ pub async fn list_databases(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<RedisConnectionRequest>, Json(req): Json<RedisConnectionRequest>,
) -> Result<Json<serde_json::Value>, AppError> { ) -> Result<Json<serde_json::Value>, AppError> {
let result = dbx_core::redis_ops::redis_list_databases_core(&state.app, &req.connection_id) let result =
.await dbx_core::redis_ops::redis_list_databases_core(&state.app, &req.connection_id).await.map_err(AppError)?;
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -96,9 +95,8 @@ pub async fn get_value(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<RedisKeyRequest>, Json(req): Json<RedisKeyRequest>,
) -> Result<Json<serde_json::Value>, AppError> { ) -> Result<Json<serde_json::Value>, AppError> {
let result = dbx_core::redis_ops::redis_get_value_core(&state.app, &req.connection_id, &req.key) let result =
.await dbx_core::redis_ops::redis_get_value_core(&state.app, &req.connection_id, &req.key).await.map_err(AppError)?;
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -106,15 +104,9 @@ pub async fn set_string(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<RedisSetStringRequest>, Json(req): Json<RedisSetStringRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
dbx_core::redis_ops::redis_set_string_core( dbx_core::redis_ops::redis_set_string_core(&state.app, &req.connection_id, &req.key, &req.value, req.ttl)
&state.app, .await
&req.connection_id, .map_err(AppError)?;
&req.key,
&req.value,
req.ttl,
)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
@@ -122,9 +114,7 @@ pub async fn delete_key(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<RedisKeyRequest>, Json(req): Json<RedisKeyRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
dbx_core::redis_ops::redis_delete_key_core(&state.app, &req.connection_id, &req.key) dbx_core::redis_ops::redis_delete_key_core(&state.app, &req.connection_id, &req.key).await.map_err(AppError)?;
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
@@ -133,15 +123,9 @@ pub async fn hash_set(
Json(req): Json<RedisHashRequest>, Json(req): Json<RedisHashRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
let value = req.value.as_deref().unwrap_or(""); let value = req.value.as_deref().unwrap_or("");
dbx_core::redis_ops::redis_hash_set_core( dbx_core::redis_ops::redis_hash_set_core(&state.app, &req.connection_id, &req.key, &req.field, value)
&state.app, .await
&req.connection_id, .map_err(AppError)?;
&req.key,
&req.field,
value,
)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
@@ -149,14 +133,9 @@ pub async fn hash_del(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<RedisHashRequest>, Json(req): Json<RedisHashRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
dbx_core::redis_ops::redis_hash_del_core( dbx_core::redis_ops::redis_hash_del_core(&state.app, &req.connection_id, &req.key, &req.field)
&state.app, .await
&req.connection_id, .map_err(AppError)?;
&req.key,
&req.field,
)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
@@ -196,13 +175,8 @@ pub async fn set_remove(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Json(req): Json<RedisSetRequest>, Json(req): Json<RedisSetRequest>,
) -> Result<Json<()>, AppError> { ) -> Result<Json<()>, AppError> {
dbx_core::redis_ops::redis_set_remove_core( dbx_core::redis_ops::redis_set_remove_core(&state.app, &req.connection_id, &req.key, &req.member)
&state.app, .await
&req.connection_id, .map_err(AppError)?;
&req.key,
&req.member,
)
.await
.map_err(AppError)?;
Ok(Json(())) Ok(Json(()))
} }
+18 -54
View File
@@ -19,9 +19,7 @@ pub async fn list_databases(
State(state): State<Arc<WebState>>, State(state): State<Arc<WebState>>,
Query(q): Query<SchemaQuery>, Query(q): Query<SchemaQuery>,
) -> Result<Json<serde_json::Value>, AppError> { ) -> Result<Json<serde_json::Value>, AppError> {
let result = dbx_core::schema::list_databases_core(&state.app, &q.connection_id) let result = dbx_core::schema::list_databases_core(&state.app, &q.connection_id).await.map_err(AppError)?;
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -30,9 +28,7 @@ pub async fn list_schemas(
Query(q): Query<SchemaQuery>, Query(q): Query<SchemaQuery>,
) -> Result<Json<Vec<String>>, AppError> { ) -> Result<Json<Vec<String>>, AppError> {
let database = q.database.as_deref().unwrap_or(""); let database = q.database.as_deref().unwrap_or("");
let result = dbx_core::schema::list_schemas_core(&state.app, &q.connection_id, database) let result = dbx_core::schema::list_schemas_core(&state.app, &q.connection_id, database).await.map_err(AppError)?;
.await
.map_err(AppError)?;
Ok(Json(result)) Ok(Json(result))
} }
@@ -43,9 +39,7 @@ pub async fn list_tables(
let database = q.database.as_deref().unwrap_or(""); let database = q.database.as_deref().unwrap_or("");
let schema = q.schema.as_deref().unwrap_or(""); let schema = q.schema.as_deref().unwrap_or("");
let result = let result =
dbx_core::schema::list_tables_core(&state.app, &q.connection_id, database, schema) dbx_core::schema::list_tables_core(&state.app, &q.connection_id, database, schema).await.map_err(AppError)?;
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -56,15 +50,9 @@ pub async fn list_columns(
let database = q.database.as_deref().unwrap_or(""); let database = q.database.as_deref().unwrap_or("");
let schema = q.schema.as_deref().unwrap_or(""); let schema = q.schema.as_deref().unwrap_or("");
let table = q.table.as_deref().unwrap_or(""); let table = q.table.as_deref().unwrap_or("");
let result = dbx_core::schema::get_columns_core( let result = dbx_core::schema::get_columns_core(&state.app, &q.connection_id, database, schema, table)
&state.app, .await
&q.connection_id, .map_err(AppError)?;
database,
schema,
table,
)
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -75,15 +63,9 @@ pub async fn list_indexes(
let database = q.database.as_deref().unwrap_or(""); let database = q.database.as_deref().unwrap_or("");
let schema = q.schema.as_deref().unwrap_or(""); let schema = q.schema.as_deref().unwrap_or("");
let table = q.table.as_deref().unwrap_or(""); let table = q.table.as_deref().unwrap_or("");
let result = dbx_core::schema::list_indexes_core( let result = dbx_core::schema::list_indexes_core(&state.app, &q.connection_id, database, schema, table)
&state.app, .await
&q.connection_id, .map_err(AppError)?;
database,
schema,
table,
)
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -94,15 +76,9 @@ pub async fn list_foreign_keys(
let database = q.database.as_deref().unwrap_or(""); let database = q.database.as_deref().unwrap_or("");
let schema = q.schema.as_deref().unwrap_or(""); let schema = q.schema.as_deref().unwrap_or("");
let table = q.table.as_deref().unwrap_or(""); let table = q.table.as_deref().unwrap_or("");
let result = dbx_core::schema::list_foreign_keys_core( let result = dbx_core::schema::list_foreign_keys_core(&state.app, &q.connection_id, database, schema, table)
&state.app, .await
&q.connection_id, .map_err(AppError)?;
database,
schema,
table,
)
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -113,15 +89,9 @@ pub async fn list_triggers(
let database = q.database.as_deref().unwrap_or(""); let database = q.database.as_deref().unwrap_or("");
let schema = q.schema.as_deref().unwrap_or(""); let schema = q.schema.as_deref().unwrap_or("");
let table = q.table.as_deref().unwrap_or(""); let table = q.table.as_deref().unwrap_or("");
let result = dbx_core::schema::list_triggers_core( let result = dbx_core::schema::list_triggers_core(&state.app, &q.connection_id, database, schema, table)
&state.app, .await
&q.connection_id, .map_err(AppError)?;
database,
schema,
table,
)
.await
.map_err(AppError)?;
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?)) Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
} }
@@ -132,14 +102,8 @@ pub async fn get_ddl(
let database = q.database.as_deref().unwrap_or(""); let database = q.database.as_deref().unwrap_or("");
let schema = q.schema.as_deref().unwrap_or(""); let schema = q.schema.as_deref().unwrap_or("");
let table = q.table.as_deref().unwrap_or(""); let table = q.table.as_deref().unwrap_or("");
let result = dbx_core::schema::get_table_ddl_core( let result = dbx_core::schema::get_table_ddl_core(&state.app, &q.connection_id, database, schema, table)
&state.app, .await
&q.connection_id, .map_err(AppError)?;
database,
schema,
table,
)
.await
.map_err(AppError)?;
Ok(Json(result)) Ok(Json(result))
} }
+6 -21
View File
@@ -3,8 +3,8 @@ use std::sync::Arc;
use axum::extract::{Multipart, Path, State}; use axum::extract::{Multipart, Path, State};
use axum::response::sse::{Event, Sse}; use axum::response::sse::{Event, Sse};
use axum::Json; use axum::Json;
use dbx_core::sql;
use dbx_core::query; use dbx_core::query;
use dbx_core::sql;
use futures::stream::Stream; use futures::stream::Stream;
use serde::Deserialize; use serde::Deserialize;
@@ -40,15 +40,8 @@ pub async fn preview_sql_file(
let tmp_dir = state.data_dir.join("tmp"); let tmp_dir = state.data_dir.join("tmp");
std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?; std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?;
while let Some(field) = multipart while let Some(field) = multipart.next_field().await.map_err(|e| AppError(e.to_string()))? {
.next_field() let file_name = field.file_name().unwrap_or("upload.sql").to_string();
.await
.map_err(|e| AppError(e.to_string()))?
{
let file_name = field
.file_name()
.unwrap_or("upload.sql")
.to_string();
let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?; let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?;
let file_path = tmp_dir.join(&file_name); let file_path = tmp_dir.join(&file_name);
@@ -77,11 +70,7 @@ pub async fn execute_sql_file(
let execution_id = req.execution_id.clone(); let execution_id = req.execution_id.clone();
let (tx, _) = tokio::sync::broadcast::channel::<String>(256); let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
state state.sse_channels.write().await.insert(execution_id.clone(), tx.clone());
.sse_channels
.write()
.await
.insert(execution_id.clone(), tx.clone());
let app = state.app.clone(); let app = state.app.clone();
let state_clone = state.clone(); let state_clone = state.clone();
@@ -149,9 +138,7 @@ pub async fn execute_sql_file(
let _ = tx.send(json); let _ = tx.send(json);
} }
match query::execute_sql_statement(&app, &req.connection_id, &req.database, stmt, None) match query::execute_sql_statement(&app, &req.connection_id, &req.database, stmt, None).await {
.await
{
Ok(result) => { Ok(result) => {
success_count += 1; success_count += 1;
total_affected += result.affected_rows; total_affected += result.affected_rows;
@@ -220,9 +207,7 @@ pub async fn sql_file_progress(
Path(execution_id): Path<String>, Path(execution_id): Path<String>,
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> { ) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
let channels = state.sse_channels.read().await; let channels = state.sse_channels.read().await;
let tx = channels let tx = channels.get(&execution_id).ok_or_else(|| AppError("Execution not found".to_string()))?;
.get(&execution_id)
.ok_or_else(|| AppError("Execution not found".to_string()))?;
let rx = tx.subscribe(); let rx = tx.subscribe();
drop(channels); drop(channels);
Ok(crate::sse::sse_from_channel(rx)) Ok(crate::sse::sse_from_channel(rx))
+9 -32
View File
@@ -3,9 +3,7 @@ use std::sync::Arc;
use axum::extract::{Multipart, Path, State}; use axum::extract::{Multipart, Path, State};
use axum::response::sse::{Event, Sse}; use axum::response::sse::{Event, Sse};
use axum::Json; use axum::Json;
use dbx_core::table_import::{ use dbx_core::table_import::{self, TableImportRequest};
self, TableImportRequest,
};
use dbx_core::transfer; use dbx_core::transfer;
use futures::stream::Stream; use futures::stream::Stream;
use serde::Deserialize; use serde::Deserialize;
@@ -32,27 +30,17 @@ pub async fn preview_import(
let tmp_dir = state.data_dir.join("tmp"); let tmp_dir = state.data_dir.join("tmp");
std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?; std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?;
while let Some(field) = multipart while let Some(field) = multipart.next_field().await.map_err(|e| AppError(e.to_string()))? {
.next_field() let file_name = field.file_name().unwrap_or("upload.csv").to_string();
.await
.map_err(|e| AppError(e.to_string()))?
{
let file_name = field
.file_name()
.unwrap_or("upload.csv")
.to_string();
let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?; let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?;
let file_path = tmp_dir.join(&file_name); let file_path = tmp_dir.join(&file_name);
std::fs::write(&file_path, &data).map_err(|e| AppError(e.to_string()))?; std::fs::write(&file_path, &data).map_err(|e| AppError(e.to_string()))?;
let file_path_str = file_path.to_string_lossy().to_string(); let file_path_str = file_path.to_string_lossy().to_string();
let preview = table_import::preview_table_import_file_core(&file_path_str) let preview = table_import::preview_table_import_file_core(&file_path_str).map_err(AppError)?;
.map_err(AppError)?;
return Ok(Json( return Ok(Json(serde_json::to_value(preview).map_err(|e| AppError(e.to_string()))?));
serde_json::to_value(preview).map_err(|e| AppError(e.to_string()))?,
));
} }
Err(AppError("No file uploaded".to_string())) Err(AppError("No file uploaded".to_string()))
@@ -66,11 +54,7 @@ pub async fn execute_import(
let import_id = req.import_id.clone(); let import_id = req.import_id.clone();
let (tx, _) = tokio::sync::broadcast::channel::<String>(256); let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
state state.sse_channels.write().await.insert(import_id.clone(), tx.clone());
.sse_channels
.write()
.await
.insert(import_id.clone(), tx.clone());
let app = state.app.clone(); let app = state.app.clone();
let state_clone = state.clone(); let state_clone = state.clone();
@@ -91,10 +75,7 @@ pub async fn execute_import(
} }
}; };
let pool_key = match app let pool_key = match app.get_or_create_pool(&req.connection_id, Some(&req.database)).await {
.get_or_create_pool(&req.connection_id, Some(&req.database))
.await
{
Ok(k) => k, Ok(k) => k,
Err(e) => { Err(e) => {
let _ = tx.send( let _ = tx.send(
@@ -118,9 +99,7 @@ pub async fn execute_import(
&pool_key, &pool_key,
|id: &str| { |id: &str| {
let id = id.to_string(); let id = id.to_string();
Box::pin(async move { Box::pin(async move { transfer::is_cancelled(&id).await })
transfer::is_cancelled(&id).await
})
}, },
|progress| { |progress| {
if let Ok(json) = serde_json::to_string(&progress) { if let Ok(json) = serde_json::to_string(&progress) {
@@ -159,9 +138,7 @@ pub async fn import_progress(
Path(import_id): Path<String>, Path(import_id): Path<String>,
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> { ) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
let channels = state.sse_channels.read().await; let channels = state.sse_channels.read().await;
let tx = channels let tx = channels.get(&import_id).ok_or_else(|| AppError("Import not found".to_string()))?;
.get(&import_id)
.ok_or_else(|| AppError("Import not found".to_string()))?;
let rx = tx.subscribe(); let rx = tx.subscribe();
drop(channels); drop(channels);
Ok(crate::sse::sse_from_channel(rx)) Ok(crate::sse::sse_from_channel(rx))
+5 -17
View File
@@ -31,11 +31,7 @@ pub async fn start_transfer(
// Create a broadcast channel for progress // Create a broadcast channel for progress
let (tx, _) = tokio::sync::broadcast::channel::<String>(256); let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
state state.sse_channels.write().await.insert(transfer_id.clone(), tx.clone());
.sse_channels
.write()
.await
.insert(transfer_id.clone(), tx.clone());
let app = state.app.clone(); let app = state.app.clone();
let state_clone = state.clone(); let state_clone = state.clone();
@@ -56,9 +52,7 @@ pub async fn start_transfer(
} }
}; };
let source_pool_key = match app let source_pool_key = match app.get_or_create_pool(&req.source_connection_id, Some(&req.source_database)).await
.get_or_create_pool(&req.source_connection_id, Some(&req.source_database))
.await
{ {
Ok(k) => k, Ok(k) => k,
Err(e) => { Err(e) => {
@@ -66,9 +60,7 @@ pub async fn start_transfer(
return; return;
} }
}; };
let target_pool_key = match app let target_pool_key = match app.get_or_create_pool(&req.target_connection_id, Some(&req.target_database)).await
.get_or_create_pool(&req.target_connection_id, Some(&req.target_database))
.await
{ {
Ok(k) => k, Ok(k) => k,
Err(e) => { Err(e) => {
@@ -151,9 +143,7 @@ pub async fn start_transfer(
state_clone.sse_channels.write().await.remove(&req.transfer_id); state_clone.sse_channels.write().await.remove(&req.transfer_id);
}); });
Ok(Json( Ok(Json(serde_json::json!({ "transferId": transfer_id })))
serde_json::json!({ "transferId": transfer_id }),
))
} }
pub async fn transfer_progress( pub async fn transfer_progress(
@@ -161,9 +151,7 @@ pub async fn transfer_progress(
Path(transfer_id): Path<String>, Path(transfer_id): Path<String>,
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> { ) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
let channels = state.sse_channels.read().await; let channels = state.sse_channels.read().await;
let tx = channels let tx = channels.get(&transfer_id).ok_or_else(|| AppError("Transfer not found".to_string()))?;
.get(&transfer_id)
.ok_or_else(|| AppError("Transfer not found".to_string()))?;
let rx = tx.subscribe(); let rx = tx.subscribe();
drop(channels); drop(channels);
Ok(crate::sse::sse_from_channel(rx)) Ok(crate::sse::sse_from_channel(rx))