use tonic::transport::{Channel, Endpoint};
use tonic::{Request, Status};
use crate::error::{CallPolicy, UdbError};
use crate::metadata::Metadata;
use crate::proto::udb::entity::v1 as entity;
use crate::proto::udb::services::v1::data_broker_client::DataBrokerClient;
#[derive(Clone, Debug)]
pub struct UdbClient {
inner: DataBrokerClient<Channel>,
meta: Metadata,
policy: Option<CallPolicy>,
}
macro_rules! rpc {
($(#[$m:meta])* $name:ident, $rpc:ident, $req:ty, $resp:ty, $path:expr) => {
$(#[$m])*
pub async fn $name(&mut self, req: $req) -> Result<$resp, UdbError> {
let policy = self.policy.unwrap_or_else(|| CallPolicy::from_contract($path));
let mut attempt: u32 = 1;
loop {
let mut request = self.request(req.clone())?;
if let Some(deadline) = policy.deadline {
request.set_timeout(deadline);
}
match self.inner.$rpc(request).await {
Ok(response) => return Ok(response.into_inner()),
Err(status) => {
let err = UdbError::from_status(status);
if policy.should_retry(attempt, &err) {
tokio::time::sleep(policy.backoff_for(attempt, &err)).await;
attempt += 1;
continue;
}
return Err(err);
}
}
}
}
};
}
impl UdbClient {
pub async fn connect(
endpoint: impl Into<String>,
meta: Metadata,
) -> Result<Self, tonic::transport::Error> {
let channel = Endpoint::from_shared(endpoint.into())?.connect().await?;
Ok(Self::with_channel(channel, meta))
}
pub fn with_channel(channel: Channel, meta: Metadata) -> Self {
Self {
inner: DataBrokerClient::new(channel),
meta,
policy: None,
}
}
pub fn metadata(&self) -> &Metadata {
&self.meta
}
pub fn with_audit(
&self,
purpose: impl Into<String>,
correlation_id: impl Into<String>,
) -> Self {
Self {
inner: self.inner.clone(),
meta: self.meta.clone().with_audit(purpose, correlation_id),
policy: self.policy,
}
}
pub fn with_bearer_token(&self, token: impl Into<String>) -> Self {
Self {
inner: self.inner.clone(),
meta: self.meta.clone().with_bearer_token(token),
policy: self.policy,
}
}
pub fn request<T>(&self, message: T) -> Result<Request<T>, Status> {
let mut req = Request::new(message);
self.meta.apply(&mut req)?;
Ok(req)
}
pub fn raw(&mut self) -> &mut DataBrokerClient<Channel> {
&mut self.inner
}
pub fn with_policy(&self, policy: CallPolicy) -> Self {
Self {
inner: self.inner.clone(),
meta: self.meta.clone(),
policy: Some(policy),
}
}
rpc!(
select,
select,
entity::SelectRequest,
entity::RecordSet,
"/udb.services.v1.DataBroker/Select"
);
rpc!(
upsert,
upsert,
entity::UpsertRequest,
entity::MutationResponse,
"/udb.services.v1.DataBroker/Upsert"
);
rpc!(
update,
update,
entity::UpdateRequest,
entity::MutationResponse,
"/udb.services.v1.DataBroker/Update"
);
rpc!(
delete,
delete,
entity::DeleteRequest,
entity::MutationResponse,
"/udb.services.v1.DataBroker/Delete"
);
rpc!(
bulk_cas,
bulk_cas,
entity::BulkCasRequest,
entity::BulkCasResponse,
"/udb.services.v1.DataBroker/BulkCas"
);
rpc!(
vector_search,
vector_search,
entity::VectorSearchRequest,
entity::VectorSet,
"/udb.services.v1.DataBroker/VectorSearch"
);
rpc!(
vector_upsert,
vector_upsert,
entity::VectorUpsertRequest,
entity::MutationResponse,
"/udb.services.v1.DataBroker/VectorUpsert"
);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::metadata::headers;
fn client() -> UdbClient {
let channel = Endpoint::from_static("http://127.0.0.1:50051").connect_lazy();
UdbClient::with_channel(
channel,
Metadata::new("tenant-1")
.with_project("proj-9")
.with_bearer_token("tok"),
)
}
#[tokio::test]
async fn request_carries_connection_identity() {
let c = client();
let req = c.request(()).expect("metadata applies");
let md = req.metadata();
assert_eq!(md.get(headers::TENANT_ID).unwrap(), "tenant-1");
assert_eq!(md.get(headers::PROJECT_ID).unwrap(), "proj-9");
assert_eq!(md.get(headers::AUTHORIZATION).unwrap(), "Bearer tok");
}
#[tokio::test]
async fn with_audit_keeps_identity_and_adds_audit() {
let c = client().with_audit("billing", "corr-7");
let req = c.request(()).expect("metadata applies");
let md = req.metadata();
assert_eq!(md.get(headers::TENANT_ID).unwrap(), "tenant-1");
assert_eq!(md.get(headers::PURPOSE).unwrap(), "billing");
assert_eq!(md.get(headers::CORRELATION_ID).unwrap(), "corr-7");
}
#[test]
fn wrapped_rpc_paths_exist_in_the_registry() {
for path in [
"/udb.services.v1.DataBroker/Select",
"/udb.services.v1.DataBroker/Upsert",
"/udb.services.v1.DataBroker/Update",
"/udb.services.v1.DataBroker/Delete",
"/udb.services.v1.DataBroker/BulkCas",
"/udb.services.v1.DataBroker/VectorSearch",
"/udb.services.v1.DataBroker/VectorUpsert",
] {
assert!(
crate::generated_rpcs::spec_for_path(path).is_some(),
"{path} is not in the generated registry"
);
}
}
#[test]
fn contract_drives_retry_and_disagrees_with_naive_naming() {
use crate::generated_rpcs::is_retry_safe;
assert!(is_retry_safe("/udb.services.v1.DataBroker/Select"));
assert!(is_retry_safe("/udb.services.v1.DataBroker/Upsert"));
assert!(is_retry_safe("/udb.services.v1.DataBroker/Delete"));
assert!(!is_retry_safe("/udb.services.v1.DataBroker/BulkCas"));
assert!(!is_retry_safe("/udb.services.v1.DataBroker/VectorUpsert"));
}
#[tokio::test]
async fn with_bearer_token_replaces_only_the_credential() {
let c = client().with_bearer_token("rotated");
assert_eq!(c.metadata().tenant_id, "tenant-1");
let req = c.request(()).expect("metadata applies");
assert_eq!(
req.metadata().get(headers::AUTHORIZATION).unwrap(),
"Bearer rotated"
);
}
}