#![allow(clippy::result_large_err)]
pub mod proto {
#![allow(clippy::all)]
#![allow(clippy::doc_lazy_continuation)]
tonic::include_proto!("statelet.v1");
}
pub mod cdc;
pub use cdc::{
CheckpointStore, CommittedChange, ConsumeError, FeedItem, FeedStream, FeedTransport,
FileCheckpointStore, SubscribeCommittedOptions,
};
use proto::statelet_client::StateletClient as GrpcClient;
use tonic::transport::Channel;
#[derive(Debug, Clone)]
pub struct VectorSearchResult {
pub id: u64,
pub distance: f32,
pub group_key: String,
}
#[derive(Debug, Clone, Default)]
pub struct GroupSpec {
pub field: String,
pub group_size: u32,
pub groups: u32,
pub overfetch: u32,
pub missing_as_own: bool,
}
#[derive(Debug, Clone)]
pub struct VectorIndexConfig {
pub dim: u32,
pub metric: i32, pub m: u32,
pub m_max0: u32,
pub ef_construction: u32,
pub ef_search: u32,
}
impl Default for VectorIndexConfig {
fn default() -> Self {
Self {
dim: 128,
metric: 0,
m: 16,
m_max0: 0,
ef_construction: 200,
ef_search: 64,
}
}
}
pub enum WriteOp {
Put {
cf: u32,
key: Vec<u8>,
value: Vec<u8>,
},
Delete {
cf: u32,
key: Vec<u8>,
},
Merge {
cf: u32,
key: Vec<u8>,
value: Vec<u8>,
},
}
#[derive(Debug, Clone, Default)]
pub struct GraphQueryOptions {
pub graph_name: String,
pub max_rows: u32,
pub as_of: u64,
pub tx_as_of: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub enum GraphValue {
Null,
Int(i64),
Double(f64),
Str(String),
Bool(bool),
Json(Vec<u8>),
}
impl GraphValue {
fn from_proto(value: proto::GraphQueryValue) -> Self {
use proto::graph_query_value::Kind;
match Kind::try_from(value.kind) {
Ok(Kind::Int) => GraphValue::Int(value.int_value),
Ok(Kind::Double) => GraphValue::Double(value.dbl_value),
Ok(Kind::String) => GraphValue::Str(value.str_value),
Ok(Kind::Bool) => GraphValue::Bool(value.bool_value),
Ok(Kind::Json) => GraphValue::Json(value.json_value),
Ok(Kind::Null) | Err(_) => GraphValue::Null,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct GraphQueryResult {
pub columns: Vec<String>,
pub rows: Vec<Vec<GraphValue>>,
pub warnings: Vec<String>,
}
pub struct StateletClient {
inner: GrpcClient<Channel>,
default_cf: u32,
}
impl StateletClient {
pub async fn connect(addr: &str) -> Result<Self, tonic::transport::Error> {
let inner = GrpcClient::connect(addr.to_string()).await?;
Ok(Self {
inner,
default_cf: 0,
})
}
pub fn set_default_cf(&mut self, cf: u32) {
self.default_cf = cf;
}
pub async fn ping(&mut self) -> Result<String, tonic::Status> {
let resp = self.inner.ping(proto::PingRequest {}).await?;
Ok(resp.into_inner().message)
}
pub async fn put(
&mut self,
key: &[u8],
value: &[u8],
cf: Option<u32>,
) -> Result<(), tonic::Status> {
self.inner
.put(proto::PutRequest {
cf: cf.unwrap_or(self.default_cf),
key: key.to_vec(),
value: value.to_vec(),
..Default::default()
})
.await?;
Ok(())
}
pub async fn get(
&mut self,
key: &[u8],
cf: Option<u32>,
) -> Result<Option<Vec<u8>>, tonic::Status> {
let resp = self
.inner
.get(proto::GetRequest {
cf: cf.unwrap_or(self.default_cf),
key: key.to_vec(),
..Default::default()
})
.await?
.into_inner();
Ok(if resp.found { Some(resp.value) } else { None })
}
pub async fn delete(&mut self, key: &[u8], cf: Option<u32>) -> Result<(), tonic::Status> {
self.inner
.delete(proto::DeleteRequest {
cf: cf.unwrap_or(self.default_cf),
key: key.to_vec(),
..Default::default()
})
.await?;
Ok(())
}
pub async fn merge(
&mut self,
key: &[u8],
value: &[u8],
cf: Option<u32>,
) -> Result<(), tonic::Status> {
self.inner
.merge(proto::MergeRequest {
cf: cf.unwrap_or(self.default_cf),
key: key.to_vec(),
value: value.to_vec(),
..Default::default()
})
.await?;
Ok(())
}
pub async fn batch_write(&mut self, ops: Vec<WriteOp>) -> Result<(), tonic::Status> {
let entries = ops
.into_iter()
.map(|op| match op {
WriteOp::Put { cf, key, value } => proto::WriteEntry {
cf,
op: proto::WriteOp::Put as i32,
key,
value,
..Default::default()
},
WriteOp::Delete { cf, key } => proto::WriteEntry {
cf,
op: proto::WriteOp::Delete as i32,
key,
value: vec![],
..Default::default()
},
WriteOp::Merge { cf, key, value } => proto::WriteEntry {
cf,
op: proto::WriteOp::Merge as i32,
key,
value,
..Default::default()
},
})
.collect();
self.inner
.batch_write(proto::BatchWriteRequest {
entries,
..Default::default()
})
.await?;
Ok(())
}
pub async fn scan(
&mut self,
prefix: &[u8],
cursor: Option<&[u8]>,
limit: u32,
cf: Option<u32>,
) -> Result<(Vec<(Vec<u8>, Vec<u8>)>, Option<Vec<u8>>), tonic::Status> {
let resp = self
.inner
.scan(proto::ScanRequest {
cf: cf.unwrap_or(self.default_cf),
prefix: prefix.to_vec(),
cursor: cursor.unwrap_or(&[]).to_vec(),
limit,
..Default::default()
})
.await?
.into_inner();
let entries = resp.entries.into_iter().map(|e| (e.key, e.value)).collect();
let next = if resp.next_cursor.is_empty() {
None
} else {
Some(resp.next_cursor)
};
Ok((entries, next))
}
pub async fn delete_by_prefix(
&mut self,
prefix: &[u8],
cf: Option<u32>,
) -> Result<u32, tonic::Status> {
let resp = self
.inner
.delete_by_prefix(proto::DeleteByPrefixRequest {
cf: cf.unwrap_or(self.default_cf),
prefix: prefix.to_vec(),
..Default::default()
})
.await?
.into_inner();
Ok(resp.deleted)
}
pub async fn create_vector_index(
&mut self,
name: &str,
config: VectorIndexConfig,
) -> Result<(), tonic::Status> {
self.inner
.create_vector_index(proto::CreateVectorIndexRequest {
index_name: name.to_string(),
config: Some(proto::VectorIndexConfig {
dim: config.dim,
metric: config.metric,
m: config.m,
m_max0: config.m_max0,
ef_construction: config.ef_construction,
ef_search: config.ef_search,
..Default::default()
}),
})
.await?;
Ok(())
}
pub async fn drop_vector_index(&mut self, name: &str) -> Result<(), tonic::Status> {
self.inner
.drop_vector_index(proto::DropVectorIndexRequest {
index_name: name.to_string(),
})
.await?;
Ok(())
}
pub async fn vector_put(
&mut self,
index_name: &str,
vector_id: u64,
vector: Vec<f32>,
) -> Result<(), tonic::Status> {
self.inner
.vector_put(proto::VectorPutRequest {
index_name: index_name.to_string(),
vector_id,
vector,
attributes: Default::default(),
})
.await?;
Ok(())
}
pub async fn vector_delete(
&mut self,
index_name: &str,
vector_id: u64,
) -> Result<(), tonic::Status> {
self.inner
.vector_delete(proto::VectorDeleteRequest {
index_name: index_name.to_string(),
vector_id,
})
.await?;
Ok(())
}
pub async fn vector_search(
&mut self,
index_name: &str,
query: Vec<f32>,
k: u32,
ef_search: Option<u32>,
) -> Result<Vec<VectorSearchResult>, tonic::Status> {
self.vector_search_reranked(index_name, query, k, ef_search, None)
.await
}
pub async fn vector_search_reranked(
&mut self,
index_name: &str,
query: Vec<f32>,
k: u32,
ef_search: Option<u32>,
rerank: Option<proto::RerankSpec>,
) -> Result<Vec<VectorSearchResult>, tonic::Status> {
let resp = self
.inner
.vector_search(proto::VectorSearchRequest {
index_name: index_name.to_string(),
query,
k,
ef_search: ef_search.unwrap_or(0),
filter: None,
query_payload: None, mmr: false, mmr_lambda: 0.0,
mmr_pool: 0,
rerank, planner_override: 0, group_field: String::new(), group_size: 0,
groups: 0,
group_overfetch: 0,
group_missing_as_own: false,
})
.await?
.into_inner();
Ok(resp
.results
.into_iter()
.map(|r| VectorSearchResult {
id: r.id,
distance: r.distance,
group_key: r.group_key,
})
.collect())
}
pub async fn vector_search_grouped(
&mut self,
index_name: &str,
query: Vec<f32>,
k: u32,
ef_search: Option<u32>,
group: GroupSpec,
) -> Result<Vec<VectorSearchResult>, tonic::Status> {
let resp = self
.inner
.vector_search(proto::VectorSearchRequest {
index_name: index_name.to_string(),
query,
k,
ef_search: ef_search.unwrap_or(0),
filter: None,
query_payload: None,
mmr: false,
mmr_lambda: 0.0,
mmr_pool: 0,
rerank: None,
planner_override: 0,
group_field: group.field,
group_size: group.group_size,
groups: group.groups,
group_overfetch: group.overfetch,
group_missing_as_own: group.missing_as_own,
})
.await?
.into_inner();
Ok(resp
.results
.into_iter()
.map(|r| VectorSearchResult {
id: r.id,
distance: r.distance,
group_key: r.group_key,
})
.collect())
}
pub async fn rerank_validate(
&mut self,
index_name: &str,
mut rerank: proto::RerankSpec,
) -> Result<(), tonic::Status> {
rerank.enabled = true;
rerank.validate_only = true;
self.inner
.vector_search(proto::VectorSearchRequest {
index_name: index_name.to_string(),
query: Vec::new(),
k: 1,
ef_search: 0,
filter: None,
query_payload: None,
mmr: false,
mmr_lambda: 0.0,
mmr_pool: 0,
rerank: Some(rerank),
planner_override: 0,
group_field: String::new(),
group_size: 0,
groups: 0,
group_overfetch: 0,
group_missing_as_own: false,
})
.await?;
Ok(())
}
pub async fn vector_get(
&mut self,
index_name: &str,
vector_id: u64,
) -> Result<Option<Vec<f32>>, tonic::Status> {
let resp = self
.inner
.vector_get(proto::VectorGetRequest {
index_name: index_name.to_string(),
vector_id,
})
.await?
.into_inner();
Ok(if resp.found { Some(resp.vector) } else { None })
}
pub async fn graph_query(
&mut self,
cypher: &str,
options: GraphQueryOptions,
) -> Result<GraphQueryResult, tonic::Status> {
let resp = self
.inner
.graph_query(proto::GraphQueryRequest {
graph_name: options.graph_name,
cypher: cypher.to_string(),
max_rows: options.max_rows,
as_of: options.as_of,
tx_as_of: options.tx_as_of,
})
.await?
.into_inner();
Ok(GraphQueryResult {
columns: resp.columns,
rows: resp
.rows
.into_iter()
.map(|row| row.values.into_iter().map(GraphValue::from_proto).collect())
.collect(),
warnings: resp.warnings,
})
}
pub async fn subscribe_committed<H, E>(
&mut self,
opts: cdc::SubscribeCommittedOptions<'_>,
handler: H,
) -> Result<(), cdc::ConsumeError<E>>
where
H: FnMut(cdc::CommittedChange) -> Result<bool, E>,
{
let default_cf = self.default_cf;
let mut sleeper = cdc::TokioSleeper;
cdc::run_consumer(self, &mut sleeper, opts, default_cf, handler).await
}
}
pub struct GrpcFeedStream {
inner: tonic::Streaming<proto::CommittedFeedItem>,
}
#[tonic::async_trait]
impl cdc::FeedStream for GrpcFeedStream {
async fn recv(&mut self) -> Result<Option<cdc::FeedItem>, tonic::Status> {
match self.inner.message().await? {
Some(item) => Ok(cdc::FeedItem::from_proto(item)),
None => Ok(None),
}
}
}
#[tonic::async_trait]
impl cdc::FeedTransport for StateletClient {
type Stream = GrpcFeedStream;
async fn open_feed(
&mut self,
shard_id: u64,
from_offset: u64,
cf: u32,
key_prefix: &[u8],
include_values: bool,
) -> Result<Self::Stream, tonic::Status> {
let resp = self
.inner
.subscribe_committed(proto::SubscribeCommittedRequest {
shard_id,
from_offset,
cf,
key_prefix: key_prefix.to_vec(),
include_values,
})
.await?;
Ok(GrpcFeedStream {
inner: resp.into_inner(),
})
}
async fn scan_page(
&mut self,
prefix: &[u8],
cursor: Option<&[u8]>,
limit: u32,
cf: u32,
) -> Result<(Vec<(Vec<u8>, Vec<u8>)>, Option<Vec<u8>>), tonic::Status> {
self.scan(prefix, cursor, limit, Some(cf)).await
}
}
#[cfg(test)]
mod graph_query_tests {
use super::*;
fn value(kind: proto::graph_query_value::Kind) -> proto::GraphQueryValue {
proto::GraphQueryValue {
kind: kind as i32,
int_value: 42,
dbl_value: 0.5,
str_value: "knows".to_string(),
bool_value: true,
json_value: br#"{"name":"ada"}"#.to_vec(),
}
}
#[test]
fn decodes_every_value_kind() {
use proto::graph_query_value::Kind;
assert_eq!(GraphValue::from_proto(value(Kind::Null)), GraphValue::Null);
assert_eq!(
GraphValue::from_proto(value(Kind::Int)),
GraphValue::Int(42)
);
assert_eq!(
GraphValue::from_proto(value(Kind::Double)),
GraphValue::Double(0.5)
);
assert_eq!(
GraphValue::from_proto(value(Kind::String)),
GraphValue::Str("knows".to_string())
);
assert_eq!(
GraphValue::from_proto(value(Kind::Bool)),
GraphValue::Bool(true)
);
assert_eq!(
GraphValue::from_proto(value(Kind::Json)),
GraphValue::Json(br#"{"name":"ada"}"#.to_vec())
);
}
#[test]
fn unknown_kind_from_a_newer_server_decodes_to_null() {
let mut v = value(proto::graph_query_value::Kind::Int);
v.kind = 99;
assert_eq!(GraphValue::from_proto(v), GraphValue::Null);
}
#[test]
fn default_options_leave_every_knob_at_the_server_default() {
let o = GraphQueryOptions::default();
assert!(o.graph_name.is_empty());
assert_eq!((o.max_rows, o.as_of, o.tx_as_of), (0, 0, 0));
}
}