Skip to content
Open
6 changes: 6 additions & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

1 change: 1 addition & 0 deletions crates/dbx-core/examples/data_transfer_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@ fn run() -> Result<(), String> {
let parsed = ParsedImportFile {
columns: columns.clone(),
rows: rows.clone(),
source_row_numbers: (1..=rows.len()).collect(),
total_rows: rows.len(),
effective_encoding: None,
};
Expand Down
81 changes: 72 additions & 9 deletions crates/dbx-core/src/db/vector_driver.rs
Original file line number Diff line number Diff line change
Expand Up @@ -749,14 +749,21 @@ pub async fn find_documents(

pub async fn execute_rest_query(client: &VectorClient, input: &str) -> Result<QueryResult, String> {
let start = Instant::now();
let flatten_single_milvus_search = client.kind == VectorDbKind::Milvus && is_milvus_entity_search(input);
let request = parse_rest_query(client, input)?;
let resp = request.send().await.map_err(|e| format!("{} request failed: {e}", client.kind.label()))?;
let status = resp.status().as_u16();
let body = resp.json::<Value>().await.unwrap_or(Value::Null);
rest_query_result(client.kind, status, body, start)
rest_query_result(client.kind, status, body, start, flatten_single_milvus_search)
}

fn rest_query_result(kind: VectorDbKind, status: u16, body: Value, start: Instant) -> Result<QueryResult, String> {
fn rest_query_result(
kind: VectorDbKind,
status: u16,
body: Value,
start: Instant,
flatten_single_milvus_search: bool,
) -> Result<QueryResult, String> {
if !(200..300).contains(&status) {
let detail = serde_json::to_string_pretty(&body).unwrap_or_else(|_| body.to_string());
return Err(format!("{} error ({status}): {detail}", kind.label()));
Expand All @@ -765,10 +772,32 @@ fn rest_query_result(kind: VectorDbKind, status: u16, body: Value, start: Instan
if let Some(error) = milvus_business_error(&body) {
return Err(error);
}
if flatten_single_milvus_search {
if let Some(result) = milvus_single_search_to_query_result(&body, start) {
return Ok(result);
}
}
}
Ok(json_to_query_result(status, body, start))
}

fn is_milvus_entity_search(input: &str) -> bool {
input
.lines()
.find(|line| !line.trim().is_empty())
.and_then(|line| line.split_whitespace().nth(1))
.is_some_and(|path| path.split('?').next() == Some("/v2/vectordb/entities/search"))
}

fn milvus_single_search_to_query_result(body: &Value, start: Instant) -> Option<QueryResult> {
let queries = body.get("data")?.as_array()?;
if queries.len() != 1 {
return None;
}
let hits = queries.first()?.as_array()?.clone();
Some(values_to_query_result(hits, start))
}

// Milvus REST uses HTTP-style code 200 for success, while some responses use gRPC-style code 0.
fn milvus_business_error(body: &Value) -> Option<String> {
let code = body.get("code").and_then(Value::as_i64)?;
Expand Down Expand Up @@ -1009,9 +1038,9 @@ fn format_reqwest_error(err: &reqwest::Error) -> String {
#[cfg(test)]
mod tests {
use super::{
chroma_get_response_to_rows, default_collection_query, milvus_collection_schema, milvus_database_names,
rename_collection, rest_query_result, starts_with_http_method, test_connection, test_connection_request,
values_to_query_result, vector_auth, weaviate_collection_names_from_schema,
chroma_get_response_to_rows, default_collection_query, is_milvus_entity_search, milvus_collection_schema,
milvus_database_names, rename_collection, 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};
Expand Down Expand Up @@ -1097,20 +1126,54 @@ mod tests {
VectorDbKind::Milvus,
200,
json!({ "code": 1100, "message": "field kind does not exist" }),
Instant::now()
Instant::now(),
false
)
.unwrap_err(),
"Milvus error (code 1100): field kind does not exist"
);
assert!(rest_query_result(VectorDbKind::Milvus, 200, json!({ "code": 0 }), Instant::now()).is_ok());
assert!(rest_query_result(VectorDbKind::Milvus, 200, json!({ "code": 0 }), Instant::now(), false).is_ok());
assert!(rest_query_result(
VectorDbKind::Milvus,
200,
json!({ "code": 200, "data": ["kb_vectors"] }),
Instant::now()
Instant::now(),
false
)
.is_ok());
assert!(rest_query_result(VectorDbKind::Milvus, 200, json!({ "data": [] }), Instant::now()).is_ok());
assert!(rest_query_result(VectorDbKind::Milvus, 200, json!({ "data": [] }), Instant::now(), false).is_ok());
}

#[test]
fn milvus_single_search_flattens_empty_and_multiple_hits_without_affecting_query() {
assert!(is_milvus_entity_search("POST /v2/vectordb/entities/search\n{}"));
assert!(!is_milvus_entity_search("POST /v2/vectordb/entities/query\n{}"));
let empty =
rest_query_result(VectorDbKind::Milvus, 200, json!({ "code": 0, "data": [[]] }), Instant::now(), true)
.unwrap();
assert!(empty.rows.is_empty());

let hits = rest_query_result(
VectorDbKind::Milvus,
200,
json!({ "code": 0, "data": [[{"card_id": "a"}, {"card_id": "b"}]] }),
Instant::now(),
true,
)
.unwrap();
assert_eq!(hits.rows.len(), 2);
assert_eq!(hits.columns, vec!["card_id"]);

let query = rest_query_result(
VectorDbKind::Milvus,
200,
json!({ "code": 0, "data": [[{"card_id": "a"}]] }),
Instant::now(),
false,
)
.unwrap();
assert_eq!(query.rows.len(), 1);
assert_eq!(query.columns, vec!["value"]);
}

#[tokio::test]
Expand Down
100 changes: 100 additions & 0 deletions crates/dbx-core/src/query_execution_sql.rs
Original file line number Diff line number Diff line change
Expand Up @@ -255,6 +255,9 @@ pub fn is_write_sql_for_database(sql: &str, database_type: DatabaseType) -> bool
if let Some(risk) = classify_search_engine_query_risk(sql, database_type) {
return risk != SearchEngineQueryRisk::ReadOnly;
}
if let Some(risk) = classify_vector_query_risk(sql, database_type) {
return risk != SearchEngineQueryRisk::ReadOnly;
}
is_write_sql_with_database_type(sql, Some(database_type))
}

Expand Down Expand Up @@ -327,6 +330,84 @@ pub(crate) fn classify_search_engine_query_risk(
}
}

/// 对向量数据库的 REST 风格查询进行风险分类。
///
/// `POST` 在 Milvus、Qdrant 和 Chroma 中既可能是检索,也可能是写入,不能仅按 HTTP
/// 方法判断。未知端点一律视为高风险,防止 MCP 通用查询绕过专用向量工具的权限边界。
pub(crate) fn classify_vector_query_risk(source: &str, database_type: DatabaseType) -> Option<SearchEngineQueryRisk> {
if !matches!(
database_type,
DatabaseType::Qdrant | DatabaseType::Milvus | DatabaseType::Weaviate | DatabaseType::ChromaDb
) {
return None;
}
let source = strip_leading_search_engine_comments(source);
let request_line = source.lines().next()?.trim();
let mut parts = request_line.split_whitespace();
let method = parts.next()?.to_ascii_uppercase();
let path = parts.next()?.split('?').next().unwrap_or("").trim_end_matches('/').to_ascii_lowercase();

if matches!(method.as_str(), "GET" | "HEAD" | "OPTIONS") {
return Some(SearchEngineQueryRisk::ReadOnly);
}

let risk = match database_type {
DatabaseType::Milvus => match (method.as_str(), path.as_str()) {
(
"POST",
"/v2/vectordb/entities/search"
| "/v2/vectordb/entities/query"
| "/v2/vectordb/entities/get"
| "/v2/vectordb/collections/list"
| "/v2/vectordb/collections/describe"
| "/v2/vectordb/databases/list"
| "/v2/vectordb/indexes/list"
| "/v2/vectordb/indexes/describe",
) => SearchEngineQueryRisk::ReadOnly,
("POST", "/v2/vectordb/entities/insert" | "/v2/vectordb/entities/upsert") => SearchEngineQueryRisk::Write,
("POST", "/v2/vectordb/entities/delete") => SearchEngineQueryRisk::Dangerous,
("POST" | "PUT" | "PATCH" | "DELETE", _) => SearchEngineQueryRisk::Dangerous,
_ => return None,
},
DatabaseType::Qdrant => {
let read_post = ["/scroll", "/search", "/query", "/recommend", "/discover", "/count"]
.iter()
.any(|suffix| path.ends_with(suffix));
if method == "POST" && read_post {
SearchEngineQueryRisk::ReadOnly
} else if method == "PUT" && path.contains("/points") {
SearchEngineQueryRisk::Write
} else if matches!(method.as_str(), "POST" | "PUT" | "PATCH" | "DELETE") {
SearchEngineQueryRisk::Dangerous
} else {
return None;
}
}
DatabaseType::ChromaDb => {
let read_post = ["/get", "/query", "/count"].iter().any(|suffix| path.ends_with(suffix));
let safe_write_post = ["/add", "/update", "/upsert"].iter().any(|suffix| path.ends_with(suffix));
if method == "POST" && read_post {
SearchEngineQueryRisk::ReadOnly
} else if method == "POST" && safe_write_post {
SearchEngineQueryRisk::Write
} else if matches!(method.as_str(), "POST" | "PUT" | "PATCH" | "DELETE") {
SearchEngineQueryRisk::Dangerous
} else {
return None;
}
}
DatabaseType::Weaviate => {
if matches!(method.as_str(), "POST" | "PUT" | "PATCH" | "DELETE") {
SearchEngineQueryRisk::Dangerous
} else {
return None;
}
}
_ => return None,
};
Some(risk)
}

fn strip_leading_search_engine_comments(input: &str) -> &str {
let mut rest = input;
loop {
Expand Down Expand Up @@ -1832,6 +1913,25 @@ mod tests {
));
}

#[test]
fn classifies_vector_rest_queries_without_treating_every_post_as_read_only() {
assert!(!is_write_sql_for_database(
"POST /v2/vectordb/entities/search\n{\"collectionName\":\"semantic_cards\"}",
DatabaseType::Milvus,
));
assert!(is_write_sql_for_database(
"POST /v2/vectordb/entities/upsert\n{\"collectionName\":\"semantic_cards\"}",
DatabaseType::Milvus,
));
assert!(is_write_sql_for_database(
"POST /v2/vectordb/entities/delete\n{\"filter\":\"semantic_batch_id == 'x'\"}",
DatabaseType::Milvus,
));
assert!(is_write_sql_for_database("POST /v2/vectordb/collections/drop\n{}", DatabaseType::Milvus,));
assert!(!is_write_sql_for_database("POST /collections/cards/points/search\n{}", DatabaseType::Qdrant,));
assert!(is_write_sql_for_database("PUT /collections/cards/points\n{}", DatabaseType::Qdrant,));
}

#[test]
fn excludes_victoriametrics_from_sql_query_paths() {
assert!(!supports_sql_query(DatabaseType::VictoriaMetrics));
Expand Down
29 changes: 26 additions & 3 deletions crates/dbx-core/src/sql_risk.rs
Original file line number Diff line number Diff line change
Expand Up @@ -684,7 +684,9 @@ pub fn classify_sql_risk(sql: &str, dialect: &str) -> Result<SqlRisk, String> {
/// Classify SQL risk using both the parser dialect and the concrete database
/// type so dialect-specific write forms cannot be mistaken for read queries.
pub fn classify_sql_risk_for_database(sql: &str, database_type: DatabaseType) -> Result<SqlRisk, String> {
if let Some(risk) = crate::query_execution_sql::classify_search_engine_query_risk(sql, database_type) {
if let Some(risk) = crate::query_execution_sql::classify_search_engine_query_risk(sql, database_type)
.or_else(|| crate::query_execution_sql::classify_vector_query_risk(sql, database_type))
{
return Ok(match risk {
crate::query_execution_sql::SearchEngineQueryRisk::ReadOnly => SqlRisk::ReadOnly,
crate::query_execution_sql::SearchEngineQueryRisk::Write => SqlRisk::Write,
Expand All @@ -701,7 +703,9 @@ pub fn classify_sql_risk_for_database(sql: &str, database_type: DatabaseType) ->
/// and single-table UPDATE/DELETE statements with an effective predicate;
/// broader or opaque mutations require central high-risk permission.
pub fn is_dangerous_sql_for_database(sql: &str, database_type: DatabaseType) -> bool {
if let Some(risk) = crate::query_execution_sql::classify_search_engine_query_risk(sql, database_type) {
if let Some(risk) = crate::query_execution_sql::classify_search_engine_query_risk(sql, database_type)
.or_else(|| crate::query_execution_sql::classify_vector_query_risk(sql, database_type))
{
return risk == crate::query_execution_sql::SearchEngineQueryRisk::Dangerous;
}
let database_type_name = format!("{database_type:?}");
Expand Down Expand Up @@ -730,7 +734,9 @@ pub fn is_dangerous_sql_for_database(sql: &str, database_type: DatabaseType) ->
/// A `USE` statement mutates pooled/session state and could redirect later SQL,
/// so it is forbidden independently of read/write and high-risk permissions.
pub fn mcp_sql_has_forbidden_database_switch(sql: &str, database_type: DatabaseType) -> bool {
if crate::query_execution_sql::classify_search_engine_query_risk(sql, database_type).is_some() {
if crate::query_execution_sql::classify_search_engine_query_risk(sql, database_type).is_some()
|| crate::query_execution_sql::classify_vector_query_risk(sql, database_type).is_some()
{
return false;
}
let database_type_name = format!("{database_type:?}");
Expand Down Expand Up @@ -1389,4 +1395,21 @@ mod tests {
assert!(is_dangerous_sql_for_database("DELETE /products", database_type));
}
}

#[test]
fn classifies_vector_rest_risk_by_method_and_path() {
assert_eq!(
classify_sql_risk_for_database("POST /v2/vectordb/entities/search\n{}", DatabaseType::Milvus).unwrap(),
SqlRisk::ReadOnly,
);
assert_eq!(
classify_sql_risk_for_database("POST /v2/vectordb/entities/upsert\n{}", DatabaseType::Milvus).unwrap(),
SqlRisk::Write,
);
assert_eq!(
classify_sql_risk_for_database("POST /v2/vectordb/entities/delete\n{}", DatabaseType::Milvus).unwrap(),
SqlRisk::Ddl,
);
assert!(is_dangerous_sql_for_database("POST /v2/vectordb/collections/drop\n{}", DatabaseType::Milvus,));
}
}
Loading