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