1#![deny(missing_docs)]
8
9#[cfg(not(any(feature = "band8", feature = "band9")))]
10compile_error!(
11 "type-bridge-typedb-runtime requires at least one band feature; enable `band8` and/or `band9` (both are default)"
12);
13
14use std::future::Future;
15use std::path::PathBuf;
16use std::pin::Pin;
17use std::sync::atomic::{AtomicU8, AtomicUsize, Ordering as AtomicOrdering};
18use std::sync::{Arc, Mutex};
19use std::time::{Duration, Instant};
20
21use futures::TryStreamExt;
22use serde::{Deserialize, Serialize};
23use tokio::sync::watch;
24use type_bridge_core_lib::version as core_version;
25use type_bridge_core_lib::version::DEFAULT_HTTP_PORT;
26
27#[cfg(feature = "band8")]
28use type_bridge_typedb_driver_b8 as driver_b8;
29#[cfg(feature = "band8")]
30use type_bridge_typedb_driver_b8::answer::QueryAnswer as B8QueryAnswer;
31#[cfg(feature = "band8")]
32use type_bridge_typedb_driver_b8::{
33 Addresses, Credentials as B8Credentials, DriverOptions, DriverTlsConfig,
34 TransactionOptions as B8TransactionOptions, TransactionType as B8TransactionType,
35 TypeDBDriver as B8Driver,
36};
37
38#[cfg(feature = "band9")]
39use typedb_driver as driver_b9;
40#[cfg(feature = "band9")]
41use typedb_driver::answer::QueryAnswer as B9QueryAnswer;
42#[cfg(feature = "band9")]
43use typedb_driver::{
44 Addresses as B9Addresses, Credentials as B9Credentials, DriverOptions as B9DriverOptions,
45 DriverTlsConfig as B9DriverTlsConfig, TransactionOptions as B9TransactionOptions,
46 TransactionType as B9TransactionType, TypeDBDriver as B9Driver,
47};
48
49pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
51
52#[derive(Clone, Copy, Debug, Eq, PartialEq)]
54pub enum CommitFailureCertainty {
55 DefinitelyAborted,
57 Unknown,
59}
60
61#[derive(Debug, thiserror::Error)]
63pub enum RuntimeError {
64 #[error("Unsupported version: {0}")]
67 UnsupportedVersion(#[from] core_version::VersionError),
68
69 #[error("Connection error: {0}")]
71 Connection(String),
72
73 #[error("Query execution error: {0}")]
75 QueryExecution(String),
76
77 #[error("Transaction error: {0}")]
79 Transaction(String),
80
81 #[error("Resource limit [{code}]: {message}")]
83 ResourceLimit {
84 code: &'static str,
86 message: &'static str,
88 },
89
90 #[error("Answer consumer rejected a streamed provider item")]
92 AnswerConsumer,
93}
94
95#[derive(Debug, thiserror::Error)]
102pub enum RuntimeCommitError {
103 #[error(transparent)]
105 Runtime(#[from] RuntimeError),
106 #[error("Transaction error: Commit failed: {message}")]
108 Driver {
109 certainty: CommitFailureCertainty,
111 message: String,
113 },
114}
115
116impl RuntimeCommitError {
117 #[must_use]
119 pub fn into_runtime_error(self) -> RuntimeError {
120 match self {
121 Self::Runtime(error) => error,
122 Self::Driver { message, .. } => {
123 RuntimeError::Transaction(format!("Commit failed: {message}"))
124 }
125 }
126 }
127}
128
129fn commit_failure(
130 certainty: CommitFailureCertainty,
131 message: impl Into<String>,
132) -> RuntimeCommitError {
133 RuntimeCommitError::Driver {
134 certainty,
135 message: message.into(),
136 }
137}
138
139#[cfg(feature = "band8")]
140fn band8_commit_failure(error: driver_b8::Error) -> RuntimeCommitError {
141 let certainty = if matches!(&error, driver_b8::Error::Server(_)) {
142 CommitFailureCertainty::DefinitelyAborted
143 } else {
144 CommitFailureCertainty::Unknown
145 };
146 commit_failure(certainty, error.to_string())
147}
148
149#[cfg(feature = "band9")]
150fn band9_commit_failure(error: driver_b9::Error) -> RuntimeCommitError {
151 let certainty = if matches!(&error, driver_b9::Error::Server(_)) {
152 CommitFailureCertainty::DefinitelyAborted
153 } else {
154 CommitFailureCertainty::Unknown
155 };
156 commit_failure(certainty, error.to_string())
157}
158
159#[cfg(test)]
160mod commit_failure_tests {
161 use super::*;
162
163 #[test]
164 fn typed_commit_failure_preserves_the_legacy_display() {
165 for certainty in [
166 CommitFailureCertainty::DefinitelyAborted,
167 CommitFailureCertainty::Unknown,
168 ] {
169 let error = commit_failure(certainty, "driver response");
170 assert_eq!(
171 error.to_string(),
172 "Transaction error: Commit failed: driver response"
173 );
174 assert!(matches!(
175 &error,
176 RuntimeCommitError::Driver {
177 certainty: actual,
178 ..
179 } if *actual == certainty
180 ));
181 assert!(matches!(
182 error.into_runtime_error(),
183 RuntimeError::Transaction(message)
184 if message == "Commit failed: driver response"
185 ));
186 }
187 }
188
189 #[cfg(feature = "band8")]
190 #[test]
191 fn band8_opaque_commit_failure_is_unknown() {
192 let error = band8_commit_failure(driver_b8::Error::Other("transport".into()));
193 assert!(matches!(
194 error,
195 RuntimeCommitError::Driver {
196 certainty: CommitFailureCertainty::Unknown,
197 ..
198 }
199 ));
200 }
201
202 #[cfg(feature = "band9")]
203 #[test]
204 fn band9_opaque_commit_failure_is_unknown() {
205 let error = band9_commit_failure(driver_b9::Error::Other("transport".into()));
206 assert!(matches!(
207 error,
208 RuntimeCommitError::Driver {
209 certainty: CommitFailureCertainty::Unknown,
210 ..
211 }
212 ));
213 }
214}
215
216pub type Result<T> = std::result::Result<T, RuntimeError>;
218
219#[derive(Debug, Clone, Serialize, Deserialize)]
221pub enum QueryResult {
222 Ok,
224 Documents(Vec<serde_json::Value>),
226 Rows(Vec<serde_json::Value>),
228}
229
230#[derive(Debug, Clone, PartialEq)]
232pub enum RuntimeAnswerItem {
233 Row(serde_json::Value),
235 Document(serde_json::Value),
237}
238
239#[derive(Debug, Clone, Copy, PartialEq, Eq)]
241pub enum RuntimeAnswerKind {
242 Ok,
244 Rows,
246 Documents,
248}
249
250#[derive(Debug, Clone, Copy, PartialEq, Eq)]
252pub enum RuntimeAnswerControl {
253 Continue,
255 Stop,
257}
258
259#[derive(Debug, Clone)]
261pub struct RuntimeAnswerCancellation {
262 cancelled: watch::Sender<bool>,
263}
264
265impl Default for RuntimeAnswerCancellation {
266 fn default() -> Self {
267 let (cancelled, _) = watch::channel(false);
268 Self { cancelled }
269 }
270}
271
272impl RuntimeAnswerCancellation {
273 pub fn from_shared(cancelled: watch::Sender<bool>) -> Self {
275 Self { cancelled }
276 }
277
278 pub fn cancel(&self) {
280 self.cancelled.send_replace(true);
281 }
282
283 pub fn is_cancelled(&self) -> bool {
285 *self.cancelled.borrow()
286 }
287
288 async fn cancelled(&self) {
289 let mut receiver = self.cancelled.subscribe();
290 if *receiver.borrow_and_update() {
291 return;
292 }
293 while receiver.changed().await.is_ok() {
294 if *receiver.borrow_and_update() {
295 return;
296 }
297 }
298 }
299}
300
301#[doc(hidden)]
308#[derive(Debug, Clone)]
309pub struct RuntimeConnectionControl {
310 max_items: u64,
311 max_bytes: u64,
312 max_statements: u32,
313 deadline: Instant,
314 cancellation: RuntimeAnswerCancellation,
315}
316
317impl RuntimeConnectionControl {
318 #[doc(hidden)]
320 #[must_use]
321 pub const fn new(
322 max_items: u64,
323 max_bytes: u64,
324 max_statements: u32,
325 deadline: Instant,
326 cancellation: RuntimeAnswerCancellation,
327 ) -> Self {
328 Self {
329 max_items,
330 max_bytes,
331 max_statements,
332 deadline,
333 cancellation,
334 }
335 }
336}
337
338#[derive(Debug, Clone)]
340pub struct RuntimeAnswerLimits {
341 pub max_items: u64,
343 pub max_bytes: u64,
345 pub deadline: Option<Instant>,
347 pub cancellation: RuntimeAnswerCancellation,
349}
350
351impl RuntimeAnswerLimits {
352 fn unbounded() -> Self {
353 Self {
354 max_items: u64::MAX,
355 max_bytes: u64::MAX,
356 deadline: None,
357 cancellation: RuntimeAnswerCancellation::default(),
358 }
359 }
360
361 fn is_unbounded(&self) -> bool {
362 self.max_items == u64::MAX && self.max_bytes == u64::MAX && self.deadline.is_none()
363 }
364}
365
366#[derive(Debug, Clone)]
371pub struct QueryV2RuntimeAnswerLimits {
372 pub answer: RuntimeAnswerLimits,
374 pub max_collection_members: u64,
376}
377
378impl Default for QueryV2RuntimeAnswerLimits {
379 fn default() -> Self {
380 Self {
381 answer: RuntimeAnswerLimits {
382 max_items: 100_000,
383 max_bytes: 64 * 1024 * 1024,
384 deadline: None,
385 cancellation: RuntimeAnswerCancellation::default(),
386 },
387 max_collection_members: 65_536,
388 }
389 }
390}
391
392#[derive(Debug, Clone, Copy, PartialEq, Eq)]
394pub struct RuntimeAnswerStats {
395 pub kind: RuntimeAnswerKind,
397 pub processed_items: u64,
399 pub response_bytes: u64,
401 pub stopped_early: bool,
403}
404
405impl RuntimeAnswerStats {
406 fn new(kind: RuntimeAnswerKind) -> Self {
407 Self {
408 kind,
409 processed_items: 0,
410 response_bytes: 0,
411 stopped_early: false,
412 }
413 }
414}
415
416#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
418pub enum TxType {
419 Read,
421 Write,
423 Schema,
425}
426
427#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
433pub enum GivenValue {
434 Empty,
436 Boolean(bool),
438 Integer(i64),
440 Double(f64),
442 String(String),
444 Date(String),
446 Datetime(String),
448 DatetimeTz(String),
450 DatetimeTzExact {
453 local: String,
455 named_zone: Option<String>,
457 effective_offset_seconds: i32,
459 },
460 Decimal(String),
462 Duration {
464 months: u32,
466 days: u32,
468 nanos: u64,
470 },
471}
472
473#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
478pub struct GivenRowsSpec {
479 pub variables: Vec<String>,
481 pub rows: Vec<Vec<GivenValue>>,
483}
484
485pub const PINNED_DRIVER_VERSION: &str = "3.11.5";
496
497pub const PINNED_DRIVER_VERSION_B9: &str = "3.12.3";
505
506const EMBEDDED_BANDS: &[u8] = &[
515 #[cfg(feature = "band8")]
516 8,
517 #[cfg(feature = "band9")]
518 9,
519];
520
521pub fn embedded_driver_versions() -> &'static [(u8, &'static str)] {
532 &[
533 #[cfg(feature = "band8")]
534 (8, PINNED_DRIVER_VERSION),
535 #[cfg(feature = "band9")]
536 (9, PINNED_DRIVER_VERSION_B9),
537 ]
538}
539
540#[derive(Debug, Clone, Copy)]
544pub struct ConnectOptions {
545 pub http_port: u16,
547 pub tls: bool,
549 pub server_version: Option<core_version::Version>,
555}
556
557impl Default for ConnectOptions {
558 fn default() -> Self {
559 Self {
560 http_port: DEFAULT_HTTP_PORT,
561 tls: false,
562 server_version: None,
563 }
564 }
565}
566
567#[derive(Debug, Clone, PartialEq, Eq, Default)]
573pub enum TlsMode {
574 #[default]
576 Disabled,
577 NativeRoots,
579 CustomRootCa(PathBuf),
581}
582
583impl TlsMode {
584 #[must_use]
586 pub const fn is_enabled(&self) -> bool {
587 !matches!(self, Self::Disabled)
588 }
589}
590
591#[derive(Debug, Clone, PartialEq, Eq)]
593pub struct SecureConnectOptions {
594 pub http_port: u16,
596 pub tls_mode: TlsMode,
598 pub server_version: Option<core_version::Version>,
603}
604
605impl Default for SecureConnectOptions {
606 fn default() -> Self {
607 Self {
608 http_port: DEFAULT_HTTP_PORT,
609 tls_mode: TlsMode::Disabled,
610 server_version: None,
611 }
612 }
613}
614
615impl SecureConnectOptions {
616 pub fn validate_transport(&self) -> SecureResult<()> {
624 self.prepare_transport().map(|_| ())
625 }
626
627 #[doc(hidden)]
634 pub fn prepare_transport(&self) -> SecureResult<PreparedSecureConnectOptions> {
635 Ok(PreparedSecureConnectOptions {
636 http_port: self.http_port,
637 server_version: self.server_version,
638 required_server_version: None,
639 resolved_tls: ResolvedTlsMode::from_configured_path(self.tls_mode.clone())?,
640 connection_control: None,
641 })
642 }
643
644 #[doc(hidden)]
651 pub fn prepare_transport_from_validated_physical_path(
652 &self,
653 ) -> SecureResult<PreparedSecureConnectOptions> {
654 Ok(PreparedSecureConnectOptions {
655 http_port: self.http_port,
656 server_version: self.server_version,
657 required_server_version: None,
658 resolved_tls: ResolvedTlsMode::from_validated_physical_path(self.tls_mode.clone())?,
659 connection_control: None,
660 })
661 }
662
663 #[doc(hidden)]
670 pub fn prepare_transport_from_captured_custom_root(
671 &self,
672 bytes: Arc<[u8]>,
673 ) -> SecureResult<PreparedSecureConnectOptions> {
674 Ok(PreparedSecureConnectOptions {
675 http_port: self.http_port,
676 server_version: self.server_version,
677 required_server_version: None,
678 resolved_tls: ResolvedTlsMode::from_captured_custom_root(self.tls_mode.clone(), bytes)?,
679 connection_control: None,
680 })
681 }
682}
683
684impl From<ConnectOptions> for SecureConnectOptions {
685 fn from(value: ConnectOptions) -> Self {
686 Self {
687 http_port: value.http_port,
688 tls_mode: if value.tls {
689 TlsMode::NativeRoots
690 } else {
691 TlsMode::Disabled
692 },
693 server_version: value.server_version,
694 }
695 }
696}
697
698#[derive(Debug, thiserror::Error)]
703#[non_exhaustive]
704pub enum SecureConnectError {
705 #[error(transparent)]
707 TlsConfiguration(#[from] core_version::TlsConfigurationError),
708 #[error(
710 "TLS configuration error [tls_driver_lowering_failed]: TypeDB driver band {band} rejected the TLS policy"
711 )]
712 DriverTlsConfiguration {
713 band: u8,
715 },
716 #[error(transparent)]
718 Runtime(#[from] RuntimeError),
719}
720
721impl SecureConnectError {
722 #[must_use]
725 pub fn configuration_code(&self) -> Option<&'static str> {
726 match self {
727 Self::TlsConfiguration(error) => Some(error.code()),
728 Self::DriverTlsConfiguration { .. } => Some("tls_driver_lowering_failed"),
729 Self::Runtime(_) => None,
730 }
731 }
732
733 #[must_use]
742 pub fn credential_safe_diagnostic(&self) -> Option<String> {
743 match self {
744 Self::TlsConfiguration(error) => Some(format!(
745 "TLS policy preparation failed [{}]; inspect the configured trust material",
746 error.code()
747 )),
748 Self::DriverTlsConfiguration { band } => Some(format!(
749 "TLS policy lowering failed [tls_driver_lowering_failed] for TypeDB driver band {band}"
750 )),
751 Self::Runtime(RuntimeError::UnsupportedVersion(
752 error @ (core_version::VersionError::Unsupported { .. }
753 | core_version::VersionError::BandMismatch { .. }
754 | core_version::VersionError::EmbeddedUnavailable { .. }
755 | core_version::VersionError::FeatureUnsupported { .. }),
756 )) => Some(format!("Unsupported version: {error}")),
757 Self::Runtime(
758 RuntimeError::UnsupportedVersion(
759 core_version::VersionError::Probe(_) | core_version::VersionError::Parse(_),
760 )
761 | RuntimeError::Connection(_)
762 | RuntimeError::QueryExecution(_)
763 | RuntimeError::Transaction(_)
764 | RuntimeError::ResourceLimit { .. }
765 | RuntimeError::AnswerConsumer,
766 ) => None,
767 }
768 }
769
770 #[must_use]
772 pub fn into_runtime_error(self) -> RuntimeError {
773 match self {
774 Self::Runtime(error) => error,
775 other => RuntimeError::Connection(other.to_string()),
776 }
777 }
778}
779
780const TRACE_CODE_DRIVER_DROP_CLOSE_FAILED: &str = "typedb_runtime_driver_drop_close_failed";
781const TRACE_CODE_VERSION_GATE_PASSED: &str = "typedb_runtime_version_gate_passed";
782#[cfg(all(feature = "band8", feature = "band9"))]
783const TRACE_CODE_BAND9_UPGRADE_SUCCEEDED: &str = "typedb_runtime_band9_upgrade_succeeded";
784#[cfg(all(feature = "band8", feature = "band9"))]
785const TRACE_CODE_BAND9_UPGRADE_FAILED: &str = "typedb_runtime_band9_upgrade_failed";
786#[cfg(feature = "band8")]
787const TRACE_CODE_BAND8_FALLBACK_CONNECTED: &str = "typedb_runtime_band8_fallback_connected";
788const TRACE_CODE_CONNECTED: &str = "typedb_runtime_connected";
789
790fn runtime_trace_failure_code(error: &RuntimeError) -> &'static str {
791 match error {
792 RuntimeError::UnsupportedVersion(_) => "typedb_runtime_unsupported_version",
793 RuntimeError::Connection(_) => "typedb_runtime_connection_failed",
794 RuntimeError::QueryExecution(_) => "typedb_runtime_query_failed",
795 RuntimeError::Transaction(_) => "typedb_runtime_transaction_failed",
796 RuntimeError::ResourceLimit { .. } => "typedb_runtime_resource_limit",
797 RuntimeError::AnswerConsumer => "typedb_runtime_answer_consumer_failed",
798 }
799}
800
801#[cfg(all(feature = "band8", feature = "band9"))]
802fn secure_trace_failure_code(error: &SecureConnectError) -> &'static str {
803 match error {
804 SecureConnectError::TlsConfiguration(error) => error.code(),
805 SecureConnectError::DriverTlsConfiguration { .. } => "tls_driver_lowering_failed",
806 SecureConnectError::Runtime(error) => runtime_trace_failure_code(error),
807 }
808}
809
810fn trace_driver_drop_close_failure(error: &RuntimeError) {
811 tracing::warn!(
812 code = TRACE_CODE_DRIVER_DROP_CLOSE_FAILED,
813 failure_code = runtime_trace_failure_code(error),
814 "TypeDB driver cleanup failed during final runtime drop"
815 );
816}
817
818fn trace_version_gate_passed(_address: &str, band: u8, server_version: core_version::Version) {
819 tracing::debug!(
820 code = TRACE_CODE_VERSION_GATE_PASSED,
821 driver_band = band,
822 %server_version,
823 "TypeDB runtime version gate passed"
824 );
825}
826
827#[cfg(all(feature = "band8", feature = "band9"))]
828fn trace_band9_upgrade_succeeded(_address: &str, server_version: core_version::Version) {
829 tracing::debug!(
830 code = TRACE_CODE_BAND9_UPGRADE_SUCCEEDED,
831 driver_band = 9_u8,
832 %server_version,
833 "TypeDB runtime upgraded its validated fallback connection"
834 );
835}
836
837#[cfg(all(feature = "band8", feature = "band9"))]
838fn trace_band9_upgrade_failed(
839 _address: &str,
840 server_version: core_version::Version,
841 error: &SecureConnectError,
842) {
843 tracing::warn!(
844 code = TRACE_CODE_BAND9_UPGRADE_FAILED,
845 failure_code = secure_trace_failure_code(error),
846 driver_band = 9_u8,
847 %server_version,
848 "TypeDB runtime could not upgrade its validated fallback connection"
849 );
850}
851
852#[cfg(feature = "band8")]
853fn trace_band8_fallback_connected(_address: &str, server_version: core_version::Version) {
854 tracing::debug!(
855 code = TRACE_CODE_BAND8_FALLBACK_CONNECTED,
856 driver_band = 8_u8,
857 %server_version,
858 "TypeDB runtime retained its validated fallback connection"
859 );
860}
861
862fn trace_runtime_connected(
863 _address: &str,
864 driver_band: u8,
865 server_version: Option<core_version::Version>,
866) {
867 match server_version {
868 Some(server_version) => tracing::info!(
869 code = TRACE_CODE_CONNECTED,
870 driver_band,
871 server_version_known = true,
872 %server_version,
873 "TypeDB runtime connected"
874 ),
875 None => tracing::info!(
876 code = TRACE_CODE_CONNECTED,
877 driver_band,
878 server_version_known = false,
879 "TypeDB runtime connected"
880 ),
881 }
882}
883
884pub type SecureResult<T> = std::result::Result<T, SecureConnectError>;
886
887#[derive(Clone, Debug)]
895struct ResolvedTlsMode {
896 probe_mode: ResolvedTlsProbeMode,
897 #[cfg(feature = "band8")]
898 band8: DriverTlsConfig,
899 #[cfg(feature = "band9")]
900 band9: B9DriverTlsConfig,
901}
902
903#[derive(Clone, Debug)]
904enum ResolvedTlsProbeMode {
905 Disabled,
906 NativeRoots,
907 CustomRootCa(core_version::RetainedCustomRootCa),
908}
909
910#[doc(hidden)]
917#[derive(Clone, Debug)]
918pub struct PreparedSecureConnectOptions {
919 http_port: u16,
920 server_version: Option<core_version::Version>,
921 required_server_version: Option<core_version::Version>,
922 resolved_tls: ResolvedTlsMode,
923 connection_control: Option<RuntimeConnectionControl>,
924}
925
926impl PreparedSecureConnectOptions {
927 #[doc(hidden)]
932 #[must_use]
933 pub fn with_connection_control(mut self, control: RuntimeConnectionControl) -> Self {
934 self.connection_control = Some(control);
935 self
936 }
937
938 #[doc(hidden)]
945 #[must_use]
946 pub fn with_generated_3_12_3_requirement(mut self) -> Self {
947 self.server_version = None;
948 self.required_server_version = Some(core_version::Version::new(3, 12, 3));
949 self
950 }
951}
952
953#[derive(Debug)]
954struct ConnectionMeter {
955 control: Option<RuntimeConnectionControl>,
956 statements: u32,
957 response_bytes: u64,
958 admitted_items: u64,
959}
960
961impl ConnectionMeter {
962 fn new(control: RuntimeConnectionControl) -> Self {
963 Self {
964 control: Some(control),
965 statements: 0,
966 response_bytes: 0,
967 admitted_items: 0,
968 }
969 }
970
971 fn legacy() -> Self {
972 Self {
973 control: None,
974 statements: 0,
975 response_bytes: 0,
976 admitted_items: 0,
977 }
978 }
979
980 fn control(&self) -> Option<&RuntimeConnectionControl> {
981 self.control.as_ref()
982 }
983
984 fn check_interruption(&self) -> Result<()> {
985 let Some(control) = &self.control else {
986 return Ok(());
987 };
988 if control.cancellation.is_cancelled() {
989 return Err(RuntimeError::ResourceLimit {
990 code: "provider_cancelled",
991 message: "provider connection was cancelled",
992 });
993 }
994 if tokio::time::Instant::now() >= tokio::time::Instant::from_std(control.deadline) {
995 return Err(RuntimeError::ResourceLimit {
996 code: "transaction_deadline_exceeded",
997 message: "provider connection deadline expired",
998 });
999 }
1000 Ok(())
1001 }
1002
1003 fn require_possible_success(&self) -> Result<()> {
1004 self.check_interruption()?;
1005 if self
1006 .control
1007 .as_ref()
1008 .is_some_and(|control| control.max_items == 0)
1009 {
1010 return Err(RuntimeError::ResourceLimit {
1011 code: "processed_item_limit",
1012 message: "provider connection item limit exceeded",
1013 });
1014 }
1015 Ok(())
1016 }
1017
1018 fn charge_statement(&mut self) -> Result<()> {
1019 self.check_interruption()?;
1020 let Some(control) = &self.control else {
1021 return Ok(());
1022 };
1023 let next = self
1024 .statements
1025 .checked_add(1)
1026 .ok_or(RuntimeError::ResourceLimit {
1027 code: "provider_statement_counter_overflow",
1028 message: "provider connection statement counter overflowed",
1029 })?;
1030 if next > control.max_statements {
1031 return Err(RuntimeError::ResourceLimit {
1032 code: "provider_statement_limit",
1033 message: "provider connection statement limit exceeded",
1034 });
1035 }
1036 self.statements = next;
1037 Ok(())
1038 }
1039
1040 fn charge_version_evidence(&mut self, version: core_version::Version) -> Result<()> {
1041 let Some(control) = &self.control else {
1042 return Ok(());
1043 };
1044 let bytes =
1049 u64::try_from(version.to_string().len()).map_err(|_| RuntimeError::ResourceLimit {
1050 code: "response_byte_counter_overflow",
1051 message: "provider connection response byte counter overflowed",
1052 })?;
1053 let next = self
1054 .response_bytes
1055 .checked_add(bytes)
1056 .ok_or(RuntimeError::ResourceLimit {
1057 code: "response_byte_counter_overflow",
1058 message: "provider connection response byte counter overflowed",
1059 })?;
1060 if next > control.max_bytes {
1061 return Err(RuntimeError::ResourceLimit {
1062 code: "response_byte_limit",
1063 message: "provider connection response byte limit exceeded",
1064 });
1065 }
1066 self.response_bytes = next;
1067 Ok(())
1068 }
1069
1070 fn admit_connection(&mut self) -> Result<()> {
1071 let Some(control) = &self.control else {
1072 return Ok(());
1073 };
1074 let next = self
1075 .admitted_items
1076 .checked_add(1)
1077 .ok_or(RuntimeError::ResourceLimit {
1078 code: "processed_item_counter_overflow",
1079 message: "provider connection item counter overflowed",
1080 })?;
1081 if next > control.max_items {
1082 return Err(RuntimeError::ResourceLimit {
1083 code: "processed_item_limit",
1084 message: "provider connection item limit exceeded",
1085 });
1086 }
1087 self.admitted_items = next;
1088 Ok(())
1089 }
1090}
1091
1092async fn await_connection_work<T>(
1093 future: impl Future<Output = T>,
1094 control: Option<&RuntimeConnectionControl>,
1095) -> Result<T> {
1096 let Some(control) = control else {
1097 return Ok(future.await);
1098 };
1099 tokio::pin!(future);
1100 let cancellation = control.cancellation.cancelled();
1101 tokio::pin!(cancellation);
1102 let deadline = tokio::time::sleep_until(tokio::time::Instant::from_std(control.deadline));
1103 tokio::pin!(deadline);
1104
1105 tokio::select! {
1106 biased;
1107 output = &mut future => Ok(output),
1108 () = &mut cancellation => Err(RuntimeError::ResourceLimit {
1109 code: "provider_cancelled",
1110 message: "provider connection was cancelled",
1111 }),
1112 () = &mut deadline => Err(RuntimeError::ResourceLimit {
1113 code: "transaction_deadline_exceeded",
1114 message: "provider connection deadline expired",
1115 }),
1116 }
1117}
1118
1119fn controlled_connection_error(error: SecureConnectError) -> SecureConnectError {
1120 match error {
1121 SecureConnectError::TlsConfiguration(
1122 core_version::TlsConfigurationError::CustomRootCaNotFile { .. }
1123 | core_version::TlsConfigurationError::CustomRootCaUnreadable { .. }
1124 | core_version::TlsConfigurationError::CustomRootCaTooLarge { .. }
1125 | core_version::TlsConfigurationError::CustomRootCaInvalidPem { .. },
1126 ) => SecureConnectError::TlsConfiguration(
1127 core_version::TlsConfigurationError::ClientConfiguration,
1128 ),
1129 SecureConnectError::Runtime(RuntimeError::UnsupportedVersion(
1130 core_version::VersionError::Probe(_) | core_version::VersionError::Parse(_),
1131 )) => SecureConnectError::Runtime(RuntimeError::Connection(
1132 "TypeDB version discovery failed".to_owned(),
1133 )),
1134 SecureConnectError::Runtime(RuntimeError::Connection(_)) => SecureConnectError::Runtime(
1135 RuntimeError::Connection("TypeDB connection failed".to_owned()),
1136 ),
1137 SecureConnectError::Runtime(RuntimeError::QueryExecution(_)) => {
1138 SecureConnectError::Runtime(RuntimeError::QueryExecution(
1139 "TypeDB connection query failed".to_owned(),
1140 ))
1141 }
1142 SecureConnectError::Runtime(RuntimeError::Transaction(_)) => SecureConnectError::Runtime(
1143 RuntimeError::Transaction("TypeDB connection transaction failed".to_owned()),
1144 ),
1145 other => other,
1146 }
1147}
1148
1149#[cfg(feature = "band8")]
1150fn is_connection_control_error(error: &SecureConnectError) -> bool {
1151 matches!(
1152 error,
1153 SecureConnectError::Runtime(RuntimeError::ResourceLimit { .. })
1154 )
1155}
1156
1157fn close_driver_preserving(driver: &DriverHandle, error: SecureConnectError) -> SecureConnectError {
1158 if let Err(close_error) = driver.force_close() {
1159 trace_driver_drop_close_failure(&close_error);
1160 }
1161 error
1162}
1163
1164impl ResolvedTlsProbeMode {
1165 #[cfg(test)]
1166 const fn is_enabled(&self) -> bool {
1167 !matches!(self, Self::Disabled)
1168 }
1169}
1170
1171impl ResolvedTlsMode {
1172 fn from_captured_custom_root(mode: TlsMode, bytes: Arc<[u8]>) -> SecureResult<Self> {
1173 let material = match &mode {
1174 TlsMode::CustomRootCa(path) => {
1175 core_version::RetainedCustomRootCa::load_captured_bytes(path, bytes)?
1176 }
1177 TlsMode::Disabled | TlsMode::NativeRoots => {
1178 return Err(core_version::TlsConfigurationError::ClientConfiguration.into());
1179 }
1180 };
1181 Self::lower(mode, ResolvedTlsProbeMode::CustomRootCa(material))
1182 }
1183
1184 fn from_validated_physical_path(mode: TlsMode) -> SecureResult<Self> {
1185 let probe_mode = match &mode {
1186 TlsMode::Disabled => ResolvedTlsProbeMode::Disabled,
1187 TlsMode::NativeRoots => ResolvedTlsProbeMode::NativeRoots,
1188 TlsMode::CustomRootCa(path) => {
1189 ResolvedTlsProbeMode::CustomRootCa(core_version::RetainedCustomRootCa::load(path)?)
1190 }
1191 };
1192 Self::lower(mode, probe_mode)
1193 }
1194
1195 fn from_configured_path(mode: TlsMode) -> SecureResult<Self> {
1196 let probe_mode = match &mode {
1197 TlsMode::Disabled => ResolvedTlsProbeMode::Disabled,
1198 TlsMode::NativeRoots => ResolvedTlsProbeMode::NativeRoots,
1199 TlsMode::CustomRootCa(path) => ResolvedTlsProbeMode::CustomRootCa(
1200 core_version::RetainedCustomRootCa::load_configured_alias(path)?,
1201 ),
1202 };
1203 Self::lower(mode, probe_mode)
1204 }
1205
1206 fn lower(mode: TlsMode, probe_mode: ResolvedTlsProbeMode) -> SecureResult<Self> {
1207 let custom_root = match &probe_mode {
1208 ResolvedTlsProbeMode::CustomRootCa(material) => Some(material),
1209 ResolvedTlsProbeMode::Disabled | ResolvedTlsProbeMode::NativeRoots => None,
1210 };
1211
1212 if matches!(mode, TlsMode::CustomRootCa(_)) != custom_root.is_some() {
1213 return Err(core_version::TlsConfigurationError::ClientConfiguration.into());
1214 }
1215
1216 #[cfg(feature = "band8")]
1217 let band8 = match &mode {
1218 TlsMode::Disabled => DriverTlsConfig::disabled(),
1219 TlsMode::NativeRoots => DriverTlsConfig::enabled_with_native_root_ca(),
1220 TlsMode::CustomRootCa(_) => custom_root
1221 .ok_or(core_version::TlsConfigurationError::ClientConfiguration)?
1222 .with_driver_root_path(DriverTlsConfig::enabled_with_root_ca)?
1223 .map_err(|_| SecureConnectError::DriverTlsConfiguration { band: 8 })?,
1224 };
1225
1226 #[cfg(feature = "band9")]
1227 let band9 = match &mode {
1228 TlsMode::Disabled => B9DriverTlsConfig::disabled(),
1229 TlsMode::NativeRoots => B9DriverTlsConfig::enabled_with_native_root_ca(),
1230 TlsMode::CustomRootCa(_) => custom_root
1231 .ok_or(core_version::TlsConfigurationError::ClientConfiguration)?
1232 .with_driver_root_path(B9DriverTlsConfig::enabled_with_root_ca)?
1233 .map_err(|_| SecureConnectError::DriverTlsConfiguration { band: 9 })?,
1234 };
1235
1236 Ok(Self {
1237 probe_mode,
1238 #[cfg(feature = "band8")]
1239 band8,
1240 #[cfg(feature = "band9")]
1241 band9,
1242 })
1243 }
1244}
1245
1246enum DriverHandleInner {
1248 #[cfg(feature = "band8")]
1249 B8(B8Driver),
1250 #[cfg(feature = "band9")]
1251 B9(B9Driver),
1252}
1253
1254fn contain_driver_shutdown(shutdown: impl FnOnce() -> Result<()>) -> Result<()> {
1255 match std::panic::catch_unwind(std::panic::AssertUnwindSafe(shutdown)) {
1256 Ok(result) => result,
1257 Err(_) => Err(RuntimeError::Connection(
1258 "Driver close failed: upstream driver panicked during shutdown".to_owned(),
1259 )),
1260 }
1261}
1262
1263#[cfg(test)]
1264fn run_driver_shutdown(
1265 state: &AtomicU8,
1266 close_lock: &Mutex<()>,
1267 shutdown: impl FnOnce() -> Result<()>,
1268) -> Result<()> {
1269 let _close_guard = match close_lock.lock() {
1273 Ok(close_guard) => close_guard,
1274 Err(poisoned) => poisoned.into_inner(),
1275 };
1276
1277 run_driver_shutdown_locked(state, shutdown)
1278}
1279
1280fn run_driver_shutdown_locked(
1281 state: &AtomicU8,
1282 shutdown: impl FnOnce() -> Result<()>,
1283) -> Result<()> {
1284 match state.load(AtomicOrdering::Acquire) {
1285 DRIVER_CLOSED => return Ok(()),
1286 DRIVER_OPEN => {
1287 state.store(DRIVER_CLOSING, AtomicOrdering::Release);
1293 }
1294 DRIVER_CLOSING => {
1295 }
1299 _ => unreachable!("driver shutdown state is internal and validated"),
1300 }
1301
1302 let result = contain_driver_shutdown(shutdown);
1303 if result.is_ok() {
1304 state.store(DRIVER_CLOSED, AtomicOrdering::Release);
1305 }
1306 result
1307}
1308
1309fn retain_driver_transaction(
1310 state: &AtomicU8,
1311 active_transactions: &AtomicUsize,
1312 close_lock: &Mutex<()>,
1313) -> Result<()> {
1314 let _close_guard = match close_lock.lock() {
1315 Ok(close_guard) => close_guard,
1316 Err(poisoned) => poisoned.into_inner(),
1317 };
1318 if state.load(AtomicOrdering::Acquire) != DRIVER_OPEN {
1319 return Err(RuntimeError::Connection(
1320 "TypeDB driver connection is closed".to_owned(),
1321 ));
1322 }
1323 active_transactions.fetch_add(1, AtomicOrdering::AcqRel);
1324 Ok(())
1325}
1326
1327pub(crate) struct DriverHandle {
1333 inner: DriverHandleInner,
1334 state: AtomicU8,
1335 active_transactions: AtomicUsize,
1336 close_lock: Mutex<()>,
1337}
1338
1339const DRIVER_OPEN: u8 = 0;
1340const DRIVER_CLOSING: u8 = 1;
1341const DRIVER_CLOSED: u8 = 2;
1342
1343impl DriverHandle {
1344 #[cfg(feature = "band8")]
1345 fn band8(driver: B8Driver) -> Self {
1346 Self {
1347 inner: DriverHandleInner::B8(driver),
1348 state: AtomicU8::new(DRIVER_OPEN),
1349 active_transactions: AtomicUsize::new(0),
1350 close_lock: Mutex::new(()),
1351 }
1352 }
1353
1354 #[cfg(feature = "band9")]
1355 fn band9(driver: B9Driver) -> Self {
1356 Self {
1357 inner: DriverHandleInner::B9(driver),
1358 state: AtomicU8::new(DRIVER_OPEN),
1359 active_transactions: AtomicUsize::new(0),
1360 close_lock: Mutex::new(()),
1361 }
1362 }
1363
1364 fn inner(&self) -> &DriverHandleInner {
1365 &self.inner
1366 }
1367
1368 fn band(&self) -> u8 {
1369 match &self.inner {
1370 #[cfg(feature = "band8")]
1371 DriverHandleInner::B8(_) => 8,
1372 #[cfg(feature = "band9")]
1373 DriverHandleInner::B9(_) => 9,
1374 }
1375 }
1376
1377 fn ensure_open(&self) -> Result<()> {
1378 if self.state.load(AtomicOrdering::Acquire) == DRIVER_OPEN {
1379 Ok(())
1380 } else {
1381 Err(RuntimeError::Connection(
1382 "TypeDB driver connection is closed".to_owned(),
1383 ))
1384 }
1385 }
1386
1387 fn shutdown_started(&self) -> bool {
1388 self.state.load(AtomicOrdering::Acquire) != DRIVER_OPEN
1389 }
1390
1391 fn force_close(&self) -> Result<()> {
1398 let _close_guard = match self.close_lock.lock() {
1399 Ok(close_guard) => close_guard,
1400 Err(poisoned) => poisoned.into_inner(),
1401 };
1402 if self.active_transactions.load(AtomicOrdering::Acquire) != 0 {
1403 return Err(RuntimeError::ResourceLimit {
1404 code: "resource_in_use",
1405 message: "connection cannot close while a transaction is active",
1406 });
1407 }
1408 run_driver_shutdown_locked(&self.state, || match &self.inner {
1409 #[cfg(feature = "band8")]
1410 DriverHandleInner::B8(driver) => driver
1411 .force_close()
1412 .map_err(|error| RuntimeError::Connection(format!("Driver close failed: {error}"))),
1413 #[cfg(feature = "band9")]
1414 DriverHandleInner::B9(driver) => driver
1415 .force_close()
1416 .map_err(|error| RuntimeError::Connection(format!("Driver close failed: {error}"))),
1417 })
1418 }
1419}
1420
1421impl Drop for DriverHandle {
1422 fn drop(&mut self) {
1423 if let Err(error) = self.force_close() {
1432 trace_driver_drop_close_failure(&error);
1433 }
1434 }
1435}
1436
1437pub struct TypeDBRuntime {
1439 driver: Arc<DriverHandle>,
1440 server_version: Option<core_version::Version>,
1446}
1447
1448impl TypeDBRuntime {
1449 pub fn force_close(&self) -> Result<()> {
1462 self.driver.force_close()
1463 }
1464}
1465
1466#[cfg(test)]
1468pub(crate) async fn gated_driver_with_probe<F>(
1469 address: &str,
1470 username: &str,
1471 password: &str,
1472 options: ConnectOptions,
1473 probe: F,
1474) -> Result<(DriverHandle, Option<core_version::Version>)>
1475where
1476 F: FnOnce(
1477 &str,
1478 u16,
1479 bool,
1480 ) -> std::result::Result<core_version::Version, core_version::VersionError>
1481 + Send
1482 + 'static,
1483{
1484 gated_driver_secure_with_probe(
1485 address,
1486 username,
1487 password,
1488 options.into(),
1489 move |probe_address, http_port, tls_mode| {
1490 probe(probe_address, http_port, tls_mode.is_enabled())
1491 .map_err(core_version::VersionProbeError::Probe)
1492 },
1493 )
1494 .await
1495 .map_err(SecureConnectError::into_runtime_error)
1496}
1497
1498async fn gated_driver_secure_with_probe<F>(
1506 address: &str,
1507 username: &str,
1508 password: &str,
1509 options: SecureConnectOptions,
1510 probe: F,
1511) -> SecureResult<(DriverHandle, Option<core_version::Version>)>
1512where
1513 F: FnOnce(
1514 &str,
1515 u16,
1516 &ResolvedTlsProbeMode,
1517 )
1518 -> std::result::Result<core_version::Version, core_version::VersionProbeError>
1519 + Send
1520 + 'static,
1521{
1522 let prepared = options.prepare_transport()?;
1523 gated_driver_prepared_with_probe(address, username, password, prepared, probe).await
1524}
1525
1526async fn gated_driver_prepared_with_probe<F>(
1527 address: &str,
1528 username: &str,
1529 password: &str,
1530 options: PreparedSecureConnectOptions,
1531 probe: F,
1532) -> SecureResult<(DriverHandle, Option<core_version::Version>)>
1533where
1534 F: FnOnce(
1535 &str,
1536 u16,
1537 &ResolvedTlsProbeMode,
1538 )
1539 -> std::result::Result<core_version::Version, core_version::VersionProbeError>
1540 + Send
1541 + 'static,
1542{
1543 let controlled = options.connection_control.is_some();
1544 let result = async {
1545 let PreparedSecureConnectOptions {
1546 http_port,
1547 server_version,
1548 required_server_version,
1549 resolved_tls,
1550 connection_control,
1551 } = options;
1552 let mut meter = connection_control
1553 .map(ConnectionMeter::new)
1554 .unwrap_or_else(ConnectionMeter::legacy);
1555 meter.require_possible_success()?;
1559
1560 let connected = if let Some(server_version) = server_version {
1561 let driver = driver_for_server_version(
1562 address,
1563 username,
1564 password,
1565 server_version,
1566 &resolved_tls,
1567 &mut meter,
1568 )
1569 .await?;
1570 (driver, Some(server_version))
1571 } else {
1572 meter.charge_statement()?;
1576 let address_owned = address.to_string();
1577 let probe_mode = resolved_tls.probe_mode.clone();
1578 let probe_task =
1579 tokio::task::spawn_blocking(move || probe(&address_owned, http_port, &probe_mode));
1580 let probe_result = await_connection_work(probe_task, meter.control())
1581 .await?
1582 .map_err(|error| {
1583 RuntimeError::Connection(if controlled {
1584 "Version probe task failed".to_owned()
1585 } else {
1586 format!("Version probe task panicked: {error}")
1587 })
1588 })?;
1589
1590 match probe_result {
1591 Ok(server_version) => {
1592 meter.charge_version_evidence(server_version)?;
1593 require_generated_server_version(server_version, required_server_version)?;
1594 let driver = driver_for_server_version(
1595 address,
1596 username,
1597 password,
1598 server_version,
1599 &resolved_tls,
1600 &mut meter,
1601 )
1602 .await?;
1603 (driver, Some(server_version))
1604 }
1605 Err(core_version::VersionProbeError::Probe(http_error)) => {
1606 grpc_fallback_driver(
1607 address,
1608 username,
1609 password,
1610 http_error,
1611 &resolved_tls,
1612 &mut meter,
1613 required_server_version,
1614 )
1615 .await?
1616 }
1617 Err(core_version::VersionProbeError::TlsConfiguration(error)) => {
1618 return Err(SecureConnectError::TlsConfiguration(error));
1619 }
1620 }
1621 };
1622
1623 if let Err(error) = meter.admit_connection() {
1624 return Err(close_driver_preserving(&connected.0, error.into()));
1625 }
1626 Ok(connected)
1627 }
1628 .await;
1629
1630 if controlled {
1631 result.map_err(controlled_connection_error)
1632 } else {
1633 result
1634 }
1635}
1636
1637fn require_generated_server_version(
1638 observed: core_version::Version,
1639 required: Option<core_version::Version>,
1640) -> Result<()> {
1641 let Some(required) = required else {
1642 return Ok(());
1643 };
1644 if observed == required {
1645 return Ok(());
1646 }
1647 Err(RuntimeError::UnsupportedVersion(
1648 core_version::VersionError::FeatureUnsupported {
1649 feature: "generated-package exact provider compatibility",
1650 server: observed,
1651 required,
1652 remediation: "connect the generated package to its exact required TypeDB release",
1653 },
1654 ))
1655}
1656
1657fn validate_server_band(server_version: &core_version::Version) -> Result<u8> {
1658 core_version::check_server_supported(server_version, EMBEDDED_BANDS)
1666 .map_err(RuntimeError::UnsupportedVersion)?;
1667
1668 Ok(
1669 core_version::negotiate_server_band(server_version, EMBEDDED_BANDS)
1670 .expect("check_server_supported accepted a server with no negotiable band"),
1671 )
1672}
1673
1674async fn driver_for_server_version(
1675 address: &str,
1676 username: &str,
1677 password: &str,
1678 server_version: core_version::Version,
1679 tls: &ResolvedTlsMode,
1680 meter: &mut ConnectionMeter,
1681) -> SecureResult<DriverHandle> {
1682 let band = validate_server_band(&server_version)?;
1683
1684 trace_version_gate_passed(address, band, server_version);
1685
1686 #[cfg(feature = "band8")]
1687 if band == 8 {
1688 return connect_band8_driver(address, username, password, tls, meter)
1689 .await
1690 .map(DriverHandle::band8);
1691 }
1692
1693 #[cfg(feature = "band9")]
1694 if band == 9 {
1695 return connect_band9_driver(address, username, password, tls, meter)
1696 .await
1697 .map(DriverHandle::band9);
1698 }
1699
1700 Err(RuntimeError::Connection(format!(
1705 "No compiled driver band supports the negotiated server band ({band})"
1706 ))
1707 .into())
1708}
1709
1710#[cfg(feature = "band8")]
1711async fn connect_band8_driver(
1712 address: &str,
1713 username: &str,
1714 password: &str,
1715 tls: &ResolvedTlsMode,
1716 meter: &mut ConnectionMeter,
1717) -> SecureResult<B8Driver> {
1718 let addresses = Addresses::try_from_address_str(address)
1719 .map_err(|e| RuntimeError::Connection(format!("Invalid TypeDB address {address}: {e}")))?;
1720 meter.charge_statement()?;
1721 let result = await_connection_work(
1722 B8Driver::new(
1723 addresses,
1724 B8Credentials::new(username, password),
1725 DriverOptions::new(tls.band8.clone()),
1726 ),
1727 meter.control(),
1728 )
1729 .await?;
1730 result.map_err(|e| {
1731 RuntimeError::Connection(format!("Failed to connect to {address}: {e}")).into()
1732 })
1733}
1734
1735#[cfg(feature = "band9")]
1736async fn connect_band9_driver(
1737 address: &str,
1738 username: &str,
1739 password: &str,
1740 tls: &ResolvedTlsMode,
1741 meter: &mut ConnectionMeter,
1742) -> SecureResult<B9Driver> {
1743 let addresses = B9Addresses::try_from_address_str(address)
1744 .map_err(|e| RuntimeError::Connection(format!("Invalid TypeDB address {address}: {e}")))?;
1745 meter.charge_statement()?;
1746 let result = await_connection_work(
1747 B9Driver::new(
1748 addresses,
1749 B9Credentials::new(username, password),
1750 B9DriverOptions::new(tls.band9.clone()),
1751 ),
1752 meter.control(),
1753 )
1754 .await?;
1755 result.map_err(|e| {
1756 RuntimeError::Connection(format!("Failed to connect to {address}: {e}")).into()
1757 })
1758}
1759
1760async fn grpc_fallback_driver(
1761 address: &str,
1762 username: &str,
1763 password: &str,
1764 http_error: core_version::VersionError,
1765 tls: &ResolvedTlsMode,
1766 meter: &mut ConnectionMeter,
1767 required_server_version: Option<core_version::Version>,
1768) -> SecureResult<(DriverHandle, Option<core_version::Version>)> {
1769 #[cfg(not(feature = "band8"))]
1770 let _ = (address, username, password, tls, meter);
1771 let mut failures = vec![format!("HTTP version probe failed: {http_error}")];
1772
1773 #[cfg(feature = "band8")]
1779 {
1780 match connect_band8_driver(address, username, password, tls, meter).await {
1781 Ok(driver) => {
1782 let driver = DriverHandle::band8(driver);
1786 let reported = match driver.inner() {
1787 DriverHandleInner::B8(driver) => {
1788 meter.charge_statement()?;
1789 await_connection_work(driver.server_version(), meter.control())
1790 .await?
1791 .map(|reported| reported.version().to_owned())
1792 .map_err(|error| error.to_string())
1793 }
1794 #[cfg(feature = "band9")]
1795 _ => {
1796 return Err(RuntimeError::Connection(
1797 "Internal TypeDB driver-band mismatch".to_owned(),
1798 )
1799 .into());
1800 }
1801 };
1802 let classification = match classify_band8_grpc_version(address, reported) {
1803 Ok(classification) => classification,
1804 Err(error) => {
1805 driver.force_close().map_err(SecureConnectError::Runtime)?;
1806 return Err(SecureConnectError::Runtime(error));
1807 }
1808 };
1809 match classification {
1810 Band8GrpcVersion::Validated(server_version) => {
1811 if let Err(error) = meter.charge_version_evidence(server_version) {
1812 return Err(close_driver_preserving(&driver, error.into()));
1813 }
1814 if let Err(error) = require_generated_server_version(
1815 server_version,
1816 required_server_version,
1817 ) {
1818 return Err(close_driver_preserving(&driver, error.into()));
1819 }
1820 #[cfg(feature = "band9")]
1825 if core_version::negotiate_server_band(&server_version, EMBEDDED_BANDS)
1826 == Some(9)
1827 {
1828 match connect_band9_driver(address, username, password, tls, meter)
1836 .await
1837 {
1838 Ok(b9_driver) => {
1839 let b9_driver = DriverHandle::band9(b9_driver);
1840 driver.force_close().map_err(SecureConnectError::Runtime)?;
1841 trace_band9_upgrade_succeeded(address, server_version);
1842 return Ok((b9_driver, Some(server_version)));
1843 }
1844 Err(error) => {
1845 if is_connection_control_error(&error) {
1846 return Err(close_driver_preserving(&driver, error));
1847 }
1848 trace_band9_upgrade_failed(address, server_version, &error);
1849 }
1850 }
1851 }
1852 trace_band8_fallback_connected(address, server_version);
1853 return Ok((driver, Some(server_version)));
1854 }
1855 Band8GrpcVersion::RetryableFailure(failure) => {
1856 driver.force_close().map_err(SecureConnectError::Runtime)?;
1857 failures.push(failure);
1858 }
1859 }
1860 }
1861 Err(error) if is_connection_control_error(&error) => return Err(error),
1862 Err(error) => failures.push(format!("band-8 gRPC attempt failed: {error}")),
1863 }
1864 }
1865
1866 #[cfg(not(feature = "band8"))]
1867 failures.push("band-8 gRPC attempt skipped: band8 feature is not compiled in".to_string());
1868
1869 Err(
1870 RuntimeError::UnsupportedVersion(core_version::VersionError::Probe(format!(
1871 "HTTP version probe and gRPC fallback both failed: {}",
1872 failures.join("; ")
1873 )))
1874 .into(),
1875 )
1876}
1877
1878#[cfg(feature = "band8")]
1879#[derive(Debug)]
1880enum Band8GrpcVersion {
1881 Validated(core_version::Version),
1882 RetryableFailure(String),
1883}
1884
1885#[cfg(feature = "band8")]
1886fn classify_band8_grpc_version(
1887 address: &str,
1888 reported: std::result::Result<String, String>,
1889) -> Result<Band8GrpcVersion> {
1890 let reported = match reported {
1891 Ok(reported) => reported,
1892 Err(error) => {
1893 return Ok(Band8GrpcVersion::RetryableFailure(format!(
1898 "band-8 gRPC version validation failed after connect to {address}: {error}"
1899 )));
1900 }
1901 };
1902
1903 let server_version = reported
1907 .parse::<core_version::Version>()
1908 .map_err(RuntimeError::UnsupportedVersion)?;
1909 validate_server_band(&server_version)?;
1910 if !core_version::server_accepted_bands(&server_version).contains(&8) {
1911 return Err(RuntimeError::UnsupportedVersion(
1912 core_version::VersionError::Probe(format!(
1913 "band-8 gRPC connection reported server version {server_version}, \
1914 which does not accept band 8"
1915 )),
1916 ));
1917 }
1918
1919 Ok(Band8GrpcVersion::Validated(server_version))
1920}
1921
1922fn probe_server_version_for_tls_mode(
1923 address: &str,
1924 http_port: u16,
1925 tls_mode: &ResolvedTlsProbeMode,
1926) -> std::result::Result<core_version::Version, core_version::VersionProbeError> {
1927 match tls_mode {
1928 ResolvedTlsProbeMode::Disabled => {
1929 core_version::server_version_plaintext(address, http_port)
1930 }
1931 ResolvedTlsProbeMode::NativeRoots => {
1932 core_version::server_version_native_roots(address, http_port)
1933 }
1934 ResolvedTlsProbeMode::CustomRootCa(material) => {
1935 core_version::server_version_retained_custom_root_ca(address, http_port, material)
1936 }
1937 }
1938}
1939
1940async fn gated_driver(
1942 address: &str,
1943 username: &str,
1944 password: &str,
1945 options: ConnectOptions,
1946) -> Result<(DriverHandle, Option<core_version::Version>)> {
1947 gated_driver_secure(address, username, password, options.into())
1948 .await
1949 .map_err(SecureConnectError::into_runtime_error)
1950}
1951
1952async fn gated_driver_secure(
1954 address: &str,
1955 username: &str,
1956 password: &str,
1957 options: SecureConnectOptions,
1958) -> SecureResult<(DriverHandle, Option<core_version::Version>)> {
1959 gated_driver_secure_with_probe(
1960 address,
1961 username,
1962 password,
1963 options,
1964 probe_server_version_for_tls_mode,
1965 )
1966 .await
1967}
1968
1969async fn gated_driver_prepared_secure(
1970 address: &str,
1971 username: &str,
1972 password: &str,
1973 options: PreparedSecureConnectOptions,
1974) -> SecureResult<(DriverHandle, Option<core_version::Version>)> {
1975 gated_driver_prepared_with_probe(
1976 address,
1977 username,
1978 password,
1979 options,
1980 probe_server_version_for_tls_mode,
1981 )
1982 .await
1983}
1984
1985impl TypeDBRuntime {
1986 pub async fn connect(
1993 address: &str,
1994 username: &str,
1995 password: &str,
1996 options: ConnectOptions,
1997 ) -> Result<Self> {
1998 let (driver, server_version) = gated_driver(address, username, password, options).await?;
1999 trace_runtime_connected(address, driver.band(), server_version);
2000 Ok(Self {
2001 driver: Arc::new(driver),
2002 server_version,
2003 })
2004 }
2005
2006 pub async fn connect_secure(
2012 address: &str,
2013 username: &str,
2014 password: &str,
2015 options: SecureConnectOptions,
2016 ) -> SecureResult<Self> {
2017 let (driver, server_version) =
2018 gated_driver_secure(address, username, password, options).await?;
2019 trace_runtime_connected(address, driver.band(), server_version);
2020 Ok(Self {
2021 driver: Arc::new(driver),
2022 server_version,
2023 })
2024 }
2025
2026 #[doc(hidden)]
2031 pub async fn connect_prepared_secure(
2032 address: &str,
2033 username: &str,
2034 password: &str,
2035 options: PreparedSecureConnectOptions,
2036 ) -> SecureResult<Self> {
2037 let (driver, server_version) =
2038 gated_driver_prepared_secure(address, username, password, options).await?;
2039 trace_runtime_connected(address, driver.band(), server_version);
2040 Ok(Self {
2041 driver: Arc::new(driver),
2042 server_version,
2043 })
2044 }
2045
2046 pub fn server_version(&self) -> Option<core_version::Version> {
2051 self.server_version
2052 }
2053
2054 pub fn supports_given_rows(&self) -> bool {
2062 match self.driver.inner() {
2063 #[cfg(feature = "band8")]
2064 DriverHandleInner::B8(_) => false,
2065 #[cfg(feature = "band9")]
2066 DriverHandleInner::B9(_) => true,
2067 }
2068 }
2069}
2070
2071fn runtime_error_diagnostic(error: RuntimeError) -> String {
2072 match error {
2073 RuntimeError::Connection(message) => message,
2074 other => other.to_string(),
2075 }
2076}
2077
2078fn finish_one_shot_database_operation<T>(operation: Result<T>, close: Result<()>) -> Result<T> {
2079 match (operation, close) {
2080 (Ok(value), Ok(())) => Ok(value),
2081 (Ok(_), Err(close_error)) => Err(close_error),
2082 (Err(operation_error), Ok(())) => Err(operation_error),
2083 (Err(operation_error), Err(close_error)) => Err(RuntimeError::Connection(format!(
2084 "{}; additionally, TypeDB driver cleanup failed: {}",
2085 runtime_error_diagnostic(operation_error),
2086 runtime_error_diagnostic(close_error),
2087 ))),
2088 }
2089}
2090
2091pub async fn ensure_database_exists(
2102 address: &str,
2103 database: &str,
2104 username: &str,
2105 password: &str,
2106 options: ConnectOptions,
2107) -> Result<()> {
2108 ensure_database_exists_secure(address, database, username, password, options.into())
2109 .await
2110 .map_err(SecureConnectError::into_runtime_error)
2111}
2112
2113pub async fn ensure_database_exists_secure(
2118 address: &str,
2119 database: &str,
2120 username: &str,
2121 password: &str,
2122 options: SecureConnectOptions,
2123) -> SecureResult<()> {
2124 ensure_database_exists_prepared_secure(
2125 address,
2126 database,
2127 username,
2128 password,
2129 options.prepare_transport()?,
2130 )
2131 .await
2132}
2133
2134#[doc(hidden)]
2136pub async fn ensure_database_exists_prepared_secure(
2137 address: &str,
2138 database: &str,
2139 username: &str,
2140 password: &str,
2141 options: PreparedSecureConnectOptions,
2142) -> SecureResult<()> {
2143 let (driver, server_version) =
2144 gated_driver_prepared_secure(address, username, password, options).await?;
2145 let runtime = TypeDBRuntime {
2146 driver: Arc::new(driver),
2147 server_version,
2148 };
2149 let operation = async {
2150 if !runtime.database_exists(database).await? {
2151 runtime.create_database(database).await?;
2152 }
2153 Ok(())
2154 }
2155 .await;
2156 let close = runtime.force_close();
2157 finish_one_shot_database_operation(operation, close).map_err(SecureConnectError::Runtime)
2158}
2159
2160pub async fn database_exists(
2165 address: &str,
2166 database: &str,
2167 username: &str,
2168 password: &str,
2169 options: ConnectOptions,
2170) -> Result<bool> {
2171 database_exists_secure(address, database, username, password, options.into())
2172 .await
2173 .map_err(SecureConnectError::into_runtime_error)
2174}
2175
2176pub async fn database_exists_secure(
2178 address: &str,
2179 database: &str,
2180 username: &str,
2181 password: &str,
2182 options: SecureConnectOptions,
2183) -> SecureResult<bool> {
2184 database_exists_prepared_secure(
2185 address,
2186 database,
2187 username,
2188 password,
2189 options.prepare_transport()?,
2190 )
2191 .await
2192}
2193
2194#[doc(hidden)]
2196pub async fn database_exists_prepared_secure(
2197 address: &str,
2198 database: &str,
2199 username: &str,
2200 password: &str,
2201 options: PreparedSecureConnectOptions,
2202) -> SecureResult<bool> {
2203 let (driver, server_version) =
2204 gated_driver_prepared_secure(address, username, password, options).await?;
2205 let runtime = TypeDBRuntime {
2206 driver: Arc::new(driver),
2207 server_version,
2208 };
2209 let operation = runtime.database_exists(database).await;
2210 let close = runtime.force_close();
2211 finish_one_shot_database_operation(operation, close).map_err(SecureConnectError::Runtime)
2212}
2213
2214pub async fn delete_database_secure(
2220 address: &str,
2221 database: &str,
2222 username: &str,
2223 password: &str,
2224 options: SecureConnectOptions,
2225) -> SecureResult<()> {
2226 delete_database_prepared_secure(
2227 address,
2228 database,
2229 username,
2230 password,
2231 options.prepare_transport()?,
2232 )
2233 .await
2234}
2235
2236#[doc(hidden)]
2238pub async fn delete_database_prepared_secure(
2239 address: &str,
2240 database: &str,
2241 username: &str,
2242 password: &str,
2243 options: PreparedSecureConnectOptions,
2244) -> SecureResult<()> {
2245 let runtime =
2246 TypeDBRuntime::connect_prepared_secure(address, username, password, options).await?;
2247 let operation = async {
2248 if runtime.database_exists(database).await? {
2249 runtime.delete_database(database).await?;
2250 }
2251 Ok(())
2252 }
2253 .await;
2254 let close = runtime.force_close();
2255 finish_one_shot_database_operation(operation, close).map_err(SecureConnectError::Runtime)
2256}
2257
2258impl TypeDBRuntime {
2259 pub fn open_transaction(
2261 &self,
2262 database: &str,
2263 tx_type: TxType,
2264 ) -> BoxFuture<'_, Result<RuntimeTransaction>> {
2265 self.open_transaction_with_optional_timeout(database, tx_type, None)
2266 }
2267
2268 pub fn open_transaction_with_timeout(
2274 &self,
2275 database: &str,
2276 tx_type: TxType,
2277 timeout: Duration,
2278 ) -> BoxFuture<'_, Result<RuntimeTransaction>> {
2279 self.open_transaction_with_optional_timeout(database, tx_type, Some(timeout))
2280 }
2281
2282 fn open_transaction_with_optional_timeout(
2283 &self,
2284 database: &str,
2285 tx_type: TxType,
2286 timeout: Option<Duration>,
2287 ) -> BoxFuture<'_, Result<RuntimeTransaction>> {
2288 let db = database.to_string();
2289 let driver_lease = Arc::clone(&self.driver);
2290 Box::pin(async move {
2291 driver_lease.ensure_open()?;
2292 match driver_lease.inner() {
2293 #[cfg(feature = "band8")]
2294 DriverHandleInner::B8(d) => {
2295 let typedb_tx_type = match tx_type {
2296 TxType::Read => B8TransactionType::Read,
2297 TxType::Write => B8TransactionType::Write,
2298 TxType::Schema => B8TransactionType::Schema,
2299 };
2300 let mut options = B8TransactionOptions::new();
2301 if let Some(timeout) = timeout {
2302 options = options
2303 .transaction_timeout(timeout)
2304 .schema_lock_acquire_timeout(timeout);
2305 }
2306 let transaction = d
2307 .transaction_with_options(&db, typedb_tx_type, options)
2308 .await
2309 .map_err(|e| {
2310 RuntimeError::Transaction(format!("Failed to open transaction: {e}"))
2311 })?;
2312 retain_driver_transaction(
2313 &driver_lease.state,
2314 &driver_lease.active_transactions,
2315 &driver_lease.close_lock,
2316 )?;
2317 Ok(RuntimeTransaction {
2318 inner: RuntimeTransactionInner::B8(Some(transaction)),
2319 driver_lease: Some(driver_lease),
2320 })
2321 }
2322 #[cfg(feature = "band9")]
2323 DriverHandleInner::B9(d) => {
2324 let typedb_tx_type = match tx_type {
2325 TxType::Read => B9TransactionType::Read,
2326 TxType::Write => B9TransactionType::Write,
2327 TxType::Schema => B9TransactionType::Schema,
2328 };
2329 let mut options = B9TransactionOptions::new();
2330 if let Some(timeout) = timeout {
2331 options = options
2332 .transaction_timeout(timeout)
2333 .schema_lock_acquire_timeout(timeout);
2334 }
2335 let transaction = d
2336 .transaction_with_options(&db, typedb_tx_type, options)
2337 .await
2338 .map_err(|e| {
2339 RuntimeError::Transaction(format!("Failed to open transaction: {e}"))
2340 })?;
2341 retain_driver_transaction(
2342 &driver_lease.state,
2343 &driver_lease.active_transactions,
2344 &driver_lease.close_lock,
2345 )?;
2346 Ok(RuntimeTransaction {
2347 inner: RuntimeTransactionInner::B9(Some(transaction)),
2348 driver_lease: Some(driver_lease),
2349 })
2350 }
2351 }
2352 })
2353 }
2354
2355 pub fn is_open(&self) -> bool {
2357 if self.driver.shutdown_started() {
2358 return false;
2359 }
2360 match self.driver.inner() {
2361 #[cfg(feature = "band8")]
2362 DriverHandleInner::B8(d) => d.is_open(),
2363 #[cfg(feature = "band9")]
2364 DriverHandleInner::B9(d) => d.is_open(),
2365 }
2366 }
2367
2368 pub fn database_exists(&self, database: &str) -> BoxFuture<'_, Result<bool>> {
2370 let database = database.to_string();
2371 Box::pin(async move {
2372 self.driver.ensure_open()?;
2373 match self.driver.inner() {
2374 #[cfg(feature = "band8")]
2375 DriverHandleInner::B8(d) => {
2376 d.databases().contains(database).await.map_err(|e| {
2377 RuntimeError::Connection(format!("Database lookup failed: {e}"))
2378 })
2379 }
2380 #[cfg(feature = "band9")]
2381 DriverHandleInner::B9(d) => {
2382 d.databases().contains(database).await.map_err(|e| {
2383 RuntimeError::Connection(format!("Database lookup failed: {e}"))
2384 })
2385 }
2386 }
2387 })
2388 }
2389
2390 pub fn create_database(&self, database: &str) -> BoxFuture<'_, Result<()>> {
2392 let database = database.to_string();
2393 Box::pin(async move {
2394 self.driver.ensure_open()?;
2395 match self.driver.inner() {
2396 #[cfg(feature = "band8")]
2397 DriverHandleInner::B8(d) => {
2398 d.databases().create(database).await.map_err(|e| {
2399 RuntimeError::Connection(format!("Database create failed: {e}"))
2400 })
2401 }
2402 #[cfg(feature = "band9")]
2403 DriverHandleInner::B9(d) => {
2404 d.databases().create(database).await.map_err(|e| {
2405 RuntimeError::Connection(format!("Database create failed: {e}"))
2406 })
2407 }
2408 }
2409 })
2410 }
2411
2412 pub fn delete_database(&self, database: &str) -> BoxFuture<'_, Result<()>> {
2414 let database = database.to_string();
2415 Box::pin(async move {
2416 self.driver.ensure_open()?;
2417 match self.driver.inner() {
2418 #[cfg(feature = "band8")]
2419 DriverHandleInner::B8(d) => {
2420 let db = d.databases().get(database).await.map_err(|e| {
2421 RuntimeError::Connection(format!("Database lookup failed: {e}"))
2422 })?;
2423 db.delete().await.map_err(|e| {
2424 RuntimeError::Connection(format!("Database delete failed: {e}"))
2425 })
2426 }
2427 #[cfg(feature = "band9")]
2428 DriverHandleInner::B9(d) => {
2429 let db = d.databases().get(database).await.map_err(|e| {
2430 RuntimeError::Connection(format!("Database lookup failed: {e}"))
2431 })?;
2432 db.delete().await.map_err(|e| {
2433 RuntimeError::Connection(format!("Database delete failed: {e}"))
2434 })
2435 }
2436 }
2437 })
2438 }
2439
2440 pub fn schema_text(&self, database: &str) -> BoxFuture<'_, Result<String>> {
2442 let database = database.to_string();
2443 Box::pin(async move {
2444 self.driver.ensure_open()?;
2445 match self.driver.inner() {
2446 #[cfg(feature = "band8")]
2447 DriverHandleInner::B8(d) => {
2448 let db = d.databases().get(&database).await.map_err(|e| {
2449 RuntimeError::Connection(format!("Database lookup failed: {e}"))
2450 })?;
2451 db.schema()
2452 .await
2453 .map_err(|e| RuntimeError::Connection(format!("Schema export failed: {e}")))
2454 }
2455 #[cfg(feature = "band9")]
2456 DriverHandleInner::B9(d) => {
2457 let db = d.databases().get(&database).await.map_err(|e| {
2458 RuntimeError::Connection(format!("Database lookup failed: {e}"))
2459 })?;
2460 db.schema()
2461 .await
2462 .map_err(|e| RuntimeError::Connection(format!("Schema export failed: {e}")))
2463 }
2464 }
2465 })
2466 }
2467}
2468
2469enum RuntimeTransactionInner {
2475 #[cfg(feature = "band8")]
2476 B8(Option<type_bridge_typedb_driver_b8::Transaction>),
2477 #[cfg(feature = "band9")]
2478 B9(Option<typedb_driver::Transaction>),
2479}
2480
2481pub struct RuntimeTransaction {
2483 inner: RuntimeTransactionInner,
2484 driver_lease: Option<Arc<DriverHandle>>,
2487}
2488
2489impl Drop for RuntimeTransaction {
2490 fn drop(&mut self) {
2491 if let Some(driver) = &self.driver_lease {
2492 let previous = driver
2493 .active_transactions
2494 .fetch_sub(1, AtomicOrdering::AcqRel);
2495 debug_assert!(previous != 0, "runtime transaction lease count underflow");
2496 }
2497 }
2498}
2499
2500fn runtime_cancelled_error() -> RuntimeError {
2501 RuntimeError::ResourceLimit {
2502 code: "provider_cancelled",
2503 message: "provider answer processing was cancelled",
2504 }
2505}
2506
2507fn runtime_deadline_error() -> RuntimeError {
2508 RuntimeError::ResourceLimit {
2509 code: "transaction_deadline_exceeded",
2510 message: "provider transaction deadline expired",
2511 }
2512}
2513
2514fn runtime_check(limits: &RuntimeAnswerLimits) -> Result<()> {
2515 if limits.cancellation.is_cancelled() {
2516 return Err(runtime_cancelled_error());
2517 }
2518 if limits.deadline.is_some_and(|deadline| {
2519 tokio::time::Instant::now() >= tokio::time::Instant::from_std(deadline)
2520 }) {
2521 return Err(runtime_deadline_error());
2522 }
2523 Ok(())
2524}
2525
2526async fn runtime_await<T>(
2527 future: impl Future<Output = Result<T>>,
2528 limits: &RuntimeAnswerLimits,
2529) -> Result<T> {
2530 runtime_check(limits)?;
2531 tokio::pin!(future);
2532 let cancellation = limits.cancellation.cancelled();
2533 tokio::pin!(cancellation);
2534
2535 if let Some(deadline) = limits.deadline {
2536 let deadline = tokio::time::sleep_until(tokio::time::Instant::from_std(deadline));
2537 tokio::pin!(deadline);
2538 tokio::select! {
2539 biased;
2540 result = &mut future => result,
2541 () = &mut cancellation => Err(runtime_cancelled_error()),
2542 () = &mut deadline => Err(runtime_deadline_error()),
2543 }
2544 } else {
2545 tokio::select! {
2546 biased;
2547 result = &mut future => result,
2548 () = &mut cancellation => Err(runtime_cancelled_error()),
2549 }
2550 }
2551}
2552
2553fn runtime_accept(
2554 limits: &RuntimeAnswerLimits,
2555 stats: &mut RuntimeAnswerStats,
2556 item: RuntimeAnswerItem,
2557 consumer: &mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2558) -> Result<RuntimeAnswerControl> {
2559 runtime_check(limits)?;
2560 let next_items = stats
2561 .processed_items
2562 .checked_add(1)
2563 .ok_or(RuntimeError::ResourceLimit {
2564 code: "processed_item_counter_overflow",
2565 message: "processed provider item counter overflowed",
2566 })?;
2567 if next_items > limits.max_items {
2568 return Err(RuntimeError::ResourceLimit {
2569 code: "processed_item_limit",
2570 message: "provider answer exceeded the processed-item ceiling",
2571 });
2572 }
2573 let value = match &item {
2574 RuntimeAnswerItem::Row(value) | RuntimeAnswerItem::Document(value) => value,
2575 };
2576 let encoded = u64::try_from(
2577 serde_json::to_vec(value)
2578 .map_err(|error| RuntimeError::QueryExecution(format!("Answer encode: {error}")))?
2579 .len(),
2580 )
2581 .map_err(|_| RuntimeError::ResourceLimit {
2582 code: "answer_byte_counter_overflow",
2583 message: "encoded provider answer length exceeds the counter range",
2584 })?;
2585 let next_bytes =
2586 stats
2587 .response_bytes
2588 .checked_add(encoded)
2589 .ok_or(RuntimeError::ResourceLimit {
2590 code: "answer_byte_counter_overflow",
2591 message: "provider answer byte counter overflowed",
2592 })?;
2593 if next_bytes > limits.max_bytes {
2594 return Err(RuntimeError::ResourceLimit {
2595 code: "response_byte_limit",
2596 message: "provider answer exceeded the response-byte ceiling",
2597 });
2598 }
2599 stats.processed_items = next_items;
2600 stats.response_bytes = next_bytes;
2601 let control = consumer(item);
2602 runtime_check(limits)?;
2603 let control = control?;
2604 if control == RuntimeAnswerControl::Stop {
2605 stats.stopped_early = true;
2606 }
2607 Ok(control)
2608}
2609
2610#[cfg(any(feature = "band9", test))]
2611async fn runtime_consume_stream<S>(
2612 mut stream: S,
2613 kind: RuntimeAnswerKind,
2614 limits: &RuntimeAnswerLimits,
2615 consumer: &mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2616) -> Result<RuntimeAnswerStats>
2617where
2618 S: futures::TryStream<Ok = RuntimeAnswerItem, Error = RuntimeError> + Send + Unpin,
2619{
2620 let mut stats = RuntimeAnswerStats::new(kind);
2621 while let Some(item) = runtime_await(async { stream.try_next().await }, limits).await? {
2622 if runtime_accept(limits, &mut stats, item, consumer)? == RuntimeAnswerControl::Stop {
2623 break;
2624 }
2625 }
2626 Ok(stats)
2627}
2628
2629fn materialize_runtime_answer(
2630 stats: RuntimeAnswerStats,
2631 items: Vec<RuntimeAnswerItem>,
2632) -> QueryResult {
2633 match stats.kind {
2634 RuntimeAnswerKind::Ok => QueryResult::Ok,
2635 RuntimeAnswerKind::Rows => QueryResult::Rows(
2636 items
2637 .into_iter()
2638 .filter_map(|item| match item {
2639 RuntimeAnswerItem::Row(value) => Some(value),
2640 RuntimeAnswerItem::Document(_) => None,
2641 })
2642 .collect(),
2643 ),
2644 RuntimeAnswerKind::Documents => QueryResult::Documents(
2645 items
2646 .into_iter()
2647 .filter_map(|item| match item {
2648 RuntimeAnswerItem::Document(value) => Some(value),
2649 RuntimeAnswerItem::Row(_) => None,
2650 })
2651 .collect(),
2652 ),
2653 }
2654}
2655
2656impl RuntimeTransaction {
2657 fn ensure_driver_open(&self) -> Result<()> {
2658 match &self.driver_lease {
2659 Some(driver) => driver.ensure_open(),
2660 None => Ok(()),
2661 }
2662 }
2663
2664 fn driver_shutdown_started(driver_lease: &Option<Arc<DriverHandle>>) -> bool {
2665 driver_lease
2666 .as_ref()
2667 .is_some_and(|driver| driver.shutdown_started())
2668 }
2669
2670 pub fn supports_given_rows(&self) -> bool {
2673 match &self.inner {
2674 #[cfg(feature = "band8")]
2675 RuntimeTransactionInner::B8(_) => false,
2676 #[cfg(feature = "band9")]
2677 RuntimeTransactionInner::B9(_) => true,
2678 }
2679 }
2680
2681 pub fn query(&mut self, typeql: &str) -> BoxFuture<'_, Result<QueryResult>> {
2685 let typeql = typeql.to_owned();
2686 Box::pin(async move {
2687 let mut items = Vec::new();
2688 let mut collect = |item| {
2689 items.push(item);
2690 Ok(RuntimeAnswerControl::Continue)
2691 };
2692 let stats = self
2693 .query_bounded(&typeql, RuntimeAnswerLimits::unbounded(), &mut collect)
2694 .await?;
2695 Ok(materialize_runtime_answer(stats, items))
2696 })
2697 }
2698
2699 pub fn query_bounded<'a>(
2701 &'a mut self,
2702 typeql: &'a str,
2703 limits: RuntimeAnswerLimits,
2704 consumer: &'a mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2705 ) -> BoxFuture<'a, Result<RuntimeAnswerStats>> {
2706 self.query_bounded_with_scalar_encoding(
2707 typeql,
2708 QueryV2RuntimeAnswerLimits {
2709 answer: limits,
2710 max_collection_members: u64::MAX,
2711 },
2712 ScalarJsonEncoding::DriverDisplay,
2713 consumer,
2714 )
2715 }
2716
2717 pub fn query_v2_bounded<'a>(
2719 &'a mut self,
2720 typeql: &'a str,
2721 limits: QueryV2RuntimeAnswerLimits,
2722 consumer: &'a mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2723 ) -> BoxFuture<'a, Result<RuntimeAnswerStats>> {
2724 self.query_bounded_with_scalar_encoding(
2725 typeql,
2726 limits,
2727 ScalarJsonEncoding::CanonicalV2,
2728 consumer,
2729 )
2730 }
2731
2732 #[doc(hidden)]
2734 pub async fn query_v2_materialized(
2735 &mut self,
2736 typeql: &str,
2737 limits: QueryV2RuntimeAnswerLimits,
2738 ) -> Result<QueryResult> {
2739 let mut items = Vec::new();
2740 let mut consumer = |item: RuntimeAnswerItem| {
2741 items.push(item);
2742 Ok(RuntimeAnswerControl::Continue)
2743 };
2744 let stats = self.query_v2_bounded(typeql, limits, &mut consumer).await?;
2745 Ok(materialize_runtime_answer(stats, items))
2746 }
2747
2748 fn query_bounded_with_scalar_encoding<'a>(
2749 &'a mut self,
2750 typeql: &'a str,
2751 limits: QueryV2RuntimeAnswerLimits,
2752 scalar_encoding: ScalarJsonEncoding,
2753 consumer: &'a mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2754 ) -> BoxFuture<'a, Result<RuntimeAnswerStats>> {
2755 let QueryV2RuntimeAnswerLimits {
2756 answer: limits,
2757 max_collection_members,
2758 } = limits;
2759 let tql = typeql.to_string();
2760 Box::pin(async move {
2761 self.ensure_driver_open()?;
2762 runtime_check(&limits)?;
2763 match &self.inner {
2764 #[cfg(feature = "band8")]
2765 RuntimeTransactionInner::B8(opt) => {
2766 let tx = opt.as_ref().ok_or_else(|| {
2767 RuntimeError::Transaction("Transaction already consumed".into())
2768 })?;
2769 let answer = runtime_await(
2770 async {
2771 let options = if limits.is_unbounded() {
2772 driver_b8::QueryOptions::new()
2773 } else {
2774 driver_b8::QueryOptions::new().prefetch_size(1)
2775 };
2776 tx.query_with_options(&tql, options)
2777 .await
2778 .map_err(|e| RuntimeError::QueryExecution(format!("{e}")))
2779 },
2780 &limits,
2781 )
2782 .await?;
2783 match answer {
2784 B8QueryAnswer::Ok(_) => Ok(RuntimeAnswerStats::new(RuntimeAnswerKind::Ok)),
2785 B8QueryAnswer::ConceptRowStream(_, mut stream) => {
2786 let mut stats = RuntimeAnswerStats::new(RuntimeAnswerKind::Rows);
2787 while let Some(row) = runtime_await(
2788 async {
2789 stream.try_next().await.map_err(|error| {
2790 RuntimeError::QueryExecution(format!("Row stream: {error}"))
2791 })
2792 },
2793 &limits,
2794 )
2795 .await?
2796 {
2797 let mut object = serde_json::Map::new();
2798 for (index, column) in row.get_column_names().iter().enumerate() {
2799 let value = row
2800 .row
2801 .get(index)
2802 .and_then(|concept| concept.as_ref())
2803 .map(|concept| concept_to_json_b8(concept, scalar_encoding))
2804 .transpose()?
2805 .unwrap_or(serde_json::Value::Null);
2806 object.insert(column.clone(), value);
2807 }
2808 if runtime_accept(
2809 &limits,
2810 &mut stats,
2811 RuntimeAnswerItem::Row(serde_json::Value::Object(object)),
2812 consumer,
2813 )? == RuntimeAnswerControl::Stop
2814 {
2815 break;
2816 }
2817 }
2818 Ok(stats)
2819 }
2820 B8QueryAnswer::ConceptDocumentStream(_, mut stream) => {
2821 let mut stats = RuntimeAnswerStats::new(RuntimeAnswerKind::Documents);
2822 let mut collection_members = 0_u64;
2823 while let Some(document) = runtime_await(
2824 async {
2825 stream.try_next().await.map_err(|error| {
2826 RuntimeError::QueryExecution(format!(
2827 "Document stream: {error}"
2828 ))
2829 })
2830 },
2831 &limits,
2832 )
2833 .await?
2834 {
2835 let value = document_to_json_b8(
2836 &document,
2837 &mut collection_members,
2838 max_collection_members,
2839 scalar_encoding,
2840 )?;
2841 if runtime_accept(
2842 &limits,
2843 &mut stats,
2844 RuntimeAnswerItem::Document(value),
2845 consumer,
2846 )? == RuntimeAnswerControl::Stop
2847 {
2848 break;
2849 }
2850 }
2851 Ok(stats)
2852 }
2853 }
2854 }
2855 #[cfg(feature = "band9")]
2856 RuntimeTransactionInner::B9(opt) => {
2857 let tx = opt.as_ref().ok_or_else(|| {
2858 RuntimeError::Transaction("Transaction already consumed".into())
2859 })?;
2860 let answer = runtime_await(
2861 async {
2862 let options = if limits.is_unbounded() {
2863 driver_b9::QueryOptions::new()
2864 } else {
2865 driver_b9::QueryOptions::new().prefetch_size(1)
2866 };
2867 tx.query_with_options(&tql, options)
2868 .await
2869 .map_err(|e| RuntimeError::QueryExecution(format!("{e}")))
2870 },
2871 &limits,
2872 )
2873 .await?;
2874 consume_answer_b9(
2875 answer,
2876 &limits,
2877 max_collection_members,
2878 scalar_encoding,
2879 consumer,
2880 )
2881 .await
2882 }
2883 }
2884 })
2885 }
2886
2887 pub fn query_with_rows(
2895 &mut self,
2896 typeql: &str,
2897 rows: GivenRowsSpec,
2898 ) -> BoxFuture<'_, Result<QueryResult>> {
2899 let typeql = typeql.to_owned();
2900 Box::pin(async move {
2901 let mut items = Vec::new();
2902 let mut collect = |item| {
2903 items.push(item);
2904 Ok(RuntimeAnswerControl::Continue)
2905 };
2906 let stats = self
2907 .query_with_rows_bounded(
2908 &typeql,
2909 rows,
2910 RuntimeAnswerLimits::unbounded(),
2911 &mut collect,
2912 )
2913 .await?;
2914 Ok(materialize_runtime_answer(stats, items))
2915 })
2916 }
2917
2918 pub fn query_with_rows_bounded<'a>(
2921 &'a mut self,
2922 typeql: &'a str,
2923 rows: GivenRowsSpec,
2924 limits: RuntimeAnswerLimits,
2925 consumer: &'a mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2926 ) -> BoxFuture<'a, Result<RuntimeAnswerStats>> {
2927 self.query_with_rows_bounded_with_scalar_encoding(
2928 typeql,
2929 rows,
2930 QueryV2RuntimeAnswerLimits {
2931 answer: limits,
2932 max_collection_members: u64::MAX,
2933 },
2934 ScalarJsonEncoding::DriverDisplay,
2935 consumer,
2936 )
2937 }
2938
2939 pub fn query_v2_with_rows_bounded<'a>(
2941 &'a mut self,
2942 typeql: &'a str,
2943 rows: GivenRowsSpec,
2944 limits: QueryV2RuntimeAnswerLimits,
2945 consumer: &'a mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2946 ) -> BoxFuture<'a, Result<RuntimeAnswerStats>> {
2947 self.query_with_rows_bounded_with_scalar_encoding(
2948 typeql,
2949 rows,
2950 limits,
2951 ScalarJsonEncoding::CanonicalV2,
2952 consumer,
2953 )
2954 }
2955
2956 fn query_with_rows_bounded_with_scalar_encoding<'a>(
2957 &'a mut self,
2958 typeql: &'a str,
2959 rows: GivenRowsSpec,
2960 limits: QueryV2RuntimeAnswerLimits,
2961 scalar_encoding: ScalarJsonEncoding,
2962 consumer: &'a mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
2963 ) -> BoxFuture<'a, Result<RuntimeAnswerStats>> {
2964 let QueryV2RuntimeAnswerLimits {
2965 answer: limits,
2966 max_collection_members,
2967 } = limits;
2968 let tql = typeql.to_owned();
2969 #[cfg(not(feature = "band9"))]
2970 let _ = (
2971 &tql,
2972 &rows,
2973 max_collection_members,
2974 scalar_encoding,
2975 &consumer,
2976 );
2977 Box::pin(async move {
2978 self.ensure_driver_open()?;
2979 runtime_check(&limits)?;
2980 match &self.inner {
2981 #[cfg(feature = "band8")]
2982 RuntimeTransactionInner::B8(_) => Err(RuntimeError::QueryExecution(
2983 "given-stage parameterized queries require the band-9 driver \
2984 (TypeDB 3.12+); this connection negotiated band 8"
2985 .into(),
2986 )),
2987 #[cfg(feature = "band9")]
2988 RuntimeTransactionInner::B9(opt) => {
2989 let tx = opt.as_ref().ok_or_else(|| {
2990 RuntimeError::Transaction("Transaction already consumed".into())
2991 })?;
2992 let given_rows = given_rows_b9(rows)?;
2993 let answer = runtime_await(
2994 async {
2995 let options = if limits.is_unbounded() {
2996 driver_b9::QueryOptions::new()
2997 } else {
2998 driver_b9::QueryOptions::new().prefetch_size(1)
2999 };
3000 tx.query_with_options_and_rows(&tql, options, Some(given_rows))
3001 .await
3002 .map_err(|e| RuntimeError::QueryExecution(format!("{e}")))
3003 },
3004 &limits,
3005 )
3006 .await?;
3007 consume_answer_b9(
3008 answer,
3009 &limits,
3010 max_collection_members,
3011 scalar_encoding,
3012 consumer,
3013 )
3014 .await
3015 }
3016 }
3017 })
3018 }
3019
3020 pub fn commit(&mut self) -> BoxFuture<'_, Result<()>> {
3022 let commit = self.commit_classified();
3023 Box::pin(async move { commit.await.map_err(RuntimeCommitError::into_runtime_error) })
3024 }
3025
3026 pub fn commit_classified(
3028 &mut self,
3029 ) -> BoxFuture<'_, std::result::Result<(), RuntimeCommitError>> {
3030 let driver_lease = self.driver_lease.clone();
3033 match &mut self.inner {
3034 #[cfg(feature = "band8")]
3035 RuntimeTransactionInner::B8(opt) => {
3036 let tx = opt.take();
3037 Box::pin(async move {
3038 if let Some(driver) = &driver_lease {
3039 driver.ensure_open().map_err(RuntimeCommitError::Runtime)?;
3040 }
3041 let t = tx.ok_or_else(|| {
3042 RuntimeCommitError::Runtime(RuntimeError::Transaction(
3043 "Transaction already consumed".into(),
3044 ))
3045 })?;
3046 t.commit().await.map_err(band8_commit_failure)
3047 })
3048 }
3049 #[cfg(feature = "band9")]
3050 RuntimeTransactionInner::B9(opt) => {
3051 let tx = opt.take();
3052 Box::pin(async move {
3053 if let Some(driver) = &driver_lease {
3054 driver.ensure_open().map_err(RuntimeCommitError::Runtime)?;
3055 }
3056 let t = tx.ok_or_else(|| {
3057 RuntimeCommitError::Runtime(RuntimeError::Transaction(
3058 "Transaction already consumed".into(),
3059 ))
3060 })?;
3061 t.commit().await.map_err(band9_commit_failure)
3062 })
3063 }
3064 }
3065 }
3066
3067 pub fn rollback(&mut self) -> BoxFuture<'_, Result<()>> {
3069 let driver_lease = self.driver_lease.clone();
3073 match &mut self.inner {
3074 #[cfg(feature = "band8")]
3075 RuntimeTransactionInner::B8(opt) => {
3076 let tx = opt.take();
3077 Box::pin(async move {
3078 if let Some(driver) = &driver_lease {
3079 driver.ensure_open()?;
3080 }
3081 let t = tx.ok_or_else(|| {
3082 RuntimeError::Transaction("Transaction already consumed".into())
3083 })?;
3084 if !t.is_open() {
3085 return Ok(());
3086 }
3087 t.rollback()
3088 .await
3089 .map_err(|e| RuntimeError::Transaction(format!("Rollback failed: {e}")))
3090 })
3091 }
3092 #[cfg(feature = "band9")]
3093 RuntimeTransactionInner::B9(opt) => {
3094 let tx = opt.take();
3095 Box::pin(async move {
3096 if let Some(driver) = &driver_lease {
3097 driver.ensure_open()?;
3098 }
3099 let t = tx.ok_or_else(|| {
3100 RuntimeError::Transaction("Transaction already consumed".into())
3101 })?;
3102 if !t.is_open() {
3103 return Ok(());
3104 }
3105 t.rollback()
3106 .await
3107 .map_err(|e| RuntimeError::Transaction(format!("Rollback failed: {e}")))
3108 })
3109 }
3110 }
3111 }
3112
3113 pub fn close(&mut self) -> BoxFuture<'_, Result<()>> {
3115 let driver_lease = self.driver_lease.clone();
3116 match &mut self.inner {
3117 #[cfg(feature = "band8")]
3118 RuntimeTransactionInner::B8(opt) => {
3119 let tx = opt.take();
3120 Box::pin(async move {
3121 let Some(t) = tx else {
3122 return Ok(());
3123 };
3124 if Self::driver_shutdown_started(&driver_lease) {
3125 drop(t);
3126 return Ok(());
3127 }
3128 t.close()
3129 .await
3130 .map_err(|e| RuntimeError::Transaction(format!("Close failed: {e}")))
3131 })
3132 }
3133 #[cfg(feature = "band9")]
3134 RuntimeTransactionInner::B9(opt) => {
3135 let tx = opt.take();
3136 Box::pin(async move {
3137 let Some(t) = tx else {
3138 return Ok(());
3139 };
3140 if Self::driver_shutdown_started(&driver_lease) {
3141 drop(t);
3142 return Ok(());
3143 }
3144 t.close()
3145 .await
3146 .map_err(|e| RuntimeError::Transaction(format!("Close failed: {e}")))
3147 })
3148 }
3149 }
3150 }
3151}
3152
3153fn document_type_to_json(kind: &str, label: &str) -> serde_json::Value {
3154 serde_json::Value::Object(serde_json::Map::from_iter([
3155 (
3156 "kind".to_owned(),
3157 serde_json::Value::String(kind.to_owned()),
3158 ),
3159 (
3160 "label".to_owned(),
3161 serde_json::Value::String(label.to_owned()),
3162 ),
3163 ]))
3164}
3165
3166fn document_attribute_type_to_json(label: &str, value_type: Option<&str>) -> serde_json::Value {
3167 serde_json::Value::Object(serde_json::Map::from_iter([
3168 (
3169 "kind".to_owned(),
3170 serde_json::Value::String("attribute".to_owned()),
3171 ),
3172 (
3173 "label".to_owned(),
3174 serde_json::Value::String(label.to_owned()),
3175 ),
3176 (
3177 "valueType".to_owned(),
3178 serde_json::Value::String(value_type.unwrap_or("none").to_owned()),
3179 ),
3180 ]))
3181}
3182
3183fn runtime_document_member_error(message: &'static str) -> RuntimeError {
3184 RuntimeError::ResourceLimit {
3185 code: "query_v2_document_member_limit",
3186 message,
3187 }
3188}
3189
3190#[derive(Clone, Copy, Debug, Eq, PartialEq)]
3191enum ScalarJsonEncoding {
3192 DriverDisplay,
3193 CanonicalV2,
3194}
3195
3196fn canonical_decimal_text(display: String) -> Result<String> {
3197 type_bridge_core_lib::decimal::parse_decimal(&display)
3198 .map(|decimal| decimal.canonical_string())
3199 .ok_or_else(|| {
3200 RuntimeError::QueryExecution("provider returned an unparseable decimal".to_owned())
3201 })
3202}
3203
3204fn checked_datetime_tz_local<Tz>(
3205 value: &chrono::DateTime<Tz>,
3206) -> Result<(chrono::NaiveDateTime, i32)>
3207where
3208 Tz: chrono::TimeZone,
3209 Tz::Offset: chrono::Offset,
3210{
3211 use chrono::Offset as _;
3212
3213 let offset_seconds = value.offset().fix().local_minus_utc();
3214 let local = value
3215 .naive_utc()
3216 .checked_add_signed(chrono::TimeDelta::seconds(i64::from(offset_seconds)))
3217 .ok_or_else(|| {
3218 RuntimeError::QueryExecution(
3219 "provider datetime-tz local value is outside the supported range".to_owned(),
3220 )
3221 })?;
3222 Ok((local, offset_seconds))
3223}
3224
3225fn append_exact_offset(rendered: &mut String, offset_seconds: i64) {
3226 use std::fmt::Write as _;
3227
3228 if offset_seconds == 0 {
3229 rendered.push('Z');
3230 return;
3231 }
3232 let sign = if offset_seconds < 0 { '-' } else { '+' };
3233 let absolute = offset_seconds.unsigned_abs();
3234 let hours = absolute / 3_600;
3235 let minutes = (absolute % 3_600) / 60;
3236 let seconds = absolute % 60;
3237 write!(rendered, "{sign}{hours:02}:{minutes:02}")
3238 .expect("writing an offset to String cannot fail");
3239 if seconds != 0 {
3240 write!(rendered, ":{seconds:02}").expect("writing an offset to String cannot fail");
3241 }
3242}
3243
3244fn exact_naive_datetime(value: &chrono::NaiveDateTime) -> String {
3245 use std::fmt::Write as _;
3246
3247 use chrono::{Datelike as _, Timelike as _};
3248
3249 let mut rendered = String::new();
3250 let year = value.year();
3251 match year {
3252 0..=9999 => write!(rendered, "{year:04}"),
3253 10_000.. => write!(rendered, "+{year}"),
3254 -9999..=-1 => write!(rendered, "-{:04}", -year),
3255 _ => write!(rendered, "{year}"),
3256 }
3257 .expect("writing a year to String cannot fail");
3258 write!(
3259 rendered,
3260 "-{:02}-{:02}T{:02}:{:02}:{:02}",
3261 value.month(),
3262 value.day(),
3263 value.hour(),
3264 value.minute(),
3265 value.second(),
3266 )
3267 .expect("writing datetime components to String cannot fail");
3268 let nanosecond = value.nanosecond();
3269 if nanosecond != 0 {
3270 let fraction = format!("{nanosecond:09}");
3271 write!(rendered, ".{}", fraction.trim_end_matches('0'))
3272 .expect("writing a datetime fraction to String cannot fail");
3273 }
3274 rendered
3275}
3276
3277fn exact_duration(months: u32, days: u32, nanos: u64) -> String {
3278 use std::fmt::Write as _;
3279
3280 let seconds = nanos / 1_000_000_000;
3281 let nanosecond = nanos % 1_000_000_000;
3282 let mut rendered = String::from("P");
3283 if months != 0 {
3284 write!(rendered, "{months}M").expect("writing duration months to String cannot fail");
3285 }
3286 if days != 0 {
3287 write!(rendered, "{days}D").expect("writing duration days to String cannot fail");
3288 }
3289 if seconds != 0 || nanosecond != 0 || (months == 0 && days == 0) {
3290 write!(rendered, "T{seconds}").expect("writing duration seconds to String cannot fail");
3291 if nanosecond != 0 {
3292 let fraction = format!("{nanosecond:09}");
3293 write!(rendered, ".{}", fraction.trim_end_matches('0'))
3294 .expect("writing a duration fraction to String cannot fail");
3295 }
3296 rendered.push('S');
3297 }
3298 rendered
3299}
3300
3301macro_rules! define_driver_json_conversion {
3302 (
3303 $feature:literal,
3304 $driver:ident,
3305 $document_fn:ident,
3306 $node_fn:ident,
3307 $leaf_fn:ident,
3308 $value_fn:ident
3309 ) => {
3310 #[cfg(feature = $feature)]
3311 fn $document_fn(
3312 document: &$driver::answer::ConceptDocument,
3313 collection_members: &mut u64,
3314 max_collection_members: u64,
3315 scalar_encoding: ScalarJsonEncoding,
3316 ) -> Result<serde_json::Value> {
3317 document
3318 .root
3319 .as_ref()
3320 .map(|node| {
3321 $node_fn(
3322 node,
3323 collection_members,
3324 max_collection_members,
3325 scalar_encoding,
3326 )
3327 })
3328 .transpose()
3329 .map(|value| value.unwrap_or(serde_json::Value::Null))
3330 }
3331
3332 #[cfg(feature = $feature)]
3333 fn $node_fn(
3334 node: &$driver::answer::concept_document::Node,
3335 collection_members: &mut u64,
3336 max_collection_members: u64,
3337 scalar_encoding: ScalarJsonEncoding,
3338 ) -> Result<serde_json::Value> {
3339 use $driver::answer::concept_document::Node;
3340
3341 match node {
3342 Node::Map(map) => map
3343 .iter()
3344 .map(|(name, node)| {
3345 $node_fn(
3346 node,
3347 collection_members,
3348 max_collection_members,
3349 scalar_encoding,
3350 )
3351 .map(|value| (name.clone(), value))
3352 })
3353 .collect::<Result<serde_json::Map<String, serde_json::Value>>>()
3354 .map(serde_json::Value::Object),
3355 Node::List(list) => {
3356 let members = u64::try_from(list.len()).map_err(|_| {
3357 runtime_document_member_error(
3358 "document list member count exceeds the supported counter range",
3359 )
3360 })?;
3361 let next = collection_members.checked_add(members).ok_or_else(|| {
3362 runtime_document_member_error("document list member counter overflowed")
3363 })?;
3364 if next > max_collection_members {
3365 return Err(runtime_document_member_error(
3366 "document lists exceed the aggregate member ceiling",
3367 ));
3368 }
3369 *collection_members = next;
3370 list.iter()
3371 .map(|node| {
3372 $node_fn(
3373 node,
3374 collection_members,
3375 max_collection_members,
3376 scalar_encoding,
3377 )
3378 })
3379 .collect::<Result<Vec<_>>>()
3380 .map(serde_json::Value::Array)
3381 }
3382 Node::Leaf(Some(leaf)) => $leaf_fn(leaf, scalar_encoding),
3383 Node::Leaf(None) => Ok(serde_json::Value::Null),
3384 }
3385 }
3386
3387 #[cfg(feature = $feature)]
3388 fn $leaf_fn(
3389 leaf: &$driver::answer::concept_document::Leaf,
3390 scalar_encoding: ScalarJsonEncoding,
3391 ) -> Result<serde_json::Value> {
3392 use $driver::answer::concept_document::Leaf;
3393 use $driver::concept::Concept;
3394
3395 match leaf {
3396 Leaf::Empty => Ok(serde_json::Value::Null),
3397 Leaf::Concept(concept) => match concept {
3398 Concept::EntityType(_) => {
3399 Ok(document_type_to_json("entity", concept.get_label()))
3400 }
3401 Concept::RelationType(_) => {
3402 Ok(document_type_to_json("relation", concept.get_label()))
3403 }
3404 Concept::RoleType(_) => {
3405 Ok(document_type_to_json("relation:role", concept.get_label()))
3406 }
3407 Concept::AttributeType(_) => {
3408 let value_type = concept.try_get_value_type();
3409 Ok(document_attribute_type_to_json(
3410 concept.get_label(),
3411 value_type.as_ref().map(|value_type| value_type.name()),
3412 ))
3413 }
3414 Concept::Attribute(_) | Concept::Value(_) => {
3415 let value = concept.try_get_value().ok_or_else(|| {
3416 RuntimeError::QueryExecution(
3417 "document value concept did not carry a value".to_owned(),
3418 )
3419 })?;
3420 $value_fn(value, scalar_encoding)
3421 }
3422 Concept::Entity(_) | Concept::Relation(_) => Err(RuntimeError::QueryExecution(
3423 "document response carried an unsupported thing instance".to_owned(),
3424 )),
3425 },
3426 Leaf::ValueType(value_type) => {
3427 Ok(serde_json::Value::String(value_type.name().to_owned()))
3428 }
3429 Leaf::Kind(kind) => Ok(serde_json::Value::String(kind.name().to_owned())),
3430 }
3431 }
3432
3433 #[cfg(feature = $feature)]
3434 fn $value_fn(
3435 value: &$driver::concept::Value,
3436 scalar_encoding: ScalarJsonEncoding,
3437 ) -> Result<serde_json::Value> {
3438 use $driver::concept::Value;
3439
3440 let converted = match value {
3441 Value::Boolean(value) => serde_json::Value::Bool(*value),
3442 Value::Integer(value) => serde_json::Value::from(*value),
3443 Value::Double(value) => serde_json::Value::from(*value),
3444 Value::String(value) => serde_json::Value::String(value.clone()),
3445 Value::Decimal(decimal) => {
3446 let rendered = match scalar_encoding {
3447 ScalarJsonEncoding::DriverDisplay => decimal.to_string(),
3448 ScalarJsonEncoding::CanonicalV2 => {
3449 canonical_decimal_text(decimal.to_string())?
3450 }
3451 };
3452 serde_json::Value::String(rendered)
3453 }
3454 Value::Date(_) => serde_json::Value::String(value.to_string()),
3455 Value::Datetime(datetime) => {
3456 let rendered = match scalar_encoding {
3457 ScalarJsonEncoding::DriverDisplay => value.to_string(),
3458 ScalarJsonEncoding::CanonicalV2 => exact_naive_datetime(datetime),
3459 };
3460 serde_json::Value::String(rendered)
3461 }
3462 Value::Duration(duration) => {
3463 let rendered = match scalar_encoding {
3464 ScalarJsonEncoding::DriverDisplay => value.to_string(),
3465 ScalarJsonEncoding::CanonicalV2 => {
3466 exact_duration(duration.months, duration.days, duration.nanos)
3467 }
3468 };
3469 serde_json::Value::String(rendered)
3470 }
3471 Value::DatetimeTZ(datetime_tz) => {
3472 let (local, offset_seconds) = checked_datetime_tz_local(datetime_tz)?;
3473 let rendered = match scalar_encoding {
3474 ScalarJsonEncoding::DriverDisplay => value.to_string(),
3475 ScalarJsonEncoding::CanonicalV2 => {
3476 let mut rendered = exact_naive_datetime(&local);
3477 append_exact_offset(&mut rendered, i64::from(offset_seconds));
3478 if let $driver::concept::value::TimeZone::IANA(timezone) =
3479 datetime_tz.timezone()
3480 {
3481 use std::fmt::Write as _;
3482 write!(rendered, "[{}]", timezone.name())
3483 .expect("writing a timezone name to String cannot fail");
3484 }
3485 rendered
3486 }
3487 };
3488 serde_json::Value::String(rendered)
3489 }
3490 Value::Struct(value, name) => {
3491 let fields = value
3492 .fields()
3493 .iter()
3494 .map(|(field, value)| -> Result<_> {
3495 Ok((
3496 field.clone(),
3497 value
3498 .as_ref()
3499 .map(|value| $value_fn(value, scalar_encoding))
3500 .transpose()?
3501 .unwrap_or(serde_json::Value::Null),
3502 ))
3503 })
3504 .collect::<Result<_>>()?;
3505 serde_json::Value::Object(serde_json::Map::from_iter([(
3506 name.clone(),
3507 serde_json::Value::Object(fields),
3508 )]))
3509 }
3510 };
3511 Ok(converted)
3512 }
3513 };
3514}
3515
3516define_driver_json_conversion!(
3517 "band8",
3518 driver_b8,
3519 document_to_json_b8,
3520 document_node_to_json_b8,
3521 document_leaf_to_json_b8,
3522 value_to_json_b8
3523);
3524define_driver_json_conversion!(
3525 "band9",
3526 driver_b9,
3527 document_to_json_b9,
3528 document_node_to_json_b9,
3529 document_leaf_to_json_b9,
3530 value_to_json_b9
3531);
3532
3533#[cfg(feature = "band8")]
3535fn concept_to_json_b8(
3536 concept: &type_bridge_typedb_driver_b8::concept::Concept,
3537 scalar_encoding: ScalarJsonEncoding,
3538) -> Result<serde_json::Value> {
3539 let mut obj = serde_json::Map::new();
3540 obj.insert(
3541 "category".into(),
3542 serde_json::Value::String(concept.get_category().name().into()),
3543 );
3544 obj.insert(
3545 "label".into(),
3546 serde_json::Value::String(concept.get_label().into()),
3547 );
3548 if let Some(iid) = concept.try_get_iid() {
3549 obj.insert("iid".into(), serde_json::Value::String(iid.to_string()));
3550 }
3551 if let Some(value) = concept.try_get_value() {
3552 obj.insert("value".into(), value_to_json_b8(value, scalar_encoding)?);
3553 }
3554 if let Some(vt) = concept.try_get_value_type() {
3555 obj.insert(
3556 "value_type".into(),
3557 serde_json::Value::String(vt.name().into()),
3558 );
3559 }
3560 Ok(serde_json::Value::Object(obj))
3561}
3562
3563#[cfg(feature = "band9")]
3568fn given_rows_b9(spec: GivenRowsSpec) -> Result<typedb_driver::given::GivenRows> {
3569 use typedb_driver::given::GivenRows;
3570
3571 let mut given = GivenRows::new(spec.variables, spec.rows.len());
3572 for row in spec.rows {
3573 let entries = row
3574 .into_iter()
3575 .map(given_entry_b9)
3576 .collect::<Result<Vec<_>>>()?;
3577 given
3578 .push_row(entries)
3579 .map_err(|e| RuntimeError::QueryExecution(format!("Invalid given row: {e}")))?;
3580 }
3581 Ok(given)
3582}
3583
3584#[cfg(feature = "band9")]
3586fn given_entry_b9(value: GivenValue) -> Result<typedb_driver::given::GivenRowEntry> {
3587 use chrono::{NaiveDate, NaiveDateTime};
3588 use typedb_driver::concept::value::TimeZone as B9TimeZone;
3589 use typedb_driver::concept::value::{Decimal, Duration};
3590 use typedb_driver::given::GivenRowEntry;
3591
3592 Ok(match value {
3593 GivenValue::Empty => GivenRowEntry::Empty,
3594 GivenValue::Boolean(b) => GivenRowEntry::from(b),
3595 GivenValue::Integer(i) => GivenRowEntry::from(i),
3596 GivenValue::Double(d) => GivenRowEntry::from(d),
3597 GivenValue::String(s) => GivenRowEntry::from(s),
3598 GivenValue::Date(s) => {
3599 let date = s.parse::<NaiveDate>().map_err(|e| {
3600 RuntimeError::QueryExecution(format!("Invalid given date {s:?}: {e}"))
3601 })?;
3602 GivenRowEntry::from(date)
3603 }
3604 GivenValue::Datetime(s) => {
3605 let dt = s.parse::<NaiveDateTime>().map_err(|e| {
3606 RuntimeError::QueryExecution(format!("Invalid given datetime {s:?}: {e}"))
3607 })?;
3608 GivenRowEntry::from(dt)
3609 }
3610 GivenValue::DatetimeTz(s) => {
3611 let dt = parse_given_datetime_tz(&s).map_err(|e| {
3612 RuntimeError::QueryExecution(format!("Invalid given datetime-tz {s:?}: {e}"))
3613 })?;
3614 let offset = *dt.offset();
3615 GivenRowEntry::from(dt.with_timezone(&B9TimeZone::Fixed(offset)))
3616 }
3617 GivenValue::DatetimeTzExact {
3618 local,
3619 named_zone,
3620 effective_offset_seconds,
3621 } => GivenRowEntry::from(exact_given_datetime_tz_b9(
3622 &local,
3623 named_zone.as_deref(),
3624 effective_offset_seconds,
3625 )?),
3626 GivenValue::Decimal(s) => {
3627 let decimal = if s == "-9223372036854775808" {
3632 Decimal::MIN
3633 } else {
3634 s.parse::<Decimal>().map_err(|e| {
3635 RuntimeError::QueryExecution(format!("Invalid given decimal {s:?}: {e}"))
3636 })?
3637 };
3638 GivenRowEntry::from(decimal)
3639 }
3640 GivenValue::Duration {
3641 months,
3642 days,
3643 nanos,
3644 } => GivenRowEntry::from(Duration::new(months, days, nanos)),
3645 })
3646}
3647
3648#[cfg(feature = "band9")]
3649fn exact_given_datetime_tz_b9(
3650 local: &str,
3651 named_zone: Option<&str>,
3652 effective_offset_seconds: i32,
3653) -> Result<chrono::DateTime<typedb_driver::concept::value::TimeZone>> {
3654 use chrono::{FixedOffset, LocalResult, NaiveDateTime, TimeZone as _};
3655 use typedb_driver::concept::value::TimeZone as B9TimeZone;
3656
3657 let local = local.parse::<NaiveDateTime>().map_err(|error| {
3658 RuntimeError::QueryExecution(format!("Invalid given exact datetime-tz local: {error}"))
3659 })?;
3660 let matches_offset = |value: &chrono::DateTime<B9TimeZone>| {
3661 value
3662 .naive_local()
3663 .signed_duration_since(value.naive_utc())
3664 .num_seconds()
3665 == i64::from(effective_offset_seconds)
3666 };
3667 if let Some(name) = named_zone {
3668 let timezone = chrono_tz::Tz::from_str_insensitive(name)
3669 .map(B9TimeZone::IANA)
3670 .map_err(|_| {
3671 RuntimeError::QueryExecution(
3672 "Invalid given exact datetime-tz named zone".to_owned(),
3673 )
3674 })?;
3675 let selected = match timezone.from_local_datetime(&local) {
3676 LocalResult::Single(value) if matches_offset(&value) => Some(value),
3677 LocalResult::Ambiguous(earlier, later) => [earlier, later]
3678 .into_iter()
3679 .find(|value| matches_offset(value)),
3680 LocalResult::Single(_) | LocalResult::None => None,
3681 };
3682 return selected.ok_or_else(|| {
3683 RuntimeError::QueryExecution(
3684 "Invalid given exact datetime-tz named-zone resolution".to_owned(),
3685 )
3686 });
3687 }
3688
3689 let offset = FixedOffset::east_opt(effective_offset_seconds).ok_or_else(|| {
3690 RuntimeError::QueryExecution("Invalid given exact datetime-tz fixed offset".to_owned())
3691 })?;
3692 B9TimeZone::Fixed(offset)
3693 .from_local_datetime(&local)
3694 .single()
3695 .ok_or_else(|| {
3696 RuntimeError::QueryExecution(
3697 "Invalid given exact datetime-tz fixed resolution".to_owned(),
3698 )
3699 })
3700}
3701
3702#[cfg(feature = "band9")]
3708fn parse_given_datetime_tz(
3709 value: &str,
3710) -> std::result::Result<chrono::DateTime<chrono::FixedOffset>, &'static str> {
3711 use chrono::{FixedOffset, NaiveDateTime, TimeZone as _};
3712
3713 let (local, offset_seconds) = if let Some(local) = value.strip_suffix('Z') {
3714 (local, 0)
3715 } else {
3716 let (local, offset) = [9_usize, 6]
3717 .into_iter()
3718 .find_map(|width| {
3719 let split = value.len().checked_sub(width)?;
3720 if !value.is_char_boundary(split) {
3721 return None;
3722 }
3723 let (local, offset) = (&value[..split], &value[split..]);
3724 parse_fixed_offset(offset).map(|seconds| (local, seconds))
3725 })
3726 .ok_or("expected a Z or signed fixed offset")?;
3727 (local, offset)
3728 };
3729 let local = local
3730 .parse::<NaiveDateTime>()
3731 .map_err(|_| "invalid local datetime")?;
3732 let offset = FixedOffset::east_opt(offset_seconds).ok_or("fixed offset is out of range")?;
3733 offset
3734 .from_local_datetime(&local)
3735 .single()
3736 .ok_or("local datetime is out of range")
3737}
3738
3739#[cfg(feature = "band9")]
3740fn parse_fixed_offset(value: &str) -> Option<i32> {
3741 let bytes = value.as_bytes();
3742 if !matches!(
3743 bytes,
3744 [b'+' | b'-', _, _, b':', _, _] | [b'+' | b'-', _, _, b':', _, _, b':', _, _]
3745 ) || !bytes
3746 .iter()
3747 .enumerate()
3748 .filter(|(index, _)| !matches!(index, 0 | 3 | 6))
3749 .all(|(_, byte)| byte.is_ascii_digit())
3750 {
3751 return None;
3752 }
3753 let component = |left: usize, right: usize| {
3754 std::str::from_utf8(&bytes[left..right])
3755 .ok()?
3756 .parse::<i32>()
3757 .ok()
3758 };
3759 let hours = component(1, 3)?;
3760 let minutes = component(4, 6)?;
3761 let seconds = if bytes.len() == 9 {
3762 component(7, 9)?
3763 } else {
3764 0
3765 };
3766 if hours > 23 || minutes > 59 || seconds > 59 {
3767 return None;
3768 }
3769 let magnitude = hours * 3_600 + minutes * 60 + seconds;
3770 Some(if bytes[0] == b'-' {
3771 -magnitude
3772 } else {
3773 magnitude
3774 })
3775}
3776
3777#[cfg(feature = "band9")]
3779async fn consume_answer_b9(
3780 answer: B9QueryAnswer,
3781 limits: &RuntimeAnswerLimits,
3782 max_collection_members: u64,
3783 scalar_encoding: ScalarJsonEncoding,
3784 consumer: &mut (dyn FnMut(RuntimeAnswerItem) -> Result<RuntimeAnswerControl> + Send),
3785) -> Result<RuntimeAnswerStats> {
3786 match answer {
3787 B9QueryAnswer::Ok(_) => Ok(RuntimeAnswerStats::new(RuntimeAnswerKind::Ok)),
3788 B9QueryAnswer::ConceptRowStream(_, stream) => {
3789 let stream = stream
3790 .map_err(|error| RuntimeError::QueryExecution(format!("Row stream: {error}")))
3791 .and_then(move |row| {
3792 let result = (|| -> Result<RuntimeAnswerItem> {
3793 let mut obj = serde_json::Map::new();
3794 for (i, col) in row.get_column_names().iter().enumerate() {
3795 let value = row
3796 .row
3797 .get(i)
3798 .and_then(|c| c.as_ref())
3799 .map(|concept| concept_to_json_b9(concept, scalar_encoding))
3800 .transpose()?
3801 .unwrap_or(serde_json::Value::Null);
3802 obj.insert(col.clone(), value);
3803 }
3804 Ok(RuntimeAnswerItem::Row(serde_json::Value::Object(obj)))
3805 })();
3806 futures::future::ready(result)
3807 });
3808 runtime_consume_stream(stream, RuntimeAnswerKind::Rows, limits, consumer).await
3809 }
3810 B9QueryAnswer::ConceptDocumentStream(_, stream) => {
3811 let mut collection_members = 0_u64;
3812 let stream = stream
3813 .map_err(|error| RuntimeError::QueryExecution(format!("Document stream: {error}")))
3814 .and_then(move |document| {
3815 let result = document_to_json_b9(
3816 &document,
3817 &mut collection_members,
3818 max_collection_members,
3819 scalar_encoding,
3820 )
3821 .map(RuntimeAnswerItem::Document);
3822 futures::future::ready(result)
3823 });
3824 runtime_consume_stream(stream, RuntimeAnswerKind::Documents, limits, consumer).await
3825 }
3826 }
3827}
3828
3829#[cfg(feature = "band9")]
3833fn concept_to_json_b9(
3834 concept: &typedb_driver::concept::Concept,
3835 scalar_encoding: ScalarJsonEncoding,
3836) -> Result<serde_json::Value> {
3837 let mut obj = serde_json::Map::new();
3838 obj.insert(
3839 "category".into(),
3840 serde_json::Value::String(concept.get_category().name().into()),
3841 );
3842 obj.insert(
3843 "label".into(),
3844 serde_json::Value::String(concept.get_label().into()),
3845 );
3846 if let Some(iid) = concept.try_get_iid() {
3847 obj.insert("iid".into(), serde_json::Value::String(iid.to_string()));
3848 }
3849 if let Some(value) = concept.try_get_value() {
3850 obj.insert("value".into(), value_to_json_b9(value, scalar_encoding)?);
3851 }
3852 if let Some(vt) = concept.try_get_value_type() {
3853 obj.insert(
3854 "value_type".into(),
3855 serde_json::Value::String(vt.name().into()),
3856 );
3857 }
3858 Ok(serde_json::Value::Object(obj))
3859}
3860
3861#[cfg(test)]
3862mod tests {
3863 use super::*;
3864 use std::io;
3865 use std::sync::atomic::{AtomicUsize, Ordering};
3866 use std::sync::{Arc, Mutex};
3867
3868 static NEXT_TLS_MATERIAL_TEST_ID: AtomicUsize = AtomicUsize::new(0);
3869 static TRACE_CAPTURE_LOCK: Mutex<()> = Mutex::new(());
3870
3871 #[derive(Clone, Default)]
3872 struct TraceCapture {
3873 bytes: Arc<Mutex<Vec<u8>>>,
3874 }
3875
3876 struct TraceCaptureWriter {
3877 bytes: Arc<Mutex<Vec<u8>>>,
3878 }
3879
3880 impl io::Write for TraceCaptureWriter {
3881 fn write(&mut self, buffer: &[u8]) -> io::Result<usize> {
3882 self.bytes
3883 .lock()
3884 .expect("trace capture lock")
3885 .extend_from_slice(buffer);
3886 Ok(buffer.len())
3887 }
3888
3889 fn flush(&mut self) -> io::Result<()> {
3890 Ok(())
3891 }
3892 }
3893
3894 impl<'writer> tracing_subscriber::fmt::MakeWriter<'writer> for TraceCapture {
3895 type Writer = TraceCaptureWriter;
3896
3897 fn make_writer(&'writer self) -> Self::Writer {
3898 TraceCaptureWriter {
3899 bytes: Arc::clone(&self.bytes),
3900 }
3901 }
3902 }
3903
3904 fn capture_traces(emit: impl FnOnce()) -> String {
3905 let _capture_guard = TRACE_CAPTURE_LOCK.lock().expect("trace capture suite lock");
3906 let capture = TraceCapture::default();
3907 let subscriber = tracing_subscriber::fmt()
3908 .without_time()
3909 .with_ansi(false)
3910 .with_target(false)
3911 .with_max_level(tracing::Level::TRACE)
3912 .with_writer(capture.clone())
3913 .finish();
3914 tracing::subscriber::with_default(subscriber, emit);
3915 let bytes = capture.bytes.lock().expect("trace capture lock").clone();
3916 String::from_utf8(bytes).expect("tracing formatter emits UTF-8")
3917 }
3918
3919 #[test]
3920 fn driver_drop_close_warning_drops_hostile_provider_text() {
3921 const PROVIDER_SENTINEL: &str = "TB_DROP_CLOSE_PROVIDER_SECRET";
3922 let error = RuntimeError::Connection(PROVIDER_SENTINEL.to_owned());
3923 let output = capture_traces(|| trace_driver_drop_close_failure(&error));
3924
3925 assert!(
3926 output.contains(TRACE_CODE_DRIVER_DROP_CLOSE_FAILED),
3927 "{output}"
3928 );
3929 assert!(
3930 output.contains("typedb_runtime_connection_failed"),
3931 "{output}"
3932 );
3933 assert!(!output.contains(PROVIDER_SENTINEL), "{output}");
3934 }
3935
3936 #[cfg(all(feature = "band8", feature = "band9"))]
3937 #[test]
3938 fn band9_upgrade_failure_warning_drops_hostile_provider_and_address_text() {
3939 const ADDRESS_SENTINEL: &str = "TB_BAND9_UPGRADE_ADDRESS_SECRET";
3940 const PROVIDER_SENTINEL: &str = "TB_BAND9_UPGRADE_PROVIDER_SECRET";
3941 let error =
3942 SecureConnectError::Runtime(RuntimeError::Connection(PROVIDER_SENTINEL.to_owned()));
3943 let output = capture_traces(|| {
3944 trace_band9_upgrade_failed(
3945 ADDRESS_SENTINEL,
3946 core_version::Version::new(3, 12, 1),
3947 &error,
3948 );
3949 });
3950
3951 assert!(output.contains(TRACE_CODE_BAND9_UPGRADE_FAILED), "{output}");
3952 assert!(
3953 output.contains("typedb_runtime_connection_failed"),
3954 "{output}"
3955 );
3956 assert!(output.contains("3.12.1"), "{output}");
3957 assert!(!output.contains(ADDRESS_SENTINEL), "{output}");
3958 assert!(!output.contains(PROVIDER_SENTINEL), "{output}");
3959 }
3960
3961 #[test]
3962 fn post_credential_connection_traces_drop_raw_address_identity() {
3963 const ADDRESS_SENTINEL: &str = "TB_POST_CREDENTIAL_ADDRESS_SECRET";
3964 let server_version = core_version::Version::new(3, 12, 1);
3965 let output = capture_traces(|| {
3966 trace_version_gate_passed(ADDRESS_SENTINEL, 9, server_version);
3967 #[cfg(all(feature = "band8", feature = "band9"))]
3968 trace_band9_upgrade_succeeded(ADDRESS_SENTINEL, server_version);
3969 #[cfg(feature = "band8")]
3970 trace_band8_fallback_connected(ADDRESS_SENTINEL, server_version);
3971 trace_runtime_connected(ADDRESS_SENTINEL, 9, Some(server_version));
3972 });
3973
3974 assert!(!output.contains(ADDRESS_SENTINEL), "{output}");
3975 assert!(output.contains(TRACE_CODE_VERSION_GATE_PASSED), "{output}");
3976 assert!(output.contains(TRACE_CODE_CONNECTED), "{output}");
3977 assert!(output.contains("3.12.1"), "{output}");
3978 #[cfg(all(feature = "band8", feature = "band9"))]
3979 assert!(
3980 output.contains(TRACE_CODE_BAND9_UPGRADE_SUCCEEDED),
3981 "{output}"
3982 );
3983 #[cfg(feature = "band8")]
3984 assert!(
3985 output.contains(TRACE_CODE_BAND8_FALLBACK_CONNECTED),
3986 "{output}"
3987 );
3988 }
3989
3990 #[test]
3991 fn retained_server_trace_is_warning_free() {
3992 for (driver_band, server_version) in [
3993 (8, Some(core_version::Version::new(3, 11, 5))),
3994 (9, Some(core_version::Version::new(3, 12, 1))),
3995 (8, Some(core_version::Version::new(3, 12, 1))),
3999 (8, None),
4000 (9, None),
4001 ] {
4002 let output =
4003 capture_traces(|| trace_runtime_connected("redacted", driver_band, server_version));
4004 assert!(output.contains(TRACE_CODE_CONNECTED), "{output}");
4005 assert!(!output.contains("WARN"), "{output}");
4006 assert!(!output.contains("deprecated"), "{output}");
4007 }
4008 }
4009
4010 #[test]
4011 fn credential_safe_secure_diagnostics_drop_provider_controlled_text() {
4012 const SENTINEL: &str = "TB_POST_CREDENTIAL_PROVIDER_SECRET";
4013 for error in [
4014 SecureConnectError::Runtime(RuntimeError::Connection(SENTINEL.to_owned())),
4015 SecureConnectError::Runtime(RuntimeError::UnsupportedVersion(
4016 core_version::VersionError::Probe(SENTINEL.to_owned()),
4017 )),
4018 SecureConnectError::Runtime(RuntimeError::UnsupportedVersion(
4019 core_version::VersionError::Parse(SENTINEL.to_owned()),
4020 )),
4021 ] {
4022 assert_eq!(error.credential_safe_diagnostic(), None);
4023 }
4024
4025 let tls = SecureConnectError::TlsConfiguration(
4026 core_version::TlsConfigurationError::CustomRootCaUnreadable {
4027 path: std::path::PathBuf::from(SENTINEL),
4028 },
4029 )
4030 .credential_safe_diagnostic()
4031 .expect("typed TLS codes remain safe");
4032 assert!(tls.contains("tls_custom_root_ca_unreadable"), "{tls}");
4033 assert!(!tls.contains(SENTINEL), "{tls}");
4034 }
4035
4036 #[test]
4037 fn credential_safe_secure_diagnostics_preserve_closed_version_data() {
4038 let error = SecureConnectError::Runtime(RuntimeError::UnsupportedVersion(
4039 core_version::VersionError::Unsupported {
4040 component: "server",
4041 found: core_version::Version::new(3, 13, 0),
4042 },
4043 ));
4044 let diagnostic = error
4045 .credential_safe_diagnostic()
4046 .expect("closed version diagnostics remain actionable");
4047 assert!(diagnostic.contains("server version 3.13.0"), "{diagnostic}");
4048 assert!(diagnostic.contains("3.11.0–3.12.x"), "{diagnostic}");
4049 }
4050
4051 #[test]
4052 fn driver_shutdown_panics_are_contained_as_connection_errors() {
4053 let error = contain_driver_shutdown(|| panic!("simulated upstream shutdown panic"))
4054 .expect_err("shutdown panic must not cross the runtime boundary");
4055
4056 assert!(matches!(
4057 error,
4058 RuntimeError::Connection(message)
4059 if message == "Driver close failed: upstream driver panicked during shutdown"
4060 ));
4061 }
4062
4063 #[test]
4064 fn failed_driver_shutdown_is_retried_without_reopening_admission() {
4065 let state = AtomicU8::new(DRIVER_OPEN);
4066 let close_lock = Mutex::new(());
4067
4068 for _ in 0..2 {
4069 let error = run_driver_shutdown(&state, &close_lock, || {
4070 Err(RuntimeError::Connection(
4071 "Driver close failed: simulated incomplete cleanup".to_owned(),
4072 ))
4073 })
4074 .expect_err("incomplete cleanup must remain observable");
4075 assert_eq!(
4076 error.to_string(),
4077 "Connection error: Driver close failed: simulated incomplete cleanup"
4078 );
4079 assert_eq!(state.load(AtomicOrdering::Acquire), DRIVER_CLOSING);
4080 }
4081
4082 run_driver_shutdown(&state, &close_lock, || Ok(()))
4083 .expect("a later close retries and completes cleanup");
4084 assert_eq!(state.load(AtomicOrdering::Acquire), DRIVER_CLOSED);
4085
4086 run_driver_shutdown(&state, &close_lock, || {
4087 panic!("successful cleanup must make later close calls no-ops")
4088 })
4089 .expect("close remains idempotent after cleanup succeeds");
4090 }
4091
4092 #[test]
4093 fn transaction_retention_and_shutdown_share_one_admission_lock() {
4094 let state = AtomicU8::new(DRIVER_OPEN);
4095 let active_transactions = AtomicUsize::new(0);
4096 let close_lock = Mutex::new(());
4097
4098 retain_driver_transaction(&state, &active_transactions, &close_lock)
4099 .expect("an open driver admits a transaction lease");
4100 assert_eq!(active_transactions.load(AtomicOrdering::Acquire), 1);
4101
4102 let _close_guard = close_lock.lock().expect("admission lock remains usable");
4103 state.store(DRIVER_CLOSING, AtomicOrdering::Release);
4104 drop(_close_guard);
4105
4106 let error = retain_driver_transaction(&state, &active_transactions, &close_lock)
4107 .expect_err("shutdown must make later transaction retention terminal");
4108 assert!(matches!(error, RuntimeError::Connection(_)));
4109 assert_eq!(active_transactions.load(AtomicOrdering::Acquire), 1);
4110 }
4111
4112 #[test]
4113 fn one_shot_database_operation_propagates_close_failure_after_success() {
4114 assert_eq!(
4115 finish_one_shot_database_operation(Ok(17_u8), Ok(())).unwrap(),
4116 17
4117 );
4118
4119 let error = finish_one_shot_database_operation(
4120 Ok(17_u8),
4121 Err(RuntimeError::Connection("close diagnosis".to_owned())),
4122 )
4123 .expect_err("a successful operation must still report driver-close failure");
4124 assert!(matches!(
4125 error,
4126 RuntimeError::Connection(message) if message == "close diagnosis"
4127 ));
4128 }
4129
4130 #[test]
4131 fn one_shot_database_operation_preserves_primary_failure_and_combines_close_failure() {
4132 let primary_only = finish_one_shot_database_operation::<()>(
4133 Err(RuntimeError::Connection("primary diagnosis".to_owned())),
4134 Ok(()),
4135 )
4136 .expect_err("the primary database-operation failure must be returned");
4137 assert!(matches!(
4138 primary_only,
4139 RuntimeError::Connection(message) if message == "primary diagnosis"
4140 ));
4141
4142 let combined = finish_one_shot_database_operation::<()>(
4143 Err(RuntimeError::Connection("primary diagnosis".to_owned())),
4144 Err(RuntimeError::Connection("close diagnosis".to_owned())),
4145 )
4146 .expect_err("both failures must produce one deterministic diagnostic");
4147 assert!(matches!(
4148 combined,
4149 RuntimeError::Connection(message)
4150 if message
4151 == "primary diagnosis; additionally, TypeDB driver cleanup failed: close diagnosis"
4152 ));
4153 }
4154
4155 macro_rules! row_and_document_scalar_regression {
4156 ($feature:literal, $name:ident, $driver:ident, $concept_fn:ident, $node_fn:ident, $decimal:expr) => {
4157 #[cfg(feature = $feature)]
4158 #[test]
4159 fn $name() {
4160 use $driver::answer::concept_document::{Leaf, Node};
4161 use $driver::concept::{Concept, Value};
4162
4163 const LARGE_INTEGER: i64 = 9_007_199_254_740_993;
4164 let integer = Node::Leaf(Some(Leaf::Concept(Concept::Value(Value::Integer(
4165 LARGE_INTEGER,
4166 )))));
4167 let mut collection_members = 0;
4168 assert_eq!(
4169 $node_fn(
4170 &integer,
4171 &mut collection_members,
4172 u64::MAX,
4173 ScalarJsonEncoding::CanonicalV2,
4174 )
4175 .unwrap(),
4176 serde_json::Value::from(LARGE_INTEGER),
4177 "concept-document integers must not cross an f64 boundary"
4178 );
4179
4180 let decimal = $decimal;
4181 let scalar_text = |value: serde_json::Value| {
4182 value
4183 .get("value")
4184 .and_then(serde_json::Value::as_str)
4185 .or_else(|| value.as_str())
4186 .map(str::to_owned)
4187 .expect("decimal conversion must expose scalar text")
4188 };
4189 let legacy_expected = decimal.to_string();
4190 let expected = "1234.56";
4191 let decimal =
4192 Node::Leaf(Some(Leaf::Concept(Concept::Value(Value::Decimal(decimal)))));
4193 let row_decimal = Concept::Value(Value::Decimal($decimal));
4194 assert_eq!(
4195 scalar_text(
4196 $node_fn(
4197 &decimal,
4198 &mut collection_members,
4199 u64::MAX,
4200 ScalarJsonEncoding::CanonicalV2,
4201 )
4202 .unwrap()
4203 ),
4204 expected,
4205 "concept-document decimals must use canonical strings"
4206 );
4207 assert_eq!(
4208 scalar_text(
4209 $concept_fn(&row_decimal, ScalarJsonEncoding::CanonicalV2).unwrap()
4210 ),
4211 expected
4212 );
4213 assert_eq!(
4214 scalar_text(
4215 $concept_fn(&row_decimal, ScalarJsonEncoding::DriverDisplay).unwrap()
4216 ),
4217 legacy_expected.clone()
4218 );
4219 let legacy_decimal = Node::Leaf(Some(Leaf::Concept(Concept::Value(
4220 Value::Decimal($decimal),
4221 ))));
4222 assert_eq!(
4223 scalar_text(
4224 $node_fn(
4225 &legacy_decimal,
4226 &mut collection_members,
4227 u64::MAX,
4228 ScalarJsonEncoding::DriverDisplay,
4229 )
4230 .unwrap()
4231 ),
4232 legacy_expected,
4233 "legacy document decimals must retain driver spelling"
4234 );
4235 }
4236 };
4237 }
4238
4239 row_and_document_scalar_regression!(
4240 "band8",
4241 band8_rows_and_documents_preserve_lossless_scalars,
4242 driver_b8,
4243 concept_to_json_b8,
4244 document_node_to_json_b8,
4245 driver_b8::concept::value::Decimal::new(1234, 5_600_000_000_000_000_000)
4246 );
4247 row_and_document_scalar_regression!(
4248 "band9",
4249 band9_rows_and_documents_preserve_lossless_scalars,
4250 driver_b9,
4251 concept_to_json_b9,
4252 document_node_to_json_b9,
4253 driver_b9::concept::value::Decimal::from_parts(1234, 5_600_000_000_000_000_000)
4254 );
4255
4256 macro_rules! datetime_tz_evidence_regression {
4257 ($feature:literal, $name:ident, $driver:ident, $concept_fn:ident, $leaf_fn:ident) => {
4258 #[cfg(feature = $feature)]
4259 #[test]
4260 fn $name() {
4261 use chrono::{NaiveDateTime, TimeZone as _};
4262 use $driver::answer::concept_document::Leaf;
4263 use $driver::concept::value::TimeZone;
4264 use $driver::concept::{Concept, Value};
4265
4266 let london = TimeZone::IANA(
4267 "Europe/London"
4268 .parse()
4269 .expect("driver timezone database contains London"),
4270 );
4271 let overlap_utc = "2024-10-27T00:30:00"
4272 .parse::<NaiveDateTime>()
4273 .expect("UTC datetime");
4274 let named = Value::DatetimeTZ(london.from_utc_datetime(&overlap_utc));
4275 assert_eq!(
4276 $concept_fn(
4277 &Concept::Value(named.clone()),
4278 ScalarJsonEncoding::CanonicalV2,
4279 )
4280 .expect("row concept")["value"],
4281 serde_json::json!("2024-10-27T01:30:00+01:00[Europe/London]"),
4282 "row evidence preserves both the IANA identity and overlap side",
4283 );
4284 assert_eq!(
4285 $leaf_fn(
4286 &Leaf::Concept(Concept::Value(named.clone())),
4287 ScalarJsonEncoding::CanonicalV2,
4288 )
4289 .expect("document leaf"),
4290 serde_json::json!("2024-10-27T01:30:00+01:00[Europe/London]"),
4291 "document evidence uses the same exact datetime-tz bridge",
4292 );
4293 assert_eq!(
4294 $concept_fn(&Concept::Value(named), ScalarJsonEncoding::DriverDisplay,)
4295 .expect("row concept")["value"],
4296 serde_json::json!("2024-10-27T01:30:00.000000000 Europe/London"),
4297 "released conversion retains the upstream driver display",
4298 );
4299
4300 let fixed = TimeZone::Fixed(
4301 chrono::FixedOffset::east_opt(1_172).expect("second-resolution offset"),
4302 );
4303 let fixed_utc = "1900-01-01T11:40:28"
4304 .parse::<NaiveDateTime>()
4305 .expect("UTC datetime");
4306 let fixed = Value::DatetimeTZ(fixed.from_utc_datetime(&fixed_utc));
4307 assert_eq!(
4308 $concept_fn(
4309 &Concept::Value(fixed.clone()),
4310 ScalarJsonEncoding::CanonicalV2,
4311 )
4312 .expect("row concept")["value"],
4313 serde_json::json!("1900-01-01T12:00:00+00:19:32"),
4314 "row evidence preserves fixed offset seconds",
4315 );
4316 assert_eq!(
4317 $leaf_fn(
4318 &Leaf::Concept(Concept::Value(fixed)),
4319 ScalarJsonEncoding::CanonicalV2,
4320 )
4321 .expect("document leaf"),
4322 serde_json::json!("1900-01-01T12:00:00+00:19:32"),
4323 "document evidence preserves fixed offset seconds",
4324 );
4325 }
4326 };
4327 }
4328
4329 datetime_tz_evidence_regression!(
4330 "band8",
4331 band8_rows_and_documents_preserve_exact_datetime_tz_evidence,
4332 driver_b8,
4333 concept_to_json_b8,
4334 document_leaf_to_json_b8
4335 );
4336 datetime_tz_evidence_regression!(
4337 "band9",
4338 band9_rows_and_documents_preserve_exact_datetime_tz_evidence,
4339 driver_b9,
4340 concept_to_json_b9,
4341 document_leaf_to_json_b9
4342 );
4343
4344 macro_rules! datetime_tz_local_range_regression {
4345 ($feature:literal, $name:ident, $driver:ident, $concept_fn:ident, $leaf_fn:ident) => {
4346 #[cfg(feature = $feature)]
4347 #[test]
4348 fn $name() {
4349 use chrono::{FixedOffset, NaiveDateTime, TimeZone as _};
4350 use $driver::answer::concept_document::Leaf;
4351 use $driver::concept::value::TimeZone;
4352 use $driver::concept::{Concept, Value};
4353
4354 let assert_range_error = |error: RuntimeError| {
4355 assert!(matches!(
4356 error,
4357 RuntimeError::QueryExecution(message)
4358 if message
4359 == "provider datetime-tz local value is outside the supported range"
4360 ));
4361 };
4362 let invalid = [
4363 (
4364 FixedOffset::east_opt(1).expect("positive offset"),
4365 NaiveDateTime::MAX,
4366 ),
4367 (
4368 FixedOffset::west_opt(1).expect("negative offset"),
4369 NaiveDateTime::MIN,
4370 ),
4371 ];
4372 for (offset, utc) in invalid {
4373 let timezone = TimeZone::Fixed(offset);
4374 let value = Value::DatetimeTZ(timezone.from_utc_datetime(&utc));
4375 for encoding in [
4376 ScalarJsonEncoding::CanonicalV2,
4377 ScalarJsonEncoding::DriverDisplay,
4378 ] {
4379 let row_error =
4380 $concept_fn(&Concept::Value(value.clone()), encoding).expect_err(
4381 "an unrepresentable provider local datetime must fail row conversion",
4382 );
4383 assert_range_error(row_error);
4384
4385 let document_error =
4386 $leaf_fn(&Leaf::Concept(Concept::Value(value.clone())), encoding)
4387 .expect_err(
4388 "an unrepresentable provider local datetime must fail document conversion",
4389 );
4390 assert_range_error(document_error);
4391 }
4392 }
4393
4394 let valid = [
4395 (
4396 FixedOffset::west_opt(1).expect("negative offset"),
4397 NaiveDateTime::MAX,
4398 ),
4399 (
4400 FixedOffset::east_opt(1).expect("positive offset"),
4401 NaiveDateTime::MIN,
4402 ),
4403 ];
4404 for (offset, utc) in valid {
4405 let timezone = TimeZone::Fixed(offset);
4406 let value = Value::DatetimeTZ(timezone.from_utc_datetime(&utc));
4407 for encoding in [
4408 ScalarJsonEncoding::CanonicalV2,
4409 ScalarJsonEncoding::DriverDisplay,
4410 ] {
4411 let row = $concept_fn(&Concept::Value(value.clone()), encoding)
4412 .expect("an inward offset must remain representable");
4413 assert!(row["value"].is_string());
4414
4415 let document =
4416 $leaf_fn(&Leaf::Concept(Concept::Value(value.clone())), encoding)
4417 .expect("document conversion accepts an inward offset");
4418 assert!(document.is_string());
4419 }
4420 }
4421 }
4422 };
4423 }
4424
4425 datetime_tz_local_range_regression!(
4426 "band8",
4427 band8_datetime_tz_local_range_errors_are_propagated,
4428 driver_b8,
4429 concept_to_json_b8,
4430 document_leaf_to_json_b8
4431 );
4432 datetime_tz_local_range_regression!(
4433 "band9",
4434 band9_datetime_tz_local_range_errors_are_propagated,
4435 driver_b9,
4436 concept_to_json_b9,
4437 document_leaf_to_json_b9
4438 );
4439
4440 macro_rules! exact_temporal_evidence_regression {
4441 ($feature:literal, $name:ident, $driver:ident, $concept_fn:ident, $leaf_fn:ident) => {
4442 #[cfg(feature = $feature)]
4443 #[test]
4444 fn $name() {
4445 use $driver::answer::concept_document::Leaf;
4446 use $driver::concept::value::Duration;
4447 use $driver::concept::{Concept, Value};
4448
4449 let duration_cases = [
4450 (Duration::new(0, 0, 0), "PT0S"),
4451 (Duration::new(12, 0, 0), "P12M"),
4452 (Duration::new(0, 0, 3_660_000_000_000), "PT3660S"),
4453 (Duration::new(0, 0, 1_500_000_000), "PT1.5S"),
4454 (
4455 Duration::new(u32::MAX, u32::MAX, u64::MAX),
4456 "P4294967295M4294967295DT18446744073.709551615S",
4457 ),
4458 ];
4459 for (duration, expected) in duration_cases {
4460 let value = Value::Duration(duration);
4461 assert_eq!(
4462 $concept_fn(
4463 &Concept::Value(value.clone()),
4464 ScalarJsonEncoding::CanonicalV2,
4465 )
4466 .expect("row concept")["value"],
4467 serde_json::Value::String(expected.to_owned()),
4468 "row evidence uses the contract's exact duration components",
4469 );
4470 assert_eq!(
4471 $leaf_fn(
4472 &Leaf::Concept(Concept::Value(value)),
4473 ScalarJsonEncoding::CanonicalV2,
4474 )
4475 .expect("document leaf"),
4476 serde_json::Value::String(expected.to_owned()),
4477 "document evidence uses the same exact duration bridge",
4478 );
4479 }
4480
4481 let year = Value::Duration(Duration::new(12, 0, 0));
4482 assert_eq!(
4483 $concept_fn(&Concept::Value(year), ScalarJsonEncoding::DriverDisplay,)
4484 .expect("row concept")["value"],
4485 serde_json::json!("P1Y"),
4486 "released conversion retains the upstream duration display",
4487 );
4488
4489 let datetime = chrono::NaiveDate::from_ymd_opt(2024, 1, 2)
4490 .expect("date")
4491 .and_hms_nano_opt(3, 4, 5, 500_000_000)
4492 .expect("datetime");
4493 let datetime = Value::Datetime(datetime);
4494 assert_eq!(
4495 $concept_fn(
4496 &Concept::Value(datetime.clone()),
4497 ScalarJsonEncoding::CanonicalV2,
4498 )
4499 .expect("row concept")["value"],
4500 serde_json::json!("2024-01-02T03:04:05.5"),
4501 "row evidence trims only insignificant datetime zeroes",
4502 );
4503 assert_eq!(
4504 $leaf_fn(
4505 &Leaf::Concept(Concept::Value(datetime.clone())),
4506 ScalarJsonEncoding::CanonicalV2,
4507 )
4508 .expect("document leaf"),
4509 serde_json::json!("2024-01-02T03:04:05.5"),
4510 "document evidence uses the same exact datetime bridge",
4511 );
4512 assert_eq!(
4513 $concept_fn(&Concept::Value(datetime), ScalarJsonEncoding::DriverDisplay,)
4514 .expect("row concept")["value"],
4515 serde_json::json!("2024-01-02T03:04:05.500000000"),
4516 "released conversion retains the upstream datetime display",
4517 );
4518 }
4519 };
4520 }
4521
4522 exact_temporal_evidence_regression!(
4523 "band8",
4524 band8_rows_and_documents_preserve_exact_datetime_and_duration_evidence,
4525 driver_b8,
4526 concept_to_json_b8,
4527 document_leaf_to_json_b8
4528 );
4529 exact_temporal_evidence_regression!(
4530 "band9",
4531 band9_rows_and_documents_preserve_exact_datetime_and_duration_evidence,
4532 driver_b9,
4533 concept_to_json_b9,
4534 document_leaf_to_json_b9
4535 );
4536
4537 #[cfg(feature = "band9")]
4538 #[test]
4539 fn document_list_limit_rejects_before_json_collection() {
4540 use driver_b9::answer::concept_document::{Leaf, Node};
4541 let document = Node::List(vec![Node::Leaf(Some(Leaf::Empty)), Node::Leaf(None)]);
4542 let mut members = 0;
4543 let error =
4544 document_node_to_json_b9(&document, &mut members, 1, ScalarJsonEncoding::CanonicalV2)
4545 .expect_err("the list exceeds its conversion budget");
4546 assert!(matches!(
4547 error,
4548 RuntimeError::ResourceLimit {
4549 code: "query_v2_document_member_limit",
4550 ..
4551 }
4552 ));
4553 assert_eq!(members, 0, "rejected lists are not partially charged");
4554 }
4555
4556 #[cfg(feature = "band9")]
4557 #[test]
4558 fn unexpected_document_thing_is_a_typed_error_not_a_panic() {
4559 use driver_b9::answer::concept_document::{Leaf, Node};
4560 use driver_b9::concept::{Concept, Entity};
4561
4562 let document = Node::Leaf(Some(Leaf::Concept(Concept::Entity(Entity {
4563 iid: vec![1_u8].into(),
4564 type_: None,
4565 }))));
4566 let mut members = 0;
4567 let error = document_node_to_json_b9(
4568 &document,
4569 &mut members,
4570 u64::MAX,
4571 ScalarJsonEncoding::CanonicalV2,
4572 )
4573 .expect_err("thing instances are not valid fetch-document leaves");
4574 assert!(matches!(
4575 error,
4576 RuntimeError::QueryExecution(message)
4577 if message == "document response carried an unsupported thing instance"
4578 ));
4579 }
4580
4581 async fn assert_pending_runtime_await_hits_deadline() {
4582 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
4583 let limits = RuntimeAnswerLimits {
4584 max_items: 10,
4585 max_bytes: 4096,
4586 deadline: Some(
4587 (tokio::time::Instant::now() + std::time::Duration::from_secs(5)).into_std(),
4588 ),
4589 cancellation: RuntimeAnswerCancellation::default(),
4590 };
4591 let execution = tokio::spawn(async move {
4592 runtime_await(
4593 async move {
4594 let _ = started_tx.send(());
4595 std::future::pending::<Result<()>>().await
4596 },
4597 &limits,
4598 )
4599 .await
4600 });
4601
4602 started_rx.await.unwrap();
4603 tokio::time::advance(std::time::Duration::from_secs(5)).await;
4604 tokio::task::yield_now().await;
4605 assert!(execution.is_finished());
4606 assert!(matches!(
4607 execution.await.unwrap().unwrap_err(),
4608 RuntimeError::ResourceLimit {
4609 code: "transaction_deadline_exceeded",
4610 ..
4611 }
4612 ));
4613 }
4614
4615 #[tokio::test(start_paused = true)]
4616 async fn runtime_deadline_interrupts_pending_query_await() {
4617 assert_pending_runtime_await_hits_deadline().await;
4618 }
4619
4620 #[tokio::test(start_paused = true)]
4621 async fn runtime_deadline_interrupts_pending_stream_poll() {
4622 assert_pending_runtime_await_hits_deadline().await;
4623 }
4624
4625 #[tokio::test]
4626 async fn runtime_cancellation_after_poll_start_interrupts_pending_await() {
4627 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
4628 let cancellation = RuntimeAnswerCancellation::default();
4629 let trigger = cancellation.clone();
4630 let limits = RuntimeAnswerLimits {
4631 max_items: 10,
4632 max_bytes: 4096,
4633 deadline: None,
4634 cancellation,
4635 };
4636 let execution = tokio::spawn(async move {
4637 runtime_await(
4638 async move {
4639 let _ = started_tx.send(());
4640 std::future::pending::<Result<()>>().await
4641 },
4642 &limits,
4643 )
4644 .await
4645 });
4646
4647 started_rx.await.unwrap();
4648 trigger.cancel();
4649 tokio::task::yield_now().await;
4650 assert!(execution.is_finished());
4651 assert!(matches!(
4652 execution.await.unwrap().unwrap_err(),
4653 RuntimeError::ResourceLimit {
4654 code: "provider_cancelled",
4655 ..
4656 }
4657 ));
4658 }
4659
4660 #[test]
4661 fn bounded_runtime_reader_stops_and_enforces_limits_before_consumer() {
4662 let limits = RuntimeAnswerLimits {
4663 max_items: 1,
4664 max_bytes: 64,
4665 deadline: None,
4666 cancellation: RuntimeAnswerCancellation::default(),
4667 };
4668 let mut stats = RuntimeAnswerStats::new(RuntimeAnswerKind::Rows);
4669 let mut accepted = 0;
4670 let mut stop = |_item| {
4671 accepted += 1;
4672 Ok(RuntimeAnswerControl::Stop)
4673 };
4674 assert_eq!(
4675 runtime_accept(
4676 &limits,
4677 &mut stats,
4678 RuntimeAnswerItem::Row(serde_json::json!({"v": 1})),
4679 &mut stop,
4680 )
4681 .unwrap(),
4682 RuntimeAnswerControl::Stop
4683 );
4684 assert_eq!(accepted, 1);
4685 assert!(stats.stopped_early);
4686
4687 let mut never_called =
4688 |_item| -> Result<RuntimeAnswerControl> { panic!("over-limit item reached consumer") };
4689 let error = runtime_accept(
4690 &limits,
4691 &mut stats,
4692 RuntimeAnswerItem::Row(serde_json::json!({"v": 2})),
4693 &mut never_called,
4694 )
4695 .unwrap_err();
4696 assert!(matches!(
4697 error,
4698 RuntimeError::ResourceLimit {
4699 code: "processed_item_limit",
4700 ..
4701 }
4702 ));
4703 }
4704
4705 #[tokio::test]
4706 async fn bounded_runtime_stream_stops_before_polling_another_item() {
4707 let polls = Arc::new(AtomicUsize::new(0));
4708 let observed_polls = Arc::clone(&polls);
4709 let mut emitted = false;
4710 let stream = futures::stream::poll_fn(move |_context| {
4711 observed_polls.fetch_add(1, Ordering::SeqCst);
4712 assert!(!emitted, "stream was polled after the consumer stopped");
4713 emitted = true;
4714 std::task::Poll::Ready(Some(Ok::<_, RuntimeError>(RuntimeAnswerItem::Row(
4715 serde_json::json!({"v": 1}),
4716 ))))
4717 });
4718 let limits = RuntimeAnswerLimits::unbounded();
4719 let mut consumer = |_item| Ok(RuntimeAnswerControl::Stop);
4720
4721 let stats = runtime_consume_stream(stream, RuntimeAnswerKind::Rows, &limits, &mut consumer)
4722 .await
4723 .unwrap();
4724
4725 assert_eq!(polls.load(Ordering::SeqCst), 1);
4726 assert_eq!(stats.processed_items, 1);
4727 assert!(stats.stopped_early);
4728 }
4729
4730 #[tokio::test]
4731 async fn bounded_runtime_stream_rejects_an_over_limit_item_before_consumer() {
4732 let stream = futures::stream::iter([
4733 Ok(RuntimeAnswerItem::Row(serde_json::json!({"v": 1}))),
4734 Ok(RuntimeAnswerItem::Row(serde_json::json!({"v": 2}))),
4735 ]);
4736 let limits = RuntimeAnswerLimits {
4737 max_items: 1,
4738 max_bytes: u64::MAX,
4739 deadline: None,
4740 cancellation: RuntimeAnswerCancellation::default(),
4741 };
4742 let mut accepted = 0;
4743 let mut consumer = |_item| {
4744 accepted += 1;
4745 Ok(RuntimeAnswerControl::Continue)
4746 };
4747
4748 let error = runtime_consume_stream(stream, RuntimeAnswerKind::Rows, &limits, &mut consumer)
4749 .await
4750 .unwrap_err();
4751
4752 assert_eq!(accepted, 1);
4753 assert!(matches!(
4754 error,
4755 RuntimeError::ResourceLimit {
4756 code: "processed_item_limit",
4757 ..
4758 }
4759 ));
4760 }
4761
4762 #[cfg(feature = "band8")]
4763 fn empty_given_rows() -> GivenRowsSpec {
4764 GivenRowsSpec {
4765 variables: vec!["value".into()],
4766 rows: Vec::new(),
4767 }
4768 }
4769
4770 #[cfg(feature = "band8")]
4771 #[tokio::test]
4772 async fn band8_transaction_rejects_bounded_given_rows_actionably() {
4773 let mut transaction = RuntimeTransaction {
4774 inner: RuntimeTransactionInner::B8(None),
4775 driver_lease: None,
4776 };
4777 let mut consumer = |_item| Ok(RuntimeAnswerControl::Continue);
4778
4779 assert!(!transaction.supports_given_rows());
4780 let error = transaction
4781 .query_with_rows_bounded(
4782 "given $value: string; match $x isa thing;",
4783 empty_given_rows(),
4784 RuntimeAnswerLimits::unbounded(),
4785 &mut consumer,
4786 )
4787 .await
4788 .unwrap_err();
4789
4790 assert!(
4791 matches!(&error, RuntimeError::QueryExecution(message) if message.contains("band-9 driver") && message.contains("band 8")),
4792 "unexpected error: {error}"
4793 );
4794 }
4795
4796 #[cfg(feature = "band9")]
4797 #[test]
4798 fn band9_transaction_reports_given_row_transport_support() {
4799 let transaction = RuntimeTransaction {
4800 inner: RuntimeTransactionInner::B9(None),
4801 driver_lease: None,
4802 };
4803
4804 assert!(transaction.supports_given_rows());
4805 }
4806
4807 #[test]
4808 fn bounded_runtime_reader_rechecks_cancellation_after_consumer_work() {
4809 let cancellation = RuntimeAnswerCancellation::default();
4810 let trigger = cancellation.clone();
4811 let limits = RuntimeAnswerLimits {
4812 max_items: 1,
4813 max_bytes: 64,
4814 deadline: None,
4815 cancellation,
4816 };
4817 let mut stats = RuntimeAnswerStats::new(RuntimeAnswerKind::Rows);
4818 let mut consumer = move |_item| {
4819 trigger.cancel();
4820 Ok(RuntimeAnswerControl::Continue)
4821 };
4822
4823 let error = runtime_accept(
4824 &limits,
4825 &mut stats,
4826 RuntimeAnswerItem::Row(serde_json::json!({"v": 1})),
4827 &mut consumer,
4828 )
4829 .unwrap_err();
4830 assert!(matches!(
4831 error,
4832 RuntimeError::ResourceLimit {
4833 code: "provider_cancelled",
4834 ..
4835 }
4836 ));
4837
4838 let cancellation = RuntimeAnswerCancellation::default();
4839 let trigger = cancellation.clone();
4840 let limits = RuntimeAnswerLimits {
4841 max_items: 1,
4842 max_bytes: 64,
4843 deadline: None,
4844 cancellation,
4845 };
4846 let mut stats = RuntimeAnswerStats::new(RuntimeAnswerKind::Rows);
4847 let mut failing_consumer = move |_item| -> Result<RuntimeAnswerControl> {
4848 trigger.cancel();
4849 Err(RuntimeError::AnswerConsumer)
4850 };
4851 let error = runtime_accept(
4852 &limits,
4853 &mut stats,
4854 RuntimeAnswerItem::Row(serde_json::json!({"v": 1})),
4855 &mut failing_consumer,
4856 )
4857 .unwrap_err();
4858 assert!(matches!(
4859 error,
4860 RuntimeError::ResourceLimit {
4861 code: "provider_cancelled",
4862 ..
4863 }
4864 ));
4865 }
4866
4867 #[test]
4868 fn connect_options_default_matches_ssot() {
4869 let options = ConnectOptions::default();
4870 assert_eq!(options.http_port, DEFAULT_HTTP_PORT);
4871 assert!(!options.tls);
4872 assert_eq!(options.server_version, None);
4873 }
4874
4875 fn test_connection_control(
4876 max_items: u64,
4877 max_bytes: u64,
4878 max_statements: u32,
4879 cancellation: RuntimeAnswerCancellation,
4880 ) -> RuntimeConnectionControl {
4881 RuntimeConnectionControl::new(
4882 max_items,
4883 max_bytes,
4884 max_statements,
4885 Instant::now() + Duration::from_secs(5),
4886 cancellation,
4887 )
4888 }
4889
4890 fn prepared_controlled_probe(
4891 control: RuntimeConnectionControl,
4892 ) -> PreparedSecureConnectOptions {
4893 SecureConnectOptions {
4894 http_port: 8123,
4895 tls_mode: TlsMode::Disabled,
4896 server_version: None,
4897 }
4898 .prepare_transport()
4899 .expect("prepare plaintext test transport")
4900 .with_connection_control(control)
4901 }
4902
4903 #[test]
4904 fn controlled_connection_meter_proves_direct_success_and_fallback_fourth_dispatch_failure() {
4905 let version = core_version::Version::new(3, 12, 1);
4906 let version_bytes = u64::try_from(version.to_string().len()).unwrap();
4907
4908 let mut direct = ConnectionMeter::new(test_connection_control(
4909 1,
4910 version_bytes,
4911 3,
4912 RuntimeAnswerCancellation::default(),
4913 ));
4914 direct.charge_statement().unwrap(); direct.charge_version_evidence(version).unwrap();
4916 direct.charge_statement().unwrap(); direct.admit_connection().unwrap();
4918 assert_eq!(direct.statements, 2);
4919 assert_eq!(direct.response_bytes, version_bytes);
4920 assert_eq!(direct.admitted_items, 1);
4921
4922 let mut fallback = ConnectionMeter::new(test_connection_control(
4923 1,
4924 version_bytes,
4925 3,
4926 RuntimeAnswerCancellation::default(),
4927 ));
4928 fallback.charge_statement().unwrap(); fallback.charge_statement().unwrap(); fallback.charge_statement().unwrap(); fallback.charge_version_evidence(version).unwrap();
4932 assert!(matches!(
4933 fallback.charge_statement(), Err(RuntimeError::ResourceLimit {
4935 code: "provider_statement_limit",
4936 ..
4937 })
4938 ));
4939 assert_eq!(fallback.statements, 3);
4940 assert_eq!(fallback.response_bytes, version_bytes);
4941 assert_eq!(fallback.admitted_items, 0);
4942
4943 let mut bytes = ConnectionMeter::new(test_connection_control(
4944 1,
4945 version_bytes,
4946 3,
4947 RuntimeAnswerCancellation::default(),
4948 ));
4949 bytes.charge_version_evidence(version).unwrap();
4950 assert!(matches!(
4951 bytes.charge_version_evidence(version),
4952 Err(RuntimeError::ResourceLimit {
4953 code: "response_byte_limit",
4954 ..
4955 })
4956 ));
4957
4958 let mut items = ConnectionMeter::new(test_connection_control(
4959 1,
4960 version_bytes,
4961 3,
4962 RuntimeAnswerCancellation::default(),
4963 ));
4964 items.admit_connection().unwrap();
4965 assert!(matches!(
4966 items.admit_connection(),
4967 Err(RuntimeError::ResourceLimit {
4968 code: "processed_item_limit",
4969 ..
4970 })
4971 ));
4972 }
4973
4974 #[test]
4975 fn legacy_connection_meter_is_an_exact_uncontrolled_adapter() {
4976 let mut meter = ConnectionMeter::legacy();
4977 for _ in 0..16 {
4978 meter.charge_statement().unwrap();
4979 meter
4980 .charge_version_evidence(core_version::Version::new(3, 12, 1))
4981 .unwrap();
4982 meter.admit_connection().unwrap();
4983 }
4984 assert_eq!(meter.statements, 0);
4985 assert_eq!(meter.response_bytes, 0);
4986 assert_eq!(meter.admitted_items, 0);
4987 }
4988
4989 #[tokio::test]
4990 async fn controlled_zero_item_and_statement_ceilings_reject_before_probe_dispatch() {
4991 for (control, expected_code) in [
4992 (
4993 test_connection_control(0, 32, 3, RuntimeAnswerCancellation::default()),
4994 "processed_item_limit",
4995 ),
4996 (
4997 test_connection_control(1, 32, 0, RuntimeAnswerCancellation::default()),
4998 "provider_statement_limit",
4999 ),
5000 ] {
5001 let probe_called = Arc::new(Mutex::new(false));
5002 let captured = Arc::clone(&probe_called);
5003 let error = gated_driver_prepared_with_probe(
5004 "endpoint must not be used",
5005 "credential must not be used",
5006 "credential must not be used",
5007 prepared_controlled_probe(control),
5008 move |_address, _port, _mode| {
5009 *captured.lock().unwrap() = true;
5010 Ok(core_version::Version::new(3, 12, 1))
5011 },
5012 )
5013 .await
5014 .err()
5015 .expect("zero applicable ceiling must reject before probe dispatch");
5016
5017 assert!(!*probe_called.lock().unwrap());
5018 assert!(matches!(
5019 error,
5020 SecureConnectError::Runtime(RuntimeError::ResourceLimit { code, .. })
5021 if code == expected_code
5022 ));
5023 }
5024 }
5025
5026 #[tokio::test]
5027 async fn controlled_probe_bytes_are_bounded_before_driver_construction() {
5028 let probe_called = Arc::new(Mutex::new(false));
5029 let captured = Arc::clone(&probe_called);
5030 let error = gated_driver_prepared_with_probe(
5031 "endpoint must not reach a driver",
5032 "credential must not be used",
5033 "credential must not be used",
5034 prepared_controlled_probe(test_connection_control(
5035 1,
5036 5,
5037 2,
5038 RuntimeAnswerCancellation::default(),
5039 )),
5040 move |_address, _port, _mode| {
5041 *captured.lock().unwrap() = true;
5042 Ok(core_version::Version::new(3, 12, 1))
5043 },
5044 )
5045 .await
5046 .err()
5047 .expect("six retained version bytes must cross a five-byte ceiling");
5048
5049 assert!(*probe_called.lock().unwrap());
5050 assert!(matches!(
5051 error,
5052 SecureConnectError::Runtime(RuntimeError::ResourceLimit {
5053 code: "response_byte_limit",
5054 ..
5055 })
5056 ));
5057 }
5058
5059 #[tokio::test]
5060 async fn generated_exact_version_rejects_after_one_probe_before_driver_construction() {
5061 let probe_calls = Arc::new(AtomicUsize::new(0));
5062 let observed_calls = Arc::clone(&probe_calls);
5063 let result = gated_driver_prepared_with_probe(
5064 "not a provider endpoint and must not reach address parsing",
5065 "credential must not be used",
5066 "credential must not be used",
5067 prepared_controlled_probe(test_connection_control(
5068 1,
5069 32,
5070 2,
5071 RuntimeAnswerCancellation::default(),
5072 ))
5073 .with_generated_3_12_3_requirement(),
5074 move |_address, _port, _mode| {
5075 observed_calls.fetch_add(1, Ordering::AcqRel);
5076 Ok(core_version::Version::new(3, 12, 0))
5077 },
5078 )
5079 .await;
5080 let error = match result {
5081 Ok(_) => panic!("a generated package requires exact TypeDB 3.12.3"),
5082 Err(error) => error,
5083 };
5084
5085 assert_eq!(probe_calls.load(Ordering::Acquire), 1);
5086 assert!(matches!(
5087 error,
5088 SecureConnectError::Runtime(RuntimeError::UnsupportedVersion(
5089 core_version::VersionError::FeatureUnsupported {
5090 server,
5091 required,
5092 ..
5093 }
5094 )) if server == core_version::Version::new(3, 12, 0)
5095 && required == core_version::Version::new(3, 12, 3)
5096 ));
5097 }
5098
5099 #[tokio::test]
5100 async fn controlled_probe_cancellation_wakes_while_blocking_work_is_in_flight() {
5101 let cancellation = RuntimeAnswerCancellation::default();
5102 let prepared =
5103 prepared_controlled_probe(test_connection_control(1, 32, 2, cancellation.clone()));
5104 let (started_sender, started_receiver) = std::sync::mpsc::channel();
5105 let (release_sender, release_receiver) = std::sync::mpsc::channel();
5106 let connection = tokio::spawn(async move {
5107 gated_driver_prepared_with_probe(
5108 "pending endpoint",
5109 "credential",
5110 "credential",
5111 prepared,
5112 move |_address, _port, _mode| {
5113 started_sender.send(()).unwrap();
5114 release_receiver.recv().unwrap();
5115 Err(core_version::VersionProbeError::Probe(
5116 core_version::VersionError::Probe("released probe".to_owned()),
5117 ))
5118 },
5119 )
5120 .await
5121 });
5122
5123 tokio::task::spawn_blocking(move || started_receiver.recv().unwrap())
5124 .await
5125 .unwrap();
5126 cancellation.cancel();
5127 let error = tokio::time::timeout(Duration::from_secs(1), connection)
5128 .await
5129 .expect("cancellation must wake the pending connection")
5130 .unwrap()
5131 .err()
5132 .expect("cancelled connection must fail");
5133 release_sender.send(()).unwrap();
5134
5135 assert!(matches!(
5136 error,
5137 SecureConnectError::Runtime(RuntimeError::ResourceLimit {
5138 code: "provider_cancelled",
5139 ..
5140 })
5141 ));
5142 }
5143
5144 #[tokio::test]
5145 async fn completed_provider_outcome_wins_an_already_ready_cancellation_branch() {
5146 let cancellation = RuntimeAnswerCancellation::default();
5147 cancellation.cancel();
5148 let control = test_connection_control(1, 32, 1, cancellation);
5149 assert_eq!(
5150 await_connection_work(std::future::ready(17_u8), Some(&control))
5151 .await
5152 .unwrap(),
5153 17
5154 );
5155 }
5156
5157 #[tokio::test(start_paused = true)]
5158 async fn controlled_deadline_wakes_pending_provider_work() {
5159 let control = RuntimeConnectionControl::new(
5160 1,
5161 32,
5162 1,
5163 Instant::now() + Duration::from_secs(1),
5164 RuntimeAnswerCancellation::default(),
5165 );
5166 let pending = tokio::spawn(async move {
5167 await_connection_work(std::future::pending::<()>(), Some(&control)).await
5168 });
5169 tokio::time::advance(Duration::from_secs(2)).await;
5170 let error = pending
5171 .await
5172 .unwrap()
5173 .expect_err("absolute deadline must wake pending provider work");
5174 assert!(matches!(
5175 error,
5176 RuntimeError::ResourceLimit {
5177 code: "transaction_deadline_exceeded",
5178 ..
5179 }
5180 ));
5181 }
5182
5183 #[test]
5184 fn controlled_connection_errors_drop_provider_endpoint_and_parse_text() {
5185 const SENTINEL: &str = "TB_CONTROLLED_CONNECT_SECRET";
5186 for error in [
5187 SecureConnectError::Runtime(RuntimeError::Connection(SENTINEL.to_owned())),
5188 SecureConnectError::Runtime(RuntimeError::UnsupportedVersion(
5189 core_version::VersionError::Parse(SENTINEL.to_owned()),
5190 )),
5191 ] {
5192 let rendered = controlled_connection_error(error).to_string();
5193 assert!(!rendered.contains(SENTINEL), "{rendered}");
5194 }
5195 }
5196
5197 #[test]
5198 fn released_boolean_options_map_to_the_typed_tls_truth_table() {
5199 let disabled = SecureConnectOptions::from(ConnectOptions {
5200 http_port: 8123,
5201 tls: false,
5202 server_version: Some(core_version::Version::new(3, 11, 5)),
5203 });
5204 assert_eq!(disabled.http_port, 8123);
5205 assert_eq!(disabled.tls_mode, TlsMode::Disabled);
5206 assert_eq!(
5207 disabled.server_version,
5208 Some(core_version::Version::new(3, 11, 5))
5209 );
5210
5211 let native = SecureConnectOptions::from(ConnectOptions {
5212 http_port: 8124,
5213 tls: true,
5214 server_version: None,
5215 });
5216 assert_eq!(native.http_port, 8124);
5217 assert_eq!(native.tls_mode, TlsMode::NativeRoots);
5218 assert_eq!(native.server_version, None);
5219 }
5220
5221 #[test]
5222 fn every_enabled_fallback_band_lowers_to_tls_without_plaintext() {
5223 let disabled = ResolvedTlsMode::from_configured_path(TlsMode::Disabled).unwrap();
5224 assert!(!disabled.probe_mode.is_enabled());
5225 #[cfg(feature = "band8")]
5226 assert!(!disabled.band8.is_enabled());
5227 #[cfg(feature = "band9")]
5228 assert!(!disabled.band9.is_enabled());
5229
5230 let native = ResolvedTlsMode::from_configured_path(TlsMode::NativeRoots).unwrap();
5231 assert!(native.probe_mode.is_enabled());
5232 #[cfg(feature = "band8")]
5233 assert!(native.band8.is_enabled());
5234 #[cfg(feature = "band9")]
5235 assert!(native.band9.is_enabled());
5236
5237 let root_ca = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/root-ca.pem");
5238 let custom =
5239 ResolvedTlsMode::from_configured_path(TlsMode::CustomRootCa(root_ca.clone())).unwrap();
5240 assert!(custom.probe_mode.is_enabled());
5241 #[cfg(any(feature = "band8", feature = "band9"))]
5242 let expected_root = std::fs::read(&root_ca).unwrap();
5243 #[cfg(feature = "band8")]
5244 {
5245 assert!(custom.band8.is_enabled());
5246 assert_eq!(
5247 std::fs::read(custom.band8.root_ca_path().unwrap()).unwrap(),
5248 expected_root
5249 );
5250 }
5251 #[cfg(feature = "band9")]
5252 {
5253 assert!(custom.band9.is_enabled());
5254 assert_eq!(
5255 std::fs::read(custom.band9.root_ca_path().unwrap()).unwrap(),
5256 expected_root
5257 );
5258 }
5259 }
5260
5261 #[test]
5262 fn every_band_lowers_from_retained_material_after_path_replacement() {
5263 let sequence = NEXT_TLS_MATERIAL_TEST_ID.fetch_add(1, Ordering::Relaxed);
5264 let directory = std::env::temp_dir().join(format!(
5265 "type-bridge-runtime-root-replacement-{}-{sequence}",
5266 std::process::id()
5267 ));
5268 std::fs::create_dir(&directory).unwrap();
5269 let directory = directory.canonicalize().unwrap();
5272 let configured = directory.join("root.pem");
5273 let moved = directory.join("loaded-root.pem");
5274 let original = include_bytes!("../tests/fixtures/root-ca.pem");
5275 std::fs::write(&configured, original).unwrap();
5276
5277 let material = core_version::RetainedCustomRootCa::load(&configured).unwrap();
5278 match std::fs::rename(&configured, &moved) {
5279 Ok(()) => {
5280 std::fs::write(&configured, b"replacement is not a certificate\n").unwrap();
5281 }
5282 Err(_) => {
5283 assert!(
5286 std::fs::write(&configured, b"replacement is not a certificate\n").is_err()
5287 );
5288 }
5289 }
5290
5291 let resolved = ResolvedTlsMode::lower(
5292 TlsMode::CustomRootCa(configured),
5293 ResolvedTlsProbeMode::CustomRootCa(material),
5294 )
5295 .expect("all bands must lower the retained original material");
5296
5297 #[cfg(feature = "band8")]
5298 assert_eq!(
5299 std::fs::read(resolved.band8.root_ca_path().unwrap()).unwrap(),
5300 original
5301 );
5302 #[cfg(feature = "band9")]
5303 assert_eq!(
5304 std::fs::read(resolved.band9.root_ca_path().unwrap()).unwrap(),
5305 original
5306 );
5307
5308 drop(resolved);
5309 std::fs::remove_dir_all(directory).unwrap();
5310 }
5311
5312 #[test]
5313 fn every_band_lowers_from_retained_material_after_parent_swap() {
5314 let sequence = NEXT_TLS_MATERIAL_TEST_ID.fetch_add(1, Ordering::Relaxed);
5315 let root = std::env::temp_dir().join(format!(
5316 "type-bridge-runtime-root-parent-swap-{}-{sequence}",
5317 std::process::id()
5318 ));
5319 std::fs::create_dir(&root).unwrap();
5320 let root = root.canonicalize().unwrap();
5323 let configured_parent = root.join("configured-parent");
5324 let moved_parent = root.join("loaded-parent");
5325 std::fs::create_dir(&configured_parent).unwrap();
5326 let configured = configured_parent.join("root.pem");
5327 let original = include_bytes!("../tests/fixtures/root-ca.pem");
5328 std::fs::write(&configured, original).unwrap();
5329
5330 let material = core_version::RetainedCustomRootCa::load(&configured).unwrap();
5331 match std::fs::rename(&configured_parent, &moved_parent) {
5332 Ok(()) => {
5333 std::fs::create_dir(&configured_parent).unwrap();
5334 std::fs::write(&configured, b"replacement is not a certificate\n").unwrap();
5335 }
5336 Err(_) => {
5337 assert_eq!(std::fs::read(&configured).unwrap(), original);
5340 }
5341 }
5342
5343 let resolved = ResolvedTlsMode::lower(
5344 TlsMode::CustomRootCa(configured),
5345 ResolvedTlsProbeMode::CustomRootCa(material),
5346 )
5347 .expect("all bands must lower the retained parent/file material");
5348
5349 #[cfg(feature = "band8")]
5350 assert_eq!(
5351 std::fs::read(resolved.band8.root_ca_path().unwrap()).unwrap(),
5352 original
5353 );
5354 #[cfg(feature = "band9")]
5355 assert_eq!(
5356 std::fs::read(resolved.band9.root_ca_path().unwrap()).unwrap(),
5357 original
5358 );
5359
5360 drop(resolved);
5361 std::fs::remove_dir_all(root).unwrap();
5362 }
5363
5364 #[test]
5365 fn every_band_lowers_from_snapshot_after_in_place_source_overwrite() {
5366 let sequence = NEXT_TLS_MATERIAL_TEST_ID.fetch_add(1, Ordering::Relaxed);
5367 let directory = std::env::temp_dir().join(format!(
5368 "type-bridge-runtime-root-overwrite-{}-{sequence}",
5369 std::process::id()
5370 ));
5371 std::fs::create_dir(&directory).unwrap();
5372 let directory = directory.canonicalize().unwrap();
5375 let configured = directory.join("root.pem");
5376 let original = include_bytes!("../tests/fixtures/root-ca.pem");
5377 std::fs::write(&configured, original).unwrap();
5378
5379 let material = core_version::RetainedCustomRootCa::load(&configured).unwrap();
5380 match std::fs::write(&configured, b"overwritten source is not a certificate\n") {
5381 Ok(()) => assert_ne!(std::fs::read(&configured).unwrap(), original),
5382 Err(_) => assert_eq!(std::fs::read(&configured).unwrap(), original),
5383 }
5384
5385 let resolved = ResolvedTlsMode::lower(
5386 TlsMode::CustomRootCa(configured),
5387 ResolvedTlsProbeMode::CustomRootCa(material),
5388 )
5389 .expect("all bands must lower the captured-byte snapshot");
5390
5391 #[cfg(feature = "band8")]
5392 assert_eq!(
5393 std::fs::read(resolved.band8.root_ca_path().unwrap()).unwrap(),
5394 original
5395 );
5396 #[cfg(feature = "band9")]
5397 assert_eq!(
5398 std::fs::read(resolved.band9.root_ca_path().unwrap()).unwrap(),
5399 original
5400 );
5401
5402 drop(resolved);
5403 std::fs::remove_dir_all(directory).unwrap();
5404 }
5405
5406 #[tokio::test]
5407 async fn prepared_transport_never_rereads_mutated_custom_root_path() {
5408 let sequence = NEXT_TLS_MATERIAL_TEST_ID.fetch_add(1, Ordering::Relaxed);
5409 let directory = std::env::temp_dir().join(format!(
5410 "type-bridge-runtime-root-prepared-{}-{sequence}",
5411 std::process::id()
5412 ));
5413 std::fs::create_dir(&directory).unwrap();
5414 let configured = directory.join("root.pem");
5415 let original = include_bytes!("../tests/fixtures/root-ca.pem");
5416 std::fs::write(&configured, original).unwrap();
5417
5418 let prepared = SecureConnectOptions {
5419 http_port: 8123,
5420 tls_mode: TlsMode::CustomRootCa(configured.clone()),
5421 server_version: None,
5422 }
5423 .prepare_transport()
5424 .expect("prepare the complete custom-root transport");
5425
5426 match std::fs::write(&configured, b"mutated after preparation\n") {
5427 Ok(()) => assert_ne!(std::fs::read(&configured).unwrap(), original),
5428 Err(_) => assert_eq!(std::fs::read(&configured).unwrap(), original),
5429 }
5430
5431 #[cfg(feature = "band8")]
5432 assert_eq!(
5433 std::fs::read(prepared.resolved_tls.band8.root_ca_path().unwrap()).unwrap(),
5434 original
5435 );
5436 #[cfg(feature = "band9")]
5437 assert_eq!(
5438 std::fs::read(prepared.resolved_tls.band9.root_ca_path().unwrap()).unwrap(),
5439 original
5440 );
5441
5442 let expected = original.to_vec();
5443 let error = gated_driver_prepared_with_probe(
5444 "host must not be constructed",
5445 "credential must not be used",
5446 "credential must not be used",
5447 prepared
5448 .clone()
5449 .with_connection_control(test_connection_control(
5450 1,
5451 32,
5452 1,
5453 RuntimeAnswerCancellation::default(),
5454 )),
5455 move |_address, port, mode| {
5456 assert_eq!(port, 8123);
5457 let ResolvedTlsProbeMode::CustomRootCa(material) = mode else {
5458 panic!("prepared probe must retain custom-root mode");
5459 };
5460 let bytes = material
5461 .with_driver_root_path(|path| std::fs::read(path))
5462 .unwrap()
5463 .unwrap();
5464 assert_eq!(bytes, expected);
5465 Err(core_version::VersionProbeError::TlsConfiguration(
5466 core_version::TlsConfigurationError::NativeRootsUnavailable,
5467 ))
5468 },
5469 )
5470 .await
5471 .err()
5472 .expect("injected terminal probe error prevents host construction");
5473 assert!(matches!(
5474 error,
5475 SecureConnectError::TlsConfiguration(
5476 core_version::TlsConfigurationError::NativeRootsUnavailable
5477 )
5478 ));
5479
5480 drop(prepared);
5481 std::fs::remove_dir_all(directory).unwrap();
5482 }
5483
5484 #[cfg(unix)]
5485 #[test]
5486 fn binding_transport_resolves_a_raw_alias_without_weakening_physical_paths() {
5487 use std::os::unix::fs::symlink;
5488
5489 let sequence = NEXT_TLS_MATERIAL_TEST_ID.fetch_add(1, Ordering::Relaxed);
5490 let directory = std::env::temp_dir().join(format!(
5491 "type-bridge-runtime-binding-alias-{}-{sequence}",
5492 std::process::id()
5493 ));
5494 std::fs::create_dir(&directory).expect("create binding-alias directory");
5495 let physical_parent = directory.join("physical");
5496 let alias_parent = directory.join("alias");
5497 std::fs::create_dir(&physical_parent).expect("create physical CA parent");
5498 std::fs::write(
5499 physical_parent.join("root.pem"),
5500 include_bytes!("../tests/fixtures/root-ca.pem"),
5501 )
5502 .expect("write binding CA");
5503 symlink(&physical_parent, &alias_parent).expect("create caller path alias");
5504 let configured = alias_parent.join("root.pem");
5505 let options = SecureConnectOptions {
5506 http_port: 8123,
5507 tls_mode: TlsMode::CustomRootCa(configured.clone()),
5508 server_version: None,
5509 };
5510
5511 assert!(matches!(
5512 options.prepare_transport_from_validated_physical_path(),
5513 Err(SecureConnectError::TlsConfiguration(
5514 core_version::TlsConfigurationError::CustomRootCaUnreadable { path }
5515 )) if path == configured
5516 ));
5517 options
5518 .prepare_transport()
5519 .expect("raw binding preparation resolves one caller alias");
5520 std::fs::remove_dir_all(directory).expect("remove binding-alias directory");
5521 }
5522
5523 #[test]
5524 fn public_transport_preflight_returns_typed_errors_without_a_host() {
5525 let missing =
5526 PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/does-not-exist.pem");
5527 let options = SecureConnectOptions {
5528 http_port: 8123,
5529 tls_mode: TlsMode::CustomRootCa(missing.clone()),
5530 server_version: Some(core_version::Version::new(3, 11, 5)),
5531 };
5532
5533 let error = options
5534 .validate_transport()
5535 .expect_err("preflight must validate trust material synchronously");
5536 assert_eq!(
5537 error.configuration_code(),
5538 Some("tls_custom_root_ca_unreadable")
5539 );
5540 assert!(matches!(
5541 error,
5542 SecureConnectError::TlsConfiguration(
5543 core_version::TlsConfigurationError::CustomRootCaUnreadable { path }
5544 ) if path == missing
5545 ));
5546 }
5547
5548 #[tokio::test]
5549 async fn invalid_custom_root_is_rejected_before_exact_version_or_probe_io() {
5550 let probe_called = Arc::new(Mutex::new(false));
5551 let captured = Arc::clone(&probe_called);
5552 let missing =
5553 PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/does-not-exist.pem");
5554
5555 let error = gated_driver_secure_with_probe(
5556 "invalid host that must not be constructed",
5557 "admin",
5558 "password",
5559 SecureConnectOptions {
5560 http_port: 8123,
5561 tls_mode: TlsMode::CustomRootCa(missing.clone()),
5562 server_version: Some(core_version::Version::new(3, 11, 5)),
5563 },
5564 move |_address, _port, _mode| {
5565 *captured.lock().unwrap() = true;
5566 Ok(core_version::Version::new(3, 11, 5))
5567 },
5568 )
5569 .await
5570 .err()
5571 .expect("invalid custom root must fail before connection construction");
5572
5573 assert!(!*probe_called.lock().unwrap());
5574 assert!(matches!(
5575 error,
5576 SecureConnectError::TlsConfiguration(
5577 core_version::TlsConfigurationError::CustomRootCaUnreadable { path }
5578 ) if path == missing
5579 ));
5580 }
5581
5582 #[tokio::test]
5583 async fn enabled_probe_tls_failure_is_terminal_before_grpc_fallback() {
5584 let error = gated_driver_secure_with_probe(
5585 "127.0.0.1:1",
5586 "admin",
5587 "password",
5588 SecureConnectOptions {
5589 http_port: 8123,
5590 tls_mode: TlsMode::NativeRoots,
5591 server_version: None,
5592 },
5593 |_address, _port, mode| {
5594 assert!(matches!(mode, ResolvedTlsProbeMode::NativeRoots));
5595 Err(core_version::VersionProbeError::TlsConfiguration(
5596 core_version::TlsConfigurationError::NativeRootsUnavailable,
5597 ))
5598 },
5599 )
5600 .await
5601 .err()
5602 .expect("TLS configuration failure must be terminal");
5603
5604 assert!(matches!(
5605 error,
5606 SecureConnectError::TlsConfiguration(
5607 core_version::TlsConfigurationError::NativeRootsUnavailable
5608 )
5609 ));
5610 }
5611
5612 #[tokio::test]
5618 async fn gated_driver_probe_receives_configured_port() {
5619 let recorded_port: Arc<Mutex<Option<u16>>> = Arc::new(Mutex::new(None));
5620 let captured = Arc::clone(&recorded_port);
5621
5622 let result = gated_driver_with_probe(
5623 "localhost:1729",
5624 "admin",
5625 "password",
5626 ConnectOptions {
5627 http_port: 9123,
5628 tls: false,
5629 server_version: None,
5630 },
5631 move |_addr, port, _tls| {
5632 *captured.lock().unwrap() = Some(port);
5633 Ok(core_version::Version::new(3, 11, 5))
5634 },
5635 )
5636 .await;
5637
5638 let observed = recorded_port.lock().unwrap().expect("probe was not called");
5639 assert_eq!(
5640 observed, 9123,
5641 "probe must receive the configured http_port (9123), got {observed}"
5642 );
5643
5644 if let Err(RuntimeError::UnsupportedVersion(_)) = result {
5645 panic!("expected a connection error (no server), not a version gate rejection")
5646 }
5647 }
5648
5649 #[tokio::test]
5653 async fn gated_driver_http_failure_reports_grpc_fallback_failures() {
5654 let result = gated_driver_with_probe(
5655 "127.0.0.1:1",
5656 "admin",
5657 "password",
5658 ConnectOptions {
5659 http_port: 9123,
5660 tls: false,
5661 server_version: None,
5662 },
5663 move |_addr, _port, _tls| {
5664 Err(core_version::VersionError::Probe(
5665 "HTTP endpoint unavailable".to_string(),
5666 ))
5667 },
5668 )
5669 .await;
5670
5671 match result {
5672 Err(RuntimeError::UnsupportedVersion(err)) => {
5673 let msg = err.to_string();
5674 assert!(msg.contains("HTTP endpoint unavailable"), "{msg}");
5675 assert!(msg.contains("band-8 gRPC attempt failed"), "{msg}");
5676 assert!(!msg.contains("band-9 gRPC attempt"), "{msg}");
5679 }
5680 Err(other) => panic!("expected aggregated version-probe failure, got {other}"),
5681 Ok(_) => panic!("expected aggregated version-probe failure, got successful connection"),
5682 }
5683 }
5684
5685 #[cfg(feature = "band8")]
5686 #[test]
5687 fn band8_lazy_validation_failure_is_retryable() {
5688 let result = classify_band8_grpc_version(
5689 "localhost:1729",
5690 Err("protocol handshake failed".to_string()),
5691 )
5692 .expect("transport validation failure should remain retryable");
5693
5694 match result {
5695 Band8GrpcVersion::RetryableFailure(failure) => {
5696 assert!(failure.contains("localhost:1729"), "{failure}");
5697 assert!(failure.contains("protocol handshake failed"), "{failure}");
5698 }
5699 Band8GrpcVersion::Validated(version) => {
5700 panic!("failed lazy connection unexpectedly validated as {version}")
5701 }
5702 }
5703 }
5704
5705 #[cfg(feature = "band8")]
5706 #[test]
5707 fn band8_reported_versions_are_authoritative() {
5708 let unsupported = classify_band8_grpc_version("localhost:1729", Ok("3.7.3".to_string()))
5709 .expect_err("reported below-window version must be terminal");
5710 assert!(matches!(
5711 unsupported,
5712 RuntimeError::UnsupportedVersion(core_version::VersionError::Unsupported {
5713 component: "server",
5714 found: core_version::Version {
5715 major: 3,
5716 minor: 7,
5717 patch: 3,
5718 },
5719 })
5720 ));
5721
5722 let retired = classify_band8_grpc_version("localhost:1729", Ok("3.10.4".to_string()))
5723 .expect_err("reported retired server version must not silently fall through");
5724 assert!(matches!(
5725 retired,
5726 RuntimeError::UnsupportedVersion(core_version::VersionError::Unsupported {
5727 component: "server",
5728 found: core_version::Version {
5729 major: 3,
5730 minor: 10,
5731 patch: 4,
5732 },
5733 })
5734 ));
5735
5736 let malformed =
5737 classify_band8_grpc_version("localhost:1729", Ok("not-a-version".to_string()))
5738 .expect_err("malformed reported version must be terminal");
5739 assert!(matches!(
5740 malformed,
5741 RuntimeError::UnsupportedVersion(core_version::VersionError::Parse(_))
5742 ));
5743
5744 let validated = classify_band8_grpc_version("localhost:1729", Ok("3.11.5".to_string()))
5745 .expect("reported band-8 version should validate");
5746 assert!(matches!(
5747 validated,
5748 Band8GrpcVersion::Validated(core_version::Version {
5749 major: 3,
5750 minor: 11,
5751 patch: 5,
5752 })
5753 ));
5754
5755 let backward_compatible =
5756 classify_band8_grpc_version("localhost:1729", Ok("3.12.0".to_string()))
5757 .expect("reported band-9 server accepts the discovery band");
5758 assert!(matches!(
5759 backward_compatible,
5760 Band8GrpcVersion::Validated(core_version::Version {
5761 major: 3,
5762 minor: 12,
5763 patch: 0,
5764 })
5765 ));
5766 }
5767
5768 #[tokio::test]
5772 async fn gated_driver_pinned_version_skips_probe() {
5773 let probe_called = Arc::new(Mutex::new(false));
5774 let captured = Arc::clone(&probe_called);
5775
5776 let result = gated_driver_with_probe(
5777 "localhost:1729",
5778 "admin",
5779 "password",
5780 ConnectOptions {
5781 http_port: 9123,
5782 tls: false,
5783 server_version: Some(core_version::Version::new(3, 11, 5)),
5784 },
5785 move |_addr, _port, _tls| {
5786 *captured.lock().unwrap() = true;
5787 Ok(core_version::Version::new(3, 11, 5))
5788 },
5789 )
5790 .await;
5791
5792 assert!(
5793 !*probe_called.lock().unwrap(),
5794 "pinned server_version must skip the HTTP probe"
5795 );
5796
5797 if let Err(RuntimeError::UnsupportedVersion(_)) = result {
5798 panic!("expected a connection error (no server), not a version gate rejection")
5799 }
5800 }
5801
5802 #[tokio::test]
5805 async fn gated_driver_rejects_unsupported_pinned_version_without_probe() {
5806 let probe_called = Arc::new(Mutex::new(false));
5807 let captured = Arc::clone(&probe_called);
5808
5809 let result = gated_driver_with_probe(
5810 "localhost:1729",
5811 "admin",
5812 "password",
5813 ConnectOptions {
5814 http_port: 9123,
5815 tls: false,
5816 server_version: Some(core_version::Version::new(3, 7, 3)),
5817 },
5818 move |_addr, _port, _tls| {
5819 *captured.lock().unwrap() = true;
5820 Ok(core_version::Version::new(3, 11, 5))
5821 },
5822 )
5823 .await;
5824
5825 assert!(
5826 !*probe_called.lock().unwrap(),
5827 "unsupported pinned server_version must skip the HTTP probe"
5828 );
5829
5830 match result {
5831 Err(RuntimeError::UnsupportedVersion(err)) => {
5832 assert!(
5833 err.to_string().contains("3.7.3"),
5834 "error should name rejected version: {err}"
5835 );
5836 }
5837 Err(other) => panic!("expected unsupported-version rejection for 3.7.3, got {other}"),
5838 Ok(_) => panic!(
5839 "expected unsupported-version rejection for 3.7.3, got successful connection"
5840 ),
5841 }
5842 }
5843
5844 #[test]
5845 fn cargo_lock_pin() {
5846 let lock_path = concat!(env!("CARGO_MANIFEST_DIR"), "/../../Cargo.lock");
5847 let lock_contents = std::fs::read_to_string(lock_path)
5848 .expect("Cargo.lock not found relative to crate root");
5849
5850 let lock_version = lock_contents
5851 .split("[[package]]")
5852 .find(|block| block.contains("name = \"type-bridge-typedb-driver-b8\""))
5853 .and_then(|block| {
5854 block
5855 .lines()
5856 .find(|line| line.trim_start().starts_with("version = "))
5857 })
5858 .and_then(|line| {
5859 let start = line.find('"')? + 1;
5860 let end = line.rfind('"')?;
5861 Some(&line[start..end])
5862 })
5863 .expect("type-bridge-typedb-driver-b8 entry not found in Cargo.lock");
5864
5865 assert_eq!(
5866 lock_version, PINNED_DRIVER_VERSION,
5867 "Cargo.lock resolves type-bridge-typedb-driver-b8 {lock_version} but \
5868 PINNED_DRIVER_VERSION \
5869 is {PINNED_DRIVER_VERSION}; update the runtime constant"
5870 );
5871
5872 let pinned: core_version::Version = PINNED_DRIVER_VERSION.parse().unwrap();
5873 assert_eq!(
5874 core_version::band(&pinned),
5875 Some(8),
5876 "pinned driver version {PINNED_DRIVER_VERSION} left protocol band 8; \
5877 review the gate expectations before accepting the bump"
5878 );
5879 }
5880
5881 #[cfg(feature = "band9")]
5884 #[test]
5885 fn given_rows_lowering() {
5886 let spec = GivenRowsSpec {
5887 variables: vec!["n".into(), "a".into()],
5888 rows: vec![
5889 vec![GivenValue::String("alice".into()), GivenValue::Integer(28)],
5890 vec![GivenValue::String("bob".into()), GivenValue::Integer(26)],
5891 ],
5892 };
5893 let rows = given_rows_b9(spec).expect("valid spec must lower");
5894 let (header, rows) = rows.into_parts();
5895 assert_eq!(header.width(), 2);
5896 assert_eq!(rows.len(), 2);
5897 }
5898
5899 #[cfg(feature = "band9")]
5900 #[test]
5901 fn given_rows_lowering_rejects_width_mismatch() {
5902 let spec = GivenRowsSpec {
5903 variables: vec!["n".into(), "a".into()],
5904 rows: vec![vec![GivenValue::String("alice".into())]],
5905 };
5906 let err = given_rows_b9(spec).expect_err("short row must be rejected");
5907 assert!(matches!(err, RuntimeError::QueryExecution(_)), "{err}");
5908 }
5909
5910 #[cfg(feature = "band9")]
5911 #[test]
5912 fn given_entry_temporal_parsing() {
5913 for value in [
5914 GivenValue::Empty,
5915 GivenValue::Date("2026-07-13".into()),
5916 GivenValue::Datetime("2026-07-13T10:30:00".into()),
5917 GivenValue::DatetimeTz("2026-07-13T10:30:00+09:00".into()),
5918 GivenValue::DatetimeTz("2026-07-13T10:30:00+00:19:32".into()),
5919 GivenValue::DatetimeTzExact {
5920 local: "2024-10-27T01:30:00".into(),
5921 named_zone: Some("Europe/London".into()),
5922 effective_offset_seconds: 3_600,
5923 },
5924 GivenValue::DatetimeTzExact {
5925 local: "2024-07-01T12:00:00".into(),
5926 named_zone: Some("europe/amsterdam".into()),
5927 effective_offset_seconds: 7_200,
5928 },
5929 GivenValue::DatetimeTzExact {
5930 local: "1900-01-01T12:00:00".into(),
5931 named_zone: None,
5932 effective_offset_seconds: 1_172,
5933 },
5934 GivenValue::Decimal("12.30dec".into()),
5935 GivenValue::Decimal("-9223372036854775808".into()),
5936 GivenValue::Duration {
5937 months: u32::MAX,
5938 days: u32::MAX,
5939 nanos: u64::MAX,
5940 },
5941 GivenValue::Boolean(true),
5942 GivenValue::Double(1.5),
5943 ] {
5944 given_entry_b9(value.clone()).unwrap_or_else(|e| panic!("{value:?} must convert: {e}"));
5945 }
5946 }
5947
5948 #[cfg(feature = "band9")]
5949 #[test]
5950 fn exact_given_entries_preserve_named_timezone_and_duration_components() {
5951 use typedb_driver::concept::Value;
5952 use typedb_driver::concept::value::TimeZone;
5953 use typedb_driver::given::GivenRowEntry;
5954
5955 let named = given_entry_b9(GivenValue::DatetimeTzExact {
5956 local: "2024-10-27T01:30:00".into(),
5957 named_zone: Some("Europe/London".into()),
5958 effective_offset_seconds: 3_600,
5959 })
5960 .expect("explicit overlap side");
5961 let GivenRowEntry::Value(Value::DatetimeTZ(named)) = named else {
5962 panic!("datetime-tz given entry")
5963 };
5964 let TimeZone::IANA(timezone) = named.timezone() else {
5965 panic!("authored IANA identity must survive transport")
5966 };
5967 assert_eq!(timezone.name(), "Europe/London");
5968 assert_eq!(named.naive_local().to_string(), "2024-10-27 01:30:00");
5969 assert_eq!(
5970 named
5971 .naive_local()
5972 .signed_duration_since(named.naive_utc())
5973 .num_seconds(),
5974 3_600
5975 );
5976
5977 let lower_case = given_entry_b9(GivenValue::DatetimeTzExact {
5978 local: "2024-07-01T12:00:00".into(),
5979 named_zone: Some("europe/amsterdam".into()),
5980 effective_offset_seconds: 7_200,
5981 })
5982 .expect("case-insensitive names admitted by schema also lower");
5983 let GivenRowEntry::Value(Value::DatetimeTZ(lower_case)) = lower_case else {
5984 panic!("datetime-tz given entry")
5985 };
5986 let TimeZone::IANA(timezone) = lower_case.timezone() else {
5987 panic!("IANA identity must survive transport")
5988 };
5989 assert_eq!(
5990 timezone.name(),
5991 "Europe/Amsterdam",
5992 "the provider canonicalizes accepted authored case"
5993 );
5994
5995 let duration = given_entry_b9(GivenValue::Duration {
5996 months: u32::MAX,
5997 days: u32::MAX,
5998 nanos: u64::MAX,
5999 })
6000 .expect("duration boundary");
6001 let GivenRowEntry::Value(Value::Duration(duration)) = duration else {
6002 panic!("duration given entry")
6003 };
6004 assert_eq!(duration.months, u32::MAX);
6005 assert_eq!(duration.days, u32::MAX);
6006 assert_eq!(duration.nanos, u64::MAX);
6007 }
6008
6009 #[cfg(feature = "band9")]
6010 #[test]
6011 fn given_entry_rejects_malformed_temporal() {
6012 for value in [
6013 GivenValue::Date("not-a-date".into()),
6014 GivenValue::Datetime("2026-13-45T99:00:00".into()),
6015 GivenValue::DatetimeTz("2026-07-13 10:30".into()),
6016 GivenValue::DatetimeTz("2026-07-13T10:30:00+24:00".into()),
6017 GivenValue::DatetimeTzExact {
6018 local: "2024-10-27T01:30:00".into(),
6019 named_zone: Some("Europe/London".into()),
6020 effective_offset_seconds: 7_200,
6021 },
6022 GivenValue::DatetimeTzExact {
6023 local: "2024-03-31T01:30:00".into(),
6024 named_zone: Some("Europe/London".into()),
6025 effective_offset_seconds: 3_600,
6026 },
6027 GivenValue::Decimal("1.00000000000000000000".into()),
6028 ] {
6029 let err = given_entry_b9(value.clone())
6030 .expect_err("malformed temporal string must be rejected");
6031 assert!(
6032 err.to_string().contains("Invalid given"),
6033 "{value:?}: {err}"
6034 );
6035 }
6036 }
6037
6038 #[test]
6039 fn cargo_lock_pin_b9() {
6040 let lock_path = concat!(env!("CARGO_MANIFEST_DIR"), "/../../Cargo.lock");
6041 let lock_contents = std::fs::read_to_string(lock_path)
6042 .expect("Cargo.lock not found relative to crate root");
6043
6044 let lock_version = lock_contents
6045 .split("[[package]]")
6046 .find(|block| block.contains("name = \"typedb-driver\""))
6047 .and_then(|block| {
6048 block
6049 .lines()
6050 .find(|line| line.trim_start().starts_with("version = "))
6051 })
6052 .and_then(|line| {
6053 let start = line.find('"')? + 1;
6054 let end = line.rfind('"')?;
6055 Some(&line[start..end])
6056 })
6057 .expect("typedb-driver entry not found in Cargo.lock");
6058
6059 assert_eq!(
6060 lock_version, PINNED_DRIVER_VERSION_B9,
6061 "Cargo.lock resolves typedb-driver {lock_version} but \
6062 PINNED_DRIVER_VERSION_B9 is {PINNED_DRIVER_VERSION_B9}; update the runtime constant"
6063 );
6064
6065 let pinned: core_version::Version = PINNED_DRIVER_VERSION_B9.parse().unwrap();
6066 assert_eq!(
6067 core_version::band(&pinned),
6068 Some(9),
6069 "pinned band-9 driver version {PINNED_DRIVER_VERSION_B9} left protocol band 9; \
6070 review the gate expectations before accepting the bump"
6071 );
6072 }
6073
6074 #[test]
6075 fn materialize_runtime_answer_preserves_empty_kinds() {
6076 let stats = |kind| RuntimeAnswerStats {
6077 kind,
6078 processed_items: 0,
6079 response_bytes: 0,
6080 stopped_early: false,
6081 };
6082 assert!(matches!(
6083 super::materialize_runtime_answer(stats(RuntimeAnswerKind::Ok), Vec::new()),
6084 QueryResult::Ok
6085 ));
6086 assert!(matches!(
6087 super::materialize_runtime_answer(stats(RuntimeAnswerKind::Rows), Vec::new()),
6088 QueryResult::Rows(values) if values.is_empty()
6089 ));
6090 assert!(matches!(
6091 super::materialize_runtime_answer(stats(RuntimeAnswerKind::Documents), Vec::new()),
6092 QueryResult::Documents(values) if values.is_empty()
6093 ));
6094 }
6095
6096 #[test]
6097 fn canonical_decimal_text_normalizes_provider_display() {
6098 assert_eq!(
6099 super::canonical_decimal_text("001234.5600dec".into()).unwrap(),
6100 "1234.56"
6101 );
6102 assert_eq!(
6103 super::canonical_decimal_text("-0.000dec".into()).unwrap(),
6104 "0"
6105 );
6106 assert_eq!(
6107 super::canonical_decimal_text("-9223372036854775808dec".into()).unwrap(),
6108 "-9223372036854775808"
6109 );
6110 assert!(matches!(
6111 super::canonical_decimal_text("not-a-decimal".into()),
6112 Err(RuntimeError::QueryExecution(message)) if message == "provider returned an unparseable decimal"
6113 ));
6114 }
6115}