Skip to main content

corium_client/
remote.rs

1//! The remote peer: the fluent API over a peer server, reached over gRPC.
2//!
3//! Presents the same surface as [`crate::LocalPeer`], but every read and
4//! write is an RPC to a hosted peer. Views are named in the request and
5//! resolved server-side; results stream back in chunks and are reassembled
6//! here.
7
8use std::collections::BTreeMap;
9use std::sync::Arc;
10use std::time::Duration;
11
12use async_trait::async_trait;
13use corium_core::{EntityId, KeywordInterner, TotalF64, Value};
14use corium_peer::server::assemble_query_result;
15use corium_protocol::auth::TokenInterceptor;
16use corium_protocol::codec;
17use corium_protocol::pb;
18use corium_protocol::pb::peer_server_client::PeerServerClient;
19use corium_query::edn::Edn;
20use tonic::service::interceptor::InterceptedService;
21use tonic::transport::{Channel, ClientTlsConfig, Endpoint};
22
23use crate::result::{QueryResult, ResultShape};
24use crate::{ClientError, DatomRow, Db, DbBackend, DbStats, Index, Peer, TxData, TxReport, View};
25
26type ServerClient = PeerServerClient<InterceptedService<Channel, TokenInterceptor>>;
27
28/// A fluent client backed by a remote peer server over gRPC.
29pub struct RemotePeer {
30    backend: Arc<RemoteDbBackend>,
31}
32
33impl RemotePeer {
34    /// Connects to a peer server hosting `db`.
35    ///
36    /// # Errors
37    /// Returns [`ClientError`] when the endpoint is unreachable.
38    pub async fn connect(
39        endpoint: impl Into<String>,
40        db: impl Into<String>,
41        token: Option<String>,
42        tls: Option<ClientTlsConfig>,
43    ) -> Result<Self, ClientError> {
44        let endpoint = endpoint.into();
45        let mut builder = Endpoint::from_shared(endpoint)
46            .map_err(|error| ClientError::Protocol(format!("bad endpoint: {error}")))?
47            .connect_timeout(Duration::from_secs(10));
48        if let Some(tls) = tls {
49            builder = builder.tls_config(tls)?;
50        }
51        let channel = builder.connect().await?;
52        let client = PeerServerClient::with_interceptor(channel, TokenInterceptor::new(token));
53        Ok(Self {
54            backend: Arc::new(RemoteDbBackend {
55                client,
56                db_name: db.into(),
57            }),
58        })
59    }
60
61    fn db_at(&self, view: View) -> Db {
62        Db::new(self.backend.clone(), view)
63    }
64}
65
66#[async_trait]
67impl Peer for RemotePeer {
68    fn db_name(&self) -> &str {
69        &self.backend.db_name
70    }
71
72    async fn db(&self) -> Result<Db, ClientError> {
73        Ok(self.db_at(View::Current))
74    }
75
76    async fn transact(&self, tx: TxData) -> Result<TxReport, ClientError> {
77        let tx_data = codec::encode_edn(&Edn::Vector(tx.into_forms()));
78        let mut client = self.backend.client.clone();
79        let response = client
80            .transact(pb::TransactRequest {
81                db: self.backend.db_name.clone(),
82                protocol_version: corium_protocol::PROTOCOL_VERSION,
83                tx_data,
84                expected_basis_t: None,
85            })
86            .await?
87            .into_inner();
88        Ok(TxReport {
89            basis_before: response.basis_before,
90            basis_t: response.basis_t,
91            tx_instant: response.tx_instant,
92            tempids: decode_tempids(&response.tempids)?,
93            // The server syncs its hosted peer to this basis before replying,
94            // so an as-of view at `basis_t` is a stable post-commit snapshot.
95            db_after: self.db_at(View::AsOf(response.basis_t)),
96        })
97    }
98
99    async fn sync(&self) -> Result<Db, ClientError> {
100        // A peer server keeps its hosted peer synced; the current view already
101        // reflects the latest applied basis.
102        Ok(self.db_at(View::Current))
103    }
104}
105
106/// A database backend that issues peer-server RPCs.
107struct RemoteDbBackend {
108    client: ServerClient,
109    db_name: String,
110}
111
112impl RemoteDbBackend {
113    fn view_spec(&self, view: View) -> pb::DbViewSpec {
114        let view = match view {
115            View::Current => None,
116            View::AsOf(t) => Some(pb::db_view_spec::View::AsOf(t)),
117            View::Since(t) => Some(pb::db_view_spec::View::Since(t)),
118            View::History => Some(pb::db_view_spec::View::History(true)),
119            View::AsOfInstant(instant) => Some(pb::db_view_spec::View::AsOfInstant(instant)),
120            View::SinceInstant(instant) => Some(pb::db_view_spec::View::SinceInstant(instant)),
121        };
122        pb::DbViewSpec {
123            db: self.db_name.clone(),
124            view,
125        }
126    }
127}
128
129#[async_trait]
130impl DbBackend for RemoteDbBackend {
131    fn db_name(&self) -> &str {
132        &self.db_name
133    }
134
135    async fn query(
136        &self,
137        view: View,
138        query: Edn,
139        args: Vec<Edn>,
140        fuel: Option<u64>,
141    ) -> Result<QueryResult, ClientError> {
142        let mut client = self.client.clone();
143        let mut stream = client
144            .query(pb::QueryRequest {
145                dbs: vec![self.view_spec(view)],
146                query: codec::encode_edn(&query),
147                args: codec::encode_edn(&Edn::Vector(args)),
148                fuel: fuel.unwrap_or(0),
149            })
150            .await?
151            .into_inner();
152        let mut chunks = Vec::new();
153        while let Some(chunk) = stream.message().await? {
154            chunks.push(chunk);
155        }
156        let shape = chunks
157            .first()
158            .map_or(ResultShape::Relation, |chunk| shape_of(chunk.shape()));
159        let value = assemble_query_result(&chunks).map_err(ClientError::Decode)?;
160        Ok(QueryResult::new(shape, value))
161    }
162
163    async fn pull(&self, view: View, pattern: Edn, eid: Edn) -> Result<Edn, ClientError> {
164        let mut client = self.client.clone();
165        let response = client
166            .pull(pb::PullRequest {
167                db: Some(self.view_spec(view)),
168                pattern: codec::encode_edn(&pattern),
169                eid: codec::encode_edn(&eid),
170            })
171            .await?
172            .into_inner();
173        Ok(codec::decode_edn(&response.result)?)
174    }
175
176    async fn datoms(
177        &self,
178        view: View,
179        index: Index,
180        components: Vec<Edn>,
181        limit: usize,
182    ) -> Result<Vec<DatomRow>, ClientError> {
183        let mut client = self.client.clone();
184        let mut stream = client
185            .datoms(pb::DatomsRequest {
186                db: Some(self.view_spec(view)),
187                index: index.as_str().to_owned(),
188                components: codec::encode_edn(&Edn::Vector(components)),
189                limit: u64::try_from(limit).unwrap_or(0),
190            })
191            .await?
192            .into_inner();
193        let mut interner = KeywordInterner::default();
194        let mut rows = Vec::new();
195        while let Some(chunk) = stream.message().await? {
196            for datom in codec::decode_datoms(&chunk.datoms, &mut interner)? {
197                rows.push(DatomRow {
198                    e: datom.e.raw(),
199                    a: datom.a.raw(),
200                    v: value_to_edn(&interner, &datom.v),
201                    tx: datom.tx.raw(),
202                    added: datom.added,
203                });
204            }
205        }
206        Ok(rows)
207    }
208
209    async fn stats(&self, view: View) -> Result<DbStats, ClientError> {
210        let mut client = self.client.clone();
211        let response = client
212            .db_stats(pb::DbStatsRequest {
213                db: Some(self.view_spec(view)),
214            })
215            .await?
216            .into_inner();
217        Ok(DbStats {
218            basis_t: response.basis_t,
219            datoms: response.datom_count,
220            entities: response.entity_count,
221            attributes: response.attribute_count,
222        })
223    }
224}
225
226fn shape_of(shape: pb::ResultShape) -> ResultShape {
227    match shape {
228        pb::ResultShape::Collection => ResultShape::Collection,
229        pb::ResultShape::Tuple => ResultShape::Tuple,
230        pb::ResultShape::Scalar => ResultShape::Scalar,
231        pb::ResultShape::Relation | pb::ResultShape::Unspecified => ResultShape::Relation,
232    }
233}
234
235/// Renders a decoded value to boundary EDN using a client-side interner for
236/// keyword names. Refs surface as longs, matching the query boundary.
237fn value_to_edn(interner: &KeywordInterner, value: &Value) -> Edn {
238    match value {
239        Value::Bool(v) => Edn::Bool(*v),
240        Value::Long(v) => Edn::Long(*v),
241        Value::Double(TotalF64(v)) => Edn::Double(TotalF64(*v)),
242        Value::Str(v) => Edn::Str(v.to_string()),
243        Value::Instant(ms) => Edn::Tagged("inst".into(), Box::new(Edn::Long(*ms))),
244        Value::Uuid(v) => Edn::Tagged("uuid".into(), Box::new(Edn::Str(format!("{v:032x}")))),
245        Value::Bytes(bytes) => Edn::Tagged(
246            "bytes".into(),
247            Box::new(Edn::Str(bytes.iter().fold(String::new(), |mut acc, b| {
248                use std::fmt::Write as _;
249                let _ = write!(acc, "{b:02x}");
250                acc
251            }))),
252        ),
253        Value::Keyword(id) => interner
254            .resolve(*id)
255            .map_or(Edn::Nil, |keyword| Edn::Keyword(keyword.clone())),
256        Value::Ref(e) => Edn::Long(i64::try_from(e.raw()).unwrap_or(i64::MAX)),
257    }
258}
259
260/// Decodes a tempid map (string name -> allocated entity id long).
261fn decode_tempids(bytes: &[u8]) -> Result<BTreeMap<String, EntityId>, ClientError> {
262    let Edn::Map(pairs) = codec::decode_edn(bytes)? else {
263        return Err(ClientError::Protocol("tempids must be a map".into()));
264    };
265    let mut tempids = BTreeMap::new();
266    for (key, value) in pairs {
267        let (Edn::Str(name), Edn::Long(raw)) = (key, value) else {
268            return Err(ClientError::Protocol("bad tempid entry".into()));
269        };
270        let raw = u64::try_from(raw).map_err(|_| ClientError::Protocol("bad entity id".into()))?;
271        tempids.insert(name, EntityId::from_raw(raw));
272    }
273    Ok(tempids)
274}