use crate::error::Result;
use crate::pb;
use crate::types::*;
pub struct GrpcClient {
client: pb::areev_service_client::AreevServiceClient<tonic::transport::Channel>,
}
impl GrpcClient {
pub async fn connect(endpoint: &str) -> Result<Self> {
let client = pb::areev_service_client::AreevServiceClient::connect(
endpoint.to_string(),
)
.await
.map_err(|e| tonic::Status::from_error(Box::new(e)))?;
Ok(Self { client })
}
pub async fn add(&mut self, req: &AddRequest) -> Result<AddResponse> {
let options = req.options.as_ref().map(|o| pb::AddOptions {
extract_event_date: o.extract_event_date,
auto_relate: o.auto_relate,
extract_memories: o.extract_memories,
sync: o.sync,
});
let response = self
.client
.add(pb::AddRequest {
grain_type: req.grain_type.to_string(),
fields_json: serde_json::to_string(&req.fields)?,
options,
})
.await?;
let inner = response.into_inner();
Ok(AddResponse { hash: inner.hash })
}
pub async fn recall(&mut self, req: &RecallRequest) -> Result<RecallResponse> {
let response = self
.client
.recall(pb::RecallRequest {
query: req.query.clone().unwrap_or_default(),
subject: req.subject.clone().unwrap_or_default(),
relation: req.relation.clone().unwrap_or_default(),
object: req.object.clone().unwrap_or_default(),
namespace: req.namespace.clone().unwrap_or_default(),
user_id: req.user_id.clone().unwrap_or_default(),
grain_type: req.grain_type.clone().unwrap_or_default(),
limit: req.limit.unwrap_or(10),
temporal_expr: req.temporal_expr.clone().unwrap_or_default(),
deduplicate: req.deduplicate.unwrap_or(false),
rerank: req.rerank.unwrap_or(false),
query_expansion: req.query_expansion.unwrap_or(false),
explanation: req.explanation.unwrap_or(false),
min_score: req.min_score,
diversity: req.diversity,
recency_weight: req.recency_weight,
entity: req.entity.clone(),
multi_hop: req.multi_hop,
..Default::default()
})
.await?;
let inner = response.into_inner();
Ok(RecallResponse {
count: inner.count,
results: inner
.results
.into_iter()
.map(|h| SearchHit {
hash: h.hash,
grain_type: h.grain_type,
score: h.score,
fields: serde_json::from_str(&h.fields_json)
.unwrap_or(serde_json::Value::Null),
source_namespace: if h.source_namespace.is_empty() {
None
} else {
Some(h.source_namespace)
},
explanation: if h.explanation.is_empty() {
None
} else {
Some(h.explanation)
},
relative_time: h.relative_time,
})
.collect(),
})
}
pub async fn remember(&mut self, req: &RememberRequest) -> Result<RememberResponse> {
let response = self
.client
.remember(pb::RememberRequest {
text: req.text.clone(),
sync: req.sync.unwrap_or(false),
keep_source: req.keep_source.unwrap_or(false),
namespace: req.namespace.clone().unwrap_or_default(),
user_id: req.user_id.clone().unwrap_or_default(),
tags: req.tags.clone().unwrap_or_default(),
source_type: req.source_type.clone().unwrap_or_default(),
created_at: req.created_at,
confidence: req.confidence,
})
.await?;
let inner = response.into_inner();
Ok(RememberResponse {
source_hash: inner.source_hash,
mode: inner.mode,
extracted_count: inner.extracted_count,
source_forgotten: inner.source_forgotten,
extracted_hashes: inner.extracted_hashes,
warnings: inner.warnings,
marker_status: if inner.marker_status.is_empty() {
None
} else {
Some(inner.marker_status)
},
})
}
pub async fn get(&mut self, hash: &str) -> Result<GetResponse> {
let response = self
.client
.get(pb::GetRequest {
hash: hash.to_string(),
})
.await?;
let inner = response.into_inner();
Ok(GetResponse {
hash: inner.hash,
grain_type: inner.grain_type,
fields: serde_json::from_str(&inner.fields_json)
.unwrap_or(serde_json::Value::Null),
})
}
pub async fn forget(&mut self, hash: &str) -> Result<()> {
self.client
.forget(pb::ForgetRequest {
hash: hash.to_string(),
})
.await?;
Ok(())
}
pub async fn health(&mut self) -> Result<HealthResponse> {
let response = self
.client
.health(pb::HealthRequest {})
.await?;
let inner = response.into_inner();
Ok(HealthResponse {
status: inner.status,
version: inner.version,
})
}
pub async fn stats(&mut self) -> Result<StatsResponse> {
let response = self
.client
.stats(pb::StatsRequest {})
.await?;
let inner = response.into_inner();
Ok(StatsResponse {
total_grains: inner.total_grains,
disk_space_bytes: inner.disk_space_bytes,
store_size: inner.store_size,
type_counts: inner.type_counts,
})
}
pub async fn flush(&mut self) -> Result<()> {
self.client.flush(pb::FlushRequest {}).await?;
Ok(())
}
}