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
|
// Stream cancel registry
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
static AI_STREAMS: LazyLock<RwLock<HashMap<String, Arc<AtomicBool>>>> =
|
static AI_STREAMS: LazyLock<RwLock<HashMap<String, Arc<AtomicBool>>>> = LazyLock::new(|| RwLock::new(HashMap::new()));
|
||||||
LazyLock::new(|| RwLock::new(HashMap::new()));
|
|
||||||
|
|
||||||
pub async fn register_stream(session_id: &str) -> Arc<AtomicBool> {
|
pub async fn register_stream(session_id: &str) -> Arc<AtomicBool> {
|
||||||
let cancelled = Arc::new(AtomicBool::new(false));
|
let cancelled = Arc::new(AtomicBool::new(false));
|
||||||
AI_STREAMS
|
AI_STREAMS.write().await.insert(session_id.to_string(), cancelled.clone());
|
||||||
.write()
|
|
||||||
.await
|
|
||||||
.insert(session_id.to_string(), cancelled.clone());
|
|
||||||
cancelled
|
cancelled
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -129,10 +125,7 @@ pub struct AiConversation {
|
|||||||
|
|
||||||
pub fn resolve_endpoint(config: &AiConfig) -> String {
|
pub fn resolve_endpoint(config: &AiConfig) -> String {
|
||||||
let ep = config.endpoint.trim().trim_end_matches('/');
|
let ep = config.endpoint.trim().trim_end_matches('/');
|
||||||
if ep.ends_with("/chat/completions")
|
if ep.ends_with("/chat/completions") || ep.ends_with("/responses") || ep.ends_with("/messages") {
|
||||||
|| ep.ends_with("/responses")
|
|
||||||
|| ep.ends_with("/messages")
|
|
||||||
{
|
|
||||||
return ep.to_string();
|
return ep.to_string();
|
||||||
}
|
}
|
||||||
match config.provider {
|
match config.provider {
|
||||||
@@ -149,11 +142,7 @@ pub fn resolve_endpoint(config: &AiConfig) -> String {
|
|||||||
|
|
||||||
pub fn stream_data_payload(line: &str) -> Option<&str> {
|
pub fn stream_data_payload(line: &str) -> Option<&str> {
|
||||||
let line = line.trim();
|
let line = line.trim();
|
||||||
if line.is_empty()
|
if line.is_empty() || line.starts_with(':') || line.starts_with("event:") || line.starts_with("id:") {
|
||||||
|| line.starts_with(':')
|
|
||||||
|| line.starts_with("event:")
|
|
||||||
|| line.starts_with("id:")
|
|
||||||
{
|
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
if let Some(data) = line.strip_prefix("data:") {
|
if let Some(data) = line.strip_prefix("data:") {
|
||||||
@@ -175,11 +164,7 @@ pub fn claude_stream_text(event: &serde_json::Value) -> Option<&str> {
|
|||||||
pub fn openai_stream_text(event: &serde_json::Value) -> Option<&str> {
|
pub fn openai_stream_text(event: &serde_json::Value) -> Option<&str> {
|
||||||
event["choices"]
|
event["choices"]
|
||||||
.get(0)
|
.get(0)
|
||||||
.and_then(|choice| {
|
.and_then(|choice| choice["delta"]["content"].as_str().or_else(|| choice["message"]["content"].as_str()))
|
||||||
choice["delta"]["content"]
|
|
||||||
.as_str()
|
|
||||||
.or_else(|| choice["message"]["content"].as_str())
|
|
||||||
})
|
|
||||||
.or_else(|| event["content"].as_str())
|
.or_else(|| event["content"].as_str())
|
||||||
.filter(|text| !text.is_empty())
|
.filter(|text| !text.is_empty())
|
||||||
}
|
}
|
||||||
@@ -196,10 +181,7 @@ pub fn responses_stream_text(event: &serde_json::Value) -> Option<&str> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_error(data: &serde_json::Value) -> Option<String> {
|
pub fn extract_error(data: &serde_json::Value) -> Option<String> {
|
||||||
data["error"]["message"]
|
data["error"]["message"].as_str().or_else(|| data["error"].as_str()).map(ToString::to_string)
|
||||||
.as_str()
|
|
||||||
.or_else(|| data["error"].as_str())
|
|
||||||
.map(ToString::to_string)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn build_responses_input(system_prompt: &str, messages: &[AiMessage]) -> serde_json::Value {
|
pub fn build_responses_input(system_prompt: &str, messages: &[AiMessage]) -> serde_json::Value {
|
||||||
@@ -240,16 +222,10 @@ fn validate_config(config: &AiConfig) -> Result<(), String> {
|
|||||||
// Non-streaming calls
|
// Non-streaming calls
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
pub async fn call_claude(
|
pub async fn call_claude(client: &reqwest::Client, request: AiCompletionRequest) -> Result<String, String> {
|
||||||
client: &reqwest::Client,
|
|
||||||
request: AiCompletionRequest,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let mut headers = HeaderMap::new();
|
let mut headers = HeaderMap::new();
|
||||||
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
||||||
headers.insert(
|
headers.insert("x-api-key", HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?);
|
||||||
"x-api-key",
|
|
||||||
HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?,
|
|
||||||
);
|
|
||||||
headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
|
headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
|
||||||
|
|
||||||
let body = json!({
|
let body = json!({
|
||||||
@@ -281,25 +257,16 @@ pub async fn call_claude(
|
|||||||
.to_string())
|
.to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn call_openai_compatible(
|
pub async fn call_openai_compatible(client: &reqwest::Client, request: AiCompletionRequest) -> Result<String, String> {
|
||||||
client: &reqwest::Client,
|
|
||||||
request: AiCompletionRequest,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let mut headers = HeaderMap::new();
|
let mut headers = HeaderMap::new();
|
||||||
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
||||||
headers.insert(
|
headers.insert(
|
||||||
AUTHORIZATION,
|
AUTHORIZATION,
|
||||||
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key))
|
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
|
||||||
.map_err(|e| e.to_string())?,
|
|
||||||
);
|
);
|
||||||
|
|
||||||
let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })];
|
let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })];
|
||||||
messages.extend(
|
messages.extend(request.messages.iter().map(|message| json!({ "role": message.role, "content": message.content })));
|
||||||
request
|
|
||||||
.messages
|
|
||||||
.iter()
|
|
||||||
.map(|message| json!({ "role": message.role, "content": message.content })),
|
|
||||||
);
|
|
||||||
|
|
||||||
let body = json!({
|
let body = json!({
|
||||||
"model": request.config.model,
|
"model": request.config.model,
|
||||||
@@ -322,22 +289,15 @@ pub async fn call_openai_compatible(
|
|||||||
return Err(extract_error(&data).unwrap_or_else(|| format!("API error: {status}")));
|
return Err(extract_error(&data).unwrap_or_else(|| format!("API error: {status}")));
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(data["choices"][0]["message"]["content"]
|
Ok(data["choices"][0]["message"]["content"].as_str().unwrap_or_default().to_string())
|
||||||
.as_str()
|
|
||||||
.unwrap_or_default()
|
|
||||||
.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn call_responses_api(
|
pub async fn call_responses_api(client: &reqwest::Client, request: AiCompletionRequest) -> Result<String, String> {
|
||||||
client: &reqwest::Client,
|
|
||||||
request: AiCompletionRequest,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let mut headers = HeaderMap::new();
|
let mut headers = HeaderMap::new();
|
||||||
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
||||||
headers.insert(
|
headers.insert(
|
||||||
AUTHORIZATION,
|
AUTHORIZATION,
|
||||||
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key))
|
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
|
||||||
.map_err(|e| e.to_string())?,
|
|
||||||
);
|
);
|
||||||
|
|
||||||
let body = json!({
|
let body = json!({
|
||||||
@@ -365,9 +325,7 @@ pub async fn call_responses_api(
|
|||||||
.as_array()
|
.as_array()
|
||||||
.and_then(|items| {
|
.and_then(|items| {
|
||||||
items.iter().find_map(|item| {
|
items.iter().find_map(|item| {
|
||||||
item["content"]
|
item["content"].as_array().and_then(|parts| parts.iter().find_map(|p| p["text"].as_str()))
|
||||||
.as_array()
|
|
||||||
.and_then(|parts| parts.iter().find_map(|p| p["text"].as_str()))
|
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
@@ -381,18 +339,13 @@ pub async fn call_responses_api(
|
|||||||
pub async fn test_connection_core(config: &AiConfig) -> Result<String, String> {
|
pub async fn test_connection_core(config: &AiConfig) -> Result<String, String> {
|
||||||
validate_config(config)?;
|
validate_config(config)?;
|
||||||
|
|
||||||
let client = reqwest::Client::builder()
|
let client =
|
||||||
.timeout(std::time::Duration::from_secs(15))
|
reqwest::Client::builder().timeout(std::time::Duration::from_secs(15)).build().map_err(|e| e.to_string())?;
|
||||||
.build()
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let request = AiCompletionRequest {
|
let request = AiCompletionRequest {
|
||||||
config: config.clone(),
|
config: config.clone(),
|
||||||
system_prompt: String::new(),
|
system_prompt: String::new(),
|
||||||
messages: vec![AiMessage {
|
messages: vec![AiMessage { role: "user".into(), content: "hi".into() }],
|
||||||
role: "user".into(),
|
|
||||||
content: "hi".into(),
|
|
||||||
}],
|
|
||||||
max_tokens: Some(1),
|
max_tokens: Some(1),
|
||||||
temperature: Some(0.0),
|
temperature: Some(0.0),
|
||||||
};
|
};
|
||||||
@@ -413,10 +366,8 @@ pub async fn test_connection_core(config: &AiConfig) -> Result<String, String> {
|
|||||||
pub async fn complete(request: &AiCompletionRequest) -> Result<String, String> {
|
pub async fn complete(request: &AiCompletionRequest) -> Result<String, String> {
|
||||||
validate_config(&request.config)?;
|
validate_config(&request.config)?;
|
||||||
|
|
||||||
let client = reqwest::Client::builder()
|
let client =
|
||||||
.timeout(std::time::Duration::from_secs(60))
|
reqwest::Client::builder().timeout(std::time::Duration::from_secs(60)).build().map_err(|e| e.to_string())?;
|
||||||
.build()
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
match request.config.provider {
|
match request.config.provider {
|
||||||
AiProvider::Claude => call_claude(&client, request.clone()).await,
|
AiProvider::Claude => call_claude(&client, request.clone()).await,
|
||||||
@@ -442,15 +393,11 @@ pub async fn stream(
|
|||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
validate_config(&request.config)?;
|
validate_config(&request.config)?;
|
||||||
|
|
||||||
let client = reqwest::Client::builder()
|
let client =
|
||||||
.timeout(std::time::Duration::from_secs(120))
|
reqwest::Client::builder().timeout(std::time::Duration::from_secs(120)).build().map_err(|e| e.to_string())?;
|
||||||
.build()
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
match request.config.provider {
|
match request.config.provider {
|
||||||
AiProvider::Claude => {
|
AiProvider::Claude => stream_claude(&client, session_id, request, cancelled, &on_chunk).await,
|
||||||
stream_claude(&client, session_id, request, cancelled, &on_chunk).await
|
|
||||||
}
|
|
||||||
AiProvider::Openai | AiProvider::Custom => {
|
AiProvider::Openai | AiProvider::Custom => {
|
||||||
if request.config.api_style == AiApiStyle::Responses {
|
if request.config.api_style == AiApiStyle::Responses {
|
||||||
stream_responses_api(&client, session_id, request, cancelled, &on_chunk).await
|
stream_responses_api(&client, session_id, request, cancelled, &on_chunk).await
|
||||||
@@ -470,10 +417,7 @@ async fn stream_claude(
|
|||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
let mut headers = HeaderMap::new();
|
let mut headers = HeaderMap::new();
|
||||||
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
||||||
headers.insert(
|
headers.insert("x-api-key", HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?);
|
||||||
"x-api-key",
|
|
||||||
HeaderValue::from_str(&request.config.api_key).map_err(|e| e.to_string())?,
|
|
||||||
);
|
|
||||||
headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
|
headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
|
||||||
|
|
||||||
let body = json!({
|
let body = json!({
|
||||||
@@ -559,17 +503,11 @@ async fn stream_openai(
|
|||||||
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
||||||
headers.insert(
|
headers.insert(
|
||||||
AUTHORIZATION,
|
AUTHORIZATION,
|
||||||
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key))
|
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
|
||||||
.map_err(|e| e.to_string())?,
|
|
||||||
);
|
);
|
||||||
|
|
||||||
let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })];
|
let mut messages = vec![json!({ "role": "system", "content": request.system_prompt })];
|
||||||
messages.extend(
|
messages.extend(request.messages.iter().map(|m| json!({ "role": m.role, "content": m.content })));
|
||||||
request
|
|
||||||
.messages
|
|
||||||
.iter()
|
|
||||||
.map(|m| json!({ "role": m.role, "content": m.content })),
|
|
||||||
);
|
|
||||||
|
|
||||||
let body = json!({
|
let body = json!({
|
||||||
"model": request.config.model,
|
"model": request.config.model,
|
||||||
@@ -661,8 +599,7 @@ async fn stream_responses_api(
|
|||||||
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
||||||
headers.insert(
|
headers.insert(
|
||||||
AUTHORIZATION,
|
AUTHORIZATION,
|
||||||
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key))
|
HeaderValue::from_str(&format!("Bearer {}", request.config.api_key)).map_err(|e| e.to_string())?,
|
||||||
.map_err(|e| e.to_string())?,
|
|
||||||
);
|
);
|
||||||
|
|
||||||
let body = json!({
|
let body = json!({
|
||||||
@@ -771,10 +708,7 @@ pub fn load_conversations(path: &Path) -> Result<Vec<AiConversation>, String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn delete_conversation(path: &Path, id: &str) -> Result<(), String> {
|
pub fn delete_conversation(path: &Path, id: &str) -> Result<(), String> {
|
||||||
let conversations: Vec<AiConversation> = read_conversations(path)?
|
let conversations: Vec<AiConversation> = read_conversations(path)?.into_iter().filter(|c| c.id != id).collect();
|
||||||
.into_iter()
|
|
||||||
.filter(|c| c.id != id)
|
|
||||||
.collect();
|
|
||||||
write_conversations(path, &conversations)
|
write_conversations(path, &conversations)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -49,20 +49,13 @@ impl AppState {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_or_create_pool(
|
pub async fn get_or_create_pool(&self, connection_id: &str, database: Option<&str>) -> Result<String, String> {
|
||||||
&self,
|
|
||||||
connection_id: &str,
|
|
||||||
database: Option<&str>,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let db_type = {
|
let db_type = {
|
||||||
let configs = self.configs.lock().await;
|
let configs = self.configs.lock().await;
|
||||||
configs.get(connection_id).map(|c| c.db_type.clone())
|
configs.get(connection_id).map(|c| c.db_type.clone())
|
||||||
};
|
};
|
||||||
|
|
||||||
let is_embedded = matches!(
|
let is_embedded = matches!(db_type, Some(DatabaseType::Sqlite) | Some(DatabaseType::DuckDb));
|
||||||
db_type,
|
|
||||||
Some(DatabaseType::Sqlite) | Some(DatabaseType::DuckDb)
|
|
||||||
);
|
|
||||||
if is_embedded {
|
if is_embedded {
|
||||||
return Ok(connection_id.to_string());
|
return Ok(connection_id.to_string());
|
||||||
}
|
}
|
||||||
@@ -84,10 +77,7 @@ impl AppState {
|
|||||||
drop(conns);
|
drop(conns);
|
||||||
|
|
||||||
let configs = self.configs.lock().await;
|
let configs = self.configs.lock().await;
|
||||||
let config = configs
|
let config = configs.get(connection_id).ok_or("Connection config not found")?.clone();
|
||||||
.get(connection_id)
|
|
||||||
.ok_or("Connection config not found")?
|
|
||||||
.clone();
|
|
||||||
drop(configs);
|
drop(configs);
|
||||||
|
|
||||||
let mut db_config = config.clone();
|
let mut db_config = config.clone();
|
||||||
@@ -108,19 +98,14 @@ impl AppState {
|
|||||||
DatabaseType::Doris | DatabaseType::StarRocks => {
|
DatabaseType::Doris | DatabaseType::StarRocks => {
|
||||||
PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true)
|
PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true)
|
||||||
}
|
}
|
||||||
DatabaseType::Postgres | DatabaseType::Redshift => {
|
DatabaseType::Postgres | DatabaseType::Redshift => PoolKind::Postgres(db::postgres::connect(&url).await?),
|
||||||
PoolKind::Postgres(db::postgres::connect(&url).await?)
|
DatabaseType::Sqlite => PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&db_config.host)).await?),
|
||||||
}
|
|
||||||
DatabaseType::Sqlite => {
|
|
||||||
PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&db_config.host)).await?)
|
|
||||||
}
|
|
||||||
DatabaseType::Redis => {
|
DatabaseType::Redis => {
|
||||||
let con = db::redis_driver::connect(&url).await?;
|
let con = db::redis_driver::connect(&url).await?;
|
||||||
PoolKind::Redis(tokio::sync::Mutex::new(con))
|
PoolKind::Redis(tokio::sync::Mutex::new(con))
|
||||||
}
|
}
|
||||||
DatabaseType::DuckDb => {
|
DatabaseType::DuckDb => {
|
||||||
let con = duckdb::Connection::open(&expand_tilde(&db_config.host))
|
let con = duckdb::Connection::open(&expand_tilde(&db_config.host)).map_err(|e| e.to_string())?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
PoolKind::DuckDb(Arc::new(std::sync::Mutex::new(con)))
|
PoolKind::DuckDb(Arc::new(std::sync::Mutex::new(con)))
|
||||||
}
|
}
|
||||||
DatabaseType::MongoDb => {
|
DatabaseType::MongoDb => {
|
||||||
@@ -129,16 +114,8 @@ impl AppState {
|
|||||||
PoolKind::MongoDb(client)
|
PoolKind::MongoDb(client)
|
||||||
}
|
}
|
||||||
DatabaseType::ClickHouse => {
|
DatabaseType::ClickHouse => {
|
||||||
let username = if db_config.username.is_empty() {
|
let username = if db_config.username.is_empty() { None } else { Some(db_config.username.clone()) };
|
||||||
None
|
let password = if db_config.password.is_empty() { None } else { Some(db_config.password.clone()) };
|
||||||
} else {
|
|
||||||
Some(db_config.username.clone())
|
|
||||||
};
|
|
||||||
let password = if db_config.password.is_empty() {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
Some(db_config.password.clone())
|
|
||||||
};
|
|
||||||
let client = db::clickhouse_driver::ChClient::new(&url, username, password);
|
let client = db::clickhouse_driver::ChClient::new(&url, username, password);
|
||||||
db::clickhouse_driver::test_connection(&client).await?;
|
db::clickhouse_driver::test_connection(&client).await?;
|
||||||
PoolKind::ClickHouse(client)
|
PoolKind::ClickHouse(client)
|
||||||
@@ -166,11 +143,8 @@ impl AppState {
|
|||||||
PoolKind::Oracle(Arc::new(tokio::sync::Mutex::new(client)))
|
PoolKind::Oracle(Arc::new(tokio::sync::Mutex::new(client)))
|
||||||
}
|
}
|
||||||
DatabaseType::Elasticsearch => {
|
DatabaseType::Elasticsearch => {
|
||||||
let client = db::elasticsearch_driver::EsClient::new(
|
let client =
|
||||||
&url,
|
db::elasticsearch_driver::EsClient::new(&url, Some(&db_config.username), Some(&db_config.password));
|
||||||
Some(&db_config.username),
|
|
||||||
Some(&db_config.password),
|
|
||||||
);
|
|
||||||
db::elasticsearch_driver::test_connection(&client).await?;
|
db::elasticsearch_driver::test_connection(&client).await?;
|
||||||
PoolKind::Elasticsearch(client)
|
PoolKind::Elasticsearch(client)
|
||||||
}
|
}
|
||||||
@@ -212,18 +186,12 @@ impl AppState {
|
|||||||
Ok(("127.0.0.1".to_string(), local_port))
|
Ok(("127.0.0.1".to_string(), local_port))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn reconnect_pool(
|
pub async fn reconnect_pool(&self, connection_id: &str, database: Option<&str>) -> Result<String, String> {
|
||||||
&self,
|
|
||||||
connection_id: &str,
|
|
||||||
database: Option<&str>,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let is_single_conn = {
|
let is_single_conn = {
|
||||||
let configs = self.configs.lock().await;
|
let configs = self.configs.lock().await;
|
||||||
configs
|
configs
|
||||||
.get(connection_id)
|
.get(connection_id)
|
||||||
.map(|c| {
|
.map(|c| c.db_type == DatabaseType::Oracle || c.db_type == DatabaseType::Elasticsearch)
|
||||||
c.db_type == DatabaseType::Oracle || c.db_type == DatabaseType::Elasticsearch
|
|
||||||
})
|
|
||||||
.unwrap_or(false)
|
.unwrap_or(false)
|
||||||
};
|
};
|
||||||
let pool_key = if is_single_conn {
|
let pool_key = if is_single_conn {
|
||||||
@@ -247,11 +215,7 @@ pub fn connection_url_for_endpoint(config: &ConnectionConfig, host: &str, port:
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn redacted_connection_url_for_endpoint(
|
pub fn redacted_connection_url_for_endpoint(config: &ConnectionConfig, host: &str, port: u16) -> String {
|
||||||
config: &ConnectionConfig,
|
|
||||||
host: &str,
|
|
||||||
port: u16,
|
|
||||||
) -> String {
|
|
||||||
if host == config.host && port == config.port {
|
if host == config.host && port == config.port {
|
||||||
config.redacted_connection_url()
|
config.redacted_connection_url()
|
||||||
} else {
|
} else {
|
||||||
@@ -259,21 +223,10 @@ pub fn redacted_connection_url_for_endpoint(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn probe_connection_endpoint(
|
pub async fn probe_connection_endpoint(config: &ConnectionConfig, host: &str, port: u16) -> Result<(), String> {
|
||||||
config: &ConnectionConfig,
|
|
||||||
host: &str,
|
|
||||||
port: u16,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
match config.db_type {
|
match config.db_type {
|
||||||
DatabaseType::Sqlite | DatabaseType::DuckDb => Ok(()),
|
DatabaseType::Sqlite | DatabaseType::DuckDb => Ok(()),
|
||||||
DatabaseType::MongoDb
|
DatabaseType::MongoDb if config.connection_string.as_deref().is_some_and(|value| !value.is_empty()) => Ok(()),
|
||||||
if config
|
|
||||||
.connection_string
|
|
||||||
.as_deref()
|
|
||||||
.is_some_and(|value| !value.is_empty()) =>
|
|
||||||
{
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
_ => db::probe_tcp_endpoint(&format!("{:?}", config.db_type), host, port).await,
|
_ => db::probe_tcp_endpoint(&format!("{:?}", config.db_type), host, port).await,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,10 +24,7 @@ impl FileSecretStore {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn read_store(&self) -> HashMap<String, String> {
|
fn read_store(&self) -> HashMap<String, String> {
|
||||||
std::fs::read_to_string(&self.path)
|
std::fs::read_to_string(&self.path).ok().and_then(|json| serde_json::from_str(&json).ok()).unwrap_or_default()
|
||||||
.ok()
|
|
||||||
.and_then(|json| serde_json::from_str(&json).ok())
|
|
||||||
.unwrap_or_default()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn write_store(&self, map: &HashMap<String, String>) -> Result<(), String> {
|
fn write_store(&self, map: &HashMap<String, String>) -> Result<(), String> {
|
||||||
@@ -64,12 +61,7 @@ pub fn save_connections_to_file(
|
|||||||
persist_secret(store, &config.id, MAIN_PASSWORD_KEY, &config.password)?;
|
persist_secret(store, &config.id, MAIN_PASSWORD_KEY, &config.password)?;
|
||||||
persist_secret(store, &config.id, SSH_PASSWORD_KEY, &config.ssh_password)?;
|
persist_secret(store, &config.id, SSH_PASSWORD_KEY, &config.ssh_password)?;
|
||||||
persist_secret(store, &config.id, SSH_KEY_PASSPHRASE_KEY, &config.ssh_key_passphrase)?;
|
persist_secret(store, &config.id, SSH_KEY_PASSPHRASE_KEY, &config.ssh_key_passphrase)?;
|
||||||
persist_optional_secret(
|
persist_optional_secret(store, &config.id, CONNECTION_STRING_KEY, config.connection_string.as_deref())?;
|
||||||
store,
|
|
||||||
&config.id,
|
|
||||||
CONNECTION_STRING_KEY,
|
|
||||||
config.connection_string.as_deref(),
|
|
||||||
)?;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
write_sanitized_connections(path, configs)
|
write_sanitized_connections(path, configs)
|
||||||
@@ -113,11 +105,7 @@ pub fn load_connections_from_file(
|
|||||||
needs_rewrite = true;
|
needs_rewrite = true;
|
||||||
}
|
}
|
||||||
|
|
||||||
match config
|
match config.connection_string.as_deref().filter(|secret| !secret.is_empty()) {
|
||||||
.connection_string
|
|
||||||
.as_deref()
|
|
||||||
.filter(|secret| !secret.is_empty())
|
|
||||||
{
|
|
||||||
Some(secret) => {
|
Some(secret) => {
|
||||||
store.set_secret(&config.id, CONNECTION_STRING_KEY, secret)?;
|
store.set_secret(&config.id, CONNECTION_STRING_KEY, secret)?;
|
||||||
needs_rewrite = true;
|
needs_rewrite = true;
|
||||||
@@ -220,8 +208,8 @@ pub fn secret_account(connection_id: &str, key: &str) -> String {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::{
|
use super::{
|
||||||
load_connections_from_file, save_connections_to_file, ConnectionSecretStore,
|
load_connections_from_file, save_connections_to_file, ConnectionSecretStore, CONNECTION_STRING_KEY,
|
||||||
CONNECTION_STRING_KEY, MAIN_PASSWORD_KEY, SSH_PASSWORD_KEY,
|
MAIN_PASSWORD_KEY, SSH_PASSWORD_KEY,
|
||||||
};
|
};
|
||||||
use crate::models::connection::{ConnectionConfig, DatabaseType};
|
use crate::models::connection::{ConnectionConfig, DatabaseType};
|
||||||
use std::cell::RefCell;
|
use std::cell::RefCell;
|
||||||
@@ -236,48 +224,31 @@ mod tests {
|
|||||||
|
|
||||||
impl MemorySecretStore {
|
impl MemorySecretStore {
|
||||||
fn set_existing(&self, connection_id: &str, key: &str, value: &str) {
|
fn set_existing(&self, connection_id: &str, key: &str, value: &str) {
|
||||||
self.values
|
self.values.borrow_mut().insert(secret_key(connection_id, key), value.to_string());
|
||||||
.borrow_mut()
|
|
||||||
.insert(secret_key(connection_id, key), value.to_string());
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn get_existing(&self, connection_id: &str, key: &str) -> Option<String> {
|
fn get_existing(&self, connection_id: &str, key: &str) -> Option<String> {
|
||||||
self.values
|
self.values.borrow().get(&secret_key(connection_id, key)).cloned()
|
||||||
.borrow()
|
|
||||||
.get(&secret_key(connection_id, key))
|
|
||||||
.cloned()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn was_deleted(&self, connection_id: &str, key: &str) -> bool {
|
fn was_deleted(&self, connection_id: &str, key: &str) -> bool {
|
||||||
self.deleted
|
self.deleted.borrow().contains(&secret_key(connection_id, key))
|
||||||
.borrow()
|
|
||||||
.contains(&secret_key(connection_id, key))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ConnectionSecretStore for MemorySecretStore {
|
impl ConnectionSecretStore for MemorySecretStore {
|
||||||
fn set_secret(&self, connection_id: &str, key: &str, secret: &str) -> Result<(), String> {
|
fn set_secret(&self, connection_id: &str, key: &str, secret: &str) -> Result<(), String> {
|
||||||
self.values
|
self.values.borrow_mut().insert(secret_key(connection_id, key), secret.to_string());
|
||||||
.borrow_mut()
|
|
||||||
.insert(secret_key(connection_id, key), secret.to_string());
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
fn get_secret(&self, connection_id: &str, key: &str) -> Result<Option<String>, String> {
|
fn get_secret(&self, connection_id: &str, key: &str) -> Result<Option<String>, String> {
|
||||||
Ok(self
|
Ok(self.values.borrow().get(&secret_key(connection_id, key)).cloned())
|
||||||
.values
|
|
||||||
.borrow()
|
|
||||||
.get(&secret_key(connection_id, key))
|
|
||||||
.cloned())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn delete_secret(&self, connection_id: &str, key: &str) -> Result<(), String> {
|
fn delete_secret(&self, connection_id: &str, key: &str) -> Result<(), String> {
|
||||||
self.values
|
self.values.borrow_mut().remove(&secret_key(connection_id, key));
|
||||||
.borrow_mut()
|
self.deleted.borrow_mut().push(secret_key(connection_id, key));
|
||||||
.remove(&secret_key(connection_id, key));
|
|
||||||
self.deleted
|
|
||||||
.borrow_mut()
|
|
||||||
.push(secret_key(connection_id, key));
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -287,10 +258,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn temp_connections_file(name: &str) -> std::path::PathBuf {
|
fn temp_connections_file(name: &str) -> std::path::PathBuf {
|
||||||
let dir = std::env::temp_dir().join(format!(
|
let dir = std::env::temp_dir().join(format!("dbx-connection-secrets-test-{}-{name}", std::process::id()));
|
||||||
"dbx-connection-secrets-test-{}-{name}",
|
|
||||||
std::process::id()
|
|
||||||
));
|
|
||||||
std::fs::create_dir_all(&dir).unwrap();
|
std::fs::create_dir_all(&dir).unwrap();
|
||||||
dir.join("connections.json")
|
dir.join("connections.json")
|
||||||
}
|
}
|
||||||
@@ -335,14 +303,8 @@ mod tests {
|
|||||||
|
|
||||||
save_connections_to_file(&path, &configs, &store).unwrap();
|
save_connections_to_file(&path, &configs, &store).unwrap();
|
||||||
|
|
||||||
assert_eq!(
|
assert_eq!(store.get_existing("main", MAIN_PASSWORD_KEY).as_deref(), Some("db-secret"));
|
||||||
store.get_existing("main", MAIN_PASSWORD_KEY).as_deref(),
|
assert_eq!(store.get_existing("main", SSH_PASSWORD_KEY).as_deref(), Some("ssh-secret"));
|
||||||
Some("db-secret")
|
|
||||||
);
|
|
||||||
assert_eq!(
|
|
||||||
store.get_existing("main", SSH_PASSWORD_KEY).as_deref(),
|
|
||||||
Some("ssh-secret")
|
|
||||||
);
|
|
||||||
let persisted = read_configs(&path);
|
let persisted = read_configs(&path);
|
||||||
assert_eq!(persisted[0].password, "");
|
assert_eq!(persisted[0].password, "");
|
||||||
assert_eq!(persisted[0].ssh_password, "");
|
assert_eq!(persisted[0].ssh_password, "");
|
||||||
@@ -374,14 +336,8 @@ mod tests {
|
|||||||
|
|
||||||
assert_eq!(loaded[0].password, "plain-db");
|
assert_eq!(loaded[0].password, "plain-db");
|
||||||
assert_eq!(loaded[0].ssh_password, "plain-ssh");
|
assert_eq!(loaded[0].ssh_password, "plain-ssh");
|
||||||
assert_eq!(
|
assert_eq!(store.get_existing("legacy", MAIN_PASSWORD_KEY).as_deref(), Some("plain-db"));
|
||||||
store.get_existing("legacy", MAIN_PASSWORD_KEY).as_deref(),
|
assert_eq!(store.get_existing("legacy", SSH_PASSWORD_KEY).as_deref(), Some("plain-ssh"));
|
||||||
Some("plain-db")
|
|
||||||
);
|
|
||||||
assert_eq!(
|
|
||||||
store.get_existing("legacy", SSH_PASSWORD_KEY).as_deref(),
|
|
||||||
Some("plain-ssh")
|
|
||||||
);
|
|
||||||
let persisted = read_configs(&path);
|
let persisted = read_configs(&path);
|
||||||
assert_eq!(persisted[0].password, "");
|
assert_eq!(persisted[0].password, "");
|
||||||
assert_eq!(persisted[0].ssh_password, "");
|
assert_eq!(persisted[0].ssh_password, "");
|
||||||
@@ -401,10 +357,7 @@ mod tests {
|
|||||||
|
|
||||||
assert!(store.was_deleted("old", MAIN_PASSWORD_KEY));
|
assert!(store.was_deleted("old", MAIN_PASSWORD_KEY));
|
||||||
assert!(store.was_deleted("old", SSH_PASSWORD_KEY));
|
assert!(store.was_deleted("old", SSH_PASSWORD_KEY));
|
||||||
assert_eq!(
|
assert_eq!(store.get_existing("kept", MAIN_PASSWORD_KEY).as_deref(), Some("new-db"));
|
||||||
store.get_existing("kept", MAIN_PASSWORD_KEY).as_deref(),
|
|
||||||
Some("new-db")
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -418,18 +371,13 @@ mod tests {
|
|||||||
save_connections_to_file(&path, &[config], &store).unwrap();
|
save_connections_to_file(&path, &[config], &store).unwrap();
|
||||||
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
store
|
store.get_existing("mongo", CONNECTION_STRING_KEY).as_deref(),
|
||||||
.get_existing("mongo", CONNECTION_STRING_KEY)
|
|
||||||
.as_deref(),
|
|
||||||
Some("mongodb://user:secret@localhost/app")
|
Some("mongodb://user:secret@localhost/app")
|
||||||
);
|
);
|
||||||
let persisted = read_configs(&path);
|
let persisted = read_configs(&path);
|
||||||
assert_eq!(persisted[0].connection_string, None);
|
assert_eq!(persisted[0].connection_string, None);
|
||||||
|
|
||||||
let loaded = load_connections_from_file(&path, &store).unwrap();
|
let loaded = load_connections_from_file(&path, &store).unwrap();
|
||||||
assert_eq!(
|
assert_eq!(loaded[0].connection_string.as_deref(), Some("mongodb://user:secret@localhost/app"));
|
||||||
loaded[0].connection_string.as_deref(),
|
|
||||||
Some("mongodb://user:secret@localhost/app")
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,16 +14,9 @@ pub struct ChClient {
|
|||||||
|
|
||||||
impl ChClient {
|
impl ChClient {
|
||||||
pub fn new(url: &str, username: Option<String>, password: Option<String>) -> Self {
|
pub fn new(url: &str, username: Option<String>, password: Option<String>) -> Self {
|
||||||
let http = HttpClient::builder()
|
let http =
|
||||||
.connect_timeout(connection_timeout())
|
HttpClient::builder().connect_timeout(connection_timeout()).build().unwrap_or_else(|_| HttpClient::new());
|
||||||
.build()
|
Self { http, base_url: url.trim_end_matches('/').to_string(), username, password }
|
||||||
.unwrap_or_else(|_| HttpClient::new());
|
|
||||||
Self {
|
|
||||||
http,
|
|
||||||
base_url: url.trim_end_matches('/').to_string(),
|
|
||||||
username,
|
|
||||||
password,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -62,55 +55,33 @@ fn build_request(client: &ChClient, req: reqwest::RequestBuilder) -> reqwest::Re
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn ch_query(
|
async fn ch_query(client: &ChClient, sql: &str, database: Option<&str>) -> Result<ChJsonResult, String> {
|
||||||
client: &ChClient,
|
|
||||||
sql: &str,
|
|
||||||
database: Option<&str>,
|
|
||||||
) -> Result<ChJsonResult, String> {
|
|
||||||
let mut url = format!("{}/?default_format=JSONCompact", client.base_url);
|
let mut url = format!("{}/?default_format=JSONCompact", client.base_url);
|
||||||
if let Some(db) = database {
|
if let Some(db) = database {
|
||||||
url.push_str(&format!("&database={}", db));
|
url.push_str(&format!("&database={}", db));
|
||||||
}
|
}
|
||||||
let req = build_request(client, client.http.post(&url).body(sql.to_string()));
|
let req = build_request(client, client.http.post(&url).body(sql.to_string()));
|
||||||
let resp = req
|
let resp = req.send().await.map_err(|e| format!("ClickHouse request failed: {e}"))?;
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("ClickHouse request failed: {e}"))?;
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
return Err(format!("ClickHouse error: {body}"));
|
return Err(format!("ClickHouse error: {body}"));
|
||||||
}
|
}
|
||||||
resp.json::<ChJsonResult>()
|
resp.json::<ChJsonResult>().await.map_err(|e| format!("ClickHouse parse error: {e}"))
|
||||||
.await
|
|
||||||
.map_err(|e| format!("ClickHouse parse error: {e}"))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn test_connection(client: &ChClient) -> Result<(), String> {
|
pub async fn test_connection(client: &ChClient) -> Result<(), String> {
|
||||||
let url = format!("{}/ping", client.base_url);
|
let url = format!("{}/ping", client.base_url);
|
||||||
let req = build_request(client, client.http.get(&url));
|
let req = build_request(client, client.http.get(&url));
|
||||||
with_connection_timeout("ClickHouse", async {
|
with_connection_timeout("ClickHouse", async {
|
||||||
req.send()
|
req.send().await.map_err(|e| format!("ClickHouse connection failed: {e}"))
|
||||||
.await
|
|
||||||
.map_err(|e| format!("ClickHouse connection failed: {e}"))
|
|
||||||
})
|
})
|
||||||
.await?;
|
.await?;
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(client: &ChClient) -> Result<Vec<DatabaseInfo>, String> {
|
pub async fn list_databases(client: &ChClient) -> Result<Vec<DatabaseInfo>, String> {
|
||||||
let result = ch_query(
|
let result = ch_query(client, "SELECT name FROM system.databases ORDER BY name", None).await?;
|
||||||
client,
|
Ok(result.data.iter().map(|row| DatabaseInfo { name: row[0].as_str().unwrap_or("").to_string() }).collect())
|
||||||
"SELECT name FROM system.databases ORDER BY name",
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
.await?;
|
|
||||||
Ok(result
|
|
||||||
.data
|
|
||||||
.iter()
|
|
||||||
.map(|row| DatabaseInfo {
|
|
||||||
name: row[0].as_str().unwrap_or("").to_string(),
|
|
||||||
})
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_tables(client: &ChClient, database: &str) -> Result<Vec<TableInfo>, String> {
|
pub async fn list_tables(client: &ChClient, database: &str) -> Result<Vec<TableInfo>, String> {
|
||||||
@@ -124,24 +95,13 @@ pub async fn list_tables(client: &ChClient, database: &str) -> Result<Vec<TableI
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
let engine = row.get(1).and_then(|v| v.as_str()).unwrap_or("");
|
let engine = row.get(1).and_then(|v| v.as_str()).unwrap_or("");
|
||||||
let table_type = if engine.contains("View") {
|
let table_type = if engine.contains("View") { "VIEW" } else { "BASE TABLE" };
|
||||||
"VIEW"
|
TableInfo { name: row[0].as_str().unwrap_or("").to_string(), table_type: table_type.to_string() }
|
||||||
} else {
|
|
||||||
"BASE TABLE"
|
|
||||||
};
|
|
||||||
TableInfo {
|
|
||||||
name: row[0].as_str().unwrap_or("").to_string(),
|
|
||||||
table_type: table_type.to_string(),
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_columns(
|
pub async fn get_columns(client: &ChClient, database: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
|
||||||
client: &ChClient,
|
|
||||||
database: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ColumnInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT name, type, default_kind, default_expression, is_in_primary_key \
|
"SELECT name, type, default_kind, default_expression, is_in_primary_key \
|
||||||
FROM system.columns WHERE database = '{}' AND table = '{}' ORDER BY position",
|
FROM system.columns WHERE database = '{}' AND table = '{}' ORDER BY position",
|
||||||
@@ -153,20 +113,12 @@ pub async fn get_columns(
|
|||||||
.data
|
.data
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
let data_type = row
|
let data_type = row.get(1).and_then(|v| v.as_str()).unwrap_or("").to_string();
|
||||||
.get(1)
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.unwrap_or("")
|
|
||||||
.to_string();
|
|
||||||
let is_nullable = data_type.starts_with("Nullable");
|
let is_nullable = data_type.starts_with("Nullable");
|
||||||
let is_pk = row.get(4).and_then(|v| v.as_u64()).unwrap_or(0) == 1;
|
let is_pk = row.get(4).and_then(|v| v.as_u64()).unwrap_or(0) == 1;
|
||||||
let default_kind = row.get(2).and_then(|v| v.as_str()).unwrap_or("");
|
let default_kind = row.get(2).and_then(|v| v.as_str()).unwrap_or("");
|
||||||
let default_expr = row.get(3).and_then(|v| v.as_str()).unwrap_or("");
|
let default_expr = row.get(3).and_then(|v| v.as_str()).unwrap_or("");
|
||||||
let column_default = if default_kind.is_empty() {
|
let column_default = if default_kind.is_empty() { None } else { Some(default_expr.to_string()) };
|
||||||
None
|
|
||||||
} else {
|
|
||||||
Some(default_expr.to_string())
|
|
||||||
};
|
|
||||||
ColumnInfo {
|
ColumnInfo {
|
||||||
name: row[0].as_str().unwrap_or("").to_string(),
|
name: row[0].as_str().unwrap_or("").to_string(),
|
||||||
data_type,
|
data_type,
|
||||||
@@ -183,11 +135,7 @@ pub async fn get_columns(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn execute_query(
|
pub async fn execute_query(client: &ChClient, database: &str, sql: &str) -> Result<QueryResult, String> {
|
||||||
client: &ChClient,
|
|
||||||
database: &str,
|
|
||||||
sql: &str,
|
|
||||||
) -> Result<QueryResult, String> {
|
|
||||||
let start = Instant::now();
|
let start = Instant::now();
|
||||||
let trimmed = sql.trim().to_uppercase();
|
let trimmed = sql.trim().to_uppercase();
|
||||||
|
|
||||||
@@ -207,15 +155,9 @@ pub async fn execute_query(
|
|||||||
truncated: false,
|
truncated: false,
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
let url = format!(
|
let url = format!("{}/?default_format=JSONCompact&database={}", client.base_url, database);
|
||||||
"{}/?default_format=JSONCompact&database={}",
|
|
||||||
client.base_url, database
|
|
||||||
);
|
|
||||||
let req = build_request(client, client.http.post(&url).body(sql.to_string()));
|
let req = build_request(client, client.http.post(&url).body(sql.to_string()));
|
||||||
let resp = req
|
let resp = req.send().await.map_err(|e| format!("ClickHouse request failed: {e}"))?;
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("ClickHouse request failed: {e}"))?;
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
return Err(format!("ClickHouse error: {body}"));
|
return Err(format!("ClickHouse error: {body}"));
|
||||||
|
|||||||
@@ -16,15 +16,9 @@ impl EsClient {
|
|||||||
(Some(u), Some(p)) if !u.is_empty() => Some((u.to_string(), p.to_string())),
|
(Some(u), Some(p)) if !u.is_empty() => Some((u.to_string(), p.to_string())),
|
||||||
_ => None,
|
_ => None,
|
||||||
};
|
};
|
||||||
let http = HttpClient::builder()
|
let http =
|
||||||
.connect_timeout(connection_timeout())
|
HttpClient::builder().connect_timeout(connection_timeout()).build().unwrap_or_else(|_| HttpClient::new());
|
||||||
.build()
|
Self { http, base_url: url.trim_end_matches('/').to_string(), auth }
|
||||||
.unwrap_or_else(|_| HttpClient::new());
|
|
||||||
Self {
|
|
||||||
http,
|
|
||||||
base_url: url.trim_end_matches('/').to_string(),
|
|
||||||
auth,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn get(&self, path: &str) -> reqwest::RequestBuilder {
|
fn get(&self, path: &str) -> reqwest::RequestBuilder {
|
||||||
@@ -58,21 +52,13 @@ impl EsClient {
|
|||||||
|
|
||||||
impl Clone for EsClient {
|
impl Clone for EsClient {
|
||||||
fn clone(&self) -> Self {
|
fn clone(&self) -> Self {
|
||||||
Self {
|
Self { http: self.http.clone(), base_url: self.base_url.clone(), auth: self.auth.clone() }
|
||||||
http: self.http.clone(),
|
|
||||||
base_url: self.base_url.clone(),
|
|
||||||
auth: self.auth.clone(),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn test_connection(client: &EsClient) -> Result<(), String> {
|
pub async fn test_connection(client: &EsClient) -> Result<(), String> {
|
||||||
let resp = with_connection_timeout("Elasticsearch", async {
|
let resp = with_connection_timeout("Elasticsearch", async {
|
||||||
client
|
client.get("/").send().await.map_err(|e| format!("Elasticsearch connection failed: {e}"))
|
||||||
.get("/")
|
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch connection failed: {e}"))
|
|
||||||
})
|
})
|
||||||
.await?;
|
.await?;
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
@@ -97,15 +83,8 @@ pub async fn list_indices(client: &EsClient) -> Result<Vec<String>, String> {
|
|||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
return Err(format!("Elasticsearch error: {body}"));
|
return Err(format!("Elasticsearch error: {body}"));
|
||||||
}
|
}
|
||||||
let indices: Vec<CatIndex> = resp
|
let indices: Vec<CatIndex> = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
|
||||||
.json()
|
let mut names: Vec<String> = indices.into_iter().filter(|i| !i.index.starts_with('.')).map(|i| i.index).collect();
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
|
|
||||||
let mut names: Vec<String> = indices
|
|
||||||
.into_iter()
|
|
||||||
.filter(|i| !i.index.starts_with('.'))
|
|
||||||
.map(|i| i.index)
|
|
||||||
.collect();
|
|
||||||
names.sort();
|
names.sort();
|
||||||
Ok(names)
|
Ok(names)
|
||||||
}
|
}
|
||||||
@@ -147,22 +126,14 @@ pub async fn find_documents(
|
|||||||
});
|
});
|
||||||
|
|
||||||
let path = format!("/{}/_search", index);
|
let path = format!("/{}/_search", index);
|
||||||
let resp = client
|
let resp = client.post(&path).json(&body).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
||||||
.post(&path)
|
|
||||||
.json(&body)
|
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
|
||||||
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
return Err(format!("Elasticsearch error: {body}"));
|
return Err(format!("Elasticsearch error: {body}"));
|
||||||
}
|
}
|
||||||
|
|
||||||
let result: SearchResponse = resp
|
let result: SearchResponse = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
|
||||||
.json()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
|
|
||||||
|
|
||||||
let documents: Vec<serde_json::Value> = result
|
let documents: Vec<serde_json::Value> = result
|
||||||
.hits
|
.hits
|
||||||
@@ -178,56 +149,29 @@ pub async fn find_documents(
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
Ok(MongoDocumentResult {
|
Ok(MongoDocumentResult { documents, total: result.hits.total.value })
|
||||||
documents,
|
|
||||||
total: result.hits.total.value,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn insert_document(
|
pub async fn insert_document(client: &EsClient, index: &str, doc_json: &str) -> Result<String, String> {
|
||||||
client: &EsClient,
|
let doc: serde_json::Value = serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
|
||||||
index: &str,
|
|
||||||
doc_json: &str,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let doc: serde_json::Value =
|
|
||||||
serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
|
|
||||||
|
|
||||||
let path = format!("/{}/_doc?refresh=true", index);
|
let path = format!("/{}/_doc?refresh=true", index);
|
||||||
let resp = client
|
let resp = client.post(&path).json(&doc).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
||||||
.post(&path)
|
|
||||||
.json(&doc)
|
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
|
||||||
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
return Err(format!("Elasticsearch error: {body}"));
|
return Err(format!("Elasticsearch error: {body}"));
|
||||||
}
|
}
|
||||||
|
|
||||||
let result: serde_json::Value = resp
|
let result: serde_json::Value = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
|
||||||
.json()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch parse error: {e}"))?;
|
|
||||||
Ok(result["_id"].as_str().unwrap_or("").to_string())
|
Ok(result["_id"].as_str().unwrap_or("").to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn update_document(
|
pub async fn update_document(client: &EsClient, index: &str, id: &str, doc_json: &str) -> Result<u64, String> {
|
||||||
client: &EsClient,
|
let doc: serde_json::Value = serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
|
||||||
index: &str,
|
|
||||||
id: &str,
|
|
||||||
doc_json: &str,
|
|
||||||
) -> Result<u64, String> {
|
|
||||||
let doc: serde_json::Value =
|
|
||||||
serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
|
|
||||||
|
|
||||||
let path = format!("/{}/_doc/{}?refresh=true", index, id);
|
let path = format!("/{}/_doc/{}?refresh=true", index, id);
|
||||||
let resp = client
|
let resp = client.put(&path).json(&doc).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
||||||
.put(&path)
|
|
||||||
.json(&doc)
|
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
|
||||||
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
@@ -239,11 +183,7 @@ pub async fn update_document(
|
|||||||
|
|
||||||
pub async fn delete_document(client: &EsClient, index: &str, id: &str) -> Result<u64, String> {
|
pub async fn delete_document(client: &EsClient, index: &str, id: &str) -> Result<u64, String> {
|
||||||
let path = format!("/{}/_doc/{}?refresh=true", index, id);
|
let path = format!("/{}/_doc/{}?refresh=true", index, id);
|
||||||
let resp = client
|
let resp = client.delete(&path).send().await.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
||||||
.delete(&path)
|
|
||||||
.send()
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Elasticsearch request failed: {e}"))?;
|
|
||||||
|
|
||||||
if !resp.status().is_success() {
|
if !resp.status().is_success() {
|
||||||
let body = resp.text().await.unwrap_or_default();
|
let body = resp.text().await.unwrap_or_default();
|
||||||
|
|||||||
@@ -36,12 +36,9 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn probe_tcp_endpoint(label: &str, host: &str, port: u16) -> Result<(), String> {
|
pub async fn probe_tcp_endpoint(label: &str, host: &str, port: u16) -> Result<(), String> {
|
||||||
tokio::time::timeout(
|
tokio::time::timeout(tcp_probe_timeout(), tokio::net::TcpStream::connect((host, port)))
|
||||||
tcp_probe_timeout(),
|
.await
|
||||||
tokio::net::TcpStream::connect((host, port)),
|
.map_err(|_| format!("{label} TCP connection timed out ({TCP_PROBE_TIMEOUT_SECS}s)"))?
|
||||||
)
|
.map(|_| ())
|
||||||
.await
|
.map_err(|e| format!("{label} TCP connection failed: {e}"))
|
||||||
.map_err(|_| format!("{label} TCP connection timed out ({TCP_PROBE_TIMEOUT_SECS}s)"))?
|
|
||||||
.map(|_| ())
|
|
||||||
.map_err(|e| format!("{label} TCP connection failed: {e}"))
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,9 +14,7 @@ pub struct MongoDocumentResult {
|
|||||||
|
|
||||||
pub async fn connect(url: &str) -> Result<Client, String> {
|
pub async fn connect(url: &str) -> Result<Client, String> {
|
||||||
with_connection_timeout("MongoDB", async {
|
with_connection_timeout("MongoDB", async {
|
||||||
Client::with_uri_str(url)
|
Client::with_uri_str(url).await.map_err(|e| format!("MongoDB connection failed: {e}"))
|
||||||
.await
|
|
||||||
.map_err(|e| format!("MongoDB connection failed: {e}"))
|
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
@@ -30,18 +28,11 @@ pub async fn test_connection(client: &Client) -> Result<(), String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(client: &Client) -> Result<Vec<String>, String> {
|
pub async fn list_databases(client: &Client) -> Result<Vec<String>, String> {
|
||||||
client
|
client.list_database_names().await.map_err(|e| e.to_string())
|
||||||
.list_database_names()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_collections(client: &Client, database: &str) -> Result<Vec<String>, String> {
|
pub async fn list_collections(client: &Client, database: &str) -> Result<Vec<String>, String> {
|
||||||
client
|
client.database(database).list_collection_names().await.map_err(|e| e.to_string())
|
||||||
.database(database)
|
|
||||||
.list_collection_names()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn find_documents(
|
pub async fn find_documents(
|
||||||
@@ -53,17 +44,9 @@ pub async fn find_documents(
|
|||||||
) -> Result<MongoDocumentResult, String> {
|
) -> Result<MongoDocumentResult, String> {
|
||||||
let col = client.database(database).collection::<Document>(collection);
|
let col = client.database(database).collection::<Document>(collection);
|
||||||
|
|
||||||
let total = col
|
let total = col.count_documents(doc! {}).await.map_err(|e| e.to_string())?;
|
||||||
.count_documents(doc! {})
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let mut cursor = col
|
let mut cursor = col.find(doc! {}).skip(skip).limit(limit).await.map_err(|e| e.to_string())?;
|
||||||
.find(doc! {})
|
|
||||||
.skip(skip)
|
|
||||||
.limit(limit)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let mut documents = Vec::new();
|
let mut documents = Vec::new();
|
||||||
while cursor.advance().await.map_err(|e| e.to_string())? {
|
while cursor.advance().await.map_err(|e| e.to_string())? {
|
||||||
@@ -94,31 +77,17 @@ pub async fn update_document(
|
|||||||
id: &str,
|
id: &str,
|
||||||
doc_json: &str,
|
doc_json: &str,
|
||||||
) -> Result<u64, String> {
|
) -> Result<u64, String> {
|
||||||
let oid = mongodb::bson::oid::ObjectId::parse_str(id)
|
let oid = mongodb::bson::oid::ObjectId::parse_str(id).map_err(|e| format!("Invalid ObjectId: {e}"))?;
|
||||||
.map_err(|e| format!("Invalid ObjectId: {e}"))?;
|
let new_doc: Document = serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
|
||||||
let new_doc: Document =
|
|
||||||
serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?;
|
|
||||||
let col = client.database(database).collection::<Document>(collection);
|
let col = client.database(database).collection::<Document>(collection);
|
||||||
let result = col
|
let result = col.replace_one(doc! { "_id": oid }, new_doc).await.map_err(|e| e.to_string())?;
|
||||||
.replace_one(doc! { "_id": oid }, new_doc)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(result.modified_count)
|
Ok(result.modified_count)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn delete_document(
|
pub async fn delete_document(client: &Client, database: &str, collection: &str, id: &str) -> Result<u64, String> {
|
||||||
client: &Client,
|
let oid = mongodb::bson::oid::ObjectId::parse_str(id).map_err(|e| format!("Invalid ObjectId: {e}"))?;
|
||||||
database: &str,
|
|
||||||
collection: &str,
|
|
||||||
id: &str,
|
|
||||||
) -> Result<u64, String> {
|
|
||||||
let oid = mongodb::bson::oid::ObjectId::parse_str(id)
|
|
||||||
.map_err(|e| format!("Invalid ObjectId: {e}"))?;
|
|
||||||
let col = client.database(database).collection::<Document>(collection);
|
let col = client.database(database).collection::<Document>(collection);
|
||||||
let result = col
|
let result = col.delete_one(doc! { "_id": oid }).await.map_err(|e| e.to_string())?;
|
||||||
.delete_one(doc! { "_id": oid })
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(result.deleted_count)
|
Ok(result.deleted_count)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+42
-159
@@ -5,9 +5,7 @@ use sqlx::{Column, Executor, Row, TypeInfo, ValueRef};
|
|||||||
use std::time::{Duration, Instant};
|
use std::time::{Duration, Instant};
|
||||||
|
|
||||||
use super::{connection_timeout, with_connection_timeout};
|
use super::{connection_timeout, with_connection_timeout};
|
||||||
use crate::types::{
|
use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
|
||||||
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
|
|
||||||
};
|
|
||||||
|
|
||||||
fn quote_value(s: &str) -> String {
|
fn quote_value(s: &str) -> String {
|
||||||
format!("'{}'", s.replace('\\', "\\\\").replace('\'', "\\'"))
|
format!("'{}'", s.replace('\\', "\\\\").replace('\'', "\\'"))
|
||||||
@@ -15,32 +13,20 @@ fn quote_value(s: &str) -> String {
|
|||||||
|
|
||||||
fn get_str(row: &MySqlRow, idx: usize) -> String {
|
fn get_str(row: &MySqlRow, idx: usize) -> String {
|
||||||
row.try_get::<String, _>(idx)
|
row.try_get::<String, _>(idx)
|
||||||
.or_else(|_| {
|
.or_else(|_| row.try_get::<Vec<u8>, _>(idx).map(|b| String::from_utf8_lossy(&b).to_string()))
|
||||||
row.try_get::<Vec<u8>, _>(idx)
|
|
||||||
.map(|b| String::from_utf8_lossy(&b).to_string())
|
|
||||||
})
|
|
||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn get_str_by_name(row: &MySqlRow, name: &str) -> String {
|
fn get_str_by_name(row: &MySqlRow, name: &str) -> String {
|
||||||
row.try_get::<String, _>(name)
|
row.try_get::<String, _>(name)
|
||||||
.or_else(|_| {
|
.or_else(|_| row.try_get::<Vec<u8>, _>(name).map(|b| String::from_utf8_lossy(&b).to_string()))
|
||||||
row.try_get::<Vec<u8>, _>(name)
|
|
||||||
.map(|b| String::from_utf8_lossy(&b).to_string())
|
|
||||||
})
|
|
||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn get_opt_str(row: &MySqlRow, name: &str) -> Option<String> {
|
fn get_opt_str(row: &MySqlRow, name: &str) -> Option<String> {
|
||||||
row.try_get::<Option<String>, _>(name)
|
row.try_get::<Option<String>, _>(name).ok().flatten().or_else(|| {
|
||||||
.ok()
|
row.try_get::<Option<Vec<u8>>, _>(name).ok().flatten().map(|b| String::from_utf8_lossy(&b).to_string())
|
||||||
.flatten()
|
})
|
||||||
.or_else(|| {
|
|
||||||
row.try_get::<Option<Vec<u8>>, _>(name)
|
|
||||||
.ok()
|
|
||||||
.flatten()
|
|
||||||
.map(|b| String::from_utf8_lossy(&b).to_string())
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn numeric_metadata_u64_to_i32(value: Option<u64>) -> Option<i32> {
|
fn numeric_metadata_u64_to_i32(value: Option<u64>) -> Option<i32> {
|
||||||
@@ -52,9 +38,7 @@ fn numeric_metadata_i64_to_i32(value: Option<i64>) -> Option<i32> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn numeric_metadata_str_to_i32(value: Option<String>) -> Option<i32> {
|
fn numeric_metadata_str_to_i32(value: Option<String>) -> Option<i32> {
|
||||||
value
|
value.and_then(|v| v.parse::<i64>().ok()).and_then(|v| i32::try_from(v).ok())
|
||||||
.and_then(|v| v.parse::<i64>().ok())
|
|
||||||
.and_then(|v| i32::try_from(v).ok())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn get_opt_i32(row: &MySqlRow, name: &str) -> Option<i32> {
|
fn get_opt_i32(row: &MySqlRow, name: &str) -> Option<i32> {
|
||||||
@@ -67,9 +51,7 @@ fn get_opt_i32(row: &MySqlRow, name: &str) -> Option<i32> {
|
|||||||
.flatten()
|
.flatten()
|
||||||
.or_else(|| numeric_metadata_i64_to_i32(row.try_get::<Option<i64>, _>(name).ok().flatten()))
|
.or_else(|| numeric_metadata_i64_to_i32(row.try_get::<Option<i64>, _>(name).ok().flatten()))
|
||||||
.or_else(|| numeric_metadata_u64_to_i32(row.try_get::<Option<u64>, _>(name).ok().flatten()))
|
.or_else(|| numeric_metadata_u64_to_i32(row.try_get::<Option<u64>, _>(name).ok().flatten()))
|
||||||
.or_else(|| {
|
.or_else(|| numeric_metadata_str_to_i32(row.try_get::<Option<String>, _>(name).ok().flatten()))
|
||||||
numeric_metadata_str_to_i32(row.try_get::<Option<String>, _>(name).ok().flatten())
|
|
||||||
})
|
|
||||||
.or_else(|| {
|
.or_else(|| {
|
||||||
row.try_get::<Option<Vec<u8>>, _>(name)
|
row.try_get::<Option<Vec<u8>>, _>(name)
|
||||||
.ok()
|
.ok()
|
||||||
@@ -107,27 +89,20 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso
|
|||||||
return v;
|
return v;
|
||||||
}
|
}
|
||||||
if let Ok(v) = row.try_get::<String, _>(idx) {
|
if let Ok(v) = row.try_get::<String, _>(idx) {
|
||||||
return serde_json::from_str::<serde_json::Value>(&v)
|
return serde_json::from_str::<serde_json::Value>(&v).unwrap_or(serde_json::Value::String(v));
|
||||||
.unwrap_or(serde_json::Value::String(v));
|
|
||||||
}
|
}
|
||||||
return serde_json::Value::Null;
|
return serde_json::Value::Null;
|
||||||
}
|
}
|
||||||
|
|
||||||
if upper_type == "BOOLEAN" {
|
if upper_type == "BOOLEAN" {
|
||||||
return row
|
return row.try_get::<bool, _>(idx).map(serde_json::Value::Bool).unwrap_or(serde_json::Value::Null);
|
||||||
.try_get::<bool, _>(idx)
|
|
||||||
.map(serde_json::Value::Bool)
|
|
||||||
.unwrap_or(serde_json::Value::Null);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if upper_type.contains("BIGINT") {
|
if upper_type.contains("BIGINT") {
|
||||||
return row
|
return row
|
||||||
.try_get::<i64, _>(idx)
|
.try_get::<i64, _>(idx)
|
||||||
.map(|v| serde_json::Value::String(v.to_string()))
|
.map(|v| serde_json::Value::String(v.to_string()))
|
||||||
.or_else(|_| {
|
.or_else(|_| row.try_get::<u64, _>(idx).map(|v| serde_json::Value::String(v.to_string())))
|
||||||
row.try_get::<u64, _>(idx)
|
|
||||||
.map(|v| serde_json::Value::String(v.to_string()))
|
|
||||||
})
|
|
||||||
.unwrap_or(serde_json::Value::Null);
|
.unwrap_or(serde_json::Value::Null);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -151,25 +126,16 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso
|
|||||||
|
|
||||||
row.try_get::<String, _>(idx)
|
row.try_get::<String, _>(idx)
|
||||||
.map(serde_json::Value::String)
|
.map(serde_json::Value::String)
|
||||||
.or_else(|_| {
|
.or_else(|_| row.try_get::<i64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
|
||||||
row.try_get::<i64, _>(idx)
|
.or_else(|_| row.try_get::<u64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
|
||||||
.map(|v| serde_json::Value::Number(v.into()))
|
|
||||||
})
|
|
||||||
.or_else(|_| {
|
|
||||||
row.try_get::<u64, _>(idx)
|
|
||||||
.map(|v| serde_json::Value::Number(v.into()))
|
|
||||||
})
|
|
||||||
.or_else(|_| {
|
.or_else(|_| {
|
||||||
row.try_get::<f64, _>(idx).map(|v| {
|
row.try_get::<f64, _>(idx).map(|v| {
|
||||||
serde_json::Number::from_f64(v)
|
serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
|
||||||
.map(serde_json::Value::Number)
|
|
||||||
.unwrap_or(serde_json::Value::Null)
|
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
.or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool))
|
.or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool))
|
||||||
.or_else(|_| {
|
.or_else(|_| {
|
||||||
row.try_get::<Vec<u8>, _>(idx)
|
row.try_get::<Vec<u8>, _>(idx).map(|b| serde_json::Value::String(String::from_utf8_lossy(&b).to_string()))
|
||||||
.map(|b| serde_json::Value::String(String::from_utf8_lossy(&b).to_string()))
|
|
||||||
})
|
})
|
||||||
.or_else(|e| mysql_temporal_to_json_value(row, idx).ok_or(e))
|
.or_else(|e| mysql_temporal_to_json_value(row, idx).ok_or(e))
|
||||||
.unwrap_or(serde_json::Value::Null)
|
.unwrap_or(serde_json::Value::Null)
|
||||||
@@ -189,14 +155,9 @@ pub async fn connect(url: &str) -> Result<MySqlPool, String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn connect_bare(url: &str) -> Result<MySqlPool, String> {
|
pub async fn connect_bare(url: &str) -> Result<MySqlPool, String> {
|
||||||
let options: sqlx::mysql::MySqlConnectOptions = url
|
let options: sqlx::mysql::MySqlConnectOptions =
|
||||||
.parse()
|
url.parse().map_err(|e: sqlx::Error| format!("Invalid MySQL URL: {e}"))?;
|
||||||
.map_err(|e: sqlx::Error| format!("Invalid MySQL URL: {e}"))?;
|
let options = options.no_engine_substitution(false).set_names(false).pipes_as_concat(false).timezone(None);
|
||||||
let options = options
|
|
||||||
.no_engine_substitution(false)
|
|
||||||
.set_names(false)
|
|
||||||
.pipes_as_concat(false)
|
|
||||||
.timezone(None);
|
|
||||||
with_connection_timeout("MySQL", async {
|
with_connection_timeout("MySQL", async {
|
||||||
MySqlPoolOptions::new()
|
MySqlPoolOptions::new()
|
||||||
.max_connections(5)
|
.max_connections(5)
|
||||||
@@ -210,18 +171,12 @@ pub async fn connect_bare(url: &str) -> Result<MySqlPool, String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(pool: &MySqlPool) -> Result<Vec<DatabaseInfo>, String> {
|
pub async fn list_databases(pool: &MySqlPool) -> Result<Vec<DatabaseInfo>, String> {
|
||||||
let rows: Vec<MySqlRow> =
|
let rows: Vec<MySqlRow> = sqlx::raw_sql("SELECT SCHEMA_NAME FROM information_schema.SCHEMATA ORDER BY SCHEMA_NAME")
|
||||||
sqlx::raw_sql("SELECT SCHEMA_NAME FROM information_schema.SCHEMATA ORDER BY SCHEMA_NAME")
|
.fetch_all(pool)
|
||||||
.fetch_all(pool)
|
.await
|
||||||
.await
|
.map_err(|e| e.to_string())?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows.iter().map(|row| DatabaseInfo { name: get_str(row, 0) }).collect())
|
||||||
.iter()
|
|
||||||
.map(|row| DatabaseInfo {
|
|
||||||
name: get_str(row, 0),
|
|
||||||
})
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableInfo>, String> {
|
pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableInfo>, String> {
|
||||||
@@ -229,10 +184,7 @@ pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableIn
|
|||||||
"SELECT TABLE_NAME, TABLE_TYPE FROM information_schema.TABLES WHERE TABLE_SCHEMA = {} ORDER BY TABLE_NAME",
|
"SELECT TABLE_NAME, TABLE_TYPE FROM information_schema.TABLES WHERE TABLE_SCHEMA = {} ORDER BY TABLE_NAME",
|
||||||
quote_value(database),
|
quote_value(database),
|
||||||
);
|
);
|
||||||
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql)
|
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -243,11 +195,7 @@ pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result<Vec<TableIn
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_columns(
|
pub async fn get_columns(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
|
||||||
pool: &MySqlPool,
|
|
||||||
database: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ColumnInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT c.COLUMN_NAME, c.COLUMN_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, c.EXTRA, c.COLUMN_COMMENT, \
|
"SELECT c.COLUMN_NAME, c.COLUMN_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, c.EXTRA, c.COLUMN_COMMENT, \
|
||||||
CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \
|
CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \
|
||||||
@@ -263,10 +211,7 @@ pub async fn get_columns(
|
|||||||
quote_value(database),
|
quote_value(database),
|
||||||
quote_value(table),
|
quote_value(table),
|
||||||
);
|
);
|
||||||
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql)
|
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -295,22 +240,11 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
|
|||||||
|| trimmed.starts_with("EXPLAIN")
|
|| trimmed.starts_with("EXPLAIN")
|
||||||
{
|
{
|
||||||
if bare {
|
if bare {
|
||||||
let rows: Vec<MySqlRow> = sqlx::raw_sql(sql)
|
let rows: Vec<MySqlRow> = sqlx::raw_sql(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let (columns, column_types) = if let Some(first) = rows.first() {
|
let (columns, column_types) = if let Some(first) = rows.first() {
|
||||||
let cols: Vec<String> = first
|
let cols: Vec<String> = first.columns().iter().map(|c| c.name().to_string()).collect();
|
||||||
.columns()
|
let types: Vec<String> = first.columns().iter().map(|c| c.type_info().name().to_string()).collect();
|
||||||
.iter()
|
|
||||||
.map(|c| c.name().to_string())
|
|
||||||
.collect();
|
|
||||||
let types: Vec<String> = first
|
|
||||||
.columns()
|
|
||||||
.iter()
|
|
||||||
.map(|c| c.type_info().name().to_string())
|
|
||||||
.collect();
|
|
||||||
(cols, types)
|
(cols, types)
|
||||||
} else {
|
} else {
|
||||||
(vec![], vec![])
|
(vec![], vec![])
|
||||||
@@ -320,13 +254,7 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
(0..row.len())
|
(0..row.len())
|
||||||
.map(|i| {
|
.map(|i| mysql_value_to_json(row, i, column_types.get(i).map(String::as_str).unwrap_or("")))
|
||||||
mysql_value_to_json(
|
|
||||||
row,
|
|
||||||
i,
|
|
||||||
column_types.get(i).map(String::as_str).unwrap_or(""),
|
|
||||||
)
|
|
||||||
})
|
|
||||||
.collect()
|
.collect()
|
||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
@@ -340,33 +268,16 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
|
|||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
|
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
|
||||||
let columns: Vec<String> = desc
|
let columns: Vec<String> = desc.columns().iter().map(|c| c.name().to_string()).collect();
|
||||||
.columns()
|
let column_types: Vec<String> = desc.columns().iter().map(|c| c.type_info().name().to_string()).collect();
|
||||||
.iter()
|
|
||||||
.map(|c| c.name().to_string())
|
|
||||||
.collect();
|
|
||||||
let column_types: Vec<String> = desc
|
|
||||||
.columns()
|
|
||||||
.iter()
|
|
||||||
.map(|c| c.type_info().name().to_string())
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
let rows: Vec<MySqlRow> = sqlx::query(sql)
|
let rows: Vec<MySqlRow> = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let result_rows: Vec<Vec<serde_json::Value>> = rows
|
let result_rows: Vec<Vec<serde_json::Value>> = rows
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
(0..row.len())
|
(0..row.len())
|
||||||
.map(|i| {
|
.map(|i| mysql_value_to_json(row, i, column_types.get(i).map(String::as_str).unwrap_or("")))
|
||||||
mysql_value_to_json(
|
|
||||||
row,
|
|
||||||
i,
|
|
||||||
column_types.get(i).map(String::as_str).unwrap_or(""),
|
|
||||||
)
|
|
||||||
})
|
|
||||||
.collect()
|
.collect()
|
||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
@@ -380,10 +291,7 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
let result = sqlx::raw_sql(sql)
|
let result = sqlx::raw_sql(sql).execute(pool).await.map_err(|e| e.to_string())?;
|
||||||
.execute(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(QueryResult {
|
Ok(QueryResult {
|
||||||
columns: vec![],
|
columns: vec![],
|
||||||
@@ -395,11 +303,7 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_indexes(
|
pub async fn list_indexes(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
|
||||||
pool: &MySqlPool,
|
|
||||||
database: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<IndexInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT INDEX_NAME, GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) AS columns, \
|
"SELECT INDEX_NAME, GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) AS columns, \
|
||||||
MIN(NON_UNIQUE) = 0 AS is_unique, INDEX_NAME = 'PRIMARY' AS is_primary, \
|
MIN(NON_UNIQUE) = 0 AS is_unique, INDEX_NAME = 'PRIMARY' AS is_primary, \
|
||||||
@@ -411,10 +315,7 @@ pub async fn list_indexes(
|
|||||||
quote_value(database),
|
quote_value(database),
|
||||||
quote_value(table),
|
quote_value(table),
|
||||||
);
|
);
|
||||||
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql)
|
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -422,11 +323,7 @@ pub async fn list_indexes(
|
|||||||
let cols_str = get_str_by_name(row, "columns");
|
let cols_str = get_str_by_name(row, "columns");
|
||||||
IndexInfo {
|
IndexInfo {
|
||||||
name: get_str_by_name(row, "INDEX_NAME"),
|
name: get_str_by_name(row, "INDEX_NAME"),
|
||||||
columns: cols_str
|
columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(),
|
||||||
.split(',')
|
|
||||||
.filter(|s| !s.is_empty())
|
|
||||||
.map(|s| s.to_string())
|
|
||||||
.collect(),
|
|
||||||
is_unique: row.get::<bool, _>("is_unique"),
|
is_unique: row.get::<bool, _>("is_unique"),
|
||||||
is_primary: row.get::<bool, _>("is_primary"),
|
is_primary: row.get::<bool, _>("is_primary"),
|
||||||
filter: None,
|
filter: None,
|
||||||
@@ -438,11 +335,7 @@ pub async fn list_indexes(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_foreign_keys(
|
pub async fn list_foreign_keys(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
|
||||||
pool: &MySqlPool,
|
|
||||||
database: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ForeignKeyInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT kcu.CONSTRAINT_NAME, kcu.COLUMN_NAME, \
|
"SELECT kcu.CONSTRAINT_NAME, kcu.COLUMN_NAME, \
|
||||||
kcu.REFERENCED_TABLE_NAME, kcu.REFERENCED_COLUMN_NAME \
|
kcu.REFERENCED_TABLE_NAME, kcu.REFERENCED_COLUMN_NAME \
|
||||||
@@ -453,10 +346,7 @@ pub async fn list_foreign_keys(
|
|||||||
quote_value(database),
|
quote_value(database),
|
||||||
quote_value(table),
|
quote_value(table),
|
||||||
);
|
);
|
||||||
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql)
|
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -469,11 +359,7 @@ pub async fn list_foreign_keys(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_triggers(
|
pub async fn list_triggers(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
|
||||||
pool: &MySqlPool,
|
|
||||||
database: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<TriggerInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT TRIGGER_NAME, EVENT_MANIPULATION, ACTION_TIMING \
|
"SELECT TRIGGER_NAME, EVENT_MANIPULATION, ACTION_TIMING \
|
||||||
FROM information_schema.TRIGGERS \
|
FROM information_schema.TRIGGERS \
|
||||||
@@ -482,10 +368,7 @@ pub async fn list_triggers(
|
|||||||
quote_value(database),
|
quote_value(database),
|
||||||
quote_value(table),
|
quote_value(table),
|
||||||
);
|
);
|
||||||
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql)
|
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
|
|||||||
@@ -2,27 +2,16 @@ use oracle_rs::{Config, Connection};
|
|||||||
use std::time::Instant;
|
use std::time::Instant;
|
||||||
|
|
||||||
use super::{connection_timeout, CONNECTION_TIMEOUT_SECS};
|
use super::{connection_timeout, CONNECTION_TIMEOUT_SECS};
|
||||||
use crate::types::{
|
use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
|
||||||
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
|
|
||||||
};
|
|
||||||
|
|
||||||
pub type OracleClient = Connection;
|
pub type OracleClient = Connection;
|
||||||
|
|
||||||
pub async fn connect(
|
pub async fn connect(host: &str, port: u16, service: &str, user: &str, pass: &str) -> Result<OracleClient, String> {
|
||||||
host: &str,
|
|
||||||
port: u16,
|
|
||||||
service: &str,
|
|
||||||
user: &str,
|
|
||||||
pass: &str,
|
|
||||||
) -> Result<OracleClient, String> {
|
|
||||||
let config = Config::new(host, port, service, user, pass);
|
let config = Config::new(host, port, service, user, pass);
|
||||||
tokio::time::timeout(
|
tokio::time::timeout(connection_timeout(), Connection::connect_with_config(config))
|
||||||
connection_timeout(),
|
.await
|
||||||
Connection::connect_with_config(config),
|
.map_err(|_| format!("Oracle connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
||||||
)
|
.map_err(|e| format!("Oracle connection failed: {e}"))
|
||||||
.await
|
|
||||||
.map_err(|_| format!("Oracle connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
|
||||||
.map_err(|e| format!("Oracle connection failed: {e}"))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value {
|
fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value {
|
||||||
@@ -30,9 +19,9 @@ fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value {
|
|||||||
oracle_rs::Value::Null => serde_json::Value::Null,
|
oracle_rs::Value::Null => serde_json::Value::Null,
|
||||||
oracle_rs::Value::String(s) => serde_json::Value::String(s.clone()),
|
oracle_rs::Value::String(s) => serde_json::Value::String(s.clone()),
|
||||||
oracle_rs::Value::Integer(n) => serde_json::Value::Number((*n).into()),
|
oracle_rs::Value::Integer(n) => serde_json::Value::Number((*n).into()),
|
||||||
oracle_rs::Value::Float(f) => serde_json::Number::from_f64(*f)
|
oracle_rs::Value::Float(f) => {
|
||||||
.map(serde_json::Value::Number)
|
serde_json::Number::from_f64(*f).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
|
||||||
.unwrap_or(serde_json::Value::Null),
|
}
|
||||||
oracle_rs::Value::Boolean(b) => serde_json::Value::Bool(*b),
|
oracle_rs::Value::Boolean(b) => serde_json::Value::Bool(*b),
|
||||||
oracle_rs::Value::Json(v) => v.clone(),
|
oracle_rs::Value::Json(v) => v.clone(),
|
||||||
_ => serde_json::Value::String(format!("{val:?}")),
|
_ => serde_json::Value::String(format!("{val:?}")),
|
||||||
@@ -40,29 +29,15 @@ fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(conn: &OracleClient) -> Result<Vec<DatabaseInfo>, String> {
|
pub async fn list_databases(conn: &OracleClient) -> Result<Vec<DatabaseInfo>, String> {
|
||||||
let result = conn
|
let result =
|
||||||
.query("SELECT username FROM all_users ORDER BY username", &[])
|
conn.query("SELECT username FROM all_users ORDER BY username", &[]).await.map_err(|e| e.to_string())?;
|
||||||
.await
|
Ok(result.rows.iter().map(|row| DatabaseInfo { name: row.get_string(0).unwrap_or("").to_string() }).collect())
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(result
|
|
||||||
.rows
|
|
||||||
.iter()
|
|
||||||
.map(|row| DatabaseInfo {
|
|
||||||
name: row.get_string(0).unwrap_or("").to_string(),
|
|
||||||
})
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_schemas(conn: &OracleClient) -> Result<Vec<String>, String> {
|
pub async fn list_schemas(conn: &OracleClient) -> Result<Vec<String>, String> {
|
||||||
let result = conn
|
let result =
|
||||||
.query("SELECT username FROM all_users ORDER BY username", &[])
|
conn.query("SELECT username FROM all_users ORDER BY username", &[]).await.map_err(|e| e.to_string())?;
|
||||||
.await
|
Ok(result.rows.iter().map(|row| row.get_string(0).unwrap_or("").to_string()).collect())
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(result
|
|
||||||
.rows
|
|
||||||
.iter()
|
|
||||||
.map(|row| row.get_string(0).unwrap_or("").to_string())
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result<Vec<TableInfo>, String> {
|
pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result<Vec<TableInfo>, String> {
|
||||||
@@ -84,37 +59,36 @@ pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result<Vec<TableI
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_columns(
|
pub async fn get_columns(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
|
||||||
conn: &OracleClient,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ColumnInfo>, String> {
|
|
||||||
let s = schema.replace('\'', "''");
|
let s = schema.replace('\'', "''");
|
||||||
let t = table.replace('\'', "''");
|
let t = table.replace('\'', "''");
|
||||||
|
|
||||||
let pk_result = conn.query(
|
let pk_result = conn
|
||||||
&format!(
|
.query(
|
||||||
"SELECT cols.COLUMN_NAME FROM ALL_CONS_COLUMNS cols \
|
&format!(
|
||||||
|
"SELECT cols.COLUMN_NAME FROM ALL_CONS_COLUMNS cols \
|
||||||
JOIN ALL_CONSTRAINTS cons ON cols.CONSTRAINT_NAME = cons.CONSTRAINT_NAME AND cols.OWNER = cons.OWNER \
|
JOIN ALL_CONSTRAINTS cons ON cols.CONSTRAINT_NAME = cons.CONSTRAINT_NAME AND cols.OWNER = cons.OWNER \
|
||||||
WHERE cons.CONSTRAINT_TYPE = 'P' AND cons.OWNER = '{s}' AND cons.TABLE_NAME = '{t}'"
|
WHERE cons.CONSTRAINT_TYPE = 'P' AND cons.OWNER = '{s}' AND cons.TABLE_NAME = '{t}'"
|
||||||
),
|
),
|
||||||
&[],
|
&[],
|
||||||
).await.map_err(|e| e.to_string())?;
|
)
|
||||||
let pk_names: std::collections::HashSet<String> = pk_result
|
.await
|
||||||
.rows
|
.map_err(|e| e.to_string())?;
|
||||||
.iter()
|
let pk_names: std::collections::HashSet<String> =
|
||||||
.filter_map(|row| row.get_string(0).map(|s| s.to_string()))
|
pk_result.rows.iter().filter_map(|row| row.get_string(0).map(|s| s.to_string())).collect();
|
||||||
.collect();
|
|
||||||
|
|
||||||
let col_result = conn.query(
|
let col_result = conn
|
||||||
&format!(
|
.query(
|
||||||
"SELECT COLUMN_NAME, DATA_TYPE, NULLABLE, DATA_PRECISION, DATA_SCALE, DATA_LENGTH, CHAR_LENGTH \
|
&format!(
|
||||||
|
"SELECT COLUMN_NAME, DATA_TYPE, NULLABLE, DATA_PRECISION, DATA_SCALE, DATA_LENGTH, CHAR_LENGTH \
|
||||||
FROM ALL_TAB_COLUMNS \
|
FROM ALL_TAB_COLUMNS \
|
||||||
WHERE OWNER = '{s}' AND TABLE_NAME = '{t}' \
|
WHERE OWNER = '{s}' AND TABLE_NAME = '{t}' \
|
||||||
ORDER BY COLUMN_ID"
|
ORDER BY COLUMN_ID"
|
||||||
),
|
),
|
||||||
&[],
|
&[],
|
||||||
).await.map_err(|e| e.to_string())?;
|
)
|
||||||
|
.await
|
||||||
|
.map_err(|e| e.to_string())?;
|
||||||
|
|
||||||
Ok(col_result
|
Ok(col_result
|
||||||
.rows
|
.rows
|
||||||
@@ -161,11 +135,7 @@ pub async fn get_columns(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_indexes(
|
pub async fn list_indexes(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
|
||||||
conn: &OracleClient,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<IndexInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT i.INDEX_NAME, \
|
"SELECT i.INDEX_NAME, \
|
||||||
LISTAGG(ic.COLUMN_NAME, ',') WITHIN GROUP (ORDER BY ic.COLUMN_POSITION) AS columns, \
|
LISTAGG(ic.COLUMN_NAME, ',') WITHIN GROUP (ORDER BY ic.COLUMN_POSITION) AS columns, \
|
||||||
@@ -189,11 +159,7 @@ pub async fn list_indexes(
|
|||||||
let cols_str = row.get_string(1).unwrap_or("");
|
let cols_str = row.get_string(1).unwrap_or("");
|
||||||
IndexInfo {
|
IndexInfo {
|
||||||
name: row.get_string(0).unwrap_or("").to_string(),
|
name: row.get_string(0).unwrap_or("").to_string(),
|
||||||
columns: cols_str
|
columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(),
|
||||||
.split(',')
|
|
||||||
.filter(|s| !s.is_empty())
|
|
||||||
.map(|s| s.to_string())
|
|
||||||
.collect(),
|
|
||||||
is_unique: row.get_string(2).unwrap_or("") == "UNIQUE",
|
is_unique: row.get_string(2).unwrap_or("") == "UNIQUE",
|
||||||
is_primary: row.get_i64(3).unwrap_or(0) == 1,
|
is_primary: row.get_i64(3).unwrap_or(0) == 1,
|
||||||
filter: None,
|
filter: None,
|
||||||
@@ -205,11 +171,7 @@ pub async fn list_indexes(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_foreign_keys(
|
pub async fn list_foreign_keys(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
|
||||||
conn: &OracleClient,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ForeignKeyInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT c.CONSTRAINT_NAME, cc.COLUMN_NAME, rc.TABLE_NAME, rcc.COLUMN_NAME \
|
"SELECT c.CONSTRAINT_NAME, cc.COLUMN_NAME, rc.TABLE_NAME, rcc.COLUMN_NAME \
|
||||||
FROM ALL_CONSTRAINTS c \
|
FROM ALL_CONSTRAINTS c \
|
||||||
@@ -218,7 +180,8 @@ pub async fn list_foreign_keys(
|
|||||||
JOIN ALL_CONS_COLUMNS rcc ON rc.CONSTRAINT_NAME = rcc.CONSTRAINT_NAME AND rc.OWNER = rcc.OWNER \
|
JOIN ALL_CONS_COLUMNS rcc ON rc.CONSTRAINT_NAME = rcc.CONSTRAINT_NAME AND rc.OWNER = rcc.OWNER \
|
||||||
WHERE c.CONSTRAINT_TYPE = 'R' AND c.OWNER = '{s}' AND c.TABLE_NAME = '{t}' \
|
WHERE c.CONSTRAINT_TYPE = 'R' AND c.OWNER = '{s}' AND c.TABLE_NAME = '{t}' \
|
||||||
ORDER BY c.CONSTRAINT_NAME",
|
ORDER BY c.CONSTRAINT_NAME",
|
||||||
s = schema.replace('\'', "''"), t = table.replace('\'', "''")
|
s = schema.replace('\'', "''"),
|
||||||
|
t = table.replace('\'', "''")
|
||||||
);
|
);
|
||||||
let result = conn.query(&sql, &[]).await.map_err(|e| e.to_string())?;
|
let result = conn.query(&sql, &[]).await.map_err(|e| e.to_string())?;
|
||||||
Ok(result
|
Ok(result
|
||||||
@@ -233,11 +196,7 @@ pub async fn list_foreign_keys(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_triggers(
|
pub async fn list_triggers(conn: &OracleClient, schema: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
|
||||||
conn: &OracleClient,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<TriggerInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT TRIGGER_NAME, TRIGGERING_EVENT, TRIGGER_TYPE \
|
"SELECT TRIGGER_NAME, TRIGGERING_EVENT, TRIGGER_TYPE \
|
||||||
FROM ALL_TRIGGERS \
|
FROM ALL_TRIGGERS \
|
||||||
@@ -276,11 +235,7 @@ pub async fn execute_query(conn: &OracleClient, sql: &str) -> Result<QueryResult
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
(0..columns.len())
|
(0..columns.len())
|
||||||
.map(|i| {
|
.map(|i| row.get(i).map(|v| value_to_json(v)).unwrap_or(serde_json::Value::Null))
|
||||||
row.get(i)
|
|
||||||
.map(|v| value_to_json(v))
|
|
||||||
.unwrap_or(serde_json::Value::Null)
|
|
||||||
})
|
|
||||||
.collect()
|
.collect()
|
||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|||||||
@@ -5,9 +5,7 @@ use sqlx::{Column, Executor, Row, TypeInfo, ValueRef};
|
|||||||
use std::time::{Duration, Instant};
|
use std::time::{Duration, Instant};
|
||||||
|
|
||||||
use super::{connection_timeout, with_connection_timeout};
|
use super::{connection_timeout, with_connection_timeout};
|
||||||
use crate::types::{
|
use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
|
||||||
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
|
|
||||||
};
|
|
||||||
|
|
||||||
fn pg_temporal_to_json_value(row: &PgRow, idx: usize) -> Option<serde_json::Value> {
|
fn pg_temporal_to_json_value(row: &PgRow, idx: usize) -> Option<serde_json::Value> {
|
||||||
if let Ok(v) = row.try_get::<DateTime<Utc>, _>(idx) {
|
if let Ok(v) = row.try_get::<DateTime<Utc>, _>(idx) {
|
||||||
@@ -37,17 +35,13 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val
|
|||||||
return v;
|
return v;
|
||||||
}
|
}
|
||||||
if let Ok(v) = row.try_get::<String, _>(idx) {
|
if let Ok(v) = row.try_get::<String, _>(idx) {
|
||||||
return serde_json::from_str::<serde_json::Value>(&v)
|
return serde_json::from_str::<serde_json::Value>(&v).unwrap_or(serde_json::Value::String(v));
|
||||||
.unwrap_or(serde_json::Value::String(v));
|
|
||||||
}
|
}
|
||||||
return serde_json::Value::Null;
|
return serde_json::Value::Null;
|
||||||
}
|
}
|
||||||
|
|
||||||
if upper == "BOOL" {
|
if upper == "BOOL" {
|
||||||
return row
|
return row.try_get::<bool, _>(idx).map(serde_json::Value::Bool).unwrap_or(serde_json::Value::Null);
|
||||||
.try_get::<bool, _>(idx)
|
|
||||||
.map(serde_json::Value::Bool)
|
|
||||||
.unwrap_or(serde_json::Value::Null);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if upper.contains("TIMESTAMP")
|
if upper.contains("TIMESTAMP")
|
||||||
@@ -70,19 +64,11 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val
|
|||||||
|
|
||||||
row.try_get::<String, _>(idx)
|
row.try_get::<String, _>(idx)
|
||||||
.map(serde_json::Value::String)
|
.map(serde_json::Value::String)
|
||||||
.or_else(|_| {
|
.or_else(|_| row.try_get::<i64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
|
||||||
row.try_get::<i64, _>(idx)
|
.or_else(|_| row.try_get::<i32, _>(idx).map(|v| serde_json::Value::Number(v.into())))
|
||||||
.map(|v| serde_json::Value::Number(v.into()))
|
|
||||||
})
|
|
||||||
.or_else(|_| {
|
|
||||||
row.try_get::<i32, _>(idx)
|
|
||||||
.map(|v| serde_json::Value::Number(v.into()))
|
|
||||||
})
|
|
||||||
.or_else(|_| {
|
.or_else(|_| {
|
||||||
row.try_get::<f64, _>(idx).map(|v| {
|
row.try_get::<f64, _>(idx).map(|v| {
|
||||||
serde_json::Number::from_f64(v)
|
serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
|
||||||
.map(serde_json::Value::Number)
|
|
||||||
.unwrap_or(serde_json::Value::Null)
|
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
.or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool))
|
.or_else(|_| row.try_get::<bool, _>(idx).map(serde_json::Value::Bool))
|
||||||
@@ -104,18 +90,12 @@ pub async fn connect(url: &str) -> Result<PgPool, String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(pool: &PgPool) -> Result<Vec<DatabaseInfo>, String> {
|
pub async fn list_databases(pool: &PgPool) -> Result<Vec<DatabaseInfo>, String> {
|
||||||
let rows: Vec<PgRow> =
|
let rows: Vec<PgRow> = sqlx::query("SELECT datname FROM pg_database WHERE datistemplate = false ORDER BY datname")
|
||||||
sqlx::query("SELECT datname FROM pg_database WHERE datistemplate = false ORDER BY datname")
|
.fetch_all(pool)
|
||||||
.fetch_all(pool)
|
.await
|
||||||
.await
|
.map_err(|e| e.to_string())?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows.iter().map(|row| DatabaseInfo { name: row.get::<String, _>("datname") }).collect())
|
||||||
.iter()
|
|
||||||
.map(|row| DatabaseInfo {
|
|
||||||
name: row.get::<String, _>("datname"),
|
|
||||||
})
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_tables(pool: &PgPool, schema: &str) -> Result<Vec<TableInfo>, String> {
|
pub async fn list_tables(pool: &PgPool, schema: &str) -> Result<Vec<TableInfo>, String> {
|
||||||
@@ -149,17 +129,10 @@ pub async fn list_schemas(pool: &PgPool) -> Result<Vec<String>, String> {
|
|||||||
.await
|
.await
|
||||||
.map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows.iter().map(|row| row.get::<String, _>("schema_name")).collect())
|
||||||
.iter()
|
|
||||||
.map(|row| row.get::<String, _>("schema_name"))
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_columns(
|
pub async fn get_columns(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
|
||||||
pool: &PgPool,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ColumnInfo>, String> {
|
|
||||||
let rows: Vec<PgRow> = sqlx::query(
|
let rows: Vec<PgRow> = sqlx::query(
|
||||||
"SELECT a.attname AS column_name, \
|
"SELECT a.attname AS column_name, \
|
||||||
format_type(a.atttypid, a.atttypmod) AS full_type, \
|
format_type(a.atttypid, a.atttypmod) AS full_type, \
|
||||||
@@ -194,9 +167,7 @@ pub async fn get_columns(
|
|||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
let full_type = row
|
let full_type = row.get::<Option<String>, _>("full_type").unwrap_or_default();
|
||||||
.get::<Option<String>, _>("full_type")
|
|
||||||
.unwrap_or_default();
|
|
||||||
ColumnInfo {
|
ColumnInfo {
|
||||||
name: row.get::<String, _>("column_name"),
|
name: row.get::<String, _>("column_name"),
|
||||||
data_type: full_type,
|
data_type: full_type,
|
||||||
@@ -223,32 +194,19 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
|
|||||||
|| trimmed.starts_with("WITH")
|
|| trimmed.starts_with("WITH")
|
||||||
|| trimmed.starts_with("TABLE")
|
|| trimmed.starts_with("TABLE")
|
||||||
{
|
{
|
||||||
let rows: Vec<PgRow> = sqlx::query(sql)
|
let rows: Vec<PgRow> = sqlx::query(sql).persistent(false).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.persistent(false)
|
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let (columns, column_types): (Vec<String>, Vec<String>) = if let Some(first) = rows.first()
|
let (columns, column_types): (Vec<String>, Vec<String>) = if let Some(first) = rows.first() {
|
||||||
{
|
|
||||||
let cols = first.columns();
|
let cols = first.columns();
|
||||||
(
|
(
|
||||||
cols.iter().map(|c| c.name().to_string()).collect(),
|
cols.iter().map(|c| c.name().to_string()).collect(),
|
||||||
cols.iter()
|
cols.iter().map(|c| c.type_info().name().to_string()).collect(),
|
||||||
.map(|c| c.type_info().name().to_string())
|
|
||||||
.collect(),
|
|
||||||
)
|
)
|
||||||
} else {
|
} else {
|
||||||
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
|
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
|
||||||
(
|
(
|
||||||
desc.columns()
|
desc.columns().iter().map(|c| c.name().to_string()).collect(),
|
||||||
.iter()
|
desc.columns().iter().map(|c| c.type_info().name().to_string()).collect(),
|
||||||
.map(|c| c.name().to_string())
|
|
||||||
.collect(),
|
|
||||||
desc.columns()
|
|
||||||
.iter()
|
|
||||||
.map(|c| c.type_info().name().to_string())
|
|
||||||
.collect(),
|
|
||||||
)
|
)
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -256,13 +214,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
(0..row.len())
|
(0..row.len())
|
||||||
.map(|i| {
|
.map(|i| pg_value_to_json(row, i, column_types.get(i).map(String::as_str).unwrap_or("")))
|
||||||
pg_value_to_json(
|
|
||||||
row,
|
|
||||||
i,
|
|
||||||
column_types.get(i).map(String::as_str).unwrap_or(""),
|
|
||||||
)
|
|
||||||
})
|
|
||||||
.collect()
|
.collect()
|
||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
@@ -275,10 +227,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
|
|||||||
truncated: false,
|
truncated: false,
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
let result = sqlx::query(sql)
|
let result = sqlx::query(sql).execute(pool).await.map_err(|e| e.to_string())?;
|
||||||
.execute(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(QueryResult {
|
Ok(QueryResult {
|
||||||
columns: vec![],
|
columns: vec![],
|
||||||
@@ -290,11 +239,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_indexes(
|
pub async fn list_indexes(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
|
||||||
pool: &PgPool,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<IndexInfo>, String> {
|
|
||||||
let rows: Vec<PgRow> = sqlx::query(
|
let rows: Vec<PgRow> = sqlx::query(
|
||||||
"SELECT i.relname AS index_name, \
|
"SELECT i.relname AS index_name, \
|
||||||
array_agg(COALESCE(a.attname, pg_get_indexdef(ix.indexrelid, k.n::int, true)) ORDER BY k.n) AS columns, \
|
array_agg(COALESCE(a.attname, pg_get_indexdef(ix.indexrelid, k.n::int, true)) ORDER BY k.n) AS columns, \
|
||||||
@@ -326,15 +271,9 @@ pub async fn list_indexes(
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
let all_cols: Vec<String> = row.get::<Vec<String>, _>("columns");
|
let all_cols: Vec<String> = row.get::<Vec<String>, _>("columns");
|
||||||
let nkeyatts = row
|
let nkeyatts = row.get::<Option<i16>, _>("nkeyatts").unwrap_or(all_cols.len() as i16) as usize;
|
||||||
.get::<Option<i16>, _>("nkeyatts")
|
|
||||||
.unwrap_or(all_cols.len() as i16) as usize;
|
|
||||||
let key_cols = all_cols[..nkeyatts].to_vec();
|
let key_cols = all_cols[..nkeyatts].to_vec();
|
||||||
let included = if nkeyatts < all_cols.len() {
|
let included = if nkeyatts < all_cols.len() { all_cols[nkeyatts..].to_vec() } else { vec![] };
|
||||||
all_cols[nkeyatts..].to_vec()
|
|
||||||
} else {
|
|
||||||
vec![]
|
|
||||||
};
|
|
||||||
IndexInfo {
|
IndexInfo {
|
||||||
name: row.get::<String, _>("index_name"),
|
name: row.get::<String, _>("index_name"),
|
||||||
columns: key_cols,
|
columns: key_cols,
|
||||||
@@ -342,22 +281,14 @@ pub async fn list_indexes(
|
|||||||
is_primary: row.get::<bool, _>("is_primary"),
|
is_primary: row.get::<bool, _>("is_primary"),
|
||||||
filter: row.get::<Option<String>, _>("filter_expr"),
|
filter: row.get::<Option<String>, _>("filter_expr"),
|
||||||
index_type: row.get::<Option<String>, _>("index_type"),
|
index_type: row.get::<Option<String>, _>("index_type"),
|
||||||
included_columns: if included.is_empty() {
|
included_columns: if included.is_empty() { None } else { Some(included) },
|
||||||
None
|
|
||||||
} else {
|
|
||||||
Some(included)
|
|
||||||
},
|
|
||||||
comment: row.get::<Option<String>, _>("index_comment"),
|
comment: row.get::<Option<String>, _>("index_comment"),
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_foreign_keys(
|
pub async fn list_foreign_keys(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
|
||||||
pool: &PgPool,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ForeignKeyInfo>, String> {
|
|
||||||
let rows: Vec<PgRow> = sqlx::query(
|
let rows: Vec<PgRow> = sqlx::query(
|
||||||
"SELECT kcu.constraint_name, kcu.column_name, \
|
"SELECT kcu.constraint_name, kcu.column_name, \
|
||||||
ccu.table_name AS ref_table, ccu.column_name AS ref_column \
|
ccu.table_name AS ref_table, ccu.column_name AS ref_column \
|
||||||
@@ -388,11 +319,7 @@ pub async fn list_foreign_keys(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_triggers(
|
pub async fn list_triggers(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
|
||||||
pool: &PgPool,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<TriggerInfo>, String> {
|
|
||||||
let rows: Vec<PgRow> = sqlx::query(
|
let rows: Vec<PgRow> = sqlx::query(
|
||||||
"SELECT trigger_name, event_manipulation, action_timing \
|
"SELECT trigger_name, event_manipulation, action_timing \
|
||||||
FROM information_schema.triggers \
|
FROM information_schema.triggers \
|
||||||
|
|||||||
@@ -29,44 +29,26 @@ pub struct RedisValue {
|
|||||||
|
|
||||||
pub async fn connect(url: &str) -> Result<redis::aio::MultiplexedConnection, String> {
|
pub async fn connect(url: &str) -> Result<redis::aio::MultiplexedConnection, String> {
|
||||||
let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?;
|
let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?;
|
||||||
let mut con = tokio::time::timeout(
|
let mut con = tokio::time::timeout(connection_timeout(), client.get_multiplexed_async_connection())
|
||||||
connection_timeout(),
|
.await
|
||||||
client.get_multiplexed_async_connection(),
|
.map_err(|_| format!("Redis connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
||||||
)
|
.map_err(|e| format!("Redis connection failed: {e}"))?;
|
||||||
.await
|
|
||||||
.map_err(|_| format!("Redis connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
|
||||||
.map_err(|e| format!("Redis connection failed: {e}"))?;
|
|
||||||
|
|
||||||
tokio::time::timeout(
|
tokio::time::timeout(connection_timeout(), redis::cmd("PING").query_async::<String>(&mut con))
|
||||||
connection_timeout(),
|
.await
|
||||||
redis::cmd("PING").query_async::<String>(&mut con),
|
.map_err(|_| format!("Redis ping timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
||||||
)
|
.map_err(|e| format!("Redis authentication failed or command rejected: {e}"))?;
|
||||||
.await
|
|
||||||
.map_err(|_| format!("Redis ping timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
|
||||||
.map_err(|e| format!("Redis authentication failed or command rejected: {e}"))?;
|
|
||||||
|
|
||||||
Ok(con)
|
Ok(con)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(
|
pub async fn list_databases(con: &mut redis::aio::MultiplexedConnection) -> Result<Vec<u32>, String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
let configured_count =
|
||||||
) -> Result<Vec<u32>, String> {
|
redis::cmd("CONFIG").arg("GET").arg("databases").query_async(con).await.ok().and_then(parse_database_count);
|
||||||
let configured_count = redis::cmd("CONFIG")
|
|
||||||
.arg("GET")
|
|
||||||
.arg("databases")
|
|
||||||
.query_async(con)
|
|
||||||
.await
|
|
||||||
.ok()
|
|
||||||
.and_then(parse_database_count);
|
|
||||||
|
|
||||||
let keyspace_dbs = list_keyspace_databases(con).await.unwrap_or_default();
|
let keyspace_dbs = list_keyspace_databases(con).await.unwrap_or_default();
|
||||||
let database_count = configured_count.unwrap_or(DEFAULT_REDIS_DATABASES);
|
let database_count = configured_count.unwrap_or(DEFAULT_REDIS_DATABASES);
|
||||||
let max_db = keyspace_dbs
|
let max_db = keyspace_dbs.iter().copied().max().map(|db| db + 1).unwrap_or(0);
|
||||||
.iter()
|
|
||||||
.copied()
|
|
||||||
.max()
|
|
||||||
.map(|db| db + 1)
|
|
||||||
.unwrap_or(0);
|
|
||||||
let visible_count = database_count.max(max_db).max(1);
|
let visible_count = database_count.max(max_db).max(1);
|
||||||
|
|
||||||
Ok((0..visible_count).collect())
|
Ok((0..visible_count).collect())
|
||||||
@@ -88,14 +70,8 @@ fn parse_database_count(value: redis::Value) -> Option<u32> {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn list_keyspace_databases(
|
async fn list_keyspace_databases(con: &mut redis::aio::MultiplexedConnection) -> Result<Vec<u32>, String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
let info: String = redis::cmd("INFO").arg("keyspace").query_async(con).await.map_err(|e| e.to_string())?;
|
||||||
) -> Result<Vec<u32>, String> {
|
|
||||||
let info: String = redis::cmd("INFO")
|
|
||||||
.arg("keyspace")
|
|
||||||
.query_async(con)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let mut dbs = Vec::new();
|
let mut dbs = Vec::new();
|
||||||
for line in info.lines() {
|
for line in info.lines() {
|
||||||
@@ -111,11 +87,7 @@ async fn list_keyspace_databases(
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn select_db(con: &mut redis::aio::MultiplexedConnection, db: u32) -> Result<(), String> {
|
pub async fn select_db(con: &mut redis::aio::MultiplexedConnection, db: u32) -> Result<(), String> {
|
||||||
redis::cmd("SELECT")
|
redis::cmd("SELECT").arg(db).query_async(con).await.map_err(|e| e.to_string())
|
||||||
.arg(db)
|
|
||||||
.query_async(con)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn scan_keys_page(
|
pub async fn scan_keys_page(
|
||||||
@@ -136,35 +108,18 @@ pub async fn scan_keys_page(
|
|||||||
|
|
||||||
let mut result = Vec::new();
|
let mut result = Vec::new();
|
||||||
for key in &keys {
|
for key in &keys {
|
||||||
let key_type: String = redis::cmd("TYPE")
|
let key_type: String =
|
||||||
.arg(key.as_str())
|
redis::cmd("TYPE").arg(key.as_str()).query_async(con).await.unwrap_or_else(|_| "unknown".to_string());
|
||||||
.query_async(con)
|
|
||||||
.await
|
|
||||||
.unwrap_or_else(|_| "unknown".to_string());
|
|
||||||
|
|
||||||
let ttl: i64 = con.ttl(key.as_str()).await.unwrap_or(-1);
|
let ttl: i64 = con.ttl(key.as_str()).await.unwrap_or(-1);
|
||||||
|
|
||||||
result.push(RedisKeyInfo {
|
result.push(RedisKeyInfo { key: key.clone(), key_type, ttl });
|
||||||
key: key.clone(),
|
|
||||||
key_type,
|
|
||||||
ttl,
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
Ok(RedisScanResult {
|
Ok(RedisScanResult { cursor: next_cursor, keys: result })
|
||||||
cursor: next_cursor,
|
|
||||||
keys: result,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_value(
|
pub async fn get_value(con: &mut redis::aio::MultiplexedConnection, key: &str) -> Result<RedisValue, String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
let key_type: String = redis::cmd("TYPE").arg(key).query_async(con).await.map_err(|e| e.to_string())?;
|
||||||
key: &str,
|
|
||||||
) -> Result<RedisValue, String> {
|
|
||||||
let key_type: String = redis::cmd("TYPE")
|
|
||||||
.arg(key)
|
|
||||||
.query_async(con)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let ttl: i64 = con.ttl(key).await.unwrap_or(-1);
|
let ttl: i64 = con.ttl(key).await.unwrap_or(-1);
|
||||||
|
|
||||||
@@ -182,33 +137,20 @@ pub async fn get_value(
|
|||||||
serde_json::json!(v)
|
serde_json::json!(v)
|
||||||
}
|
}
|
||||||
"zset" => {
|
"zset" => {
|
||||||
let v: Vec<(String, f64)> = con
|
let v: Vec<(String, f64)> = con.zrange_withscores(key, 0, -1).await.map_err(|e| e.to_string())?;
|
||||||
.zrange_withscores(key, 0, -1)
|
serde_json::json!(v.iter().map(|(m, s)| serde_json::json!({"member": m, "score": s})).collect::<Vec<_>>())
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
serde_json::json!(v
|
|
||||||
.iter()
|
|
||||||
.map(|(m, s)| serde_json::json!({"member": m, "score": s}))
|
|
||||||
.collect::<Vec<_>>())
|
|
||||||
}
|
}
|
||||||
"hash" => {
|
"hash" => {
|
||||||
let v: Vec<(String, String)> = con.hgetall(key).await.map_err(|e| e.to_string())?;
|
let v: Vec<(String, String)> = con.hgetall(key).await.map_err(|e| e.to_string())?;
|
||||||
let map: serde_json::Map<String, serde_json::Value> = v
|
let map: serde_json::Map<String, serde_json::Value> =
|
||||||
.into_iter()
|
v.into_iter().map(|(k, v)| (k, serde_json::Value::String(v))).collect();
|
||||||
.map(|(k, v)| (k, serde_json::Value::String(v)))
|
|
||||||
.collect();
|
|
||||||
serde_json::Value::Object(map)
|
serde_json::Value::Object(map)
|
||||||
}
|
}
|
||||||
"stream" => get_stream_entries(con, key).await?,
|
"stream" => get_stream_entries(con, key).await?,
|
||||||
_ => serde_json::Value::Null,
|
_ => serde_json::Value::Null,
|
||||||
};
|
};
|
||||||
|
|
||||||
Ok(RedisValue {
|
Ok(RedisValue { key: key.to_string(), key_type, ttl, value })
|
||||||
key: key.to_string(),
|
|
||||||
key_type,
|
|
||||||
ttl,
|
|
||||||
value,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn get_stream_entries(
|
async fn get_stream_entries(
|
||||||
@@ -286,23 +228,16 @@ pub async fn set_string(
|
|||||||
value: &str,
|
value: &str,
|
||||||
ttl: Option<i64>,
|
ttl: Option<i64>,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
con.set::<_, _, ()>(key, value)
|
con.set::<_, _, ()>(key, value).await.map_err(|e| e.to_string())?;
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
if let Some(t) = ttl {
|
if let Some(t) = ttl {
|
||||||
if t > 0 {
|
if t > 0 {
|
||||||
con.expire::<_, ()>(key, t)
|
con.expire::<_, ()>(key, t).await.map_err(|e| e.to_string())?;
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn delete_key(
|
pub async fn delete_key(con: &mut redis::aio::MultiplexedConnection, key: &str) -> Result<(), String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
|
||||||
key: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
con.del::<_, ()>(key).await.map_err(|e| e.to_string())
|
con.del::<_, ()>(key).await.map_err(|e| e.to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -312,67 +247,29 @@ pub async fn hash_set(
|
|||||||
field: &str,
|
field: &str,
|
||||||
value: &str,
|
value: &str,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
con.hset::<_, _, _, ()>(key, field, value)
|
con.hset::<_, _, _, ()>(key, field, value).await.map_err(|e| e.to_string())
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn hash_del(
|
pub async fn hash_del(con: &mut redis::aio::MultiplexedConnection, key: &str, field: &str) -> Result<(), String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
con.hdel::<_, _, ()>(key, field).await.map_err(|e| e.to_string())
|
||||||
key: &str,
|
|
||||||
field: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
con.hdel::<_, _, ()>(key, field)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_push(
|
pub async fn list_push(con: &mut redis::aio::MultiplexedConnection, key: &str, value: &str) -> Result<(), String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
con.rpush::<_, _, ()>(key, value).await.map_err(|e| e.to_string())
|
||||||
key: &str,
|
|
||||||
value: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
con.rpush::<_, _, ()>(key, value)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_remove(
|
pub async fn list_remove(con: &mut redis::aio::MultiplexedConnection, key: &str, index: i64) -> Result<(), String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
|
||||||
key: &str,
|
|
||||||
index: i64,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let placeholder = "__DELETED_PLACEHOLDER__";
|
let placeholder = "__DELETED_PLACEHOLDER__";
|
||||||
redis::cmd("LSET")
|
redis::cmd("LSET").arg(key).arg(index).arg(placeholder).query_async::<()>(con).await.map_err(|e| e.to_string())?;
|
||||||
.arg(key)
|
con.lrem::<_, _, ()>(key, 1, placeholder).await.map_err(|e| e.to_string())
|
||||||
.arg(index)
|
|
||||||
.arg(placeholder)
|
|
||||||
.query_async::<()>(con)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
con.lrem::<_, _, ()>(key, 1, placeholder)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn set_add(
|
pub async fn set_add(con: &mut redis::aio::MultiplexedConnection, key: &str, member: &str) -> Result<(), String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
con.sadd::<_, _, ()>(key, member).await.map_err(|e| e.to_string())
|
||||||
key: &str,
|
|
||||||
member: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
con.sadd::<_, _, ()>(key, member)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn set_remove(
|
pub async fn set_remove(con: &mut redis::aio::MultiplexedConnection, key: &str, member: &str) -> Result<(), String> {
|
||||||
con: &mut redis::aio::MultiplexedConnection,
|
con.srem::<_, _, ()>(key, member).await.map_err(|e| e.to_string())
|
||||||
key: &str,
|
|
||||||
member: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
con.srem::<_, _, ()>(key, member)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -387,12 +284,7 @@ mod tests {
|
|||||||
fn parses_stream_entries() {
|
fn parses_stream_entries() {
|
||||||
let raw = RedisRawValue::Array(vec![RedisRawValue::Array(vec![
|
let raw = RedisRawValue::Array(vec![RedisRawValue::Array(vec![
|
||||||
bulk("1714470000000-0"),
|
bulk("1714470000000-0"),
|
||||||
RedisRawValue::Array(vec![
|
RedisRawValue::Array(vec![bulk("event"), bulk("login"), bulk("user_id"), bulk("42")]),
|
||||||
bulk("event"),
|
|
||||||
bulk("login"),
|
|
||||||
bulk("user_id"),
|
|
||||||
bulk("42"),
|
|
||||||
]),
|
|
||||||
])]);
|
])]);
|
||||||
|
|
||||||
let parsed = parse_stream_entries(raw);
|
let parsed = parse_stream_entries(raw);
|
||||||
|
|||||||
@@ -2,14 +2,10 @@ use sqlx::sqlite::{SqliteConnectOptions, SqlitePool, SqlitePoolOptions, SqliteRo
|
|||||||
use sqlx::{Column, Executor, Row};
|
use sqlx::{Column, Executor, Row};
|
||||||
use std::time::{Duration, Instant};
|
use std::time::{Duration, Instant};
|
||||||
|
|
||||||
use crate::types::{
|
use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
|
||||||
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
|
|
||||||
};
|
|
||||||
|
|
||||||
pub async fn connect_path(path: &str) -> Result<SqlitePool, String> {
|
pub async fn connect_path(path: &str) -> Result<SqlitePool, String> {
|
||||||
let mut options = SqliteConnectOptions::new()
|
let mut options = SqliteConnectOptions::new().filename(path).create_if_missing(true);
|
||||||
.filename(path)
|
|
||||||
.create_if_missing(true);
|
|
||||||
|
|
||||||
if is_network_path(path) {
|
if is_network_path(path) {
|
||||||
options = options.vfs("unix-nolock");
|
options = options.vfs("unix-nolock");
|
||||||
@@ -25,16 +21,11 @@ pub async fn connect_path(path: &str) -> Result<SqlitePool, String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn is_network_path(path: &str) -> bool {
|
fn is_network_path(path: &str) -> bool {
|
||||||
path.starts_with("\\\\")
|
path.starts_with("\\\\") || path.starts_with("//") || path.contains("wsl.localhost") || path.contains("wsl$")
|
||||||
|| path.starts_with("//")
|
|
||||||
|| path.contains("wsl.localhost")
|
|
||||||
|| path.contains("wsl$")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(_pool: &SqlitePool) -> Result<Vec<DatabaseInfo>, String> {
|
pub async fn list_databases(_pool: &SqlitePool) -> Result<Vec<DatabaseInfo>, String> {
|
||||||
Ok(vec![DatabaseInfo {
|
Ok(vec![DatabaseInfo { name: "main".to_string() }])
|
||||||
name: "main".to_string(),
|
|
||||||
}])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result<Vec<TableInfo>, String> {
|
pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result<Vec<TableInfo>, String> {
|
||||||
@@ -51,25 +42,15 @@ pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result<Vec<TableIn
|
|||||||
let t: String = row.get("type");
|
let t: String = row.get("type");
|
||||||
TableInfo {
|
TableInfo {
|
||||||
name: row.get::<String, _>("name"),
|
name: row.get::<String, _>("name"),
|
||||||
table_type: if t == "view" {
|
table_type: if t == "view" { "VIEW".to_string() } else { "BASE TABLE".to_string() },
|
||||||
"VIEW".to_string()
|
|
||||||
} else {
|
|
||||||
"BASE TABLE".to_string()
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_columns(
|
pub async fn get_columns(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
|
||||||
pool: &SqlitePool,
|
let rows: Vec<SqliteRow> =
|
||||||
_schema: &str,
|
sqlx::query(&format!("PRAGMA table_info(\"{}\")", table)).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ColumnInfo>, String> {
|
|
||||||
let rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA table_info(\"{}\")", table))
|
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -88,11 +69,7 @@ pub async fn get_columns(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_indexes(
|
pub async fn list_indexes(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
|
||||||
pool: &SqlitePool,
|
|
||||||
_schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<IndexInfo>, String> {
|
|
||||||
let safe_table = table.replace('"', "\"\"");
|
let safe_table = table.replace('"', "\"\"");
|
||||||
let idx_rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA index_list(\"{safe_table}\")"))
|
let idx_rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA index_list(\"{safe_table}\")"))
|
||||||
.fetch_all(pool)
|
.fetch_all(pool)
|
||||||
@@ -112,10 +89,7 @@ pub async fn list_indexes(
|
|||||||
.await
|
.await
|
||||||
.map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
|
|
||||||
let columns: Vec<String> = col_rows
|
let columns: Vec<String> = col_rows.iter().map(|r| r.get::<String, _>("name")).collect();
|
||||||
.iter()
|
|
||||||
.map(|r| r.get::<String, _>("name"))
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
indexes.push(IndexInfo {
|
indexes.push(IndexInfo {
|
||||||
name,
|
name,
|
||||||
@@ -131,11 +105,7 @@ pub async fn list_indexes(
|
|||||||
Ok(indexes)
|
Ok(indexes)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_foreign_keys(
|
pub async fn list_foreign_keys(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<ForeignKeyInfo>, String> {
|
||||||
pool: &SqlitePool,
|
|
||||||
_schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ForeignKeyInfo>, String> {
|
|
||||||
let rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA foreign_key_list(\"{}\")", table))
|
let rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA foreign_key_list(\"{}\")", table))
|
||||||
.fetch_all(pool)
|
.fetch_all(pool)
|
||||||
.await
|
.await
|
||||||
@@ -152,18 +122,13 @@ pub async fn list_foreign_keys(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_triggers(
|
pub async fn list_triggers(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<TriggerInfo>, String> {
|
||||||
pool: &SqlitePool,
|
let rows: Vec<SqliteRow> =
|
||||||
_schema: &str,
|
sqlx::query("SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND tbl_name = ? ORDER BY name")
|
||||||
table: &str,
|
.bind(table)
|
||||||
) -> Result<Vec<TriggerInfo>, String> {
|
.fetch_all(pool)
|
||||||
let rows: Vec<SqliteRow> = sqlx::query(
|
.await
|
||||||
"SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND tbl_name = ? ORDER BY name",
|
.map_err(|e| e.to_string())?;
|
||||||
)
|
|
||||||
.bind(table)
|
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -184,11 +149,7 @@ pub async fn list_triggers(
|
|||||||
} else {
|
} else {
|
||||||
"DELETE"
|
"DELETE"
|
||||||
};
|
};
|
||||||
TriggerInfo {
|
TriggerInfo { name: row.get::<String, _>("name"), event: event.to_string(), timing: timing.to_string() }
|
||||||
name: row.get::<String, _>("name"),
|
|
||||||
event: event.to_string(),
|
|
||||||
timing: timing.to_string(),
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
@@ -203,16 +164,9 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
|
|||||||
|| trimmed.starts_with("WITH")
|
|| trimmed.starts_with("WITH")
|
||||||
{
|
{
|
||||||
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
|
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
|
||||||
let columns: Vec<String> = desc
|
let columns: Vec<String> = desc.columns().iter().map(|c| c.name().to_string()).collect();
|
||||||
.columns()
|
|
||||||
.iter()
|
|
||||||
.map(|c| c.name().to_string())
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
let rows: Vec<SqliteRow> = sqlx::query(sql)
|
let rows: Vec<SqliteRow> = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_all(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let result_rows: Vec<Vec<serde_json::Value>> = rows
|
let result_rows: Vec<Vec<serde_json::Value>> = rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -221,10 +175,7 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
|
|||||||
.map(|i| {
|
.map(|i| {
|
||||||
row.try_get::<String, _>(i)
|
row.try_get::<String, _>(i)
|
||||||
.map(serde_json::Value::String)
|
.map(serde_json::Value::String)
|
||||||
.or_else(|_| {
|
.or_else(|_| row.try_get::<i64, _>(i).map(|v| serde_json::Value::Number(v.into())))
|
||||||
row.try_get::<i64, _>(i)
|
|
||||||
.map(|v| serde_json::Value::Number(v.into()))
|
|
||||||
})
|
|
||||||
.or_else(|_| {
|
.or_else(|_| {
|
||||||
row.try_get::<f64, _>(i).map(|v| {
|
row.try_get::<f64, _>(i).map(|v| {
|
||||||
serde_json::Number::from_f64(v)
|
serde_json::Number::from_f64(v)
|
||||||
@@ -247,10 +198,7 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
|
|||||||
truncated: false,
|
truncated: false,
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
let result = sqlx::query(sql)
|
let result = sqlx::query(sql).execute(pool).await.map_err(|e| e.to_string())?;
|
||||||
.execute(pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
Ok(QueryResult {
|
Ok(QueryResult {
|
||||||
columns: vec![],
|
columns: vec![],
|
||||||
|
|||||||
@@ -5,9 +5,7 @@ use tokio::net::TcpStream;
|
|||||||
use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt};
|
use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt};
|
||||||
|
|
||||||
use super::{connection_timeout, CONNECTION_TIMEOUT_SECS};
|
use super::{connection_timeout, CONNECTION_TIMEOUT_SECS};
|
||||||
use crate::types::{
|
use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo};
|
||||||
ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo,
|
|
||||||
};
|
|
||||||
|
|
||||||
pub type SqlServerClient = Client<Compat<TcpStream>>;
|
pub type SqlServerClient = Client<Compat<TcpStream>>;
|
||||||
|
|
||||||
@@ -48,13 +46,10 @@ async fn try_connect(
|
|||||||
.await
|
.await
|
||||||
.map_err(|_| format!("SQL Server connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
.map_err(|_| format!("SQL Server connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
||||||
.map_err(|e| format!("SQL Server connection failed: {e}"))?;
|
.map_err(|e| format!("SQL Server connection failed: {e}"))?;
|
||||||
tokio::time::timeout(
|
tokio::time::timeout(connection_timeout(), Client::connect(config, tcp.compat_write()))
|
||||||
connection_timeout(),
|
.await
|
||||||
Client::connect(config, tcp.compat_write()),
|
.map_err(|_| format!("SQL Server handshake timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
||||||
)
|
.map_err(|e| format!("SQL Server connection failed: {e}"))
|
||||||
.await
|
|
||||||
.map_err(|_| format!("SQL Server handshake timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
|
||||||
.map_err(|e| format!("SQL Server connection failed: {e}"))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> {
|
fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> {
|
||||||
@@ -69,9 +64,7 @@ fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> {
|
|||||||
} else if let Some(v) = row.try_get::<i64, _>(i).ok().flatten() {
|
} else if let Some(v) = row.try_get::<i64, _>(i).ok().flatten() {
|
||||||
serde_json::Value::Number(v.into())
|
serde_json::Value::Number(v.into())
|
||||||
} else if let Some(v) = row.try_get::<f64, _>(i).ok().flatten() {
|
} else if let Some(v) = row.try_get::<f64, _>(i).ok().flatten() {
|
||||||
serde_json::Number::from_f64(v)
|
serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
|
||||||
.map(serde_json::Value::Number)
|
|
||||||
.unwrap_or(serde_json::Value::Null)
|
|
||||||
} else if let Some(v) = row.try_get::<bool, _>(i).ok().flatten() {
|
} else if let Some(v) = row.try_get::<bool, _>(i).ok().flatten() {
|
||||||
serde_json::Value::Bool(v)
|
serde_json::Value::Bool(v)
|
||||||
} else {
|
} else {
|
||||||
@@ -82,20 +75,9 @@ fn row_to_json(row: &tiberius::Row) -> Vec<serde_json::Value> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases(client: &mut SqlServerClient) -> Result<Vec<DatabaseInfo>, String> {
|
pub async fn list_databases(client: &mut SqlServerClient) -> Result<Vec<DatabaseInfo>, String> {
|
||||||
let stream = client
|
let stream = client.query("SELECT name FROM sys.databases ORDER BY name", &[]).await.map_err(|e| e.to_string())?;
|
||||||
.query("SELECT name FROM sys.databases ORDER BY name", &[])
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.await
|
Ok(rows.iter().map(|row| DatabaseInfo { name: row.get::<&str, _>(0).unwrap_or("").to_string() }).collect())
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
let rows = stream
|
|
||||||
.into_first_result()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows
|
|
||||||
.iter()
|
|
||||||
.map(|row| DatabaseInfo {
|
|
||||||
name: row.get::<&str, _>(0).unwrap_or("").to_string(),
|
|
||||||
})
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_schemas(client: &mut SqlServerClient) -> Result<Vec<String>, String> {
|
pub async fn list_schemas(client: &mut SqlServerClient) -> Result<Vec<String>, String> {
|
||||||
@@ -108,29 +90,17 @@ pub async fn list_schemas(client: &mut SqlServerClient) -> Result<Vec<String>, S
|
|||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
let rows = stream
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.into_first_result()
|
Ok(rows.iter().map(|row| row.get::<&str, _>(0).unwrap_or("").to_string()).collect())
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows
|
|
||||||
.iter()
|
|
||||||
.map(|row| row.get::<&str, _>(0).unwrap_or("").to_string())
|
|
||||||
.collect())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_tables(
|
pub async fn list_tables(client: &mut SqlServerClient, schema: &str) -> Result<Vec<TableInfo>, String> {
|
||||||
client: &mut SqlServerClient,
|
|
||||||
schema: &str,
|
|
||||||
) -> Result<Vec<TableInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT TABLE_NAME, TABLE_TYPE FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA = '{}' ORDER BY TABLE_NAME",
|
"SELECT TABLE_NAME, TABLE_TYPE FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA = '{}' ORDER BY TABLE_NAME",
|
||||||
schema.replace('\'', "''")
|
schema.replace('\'', "''")
|
||||||
);
|
);
|
||||||
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
||||||
let rows = stream
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.into_first_result()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| TableInfo {
|
.map(|row| TableInfo {
|
||||||
@@ -140,11 +110,7 @@ pub async fn list_tables(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_columns(
|
pub async fn get_columns(client: &mut SqlServerClient, schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
|
||||||
client: &mut SqlServerClient,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<ColumnInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT c.COLUMN_NAME, c.DATA_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, \
|
"SELECT c.COLUMN_NAME, c.DATA_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, \
|
||||||
CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \
|
CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \
|
||||||
@@ -158,10 +124,7 @@ pub async fn get_columns(
|
|||||||
s = schema.replace('\'', "''"), t = table.replace('\'', "''")
|
s = schema.replace('\'', "''"), t = table.replace('\'', "''")
|
||||||
);
|
);
|
||||||
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
||||||
let rows = stream
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.into_first_result()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
@@ -216,11 +179,7 @@ pub async fn get_columns(
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_indexes(
|
pub async fn list_indexes(client: &mut SqlServerClient, schema: &str, table: &str) -> Result<Vec<IndexInfo>, String> {
|
||||||
client: &mut SqlServerClient,
|
|
||||||
schema: &str,
|
|
||||||
table: &str,
|
|
||||||
) -> Result<Vec<IndexInfo>, String> {
|
|
||||||
let sql = format!(
|
let sql = format!(
|
||||||
"SELECT i.name, \
|
"SELECT i.name, \
|
||||||
STRING_AGG(CASE WHEN ic.is_included_column = 0 THEN c.name END, ',') WITHIN GROUP (ORDER BY ic.key_ordinal) AS columns, \
|
STRING_AGG(CASE WHEN ic.is_included_column = 0 THEN c.name END, ',') WITHIN GROUP (ORDER BY ic.key_ordinal) AS columns, \
|
||||||
@@ -236,10 +195,7 @@ pub async fn list_indexes(
|
|||||||
s = schema.replace('\'', "''"), t = table.replace('\'', "''")
|
s = schema.replace('\'', "''"), t = table.replace('\'', "''")
|
||||||
);
|
);
|
||||||
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
||||||
let rows = stream
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.into_first_result()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| {
|
.map(|row| {
|
||||||
@@ -247,11 +203,7 @@ pub async fn list_indexes(
|
|||||||
let inc_str = row.get::<&str, _>(5).unwrap_or("");
|
let inc_str = row.get::<&str, _>(5).unwrap_or("");
|
||||||
IndexInfo {
|
IndexInfo {
|
||||||
name: row.get::<&str, _>(0).unwrap_or("").to_string(),
|
name: row.get::<&str, _>(0).unwrap_or("").to_string(),
|
||||||
columns: cols_str
|
columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(),
|
||||||
.split(',')
|
|
||||||
.filter(|s| !s.is_empty())
|
|
||||||
.map(|s| s.to_string())
|
|
||||||
.collect(),
|
|
||||||
is_unique: row.get::<bool, _>(2).unwrap_or(false),
|
is_unique: row.get::<bool, _>(2).unwrap_or(false),
|
||||||
is_primary: row.get::<bool, _>(3).unwrap_or(false),
|
is_primary: row.get::<bool, _>(3).unwrap_or(false),
|
||||||
filter: row.get::<&str, _>(6).map(|s| s.to_string()),
|
filter: row.get::<&str, _>(6).map(|s| s.to_string()),
|
||||||
@@ -281,13 +233,11 @@ pub async fn list_foreign_keys(
|
|||||||
JOIN sys.columns rc ON fkc.referenced_object_id = rc.object_id AND fkc.referenced_column_id = rc.column_id \
|
JOIN sys.columns rc ON fkc.referenced_object_id = rc.object_id AND fkc.referenced_column_id = rc.column_id \
|
||||||
WHERE fk.parent_object_id = OBJECT_ID('{s}.{t}') \
|
WHERE fk.parent_object_id = OBJECT_ID('{s}.{t}') \
|
||||||
ORDER BY fk.name",
|
ORDER BY fk.name",
|
||||||
s = schema.replace('\'', "''"), t = table.replace('\'', "''")
|
s = schema.replace('\'', "''"),
|
||||||
|
t = table.replace('\'', "''")
|
||||||
);
|
);
|
||||||
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
||||||
let rows = stream
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.into_first_result()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| ForeignKeyInfo {
|
.map(|row| ForeignKeyInfo {
|
||||||
@@ -310,13 +260,11 @@ pub async fn list_triggers(
|
|||||||
JOIN sys.trigger_events te ON t.object_id = te.object_id \
|
JOIN sys.trigger_events te ON t.object_id = te.object_id \
|
||||||
WHERE t.parent_id = OBJECT_ID('{s}.{t}') \
|
WHERE t.parent_id = OBJECT_ID('{s}.{t}') \
|
||||||
ORDER BY t.name",
|
ORDER BY t.name",
|
||||||
s = schema.replace('\'', "''"), t = table.replace('\'', "''")
|
s = schema.replace('\'', "''"),
|
||||||
|
t = table.replace('\'', "''")
|
||||||
);
|
);
|
||||||
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?;
|
||||||
let rows = stream
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.into_first_result()
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows
|
Ok(rows
|
||||||
.iter()
|
.iter()
|
||||||
.map(|row| TriggerInfo {
|
.map(|row| TriggerInfo {
|
||||||
@@ -341,19 +289,11 @@ pub async fn execute_query(client: &mut SqlServerClient, sql: &str) -> Result<Qu
|
|||||||
.columns()
|
.columns()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| e.to_string())?
|
.map_err(|e| e.to_string())?
|
||||||
.map(|cols| {
|
.map(|cols| cols.iter().map(|c| c.name().to_string()).collect::<Vec<_>>())
|
||||||
cols.iter()
|
|
||||||
.map(|c| c.name().to_string())
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
})
|
|
||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
|
|
||||||
let rows = stream
|
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
|
||||||
.into_first_result()
|
let result_rows: Vec<Vec<serde_json::Value>> = rows.iter().map(|row| row_to_json(row)).collect();
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
let result_rows: Vec<Vec<serde_json::Value>> =
|
|
||||||
rows.iter().map(|row| row_to_json(row)).collect();
|
|
||||||
|
|
||||||
Ok(QueryResult {
|
Ok(QueryResult {
|
||||||
columns: columns_meta,
|
columns: columns_meta,
|
||||||
|
|||||||
@@ -32,39 +32,24 @@ async fn connect_and_authenticate(
|
|||||||
ssh_key_path: &str,
|
ssh_key_path: &str,
|
||||||
ssh_key_passphrase: &str,
|
ssh_key_passphrase: &str,
|
||||||
) -> Result<Handle<SshClient>, String> {
|
) -> Result<Handle<SshClient>, String> {
|
||||||
let config = Arc::new(Config {
|
let config = Arc::new(Config { nodelay: true, ..Default::default() });
|
||||||
nodelay: true,
|
|
||||||
..Default::default()
|
|
||||||
});
|
|
||||||
|
|
||||||
let mut session = tokio::time::timeout(
|
let mut session =
|
||||||
connection_timeout(),
|
tokio::time::timeout(connection_timeout(), client::connect(config, (ssh_host, ssh_port), SshClient {}))
|
||||||
client::connect(config, (ssh_host, ssh_port), SshClient {}),
|
.await
|
||||||
)
|
.map_err(|_| format!("SSH connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
||||||
.await
|
.map_err(|e| format!("SSH connection failed: {e}"))?;
|
||||||
.map_err(|_| format!("SSH connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
|
||||||
.map_err(|e| format!("SSH connection failed: {e}"))?;
|
|
||||||
|
|
||||||
if !ssh_key_path.is_empty() {
|
if !ssh_key_path.is_empty() {
|
||||||
let passphrase = if ssh_key_passphrase.is_empty() {
|
let passphrase = if ssh_key_passphrase.is_empty() { None } else { Some(ssh_key_passphrase) };
|
||||||
None
|
let key_pair = load_secret_key(ssh_key_path, passphrase).map_err(|e| format!("Failed to load SSH key: {e}"))?;
|
||||||
} else {
|
|
||||||
Some(ssh_key_passphrase)
|
|
||||||
};
|
|
||||||
let key_pair = load_secret_key(ssh_key_path, passphrase)
|
|
||||||
.map_err(|e| format!("Failed to load SSH key: {e}"))?;
|
|
||||||
let auth_res = tokio::time::timeout(
|
let auth_res = tokio::time::timeout(
|
||||||
connection_timeout(),
|
connection_timeout(),
|
||||||
session.authenticate_publickey(
|
session.authenticate_publickey(
|
||||||
ssh_user,
|
ssh_user,
|
||||||
PrivateKeyWithHashAlg::new(
|
PrivateKeyWithHashAlg::new(
|
||||||
Arc::new(key_pair),
|
Arc::new(key_pair),
|
||||||
session
|
session.best_supported_rsa_hash().await.ok().flatten().flatten(),
|
||||||
.best_supported_rsa_hash()
|
|
||||||
.await
|
|
||||||
.ok()
|
|
||||||
.flatten()
|
|
||||||
.flatten(),
|
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -75,13 +60,11 @@ async fn connect_and_authenticate(
|
|||||||
return Err("SSH public key authentication failed".to_string());
|
return Err("SSH public key authentication failed".to_string());
|
||||||
}
|
}
|
||||||
} else if !ssh_password.is_empty() {
|
} else if !ssh_password.is_empty() {
|
||||||
let auth_res = tokio::time::timeout(
|
let auth_res =
|
||||||
connection_timeout(),
|
tokio::time::timeout(connection_timeout(), session.authenticate_password(ssh_user, ssh_password))
|
||||||
session.authenticate_password(ssh_user, ssh_password),
|
.await
|
||||||
)
|
.map_err(|_| format!("SSH password auth timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
||||||
.await
|
.map_err(|e| format!("SSH password auth failed: {e}"))?;
|
||||||
.map_err(|_| format!("SSH password auth timed out ({CONNECTION_TIMEOUT_SECS}s)"))?
|
|
||||||
.map_err(|e| format!("SSH password auth failed: {e}"))?;
|
|
||||||
if !auth_res.success() {
|
if !auth_res.success() {
|
||||||
return Err("SSH password authentication failed".to_string());
|
return Err("SSH password authentication failed".to_string());
|
||||||
}
|
}
|
||||||
@@ -92,12 +75,7 @@ async fn connect_and_authenticate(
|
|||||||
Ok(session)
|
Ok(session)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn forward_loop(
|
async fn forward_loop(session: Handle<SshClient>, listener: TcpListener, remote_host: String, remote_port: u16) {
|
||||||
session: Handle<SshClient>,
|
|
||||||
listener: TcpListener,
|
|
||||||
remote_host: String,
|
|
||||||
remote_port: u16,
|
|
||||||
) {
|
|
||||||
loop {
|
loop {
|
||||||
let (mut stream, peer_addr) = match listener.accept().await {
|
let (mut stream, peer_addr) = match listener.accept().await {
|
||||||
Ok(v) => v,
|
Ok(v) => v,
|
||||||
@@ -163,9 +141,7 @@ pub struct TunnelManager {
|
|||||||
|
|
||||||
impl TunnelManager {
|
impl TunnelManager {
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self { tunnels: Mutex::new(HashMap::new()) }
|
||||||
tunnels: Mutex::new(HashMap::new()),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn start_tunnel(
|
pub async fn start_tunnel(
|
||||||
@@ -183,42 +159,24 @@ impl TunnelManager {
|
|||||||
) -> Result<u16, String> {
|
) -> Result<u16, String> {
|
||||||
let local_port = portpicker::pick_unused_port().ok_or("No available port")?;
|
let local_port = portpicker::pick_unused_port().ok_or("No available port")?;
|
||||||
|
|
||||||
let session = connect_and_authenticate(
|
let session =
|
||||||
ssh_host,
|
connect_and_authenticate(ssh_host, ssh_port, ssh_user, ssh_password, ssh_key_path, ssh_key_passphrase)
|
||||||
ssh_port,
|
.await?;
|
||||||
ssh_user,
|
|
||||||
ssh_password,
|
|
||||||
ssh_key_path,
|
|
||||||
ssh_key_passphrase,
|
|
||||||
)
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
let bind_addr = if expose_to_lan {
|
let bind_addr = if expose_to_lan { "0.0.0.0" } else { "127.0.0.1" };
|
||||||
"0.0.0.0"
|
let listener =
|
||||||
} else {
|
TcpListener::bind((bind_addr, local_port)).await.map_err(|e| format!("Failed to bind local port: {e}"))?;
|
||||||
"127.0.0.1"
|
|
||||||
};
|
|
||||||
let listener = TcpListener::bind((bind_addr, local_port))
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Failed to bind local port: {e}"))?;
|
|
||||||
|
|
||||||
let remote_host = remote_host.to_string();
|
let remote_host = remote_host.to_string();
|
||||||
let handle = tokio::spawn(forward_loop(session, listener, remote_host, remote_port));
|
let handle = tokio::spawn(forward_loop(session, listener, remote_host, remote_port));
|
||||||
|
|
||||||
self.tunnels
|
self.tunnels.lock().await.insert(connection_id.to_string(), (handle, local_port));
|
||||||
.lock()
|
|
||||||
.await
|
|
||||||
.insert(connection_id.to_string(), (handle, local_port));
|
|
||||||
|
|
||||||
Ok(local_port)
|
Ok(local_port)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn local_port(&self, connection_id: &str) -> Option<u16> {
|
pub async fn local_port(&self, connection_id: &str) -> Option<u16> {
|
||||||
self.tunnels
|
self.tunnels.lock().await.get(connection_id).map(|(_, port)| *port)
|
||||||
.lock()
|
|
||||||
.await
|
|
||||||
.get(connection_id)
|
|
||||||
.map(|(_, port)| *port)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn stop_tunnel(&self, connection_id: &str) {
|
pub async fn stop_tunnel(&self, connection_id: &str) {
|
||||||
|
|||||||
@@ -45,9 +45,6 @@ pub fn clear_history_entries(path: &Path) -> Result<(), String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn delete_history_entry_by_id(path: &Path, id: &str) -> Result<(), String> {
|
pub fn delete_history_entry_by_id(path: &Path, id: &str) -> Result<(), String> {
|
||||||
let entries: Vec<HistoryEntry> = read_all(path)?
|
let entries: Vec<HistoryEntry> = read_all(path)?.into_iter().filter(|e| e.id != id).collect();
|
||||||
.into_iter()
|
|
||||||
.filter(|e| e.id != id)
|
|
||||||
.collect();
|
|
||||||
write_all(path, &entries)
|
write_all(path, &entries)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -73,7 +73,10 @@ pub enum DatabaseType {
|
|||||||
impl ConnectionConfig {
|
impl ConnectionConfig {
|
||||||
pub fn needs_bare_mysql(&self) -> bool {
|
pub fn needs_bare_mysql(&self) -> bool {
|
||||||
matches!(self.db_type, DatabaseType::Doris | DatabaseType::StarRocks)
|
matches!(self.db_type, DatabaseType::Doris | DatabaseType::StarRocks)
|
||||||
|| self.driver_profile.as_deref().map(|p| p.to_lowercase())
|
|| self
|
||||||
|
.driver_profile
|
||||||
|
.as_deref()
|
||||||
|
.map(|p| p.to_lowercase())
|
||||||
.is_some_and(|p| matches!(p.as_str(), "doris" | "starrocks" | "selectdb" | "tdengine"))
|
.is_some_and(|p| matches!(p.as_str(), "doris" | "starrocks" | "selectdb" | "tdengine"))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -103,20 +106,17 @@ impl ConnectionConfig {
|
|||||||
let scheme = if self.ssl { "rediss" } else { "redis" };
|
let scheme = if self.ssl { "rediss" } else { "redis" };
|
||||||
format!("{scheme}://{host}:{port}/")
|
format!("{scheme}://{host}:{port}/")
|
||||||
}
|
}
|
||||||
DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => format!("mysql://{host}:{port}{db_part}?{params}"),
|
DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => {
|
||||||
|
format!("mysql://{host}:{port}{db_part}?{params}")
|
||||||
|
}
|
||||||
DatabaseType::Postgres | DatabaseType::Redshift => {
|
DatabaseType::Postgres | DatabaseType::Redshift => {
|
||||||
let suffix = if params.is_empty() {
|
let suffix = if params.is_empty() { String::new() } else { format!("?{params}") };
|
||||||
String::new()
|
|
||||||
} else {
|
|
||||||
format!("?{params}")
|
|
||||||
};
|
|
||||||
format!("postgres://{host}:{port}{db_part}{suffix}")
|
format!("postgres://{host}:{port}{db_part}{suffix}")
|
||||||
}
|
}
|
||||||
DatabaseType::ClickHouse => format!("http://{host}:{port}"),
|
DatabaseType::ClickHouse => format!("http://{host}:{port}"),
|
||||||
DatabaseType::SqlServer => format!(
|
DatabaseType::SqlServer => {
|
||||||
"server=tcp:{host},{port};database={}",
|
format!("server=tcp:{host},{port};database={}", self.database.as_deref().unwrap_or("master"))
|
||||||
self.database.as_deref().unwrap_or("master")
|
}
|
||||||
),
|
|
||||||
DatabaseType::MongoDb => {
|
DatabaseType::MongoDb => {
|
||||||
if let Some(cs) = self.connection_string.as_deref().filter(|s| !s.is_empty()) {
|
if let Some(cs) = self.connection_string.as_deref().filter(|s| !s.is_empty()) {
|
||||||
return cs.to_string();
|
return cs.to_string();
|
||||||
@@ -154,20 +154,12 @@ impl ConnectionConfig {
|
|||||||
format!("{scheme}://{username}:{password}@{host}:{port}/")
|
format!("{scheme}://{username}:{password}@{host}:{port}/")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => format!(
|
DatabaseType::Mysql | DatabaseType::Doris | DatabaseType::StarRocks => {
|
||||||
"mysql://{}:{}@{host}:{port}{db_part}?{params}",
|
format!("mysql://{}:{}@{host}:{port}{db_part}?{params}", username, password)
|
||||||
username, password
|
}
|
||||||
),
|
|
||||||
DatabaseType::Postgres | DatabaseType::Redshift => {
|
DatabaseType::Postgres | DatabaseType::Redshift => {
|
||||||
let suffix = if params.is_empty() {
|
let suffix = if params.is_empty() { String::new() } else { format!("?{params}") };
|
||||||
String::new()
|
format!("postgres://{}:{}@{host}:{port}{db_part}{suffix}", username, password)
|
||||||
} else {
|
|
||||||
format!("?{params}")
|
|
||||||
};
|
|
||||||
format!(
|
|
||||||
"postgres://{}:{}@{host}:{port}{db_part}{suffix}",
|
|
||||||
username, password
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
DatabaseType::ClickHouse => format!("http://{host}:{port}"),
|
DatabaseType::ClickHouse => format!("http://{host}:{port}"),
|
||||||
DatabaseType::SqlServer => format!(
|
DatabaseType::SqlServer => format!(
|
||||||
@@ -197,10 +189,15 @@ impl ConnectionConfig {
|
|||||||
let value = self.url_params.as_deref().unwrap_or("").trim();
|
let value = self.url_params.as_deref().unwrap_or("").trim();
|
||||||
if self.needs_bare_mysql() {
|
if self.needs_bare_mysql() {
|
||||||
let v = value.trim_start_matches('?');
|
let v = value.trim_start_matches('?');
|
||||||
let filtered: Vec<&str> = v.split('&')
|
let filtered: Vec<&str> = v
|
||||||
|
.split('&')
|
||||||
.filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred"))
|
.filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred"))
|
||||||
.collect();
|
.collect();
|
||||||
return if filtered.is_empty() { "ssl-mode=disabled".to_string() } else { format!("ssl-mode=disabled&{}", filtered.join("&")) };
|
return if filtered.is_empty() {
|
||||||
|
"ssl-mode=disabled".to_string()
|
||||||
|
} else {
|
||||||
|
format!("ssl-mode=disabled&{}", filtered.join("&"))
|
||||||
|
};
|
||||||
}
|
}
|
||||||
match self.db_type {
|
match self.db_type {
|
||||||
DatabaseType::Mysql => {
|
DatabaseType::Mysql => {
|
||||||
@@ -209,18 +206,31 @@ impl ConnectionConfig {
|
|||||||
base.to_string()
|
base.to_string()
|
||||||
} else if value.contains("ssl-mode=") {
|
} else if value.contains("ssl-mode=") {
|
||||||
let v = value.trim_start_matches('?');
|
let v = value.trim_start_matches('?');
|
||||||
if v.contains("charset=") { v.to_string() } else { format!("{v}&charset=utf8mb4") }
|
if v.contains("charset=") {
|
||||||
|
v.to_string()
|
||||||
|
} else {
|
||||||
|
format!("{v}&charset=utf8mb4")
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
let v = value.trim_start_matches('?');
|
let v = value.trim_start_matches('?');
|
||||||
if v.contains("charset=") { format!("ssl-mode=preferred&{v}") } else { format!("{base}&{v}") }
|
if v.contains("charset=") {
|
||||||
|
format!("ssl-mode=preferred&{v}")
|
||||||
|
} else {
|
||||||
|
format!("{base}&{v}")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
DatabaseType::Doris | DatabaseType::StarRocks => {
|
DatabaseType::Doris | DatabaseType::StarRocks => {
|
||||||
let v = value.trim_start_matches('?');
|
let v = value.trim_start_matches('?');
|
||||||
let filtered: Vec<&str> = v.split('&')
|
let filtered: Vec<&str> = v
|
||||||
|
.split('&')
|
||||||
.filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred"))
|
.filter(|p| !p.is_empty() && !p.starts_with("charset=") && !p.starts_with("ssl-mode=preferred"))
|
||||||
.collect();
|
.collect();
|
||||||
if filtered.is_empty() { "ssl-mode=disabled".to_string() } else { format!("ssl-mode=disabled&{}", filtered.join("&")) }
|
if filtered.is_empty() {
|
||||||
|
"ssl-mode=disabled".to_string()
|
||||||
|
} else {
|
||||||
|
format!("ssl-mode=disabled&{}", filtered.join("&"))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
DatabaseType::Postgres | DatabaseType::Redshift => value.trim_start_matches('?').to_string(),
|
DatabaseType::Postgres | DatabaseType::Redshift => value.trim_start_matches('?').to_string(),
|
||||||
_ => value.trim_start_matches('?').to_string(),
|
_ => value.trim_start_matches('?').to_string(),
|
||||||
@@ -308,10 +318,7 @@ mod tests {
|
|||||||
config.db_type = DatabaseType::Postgres;
|
config.db_type = DatabaseType::Postgres;
|
||||||
config.url_params = Some("sslmode=disable".to_string());
|
config.url_params = Some("sslmode=disable".to_string());
|
||||||
|
|
||||||
assert_eq!(
|
assert_eq!(config.connection_url(), "postgres://postgres:secret@10.1.2.3:2883/test?sslmode=disable");
|
||||||
config.connection_url(),
|
|
||||||
"postgres://postgres:secret@10.1.2.3:2883/test?sslmode=disable"
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -1,11 +1,8 @@
|
|||||||
use crate::connection::{AppState, PoolKind};
|
use crate::connection::{AppState, PoolKind};
|
||||||
use crate::db::mongo_driver::{self, MongoDocumentResult};
|
|
||||||
use crate::db::elasticsearch_driver;
|
use crate::db::elasticsearch_driver;
|
||||||
|
use crate::db::mongo_driver::{self, MongoDocumentResult};
|
||||||
|
|
||||||
pub async fn mongo_list_databases_core(
|
pub async fn mongo_list_databases_core(state: &AppState, connection_id: &str) -> Result<Vec<String>, String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
) -> Result<Vec<String>, String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::MongoDb(client) => mongo_driver::list_databases(client).await,
|
PoolKind::MongoDb(client) => mongo_driver::list_databases(client).await,
|
||||||
@@ -37,9 +34,7 @@ pub async fn mongo_find_documents_core(
|
|||||||
) -> Result<MongoDocumentResult, String> {
|
) -> Result<MongoDocumentResult, String> {
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::MongoDb(client) => {
|
PoolKind::MongoDb(client) => mongo_driver::find_documents(client, database, collection, skip, limit).await,
|
||||||
mongo_driver::find_documents(client, database, collection, skip, limit).await
|
|
||||||
}
|
|
||||||
PoolKind::Elasticsearch(client) => {
|
PoolKind::Elasticsearch(client) => {
|
||||||
let client = client.clone();
|
let client = client.clone();
|
||||||
drop(connections);
|
drop(connections);
|
||||||
@@ -58,9 +53,7 @@ pub async fn mongo_insert_document_core(
|
|||||||
) -> Result<String, String> {
|
) -> Result<String, String> {
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::MongoDb(client) => {
|
PoolKind::MongoDb(client) => mongo_driver::insert_document(client, database, collection, doc_json).await,
|
||||||
mongo_driver::insert_document(client, database, collection, doc_json).await
|
|
||||||
}
|
|
||||||
PoolKind::Elasticsearch(client) => {
|
PoolKind::Elasticsearch(client) => {
|
||||||
let client = client.clone();
|
let client = client.clone();
|
||||||
drop(connections);
|
drop(connections);
|
||||||
@@ -80,9 +73,7 @@ pub async fn mongo_update_document_core(
|
|||||||
) -> Result<u64, String> {
|
) -> Result<u64, String> {
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::MongoDb(client) => {
|
PoolKind::MongoDb(client) => mongo_driver::update_document(client, database, collection, id, doc_json).await,
|
||||||
mongo_driver::update_document(client, database, collection, id, doc_json).await
|
|
||||||
}
|
|
||||||
PoolKind::Elasticsearch(client) => {
|
PoolKind::Elasticsearch(client) => {
|
||||||
let client = client.clone();
|
let client = client.clone();
|
||||||
drop(connections);
|
drop(connections);
|
||||||
@@ -101,9 +92,7 @@ pub async fn mongo_delete_document_core(
|
|||||||
) -> Result<u64, String> {
|
) -> Result<u64, String> {
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::MongoDb(client) => {
|
PoolKind::MongoDb(client) => mongo_driver::delete_document(client, database, collection, id).await,
|
||||||
mongo_driver::delete_document(client, database, collection, id).await
|
|
||||||
}
|
|
||||||
PoolKind::Elasticsearch(client) => {
|
PoolKind::Elasticsearch(client) => {
|
||||||
let client = client.clone();
|
let client = client.clone();
|
||||||
drop(connections);
|
drop(connections);
|
||||||
|
|||||||
@@ -15,8 +15,12 @@ pub fn duckdb_execute(con: &duckdb::Connection, sql: &str) -> Result<db::QueryRe
|
|||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
let trimmed = sql.trim().to_uppercase();
|
let trimmed = sql.trim().to_uppercase();
|
||||||
|
|
||||||
if trimmed.starts_with("SELECT") || trimmed.starts_with("SHOW") || trimmed.starts_with("DESCRIBE")
|
if trimmed.starts_with("SELECT")
|
||||||
|| trimmed.starts_with("EXPLAIN") || trimmed.starts_with("WITH") || trimmed.starts_with("PRAGMA")
|
|| trimmed.starts_with("SHOW")
|
||||||
|
|| trimmed.starts_with("DESCRIBE")
|
||||||
|
|| trimmed.starts_with("EXPLAIN")
|
||||||
|
|| trimmed.starts_with("WITH")
|
||||||
|
|| trimmed.starts_with("PRAGMA")
|
||||||
{
|
{
|
||||||
let mut stmt = con.prepare(sql).map_err(|e| e.to_string())?;
|
let mut stmt = con.prepare(sql).map_err(|e| e.to_string())?;
|
||||||
let mut rows = stmt.query([]).map_err(|e| e.to_string())?;
|
let mut rows = stmt.query([]).map_err(|e| e.to_string())?;
|
||||||
@@ -28,27 +32,45 @@ pub fn duckdb_execute(con: &duckdb::Connection, sql: &str) -> Result<db::QueryRe
|
|||||||
|
|
||||||
let mut result_rows = Vec::new();
|
let mut result_rows = Vec::new();
|
||||||
while let Some(row) = rows.next().map_err(|e| e.to_string())? {
|
while let Some(row) = rows.next().map_err(|e| e.to_string())? {
|
||||||
if result_rows.len() >= MAX_ROWS { break; }
|
if result_rows.len() >= MAX_ROWS {
|
||||||
let vals: Vec<serde_json::Value> = (0..col_count).map(|i| {
|
break;
|
||||||
row.get::<_, String>(i)
|
}
|
||||||
.map(serde_json::Value::String)
|
let vals: Vec<serde_json::Value> = (0..col_count)
|
||||||
.or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into())))
|
.map(|i| {
|
||||||
.or_else(|_| row.get::<_, f64>(i).map(|v| {
|
row.get::<_, String>(i)
|
||||||
serde_json::Number::from_f64(v)
|
.map(serde_json::Value::String)
|
||||||
.map(serde_json::Value::Number)
|
.or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into())))
|
||||||
.unwrap_or(serde_json::Value::Null)
|
.or_else(|_| {
|
||||||
}))
|
row.get::<_, f64>(i).map(|v| {
|
||||||
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool))
|
serde_json::Number::from_f64(v)
|
||||||
.unwrap_or(serde_json::Value::Null)
|
.map(serde_json::Value::Number)
|
||||||
}).collect();
|
.unwrap_or(serde_json::Value::Null)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool))
|
||||||
|
.unwrap_or(serde_json::Value::Null)
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
result_rows.push(vals);
|
result_rows.push(vals);
|
||||||
}
|
}
|
||||||
|
|
||||||
let truncated = result_rows.len() >= MAX_ROWS;
|
let truncated = result_rows.len() >= MAX_ROWS;
|
||||||
Ok(db::QueryResult { columns, rows: result_rows, affected_rows: 0, execution_time_ms: start.elapsed().as_millis(), truncated })
|
Ok(db::QueryResult {
|
||||||
|
columns,
|
||||||
|
rows: result_rows,
|
||||||
|
affected_rows: 0,
|
||||||
|
execution_time_ms: start.elapsed().as_millis(),
|
||||||
|
truncated,
|
||||||
|
})
|
||||||
} else {
|
} else {
|
||||||
let affected = con.execute(sql, []).map_err(|e| e.to_string())?;
|
let affected = con.execute(sql, []).map_err(|e| e.to_string())?;
|
||||||
Ok(db::QueryResult { columns: vec![], rows: vec![], affected_rows: affected as u64, execution_time_ms: start.elapsed().as_millis(), truncated: false })
|
Ok(db::QueryResult {
|
||||||
|
columns: vec![],
|
||||||
|
rows: vec![],
|
||||||
|
affected_rows: affected as u64,
|
||||||
|
execution_time_ms: start.elapsed().as_millis(),
|
||||||
|
truncated: false,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -79,16 +101,10 @@ pub fn canceled_error() -> String {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn is_canceled(cancel_token: &Option<CancellationToken>) -> bool {
|
pub fn is_canceled(cancel_token: &Option<CancellationToken>) -> bool {
|
||||||
cancel_token
|
cancel_token.as_ref().map(|token| token.is_cancelled()).unwrap_or(false)
|
||||||
.as_ref()
|
|
||||||
.map(|token| token.is_cancelled())
|
|
||||||
.unwrap_or(false)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn wait_for_query<F>(
|
pub async fn wait_for_query<F>(cancel_token: Option<CancellationToken>, future: F) -> Result<db::QueryResult, String>
|
||||||
cancel_token: Option<CancellationToken>,
|
|
||||||
future: F,
|
|
||||||
) -> Result<db::QueryResult, String>
|
|
||||||
where
|
where
|
||||||
F: Future<Output = Result<db::QueryResult, String>>,
|
F: Future<Output = Result<db::QueryResult, String>>,
|
||||||
{
|
{
|
||||||
@@ -110,9 +126,7 @@ where
|
|||||||
result = timeout(timeout_duration, future) => result.map_err(|_| timeout_error())?,
|
result = timeout(timeout_duration, future) => result.map_err(|_| timeout_error())?,
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
timeout(timeout_duration, future)
|
timeout(timeout_duration, future).await.map_err(|_| timeout_error())?
|
||||||
.await
|
|
||||||
.map_err(|_| timeout_error())?
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -143,23 +157,17 @@ pub async fn do_execute(
|
|||||||
let p = p.clone();
|
let p = p.clone();
|
||||||
let bare = *bare;
|
let bare = *bare;
|
||||||
drop(connections);
|
drop(connections);
|
||||||
wait_for_query(cancel_token, db::mysql::execute_query(&p, sql, bare))
|
wait_for_query(cancel_token, db::mysql::execute_query(&p, sql, bare)).await.map(truncate_result)
|
||||||
.await
|
|
||||||
.map(truncate_result)
|
|
||||||
}
|
}
|
||||||
PoolKind::Postgres(p) => {
|
PoolKind::Postgres(p) => {
|
||||||
let p = p.clone();
|
let p = p.clone();
|
||||||
drop(connections);
|
drop(connections);
|
||||||
wait_for_query(cancel_token, db::postgres::execute_query(&p, sql))
|
wait_for_query(cancel_token, db::postgres::execute_query(&p, sql)).await.map(truncate_result)
|
||||||
.await
|
|
||||||
.map(truncate_result)
|
|
||||||
}
|
}
|
||||||
PoolKind::Sqlite(p) => {
|
PoolKind::Sqlite(p) => {
|
||||||
let p = p.clone();
|
let p = p.clone();
|
||||||
drop(connections);
|
drop(connections);
|
||||||
wait_for_query(cancel_token, db::sqlite::execute_query(&p, sql))
|
wait_for_query(cancel_token, db::sqlite::execute_query(&p, sql)).await.map(truncate_result)
|
||||||
.await
|
|
||||||
.map(truncate_result)
|
|
||||||
}
|
}
|
||||||
PoolKind::ClickHouse(client) => {
|
PoolKind::ClickHouse(client) => {
|
||||||
let client = client.clone();
|
let client = client.clone();
|
||||||
@@ -180,9 +188,7 @@ pub async fn do_execute(
|
|||||||
},
|
},
|
||||||
None => client.lock().await,
|
None => client.lock().await,
|
||||||
};
|
};
|
||||||
wait_for_query(cancel_token, db::sqlserver::execute_query(&mut client, sql))
|
wait_for_query(cancel_token, db::sqlserver::execute_query(&mut client, sql)).await.map(truncate_result)
|
||||||
.await
|
|
||||||
.map(truncate_result)
|
|
||||||
}
|
}
|
||||||
PoolKind::Oracle(client) => {
|
PoolKind::Oracle(client) => {
|
||||||
let client = client.clone();
|
let client = client.clone();
|
||||||
@@ -195,9 +201,7 @@ pub async fn do_execute(
|
|||||||
},
|
},
|
||||||
None => client.lock().await,
|
None => client.lock().await,
|
||||||
};
|
};
|
||||||
wait_for_query(cancel_token, db::oracle_driver::execute_query(&*client, sql))
|
wait_for_query(cancel_token, db::oracle_driver::execute_query(&*client, sql)).await.map(truncate_result)
|
||||||
.await
|
|
||||||
.map(truncate_result)
|
|
||||||
}
|
}
|
||||||
PoolKind::Elasticsearch(_) => Err("Use document browser for Elasticsearch".to_string()),
|
PoolKind::Elasticsearch(_) => Err("Use document browser for Elasticsearch".to_string()),
|
||||||
PoolKind::Redis(_) => Err("Use Redis-specific commands".to_string()),
|
PoolKind::Redis(_) => Err("Use Redis-specific commands".to_string()),
|
||||||
@@ -244,9 +248,7 @@ pub async fn execute_multi_core(
|
|||||||
let statements = split_sql_statements(sql);
|
let statements = split_sql_statements(sql);
|
||||||
if statements.len() <= 1 {
|
if statements.len() <= 1 {
|
||||||
let single_sql = statements.into_iter().next().unwrap_or_default();
|
let single_sql = statements.into_iter().next().unwrap_or_default();
|
||||||
let result = execute_sql_statement(
|
let result = execute_sql_statement(state, connection_id, database, &single_sql, cancel_token).await?;
|
||||||
state, connection_id, database, &single_sql, cancel_token,
|
|
||||||
).await?;
|
|
||||||
return Ok(vec![result]);
|
return Ok(vec![result]);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -262,9 +264,7 @@ pub async fn execute_multi_core(
|
|||||||
});
|
});
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
match execute_sql_statement(
|
match execute_sql_statement(state, connection_id, database, stmt, cancel_token.clone()).await {
|
||||||
state, connection_id, database, stmt, cancel_token.clone(),
|
|
||||||
).await {
|
|
||||||
Ok(r) => results.push(r),
|
Ok(r) => results.push(r),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
results.push(db::QueryResult {
|
results.push(db::QueryResult {
|
||||||
@@ -308,7 +308,9 @@ pub async fn execute_statements(
|
|||||||
}
|
}
|
||||||
return Err(format!(
|
return Err(format!(
|
||||||
"Statement {} failed: {}. Previous {} statement(s) may have been committed.",
|
"Statement {} failed: {}. Previous {} statement(s) may have been committed.",
|
||||||
i + 1, e, i
|
i + 1,
|
||||||
|
e,
|
||||||
|
i
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,25 +10,13 @@ pub struct RunningQueries {
|
|||||||
impl RunningQueries {
|
impl RunningQueries {
|
||||||
pub fn register(&self, execution_id: String) -> RegisteredQuery {
|
pub fn register(&self, execution_id: String) -> RegisteredQuery {
|
||||||
let token = CancellationToken::new();
|
let token = CancellationToken::new();
|
||||||
self.inner
|
self.inner.lock().expect("running query registry poisoned").insert(execution_id.clone(), token.clone());
|
||||||
.lock()
|
|
||||||
.expect("running query registry poisoned")
|
|
||||||
.insert(execution_id.clone(), token.clone());
|
|
||||||
|
|
||||||
RegisteredQuery {
|
RegisteredQuery { execution_id, token, running_queries: self.clone() }
|
||||||
execution_id,
|
|
||||||
token,
|
|
||||||
running_queries: self.clone(),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn cancel(&self, execution_id: &str) -> bool {
|
pub fn cancel(&self, execution_id: &str) -> bool {
|
||||||
let token = self
|
let token = self.inner.lock().expect("running query registry poisoned").get(execution_id).cloned();
|
||||||
.inner
|
|
||||||
.lock()
|
|
||||||
.expect("running query registry poisoned")
|
|
||||||
.get(execution_id)
|
|
||||||
.cloned();
|
|
||||||
|
|
||||||
if let Some(token) = token {
|
if let Some(token) = token {
|
||||||
token.cancel();
|
token.cancel();
|
||||||
@@ -40,17 +28,11 @@ impl RunningQueries {
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
pub fn has(&self, execution_id: &str) -> bool {
|
pub fn has(&self, execution_id: &str) -> bool {
|
||||||
self.inner
|
self.inner.lock().expect("running query registry poisoned").contains_key(execution_id)
|
||||||
.lock()
|
|
||||||
.expect("running query registry poisoned")
|
|
||||||
.contains_key(execution_id)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn remove(&self, execution_id: &str) {
|
fn remove(&self, execution_id: &str) {
|
||||||
self.inner
|
self.inner.lock().expect("running query registry poisoned").remove(execution_id);
|
||||||
.lock()
|
|
||||||
.expect("running query registry poisoned")
|
|
||||||
.remove(execution_id);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,7 @@
|
|||||||
use crate::connection::{AppState, PoolKind};
|
use crate::connection::{AppState, PoolKind};
|
||||||
use crate::db::redis_driver::{self, RedisScanResult, RedisValue};
|
use crate::db::redis_driver::{self, RedisScanResult, RedisValue};
|
||||||
|
|
||||||
pub async fn redis_list_databases_core(
|
pub async fn redis_list_databases_core(state: &AppState, connection_id: &str) -> Result<Vec<u32>, String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
) -> Result<Vec<u32>, String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
let pool = connections.get(connection_id).ok_or("Connection not found")?;
|
let pool = connections.get(connection_id).ok_or("Connection not found")?;
|
||||||
match pool {
|
match pool {
|
||||||
@@ -36,11 +33,7 @@ pub async fn redis_scan_keys_core(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn redis_get_value_core(
|
pub async fn redis_get_value_core(state: &AppState, connection_id: &str, key: &str) -> Result<RedisValue, String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
key: &str,
|
|
||||||
) -> Result<RedisValue, String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
let pool = connections.get(connection_id).ok_or("Connection not found")?;
|
let pool = connections.get(connection_id).ok_or("Connection not found")?;
|
||||||
match pool {
|
match pool {
|
||||||
@@ -70,11 +63,7 @@ pub async fn redis_set_string_core(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn redis_delete_key_core(
|
pub async fn redis_delete_key_core(state: &AppState, connection_id: &str, key: &str) -> Result<(), String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
key: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
let pool = connections.get(connection_id).ok_or("Connection not found")?;
|
let pool = connections.get(connection_id).ok_or("Connection not found")?;
|
||||||
match pool {
|
match pool {
|
||||||
@@ -100,12 +89,7 @@ pub async fn redis_hash_set_core(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn redis_hash_del_core(
|
pub async fn redis_hash_del_core(state: &AppState, connection_id: &str, key: &str, field: &str) -> Result<(), String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
key: &str,
|
|
||||||
field: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::Redis(con) => redis_driver::hash_del(&mut *con.lock().await, key, field).await,
|
PoolKind::Redis(con) => redis_driver::hash_del(&mut *con.lock().await, key, field).await,
|
||||||
@@ -113,12 +97,7 @@ pub async fn redis_hash_del_core(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn redis_list_push_core(
|
pub async fn redis_list_push_core(state: &AppState, connection_id: &str, key: &str, value: &str) -> Result<(), String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
key: &str,
|
|
||||||
value: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::Redis(con) => redis_driver::list_push(&mut *con.lock().await, key, value).await,
|
PoolKind::Redis(con) => redis_driver::list_push(&mut *con.lock().await, key, value).await,
|
||||||
@@ -139,12 +118,7 @@ pub async fn redis_list_remove_core(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn redis_set_add_core(
|
pub async fn redis_set_add_core(state: &AppState, connection_id: &str, key: &str, member: &str) -> Result<(), String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
key: &str,
|
|
||||||
member: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
match connections.get(connection_id).ok_or("Not found")? {
|
match connections.get(connection_id).ok_or("Not found")? {
|
||||||
PoolKind::Redis(con) => redis_driver::set_add(&mut *con.lock().await, key, member).await,
|
PoolKind::Redis(con) => redis_driver::set_add(&mut *con.lock().await, key, member).await,
|
||||||
|
|||||||
+165
-80
@@ -8,18 +8,16 @@ pub fn duckdb_query_tables(con: &duckdb::Connection) -> Result<Vec<db::TableInfo
|
|||||||
let mut stmt = con.prepare(
|
let mut stmt = con.prepare(
|
||||||
"SELECT table_name, table_type FROM information_schema.tables WHERE table_schema = 'main' ORDER BY table_name"
|
"SELECT table_name, table_type FROM information_schema.tables WHERE table_schema = 'main' ORDER BY table_name"
|
||||||
).map_err(|e| e.to_string())?;
|
).map_err(|e| e.to_string())?;
|
||||||
let rows = stmt.query_map([], |row| {
|
let rows = stmt
|
||||||
Ok(db::TableInfo {
|
.query_map([], |row| Ok(db::TableInfo { name: row.get::<_, String>(0)?, table_type: row.get::<_, String>(1)? }))
|
||||||
name: row.get::<_, String>(0)?,
|
.map_err(|e| e.to_string())?;
|
||||||
table_type: row.get::<_, String>(1)?,
|
|
||||||
})
|
|
||||||
}).map_err(|e| e.to_string())?;
|
|
||||||
Ok(rows.filter_map(|r| r.ok()).collect())
|
Ok(rows.filter_map(|r| r.ok()).collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn duckdb_query_columns(con: &duckdb::Connection, table: &str) -> Result<Vec<db::ColumnInfo>, String> {
|
pub fn duckdb_query_columns(con: &duckdb::Connection, table: &str) -> Result<Vec<db::ColumnInfo>, String> {
|
||||||
let mut pk_stmt = con.prepare(
|
let mut pk_stmt = con
|
||||||
"SELECT kcu.column_name
|
.prepare(
|
||||||
|
"SELECT kcu.column_name
|
||||||
FROM information_schema.table_constraints tc
|
FROM information_schema.table_constraints tc
|
||||||
JOIN information_schema.key_column_usage kcu
|
JOIN information_schema.key_column_usage kcu
|
||||||
ON tc.constraint_name = kcu.constraint_name
|
ON tc.constraint_name = kcu.constraint_name
|
||||||
@@ -28,67 +26,81 @@ pub fn duckdb_query_columns(con: &duckdb::Connection, table: &str) -> Result<Vec
|
|||||||
WHERE tc.constraint_type = 'PRIMARY KEY'
|
WHERE tc.constraint_type = 'PRIMARY KEY'
|
||||||
AND tc.table_schema = 'main'
|
AND tc.table_schema = 'main'
|
||||||
AND tc.table_name = ?
|
AND tc.table_name = ?
|
||||||
ORDER BY kcu.ordinal_position"
|
ORDER BY kcu.ordinal_position",
|
||||||
).map_err(|e| e.to_string())?;
|
)
|
||||||
let pk_rows = pk_stmt.query_map([table], |row| row.get::<_, String>(0))
|
|
||||||
.map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
|
let pk_rows = pk_stmt.query_map([table], |row| row.get::<_, String>(0)).map_err(|e| e.to_string())?;
|
||||||
let primary_keys: std::collections::HashSet<String> = pk_rows.filter_map(|r| r.ok()).collect();
|
let primary_keys: std::collections::HashSet<String> = pk_rows.filter_map(|r| r.ok()).collect();
|
||||||
|
|
||||||
let mut stmt = con.prepare(
|
let mut stmt = con
|
||||||
"SELECT column_name, data_type, is_nullable, column_default
|
.prepare(
|
||||||
|
"SELECT column_name, data_type, is_nullable, column_default
|
||||||
FROM information_schema.columns
|
FROM information_schema.columns
|
||||||
WHERE table_schema = 'main' AND table_name = ?
|
WHERE table_schema = 'main' AND table_name = ?
|
||||||
ORDER BY ordinal_position"
|
ORDER BY ordinal_position",
|
||||||
).map_err(|e| e.to_string())?;
|
)
|
||||||
let rows = stmt.query_map([table], |row| {
|
.map_err(|e| e.to_string())?;
|
||||||
let name = row.get::<_, String>(0)?;
|
let rows = stmt
|
||||||
Ok(db::ColumnInfo {
|
.query_map([table], |row| {
|
||||||
is_primary_key: primary_keys.contains(&name),
|
let name = row.get::<_, String>(0)?;
|
||||||
name,
|
Ok(db::ColumnInfo {
|
||||||
data_type: row.get::<_, String>(1)?,
|
is_primary_key: primary_keys.contains(&name),
|
||||||
is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES",
|
name,
|
||||||
column_default: row.get::<_, Option<String>>(3)?,
|
data_type: row.get::<_, String>(1)?,
|
||||||
extra: None, comment: None,
|
is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES",
|
||||||
numeric_precision: None,
|
column_default: row.get::<_, Option<String>>(3)?,
|
||||||
numeric_scale: None,
|
extra: None,
|
||||||
character_maximum_length: None,
|
comment: None,
|
||||||
|
numeric_precision: None,
|
||||||
|
numeric_scale: None,
|
||||||
|
character_maximum_length: None,
|
||||||
|
})
|
||||||
})
|
})
|
||||||
}).map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
Ok(rows.filter_map(|r| r.ok()).collect())
|
Ok(rows.filter_map(|r| r.ok()).collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_duckdb(connections: &HashMap<String, PoolKind>, key: &str) -> Option<Arc<std::sync::Mutex<duckdb::Connection>>> {
|
pub fn extract_duckdb(
|
||||||
|
connections: &HashMap<String, PoolKind>,
|
||||||
|
key: &str,
|
||||||
|
) -> Option<Arc<std::sync::Mutex<duckdb::Connection>>> {
|
||||||
match connections.get(key)? {
|
match connections.get(key)? {
|
||||||
PoolKind::DuckDb(con) => Some(con.clone()),
|
PoolKind::DuckDb(con) => Some(con.clone()),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_sqlserver(connections: &HashMap<String, PoolKind>, key: &str) -> Option<Arc<tokio::sync::Mutex<db::sqlserver::SqlServerClient>>> {
|
pub fn extract_sqlserver(
|
||||||
|
connections: &HashMap<String, PoolKind>,
|
||||||
|
key: &str,
|
||||||
|
) -> Option<Arc<tokio::sync::Mutex<db::sqlserver::SqlServerClient>>> {
|
||||||
match connections.get(key)? {
|
match connections.get(key)? {
|
||||||
PoolKind::SqlServer(client) => Some(client.clone()),
|
PoolKind::SqlServer(client) => Some(client.clone()),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_clickhouse(connections: &HashMap<String, PoolKind>, key: &str) -> Option<db::clickhouse_driver::ChClient> {
|
pub fn extract_clickhouse(
|
||||||
|
connections: &HashMap<String, PoolKind>,
|
||||||
|
key: &str,
|
||||||
|
) -> Option<db::clickhouse_driver::ChClient> {
|
||||||
match connections.get(key)? {
|
match connections.get(key)? {
|
||||||
PoolKind::ClickHouse(client) => Some(client.clone()),
|
PoolKind::ClickHouse(client) => Some(client.clone()),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_oracle(connections: &HashMap<String, PoolKind>, key: &str) -> Option<Arc<tokio::sync::Mutex<db::oracle_driver::OracleClient>>> {
|
pub fn extract_oracle(
|
||||||
|
connections: &HashMap<String, PoolKind>,
|
||||||
|
key: &str,
|
||||||
|
) -> Option<Arc<tokio::sync::Mutex<db::oracle_driver::OracleClient>>> {
|
||||||
match connections.get(key)? {
|
match connections.get(key)? {
|
||||||
PoolKind::Oracle(client) => Some(client.clone()),
|
PoolKind::Oracle(client) => Some(client.clone()),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_databases_core(
|
pub async fn list_databases_core(state: &AppState, connection_id: &str) -> Result<Vec<db::DatabaseInfo>, String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
) -> Result<Vec<db::DatabaseInfo>, String> {
|
|
||||||
{
|
{
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
if let Some(client) = extract_clickhouse(&connections, connection_id) {
|
if let Some(client) = extract_clickhouse(&connections, connection_id) {
|
||||||
@@ -119,11 +131,7 @@ pub async fn list_databases_core(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_schemas_core(
|
pub async fn list_schemas_core(state: &AppState, connection_id: &str, database: &str) -> Result<Vec<String>, String> {
|
||||||
state: &AppState,
|
|
||||||
connection_id: &str,
|
|
||||||
database: &str,
|
|
||||||
) -> Result<Vec<String>, String> {
|
|
||||||
let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?;
|
let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?;
|
||||||
|
|
||||||
{
|
{
|
||||||
@@ -351,7 +359,8 @@ pub async fn get_table_ddl_core(
|
|||||||
drop(connections);
|
drop(connections);
|
||||||
let tbl = table.replace('\'', "''");
|
let tbl = table.replace('\'', "''");
|
||||||
let con = con.lock().map_err(|e| e.to_string())?;
|
let con = con.lock().map_err(|e| e.to_string())?;
|
||||||
let mut stmt = con.prepare(&format!("SELECT sql FROM duckdb_tables() WHERE table_name = '{tbl}'"))
|
let mut stmt = con
|
||||||
|
.prepare(&format!("SELECT sql FROM duckdb_tables() WHERE table_name = '{tbl}'"))
|
||||||
.map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
let mut rows = stmt.query([]).map_err(|e| e.to_string())?;
|
let mut rows = stmt.query([]).map_err(|e| e.to_string())?;
|
||||||
if let Some(row) = rows.next().map_err(|e| e.to_string())? {
|
if let Some(row) = rows.next().map_err(|e| e.to_string())? {
|
||||||
@@ -361,8 +370,12 @@ pub async fn get_table_ddl_core(
|
|||||||
}
|
}
|
||||||
if let Some(client) = extract_clickhouse(&connections, &pool_key) {
|
if let Some(client) = extract_clickhouse(&connections, &pool_key) {
|
||||||
drop(connections);
|
drop(connections);
|
||||||
let result = db::clickhouse_driver::execute_query(&client, database, &format!("SHOW CREATE TABLE `{table}`")).await?;
|
let result =
|
||||||
return result.rows.first()
|
db::clickhouse_driver::execute_query(&client, database, &format!("SHOW CREATE TABLE `{table}`"))
|
||||||
|
.await?;
|
||||||
|
return result
|
||||||
|
.rows
|
||||||
|
.first()
|
||||||
.and_then(|r| r.first())
|
.and_then(|r| r.first())
|
||||||
.and_then(|v| v.as_str())
|
.and_then(|v| v.as_str())
|
||||||
.map(|s| s.to_string())
|
.map(|s| s.to_string())
|
||||||
@@ -394,8 +407,7 @@ pub async fn get_table_ddl_core(
|
|||||||
pub async fn mysql_ddl(pool: &sqlx::mysql::MySqlPool, table: &str) -> Result<String, String> {
|
pub async fn mysql_ddl(pool: &sqlx::mysql::MySqlPool, table: &str) -> Result<String, String> {
|
||||||
use sqlx::Row;
|
use sqlx::Row;
|
||||||
let sql = format!("SHOW CREATE TABLE `{}`", table.replace('`', "``"));
|
let sql = format!("SHOW CREATE TABLE `{}`", table.replace('`', "``"));
|
||||||
let row: sqlx::mysql::MySqlRow = sqlx::raw_sql(&sql)
|
let row: sqlx::mysql::MySqlRow = sqlx::raw_sql(&sql).fetch_one(pool).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_one(pool).await.map_err(|e| e.to_string())?;
|
|
||||||
row.try_get::<String, _>(1)
|
row.try_get::<String, _>(1)
|
||||||
.or_else(|_| row.try_get::<Vec<u8>, _>(1).map(|b| String::from_utf8_lossy(&b).to_string()))
|
.or_else(|_| row.try_get::<Vec<u8>, _>(1).map(|b| String::from_utf8_lossy(&b).to_string()))
|
||||||
.map_err(|e| e.to_string())
|
.map_err(|e| e.to_string())
|
||||||
@@ -405,7 +417,9 @@ pub async fn sqlite_ddl(pool: &sqlx::sqlite::SqlitePool, table: &str) -> Result<
|
|||||||
use sqlx::Row;
|
use sqlx::Row;
|
||||||
let row: sqlx::sqlite::SqliteRow = sqlx::query("SELECT sql FROM sqlite_master WHERE type='table' AND name=?")
|
let row: sqlx::sqlite::SqliteRow = sqlx::query("SELECT sql FROM sqlite_master WHERE type='table' AND name=?")
|
||||||
.bind(table)
|
.bind(table)
|
||||||
.fetch_one(pool).await.map_err(|e| e.to_string())?;
|
.fetch_one(pool)
|
||||||
|
.await
|
||||||
|
.map_err(|e| e.to_string())?;
|
||||||
row.try_get::<String, _>(0).map_err(|e| e.to_string())
|
row.try_get::<String, _>(0).map_err(|e| e.to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -415,31 +429,56 @@ pub async fn pg_ddl(pool: &sqlx::postgres::PgPool, schema: &str, table: &str) ->
|
|||||||
let fkeys = db::postgres::list_foreign_keys(pool, schema, table).await?;
|
let fkeys = db::postgres::list_foreign_keys(pool, schema, table).await?;
|
||||||
|
|
||||||
let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n");
|
let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n");
|
||||||
let col_lines: Vec<String> = columns.iter().map(|c| {
|
let col_lines: Vec<String> = columns
|
||||||
let mut line = format!(" \"{}\" {}", c.name, c.data_type);
|
.iter()
|
||||||
if !c.is_nullable { line.push_str(" NOT NULL"); }
|
.map(|c| {
|
||||||
if let Some(ref def) = c.column_default { line.push_str(&format!(" DEFAULT {def}")); }
|
let mut line = format!(" \"{}\" {}", c.name, c.data_type);
|
||||||
line
|
if !c.is_nullable {
|
||||||
}).collect();
|
line.push_str(" NOT NULL");
|
||||||
|
}
|
||||||
|
if let Some(ref def) = c.column_default {
|
||||||
|
line.push_str(&format!(" DEFAULT {def}"));
|
||||||
|
}
|
||||||
|
line
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
ddl.push_str(&col_lines.join(",\n"));
|
ddl.push_str(&col_lines.join(",\n"));
|
||||||
|
|
||||||
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
|
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
|
||||||
if !pks.is_empty() {
|
if !pks.is_empty() {
|
||||||
ddl.push_str(&format!(",\n PRIMARY KEY ({})", pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", ")));
|
ddl.push_str(&format!(
|
||||||
|
",\n PRIMARY KEY ({})",
|
||||||
|
pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", ")
|
||||||
|
));
|
||||||
}
|
}
|
||||||
for fk in &fkeys {
|
for fk in &fkeys {
|
||||||
ddl.push_str(&format!(",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")", fk.name, fk.column, fk.ref_table, fk.ref_column));
|
ddl.push_str(&format!(
|
||||||
|
",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")",
|
||||||
|
fk.name, fk.column, fk.ref_table, fk.ref_column
|
||||||
|
));
|
||||||
}
|
}
|
||||||
ddl.push_str("\n);\n");
|
ddl.push_str("\n);\n");
|
||||||
|
|
||||||
for idx in &indexes {
|
for idx in &indexes {
|
||||||
if idx.is_primary { continue; }
|
if idx.is_primary {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
let unique = if idx.is_unique { "UNIQUE " } else { "" };
|
let unique = if idx.is_unique { "UNIQUE " } else { "" };
|
||||||
let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", ");
|
let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", ");
|
||||||
let using = idx.index_type.as_deref().map(|t| format!(" USING {t}")).unwrap_or_default();
|
let using = idx.index_type.as_deref().map(|t| format!(" USING {t}")).unwrap_or_default();
|
||||||
let include = idx.included_columns.as_deref().filter(|c| !c.is_empty()).map(|cols| format!(" INCLUDE ({})", cols.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", "))).unwrap_or_default();
|
let include = idx
|
||||||
|
.included_columns
|
||||||
|
.as_deref()
|
||||||
|
.filter(|c| !c.is_empty())
|
||||||
|
.map(|cols| {
|
||||||
|
format!(" INCLUDE ({})", cols.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", "))
|
||||||
|
})
|
||||||
|
.unwrap_or_default();
|
||||||
let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default();
|
let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default();
|
||||||
ddl.push_str(&format!("\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\"{using} ({cols}){include}{filter};", idx.name));
|
ddl.push_str(&format!(
|
||||||
|
"\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\"{using} ({cols}){include}{filter};",
|
||||||
|
idx.name
|
||||||
|
));
|
||||||
if let Some(ref c) = idx.comment {
|
if let Some(ref c) = idx.comment {
|
||||||
ddl.push_str(&format!("\nCOMMENT ON INDEX \"{schema}\".\"{}\" IS '{}';", idx.name, c.replace('\'', "''")));
|
ddl.push_str(&format!("\nCOMMENT ON INDEX \"{schema}\".\"{}\" IS '{}';", idx.name, c.replace('\'', "''")));
|
||||||
}
|
}
|
||||||
@@ -447,66 +486,112 @@ pub async fn pg_ddl(pool: &sqlx::postgres::PgPool, schema: &str, table: &str) ->
|
|||||||
Ok(ddl)
|
Ok(ddl)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn build_sqlserver_ddl(client: &mut db::sqlserver::SqlServerClient, schema: &str, table: &str) -> Result<String, String> {
|
pub async fn build_sqlserver_ddl(
|
||||||
|
client: &mut db::sqlserver::SqlServerClient,
|
||||||
|
schema: &str,
|
||||||
|
table: &str,
|
||||||
|
) -> Result<String, String> {
|
||||||
let columns = db::sqlserver::get_columns(client, schema, table).await?;
|
let columns = db::sqlserver::get_columns(client, schema, table).await?;
|
||||||
let indexes = db::sqlserver::list_indexes(client, schema, table).await?;
|
let indexes = db::sqlserver::list_indexes(client, schema, table).await?;
|
||||||
let fkeys = db::sqlserver::list_foreign_keys(client, schema, table).await?;
|
let fkeys = db::sqlserver::list_foreign_keys(client, schema, table).await?;
|
||||||
|
|
||||||
let mut ddl = format!("CREATE TABLE [{schema}].[{table}] (\n");
|
let mut ddl = format!("CREATE TABLE [{schema}].[{table}] (\n");
|
||||||
let col_lines: Vec<String> = columns.iter().map(|c| {
|
let col_lines: Vec<String> = columns
|
||||||
let mut line = format!(" [{}] {}", c.name, c.data_type);
|
.iter()
|
||||||
if !c.is_nullable { line.push_str(" NOT NULL"); }
|
.map(|c| {
|
||||||
if let Some(ref def) = c.column_default { line.push_str(&format!(" DEFAULT {def}")); }
|
let mut line = format!(" [{}] {}", c.name, c.data_type);
|
||||||
line
|
if !c.is_nullable {
|
||||||
}).collect();
|
line.push_str(" NOT NULL");
|
||||||
|
}
|
||||||
|
if let Some(ref def) = c.column_default {
|
||||||
|
line.push_str(&format!(" DEFAULT {def}"));
|
||||||
|
}
|
||||||
|
line
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
ddl.push_str(&col_lines.join(",\n"));
|
ddl.push_str(&col_lines.join(",\n"));
|
||||||
|
|
||||||
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
|
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
|
||||||
if !pks.is_empty() {
|
if !pks.is_empty() {
|
||||||
ddl.push_str(&format!(",\n PRIMARY KEY ({})", pks.iter().map(|k| format!("[{k}]")).collect::<Vec<_>>().join(", ")));
|
ddl.push_str(&format!(
|
||||||
|
",\n PRIMARY KEY ({})",
|
||||||
|
pks.iter().map(|k| format!("[{k}]")).collect::<Vec<_>>().join(", ")
|
||||||
|
));
|
||||||
}
|
}
|
||||||
for fk in &fkeys {
|
for fk in &fkeys {
|
||||||
ddl.push_str(&format!(",\n CONSTRAINT [{}] FOREIGN KEY ([{}]) REFERENCES [{}]([{}])", fk.name, fk.column, fk.ref_table, fk.ref_column));
|
ddl.push_str(&format!(
|
||||||
|
",\n CONSTRAINT [{}] FOREIGN KEY ([{}]) REFERENCES [{}]([{}])",
|
||||||
|
fk.name, fk.column, fk.ref_table, fk.ref_column
|
||||||
|
));
|
||||||
}
|
}
|
||||||
ddl.push_str("\n);\n");
|
ddl.push_str("\n);\n");
|
||||||
|
|
||||||
for idx in &indexes {
|
for idx in &indexes {
|
||||||
if idx.is_primary { continue; }
|
if idx.is_primary {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
let unique = if idx.is_unique { "UNIQUE " } else { "" };
|
let unique = if idx.is_unique { "UNIQUE " } else { "" };
|
||||||
let idx_type = idx.index_type.as_deref().map(|t| format!("{t} ")).unwrap_or_default();
|
let idx_type = idx.index_type.as_deref().map(|t| format!("{t} ")).unwrap_or_default();
|
||||||
let cols = idx.columns.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", ");
|
let cols = idx.columns.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", ");
|
||||||
let include = idx.included_columns.as_deref().filter(|c| !c.is_empty()).map(|cols| format!(" INCLUDE ({})", cols.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", "))).unwrap_or_default();
|
let include = idx
|
||||||
|
.included_columns
|
||||||
|
.as_deref()
|
||||||
|
.filter(|c| !c.is_empty())
|
||||||
|
.map(|cols| format!(" INCLUDE ({})", cols.iter().map(|c| format!("[{c}]")).collect::<Vec<_>>().join(", ")))
|
||||||
|
.unwrap_or_default();
|
||||||
let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default();
|
let filter = idx.filter.as_deref().map(|f| format!(" WHERE {f}")).unwrap_or_default();
|
||||||
ddl.push_str(&format!("\nCREATE {unique}{idx_type}INDEX [{}] ON [{schema}].[{table}] ({cols}){include}{filter};", idx.name));
|
ddl.push_str(&format!(
|
||||||
|
"\nCREATE {unique}{idx_type}INDEX [{}] ON [{schema}].[{table}] ({cols}){include}{filter};",
|
||||||
|
idx.name
|
||||||
|
));
|
||||||
}
|
}
|
||||||
Ok(ddl)
|
Ok(ddl)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn build_oracle_ddl(client: &db::oracle_driver::OracleClient, schema: &str, table: &str) -> Result<String, String> {
|
pub async fn build_oracle_ddl(
|
||||||
|
client: &db::oracle_driver::OracleClient,
|
||||||
|
schema: &str,
|
||||||
|
table: &str,
|
||||||
|
) -> Result<String, String> {
|
||||||
let columns = db::oracle_driver::get_columns(client, schema, table).await?;
|
let columns = db::oracle_driver::get_columns(client, schema, table).await?;
|
||||||
let indexes = db::oracle_driver::list_indexes(client, schema, table).await?;
|
let indexes = db::oracle_driver::list_indexes(client, schema, table).await?;
|
||||||
let fkeys = db::oracle_driver::list_foreign_keys(client, schema, table).await?;
|
let fkeys = db::oracle_driver::list_foreign_keys(client, schema, table).await?;
|
||||||
|
|
||||||
let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n");
|
let mut ddl = format!("CREATE TABLE \"{schema}\".\"{table}\" (\n");
|
||||||
let col_lines: Vec<String> = columns.iter().map(|c| {
|
let col_lines: Vec<String> = columns
|
||||||
let mut line = format!(" \"{}\" {}", c.name, c.data_type);
|
.iter()
|
||||||
if !c.is_nullable { line.push_str(" NOT NULL"); }
|
.map(|c| {
|
||||||
if let Some(ref def) = c.column_default { line.push_str(&format!(" DEFAULT {def}")); }
|
let mut line = format!(" \"{}\" {}", c.name, c.data_type);
|
||||||
line
|
if !c.is_nullable {
|
||||||
}).collect();
|
line.push_str(" NOT NULL");
|
||||||
|
}
|
||||||
|
if let Some(ref def) = c.column_default {
|
||||||
|
line.push_str(&format!(" DEFAULT {def}"));
|
||||||
|
}
|
||||||
|
line
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
ddl.push_str(&col_lines.join(",\n"));
|
ddl.push_str(&col_lines.join(",\n"));
|
||||||
|
|
||||||
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
|
let pks: Vec<&str> = columns.iter().filter(|c| c.is_primary_key).map(|c| c.name.as_str()).collect();
|
||||||
if !pks.is_empty() {
|
if !pks.is_empty() {
|
||||||
ddl.push_str(&format!(",\n PRIMARY KEY ({})", pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", ")));
|
ddl.push_str(&format!(
|
||||||
|
",\n PRIMARY KEY ({})",
|
||||||
|
pks.iter().map(|k| format!("\"{k}\"")).collect::<Vec<_>>().join(", ")
|
||||||
|
));
|
||||||
}
|
}
|
||||||
for fk in &fkeys {
|
for fk in &fkeys {
|
||||||
ddl.push_str(&format!(",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")", fk.name, fk.column, fk.ref_table, fk.ref_column));
|
ddl.push_str(&format!(
|
||||||
|
",\n CONSTRAINT \"{}\" FOREIGN KEY (\"{}\") REFERENCES \"{}\"(\"{}\")",
|
||||||
|
fk.name, fk.column, fk.ref_table, fk.ref_column
|
||||||
|
));
|
||||||
}
|
}
|
||||||
ddl.push_str("\n);\n");
|
ddl.push_str("\n);\n");
|
||||||
|
|
||||||
for idx in &indexes {
|
for idx in &indexes {
|
||||||
if idx.is_primary { continue; }
|
if idx.is_primary {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
let unique = if idx.is_unique { "UNIQUE " } else { "" };
|
let unique = if idx.is_unique { "UNIQUE " } else { "" };
|
||||||
let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", ");
|
let cols = idx.columns.iter().map(|c| format!("\"{c}\"")).collect::<Vec<_>>().join(", ");
|
||||||
ddl.push_str(&format!("\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\" ({cols});", idx.name));
|
ddl.push_str(&format!("\nCREATE {unique}INDEX \"{}\" ON \"{schema}\".\"{table}\" ({cols});", idx.name));
|
||||||
|
|||||||
@@ -147,17 +147,11 @@ impl SqlStatementSplitter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
match ch {
|
match ch {
|
||||||
'\'' if !self.in_double_quote
|
'\'' if !self.in_double_quote && !self.in_backtick && self.previous != Some('\\') => {
|
||||||
&& !self.in_backtick
|
|
||||||
&& self.previous != Some('\\') =>
|
|
||||||
{
|
|
||||||
self.in_single_quote = !self.in_single_quote;
|
self.in_single_quote = !self.in_single_quote;
|
||||||
self.buffer.push(ch);
|
self.buffer.push(ch);
|
||||||
}
|
}
|
||||||
'"' if !self.in_single_quote
|
'"' if !self.in_single_quote && !self.in_backtick && self.previous != Some('\\') => {
|
||||||
&& !self.in_backtick
|
|
||||||
&& self.previous != Some('\\') =>
|
|
||||||
{
|
|
||||||
self.in_double_quote = !self.in_double_quote;
|
self.in_double_quote = !self.in_double_quote;
|
||||||
self.buffer.push(ch);
|
self.buffer.push(ch);
|
||||||
}
|
}
|
||||||
@@ -351,10 +345,7 @@ mod tests {
|
|||||||
let mut splitter = SqlStatementSplitter::default();
|
let mut splitter = SqlStatementSplitter::default();
|
||||||
|
|
||||||
assert_eq!(splitter.push_chunk("SELECT 1; -"), vec!["SELECT 1"]);
|
assert_eq!(splitter.push_chunk("SELECT 1; -"), vec!["SELECT 1"]);
|
||||||
assert_eq!(
|
assert_eq!(splitter.push_chunk("- comment ; ignored\nSELECT 2;"), vec!["-- comment ; ignored\nSELECT 2"]);
|
||||||
splitter.push_chunk("- comment ; ignored\nSELECT 2;"),
|
|
||||||
vec!["-- comment ; ignored\nSELECT 2"]
|
|
||||||
);
|
|
||||||
assert_eq!(splitter.finish(), Vec::<String>::new());
|
assert_eq!(splitter.finish(), Vec::<String>::new());
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -363,10 +354,7 @@ mod tests {
|
|||||||
let mut splitter = SqlStatementSplitter::default();
|
let mut splitter = SqlStatementSplitter::default();
|
||||||
|
|
||||||
assert_eq!(splitter.push_chunk("SELECT 1; /"), vec!["SELECT 1"]);
|
assert_eq!(splitter.push_chunk("SELECT 1; /"), vec!["SELECT 1"]);
|
||||||
assert_eq!(
|
assert_eq!(splitter.push_chunk("* comment ; ignored */\nSELECT 2;"), vec!["/* comment ; ignored */\nSELECT 2"]);
|
||||||
splitter.push_chunk("* comment ; ignored */\nSELECT 2;"),
|
|
||||||
vec!["/* comment ; ignored */\nSELECT 2"]
|
|
||||||
);
|
|
||||||
assert_eq!(splitter.finish(), Vec::<String>::new());
|
assert_eq!(splitter.finish(), Vec::<String>::new());
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -402,14 +390,8 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn keeps_mysql_executable_comments_as_statements() {
|
fn keeps_mysql_executable_comments_as_statements() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
split_sql_script(
|
split_sql_script("/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */;\nSELECT 1;",).unwrap(),
|
||||||
"/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */;\nSELECT 1;",
|
vec!["/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */", "SELECT 1",]
|
||||||
)
|
|
||||||
.unwrap(),
|
|
||||||
vec![
|
|
||||||
"/*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */",
|
|
||||||
"SELECT 1",
|
|
||||||
]
|
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+63
-149
@@ -58,20 +58,12 @@ const SCHEMA_STATEMENTS: &[&str] = &[
|
|||||||
impl Storage {
|
impl Storage {
|
||||||
pub async fn open(db_path: &Path) -> Result<Self, String> {
|
pub async fn open(db_path: &Path) -> Result<Self, String> {
|
||||||
let url = format!("sqlite:{}?mode=rwc", db_path.display());
|
let url = format!("sqlite:{}?mode=rwc", db_path.display());
|
||||||
let options = SqliteConnectOptions::from_str(&url)
|
let options = SqliteConnectOptions::from_str(&url).map_err(|e| e.to_string())?.create_if_missing(true);
|
||||||
.map_err(|e| e.to_string())?
|
let pool =
|
||||||
.create_if_missing(true);
|
SqlitePoolOptions::new().max_connections(5).connect_with(options).await.map_err(|e| e.to_string())?;
|
||||||
let pool = SqlitePoolOptions::new()
|
|
||||||
.max_connections(5)
|
|
||||||
.connect_with(options)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
for statement in SCHEMA_STATEMENTS {
|
for statement in SCHEMA_STATEMENTS {
|
||||||
sqlx::query(statement)
|
sqlx::query(statement).execute(&pool).await.map_err(|e| e.to_string())?;
|
||||||
.execute(&pool)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(Self { db: pool })
|
Ok(Self { db: pool })
|
||||||
@@ -125,11 +117,7 @@ impl Storage {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_history_entries(
|
pub async fn load_history_entries(&self, limit: usize, offset: usize) -> Result<Vec<HistoryEntry>, String> {
|
||||||
&self,
|
|
||||||
limit: usize,
|
|
||||||
offset: usize,
|
|
||||||
) -> Result<Vec<HistoryEntry>, String> {
|
|
||||||
let rows: Vec<HistoryRow> = sqlx::query_as(
|
let rows: Vec<HistoryRow> = sqlx::query_as(
|
||||||
"SELECT id, connection_name, database, sql_text, executed_at, \
|
"SELECT id, connection_name, database, sql_text, executed_at, \
|
||||||
execution_time_ms, success, error \
|
execution_time_ms, success, error \
|
||||||
@@ -157,19 +145,12 @@ impl Storage {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn clear_history(&self) -> Result<(), String> {
|
pub async fn clear_history(&self) -> Result<(), String> {
|
||||||
sqlx::query("DELETE FROM history")
|
sqlx::query("DELETE FROM history").execute(&self.db).await.map_err(|e| e.to_string())?;
|
||||||
.execute(&self.db)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn delete_history_entry(&self, id: &str) -> Result<(), String> {
|
pub async fn delete_history_entry(&self, id: &str) -> Result<(), String> {
|
||||||
sqlx::query("DELETE FROM history WHERE id = ?")
|
sqlx::query("DELETE FROM history WHERE id = ?").bind(id).execute(&self.db).await.map_err(|e| e.to_string())?;
|
||||||
.bind(id)
|
|
||||||
.execute(&self.db)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -190,15 +171,12 @@ impl Storage {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_ai_config(&self) -> Result<Option<AiConfig>, String> {
|
pub async fn load_ai_config(&self) -> Result<Option<AiConfig>, String> {
|
||||||
let row: Option<(String,)> =
|
let row: Option<(String,)> = sqlx::query_as("SELECT config_json FROM ai_config WHERE id = 1")
|
||||||
sqlx::query_as("SELECT config_json FROM ai_config WHERE id = 1")
|
.fetch_optional(&self.db)
|
||||||
.fetch_optional(&self.db)
|
.await
|
||||||
.await
|
.map_err(|e| e.to_string())?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
match row {
|
match row {
|
||||||
Some((json,)) => serde_json::from_str(&json)
|
Some((json,)) => serde_json::from_str(&json).map(Some).map_err(|e| e.to_string()),
|
||||||
.map(Some)
|
|
||||||
.map_err(|e| e.to_string()),
|
|
||||||
None => Ok(None),
|
None => Ok(None),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -221,8 +199,7 @@ struct AiConversationRow {
|
|||||||
|
|
||||||
impl Storage {
|
impl Storage {
|
||||||
pub async fn save_ai_conversation(&self, conv: &AiConversation) -> Result<(), String> {
|
pub async fn save_ai_conversation(&self, conv: &AiConversation) -> Result<(), String> {
|
||||||
let messages_json =
|
let messages_json = serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
|
||||||
serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
|
|
||||||
sqlx::query(
|
sqlx::query(
|
||||||
"INSERT OR REPLACE INTO ai_conversations \
|
"INSERT OR REPLACE INTO ai_conversations \
|
||||||
(id, title, connection_name, database, messages_json, created_at, updated_at) \
|
(id, title, connection_name, database, messages_json, created_at, updated_at) \
|
||||||
@@ -263,8 +240,7 @@ impl Storage {
|
|||||||
|
|
||||||
rows.into_iter()
|
rows.into_iter()
|
||||||
.map(|r| {
|
.map(|r| {
|
||||||
let messages: Vec<AiChatMessage> =
|
let messages: Vec<AiChatMessage> = serde_json::from_str(&r.messages_json).map_err(|e| e.to_string())?;
|
||||||
serde_json::from_str(&r.messages_json).map_err(|e| e.to_string())?;
|
|
||||||
Ok(AiConversation {
|
Ok(AiConversation {
|
||||||
id: r.id,
|
id: r.id,
|
||||||
title: r.title,
|
title: r.title,
|
||||||
@@ -296,10 +272,7 @@ impl Storage {
|
|||||||
pub async fn save_connections(&self, configs: &[ConnectionConfig]) -> Result<(), String> {
|
pub async fn save_connections(&self, configs: &[ConnectionConfig]) -> Result<(), String> {
|
||||||
let mut tx = self.db.begin().await.map_err(|e| e.to_string())?;
|
let mut tx = self.db.begin().await.map_err(|e| e.to_string())?;
|
||||||
|
|
||||||
sqlx::query("DELETE FROM connections")
|
sqlx::query("DELETE FROM connections").execute(&mut *tx).await.map_err(|e| e.to_string())?;
|
||||||
.execute(&mut *tx)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
for config in configs {
|
for config in configs {
|
||||||
// Store config without secrets
|
// Store config without secrets
|
||||||
@@ -319,49 +292,31 @@ impl Storage {
|
|||||||
|
|
||||||
// Store secrets
|
// Store secrets
|
||||||
persist_secret_in_tx(&mut tx, &config.id, "password", &config.password).await?;
|
persist_secret_in_tx(&mut tx, &config.id, "password", &config.password).await?;
|
||||||
persist_secret_in_tx(&mut tx, &config.id, "ssh_password", &config.ssh_password)
|
persist_secret_in_tx(&mut tx, &config.id, "ssh_password", &config.ssh_password).await?;
|
||||||
.await?;
|
persist_secret_in_tx(&mut tx, &config.id, "ssh_key_passphrase", &config.ssh_key_passphrase).await?;
|
||||||
persist_secret_in_tx(
|
|
||||||
&mut tx,
|
|
||||||
&config.id,
|
|
||||||
"ssh_key_passphrase",
|
|
||||||
&config.ssh_key_passphrase,
|
|
||||||
)
|
|
||||||
.await?;
|
|
||||||
if let Some(cs) = &config.connection_string {
|
if let Some(cs) = &config.connection_string {
|
||||||
persist_secret_in_tx(&mut tx, &config.id, "connection_string", cs).await?;
|
persist_secret_in_tx(&mut tx, &config.id, "connection_string", cs).await?;
|
||||||
} else {
|
} else {
|
||||||
sqlx::query(
|
sqlx::query("DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?")
|
||||||
"DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?",
|
.bind(&config.id)
|
||||||
)
|
.bind("connection_string")
|
||||||
.bind(&config.id)
|
.execute(&mut *tx)
|
||||||
.bind("connection_string")
|
.await
|
||||||
.execute(&mut *tx)
|
.map_err(|e| e.to_string())?;
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Remove secrets for connections that no longer exist
|
// Remove secrets for connections that no longer exist
|
||||||
if configs.is_empty() {
|
if configs.is_empty() {
|
||||||
sqlx::query("DELETE FROM connection_secrets")
|
sqlx::query("DELETE FROM connection_secrets").execute(&mut *tx).await.map_err(|e| e.to_string())?;
|
||||||
.execute(&mut *tx)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
} else {
|
} else {
|
||||||
let placeholders: Vec<&str> = configs.iter().map(|_| "?").collect();
|
let placeholders: Vec<&str> = configs.iter().map(|_| "?").collect();
|
||||||
let sql = format!(
|
let sql = format!("DELETE FROM connection_secrets WHERE connection_id NOT IN ({})", placeholders.join(","));
|
||||||
"DELETE FROM connection_secrets WHERE connection_id NOT IN ({})",
|
|
||||||
placeholders.join(",")
|
|
||||||
);
|
|
||||||
let mut query = sqlx::query(&sql);
|
let mut query = sqlx::query(&sql);
|
||||||
for config in configs {
|
for config in configs {
|
||||||
query = query.bind(&config.id);
|
query = query.bind(&config.id);
|
||||||
}
|
}
|
||||||
query
|
query.execute(&mut *tx).await.map_err(|e| e.to_string())?;
|
||||||
.execute(&mut *tx)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
tx.commit().await.map_err(|e| e.to_string())?;
|
tx.commit().await.map_err(|e| e.to_string())?;
|
||||||
@@ -369,28 +324,17 @@ impl Storage {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_connections(&self) -> Result<Vec<ConnectionConfig>, String> {
|
pub async fn load_connections(&self) -> Result<Vec<ConnectionConfig>, String> {
|
||||||
let rows: Vec<(String, String)> =
|
let rows: Vec<(String, String)> = sqlx::query_as("SELECT id, config_json FROM connections")
|
||||||
sqlx::query_as("SELECT id, config_json FROM connections")
|
.fetch_all(&self.db)
|
||||||
.fetch_all(&self.db)
|
.await
|
||||||
.await
|
.map_err(|e| e.to_string())?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
let mut configs = Vec::new();
|
let mut configs = Vec::new();
|
||||||
for (id, json) in rows {
|
for (id, json) in rows {
|
||||||
let mut config: ConnectionConfig =
|
let mut config: ConnectionConfig = serde_json::from_str(&json).map_err(|e| e.to_string())?;
|
||||||
serde_json::from_str(&json).map_err(|e| e.to_string())?;
|
config.password = self.get_secret(&id, "password").await?.unwrap_or_default();
|
||||||
config.password = self
|
config.ssh_password = self.get_secret(&id, "ssh_password").await?.unwrap_or_default();
|
||||||
.get_secret(&id, "password")
|
config.ssh_key_passphrase = self.get_secret(&id, "ssh_key_passphrase").await?.unwrap_or_default();
|
||||||
.await?
|
|
||||||
.unwrap_or_default();
|
|
||||||
config.ssh_password = self
|
|
||||||
.get_secret(&id, "ssh_password")
|
|
||||||
.await?
|
|
||||||
.unwrap_or_default();
|
|
||||||
config.ssh_key_passphrase = self
|
|
||||||
.get_secret(&id, "ssh_key_passphrase")
|
|
||||||
.await?
|
|
||||||
.unwrap_or_default();
|
|
||||||
config.connection_string = self.get_secret(&id, "connection_string").await?;
|
config.connection_string = self.get_secret(&id, "connection_string").await?;
|
||||||
configs.push(config);
|
configs.push(config);
|
||||||
}
|
}
|
||||||
@@ -403,28 +347,18 @@ impl Storage {
|
|||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
impl Storage {
|
impl Storage {
|
||||||
pub async fn get_secret(
|
pub async fn get_secret(&self, connection_id: &str, key: &str) -> Result<Option<String>, String> {
|
||||||
&self,
|
let row: Option<(String,)> =
|
||||||
connection_id: &str,
|
sqlx::query_as("SELECT secret FROM connection_secrets WHERE connection_id = ? AND key = ?")
|
||||||
key: &str,
|
.bind(connection_id)
|
||||||
) -> Result<Option<String>, String> {
|
.bind(key)
|
||||||
let row: Option<(String,)> = sqlx::query_as(
|
.fetch_optional(&self.db)
|
||||||
"SELECT secret FROM connection_secrets WHERE connection_id = ? AND key = ?",
|
.await
|
||||||
)
|
.map_err(|e| e.to_string())?;
|
||||||
.bind(connection_id)
|
|
||||||
.bind(key)
|
|
||||||
.fetch_optional(&self.db)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
Ok(row.map(|(s,)| s))
|
Ok(row.map(|(s,)| s))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn set_secret(
|
pub async fn set_secret(&self, connection_id: &str, key: &str, secret: &str) -> Result<(), String> {
|
||||||
&self,
|
|
||||||
connection_id: &str,
|
|
||||||
key: &str,
|
|
||||||
secret: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
sqlx::query(
|
sqlx::query(
|
||||||
"INSERT OR REPLACE INTO connection_secrets (connection_id, key, secret) \
|
"INSERT OR REPLACE INTO connection_secrets (connection_id, key, secret) \
|
||||||
VALUES (?, ?, ?)",
|
VALUES (?, ?, ?)",
|
||||||
@@ -438,11 +372,7 @@ impl Storage {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn delete_secret(
|
pub async fn delete_secret(&self, connection_id: &str, key: &str) -> Result<(), String> {
|
||||||
&self,
|
|
||||||
connection_id: &str,
|
|
||||||
key: &str,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
sqlx::query("DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?")
|
sqlx::query("DELETE FROM connection_secrets WHERE connection_id = ? AND key = ?")
|
||||||
.bind(connection_id)
|
.bind(connection_id)
|
||||||
.bind(key)
|
.bind(key)
|
||||||
@@ -458,10 +388,7 @@ impl Storage {
|
|||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
impl Storage {
|
impl Storage {
|
||||||
pub async fn save_sidebar_layout(
|
pub async fn save_sidebar_layout(&self, layout: &serde_json::Value) -> Result<(), String> {
|
||||||
&self,
|
|
||||||
layout: &serde_json::Value,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let json = serde_json::to_string(layout).map_err(|e| e.to_string())?;
|
let json = serde_json::to_string(layout).map_err(|e| e.to_string())?;
|
||||||
sqlx::query("INSERT OR REPLACE INTO sidebar_layout (id, layout_json) VALUES (1, ?)")
|
sqlx::query("INSERT OR REPLACE INTO sidebar_layout (id, layout_json) VALUES (1, ?)")
|
||||||
.bind(&json)
|
.bind(&json)
|
||||||
@@ -472,15 +399,12 @@ impl Storage {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_sidebar_layout(&self) -> Result<Option<serde_json::Value>, String> {
|
pub async fn load_sidebar_layout(&self) -> Result<Option<serde_json::Value>, String> {
|
||||||
let row: Option<(String,)> =
|
let row: Option<(String,)> = sqlx::query_as("SELECT layout_json FROM sidebar_layout WHERE id = 1")
|
||||||
sqlx::query_as("SELECT layout_json FROM sidebar_layout WHERE id = 1")
|
.fetch_optional(&self.db)
|
||||||
.fetch_optional(&self.db)
|
.await
|
||||||
.await
|
.map_err(|e| e.to_string())?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
match row {
|
match row {
|
||||||
Some((json,)) => serde_json::from_str(&json)
|
Some((json,)) => serde_json::from_str(&json).map(Some).map_err(|e| e.to_string()),
|
||||||
.map(Some)
|
|
||||||
.map_err(|e| e.to_string()),
|
|
||||||
None => Ok(None),
|
None => Ok(None),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -509,8 +433,7 @@ impl Storage {
|
|||||||
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
||||||
let configs: Vec<ConnectionConfig> = serde_json::from_str(&json).unwrap_or_default();
|
let configs: Vec<ConnectionConfig> = serde_json::from_str(&json).unwrap_or_default();
|
||||||
for config in &configs {
|
for config in &configs {
|
||||||
let config_json =
|
let config_json = serde_json::to_string(config).map_err(|e| e.to_string())?;
|
||||||
serde_json::to_string(config).map_err(|e| e.to_string())?;
|
|
||||||
sqlx::query("INSERT OR IGNORE INTO connections (id, config_json) VALUES (?, ?)")
|
sqlx::query("INSERT OR IGNORE INTO connections (id, config_json) VALUES (?, ?)")
|
||||||
.bind(&config.id)
|
.bind(&config.id)
|
||||||
.bind(&config_json)
|
.bind(&config_json)
|
||||||
@@ -528,8 +451,7 @@ impl Storage {
|
|||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
||||||
let secrets: std::collections::HashMap<String, String> =
|
let secrets: std::collections::HashMap<String, String> = serde_json::from_str(&json).unwrap_or_default();
|
||||||
serde_json::from_str(&json).unwrap_or_default();
|
|
||||||
for (key, secret) in &secrets {
|
for (key, secret) in &secrets {
|
||||||
// key format: "connection:{id}:{field}"
|
// key format: "connection:{id}:{field}"
|
||||||
let parts: Vec<&str> = key.splitn(3, ':').collect();
|
let parts: Vec<&str> = key.splitn(3, ':').collect();
|
||||||
@@ -588,10 +510,7 @@ impl Storage {
|
|||||||
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
||||||
// Only migrate if the table is empty
|
// Only migrate if the table is empty
|
||||||
let count: (i64,) =
|
let count: (i64,) =
|
||||||
sqlx::query_as("SELECT COUNT(*) FROM ai_config")
|
sqlx::query_as("SELECT COUNT(*) FROM ai_config").fetch_one(&self.db).await.map_err(|e| e.to_string())?;
|
||||||
.fetch_one(&self.db)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
if count.0 == 0 {
|
if count.0 == 0 {
|
||||||
sqlx::query("INSERT OR IGNORE INTO ai_config (id, config_json) VALUES (1, ?)")
|
sqlx::query("INSERT OR IGNORE INTO ai_config (id, config_json) VALUES (1, ?)")
|
||||||
.bind(&json)
|
.bind(&json)
|
||||||
@@ -609,11 +528,9 @@ impl Storage {
|
|||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
||||||
let conversations: Vec<AiConversation> =
|
let conversations: Vec<AiConversation> = serde_json::from_str(&json).unwrap_or_default();
|
||||||
serde_json::from_str(&json).unwrap_or_default();
|
|
||||||
for conv in &conversations {
|
for conv in &conversations {
|
||||||
let messages_json =
|
let messages_json = serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
|
||||||
serde_json::to_string(&conv.messages).map_err(|e| e.to_string())?;
|
|
||||||
sqlx::query(
|
sqlx::query(
|
||||||
"INSERT OR IGNORE INTO ai_conversations \
|
"INSERT OR IGNORE INTO ai_conversations \
|
||||||
(id, title, connection_name, database, messages_json, \
|
(id, title, connection_name, database, messages_json, \
|
||||||
@@ -641,19 +558,16 @@ impl Storage {
|
|||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
let json = std::fs::read_to_string(&path).map_err(|e| e.to_string())?;
|
||||||
let count: (i64,) =
|
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM sidebar_layout")
|
||||||
sqlx::query_as("SELECT COUNT(*) FROM sidebar_layout")
|
.fetch_one(&self.db)
|
||||||
.fetch_one(&self.db)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
if count.0 == 0 {
|
|
||||||
sqlx::query(
|
|
||||||
"INSERT OR IGNORE INTO sidebar_layout (id, layout_json) VALUES (1, ?)",
|
|
||||||
)
|
|
||||||
.bind(&json)
|
|
||||||
.execute(&self.db)
|
|
||||||
.await
|
.await
|
||||||
.map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
|
if count.0 == 0 {
|
||||||
|
sqlx::query("INSERT OR IGNORE INTO sidebar_layout (id, layout_json) VALUES (1, ?)")
|
||||||
|
.bind(&json)
|
||||||
|
.execute(&self.db)
|
||||||
|
.await
|
||||||
|
.map_err(|e| e.to_string())?;
|
||||||
}
|
}
|
||||||
std::fs::rename(&path, data_dir.join("sidebar_layout.json.bak")).ok();
|
std::fs::rename(&path, data_dir.join("sidebar_layout.json.bak")).ok();
|
||||||
Ok(())
|
Ok(())
|
||||||
|
|||||||
@@ -4,9 +4,9 @@ use std::path::Path;
|
|||||||
use calamine::{open_workbook_auto, Data, Reader};
|
use calamine::{open_workbook_auto, Data, Reader};
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
|
use crate::connection::AppState;
|
||||||
use crate::models::connection::DatabaseType;
|
use crate::models::connection::DatabaseType;
|
||||||
use crate::transfer::{execute_on_pool, generate_insert, qualified_table};
|
use crate::transfer::{execute_on_pool, generate_insert, qualified_table};
|
||||||
use crate::connection::AppState;
|
|
||||||
|
|
||||||
pub const DEFAULT_PREVIEW_LIMIT: usize = 50;
|
pub const DEFAULT_PREVIEW_LIMIT: usize = 50;
|
||||||
pub const DEFAULT_BATCH_SIZE: usize = 500;
|
pub const DEFAULT_BATCH_SIZE: usize = 500;
|
||||||
@@ -142,15 +142,8 @@ pub fn csv_value(value: &str) -> serde_json::Value {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn parse_delimited_bytes(
|
pub fn parse_delimited_bytes(bytes: &[u8], delimiter: u8, preview_limit: usize) -> Result<ParsedImportFile, String> {
|
||||||
bytes: &[u8],
|
let mut reader = csv::ReaderBuilder::new().delimiter(delimiter).flexible(true).from_reader(bytes);
|
||||||
delimiter: u8,
|
|
||||||
preview_limit: usize,
|
|
||||||
) -> Result<ParsedImportFile, String> {
|
|
||||||
let mut reader = csv::ReaderBuilder::new()
|
|
||||||
.delimiter(delimiter)
|
|
||||||
.flexible(true)
|
|
||||||
.from_reader(bytes);
|
|
||||||
let columns = reader
|
let columns = reader
|
||||||
.headers()
|
.headers()
|
||||||
.map_err(|e| e.to_string())?
|
.map_err(|e| e.to_string())?
|
||||||
@@ -172,21 +165,12 @@ pub fn parse_delimited_bytes(
|
|||||||
}
|
}
|
||||||
let mut row = Vec::with_capacity(columns.len());
|
let mut row = Vec::with_capacity(columns.len());
|
||||||
for index in 0..columns.len() {
|
for index in 0..columns.len() {
|
||||||
row.push(
|
row.push(record.get(index).map(csv_value).unwrap_or(serde_json::Value::Null));
|
||||||
record
|
|
||||||
.get(index)
|
|
||||||
.map(csv_value)
|
|
||||||
.unwrap_or(serde_json::Value::Null),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
rows.push(row);
|
rows.push(row);
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(ParsedImportFile {
|
Ok(ParsedImportFile { columns, rows, total_rows })
|
||||||
columns,
|
|
||||||
rows,
|
|
||||||
total_rows,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn parse_csv_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImportFile, String> {
|
pub fn parse_csv_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImportFile, String> {
|
||||||
@@ -229,25 +213,15 @@ pub fn parse_json_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImpo
|
|||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
})
|
})
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
return Ok(ParsedImportFile {
|
return Ok(ParsedImportFile { columns, rows, total_rows: items.len() });
|
||||||
columns,
|
|
||||||
rows,
|
|
||||||
total_rows: items.len(),
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if items.iter().all(|item| item.is_array()) {
|
if items.iter().all(|item| item.is_array()) {
|
||||||
let max_cols = items
|
let max_cols = items.iter().filter_map(|item| item.as_array().map(|row| row.len())).max().unwrap_or(0);
|
||||||
.iter()
|
|
||||||
.filter_map(|item| item.as_array().map(|row| row.len()))
|
|
||||||
.max()
|
|
||||||
.unwrap_or(0);
|
|
||||||
if max_cols == 0 {
|
if max_cols == 0 {
|
||||||
return Err("Import file has no columns".to_string());
|
return Err("Import file has no columns".to_string());
|
||||||
}
|
}
|
||||||
let columns = (0..max_cols)
|
let columns = (0..max_cols).map(|index| format!("column_{}", index + 1)).collect::<Vec<_>>();
|
||||||
.map(|index| format!("column_{}", index + 1))
|
|
||||||
.collect::<Vec<_>>();
|
|
||||||
let rows = items
|
let rows = items
|
||||||
.iter()
|
.iter()
|
||||||
.take(preview_limit)
|
.take(preview_limit)
|
||||||
@@ -258,11 +232,7 @@ pub fn parse_json_bytes(bytes: &[u8], preview_limit: usize) -> Result<ParsedImpo
|
|||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
})
|
})
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
return Ok(ParsedImportFile {
|
return Ok(ParsedImportFile { columns, rows, total_rows: items.len() });
|
||||||
columns,
|
|
||||||
rows,
|
|
||||||
total_rows: items.len(),
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Err("JSON rows must all be objects or all be arrays".to_string())
|
Err("JSON rows must all be objects or all be arrays".to_string())
|
||||||
@@ -272,9 +242,9 @@ pub fn xlsx_cell_value(cell: &Data) -> serde_json::Value {
|
|||||||
match cell {
|
match cell {
|
||||||
Data::Empty => serde_json::Value::Null,
|
Data::Empty => serde_json::Value::Null,
|
||||||
Data::String(s) => csv_value(s),
|
Data::String(s) => csv_value(s),
|
||||||
Data::Float(n) => serde_json::Number::from_f64(*n)
|
Data::Float(n) => {
|
||||||
.map(serde_json::Value::Number)
|
serde_json::Number::from_f64(*n).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
|
||||||
.unwrap_or(serde_json::Value::Null),
|
}
|
||||||
Data::Int(n) => serde_json::Value::Number((*n).into()),
|
Data::Int(n) => serde_json::Value::Number((*n).into()),
|
||||||
Data::Bool(v) => serde_json::Value::Bool(*v),
|
Data::Bool(v) => serde_json::Value::Bool(*v),
|
||||||
Data::DateTime(v) => serde_json::Value::String(v.to_string()),
|
Data::DateTime(v) => serde_json::Value::String(v.to_string()),
|
||||||
@@ -300,18 +270,10 @@ pub fn xlsx_cell_label(cell: &Data) -> String {
|
|||||||
|
|
||||||
pub fn parse_xlsx_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> {
|
pub fn parse_xlsx_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> {
|
||||||
let mut workbook = open_workbook_auto(path).map_err(|e| e.to_string())?;
|
let mut workbook = open_workbook_auto(path).map_err(|e| e.to_string())?;
|
||||||
let sheet_name = workbook
|
let sheet_name = workbook.sheet_names().first().cloned().ok_or_else(|| "Workbook has no sheets".to_string())?;
|
||||||
.sheet_names()
|
let range = workbook.worksheet_range(&sheet_name).map_err(|e| e.to_string())?;
|
||||||
.first()
|
|
||||||
.cloned()
|
|
||||||
.ok_or_else(|| "Workbook has no sheets".to_string())?;
|
|
||||||
let range = workbook
|
|
||||||
.worksheet_range(&sheet_name)
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
let mut rows_iter = range.rows();
|
let mut rows_iter = range.rows();
|
||||||
let header = rows_iter
|
let header = rows_iter.next().ok_or_else(|| "Import file has no rows".to_string())?;
|
||||||
.next()
|
|
||||||
.ok_or_else(|| "Import file has no rows".to_string())?;
|
|
||||||
let columns = header
|
let columns = header
|
||||||
.iter()
|
.iter()
|
||||||
.enumerate()
|
.enumerate()
|
||||||
@@ -330,21 +292,12 @@ pub fn parse_xlsx_file(path: &str, preview_limit: usize) -> Result<ParsedImportF
|
|||||||
}
|
}
|
||||||
let mut row = Vec::with_capacity(columns.len());
|
let mut row = Vec::with_capacity(columns.len());
|
||||||
for index in 0..columns.len() {
|
for index in 0..columns.len() {
|
||||||
row.push(
|
row.push(source_row.get(index).map(xlsx_cell_value).unwrap_or(serde_json::Value::Null));
|
||||||
source_row
|
|
||||||
.get(index)
|
|
||||||
.map(xlsx_cell_value)
|
|
||||||
.unwrap_or(serde_json::Value::Null),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
rows.push(row);
|
rows.push(row);
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(ParsedImportFile {
|
Ok(ParsedImportFile { columns, rows, total_rows })
|
||||||
columns,
|
|
||||||
rows,
|
|
||||||
total_rows,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn parse_import_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> {
|
pub fn parse_import_file(path: &str, preview_limit: usize) -> Result<ParsedImportFile, String> {
|
||||||
@@ -384,10 +337,7 @@ pub fn mapping_indexes(
|
|||||||
return Err("Target column cannot be empty".to_string());
|
return Err("Target column cannot be empty".to_string());
|
||||||
}
|
}
|
||||||
if !target_seen.insert(mapping.target_column.clone()) {
|
if !target_seen.insert(mapping.target_column.clone()) {
|
||||||
return Err(format!(
|
return Err(format!("Target column mapped more than once: {}", mapping.target_column));
|
||||||
"Target column mapped more than once: {}",
|
|
||||||
mapping.target_column
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
mapped.push((source_index, mapping.target_column.clone()));
|
mapped.push((source_index, mapping.target_column.clone()));
|
||||||
}
|
}
|
||||||
@@ -403,10 +353,7 @@ pub fn build_import_insert_batches(
|
|||||||
batch_size: usize,
|
batch_size: usize,
|
||||||
) -> Result<Vec<ImportSqlBatch>, String> {
|
) -> Result<Vec<ImportSqlBatch>, String> {
|
||||||
let mapped = mapping_indexes(data, mappings)?;
|
let mapped = mapping_indexes(data, mappings)?;
|
||||||
let columns = mapped
|
let columns = mapped.iter().map(|(_, target)| target.clone()).collect::<Vec<_>>();
|
||||||
.iter()
|
|
||||||
.map(|(_, target)| target.clone())
|
|
||||||
.collect::<Vec<_>>();
|
|
||||||
let batch_size = batch_size.max(1);
|
let batch_size = batch_size.max(1);
|
||||||
let mut batches = Vec::new();
|
let mut batches = Vec::new();
|
||||||
|
|
||||||
@@ -416,20 +363,13 @@ pub fn build_import_insert_batches(
|
|||||||
.map(|row| {
|
.map(|row| {
|
||||||
mapped
|
mapped
|
||||||
.iter()
|
.iter()
|
||||||
.map(|(source_index, _)| {
|
.map(|(source_index, _)| row.get(*source_index).cloned().unwrap_or(serde_json::Value::Null))
|
||||||
row.get(*source_index)
|
|
||||||
.cloned()
|
|
||||||
.unwrap_or(serde_json::Value::Null)
|
|
||||||
})
|
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
})
|
})
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
let sql = generate_insert(&columns, &rows, table, schema, db_type);
|
let sql = generate_insert(&columns, &rows, table, schema, db_type);
|
||||||
if !sql.trim().is_empty() {
|
if !sql.trim().is_empty() {
|
||||||
batches.push(ImportSqlBatch {
|
batches.push(ImportSqlBatch { sql, row_count: chunk.len() });
|
||||||
sql,
|
|
||||||
row_count: chunk.len(),
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -448,11 +388,7 @@ pub fn preview_table_import_file_core(file_path: &str) -> Result<TableImportPrev
|
|||||||
let kind = import_file_kind(file_path)?;
|
let kind = import_file_kind(file_path)?;
|
||||||
let parsed = parse_import_file(file_path, DEFAULT_PREVIEW_LIMIT)?;
|
let parsed = parse_import_file(file_path, DEFAULT_PREVIEW_LIMIT)?;
|
||||||
let metadata = std::fs::metadata(file_path).map_err(|e| e.to_string())?;
|
let metadata = std::fs::metadata(file_path).map_err(|e| e.to_string())?;
|
||||||
let file_name = Path::new(file_path)
|
let file_name = Path::new(file_path).file_name().and_then(|name| name.to_str()).unwrap_or(file_path).to_string();
|
||||||
.file_name()
|
|
||||||
.and_then(|name| name.to_str())
|
|
||||||
.unwrap_or(file_path)
|
|
||||||
.to_string();
|
|
||||||
|
|
||||||
Ok(TableImportPreview {
|
Ok(TableImportPreview {
|
||||||
file_name,
|
file_name,
|
||||||
@@ -478,11 +414,7 @@ pub async fn import_table_file_core<F>(
|
|||||||
where
|
where
|
||||||
F: FnMut(TableImportProgress),
|
F: FnMut(TableImportProgress),
|
||||||
{
|
{
|
||||||
let batch_size = if request.batch_size == 0 {
|
let batch_size = if request.batch_size == 0 { DEFAULT_BATCH_SIZE } else { request.batch_size };
|
||||||
DEFAULT_BATCH_SIZE
|
|
||||||
} else {
|
|
||||||
request.batch_size
|
|
||||||
};
|
|
||||||
|
|
||||||
let parsed = match parse_import_file(&request.file_path, usize::MAX) {
|
let parsed = match parse_import_file(&request.file_path, usize::MAX) {
|
||||||
Ok(parsed) => parsed,
|
Ok(parsed) => parsed,
|
||||||
@@ -583,11 +515,7 @@ where
|
|||||||
error: None,
|
error: None,
|
||||||
});
|
});
|
||||||
|
|
||||||
Ok(TableImportSummary {
|
Ok(TableImportSummary { import_id: request.import_id.clone(), rows_imported, total_rows })
|
||||||
import_id: request.import_id.clone(),
|
|
||||||
rows_imported,
|
|
||||||
total_rows,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -627,81 +555,38 @@ mod tests {
|
|||||||
assert_eq!(parsed.total_rows, 1);
|
assert_eq!(parsed.total_rows, 1);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
parsed.rows[0],
|
parsed.rows[0],
|
||||||
vec![
|
vec![serde_json::Value::String("1".to_string()), serde_json::Value::String("Ada".to_string()),]
|
||||||
serde_json::Value::String("1".to_string()),
|
|
||||||
serde_json::Value::String("Ada".to_string()),
|
|
||||||
]
|
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn parses_json_array_objects_with_union_columns() {
|
fn parses_json_array_objects_with_union_columns() {
|
||||||
let parsed =
|
let parsed = parse_json_bytes(br#"[{"id":1,"name":"Ada"},{"id":2,"active":true}]"#, 10).unwrap();
|
||||||
parse_json_bytes(br#"[{"id":1,"name":"Ada"},{"id":2,"active":true}]"#, 10).unwrap();
|
|
||||||
|
|
||||||
assert_eq!(parsed.columns, vec!["id", "name", "active"]);
|
assert_eq!(parsed.columns, vec!["id", "name", "active"]);
|
||||||
assert_eq!(parsed.total_rows, 2);
|
assert_eq!(parsed.total_rows, 2);
|
||||||
assert_eq!(
|
assert_eq!(parsed.rows[0], vec![serde_json::json!(1), serde_json::json!("Ada"), serde_json::Value::Null,]);
|
||||||
parsed.rows[0],
|
assert_eq!(parsed.rows[1], vec![serde_json::json!(2), serde_json::Value::Null, serde_json::json!(true),]);
|
||||||
vec![
|
|
||||||
serde_json::json!(1),
|
|
||||||
serde_json::json!("Ada"),
|
|
||||||
serde_json::Value::Null,
|
|
||||||
]
|
|
||||||
);
|
|
||||||
assert_eq!(
|
|
||||||
parsed.rows[1],
|
|
||||||
vec![
|
|
||||||
serde_json::json!(2),
|
|
||||||
serde_json::Value::Null,
|
|
||||||
serde_json::json!(true),
|
|
||||||
]
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn builds_import_insert_batches_from_mapped_columns() {
|
fn builds_import_insert_batches_from_mapped_columns() {
|
||||||
let mappings = vec![
|
let mappings = vec![
|
||||||
TableImportColumnMapping {
|
TableImportColumnMapping { source_column: "id".to_string(), target_column: "user_id".to_string() },
|
||||||
source_column: "id".to_string(),
|
TableImportColumnMapping { source_column: "name".to_string(), target_column: "display_name".to_string() },
|
||||||
target_column: "user_id".to_string(),
|
|
||||||
},
|
|
||||||
TableImportColumnMapping {
|
|
||||||
source_column: "name".to_string(),
|
|
||||||
target_column: "display_name".to_string(),
|
|
||||||
},
|
|
||||||
];
|
];
|
||||||
let data = ParsedImportFile {
|
let data = ParsedImportFile {
|
||||||
columns: vec!["id".to_string(), "name".to_string(), "ignored".to_string()],
|
columns: vec!["id".to_string(), "name".to_string(), "ignored".to_string()],
|
||||||
rows: vec![
|
rows: vec![
|
||||||
vec![
|
vec![serde_json::json!(1), serde_json::json!("Ada"), serde_json::json!("x")],
|
||||||
serde_json::json!(1),
|
vec![serde_json::json!(2), serde_json::json!("O'Hara"), serde_json::json!("y")],
|
||||||
serde_json::json!("Ada"),
|
vec![serde_json::json!(3), serde_json::Value::Null, serde_json::json!("z")],
|
||||||
serde_json::json!("x"),
|
|
||||||
],
|
|
||||||
vec![
|
|
||||||
serde_json::json!(2),
|
|
||||||
serde_json::json!("O'Hara"),
|
|
||||||
serde_json::json!("y"),
|
|
||||||
],
|
|
||||||
vec![
|
|
||||||
serde_json::json!(3),
|
|
||||||
serde_json::Value::Null,
|
|
||||||
serde_json::json!("z"),
|
|
||||||
],
|
|
||||||
],
|
],
|
||||||
total_rows: 3,
|
total_rows: 3,
|
||||||
};
|
};
|
||||||
|
|
||||||
let batches = build_import_insert_batches(
|
let batches =
|
||||||
&data,
|
build_import_insert_batches(&data, &mappings, "users", "public", &DatabaseType::Postgres, 2).unwrap();
|
||||||
&mappings,
|
|
||||||
"users",
|
|
||||||
"public",
|
|
||||||
&DatabaseType::Postgres,
|
|
||||||
2,
|
|
||||||
)
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(batches, vec![
|
assert_eq!(batches, vec![
|
||||||
ImportSqlBatch {
|
ImportSqlBatch {
|
||||||
|
|||||||
+110
-80
@@ -1,5 +1,5 @@
|
|||||||
use std::collections::HashSet;
|
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
|
use std::collections::HashSet;
|
||||||
use tokio::sync::RwLock;
|
use tokio::sync::RwLock;
|
||||||
|
|
||||||
use crate::connection::{AppState, PoolKind};
|
use crate::connection::{AppState, PoolKind};
|
||||||
@@ -50,7 +50,9 @@ pub enum TransferStatus {
|
|||||||
|
|
||||||
pub fn quote_identifier(name: &str, db_type: &DatabaseType) -> String {
|
pub fn quote_identifier(name: &str, db_type: &DatabaseType) -> String {
|
||||||
match db_type {
|
match db_type {
|
||||||
DatabaseType::Mysql | DatabaseType::ClickHouse | DatabaseType::Doris | DatabaseType::StarRocks => format!("`{}`", name.replace('`', "``")),
|
DatabaseType::Mysql | DatabaseType::ClickHouse | DatabaseType::Doris | DatabaseType::StarRocks => {
|
||||||
|
format!("`{}`", name.replace('`', "``"))
|
||||||
|
}
|
||||||
DatabaseType::SqlServer => format!("[{}]", name.replace(']', "]]")),
|
DatabaseType::SqlServer => format!("[{}]", name.replace(']', "]]")),
|
||||||
_ => format!("\"{}\"", name.replace('"', "\"\"")),
|
_ => format!("\"{}\"", name.replace('"', "\"\"")),
|
||||||
}
|
}
|
||||||
@@ -69,10 +71,24 @@ pub fn escape_value(val: &serde_json::Value, db_type: &DatabaseType) -> String {
|
|||||||
match val {
|
match val {
|
||||||
serde_json::Value::Null => "NULL".to_string(),
|
serde_json::Value::Null => "NULL".to_string(),
|
||||||
serde_json::Value::Bool(b) => match db_type {
|
serde_json::Value::Bool(b) => match db_type {
|
||||||
DatabaseType::Mysql | DatabaseType::Sqlite | DatabaseType::DuckDb | DatabaseType::Doris | DatabaseType::StarRocks => {
|
DatabaseType::Mysql
|
||||||
if *b { "1".to_string() } else { "0".to_string() }
|
| DatabaseType::Sqlite
|
||||||
|
| DatabaseType::DuckDb
|
||||||
|
| DatabaseType::Doris
|
||||||
|
| DatabaseType::StarRocks => {
|
||||||
|
if *b {
|
||||||
|
"1".to_string()
|
||||||
|
} else {
|
||||||
|
"0".to_string()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
if *b {
|
||||||
|
"TRUE".to_string()
|
||||||
|
} else {
|
||||||
|
"FALSE".to_string()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
_ => if *b { "TRUE".to_string() } else { "FALSE".to_string() },
|
|
||||||
},
|
},
|
||||||
serde_json::Value::Number(n) => n.to_string(),
|
serde_json::Value::Number(n) => n.to_string(),
|
||||||
serde_json::Value::String(s) => {
|
serde_json::Value::String(s) => {
|
||||||
@@ -160,20 +176,17 @@ pub fn map_column_type(source_type: &str, _source_db: &DatabaseType, target_db:
|
|||||||
DatabaseType::Postgres => "TIMESTAMP".into(),
|
DatabaseType::Postgres => "TIMESTAMP".into(),
|
||||||
_ => "DATETIME".into(),
|
_ => "DATETIME".into(),
|
||||||
},
|
},
|
||||||
"timestamp" | "timestamptz" | "timestamp with time zone"
|
"timestamp" | "timestamptz" | "timestamp with time zone" | "timestamp without time zone" => match target_db {
|
||||||
| "timestamp without time zone" => match target_db {
|
|
||||||
DatabaseType::Mysql => "DATETIME".into(),
|
DatabaseType::Mysql => "DATETIME".into(),
|
||||||
DatabaseType::SqlServer => "DATETIME2".into(),
|
DatabaseType::SqlServer => "DATETIME2".into(),
|
||||||
_ => "TIMESTAMP".into(),
|
_ => "TIMESTAMP".into(),
|
||||||
},
|
},
|
||||||
"blob" | "longblob" | "mediumblob" | "tinyblob" | "binary" | "varbinary" | "image" => {
|
"blob" | "longblob" | "mediumblob" | "tinyblob" | "binary" | "varbinary" | "image" => match target_db {
|
||||||
match target_db {
|
DatabaseType::Postgres => "BYTEA".into(),
|
||||||
DatabaseType::Postgres => "BYTEA".into(),
|
DatabaseType::Mysql => "BLOB".into(),
|
||||||
DatabaseType::Mysql => "BLOB".into(),
|
DatabaseType::SqlServer => "VARBINARY(MAX)".into(),
|
||||||
DatabaseType::SqlServer => "VARBINARY(MAX)".into(),
|
_ => "BLOB".into(),
|
||||||
_ => "BLOB".into(),
|
},
|
||||||
}
|
|
||||||
}
|
|
||||||
"bytea" => match target_db {
|
"bytea" => match target_db {
|
||||||
DatabaseType::Postgres => "BYTEA".into(),
|
DatabaseType::Postgres => "BYTEA".into(),
|
||||||
DatabaseType::Mysql => "BLOB".into(),
|
DatabaseType::Mysql => "BLOB".into(),
|
||||||
@@ -217,14 +230,13 @@ pub fn generate_create_table_ddl(
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
let pks: Vec<String> = columns
|
let pks: Vec<String> =
|
||||||
.iter()
|
columns.iter().filter(|c| c.is_primary_key).map(|c| quote_identifier(&c.name, target_db)).collect();
|
||||||
.filter(|c| c.is_primary_key)
|
|
||||||
.map(|c| quote_identifier(&c.name, target_db))
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
let mut ddl = match target_db {
|
let mut ddl = match target_db {
|
||||||
DatabaseType::SqlServer => format!("IF NOT EXISTS (SELECT * FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_NAME = '{table}')\n"),
|
DatabaseType::SqlServer => {
|
||||||
|
format!("IF NOT EXISTS (SELECT * FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_NAME = '{table}')\n")
|
||||||
|
}
|
||||||
_ => String::new(),
|
_ => String::new(),
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -261,11 +273,7 @@ pub fn generate_insert(
|
|||||||
}
|
}
|
||||||
|
|
||||||
let full_table = qualified_table(table, schema, db_type);
|
let full_table = qualified_table(table, schema, db_type);
|
||||||
let col_list = columns
|
let col_list = columns.iter().map(|c| quote_identifier(c, db_type)).collect::<Vec<_>>().join(", ");
|
||||||
.iter()
|
|
||||||
.map(|c| quote_identifier(c, db_type))
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
.join(", ");
|
|
||||||
|
|
||||||
let value_rows: Vec<String> = rows
|
let value_rows: Vec<String> = rows
|
||||||
.iter()
|
.iter()
|
||||||
@@ -287,11 +295,7 @@ pub fn pagination_sql(
|
|||||||
limit: usize,
|
limit: usize,
|
||||||
) -> String {
|
) -> String {
|
||||||
let full_table = qualified_table(table, schema, db_type);
|
let full_table = qualified_table(table, schema, db_type);
|
||||||
let col_list = columns
|
let col_list = columns.iter().map(|c| quote_identifier(c, db_type)).collect::<Vec<_>>().join(", ");
|
||||||
.iter()
|
|
||||||
.map(|c| quote_identifier(c, db_type))
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
.join(", ");
|
|
||||||
|
|
||||||
match db_type {
|
match db_type {
|
||||||
DatabaseType::SqlServer | DatabaseType::Oracle => {
|
DatabaseType::SqlServer | DatabaseType::Oracle => {
|
||||||
@@ -310,11 +314,7 @@ pub fn count_sql(table: &str, schema: &str, db_type: &DatabaseType) -> String {
|
|||||||
format!("SELECT COUNT(*) FROM {full_table}")
|
format!("SELECT COUNT(*) FROM {full_table}")
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn execute_on_pool(
|
pub async fn execute_on_pool(state: &AppState, pool_key: &str, sql: &str) -> Result<db::QueryResult, String> {
|
||||||
state: &AppState,
|
|
||||||
pool_key: &str,
|
|
||||||
sql: &str,
|
|
||||||
) -> Result<db::QueryResult, String> {
|
|
||||||
let connections = state.connections.lock().await;
|
let connections = state.connections.lock().await;
|
||||||
let pool = connections.get(pool_key).ok_or("Connection not found")?;
|
let pool = connections.get(pool_key).ok_or("Connection not found")?;
|
||||||
|
|
||||||
@@ -361,8 +361,10 @@ pub async fn execute_on_pool(
|
|||||||
let con = con.lock().map_err(|e| e.to_string())?;
|
let con = con.lock().map_err(|e| e.to_string())?;
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
let trimmed = sql.trim().to_uppercase();
|
let trimmed = sql.trim().to_uppercase();
|
||||||
if trimmed.starts_with("SELECT") || trimmed.starts_with("SHOW")
|
if trimmed.starts_with("SELECT")
|
||||||
|| trimmed.starts_with("DESCRIBE") || trimmed.starts_with("WITH")
|
|| trimmed.starts_with("SHOW")
|
||||||
|
|| trimmed.starts_with("DESCRIBE")
|
||||||
|
|| trimmed.starts_with("WITH")
|
||||||
|| trimmed.starts_with("PRAGMA")
|
|| trimmed.starts_with("PRAGMA")
|
||||||
{
|
{
|
||||||
let mut stmt = con.prepare(&sql).map_err(|e| e.to_string())?;
|
let mut stmt = con.prepare(&sql).map_err(|e| e.to_string())?;
|
||||||
@@ -374,21 +376,40 @@ pub async fn execute_on_pool(
|
|||||||
.collect();
|
.collect();
|
||||||
let mut result_rows = Vec::new();
|
let mut result_rows = Vec::new();
|
||||||
while let Some(row) = rows.next().map_err(|e| e.to_string())? {
|
while let Some(row) = rows.next().map_err(|e| e.to_string())? {
|
||||||
let vals: Vec<serde_json::Value> = (0..col_count).map(|i| {
|
let vals: Vec<serde_json::Value> = (0..col_count)
|
||||||
row.get::<_, String>(i).map(serde_json::Value::String)
|
.map(|i| {
|
||||||
.or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into())))
|
row.get::<_, String>(i)
|
||||||
.or_else(|_| row.get::<_, f64>(i).map(|v| {
|
.map(serde_json::Value::String)
|
||||||
serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null)
|
.or_else(|_| row.get::<_, i64>(i).map(|v| serde_json::Value::Number(v.into())))
|
||||||
}))
|
.or_else(|_| {
|
||||||
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool))
|
row.get::<_, f64>(i).map(|v| {
|
||||||
.unwrap_or(serde_json::Value::Null)
|
serde_json::Number::from_f64(v)
|
||||||
}).collect();
|
.map(serde_json::Value::Number)
|
||||||
|
.unwrap_or(serde_json::Value::Null)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
.or_else(|_| row.get::<_, bool>(i).map(serde_json::Value::Bool))
|
||||||
|
.unwrap_or(serde_json::Value::Null)
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
result_rows.push(vals);
|
result_rows.push(vals);
|
||||||
}
|
}
|
||||||
Ok(db::QueryResult { columns, rows: result_rows, affected_rows: 0, execution_time_ms: start.elapsed().as_millis(), truncated: false })
|
Ok(db::QueryResult {
|
||||||
|
columns,
|
||||||
|
rows: result_rows,
|
||||||
|
affected_rows: 0,
|
||||||
|
execution_time_ms: start.elapsed().as_millis(),
|
||||||
|
truncated: false,
|
||||||
|
})
|
||||||
} else {
|
} else {
|
||||||
let affected = con.execute(&sql, []).map_err(|e| e.to_string())?;
|
let affected = con.execute(&sql, []).map_err(|e| e.to_string())?;
|
||||||
Ok(db::QueryResult { columns: vec![], rows: vec![], affected_rows: affected as u64, execution_time_ms: start.elapsed().as_millis(), truncated: false })
|
Ok(db::QueryResult {
|
||||||
|
columns: vec![],
|
||||||
|
rows: vec![],
|
||||||
|
affected_rows: affected as u64,
|
||||||
|
execution_time_ms: start.elapsed().as_millis(),
|
||||||
|
truncated: false,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
@@ -422,28 +443,34 @@ pub async fn get_columns_for_transfer(
|
|||||||
let table = table.to_string();
|
let table = table.to_string();
|
||||||
return tokio::task::spawn_blocking(move || {
|
return tokio::task::spawn_blocking(move || {
|
||||||
let con = con.lock().map_err(|e| e.to_string())?;
|
let con = con.lock().map_err(|e| e.to_string())?;
|
||||||
let mut stmt = con.prepare(
|
let mut stmt = con
|
||||||
"SELECT column_name, data_type, is_nullable, column_default
|
.prepare(
|
||||||
|
"SELECT column_name, data_type, is_nullable, column_default
|
||||||
FROM information_schema.columns
|
FROM information_schema.columns
|
||||||
WHERE table_schema = 'main' AND table_name = ?
|
WHERE table_schema = 'main' AND table_name = ?
|
||||||
ORDER BY ordinal_position"
|
ORDER BY ordinal_position",
|
||||||
).map_err(|e| e.to_string())?;
|
)
|
||||||
let rows = stmt.query_map([&table], |row| {
|
.map_err(|e| e.to_string())?;
|
||||||
Ok(db::ColumnInfo {
|
let rows = stmt
|
||||||
name: row.get::<_, String>(0)?,
|
.query_map([&table], |row| {
|
||||||
data_type: row.get::<_, String>(1)?,
|
Ok(db::ColumnInfo {
|
||||||
is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES",
|
name: row.get::<_, String>(0)?,
|
||||||
column_default: row.get::<_, Option<String>>(3)?,
|
data_type: row.get::<_, String>(1)?,
|
||||||
is_primary_key: false,
|
is_nullable: row.get::<_, String>(2).unwrap_or_default() == "YES",
|
||||||
extra: None,
|
column_default: row.get::<_, Option<String>>(3)?,
|
||||||
comment: None,
|
is_primary_key: false,
|
||||||
numeric_precision: None,
|
extra: None,
|
||||||
numeric_scale: None,
|
comment: None,
|
||||||
character_maximum_length: None,
|
numeric_precision: None,
|
||||||
|
numeric_scale: None,
|
||||||
|
character_maximum_length: None,
|
||||||
|
})
|
||||||
})
|
})
|
||||||
}).map_err(|e| e.to_string())?;
|
.map_err(|e| e.to_string())?;
|
||||||
Ok(rows.filter_map(|r| r.ok()).collect())
|
Ok(rows.filter_map(|r| r.ok()).collect())
|
||||||
}).await.map_err(|e| e.to_string())?;
|
})
|
||||||
|
.await
|
||||||
|
.map_err(|e| e.to_string())?;
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(PoolKind::ClickHouse(client)) = connections.get(pool_key) {
|
if let Some(PoolKind::ClickHouse(client)) = connections.get(pool_key) {
|
||||||
@@ -526,9 +553,14 @@ where
|
|||||||
// Get source columns (deduplicate by name)
|
// Get source columns (deduplicate by name)
|
||||||
let columns = {
|
let columns = {
|
||||||
let raw = get_columns_for_transfer(
|
let raw = get_columns_for_transfer(
|
||||||
state, source_pool_key, &request.source_connection_id,
|
state,
|
||||||
&request.source_database, &request.source_schema, table,
|
source_pool_key,
|
||||||
).await?;
|
&request.source_connection_id,
|
||||||
|
&request.source_database,
|
||||||
|
&request.source_schema,
|
||||||
|
table,
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
let mut seen = std::collections::HashSet::new();
|
let mut seen = std::collections::HashSet::new();
|
||||||
raw.into_iter().filter(|c| seen.insert(c.name.clone())).collect::<Vec<_>>()
|
raw.into_iter().filter(|c| seen.insert(c.name.clone())).collect::<Vec<_>>()
|
||||||
};
|
};
|
||||||
@@ -544,13 +576,11 @@ where
|
|||||||
let total_rows = {
|
let total_rows = {
|
||||||
let sql = count_sql(table, &request.source_schema, source_db_type);
|
let sql = count_sql(table, &request.source_schema, source_db_type);
|
||||||
match execute_on_pool(state, source_pool_key, &sql).await {
|
match execute_on_pool(state, source_pool_key, &sql).await {
|
||||||
Ok(result) => result.rows.first()
|
Ok(result) => result.rows.first().and_then(|r| r.first()).and_then(|v| match v {
|
||||||
.and_then(|r| r.first())
|
serde_json::Value::Number(n) => n.as_u64(),
|
||||||
.and_then(|v| match v {
|
serde_json::Value::String(s) => s.parse::<u64>().ok(),
|
||||||
serde_json::Value::Number(n) => n.as_u64(),
|
_ => None,
|
||||||
serde_json::Value::String(s) => s.parse::<u64>().ok(),
|
}),
|
||||||
_ => None,
|
|
||||||
}),
|
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
log::warn!("[transfer] count failed for {}: {}", table, e);
|
log::warn!("[transfer] count failed for {}: {}", table, e);
|
||||||
None
|
None
|
||||||
@@ -578,8 +608,7 @@ where
|
|||||||
DatabaseType::Sqlite | DatabaseType::DuckDb => format!("DELETE FROM {full_table}"),
|
DatabaseType::Sqlite | DatabaseType::DuckDb => format!("DELETE FROM {full_table}"),
|
||||||
_ => format!("TRUNCATE TABLE {full_table}"),
|
_ => format!("TRUNCATE TABLE {full_table}"),
|
||||||
};
|
};
|
||||||
execute_on_pool(state, target_pool_key, &truncate_sql).await
|
execute_on_pool(state, target_pool_key, &truncate_sql).await.map_err(|e| format!("Failed to truncate: {e}"))?;
|
||||||
.map_err(|e| format!("Failed to truncate: {e}"))?;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Transfer data in batches
|
// Transfer data in batches
|
||||||
@@ -602,7 +631,8 @@ where
|
|||||||
|
|
||||||
let insert_sql = generate_insert(&col_names, &result.rows, table, &request.target_schema, target_db_type);
|
let insert_sql = generate_insert(&col_names, &result.rows, table, &request.target_schema, target_db_type);
|
||||||
if !insert_sql.is_empty() {
|
if !insert_sql.is_empty() {
|
||||||
execute_on_pool(state, target_pool_key, &insert_sql).await
|
execute_on_pool(state, target_pool_key, &insert_sql)
|
||||||
|
.await
|
||||||
.map_err(|e| format!("Insert failed at offset {offset}: {e}"))?;
|
.map_err(|e| format!("Insert failed at offset {offset}: {e}"))?;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -53,10 +53,7 @@ pub fn normalize_version(version: &str) -> String {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn parse_version(version: &str) -> Vec<u64> {
|
pub fn parse_version(version: &str) -> Vec<u64> {
|
||||||
normalize_version(version)
|
normalize_version(version).split(['.', '-', '+']).map(|part| part.parse::<u64>().unwrap_or(0)).collect()
|
||||||
.split(['.', '-', '+'])
|
|
||||||
.map(|part| part.parse::<u64>().unwrap_or(0))
|
|
||||||
.collect()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn is_newer_version(latest: &str, current: &str) -> bool {
|
pub fn is_newer_version(latest: &str, current: &str) -> bool {
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
edition = "2021"
|
||||||
|
max_width = 120
|
||||||
|
use_small_heuristics = "Max"
|
||||||
+1
-1
@@ -1,3 +1,3 @@
|
|||||||
fn main() {
|
fn main() {
|
||||||
tauri_build::build()
|
tauri_build::build()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use tauri::{Emitter, AppHandle, State};
|
use tauri::{AppHandle, Emitter, State};
|
||||||
|
|
||||||
use super::connection::AppState;
|
use super::connection::AppState;
|
||||||
pub use dbx_core::ai::*;
|
pub use dbx_core::ai::*;
|
||||||
@@ -10,17 +10,12 @@ pub async fn ai_test_connection(config: AiConfig) -> Result<String, String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn save_ai_config(
|
pub async fn save_ai_config(state: State<'_, Arc<AppState>>, config: AiConfig) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
config: AiConfig,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
state.storage.save_ai_config(&config).await
|
state.storage.save_ai_config(&config).await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn load_ai_config(
|
pub async fn load_ai_config(state: State<'_, Arc<AppState>>) -> Result<Option<AiConfig>, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
) -> Result<Option<AiConfig>, String> {
|
|
||||||
state.storage.load_ai_config().await
|
state.storage.load_ai_config().await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -30,11 +25,7 @@ pub async fn ai_complete(request: AiCompletionRequest) -> Result<String, String>
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn ai_stream(
|
pub async fn ai_stream(app: AppHandle, session_id: String, request: AiCompletionRequest) -> Result<(), String> {
|
||||||
app: AppHandle,
|
|
||||||
session_id: String,
|
|
||||||
request: AiCompletionRequest,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let cancelled = dbx_core::ai::register_stream(&session_id).await;
|
let cancelled = dbx_core::ai::register_stream(&session_id).await;
|
||||||
|
|
||||||
let result = dbx_core::ai::stream(&session_id, &request, &cancelled, |chunk| {
|
let result = dbx_core::ai::stream(&session_id, &request, &cancelled, |chunk| {
|
||||||
@@ -52,24 +43,16 @@ pub async fn ai_cancel_stream(session_id: String) -> Result<bool, String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn save_ai_conversation(
|
pub async fn save_ai_conversation(state: State<'_, Arc<AppState>>, conversation: AiConversation) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
conversation: AiConversation,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
state.storage.save_ai_conversation(&conversation).await
|
state.storage.save_ai_conversation(&conversation).await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn load_ai_conversations(
|
pub async fn load_ai_conversations(state: State<'_, Arc<AppState>>) -> Result<Vec<AiConversation>, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
) -> Result<Vec<AiConversation>, String> {
|
|
||||||
state.storage.load_ai_conversations().await
|
state.storage.load_ai_conversations().await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn delete_ai_conversation(
|
pub async fn delete_ai_conversation(state: State<'_, Arc<AppState>>, id: String) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
id: String,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
state.storage.delete_ai_conversation(&id).await
|
state.storage.delete_ai_conversation(&id).await
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,74 +2,52 @@ use std::sync::Arc;
|
|||||||
use tauri::State;
|
use tauri::State;
|
||||||
|
|
||||||
pub use dbx_core::connection::{
|
pub use dbx_core::connection::{
|
||||||
connection_url_for_endpoint, expand_tilde, probe_connection_endpoint,
|
connection_url_for_endpoint, expand_tilde, probe_connection_endpoint, redacted_connection_url_for_endpoint,
|
||||||
redacted_connection_url_for_endpoint, AppState, PoolKind,
|
AppState, PoolKind,
|
||||||
};
|
};
|
||||||
use dbx_core::db;
|
use dbx_core::db;
|
||||||
use dbx_core::models::connection::{ConnectionConfig, DatabaseType};
|
use dbx_core::models::connection::{ConnectionConfig, DatabaseType};
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn save_connections(
|
pub async fn save_connections(state: State<'_, Arc<AppState>>, configs: Vec<ConnectionConfig>) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
configs: Vec<ConnectionConfig>,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
state.storage.save_connections(&configs).await
|
state.storage.save_connections(&configs).await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn load_connections(
|
pub async fn load_connections(state: State<'_, Arc<AppState>>) -> Result<Vec<ConnectionConfig>, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
) -> Result<Vec<ConnectionConfig>, String> {
|
|
||||||
state.storage.load_connections().await
|
state.storage.load_connections().await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn save_sidebar_layout(
|
pub async fn save_sidebar_layout(state: State<'_, Arc<AppState>>, layout: serde_json::Value) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
layout: serde_json::Value,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
state.storage.save_sidebar_layout(&layout).await
|
state.storage.save_sidebar_layout(&layout).await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn load_sidebar_layout(
|
pub async fn load_sidebar_layout(state: State<'_, Arc<AppState>>) -> Result<Option<serde_json::Value>, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
) -> Result<Option<serde_json::Value>, String> {
|
|
||||||
state.storage.load_sidebar_layout().await
|
state.storage.load_sidebar_layout().await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn test_connection(
|
pub async fn test_connection(state: State<'_, Arc<AppState>>, config: ConnectionConfig) -> Result<String, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
config: ConnectionConfig,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let tunnel_id = format!("{}:test", config.id);
|
let tunnel_id = format!("{}:test", config.id);
|
||||||
let connection_id = if config.ssh_enabled && !config.ssh_host.is_empty() {
|
let connection_id =
|
||||||
tunnel_id.as_str()
|
if config.ssh_enabled && !config.ssh_host.is_empty() { tunnel_id.as_str() } else { config.id.as_str() };
|
||||||
} else {
|
|
||||||
config.id.as_str()
|
|
||||||
};
|
|
||||||
let (host, port) = state.connection_host_port(connection_id, &config).await?;
|
let (host, port) = state.connection_host_port(connection_id, &config).await?;
|
||||||
let probe_result = probe_connection_endpoint(&config, &host, port).await;
|
let probe_result = probe_connection_endpoint(&config, &host, port).await;
|
||||||
let url = connection_url_for_endpoint(&config, &host, port);
|
let url = connection_url_for_endpoint(&config, &host, port);
|
||||||
let target = redacted_connection_url_for_endpoint(&config, &host, port);
|
let target = redacted_connection_url_for_endpoint(&config, &host, port);
|
||||||
log::info!(
|
log::info!("[test_connection] db_type={:?} target={}", config.db_type, target);
|
||||||
"[test_connection] db_type={:?} target={}",
|
|
||||||
config.db_type,
|
|
||||||
target
|
|
||||||
);
|
|
||||||
let result = match probe_result {
|
let result = match probe_result {
|
||||||
Err(e) => Err(e),
|
Err(e) => Err(e),
|
||||||
Ok(()) => match config.db_type {
|
Ok(()) => match config.db_type {
|
||||||
DatabaseType::Mysql if config.needs_bare_mysql() => {
|
DatabaseType::Mysql if config.needs_bare_mysql() => match db::mysql::connect_bare(&url).await {
|
||||||
match db::mysql::connect_bare(&url).await {
|
Ok(pool) => {
|
||||||
Ok(pool) => {
|
pool.close().await;
|
||||||
pool.close().await;
|
Ok("Connection successful".to_string())
|
||||||
Ok("Connection successful".to_string())
|
|
||||||
}
|
|
||||||
Err(e) => Err(e),
|
|
||||||
}
|
}
|
||||||
}
|
Err(e) => Err(e),
|
||||||
|
},
|
||||||
DatabaseType::Mysql => match db::mysql::connect(&url).await {
|
DatabaseType::Mysql => match db::mysql::connect(&url).await {
|
||||||
Ok(pool) => {
|
Ok(pool) => {
|
||||||
pool.close().await;
|
pool.close().await;
|
||||||
@@ -77,70 +55,48 @@ pub async fn test_connection(
|
|||||||
}
|
}
|
||||||
Err(e) => Err(e),
|
Err(e) => Err(e),
|
||||||
},
|
},
|
||||||
DatabaseType::Doris | DatabaseType::StarRocks => {
|
DatabaseType::Doris | DatabaseType::StarRocks => match db::mysql::connect_bare(&url).await {
|
||||||
match db::mysql::connect_bare(&url).await {
|
Ok(pool) => {
|
||||||
Ok(pool) => {
|
pool.close().await;
|
||||||
pool.close().await;
|
Ok("Connection successful".to_string())
|
||||||
Ok("Connection successful".to_string())
|
|
||||||
}
|
|
||||||
Err(e) => Err(e),
|
|
||||||
}
|
}
|
||||||
}
|
Err(e) => Err(e),
|
||||||
DatabaseType::Postgres | DatabaseType::Redshift => {
|
},
|
||||||
match db::postgres::connect(&url).await {
|
DatabaseType::Postgres | DatabaseType::Redshift => match db::postgres::connect(&url).await {
|
||||||
Ok(pool) => {
|
Ok(pool) => {
|
||||||
pool.close().await;
|
pool.close().await;
|
||||||
Ok("Connection successful".to_string())
|
Ok("Connection successful".to_string())
|
||||||
}
|
|
||||||
Err(e) => Err(e),
|
|
||||||
}
|
}
|
||||||
}
|
Err(e) => Err(e),
|
||||||
DatabaseType::Sqlite => {
|
},
|
||||||
match db::sqlite::connect_path(&expand_tilde(&config.host)).await {
|
DatabaseType::Sqlite => match db::sqlite::connect_path(&expand_tilde(&config.host)).await {
|
||||||
Ok(pool) => {
|
Ok(pool) => {
|
||||||
pool.close().await;
|
pool.close().await;
|
||||||
Ok("Connection successful".to_string())
|
Ok("Connection successful".to_string())
|
||||||
}
|
|
||||||
Err(e) => Err(e),
|
|
||||||
}
|
}
|
||||||
}
|
Err(e) => Err(e),
|
||||||
DatabaseType::Redis => db::redis_driver::connect(&url)
|
},
|
||||||
.await
|
DatabaseType::Redis => db::redis_driver::connect(&url).await.map(|_| "Connection successful".to_string()),
|
||||||
.map(|_| "Connection successful".to_string()),
|
|
||||||
DatabaseType::DuckDb => duckdb::Connection::open(&expand_tilde(&config.host))
|
DatabaseType::DuckDb => duckdb::Connection::open(&expand_tilde(&config.host))
|
||||||
.map(|_| "Connection successful".to_string())
|
.map(|_| "Connection successful".to_string())
|
||||||
.map_err(|e| e.to_string()),
|
.map_err(|e| e.to_string()),
|
||||||
DatabaseType::MongoDb => match db::mongo_driver::connect(&url).await {
|
DatabaseType::MongoDb => match db::mongo_driver::connect(&url).await {
|
||||||
Ok(client) => db::mongo_driver::test_connection(&client)
|
Ok(client) => {
|
||||||
.await
|
db::mongo_driver::test_connection(&client).await.map(|_| "Connection successful".to_string())
|
||||||
.map(|_| "Connection successful".to_string()),
|
}
|
||||||
Err(e) => Err(e.to_string()),
|
Err(e) => Err(e.to_string()),
|
||||||
},
|
},
|
||||||
DatabaseType::ClickHouse => {
|
DatabaseType::ClickHouse => {
|
||||||
let username = if config.username.is_empty() {
|
let username = if config.username.is_empty() { None } else { Some(config.username.clone()) };
|
||||||
None
|
let password = if config.password.is_empty() { None } else { Some(config.password.clone()) };
|
||||||
} else {
|
|
||||||
Some(config.username.clone())
|
|
||||||
};
|
|
||||||
let password = if config.password.is_empty() {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
Some(config.password.clone())
|
|
||||||
};
|
|
||||||
let client = db::clickhouse_driver::ChClient::new(&url, username, password);
|
let client = db::clickhouse_driver::ChClient::new(&url, username, password);
|
||||||
db::clickhouse_driver::test_connection(&client)
|
db::clickhouse_driver::test_connection(&client).await.map(|_| "Connection successful".to_string())
|
||||||
|
}
|
||||||
|
DatabaseType::SqlServer => {
|
||||||
|
db::sqlserver::connect(&host, port, &config.username, &config.password, config.database.as_deref())
|
||||||
.await
|
.await
|
||||||
.map(|_| "Connection successful".to_string())
|
.map(|_| "Connection successful".to_string())
|
||||||
}
|
}
|
||||||
DatabaseType::SqlServer => db::sqlserver::connect(
|
|
||||||
&host,
|
|
||||||
port,
|
|
||||||
&config.username,
|
|
||||||
&config.password,
|
|
||||||
config.database.as_deref(),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map(|_| "Connection successful".to_string()),
|
|
||||||
DatabaseType::Oracle => db::oracle_driver::connect(
|
DatabaseType::Oracle => db::oracle_driver::connect(
|
||||||
&host,
|
&host,
|
||||||
port,
|
port,
|
||||||
@@ -151,14 +107,9 @@ pub async fn test_connection(
|
|||||||
.await
|
.await
|
||||||
.map(|_| "Connection successful".to_string()),
|
.map(|_| "Connection successful".to_string()),
|
||||||
DatabaseType::Elasticsearch => {
|
DatabaseType::Elasticsearch => {
|
||||||
let client = db::elasticsearch_driver::EsClient::new(
|
let client =
|
||||||
&url,
|
db::elasticsearch_driver::EsClient::new(&url, Some(&config.username), Some(&config.password));
|
||||||
Some(&config.username),
|
db::elasticsearch_driver::test_connection(&client).await.map(|_| "Connection successful".to_string())
|
||||||
Some(&config.password),
|
|
||||||
);
|
|
||||||
db::elasticsearch_driver::test_connection(&client)
|
|
||||||
.await
|
|
||||||
.map(|_| "Connection successful".to_string())
|
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
@@ -171,10 +122,7 @@ pub async fn test_connection(
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn connect_db(
|
pub async fn connect_db(state: State<'_, Arc<AppState>>, config: ConnectionConfig) -> Result<String, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
config: ConnectionConfig,
|
|
||||||
) -> Result<String, String> {
|
|
||||||
let id = config.id.clone();
|
let id = config.id.clone();
|
||||||
|
|
||||||
let (host, port) = state.connection_host_port(&id, &config).await?;
|
let (host, port) = state.connection_host_port(&id, &config).await?;
|
||||||
@@ -182,26 +130,17 @@ pub async fn connect_db(
|
|||||||
let url = connection_url_for_endpoint(&config, &host, port);
|
let url = connection_url_for_endpoint(&config, &host, port);
|
||||||
|
|
||||||
let pool = match config.db_type {
|
let pool = match config.db_type {
|
||||||
DatabaseType::Mysql if config.needs_bare_mysql() => {
|
DatabaseType::Mysql if config.needs_bare_mysql() => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true),
|
||||||
PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true)
|
|
||||||
}
|
|
||||||
DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?, false),
|
DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?, false),
|
||||||
DatabaseType::Doris | DatabaseType::StarRocks => {
|
DatabaseType::Doris | DatabaseType::StarRocks => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true),
|
||||||
PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true)
|
DatabaseType::Postgres | DatabaseType::Redshift => PoolKind::Postgres(db::postgres::connect(&url).await?),
|
||||||
}
|
DatabaseType::Sqlite => PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&config.host)).await?),
|
||||||
DatabaseType::Postgres | DatabaseType::Redshift => {
|
|
||||||
PoolKind::Postgres(db::postgres::connect(&url).await?)
|
|
||||||
}
|
|
||||||
DatabaseType::Sqlite => {
|
|
||||||
PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&config.host)).await?)
|
|
||||||
}
|
|
||||||
DatabaseType::Redis => {
|
DatabaseType::Redis => {
|
||||||
let con = db::redis_driver::connect(&url).await?;
|
let con = db::redis_driver::connect(&url).await?;
|
||||||
PoolKind::Redis(tokio::sync::Mutex::new(con))
|
PoolKind::Redis(tokio::sync::Mutex::new(con))
|
||||||
}
|
}
|
||||||
DatabaseType::DuckDb => {
|
DatabaseType::DuckDb => {
|
||||||
let con =
|
let con = duckdb::Connection::open(&expand_tilde(&config.host)).map_err(|e| e.to_string())?;
|
||||||
duckdb::Connection::open(&expand_tilde(&config.host)).map_err(|e| e.to_string())?;
|
|
||||||
PoolKind::DuckDb(std::sync::Arc::new(std::sync::Mutex::new(con)))
|
PoolKind::DuckDb(std::sync::Arc::new(std::sync::Mutex::new(con)))
|
||||||
}
|
}
|
||||||
DatabaseType::MongoDb => {
|
DatabaseType::MongoDb => {
|
||||||
@@ -210,29 +149,16 @@ pub async fn connect_db(
|
|||||||
PoolKind::MongoDb(client)
|
PoolKind::MongoDb(client)
|
||||||
}
|
}
|
||||||
DatabaseType::ClickHouse => {
|
DatabaseType::ClickHouse => {
|
||||||
let username = if config.username.is_empty() {
|
let username = if config.username.is_empty() { None } else { Some(config.username.clone()) };
|
||||||
None
|
let password = if config.password.is_empty() { None } else { Some(config.password.clone()) };
|
||||||
} else {
|
|
||||||
Some(config.username.clone())
|
|
||||||
};
|
|
||||||
let password = if config.password.is_empty() {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
Some(config.password.clone())
|
|
||||||
};
|
|
||||||
let client = db::clickhouse_driver::ChClient::new(&url, username, password);
|
let client = db::clickhouse_driver::ChClient::new(&url, username, password);
|
||||||
db::clickhouse_driver::test_connection(&client).await?;
|
db::clickhouse_driver::test_connection(&client).await?;
|
||||||
PoolKind::ClickHouse(client)
|
PoolKind::ClickHouse(client)
|
||||||
}
|
}
|
||||||
DatabaseType::SqlServer => {
|
DatabaseType::SqlServer => {
|
||||||
let client = db::sqlserver::connect(
|
let client =
|
||||||
&host,
|
db::sqlserver::connect(&host, port, &config.username, &config.password, config.database.as_deref())
|
||||||
port,
|
.await?;
|
||||||
&config.username,
|
|
||||||
&config.password,
|
|
||||||
config.database.as_deref(),
|
|
||||||
)
|
|
||||||
.await?;
|
|
||||||
PoolKind::SqlServer(std::sync::Arc::new(tokio::sync::Mutex::new(client)))
|
PoolKind::SqlServer(std::sync::Arc::new(tokio::sync::Mutex::new(client)))
|
||||||
}
|
}
|
||||||
DatabaseType::Oracle => {
|
DatabaseType::Oracle => {
|
||||||
@@ -247,11 +173,7 @@ pub async fn connect_db(
|
|||||||
PoolKind::Oracle(std::sync::Arc::new(tokio::sync::Mutex::new(client)))
|
PoolKind::Oracle(std::sync::Arc::new(tokio::sync::Mutex::new(client)))
|
||||||
}
|
}
|
||||||
DatabaseType::Elasticsearch => {
|
DatabaseType::Elasticsearch => {
|
||||||
let client = db::elasticsearch_driver::EsClient::new(
|
let client = db::elasticsearch_driver::EsClient::new(&url, Some(&config.username), Some(&config.password));
|
||||||
&url,
|
|
||||||
Some(&config.username),
|
|
||||||
Some(&config.password),
|
|
||||||
);
|
|
||||||
db::elasticsearch_driver::test_connection(&client).await?;
|
db::elasticsearch_driver::test_connection(&client).await?;
|
||||||
PoolKind::Elasticsearch(client)
|
PoolKind::Elasticsearch(client)
|
||||||
}
|
}
|
||||||
@@ -264,16 +186,10 @@ pub async fn connect_db(
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn disconnect_db(
|
pub async fn disconnect_db(state: State<'_, Arc<AppState>>, connection_id: String) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
connection_id: String,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
let mut conns = state.connections.lock().await;
|
let mut conns = state.connections.lock().await;
|
||||||
let keys_to_remove: Vec<String> = conns
|
let keys_to_remove: Vec<String> =
|
||||||
.keys()
|
conns.keys().filter(|k| *k == &connection_id || k.starts_with(&format!("{connection_id}:"))).cloned().collect();
|
||||||
.filter(|k| *k == &connection_id || k.starts_with(&format!("{connection_id}:")))
|
|
||||||
.cloned()
|
|
||||||
.collect();
|
|
||||||
for key in keys_to_remove {
|
for key in keys_to_remove {
|
||||||
if let Some(pool) = conns.remove(&key) {
|
if let Some(pool) = conns.remove(&key) {
|
||||||
match pool {
|
match pool {
|
||||||
|
|||||||
@@ -5,10 +5,7 @@ use super::connection::AppState;
|
|||||||
pub use dbx_core::history::HistoryEntry;
|
pub use dbx_core::history::HistoryEntry;
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn save_history(
|
pub async fn save_history(state: State<'_, Arc<AppState>>, entry: HistoryEntry) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
entry: HistoryEntry,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
state.storage.save_history_entry(&entry).await
|
state.storage.save_history_entry(&entry).await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -27,9 +24,6 @@ pub async fn clear_history(state: State<'_, Arc<AppState>>) -> Result<(), String
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn delete_history_entry(
|
pub async fn delete_history_entry(state: State<'_, Arc<AppState>>, id: String) -> Result<(), String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
id: String,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
state.storage.delete_history_entry(&id).await
|
state.storage.delete_history_entry(&id).await
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use tauri::{Emitter, AppHandle, Manager};
|
use tauri::{AppHandle, Emitter, Manager};
|
||||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||||
use tokio::net::TcpListener;
|
use tokio::net::TcpListener;
|
||||||
|
|
||||||
@@ -82,7 +82,10 @@ pub fn start(app_handle: AppHandle, state: Arc<AppState>) {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
fn find_config_by_name<'a>(configs: &'a [crate::models::connection::ConnectionConfig], name: &str) -> Option<&'a crate::models::connection::ConnectionConfig> {
|
fn find_config_by_name<'a>(
|
||||||
|
configs: &'a [crate::models::connection::ConnectionConfig],
|
||||||
|
name: &str,
|
||||||
|
) -> Option<&'a crate::models::connection::ConnectionConfig> {
|
||||||
configs.iter().find(|c| c.name.eq_ignore_ascii_case(name))
|
configs.iter().find(|c| c.name.eq_ignore_ascii_case(name))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -94,14 +97,21 @@ async fn respond(stream: &mut tokio::net::TcpStream, status: &str, body: &str) {
|
|||||||
async fn handle_open_table(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) {
|
async fn handle_open_table(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) {
|
||||||
let req: OpenTableRequest = match serde_json::from_str(body) {
|
let req: OpenTableRequest = match serde_json::from_str(body) {
|
||||||
Ok(r) => r,
|
Ok(r) => r,
|
||||||
Err(_) => { respond(stream, "400 Bad Request", "").await; return; }
|
Err(_) => {
|
||||||
|
respond(stream, "400 Bad Request", "").await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
};
|
};
|
||||||
let configs = match state.storage.load_connections().await {
|
let configs = match state.storage.load_connections().await {
|
||||||
Ok(c) => c,
|
Ok(c) => c,
|
||||||
Err(_) => { respond(stream, "500 Internal Server Error", "").await; return; }
|
Err(_) => {
|
||||||
|
respond(stream, "500 Internal Server Error", "").await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
};
|
};
|
||||||
let Some(config) = find_config_by_name(&configs, &req.connection_name) else {
|
let Some(config) = find_config_by_name(&configs, &req.connection_name) else {
|
||||||
respond(stream, "404 Not Found", "Connection not found").await; return;
|
respond(stream, "404 Not Found", "Connection not found").await;
|
||||||
|
return;
|
||||||
};
|
};
|
||||||
let event = McpOpenTableEvent {
|
let event = McpOpenTableEvent {
|
||||||
connection_id: config.id.clone(),
|
connection_id: config.id.clone(),
|
||||||
@@ -116,14 +126,21 @@ async fn handle_open_table(app: &AppHandle, state: &Arc<AppState>, body: &str, s
|
|||||||
async fn handle_execute_query(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) {
|
async fn handle_execute_query(app: &AppHandle, state: &Arc<AppState>, body: &str, stream: &mut tokio::net::TcpStream) {
|
||||||
let req: ExecuteQueryRequest = match serde_json::from_str(body) {
|
let req: ExecuteQueryRequest = match serde_json::from_str(body) {
|
||||||
Ok(r) => r,
|
Ok(r) => r,
|
||||||
Err(_) => { respond(stream, "400 Bad Request", "").await; return; }
|
Err(_) => {
|
||||||
|
respond(stream, "400 Bad Request", "").await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
};
|
};
|
||||||
let configs = match state.storage.load_connections().await {
|
let configs = match state.storage.load_connections().await {
|
||||||
Ok(c) => c,
|
Ok(c) => c,
|
||||||
Err(_) => { respond(stream, "500 Internal Server Error", "").await; return; }
|
Err(_) => {
|
||||||
|
respond(stream, "500 Internal Server Error", "").await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
};
|
};
|
||||||
let Some(config) = find_config_by_name(&configs, &req.connection_name) else {
|
let Some(config) = find_config_by_name(&configs, &req.connection_name) else {
|
||||||
respond(stream, "404 Not Found", "Connection not found").await; return;
|
respond(stream, "404 Not Found", "Connection not found").await;
|
||||||
|
return;
|
||||||
};
|
};
|
||||||
let event = McpExecuteQueryEvent {
|
let event = McpExecuteQueryEvent {
|
||||||
connection_id: config.id.clone(),
|
connection_id: config.id.clone(),
|
||||||
|
|||||||
@@ -53,7 +53,8 @@ pub async fn mongo_update_document(
|
|||||||
id: String,
|
id: String,
|
||||||
doc_json: String,
|
doc_json: String,
|
||||||
) -> Result<u64, String> {
|
) -> Result<u64, String> {
|
||||||
dbx_core::mongo_ops::mongo_update_document_core(&state, &connection_id, &database, &collection, &id, &doc_json).await
|
dbx_core::mongo_ops::mongo_update_document_core(&state, &connection_id, &database, &collection, &id, &doc_json)
|
||||||
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
|
|||||||
@@ -16,20 +16,11 @@ pub async fn execute_query(
|
|||||||
sql: String,
|
sql: String,
|
||||||
execution_id: Option<String>,
|
execution_id: Option<String>,
|
||||||
) -> Result<db::QueryResult, String> {
|
) -> Result<db::QueryResult, String> {
|
||||||
let registered_query = execution_id
|
let registered_query =
|
||||||
.as_ref()
|
execution_id.as_ref().filter(|id| !id.trim().is_empty()).map(|id| state.running_queries.register(id.clone()));
|
||||||
.filter(|id| !id.trim().is_empty())
|
|
||||||
.map(|id| state.running_queries.register(id.clone()));
|
|
||||||
let cancel_token = registered_query.as_ref().map(|query| query.token());
|
let cancel_token = registered_query.as_ref().map(|query| query.token());
|
||||||
|
|
||||||
dbx_core::query::execute_sql_statement(
|
dbx_core::query::execute_sql_statement(&state, &connection_id, &database, &sql, cancel_token).await
|
||||||
&state,
|
|
||||||
&connection_id,
|
|
||||||
&database,
|
|
||||||
&sql,
|
|
||||||
cancel_token,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
@@ -40,27 +31,15 @@ pub async fn execute_multi(
|
|||||||
sql: String,
|
sql: String,
|
||||||
execution_id: Option<String>,
|
execution_id: Option<String>,
|
||||||
) -> Result<Vec<db::QueryResult>, String> {
|
) -> Result<Vec<db::QueryResult>, String> {
|
||||||
let registered_query = execution_id
|
let registered_query =
|
||||||
.as_ref()
|
execution_id.as_ref().filter(|id| !id.trim().is_empty()).map(|id| state.running_queries.register(id.clone()));
|
||||||
.filter(|id| !id.trim().is_empty())
|
|
||||||
.map(|id| state.running_queries.register(id.clone()));
|
|
||||||
let cancel_token = registered_query.as_ref().map(|query| query.token());
|
let cancel_token = registered_query.as_ref().map(|query| query.token());
|
||||||
|
|
||||||
dbx_core::query::execute_multi_core(
|
dbx_core::query::execute_multi_core(&state, &connection_id, &database, &sql, cancel_token).await
|
||||||
&state,
|
|
||||||
&connection_id,
|
|
||||||
&database,
|
|
||||||
&sql,
|
|
||||||
cancel_token,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn cancel_query(
|
pub async fn cancel_query(state: State<'_, Arc<AppState>>, execution_id: String) -> Result<bool, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
execution_id: String,
|
|
||||||
) -> Result<bool, String> {
|
|
||||||
Ok(state.running_queries.cancel(&execution_id))
|
Ok(state.running_queries.cancel(&execution_id))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,10 +5,7 @@ use crate::commands::connection::AppState;
|
|||||||
use dbx_core::db::redis_driver::{RedisScanResult, RedisValue};
|
use dbx_core::db::redis_driver::{RedisScanResult, RedisValue};
|
||||||
|
|
||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn redis_list_databases(
|
pub async fn redis_list_databases(state: State<'_, Arc<AppState>>, connection_id: String) -> Result<Vec<u32>, String> {
|
||||||
state: State<'_, Arc<AppState>>,
|
|
||||||
connection_id: String,
|
|
||||||
) -> Result<Vec<u32>, String> {
|
|
||||||
dbx_core::redis_ops::redis_list_databases_core(&state, &connection_id).await
|
dbx_core::redis_ops::redis_list_databases_core(&state, &connection_id).await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -56,7 +53,10 @@ pub async fn redis_delete_key(
|
|||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn redis_hash_set(
|
pub async fn redis_hash_set(
|
||||||
state: State<'_, Arc<AppState>>,
|
state: State<'_, Arc<AppState>>,
|
||||||
connection_id: String, key: String, field: String, value: String,
|
connection_id: String,
|
||||||
|
key: String,
|
||||||
|
field: String,
|
||||||
|
value: String,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
dbx_core::redis_ops::redis_hash_set_core(&state, &connection_id, &key, &field, &value).await
|
dbx_core::redis_ops::redis_hash_set_core(&state, &connection_id, &key, &field, &value).await
|
||||||
}
|
}
|
||||||
@@ -64,7 +64,9 @@ pub async fn redis_hash_set(
|
|||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn redis_hash_del(
|
pub async fn redis_hash_del(
|
||||||
state: State<'_, Arc<AppState>>,
|
state: State<'_, Arc<AppState>>,
|
||||||
connection_id: String, key: String, field: String,
|
connection_id: String,
|
||||||
|
key: String,
|
||||||
|
field: String,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
dbx_core::redis_ops::redis_hash_del_core(&state, &connection_id, &key, &field).await
|
dbx_core::redis_ops::redis_hash_del_core(&state, &connection_id, &key, &field).await
|
||||||
}
|
}
|
||||||
@@ -72,7 +74,9 @@ pub async fn redis_hash_del(
|
|||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn redis_list_push(
|
pub async fn redis_list_push(
|
||||||
state: State<'_, Arc<AppState>>,
|
state: State<'_, Arc<AppState>>,
|
||||||
connection_id: String, key: String, value: String,
|
connection_id: String,
|
||||||
|
key: String,
|
||||||
|
value: String,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
dbx_core::redis_ops::redis_list_push_core(&state, &connection_id, &key, &value).await
|
dbx_core::redis_ops::redis_list_push_core(&state, &connection_id, &key, &value).await
|
||||||
}
|
}
|
||||||
@@ -80,7 +84,9 @@ pub async fn redis_list_push(
|
|||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn redis_list_remove(
|
pub async fn redis_list_remove(
|
||||||
state: State<'_, Arc<AppState>>,
|
state: State<'_, Arc<AppState>>,
|
||||||
connection_id: String, key: String, index: i64,
|
connection_id: String,
|
||||||
|
key: String,
|
||||||
|
index: i64,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
dbx_core::redis_ops::redis_list_remove_core(&state, &connection_id, &key, index).await
|
dbx_core::redis_ops::redis_list_remove_core(&state, &connection_id, &key, index).await
|
||||||
}
|
}
|
||||||
@@ -88,7 +94,9 @@ pub async fn redis_list_remove(
|
|||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn redis_set_add(
|
pub async fn redis_set_add(
|
||||||
state: State<'_, Arc<AppState>>,
|
state: State<'_, Arc<AppState>>,
|
||||||
connection_id: String, key: String, member: String,
|
connection_id: String,
|
||||||
|
key: String,
|
||||||
|
member: String,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
dbx_core::redis_ops::redis_set_add_core(&state, &connection_id, &key, &member).await
|
dbx_core::redis_ops::redis_set_add_core(&state, &connection_id, &key, &member).await
|
||||||
}
|
}
|
||||||
@@ -96,7 +104,9 @@ pub async fn redis_set_add(
|
|||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn redis_set_remove(
|
pub async fn redis_set_remove(
|
||||||
state: State<'_, Arc<AppState>>,
|
state: State<'_, Arc<AppState>>,
|
||||||
connection_id: String, key: String, member: String,
|
connection_id: String,
|
||||||
|
key: String,
|
||||||
|
member: String,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
dbx_core::redis_ops::redis_set_remove_core(&state, &connection_id, &key, &member).await
|
dbx_core::redis_ops::redis_set_remove_core(&state, &connection_id, &key, &member).await
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,8 +12,7 @@ use crate::commands::connection::AppState;
|
|||||||
use crate::commands::query::execute_sql_statement;
|
use crate::commands::query::execute_sql_statement;
|
||||||
|
|
||||||
pub use dbx_core::sql::{
|
pub use dbx_core::sql::{
|
||||||
statement_summary, SqlFilePreview, SqlFileProgress, SqlFileRequest,
|
statement_summary, SqlFilePreview, SqlFileProgress, SqlFileRequest, SqlFileStatus, SqlStatementSplitter,
|
||||||
SqlFileStatus, SqlStatementSplitter,
|
|
||||||
};
|
};
|
||||||
|
|
||||||
static SQL_FILE_EXECUTIONS: std::sync::LazyLock<RwLock<HashMap<String, CancellationToken>>> =
|
static SQL_FILE_EXECUTIONS: std::sync::LazyLock<RwLock<HashMap<String, CancellationToken>>> =
|
||||||
@@ -38,25 +37,15 @@ struct SqlFileSummary {
|
|||||||
#[tauri::command]
|
#[tauri::command]
|
||||||
pub async fn preview_sql_file(file_path: String) -> Result<SqlFilePreview, String> {
|
pub async fn preview_sql_file(file_path: String) -> Result<SqlFilePreview, String> {
|
||||||
let path = PathBuf::from(&file_path);
|
let path = PathBuf::from(&file_path);
|
||||||
let metadata = tokio::fs::metadata(&path)
|
let metadata = tokio::fs::metadata(&path).await.map_err(|e| e.to_string())?;
|
||||||
.await
|
let mut file = tokio::fs::File::open(&path).await.map_err(|e| e.to_string())?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
let mut file = tokio::fs::File::open(&path)
|
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
let mut buffer = vec![0; 4096];
|
let mut buffer = vec![0; 4096];
|
||||||
let bytes_read = tokio::io::AsyncReadExt::read(&mut file, &mut buffer)
|
let bytes_read = tokio::io::AsyncReadExt::read(&mut file, &mut buffer).await.map_err(|e| e.to_string())?;
|
||||||
.await
|
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
buffer.truncate(bytes_read);
|
buffer.truncate(bytes_read);
|
||||||
let preview = String::from_utf8_lossy(&buffer).to_string();
|
let preview = String::from_utf8_lossy(&buffer).to_string();
|
||||||
|
|
||||||
Ok(SqlFilePreview {
|
Ok(SqlFilePreview {
|
||||||
file_name: path
|
file_name: path.file_name().and_then(|name| name.to_str()).unwrap_or("script.sql").to_string(),
|
||||||
.file_name()
|
|
||||||
.and_then(|name| name.to_str())
|
|
||||||
.unwrap_or("script.sql")
|
|
||||||
.to_string(),
|
|
||||||
file_path,
|
file_path,
|
||||||
size_bytes: metadata.len(),
|
size_bytes: metadata.len(),
|
||||||
preview,
|
preview,
|
||||||
@@ -76,18 +65,7 @@ pub async fn execute_sql_file(
|
|||||||
}
|
}
|
||||||
|
|
||||||
let started_at = Instant::now();
|
let started_at = Instant::now();
|
||||||
emit_progress(
|
emit_progress(&app, &request.execution_id, SqlFileStatus::Started, 0, 0, 0, 0, started_at, "", None);
|
||||||
&app,
|
|
||||||
&request.execution_id,
|
|
||||||
SqlFileStatus::Started,
|
|
||||||
0,
|
|
||||||
0,
|
|
||||||
0,
|
|
||||||
0,
|
|
||||||
started_at,
|
|
||||||
"",
|
|
||||||
None,
|
|
||||||
);
|
|
||||||
|
|
||||||
let result = execute_sql_file_inner(&app, &state, &request, token, started_at).await;
|
let result = execute_sql_file_inner(&app, &state, &request, token, started_at).await;
|
||||||
{
|
{
|
||||||
@@ -334,19 +312,14 @@ fn register_sql_file_execution(
|
|||||||
token: CancellationToken,
|
token: CancellationToken,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
if executions.contains_key(&execution_id) {
|
if executions.contains_key(&execution_id) {
|
||||||
return Err(format!(
|
return Err(format!("SQL file execution '{execution_id}' already exists"));
|
||||||
"SQL file execution '{execution_id}' already exists"
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
executions.insert(execution_id, token);
|
executions.insert(execution_id, token);
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
fn remove_sql_file_execution(
|
fn remove_sql_file_execution(executions: &mut HashMap<String, CancellationToken>, execution_id: &str) {
|
||||||
executions: &mut HashMap<String, CancellationToken>,
|
|
||||||
execution_id: &str,
|
|
||||||
) {
|
|
||||||
executions.remove(execution_id);
|
executions.remove(execution_id);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -394,11 +367,7 @@ fn statement_error_decision(
|
|||||||
);
|
);
|
||||||
|
|
||||||
if continue_on_error {
|
if continue_on_error {
|
||||||
return StatementErrorDecision {
|
return StatementErrorDecision { progress: vec![statement_failed], failure_count, result: Ok(false) };
|
||||||
progress: vec![statement_failed],
|
|
||||||
failure_count,
|
|
||||||
result: Ok(false),
|
|
||||||
};
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let terminal_error = sql_file_progress(
|
let terminal_error = sql_file_progress(
|
||||||
@@ -413,11 +382,7 @@ fn statement_error_decision(
|
|||||||
Some(error.clone()),
|
Some(error.clone()),
|
||||||
);
|
);
|
||||||
|
|
||||||
StatementErrorDecision {
|
StatementErrorDecision { progress: vec![statement_failed, terminal_error], failure_count, result: Err(error) }
|
||||||
progress: vec![statement_failed, terminal_error],
|
|
||||||
failure_count,
|
|
||||||
result: Err(error),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn emit_progress(
|
fn emit_progress(
|
||||||
@@ -559,11 +524,7 @@ async fn run_statements_for_test(
|
|||||||
}
|
}
|
||||||
|
|
||||||
SqlFileSummary {
|
SqlFileSummary {
|
||||||
status: if token.is_cancelled() {
|
status: if token.is_cancelled() { SqlFileStatus::Cancelled } else { SqlFileStatus::Done },
|
||||||
SqlFileStatus::Cancelled
|
|
||||||
} else {
|
|
||||||
SqlFileStatus::Done
|
|
||||||
},
|
|
||||||
success_count,
|
success_count,
|
||||||
failure_count,
|
failure_count,
|
||||||
failed_statement_index,
|
failed_statement_index,
|
||||||
@@ -586,12 +547,7 @@ mod execution_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn stops_on_first_failure_by_default() {
|
async fn stops_on_first_failure_by_default() {
|
||||||
let summary = run_fake_script(
|
let summary = run_fake_script(vec!["ok 1".into(), "fail 2".into(), "ok 3".into()], false, None).await;
|
||||||
vec!["ok 1".into(), "fail 2".into(), "ok 3".into()],
|
|
||||||
false,
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
assert_eq!(summary.success_count, 1);
|
assert_eq!(summary.success_count, 1);
|
||||||
assert_eq!(summary.failure_count, 1);
|
assert_eq!(summary.failure_count, 1);
|
||||||
@@ -601,12 +557,7 @@ mod execution_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn continues_after_failure_when_enabled() {
|
async fn continues_after_failure_when_enabled() {
|
||||||
let summary = run_fake_script(
|
let summary = run_fake_script(vec!["ok 1".into(), "fail 2".into(), "ok 3".into()], true, None).await;
|
||||||
vec!["ok 1".into(), "fail 2".into(), "ok 3".into()],
|
|
||||||
true,
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
assert_eq!(summary.success_count, 2);
|
assert_eq!(summary.success_count, 2);
|
||||||
assert_eq!(summary.failure_count, 1);
|
assert_eq!(summary.failure_count, 1);
|
||||||
@@ -615,12 +566,7 @@ mod execution_tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn cancellation_stops_before_next_statement() {
|
async fn cancellation_stops_before_next_statement() {
|
||||||
let summary = run_fake_script(
|
let summary = run_fake_script(vec!["ok 1".into(), "ok 2".into(), "ok 3".into()], true, Some(1)).await;
|
||||||
vec!["ok 1".into(), "ok 2".into(), "ok 3".into()],
|
|
||||||
true,
|
|
||||||
Some(1),
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
assert_eq!(summary.success_count, 1);
|
assert_eq!(summary.success_count, 1);
|
||||||
assert_eq!(summary.status, SqlFileStatus::Cancelled);
|
assert_eq!(summary.status, SqlFileStatus::Cancelled);
|
||||||
@@ -628,15 +574,7 @@ mod execution_tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn file_io_errors_build_terminal_error_progress() {
|
fn file_io_errors_build_terminal_error_progress() {
|
||||||
let progress = file_io_error_progress(
|
let progress = file_io_error_progress("exec-1", 4, 2, 1, 17, Instant::now(), "read failed".to_string());
|
||||||
"exec-1",
|
|
||||||
4,
|
|
||||||
2,
|
|
||||||
1,
|
|
||||||
17,
|
|
||||||
Instant::now(),
|
|
||||||
"read failed".to_string(),
|
|
||||||
);
|
|
||||||
|
|
||||||
assert_eq!(progress.execution_id, "exec-1");
|
assert_eq!(progress.execution_id, "exec-1");
|
||||||
assert_eq!(progress.status, SqlFileStatus::Error);
|
assert_eq!(progress.status, SqlFileStatus::Error);
|
||||||
@@ -655,13 +593,9 @@ mod execution_tests {
|
|||||||
let replacement = CancellationToken::new();
|
let replacement = CancellationToken::new();
|
||||||
executions.insert("dup".to_string(), original.clone());
|
executions.insert("dup".to_string(), original.clone());
|
||||||
|
|
||||||
let result =
|
let result = register_sql_file_execution(&mut executions, "dup".to_string(), replacement.clone());
|
||||||
register_sql_file_execution(&mut executions, "dup".to_string(), replacement.clone());
|
|
||||||
|
|
||||||
assert_eq!(
|
assert_eq!(result.unwrap_err(), "SQL file execution 'dup' already exists");
|
||||||
result.unwrap_err(),
|
|
||||||
"SQL file execution 'dup' already exists"
|
|
||||||
);
|
|
||||||
assert_eq!(executions.len(), 1);
|
assert_eq!(executions.len(), 1);
|
||||||
|
|
||||||
executions.get("dup").unwrap().cancel();
|
executions.get("dup").unwrap().cancel();
|
||||||
@@ -720,17 +654,8 @@ mod execution_tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn progress_payload_serializes_camel_case_status() {
|
fn progress_payload_serializes_camel_case_status() {
|
||||||
let progress = sql_file_progress(
|
let progress =
|
||||||
"exec-1",
|
sql_file_progress("exec-1", SqlFileStatus::StatementDone, 1, 1, 0, 3, Instant::now(), "select 1", None);
|
||||||
SqlFileStatus::StatementDone,
|
|
||||||
1,
|
|
||||||
1,
|
|
||||||
0,
|
|
||||||
3,
|
|
||||||
Instant::now(),
|
|
||||||
"select 1",
|
|
||||||
None,
|
|
||||||
);
|
|
||||||
|
|
||||||
let value = serde_json::to_value(progress).unwrap();
|
let value = serde_json::to_value(progress).unwrap();
|
||||||
|
|
||||||
|
|||||||
@@ -8,10 +8,7 @@ use crate::commands::connection::AppState;
|
|||||||
use crate::commands::transfer::get_db_type;
|
use crate::commands::transfer::get_db_type;
|
||||||
|
|
||||||
// Re-export types for backward compatibility
|
// Re-export types for backward compatibility
|
||||||
pub use dbx_core::table_import::{
|
pub use dbx_core::table_import::{TableImportPreview, TableImportProgress, TableImportRequest, TableImportSummary};
|
||||||
TableImportPreview, TableImportProgress,
|
|
||||||
TableImportRequest, TableImportSummary,
|
|
||||||
};
|
|
||||||
|
|
||||||
static CANCELLED_IMPORTS: std::sync::LazyLock<RwLock<HashSet<String>>> =
|
static CANCELLED_IMPORTS: std::sync::LazyLock<RwLock<HashSet<String>>> =
|
||||||
std::sync::LazyLock::new(|| RwLock::new(HashSet::new()));
|
std::sync::LazyLock::new(|| RwLock::new(HashSet::new()));
|
||||||
@@ -44,9 +41,7 @@ pub async fn import_table_file(
|
|||||||
let pool_key = if request.database.is_empty() {
|
let pool_key = if request.database.is_empty() {
|
||||||
request.connection_id.clone()
|
request.connection_id.clone()
|
||||||
} else {
|
} else {
|
||||||
state
|
state.get_or_create_pool(&request.connection_id, Some(&request.database)).await?
|
||||||
.get_or_create_pool(&request.connection_id, Some(&request.database))
|
|
||||||
.await?
|
|
||||||
};
|
};
|
||||||
|
|
||||||
let result = dbx_core::table_import::import_table_file_core(
|
let result = dbx_core::table_import::import_table_file_core(
|
||||||
|
|||||||
@@ -4,10 +4,7 @@ use tauri::{AppHandle, Emitter, State};
|
|||||||
use crate::commands::connection::AppState;
|
use crate::commands::connection::AppState;
|
||||||
|
|
||||||
// Re-export types and functions used by other modules
|
// Re-export types and functions used by other modules
|
||||||
pub use dbx_core::transfer::{
|
pub use dbx_core::transfer::{get_db_type, TransferProgress, TransferRequest, TransferStatus};
|
||||||
get_db_type,
|
|
||||||
TransferProgress, TransferRequest, TransferStatus,
|
|
||||||
};
|
|
||||||
|
|
||||||
fn emit_progress(app: &AppHandle, progress: TransferProgress) {
|
fn emit_progress(app: &AppHandle, progress: TransferProgress) {
|
||||||
let _ = app.emit("transfer-progress", progress);
|
let _ = app.emit("transfer-progress", progress);
|
||||||
@@ -27,12 +24,10 @@ pub async fn start_transfer(
|
|||||||
let target_db_type = get_db_type(&state, &request.target_connection_id).await?;
|
let target_db_type = get_db_type(&state, &request.target_connection_id).await?;
|
||||||
|
|
||||||
// Ensure pools
|
// Ensure pools
|
||||||
let source_pool_key = state
|
let source_pool_key =
|
||||||
.get_or_create_pool(&request.source_connection_id, Some(&request.source_database))
|
state.get_or_create_pool(&request.source_connection_id, Some(&request.source_database)).await?;
|
||||||
.await?;
|
let target_pool_key =
|
||||||
let target_pool_key = state
|
state.get_or_create_pool(&request.target_connection_id, Some(&request.target_database)).await?;
|
||||||
.get_or_create_pool(&request.target_connection_id, Some(&request.target_database))
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
let total_tables = request.tables.len();
|
let total_tables = request.tables.len();
|
||||||
@@ -40,16 +35,19 @@ pub async fn start_transfer(
|
|||||||
|
|
||||||
for (i, table) in request.tables.iter().enumerate() {
|
for (i, table) in request.tables.iter().enumerate() {
|
||||||
if dbx_core::transfer::is_cancelled(&transfer_id).await {
|
if dbx_core::transfer::is_cancelled(&transfer_id).await {
|
||||||
emit_progress(&app, TransferProgress {
|
emit_progress(
|
||||||
transfer_id: transfer_id.clone(),
|
&app,
|
||||||
table: table.clone(),
|
TransferProgress {
|
||||||
table_index: i,
|
transfer_id: transfer_id.clone(),
|
||||||
total_tables,
|
table: table.clone(),
|
||||||
rows_transferred: 0,
|
table_index: i,
|
||||||
total_rows: None,
|
total_tables,
|
||||||
status: TransferStatus::Cancelled,
|
rows_transferred: 0,
|
||||||
error: None,
|
total_rows: None,
|
||||||
});
|
status: TransferStatus::Cancelled,
|
||||||
|
error: None,
|
||||||
|
},
|
||||||
|
);
|
||||||
dbx_core::transfer::clear_cancelled(&transfer_id).await;
|
dbx_core::transfer::clear_cancelled(&transfer_id).await;
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -57,48 +55,68 @@ pub async fn start_transfer(
|
|||||||
log::info!("[transfer] table {}/{}: {}", i + 1, total_tables, table);
|
log::info!("[transfer] table {}/{}: {}", i + 1, total_tables, table);
|
||||||
|
|
||||||
match dbx_core::transfer::transfer_table(
|
match dbx_core::transfer::transfer_table(
|
||||||
&state, &request, table, i,
|
&state,
|
||||||
&source_db_type, &target_db_type,
|
&request,
|
||||||
&source_pool_key, &target_pool_key,
|
table,
|
||||||
|
i,
|
||||||
|
&source_db_type,
|
||||||
|
&target_db_type,
|
||||||
|
&source_pool_key,
|
||||||
|
&target_pool_key,
|
||||||
|progress| emit_progress(&app, progress),
|
|progress| emit_progress(&app, progress),
|
||||||
).await {
|
)
|
||||||
|
.await
|
||||||
|
{
|
||||||
Ok(rows) => {
|
Ok(rows) => {
|
||||||
emit_progress(&app, TransferProgress {
|
emit_progress(
|
||||||
transfer_id: transfer_id.clone(),
|
&app,
|
||||||
table: table.clone(),
|
TransferProgress {
|
||||||
table_index: i,
|
transfer_id: transfer_id.clone(),
|
||||||
total_tables,
|
table: table.clone(),
|
||||||
rows_transferred: rows,
|
table_index: i,
|
||||||
total_rows: Some(rows),
|
total_tables,
|
||||||
status: if i == total_tables - 1 { TransferStatus::Done } else { TransferStatus::TableDone },
|
rows_transferred: rows,
|
||||||
error: None,
|
total_rows: Some(rows),
|
||||||
});
|
status: if i == total_tables - 1 {
|
||||||
|
TransferStatus::Done
|
||||||
|
} else {
|
||||||
|
TransferStatus::TableDone
|
||||||
|
},
|
||||||
|
error: None,
|
||||||
|
},
|
||||||
|
);
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
if e == "Cancelled" {
|
if e == "Cancelled" {
|
||||||
emit_progress(&app, TransferProgress {
|
emit_progress(
|
||||||
|
&app,
|
||||||
|
TransferProgress {
|
||||||
|
transfer_id: transfer_id.clone(),
|
||||||
|
table: table.clone(),
|
||||||
|
table_index: i,
|
||||||
|
total_tables,
|
||||||
|
rows_transferred: 0,
|
||||||
|
total_rows: None,
|
||||||
|
status: TransferStatus::Cancelled,
|
||||||
|
error: None,
|
||||||
|
},
|
||||||
|
);
|
||||||
|
dbx_core::transfer::clear_cancelled(&transfer_id).await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
emit_progress(
|
||||||
|
&app,
|
||||||
|
TransferProgress {
|
||||||
transfer_id: transfer_id.clone(),
|
transfer_id: transfer_id.clone(),
|
||||||
table: table.clone(),
|
table: table.clone(),
|
||||||
table_index: i,
|
table_index: i,
|
||||||
total_tables,
|
total_tables,
|
||||||
rows_transferred: 0,
|
rows_transferred: 0,
|
||||||
total_rows: None,
|
total_rows: None,
|
||||||
status: TransferStatus::Cancelled,
|
status: TransferStatus::Error,
|
||||||
error: None,
|
error: Some(e),
|
||||||
});
|
},
|
||||||
dbx_core::transfer::clear_cancelled(&transfer_id).await;
|
);
|
||||||
return;
|
|
||||||
}
|
|
||||||
emit_progress(&app, TransferProgress {
|
|
||||||
transfer_id: transfer_id.clone(),
|
|
||||||
table: table.clone(),
|
|
||||||
table_index: i,
|
|
||||||
total_tables,
|
|
||||||
rows_transferred: 0,
|
|
||||||
total_rows: None,
|
|
||||||
status: TransferStatus::Error,
|
|
||||||
error: Some(e),
|
|
||||||
});
|
|
||||||
dbx_core::transfer::clear_cancelled(&transfer_id).await;
|
dbx_core::transfer::clear_cancelled(&transfer_id).await;
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|||||||
+5
-16
@@ -9,9 +9,7 @@ use tauri::{Manager, RunEvent};
|
|||||||
|
|
||||||
#[cfg_attr(mobile, tauri::mobile_entry_point)]
|
#[cfg_attr(mobile, tauri::mobile_entry_point)]
|
||||||
pub fn run() {
|
pub fn run() {
|
||||||
rustls::crypto::aws_lc_rs::default_provider()
|
rustls::crypto::aws_lc_rs::default_provider().install_default().expect("Failed to install rustls crypto provider");
|
||||||
.install_default()
|
|
||||||
.expect("Failed to install rustls crypto provider");
|
|
||||||
|
|
||||||
tauri::Builder::default()
|
tauri::Builder::default()
|
||||||
.plugin(tauri_plugin_dialog::init())
|
.plugin(tauri_plugin_dialog::init())
|
||||||
@@ -21,26 +19,17 @@ pub fn run() {
|
|||||||
.plugin(tauri_plugin_process::init())
|
.plugin(tauri_plugin_process::init())
|
||||||
.setup(|app| {
|
.setup(|app| {
|
||||||
if cfg!(debug_assertions) {
|
if cfg!(debug_assertions) {
|
||||||
app.handle().plugin(
|
app.handle().plugin(tauri_plugin_log::Builder::default().level(log::LevelFilter::Info).build())?;
|
||||||
tauri_plugin_log::Builder::default()
|
|
||||||
.level(log::LevelFilter::Info)
|
|
||||||
.build(),
|
|
||||||
)?;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let data_dir = app
|
let data_dir =
|
||||||
.path()
|
app.path().app_data_dir().map_err(|e| e.to_string()).expect("Failed to resolve app data dir");
|
||||||
.app_data_dir()
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
.expect("Failed to resolve app data dir");
|
|
||||||
std::fs::create_dir_all(&data_dir).expect("Failed to create data dir");
|
std::fs::create_dir_all(&data_dir).expect("Failed to create data dir");
|
||||||
let db_path = data_dir.join("dbx.db");
|
let db_path = data_dir.join("dbx.db");
|
||||||
|
|
||||||
let storage = tauri::async_runtime::block_on(async {
|
let storage = tauri::async_runtime::block_on(async {
|
||||||
let s = Storage::open(&db_path).await.expect("Failed to open storage");
|
let s = Storage::open(&db_path).await.expect("Failed to open storage");
|
||||||
s.migrate_from_json(&data_dir)
|
s.migrate_from_json(&data_dir).await.expect("Failed to migrate JSON data");
|
||||||
.await
|
|
||||||
.expect("Failed to migrate JSON data");
|
|
||||||
s
|
s
|
||||||
});
|
});
|
||||||
|
|
||||||
|
|||||||
@@ -2,5 +2,5 @@
|
|||||||
#![cfg_attr(not(debug_assertions), windows_subsystem = "windows")]
|
#![cfg_attr(not(debug_assertions), windows_subsystem = "windows")]
|
||||||
|
|
||||||
fn main() {
|
fn main() {
|
||||||
dbx_lib::run();
|
dbx_lib::run();
|
||||||
}
|
}
|
||||||
|
|||||||
+6
-22
@@ -24,18 +24,11 @@ pub struct AuthCheckResponse {
|
|||||||
const MAX_ATTEMPTS: u32 = 5;
|
const MAX_ATTEMPTS: u32 = 5;
|
||||||
const LOCKOUT_SECS: u64 = 60;
|
const LOCKOUT_SECS: u64 = 60;
|
||||||
|
|
||||||
pub async fn login(
|
pub async fn login(State(state): State<Arc<WebState>>, Json(body): Json<LoginRequest>) -> Result<Response, StatusCode> {
|
||||||
State(state): State<Arc<WebState>>,
|
|
||||||
Json(body): Json<LoginRequest>,
|
|
||||||
) -> Result<Response, StatusCode> {
|
|
||||||
let password_hash = match &state.password_hash {
|
let password_hash = match &state.password_hash {
|
||||||
Some(h) => h,
|
Some(h) => h,
|
||||||
None => {
|
None => {
|
||||||
return Ok((
|
return Ok((StatusCode::OK, Json(serde_json::json!({"ok": true}))).into_response());
|
||||||
StatusCode::OK,
|
|
||||||
Json(serde_json::json!({"ok": true})),
|
|
||||||
)
|
|
||||||
.into_response());
|
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -48,7 +41,8 @@ pub async fn login(
|
|||||||
return Ok((
|
return Ok((
|
||||||
StatusCode::TOO_MANY_REQUESTS,
|
StatusCode::TOO_MANY_REQUESTS,
|
||||||
Json(serde_json::json!({"error": format!("请 {remaining} 秒后再试")})),
|
Json(serde_json::json!({"error": format!("请 {remaining} 秒后再试")})),
|
||||||
).into_response());
|
)
|
||||||
|
.into_response());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -78,12 +72,7 @@ pub async fn login(
|
|||||||
state.sessions.write().await.insert(token.clone());
|
state.sessions.write().await.insert(token.clone());
|
||||||
|
|
||||||
let cookie = format!("dbx_session={token}; Path=/; HttpOnly; SameSite=Lax");
|
let cookie = format!("dbx_session={token}; Path=/; HttpOnly; SameSite=Lax");
|
||||||
Ok((
|
Ok((StatusCode::OK, [("set-cookie", cookie.as_str())], Json(serde_json::json!({"ok": true}))).into_response())
|
||||||
StatusCode::OK,
|
|
||||||
[("set-cookie", cookie.as_str())],
|
|
||||||
Json(serde_json::json!({"ok": true})),
|
|
||||||
)
|
|
||||||
.into_response())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn check(State(state): State<Arc<WebState>>, req: Request<axum::body::Body>) -> Json<AuthCheckResponse> {
|
pub async fn check(State(state): State<Arc<WebState>>, req: Request<axum::body::Body>) -> Json<AuthCheckResponse> {
|
||||||
@@ -102,12 +91,7 @@ pub async fn logout(State(state): State<Arc<WebState>>, req: Request<axum::body:
|
|||||||
state.sessions.write().await.remove(&token);
|
state.sessions.write().await.remove(&token);
|
||||||
}
|
}
|
||||||
let cookie = "dbx_session=; Path=/; HttpOnly; Max-Age=0";
|
let cookie = "dbx_session=; Path=/; HttpOnly; Max-Age=0";
|
||||||
(
|
(StatusCode::OK, [("set-cookie", cookie)], Json(serde_json::json!({"ok": true}))).into_response()
|
||||||
StatusCode::OK,
|
|
||||||
[("set-cookie", cookie)],
|
|
||||||
Json(serde_json::json!({"ok": true})),
|
|
||||||
)
|
|
||||||
.into_response()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn extract_session_token<B>(req: &Request<B>) -> Option<String> {
|
fn extract_session_token<B>(req: &Request<B>) -> Option<String> {
|
||||||
|
|||||||
+14
-40
@@ -28,28 +28,19 @@ async fn main() {
|
|||||||
)
|
)
|
||||||
.init();
|
.init();
|
||||||
|
|
||||||
rustls::crypto::aws_lc_rs::default_provider()
|
rustls::crypto::aws_lc_rs::default_provider().install_default().expect("Failed to install rustls crypto provider");
|
||||||
.install_default()
|
|
||||||
.expect("Failed to install rustls crypto provider");
|
|
||||||
|
|
||||||
// Data directory
|
// Data directory
|
||||||
let data_dir = std::env::var("DBX_DATA_DIR")
|
let data_dir = std::env::var("DBX_DATA_DIR").map(std::path::PathBuf::from).unwrap_or_else(|_| {
|
||||||
.map(std::path::PathBuf::from)
|
let home = std::env::var("HOME").unwrap_or_else(|_| ".".to_string());
|
||||||
.unwrap_or_else(|_| {
|
std::path::PathBuf::from(home).join(".dbx-web")
|
||||||
let home = std::env::var("HOME").unwrap_or_else(|_| ".".to_string());
|
});
|
||||||
std::path::PathBuf::from(home).join(".dbx-web")
|
|
||||||
});
|
|
||||||
std::fs::create_dir_all(&data_dir).expect("Failed to create data directory");
|
std::fs::create_dir_all(&data_dir).expect("Failed to create data directory");
|
||||||
|
|
||||||
let app_state = {
|
let app_state = {
|
||||||
let db_path = data_dir.join("dbx.db");
|
let db_path = data_dir.join("dbx.db");
|
||||||
let storage = Storage::open(&db_path)
|
let storage = Storage::open(&db_path).await.expect("Failed to open storage");
|
||||||
.await
|
storage.migrate_from_json(&data_dir).await.expect("Failed to migrate JSON data");
|
||||||
.expect("Failed to open storage");
|
|
||||||
storage
|
|
||||||
.migrate_from_json(&data_dir)
|
|
||||||
.await
|
|
||||||
.expect("Failed to migrate JSON data");
|
|
||||||
Arc::new(AppState::new(storage))
|
Arc::new(AppState::new(storage))
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -66,17 +57,11 @@ async fn main() {
|
|||||||
password_hash,
|
password_hash,
|
||||||
sessions: RwLock::new(HashSet::new()),
|
sessions: RwLock::new(HashSet::new()),
|
||||||
sse_channels: RwLock::new(HashMap::new()),
|
sse_channels: RwLock::new(HashMap::new()),
|
||||||
login_rate_limit: tokio::sync::Mutex::new(state::LoginRateLimit {
|
login_rate_limit: tokio::sync::Mutex::new(state::LoginRateLimit { fail_count: 0, locked_until: None }),
|
||||||
fail_count: 0,
|
|
||||||
locked_until: None,
|
|
||||||
}),
|
|
||||||
});
|
});
|
||||||
|
|
||||||
// CORS
|
// CORS
|
||||||
let cors = CorsLayer::new()
|
let cors = CorsLayer::new().allow_origin(Any).allow_methods(Any).allow_headers(Any);
|
||||||
.allow_origin(Any)
|
|
||||||
.allow_methods(Any)
|
|
||||||
.allow_headers(Any);
|
|
||||||
|
|
||||||
// API routes
|
// API routes
|
||||||
let api = Router::new()
|
let api = Router::new()
|
||||||
@@ -159,25 +144,18 @@ async fn main() {
|
|||||||
.with_state(web_state.clone());
|
.with_state(web_state.clone());
|
||||||
|
|
||||||
// Build app
|
// Build app
|
||||||
let mut app = Router::new()
|
let mut app = Router::new().nest("/api", api).layer(tower_http::trace::TraceLayer::new_for_http()).layer(cors);
|
||||||
.nest("/api", api)
|
|
||||||
.layer(tower_http::trace::TraceLayer::new_for_http())
|
|
||||||
.layer(cors);
|
|
||||||
|
|
||||||
// Static file serving
|
// Static file serving
|
||||||
if let Ok(static_dir) = std::env::var("DBX_STATIC_DIR") {
|
if let Ok(static_dir) = std::env::var("DBX_STATIC_DIR") {
|
||||||
use tower_http::services::{ServeDir, ServeFile};
|
use tower_http::services::{ServeDir, ServeFile};
|
||||||
let index_path = format!("{}/index.html", static_dir);
|
let index_path = format!("{}/index.html", static_dir);
|
||||||
let serve_dir = ServeDir::new(&static_dir)
|
let serve_dir = ServeDir::new(&static_dir).not_found_service(ServeFile::new(&index_path));
|
||||||
.not_found_service(ServeFile::new(&index_path));
|
|
||||||
app = app.fallback_service(serve_dir);
|
app = app.fallback_service(serve_dir);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Bind address
|
// Bind address
|
||||||
let port: u16 = std::env::var("DBX_PORT")
|
let port: u16 = std::env::var("DBX_PORT").ok().and_then(|p| p.parse().ok()).unwrap_or(4224);
|
||||||
.ok()
|
|
||||||
.and_then(|p| p.parse().ok())
|
|
||||||
.unwrap_or(4224);
|
|
||||||
let addr = SocketAddr::from(([0, 0, 0, 0], port));
|
let addr = SocketAddr::from(([0, 0, 0, 0], port));
|
||||||
|
|
||||||
tracing::info!("DBX Web server starting on http://{}", addr);
|
tracing::info!("DBX Web server starting on http://{}", addr);
|
||||||
@@ -185,10 +163,6 @@ async fn main() {
|
|||||||
tracing::info!("Password protection is enabled");
|
tracing::info!("Password protection is enabled");
|
||||||
}
|
}
|
||||||
|
|
||||||
let listener = tokio::net::TcpListener::bind(addr)
|
let listener = tokio::net::TcpListener::bind(addr).await.expect("Failed to bind address");
|
||||||
.await
|
axum::serve(listener, app).await.expect("Server error");
|
||||||
.expect("Failed to bind address");
|
|
||||||
axum::serve(listener, app)
|
|
||||||
.await
|
|
||||||
.expect("Server error");
|
|
||||||
}
|
}
|
||||||
|
|||||||
+15
-60
@@ -6,9 +6,7 @@ use axum::Json;
|
|||||||
use futures::stream::Stream;
|
use futures::stream::Stream;
|
||||||
use serde::Deserialize;
|
use serde::Deserialize;
|
||||||
|
|
||||||
use dbx_core::ai::{
|
use dbx_core::ai::{AiCompletionRequest, AiConfig, AiConversation, AiStreamChunk};
|
||||||
AiCompletionRequest, AiConfig, AiConversation, AiStreamChunk,
|
|
||||||
};
|
|
||||||
|
|
||||||
use crate::error::AppError;
|
use crate::error::AppError;
|
||||||
use crate::state::WebState;
|
use crate::state::WebState;
|
||||||
@@ -62,24 +60,12 @@ pub async fn save_ai_config(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(body): Json<SaveAiConfigRequest>,
|
Json(body): Json<SaveAiConfigRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
state
|
state.app.storage.save_ai_config(&body.config).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.save_ai_config(&body.config)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_ai_config(
|
pub async fn load_ai_config(State(state): State<Arc<WebState>>) -> Result<Json<Option<AiConfig>>, AppError> {
|
||||||
State(state): State<Arc<WebState>>,
|
let config = state.app.storage.load_ai_config().await.map_err(AppError)?;
|
||||||
) -> Result<Json<Option<AiConfig>>, AppError> {
|
|
||||||
let config = state
|
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.load_ai_config()
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(config))
|
Ok(Json(config))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -91,24 +77,12 @@ pub async fn save_ai_conversation(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(body): Json<SaveAiConversationRequest>,
|
Json(body): Json<SaveAiConversationRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
state
|
state.app.storage.save_ai_conversation(&body.conversation).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.save_ai_conversation(&body.conversation)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_ai_conversations(
|
pub async fn load_ai_conversations(State(state): State<Arc<WebState>>) -> Result<Json<Vec<AiConversation>>, AppError> {
|
||||||
State(state): State<Arc<WebState>>,
|
let conversations = state.app.storage.load_ai_conversations().await.map_err(AppError)?;
|
||||||
) -> Result<Json<Vec<AiConversation>>, AppError> {
|
|
||||||
let conversations = state
|
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.load_ai_conversations()
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(conversations))
|
Ok(Json(conversations))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -116,12 +90,7 @@ pub async fn delete_ai_conversation(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
state
|
state.app.storage.delete_ai_conversation(&id).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.delete_ai_conversation(&id)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -129,12 +98,8 @@ pub async fn delete_ai_conversation(
|
|||||||
// AI complete (non-streaming)
|
// AI complete (non-streaming)
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
pub async fn ai_complete(
|
pub async fn ai_complete(Json(body): Json<AiCompleteRequest>) -> Result<Json<String>, AppError> {
|
||||||
Json(body): Json<AiCompleteRequest>,
|
let result = dbx_core::ai::complete(&body.request).await.map_err(AppError)?;
|
||||||
) -> Result<Json<String>, AppError> {
|
|
||||||
let result = dbx_core::ai::complete(&body.request)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(result))
|
Ok(Json(result))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -142,12 +107,8 @@ pub async fn ai_complete(
|
|||||||
// AI test connection
|
// AI test connection
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
pub async fn ai_test_connection(
|
pub async fn ai_test_connection(Json(body): Json<AiTestConnectionRequest>) -> Result<Json<String>, AppError> {
|
||||||
Json(body): Json<AiTestConnectionRequest>,
|
let result = dbx_core::ai::test_connection_core(&body.config).await.map_err(AppError)?;
|
||||||
) -> Result<Json<String>, AppError> {
|
|
||||||
let result = dbx_core::ai::test_connection_core(&body.config)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(result))
|
Ok(Json(result))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -155,9 +116,7 @@ pub async fn ai_test_connection(
|
|||||||
// AI cancel stream
|
// AI cancel stream
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
pub async fn ai_cancel_stream(
|
pub async fn ai_cancel_stream(Json(body): Json<AiCancelStreamRequest>) -> Result<Json<bool>, AppError> {
|
||||||
Json(body): Json<AiCancelStreamRequest>,
|
|
||||||
) -> Result<Json<bool>, AppError> {
|
|
||||||
let result = dbx_core::ai::cancel_stream(&body.session_id).await;
|
let result = dbx_core::ai::cancel_stream(&body.session_id).await;
|
||||||
Ok(Json(result))
|
Ok(Json(result))
|
||||||
}
|
}
|
||||||
@@ -184,12 +143,8 @@ pub async fn ai_stream(
|
|||||||
.await;
|
.await;
|
||||||
|
|
||||||
if let Err(_e) = result {
|
if let Err(_e) = result {
|
||||||
let error_chunk = AiStreamChunk {
|
let error_chunk =
|
||||||
session_id: sid.clone(),
|
AiStreamChunk { session_id: sid.clone(), delta: String::new(), reasoning_delta: None, done: true };
|
||||||
delta: String::new(),
|
|
||||||
reasoning_delta: None,
|
|
||||||
done: true,
|
|
||||||
};
|
|
||||||
let _ = tx.send(serde_json::to_string(&error_chunk).unwrap_or_default());
|
let _ = tx.send(serde_json::to_string(&error_chunk).unwrap_or_default());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -35,24 +35,15 @@ pub async fn test_connection(
|
|||||||
|
|
||||||
// Store config temporarily
|
// Store config temporarily
|
||||||
let temp_id = format!("__test_{}", uuid::Uuid::new_v4());
|
let temp_id = format!("__test_{}", uuid::Uuid::new_v4());
|
||||||
app.configs
|
app.configs.lock().await.insert(temp_id.clone(), config.clone());
|
||||||
.lock()
|
|
||||||
.await
|
|
||||||
.insert(temp_id.clone(), config.clone());
|
|
||||||
|
|
||||||
// Try to connect
|
// Try to connect
|
||||||
let result = app
|
let result = app.get_or_create_pool(&temp_id, config.database.as_deref()).await;
|
||||||
.get_or_create_pool(&temp_id, config.database.as_deref())
|
|
||||||
.await;
|
|
||||||
|
|
||||||
// Clean up any pool keys created for the temporary connection, including
|
// Clean up any pool keys created for the temporary connection, including
|
||||||
// database-scoped keys like "__test_uuid:database".
|
// database-scoped keys like "__test_uuid:database".
|
||||||
let mut connections = app.connections.lock().await;
|
let mut connections = app.connections.lock().await;
|
||||||
let temp_keys: Vec<String> = connections
|
let temp_keys: Vec<String> = connections.keys().filter(|key| key.starts_with(&temp_id)).cloned().collect();
|
||||||
.keys()
|
|
||||||
.filter(|key| key.starts_with(&temp_id))
|
|
||||||
.cloned()
|
|
||||||
.collect();
|
|
||||||
for key in temp_keys {
|
for key in temp_keys {
|
||||||
connections.remove(&key);
|
connections.remove(&key);
|
||||||
}
|
}
|
||||||
@@ -73,15 +64,9 @@ pub async fn connect_db(
|
|||||||
let app = &state.app;
|
let app = &state.app;
|
||||||
let connection_id = config.id.clone();
|
let connection_id = config.id.clone();
|
||||||
|
|
||||||
app.configs
|
app.configs.lock().await.insert(connection_id.clone(), config.clone());
|
||||||
.lock()
|
|
||||||
.await
|
|
||||||
.insert(connection_id.clone(), config.clone());
|
|
||||||
|
|
||||||
let pool_key = app
|
let pool_key = app.get_or_create_pool(&connection_id, config.database.as_deref()).await.map_err(AppError)?;
|
||||||
.get_or_create_pool(&connection_id, config.database.as_deref())
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
|
|
||||||
Ok(Json(pool_key))
|
Ok(Json(pool_key))
|
||||||
}
|
}
|
||||||
@@ -94,11 +79,8 @@ pub async fn disconnect_db(
|
|||||||
let mut connections = app.connections.lock().await;
|
let mut connections = app.connections.lock().await;
|
||||||
|
|
||||||
// Remove all pool keys that start with this connection_id
|
// Remove all pool keys that start with this connection_id
|
||||||
let keys_to_remove: Vec<String> = connections
|
let keys_to_remove: Vec<String> =
|
||||||
.keys()
|
connections.keys().filter(|k| k.starts_with(&body.connection_id)).cloned().collect();
|
||||||
.filter(|k| k.starts_with(&body.connection_id))
|
|
||||||
.cloned()
|
|
||||||
.collect();
|
|
||||||
for key in keys_to_remove {
|
for key in keys_to_remove {
|
||||||
connections.remove(&key);
|
connections.remove(&key);
|
||||||
}
|
}
|
||||||
@@ -114,23 +96,11 @@ pub async fn save_connections(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(body): Json<SaveConnectionsRequest>,
|
Json(body): Json<SaveConnectionsRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
state
|
state.app.storage.save_connections(&body.configs).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.save_connections(&body.configs)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_connections(
|
pub async fn load_connections(State(state): State<Arc<WebState>>) -> Result<Json<Vec<ConnectionConfig>>, AppError> {
|
||||||
State(state): State<Arc<WebState>>,
|
let configs = state.app.storage.load_connections().await.map_err(AppError)?;
|
||||||
) -> Result<Json<Vec<ConnectionConfig>>, AppError> {
|
|
||||||
let configs = state
|
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.load_connections()
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(configs))
|
Ok(Json(configs))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,12 +24,7 @@ pub async fn save_history(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(body): Json<SaveHistoryRequest>,
|
Json(body): Json<SaveHistoryRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
state
|
state.app.storage.save_history_entry(&body.entry).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.save_history_entry(&body.entry)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -39,12 +34,7 @@ pub async fn load_history(
|
|||||||
) -> Result<Json<Vec<HistoryEntry>>, AppError> {
|
) -> Result<Json<Vec<HistoryEntry>>, AppError> {
|
||||||
let limit = q.limit.unwrap_or(100);
|
let limit = q.limit.unwrap_or(100);
|
||||||
let offset = q.offset.unwrap_or(0);
|
let offset = q.offset.unwrap_or(0);
|
||||||
let entries = state
|
let entries = state.app.storage.load_history_entries(limit, offset).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.load_history_entries(limit, offset)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(entries))
|
Ok(Json(entries))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -57,11 +47,6 @@ pub async fn delete_history_entry(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
state
|
state.app.storage.delete_history_entry(&id).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.delete_history_entry(&id)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,23 +17,11 @@ pub async fn save_sidebar_layout(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(body): Json<SaveLayoutRequest>,
|
Json(body): Json<SaveLayoutRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
state
|
state.app.storage.save_sidebar_layout(&body.layout).await.map_err(AppError)?;
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.save_sidebar_layout(&body.layout)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn load_sidebar_layout(
|
pub async fn load_sidebar_layout(State(state): State<Arc<WebState>>) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
State(state): State<Arc<WebState>>,
|
let layout = state.app.storage.load_sidebar_layout().await.map_err(AppError)?;
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
|
||||||
let layout = state
|
|
||||||
.app
|
|
||||||
.storage
|
|
||||||
.load_sidebar_layout()
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(layout.unwrap_or(serde_json::json!(null))))
|
Ok(Json(layout.unwrap_or(serde_json::json!(null))))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -62,9 +62,8 @@ pub async fn list_databases(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<MongoConnectionRequest>,
|
Json(req): Json<MongoConnectionRequest>,
|
||||||
) -> Result<Json<Vec<String>>, AppError> {
|
) -> Result<Json<Vec<String>>, AppError> {
|
||||||
let result = dbx_core::mongo_ops::mongo_list_databases_core(&state.app, &req.connection_id)
|
let result =
|
||||||
.await
|
dbx_core::mongo_ops::mongo_list_databases_core(&state.app, &req.connection_id).await.map_err(AppError)?;
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(result))
|
Ok(Json(result))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -72,13 +71,9 @@ pub async fn list_collections(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<MongoCollectionRequest>,
|
Json(req): Json<MongoCollectionRequest>,
|
||||||
) -> Result<Json<Vec<String>>, AppError> {
|
) -> Result<Json<Vec<String>>, AppError> {
|
||||||
let result = dbx_core::mongo_ops::mongo_list_collections_core(
|
let result = dbx_core::mongo_ops::mongo_list_collections_core(&state.app, &req.connection_id, &req.database)
|
||||||
&state.app,
|
.await
|
||||||
&req.connection_id,
|
.map_err(AppError)?;
|
||||||
&req.database,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(result))
|
Ok(Json(result))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -34,9 +34,7 @@ pub async fn execute_query(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<ExecuteQueryRequest>,
|
Json(req): Json<ExecuteQueryRequest>,
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
let execution_id = req
|
let execution_id = req.execution_id.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
|
||||||
.execution_id
|
|
||||||
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
|
|
||||||
|
|
||||||
let registered = state.app.running_queries.register(execution_id);
|
let registered = state.app.running_queries.register(execution_id);
|
||||||
let cancel_token = registered.token();
|
let cancel_token = registered.token();
|
||||||
@@ -59,9 +57,7 @@ pub async fn execute_multi(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<ExecuteQueryRequest>,
|
Json(req): Json<ExecuteQueryRequest>,
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
let execution_id = req
|
let execution_id = req.execution_id.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
|
||||||
.execution_id
|
|
||||||
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
|
|
||||||
|
|
||||||
let registered = state.app.running_queries.register(execution_id);
|
let registered = state.app.running_queries.register(execution_id);
|
||||||
let cancel_token = registered.token();
|
let cancel_token = registered.token();
|
||||||
@@ -84,14 +80,9 @@ pub async fn execute_batch(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<ExecuteBatchRequest>,
|
Json(req): Json<ExecuteBatchRequest>,
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
let result = dbx_core::query::execute_statements(
|
let result = dbx_core::query::execute_statements(&state.app, &req.connection_id, &req.database, &req.statements)
|
||||||
&state.app,
|
.await
|
||||||
&req.connection_id,
|
.map_err(AppError)?;
|
||||||
&req.database,
|
|
||||||
&req.statements,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
@@ -109,14 +100,9 @@ pub async fn execute_script(
|
|||||||
Json(req): Json<ExecuteQueryRequest>,
|
Json(req): Json<ExecuteQueryRequest>,
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
let statements = dbx_core::sql::split_sql_statements(&req.sql);
|
let statements = dbx_core::sql::split_sql_statements(&req.sql);
|
||||||
let result = dbx_core::query::execute_statements(
|
let result = dbx_core::query::execute_statements(&state.app, &req.connection_id, &req.database, &statements)
|
||||||
&state.app,
|
.await
|
||||||
&req.connection_id,
|
.map_err(AppError)?;
|
||||||
&req.database,
|
|
||||||
&statements,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|||||||
+17
-43
@@ -69,9 +69,8 @@ pub async fn list_databases(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<RedisConnectionRequest>,
|
Json(req): Json<RedisConnectionRequest>,
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
let result = dbx_core::redis_ops::redis_list_databases_core(&state.app, &req.connection_id)
|
let result =
|
||||||
.await
|
dbx_core::redis_ops::redis_list_databases_core(&state.app, &req.connection_id).await.map_err(AppError)?;
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -96,9 +95,8 @@ pub async fn get_value(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<RedisKeyRequest>,
|
Json(req): Json<RedisKeyRequest>,
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
let result = dbx_core::redis_ops::redis_get_value_core(&state.app, &req.connection_id, &req.key)
|
let result =
|
||||||
.await
|
dbx_core::redis_ops::redis_get_value_core(&state.app, &req.connection_id, &req.key).await.map_err(AppError)?;
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -106,15 +104,9 @@ pub async fn set_string(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<RedisSetStringRequest>,
|
Json(req): Json<RedisSetStringRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
dbx_core::redis_ops::redis_set_string_core(
|
dbx_core::redis_ops::redis_set_string_core(&state.app, &req.connection_id, &req.key, &req.value, req.ttl)
|
||||||
&state.app,
|
.await
|
||||||
&req.connection_id,
|
.map_err(AppError)?;
|
||||||
&req.key,
|
|
||||||
&req.value,
|
|
||||||
req.ttl,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -122,9 +114,7 @@ pub async fn delete_key(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<RedisKeyRequest>,
|
Json(req): Json<RedisKeyRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
dbx_core::redis_ops::redis_delete_key_core(&state.app, &req.connection_id, &req.key)
|
dbx_core::redis_ops::redis_delete_key_core(&state.app, &req.connection_id, &req.key).await.map_err(AppError)?;
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -133,15 +123,9 @@ pub async fn hash_set(
|
|||||||
Json(req): Json<RedisHashRequest>,
|
Json(req): Json<RedisHashRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
let value = req.value.as_deref().unwrap_or("");
|
let value = req.value.as_deref().unwrap_or("");
|
||||||
dbx_core::redis_ops::redis_hash_set_core(
|
dbx_core::redis_ops::redis_hash_set_core(&state.app, &req.connection_id, &req.key, &req.field, value)
|
||||||
&state.app,
|
.await
|
||||||
&req.connection_id,
|
.map_err(AppError)?;
|
||||||
&req.key,
|
|
||||||
&req.field,
|
|
||||||
value,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -149,14 +133,9 @@ pub async fn hash_del(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<RedisHashRequest>,
|
Json(req): Json<RedisHashRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
dbx_core::redis_ops::redis_hash_del_core(
|
dbx_core::redis_ops::redis_hash_del_core(&state.app, &req.connection_id, &req.key, &req.field)
|
||||||
&state.app,
|
.await
|
||||||
&req.connection_id,
|
.map_err(AppError)?;
|
||||||
&req.key,
|
|
||||||
&req.field,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -196,13 +175,8 @@ pub async fn set_remove(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Json(req): Json<RedisSetRequest>,
|
Json(req): Json<RedisSetRequest>,
|
||||||
) -> Result<Json<()>, AppError> {
|
) -> Result<Json<()>, AppError> {
|
||||||
dbx_core::redis_ops::redis_set_remove_core(
|
dbx_core::redis_ops::redis_set_remove_core(&state.app, &req.connection_id, &req.key, &req.member)
|
||||||
&state.app,
|
.await
|
||||||
&req.connection_id,
|
.map_err(AppError)?;
|
||||||
&req.key,
|
|
||||||
&req.member,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(()))
|
Ok(Json(()))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -19,9 +19,7 @@ pub async fn list_databases(
|
|||||||
State(state): State<Arc<WebState>>,
|
State(state): State<Arc<WebState>>,
|
||||||
Query(q): Query<SchemaQuery>,
|
Query(q): Query<SchemaQuery>,
|
||||||
) -> Result<Json<serde_json::Value>, AppError> {
|
) -> Result<Json<serde_json::Value>, AppError> {
|
||||||
let result = dbx_core::schema::list_databases_core(&state.app, &q.connection_id)
|
let result = dbx_core::schema::list_databases_core(&state.app, &q.connection_id).await.map_err(AppError)?;
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -30,9 +28,7 @@ pub async fn list_schemas(
|
|||||||
Query(q): Query<SchemaQuery>,
|
Query(q): Query<SchemaQuery>,
|
||||||
) -> Result<Json<Vec<String>>, AppError> {
|
) -> Result<Json<Vec<String>>, AppError> {
|
||||||
let database = q.database.as_deref().unwrap_or("");
|
let database = q.database.as_deref().unwrap_or("");
|
||||||
let result = dbx_core::schema::list_schemas_core(&state.app, &q.connection_id, database)
|
let result = dbx_core::schema::list_schemas_core(&state.app, &q.connection_id, database).await.map_err(AppError)?;
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(result))
|
Ok(Json(result))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -43,9 +39,7 @@ pub async fn list_tables(
|
|||||||
let database = q.database.as_deref().unwrap_or("");
|
let database = q.database.as_deref().unwrap_or("");
|
||||||
let schema = q.schema.as_deref().unwrap_or("");
|
let schema = q.schema.as_deref().unwrap_or("");
|
||||||
let result =
|
let result =
|
||||||
dbx_core::schema::list_tables_core(&state.app, &q.connection_id, database, schema)
|
dbx_core::schema::list_tables_core(&state.app, &q.connection_id, database, schema).await.map_err(AppError)?;
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -56,15 +50,9 @@ pub async fn list_columns(
|
|||||||
let database = q.database.as_deref().unwrap_or("");
|
let database = q.database.as_deref().unwrap_or("");
|
||||||
let schema = q.schema.as_deref().unwrap_or("");
|
let schema = q.schema.as_deref().unwrap_or("");
|
||||||
let table = q.table.as_deref().unwrap_or("");
|
let table = q.table.as_deref().unwrap_or("");
|
||||||
let result = dbx_core::schema::get_columns_core(
|
let result = dbx_core::schema::get_columns_core(&state.app, &q.connection_id, database, schema, table)
|
||||||
&state.app,
|
.await
|
||||||
&q.connection_id,
|
.map_err(AppError)?;
|
||||||
database,
|
|
||||||
schema,
|
|
||||||
table,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -75,15 +63,9 @@ pub async fn list_indexes(
|
|||||||
let database = q.database.as_deref().unwrap_or("");
|
let database = q.database.as_deref().unwrap_or("");
|
||||||
let schema = q.schema.as_deref().unwrap_or("");
|
let schema = q.schema.as_deref().unwrap_or("");
|
||||||
let table = q.table.as_deref().unwrap_or("");
|
let table = q.table.as_deref().unwrap_or("");
|
||||||
let result = dbx_core::schema::list_indexes_core(
|
let result = dbx_core::schema::list_indexes_core(&state.app, &q.connection_id, database, schema, table)
|
||||||
&state.app,
|
.await
|
||||||
&q.connection_id,
|
.map_err(AppError)?;
|
||||||
database,
|
|
||||||
schema,
|
|
||||||
table,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -94,15 +76,9 @@ pub async fn list_foreign_keys(
|
|||||||
let database = q.database.as_deref().unwrap_or("");
|
let database = q.database.as_deref().unwrap_or("");
|
||||||
let schema = q.schema.as_deref().unwrap_or("");
|
let schema = q.schema.as_deref().unwrap_or("");
|
||||||
let table = q.table.as_deref().unwrap_or("");
|
let table = q.table.as_deref().unwrap_or("");
|
||||||
let result = dbx_core::schema::list_foreign_keys_core(
|
let result = dbx_core::schema::list_foreign_keys_core(&state.app, &q.connection_id, database, schema, table)
|
||||||
&state.app,
|
.await
|
||||||
&q.connection_id,
|
.map_err(AppError)?;
|
||||||
database,
|
|
||||||
schema,
|
|
||||||
table,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -113,15 +89,9 @@ pub async fn list_triggers(
|
|||||||
let database = q.database.as_deref().unwrap_or("");
|
let database = q.database.as_deref().unwrap_or("");
|
||||||
let schema = q.schema.as_deref().unwrap_or("");
|
let schema = q.schema.as_deref().unwrap_or("");
|
||||||
let table = q.table.as_deref().unwrap_or("");
|
let table = q.table.as_deref().unwrap_or("");
|
||||||
let result = dbx_core::schema::list_triggers_core(
|
let result = dbx_core::schema::list_triggers_core(&state.app, &q.connection_id, database, schema, table)
|
||||||
&state.app,
|
.await
|
||||||
&q.connection_id,
|
.map_err(AppError)?;
|
||||||
database,
|
|
||||||
schema,
|
|
||||||
table,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
Ok(Json(serde_json::to_value(result).map_err(|e| AppError(e.to_string()))?))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -132,14 +102,8 @@ pub async fn get_ddl(
|
|||||||
let database = q.database.as_deref().unwrap_or("");
|
let database = q.database.as_deref().unwrap_or("");
|
||||||
let schema = q.schema.as_deref().unwrap_or("");
|
let schema = q.schema.as_deref().unwrap_or("");
|
||||||
let table = q.table.as_deref().unwrap_or("");
|
let table = q.table.as_deref().unwrap_or("");
|
||||||
let result = dbx_core::schema::get_table_ddl_core(
|
let result = dbx_core::schema::get_table_ddl_core(&state.app, &q.connection_id, database, schema, table)
|
||||||
&state.app,
|
.await
|
||||||
&q.connection_id,
|
.map_err(AppError)?;
|
||||||
database,
|
|
||||||
schema,
|
|
||||||
table,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(AppError)?;
|
|
||||||
Ok(Json(result))
|
Ok(Json(result))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,8 +3,8 @@ use std::sync::Arc;
|
|||||||
use axum::extract::{Multipart, Path, State};
|
use axum::extract::{Multipart, Path, State};
|
||||||
use axum::response::sse::{Event, Sse};
|
use axum::response::sse::{Event, Sse};
|
||||||
use axum::Json;
|
use axum::Json;
|
||||||
use dbx_core::sql;
|
|
||||||
use dbx_core::query;
|
use dbx_core::query;
|
||||||
|
use dbx_core::sql;
|
||||||
use futures::stream::Stream;
|
use futures::stream::Stream;
|
||||||
use serde::Deserialize;
|
use serde::Deserialize;
|
||||||
|
|
||||||
@@ -40,15 +40,8 @@ pub async fn preview_sql_file(
|
|||||||
let tmp_dir = state.data_dir.join("tmp");
|
let tmp_dir = state.data_dir.join("tmp");
|
||||||
std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?;
|
std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?;
|
||||||
|
|
||||||
while let Some(field) = multipart
|
while let Some(field) = multipart.next_field().await.map_err(|e| AppError(e.to_string()))? {
|
||||||
.next_field()
|
let file_name = field.file_name().unwrap_or("upload.sql").to_string();
|
||||||
.await
|
|
||||||
.map_err(|e| AppError(e.to_string()))?
|
|
||||||
{
|
|
||||||
let file_name = field
|
|
||||||
.file_name()
|
|
||||||
.unwrap_or("upload.sql")
|
|
||||||
.to_string();
|
|
||||||
let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?;
|
let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?;
|
||||||
|
|
||||||
let file_path = tmp_dir.join(&file_name);
|
let file_path = tmp_dir.join(&file_name);
|
||||||
@@ -77,11 +70,7 @@ pub async fn execute_sql_file(
|
|||||||
let execution_id = req.execution_id.clone();
|
let execution_id = req.execution_id.clone();
|
||||||
|
|
||||||
let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
|
let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
|
||||||
state
|
state.sse_channels.write().await.insert(execution_id.clone(), tx.clone());
|
||||||
.sse_channels
|
|
||||||
.write()
|
|
||||||
.await
|
|
||||||
.insert(execution_id.clone(), tx.clone());
|
|
||||||
|
|
||||||
let app = state.app.clone();
|
let app = state.app.clone();
|
||||||
let state_clone = state.clone();
|
let state_clone = state.clone();
|
||||||
@@ -149,9 +138,7 @@ pub async fn execute_sql_file(
|
|||||||
let _ = tx.send(json);
|
let _ = tx.send(json);
|
||||||
}
|
}
|
||||||
|
|
||||||
match query::execute_sql_statement(&app, &req.connection_id, &req.database, stmt, None)
|
match query::execute_sql_statement(&app, &req.connection_id, &req.database, stmt, None).await {
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(result) => {
|
Ok(result) => {
|
||||||
success_count += 1;
|
success_count += 1;
|
||||||
total_affected += result.affected_rows;
|
total_affected += result.affected_rows;
|
||||||
@@ -220,9 +207,7 @@ pub async fn sql_file_progress(
|
|||||||
Path(execution_id): Path<String>,
|
Path(execution_id): Path<String>,
|
||||||
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
|
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
|
||||||
let channels = state.sse_channels.read().await;
|
let channels = state.sse_channels.read().await;
|
||||||
let tx = channels
|
let tx = channels.get(&execution_id).ok_or_else(|| AppError("Execution not found".to_string()))?;
|
||||||
.get(&execution_id)
|
|
||||||
.ok_or_else(|| AppError("Execution not found".to_string()))?;
|
|
||||||
let rx = tx.subscribe();
|
let rx = tx.subscribe();
|
||||||
drop(channels);
|
drop(channels);
|
||||||
Ok(crate::sse::sse_from_channel(rx))
|
Ok(crate::sse::sse_from_channel(rx))
|
||||||
|
|||||||
@@ -3,9 +3,7 @@ use std::sync::Arc;
|
|||||||
use axum::extract::{Multipart, Path, State};
|
use axum::extract::{Multipart, Path, State};
|
||||||
use axum::response::sse::{Event, Sse};
|
use axum::response::sse::{Event, Sse};
|
||||||
use axum::Json;
|
use axum::Json;
|
||||||
use dbx_core::table_import::{
|
use dbx_core::table_import::{self, TableImportRequest};
|
||||||
self, TableImportRequest,
|
|
||||||
};
|
|
||||||
use dbx_core::transfer;
|
use dbx_core::transfer;
|
||||||
use futures::stream::Stream;
|
use futures::stream::Stream;
|
||||||
use serde::Deserialize;
|
use serde::Deserialize;
|
||||||
@@ -32,27 +30,17 @@ pub async fn preview_import(
|
|||||||
let tmp_dir = state.data_dir.join("tmp");
|
let tmp_dir = state.data_dir.join("tmp");
|
||||||
std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?;
|
std::fs::create_dir_all(&tmp_dir).map_err(|e| AppError(e.to_string()))?;
|
||||||
|
|
||||||
while let Some(field) = multipart
|
while let Some(field) = multipart.next_field().await.map_err(|e| AppError(e.to_string()))? {
|
||||||
.next_field()
|
let file_name = field.file_name().unwrap_or("upload.csv").to_string();
|
||||||
.await
|
|
||||||
.map_err(|e| AppError(e.to_string()))?
|
|
||||||
{
|
|
||||||
let file_name = field
|
|
||||||
.file_name()
|
|
||||||
.unwrap_or("upload.csv")
|
|
||||||
.to_string();
|
|
||||||
let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?;
|
let data = field.bytes().await.map_err(|e| AppError(e.to_string()))?;
|
||||||
|
|
||||||
let file_path = tmp_dir.join(&file_name);
|
let file_path = tmp_dir.join(&file_name);
|
||||||
std::fs::write(&file_path, &data).map_err(|e| AppError(e.to_string()))?;
|
std::fs::write(&file_path, &data).map_err(|e| AppError(e.to_string()))?;
|
||||||
|
|
||||||
let file_path_str = file_path.to_string_lossy().to_string();
|
let file_path_str = file_path.to_string_lossy().to_string();
|
||||||
let preview = table_import::preview_table_import_file_core(&file_path_str)
|
let preview = table_import::preview_table_import_file_core(&file_path_str).map_err(AppError)?;
|
||||||
.map_err(AppError)?;
|
|
||||||
|
|
||||||
return Ok(Json(
|
return Ok(Json(serde_json::to_value(preview).map_err(|e| AppError(e.to_string()))?));
|
||||||
serde_json::to_value(preview).map_err(|e| AppError(e.to_string()))?,
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Err(AppError("No file uploaded".to_string()))
|
Err(AppError("No file uploaded".to_string()))
|
||||||
@@ -66,11 +54,7 @@ pub async fn execute_import(
|
|||||||
let import_id = req.import_id.clone();
|
let import_id = req.import_id.clone();
|
||||||
|
|
||||||
let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
|
let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
|
||||||
state
|
state.sse_channels.write().await.insert(import_id.clone(), tx.clone());
|
||||||
.sse_channels
|
|
||||||
.write()
|
|
||||||
.await
|
|
||||||
.insert(import_id.clone(), tx.clone());
|
|
||||||
|
|
||||||
let app = state.app.clone();
|
let app = state.app.clone();
|
||||||
let state_clone = state.clone();
|
let state_clone = state.clone();
|
||||||
@@ -91,10 +75,7 @@ pub async fn execute_import(
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
let pool_key = match app
|
let pool_key = match app.get_or_create_pool(&req.connection_id, Some(&req.database)).await {
|
||||||
.get_or_create_pool(&req.connection_id, Some(&req.database))
|
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(k) => k,
|
Ok(k) => k,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
let _ = tx.send(
|
let _ = tx.send(
|
||||||
@@ -118,9 +99,7 @@ pub async fn execute_import(
|
|||||||
&pool_key,
|
&pool_key,
|
||||||
|id: &str| {
|
|id: &str| {
|
||||||
let id = id.to_string();
|
let id = id.to_string();
|
||||||
Box::pin(async move {
|
Box::pin(async move { transfer::is_cancelled(&id).await })
|
||||||
transfer::is_cancelled(&id).await
|
|
||||||
})
|
|
||||||
},
|
},
|
||||||
|progress| {
|
|progress| {
|
||||||
if let Ok(json) = serde_json::to_string(&progress) {
|
if let Ok(json) = serde_json::to_string(&progress) {
|
||||||
@@ -159,9 +138,7 @@ pub async fn import_progress(
|
|||||||
Path(import_id): Path<String>,
|
Path(import_id): Path<String>,
|
||||||
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
|
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
|
||||||
let channels = state.sse_channels.read().await;
|
let channels = state.sse_channels.read().await;
|
||||||
let tx = channels
|
let tx = channels.get(&import_id).ok_or_else(|| AppError("Import not found".to_string()))?;
|
||||||
.get(&import_id)
|
|
||||||
.ok_or_else(|| AppError("Import not found".to_string()))?;
|
|
||||||
let rx = tx.subscribe();
|
let rx = tx.subscribe();
|
||||||
drop(channels);
|
drop(channels);
|
||||||
Ok(crate::sse::sse_from_channel(rx))
|
Ok(crate::sse::sse_from_channel(rx))
|
||||||
|
|||||||
@@ -31,11 +31,7 @@ pub async fn start_transfer(
|
|||||||
|
|
||||||
// Create a broadcast channel for progress
|
// Create a broadcast channel for progress
|
||||||
let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
|
let (tx, _) = tokio::sync::broadcast::channel::<String>(256);
|
||||||
state
|
state.sse_channels.write().await.insert(transfer_id.clone(), tx.clone());
|
||||||
.sse_channels
|
|
||||||
.write()
|
|
||||||
.await
|
|
||||||
.insert(transfer_id.clone(), tx.clone());
|
|
||||||
|
|
||||||
let app = state.app.clone();
|
let app = state.app.clone();
|
||||||
let state_clone = state.clone();
|
let state_clone = state.clone();
|
||||||
@@ -56,9 +52,7 @@ pub async fn start_transfer(
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
let source_pool_key = match app
|
let source_pool_key = match app.get_or_create_pool(&req.source_connection_id, Some(&req.source_database)).await
|
||||||
.get_or_create_pool(&req.source_connection_id, Some(&req.source_database))
|
|
||||||
.await
|
|
||||||
{
|
{
|
||||||
Ok(k) => k,
|
Ok(k) => k,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -66,9 +60,7 @@ pub async fn start_transfer(
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let target_pool_key = match app
|
let target_pool_key = match app.get_or_create_pool(&req.target_connection_id, Some(&req.target_database)).await
|
||||||
.get_or_create_pool(&req.target_connection_id, Some(&req.target_database))
|
|
||||||
.await
|
|
||||||
{
|
{
|
||||||
Ok(k) => k,
|
Ok(k) => k,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -151,9 +143,7 @@ pub async fn start_transfer(
|
|||||||
state_clone.sse_channels.write().await.remove(&req.transfer_id);
|
state_clone.sse_channels.write().await.remove(&req.transfer_id);
|
||||||
});
|
});
|
||||||
|
|
||||||
Ok(Json(
|
Ok(Json(serde_json::json!({ "transferId": transfer_id })))
|
||||||
serde_json::json!({ "transferId": transfer_id }),
|
|
||||||
))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn transfer_progress(
|
pub async fn transfer_progress(
|
||||||
@@ -161,9 +151,7 @@ pub async fn transfer_progress(
|
|||||||
Path(transfer_id): Path<String>,
|
Path(transfer_id): Path<String>,
|
||||||
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
|
) -> Result<Sse<impl Stream<Item = Result<Event, std::convert::Infallible>>>, AppError> {
|
||||||
let channels = state.sse_channels.read().await;
|
let channels = state.sse_channels.read().await;
|
||||||
let tx = channels
|
let tx = channels.get(&transfer_id).ok_or_else(|| AppError("Transfer not found".to_string()))?;
|
||||||
.get(&transfer_id)
|
|
||||||
.ok_or_else(|| AppError("Transfer not found".to_string()))?;
|
|
||||||
let rx = tx.subscribe();
|
let rx = tx.subscribe();
|
||||||
drop(channels);
|
drop(channels);
|
||||||
Ok(crate::sse::sse_from_channel(rx))
|
Ok(crate::sse::sse_from_channel(rx))
|
||||||
|
|||||||
Reference in New Issue
Block a user