1#![cfg_attr(docsrs, feature(doc_cfg))]
2#![warn(missing_docs)] #[macro_use]
9extern crate tracing;
10
11mod activity;
12mod async_activity_handle;
13pub mod callback_based;
14mod dns;
15#[cfg(feature = "envconfig")]
17pub mod envconfig;
18pub mod errors;
19pub mod grpc;
20pub mod interceptors;
22mod metrics;
23mod options_structs;
24#[cfg(feature = "experimental")]
25pub mod plugins;
27mod proxy;
28mod replaceable;
29pub mod request_extensions;
30mod retry;
31mod rpc_options;
32pub mod schedules;
34#[cfg(test)]
35mod test_helpers;
36pub mod worker;
37mod workflow_handle;
38mod workflow_status;
39
40pub use crate::{proxy::HttpConnectProxyOptions, request_extensions::PayloadErrorLimits};
41pub use activity::*;
42pub use async_activity_handle::{
43 ActivityHeartbeatResponse, ActivityIdentifier, AsyncActivityHandle,
44};
45pub(crate) use retry::CallType;
46#[doc(hidden)]
47pub use retry::jittered;
48
49pub use interceptors::{
50 BackfillScheduleInput, CancelWorkflowInput, ClientInterceptor, CompleteAsyncActivityInput,
51 CountWorkflowsInput, CountWorkflowsOutput, CreateScheduleInput, CreateScheduleOutput,
52 DeleteScheduleInput, DescribeScheduleInput, DescribeScheduleOutput, DescribeWorkflowInput,
53 DescribeWorkflowOutput, FailAsyncActivityInput, FetchWorkflowHistoryPageInput,
54 FetchWorkflowHistoryPageOutput, HasArgs, HeartbeatAsyncActivityInput, ListSchedulesPageInput,
55 ListSchedulesPageOutput, ListWorkflowsPageInput, ListWorkflowsPageOutput, Next,
56 PauseScheduleInput, PollWorkflowUpdateInput, PollWorkflowUpdateOutput, QueryWorkflowInput,
57 QueryWorkflowOutput, ReportAsyncActivityCancellationInput, SendScheduleUpdateInput,
58 SignalWithStartWorkflowInput, SignalWorkflowInput, StartWorkflowInput, StartWorkflowOutput,
59 StartWorkflowUpdateInput, StartWorkflowUpdateOutput, TemporalClientValue,
60 TerminateWorkflowInput, TriggerScheduleInput, UnpauseScheduleInput, UpdateScheduleInput,
61 UpdateWithStartWorkflowInput, UpdateWithStartWorkflowOutput,
62};
63pub use metrics::{LONG_REQUEST_LATENCY_HISTOGRAM_NAME, REQUEST_LATENCY_HISTOGRAM_NAME};
64pub use options_structs::*;
65#[cfg(feature = "experimental")]
66pub use plugins::{
67 ClientPlugin, ErasedClientPlugin, PluginApplyError, PluginError, PluginTarget, WorkerPluginData,
68};
69pub use replaceable::SharedReplaceableClient;
70pub use retry::RetryOptions;
71pub use rpc_options::{RpcMetadata, RpcMetadataError, RpcOptions};
72pub use temporalio_common::{Memo, RetryPolicy};
73pub use url::Url;
74pub mod danger {
76 pub use tokio_rustls::rustls::client::danger::ServerCertVerifier;
80}
81#[cfg(feature = "dynamic-tls")]
82pub use tokio_rustls::rustls::SignatureScheme;
85#[cfg(feature = "dynamic-tls")]
86pub use tokio_rustls::rustls::client::ResolvesClientCert;
95#[cfg(feature = "dynamic-tls")]
96pub use tokio_rustls::rustls::sign::CertifiedKey;
99pub use tonic;
100pub use workflow_handle::{
101 UntypedQuery, UntypedSignal, UntypedUpdate, UntypedWorkflow, UntypedWorkflowHandle,
102 WorkflowExecutionDescription, WorkflowExecutionInfo, WorkflowExecutionResult, WorkflowHandle,
103 WorkflowHistory, WorkflowHistoryError, WorkflowResultDetails, WorkflowUpdateHandle,
104};
105pub use workflow_status::WorkflowExecutionStatus;
106
107use crate::{
108 grpc::{
109 AttachMetricLabels, CloudService, HealthService, OperatorService, TestService,
110 WorkflowService,
111 },
112 metrics::{ChannelOrGrpcOverride, GrpcMetricSvc, MetricsContext},
113 request_extensions::RequestExt,
114 worker::ClientWorkerSet,
115};
116use errors::*;
117use futures_util::{
118 future::{BoxFuture, try_join},
119 stream,
120 stream::Stream,
121};
122use http::Uri;
123use parking_lot::RwLock;
124use std::{
125 collections::{HashMap, VecDeque},
126 error::Error,
127 fmt::Debug,
128 pin::Pin,
129 str::FromStr,
130 sync::{Arc, OnceLock},
131 task::{Context, Poll},
132 time::{Duration, SystemTime},
133};
134use temporalio_common::{
135 ActivityDefinition, HasWorkflowDefinition, SignalDefinition, UntypedActivity, UpdateDefinition,
136 data_converters::{
137 ActivitySerializationContext, DataConverter, SerializationContext,
138 SerializationContextData, WorkflowSerializationContext,
139 },
140 payload_visitor::decode_payloads,
141 protos::{
142 coresdk::IntoPayloadsExt,
143 grpc::health::v1::health_client::HealthClient,
144 proto_ts_to_system_time,
145 temporal::api::{
146 cloud::cloudservice::v1::cloud_service_client::CloudServiceClient,
147 common::v1::{ActivityType, Memo as ProtoMemo, Payloads, WorkflowType},
148 enums::v1::{
149 ActivityIdConflictPolicy as ProtoActivityIdConflictPolicy,
150 ActivityIdReusePolicy as ProtoActivityIdReusePolicy, TaskQueueKind,
151 UpdateWorkflowExecutionLifecycleStage,
152 WorkflowIdConflictPolicy as ProtoWorkflowIdConflictPolicy,
153 WorkflowIdReusePolicy as ProtoWorkflowIdReusePolicy,
154 },
155 operatorservice::v1::operator_service_client::OperatorServiceClient,
156 sdk::v1::UserMetadata,
157 taskqueue::v1::TaskQueue,
158 testservice::v1::test_service_client::TestServiceClient,
159 workflow::v1 as workflow,
160 workflowservice::v1::{
161 count_workflow_executions_response,
162 execute_multi_operation_request::operation::Operation as MultiOperationRequest,
163 execute_multi_operation_response::response::Response as MultiOperationResponse,
164 workflow_service_client::WorkflowServiceClient, *,
165 },
166 },
167 },
168 search_attributes::{SearchAttributeError, SearchAttributeValue, SearchAttributes},
169};
170use tonic::{
171 Code, IntoRequest,
172 body::Body,
173 client::GrpcService,
174 codec::CompressionEncoding,
175 codegen::InterceptedService,
176 metadata::{
177 AsciiMetadataKey, AsciiMetadataValue, BinaryMetadataKey, BinaryMetadataValue, MetadataMap,
178 MetadataValue,
179 },
180 service::Interceptor,
181 transport::{Certificate, Endpoint, Identity},
182};
183use tower::ServiceBuilder;
184use uuid::Uuid;
185
186static CLIENT_NAME_HEADER_KEY: &str = "client-name";
187static CLIENT_VERSION_HEADER_KEY: &str = "client-version";
188static TEMPORAL_NAMESPACE_HEADER_KEY: &str = "temporal-namespace";
189
190#[doc(hidden)]
191pub static MESSAGE_TOO_LARGE_KEY: &str = "message-too-large";
193#[doc(hidden)]
194pub static ERROR_RETURNED_DUE_TO_SHORT_CIRCUIT: &str = "short-circuit";
196
197const LONG_POLL_TIMEOUT: Duration = Duration::from_secs(70);
199const OTHER_CALL_TIMEOUT: Duration = Duration::from_secs(30);
200const VERSION: &str = env!("CARGO_PKG_VERSION");
201
202#[derive(Clone, Debug)]
207pub struct Connection {
208 inner: Arc<ConnectionInner>,
209}
210
211#[derive(Clone, derive_more::Debug)]
212struct ConnectionInner {
213 #[debug(skip)]
214 service: TemporalServiceClient,
215 retry_options: RetryOptions,
216 identity: String,
217 headers: Arc<RwLock<ClientHeaders>>,
218 client_name: String,
219 client_version: String,
220 capabilities: Option<get_system_info_response::Capabilities>,
222 workers: Arc<ClientWorkerSet>,
223 _dns_task: Option<Arc<dns::DnsReresolutionHandle>>,
224 payloads_warn_size: usize,
226 memo_warn_size: usize,
227}
228
229fn resolve_warn_threshold(option: &'static str, bytes: u64) -> usize {
234 usize::try_from(bytes).unwrap_or_else(|_| {
235 warn!(
236 option,
237 configured_bytes = bytes,
238 "Configured payload size warning threshold exceeds the maximum addressable size on this \
239 platform; disabling this warning"
240 );
241 0
242 })
243}
244
245impl Connection {
246 pub async fn connect(mut options: ConnectionOptions) -> Result<Self, ClientConnectError> {
248 if options.service_override.is_some() {
249 options.grpc_compression = GrpcCompression::None;
250 }
251
252 let first_result = Self::connect_once(&options).await;
253 if options.grpc_compression == GrpcCompression::Gzip
254 && let Err(ClientConnectError::SystemInfoCallError(status)) = &first_result
255 && status.code() == Code::Unimplemented
256 && {
257 let msg = status.message().to_lowercase();
258 msg.contains("decompress")
259 || msg.contains("grpc-encoding")
260 || msg.contains("compressor")
261 }
262 {
263 options.grpc_compression = GrpcCompression::None;
264 return Self::connect_once(&options).await;
265 }
266 first_result
267 }
268
269 async fn connect_once(options: &ConnectionOptions) -> Result<Self, ClientConnectError> {
270 let dns_lb_opts = dns::validate_and_get_dns_lb(options)?.cloned();
271 let (service, dns_task) = if let Some(service_override) = options.service_override.clone() {
272 (
273 GrpcMetricSvc {
274 inner: ChannelOrGrpcOverride::GrpcOverride(service_override),
275 metrics: options.metrics_meter.clone().map(MetricsContext::new),
276 disable_errcode_label: options.disable_error_code_metric_tags,
277 },
278 None,
279 )
280 } else if let Some(dns_opts) = &dns_lb_opts {
281 let (channel, sender) = dns::create_balanced_channel(options).await?;
282 let handle = dns::spawn_dns_reresolution(
283 sender,
284 options.target.clone(),
285 options.tls_options.clone(),
286 options.keep_alive.clone(),
287 options.override_origin.clone(),
288 dns_opts.resolution_interval,
289 options.connect_timeout,
290 );
291 (
292 ServiceBuilder::new()
293 .layer_fn(move |channel| GrpcMetricSvc {
294 inner: ChannelOrGrpcOverride::Channel(channel),
295 metrics: options.metrics_meter.clone().map(MetricsContext::new),
296 disable_errcode_label: options.disable_error_code_metric_tags,
297 })
298 .service(channel),
299 Some(handle),
300 )
301 } else {
302 let endpoint = Endpoint::from_shared(options.target.to_string())?;
303 let endpoint = if let Some(timeout) = options.connect_timeout {
304 endpoint.connect_timeout(timeout)
305 } else {
306 endpoint
307 };
308 let tls_result = add_tls_to_channel(options.tls_options.as_ref(), endpoint).await?;
309
310 #[cfg(feature = "dynamic-tls")]
311 let (channel, custom_connector_info) = match tls_result {
312 TlsConfigResult::Standard(ep) => (
313 ep,
314 None::<(Arc<tokio_rustls::rustls::ClientConfig>, String)>,
315 ),
316 TlsConfigResult::CustomConnector {
317 endpoint: ep,
318 rustls_config,
319 domain,
320 } => (ep, Some((rustls_config, domain))),
321 };
322 #[cfg(not(feature = "dynamic-tls"))]
323 let channel = match tls_result {
324 TlsConfigResult::Standard(ep) => ep,
325 };
326
327 let channel = if let Some(keep_alive) = options.keep_alive.as_ref() {
328 channel
329 .keep_alive_while_idle(true)
330 .http2_keep_alive_interval(keep_alive.interval)
331 .keep_alive_timeout(keep_alive.timeout)
332 } else {
333 channel
334 };
335 let channel = if let Some(origin) = options.override_origin.clone() {
336 channel.origin(origin)
337 } else {
338 channel
339 };
340 #[cfg(feature = "dynamic-tls")]
342 if options.http_connect_proxy.is_some() && custom_connector_info.is_some() {
343 return Err(ClientConnectError::InvalidConfig(
344 "client_cert_resolver is not yet supported with http_connect_proxy. \
345 Use static client_tls_options when using a proxy, or remove the proxy."
346 .to_owned(),
347 ));
348 }
349 let channel = if let Some(proxy) = options.http_connect_proxy.as_ref() {
351 proxy.connect_endpoint(&channel).await?
352 } else {
353 #[cfg(feature = "dynamic-tls")]
354 if let Some((rustls_config, domain)) = custom_connector_info {
355 let server_name =
356 tokio_rustls::rustls::pki_types::ServerName::try_from(domain.as_str())
357 .map_err(|e| {
358 ClientConnectError::InvalidConfig(format!(
359 "Invalid TLS domain name '{domain}': {e}"
360 ))
361 })?
362 .to_owned();
363 let connector = DynamicTlsConnector {
364 tls: tokio_rustls::TlsConnector::from(rustls_config),
365 domain: Arc::new(server_name),
366 };
367 channel.connect_with_connector(connector).await?
368 } else {
369 channel.connect().await?
370 }
371 #[cfg(not(feature = "dynamic-tls"))]
372 channel.connect().await?
373 };
374 (
375 ServiceBuilder::new()
376 .layer_fn(move |channel| GrpcMetricSvc {
377 inner: ChannelOrGrpcOverride::Channel(channel),
378 metrics: options.metrics_meter.clone().map(MetricsContext::new),
379 disable_errcode_label: options.disable_error_code_metric_tags,
380 })
381 .service(channel),
382 None,
383 )
384 };
385
386 let headers = Arc::new(RwLock::new(ClientHeaders {
387 user_headers: parse_ascii_headers(options.headers.clone().unwrap_or_default())?,
388 user_binary_headers: parse_binary_headers(
389 options.binary_headers.clone().unwrap_or_default(),
390 )?,
391 api_key: options.api_key.clone(),
392 }));
393 let interceptor = ServiceCallInterceptor {
394 client_name: options.client_name.clone(),
395 client_version: options.client_version.clone(),
396 headers: headers.clone(),
397 };
398 let svc = InterceptedService::new(service, interceptor);
399 let mut svc_client = TemporalServiceClient::new(svc, options.grpc_compression);
400
401 let capabilities = if !options.skip_get_system_info {
402 match svc_client
403 .get_system_info(GetSystemInfoRequest::default().into_request())
404 .await
405 {
406 Ok(sysinfo) => sysinfo.into_inner().capabilities,
407 Err(status) => match status.code() {
408 Code::Unimplemented
409 if {
410 let msg = status.message().to_lowercase();
411 msg.contains("unknown method")
412 || msg.contains("unknown service")
413 || msg.contains("method not found")
414 || (msg.contains("getsysteminfo")
415 && (msg.contains("is unimplemented")
416 || msg.contains("not implement")))
417 } =>
418 {
419 None
420 }
421 _ => return Err(ClientConnectError::SystemInfoCallError(status)),
422 },
423 }
424 } else {
425 None
426 };
427 #[cfg(feature = "experimental")]
428 let payloads_warn_size = options.payload_limits.payloads_warn_size;
429 #[cfg(not(feature = "experimental"))]
430 let payloads_warn_size = options_structs::DEFAULT_PAYLOADS_WARN_SIZE;
431 #[cfg(feature = "experimental")]
432 let memo_warn_size = options.payload_limits.memo_warn_size;
433 #[cfg(not(feature = "experimental"))]
434 let memo_warn_size = options_structs::DEFAULT_MEMO_WARN_SIZE;
435 Ok(Self {
436 inner: Arc::new(ConnectionInner {
437 service: svc_client,
438 retry_options: options.retry_options.clone(),
439 identity: options.identity.clone(),
440 headers,
441 client_name: options.client_name.clone(),
442 client_version: options.client_version.clone(),
443 capabilities,
444 workers: Arc::new(ClientWorkerSet::new()),
445 _dns_task: dns_task,
446 payloads_warn_size: resolve_warn_threshold(
447 "payloads_warn_size",
448 payloads_warn_size,
449 ),
450 memo_warn_size: resolve_warn_threshold("memo_warn_size", memo_warn_size),
451 }),
452 })
453 }
454
455 pub fn set_api_key(&self, api_key: Option<String>) {
457 self.inner.headers.write().api_key = api_key;
458 }
459
460 pub fn set_headers(&self, headers: HashMap<String, String>) -> Result<(), InvalidHeaderError> {
469 self.inner.headers.write().user_headers = parse_ascii_headers(headers)?;
470 Ok(())
471 }
472
473 pub fn set_binary_headers(
482 &self,
483 binary_headers: HashMap<String, Vec<u8>>,
484 ) -> Result<(), InvalidHeaderError> {
485 self.inner.headers.write().user_binary_headers = parse_binary_headers(binary_headers)?;
486 Ok(())
487 }
488
489 pub fn client_name(&self) -> &str {
491 &self.inner.client_name
492 }
493
494 pub fn client_version(&self) -> &str {
496 &self.inner.client_version
497 }
498
499 pub fn capabilities(&self) -> Option<&get_system_info_response::Capabilities> {
502 self.inner.capabilities.as_ref()
503 }
504
505 pub fn retry_options_mut(&mut self) -> &mut RetryOptions {
510 &mut Arc::make_mut(&mut self.inner).retry_options
511 }
512
513 pub fn identity(&self) -> &str {
515 &self.inner.identity
516 }
517
518 pub fn identity_mut(&mut self) -> &mut String {
523 &mut Arc::make_mut(&mut self.inner).identity
524 }
525
526 pub fn workers(&self) -> Arc<ClientWorkerSet> {
528 self.inner.workers.clone()
529 }
530
531 pub fn worker_grouping_key(&self) -> Uuid {
533 self.inner.workers.worker_grouping_key()
534 }
535
536 pub fn workflow_service(&self) -> Box<dyn WorkflowService> {
538 self.inner.service.workflow_service()
539 }
540
541 pub fn operator_service(&self) -> Box<dyn OperatorService> {
543 self.inner.service.operator_service()
544 }
545
546 pub fn cloud_service(&self) -> Box<dyn CloudService> {
548 self.inner.service.cloud_service()
549 }
550
551 pub fn test_service(&self) -> Box<dyn TestService> {
553 self.inner.service.test_service()
554 }
555
556 pub fn health_service(&self) -> Box<dyn HealthService> {
558 self.inner.service.health_service()
559 }
560}
561
562#[derive(Debug)]
563struct ClientHeaders {
564 user_headers: HashMap<AsciiMetadataKey, AsciiMetadataValue>,
565 user_binary_headers: HashMap<BinaryMetadataKey, BinaryMetadataValue>,
566 api_key: Option<String>,
567}
568
569impl ClientHeaders {
570 fn apply_to_metadata(&self, metadata: &mut MetadataMap) {
571 for (key, val) in self.user_headers.iter() {
572 if !metadata.contains_key(key) {
574 metadata.insert(key, val.clone());
575 }
576 }
577 for (key, val) in self.user_binary_headers.iter() {
578 if !metadata.contains_key(key) {
580 metadata.insert_bin(key, val.clone());
581 }
582 }
583 if let Some(api_key) = &self.api_key {
584 if !metadata.contains_key("authorization")
586 && let Ok(val) = format!("Bearer {api_key}").parse()
587 {
588 metadata.insert("authorization", val);
589 }
590 }
591 }
592}
593
594#[derive(Debug)]
597enum TlsConfigResult {
598 Standard(Endpoint),
600 #[cfg(feature = "dynamic-tls")]
605 CustomConnector {
606 endpoint: Endpoint,
607 rustls_config: Arc<tokio_rustls::rustls::ClientConfig>,
608 domain: String,
609 },
610}
611
612async fn add_tls_to_channel(
620 tls_options: Option<&TlsOptions>,
621 mut channel: Endpoint,
622) -> Result<TlsConfigResult, ClientConnectError> {
623 if let Some(tls_cfg) = tls_options {
624 if tls_cfg.server_cert_verifier.is_some() && tls_cfg.server_root_ca_cert.is_some() {
625 return Err(ClientConnectError::InvalidConfig(
626 "Cannot set both `server_root_ca_cert` and `server_cert_verifier`".to_owned(),
627 ));
628 }
629
630 #[cfg(feature = "dynamic-tls")]
631 if tls_cfg.client_tls_options.is_some() && tls_cfg.client_cert_resolver.is_some() {
632 return Err(ClientConnectError::InvalidConfig(
633 "Cannot set both `client_tls_options` and `client_cert_resolver`. \
634 Use `client_tls_options` for static certificates or \
635 `client_cert_resolver` for dynamic certificate resolution, but not both."
636 .to_owned(),
637 ));
638 }
639
640 let domain_override = tls_cfg.domain.clone();
642 if let Some(domain) = &domain_override {
643 let uri: Uri = format!("https://{domain}").parse()?;
644 channel = channel.origin(uri);
645 }
646
647 #[cfg(feature = "dynamic-tls")]
649 if let Some(resolver) = &tls_cfg.client_cert_resolver {
650 let rustls_config = build_custom_rustls_config(tls_cfg, Some(resolver.clone()))?;
651 let sni_domain = domain_override
654 .or_else(|| {
655 channel
656 .uri()
657 .host()
658 .map(|h| h.trim_matches(|c| c == '[' || c == ']').to_owned())
659 })
660 .ok_or_else(|| {
661 ClientConnectError::InvalidConfig(
662 "Cannot determine TLS server name for dynamic cert resolution: \
663 set 'domain' in TlsOptions or use a URL with a hostname"
664 .to_owned(),
665 )
666 })?;
667 return Ok(TlsConfigResult::CustomConnector {
668 endpoint: channel,
669 rustls_config: Arc::new(rustls_config),
670 domain: sni_domain,
671 });
672 }
673
674 let mut tls = tonic::transport::ClientTlsConfig::new();
676
677 if tls_cfg.server_cert_verifier.is_none() {
678 if let Some(root_cert) = &tls_cfg.server_root_ca_cert {
679 let server_root_ca_cert = Certificate::from_pem(root_cert);
680 tls = tls.ca_certificate(server_root_ca_cert);
681 } else {
682 tls = tls.with_native_roots();
683 }
684 }
685
686 if let Some(domain) = &tls_cfg.domain {
687 tls = tls.domain_name(domain);
688 }
689
690 if let Some(client_opts) = &tls_cfg.client_tls_options {
691 let client_identity =
692 Identity::from_pem(&client_opts.client_cert, &client_opts.client_private_key);
693 tls = tls.identity(client_identity);
694 }
695
696 let endpoint = if let Some(verifier) = &tls_cfg.server_cert_verifier {
697 channel
698 .tls_config_with_verifier(tls, verifier.clone())
699 .map_err(ClientConnectError::from)?
700 } else {
701 channel.tls_config(tls).map_err(ClientConnectError::from)?
702 };
703 return Ok(TlsConfigResult::Standard(endpoint));
704 }
705 Ok(TlsConfigResult::Standard(channel))
706}
707
708#[cfg(feature = "dynamic-tls")]
709fn build_custom_rustls_config(
714 tls_cfg: &TlsOptions,
715 client_cert_resolver: Option<Arc<dyn tokio_rustls::rustls::client::ResolvesClientCert>>,
716) -> Result<tokio_rustls::rustls::ClientConfig, ClientConnectError> {
717 use tokio_rustls::rustls::{ClientConfig, RootCertStore, crypto};
718
719 let provider = crypto::CryptoProvider::get_default()
721 .cloned()
722 .or_else(|| {
723 #[cfg(feature = "tls-ring")]
725 {
726 return Some(Arc::new(crypto::ring::default_provider()));
727 }
728 #[cfg(feature = "tls-aws-lc")]
729 #[allow(unreachable_code)]
730 {
731 return Some(Arc::new(crypto::aws_lc_rs::default_provider()));
732 }
733 #[allow(unreachable_code)]
734 None
735 })
736 .ok_or_else(|| {
737 ClientConnectError::InvalidConfig(
738 "No TLS crypto provider available. Enable the `tls-ring` or `tls-aws-lc` feature."
739 .to_owned(),
740 )
741 })?;
742
743 let builder = ClientConfig::builder_with_provider(provider)
744 .with_safe_default_protocol_versions()
745 .map_err(|e| {
746 ClientConnectError::InvalidConfig(format!("Failed to configure TLS protocols: {e}"))
747 })?;
748
749 let builder = if let Some(verifier) = &tls_cfg.server_cert_verifier {
751 builder
752 .dangerous()
753 .with_custom_certificate_verifier(verifier.clone())
754 } else {
755 use std::io::Cursor;
756 use tokio_rustls::rustls::pki_types::{CertificateDer, pem::PemObject as _};
757
758 let mut roots = RootCertStore::empty();
759 if let Some(ca_cert) = &tls_cfg.server_root_ca_cert {
760 let certs: Vec<CertificateDer<'static>> =
761 CertificateDer::pem_reader_iter(&mut Cursor::new(ca_cert))
762 .collect::<Result<Vec<_>, _>>()
763 .map_err(|e| {
764 ClientConnectError::InvalidConfig(format!(
765 "Failed to parse CA certificate PEM: {e}"
766 ))
767 })?;
768 roots.add_parsable_certificates(certs);
769 if roots.is_empty() {
770 return Err(ClientConnectError::InvalidConfig(
771 "None of the provided CA certificates could be parsed. \
772 Ensure the PEM data contains valid X.509 certificates."
773 .to_owned(),
774 ));
775 }
776 } else {
777 let native_result = rustls_native_certs::load_native_certs();
779 if !native_result.errors.is_empty() {
780 warn!(
781 "errors occurred when loading native certs: {:?}",
782 native_result.errors
783 );
784 }
785 if native_result.certs.is_empty() {
786 return Err(ClientConnectError::InvalidConfig(
787 "No native TLS root certificates found".to_owned(),
788 ));
789 }
790 roots.add_parsable_certificates(native_result.certs);
791 if roots.is_empty() {
792 return Err(ClientConnectError::InvalidConfig(
793 "Native TLS root certificates were found but none could be parsed".to_owned(),
794 ));
795 }
796 }
797 builder.with_root_certificates(roots)
798 };
799
800 let mut config = if let Some(resolver) = client_cert_resolver {
802 builder.with_client_cert_resolver(resolver)
803 } else {
804 builder.with_no_client_auth()
805 };
806
807 config.alpn_protocols.push(b"h2".to_vec());
809
810 Ok(config)
811}
812
813#[cfg(feature = "dynamic-tls")]
814const DYNAMIC_TLS_CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
818
819#[cfg(feature = "dynamic-tls")]
820#[derive(Clone)]
823struct DynamicTlsConnector {
824 tls: tokio_rustls::TlsConnector,
825 domain: Arc<tokio_rustls::rustls::pki_types::ServerName<'static>>,
826}
827
828#[cfg(feature = "dynamic-tls")]
829impl std::fmt::Debug for DynamicTlsConnector {
830 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
831 f.debug_struct("DynamicTlsConnector")
832 .field("domain", &self.domain)
833 .finish()
834 }
835}
836
837#[cfg(feature = "dynamic-tls")]
838impl tower::Service<Uri> for DynamicTlsConnector {
839 type Response = hyper_util::rt::TokioIo<tokio_rustls::client::TlsStream<tokio::net::TcpStream>>;
840 type Error = Box<dyn std::error::Error + Send + Sync>;
841 type Future =
842 Pin<Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>> + Send>>;
843
844 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
845 Poll::Ready(Ok(()))
846 }
847
848 fn call(&mut self, uri: Uri) -> Self::Future {
849 let tls = self.tls.clone();
850 let domain = self.domain.clone();
851
852 Box::pin(async move {
853 let host = uri
854 .host()
855 .ok_or_else(|| -> Box<dyn std::error::Error + Send + Sync> {
856 format!("URI has no host for TLS connection: {uri}").into()
857 })?;
858 let port = uri.port_u16().unwrap_or(443);
859 let addr_display = format!("{}:{}", host, port);
862
863 debug!(target: "temporal_client", %uri, addr = %addr_display, "DynamicTlsConnector: establishing TCP+TLS connection");
864
865 let tcp = tokio::time::timeout(
869 DYNAMIC_TLS_CONNECT_TIMEOUT,
870 tokio::net::TcpStream::connect((host, port)),
871 )
872 .await
873 .map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
874 format!(
875 "TCP connect to {addr_display} timed out after {}s",
876 DYNAMIC_TLS_CONNECT_TIMEOUT.as_secs()
877 )
878 .into()
879 })?
880 .map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
881 format!("TCP connect to {addr_display} failed: {e}").into()
882 })?;
883
884 tcp.set_nodelay(true)?;
886
887 let tls_stream = tls.connect(domain.as_ref().to_owned(), tcp).await?;
888 debug!(target: "temporal_client", addr = %addr_display, "DynamicTlsConnector: TLS handshake complete");
889 Ok(hyper_util::rt::TokioIo::new(tls_stream))
890 })
891 }
892}
893
894fn parse_ascii_headers(
895 headers: HashMap<String, String>,
896) -> Result<HashMap<AsciiMetadataKey, AsciiMetadataValue>, InvalidHeaderError> {
897 let mut parsed_headers = HashMap::with_capacity(headers.len());
898 for (k, v) in headers.into_iter() {
899 let key = match AsciiMetadataKey::from_str(&k) {
900 Ok(key) => key,
901 Err(err) => {
902 return Err(InvalidHeaderError::InvalidAsciiHeaderKey {
903 key: k,
904 source: err,
905 });
906 }
907 };
908 let value = match MetadataValue::from_str(&v) {
909 Ok(value) => value,
910 Err(err) => {
911 return Err(InvalidHeaderError::InvalidAsciiHeaderValue {
912 key: k,
913 value: v,
914 source: err,
915 });
916 }
917 };
918 parsed_headers.insert(key, value);
919 }
920
921 Ok(parsed_headers)
922}
923
924fn parse_binary_headers(
925 headers: HashMap<String, Vec<u8>>,
926) -> Result<HashMap<BinaryMetadataKey, BinaryMetadataValue>, InvalidHeaderError> {
927 let mut parsed_headers = HashMap::with_capacity(headers.len());
928 for (k, v) in headers.into_iter() {
929 let key = match BinaryMetadataKey::from_str(&k) {
930 Ok(key) => key,
931 Err(err) => {
932 return Err(InvalidHeaderError::InvalidBinaryHeaderKey {
933 key: k,
934 source: err,
935 });
936 }
937 };
938 let value = BinaryMetadataValue::from_bytes(&v);
939 parsed_headers.insert(key, value);
940 }
941
942 Ok(parsed_headers)
943}
944
945#[derive(Clone)]
947pub struct ServiceCallInterceptor {
948 client_name: String,
949 client_version: String,
950 headers: Arc<RwLock<ClientHeaders>>,
952}
953
954impl Interceptor for ServiceCallInterceptor {
955 fn call(
958 &mut self,
959 mut request: tonic::Request<()>,
960 ) -> Result<tonic::Request<()>, tonic::Status> {
961 let metadata = request.metadata_mut();
962 if !metadata.contains_key(CLIENT_NAME_HEADER_KEY) {
963 metadata.insert(
964 CLIENT_NAME_HEADER_KEY,
965 self.client_name
966 .parse()
967 .unwrap_or_else(|_| MetadataValue::from_static("")),
968 );
969 }
970 if !metadata.contains_key(CLIENT_VERSION_HEADER_KEY) {
971 metadata.insert(
972 CLIENT_VERSION_HEADER_KEY,
973 self.client_version
974 .parse()
975 .unwrap_or_else(|_| MetadataValue::from_static("")),
976 );
977 }
978 self.headers.read().apply_to_metadata(metadata);
979 request.set_default_timeout(OTHER_CALL_TIMEOUT);
980
981 Ok(request)
982 }
983}
984
985#[derive(Clone)]
987pub struct TemporalServiceClient {
988 workflow_svc_client: Box<dyn WorkflowService>,
989 operator_svc_client: Box<dyn OperatorService>,
990 cloud_svc_client: Box<dyn CloudService>,
991 test_svc_client: Box<dyn TestService>,
992 health_svc_client: Box<dyn HealthService>,
993}
994
995fn get_decode_max_size() -> usize {
998 static _DECODE_MAX_SIZE: OnceLock<usize> = OnceLock::new();
999 *_DECODE_MAX_SIZE.get_or_init(|| {
1000 std::env::var("TEMPORAL_MAX_INCOMING_GRPC_BYTES")
1001 .ok()
1002 .and_then(|s| s.parse().ok())
1003 .unwrap_or(128 * 1024 * 1024)
1004 })
1005}
1006
1007impl TemporalServiceClient {
1008 fn new<T>(svc: T, compression: GrpcCompression) -> Self
1009 where
1010 T: GrpcService<Body> + Send + Sync + Clone + 'static,
1011 T::ResponseBody: tonic::codegen::Body<Data = tonic::codegen::Bytes> + Send + 'static,
1012 T::Error: Into<tonic::codegen::StdError>,
1013 <T::ResponseBody as tonic::codegen::Body>::Error: Into<tonic::codegen::StdError> + Send,
1014 <T as GrpcService<Body>>::Future: Send,
1015 {
1016 macro_rules! configure {
1019 ($client:expr) => {{
1020 let client = $client.max_decoding_message_size(get_decode_max_size());
1021 match compression {
1022 GrpcCompression::Gzip => client
1023 .send_compressed(CompressionEncoding::Gzip)
1024 .accept_compressed(CompressionEncoding::Gzip),
1025 GrpcCompression::None => client,
1026 }
1027 }};
1028 }
1029
1030 let workflow_svc_client = Box::new(configure!(WorkflowServiceClient::new(svc.clone())));
1031 let operator_svc_client = Box::new(configure!(OperatorServiceClient::new(svc.clone())));
1032 let cloud_svc_client = Box::new(configure!(CloudServiceClient::new(svc.clone())));
1033 let test_svc_client = Box::new(configure!(TestServiceClient::new(svc.clone())));
1034 let health_svc_client = Box::new(configure!(HealthClient::new(svc.clone())));
1035
1036 Self {
1037 workflow_svc_client,
1038 operator_svc_client,
1039 cloud_svc_client,
1040 test_svc_client,
1041 health_svc_client,
1042 }
1043 }
1044
1045 pub fn from_services(
1048 workflow: Box<dyn WorkflowService>,
1049 operator: Box<dyn OperatorService>,
1050 cloud: Box<dyn CloudService>,
1051 test: Box<dyn TestService>,
1052 health: Box<dyn HealthService>,
1053 ) -> Self {
1054 Self {
1055 workflow_svc_client: workflow,
1056 operator_svc_client: operator,
1057 cloud_svc_client: cloud,
1058 test_svc_client: test,
1059 health_svc_client: health,
1060 }
1061 }
1062
1063 pub fn workflow_service(&self) -> Box<dyn WorkflowService> {
1065 self.workflow_svc_client.clone()
1066 }
1067 pub fn operator_service(&self) -> Box<dyn OperatorService> {
1069 self.operator_svc_client.clone()
1070 }
1071 pub fn cloud_service(&self) -> Box<dyn CloudService> {
1073 self.cloud_svc_client.clone()
1074 }
1075 pub fn test_service(&self) -> Box<dyn TestService> {
1077 self.test_svc_client.clone()
1078 }
1079 pub fn health_service(&self) -> Box<dyn HealthService> {
1081 self.health_svc_client.clone()
1082 }
1083}
1084
1085#[derive(Clone, Debug)]
1088pub struct Client {
1089 connection: Connection,
1090 options: Arc<ClientOptions>,
1091}
1092
1093impl Client {
1094 pub async fn connect(
1097 connection_options: ConnectionOptions,
1098 client_options: ClientOptions,
1099 ) -> Result<Self, ClientConnectError> {
1100 #[cfg(feature = "experimental")]
1101 let mut connection_options = connection_options;
1102 #[cfg(feature = "experimental")]
1103 plugins::apply_connection_plugins(&client_options, &mut connection_options)?;
1104 let connection = Connection::connect(connection_options).await?;
1105 Ok(Self::new(connection, client_options)?)
1106 }
1107
1108 pub fn new(connection: Connection, options: ClientOptions) -> Result<Self, ClientNewError> {
1113 #[cfg(feature = "experimental")]
1114 let mut options = options;
1115 #[cfg(feature = "experimental")]
1116 plugins::apply_client_plugins(&mut options)?;
1117 Ok(Client {
1118 connection,
1119 options: Arc::new(options),
1120 })
1121 }
1122
1123 pub fn options(&self) -> &ClientOptions {
1125 &self.options
1126 }
1127
1128 pub fn options_mut(&mut self) -> &mut ClientOptions {
1133 Arc::make_mut(&mut self.options)
1134 }
1135
1136 pub fn connection(&self) -> &Connection {
1138 &self.connection
1139 }
1140
1141 pub fn connection_mut(&mut self) -> &mut Connection {
1143 &mut self.connection
1144 }
1145}
1146
1147impl Client {
1151 pub async fn start_workflow<W>(
1156 &self,
1157 workflow: W,
1158 input: W::Input,
1159 options: WorkflowStartOptions,
1160 ) -> Result<WorkflowHandle<Self, W>, WorkflowStartError>
1161 where
1162 W: HasWorkflowDefinition,
1163 W::Input: Send,
1164 {
1165 WorkflowClientTrait::start_workflow(self, workflow, input, options).await
1166 }
1167
1168 pub async fn signal_with_start_workflow<W, S>(
1172 &self,
1173 workflow: W,
1174 workflow_input: W::Input,
1175 signal: S,
1176 signal_input: S::Input,
1177 options: WorkflowStartOptions,
1178 ) -> Result<WorkflowHandle<Self, W>, WorkflowStartError>
1179 where
1180 W: HasWorkflowDefinition,
1181 W::Input: Send,
1182 S: SignalDefinition<Workflow = W::Run>,
1183 S::Input: Send,
1184 {
1185 WorkflowClientTrait::signal_with_start_workflow(
1186 self,
1187 workflow,
1188 workflow_input,
1189 signal,
1190 signal_input,
1191 options,
1192 )
1193 .await
1194 }
1195
1196 pub async fn start_update_with_start_workflow<W, U>(
1201 &self,
1202 workflow: W,
1203 workflow_input: W::Input,
1204 update: U,
1205 update_input: U::Input,
1206 options: WorkflowUpdateWithStartOptions,
1207 ) -> Result<WorkflowUpdateHandle<Self, U::Output>, WorkflowUpdateWithStartError>
1208 where
1209 W: HasWorkflowDefinition,
1210 W::Input: Send,
1211 U: UpdateDefinition<Workflow = W::Run>,
1212 U::Input: Send,
1213 {
1214 WorkflowClientTrait::start_update_with_start_workflow(
1215 self,
1216 workflow,
1217 workflow_input,
1218 update,
1219 update_input,
1220 options,
1221 )
1222 .await
1223 }
1224
1225 pub async fn execute_update_with_start_workflow<W, U>(
1230 &self,
1231 workflow: W,
1232 workflow_input: W::Input,
1233 update: U,
1234 update_input: U::Input,
1235 options: WorkflowUpdateWithStartOptions,
1236 ) -> Result<U::Output, WorkflowUpdateWithStartError>
1237 where
1238 W: HasWorkflowDefinition,
1239 W::Input: Send,
1240 U: UpdateDefinition<Workflow = W::Run>,
1241 U::Input: Send,
1242 {
1243 WorkflowClientTrait::execute_update_with_start_workflow(
1244 self,
1245 workflow,
1246 workflow_input,
1247 update,
1248 update_input,
1249 options,
1250 )
1251 .await
1252 }
1253
1254 pub fn get_workflow_handle<W: HasWorkflowDefinition>(
1258 &self,
1259 workflow_id: impl Into<String>,
1260 ) -> WorkflowHandle<Self, W> {
1261 WorkflowClientTrait::get_workflow_handle(self, workflow_id)
1262 }
1263
1264 pub fn list_workflows(
1269 &self,
1270 query: impl Into<String>,
1271 opts: WorkflowListOptions,
1272 ) -> ListWorkflowsStream {
1273 WorkflowClientTrait::list_workflows(self, query, opts)
1274 }
1275
1276 pub async fn count_workflows(
1278 &self,
1279 query: impl Into<String>,
1280 opts: WorkflowCountOptions,
1281 ) -> Result<WorkflowExecutionCount, ClientError> {
1282 WorkflowClientTrait::count_workflows(self, query, opts).await
1283 }
1284
1285 pub fn get_async_activity_handle(
1292 &self,
1293 identifier: ActivityIdentifier,
1294 ) -> AsyncActivityHandle<Self> {
1295 WorkflowClientTrait::get_async_activity_handle(self, identifier)
1296 }
1297
1298 pub async fn start_activity<A>(
1303 &self,
1304 activity: A,
1305 input: A::Input,
1306 options: ActivityStartOptions,
1307 ) -> Result<ActivityHandle<Self, A>, StartActivityError>
1308 where
1309 A: ActivityDefinition,
1310 {
1311 WorkflowClientTrait::start_activity(self, activity, input, options).await
1312 }
1313
1314 pub fn get_activity_handle<A>(
1325 &self,
1326 activity: A,
1327 id: impl Into<String>,
1328 run_id: Option<String>,
1329 ) -> ActivityHandle<Self, A>
1330 where
1331 Self: Sized,
1332 A: ActivityDefinition,
1333 {
1334 WorkflowClientTrait::get_activity_handle(self, activity, id, run_id)
1335 }
1336
1337 pub fn get_untyped_activity_handle(
1348 &self,
1349 id: impl Into<String>,
1350 run_id: Option<String>,
1351 ) -> ActivityHandle<Self, UntypedActivity>
1352 where
1353 Self: Sized,
1354 {
1355 WorkflowClientTrait::get_untyped_activity_handle(self, id, run_id)
1356 }
1357
1358 pub fn list_activities(
1360 &self,
1361 query: impl Into<String>,
1362 options: ActivityListOptions,
1363 ) -> ListActivitiesStream {
1364 WorkflowClientTrait::list_activities(self, query, options)
1365 }
1366
1367 pub async fn count_activities(
1369 &self,
1370 query: impl Into<String>,
1371 options: ActivityCountOptions,
1372 ) -> Result<ActivityExecutionCount, ClientError> {
1373 WorkflowClientTrait::count_activities(self, query, options).await
1374 }
1375}
1376
1377impl NamespacedClient for Client {
1378 fn namespace(&self) -> String {
1379 self.options.namespace.clone()
1380 }
1381
1382 fn identity(&self) -> String {
1383 self.connection.identity().to_owned()
1384 }
1385
1386 fn data_converter(&self) -> &DataConverter {
1387 &self.options.data_converter
1388 }
1389
1390 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
1391 &self.options.client_interceptors
1392 }
1393}
1394
1395#[derive(Clone)]
1397pub enum Namespace {
1398 Name(String),
1400 Id(String),
1402}
1403
1404pub(crate) trait WorkflowClientTrait: NamespacedClient {
1407 fn start_workflow<W>(
1409 &self,
1410 workflow: W,
1411 input: W::Input,
1412 options: WorkflowStartOptions,
1413 ) -> impl Future<Output = Result<WorkflowHandle<Self, W>, WorkflowStartError>>
1414 where
1415 Self: Sized,
1416 W: HasWorkflowDefinition,
1417 W::Input: Send;
1418
1419 fn signal_with_start_workflow<W, S>(
1421 &self,
1422 workflow: W,
1423 workflow_input: W::Input,
1424 signal: S,
1425 signal_input: S::Input,
1426 options: WorkflowStartOptions,
1427 ) -> impl Future<Output = Result<WorkflowHandle<Self, W>, WorkflowStartError>>
1428 where
1429 Self: Sized,
1430 W: HasWorkflowDefinition,
1431 W::Input: Send,
1432 S: SignalDefinition<Workflow = W::Run>,
1433 S::Input: Send;
1434
1435 fn start_update_with_start_workflow<W, U>(
1438 &self,
1439 workflow: W,
1440 workflow_input: W::Input,
1441 update: U,
1442 update_input: U::Input,
1443 options: WorkflowUpdateWithStartOptions,
1444 ) -> impl Future<Output = Result<WorkflowUpdateHandle<Self, U::Output>, WorkflowUpdateWithStartError>>
1445 where
1446 Self: Sized,
1447 W: HasWorkflowDefinition,
1448 W::Input: Send,
1449 U: UpdateDefinition<Workflow = W::Run>,
1450 U::Input: Send;
1451
1452 fn execute_update_with_start_workflow<W, U>(
1455 &self,
1456 workflow: W,
1457 workflow_input: W::Input,
1458 update: U,
1459 update_input: U::Input,
1460 options: WorkflowUpdateWithStartOptions,
1461 ) -> impl Future<Output = Result<U::Output, WorkflowUpdateWithStartError>>
1462 where
1463 Self: Sized,
1464 W: HasWorkflowDefinition,
1465 W::Input: Send,
1466 U: UpdateDefinition<Workflow = W::Run>,
1467 U::Input: Send;
1468
1469 fn get_workflow_handle<W: HasWorkflowDefinition>(
1476 &self,
1477 workflow_id: impl Into<String>,
1478 ) -> WorkflowHandle<Self, W>
1479 where
1480 Self: Sized;
1481
1482 fn list_workflows(
1486 &self,
1487 query: impl Into<String>,
1488 opts: WorkflowListOptions,
1489 ) -> ListWorkflowsStream;
1490
1491 fn count_workflows(
1493 &self,
1494 query: impl Into<String>,
1495 opts: WorkflowCountOptions,
1496 ) -> impl Future<Output = Result<WorkflowExecutionCount, ClientError>>;
1497
1498 fn get_async_activity_handle(
1502 &self,
1503 identifier: ActivityIdentifier,
1504 ) -> AsyncActivityHandle<Self>
1505 where
1506 Self: Sized;
1507
1508 fn start_activity<A>(
1510 &self,
1511 activity: A,
1512 input: A::Input,
1513 options: ActivityStartOptions,
1514 ) -> impl Future<Output = Result<ActivityHandle<Self, A>, StartActivityError>>
1515 where
1516 Self: Sized,
1517 A: ActivityDefinition;
1518
1519 fn get_activity_handle<A>(
1521 &self,
1522 activity: A,
1523 id: impl Into<String>,
1524 run_id: Option<String>,
1525 ) -> ActivityHandle<Self, A>
1526 where
1527 Self: Sized,
1528 A: ActivityDefinition;
1529
1530 fn get_untyped_activity_handle(
1532 &self,
1533 id: impl Into<String>,
1534 run_id: Option<String>,
1535 ) -> ActivityHandle<Self, UntypedActivity>
1536 where
1537 Self: Sized;
1538
1539 fn list_activities(
1541 &self,
1542 query: impl Into<String>,
1543 _options: ActivityListOptions,
1544 ) -> ListActivitiesStream;
1545
1546 fn count_activities(
1548 &self,
1549 query: impl Into<String>,
1550 _options: ActivityCountOptions,
1551 ) -> impl Future<Output = Result<ActivityExecutionCount, ClientError>>;
1552}
1553
1554pub trait NamespacedClient {
1556 fn namespace(&self) -> String;
1558 fn identity(&self) -> String;
1560 fn data_converter(&self) -> &DataConverter {
1563 static DEFAULT: OnceLock<DataConverter> = OnceLock::new();
1564 DEFAULT.get_or_init(DataConverter::default)
1565 }
1566 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
1574 &[]
1575 }
1576}
1577
1578#[derive(Debug, Clone)]
1581pub struct WorkflowExecution {
1582 raw: workflow::WorkflowExecutionInfo,
1583 data_converter: DataConverter,
1584}
1585
1586impl WorkflowExecution {
1587 fn new_with_data_converter(
1588 raw: workflow::WorkflowExecutionInfo,
1589 data_converter: DataConverter,
1590 ) -> Self {
1591 Self {
1592 raw,
1593 data_converter,
1594 }
1595 }
1596
1597 pub fn id(&self) -> &str {
1599 self.raw
1600 .execution
1601 .as_ref()
1602 .map(|e| e.workflow_id.as_str())
1603 .unwrap_or("")
1604 }
1605
1606 pub fn run_id(&self) -> &str {
1608 self.raw
1609 .execution
1610 .as_ref()
1611 .map(|e| e.run_id.as_str())
1612 .unwrap_or("")
1613 }
1614
1615 pub fn workflow_type(&self) -> &str {
1617 self.raw
1618 .r#type
1619 .as_ref()
1620 .map(|t| t.name.as_str())
1621 .unwrap_or("")
1622 }
1623
1624 pub fn status(&self) -> WorkflowExecutionStatus {
1626 WorkflowExecutionStatus::from_raw(self.raw.status)
1627 }
1628
1629 pub fn start_time(&self) -> Option<SystemTime> {
1631 self.raw
1632 .start_time
1633 .as_ref()
1634 .and_then(proto_ts_to_system_time)
1635 }
1636
1637 pub fn execution_time(&self) -> Option<SystemTime> {
1639 self.raw
1640 .execution_time
1641 .as_ref()
1642 .and_then(proto_ts_to_system_time)
1643 }
1644
1645 pub fn close_time(&self) -> Option<SystemTime> {
1647 self.raw
1648 .close_time
1649 .as_ref()
1650 .and_then(proto_ts_to_system_time)
1651 }
1652
1653 pub fn task_queue(&self) -> &str {
1655 &self.raw.task_queue
1656 }
1657
1658 pub fn history_length(&self) -> i64 {
1660 self.raw.history_length
1661 }
1662
1663 pub fn memo(&self) -> Memo {
1665 Memo::from_raw(
1666 self.raw.memo.clone(),
1667 self.data_converter.payload_converter().clone(),
1668 SerializationContextData::Workflow(WorkflowSerializationContext::new()),
1669 )
1670 }
1671
1672 pub fn parent_id(&self) -> Option<&str> {
1674 self.raw
1675 .parent_execution
1676 .as_ref()
1677 .map(|e| e.workflow_id.as_str())
1678 }
1679
1680 pub fn parent_run_id(&self) -> Option<&str> {
1682 self.raw
1683 .parent_execution
1684 .as_ref()
1685 .map(|e| e.run_id.as_str())
1686 }
1687
1688 pub fn search_attributes(&self) -> SearchAttributes {
1690 self.raw
1691 .search_attributes
1692 .as_ref()
1693 .map(SearchAttributes::from_proto)
1694 .unwrap_or_default()
1695 }
1696
1697 pub fn raw(&self) -> &workflow::WorkflowExecutionInfo {
1699 &self.raw
1700 }
1701
1702 pub fn into_raw(self) -> workflow::WorkflowExecutionInfo {
1704 self.raw
1705 }
1706}
1707
1708pub struct ListWorkflowsStream {
1711 inner: Pin<Box<dyn Stream<Item = Result<WorkflowExecution, ClientError>> + Send>>,
1712}
1713
1714impl ListWorkflowsStream {
1715 fn new(
1716 inner: Pin<Box<dyn Stream<Item = Result<WorkflowExecution, ClientError>> + Send>>,
1717 ) -> Self {
1718 Self { inner }
1719 }
1720}
1721
1722impl Stream for ListWorkflowsStream {
1723 type Item = Result<WorkflowExecution, ClientError>;
1724
1725 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1726 self.inner.as_mut().poll_next(cx)
1727 }
1728}
1729
1730#[derive(Debug, Clone)]
1735pub struct WorkflowExecutionCount {
1736 count: usize,
1737 groups: Vec<WorkflowCountAggregationGroup>,
1738}
1739
1740impl WorkflowExecutionCount {
1741 pub(crate) fn from_response(resp: CountWorkflowExecutionsResponse) -> Self {
1742 Self {
1743 count: resp.count as usize,
1744 groups: resp
1745 .groups
1746 .into_iter()
1747 .map(WorkflowCountAggregationGroup::from_proto)
1748 .collect(),
1749 }
1750 }
1751
1752 pub fn count(&self) -> usize {
1755 self.count
1756 }
1757
1758 pub fn groups(&self) -> &[WorkflowCountAggregationGroup] {
1760 &self.groups
1761 }
1762}
1763
1764#[derive(Debug, Clone)]
1766pub struct WorkflowCountAggregationGroup {
1767 raw: count_workflow_executions_response::AggregationGroup,
1768}
1769
1770impl WorkflowCountAggregationGroup {
1771 fn from_proto(proto: count_workflow_executions_response::AggregationGroup) -> Self {
1772 Self { raw: proto }
1773 }
1774
1775 pub fn get<T: SearchAttributeValue>(&self, index: usize) -> Option<T> {
1780 self.try_get(index).ok().flatten()
1781 }
1782
1783 pub fn try_get<T: SearchAttributeValue>(
1789 &self,
1790 index: usize,
1791 ) -> Result<Option<T>, SearchAttributeError> {
1792 match self.raw.group_values.get(index) {
1793 Some(payload) => T::from_search_attribute_payload(payload).map(Some),
1794 None => Ok(None),
1795 }
1796 }
1797
1798 pub fn count(&self) -> usize {
1800 self.raw.count as usize
1801 }
1802}
1803
1804fn build_start_workflow_request(
1807 client: &impl NamespacedClient,
1808 workflow_type: String,
1809 input: Option<Payloads>,
1810 memo: Option<ProtoMemo>,
1811 options: WorkflowStartOptions,
1812) -> StartWorkflowExecutionRequest {
1813 let user_metadata = options.user_metadata();
1814 let request_eager_execution = options.enable_eager_workflow_start;
1815 StartWorkflowExecutionRequest {
1816 namespace: client.namespace(),
1817 input,
1818 workflow_id: options.workflow_id,
1819 workflow_type: Some(WorkflowType {
1820 name: workflow_type,
1821 }),
1822 task_queue: Some(TaskQueue {
1823 name: options.task_queue,
1824 kind: TaskQueueKind::Unspecified as i32,
1825 normal_name: String::new(),
1826 }),
1827 identity: client.identity(),
1828 request_id: Uuid::new_v4().to_string(),
1829 workflow_id_reuse_policy: ProtoWorkflowIdReusePolicy::from(options.id_reuse_policy) as i32,
1830 workflow_id_conflict_policy: ProtoWorkflowIdConflictPolicy::from(options.id_conflict_policy)
1831 as i32,
1832 workflow_execution_timeout: options
1833 .execution_timeout
1834 .and_then(|duration| duration.try_into().ok()),
1835 workflow_run_timeout: options
1836 .run_timeout
1837 .and_then(|duration| duration.try_into().ok()),
1838 workflow_task_timeout: options
1839 .task_timeout
1840 .and_then(|duration| duration.try_into().ok()),
1841 search_attributes: options
1842 .search_attributes
1843 .map(|attributes| attributes.into_proto()),
1844 cron_schedule: options.cron_schedule.unwrap_or_default(),
1845 request_eager_execution,
1846 retry_policy: options.retry_policy.map(Into::into),
1847 links: options.links,
1848 completion_callbacks: options.completion_callbacks,
1849 priority: Some(options.priority.into()),
1850 memo,
1851 header: options.header,
1852 user_metadata,
1853 ..Default::default()
1854 }
1855}
1856
1857impl<T> WorkflowClientTrait for T
1858where
1859 T: WorkflowService + NamespacedClient + Clone + Send + Sync + 'static,
1860{
1861 async fn start_workflow<W>(
1862 &self,
1863 workflow: W,
1864 input: W::Input,
1865 options: WorkflowStartOptions,
1866 ) -> Result<WorkflowHandle<Self, W>, WorkflowStartError>
1867 where
1868 W: HasWorkflowDefinition,
1869 W::Input: Send,
1870 {
1871 let namespace = self.namespace();
1872 let interceptor_output = interceptors::call_start_workflow(
1873 self.client_interceptors(),
1874 StartWorkflowInput::new(workflow.name().to_owned(), input, options),
1875 Next::new({
1876 let client = (*self).clone();
1877 move |input: StartWorkflowInput| -> BoxFuture<
1878 '_,
1879 Result<StartWorkflowOutput, WorkflowStartError>,
1880 > {
1881 let mut client = client;
1882 Box::pin(async move {
1883 let (workflow_type, args, options, rpc_options) = input.into_parts();
1884 let data_converter = client.data_converter().clone();
1885 let unencoded_payloads = {
1886 let payload_converter = data_converter.payload_converter();
1887 let context_data = SerializationContextData::Workflow(
1888 WorkflowSerializationContext::new(),
1889 );
1890 let context =
1891 SerializationContext::new(&context_data, payload_converter);
1892 args.serialize_payloads(&context)
1893 };
1894 drop(args);
1895
1896 let payloads = data_converter
1897 .codec()
1898 .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), unencoded_payloads?)
1899 .await?;
1900 let workflow_id = options.workflow_id.clone();
1901 let memo = options.encoded_memo(&data_converter).await?;
1902 let mut request = build_start_workflow_request(
1903 &client,
1904 workflow_type,
1905 payloads.into_payloads(),
1906 memo,
1907 options,
1908 )
1909 .into_request();
1910 rpc_options.apply_to(&mut request);
1911 let run_id = client
1912 .start_workflow_execution(request)
1913 .await
1914 .map_err(WorkflowStartError::from_status)?
1915 .into_inner()
1916 .run_id;
1917
1918 Ok(StartWorkflowOutput::new(workflow_id, run_id))
1919 })
1920 }
1921 }),
1922 )
1923 .await?;
1924 let StartWorkflowOutput {
1925 workflow_id,
1926 run_id,
1927 } = interceptor_output;
1928
1929 Ok(WorkflowHandle::new(
1930 self.clone(),
1931 WorkflowExecutionInfo {
1932 namespace,
1933 workflow_id,
1934 run_id: Some(run_id.clone()),
1935 first_execution_run_id: Some(run_id),
1936 },
1937 ))
1938 }
1939
1940 async fn signal_with_start_workflow<W, S>(
1941 &self,
1942 workflow: W,
1943 workflow_input: W::Input,
1944 signal: S,
1945 signal_input: S::Input,
1946 options: WorkflowStartOptions,
1947 ) -> Result<WorkflowHandle<Self, W>, WorkflowStartError>
1948 where
1949 W: HasWorkflowDefinition,
1950 W::Input: Send,
1951 S: SignalDefinition<Workflow = W::Run>,
1952 S::Input: Send,
1953 {
1954 let namespace = self.namespace();
1955 let interceptor_output = interceptors::call_signal_with_start_workflow(
1956 self.client_interceptors(),
1957 SignalWithStartWorkflowInput::new(
1958 workflow.name().to_owned(),
1959 workflow_input,
1960 signal.name().to_owned(),
1961 signal_input,
1962 options,
1963 ),
1964 Next::new({
1965 let client = (*self).clone();
1966 move |input: SignalWithStartWorkflowInput| -> BoxFuture<
1967 '_,
1968 Result<StartWorkflowOutput, WorkflowStartError>,
1969 > {
1970 let mut client = client;
1971 Box::pin(async move {
1972 let (
1973 workflow_type,
1974 workflow_args,
1975 signal_name,
1976 signal_args,
1977 options,
1978 rpc_options,
1979 ) = input.into_parts();
1980 let data_converter = client.data_converter().clone();
1981 let payload_converter = data_converter.payload_converter();
1982 let context_data = SerializationContextData::Workflow(
1983 WorkflowSerializationContext::new(),
1984 );
1985 let context = SerializationContext::new(&context_data, payload_converter);
1986 let workflow_payloads = workflow_args.serialize_payloads(&context);
1987 let signal_payloads = signal_args.serialize_payloads(&context);
1988 drop(workflow_args);
1989 drop(signal_args);
1990 let workflow_payloads = data_converter
1991 .codec()
1992 .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), workflow_payloads?)
1993 .await?;
1994 let signal_payloads = data_converter
1995 .codec()
1996 .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), signal_payloads?)
1997 .await?;
1998 let workflow_id = options.workflow_id.clone();
1999 let memo = options.encoded_memo(&data_converter).await?;
2000 let mut start_request = build_start_workflow_request(
2001 &client,
2002 workflow_type,
2003 workflow_payloads.into_payloads(),
2004 memo,
2005 options,
2006 );
2007 if let Some(task_queue) = &mut start_request.task_queue {
2008 task_queue.kind = TaskQueueKind::Normal as i32;
2009 }
2010 let mut request = SignalWithStartWorkflowExecutionRequest {
2011 namespace: start_request.namespace,
2012 workflow_id: start_request.workflow_id,
2013 workflow_type: start_request.workflow_type,
2014 task_queue: start_request.task_queue,
2015 input: start_request.input,
2016 workflow_execution_timeout: start_request.workflow_execution_timeout,
2017 workflow_run_timeout: start_request.workflow_run_timeout,
2018 workflow_task_timeout: start_request.workflow_task_timeout,
2019 identity: start_request.identity,
2020 request_id: start_request.request_id,
2021 workflow_id_reuse_policy: start_request.workflow_id_reuse_policy,
2022 workflow_id_conflict_policy: start_request.workflow_id_conflict_policy,
2023 signal_name,
2024 signal_input: Some(Payloads {
2025 payloads: signal_payloads,
2026 }),
2027 retry_policy: start_request.retry_policy,
2028 cron_schedule: start_request.cron_schedule,
2029 memo: start_request.memo,
2030 search_attributes: start_request.search_attributes,
2031 header: start_request.header,
2032 workflow_start_delay: start_request.workflow_start_delay,
2033 user_metadata: start_request.user_metadata,
2034 links: start_request.links,
2035 versioning_override: start_request.versioning_override,
2036 priority: start_request.priority,
2037 time_skipping_config: start_request.time_skipping_config,
2038 ..Default::default()
2039 }
2040 .into_request();
2041 rpc_options.apply_to(&mut request);
2042 let run_id = WorkflowService::signal_with_start_workflow_execution(
2043 &mut client,
2044 request,
2045 )
2046 .await?
2047 .into_inner()
2048 .run_id;
2049 Ok(StartWorkflowOutput::new(workflow_id, run_id))
2050 })
2051 }
2052 }),
2053 )
2054 .await?;
2055 let StartWorkflowOutput {
2056 workflow_id,
2057 run_id,
2058 } = interceptor_output;
2059
2060 Ok(WorkflowHandle::new(
2061 self.clone(),
2062 WorkflowExecutionInfo {
2063 namespace,
2064 workflow_id,
2065 run_id: Some(run_id.clone()),
2066 first_execution_run_id: Some(run_id),
2067 },
2068 ))
2069 }
2070
2071 async fn start_update_with_start_workflow<W, U>(
2072 &self,
2073 workflow: W,
2074 workflow_input: W::Input,
2075 update: U,
2076 update_input: U::Input,
2077 options: WorkflowUpdateWithStartOptions,
2078 ) -> Result<WorkflowUpdateHandle<Self, U::Output>, WorkflowUpdateWithStartError>
2079 where
2080 W: HasWorkflowDefinition,
2081 W::Input: Send,
2082 U: UpdateDefinition<Workflow = W::Run>,
2083 U::Input: Send,
2084 {
2085 let output = interceptors::call_update_with_start_workflow(
2086 self.client_interceptors(),
2087 UpdateWithStartWorkflowInput::new(
2088 workflow.name().to_owned(),
2089 workflow_input,
2090 update.name().to_owned(),
2091 update_input,
2092 options,
2093 ),
2094 Next::new({
2095 let client = (*self).clone();
2096 move |input: UpdateWithStartWorkflowInput| -> BoxFuture<
2097 '_,
2098 Result<UpdateWithStartWorkflowOutput, WorkflowUpdateWithStartError>,
2099 > {
2100 let mut client = client;
2101 Box::pin(async move {
2102 let UpdateWithStartWorkflowInput {
2103 workflow_type,
2104 update_name,
2105 options,
2106 rpc_options,
2107 workflow_args,
2108 update_args,
2109 } = input;
2110 let (start_options, update_id, update_header) = options.into_parts();
2111
2112 let data_converter = client.data_converter().clone();
2113 let (unencoded_workflow_payloads, unencoded_update_payloads) = {
2114 let payload_converter = data_converter.payload_converter();
2115 let context_data = SerializationContextData::Workflow(
2116 WorkflowSerializationContext::new(),
2117 );
2118 let context =
2119 SerializationContext::new(&context_data, payload_converter);
2120 (
2121 workflow_args.serialize_payloads(&context),
2122 update_args.serialize_payloads(&context),
2123 )
2124 };
2125 drop(workflow_args);
2126 drop(update_args);
2127 let (workflow_payloads, update_payloads) = try_join(
2130 data_converter.codec().encode(
2131 &SerializationContextData::Workflow(
2132 WorkflowSerializationContext::new(),
2133 ),
2134 unencoded_workflow_payloads?,
2135 ),
2136 data_converter.codec().encode(
2137 &SerializationContextData::Workflow(
2138 WorkflowSerializationContext::new(),
2139 ),
2140 unencoded_update_payloads?,
2141 ),
2142 )
2143 .await?;
2144
2145 let namespace = client.namespace();
2146 let workflow_id = start_options.workflow_id.clone();
2147 let memo = start_options.encoded_memo(&data_converter).await?;
2148 let start_request = build_start_workflow_request(
2149 &client,
2150 workflow_type,
2151 workflow_payloads.into_payloads(),
2152 memo,
2153 start_options,
2154 );
2155
2156 let update_id = update_id.unwrap_or_else(|| Uuid::new_v4().to_string());
2157 let update_request = workflow_handle::build_update_workflow_request(
2158 namespace.clone(),
2159 client.identity(),
2160 workflow_id.clone(),
2161 String::new(),
2162 update_id.clone(),
2163 update_name,
2164 update_header,
2165 update_payloads,
2166 );
2167
2168 let request = ExecuteMultiOperationRequest {
2169 namespace,
2170 operations: vec![
2171 execute_multi_operation_request::Operation {
2172 operation: Some(MultiOperationRequest::StartWorkflow(
2173 start_request,
2174 )),
2175 },
2176 execute_multi_operation_request::Operation {
2177 operation: Some(MultiOperationRequest::UpdateWorkflow(
2178 update_request,
2179 )),
2180 },
2181 ],
2182 resource_id: workflow_id.clone(),
2183 };
2184
2185 let (start_response, update_response) = loop {
2186 let mut rpc_request = request.clone().into_request();
2187 rpc_options.apply_to(&mut rpc_request);
2188 let response =
2189 WorkflowService::execute_multi_operation(&mut client, rpc_request)
2190 .await
2191 .map_err(WorkflowUpdateWithStartError::from_status)?
2192 .into_inner();
2193
2194 let [start_response, update_response]: [_; 2] =
2195 response.responses.try_into().map_err(|_| {
2196 WorkflowUpdateWithStartError::Other(
2197 "Server response did not include exactly two operation \
2198 responses"
2199 .into(),
2200 )
2201 })?;
2202 let (
2203 Some(MultiOperationResponse::StartWorkflow(start_response)),
2204 Some(MultiOperationResponse::UpdateWorkflow(update_response)),
2205 ) = (start_response.response, update_response.response)
2206 else {
2207 return Err(WorkflowUpdateWithStartError::Other(
2208 "Server response did not include start and update operation \
2209 responses in request order"
2210 .into(),
2211 ));
2212 };
2213
2214 if update_response.stage
2215 < UpdateWorkflowExecutionLifecycleStage::Accepted as i32
2216 {
2217 continue;
2218 }
2219 break (start_response, update_response);
2220 };
2221
2222 let run_id = update_response
2223 .update_ref
2224 .as_ref()
2225 .and_then(|reference| reference.workflow_execution.as_ref())
2226 .map(|execution| execution.run_id.clone())
2227 .filter(|run_id| !run_id.is_empty())
2228 .or_else(|| {
2229 (!start_response.run_id.is_empty()).then_some(start_response.run_id)
2230 });
2231 Ok(UpdateWithStartWorkflowOutput::new(
2232 workflow_id,
2233 update_id,
2234 run_id,
2235 update_response.outcome,
2236 ))
2237 })
2238 }
2239 }),
2240 )
2241 .await?;
2242 Ok(WorkflowUpdateHandle::new(
2243 self.clone(),
2244 output.update_id,
2245 output.workflow_id,
2246 output.run_id,
2247 output.known_outcome,
2248 ))
2249 }
2250
2251 async fn execute_update_with_start_workflow<W, U>(
2252 &self,
2253 workflow: W,
2254 workflow_input: W::Input,
2255 update: U,
2256 update_input: U::Input,
2257 options: WorkflowUpdateWithStartOptions,
2258 ) -> Result<U::Output, WorkflowUpdateWithStartError>
2259 where
2260 W: HasWorkflowDefinition,
2261 W::Input: Send,
2262 U: UpdateDefinition<Workflow = W::Run>,
2263 U::Input: Send,
2264 {
2265 let rpc_options = options.rpc_options.clone();
2266 let update_handle = WorkflowClientTrait::start_update_with_start_workflow(
2267 self,
2268 workflow,
2269 workflow_input,
2270 update,
2271 update_input,
2272 options,
2273 )
2274 .await?;
2275 let result = update_handle
2276 .get_result(rpc_options)
2277 .await
2278 .map_err(WorkflowUpdateWithStartError::Update)?;
2279 Ok(result)
2280 }
2281
2282 fn get_workflow_handle<W: HasWorkflowDefinition>(
2283 &self,
2284 workflow_id: impl Into<String>,
2285 ) -> WorkflowHandle<Self, W>
2286 where
2287 Self: Sized,
2288 {
2289 WorkflowHandle::new(
2290 self.clone(),
2291 WorkflowExecutionInfo {
2292 namespace: self.namespace(),
2293 workflow_id: workflow_id.into(),
2294 run_id: None,
2295 first_execution_run_id: None,
2296 },
2297 )
2298 }
2299
2300 fn list_workflows(
2301 &self,
2302 query: impl Into<String>,
2303 opts: WorkflowListOptions,
2304 ) -> ListWorkflowsStream {
2305 let client = self.clone();
2306 let namespace = self.namespace();
2307 let query = query.into();
2308 let limit = opts.limit;
2309 let rpc_options = opts.rpc_options;
2310
2311 let initial_state = (Vec::new(), VecDeque::new(), 0, false);
2313
2314 let stream = stream::unfold(
2315 initial_state,
2316 move |(next_page_token, mut buffer, mut yielded, exhausted)| {
2317 let client = client.clone();
2318 let namespace = namespace.clone();
2319 let query = query.clone();
2320 let rpc_options = rpc_options.clone();
2321
2322 async move {
2323 if let Some(l) = limit
2324 && yielded >= l
2325 {
2326 return None;
2327 }
2328
2329 if let Some(exec) = buffer.pop_front() {
2330 yielded += 1;
2331 return Some((Ok(exec), (next_page_token, buffer, yielded, exhausted)));
2332 }
2333
2334 if exhausted {
2335 return None;
2336 }
2337
2338 let response = interceptors::call_list_workflows_page(
2339 client.client_interceptors(),
2340 ListWorkflowsPageInput {
2341 query,
2342 next_page_token: next_page_token.clone(),
2343 rpc_options,
2344 },
2345 Next::new({
2346 let mut rpc_client = client.clone();
2347 move |input: ListWorkflowsPageInput| -> BoxFuture<
2348 '_,
2349 Result<ListWorkflowsPageOutput, ClientError>,
2350 > {
2351 Box::pin(async move {
2352 let mut request = ListWorkflowExecutionsRequest {
2353 namespace,
2354 page_size: 0,
2355 next_page_token: input.next_page_token,
2356 query: input.query,
2357 }
2358 .into_request();
2359 input.rpc_options.apply_to(&mut request);
2360 let response = WorkflowService::list_workflow_executions(
2361 &mut rpc_client,
2362 request,
2363 )
2364 .await?
2365 .into_inner();
2366 Ok(ListWorkflowsPageOutput::new(
2367 response.executions,
2368 response.next_page_token,
2369 ))
2370 })
2371 }
2372 }),
2373 )
2374 .await;
2375
2376 match response {
2377 Ok(mut output) => {
2378 let new_exhausted = output.next_page_token.is_empty();
2379 let new_token = output.next_page_token;
2380
2381 let data_converter = client.data_converter().clone();
2382 for execution in &mut output.executions {
2383 if let Some(memo) = execution.memo.as_mut()
2384 && let Err(err) = decode_payloads(
2385 memo,
2386 data_converter.codec(),
2387 &SerializationContextData::Workflow(
2388 WorkflowSerializationContext::new(),
2389 ),
2390 )
2391 .await
2392 {
2393 return Some((
2394 Err(ClientError::from(err)),
2395 (new_token, buffer, yielded, true),
2396 ));
2397 }
2398 }
2399 buffer = output
2400 .executions
2401 .into_iter()
2402 .map(|raw| {
2403 WorkflowExecution::new_with_data_converter(
2404 raw,
2405 data_converter.clone(),
2406 )
2407 })
2408 .collect();
2409
2410 if let Some(exec) = buffer.pop_front() {
2411 yielded += 1;
2412 Some((Ok(exec), (new_token, buffer, yielded, new_exhausted)))
2413 } else {
2414 None
2415 }
2416 }
2417 Err(e) => Some((Err(e), (next_page_token, buffer, yielded, true))),
2418 }
2419 }
2420 },
2421 );
2422
2423 ListWorkflowsStream::new(Box::pin(stream))
2424 }
2425
2426 async fn count_workflows(
2427 &self,
2428 query: impl Into<String>,
2429 opts: WorkflowCountOptions,
2430 ) -> Result<WorkflowExecutionCount, ClientError> {
2431 let output = interceptors::call_count_workflows(
2432 self.client_interceptors(),
2433 CountWorkflowsInput {
2434 query: query.into(),
2435 options: opts,
2436 },
2437 Next::new({
2438 let mut client = (*self).clone();
2439 move |input: CountWorkflowsInput| -> BoxFuture<
2440 '_,
2441 Result<CountWorkflowsOutput, ClientError>,
2442 > {
2443 Box::pin(async move {
2444 let mut request = CountWorkflowExecutionsRequest {
2445 namespace: client.namespace(),
2446 query: input.query,
2447 }
2448 .into_request();
2449 input.options.rpc_options.apply_to(&mut request);
2450 let response = WorkflowService::count_workflow_executions(
2451 &mut client,
2452 request,
2453 )
2454 .await?
2455 .into_inner();
2456 Ok(CountWorkflowsOutput::new(response))
2457 })
2458 }
2459 }),
2460 )
2461 .await?;
2462
2463 Ok(WorkflowExecutionCount::from_response(output.response))
2464 }
2465
2466 fn get_async_activity_handle(&self, identifier: ActivityIdentifier) -> AsyncActivityHandle<Self>
2467 where
2468 Self: Sized,
2469 {
2470 AsyncActivityHandle::new(self.clone(), identifier)
2471 }
2472
2473 async fn start_activity<A>(
2474 &self,
2475 activity: A,
2476 input: A::Input,
2477 options: ActivityStartOptions,
2478 ) -> Result<ActivityHandle<Self, A>, StartActivityError>
2479 where
2480 Self: Sized,
2481 A: ActivityDefinition,
2482 {
2483 let mut client = self.clone();
2484 let dc = client.data_converter();
2485 let sc = &SerializationContextData::Activity(ActivitySerializationContext::new());
2486
2487 let user_metadata = {
2488 let summary = match &options.summary {
2489 Some(summary) => Some(dc.to_payload(sc, summary).await?),
2490 None => None,
2491 };
2492 let details = match &options.static_details {
2493 Some(details) => Some(dc.to_payload(sc, details).await?),
2494 None => None,
2495 };
2496 (summary.is_some() || details.is_some()).then_some(UserMetadata { summary, details })
2497 };
2498
2499 let resp = client
2500 .start_activity_execution(
2501 StartActivityExecutionRequest {
2502 namespace: client.namespace(),
2503 identity: client.identity(),
2504 request_id: Uuid::new_v4().to_string(),
2505 activity_id: options.id.clone(),
2506 activity_type: Some(ActivityType {
2507 name: activity.name().to_string(),
2508 }),
2509 task_queue: Some(TaskQueue {
2510 name: options.task_queue,
2511 kind: TaskQueueKind::Normal.into(),
2512 normal_name: "".to_string(),
2513 }),
2514 schedule_to_close_timeout: try_into_or_box_err(
2515 options.close_timeouts.schedule_to_close(),
2516 StartActivityError::Other,
2517 )?,
2518 schedule_to_start_timeout: try_into_or_box_err(
2519 options.schedule_to_start_timeout,
2520 StartActivityError::Other,
2521 )?,
2522 start_to_close_timeout: try_into_or_box_err(
2523 options.close_timeouts.start_to_close(),
2524 StartActivityError::Other,
2525 )?,
2526 heartbeat_timeout: try_into_or_box_err(
2527 options.heartbeat_timeout,
2528 StartActivityError::Other,
2529 )?,
2530 retry_policy: options.retry_policy.map(Into::into),
2531 input: dc.to_payloads(sc, &input).await?.into_payloads(),
2532 id_reuse_policy: ProtoActivityIdReusePolicy::from(options.id_reuse_policy)
2533 .into(),
2534 id_conflict_policy: ProtoActivityIdConflictPolicy::from(
2535 options.id_conflict_policy,
2536 )
2537 .into(),
2538 search_attributes: options.search_attributes.map(SearchAttributes::into_proto),
2539 header: options.header,
2540 user_metadata,
2541 priority: Some(options.priority.into()),
2542 start_delay: try_into_or_box_err(
2543 options.start_delay,
2544 StartActivityError::Other,
2545 )?,
2546 ..Default::default()
2547 }
2548 .into_request(),
2549 )
2550 .await?
2551 .into_inner();
2552
2553 Ok(ActivityHandle::new(
2554 client,
2555 options.id,
2556 (!resp.run_id.is_empty()).then_some(resp.run_id),
2557 ))
2558 }
2559
2560 fn get_activity_handle<A>(
2561 &self,
2562 _activity: A,
2563 id: impl Into<String>,
2564 run_id: Option<String>,
2565 ) -> ActivityHandle<Self, A>
2566 where
2567 Self: Sized,
2568 A: ActivityDefinition,
2569 {
2570 ActivityHandle::new(self.clone(), id.into(), run_id)
2571 }
2572
2573 fn get_untyped_activity_handle(
2574 &self,
2575 id: impl Into<String>,
2576 run_id: Option<String>,
2577 ) -> ActivityHandle<Self, UntypedActivity>
2578 where
2579 Self: Sized,
2580 {
2581 ActivityHandle::new(self.clone(), id.into(), run_id)
2582 }
2583
2584 fn list_activities(
2585 &self,
2586 query: impl Into<String>,
2587 _options: ActivityListOptions,
2588 ) -> ListActivitiesStream {
2589 let client = self.clone();
2590 let namespace = client.namespace();
2591 let query = query.into();
2592
2593 ListActivitiesStream::new(stream::unfold(
2594 Some(vec![]), move |next_page_token| {
2596 let mut client = client.clone();
2597 let namespace = namespace.clone();
2598 let query = query.clone();
2599
2600 async move {
2601 #[allow(clippy::question_mark)]
2603 let Some(token): Option<Vec<u8>> = next_page_token else {
2604 return None;
2605 };
2606
2607 match WorkflowService::list_activity_executions(
2608 &mut client,
2609 ListActivityExecutionsRequest {
2610 namespace,
2611 page_size: 0, next_page_token: token.clone(),
2613 query,
2614 }
2615 .into_request(),
2616 )
2617 .await
2618 .map(|r| r.into_inner())
2619 {
2620 Ok(resp) => Some((
2621 Ok(resp.executions),
2622 (!resp.next_page_token.is_empty()).then_some(resp.next_page_token),
2623 )),
2624 Err(e) => Some((Err(e.into()), Some(token))),
2625 }
2626 }
2627 },
2628 ))
2629 }
2630
2631 async fn count_activities(
2632 &self,
2633 query: impl Into<String>,
2634 _options: ActivityCountOptions,
2635 ) -> Result<ActivityExecutionCount, ClientError> {
2636 let mut client = self.clone();
2637 let resp = client
2638 .count_activity_executions(
2639 CountActivityExecutionsRequest {
2640 namespace: client.namespace(),
2641 query: query.into(),
2642 }
2643 .into_request(),
2644 )
2645 .await?
2646 .into_inner();
2647 Ok(ActivityExecutionCount::from_response(resp))
2648 }
2649}
2650
2651macro_rules! dbg_panic {
2652 ($($arg:tt)*) => {
2653 use tracing::error;
2654 error!($($arg)*);
2655 debug_assert!(false, $($arg)*);
2656 };
2657}
2658pub(crate) use dbg_panic;
2659
2660fn try_into_or_box_err<A, B, E, MapErr>(val: Option<A>, map_err: MapErr) -> Result<Option<B>, E>
2661where
2662 A: TryInto<B>,
2663 <A as TryInto<B>>::Error: Error + Send + Sync + 'static,
2664 MapErr: FnOnce(Box<dyn Error + Send + Sync + 'static>) -> E,
2665{
2666 val.map(TryInto::try_into)
2667 .transpose()
2668 .map_err(|e| map_err(Box::from(e)))
2669}
2670
2671#[cfg(test)]
2672mod tests {
2673 use super::*;
2674 use crate::callback_based::CallbackBasedGrpcService;
2675 use std::{
2676 sync::atomic::{AtomicUsize, Ordering},
2677 time::Instant,
2678 };
2679 use temporalio_common::search_attributes::SearchAttributeKey;
2680 use tonic::{Status, metadata::Ascii};
2681 use url::Url;
2682
2683 #[test]
2684 fn count_aggregation_group_gets_typed_value() {
2685 let attrs = SearchAttributes::new([SearchAttributeKey::int("group").value_set(42)]);
2686 let group = WorkflowCountAggregationGroup {
2687 raw: count_workflow_executions_response::AggregationGroup {
2688 group_values: vec![attrs.raw_payload("group").unwrap().clone()],
2689 count: 1,
2690 },
2691 };
2692
2693 assert_eq!(group.get::<i64>(0), Some(42));
2694 assert_eq!(group.get::<i64>(1), None);
2695 assert!(group.try_get::<String>(0).is_err());
2696 assert_eq!(group.try_get::<i64>(1).unwrap(), None);
2697 }
2698
2699 fn connection_options_for_system_info_test(
2700 service_override: CallbackBasedGrpcService,
2701 ) -> ConnectionOptions {
2702 ConnectionOptions::new(Url::parse("http://localhost:7233").unwrap())
2703 .service_override(service_override)
2704 .dns_load_balancing(None)
2705 .build()
2706 }
2707
2708 #[test]
2709 fn applies_headers() {
2710 let headers = Arc::new(RwLock::new(ClientHeaders {
2712 user_headers: HashMap::new(),
2713 user_binary_headers: HashMap::new(),
2714 api_key: Some("my-api-key".to_owned()),
2715 }));
2716 headers.clone().write().user_headers.insert(
2717 "my-meta-key".parse().unwrap(),
2718 "my-meta-val".parse().unwrap(),
2719 );
2720 headers.clone().write().user_binary_headers.insert(
2721 "my-bin-meta-key-bin".parse().unwrap(),
2722 vec![1, 2, 3].try_into().unwrap(),
2723 );
2724 let mut interceptor = ServiceCallInterceptor {
2725 client_name: "cute-kitty".to_string(),
2726 client_version: "0.1.0".to_string(),
2727 headers: headers.clone(),
2728 };
2729
2730 let req = interceptor.call(tonic::Request::new(())).unwrap();
2732 assert_eq!(req.metadata().get("my-meta-key").unwrap(), "my-meta-val");
2733 assert_eq!(
2734 req.metadata().get("authorization").unwrap(),
2735 "Bearer my-api-key"
2736 );
2737 assert_eq!(
2738 req.metadata().get_bin("my-bin-meta-key-bin").unwrap(),
2739 vec![1, 2, 3].as_slice()
2740 );
2741
2742 let mut req = tonic::Request::new(());
2744 req.metadata_mut()
2745 .insert("my-meta-key", "my-meta-val2".parse().unwrap());
2746 req.metadata_mut()
2747 .insert("authorization", "my-api-key2".parse().unwrap());
2748 req.metadata_mut()
2749 .insert_bin("my-bin-meta-key-bin", vec![4, 5, 6].try_into().unwrap());
2750 let req = interceptor.call(req).unwrap();
2751 assert_eq!(req.metadata().get("my-meta-key").unwrap(), "my-meta-val2");
2752 assert_eq!(req.metadata().get("authorization").unwrap(), "my-api-key2");
2753 assert_eq!(
2754 req.metadata().get_bin("my-bin-meta-key-bin").unwrap(),
2755 vec![4, 5, 6].as_slice()
2756 );
2757
2758 headers.clone().write().user_headers.insert(
2760 "authorization".parse().unwrap(),
2761 "my-api-key3".parse().unwrap(),
2762 );
2763 let req = interceptor.call(tonic::Request::new(())).unwrap();
2764 assert_eq!(req.metadata().get("my-meta-key").unwrap(), "my-meta-val");
2765 assert_eq!(req.metadata().get("authorization").unwrap(), "my-api-key3");
2766
2767 headers.clone().write().user_headers.clear();
2769 headers.clone().write().user_binary_headers.clear();
2770 headers.clone().write().api_key.take();
2771 let req = interceptor.call(tonic::Request::new(())).unwrap();
2772 assert!(!req.metadata().contains_key("my-meta-key"));
2773 assert!(!req.metadata().contains_key("authorization"));
2774 assert!(!req.metadata().contains_key("my-bin-meta-key-bin"));
2775
2776 let mut req = tonic::Request::new(());
2778 req.metadata_mut()
2779 .insert("grpc-timeout", "1S".parse().unwrap());
2780 let req = interceptor.call(req).unwrap();
2781 assert_eq!(
2782 req.metadata().get("grpc-timeout").unwrap(),
2783 "1S".parse::<MetadataValue<Ascii>>().unwrap()
2784 );
2785 }
2786
2787 #[test]
2788 fn invalid_ascii_header_key() {
2789 let invalid_headers = {
2790 let mut h = HashMap::new();
2791 h.insert("x-binary-key-bin".to_owned(), "value".to_owned());
2792 h
2793 };
2794
2795 let result = parse_ascii_headers(invalid_headers);
2796 assert!(result.is_err());
2797 assert_eq!(
2798 result.err().unwrap().to_string(),
2799 "Invalid ASCII header key 'x-binary-key-bin': invalid gRPC metadata key name"
2800 );
2801 }
2802
2803 #[test]
2804 fn invalid_ascii_header_value() {
2805 let invalid_headers = {
2806 let mut h = HashMap::new();
2807 h.insert("x-ascii-key".to_owned(), "\x00value".to_owned());
2809 h
2810 };
2811
2812 let result = parse_ascii_headers(invalid_headers);
2813 assert!(result.is_err());
2814 assert_eq!(
2815 result.err().unwrap().to_string(),
2816 "Invalid ASCII header value for key 'x-ascii-key': failed to parse metadata value"
2817 );
2818 }
2819
2820 #[test]
2821 fn invalid_binary_header_key() {
2822 let invalid_headers = {
2823 let mut h = HashMap::new();
2824 h.insert("x-ascii-key".to_owned(), vec![1, 2, 3]);
2825 h
2826 };
2827
2828 let result = parse_binary_headers(invalid_headers);
2829 assert!(result.is_err());
2830 assert_eq!(
2831 result.err().unwrap().to_string(),
2832 "Invalid binary header key 'x-ascii-key': invalid gRPC metadata key name"
2833 );
2834 }
2835
2836 #[test]
2837 fn keep_alive_defaults() {
2838 let opts = ConnectionOptions::new(Url::parse("https://smolkitty").unwrap())
2839 .identity("enchicat".to_string())
2840 .client_name("cute-kitty".to_string())
2841 .client_version("0.1.0".to_string())
2842 .build();
2843 assert_eq!(
2844 opts.keep_alive.clone().unwrap().interval,
2845 ClientKeepAliveOptions::default().interval
2846 );
2847 assert_eq!(
2848 opts.keep_alive.clone().unwrap().timeout,
2849 ClientKeepAliveOptions::default().timeout
2850 );
2851
2852 let opts = ConnectionOptions::new(Url::parse("https://smolkitty").unwrap())
2854 .identity("enchicat".to_string())
2855 .client_name("cute-kitty".to_string())
2856 .client_version("0.1.0".to_string())
2857 .keep_alive(None)
2858 .build();
2859 dbg!(&opts.keep_alive);
2860 assert!(opts.keep_alive.is_none());
2861 }
2862
2863 #[rstest::rstest]
2864 #[case(
2865 "unknown method GetSystemInfo for service temporal.api.workflowservice.v1.WorkflowService"
2866 )]
2867 #[case("Method temporal.api.workflowservice.v1.WorkflowService/GetSystemInfo is unimplemented")]
2868 #[case(
2869 "The server does not implement the method /temporal.api.workflowservice.v1.WorkflowService/GetSystemInfo"
2870 )]
2871 #[tokio::test]
2872 async fn get_system_info_missing_method_falls_back_to_empty_capabilities(
2873 #[case] message: &'static str,
2874 ) {
2875 let attempts = Arc::new(AtomicUsize::new(0));
2876 let attempts_clone = attempts.clone();
2877 let service_override = CallbackBasedGrpcService {
2878 callback: Arc::new(move |req| {
2879 let attempts = attempts_clone.clone();
2880 Box::pin(async move {
2881 assert_eq!(req.rpc, "GetSystemInfo");
2882 attempts.fetch_add(1, Ordering::SeqCst);
2883 Err(Status::unimplemented(message))
2884 })
2885 }),
2886 };
2887
2888 let connection =
2889 Connection::connect(connection_options_for_system_info_test(service_override))
2890 .await
2891 .unwrap();
2892
2893 assert!(connection.capabilities().is_none());
2894 assert_eq!(attempts.load(Ordering::SeqCst), 1);
2895 }
2896
2897 #[tokio::test]
2898 async fn get_system_info_non_missing_unimplemented_fails_connect() {
2899 let attempts = Arc::new(AtomicUsize::new(0));
2900 let attempts_clone = attempts.clone();
2901 let service_override = CallbackBasedGrpcService {
2902 callback: Arc::new(move |req| {
2903 let attempts = attempts_clone.clone();
2904 Box::pin(async move {
2905 assert_eq!(req.rpc, "GetSystemInfo");
2906 attempts.fetch_add(1, Ordering::SeqCst);
2907 Err(Status::unimplemented("backend temporarily unimplemented"))
2908 })
2909 }),
2910 };
2911
2912 let err =
2913 match Connection::connect(connection_options_for_system_info_test(service_override))
2914 .await
2915 {
2916 Ok(_) => panic!("connection should fail"),
2917 Err(err) => err,
2918 };
2919
2920 assert!(matches!(
2921 err,
2922 ClientConnectError::SystemInfoCallError(status)
2923 if status.code() == Code::Unimplemented
2924 && status.message() == "backend temporarily unimplemented"
2925 ));
2926 assert_eq!(attempts.load(Ordering::SeqCst), 1);
2927 }
2928
2929 #[tokio::test]
2930 async fn connect_timeout_bounds_connection_attempt() {
2931 let url = Url::parse("http://10.255.255.1:7233").unwrap();
2932 let opts = ConnectionOptions::new(url)
2933 .connect_timeout(Duration::from_millis(500))
2934 .build();
2935 let start = Instant::now();
2936 let result = Connection::connect(opts).await;
2937 assert!(result.is_err(), "connection should fail");
2938 assert!(start.elapsed() < Duration::from_secs(2));
2939 }
2940
2941 mod tls_custom_verifier_tests {
2942 use super::*;
2943 use tokio_rustls::rustls::{
2944 DigitallySignedStruct, Error as RustlsError, SignatureScheme,
2945 client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier},
2946 pki_types::{CertificateDer, ServerName, UnixTime},
2947 };
2948
2949 #[derive(Debug)]
2952 struct MockVerifier;
2953
2954 impl ServerCertVerifier for MockVerifier {
2955 fn verify_server_cert(
2956 &self,
2957 _end_entity: &CertificateDer<'_>,
2958 _intermediates: &[CertificateDer<'_>],
2959 _server_name: &ServerName<'_>,
2960 _ocsp_response: &[u8],
2961 _now: UnixTime,
2962 ) -> Result<ServerCertVerified, RustlsError> {
2963 Ok(ServerCertVerified::assertion())
2964 }
2965
2966 fn verify_tls12_signature(
2967 &self,
2968 _message: &[u8],
2969 _cert: &CertificateDer<'_>,
2970 _dss: &DigitallySignedStruct,
2971 ) -> Result<HandshakeSignatureValid, RustlsError> {
2972 Ok(HandshakeSignatureValid::assertion())
2973 }
2974
2975 fn verify_tls13_signature(
2976 &self,
2977 _message: &[u8],
2978 _cert: &CertificateDer<'_>,
2979 _dss: &DigitallySignedStruct,
2980 ) -> Result<HandshakeSignatureValid, RustlsError> {
2981 Ok(HandshakeSignatureValid::assertion())
2982 }
2983
2984 fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
2985 vec![
2986 SignatureScheme::ECDSA_NISTP256_SHA256,
2987 SignatureScheme::RSA_PSS_SHA256,
2988 ]
2989 }
2990 }
2991
2992 #[tokio::test]
2993 async fn add_tls_to_channel_with_custom_verifier() {
2994 let tls_opts = TlsOptions::builder()
2995 .server_cert_verifier(Arc::new(MockVerifier))
2996 .domain("test.temporal.io".to_string())
2997 .build();
2998 let endpoint = tonic::transport::Channel::from_static("https://test.temporal.io:7233");
2999 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3000 assert!(
3001 matches!(&result, Ok(TlsConfigResult::Standard(_))),
3002 "add_tls_to_channel should succeed with a custom verifier: {:?}",
3003 result.err()
3004 );
3005 }
3006
3007 #[tokio::test]
3008 async fn add_tls_to_channel_with_verifier_and_ca_cert_fails() {
3009 let tls_opts = TlsOptions::builder()
3012 .server_root_ca_cert(b"some-ca-cert-bytes".to_vec())
3013 .server_cert_verifier(Arc::new(MockVerifier))
3014 .domain("test.temporal.io".to_string())
3015 .build();
3016 let endpoint = tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3017 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3018 assert!(
3019 matches!(result, Err(ClientConnectError::InvalidConfig(_))),
3020 "add_tls_to_channel should fail with InvalidConfig when both CA cert and verifier are set: {:?}",
3021 result
3022 );
3023 }
3024
3025 #[tokio::test]
3026 async fn add_tls_to_channel_without_verifier_still_works() {
3027 let tls_opts = TlsOptions::builder()
3029 .domain("test.temporal.io".to_string())
3030 .build();
3031 let endpoint = tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3032 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3033 assert!(
3034 matches!(&result, Ok(TlsConfigResult::Standard(_))),
3035 "add_tls_to_channel should succeed without a verifier (native roots): {:?}",
3036 result.err()
3037 );
3038 }
3039
3040 #[cfg(feature = "dynamic-tls")]
3043 mod dynamic_cert_tests {
3044 use super::*;
3045
3046 #[derive(Debug)]
3049 struct MockClientCertResolver;
3050
3051 impl tokio_rustls::rustls::client::ResolvesClientCert for MockClientCertResolver {
3052 fn resolve(
3053 &self,
3054 _acceptable_issuers: &[&[u8]],
3055 _sigschemes: &[tokio_rustls::rustls::SignatureScheme],
3056 ) -> Option<Arc<tokio_rustls::rustls::sign::CertifiedKey>> {
3057 None }
3059
3060 fn has_certs(&self) -> bool {
3061 false
3062 }
3063 }
3064
3065 #[tokio::test]
3066 async fn add_tls_with_client_cert_resolver_returns_custom_connector() {
3067 let resolver = Arc::new(MockClientCertResolver);
3068 let tls_opts = TlsOptions {
3069 client_cert_resolver: Some(resolver),
3070 domain: Some("test.temporal.io".to_string()),
3071 ..Default::default()
3072 };
3073 let endpoint =
3074 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3075 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3076 match result {
3077 Ok(TlsConfigResult::CustomConnector {
3078 domain,
3079 rustls_config,
3080 ..
3081 }) => {
3082 assert_eq!(domain, "test.temporal.io");
3083 assert_eq!(rustls_config.alpn_protocols, vec![b"h2".to_vec()]);
3085 }
3086 other => panic!(
3087 "Expected TlsConfigResult::CustomConnector, got {:?}",
3088 other.err()
3089 ),
3090 }
3091 }
3092
3093 #[tokio::test]
3094 async fn add_tls_with_client_cert_resolver_inherits_domain_from_endpoint() {
3095 let resolver = Arc::new(MockClientCertResolver);
3096 let tls_opts = TlsOptions {
3097 client_cert_resolver: Some(resolver),
3098 ..Default::default()
3100 };
3101 let endpoint =
3102 tonic::transport::Channel::from_static("https://my-server.example.com:7233");
3103 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3104 match result {
3105 Ok(TlsConfigResult::CustomConnector { domain, .. }) => {
3106 assert_eq!(domain, "my-server.example.com");
3107 }
3108 other => panic!(
3109 "Expected TlsConfigResult::CustomConnector, got {:?}",
3110 other.err()
3111 ),
3112 }
3113 }
3114
3115 #[tokio::test]
3116 async fn add_tls_with_resolver_and_custom_verifier() {
3117 let resolver = Arc::new(MockClientCertResolver);
3118 let tls_opts = TlsOptions {
3119 client_cert_resolver: Some(resolver),
3120 server_cert_verifier: Some(Arc::new(MockVerifier)),
3121 domain: Some("test.temporal.io".to_string()),
3122 ..Default::default()
3123 };
3124 let endpoint =
3125 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3126 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3127 assert!(
3128 matches!(&result, Ok(TlsConfigResult::CustomConnector { .. })),
3129 "Should succeed when combining cert resolver with custom server verifier: {:?}",
3130 result.err()
3131 );
3132 }
3133
3134 #[tokio::test]
3135 async fn add_tls_with_resolver_and_custom_ca_cert() {
3136 let ca_pem = include_bytes!("../tests/testdata/ca.pem");
3138 let resolver = Arc::new(MockClientCertResolver);
3139 let tls_opts = TlsOptions {
3140 client_cert_resolver: Some(resolver),
3141 server_root_ca_cert: Some(ca_pem.to_vec()),
3142 domain: Some("test.temporal.io".to_string()),
3143 ..Default::default()
3144 };
3145 let endpoint =
3146 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3147 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3148 assert!(
3149 matches!(&result, Ok(TlsConfigResult::CustomConnector { .. })),
3150 "Should succeed when combining cert resolver with custom CA cert: {:?}",
3151 result.err()
3152 );
3153 }
3154
3155 #[tokio::test]
3156 async fn add_tls_both_static_and_dynamic_client_cert_fails() {
3157 let resolver = Arc::new(MockClientCertResolver);
3158 let tls_opts = TlsOptions {
3159 client_tls_options: Some(ClientTlsOptions {
3160 client_cert: b"some-cert".to_vec(),
3161 client_private_key: b"some-key".to_vec(),
3162 }),
3163 client_cert_resolver: Some(resolver),
3164 domain: Some("test.temporal.io".to_string()),
3165 ..Default::default()
3166 };
3167 let endpoint =
3168 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3169 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3170 assert!(
3171 matches!(result, Err(ClientConnectError::InvalidConfig(msg)) if msg.contains("client_tls_options") && msg.contains("client_cert_resolver")),
3172 "Should fail with InvalidConfig when both static and dynamic client certs are set"
3173 );
3174 }
3175
3176 #[tokio::test]
3177 async fn add_tls_no_options_returns_standard_passthrough() {
3178 let endpoint = tonic::transport::Channel::from_static("http://localhost:7233");
3179 let result = add_tls_to_channel(None, endpoint).await;
3180 assert!(
3181 matches!(&result, Ok(TlsConfigResult::Standard(_))),
3182 "Should return Standard when no TLS options are set"
3183 );
3184 }
3185
3186 #[test]
3187 fn build_custom_rustls_config_with_resolver() {
3188 let resolver = Arc::new(MockClientCertResolver);
3189 let tls_opts = TlsOptions {
3190 domain: Some("test.temporal.io".to_string()),
3191 ..Default::default()
3192 };
3193 let config = build_custom_rustls_config(&tls_opts, Some(resolver));
3194 assert!(config.is_ok(), "Should build config: {:?}", config.err());
3195 let config = config.unwrap();
3196 assert_eq!(config.alpn_protocols, vec![b"h2".to_vec()]);
3197 }
3198
3199 #[test]
3200 fn build_custom_rustls_config_without_resolver() {
3201 let tls_opts = TlsOptions {
3202 domain: Some("test.temporal.io".to_string()),
3203 ..Default::default()
3204 };
3205 let config = build_custom_rustls_config(&tls_opts, None);
3206 assert!(config.is_ok(), "Should build config: {:?}", config.err());
3207 }
3208
3209 #[test]
3210 fn build_custom_rustls_config_with_custom_verifier_and_resolver() {
3211 let resolver = Arc::new(MockClientCertResolver);
3212 let tls_opts = TlsOptions {
3213 server_cert_verifier: Some(Arc::new(MockVerifier)),
3214 domain: Some("test.temporal.io".to_string()),
3215 ..Default::default()
3216 };
3217 let config = build_custom_rustls_config(&tls_opts, Some(resolver));
3218 assert!(
3219 config.is_ok(),
3220 "Should build config with custom verifier + resolver: {:?}",
3221 config.err()
3222 );
3223 }
3224
3225 #[test]
3226 fn tls_options_debug_shows_custom_for_resolver() {
3227 let resolver = Arc::new(MockClientCertResolver);
3228 let tls_opts = TlsOptions {
3229 client_cert_resolver: Some(resolver),
3230 ..Default::default()
3231 };
3232 let debug_str = format!("{:?}", tls_opts);
3233 assert!(
3234 debug_str.contains("\"<custom>\""),
3235 "Debug should show <custom> for client_cert_resolver: {debug_str}"
3236 );
3237 assert!(
3238 debug_str.contains("client_cert_resolver"),
3239 "Debug should contain field name: {debug_str}"
3240 );
3241 }
3242
3243 #[test]
3244 fn tls_options_default_has_no_resolver() {
3245 let tls_opts = TlsOptions::default();
3246 assert!(tls_opts.client_cert_resolver.is_none());
3247 assert!(tls_opts.client_tls_options.is_none());
3248 assert!(tls_opts.server_cert_verifier.is_none());
3249 }
3250
3251 #[tokio::test]
3252 async fn add_tls_resolver_with_ip_host_uses_ip_as_domain() {
3253 let resolver = Arc::new(MockClientCertResolver);
3256 let tls_opts = TlsOptions {
3257 client_cert_resolver: Some(resolver),
3258 ..Default::default()
3260 };
3261 let endpoint = tonic::transport::Channel::from_static("https://192.168.1.100:7233");
3262 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3263 match result {
3264 Ok(TlsConfigResult::CustomConnector { domain, .. }) => {
3265 assert_eq!(domain, "192.168.1.100");
3266 }
3267 other => panic!(
3268 "Expected CustomConnector with IP domain, got {:?}",
3269 other.err()
3270 ),
3271 }
3272 }
3273 }
3274 }
3275
3276 mod start_workflow_interceptor_tests {
3277 use super::*;
3278 use crate::{request_extensions::RetryConfigForCall, test_helpers::XorCodec};
3279 use parking_lot::Mutex;
3280 use std::sync::atomic::{AtomicUsize, Ordering};
3281 use temporalio_common::{
3282 MemoValues, SignalDefinition,
3283 data_converters::{
3284 DefaultFailureConverter, PayloadCodec, PayloadConversionError, PayloadConverter,
3285 SerializationContext, SerializationContextData, TemporalDeserializable,
3286 TemporalSerializable,
3287 },
3288 protos::temporal::api::common::v1::{
3289 Link, Memo as ProtoMemo, Payload, Priority as ProtoPriority,
3290 },
3291 };
3292 use temporalio_macros::{workflow, workflow_methods};
3293 use temporalio_workflow::{SyncWorkflowContext, WorkflowContext, WorkflowResult};
3294 use tonic::{Request, Response};
3295
3296 #[workflow]
3297 #[derive(Default)]
3298 struct TestWorkflow;
3299
3300 #[workflow_methods]
3301 impl TestWorkflow {
3302 #[run]
3303 async fn run(
3304 _ctx: &mut WorkflowContext<Self>,
3305 _input: Vec<String>,
3306 ) -> WorkflowResult<()> {
3307 Ok(())
3308 }
3309
3310 #[signal]
3311 fn test_signal(&mut self, _ctx: &mut SyncWorkflowContext<Self>, _input: Vec<String>) {}
3312 }
3313
3314 #[derive(Default)]
3315 struct RecordedStart {
3316 calls: usize,
3317 workflow_type: String,
3318 memo: Option<ProtoMemo>,
3319 payloads: Vec<Payload>,
3320 signal_name: String,
3321 signal_payloads: Vec<Payload>,
3322 identity: String,
3323 links: Vec<Link>,
3324 priority: Option<ProtoPriority>,
3325 ascii_metadata: Option<String>,
3326 binary_metadata: Option<Vec<u8>>,
3327 grpc_timeout: Option<String>,
3328 retry_options: Option<RetryOptions>,
3329 }
3330
3331 struct CountingCodec {
3332 encode_calls: Arc<AtomicUsize>,
3333 }
3334
3335 impl PayloadCodec for CountingCodec {
3336 fn encode(
3337 &self,
3338 _context: &SerializationContextData,
3339 payloads: Vec<Payload>,
3340 ) -> futures_util::future::BoxFuture<
3341 'static,
3342 Result<Vec<Payload>, PayloadConversionError>,
3343 > {
3344 self.encode_calls.fetch_add(1, Ordering::SeqCst);
3345 Box::pin(async move { Ok(payloads) })
3346 }
3347
3348 fn decode(
3349 &self,
3350 _context: &SerializationContextData,
3351 payloads: Vec<Payload>,
3352 ) -> futures_util::future::BoxFuture<
3353 'static,
3354 Result<Vec<Payload>, PayloadConversionError>,
3355 > {
3356 Box::pin(async move { Ok(payloads) })
3357 }
3358 }
3359
3360 #[derive(Clone)]
3361 struct MockStartWorkflowClient {
3362 recorded: Arc<Mutex<RecordedStart>>,
3363 data_converter: DataConverter,
3364 }
3365
3366 impl NamespacedClient for MockStartWorkflowClient {
3367 fn namespace(&self) -> String {
3368 "test-namespace".to_owned()
3369 }
3370
3371 fn identity(&self) -> String {
3372 "test-identity".to_owned()
3373 }
3374
3375 fn data_converter(&self) -> &DataConverter {
3376 &self.data_converter
3377 }
3378 }
3379
3380 impl WorkflowService for MockStartWorkflowClient {
3381 fn start_workflow_execution(
3382 &mut self,
3383 request: Request<StartWorkflowExecutionRequest>,
3384 ) -> futures_util::future::BoxFuture<
3385 '_,
3386 Result<Response<StartWorkflowExecutionResponse>, tonic::Status>,
3387 > {
3388 let ascii_metadata = request
3389 .metadata()
3390 .get("call-meta")
3391 .map(|value| value.to_str().unwrap().to_owned());
3392 let binary_metadata = request
3393 .metadata()
3394 .get_bin("call-meta-bin")
3395 .map(|value| value.to_bytes().unwrap().to_vec());
3396 let grpc_timeout = request
3397 .metadata()
3398 .get("grpc-timeout")
3399 .map(|value| value.to_str().unwrap().to_owned());
3400 let retry_options = request
3401 .extensions()
3402 .get::<RetryConfigForCall>()
3403 .map(|config| config.0.clone());
3404 let request = request.into_inner();
3405 let mut recorded = self.recorded.lock();
3406 recorded.calls += 1;
3407 recorded.workflow_type = request.workflow_type.unwrap().name;
3408 recorded.memo = request.memo;
3409 recorded.payloads = request.input.unwrap_or_default().payloads;
3410 recorded.identity = request.identity;
3411 recorded.links = request.links;
3412 recorded.priority = request.priority;
3413 recorded.ascii_metadata = ascii_metadata;
3414 recorded.binary_metadata = binary_metadata;
3415 recorded.grpc_timeout = grpc_timeout;
3416 recorded.retry_options = retry_options;
3417
3418 Box::pin(async {
3419 Ok(Response::new(StartWorkflowExecutionResponse {
3420 run_id: "server-run-id".to_owned(),
3421 ..Default::default()
3422 }))
3423 })
3424 }
3425
3426 fn signal_with_start_workflow_execution(
3427 &mut self,
3428 request: Request<SignalWithStartWorkflowExecutionRequest>,
3429 ) -> futures_util::future::BoxFuture<
3430 '_,
3431 Result<Response<SignalWithStartWorkflowExecutionResponse>, tonic::Status>,
3432 > {
3433 let ascii_metadata = request
3434 .metadata()
3435 .get("call-meta")
3436 .map(|value| value.to_str().unwrap().to_owned());
3437 let binary_metadata = request
3438 .metadata()
3439 .get_bin("call-meta-bin")
3440 .map(|value| value.to_bytes().unwrap().to_vec());
3441 let grpc_timeout = request
3442 .metadata()
3443 .get("grpc-timeout")
3444 .map(|value| value.to_str().unwrap().to_owned());
3445 let retry_options = request
3446 .extensions()
3447 .get::<RetryConfigForCall>()
3448 .map(|config| config.0.clone());
3449 let request = request.into_inner();
3450 let mut recorded = self.recorded.lock();
3451 recorded.calls += 1;
3452 recorded.workflow_type = request.workflow_type.unwrap().name;
3453 recorded.memo = request.memo;
3454 recorded.payloads = request.input.unwrap_or_default().payloads;
3455 recorded.signal_name = request.signal_name;
3456 recorded.signal_payloads = request.signal_input.unwrap_or_default().payloads;
3457 recorded.identity = request.identity;
3458 recorded.links = request.links;
3459 recorded.priority = request.priority;
3460 recorded.ascii_metadata = ascii_metadata;
3461 recorded.binary_metadata = binary_metadata;
3462 recorded.grpc_timeout = grpc_timeout;
3463 recorded.retry_options = retry_options;
3464
3465 Box::pin(async {
3466 Ok(Response::new(SignalWithStartWorkflowExecutionResponse {
3467 run_id: "signal-server-run-id".to_owned(),
3468 ..Default::default()
3469 }))
3470 })
3471 }
3472 }
3473
3474 #[derive(Clone)]
3475 struct InterceptedClient {
3476 inner: MockStartWorkflowClient,
3477 interceptors: Vec<Arc<dyn ClientInterceptor>>,
3478 }
3479
3480 impl NamespacedClient for InterceptedClient {
3481 fn namespace(&self) -> String {
3482 self.inner.namespace()
3483 }
3484
3485 fn identity(&self) -> String {
3486 self.inner.identity()
3487 }
3488
3489 fn data_converter(&self) -> &DataConverter {
3490 self.inner.data_converter()
3491 }
3492
3493 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
3494 &self.interceptors
3495 }
3496 }
3497
3498 impl WorkflowService for InterceptedClient {
3499 fn start_workflow_execution(
3500 &mut self,
3501 request: Request<StartWorkflowExecutionRequest>,
3502 ) -> futures_util::future::BoxFuture<
3503 '_,
3504 Result<Response<StartWorkflowExecutionResponse>, tonic::Status>,
3505 > {
3506 self.inner.start_workflow_execution(request)
3507 }
3508
3509 fn signal_with_start_workflow_execution(
3510 &mut self,
3511 request: Request<SignalWithStartWorkflowExecutionRequest>,
3512 ) -> futures_util::future::BoxFuture<
3513 '_,
3514 Result<Response<SignalWithStartWorkflowExecutionResponse>, tonic::Status>,
3515 > {
3516 self.inner.signal_with_start_workflow_execution(request)
3517 }
3518 }
3519
3520 struct OrderedInterceptor {
3521 name: &'static str,
3522 events: Arc<Mutex<Vec<String>>>,
3523 encode_calls: Arc<AtomicUsize>,
3524 }
3525
3526 impl ClientInterceptor for OrderedInterceptor {
3527 fn start_workflow<'a>(
3528 &'a self,
3529 mut input: StartWorkflowInput,
3530 next: Next<
3531 'a,
3532 StartWorkflowInput,
3533 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3534 >,
3535 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3536 Box::pin(async move {
3537 assert_eq!(self.encode_calls.load(Ordering::SeqCst), 0);
3538 self.events.lock().push(format!("{}-pre", self.name));
3539 tokio::task::yield_now().await;
3540 if self.name == "outer" {
3541 input
3542 .args_mut::<Vec<String>>()
3543 .unwrap()
3544 .push("mutated".to_owned());
3545 } else {
3546 assert_eq!(
3547 input.args_ref::<Vec<String>>().unwrap(),
3548 &["initial".to_owned(), "mutated".to_owned()]
3549 );
3550 input.replace_args("replacement".to_owned());
3551 input.workflow_type = "replacement-workflow".to_owned();
3552 }
3553 let result = next.run(input).await;
3554 tokio::task::yield_now().await;
3555 self.events.lock().push(format!("{}-post", self.name));
3556 result
3557 })
3558 }
3559 }
3560
3561 struct ShortCircuitInterceptor;
3562
3563 impl ClientInterceptor for ShortCircuitInterceptor {
3564 fn start_workflow<'a>(
3565 &'a self,
3566 input: StartWorkflowInput,
3567 _next: Next<
3568 'a,
3569 StartWorkflowInput,
3570 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3571 >,
3572 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3573 assert_eq!(
3574 input.args_ref::<Vec<String>>().unwrap(),
3575 &["initial".to_owned()]
3576 );
3577 Box::pin(async {
3578 Ok(StartWorkflowOutput::new(
3579 "short-circuit-workflow-id",
3580 "short-circuit-run-id",
3581 ))
3582 })
3583 }
3584 }
3585
3586 struct CountingInput {
3587 conversion_calls: Arc<AtomicUsize>,
3588 }
3589
3590 impl TemporalSerializable for CountingInput {
3591 fn to_payloads(
3592 &self,
3593 _context: &SerializationContext<'_>,
3594 ) -> Result<Vec<Payload>, PayloadConversionError> {
3595 self.conversion_calls.fetch_add(1, Ordering::SeqCst);
3596 Ok(vec![Payload::default()])
3597 }
3598 }
3599
3600 struct ConversionTimingInterceptor {
3601 conversion_calls: Arc<AtomicUsize>,
3602 }
3603
3604 impl ClientInterceptor for ConversionTimingInterceptor {
3605 fn start_workflow<'a>(
3606 &'a self,
3607 mut input: StartWorkflowInput,
3608 next: Next<
3609 'a,
3610 StartWorkflowInput,
3611 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3612 >,
3613 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3614 input.replace_args(CountingInput {
3615 conversion_calls: self.conversion_calls.clone(),
3616 });
3617 let future = next.run(input);
3618 assert_eq!(self.conversion_calls.load(Ordering::SeqCst), 0);
3619 future
3620 }
3621 }
3622
3623 struct ReplacingSignalWithStartInterceptor;
3624
3625 impl ClientInterceptor for ReplacingSignalWithStartInterceptor {
3626 fn signal_with_start_workflow<'a>(
3627 &'a self,
3628 mut input: SignalWithStartWorkflowInput,
3629 next: Next<
3630 'a,
3631 SignalWithStartWorkflowInput,
3632 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3633 >,
3634 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3635 assert_eq!(
3636 input.workflow_args_ref::<Vec<String>>().unwrap(),
3637 &["workflow".to_owned()]
3638 );
3639 assert_eq!(
3640 input.signal_args_ref::<Vec<String>>().unwrap(),
3641 &["signal".to_owned()]
3642 );
3643 input.replace_workflow_args(vec!["replaced-workflow".to_owned()]);
3644 input.replace_signal_args(vec!["replaced-signal".to_owned()]);
3645 next.run(input)
3646 }
3647 }
3648
3649 struct FailingSignal;
3650
3651 impl SignalDefinition for FailingSignal {
3652 type Workflow = test_workflow::Run;
3653 type Input = FailingSignalInput;
3654
3655 fn name(&self) -> &str {
3656 "failing-signal"
3657 }
3658 }
3659
3660 struct FailingSignalInput;
3661
3662 impl TemporalDeserializable for FailingSignalInput {}
3663
3664 impl TemporalSerializable for FailingSignalInput {
3665 fn to_payloads(
3666 &self,
3667 _context: &SerializationContext<'_>,
3668 ) -> Result<Vec<Payload>, PayloadConversionError> {
3669 Err(PayloadConversionError::WrongEncoding)
3670 }
3671 }
3672
3673 fn mock_client(
3674 interceptors: Vec<Arc<dyn ClientInterceptor>>,
3675 encode_calls: Arc<AtomicUsize>,
3676 ) -> (InterceptedClient, Arc<Mutex<RecordedStart>>) {
3677 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
3678 let data_converter = DataConverter::new(
3679 PayloadConverter::default(),
3680 DefaultFailureConverter::default(),
3681 CountingCodec {
3682 encode_calls: encode_calls.clone(),
3683 },
3684 );
3685 (
3686 InterceptedClient {
3687 inner: MockStartWorkflowClient {
3688 recorded: recorded.clone(),
3689 data_converter,
3690 },
3691 interceptors,
3692 },
3693 recorded,
3694 )
3695 }
3696
3697 fn mock_client_with_codec(
3700 codec: impl PayloadCodec + Send + Sync + 'static,
3701 ) -> (MockStartWorkflowClient, Arc<Mutex<RecordedStart>>) {
3702 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
3703 let data_converter = DataConverter::new(
3704 PayloadConverter::default(),
3705 DefaultFailureConverter::default(),
3706 codec,
3707 );
3708 (
3709 MockStartWorkflowClient {
3710 recorded: recorded.clone(),
3711 data_converter,
3712 },
3713 recorded,
3714 )
3715 }
3716
3717 async fn read_back(sent: ProtoMemo) -> Memo {
3719 let mut sent = sent;
3720 decode_payloads(
3721 &mut sent,
3722 &XorCodec,
3723 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
3724 )
3725 .await
3726 .unwrap();
3727 Memo::from_raw(
3728 Some(sent),
3729 PayloadConverter::default(),
3730 SerializationContextData::Workflow(WorkflowSerializationContext::new()),
3731 )
3732 }
3733
3734 #[tokio::test]
3735 async fn start_workflow_encodes_memo_with_payload_converter_and_codec() {
3736 let (client, recorded) = mock_client_with_codec(XorCodec);
3737 let mut memo = MemoValues::new();
3738 memo.insert("memo-key", "memo-value".to_owned());
3739
3740 client
3741 .start_workflow(
3742 TestWorkflow::run,
3743 vec!["initial".to_owned()],
3744 WorkflowStartOptions::new("task-queue", "workflow-id")
3745 .memo(memo)
3746 .build(),
3747 )
3748 .await
3749 .unwrap();
3750
3751 let sent = recorded.lock().memo.clone().expect("memo should be sent");
3752 assert_eq!(
3753 read_back(sent).await.get::<String>("memo-key").unwrap(),
3754 Some("memo-value".to_owned())
3755 );
3756 }
3757
3758 #[tokio::test]
3759 async fn signal_with_start_workflow_encodes_memo() {
3760 let (client, recorded) = mock_client_with_codec(XorCodec);
3761 let mut memo = MemoValues::new();
3762 memo.insert("memo-key", "memo-value".to_owned());
3763
3764 client
3765 .signal_with_start_workflow(
3766 TestWorkflow::run,
3767 vec!["initial".to_owned()],
3768 TestWorkflow::test_signal,
3769 vec!["signal".to_owned()],
3770 WorkflowStartOptions::new("task-queue", "workflow-id")
3771 .memo(memo)
3772 .build(),
3773 )
3774 .await
3775 .unwrap();
3776
3777 let sent = recorded.lock().memo.clone().expect("memo should be sent");
3778 assert_eq!(
3779 read_back(sent).await.get::<String>("memo-key").unwrap(),
3780 Some("memo-value".to_owned())
3781 );
3782 }
3783
3784 #[tokio::test]
3785 async fn start_workflow_without_memo_sends_none() {
3786 let (client, recorded) = mock_client_with_codec(XorCodec);
3787
3788 client
3789 .start_workflow(
3790 TestWorkflow::run,
3791 vec!["initial".to_owned()],
3792 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
3793 )
3794 .await
3795 .unwrap();
3796
3797 assert_eq!(recorded.lock().memo, None);
3798 }
3799
3800 #[tokio::test]
3801 async fn start_workflow_reports_memo_serialization_errors() {
3802 #[derive(Debug)]
3803 struct FailingMemoValue;
3804
3805 impl TemporalSerializable for FailingMemoValue {
3806 fn to_payload(
3807 &self,
3808 _ctx: &SerializationContext<'_>,
3809 ) -> Result<Payload, PayloadConversionError> {
3810 Err(PayloadConversionError::EncodingError(
3811 std::io::Error::other("memo serialization failure").into(),
3812 ))
3813 }
3814 }
3815
3816 let (client, recorded) = mock_client_with_codec(XorCodec);
3817 let mut memo = MemoValues::new();
3818 memo.insert("invalid", FailingMemoValue);
3819
3820 let err = client
3821 .start_workflow(
3822 TestWorkflow::run,
3823 vec!["initial".to_owned()],
3824 WorkflowStartOptions::new("task-queue", "workflow-id")
3825 .memo(memo)
3826 .build(),
3827 )
3828 .await
3829 .map(|_| ())
3830 .expect_err("memo serialization errors should be surfaced");
3831
3832 assert!(
3833 matches!(err, WorkflowStartError::PayloadConversion(_)),
3834 "expected a payload conversion error, got {err:?}"
3835 );
3836 assert!(
3837 err.to_string().contains("memo serialization failure"),
3838 "error should surface the underlying cause, got {err}"
3839 );
3840 assert_eq!(recorded.lock().calls, 0);
3842 }
3843
3844 #[tokio::test]
3845 async fn interceptors_order_mutate_replace_and_defer_conversion() {
3846 let events = Arc::new(Mutex::new(Vec::new()));
3847 let encode_calls = Arc::new(AtomicUsize::new(0));
3848 let interceptors: Vec<Arc<dyn ClientInterceptor>> = vec![
3849 Arc::new(OrderedInterceptor {
3850 name: "outer",
3851 events: events.clone(),
3852 encode_calls: encode_calls.clone(),
3853 }),
3854 Arc::new(OrderedInterceptor {
3855 name: "inner",
3856 events: events.clone(),
3857 encode_calls: encode_calls.clone(),
3858 }),
3859 ];
3860 let (client, recorded) = mock_client(interceptors, encode_calls.clone());
3861
3862 let handle = client
3863 .start_workflow(
3864 TestWorkflow::run,
3865 vec!["initial".to_owned()],
3866 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
3867 )
3868 .await
3869 .unwrap();
3870
3871 assert_eq!(
3872 events.lock().as_slice(),
3873 ["outer-pre", "inner-pre", "inner-post", "outer-post"]
3874 );
3875 assert_eq!(encode_calls.load(Ordering::SeqCst), 1);
3876 assert_eq!(handle.run_id(), Some("server-run-id"));
3877 let payloads = {
3878 let recorded = recorded.lock();
3879 assert_eq!(recorded.calls, 1);
3880 assert_eq!(recorded.workflow_type, "replacement-workflow");
3881 recorded.payloads.clone()
3882 };
3883 let replacement: String = client
3884 .data_converter()
3885 .from_payloads(
3886 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
3887 payloads,
3888 )
3889 .await
3890 .unwrap();
3891 assert_eq!(replacement, "replacement");
3892 }
3893
3894 #[tokio::test]
3895 async fn interceptor_can_short_circuit() {
3896 let encode_calls = Arc::new(AtomicUsize::new(0));
3897 let (client, recorded) = mock_client(
3898 vec![Arc::new(ShortCircuitInterceptor)],
3899 encode_calls.clone(),
3900 );
3901 let handle = client
3902 .start_workflow(
3903 TestWorkflow::run,
3904 vec!["initial".to_owned()],
3905 WorkflowStartOptions::new("task-queue", "ignored-workflow-id").build(),
3906 )
3907 .await
3908 .unwrap();
3909
3910 assert_eq!(handle.info().workflow_id, "short-circuit-workflow-id");
3911 assert_eq!(handle.run_id(), Some("short-circuit-run-id"));
3912 assert_eq!(recorded.lock().calls, 0);
3913 assert_eq!(encode_calls.load(Ordering::SeqCst), 0);
3914 }
3915
3916 #[tokio::test]
3917 async fn payload_conversion_waits_for_next_future_poll() {
3918 let conversion_calls = Arc::new(AtomicUsize::new(0));
3919 let encode_calls = Arc::new(AtomicUsize::new(0));
3920 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
3921 let data_converter = DataConverter::new(
3922 PayloadConverter::UseWrappers,
3923 DefaultFailureConverter::default(),
3924 CountingCodec {
3925 encode_calls: encode_calls.clone(),
3926 },
3927 );
3928 let client = InterceptedClient {
3929 inner: MockStartWorkflowClient {
3930 recorded: recorded.clone(),
3931 data_converter,
3932 },
3933 interceptors: vec![Arc::new(ConversionTimingInterceptor {
3934 conversion_calls: conversion_calls.clone(),
3935 })],
3936 };
3937
3938 client
3939 .start_workflow(
3940 TestWorkflow::run,
3941 vec!["initial".to_owned()],
3942 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
3943 )
3944 .await
3945 .unwrap();
3946
3947 assert_eq!(conversion_calls.load(Ordering::SeqCst), 1);
3948 assert_eq!(encode_calls.load(Ordering::SeqCst), 1);
3949 assert_eq!(recorded.lock().calls, 1);
3950 }
3951
3952 #[tokio::test]
3953 async fn custom_client_defaults_to_empty_chain() {
3954 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
3955 let client = MockStartWorkflowClient {
3956 recorded: recorded.clone(),
3957 data_converter: DataConverter::default(),
3958 };
3959 assert!(client.client_interceptors().is_empty());
3960
3961 client
3962 .start_workflow(
3963 TestWorkflow::run,
3964 vec!["initial".to_owned()],
3965 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
3966 )
3967 .await
3968 .unwrap();
3969 assert_eq!(recorded.lock().calls, 1);
3970 }
3971
3972 #[tokio::test]
3973 async fn rpc_options_reach_the_request() {
3974 let (client, recorded) = mock_client(Vec::new(), Arc::new(AtomicUsize::new(0)));
3975 let mut metadata = RpcMetadata::new();
3976 metadata.insert("call-meta", "call-value").unwrap();
3977 metadata
3978 .insert_binary("call-meta-bin", vec![0, 255])
3979 .unwrap();
3980 let rpc_options = RpcOptions::builder()
3981 .metadata(metadata)
3982 .timeout(Duration::from_millis(250))
3983 .retry_options(RetryOptions::no_retries())
3984 .build();
3985 let mut options = WorkflowStartOptions::new("task-queue", "workflow-id").build();
3986 options.rpc_options = rpc_options.clone();
3987
3988 client
3989 .start_workflow(TestWorkflow::run, vec!["initial".to_owned()], options)
3990 .await
3991 .unwrap();
3992
3993 {
3994 let recorded = recorded.lock();
3995 assert_eq!(recorded.ascii_metadata.as_deref(), Some("call-value"));
3996 assert_eq!(recorded.binary_metadata.as_deref(), Some(&[0, 255][..]));
3997 assert_eq!(recorded.grpc_timeout.as_deref(), Some("250000u"));
3998 assert_eq!(recorded.retry_options, Some(RetryOptions::no_retries()));
3999 }
4000
4001 let mut options = WorkflowStartOptions::new("task-queue", "signal-workflow-id").build();
4002 options.rpc_options = rpc_options;
4003 let handle = client
4004 .signal_with_start_workflow(
4005 TestWorkflow::run,
4006 vec!["initial".to_owned()],
4007 TestWorkflow::test_signal,
4008 vec!["signal".to_owned()],
4009 options,
4010 )
4011 .await
4012 .unwrap();
4013
4014 let recorded = recorded.lock();
4015 assert_eq!(recorded.calls, 2);
4016 assert_eq!(recorded.ascii_metadata.as_deref(), Some("call-value"));
4017 assert_eq!(recorded.binary_metadata.as_deref(), Some(&[0, 255][..]));
4018 assert_eq!(recorded.grpc_timeout.as_deref(), Some("250000u"));
4019 assert_eq!(recorded.retry_options, Some(RetryOptions::no_retries()));
4020 assert_eq!(recorded.signal_name, "test_signal");
4021 assert_eq!(recorded.signal_payloads.len(), 1);
4022 assert_eq!(handle.run_id(), Some("signal-server-run-id"));
4023 }
4024
4025 #[tokio::test]
4026 async fn signal_with_start_interceptor_can_replace_both_argument_sets() {
4027 let (client, recorded) = mock_client(
4028 vec![Arc::new(ReplacingSignalWithStartInterceptor)],
4029 Arc::new(AtomicUsize::new(0)),
4030 );
4031
4032 client
4033 .signal_with_start_workflow(
4034 TestWorkflow::run,
4035 vec!["workflow".to_owned()],
4036 TestWorkflow::test_signal,
4037 vec!["signal".to_owned()],
4038 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
4039 )
4040 .await
4041 .unwrap();
4042
4043 let data_converter = DataConverter::default();
4044 let (workflow_payloads, signal_payloads) = {
4045 let recorded = recorded.lock();
4046 (recorded.payloads.clone(), recorded.signal_payloads.clone())
4047 };
4048 assert_eq!(
4049 data_converter
4050 .from_payloads::<Vec<String>>(
4051 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4052 workflow_payloads,
4053 )
4054 .await
4055 .unwrap(),
4056 vec!["replaced-workflow".to_owned()]
4057 );
4058 assert_eq!(
4059 data_converter
4060 .from_payloads::<Vec<String>>(
4061 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4062 signal_payloads,
4063 )
4064 .await
4065 .unwrap(),
4066 vec!["replaced-signal".to_owned()]
4067 );
4068 }
4069
4070 #[tokio::test]
4071 async fn signal_with_start_payload_conversion_failure_does_not_call_service() {
4072 let (client, recorded) = mock_client(Vec::new(), Arc::new(AtomicUsize::new(0)));
4073
4074 let result = client
4075 .signal_with_start_workflow(
4076 TestWorkflow::run,
4077 vec!["workflow".to_owned()],
4078 FailingSignal,
4079 FailingSignalInput,
4080 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
4081 )
4082 .await;
4083
4084 assert!(matches!(
4085 result,
4086 Err(WorkflowStartError::PayloadConversion(_))
4087 ));
4088 assert_eq!(recorded.lock().calls, 0);
4089 }
4090
4091 #[test]
4092 fn rpc_metadata_combines_with_and_overrides_connection_defaults() {
4093 let headers = Arc::new(RwLock::new(ClientHeaders {
4094 user_headers: HashMap::from([
4095 (
4096 "shared-meta".parse().unwrap(),
4097 "connection-value".parse().unwrap(),
4098 ),
4099 (
4100 "connection-meta".parse().unwrap(),
4101 "connection-only".parse().unwrap(),
4102 ),
4103 ]),
4104 user_binary_headers: HashMap::from([
4105 (
4106 "shared-meta-bin".parse().unwrap(),
4107 BinaryMetadataValue::from_bytes(&[1]),
4108 ),
4109 (
4110 "connection-meta-bin".parse().unwrap(),
4111 BinaryMetadataValue::from_bytes(&[2]),
4112 ),
4113 ]),
4114 api_key: None,
4115 }));
4116 let mut service_interceptor = ServiceCallInterceptor {
4117 client_name: "test-client".to_owned(),
4118 client_version: "test-version".to_owned(),
4119 headers,
4120 };
4121 let mut rpc_options = RpcOptions::default();
4122 rpc_options
4123 .metadata
4124 .insert("shared-meta", "call-value")
4125 .unwrap();
4126 rpc_options
4127 .metadata
4128 .insert("call-meta", "call-only")
4129 .unwrap();
4130 rpc_options
4131 .metadata
4132 .insert_binary("shared-meta-bin", vec![3])
4133 .unwrap();
4134 rpc_options
4135 .metadata
4136 .insert_binary("call-meta-bin", vec![4])
4137 .unwrap();
4138 let mut request = Request::new(());
4139 rpc_options.apply_to(&mut request);
4140
4141 let request = service_interceptor.call(request).unwrap();
4142 assert_eq!(request.metadata().get("shared-meta").unwrap(), "call-value");
4143 assert_eq!(request.metadata().get("call-meta").unwrap(), "call-only");
4144 assert_eq!(
4145 request.metadata().get("connection-meta").unwrap(),
4146 "connection-only"
4147 );
4148 assert_eq!(
4149 request.metadata().get_bin("shared-meta-bin").unwrap(),
4150 &[3][..]
4151 );
4152 assert_eq!(
4153 request.metadata().get_bin("call-meta-bin").unwrap(),
4154 &[4][..]
4155 );
4156 assert_eq!(
4157 request.metadata().get_bin("connection-meta-bin").unwrap(),
4158 &[2][..]
4159 );
4160 }
4161 }
4162
4163 mod update_with_start_tests {
4164 use super::*;
4165 use assert_matches::assert_matches;
4166 use parking_lot::Mutex;
4167 use std::collections::VecDeque;
4168 use temporalio_common::{
4169 UpdateDefinition, WorkflowDefinition,
4170 data_converters::{GenericPayloadConverter, PayloadConverter},
4171 protos::temporal::api::{
4172 common::v1::{
4173 Header, Payload, Payloads, WorkflowExecution as ProtoWorkflowExecution,
4174 },
4175 enums::v1::{
4176 UpdateWorkflowExecutionLifecycleStage,
4177 WorkflowIdConflictPolicy as ProtoWorkflowIdConflictPolicy,
4178 },
4179 update::v1::{
4180 Input as UpdateInput, Meta as UpdateMeta, Outcome, Request as UpdateRequest,
4181 UpdateRef, WaitPolicy, outcome,
4182 },
4183 },
4184 };
4185 use tonic::{Request, Response};
4186
4187 struct TestWorkflow;
4188
4189 impl WorkflowDefinition for TestWorkflow {
4190 type Input = String;
4191 type Output = ();
4192
4193 fn name(&self) -> &str {
4194 "test-workflow"
4195 }
4196 }
4197
4198 impl HasWorkflowDefinition for TestWorkflow {
4199 type Run = Self;
4200 }
4201
4202 struct TestUpdate;
4203
4204 impl UpdateDefinition for TestUpdate {
4205 type Workflow = TestWorkflow;
4206 type Input = String;
4207 type Output = String;
4208
4209 fn name(&self) -> &str {
4210 "test-update"
4211 }
4212 }
4213
4214 fn successful_multi_operation_response(
4215 stage: UpdateWorkflowExecutionLifecycleStage,
4216 ) -> ExecuteMultiOperationResponse {
4217 let outcome = (stage == UpdateWorkflowExecutionLifecycleStage::Completed).then(|| {
4218 let payload_converter = PayloadConverter::default();
4219 let result_payloads =
4220 payload_converter
4221 .to_payloads(
4222 &SerializationContext::new(
4223 &SerializationContextData::Workflow(
4224 WorkflowSerializationContext::new(),
4225 ),
4226 &payload_converter,
4227 ),
4228 &"update-result".to_owned(),
4229 )
4230 .unwrap();
4231 Outcome {
4232 value: Some(outcome::Value::Success(Payloads {
4233 payloads: result_payloads,
4234 })),
4235 }
4236 });
4237 ExecuteMultiOperationResponse {
4238 responses: vec![
4239 execute_multi_operation_response::Response {
4240 response: Some(MultiOperationResponse::StartWorkflow(
4241 StartWorkflowExecutionResponse {
4242 run_id: "started-run-id".to_owned(),
4243 first_execution_run_id: "first-run-id".to_owned(),
4244 started: true,
4245 ..Default::default()
4246 },
4247 )),
4248 },
4249 execute_multi_operation_response::Response {
4250 response: Some(MultiOperationResponse::UpdateWorkflow(
4251 UpdateWorkflowExecutionResponse {
4252 update_ref: Some(UpdateRef {
4253 workflow_execution: Some(ProtoWorkflowExecution {
4254 workflow_id: "workflow-id".to_owned(),
4255 run_id: "update-run-id".to_owned(),
4256 }),
4257 update_id: "server-update-id".to_owned(),
4258 }),
4259 outcome,
4260 stage: stage as i32,
4261 ..Default::default()
4262 },
4263 )),
4264 },
4265 ],
4266 }
4267 }
4268
4269 #[derive(Clone)]
4270 struct MockMultiOperationClient {
4271 recorded: Arc<Mutex<Option<ExecuteMultiOperationRequest>>>,
4272 responses: Arc<Mutex<VecDeque<ExecuteMultiOperationResponse>>>,
4273 call_count: Arc<Mutex<usize>>,
4274 interceptors: Vec<Arc<dyn ClientInterceptor>>,
4275 }
4276
4277 impl MockMultiOperationClient {
4278 fn new(
4279 interceptors: Vec<Arc<dyn ClientInterceptor>>,
4280 responses: impl IntoIterator<Item = ExecuteMultiOperationResponse>,
4281 ) -> Self {
4282 Self {
4283 recorded: Arc::new(Mutex::new(None)),
4284 responses: Arc::new(Mutex::new(responses.into_iter().collect())),
4285 call_count: Arc::new(Mutex::new(0)),
4286 interceptors,
4287 }
4288 }
4289 }
4290
4291 impl NamespacedClient for MockMultiOperationClient {
4292 fn namespace(&self) -> String {
4293 "test-namespace".to_owned()
4294 }
4295
4296 fn identity(&self) -> String {
4297 "test-identity".to_owned()
4298 }
4299
4300 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
4301 &self.interceptors
4302 }
4303 }
4304
4305 impl WorkflowService for MockMultiOperationClient {
4306 fn execute_multi_operation(
4307 &mut self,
4308 request: Request<ExecuteMultiOperationRequest>,
4309 ) -> futures_util::future::BoxFuture<
4310 '_,
4311 Result<Response<ExecuteMultiOperationResponse>, tonic::Status>,
4312 > {
4313 *self.recorded.lock() = Some(request.into_inner());
4314 *self.call_count.lock() += 1;
4315 let response = self.responses.lock().pop_front().unwrap_or_else(|| {
4316 successful_multi_operation_response(
4317 UpdateWorkflowExecutionLifecycleStage::Completed,
4318 )
4319 });
4320 Box::pin(async { Ok(Response::new(response)) })
4321 }
4322 }
4323
4324 fn update_with_start_options(
4325 conflict_policy: WorkflowIdConflictPolicy,
4326 ) -> WorkflowUpdateWithStartOptions {
4327 WorkflowUpdateWithStartOptions::new("task-queue", "workflow-id", conflict_policy)
4328 .build()
4329 }
4330
4331 #[tokio::test]
4332 async fn update_with_start_builds_multi_operation_request() {
4333 let client = MockMultiOperationClient::new(Vec::new(), []);
4334 let recorded = client.recorded.clone();
4335
4336 let start_header = Header {
4337 fields: HashMap::from([("start-header".to_owned(), Payload::default())]),
4338 };
4339 let update_header = Header {
4340 fields: HashMap::from([("update-header".to_owned(), Payload::default())]),
4341 };
4342 let update_handle = client
4343 .start_update_with_start_workflow(
4344 TestWorkflow,
4345 "workflow-input".to_owned(),
4346 TestUpdate,
4347 "update-input".to_owned(),
4348 WorkflowUpdateWithStartOptions::new(
4349 "task-queue",
4350 "workflow-id",
4351 WorkflowIdConflictPolicy::UseExisting,
4352 )
4353 .update_id("my-update-id".to_owned())
4354 .start_header(start_header.clone())
4355 .update_header(update_header.clone())
4356 .build(),
4357 )
4358 .await
4359 .unwrap();
4360
4361 let payload_converter = PayloadConverter::default();
4362 let context_data =
4363 SerializationContextData::Workflow(WorkflowSerializationContext::new());
4364 let context = SerializationContext::new(&context_data, &payload_converter);
4365 let workflow_payloads = payload_converter
4366 .to_payloads(&context, &"workflow-input".to_owned())
4367 .unwrap();
4368 let update_payloads = payload_converter
4369 .to_payloads(&context, &"update-input".to_owned())
4370 .unwrap();
4371
4372 let request = recorded.lock().take().unwrap();
4373 let request_id = assert_matches!(
4374 &request.operations[0].operation,
4375 Some(execute_multi_operation_request::operation::Operation::StartWorkflow(r)) => r
4376 )
4377 .request_id
4378 .clone();
4379 assert_eq!(
4380 request,
4381 ExecuteMultiOperationRequest {
4382 namespace: "test-namespace".to_owned(),
4383 operations: vec![
4384 execute_multi_operation_request::Operation {
4385 operation: Some(MultiOperationRequest::StartWorkflow(
4386 StartWorkflowExecutionRequest {
4387 namespace: "test-namespace".to_owned(),
4388 workflow_id: "workflow-id".to_owned(),
4389 workflow_type: Some(WorkflowType {
4390 name: "test-workflow".to_owned(),
4391 }),
4392 task_queue: Some(TaskQueue {
4393 name: "task-queue".to_owned(),
4394 ..Default::default()
4395 }),
4396 input: Some(Payloads {
4397 payloads: workflow_payloads,
4398 }),
4399 request_id,
4400 identity: "test-identity".to_owned(),
4401 workflow_id_conflict_policy:
4402 ProtoWorkflowIdConflictPolicy::UseExisting as i32,
4403 header: Some(start_header),
4404 priority: Some(Default::default()),
4405 ..Default::default()
4406 },
4407 )),
4408 },
4409 execute_multi_operation_request::Operation {
4410 operation: Some(MultiOperationRequest::UpdateWorkflow(
4411 UpdateWorkflowExecutionRequest {
4412 namespace: "test-namespace".to_owned(),
4413 workflow_execution: Some(ProtoWorkflowExecution {
4414 workflow_id: "workflow-id".to_owned(),
4415 run_id: String::new(),
4416 }),
4417 wait_policy: Some(WaitPolicy {
4418 lifecycle_stage:
4419 UpdateWorkflowExecutionLifecycleStage::Accepted as i32,
4420 }),
4421 request: Some(UpdateRequest {
4422 meta: Some(UpdateMeta {
4423 update_id: "my-update-id".to_owned(),
4424 identity: "test-identity".to_owned(),
4425 }),
4426 input: Some(UpdateInput {
4427 header: Some(update_header),
4428 name: "test-update".to_owned(),
4429 args: Some(Payloads {
4430 payloads: update_payloads,
4431 }),
4432 }),
4433 ..Default::default()
4434 }),
4435 ..Default::default()
4436 },
4437 )),
4438 },
4439 ],
4440 resource_id: "workflow-id".to_owned(),
4441 }
4442 );
4443
4444 assert_eq!(update_handle.id(), "my-update-id");
4445 assert_eq!(update_handle.workflow_run_id(), Some("update-run-id"));
4446 let result: String = update_handle
4449 .get_result(RpcOptions::default())
4450 .await
4451 .unwrap();
4452 assert_eq!(result, "update-result");
4453 }
4454
4455 #[tokio::test]
4456 async fn update_with_start_retries_until_update_is_accepted() {
4457 let client = MockMultiOperationClient::new(
4458 Vec::new(),
4459 [
4460 successful_multi_operation_response(
4461 UpdateWorkflowExecutionLifecycleStage::Unspecified,
4462 ),
4463 successful_multi_operation_response(
4464 UpdateWorkflowExecutionLifecycleStage::Accepted,
4465 ),
4466 ],
4467 );
4468 let call_count = client.call_count.clone();
4469
4470 let update_handle = client
4471 .start_update_with_start_workflow(
4472 TestWorkflow,
4473 "workflow-input".to_owned(),
4474 TestUpdate,
4475 "update-input".to_owned(),
4476 update_with_start_options(WorkflowIdConflictPolicy::Fail),
4477 )
4478 .await
4479 .unwrap();
4480
4481 assert_eq!(*call_count.lock(), 2);
4482 assert_eq!(update_handle.workflow_run_id(), Some("update-run-id"));
4483 }
4484
4485 #[tokio::test]
4486 async fn update_with_start_rejects_malformed_operation_responses() {
4487 let mut missing_response = successful_multi_operation_response(
4488 UpdateWorkflowExecutionLifecycleStage::Accepted,
4489 );
4490 missing_response.responses[0] = execute_multi_operation_response::Response::default();
4491 let mut extra_response = successful_multi_operation_response(
4492 UpdateWorkflowExecutionLifecycleStage::Accepted,
4493 );
4494 extra_response
4495 .responses
4496 .push(execute_multi_operation_response::Response::default());
4497 let mut wrong_order = successful_multi_operation_response(
4498 UpdateWorkflowExecutionLifecycleStage::Accepted,
4499 );
4500 wrong_order.responses.swap(0, 1);
4501
4502 for response in [missing_response, extra_response, wrong_order] {
4503 let client = MockMultiOperationClient::new(Vec::new(), [response]);
4504 let result = client
4505 .start_update_with_start_workflow(
4506 TestWorkflow,
4507 "workflow-input".to_owned(),
4508 TestUpdate,
4509 "update-input".to_owned(),
4510 update_with_start_options(WorkflowIdConflictPolicy::Fail),
4511 )
4512 .await;
4513 assert!(matches!(
4514 result,
4515 Err(WorkflowUpdateWithStartError::Other(_))
4516 ));
4517 }
4518 }
4519
4520 #[tokio::test]
4521 async fn update_with_start_interceptor_can_mutate_args() {
4522 struct ReplaceArgsInterceptor;
4523
4524 impl ClientInterceptor for ReplaceArgsInterceptor {
4525 fn update_with_start_workflow<'a>(
4526 &'a self,
4527 mut input: UpdateWithStartWorkflowInput,
4528 next: Next<
4529 'a,
4530 UpdateWithStartWorkflowInput,
4531 BoxFuture<
4532 'a,
4533 Result<UpdateWithStartWorkflowOutput, WorkflowUpdateWithStartError>,
4534 >,
4535 >,
4536 ) -> BoxFuture<
4537 'a,
4538 Result<UpdateWithStartWorkflowOutput, WorkflowUpdateWithStartError>,
4539 > {
4540 assert_eq!(
4541 input.workflow_args_ref::<String>().unwrap(),
4542 "workflow-input"
4543 );
4544 input.replace_workflow_args("replaced-workflow-input".to_owned());
4545 *input.update_args_mut::<String>().unwrap() =
4546 "replaced-update-input".to_owned();
4547 next.run(input)
4548 }
4549 }
4550
4551 let client = MockMultiOperationClient::new(vec![Arc::new(ReplaceArgsInterceptor)], []);
4552 let recorded = client.recorded.clone();
4553
4554 client
4555 .start_update_with_start_workflow(
4556 TestWorkflow,
4557 "workflow-input".to_owned(),
4558 TestUpdate,
4559 "update-input".to_owned(),
4560 update_with_start_options(WorkflowIdConflictPolicy::Fail),
4561 )
4562 .await
4563 .unwrap();
4564
4565 let request = recorded.lock().take().unwrap();
4566 let start_request = assert_matches!(
4567 &request.operations[0].operation,
4568 Some(execute_multi_operation_request::operation::Operation::StartWorkflow(r)) => r
4569 );
4570 let workflow_input: String = client
4571 .data_converter()
4572 .from_payloads(
4573 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4574 start_request.input.clone().unwrap().payloads,
4575 )
4576 .await
4577 .unwrap();
4578 assert_eq!(workflow_input, "replaced-workflow-input");
4579 let update_request = assert_matches!(
4580 &request.operations[1].operation,
4581 Some(execute_multi_operation_request::operation::Operation::UpdateWorkflow(r)) => r
4582 );
4583 let update_input: String = client
4584 .data_converter()
4585 .from_payloads(
4586 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4587 update_request
4588 .request
4589 .clone()
4590 .unwrap()
4591 .input
4592 .unwrap()
4593 .args
4594 .unwrap()
4595 .payloads,
4596 )
4597 .await
4598 .unwrap();
4599 assert_eq!(update_input, "replaced-update-input");
4600 }
4601 }
4602
4603 mod list_workflows_tests {
4604 use super::*;
4605 use crate::test_helpers::{FailingCodec, XorCodec};
4606 use futures_util::{FutureExt, StreamExt};
4607 use std::sync::atomic::{AtomicUsize, Ordering};
4608 use temporalio_common::{
4609 data_converters::{DefaultFailureConverter, PayloadConverter},
4610 protos::temporal::api::common::v1::{
4611 Memo as ProtoMemo, Payload, WorkflowExecution as ProtoWorkflowExecution,
4612 },
4613 };
4614 use tonic::{Request, Response};
4615
4616 #[derive(Clone)]
4617 struct MockListWorkflowsClient {
4618 call_count: Arc<AtomicUsize>,
4619 page_size: usize,
4621 total_workflows: usize,
4623 data_converter: DataConverter,
4624 memo_payload: Option<Payload>,
4625 interceptors: Vec<Arc<dyn ClientInterceptor>>,
4626 }
4627
4628 impl NamespacedClient for MockListWorkflowsClient {
4629 fn namespace(&self) -> String {
4630 "test-namespace".to_string()
4631 }
4632 fn identity(&self) -> String {
4633 "test-identity".to_string()
4634 }
4635 fn data_converter(&self) -> &DataConverter {
4636 &self.data_converter
4637 }
4638 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
4639 &self.interceptors
4640 }
4641 }
4642
4643 struct CountingListInterceptor {
4644 calls: Arc<AtomicUsize>,
4645 }
4646
4647 impl ClientInterceptor for CountingListInterceptor {
4648 fn list_workflows_page<'a>(
4649 &'a self,
4650 input: ListWorkflowsPageInput,
4651 next: Next<
4652 'a,
4653 ListWorkflowsPageInput,
4654 BoxFuture<'a, Result<ListWorkflowsPageOutput, ClientError>>,
4655 >,
4656 ) -> BoxFuture<'a, Result<ListWorkflowsPageOutput, ClientError>> {
4657 self.calls.fetch_add(1, Ordering::SeqCst);
4658 next.run(input)
4659 }
4660 }
4661
4662 impl WorkflowService for MockListWorkflowsClient {
4663 fn list_workflow_executions(
4664 &mut self,
4665 request: Request<ListWorkflowExecutionsRequest>,
4666 ) -> futures_util::future::BoxFuture<
4667 '_,
4668 Result<Response<ListWorkflowExecutionsResponse>, tonic::Status>,
4669 > {
4670 self.call_count.fetch_add(1, Ordering::SeqCst);
4671 let req = request.into_inner();
4672
4673 let offset: usize = if req.next_page_token.is_empty() {
4675 0
4676 } else {
4677 String::from_utf8(req.next_page_token)
4678 .unwrap()
4679 .parse()
4680 .unwrap()
4681 };
4682
4683 let remaining = self.total_workflows.saturating_sub(offset);
4684 let count = remaining.min(self.page_size);
4685 let new_offset = offset + count;
4686
4687 let executions: Vec<_> = (offset..offset + count)
4688 .map(|i| workflow::WorkflowExecutionInfo {
4689 execution: Some(ProtoWorkflowExecution {
4690 workflow_id: format!("wf-{i}"),
4691 run_id: format!("run-{i}"),
4692 }),
4693 r#type: Some(WorkflowType {
4694 name: "TestWorkflow".to_string(),
4695 }),
4696 task_queue: "test-queue".to_string(),
4697 memo: self.memo_payload.clone().map(|payload| ProtoMemo {
4698 fields: HashMap::from([("memo-key".to_owned(), payload)]),
4699 }),
4700 ..Default::default()
4701 })
4702 .collect();
4703
4704 let next_page_token = if new_offset < self.total_workflows {
4705 new_offset.to_string().into_bytes()
4706 } else {
4707 vec![]
4708 };
4709
4710 async move {
4711 Ok(Response::new(ListWorkflowExecutionsResponse {
4712 executions,
4713 next_page_token,
4714 }))
4715 }
4716 .boxed()
4717 }
4718 }
4719
4720 #[tokio::test]
4721 async fn list_workflows_paginates_through_all_results() {
4722 let call_count = Arc::new(AtomicUsize::new(0));
4723 let interceptor_calls = Arc::new(AtomicUsize::new(0));
4724 let client = MockListWorkflowsClient {
4725 call_count: call_count.clone(),
4726 page_size: 3,
4727 total_workflows: 10,
4728 data_converter: DataConverter::default(),
4729 memo_payload: None,
4730 interceptors: vec![Arc::new(CountingListInterceptor {
4731 calls: interceptor_calls.clone(),
4732 })],
4733 };
4734
4735 let stream = client.list_workflows("", WorkflowListOptions::default());
4736 let results: Vec<_> = stream.collect().await;
4737
4738 assert_eq!(results.len(), 10);
4739 for (i, result) in results.iter().enumerate() {
4740 let wf = result.as_ref().unwrap();
4741 assert_eq!(wf.id(), format!("wf-{i}"));
4742 assert_eq!(wf.run_id(), format!("run-{i}"));
4743 }
4744 assert_eq!(call_count.load(Ordering::SeqCst), 4);
4746 assert_eq!(interceptor_calls.load(Ordering::SeqCst), 4);
4747 }
4748
4749 #[tokio::test]
4750 async fn list_workflows_respects_limit() {
4751 let call_count = Arc::new(AtomicUsize::new(0));
4752 let client = MockListWorkflowsClient {
4753 call_count: call_count.clone(),
4754 page_size: 3,
4755 total_workflows: 10,
4756 data_converter: DataConverter::default(),
4757 memo_payload: None,
4758 interceptors: Vec::new(),
4759 };
4760
4761 let opts = WorkflowListOptions::builder().limit(5).build();
4762 let stream = client.list_workflows("", opts);
4763 let results: Vec<_> = stream.collect().await;
4764
4765 assert_eq!(results.len(), 5);
4766 for (i, result) in results.iter().enumerate() {
4767 let wf = result.as_ref().unwrap();
4768 assert_eq!(wf.id(), format!("wf-{i}"));
4769 }
4770 assert_eq!(call_count.load(Ordering::SeqCst), 2);
4772 }
4773
4774 #[tokio::test]
4775 async fn list_workflows_limit_less_than_page_size() {
4776 let call_count = Arc::new(AtomicUsize::new(0));
4777 let client = MockListWorkflowsClient {
4778 call_count: call_count.clone(),
4779 page_size: 10,
4780 total_workflows: 100,
4781 data_converter: DataConverter::default(),
4782 memo_payload: None,
4783 interceptors: Vec::new(),
4784 };
4785
4786 let opts = WorkflowListOptions::builder().limit(3).build();
4787 let stream = client.list_workflows("", opts);
4788 let results: Vec<_> = stream.collect().await;
4789
4790 assert_eq!(results.len(), 3);
4791 assert_eq!(call_count.load(Ordering::SeqCst), 1);
4793 }
4794
4795 #[tokio::test]
4796 async fn list_workflows_empty_results() {
4797 let call_count = Arc::new(AtomicUsize::new(0));
4798 let client = MockListWorkflowsClient {
4799 call_count: call_count.clone(),
4800 page_size: 10,
4801 total_workflows: 0,
4802 data_converter: DataConverter::default(),
4803 memo_payload: None,
4804 interceptors: Vec::new(),
4805 };
4806
4807 let stream = client.list_workflows("", WorkflowListOptions::default());
4808 let results: Vec<_> = stream.collect().await;
4809
4810 assert_eq!(results.len(), 0);
4811 assert_eq!(call_count.load(Ordering::SeqCst), 1);
4812 }
4813
4814 #[tokio::test]
4815 async fn list_workflows_exposes_typed_memo() {
4816 let data_converter = DataConverter::new(
4817 PayloadConverter::default(),
4818 DefaultFailureConverter::default(),
4819 XorCodec,
4820 );
4821 let memo_payload = data_converter
4822 .to_payload(
4823 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4824 &"memo-value".to_owned(),
4825 )
4826 .await
4827 .unwrap();
4828 let client = MockListWorkflowsClient {
4829 call_count: Arc::new(AtomicUsize::new(0)),
4830 page_size: 1,
4831 total_workflows: 1,
4832 data_converter,
4833 memo_payload: Some(memo_payload),
4834 interceptors: Vec::new(),
4835 };
4836
4837 let workflow = client
4838 .list_workflows("", WorkflowListOptions::default())
4839 .next()
4840 .await
4841 .unwrap()
4842 .unwrap();
4843
4844 assert_eq!(
4845 workflow.memo().get::<String>("memo-key").unwrap(),
4846 Some("memo-value".to_owned())
4847 );
4848 }
4849
4850 #[tokio::test]
4851 async fn list_workflows_yields_codec_error_then_ends() {
4852 let client = MockListWorkflowsClient {
4853 call_count: Arc::new(AtomicUsize::new(0)),
4854 page_size: 1,
4855 total_workflows: 1,
4856 data_converter: DataConverter::new(
4857 PayloadConverter::default(),
4858 DefaultFailureConverter::default(),
4859 FailingCodec,
4860 ),
4861 memo_payload: Some(Payload::default()),
4862 interceptors: Vec::new(),
4863 };
4864 let mut stream = client.list_workflows("", WorkflowListOptions::default());
4865
4866 let err = stream.next().await.unwrap().unwrap_err();
4867
4868 assert!(matches!(err, ClientError::PayloadConversion(_)));
4869 assert!(stream.next().await.is_none());
4870 }
4871 }
4872}