use crate::{
EsDocument, EsError, EsHit, EsQuery, EsSearchRequest, EsSearchResult, EsSortOrder, EsSync,
EsSyncResult,
};
use elasticsearch::{
auth::Credentials,
http::{
transport::{SingleNodeConnectionPool, TransportBuilder},
Url,
},
BulkOperation, BulkParts, DeleteParts, Elasticsearch, IndexParts, SearchParts, UpdateParts,
};
pub struct RealEsClient {
client: Elasticsearch,
runtime: tokio::runtime::Runtime,
}
impl RealEsClient {
pub fn new(url: &str) -> Result<Self, EsError> {
let transport = build_transport(url, None)?;
let client = Elasticsearch::new(transport);
let runtime = tokio::runtime::Runtime::new()
.map_err(|e| EsError::ConnectionFailed(format!("创建 tokio 运行时失败: {}", e)))?;
Ok(Self { client, runtime })
}
pub fn with_auth(url: &str, username: &str, password: &str) -> Result<Self, EsError> {
let credentials = Credentials::Basic(username.to_string(), password.to_string());
let transport = build_transport(url, Some(credentials))?;
let client = Elasticsearch::new(transport);
let runtime = tokio::runtime::Runtime::new()
.map_err(|e| EsError::ConnectionFailed(format!("创建 tokio 运行时失败: {}", e)))?;
Ok(Self { client, runtime })
}
pub fn index_doc(
&self,
index: &str,
id: &str,
doc: &serde_json::Value,
) -> Result<(), EsError> {
self.runtime.block_on(async {
let response = self
.client
.index(IndexParts::IndexId(index, id))
.body(doc.clone())
.send()
.await
.map_err(|e| EsError::SyncError(e.to_string()))?;
if !response.status_code().is_success() {
return Err(EsError::SyncError(format!(
"索引文档失败, HTTP {}",
response.status_code()
)));
}
Ok(())
})
}
pub fn search(
&self,
index: &str,
query: &serde_json::Value,
) -> Result<serde_json::Value, EsError> {
self.runtime.block_on(async {
let response = self
.client
.search(SearchParts::Index(&[index]))
.body(query.clone())
.send()
.await
.map_err(|e| EsError::QueryError(e.to_string()))?;
if !response.status_code().is_success() {
return Err(EsError::QueryError(format!(
"搜索失败, HTTP {}",
response.status_code()
)));
}
let body: serde_json::Value = response
.json()
.await
.map_err(|e| EsError::QueryError(e.to_string()))?;
Ok(body)
})
}
pub fn update_doc(
&self,
index: &str,
id: &str,
doc: &serde_json::Value,
) -> Result<(), EsError> {
self.runtime.block_on(async {
let body = serde_json::json!({ "doc": doc });
let response = self
.client
.update(UpdateParts::IndexId(index, id))
.body(body)
.send()
.await
.map_err(|e| EsError::SyncError(e.to_string()))?;
if !response.status_code().is_success() {
return Err(EsError::SyncError(format!(
"更新文档失败, HTTP {}",
response.status_code()
)));
}
Ok(())
})
}
pub fn delete_doc(&self, index: &str, id: &str) -> Result<(), EsError> {
self.runtime.block_on(async {
let response = self
.client
.delete(DeleteParts::IndexId(index, id))
.send()
.await
.map_err(|e| EsError::SyncError(e.to_string()))?;
if !response.status_code().is_success() && response.status_code() != 404 {
return Err(EsError::SyncError(format!(
"删除文档失败, HTTP {}",
response.status_code()
)));
}
Ok(())
})
}
pub fn bulk_index(
&self,
index: &str,
docs: &[(String, serde_json::Value)],
) -> Result<(), EsError> {
if docs.is_empty() {
return Ok(());
}
self.runtime.block_on(async {
let ops: Vec<BulkOperation<serde_json::Value>> = docs
.iter()
.map(|(id, doc)| {
BulkOperation::index(doc.clone())
.id(id.clone())
.into()
})
.collect();
let response = self
.client
.bulk(BulkParts::Index(index))
.body(ops)
.send()
.await
.map_err(|e| EsError::SyncError(e.to_string()))?;
if !response.status_code().is_success() {
return Err(EsError::SyncError(format!(
"批量索引失败, HTTP {}",
response.status_code()
)));
}
let body: serde_json::Value = response
.json()
.await
.map_err(|e| EsError::SyncError(e.to_string()))?;
if body
.get("errors")
.and_then(|e| e.as_bool())
.unwrap_or(false)
{
let error_msgs: Vec<String> = body
.get("items")
.and_then(|i| i.as_array())
.map(|items| {
items
.iter()
.filter_map(|item| {
let idx = item.get("index")?;
if idx.get("error").is_some() {
let id = idx
.get("_id")
.and_then(|v| v.as_str())
.unwrap_or("unknown");
let error = idx
.get("error")
.and_then(|e| e.get("reason"))
.and_then(|r| r.as_str())
.unwrap_or("unknown error");
Some(format!("文档 {}: {}", id, error))
} else {
None
}
})
.collect()
})
.unwrap_or_default();
return Err(EsError::SyncError(format!(
"批量索引部分失败 ({}): {}",
error_msgs.len(),
error_msgs.join("; ")
)));
}
Ok(())
})
}
}
fn build_transport(
url: &str,
credentials: Option<Credentials>,
) -> Result<elasticsearch::http::transport::Transport, EsError> {
let parsed_url = Url::parse(url)
.map_err(|e| EsError::ConnectionFailed(format!("无效的 URL: {}", e)))?;
let conn_pool = SingleNodeConnectionPool::new(parsed_url);
let mut builder = TransportBuilder::new(conn_pool);
if let Some(creds) = credentials {
builder = builder.auth(creds);
}
builder
.build()
.map_err(|e| EsError::ConnectionFailed(format!("构建 Transport 失败: {}", e)))
}
impl EsSync for RealEsClient {
fn sync_to_es(&self, documents: Vec<EsDocument>) -> Result<EsSyncResult, EsError> {
if documents.is_empty() {
return Ok(EsSyncResult::success(0));
}
let mut by_index: std::collections::HashMap<String, Vec<(String, serde_json::Value)>> =
std::collections::HashMap::new();
for (i, doc) in documents.into_iter().enumerate() {
if doc.index.is_empty() {
continue;
}
let id = doc
.id
.clone()
.unwrap_or_else(|| format!("auto-{}-{}", doc.timestamp, i));
by_index
.entry(doc.index.clone())
.or_default()
.push((id, doc.source));
}
let mut total_indexed = 0usize;
let mut errors: Vec<String> = Vec::new();
for (index, docs) in by_index {
match self.bulk_index(&index, &docs) {
Ok(()) => total_indexed += docs.len(),
Err(e) => {
errors.push(format!("索引 {} 失败: {}", index, e));
}
}
}
if errors.is_empty() {
Ok(EsSyncResult::success(total_indexed))
} else {
Ok(EsSyncResult::with_errors(total_indexed, errors))
}
}
fn delete_from_es(&self, index: &str, ids: Vec<String>) -> Result<EsSyncResult, EsError> {
let mut deleted = 0usize;
let mut errors: Vec<String> = Vec::new();
for id in &ids {
match self.delete_doc(index, id) {
Ok(()) => deleted += 1,
Err(e) => errors.push(format!("文档 {}: {}", id, e)),
}
}
if errors.is_empty() {
Ok(EsSyncResult::success(deleted))
} else {
Ok(EsSyncResult::with_errors(deleted, errors))
}
}
fn search(&self, request: EsSearchRequest) -> Result<EsSearchResult, EsError> {
let dsl = build_es_dsl(&request);
let response = RealEsClient::search(self, &request.index, &dsl)?;
parse_search_response(response)
}
}
fn build_es_dsl(request: &EsSearchRequest) -> serde_json::Value {
let query = es_query_to_dsl(&request.query);
let mut dsl = serde_json::json!({
"query": query,
"from": request.from,
"size": request.size,
});
if !request.sort.is_empty() {
let sort: Vec<serde_json::Value> = request
.sort
.iter()
.map(|s| {
let order = match s.order {
EsSortOrder::Asc => "asc",
EsSortOrder::Desc => "desc",
};
serde_json::json!({ &s.field: { "order": order } })
})
.collect();
dsl["sort"] = serde_json::Value::Array(sort);
}
dsl
}
fn es_query_to_dsl(query: &EsQuery) -> serde_json::Value {
match query {
EsQuery::MatchAll => serde_json::json!({ "match_all": {} }),
EsQuery::Term(terms) => {
if terms.len() == 1 {
let (field, value) = terms.iter().next().expect("len()==1 guarantees non-empty");
serde_json::json!({ "term": { field: value } })
} else {
let must: Vec<serde_json::Value> = terms
.iter()
.map(|(field, value)| {
serde_json::json!({ "term": { field: value } })
})
.collect();
serde_json::json!({ "bool": { "must": must } })
}
}
EsQuery::Terms(terms) => {
if terms.len() == 1 {
let (field, values) = terms.iter().next().expect("len()==1 guarantees non-empty");
serde_json::json!({ "terms": { field: values } })
} else {
let must: Vec<serde_json::Value> = terms
.iter()
.map(|(field, values)| {
serde_json::json!({ "terms": { field: values } })
})
.collect();
serde_json::json!({ "bool": { "must": must } })
}
}
EsQuery::Range(ranges) => {
if ranges.len() == 1 {
let (field, range) = ranges.iter().next().expect("len()==1 guarantees non-empty");
let range_obj = build_range_obj(range);
serde_json::json!({ "range": { field: range_obj } })
} else {
let must: Vec<serde_json::Value> = ranges
.iter()
.map(|(field, range)| {
let range_obj = build_range_obj(range);
serde_json::json!({ "range": { field: range_obj } })
})
.collect();
serde_json::json!({ "bool": { "must": must } })
}
}
EsQuery::Bool(b) => {
let mut bool_obj = serde_json::Map::new();
if let Some(must) = &b.must {
bool_obj.insert(
"must".to_string(),
serde_json::Value::Array(must.iter().map(es_query_to_dsl).collect()),
);
}
if let Some(should) = &b.should {
bool_obj.insert(
"should".to_string(),
serde_json::Value::Array(should.iter().map(es_query_to_dsl).collect()),
);
}
if let Some(filter) = &b.filter {
bool_obj.insert(
"filter".to_string(),
serde_json::Value::Array(filter.iter().map(es_query_to_dsl).collect()),
);
}
if let Some(must_not) = &b.must_not {
bool_obj.insert(
"must_not".to_string(),
serde_json::Value::Array(must_not.iter().map(es_query_to_dsl).collect()),
);
}
if let Some(min_match) = &b.minimum_should_match {
bool_obj.insert(
"minimum_should_match".to_string(),
serde_json::Value::from(*min_match as u64),
);
}
serde_json::json!({ "bool": bool_obj })
}
}
}
fn build_range_obj(range: &crate::EsRangeQuery) -> serde_json::Value {
let mut obj = serde_json::Map::new();
if let Some(gt) = &range.gt {
obj.insert("gt".to_string(), gt.clone());
}
if let Some(gte) = &range.gte {
obj.insert("gte".to_string(), gte.clone());
}
if let Some(lt) = &range.lt {
obj.insert("lt".to_string(), lt.clone());
}
if let Some(lte) = &range.lte {
obj.insert("lte".to_string(), lte.clone());
}
serde_json::Value::Object(obj)
}
fn parse_search_response(response: serde_json::Value) -> Result<EsSearchResult, EsError> {
let took = response
.get("took")
.and_then(|v| v.as_i64())
.unwrap_or(0);
let total = response
.get("hits")
.and_then(|h| h.get("total"))
.and_then(|t| t.get("value"))
.and_then(|v| v.as_u64())
.unwrap_or(0) as usize;
let hits: Vec<EsHit> = response
.get("hits")
.and_then(|h| h.get("hits"))
.and_then(|h| h.as_array())
.map(|arr| {
arr.iter()
.filter_map(|hit| {
let id = hit.get("_id")?.as_str()?.to_string();
let score = hit
.get("_score")
.and_then(|s| s.as_f64())
.unwrap_or(0.0);
let source = hit.get("_source")?.clone();
Some(EsHit { id, score, source })
})
.collect()
})
.unwrap_or_default();
Ok(EsSearchResult { total, hits, took })
}
#[cfg(test)]
mod tests {
use super::*;
use crate::EsQuery;
use serde_json::json;
#[test]
fn test_real_es_client_new() {
let client = RealEsClient::new("http://localhost:9200");
assert!(client.is_ok(), "构造客户端应成功: {:?}", client.err());
}
#[test]
fn test_real_es_client_with_auth() {
let client = RealEsClient::with_auth("http://localhost:9200", "user", "pass");
assert!(client.is_ok(), "构造带认证客户端应成功: {:?}", client.err());
}
#[test]
fn test_real_es_client_invalid_url() {
let result = RealEsClient::new("");
assert!(result.is_err(), "空 URL 应构造失败");
let result = RealEsClient::with_auth("not-a-url", "u", "p");
assert!(result.is_err(), "无效 URL 应构造失败");
}
#[test]
fn test_build_dsl_match_all() {
let req = EsSearchRequest::new("idx", EsQuery::match_all())
.with_pagination(0, 20);
let dsl = build_es_dsl(&req);
assert_eq!(dsl["query"]["match_all"], json!({}));
assert_eq!(dsl["from"], 0);
assert_eq!(dsl["size"], 20);
}
#[test]
fn test_build_dsl_term_with_sort() {
let req = EsSearchRequest::new("idx", EsQuery::term("status", json!("active")))
.with_pagination(10, 5)
.with_sort("date", EsSortOrder::Desc);
let dsl = build_es_dsl(&req);
assert_eq!(dsl["query"]["term"]["status"], json!("active"));
assert_eq!(dsl["from"], 10);
assert_eq!(dsl["size"], 5);
assert_eq!(dsl["sort"][0]["date"]["order"], "desc");
}
#[test]
fn test_build_dsl_bool_query() {
let bool_q = EsQuery::must(vec![
EsQuery::term("status", json!("active")),
EsQuery::range(
"age",
crate::EsRangeQuery::new().gte(json!(18)),
),
]);
let req = EsSearchRequest::new("idx", bool_q);
let dsl = build_es_dsl(&req);
assert!(dsl["query"]["bool"]["must"].is_array());
assert_eq!(dsl["query"]["bool"]["must"].as_array().unwrap().len(), 2);
}
#[test]
fn test_parse_search_response() {
let response = json!({
"took": 5,
"hits": {
"total": { "value": 2, "relation": "eq" },
"hits": [
{
"_id": "1",
"_score": 1.5,
"_source": { "name": "alice" }
},
{
"_id": "2",
"_score": 0.8,
"_source": { "name": "bob" }
}
]
}
});
let result = parse_search_response(response).unwrap();
assert_eq!(result.took, 5);
assert_eq!(result.total, 2);
assert_eq!(result.hits.len(), 2);
assert_eq!(result.hits[0].id, "1");
assert_eq!(result.hits[0].score, 1.5);
assert_eq!(result.hits[0].source["name"], "alice");
}
#[test]
fn test_parse_empty_search_response() {
let response = json!({
"took": 0,
"hits": {
"total": { "value": 0, "relation": "eq" },
"hits": []
}
});
let result = parse_search_response(response).unwrap();
assert_eq!(result.total, 0);
assert!(result.hits.is_empty());
}
#[test]
fn test_es_document_serde() {
let doc = EsDocument::new("test-index", json!({"name": "test", "value": 42}))
.with_id("doc1");
let serialized = serde_json::to_string(&doc).unwrap();
let deserialized: EsDocument = serde_json::from_str(&serialized).unwrap();
assert_eq!(deserialized.index, "test-index");
assert_eq!(deserialized.id, Some("doc1".to_string()));
assert_eq!(deserialized.source["name"], "test");
assert_eq!(deserialized.source["value"], 42);
}
#[test]
fn test_real_es_client_as_trait_object() {
let client = RealEsClient::new("http://localhost:9200").unwrap();
let _boxed: Box<dyn EsSync> = Box::new(client);
}
}