1use std::fmt::{Display, Formatter};
2use std::num::NonZeroI64;
3use std::str::FromStr;
4use std::sync::Arc;
5use std::time::Duration;
6
7use crate::auth::{CallError, CheckResult, Principal};
8use crate::error::ReadError;
9use chrono::{DateTime, Utc};
10use error::WriteError;
11pub use error::{ConnectError, ParseError};
12use http::Uri;
13use tonic::transport::{Channel, ClientTlsConfig};
14
15pub mod auth;
16#[cfg(feature = "axum")]
17pub mod axum;
18mod error;
19pub mod memo;
20pub mod session;
21
22#[derive(Clone, Debug, PartialEq, Eq, Hash)]
24pub struct Namespace(pub String);
25
26impl Namespace {
29 pub const IAM: &'static str = "iam";
30 pub const SERVICEACCOUNT: &'static str = "serviceaccount";
31
32 pub fn iam() -> Namespace {
33 Namespace(Self::IAM.into())
34 }
35 pub fn serviceaccount() -> Namespace {
36 Namespace(Self::SERVICEACCOUNT.into())
37 }
38}
39
40#[derive(Clone, Debug, PartialEq, Eq, Hash)]
42pub struct Rel(pub String);
43
44impl Rel {
48 pub const IS: &'static str = "is";
49 pub const UNSPECIFIED: &'static str = "...";
50 pub const PARENT: &'static str = "parent";
51
52 pub const ADMIN: &'static str = "admin";
53 pub const EDITOR: &'static str = "editor";
54 pub const VIEWER: &'static str = "viewer";
55
56 pub const IAM_GET: &'static str = "iam.get";
57 pub const IAM_UPDATE: &'static str = "iam.update";
58 pub const IAM_DELETE: &'static str = "iam.delete";
59
60 pub const SERVICEACCOUNT_GET: &'static str = "serviceaccount.get";
61 pub const SERVICEACCOUNT_CREATE: &'static str = "serviceaccount.create";
62 pub const SERVICEACCOUNT_UPDATE: &'static str = "serviceaccount.update";
63 pub const SERVICEACCOUNT_CREATE_TOKEN: &'static str = "serviceaccount.createToken";
64 pub const SERVICEACCOUNT_KEY_CREATE: &'static str = "serviceaccount.key.create";
65 pub const SERVICEACCOUNT_KEY_GET: &'static str = "serviceaccount.key.get";
66
67 pub const USER_CREATE: &'static str = "user.create";
68
69 pub const IMPOSSIBLE: &'static str = "impossible";
72
73 pub fn is() -> Rel {
74 Rel(Self::IS.into())
75 }
76 pub fn unspecified() -> Rel {
77 Rel(Self::UNSPECIFIED.into())
78 }
79 pub fn parent() -> Rel {
80 Rel(Self::PARENT.into())
81 }
82 pub fn admin() -> Rel {
83 Rel(Self::ADMIN.into())
84 }
85 pub fn editor() -> Rel {
86 Rel(Self::EDITOR.into())
87 }
88 pub fn viewer() -> Rel {
89 Rel(Self::VIEWER.into())
90 }
91 pub fn iam_get() -> Rel {
92 Rel(Self::IAM_GET.into())
93 }
94 pub fn iam_update() -> Rel {
95 Rel(Self::IAM_UPDATE.into())
96 }
97 pub fn iam_delete() -> Rel {
98 Rel(Self::IAM_DELETE.into())
99 }
100 pub fn serviceaccount_get() -> Rel {
101 Rel(Self::SERVICEACCOUNT_GET.into())
102 }
103 pub fn serviceaccount_create() -> Rel {
104 Rel(Self::SERVICEACCOUNT_CREATE.into())
105 }
106 pub fn serviceaccount_update() -> Rel {
107 Rel(Self::SERVICEACCOUNT_UPDATE.into())
108 }
109 pub fn serviceaccount_create_token() -> Rel {
110 Rel(Self::SERVICEACCOUNT_CREATE_TOKEN.into())
111 }
112 pub fn serviceaccount_key_create() -> Rel {
113 Rel(Self::SERVICEACCOUNT_KEY_CREATE.into())
114 }
115 pub fn serviceaccount_key_get() -> Rel {
116 Rel(Self::SERVICEACCOUNT_KEY_GET.into())
117 }
118 pub fn user_create() -> Rel {
119 Rel(Self::USER_CREATE.into())
120 }
121 pub fn impossible() -> Rel {
122 Rel(Self::IMPOSSIBLE.into())
123 }
124}
125
126impl From<&str> for Rel {
127 fn from(value: &str) -> Self {
128 Rel(value.to_string())
129 }
130}
131
132#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord)]
134pub struct UserId(NonZeroI64);
135
136impl UserId {
137 pub fn get(self) -> i64 {
138 self.0.get()
139 }
140}
141
142impl TryFrom<i64> for UserId {
143 type Error = ParseError;
144
145 fn try_from(n: i64) -> Result<Self, Self::Error> {
146 match NonZeroI64::new(n) {
147 Some(nz) if n > 0 => Ok(UserId(nz)),
148 _ => Err(ParseError::invalid_syntax(
149 "UserId",
150 n.to_string(),
151 "must be a positive 64-bit integer",
152 )),
153 }
154 }
155}
156
157impl FromStr for UserId {
158 type Err = ParseError;
159
160 fn from_str(s: &str) -> Result<Self, Self::Err> {
161 let decimal = !s.is_empty()
162 && s.len() <= 19
163 && !s.starts_with('0')
164 && s.bytes().all(|b| b.is_ascii_digit());
165 if !decimal {
166 return Err(ParseError::invalid_syntax(
167 "UserId",
168 s,
169 "must be a decimal integer from 1 to 9223372036854775807",
170 ));
171 }
172 let n: i64 = s
173 .parse()
174 .map_err(|_| ParseError::invalid_syntax("UserId", s, "exceeds 9223372036854775807"))?;
175 UserId::try_from(n)
176 }
177}
178
179impl TryFrom<String> for UserId {
180 type Error = ParseError;
181
182 fn try_from(value: String) -> Result<Self, Self::Error> {
183 UserId::from_str(&value)
184 }
185}
186
187impl Display for UserId {
188 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
189 write!(f, "{}", self.0)
190 }
191}
192
193#[derive(Clone, Debug, PartialEq, Eq)]
197pub struct Timestamp(pub String);
198
199impl Timestamp {
200 pub const EMPTY: &'static str = "AQAAAAAAAA==";
203
204 pub fn empty() -> Self {
205 Timestamp(Self::EMPTY.into())
206 }
207}
208
209#[derive(Clone, Debug, PartialEq, Eq, Hash)]
211pub struct Obj(pub String);
212
213impl Obj {
216 pub const ROOT: &'static str = "root";
217 pub const UNSPECIFIED: &'static str = "...";
218
219 pub fn root() -> Obj {
220 Obj(Self::ROOT.into())
221 }
222 pub fn unspecified() -> Obj {
223 Obj(Self::UNSPECIFIED.into())
224 }
225}
226
227impl FromStr for Obj {
228 type Err = ParseError;
229
230 fn from_str(s: &str) -> Result<Self, Self::Err> {
231 Ok(Obj(s.into()))
232 }
233}
234
235impl TryFrom<String> for Obj {
236 type Error = ParseError;
237 fn try_from(value: String) -> Result<Self, Self::Error> {
238 Obj::from_str(&value)
239 }
240}
241
242#[derive(Clone, Debug, PartialEq, Eq)]
244pub struct UserSet {
245 pub ns: Namespace,
246 pub obj: Obj,
247 pub rel: Rel,
248}
249
250#[derive(Clone, Debug, PartialEq, Eq, Hash)]
252pub enum User {
253 UserId(UserId),
254 AllUsers,
255 AuthenticatedUsers,
256 UserSet { ns: Namespace, obj: Obj, rel: Rel },
257}
258
259impl User {
260 pub const ALL_USERS: &'static str = "allUsers";
261 pub const AUTHENTICATED_USERS: &'static str = "authenticatedUsers";
262
263 fn from_id_str(s: &str) -> Result<User, ParseError> {
264 match s {
265 Self::ALL_USERS => Ok(User::AllUsers),
266 Self::AUTHENTICATED_USERS => Ok(User::AuthenticatedUsers),
267 _ => Ok(User::UserId(UserId::from_str(s)?)),
268 }
269 }
270}
271
272impl FromStr for User {
273 type Err = ParseError;
274
275 fn from_str(s: &str) -> Result<Self, Self::Err> {
276 let Some((ns_obj, rel)) = s.split_once('#') else {
277 return User::from_id_str(s);
278 };
279 let Some((ns, obj)) = ns_obj.split_once(':') else {
280 return Err(ParseError::invalid_syntax(
281 "User::UserSet",
282 s,
283 "wrong pattern for userset: missing ':' delimiter",
284 ));
285 };
286 Ok(User::UserSet {
287 ns: Namespace(ns.into()),
288 obj: Obj(obj.into()),
289 rel: Rel(rel.into()),
290 })
291 }
292}
293
294impl Display for User {
295 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
296 match self {
297 User::UserId(id) => write!(f, "{id}"),
298 User::AllUsers => f.write_str(Self::ALL_USERS),
299 User::AuthenticatedUsers => f.write_str(Self::AUTHENTICATED_USERS),
300 User::UserSet { ns, obj, rel } => write!(f, "{}:{}#{}", ns.0, obj.0, rel.0),
301 }
302 }
303}
304
305impl TryFrom<String> for User {
306 type Error = ParseError;
307
308 fn try_from(value: String) -> Result<Self, Self::Error> {
309 User::from_str(&value)
310 }
311}
312
313impl From<UserId> for User {
314 fn from(value: UserId) -> Self {
315 User::UserId(value)
316 }
317}
318
319#[derive(Clone, Debug)]
320pub enum Condition {
321 Expires(DateTime<Utc>),
322}
323
324#[derive(Clone, Debug)]
326pub struct Tuple {
327 pub ns: Namespace,
328 pub obj: Obj,
329 pub rel: Rel,
330 pub sbj: User,
331 pub condition: Option<Condition>,
332}
333
334impl Tuple {
335 pub fn new(ns: Namespace, obj: Obj, rel: Rel, sbj: User) -> Tuple {
336 Tuple {
337 ns,
338 obj,
339 rel,
340 sbj,
341 condition: None,
342 }
343 }
344
345 pub fn with_expires(mut self, expires: DateTime<Utc>) -> Tuple {
347 self.condition = Some(Condition::Expires(expires));
348 self
349 }
350}
351
352impl Display for Tuple {
353 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
354 write!(
355 f,
356 "Tuple({}:{}#{}@{})",
357 self.ns.0, self.obj.0, self.rel.0, self.sbj
358 )
359 }
360}
361
362mod pb {
363 tonic::include_proto!("am");
364}
365
366#[doc(hidden)]
367pub mod wire {
368 pub use crate::pb::*;
371}
372
373#[derive(Clone, Debug)]
377pub struct ListResult {
378 pub ts: Timestamp,
379 pub objs: Vec<String>,
380}
381
382#[derive(Clone, Debug)]
387pub struct ExpandResult {
388 pub ts: Timestamp,
389 pub user_ids: Vec<UserId>,
390 pub all_users: bool,
391 pub authenticated_users: bool,
392 pub usersets: Vec<UserSet>,
393}
394
395#[derive(Clone, Debug)]
399pub struct ReadResult {
400 pub ts: Timestamp,
401 pub tuples: Vec<Tuple>,
402}
403
404#[derive(Clone, Debug)]
408pub struct ContentChangeCheckResult {
409 pub ok: bool,
410 pub ts: Timestamp,
411}
412
413#[derive(Clone, Debug, PartialEq, Eq)]
416pub struct RelationMeta {
417 pub name: String,
418 pub kind: String,
419}
420
421#[derive(Clone, Debug, PartialEq, Eq)]
423pub struct NamespaceMeta {
424 pub name: String,
425 pub relations: Vec<RelationMeta>,
426}
427
428#[derive(Clone, Debug)]
430pub struct WatchUpdate {
431 pub tuple: Tuple,
432 pub deleted: bool,
433}
434
435#[derive(Clone, Debug)]
440pub struct WatchEvent {
441 pub ts: Timestamp,
442 pub updates: Vec<WatchUpdate>,
443}
444
445pub struct WatchStream {
448 inner: tonic::Streaming<pb::WatchResponse>,
449}
450
451impl WatchStream {
452 pub async fn recv(&mut self) -> Result<Option<WatchEvent>, ReadError> {
455 match self.inner.message().await {
456 Ok(None) => Ok(None),
457 Ok(Some(resp)) => watch_event_from_pb(resp).map(Some),
458 Err(status) => Err(status.into()),
459 }
460 }
461}
462
463#[derive(Clone, Debug)]
467pub struct ReadFilter {
468 set: pb::TupleSet,
469}
470
471impl ReadFilter {
472 pub fn by_object(ns: Namespace, obj: Obj, rel: Option<Rel>) -> ReadFilter {
474 ReadFilter {
475 set: pb::TupleSet {
476 ns: ns.0,
477 spec: Some(pb::tuple_set::Spec::ObjectSpec(pb::tuple_set::ObjectSpec {
478 obj: obj.0,
479 rel: rel.map(|r| r.0),
480 })),
481 },
482 }
483 }
484
485 pub fn by_user(ns: Namespace, user: User, rel: Option<Rel>) -> ReadFilter {
489 use pb::tuple_set::user_set_spec::User as Pb;
490 ReadFilter {
491 set: pb::TupleSet {
492 ns: ns.0,
493 spec: Some(pb::tuple_set::Spec::UsersetSpec(
494 pb::tuple_set::UserSetSpec {
495 user: Some(user_to_pb(user, Pb::UserId, Pb::UserSet, Pb::Wildcard)),
496 rel: rel.map(|r| r.0),
497 },
498 )),
499 },
500 }
501 }
502
503 pub fn by_user_set(ns: Namespace, user_set: UserSet, rel: Option<Rel>) -> ReadFilter {
506 let user = User::UserSet {
507 ns: user_set.ns,
508 obj: user_set.obj,
509 rel: user_set.rel,
510 };
511 Self::by_user(ns, user, rel)
512 }
513}
514
515pub type ObserveCheckFn =
516 Arc<dyn Fn(&Namespace, &Obj, &Rel, &UserId, Duration, bool, bool) + Send + Sync>;
517pub type ObserveListFn = Arc<dyn Fn(&Namespace, &Rel, &UserId, Duration, bool) + Send + Sync>;
518
519const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(30);
522const KEEPALIVE_TIMEOUT: Duration = Duration::from_secs(10);
523const KEEPALIVE_WHILE_IDLE: bool = true;
524
525pub async fn connect_channel(
530 uri: Uri,
531 tls_config: Option<ClientTlsConfig>,
532) -> Result<Channel, ConnectError> {
533 let mut builder = Channel::builder(uri)
534 .http2_keep_alive_interval(KEEPALIVE_INTERVAL)
535 .keep_alive_timeout(KEEPALIVE_TIMEOUT)
536 .keep_alive_while_idle(KEEPALIVE_WHILE_IDLE);
537 if let Some(tls) = tls_config {
538 builder = builder.tls_config(tls).map_err(ConnectError)?;
539 }
540 builder.connect().await.map_err(ConnectError)
541}
542
543#[derive(Clone)]
547pub struct CheckClient {
548 check: pb::check_service_client::CheckServiceClient<Channel>,
549 ns: pb::namespace_service_client::NamespaceServiceClient<Channel>,
550 observe_check: Option<ObserveCheckFn>,
551 observe_list: Option<ObserveListFn>,
552}
553
554impl std::fmt::Debug for CheckClient {
555 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
556 f.debug_struct("CheckClient").finish_non_exhaustive()
557 }
558}
559
560impl CheckClient {
561 pub async fn create(uri: Uri) -> Result<Self, ConnectError> {
562 Self::create_with_tls(uri, None).await
563 }
564
565 pub async fn create_with_tls(
566 uri: Uri,
567 tls_config: Option<ClientTlsConfig>,
568 ) -> Result<Self, ConnectError> {
569 let channel = connect_channel(uri, tls_config).await?;
570 Ok(Self::from_channel(channel))
571 }
572
573 pub fn from_channel(channel: Channel) -> Self {
574 CheckClient {
575 check: pb::check_service_client::CheckServiceClient::new(channel.clone()),
576 ns: pb::namespace_service_client::NamespaceServiceClient::new(channel),
577 observe_check: None,
578 observe_list: None,
579 }
580 }
581
582 pub fn with_observe_check(mut self, f: ObserveCheckFn) -> Self {
585 self.observe_check = Some(f);
586 self
587 }
588
589 pub fn with_observe_list(mut self, f: ObserveListFn) -> Self {
592 self.observe_list = Some(f);
593 self
594 }
595
596 pub async fn check(
605 &mut self,
606 ns: Namespace,
607 obj: Obj,
608 rel: Rel,
609 user_id: UserId,
610 timestamp: Option<Timestamp>,
611 ) -> Result<CheckResult, CallError> {
612 if rel.0 == Rel::IMPOSSIBLE {
613 return Ok(CheckResult::Forbidden(user_id.into()));
614 }
615 let r = pb::CheckRequest {
616 ns: ns.0.clone(),
617 obj: obj.0.clone(),
618 rel: rel.0.clone(),
619 user: Some(pb::check_request::User::UserId(user_id.get())),
620 ts: timestamp.map(|t| t.0),
621 };
622 let started = std::time::Instant::now();
623 let result = self.check.check(r).await;
624 if let Some(observe) = &self.observe_check {
625 let ok = result.as_ref().map(|r| r.get_ref().ok).unwrap_or(false);
626 observe(
627 &ns,
628 &obj,
629 &rel,
630 &user_id,
631 started.elapsed(),
632 ok,
633 result.is_err(),
634 );
635 }
636 let response = result?.into_inner();
637 let Some(principal) = response
638 .principal
639 .and_then(|p| UserId::try_from(p.id).ok())
640 .map(Principal::from)
641 else {
642 return Err(CallError::UnexpectedResponseFormat);
643 };
644 Ok(match response.ok {
645 true => CheckResult::Ok(principal),
646 false => CheckResult::Forbidden(principal),
647 })
648 }
649
650 pub async fn list(
656 &mut self,
657 ns: Namespace,
658 rel: Rel,
659 user_id: UserId,
660 timestamp: Option<Timestamp>,
661 ) -> Result<ListResult, CallError> {
662 let r = pb::ListRequest {
663 ns: ns.0.clone(),
664 rel: rel.0.clone(),
665 user: Some(pb::list_request::User::UserId(user_id.get())),
666 ts: timestamp.map(|t| t.0),
667 };
668 let started = std::time::Instant::now();
669 let result = self.check.list(r).await;
670 if let Some(observe) = &self.observe_list {
671 observe(&ns, &rel, &user_id, started.elapsed(), result.is_err());
672 }
673 match result.map(|r| r.into_inner()) {
674 Ok(response) => Ok(ListResult {
675 ts: Timestamp(response.ts),
676 objs: response.objs,
677 }),
678 Err(e) => Err(e.into()),
679 }
680 }
681
682 pub async fn expand(
688 &mut self,
689 ns: Namespace,
690 obj: Obj,
691 rel: Rel,
692 timestamp: Option<Timestamp>,
693 ) -> Result<ExpandResult, ReadError> {
694 let r = pb::ExpandRequest {
695 ns: ns.0,
696 obj: obj.0,
697 rel: rel.0,
698 ts: timestamp.map(|t| t.0),
699 };
700 let response = self.check.expand(r).await?.into_inner();
701 let user_ids = response
702 .user_ids
703 .into_iter()
704 .map(UserId::try_from)
705 .collect::<Result<_, _>>()
706 .map_err(|e| ReadError::invalid_response(e.to_string()))?;
707 let mut all_users = false;
708 let mut authenticated_users = false;
709 for w in response.wildcards {
710 match wildcard_from_pb(w)? {
711 Wildcard::AllUsers => all_users = true,
712 Wildcard::AuthenticatedUsers => authenticated_users = true,
713 }
714 }
715 Ok(ExpandResult {
716 ts: Timestamp(response.ts),
717 user_ids,
718 all_users,
719 authenticated_users,
720 usersets: response
721 .usersets
722 .into_iter()
723 .map(|us| UserSet {
724 ns: Namespace(us.ns),
725 obj: Obj(us.obj),
726 rel: Rel(us.rel),
727 })
728 .collect(),
729 })
730 }
731
732 pub async fn content_change_check(
736 &mut self,
737 ns: Namespace,
738 obj: Obj,
739 rel: Rel,
740 user_id: UserId,
741 ) -> Result<ContentChangeCheckResult, CallError> {
742 let r = pb::ContentChangeCheckRequest {
743 ns: ns.0,
744 obj: obj.0,
745 rel: rel.0,
746 user: Some(pb::content_change_check_request::User::UserId(
747 user_id.get(),
748 )),
749 };
750 match self
751 .check
752 .content_change_check(r)
753 .await
754 .map(|r| r.into_inner())
755 {
756 Ok(response) => Ok(ContentChangeCheckResult {
757 ok: response.ok,
758 ts: Timestamp(response.ts),
759 }),
760 Err(status) => Err(status.into()),
761 }
762 }
763
764 pub async fn watch(
770 &mut self,
771 ns: Namespace,
772 start_ts: Timestamp,
773 ) -> Result<WatchStream, CallError> {
774 let r = pb::WatchRequest {
775 ns: ns.0,
776 start_ts: start_ts.0,
777 };
778 match self.check.watch(r).await {
779 Ok(response) => Ok(WatchStream {
780 inner: response.into_inner(),
781 }),
782 Err(status) => Err(status.into()),
783 }
784 }
785
786 pub async fn list_namespaces(&mut self) -> Result<Vec<NamespaceMeta>, ReadError> {
790 match self.ns.list_namespaces(()).await.map(|r| r.into_inner()) {
791 Ok(resp) => Ok(resp
792 .namespaces
793 .into_iter()
794 .map(|ns| NamespaceMeta {
795 name: ns.name,
796 relations: ns
797 .relations
798 .into_iter()
799 .map(|r| RelationMeta {
800 name: r.name,
801 kind: r.kind,
802 })
803 .collect(),
804 })
805 .collect()),
806 Err(status) => Err(status.into()),
807 }
808 }
809
810 pub async fn get_all(&mut self, ns: &Namespace, obj: &Obj) -> Result<ReadResult, ReadError> {
813 self.read(vec![ReadFilter::by_object(ns.clone(), obj.clone(), None)])
814 .await
815 }
816
817 pub async fn get_all_rel(
819 &mut self,
820 ns: &Namespace,
821 obj: &Obj,
822 rel: &Rel,
823 ) -> Result<ReadResult, ReadError> {
824 self.read(vec![ReadFilter::by_object(
825 ns.clone(),
826 obj.clone(),
827 Some(rel.clone()),
828 )])
829 .await
830 }
831
832 pub async fn read_by_user(
835 &mut self,
836 ns: &Namespace,
837 user: &User,
838 rel: Option<Rel>,
839 ) -> Result<ReadResult, ReadError> {
840 self.read(vec![ReadFilter::by_user(ns.clone(), user.clone(), rel)])
841 .await
842 }
843
844 pub async fn read_by_user_set(
847 &mut self,
848 ns: &Namespace,
849 user_set: &UserSet,
850 rel: Option<Rel>,
851 ) -> Result<ReadResult, ReadError> {
852 self.read(vec![ReadFilter::by_user_set(
853 ns.clone(),
854 user_set.clone(),
855 rel,
856 )])
857 .await
858 }
859
860 pub async fn read(&mut self, filters: Vec<ReadFilter>) -> Result<ReadResult, ReadError> {
862 self.read_with_timestamp(Timestamp::empty(), filters).await
863 }
864
865 pub async fn read_with_timestamp(
868 &mut self,
869 ts: Timestamp,
870 filters: Vec<ReadFilter>,
871 ) -> Result<ReadResult, ReadError> {
872 if filters.is_empty() {
873 return Err(ReadError::invalid_response(
874 "read: at least one filter required",
875 ));
876 }
877 let request = pb::ReadRequest {
878 ts: (ts != Timestamp::empty()).then_some(ts.0),
879 tuple_sets: filters.into_iter().map(|f| f.set).collect(),
880 };
881 let response = self.check.read(request).await?.into_inner();
882 let mut tuples = Vec::with_capacity(response.tuples.len());
883 for tup in response.tuples {
884 tuples.push(tuple_from_pb(tup)?);
885 }
886 Ok(ReadResult {
887 ts: Timestamp(response.ts),
888 tuples,
889 })
890 }
891
892 pub async fn write(
896 &mut self,
897 add: Vec<Tuple>,
898 del: Vec<Tuple>,
899 precondition: Option<Timestamp>,
900 ) -> Result<Timestamp, WriteError> {
901 let request = pb::WriteRequest {
902 ts: precondition.map(|t| t.0),
903 add_tuples: add.into_iter().map(tuple_to_pb).collect(),
904 del_tuples: del.into_iter().map(tuple_to_pb).collect(),
905 };
906 self.check
907 .write(request)
908 .await
909 .map(|r| Timestamp(r.into_inner().ts))
910 .map_err(Into::into)
911 }
912
913 pub async fn add_one(&mut self, tuple: Tuple) -> Result<Timestamp, WriteError> {
915 self.write(vec![tuple], vec![], None).await
916 }
917
918 pub async fn add_many(&mut self, tuples: Vec<Tuple>) -> Result<Timestamp, WriteError> {
920 self.write(tuples, vec![], None).await
921 }
922
923 pub async fn add_parent(
927 &mut self,
928 ns: Namespace,
929 obj: Obj,
930 parent_ns: Namespace,
931 parent_obj: Obj,
932 ) -> Result<Timestamp, WriteError> {
933 self.add_one(Tuple::new(
934 ns,
935 obj,
936 Rel::parent(),
937 User::UserSet {
938 ns: parent_ns,
939 obj: parent_obj,
940 rel: Rel::unspecified(),
941 },
942 ))
943 .await
944 }
945
946 pub async fn delete_one(&mut self, tuple: Tuple) -> Result<Timestamp, WriteError> {
948 self.write(vec![], vec![tuple], None).await
949 }
950}
951
952fn user_to_pb<T>(
955 user: User,
956 user_id: fn(i64) -> T,
957 user_set: fn(pb::UserSet) -> T,
958 wildcard: fn(i32) -> T,
959) -> T {
960 match user {
961 User::UserId(id) => user_id(id.get()),
962 User::AllUsers => wildcard(pb::Wildcard::AllUsers.into()),
963 User::AuthenticatedUsers => wildcard(pb::Wildcard::AuthenticatedUsers.into()),
964 User::UserSet { ns, obj, rel } => user_set(pb::UserSet {
965 ns: ns.0,
966 obj: obj.0,
967 rel: rel.0,
968 }),
969 }
970}
971
972enum Wildcard {
973 AllUsers,
974 AuthenticatedUsers,
975}
976
977#[allow(clippy::result_large_err)]
978fn wildcard_from_pb(w: i32) -> Result<Wildcard, ReadError> {
979 match pb::Wildcard::try_from(w) {
980 Ok(pb::Wildcard::AllUsers) => Ok(Wildcard::AllUsers),
981 Ok(pb::Wildcard::AuthenticatedUsers) => Ok(Wildcard::AuthenticatedUsers),
982 Ok(pb::Wildcard::Unspecified) | Err(_) => {
983 Err(ReadError::invalid_response(format!("unknown wildcard {w}")))
984 }
985 }
986}
987
988#[allow(clippy::result_large_err)]
989fn user_from_pb(user: pb::tuple::User) -> Result<User, ReadError> {
990 match user {
991 pb::tuple::User::UserId(id) => UserId::try_from(id)
992 .map(User::UserId)
993 .map_err(|e| ReadError::invalid_response(e.to_string())),
994 pb::tuple::User::Wildcard(w) => Ok(match wildcard_from_pb(w)? {
995 Wildcard::AllUsers => User::AllUsers,
996 Wildcard::AuthenticatedUsers => User::AuthenticatedUsers,
997 }),
998 pb::tuple::User::UserSet(pb::UserSet { ns, obj, rel }) => Ok(User::UserSet {
999 ns: Namespace(ns),
1000 obj: Obj(obj),
1001 rel: Rel(rel),
1002 }),
1003 }
1004}
1005
1006fn tuple_to_pb(t: Tuple) -> pb::Tuple {
1007 use pb::tuple::User as Pb;
1008 pb::Tuple {
1009 ns: t.ns.0,
1010 obj: t.obj.0,
1011 rel: t.rel.0,
1012 user: Some(user_to_pb(t.sbj, Pb::UserId, Pb::UserSet, Pb::Wildcard)),
1013 condition: t.condition.map(|c| match c {
1014 Condition::Expires(exp) => pb::tuple::Condition::Expires(exp.timestamp()),
1015 }),
1016 }
1017}
1018
1019#[allow(clippy::result_large_err)] fn tuple_from_pb(tup: pb::Tuple) -> Result<Tuple, ReadError> {
1023 let Some(user) = tup.user else {
1024 return Err(ReadError::invalid_response(format!(
1025 "tuple {}:{}#{} missing user field",
1026 tup.ns, tup.obj, tup.rel
1027 )));
1028 };
1029 let sbj = user_from_pb(user)?;
1030 let condition = match tup.condition {
1031 None => None,
1032 Some(pb::tuple::Condition::Expires(secs)) => match DateTime::from_timestamp(secs, 0) {
1033 Some(dt) => Some(Condition::Expires(dt)),
1034 None => {
1035 return Err(ReadError::invalid_response(format!(
1036 "tuple {}:{}#{} expires out of range: {}",
1037 tup.ns, tup.obj, tup.rel, secs
1038 )));
1039 }
1040 },
1041 };
1042 Ok(Tuple {
1043 ns: Namespace(tup.ns),
1044 obj: Obj(tup.obj),
1045 rel: Rel(tup.rel),
1046 sbj,
1047 condition,
1048 })
1049}
1050
1051#[allow(clippy::result_large_err)] fn watch_event_from_pb(resp: pb::WatchResponse) -> Result<WatchEvent, ReadError> {
1053 let mut updates = Vec::with_capacity(resp.updates.len());
1054 for (i, u) in resp.updates.into_iter().enumerate() {
1055 let tuple = match u.tuple {
1056 None => {
1057 return Err(ReadError::invalid_response(format!(
1058 "watch update[{i}]: missing tuple"
1059 )));
1060 }
1061 Some(t) => tuple_from_pb(t)?,
1062 };
1063 updates.push(WatchUpdate {
1064 tuple,
1065 deleted: u.deleted,
1066 });
1067 }
1068 Ok(WatchEvent {
1069 ts: Timestamp(resp.ts),
1070 updates,
1071 })
1072}
1073
1074#[cfg(test)]
1075mod tests {
1076 use super::*;
1077
1078 fn uid(n: i64) -> UserId {
1079 UserId::try_from(n).unwrap()
1080 }
1081
1082 fn tuple_with(user: pb::tuple::User) -> pb::Tuple {
1083 pb::Tuple {
1084 ns: "doc".into(),
1085 obj: "1".into(),
1086 rel: "viewer".into(),
1087 user: Some(user),
1088 condition: None,
1089 }
1090 }
1091
1092 #[test]
1093 fn timestamp_empty_is_packed_empty_zookie() {
1094 assert_eq!(Timestamp::empty().0, "AQAAAAAAAA==");
1095 assert_eq!(Timestamp::EMPTY, "AQAAAAAAAA==");
1096 }
1097
1098 #[test]
1101 #[allow(clippy::assertions_on_constants)]
1102 fn keepalive_matches_nio_check_client() {
1103 assert_eq!(KEEPALIVE_INTERVAL, Duration::from_secs(30));
1104 assert_eq!(KEEPALIVE_TIMEOUT, Duration::from_secs(10));
1105 assert!(KEEPALIVE_WHILE_IDLE);
1106 }
1107
1108 #[test]
1111 fn domain_constants_match_nio() {
1112 assert_eq!(Namespace::iam().0, "iam");
1113 assert_eq!(Namespace::serviceaccount().0, "serviceaccount");
1114 assert_eq!(Obj::root().0, "root");
1115 assert_eq!(Obj::unspecified().0, "...");
1116 assert_eq!(Rel::is().0, "is");
1117 assert_eq!(Rel::unspecified().0, "...");
1118 assert_eq!(Rel::parent().0, "parent");
1119 assert_eq!(Rel::admin().0, "admin");
1120 assert_eq!(Rel::editor().0, "editor");
1121 assert_eq!(Rel::viewer().0, "viewer");
1122 assert_eq!(Rel::iam_get().0, "iam.get");
1123 assert_eq!(Rel::iam_update().0, "iam.update");
1124 assert_eq!(Rel::iam_delete().0, "iam.delete");
1125 assert_eq!(Rel::serviceaccount_get().0, "serviceaccount.get");
1126 assert_eq!(Rel::serviceaccount_create().0, "serviceaccount.create");
1127 assert_eq!(Rel::serviceaccount_update().0, "serviceaccount.update");
1128 assert_eq!(
1129 Rel::serviceaccount_create_token().0,
1130 "serviceaccount.createToken"
1131 );
1132 assert_eq!(
1133 Rel::serviceaccount_key_create().0,
1134 "serviceaccount.key.create"
1135 );
1136 assert_eq!(Rel::serviceaccount_key_get().0, "serviceaccount.key.get");
1137 assert_eq!(Rel::user_create().0, "user.create");
1138 assert_eq!(User::AllUsers.to_string(), "allUsers");
1139 assert_eq!(User::AuthenticatedUsers.to_string(), "authenticatedUsers");
1140 }
1141
1142 #[test]
1143 fn user_id_accepts_max_i64() {
1144 let id = UserId::from_str("9223372036854775807").unwrap();
1145 assert_eq!(id.get(), 9223372036854775807);
1146 assert_eq!(id.to_string(), "9223372036854775807");
1147 assert_eq!(UserId::try_from("42".to_string()).unwrap().get(), 42);
1148 }
1149
1150 #[test]
1151 fn user_id_rejects_non_canonical_decimals() {
1152 let not_decimal =
1153 "invalid syntax: 'must be a decimal integer from 1 to 9223372036854775807'";
1154 let error_of = |s: &str| UserId::from_str(s).unwrap_err().to_string();
1155 assert_eq!(
1156 error_of("0"),
1157 format!("'0' has invalid syntax for UserId {not_decimal}")
1158 );
1159 assert_eq!(
1160 error_of("-5"),
1161 format!("'-5' has invalid syntax for UserId {not_decimal}")
1162 );
1163 assert_eq!(
1164 error_of("01"),
1165 format!("'01' has invalid syntax for UserId {not_decimal}")
1166 );
1167 assert_eq!(
1168 error_of("9223372036854775808"),
1169 "'9223372036854775808' has invalid syntax for UserId invalid syntax: 'exceeds 9223372036854775807'"
1170 );
1171 let uuid = ["812eebc6", "480b", "4527", "bed4", "057e4d2fd1e3"].join("-");
1172 assert_eq!(
1173 error_of(&uuid),
1174 format!("'{uuid}' has invalid syntax for UserId {not_decimal}")
1175 );
1176 }
1177
1178 #[test]
1179 fn user_id_try_from_i64_rejects_non_positive() {
1180 let positive = "invalid syntax: 'must be a positive 64-bit integer'";
1181 assert_eq!(
1182 UserId::try_from(0).unwrap_err().to_string(),
1183 format!("'0' has invalid syntax for UserId {positive}")
1184 );
1185 assert_eq!(
1186 UserId::try_from(-5).unwrap_err().to_string(),
1187 format!("'-5' has invalid syntax for UserId {positive}")
1188 );
1189 }
1190
1191 #[test]
1192 fn user_parses_and_displays_every_subject_kind() {
1193 let cases = [
1194 ("42", User::UserId(uid(42))),
1195 ("allUsers", User::AllUsers),
1196 ("authenticatedUsers", User::AuthenticatedUsers),
1197 (
1198 "group:eng#member",
1199 User::UserSet {
1200 ns: Namespace("group".into()),
1201 obj: Obj("eng".into()),
1202 rel: Rel("member".into()),
1203 },
1204 ),
1205 ];
1206 for (text, user) in cases {
1207 assert_eq!(User::from_str(text).unwrap(), user);
1208 assert_eq!(user.to_string(), text);
1209 }
1210 }
1211
1212 #[test]
1213 fn user_rejects_bad_subjects() {
1214 assert_eq!(
1215 User::from_str("group#member").unwrap_err().to_string(),
1216 "'group#member' has invalid syntax for User::UserSet invalid syntax: 'wrong pattern for userset: missing ':' delimiter'"
1217 );
1218 assert!(User::from_str("alice").is_err());
1219 }
1220
1221 #[test]
1222 fn tuple_display_uses_subject_text() {
1223 let t = Tuple::new(
1224 Namespace("doc".into()),
1225 Obj("1".into()),
1226 Rel("viewer".into()),
1227 User::AllUsers,
1228 );
1229 assert_eq!(t.to_string(), "Tuple(doc:1#viewer@allUsers)");
1230 }
1231
1232 #[test]
1233 fn tuple_to_pb_encodes_each_subject_kind() {
1234 let encode = |sbj: User| {
1235 tuple_to_pb(Tuple::new(
1236 Namespace("doc".into()),
1237 Obj("1".into()),
1238 Rel("viewer".into()),
1239 sbj,
1240 ))
1241 };
1242 let pt = encode(User::UserId(uid(42)));
1243 assert_eq!(pt.ns, "doc");
1244 assert_eq!(pt.obj, "1");
1245 assert_eq!(pt.rel, "viewer");
1246 assert_eq!(pt.user, Some(pb::tuple::User::UserId(42)));
1247 assert!(pt.condition.is_none());
1248 assert_eq!(
1249 encode(User::AllUsers).user,
1250 Some(pb::tuple::User::Wildcard(1))
1251 );
1252 assert_eq!(
1253 encode(User::AuthenticatedUsers).user,
1254 Some(pb::tuple::User::Wildcard(2))
1255 );
1256 assert_eq!(
1257 encode(User::UserSet {
1258 ns: Namespace("group".into()),
1259 obj: Obj("eng".into()),
1260 rel: Rel("member".into()),
1261 })
1262 .user,
1263 Some(pb::tuple::User::UserSet(pb::UserSet {
1264 ns: "group".into(),
1265 obj: "eng".into(),
1266 rel: "member".into(),
1267 }))
1268 );
1269 }
1270
1271 #[test]
1272 fn tuple_to_pb_expires() {
1273 let exp = DateTime::from_timestamp(1894785600, 0).unwrap();
1274 let pt = tuple_to_pb(
1275 Tuple::new(
1276 Namespace("doc".into()),
1277 Obj("1".into()),
1278 Rel("viewer".into()),
1279 User::UserId(uid(1)),
1280 )
1281 .with_expires(exp),
1282 );
1283 assert_eq!(
1284 pt.condition,
1285 Some(pb::tuple::Condition::Expires(1894785600))
1286 );
1287 }
1288
1289 #[test]
1290 fn tuple_from_pb_maps_each_subject_kind() {
1291 let t = tuple_from_pb(tuple_with(pb::tuple::User::UserId(42))).unwrap();
1292 assert_eq!(t.ns.0, "doc");
1293 assert_eq!(t.obj.0, "1");
1294 assert_eq!(t.rel.0, "viewer");
1295 assert_eq!(t.sbj, User::UserId(uid(42)));
1296
1297 let t = tuple_from_pb(tuple_with(pb::tuple::User::Wildcard(
1298 pb::Wildcard::AllUsers.into(),
1299 )))
1300 .unwrap();
1301 assert_eq!(t.sbj, User::AllUsers);
1302
1303 let t = tuple_from_pb(tuple_with(pb::tuple::User::Wildcard(
1304 pb::Wildcard::AuthenticatedUsers.into(),
1305 )))
1306 .unwrap();
1307 assert_eq!(t.sbj, User::AuthenticatedUsers);
1308
1309 let t = tuple_from_pb(tuple_with(pb::tuple::User::UserSet(pb::UserSet {
1310 ns: "grp".into(),
1311 obj: "eng".into(),
1312 rel: "member".into(),
1313 })))
1314 .unwrap();
1315 assert_eq!(
1316 t.sbj,
1317 User::UserSet {
1318 ns: Namespace("grp".into()),
1319 obj: Obj("eng".into()),
1320 rel: Rel("member".into()),
1321 }
1322 );
1323 }
1324
1325 #[test]
1326 fn tuple_from_pb_rejects_unspecified_and_unknown_wildcards() {
1327 let message = |user| match tuple_from_pb(tuple_with(user)) {
1328 Err(ReadError::InvalidResponse(msg)) => msg,
1329 other => panic!("expected InvalidResponse, got {other:?}"),
1330 };
1331 assert_eq!(
1332 message(pb::tuple::User::Wildcard(pb::Wildcard::Unspecified.into())),
1333 "unknown wildcard 0"
1334 );
1335 assert_eq!(
1336 message(pb::tuple::User::Wildcard(99)),
1337 "unknown wildcard 99"
1338 );
1339 assert_eq!(
1340 message(pb::tuple::User::UserId(0)),
1341 "'0' has invalid syntax for UserId invalid syntax: 'must be a positive 64-bit integer'"
1342 );
1343 }
1344
1345 #[test]
1346 fn tuple_from_pb_missing_user_is_invalid_response_not_panic() {
1347 let bare = pb::Tuple {
1348 ns: "coll".into(),
1349 obj: "uk".into(),
1350 rel: "owner".into(),
1351 user: None,
1352 condition: None,
1353 };
1354 match tuple_from_pb(bare) {
1355 Err(ReadError::InvalidResponse(msg)) => {
1356 assert_eq!(msg, "tuple coll:uk#owner missing user field")
1357 }
1358 other => panic!("expected InvalidResponse, got {other:?}"),
1359 }
1360 }
1361
1362 #[test]
1363 fn tuple_round_trip_expires() {
1364 let exp = DateTime::from_timestamp(1894785600, 0).unwrap();
1365 let t = Tuple::new(
1366 Namespace("doc".into()),
1367 Obj("1".into()),
1368 Rel("viewer".into()),
1369 User::UserId(uid(1)),
1370 )
1371 .with_expires(exp);
1372 let back = tuple_from_pb(tuple_to_pb(t)).expect("round trip");
1373 match back.condition {
1374 Some(Condition::Expires(dt)) => assert_eq!(dt, exp),
1375 None => panic!("expected expires condition"),
1376 }
1377 }
1378
1379 #[test]
1380 fn filter_by_object() {
1381 let f = ReadFilter::by_object(Namespace("doc".into()), Obj("1".into()), None);
1382 assert_eq!(f.set.ns, "doc");
1383 assert_eq!(
1384 f.set.spec,
1385 Some(pb::tuple_set::Spec::ObjectSpec(pb::tuple_set::ObjectSpec {
1386 obj: "1".into(),
1387 rel: None,
1388 }))
1389 );
1390
1391 let f = ReadFilter::by_object(
1392 Namespace("doc".into()),
1393 Obj("1".into()),
1394 Some(Rel::viewer()),
1395 );
1396 assert_eq!(
1397 f.set.spec,
1398 Some(pb::tuple_set::Spec::ObjectSpec(pb::tuple_set::ObjectSpec {
1399 obj: "1".into(),
1400 rel: Some("viewer".into()),
1401 }))
1402 );
1403 }
1404
1405 #[test]
1406 fn filter_by_user() {
1407 use pb::tuple_set::user_set_spec::User as Pb;
1408 let spec = |user, rel| {
1409 ReadFilter::by_user(Namespace("doc".into()), user, rel)
1410 .set
1411 .spec
1412 };
1413 let expected = |user, rel: Option<&str>| {
1414 Some(pb::tuple_set::Spec::UsersetSpec(
1415 pb::tuple_set::UserSetSpec {
1416 user: Some(user),
1417 rel: rel.map(String::from),
1418 },
1419 ))
1420 };
1421 assert_eq!(
1422 spec(User::UserId(uid(42)), None),
1423 expected(Pb::UserId(42), None)
1424 );
1425 assert_eq!(
1426 spec(User::AuthenticatedUsers, Some(Rel::editor())),
1427 expected(Pb::Wildcard(2), Some("editor"))
1428 );
1429 }
1430
1431 #[test]
1432 fn filter_by_user_set() {
1433 let f = ReadFilter::by_user_set(
1434 Namespace("doc".into()),
1435 UserSet {
1436 ns: Namespace("grp".into()),
1437 obj: Obj("eng".into()),
1438 rel: Rel("member".into()),
1439 },
1440 None,
1441 );
1442 assert_eq!(
1443 f.set.spec,
1444 Some(pb::tuple_set::Spec::UsersetSpec(
1445 pb::tuple_set::UserSetSpec {
1446 user: Some(pb::tuple_set::user_set_spec::User::UserSet(pb::UserSet {
1447 ns: "grp".into(),
1448 obj: "eng".into(),
1449 rel: "member".into(),
1450 })),
1451 rel: None,
1452 },
1453 ))
1454 );
1455 }
1456
1457 #[test]
1458 fn watch_event_from_pb_heartbeat() {
1459 let ev = watch_event_from_pb(pb::WatchResponse {
1460 ts: "AQAAAAAAAA==".into(),
1461 updates: vec![],
1462 })
1463 .expect("heartbeat");
1464 assert_eq!(ev.ts, Timestamp::empty());
1465 assert!(ev.updates.is_empty());
1466 }
1467
1468 #[test]
1469 fn watch_event_from_pb_atomic_write() {
1470 let mut editor = tuple_with(pb::tuple::User::UserId(7));
1471 editor.rel = "editor".into();
1472 let ev = watch_event_from_pb(pb::WatchResponse {
1473 ts: "commit-ts".into(),
1474 updates: vec![
1475 pb::Update {
1476 tuple: Some(tuple_with(pb::tuple::User::UserId(7))),
1477 deleted: false,
1478 },
1479 pb::Update {
1480 tuple: Some(editor),
1481 deleted: true,
1482 },
1483 ],
1484 })
1485 .expect("atomic write");
1486 assert_eq!(ev.ts.0, "commit-ts");
1487 assert_eq!(ev.updates.len(), 2);
1488 assert!(!ev.updates[0].deleted);
1489 assert_eq!(ev.updates[0].tuple.sbj, User::UserId(uid(7)));
1490 assert!(ev.updates[1].deleted);
1491 assert_eq!(ev.updates[1].tuple.rel.0, "editor");
1492 }
1493
1494 #[test]
1495 fn watch_event_from_pb_missing_tuple_user() {
1496 let mut bare = tuple_with(pb::tuple::User::UserId(7));
1497 bare.user = None;
1498 let err = watch_event_from_pb(pb::WatchResponse {
1499 ts: "t".into(),
1500 updates: vec![pb::Update {
1501 tuple: Some(bare),
1502 deleted: false,
1503 }],
1504 })
1505 .expect_err("missing user must fail");
1506 assert!(matches!(err, ReadError::InvalidResponse(_)));
1507 }
1508
1509 #[test]
1510 fn watch_event_from_pb_missing_tuple() {
1511 let err = watch_event_from_pb(pb::WatchResponse {
1512 ts: "t".into(),
1513 updates: vec![pb::Update {
1514 tuple: None,
1515 deleted: false,
1516 }],
1517 })
1518 .expect_err("missing tuple must fail");
1519 assert!(matches!(err, ReadError::InvalidResponse(_)));
1520 }
1521}