use crate::error::Result;
use crate::pb;
use crate::types::{RecallRequest, RecallResponse, RememberRequest, RememberResponse, SearchHit};
#[derive(Debug, Clone)]
pub struct GrpcGrain {
pub blob_hash: String,
pub grain_type: String,
pub fields: serde_json::Value,
}
#[derive(Debug, Clone)]
pub struct GrpcHealth {
pub status: String,
pub version: String,
}
#[derive(Debug, Clone)]
pub struct GrpcStats {
pub total_grains: u64,
pub disk_space_bytes: u64,
pub store_size: String,
pub type_counts: std::collections::HashMap<String, u64>,
}
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, grain_type: &str, fields: &serde_json::Value) -> Result<String> {
let response = self
.client
.add(pb::AddRequest {
grain_type: grain_type.to_string(),
fields_json: serde_json::to_string(fields)?,
options: None,
})
.await?;
Ok(response.into_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,
total: inner.count,
offset: None,
has_more: false,
sources: vec![],
total_candidates: None,
results: inner
.results
.into_iter()
.map(|h| SearchHit {
blob_hash: h.hash,
grain_type: h.grain_type,
score: h.score,
subject: None,
relation: None,
object: None,
confidence: None,
namespace: if h.source_namespace.is_empty() {
None
} else {
Some(h.source_namespace)
},
tags: vec![],
temporal_type: None,
created_at: None,
fields: serde_json::from_str(&h.fields_json).unwrap_or(serde_json::Value::Null),
score_breakdown: None,
explanation: if h.explanation.is_empty() {
None
} else {
Some(h.explanation)
},
relative_time: h.relative_time,
conflict_status: None,
supersession_status: None,
recall_source: None,
})
.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,
..Default::default()
})
.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, blob_hash: &str) -> Result<GrpcGrain> {
let response = self
.client
.get(pb::GetRequest {
hash: blob_hash.to_string(),
})
.await?;
let inner = response.into_inner();
Ok(GrpcGrain {
blob_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, blob_hash: &str) -> Result<()> {
self.client
.forget(pb::ForgetRequest {
hash: blob_hash.to_string(),
})
.await?;
Ok(())
}
pub async fn health(&mut self) -> Result<GrpcHealth> {
let response = self.client.health(pb::HealthRequest {}).await?;
let inner = response.into_inner();
Ok(GrpcHealth {
status: inner.status,
version: inner.version,
})
}
pub async fn stats(&mut self) -> Result<GrpcStats> {
let response = self.client.stats(pb::StatsRequest {}).await?;
let inner = response.into_inner();
Ok(GrpcStats {
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(())
}
}