From 1648c0b60294cc2b140faa9636efb3a028524055 Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Thu, 6 Aug 2026 23:40:58 +0800 Subject: [PATCH] fix(milvus): honor configured database in connection checks Closes #5547 --- crates/dbx-core/src/connection.rs | 3 +- crates/dbx-core/src/db/vector_driver.rs | 191 +++++++++++++++++++----- 2 files changed, 152 insertions(+), 42 deletions(-) diff --git a/crates/dbx-core/src/connection.rs b/crates/dbx-core/src/connection.rs index ec7e9178c..8ce48acb9 100644 --- a/crates/dbx-core/src/connection.rs +++ b/crates/dbx-core/src/connection.rs @@ -2014,7 +2014,8 @@ impl AppState { Some(&db_config.password), db_config.ssl, connect_timeout, - ); + ) + .with_database(db_config.database.as_deref()); db::vector_driver::test_connection(&client, connect_timeout).await?; PoolKind::VectorDb(client) } diff --git a/crates/dbx-core/src/db/vector_driver.rs b/crates/dbx-core/src/db/vector_driver.rs index 58b377ad0..661175e6d 100644 --- a/crates/dbx-core/src/db/vector_driver.rs +++ b/crates/dbx-core/src/db/vector_driver.rs @@ -89,6 +89,7 @@ pub struct VectorClient { http: HttpClient, base_url: String, auth: Option, + database: Option, } #[derive(Clone, Debug, PartialEq, Eq)] @@ -112,7 +113,16 @@ impl VectorClient { let auth = vector_auth(kind, username, password); let builder = http_client_builder(timeout).danger_accept_invalid_certs(accept_invalid_certs); let http = builder.build().unwrap_or_else(|_| HttpClient::new()); - Self { kind, http, base_url, auth } + Self { kind, http, base_url, auth, database: None } + } + + pub fn with_database(mut self, database: Option<&str>) -> Self { + self.database = database.map(str::trim).filter(|database| !database.is_empty()).map(str::to_string); + self + } + + fn database_or_default(&self) -> &str { + self.database.as_deref().unwrap_or("default") } fn get(&self, path: &str) -> reqwest::RequestBuilder { @@ -142,6 +152,21 @@ impl VectorClient { } } +fn test_connection_request(client: &VectorClient) -> reqwest::RequestBuilder { + let path = match client.kind { + VectorDbKind::Qdrant => "/collections", + VectorDbKind::Milvus => "/v2/vectordb/collections/list", + VectorDbKind::Weaviate => "/v1/meta", + VectorDbKind::ChromaDb => "/api/v2/heartbeat", + }; + match client.kind { + VectorDbKind::Qdrant => client.get(path), + VectorDbKind::Milvus => client.post(path).json(&serde_json::json!({ "dbName": client.database_or_default() })), + VectorDbKind::Weaviate => client.get(path), + VectorDbKind::ChromaDb => client.get(path), + } +} + fn vector_auth(kind: VectorDbKind, username: Option<&str>, password: Option<&str>) -> Option { let username = username.unwrap_or("").trim(); let password = password.unwrap_or(""); @@ -162,23 +187,12 @@ fn vector_auth(kind: VectorDbKind, username: Option<&str>, password: Option<&str pub async fn test_connection(client: &VectorClient, timeout: Duration) -> Result<(), String> { let label = client.kind.label(); - let path = match client.kind { - VectorDbKind::Qdrant => "/collections", - VectorDbKind::Milvus => "/v2/vectordb/collections/list", - VectorDbKind::Weaviate => "/v1/meta", - VectorDbKind::ChromaDb => "/api/v2/heartbeat", - }; - let request = match client.kind { - VectorDbKind::Qdrant => client.get(path), - VectorDbKind::Milvus => client.post(path).json(&serde_json::json!({ "dbName": "default" })), - VectorDbKind::Weaviate => client.get(path), - VectorDbKind::ChromaDb => client.get(path), - }; - let resp = with_connection_timeout(label, timeout, async { - request.send().await.map_err(|e| format!("{label} connection failed: {}", format_reqwest_error(&e))) + let request = test_connection_request(client); + with_connection_timeout(label, timeout, async { + send_json(request, client.kind).await.map_err(|error| error.replacen("request failed", "connection failed", 1)) }) - .await?; - ensure_success(label, resp).await.map(|_| ()) + .await + .map(|_| ()) } pub async fn list_collections(client: &VectorClient) -> Result, String> { @@ -208,25 +222,30 @@ pub async fn list_databases(client: &VectorClient) -> Result, String } async fn list_milvus_databases(client: &VectorClient) -> Result, String> { - // Older Milvus versions (pre-2.2) do not expose the databases endpoint; fall back to "default" - // so the connection stays browsable instead of failing the whole tree load. + // Older Milvus versions (pre-2.2) do not expose the databases endpoint; fall back to the + // configured database (or "default") so the connection stays browsable instead of failing the whole tree load. // // The endpoint rejects a bodyless POST with `{"code":1801,...}` (HTTP 200, no `data` field), // so send an empty JSON object like every other Milvus v2 endpoint. - let body = match send_json(client.post("/v2/vectordb/databases/list").json(&serde_json::json!({})), "Milvus").await - { - Ok(body) => body, - Err(_) => return Ok(vec!["default".to_string()]), - }; + let body = + match send_json(client.post("/v2/vectordb/databases/list").json(&serde_json::json!({})), client.kind).await { + Ok(body) => body, + Err(_) => return Ok(vec![client.database_or_default().to_string()]), + }; + Ok(milvus_database_names(&body, client.database_or_default())) +} + +fn milvus_database_names(body: &Value, configured_database: &str) -> Vec { let mut names: Vec = match body.get("data") { Some(Value::Array(items)) => items.iter().filter_map(milvus_database_name_from_item).collect(), _ => Vec::new(), }; - if !names.iter().any(|name| name == "default") { - names.push("default".to_string()); + if !names.iter().any(|name| name == configured_database) { + names.push(configured_database.to_string()); } names.sort(); - Ok(names) + names.dedup(); + names } fn milvus_database_name_from_item(item: &Value) -> Option { @@ -237,7 +256,7 @@ fn milvus_database_name_from_item(item: &Value) -> Option { } async fn list_qdrant_collections(client: &VectorClient) -> Result, String> { - let body = send_json(client.get("/collections"), "Qdrant").await?; + let body = send_json(client.get("/collections"), client.kind).await?; let mut infos: Vec = body .pointer("/result/collections") .and_then(Value::as_array) @@ -256,7 +275,7 @@ async fn list_milvus_collections(client: &VectorClient, database: &str) -> Resul let db_name = if database.is_empty() { "default" } else { database }; let body = send_json( client.post("/v2/vectordb/collections/list").json(&serde_json::json!({ "dbName": db_name })), - "Milvus", + client.kind, ) .await?; let mut infos: Vec = match body.get("data") { @@ -281,7 +300,7 @@ fn collection_name_from_milvus_item(item: &Value) -> Option { } async fn list_weaviate_collections(client: &VectorClient) -> Result, String> { - let body = send_json(client.get("/v1/schema"), "Weaviate").await?; + let body = send_json(client.get("/v1/schema"), client.kind).await?; let mut infos: Vec = weaviate_collection_names_from_schema(&body) .into_iter() .map(|name| CollectionInfo { name: name.clone(), id: name, ..Default::default() }) @@ -292,7 +311,7 @@ async fn list_weaviate_collections(client: &VectorClient) -> Result Result, String> { let body = - send_json(client.get("/api/v2/tenants/default_tenant/databases/default_database/collections"), "ChromaDB") + send_json(client.get("/api/v2/tenants/default_tenant/databases/default_database/collections"), client.kind) .await?; let mut infos: Vec = body .as_array() @@ -325,7 +344,7 @@ pub async fn get_collection_detail( async fn get_weaviate_collection_detail(client: &VectorClient, collection: &str) -> Result { let query = format!("{{ Get {{ {collection}(limit: 1) {{ _additional {{ vector }} }} }} }}"); let dimension = - match send_json(client.post("/v1/graphql").json(&serde_json::json!({ "query": query })), "Weaviate").await { + match send_json(client.post("/v1/graphql").json(&serde_json::json!({ "query": query })), client.kind).await { Ok(body) => weaviate_vector_dimension_from_graphql(&body, collection), Err(_) => None, }; @@ -333,7 +352,7 @@ async fn get_weaviate_collection_detail(client: &VectorClient, collection: &str) } async fn get_qdrant_collection_detail(client: &VectorClient, collection: &str) -> Result { - let body = send_json(client.get(&format!("/collections/{}", path_segment(collection))), "Qdrant").await?; + let body = send_json(client.get(&format!("/collections/{}", path_segment(collection))), client.kind).await?; let dim = body .pointer("/result/config/params/vectors/size") .and_then(Value::as_u64) @@ -431,7 +450,7 @@ async fn get_milvus_collection_detail( client .post("/v2/vectordb/collections/describe") .json(&serde_json::json!({ "dbName": db_name, "collectionName": collection })), - "Milvus", + client.kind, ) .await?; if body.get("code").and_then(Value::as_i64) != Some(0) { @@ -459,7 +478,7 @@ async fn get_chroma_collection_detail(client: &VectorClient, collection: &str) - "/api/v2/tenants/default_tenant/databases/default_database/collections/{}", path_segment(collection) )), - "ChromaDB", + client.kind, ) .await?; let name = body.get("name").and_then(Value::as_str).unwrap_or(collection); @@ -754,10 +773,17 @@ pub(crate) fn query_value(value: &str) -> String { utf8_percent_encode(value, QUERY_VALUE_ENCODE_SET).to_string() } -async fn send_json(req: reqwest::RequestBuilder, label: &str) -> Result { +async fn send_json(req: reqwest::RequestBuilder, kind: VectorDbKind) -> Result { + let label = kind.label(); let resp = req.send().await.map_err(|e| format!("{label} request failed: {e}"))?; let resp = ensure_success(label, resp).await?; - resp.json().await.map_err(|e| format!("{label} parse error: {e}")) + let body = resp.json().await.map_err(|e| format!("{label} parse error: {e}"))?; + if kind == VectorDbKind::Milvus { + if let Some(error) = milvus_business_error(&body) { + return Err(error); + } + } + Ok(body) } async fn ensure_success(label: &str, resp: reqwest::Response) -> Result { @@ -881,12 +907,33 @@ fn format_reqwest_error(err: &reqwest::Error) -> String { #[cfg(test)] mod tests { use super::{ - chroma_get_response_to_rows, milvus_collection_schema, rest_query_result, starts_with_http_method, - values_to_query_result, vector_auth, weaviate_collection_names_from_schema, - weaviate_vector_dimension_from_graphql, CollectionInfo, VectorAuth, VectorDbKind, + chroma_get_response_to_rows, milvus_collection_schema, milvus_database_names, rest_query_result, + starts_with_http_method, test_connection, test_connection_request, values_to_query_result, vector_auth, + weaviate_collection_names_from_schema, weaviate_vector_dimension_from_graphql, CollectionInfo, VectorAuth, + VectorClient, VectorDbKind, }; use serde_json::{json, Value}; - use std::time::Instant; + use std::time::{Duration, Instant}; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + async fn spawn_json_response_server(body: Value) -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let body = body.to_string(); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut request = [0_u8; 2048]; + stream.read(&mut request).await.unwrap(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + body.len(), + body + ); + stream.write_all(response.as_bytes()).await.unwrap(); + }); + (format!("http://{address}"), server) + } #[test] fn detects_rest_queries_case_insensitively() { @@ -910,6 +957,68 @@ mod tests { assert!(rest_query_result(VectorDbKind::Milvus, 200, json!({ "code": 0 }), Instant::now()).is_ok()); } + #[tokio::test] + async fn milvus_connection_test_rejects_business_errors() { + let (url, server) = spawn_json_response_server(json!({ + "code": 800, + "message": "database not found[database=resume_test]" + })) + .await; + let client = VectorClient::new(VectorDbKind::Milvus, &url, None, None, false, Duration::from_secs(1)) + .with_database(Some("resume_test")); + + let error = test_connection(&client, Duration::from_secs(1)).await.unwrap_err(); + + assert_eq!(error, "Milvus error (code 800): database not found[database=resume_test]"); + server.await.unwrap(); + } + + #[test] + fn milvus_connection_test_uses_configured_database() { + let client = VectorClient::new( + VectorDbKind::Milvus, + "http://localhost:19530", + None, + None, + false, + Duration::from_secs(1), + ) + .with_database(Some(" resume_test ")); + let request = test_connection_request(&client).build().unwrap(); + let body = request.body().and_then(reqwest::Body::as_bytes).unwrap(); + + assert_eq!(serde_json::from_slice::(body).unwrap(), json!({ "dbName": "resume_test" })); + } + + #[test] + fn milvus_connection_test_defaults_empty_database() { + let client = VectorClient::new( + VectorDbKind::Milvus, + "http://localhost:19530", + None, + None, + false, + Duration::from_secs(1), + ) + .with_database(Some(" ")); + let request = test_connection_request(&client).build().unwrap(); + let body = request.body().and_then(reqwest::Body::as_bytes).unwrap(); + + assert_eq!(serde_json::from_slice::(body).unwrap(), json!({ "dbName": "default" })); + } + + #[test] + fn milvus_database_list_keeps_the_configured_database() { + assert_eq!( + milvus_database_names(&json!({ "data": ["default"] }), "resume_test"), + vec!["default".to_string(), "resume_test".to_string()] + ); + assert_eq!( + milvus_database_names(&json!({ "data": ["resume_test"] }), "resume_test"), + vec!["resume_test".to_string()] + ); + } + #[test] fn flattens_qdrant_payload_columns() { let result =