refactor(agent): share driver service logic

This commit is contained in:
t8y2 2026-05-18 22:08:10 +08:00
parent 8a81a01fac
commit 94eeb38d63
5 changed files with 233 additions and 277 deletions

View File

@ -0,0 +1,128 @@
use std::path::PathBuf;
use crate::agent_manager::{AgentDriverInfo, AgentManager, AgentRegistry, InstalledDriver, DEFAULT_JRE_KEY};
const REGISTRY_PATH: &str = "https://github.com/t8y2/dbx-agents/releases/latest/download/agent-registry.json";
const REGISTRY_R2_PATH: &str = "agents/agent-registry.json";
static REGISTRY_CACHE: std::sync::LazyLock<tokio::sync::Mutex<Option<(std::time::Instant, AgentRegistry)>>> =
std::sync::LazyLock::new(|| tokio::sync::Mutex::new(None));
pub const AGENT_TYPES: &[(&str, &str)] = &[
("dameng", "达梦 DM8"),
("kingbase", "人大金仓 KingbaseES"),
("highgo", "瀚高 HighGo"),
("vastbase", "Vastbase"),
("goldendb", "GoldenDB"),
("access", "Microsoft Access"),
("oracle", "Oracle"),
("oracle-10g", "Oracle 10g"),
("h2", "H2"),
("snowflake", "Snowflake"),
("trino", "Trino (Presto)"),
("hive", "Apache Hive"),
("db2", "IBM DB2"),
("informix", "IBM Informix"),
("neo4j", "Neo4j"),
("cassandra", "Apache Cassandra"),
("bigquery", "Google BigQuery"),
("kylin", "Apache Kylin"),
("sundb", "SunDB"),
("gaussdb", "GaussDB"),
("yashandb", "崖山 YashanDB"),
("tdengine", "TDengine"),
("mongodb", "MongoDB (Legacy)"),
];
pub fn build_agent_list(am: &AgentManager, registry: Option<&AgentRegistry>) -> Vec<AgentDriverInfo> {
let local_state = am.load_state();
AGENT_TYPES
.iter()
.map(|(key, label)| {
let installed = am.is_driver_installed(key);
let local = local_state.installed_drivers.get(*key);
let remote = registry.and_then(|r| r.drivers.get(*key));
let jre_key = remote
.map(|r| r.jre.clone())
.or_else(|| local.map(|l| l.jre.clone()))
.unwrap_or_else(|| DEFAULT_JRE_KEY.to_string());
AgentDriverInfo {
db_type: key.to_string(),
label: label.to_string(),
version: remote.map(|r| r.version.clone()).unwrap_or_default(),
size: remote.map(|r| r.jar.size).unwrap_or(0),
installed,
installed_version: local.map(|l| l.version.clone()),
update_available: match (local, remote) {
(Some(l), Some(r)) => l.version != r.version,
_ => false,
},
jre: jre_key.clone(),
jre_installed: am.is_jre_installed(&jre_key),
}
})
.collect()
}
pub fn local_agent_jar_candidates(db_type: &str) -> Vec<PathBuf> {
let jar_name = format!("dbx-agent-{db_type}.jar");
let relative = PathBuf::from("..").join("dbx-agents").join(db_type).join("build").join("libs").join(&jar_name);
let nested = PathBuf::from("dbx-agents").join(db_type).join("build").join("libs").join(&jar_name);
vec![relative, nested]
}
pub fn find_local_agent_jar(db_type: &str) -> Option<PathBuf> {
local_agent_jar_candidates(db_type).into_iter().find(|path| path.exists())
}
pub fn install_local_agent(am: &AgentManager, db_type: &str, source: PathBuf) -> Result<(), String> {
let jar_path = am.driver_jar_path(db_type);
let parent = jar_path.parent().ok_or_else(|| format!("Invalid driver path: {}", jar_path.display()))?;
std::fs::create_dir_all(parent).map_err(|e| e.to_string())?;
std::fs::copy(&source, &jar_path).map_err(|e| format!("Failed to copy local agent jar: {e}"))?;
let mut local_state = am.load_state();
local_state.installed_drivers.insert(
db_type.to_string(),
InstalledDriver {
version: "0.1.0-local".to_string(),
installed_at: chrono::Utc::now().to_rfc3339(),
jre: DEFAULT_JRE_KEY.to_string(),
},
);
am.save_state(&local_state)
}
pub async fn fetch_registry() -> Result<AgentRegistry, String> {
{
let cache = REGISTRY_CACHE.lock().await;
if let Some((ts, registry)) = cache.as_ref() {
if ts.elapsed() < std::time::Duration::from_secs(300) {
return Ok(registry.clone());
}
}
}
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|err| format!("Failed to create HTTP client: {err}"))?;
let resp = crate::race_download(&client, REGISTRY_PATH, REGISTRY_R2_PATH, "dbx-agent-manager")
.await
.map_err(|err| format!("Failed to fetch agent registry: {err}"))?;
let registry: AgentRegistry = resp.json().await.map_err(|err| format!("Failed to parse registry: {err}"))?;
*REGISTRY_CACHE.lock().await = Some((std::time::Instant::now(), registry.clone()));
Ok(registry)
}
pub async fn invalidate_registry_cache() {
*REGISTRY_CACHE.lock().await = None;
}
pub fn github_url_to_r2_path(github_url: &str, category: &str) -> String {
let filename = github_url.rsplit('/').next().unwrap_or(github_url);
match category {
"jre" => format!("agents/jre/{filename}"),
"driver" => format!("agents/drivers/{filename}"),
_ => format!("agents/{filename}"),
}
}

View File

@ -1,4 +1,5 @@
pub mod agent_manager;
pub mod agent_service;
pub mod ai;
pub mod connection;
pub mod connection_secrets;

View File

@ -0,0 +1,92 @@
use dbx_core::agent_manager::{
AgentManager, AgentRegistry, ArtifactInfo, DriverInfo, InstalledDriver, DEFAULT_JRE_KEY,
};
use dbx_core::agent_service::{build_agent_list, github_url_to_r2_path, local_agent_jar_candidates};
fn test_manager(name: &str) -> AgentManager {
let dir = std::env::temp_dir().join(format!("dbx-agent-service-{name}-{}", uuid::Uuid::new_v4()));
AgentManager::new_with_base_dir(dir)
}
fn registry_with_driver(db_type: &str, version: &str, jre: &str) -> AgentRegistry {
let mut drivers = std::collections::HashMap::new();
drivers.insert(
db_type.to_string(),
DriverInfo {
version: version.to_string(),
label: db_type.to_string(),
min_app_version: "0.1.0".to_string(),
jre: jre.to_string(),
jar: ArtifactInfo {
url: format!("https://example.com/dbx-agent-{db_type}.jar"),
sha256: "sha".to_string(),
size: 42,
},
},
);
AgentRegistry { jre: None, jres: std::collections::HashMap::new(), drivers }
}
#[test]
fn built_in_agent_list_includes_expected_driver_labels() {
let manager = test_manager("labels");
let agents = build_agent_list(&manager, None);
assert!(agents.iter().any(|agent| agent.db_type == "tdengine" && agent.label == "TDengine"));
assert!(agents.iter().any(|agent| agent.db_type == "yashandb" && agent.label == "崖山 YashanDB"));
assert!(agents.iter().any(|agent| agent.db_type == "access" && agent.label == "Microsoft Access"));
}
#[test]
fn agent_list_marks_installed_driver_update_when_registry_version_differs() {
let manager = test_manager("update");
let jar_path = manager.driver_jar_path("h2");
std::fs::create_dir_all(jar_path.parent().unwrap()).unwrap();
std::fs::write(&jar_path, b"jar").unwrap();
manager
.save_state(&dbx_core::agent_manager::AgentState {
installed_drivers: [(
"h2".to_string(),
InstalledDriver {
version: "0.1.0".to_string(),
installed_at: "2026-05-18T00:00:00Z".to_string(),
jre: DEFAULT_JRE_KEY.to_string(),
},
)]
.into_iter()
.collect(),
..Default::default()
})
.unwrap();
let registry = registry_with_driver("h2", "0.2.0", "21");
let agents = build_agent_list(&manager, Some(&registry));
let h2 = agents.iter().find(|agent| agent.db_type == "h2").unwrap();
assert!(h2.installed);
assert_eq!(h2.installed_version.as_deref(), Some("0.1.0"));
assert_eq!(h2.version, "0.2.0");
assert_eq!(h2.size, 42);
assert_eq!(h2.jre, "21");
assert!(h2.update_available);
}
#[test]
fn local_agent_jar_candidates_include_sibling_build_output() {
let candidates = local_agent_jar_candidates("tdengine");
assert!(candidates.iter().any(|path| path.ends_with("dbx-agents/tdengine/build/libs/dbx-agent-tdengine.jar")));
}
#[test]
fn github_agent_asset_urls_map_to_r2_paths_by_category() {
assert_eq!(
github_url_to_r2_path("https://github.com/t8y2/dbx-agents/releases/download/v1/jre-17.tar.gz", "jre"),
"agents/jre/jre-17.tar.gz"
);
assert_eq!(
github_url_to_r2_path("https://github.com/t8y2/dbx-agents/releases/download/v1/dbx-agent-h2.jar", "driver"),
"agents/drivers/dbx-agent-h2.jar"
);
}

View File

@ -1,4 +1,3 @@
use std::path::PathBuf;
use std::sync::Arc;
use axum::extract::{Path, State};
@ -8,45 +7,17 @@ use dbx_core::agent_manager::{
AgentDriverInfo, AgentManager, AgentRegistry, AgentState, InstalledDriver, JavaRuntimeConfig, JavaRuntimeMode,
DEFAULT_JRE_KEY,
};
use dbx_core::agent_service::{
build_agent_list, fetch_registry, find_local_agent_jar, github_url_to_r2_path, install_local_agent,
invalidate_registry_cache,
};
use futures::Stream;
use serde::Deserialize;
use tokio::sync::{broadcast, Mutex};
use tokio::sync::broadcast;
use crate::error::AppError;
use crate::state::WebState;
const REGISTRY_PATH: &str = "https://github.com/t8y2/dbx-agents/releases/latest/download/agent-registry.json";
const REGISTRY_R2_PATH: &str = "agents/agent-registry.json";
static REGISTRY_CACHE: std::sync::LazyLock<Mutex<Option<(std::time::Instant, AgentRegistry)>>> =
std::sync::LazyLock::new(|| Mutex::new(None));
const AGENT_TYPES: &[(&str, &str)] = &[
("dameng", "达梦 DM8"),
("kingbase", "人大金仓 KingbaseES"),
("highgo", "瀚高 HighGo"),
("vastbase", "Vastbase"),
("goldendb", "GoldenDB"),
("access", "Microsoft Access"),
("oracle", "Oracle"),
("oracle-10g", "Oracle 10g"),
("h2", "H2"),
("snowflake", "Snowflake"),
("trino", "Trino (Presto)"),
("hive", "Apache Hive"),
("db2", "IBM DB2"),
("informix", "IBM Informix"),
("neo4j", "Neo4j"),
("cassandra", "Apache Cassandra"),
("bigquery", "Google BigQuery"),
("kylin", "Apache Kylin"),
("sundb", "SunDB"),
("gaussdb", "GaussDB"),
("yashandb", "崖山 YashanDB"),
("tdengine", "TDengine"),
("mongodb", "MongoDB (Legacy)"),
];
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AgentTypeRequest {
@ -65,36 +36,6 @@ pub struct JavaRuntimeRequest {
pub config: JavaRuntimeConfig,
}
fn build_agent_list(am: &AgentManager, registry: Option<&AgentRegistry>) -> Vec<AgentDriverInfo> {
let local_state = am.load_state();
AGENT_TYPES
.iter()
.map(|(key, label)| {
let installed = am.is_driver_installed(key);
let local = local_state.installed_drivers.get(*key);
let remote = registry.and_then(|r| r.drivers.get(*key));
let jre_key = remote
.map(|r| r.jre.clone())
.or_else(|| local.map(|l| l.jre.clone()))
.unwrap_or_else(|| DEFAULT_JRE_KEY.to_string());
AgentDriverInfo {
db_type: key.to_string(),
label: label.to_string(),
version: remote.map(|r| r.version.clone()).unwrap_or_default(),
size: remote.map(|r| r.jar.size).unwrap_or(0),
installed,
installed_version: local.map(|l| l.version.clone()),
update_available: match (local, remote) {
(Some(l), Some(r)) => l.version != r.version,
_ => false,
},
jre: jre_key.clone(),
jre_installed: am.is_jre_installed(&jre_key),
}
})
.collect()
}
pub async fn list_installed_agents_local(
State(state): State<Arc<WebState>>,
) -> Result<Json<Vec<AgentDriverInfo>>, AppError> {
@ -182,7 +123,7 @@ pub async fn set_agent_java_runtime_config(
}
pub async fn invalidate_agent_registry_cache() -> Result<Json<serde_json::Value>, AppError> {
*REGISTRY_CACHE.lock().await = None;
invalidate_registry_cache().await;
Ok(Json(serde_json::json!({ "ok": true })))
}
@ -370,65 +311,6 @@ async fn reinstall_jre_core(am: &AgentManager, jre_key: &str, tx: &broadcast::Se
Ok(())
}
fn local_agent_jar_candidates(db_type: &str) -> Vec<PathBuf> {
let jar_name = format!("dbx-agent-{db_type}.jar");
let relative = PathBuf::from("..").join("dbx-agents").join(db_type).join("build").join("libs").join(&jar_name);
let nested = PathBuf::from("dbx-agents").join(db_type).join("build").join("libs").join(&jar_name);
vec![relative, nested]
}
fn find_local_agent_jar(db_type: &str) -> Option<PathBuf> {
local_agent_jar_candidates(db_type).into_iter().find(|path| path.exists())
}
fn install_local_agent(am: &AgentManager, db_type: &str, source: PathBuf) -> Result<(), String> {
let jar_path = am.driver_jar_path(db_type);
let parent = jar_path.parent().ok_or_else(|| format!("Invalid driver path: {}", jar_path.display()))?;
std::fs::create_dir_all(parent).map_err(|err| err.to_string())?;
std::fs::copy(&source, &jar_path).map_err(|err| format!("Failed to copy local agent jar: {err}"))?;
let mut local_state = am.load_state();
local_state.installed_drivers.insert(
db_type.to_string(),
InstalledDriver {
version: "0.1.0-local".to_string(),
installed_at: chrono::Utc::now().to_rfc3339(),
jre: DEFAULT_JRE_KEY.to_string(),
},
);
am.save_state(&local_state)
}
async fn fetch_registry() -> Result<AgentRegistry, String> {
{
let cache = REGISTRY_CACHE.lock().await;
if let Some((ts, registry)) = cache.as_ref() {
if ts.elapsed() < std::time::Duration::from_secs(300) {
return Ok(registry.clone());
}
}
}
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|err| format!("Failed to create HTTP client: {err}"))?;
let resp = dbx_core::race_download(&client, REGISTRY_PATH, REGISTRY_R2_PATH, "dbx-agent-manager")
.await
.map_err(|err| format!("Failed to fetch agent registry: {err}"))?;
let registry: AgentRegistry = resp.json().await.map_err(|err| format!("Failed to parse registry: {err}"))?;
*REGISTRY_CACHE.lock().await = Some((std::time::Instant::now(), registry.clone()));
Ok(registry)
}
fn github_url_to_r2_path(github_url: &str, category: &str) -> String {
let filename = github_url.rsplit('/').next().unwrap_or(github_url);
match category {
"jre" => format!("agents/jre/{filename}"),
"driver" => format!("agents/drivers/{filename}"),
_ => format!("agents/{filename}"),
}
}
async fn download_with_progress(
tx: &broadcast::Sender<String>,
step: &str,

View File

@ -1,133 +1,16 @@
use std::path::PathBuf;
use std::sync::Arc;
use tauri::{Emitter, State};
use tokio::sync::Mutex;
use dbx_core::agent_manager::{
AgentDriverInfo, AgentManager, AgentRegistry, InstalledDriver, JavaRuntimeConfig, JavaRuntimeMode, DEFAULT_JRE_KEY,
AgentDriverInfo, AgentManager, InstalledDriver, JavaRuntimeConfig, JavaRuntimeMode, DEFAULT_JRE_KEY,
};
use dbx_core::agent_service::{
build_agent_list, fetch_registry, find_local_agent_jar, github_url_to_r2_path, install_local_agent,
invalidate_registry_cache,
};
use dbx_core::connection::AppState;
const REGISTRY_PATH: &str = "https://github.com/t8y2/dbx-agents/releases/latest/download/agent-registry.json";
const REGISTRY_R2_PATH: &str = "agents/agent-registry.json";
static REGISTRY_CACHE: std::sync::LazyLock<Mutex<Option<(std::time::Instant, AgentRegistry)>>> =
std::sync::LazyLock::new(|| Mutex::new(None));
const AGENT_TYPES: &[(&str, &str)] = &[
("dameng", "达梦 DM8"),
("kingbase", "人大金仓 KingbaseES"),
("highgo", "瀚高 HighGo"),
("vastbase", "Vastbase"),
("goldendb", "GoldenDB"),
("access", "Microsoft Access"),
("oracle", "Oracle"),
("oracle-10g", "Oracle 10g"),
("h2", "H2"),
("snowflake", "Snowflake"),
("trino", "Trino (Presto)"),
("hive", "Apache Hive"),
("db2", "IBM DB2"),
("informix", "IBM Informix"),
("neo4j", "Neo4j"),
("cassandra", "Apache Cassandra"),
("bigquery", "Google BigQuery"),
("kylin", "Apache Kylin"),
("sundb", "SunDB"),
("gaussdb", "GaussDB"),
("yashandb", "崖山 YashanDB"),
("tdengine", "TDengine"),
("mongodb", "MongoDB (Legacy)"),
];
fn build_agent_list(am: &AgentManager, registry: Option<&AgentRegistry>) -> Vec<AgentDriverInfo> {
let local_state = am.load_state();
AGENT_TYPES
.iter()
.map(|(key, label)| {
let installed = am.is_driver_installed(key);
let local = local_state.installed_drivers.get(*key);
let remote = registry.and_then(|r| r.drivers.get(*key));
let jre_key = remote
.map(|r| r.jre.clone())
.or_else(|| local.map(|l| l.jre.clone()))
.unwrap_or_else(|| DEFAULT_JRE_KEY.to_string());
AgentDriverInfo {
db_type: key.to_string(),
label: label.to_string(),
version: remote.map(|r| r.version.clone()).unwrap_or_default(),
size: remote.map(|r| r.jar.size).unwrap_or(0),
installed,
installed_version: local.map(|l| l.version.clone()),
update_available: match (local, remote) {
(Some(l), Some(r)) => l.version != r.version,
_ => false,
},
jre: jre_key.clone(),
jre_installed: am.is_jre_installed(&jre_key),
}
})
.collect()
}
fn local_agent_jar_candidates(db_type: &str) -> Vec<PathBuf> {
let jar_name = format!("dbx-agent-{db_type}.jar");
let relative = PathBuf::from("..").join("dbx-agents").join(db_type).join("build").join("libs").join(&jar_name);
let nested = PathBuf::from("dbx-agents").join(db_type).join("build").join("libs").join(&jar_name);
vec![relative, nested]
}
fn find_local_agent_jar(db_type: &str) -> Option<PathBuf> {
local_agent_jar_candidates(db_type).into_iter().find(|path| path.exists())
}
fn install_local_agent(am: &AgentManager, db_type: &str, source: PathBuf) -> Result<(), String> {
let jar_path = am.driver_jar_path(db_type);
let parent = jar_path.parent().ok_or_else(|| format!("Invalid driver path: {}", jar_path.display()))?;
std::fs::create_dir_all(parent).map_err(|e| e.to_string())?;
std::fs::copy(&source, &jar_path).map_err(|e| format!("Failed to copy local agent jar: {e}"))?;
let mut local_state = am.load_state();
local_state.installed_drivers.insert(
db_type.to_string(),
InstalledDriver {
version: "0.1.0-local".to_string(),
installed_at: chrono::Utc::now().to_rfc3339(),
jre: DEFAULT_JRE_KEY.to_string(),
},
);
am.save_state(&local_state)
}
#[cfg(test)]
mod tests {
use super::build_agent_list;
use dbx_core::agent_manager::AgentManager;
#[test]
fn built_in_agent_list_includes_access_and_tdengine() {
let dir = std::env::temp_dir().join(format!("dbx-agent-list-test-{}", uuid::Uuid::new_v4()));
let manager = AgentManager::new_with_base_dir(dir.clone());
let agents = build_agent_list(&manager, None);
assert!(agents.iter().any(|agent| agent.db_type == "tdengine" && agent.label == "TDengine"));
assert!(agents.iter().any(|agent| agent.db_type == "yashandb" && agent.label == "崖山 YashanDB"));
assert!(agents.iter().any(|agent| agent.db_type == "access" && agent.label == "Microsoft Access"));
let _ = std::fs::remove_dir_all(dir);
}
#[test]
fn local_agent_jar_candidates_include_sibling_dbx_agents_build_output() {
let candidates = super::local_agent_jar_candidates("tdengine");
assert!(candidates
.iter()
.any(|path| { path.ends_with("dbx-agents/tdengine/build/libs/dbx-agent-tdengine.jar") }));
}
}
#[tauri::command]
pub async fn list_installed_agents_local(state: State<'_, Arc<AppState>>) -> Result<Vec<AgentDriverInfo>, String> {
Ok(build_agent_list(&state.agent_manager, None))
@ -401,7 +284,7 @@ pub async fn uninstall_jre(state: State<'_, Arc<AppState>>, jre_key: String) ->
#[tauri::command]
pub async fn invalidate_agent_registry_cache() -> Result<(), String> {
*REGISTRY_CACHE.lock().await = None;
invalidate_registry_cache().await;
Ok(())
}
@ -441,36 +324,6 @@ pub async fn reinstall_jre(
Ok(())
}
async fn fetch_registry() -> Result<AgentRegistry, String> {
{
let cache = REGISTRY_CACHE.lock().await;
if let Some((ts, reg)) = cache.as_ref() {
if ts.elapsed() < std::time::Duration::from_secs(300) {
return Ok(reg.clone());
}
}
}
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|e| format!("Failed to create HTTP client: {e}"))?;
let resp = dbx_core::race_download(&client, REGISTRY_PATH, REGISTRY_R2_PATH, "dbx-agent-manager")
.await
.map_err(|e| format!("Failed to fetch agent registry: {e}"))?;
let reg: AgentRegistry = resp.json().await.map_err(|e| format!("Failed to parse registry: {e}"))?;
*REGISTRY_CACHE.lock().await = Some((std::time::Instant::now(), reg.clone()));
Ok(reg)
}
fn github_url_to_r2_path(github_url: &str, category: &str) -> String {
let filename = github_url.rsplit('/').next().unwrap_or(github_url);
match category {
"jre" => format!("agents/jre/{filename}"),
"driver" => format!("agents/drivers/{filename}"),
_ => format!("agents/{filename}"),
}
}
async fn download_with_progress(
app: &tauri::AppHandle,
step: &str,