fix(milvus): honor configured database in connection checks

Closes #5547
This commit is contained in:
t8y2 2026-08-06 23:40:58 +08:00
parent ba68892759
commit 1648c0b602
No known key found for this signature in database
2 changed files with 152 additions and 42 deletions

View File

@ -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)
}

View File

@ -89,6 +89,7 @@ pub struct VectorClient {
http: HttpClient,
base_url: String,
auth: Option<VectorAuth>,
database: Option<String>,
}
#[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<VectorAuth> {
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<Vec<CollectionInfo>, String> {
@ -208,25 +222,30 @@ pub async fn list_databases(client: &VectorClient) -> Result<Vec<String>, String
}
async fn list_milvus_databases(client: &VectorClient) -> Result<Vec<String>, 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<String> {
let mut names: Vec<String> = 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<String> {
@ -237,7 +256,7 @@ fn milvus_database_name_from_item(item: &Value) -> Option<String> {
}
async fn list_qdrant_collections(client: &VectorClient) -> Result<Vec<CollectionInfo>, String> {
let body = send_json(client.get("/collections"), "Qdrant").await?;
let body = send_json(client.get("/collections"), client.kind).await?;
let mut infos: Vec<CollectionInfo> = 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<CollectionInfo> = match body.get("data") {
@ -281,7 +300,7 @@ fn collection_name_from_milvus_item(item: &Value) -> Option<String> {
}
async fn list_weaviate_collections(client: &VectorClient) -> Result<Vec<CollectionInfo>, 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<CollectionInfo> = 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<Vec<Collecti
async fn list_chroma_collections(client: &VectorClient) -> Result<Vec<CollectionInfo>, 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<CollectionInfo> = body
.as_array()
@ -325,7 +344,7 @@ pub async fn get_collection_detail(
async fn get_weaviate_collection_detail(client: &VectorClient, collection: &str) -> Result<CollectionInfo, String> {
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<CollectionInfo, String> {
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<Value, String> {
async fn send_json(req: reqwest::RequestBuilder, kind: VectorDbKind) -> Result<Value, String> {
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<reqwest::Response, String> {
@ -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::<Value>(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::<Value>(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 =