1use 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
28pub struct RemotePeer {
30 backend: Arc<RemoteDbBackend>,
31}
32
33impl RemotePeer {
34 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 db_after: self.db_at(View::AsOf(response.basis_t)),
96 })
97 }
98
99 async fn sync(&self) -> Result<Db, ClientError> {
100 Ok(self.db_at(View::Current))
103 }
104}
105
106struct 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
235fn 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
260fn 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}