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, VersioningOverride};
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 versioning_override: options.versioning_override.map(Into::into),
1851 memo,
1852 header: options.header,
1853 user_metadata,
1854 ..Default::default()
1855 }
1856}
1857
1858impl<T> WorkflowClientTrait for T
1859where
1860 T: WorkflowService + NamespacedClient + Clone + Send + Sync + 'static,
1861{
1862 async fn start_workflow<W>(
1863 &self,
1864 workflow: W,
1865 input: W::Input,
1866 options: WorkflowStartOptions,
1867 ) -> Result<WorkflowHandle<Self, W>, WorkflowStartError>
1868 where
1869 W: HasWorkflowDefinition,
1870 W::Input: Send,
1871 {
1872 let namespace = self.namespace();
1873 let interceptor_output = interceptors::call_start_workflow(
1874 self.client_interceptors(),
1875 StartWorkflowInput::new(workflow.name().to_owned(), input, options),
1876 Next::new({
1877 let client = (*self).clone();
1878 move |input: StartWorkflowInput| -> BoxFuture<
1879 '_,
1880 Result<StartWorkflowOutput, WorkflowStartError>,
1881 > {
1882 let mut client = client;
1883 Box::pin(async move {
1884 let (workflow_type, args, options, rpc_options) = input.into_parts();
1885 let data_converter = client.data_converter().clone();
1886 let unencoded_payloads = {
1887 let payload_converter = data_converter.payload_converter();
1888 let context_data = SerializationContextData::Workflow(
1889 WorkflowSerializationContext::new(),
1890 );
1891 let context =
1892 SerializationContext::new(&context_data, payload_converter);
1893 args.serialize_payloads(&context)
1894 };
1895 drop(args);
1896
1897 let payloads = data_converter
1898 .codec()
1899 .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), unencoded_payloads?)
1900 .await?;
1901 let workflow_id = options.workflow_id.clone();
1902 let memo = options.encoded_memo(&data_converter).await?;
1903 let mut request = build_start_workflow_request(
1904 &client,
1905 workflow_type,
1906 payloads.into_payloads(),
1907 memo,
1908 options,
1909 )
1910 .into_request();
1911 rpc_options.apply_to(&mut request);
1912 let run_id = client
1913 .start_workflow_execution(request)
1914 .await
1915 .map_err(WorkflowStartError::from_status)?
1916 .into_inner()
1917 .run_id;
1918
1919 Ok(StartWorkflowOutput::new(workflow_id, run_id))
1920 })
1921 }
1922 }),
1923 )
1924 .await?;
1925 let StartWorkflowOutput {
1926 workflow_id,
1927 run_id,
1928 } = interceptor_output;
1929
1930 Ok(WorkflowHandle::new(
1931 self.clone(),
1932 WorkflowExecutionInfo {
1933 namespace,
1934 workflow_id,
1935 run_id: Some(run_id.clone()),
1936 first_execution_run_id: Some(run_id),
1937 },
1938 ))
1939 }
1940
1941 async fn signal_with_start_workflow<W, S>(
1942 &self,
1943 workflow: W,
1944 workflow_input: W::Input,
1945 signal: S,
1946 signal_input: S::Input,
1947 options: WorkflowStartOptions,
1948 ) -> Result<WorkflowHandle<Self, W>, WorkflowStartError>
1949 where
1950 W: HasWorkflowDefinition,
1951 W::Input: Send,
1952 S: SignalDefinition<Workflow = W::Run>,
1953 S::Input: Send,
1954 {
1955 let namespace = self.namespace();
1956 let interceptor_output = interceptors::call_signal_with_start_workflow(
1957 self.client_interceptors(),
1958 SignalWithStartWorkflowInput::new(
1959 workflow.name().to_owned(),
1960 workflow_input,
1961 signal.name().to_owned(),
1962 signal_input,
1963 options,
1964 ),
1965 Next::new({
1966 let client = (*self).clone();
1967 move |input: SignalWithStartWorkflowInput| -> BoxFuture<
1968 '_,
1969 Result<StartWorkflowOutput, WorkflowStartError>,
1970 > {
1971 let mut client = client;
1972 Box::pin(async move {
1973 let (
1974 workflow_type,
1975 workflow_args,
1976 signal_name,
1977 signal_args,
1978 options,
1979 rpc_options,
1980 ) = input.into_parts();
1981 let data_converter = client.data_converter().clone();
1982 let payload_converter = data_converter.payload_converter();
1983 let context_data = SerializationContextData::Workflow(
1984 WorkflowSerializationContext::new(),
1985 );
1986 let context = SerializationContext::new(&context_data, payload_converter);
1987 let workflow_payloads = workflow_args.serialize_payloads(&context);
1988 let signal_payloads = signal_args.serialize_payloads(&context);
1989 drop(workflow_args);
1990 drop(signal_args);
1991 let workflow_payloads = data_converter
1992 .codec()
1993 .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), workflow_payloads?)
1994 .await?;
1995 let signal_payloads = data_converter
1996 .codec()
1997 .encode(&SerializationContextData::Workflow(WorkflowSerializationContext::new()), signal_payloads?)
1998 .await?;
1999 let workflow_id = options.workflow_id.clone();
2000 let memo = options.encoded_memo(&data_converter).await?;
2001 let mut start_request = build_start_workflow_request(
2002 &client,
2003 workflow_type,
2004 workflow_payloads.into_payloads(),
2005 memo,
2006 options,
2007 );
2008 if let Some(task_queue) = &mut start_request.task_queue {
2009 task_queue.kind = TaskQueueKind::Normal as i32;
2010 }
2011 let mut request = SignalWithStartWorkflowExecutionRequest {
2012 namespace: start_request.namespace,
2013 workflow_id: start_request.workflow_id,
2014 workflow_type: start_request.workflow_type,
2015 task_queue: start_request.task_queue,
2016 input: start_request.input,
2017 workflow_execution_timeout: start_request.workflow_execution_timeout,
2018 workflow_run_timeout: start_request.workflow_run_timeout,
2019 workflow_task_timeout: start_request.workflow_task_timeout,
2020 identity: start_request.identity,
2021 request_id: start_request.request_id,
2022 workflow_id_reuse_policy: start_request.workflow_id_reuse_policy,
2023 workflow_id_conflict_policy: start_request.workflow_id_conflict_policy,
2024 signal_name,
2025 signal_input: Some(Payloads {
2026 payloads: signal_payloads,
2027 }),
2028 retry_policy: start_request.retry_policy,
2029 cron_schedule: start_request.cron_schedule,
2030 memo: start_request.memo,
2031 search_attributes: start_request.search_attributes,
2032 header: start_request.header,
2033 workflow_start_delay: start_request.workflow_start_delay,
2034 user_metadata: start_request.user_metadata,
2035 links: start_request.links,
2036 versioning_override: start_request.versioning_override,
2037 priority: start_request.priority,
2038 time_skipping_config: start_request.time_skipping_config,
2039 ..Default::default()
2040 }
2041 .into_request();
2042 rpc_options.apply_to(&mut request);
2043 let run_id = WorkflowService::signal_with_start_workflow_execution(
2044 &mut client,
2045 request,
2046 )
2047 .await?
2048 .into_inner()
2049 .run_id;
2050 Ok(StartWorkflowOutput::new(workflow_id, run_id))
2051 })
2052 }
2053 }),
2054 )
2055 .await?;
2056 let StartWorkflowOutput {
2057 workflow_id,
2058 run_id,
2059 } = interceptor_output;
2060
2061 Ok(WorkflowHandle::new(
2062 self.clone(),
2063 WorkflowExecutionInfo {
2064 namespace,
2065 workflow_id,
2066 run_id: Some(run_id.clone()),
2067 first_execution_run_id: Some(run_id),
2068 },
2069 ))
2070 }
2071
2072 async fn start_update_with_start_workflow<W, U>(
2073 &self,
2074 workflow: W,
2075 workflow_input: W::Input,
2076 update: U,
2077 update_input: U::Input,
2078 options: WorkflowUpdateWithStartOptions,
2079 ) -> Result<WorkflowUpdateHandle<Self, U::Output>, WorkflowUpdateWithStartError>
2080 where
2081 W: HasWorkflowDefinition,
2082 W::Input: Send,
2083 U: UpdateDefinition<Workflow = W::Run>,
2084 U::Input: Send,
2085 {
2086 let output = interceptors::call_update_with_start_workflow(
2087 self.client_interceptors(),
2088 UpdateWithStartWorkflowInput::new(
2089 workflow.name().to_owned(),
2090 workflow_input,
2091 update.name().to_owned(),
2092 update_input,
2093 options,
2094 ),
2095 Next::new({
2096 let client = (*self).clone();
2097 move |input: UpdateWithStartWorkflowInput| -> BoxFuture<
2098 '_,
2099 Result<UpdateWithStartWorkflowOutput, WorkflowUpdateWithStartError>,
2100 > {
2101 let mut client = client;
2102 Box::pin(async move {
2103 let UpdateWithStartWorkflowInput {
2104 workflow_type,
2105 update_name,
2106 options,
2107 rpc_options,
2108 workflow_args,
2109 update_args,
2110 } = input;
2111 let (start_options, update_id, update_header) = options.into_parts();
2112
2113 let data_converter = client.data_converter().clone();
2114 let (unencoded_workflow_payloads, unencoded_update_payloads) = {
2115 let payload_converter = data_converter.payload_converter();
2116 let context_data = SerializationContextData::Workflow(
2117 WorkflowSerializationContext::new(),
2118 );
2119 let context =
2120 SerializationContext::new(&context_data, payload_converter);
2121 (
2122 workflow_args.serialize_payloads(&context),
2123 update_args.serialize_payloads(&context),
2124 )
2125 };
2126 drop(workflow_args);
2127 drop(update_args);
2128 let (workflow_payloads, update_payloads) = try_join(
2131 data_converter.codec().encode(
2132 &SerializationContextData::Workflow(
2133 WorkflowSerializationContext::new(),
2134 ),
2135 unencoded_workflow_payloads?,
2136 ),
2137 data_converter.codec().encode(
2138 &SerializationContextData::Workflow(
2139 WorkflowSerializationContext::new(),
2140 ),
2141 unencoded_update_payloads?,
2142 ),
2143 )
2144 .await?;
2145
2146 let namespace = client.namespace();
2147 let workflow_id = start_options.workflow_id.clone();
2148 let memo = start_options.encoded_memo(&data_converter).await?;
2149 let start_request = build_start_workflow_request(
2150 &client,
2151 workflow_type,
2152 workflow_payloads.into_payloads(),
2153 memo,
2154 start_options,
2155 );
2156
2157 let update_id = update_id.unwrap_or_else(|| Uuid::new_v4().to_string());
2158 let update_request = workflow_handle::build_update_workflow_request(
2159 namespace.clone(),
2160 client.identity(),
2161 workflow_id.clone(),
2162 String::new(),
2163 update_id.clone(),
2164 update_name,
2165 update_header,
2166 update_payloads,
2167 );
2168
2169 let request = ExecuteMultiOperationRequest {
2170 namespace,
2171 operations: vec![
2172 execute_multi_operation_request::Operation {
2173 operation: Some(MultiOperationRequest::StartWorkflow(
2174 start_request,
2175 )),
2176 },
2177 execute_multi_operation_request::Operation {
2178 operation: Some(MultiOperationRequest::UpdateWorkflow(
2179 update_request,
2180 )),
2181 },
2182 ],
2183 resource_id: workflow_id.clone(),
2184 };
2185
2186 let (start_response, update_response) = loop {
2187 let mut rpc_request = request.clone().into_request();
2188 rpc_options.apply_to(&mut rpc_request);
2189 let response =
2190 WorkflowService::execute_multi_operation(&mut client, rpc_request)
2191 .await
2192 .map_err(WorkflowUpdateWithStartError::from_status)?
2193 .into_inner();
2194
2195 let [start_response, update_response]: [_; 2] =
2196 response.responses.try_into().map_err(|_| {
2197 WorkflowUpdateWithStartError::Other(
2198 "Server response did not include exactly two operation \
2199 responses"
2200 .into(),
2201 )
2202 })?;
2203 let (
2204 Some(MultiOperationResponse::StartWorkflow(start_response)),
2205 Some(MultiOperationResponse::UpdateWorkflow(update_response)),
2206 ) = (start_response.response, update_response.response)
2207 else {
2208 return Err(WorkflowUpdateWithStartError::Other(
2209 "Server response did not include start and update operation \
2210 responses in request order"
2211 .into(),
2212 ));
2213 };
2214
2215 if update_response.stage
2216 < UpdateWorkflowExecutionLifecycleStage::Accepted as i32
2217 {
2218 continue;
2219 }
2220 break (start_response, update_response);
2221 };
2222
2223 let run_id = update_response
2224 .update_ref
2225 .as_ref()
2226 .and_then(|reference| reference.workflow_execution.as_ref())
2227 .map(|execution| execution.run_id.clone())
2228 .filter(|run_id| !run_id.is_empty())
2229 .or_else(|| {
2230 (!start_response.run_id.is_empty()).then_some(start_response.run_id)
2231 });
2232 Ok(UpdateWithStartWorkflowOutput::new(
2233 workflow_id,
2234 update_id,
2235 run_id,
2236 update_response.outcome,
2237 ))
2238 })
2239 }
2240 }),
2241 )
2242 .await?;
2243 Ok(WorkflowUpdateHandle::new(
2244 self.clone(),
2245 output.update_id,
2246 output.workflow_id,
2247 output.run_id,
2248 output.known_outcome,
2249 ))
2250 }
2251
2252 async fn execute_update_with_start_workflow<W, U>(
2253 &self,
2254 workflow: W,
2255 workflow_input: W::Input,
2256 update: U,
2257 update_input: U::Input,
2258 options: WorkflowUpdateWithStartOptions,
2259 ) -> Result<U::Output, WorkflowUpdateWithStartError>
2260 where
2261 W: HasWorkflowDefinition,
2262 W::Input: Send,
2263 U: UpdateDefinition<Workflow = W::Run>,
2264 U::Input: Send,
2265 {
2266 let rpc_options = options.rpc_options.clone();
2267 let update_handle = WorkflowClientTrait::start_update_with_start_workflow(
2268 self,
2269 workflow,
2270 workflow_input,
2271 update,
2272 update_input,
2273 options,
2274 )
2275 .await?;
2276 let result = update_handle
2277 .get_result(rpc_options)
2278 .await
2279 .map_err(WorkflowUpdateWithStartError::Update)?;
2280 Ok(result)
2281 }
2282
2283 fn get_workflow_handle<W: HasWorkflowDefinition>(
2284 &self,
2285 workflow_id: impl Into<String>,
2286 ) -> WorkflowHandle<Self, W>
2287 where
2288 Self: Sized,
2289 {
2290 WorkflowHandle::new(
2291 self.clone(),
2292 WorkflowExecutionInfo {
2293 namespace: self.namespace(),
2294 workflow_id: workflow_id.into(),
2295 run_id: None,
2296 first_execution_run_id: None,
2297 },
2298 )
2299 }
2300
2301 fn list_workflows(
2302 &self,
2303 query: impl Into<String>,
2304 opts: WorkflowListOptions,
2305 ) -> ListWorkflowsStream {
2306 let client = self.clone();
2307 let namespace = self.namespace();
2308 let query = query.into();
2309 let limit = opts.limit;
2310 let rpc_options = opts.rpc_options;
2311
2312 let initial_state = (Vec::new(), VecDeque::new(), 0, false);
2314
2315 let stream = stream::unfold(
2316 initial_state,
2317 move |(next_page_token, mut buffer, mut yielded, exhausted)| {
2318 let client = client.clone();
2319 let namespace = namespace.clone();
2320 let query = query.clone();
2321 let rpc_options = rpc_options.clone();
2322
2323 async move {
2324 if let Some(l) = limit
2325 && yielded >= l
2326 {
2327 return None;
2328 }
2329
2330 if let Some(exec) = buffer.pop_front() {
2331 yielded += 1;
2332 return Some((Ok(exec), (next_page_token, buffer, yielded, exhausted)));
2333 }
2334
2335 if exhausted {
2336 return None;
2337 }
2338
2339 let response = interceptors::call_list_workflows_page(
2340 client.client_interceptors(),
2341 ListWorkflowsPageInput {
2342 query,
2343 next_page_token: next_page_token.clone(),
2344 rpc_options,
2345 },
2346 Next::new({
2347 let mut rpc_client = client.clone();
2348 move |input: ListWorkflowsPageInput| -> BoxFuture<
2349 '_,
2350 Result<ListWorkflowsPageOutput, ClientError>,
2351 > {
2352 Box::pin(async move {
2353 let mut request = ListWorkflowExecutionsRequest {
2354 namespace,
2355 page_size: 0,
2356 next_page_token: input.next_page_token,
2357 query: input.query,
2358 }
2359 .into_request();
2360 input.rpc_options.apply_to(&mut request);
2361 let response = WorkflowService::list_workflow_executions(
2362 &mut rpc_client,
2363 request,
2364 )
2365 .await?
2366 .into_inner();
2367 Ok(ListWorkflowsPageOutput::new(
2368 response.executions,
2369 response.next_page_token,
2370 ))
2371 })
2372 }
2373 }),
2374 )
2375 .await;
2376
2377 match response {
2378 Ok(mut output) => {
2379 let new_exhausted = output.next_page_token.is_empty();
2380 let new_token = output.next_page_token;
2381
2382 let data_converter = client.data_converter().clone();
2383 for execution in &mut output.executions {
2384 if let Some(memo) = execution.memo.as_mut()
2385 && let Err(err) = decode_payloads(
2386 memo,
2387 data_converter.codec(),
2388 &SerializationContextData::Workflow(
2389 WorkflowSerializationContext::new(),
2390 ),
2391 )
2392 .await
2393 {
2394 return Some((
2395 Err(ClientError::from(err)),
2396 (new_token, buffer, yielded, true),
2397 ));
2398 }
2399 }
2400 buffer = output
2401 .executions
2402 .into_iter()
2403 .map(|raw| {
2404 WorkflowExecution::new_with_data_converter(
2405 raw,
2406 data_converter.clone(),
2407 )
2408 })
2409 .collect();
2410
2411 if let Some(exec) = buffer.pop_front() {
2412 yielded += 1;
2413 Some((Ok(exec), (new_token, buffer, yielded, new_exhausted)))
2414 } else {
2415 None
2416 }
2417 }
2418 Err(e) => Some((Err(e), (next_page_token, buffer, yielded, true))),
2419 }
2420 }
2421 },
2422 );
2423
2424 ListWorkflowsStream::new(Box::pin(stream))
2425 }
2426
2427 async fn count_workflows(
2428 &self,
2429 query: impl Into<String>,
2430 opts: WorkflowCountOptions,
2431 ) -> Result<WorkflowExecutionCount, ClientError> {
2432 let output = interceptors::call_count_workflows(
2433 self.client_interceptors(),
2434 CountWorkflowsInput {
2435 query: query.into(),
2436 options: opts,
2437 },
2438 Next::new({
2439 let mut client = (*self).clone();
2440 move |input: CountWorkflowsInput| -> BoxFuture<
2441 '_,
2442 Result<CountWorkflowsOutput, ClientError>,
2443 > {
2444 Box::pin(async move {
2445 let mut request = CountWorkflowExecutionsRequest {
2446 namespace: client.namespace(),
2447 query: input.query,
2448 }
2449 .into_request();
2450 input.options.rpc_options.apply_to(&mut request);
2451 let response = WorkflowService::count_workflow_executions(
2452 &mut client,
2453 request,
2454 )
2455 .await?
2456 .into_inner();
2457 Ok(CountWorkflowsOutput::new(response))
2458 })
2459 }
2460 }),
2461 )
2462 .await?;
2463
2464 Ok(WorkflowExecutionCount::from_response(output.response))
2465 }
2466
2467 fn get_async_activity_handle(&self, identifier: ActivityIdentifier) -> AsyncActivityHandle<Self>
2468 where
2469 Self: Sized,
2470 {
2471 AsyncActivityHandle::new(self.clone(), identifier)
2472 }
2473
2474 async fn start_activity<A>(
2475 &self,
2476 activity: A,
2477 input: A::Input,
2478 options: ActivityStartOptions,
2479 ) -> Result<ActivityHandle<Self, A>, StartActivityError>
2480 where
2481 Self: Sized,
2482 A: ActivityDefinition,
2483 {
2484 let mut client = self.clone();
2485 let dc = client.data_converter();
2486 let sc = &SerializationContextData::Activity(ActivitySerializationContext::new());
2487
2488 let user_metadata = {
2489 let summary = match &options.summary {
2490 Some(summary) => Some(dc.to_payload(sc, summary).await?),
2491 None => None,
2492 };
2493 let details = match &options.static_details {
2494 Some(details) => Some(dc.to_payload(sc, details).await?),
2495 None => None,
2496 };
2497 (summary.is_some() || details.is_some()).then_some(UserMetadata { summary, details })
2498 };
2499
2500 let resp = client
2501 .start_activity_execution(
2502 StartActivityExecutionRequest {
2503 namespace: client.namespace(),
2504 identity: client.identity(),
2505 request_id: Uuid::new_v4().to_string(),
2506 activity_id: options.id.clone(),
2507 activity_type: Some(ActivityType {
2508 name: activity.name().to_string(),
2509 }),
2510 task_queue: Some(TaskQueue {
2511 name: options.task_queue,
2512 kind: TaskQueueKind::Normal.into(),
2513 normal_name: "".to_string(),
2514 }),
2515 schedule_to_close_timeout: try_into_or_box_err(
2516 options.close_timeouts.schedule_to_close(),
2517 StartActivityError::Other,
2518 )?,
2519 schedule_to_start_timeout: try_into_or_box_err(
2520 options.schedule_to_start_timeout,
2521 StartActivityError::Other,
2522 )?,
2523 start_to_close_timeout: try_into_or_box_err(
2524 options.close_timeouts.start_to_close(),
2525 StartActivityError::Other,
2526 )?,
2527 heartbeat_timeout: try_into_or_box_err(
2528 options.heartbeat_timeout,
2529 StartActivityError::Other,
2530 )?,
2531 retry_policy: options.retry_policy.map(Into::into),
2532 input: dc.to_payloads(sc, &input).await?.into_payloads(),
2533 id_reuse_policy: ProtoActivityIdReusePolicy::from(options.id_reuse_policy)
2534 .into(),
2535 id_conflict_policy: ProtoActivityIdConflictPolicy::from(
2536 options.id_conflict_policy,
2537 )
2538 .into(),
2539 search_attributes: options.search_attributes.map(SearchAttributes::into_proto),
2540 header: options.header,
2541 user_metadata,
2542 priority: Some(options.priority.into()),
2543 start_delay: try_into_or_box_err(
2544 options.start_delay,
2545 StartActivityError::Other,
2546 )?,
2547 ..Default::default()
2548 }
2549 .into_request(),
2550 )
2551 .await?
2552 .into_inner();
2553
2554 Ok(ActivityHandle::new(
2555 client,
2556 options.id,
2557 (!resp.run_id.is_empty()).then_some(resp.run_id),
2558 ))
2559 }
2560
2561 fn get_activity_handle<A>(
2562 &self,
2563 _activity: A,
2564 id: impl Into<String>,
2565 run_id: Option<String>,
2566 ) -> ActivityHandle<Self, A>
2567 where
2568 Self: Sized,
2569 A: ActivityDefinition,
2570 {
2571 ActivityHandle::new(self.clone(), id.into(), run_id)
2572 }
2573
2574 fn get_untyped_activity_handle(
2575 &self,
2576 id: impl Into<String>,
2577 run_id: Option<String>,
2578 ) -> ActivityHandle<Self, UntypedActivity>
2579 where
2580 Self: Sized,
2581 {
2582 ActivityHandle::new(self.clone(), id.into(), run_id)
2583 }
2584
2585 fn list_activities(
2586 &self,
2587 query: impl Into<String>,
2588 _options: ActivityListOptions,
2589 ) -> ListActivitiesStream {
2590 let client = self.clone();
2591 let namespace = client.namespace();
2592 let query = query.into();
2593
2594 ListActivitiesStream::new(stream::unfold(
2595 Some(vec![]), move |next_page_token| {
2597 let mut client = client.clone();
2598 let namespace = namespace.clone();
2599 let query = query.clone();
2600
2601 async move {
2602 #[allow(clippy::question_mark)]
2604 let Some(token): Option<Vec<u8>> = next_page_token else {
2605 return None;
2606 };
2607
2608 match WorkflowService::list_activity_executions(
2609 &mut client,
2610 ListActivityExecutionsRequest {
2611 namespace,
2612 page_size: 0, next_page_token: token.clone(),
2614 query,
2615 }
2616 .into_request(),
2617 )
2618 .await
2619 .map(|r| r.into_inner())
2620 {
2621 Ok(resp) => Some((
2622 Ok(resp.executions),
2623 (!resp.next_page_token.is_empty()).then_some(resp.next_page_token),
2624 )),
2625 Err(e) => Some((Err(e.into()), Some(token))),
2626 }
2627 }
2628 },
2629 ))
2630 }
2631
2632 async fn count_activities(
2633 &self,
2634 query: impl Into<String>,
2635 _options: ActivityCountOptions,
2636 ) -> Result<ActivityExecutionCount, ClientError> {
2637 let mut client = self.clone();
2638 let resp = client
2639 .count_activity_executions(
2640 CountActivityExecutionsRequest {
2641 namespace: client.namespace(),
2642 query: query.into(),
2643 }
2644 .into_request(),
2645 )
2646 .await?
2647 .into_inner();
2648 Ok(ActivityExecutionCount::from_response(resp))
2649 }
2650}
2651
2652macro_rules! dbg_panic {
2653 ($($arg:tt)*) => {
2654 use tracing::error;
2655 error!($($arg)*);
2656 debug_assert!(false, $($arg)*);
2657 };
2658}
2659pub(crate) use dbg_panic;
2660
2661fn try_into_or_box_err<A, B, E, MapErr>(val: Option<A>, map_err: MapErr) -> Result<Option<B>, E>
2662where
2663 A: TryInto<B>,
2664 <A as TryInto<B>>::Error: Error + Send + Sync + 'static,
2665 MapErr: FnOnce(Box<dyn Error + Send + Sync + 'static>) -> E,
2666{
2667 val.map(TryInto::try_into)
2668 .transpose()
2669 .map_err(|e| map_err(Box::from(e)))
2670}
2671
2672#[cfg(test)]
2673mod tests {
2674 use super::*;
2675 use crate::callback_based::CallbackBasedGrpcService;
2676 use std::{
2677 sync::atomic::{AtomicUsize, Ordering},
2678 time::Instant,
2679 };
2680 use temporalio_common::search_attributes::SearchAttributeKey;
2681 use tonic::{Status, metadata::Ascii};
2682 use url::Url;
2683
2684 #[test]
2685 fn count_aggregation_group_gets_typed_value() {
2686 let attrs = SearchAttributes::new([SearchAttributeKey::int("group").value_set(42)]);
2687 let group = WorkflowCountAggregationGroup {
2688 raw: count_workflow_executions_response::AggregationGroup {
2689 group_values: vec![attrs.raw_payload("group").unwrap().clone()],
2690 count: 1,
2691 },
2692 };
2693
2694 assert_eq!(group.get::<i64>(0), Some(42));
2695 assert_eq!(group.get::<i64>(1), None);
2696 assert!(group.try_get::<String>(0).is_err());
2697 assert_eq!(group.try_get::<i64>(1).unwrap(), None);
2698 }
2699
2700 fn connection_options_for_system_info_test(
2701 service_override: CallbackBasedGrpcService,
2702 ) -> ConnectionOptions {
2703 ConnectionOptions::new(Url::parse("http://localhost:7233").unwrap())
2704 .service_override(service_override)
2705 .dns_load_balancing(None)
2706 .build()
2707 }
2708
2709 #[test]
2710 fn applies_headers() {
2711 let headers = Arc::new(RwLock::new(ClientHeaders {
2713 user_headers: HashMap::new(),
2714 user_binary_headers: HashMap::new(),
2715 api_key: Some("my-api-key".to_owned()),
2716 }));
2717 headers.clone().write().user_headers.insert(
2718 "my-meta-key".parse().unwrap(),
2719 "my-meta-val".parse().unwrap(),
2720 );
2721 headers.clone().write().user_binary_headers.insert(
2722 "my-bin-meta-key-bin".parse().unwrap(),
2723 vec![1, 2, 3].try_into().unwrap(),
2724 );
2725 let mut interceptor = ServiceCallInterceptor {
2726 client_name: "cute-kitty".to_string(),
2727 client_version: "0.1.0".to_string(),
2728 headers: headers.clone(),
2729 };
2730
2731 let req = interceptor.call(tonic::Request::new(())).unwrap();
2733 assert_eq!(req.metadata().get("my-meta-key").unwrap(), "my-meta-val");
2734 assert_eq!(
2735 req.metadata().get("authorization").unwrap(),
2736 "Bearer my-api-key"
2737 );
2738 assert_eq!(
2739 req.metadata().get_bin("my-bin-meta-key-bin").unwrap(),
2740 vec![1, 2, 3].as_slice()
2741 );
2742
2743 let mut req = tonic::Request::new(());
2745 req.metadata_mut()
2746 .insert("my-meta-key", "my-meta-val2".parse().unwrap());
2747 req.metadata_mut()
2748 .insert("authorization", "my-api-key2".parse().unwrap());
2749 req.metadata_mut()
2750 .insert_bin("my-bin-meta-key-bin", vec![4, 5, 6].try_into().unwrap());
2751 let req = interceptor.call(req).unwrap();
2752 assert_eq!(req.metadata().get("my-meta-key").unwrap(), "my-meta-val2");
2753 assert_eq!(req.metadata().get("authorization").unwrap(), "my-api-key2");
2754 assert_eq!(
2755 req.metadata().get_bin("my-bin-meta-key-bin").unwrap(),
2756 vec![4, 5, 6].as_slice()
2757 );
2758
2759 headers.clone().write().user_headers.insert(
2761 "authorization".parse().unwrap(),
2762 "my-api-key3".parse().unwrap(),
2763 );
2764 let req = interceptor.call(tonic::Request::new(())).unwrap();
2765 assert_eq!(req.metadata().get("my-meta-key").unwrap(), "my-meta-val");
2766 assert_eq!(req.metadata().get("authorization").unwrap(), "my-api-key3");
2767
2768 headers.clone().write().user_headers.clear();
2770 headers.clone().write().user_binary_headers.clear();
2771 headers.clone().write().api_key.take();
2772 let req = interceptor.call(tonic::Request::new(())).unwrap();
2773 assert!(!req.metadata().contains_key("my-meta-key"));
2774 assert!(!req.metadata().contains_key("authorization"));
2775 assert!(!req.metadata().contains_key("my-bin-meta-key-bin"));
2776
2777 let mut req = tonic::Request::new(());
2779 req.metadata_mut()
2780 .insert("grpc-timeout", "1S".parse().unwrap());
2781 let req = interceptor.call(req).unwrap();
2782 assert_eq!(
2783 req.metadata().get("grpc-timeout").unwrap(),
2784 "1S".parse::<MetadataValue<Ascii>>().unwrap()
2785 );
2786 }
2787
2788 #[test]
2789 fn invalid_ascii_header_key() {
2790 let invalid_headers = {
2791 let mut h = HashMap::new();
2792 h.insert("x-binary-key-bin".to_owned(), "value".to_owned());
2793 h
2794 };
2795
2796 let result = parse_ascii_headers(invalid_headers);
2797 assert!(result.is_err());
2798 assert_eq!(
2799 result.err().unwrap().to_string(),
2800 "Invalid ASCII header key 'x-binary-key-bin': invalid gRPC metadata key name"
2801 );
2802 }
2803
2804 #[test]
2805 fn invalid_ascii_header_value() {
2806 let invalid_headers = {
2807 let mut h = HashMap::new();
2808 h.insert("x-ascii-key".to_owned(), "\x00value".to_owned());
2810 h
2811 };
2812
2813 let result = parse_ascii_headers(invalid_headers);
2814 assert!(result.is_err());
2815 assert_eq!(
2816 result.err().unwrap().to_string(),
2817 "Invalid ASCII header value for key 'x-ascii-key': failed to parse metadata value"
2818 );
2819 }
2820
2821 #[test]
2822 fn invalid_binary_header_key() {
2823 let invalid_headers = {
2824 let mut h = HashMap::new();
2825 h.insert("x-ascii-key".to_owned(), vec![1, 2, 3]);
2826 h
2827 };
2828
2829 let result = parse_binary_headers(invalid_headers);
2830 assert!(result.is_err());
2831 assert_eq!(
2832 result.err().unwrap().to_string(),
2833 "Invalid binary header key 'x-ascii-key': invalid gRPC metadata key name"
2834 );
2835 }
2836
2837 #[test]
2838 fn keep_alive_defaults() {
2839 let opts = ConnectionOptions::new(Url::parse("https://smolkitty").unwrap())
2840 .identity("enchicat".to_string())
2841 .client_name("cute-kitty".to_string())
2842 .client_version("0.1.0".to_string())
2843 .build();
2844 assert_eq!(
2845 opts.keep_alive.clone().unwrap().interval,
2846 ClientKeepAliveOptions::default().interval
2847 );
2848 assert_eq!(
2849 opts.keep_alive.clone().unwrap().timeout,
2850 ClientKeepAliveOptions::default().timeout
2851 );
2852
2853 let opts = ConnectionOptions::new(Url::parse("https://smolkitty").unwrap())
2855 .identity("enchicat".to_string())
2856 .client_name("cute-kitty".to_string())
2857 .client_version("0.1.0".to_string())
2858 .keep_alive(None)
2859 .build();
2860 dbg!(&opts.keep_alive);
2861 assert!(opts.keep_alive.is_none());
2862 }
2863
2864 #[rstest::rstest]
2865 #[case(
2866 "unknown method GetSystemInfo for service temporal.api.workflowservice.v1.WorkflowService"
2867 )]
2868 #[case("Method temporal.api.workflowservice.v1.WorkflowService/GetSystemInfo is unimplemented")]
2869 #[case(
2870 "The server does not implement the method /temporal.api.workflowservice.v1.WorkflowService/GetSystemInfo"
2871 )]
2872 #[tokio::test]
2873 async fn get_system_info_missing_method_falls_back_to_empty_capabilities(
2874 #[case] message: &'static str,
2875 ) {
2876 let attempts = Arc::new(AtomicUsize::new(0));
2877 let attempts_clone = attempts.clone();
2878 let service_override = CallbackBasedGrpcService {
2879 callback: Arc::new(move |req| {
2880 let attempts = attempts_clone.clone();
2881 Box::pin(async move {
2882 assert_eq!(req.rpc, "GetSystemInfo");
2883 attempts.fetch_add(1, Ordering::SeqCst);
2884 Err(Status::unimplemented(message))
2885 })
2886 }),
2887 };
2888
2889 let connection =
2890 Connection::connect(connection_options_for_system_info_test(service_override))
2891 .await
2892 .unwrap();
2893
2894 assert!(connection.capabilities().is_none());
2895 assert_eq!(attempts.load(Ordering::SeqCst), 1);
2896 }
2897
2898 #[tokio::test]
2899 async fn get_system_info_non_missing_unimplemented_fails_connect() {
2900 let attempts = Arc::new(AtomicUsize::new(0));
2901 let attempts_clone = attempts.clone();
2902 let service_override = CallbackBasedGrpcService {
2903 callback: Arc::new(move |req| {
2904 let attempts = attempts_clone.clone();
2905 Box::pin(async move {
2906 assert_eq!(req.rpc, "GetSystemInfo");
2907 attempts.fetch_add(1, Ordering::SeqCst);
2908 Err(Status::unimplemented("backend temporarily unimplemented"))
2909 })
2910 }),
2911 };
2912
2913 let err =
2914 match Connection::connect(connection_options_for_system_info_test(service_override))
2915 .await
2916 {
2917 Ok(_) => panic!("connection should fail"),
2918 Err(err) => err,
2919 };
2920
2921 assert!(matches!(
2922 err,
2923 ClientConnectError::SystemInfoCallError(status)
2924 if status.code() == Code::Unimplemented
2925 && status.message() == "backend temporarily unimplemented"
2926 ));
2927 assert_eq!(attempts.load(Ordering::SeqCst), 1);
2928 }
2929
2930 #[tokio::test]
2931 async fn connect_timeout_bounds_connection_attempt() {
2932 let url = Url::parse("http://10.255.255.1:7233").unwrap();
2933 let opts = ConnectionOptions::new(url)
2934 .connect_timeout(Duration::from_millis(500))
2935 .build();
2936 let start = Instant::now();
2937 let result = Connection::connect(opts).await;
2938 assert!(result.is_err(), "connection should fail");
2939 assert!(start.elapsed() < Duration::from_secs(2));
2940 }
2941
2942 mod tls_custom_verifier_tests {
2943 use super::*;
2944 use tokio_rustls::rustls::{
2945 DigitallySignedStruct, Error as RustlsError, SignatureScheme,
2946 client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier},
2947 pki_types::{CertificateDer, ServerName, UnixTime},
2948 };
2949
2950 #[derive(Debug)]
2953 struct MockVerifier;
2954
2955 impl ServerCertVerifier for MockVerifier {
2956 fn verify_server_cert(
2957 &self,
2958 _end_entity: &CertificateDer<'_>,
2959 _intermediates: &[CertificateDer<'_>],
2960 _server_name: &ServerName<'_>,
2961 _ocsp_response: &[u8],
2962 _now: UnixTime,
2963 ) -> Result<ServerCertVerified, RustlsError> {
2964 Ok(ServerCertVerified::assertion())
2965 }
2966
2967 fn verify_tls12_signature(
2968 &self,
2969 _message: &[u8],
2970 _cert: &CertificateDer<'_>,
2971 _dss: &DigitallySignedStruct,
2972 ) -> Result<HandshakeSignatureValid, RustlsError> {
2973 Ok(HandshakeSignatureValid::assertion())
2974 }
2975
2976 fn verify_tls13_signature(
2977 &self,
2978 _message: &[u8],
2979 _cert: &CertificateDer<'_>,
2980 _dss: &DigitallySignedStruct,
2981 ) -> Result<HandshakeSignatureValid, RustlsError> {
2982 Ok(HandshakeSignatureValid::assertion())
2983 }
2984
2985 fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
2986 vec![
2987 SignatureScheme::ECDSA_NISTP256_SHA256,
2988 SignatureScheme::RSA_PSS_SHA256,
2989 ]
2990 }
2991 }
2992
2993 #[tokio::test]
2994 async fn add_tls_to_channel_with_custom_verifier() {
2995 let tls_opts = TlsOptions::builder()
2996 .server_cert_verifier(Arc::new(MockVerifier))
2997 .domain("test.temporal.io".to_string())
2998 .build();
2999 let endpoint = tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3000 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3001 assert!(
3002 matches!(&result, Ok(TlsConfigResult::Standard(_))),
3003 "add_tls_to_channel should succeed with a custom verifier: {:?}",
3004 result.err()
3005 );
3006 }
3007
3008 #[tokio::test]
3009 async fn add_tls_to_channel_with_verifier_and_ca_cert_fails() {
3010 let tls_opts = TlsOptions::builder()
3013 .server_root_ca_cert(b"some-ca-cert-bytes".to_vec())
3014 .server_cert_verifier(Arc::new(MockVerifier))
3015 .domain("test.temporal.io".to_string())
3016 .build();
3017 let endpoint = tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3018 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3019 assert!(
3020 matches!(result, Err(ClientConnectError::InvalidConfig(_))),
3021 "add_tls_to_channel should fail with InvalidConfig when both CA cert and verifier are set: {:?}",
3022 result
3023 );
3024 }
3025
3026 #[tokio::test]
3027 async fn add_tls_to_channel_without_verifier_still_works() {
3028 let tls_opts = TlsOptions::builder()
3030 .domain("test.temporal.io".to_string())
3031 .build();
3032 let endpoint = tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3033 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3034 assert!(
3035 matches!(&result, Ok(TlsConfigResult::Standard(_))),
3036 "add_tls_to_channel should succeed without a verifier (native roots): {:?}",
3037 result.err()
3038 );
3039 }
3040
3041 #[cfg(feature = "dynamic-tls")]
3044 mod dynamic_cert_tests {
3045 use super::*;
3046
3047 #[derive(Debug)]
3050 struct MockClientCertResolver;
3051
3052 impl tokio_rustls::rustls::client::ResolvesClientCert for MockClientCertResolver {
3053 fn resolve(
3054 &self,
3055 _acceptable_issuers: &[&[u8]],
3056 _sigschemes: &[tokio_rustls::rustls::SignatureScheme],
3057 ) -> Option<Arc<tokio_rustls::rustls::sign::CertifiedKey>> {
3058 None }
3060
3061 fn has_certs(&self) -> bool {
3062 false
3063 }
3064 }
3065
3066 #[tokio::test]
3067 async fn add_tls_with_client_cert_resolver_returns_custom_connector() {
3068 let resolver = Arc::new(MockClientCertResolver);
3069 let tls_opts = TlsOptions {
3070 client_cert_resolver: Some(resolver),
3071 domain: Some("test.temporal.io".to_string()),
3072 ..Default::default()
3073 };
3074 let endpoint =
3075 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3076 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3077 match result {
3078 Ok(TlsConfigResult::CustomConnector {
3079 domain,
3080 rustls_config,
3081 ..
3082 }) => {
3083 assert_eq!(domain, "test.temporal.io");
3084 assert_eq!(rustls_config.alpn_protocols, vec![b"h2".to_vec()]);
3086 }
3087 other => panic!(
3088 "Expected TlsConfigResult::CustomConnector, got {:?}",
3089 other.err()
3090 ),
3091 }
3092 }
3093
3094 #[tokio::test]
3095 async fn add_tls_with_client_cert_resolver_inherits_domain_from_endpoint() {
3096 let resolver = Arc::new(MockClientCertResolver);
3097 let tls_opts = TlsOptions {
3098 client_cert_resolver: Some(resolver),
3099 ..Default::default()
3101 };
3102 let endpoint =
3103 tonic::transport::Channel::from_static("https://my-server.example.com:7233");
3104 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3105 match result {
3106 Ok(TlsConfigResult::CustomConnector { domain, .. }) => {
3107 assert_eq!(domain, "my-server.example.com");
3108 }
3109 other => panic!(
3110 "Expected TlsConfigResult::CustomConnector, got {:?}",
3111 other.err()
3112 ),
3113 }
3114 }
3115
3116 #[tokio::test]
3117 async fn add_tls_with_resolver_and_custom_verifier() {
3118 let resolver = Arc::new(MockClientCertResolver);
3119 let tls_opts = TlsOptions {
3120 client_cert_resolver: Some(resolver),
3121 server_cert_verifier: Some(Arc::new(MockVerifier)),
3122 domain: Some("test.temporal.io".to_string()),
3123 ..Default::default()
3124 };
3125 let endpoint =
3126 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3127 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3128 assert!(
3129 matches!(&result, Ok(TlsConfigResult::CustomConnector { .. })),
3130 "Should succeed when combining cert resolver with custom server verifier: {:?}",
3131 result.err()
3132 );
3133 }
3134
3135 #[tokio::test]
3136 async fn add_tls_with_resolver_and_custom_ca_cert() {
3137 let ca_pem = include_bytes!("../tests/testdata/ca.pem");
3139 let resolver = Arc::new(MockClientCertResolver);
3140 let tls_opts = TlsOptions {
3141 client_cert_resolver: Some(resolver),
3142 server_root_ca_cert: Some(ca_pem.to_vec()),
3143 domain: Some("test.temporal.io".to_string()),
3144 ..Default::default()
3145 };
3146 let endpoint =
3147 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3148 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3149 assert!(
3150 matches!(&result, Ok(TlsConfigResult::CustomConnector { .. })),
3151 "Should succeed when combining cert resolver with custom CA cert: {:?}",
3152 result.err()
3153 );
3154 }
3155
3156 #[tokio::test]
3157 async fn add_tls_both_static_and_dynamic_client_cert_fails() {
3158 let resolver = Arc::new(MockClientCertResolver);
3159 let tls_opts = TlsOptions {
3160 client_tls_options: Some(ClientTlsOptions {
3161 client_cert: b"some-cert".to_vec(),
3162 client_private_key: b"some-key".to_vec(),
3163 }),
3164 client_cert_resolver: Some(resolver),
3165 domain: Some("test.temporal.io".to_string()),
3166 ..Default::default()
3167 };
3168 let endpoint =
3169 tonic::transport::Channel::from_static("https://test.temporal.io:7233");
3170 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3171 assert!(
3172 matches!(result, Err(ClientConnectError::InvalidConfig(msg)) if msg.contains("client_tls_options") && msg.contains("client_cert_resolver")),
3173 "Should fail with InvalidConfig when both static and dynamic client certs are set"
3174 );
3175 }
3176
3177 #[tokio::test]
3178 async fn add_tls_no_options_returns_standard_passthrough() {
3179 let endpoint = tonic::transport::Channel::from_static("http://localhost:7233");
3180 let result = add_tls_to_channel(None, endpoint).await;
3181 assert!(
3182 matches!(&result, Ok(TlsConfigResult::Standard(_))),
3183 "Should return Standard when no TLS options are set"
3184 );
3185 }
3186
3187 #[test]
3188 fn build_custom_rustls_config_with_resolver() {
3189 let resolver = Arc::new(MockClientCertResolver);
3190 let tls_opts = TlsOptions {
3191 domain: Some("test.temporal.io".to_string()),
3192 ..Default::default()
3193 };
3194 let config = build_custom_rustls_config(&tls_opts, Some(resolver));
3195 assert!(config.is_ok(), "Should build config: {:?}", config.err());
3196 let config = config.unwrap();
3197 assert_eq!(config.alpn_protocols, vec![b"h2".to_vec()]);
3198 }
3199
3200 #[test]
3201 fn build_custom_rustls_config_without_resolver() {
3202 let tls_opts = TlsOptions {
3203 domain: Some("test.temporal.io".to_string()),
3204 ..Default::default()
3205 };
3206 let config = build_custom_rustls_config(&tls_opts, None);
3207 assert!(config.is_ok(), "Should build config: {:?}", config.err());
3208 }
3209
3210 #[test]
3211 fn build_custom_rustls_config_with_custom_verifier_and_resolver() {
3212 let resolver = Arc::new(MockClientCertResolver);
3213 let tls_opts = TlsOptions {
3214 server_cert_verifier: Some(Arc::new(MockVerifier)),
3215 domain: Some("test.temporal.io".to_string()),
3216 ..Default::default()
3217 };
3218 let config = build_custom_rustls_config(&tls_opts, Some(resolver));
3219 assert!(
3220 config.is_ok(),
3221 "Should build config with custom verifier + resolver: {:?}",
3222 config.err()
3223 );
3224 }
3225
3226 #[test]
3227 fn tls_options_debug_shows_custom_for_resolver() {
3228 let resolver = Arc::new(MockClientCertResolver);
3229 let tls_opts = TlsOptions {
3230 client_cert_resolver: Some(resolver),
3231 ..Default::default()
3232 };
3233 let debug_str = format!("{:?}", tls_opts);
3234 assert!(
3235 debug_str.contains("\"<custom>\""),
3236 "Debug should show <custom> for client_cert_resolver: {debug_str}"
3237 );
3238 assert!(
3239 debug_str.contains("client_cert_resolver"),
3240 "Debug should contain field name: {debug_str}"
3241 );
3242 }
3243
3244 #[test]
3245 fn tls_options_default_has_no_resolver() {
3246 let tls_opts = TlsOptions::default();
3247 assert!(tls_opts.client_cert_resolver.is_none());
3248 assert!(tls_opts.client_tls_options.is_none());
3249 assert!(tls_opts.server_cert_verifier.is_none());
3250 }
3251
3252 #[tokio::test]
3253 async fn add_tls_resolver_with_ip_host_uses_ip_as_domain() {
3254 let resolver = Arc::new(MockClientCertResolver);
3257 let tls_opts = TlsOptions {
3258 client_cert_resolver: Some(resolver),
3259 ..Default::default()
3261 };
3262 let endpoint = tonic::transport::Channel::from_static("https://192.168.1.100:7233");
3263 let result = add_tls_to_channel(Some(&tls_opts), endpoint).await;
3264 match result {
3265 Ok(TlsConfigResult::CustomConnector { domain, .. }) => {
3266 assert_eq!(domain, "192.168.1.100");
3267 }
3268 other => panic!(
3269 "Expected CustomConnector with IP domain, got {:?}",
3270 other.err()
3271 ),
3272 }
3273 }
3274 }
3275 }
3276
3277 mod start_workflow_interceptor_tests {
3278 use super::*;
3279 use crate::{request_extensions::RetryConfigForCall, test_helpers::XorCodec};
3280 use parking_lot::Mutex;
3281 use std::sync::atomic::{AtomicUsize, Ordering};
3282 use temporalio_common::{
3283 MemoValues, SignalDefinition, WorkerDeploymentVersion,
3284 data_converters::{
3285 DefaultFailureConverter, PayloadCodec, PayloadConversionError, PayloadConverter,
3286 SerializationContext, SerializationContextData, TemporalDeserializable,
3287 TemporalSerializable,
3288 },
3289 protos::temporal::api::{
3290 common::v1::{Link, Memo as ProtoMemo, Payload, Priority as ProtoPriority},
3291 workflow::v1::VersioningOverride as ProtoVersioningOverride,
3292 },
3293 };
3294 use temporalio_macros::{workflow, workflow_methods};
3295 use temporalio_workflow::{SyncWorkflowContext, WorkflowContext, WorkflowResult};
3296 use tonic::{Request, Response};
3297
3298 #[workflow]
3299 #[derive(Default)]
3300 struct TestWorkflow;
3301
3302 #[workflow_methods]
3303 impl TestWorkflow {
3304 #[run]
3305 async fn run(
3306 _ctx: &mut WorkflowContext<Self>,
3307 _input: Vec<String>,
3308 ) -> WorkflowResult<()> {
3309 Ok(())
3310 }
3311
3312 #[signal]
3313 fn test_signal(&mut self, _ctx: &mut SyncWorkflowContext<Self>, _input: Vec<String>) {}
3314 }
3315
3316 #[derive(Default)]
3317 struct RecordedStart {
3318 calls: usize,
3319 workflow_type: String,
3320 memo: Option<ProtoMemo>,
3321 payloads: Vec<Payload>,
3322 signal_name: String,
3323 signal_payloads: Vec<Payload>,
3324 identity: String,
3325 links: Vec<Link>,
3326 priority: Option<ProtoPriority>,
3327 versioning_override: Option<ProtoVersioningOverride>,
3328 ascii_metadata: Option<String>,
3329 binary_metadata: Option<Vec<u8>>,
3330 grpc_timeout: Option<String>,
3331 retry_options: Option<RetryOptions>,
3332 }
3333
3334 struct CountingCodec {
3335 encode_calls: Arc<AtomicUsize>,
3336 }
3337
3338 impl PayloadCodec for CountingCodec {
3339 fn encode(
3340 &self,
3341 _context: &SerializationContextData,
3342 payloads: Vec<Payload>,
3343 ) -> futures_util::future::BoxFuture<
3344 'static,
3345 Result<Vec<Payload>, PayloadConversionError>,
3346 > {
3347 self.encode_calls.fetch_add(1, Ordering::SeqCst);
3348 Box::pin(async move { Ok(payloads) })
3349 }
3350
3351 fn decode(
3352 &self,
3353 _context: &SerializationContextData,
3354 payloads: Vec<Payload>,
3355 ) -> futures_util::future::BoxFuture<
3356 'static,
3357 Result<Vec<Payload>, PayloadConversionError>,
3358 > {
3359 Box::pin(async move { Ok(payloads) })
3360 }
3361 }
3362
3363 #[derive(Clone)]
3364 struct MockStartWorkflowClient {
3365 recorded: Arc<Mutex<RecordedStart>>,
3366 data_converter: DataConverter,
3367 }
3368
3369 impl NamespacedClient for MockStartWorkflowClient {
3370 fn namespace(&self) -> String {
3371 "test-namespace".to_owned()
3372 }
3373
3374 fn identity(&self) -> String {
3375 "test-identity".to_owned()
3376 }
3377
3378 fn data_converter(&self) -> &DataConverter {
3379 &self.data_converter
3380 }
3381 }
3382
3383 impl WorkflowService for MockStartWorkflowClient {
3384 fn start_workflow_execution(
3385 &mut self,
3386 request: Request<StartWorkflowExecutionRequest>,
3387 ) -> futures_util::future::BoxFuture<
3388 '_,
3389 Result<Response<StartWorkflowExecutionResponse>, tonic::Status>,
3390 > {
3391 let ascii_metadata = request
3392 .metadata()
3393 .get("call-meta")
3394 .map(|value| value.to_str().unwrap().to_owned());
3395 let binary_metadata = request
3396 .metadata()
3397 .get_bin("call-meta-bin")
3398 .map(|value| value.to_bytes().unwrap().to_vec());
3399 let grpc_timeout = request
3400 .metadata()
3401 .get("grpc-timeout")
3402 .map(|value| value.to_str().unwrap().to_owned());
3403 let retry_options = request
3404 .extensions()
3405 .get::<RetryConfigForCall>()
3406 .map(|config| config.0.clone());
3407 let request = request.into_inner();
3408 let mut recorded = self.recorded.lock();
3409 recorded.calls += 1;
3410 recorded.workflow_type = request.workflow_type.unwrap().name;
3411 recorded.memo = request.memo;
3412 recorded.payloads = request.input.unwrap_or_default().payloads;
3413 recorded.identity = request.identity;
3414 recorded.links = request.links;
3415 recorded.priority = request.priority;
3416 recorded.versioning_override = request.versioning_override;
3417 recorded.ascii_metadata = ascii_metadata;
3418 recorded.binary_metadata = binary_metadata;
3419 recorded.grpc_timeout = grpc_timeout;
3420 recorded.retry_options = retry_options;
3421
3422 Box::pin(async {
3423 Ok(Response::new(StartWorkflowExecutionResponse {
3424 run_id: "server-run-id".to_owned(),
3425 ..Default::default()
3426 }))
3427 })
3428 }
3429
3430 fn signal_with_start_workflow_execution(
3431 &mut self,
3432 request: Request<SignalWithStartWorkflowExecutionRequest>,
3433 ) -> futures_util::future::BoxFuture<
3434 '_,
3435 Result<Response<SignalWithStartWorkflowExecutionResponse>, tonic::Status>,
3436 > {
3437 let ascii_metadata = request
3438 .metadata()
3439 .get("call-meta")
3440 .map(|value| value.to_str().unwrap().to_owned());
3441 let binary_metadata = request
3442 .metadata()
3443 .get_bin("call-meta-bin")
3444 .map(|value| value.to_bytes().unwrap().to_vec());
3445 let grpc_timeout = request
3446 .metadata()
3447 .get("grpc-timeout")
3448 .map(|value| value.to_str().unwrap().to_owned());
3449 let retry_options = request
3450 .extensions()
3451 .get::<RetryConfigForCall>()
3452 .map(|config| config.0.clone());
3453 let request = request.into_inner();
3454 let mut recorded = self.recorded.lock();
3455 recorded.calls += 1;
3456 recorded.workflow_type = request.workflow_type.unwrap().name;
3457 recorded.memo = request.memo;
3458 recorded.payloads = request.input.unwrap_or_default().payloads;
3459 recorded.signal_name = request.signal_name;
3460 recorded.signal_payloads = request.signal_input.unwrap_or_default().payloads;
3461 recorded.identity = request.identity;
3462 recorded.links = request.links;
3463 recorded.priority = request.priority;
3464 recorded.versioning_override = request.versioning_override;
3465 recorded.ascii_metadata = ascii_metadata;
3466 recorded.binary_metadata = binary_metadata;
3467 recorded.grpc_timeout = grpc_timeout;
3468 recorded.retry_options = retry_options;
3469
3470 Box::pin(async {
3471 Ok(Response::new(SignalWithStartWorkflowExecutionResponse {
3472 run_id: "signal-server-run-id".to_owned(),
3473 ..Default::default()
3474 }))
3475 })
3476 }
3477 }
3478
3479 #[derive(Clone)]
3480 struct InterceptedClient {
3481 inner: MockStartWorkflowClient,
3482 interceptors: Vec<Arc<dyn ClientInterceptor>>,
3483 }
3484
3485 impl NamespacedClient for InterceptedClient {
3486 fn namespace(&self) -> String {
3487 self.inner.namespace()
3488 }
3489
3490 fn identity(&self) -> String {
3491 self.inner.identity()
3492 }
3493
3494 fn data_converter(&self) -> &DataConverter {
3495 self.inner.data_converter()
3496 }
3497
3498 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
3499 &self.interceptors
3500 }
3501 }
3502
3503 impl WorkflowService for InterceptedClient {
3504 fn start_workflow_execution(
3505 &mut self,
3506 request: Request<StartWorkflowExecutionRequest>,
3507 ) -> futures_util::future::BoxFuture<
3508 '_,
3509 Result<Response<StartWorkflowExecutionResponse>, tonic::Status>,
3510 > {
3511 self.inner.start_workflow_execution(request)
3512 }
3513
3514 fn signal_with_start_workflow_execution(
3515 &mut self,
3516 request: Request<SignalWithStartWorkflowExecutionRequest>,
3517 ) -> futures_util::future::BoxFuture<
3518 '_,
3519 Result<Response<SignalWithStartWorkflowExecutionResponse>, tonic::Status>,
3520 > {
3521 self.inner.signal_with_start_workflow_execution(request)
3522 }
3523 }
3524
3525 struct OrderedInterceptor {
3526 name: &'static str,
3527 events: Arc<Mutex<Vec<String>>>,
3528 encode_calls: Arc<AtomicUsize>,
3529 }
3530
3531 impl ClientInterceptor for OrderedInterceptor {
3532 fn start_workflow<'a>(
3533 &'a self,
3534 mut input: StartWorkflowInput,
3535 next: Next<
3536 'a,
3537 StartWorkflowInput,
3538 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3539 >,
3540 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3541 Box::pin(async move {
3542 assert_eq!(self.encode_calls.load(Ordering::SeqCst), 0);
3543 self.events.lock().push(format!("{}-pre", self.name));
3544 tokio::task::yield_now().await;
3545 if self.name == "outer" {
3546 input
3547 .args_mut::<Vec<String>>()
3548 .unwrap()
3549 .push("mutated".to_owned());
3550 } else {
3551 assert_eq!(
3552 input.args_ref::<Vec<String>>().unwrap(),
3553 &["initial".to_owned(), "mutated".to_owned()]
3554 );
3555 input.replace_args("replacement".to_owned());
3556 input.workflow_type = "replacement-workflow".to_owned();
3557 }
3558 let result = next.run(input).await;
3559 tokio::task::yield_now().await;
3560 self.events.lock().push(format!("{}-post", self.name));
3561 result
3562 })
3563 }
3564 }
3565
3566 struct ShortCircuitInterceptor;
3567
3568 impl ClientInterceptor for ShortCircuitInterceptor {
3569 fn start_workflow<'a>(
3570 &'a self,
3571 input: StartWorkflowInput,
3572 _next: Next<
3573 'a,
3574 StartWorkflowInput,
3575 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3576 >,
3577 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3578 assert_eq!(
3579 input.args_ref::<Vec<String>>().unwrap(),
3580 &["initial".to_owned()]
3581 );
3582 Box::pin(async {
3583 Ok(StartWorkflowOutput::new(
3584 "short-circuit-workflow-id",
3585 "short-circuit-run-id",
3586 ))
3587 })
3588 }
3589 }
3590
3591 struct CountingInput {
3592 conversion_calls: Arc<AtomicUsize>,
3593 }
3594
3595 impl TemporalSerializable for CountingInput {
3596 fn to_payloads(
3597 &self,
3598 _context: &SerializationContext<'_>,
3599 ) -> Result<Vec<Payload>, PayloadConversionError> {
3600 self.conversion_calls.fetch_add(1, Ordering::SeqCst);
3601 Ok(vec![Payload::default()])
3602 }
3603 }
3604
3605 struct ConversionTimingInterceptor {
3606 conversion_calls: Arc<AtomicUsize>,
3607 }
3608
3609 impl ClientInterceptor for ConversionTimingInterceptor {
3610 fn start_workflow<'a>(
3611 &'a self,
3612 mut input: StartWorkflowInput,
3613 next: Next<
3614 'a,
3615 StartWorkflowInput,
3616 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3617 >,
3618 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3619 input.replace_args(CountingInput {
3620 conversion_calls: self.conversion_calls.clone(),
3621 });
3622 let future = next.run(input);
3623 assert_eq!(self.conversion_calls.load(Ordering::SeqCst), 0);
3624 future
3625 }
3626 }
3627
3628 struct ReplacingSignalWithStartInterceptor;
3629
3630 impl ClientInterceptor for ReplacingSignalWithStartInterceptor {
3631 fn signal_with_start_workflow<'a>(
3632 &'a self,
3633 mut input: SignalWithStartWorkflowInput,
3634 next: Next<
3635 'a,
3636 SignalWithStartWorkflowInput,
3637 BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>>,
3638 >,
3639 ) -> BoxFuture<'a, Result<StartWorkflowOutput, WorkflowStartError>> {
3640 assert_eq!(
3641 input.workflow_args_ref::<Vec<String>>().unwrap(),
3642 &["workflow".to_owned()]
3643 );
3644 assert_eq!(
3645 input.signal_args_ref::<Vec<String>>().unwrap(),
3646 &["signal".to_owned()]
3647 );
3648 input.replace_workflow_args(vec!["replaced-workflow".to_owned()]);
3649 input.replace_signal_args(vec!["replaced-signal".to_owned()]);
3650 next.run(input)
3651 }
3652 }
3653
3654 struct FailingSignal;
3655
3656 impl SignalDefinition for FailingSignal {
3657 type Workflow = test_workflow::Run;
3658 type Input = FailingSignalInput;
3659
3660 fn name(&self) -> &str {
3661 "failing-signal"
3662 }
3663 }
3664
3665 struct FailingSignalInput;
3666
3667 impl TemporalDeserializable for FailingSignalInput {}
3668
3669 impl TemporalSerializable for FailingSignalInput {
3670 fn to_payloads(
3671 &self,
3672 _context: &SerializationContext<'_>,
3673 ) -> Result<Vec<Payload>, PayloadConversionError> {
3674 Err(PayloadConversionError::WrongEncoding)
3675 }
3676 }
3677
3678 fn mock_client(
3679 interceptors: Vec<Arc<dyn ClientInterceptor>>,
3680 encode_calls: Arc<AtomicUsize>,
3681 ) -> (InterceptedClient, Arc<Mutex<RecordedStart>>) {
3682 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
3683 let data_converter = DataConverter::new(
3684 PayloadConverter::default(),
3685 DefaultFailureConverter::default(),
3686 CountingCodec {
3687 encode_calls: encode_calls.clone(),
3688 },
3689 );
3690 (
3691 InterceptedClient {
3692 inner: MockStartWorkflowClient {
3693 recorded: recorded.clone(),
3694 data_converter,
3695 },
3696 interceptors,
3697 },
3698 recorded,
3699 )
3700 }
3701
3702 fn mock_client_with_codec(
3705 codec: impl PayloadCodec + Send + Sync + 'static,
3706 ) -> (MockStartWorkflowClient, Arc<Mutex<RecordedStart>>) {
3707 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
3708 let data_converter = DataConverter::new(
3709 PayloadConverter::default(),
3710 DefaultFailureConverter::default(),
3711 codec,
3712 );
3713 (
3714 MockStartWorkflowClient {
3715 recorded: recorded.clone(),
3716 data_converter,
3717 },
3718 recorded,
3719 )
3720 }
3721
3722 async fn read_back(sent: ProtoMemo) -> Memo {
3724 let mut sent = sent;
3725 decode_payloads(
3726 &mut sent,
3727 &XorCodec,
3728 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
3729 )
3730 .await
3731 .unwrap();
3732 Memo::from_raw(
3733 Some(sent),
3734 PayloadConverter::default(),
3735 SerializationContextData::Workflow(WorkflowSerializationContext::new()),
3736 )
3737 }
3738
3739 #[rstest::rstest]
3740 #[case::unset(None)]
3741 #[case::pinned(Some(VersioningOverride::Pinned(
3742 WorkerDeploymentVersion::builder().deployment_name("deployment").build_id("build").build()
3743 )))]
3744 #[case::auto_upgrade(Some(VersioningOverride::AutoUpgrade))]
3745 #[case::one_time(Some(VersioningOverride::OneTime(
3746 WorkerDeploymentVersion::builder().deployment_name("deployment").build_id("build").build()
3747 )))]
3748 #[tokio::test]
3749 async fn start_workflow_sends_versioning_override(
3750 #[case] versioning_override: Option<VersioningOverride>,
3751 #[values(false, true)] signal_with_start: bool,
3752 ) {
3753 let (client, recorded) = mock_client_with_codec(XorCodec);
3754 let options = WorkflowStartOptions::new("task-queue", "workflow-id")
3755 .maybe_versioning_override(versioning_override.clone())
3756 .build();
3757 if signal_with_start {
3758 client
3759 .signal_with_start_workflow(
3760 TestWorkflow::run,
3761 vec!["initial".to_owned()],
3762 TestWorkflow::test_signal,
3763 vec!["signal".to_owned()],
3764 options,
3765 )
3766 .await
3767 .unwrap();
3768 } else {
3769 client
3770 .start_workflow(TestWorkflow::run, vec!["initial".to_owned()], options)
3771 .await
3772 .unwrap();
3773 }
3774 let recorded = recorded.lock();
3775 assert_eq!(recorded.calls, 1);
3776 assert_eq!(
3777 recorded.versioning_override,
3778 versioning_override.map(Into::into)
3779 );
3780 }
3781
3782 #[tokio::test]
3783 async fn start_workflow_encodes_memo_with_payload_converter_and_codec() {
3784 let (client, recorded) = mock_client_with_codec(XorCodec);
3785 let mut memo = MemoValues::new();
3786 memo.insert("memo-key", "memo-value".to_owned());
3787
3788 client
3789 .start_workflow(
3790 TestWorkflow::run,
3791 vec!["initial".to_owned()],
3792 WorkflowStartOptions::new("task-queue", "workflow-id")
3793 .memo(memo)
3794 .build(),
3795 )
3796 .await
3797 .unwrap();
3798
3799 let sent = recorded.lock().memo.clone().expect("memo should be sent");
3800 assert_eq!(
3801 read_back(sent).await.get::<String>("memo-key").unwrap(),
3802 Some("memo-value".to_owned())
3803 );
3804 }
3805
3806 #[tokio::test]
3807 async fn signal_with_start_workflow_encodes_memo() {
3808 let (client, recorded) = mock_client_with_codec(XorCodec);
3809 let mut memo = MemoValues::new();
3810 memo.insert("memo-key", "memo-value".to_owned());
3811
3812 client
3813 .signal_with_start_workflow(
3814 TestWorkflow::run,
3815 vec!["initial".to_owned()],
3816 TestWorkflow::test_signal,
3817 vec!["signal".to_owned()],
3818 WorkflowStartOptions::new("task-queue", "workflow-id")
3819 .memo(memo)
3820 .build(),
3821 )
3822 .await
3823 .unwrap();
3824
3825 let sent = recorded.lock().memo.clone().expect("memo should be sent");
3826 assert_eq!(
3827 read_back(sent).await.get::<String>("memo-key").unwrap(),
3828 Some("memo-value".to_owned())
3829 );
3830 }
3831
3832 #[tokio::test]
3833 async fn start_workflow_without_memo_sends_none() {
3834 let (client, recorded) = mock_client_with_codec(XorCodec);
3835
3836 client
3837 .start_workflow(
3838 TestWorkflow::run,
3839 vec!["initial".to_owned()],
3840 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
3841 )
3842 .await
3843 .unwrap();
3844
3845 assert_eq!(recorded.lock().memo, None);
3846 }
3847
3848 #[tokio::test]
3849 async fn start_workflow_reports_memo_serialization_errors() {
3850 #[derive(Debug)]
3851 struct FailingMemoValue;
3852
3853 impl TemporalSerializable for FailingMemoValue {
3854 fn to_payload(
3855 &self,
3856 _ctx: &SerializationContext<'_>,
3857 ) -> Result<Payload, PayloadConversionError> {
3858 Err(PayloadConversionError::EncodingError(
3859 std::io::Error::other("memo serialization failure").into(),
3860 ))
3861 }
3862 }
3863
3864 let (client, recorded) = mock_client_with_codec(XorCodec);
3865 let mut memo = MemoValues::new();
3866 memo.insert("invalid", FailingMemoValue);
3867
3868 let err = client
3869 .start_workflow(
3870 TestWorkflow::run,
3871 vec!["initial".to_owned()],
3872 WorkflowStartOptions::new("task-queue", "workflow-id")
3873 .memo(memo)
3874 .build(),
3875 )
3876 .await
3877 .map(|_| ())
3878 .expect_err("memo serialization errors should be surfaced");
3879
3880 assert!(
3881 matches!(err, WorkflowStartError::PayloadConversion(_)),
3882 "expected a payload conversion error, got {err:?}"
3883 );
3884 assert!(
3885 err.to_string().contains("memo serialization failure"),
3886 "error should surface the underlying cause, got {err}"
3887 );
3888 assert_eq!(recorded.lock().calls, 0);
3890 }
3891
3892 #[tokio::test]
3893 async fn interceptors_order_mutate_replace_and_defer_conversion() {
3894 let events = Arc::new(Mutex::new(Vec::new()));
3895 let encode_calls = Arc::new(AtomicUsize::new(0));
3896 let interceptors: Vec<Arc<dyn ClientInterceptor>> = vec![
3897 Arc::new(OrderedInterceptor {
3898 name: "outer",
3899 events: events.clone(),
3900 encode_calls: encode_calls.clone(),
3901 }),
3902 Arc::new(OrderedInterceptor {
3903 name: "inner",
3904 events: events.clone(),
3905 encode_calls: encode_calls.clone(),
3906 }),
3907 ];
3908 let (client, recorded) = mock_client(interceptors, encode_calls.clone());
3909
3910 let handle = client
3911 .start_workflow(
3912 TestWorkflow::run,
3913 vec!["initial".to_owned()],
3914 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
3915 )
3916 .await
3917 .unwrap();
3918
3919 assert_eq!(
3920 events.lock().as_slice(),
3921 ["outer-pre", "inner-pre", "inner-post", "outer-post"]
3922 );
3923 assert_eq!(encode_calls.load(Ordering::SeqCst), 1);
3924 assert_eq!(handle.run_id(), Some("server-run-id"));
3925 let payloads = {
3926 let recorded = recorded.lock();
3927 assert_eq!(recorded.calls, 1);
3928 assert_eq!(recorded.workflow_type, "replacement-workflow");
3929 recorded.payloads.clone()
3930 };
3931 let replacement: String = client
3932 .data_converter()
3933 .from_payloads(
3934 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
3935 payloads,
3936 )
3937 .await
3938 .unwrap();
3939 assert_eq!(replacement, "replacement");
3940 }
3941
3942 #[tokio::test]
3943 async fn interceptor_can_short_circuit() {
3944 let encode_calls = Arc::new(AtomicUsize::new(0));
3945 let (client, recorded) = mock_client(
3946 vec![Arc::new(ShortCircuitInterceptor)],
3947 encode_calls.clone(),
3948 );
3949 let handle = client
3950 .start_workflow(
3951 TestWorkflow::run,
3952 vec!["initial".to_owned()],
3953 WorkflowStartOptions::new("task-queue", "ignored-workflow-id").build(),
3954 )
3955 .await
3956 .unwrap();
3957
3958 assert_eq!(handle.info().workflow_id, "short-circuit-workflow-id");
3959 assert_eq!(handle.run_id(), Some("short-circuit-run-id"));
3960 assert_eq!(recorded.lock().calls, 0);
3961 assert_eq!(encode_calls.load(Ordering::SeqCst), 0);
3962 }
3963
3964 #[tokio::test]
3965 async fn payload_conversion_waits_for_next_future_poll() {
3966 let conversion_calls = Arc::new(AtomicUsize::new(0));
3967 let encode_calls = Arc::new(AtomicUsize::new(0));
3968 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
3969 let data_converter = DataConverter::new(
3970 PayloadConverter::UseWrappers,
3971 DefaultFailureConverter::default(),
3972 CountingCodec {
3973 encode_calls: encode_calls.clone(),
3974 },
3975 );
3976 let client = InterceptedClient {
3977 inner: MockStartWorkflowClient {
3978 recorded: recorded.clone(),
3979 data_converter,
3980 },
3981 interceptors: vec![Arc::new(ConversionTimingInterceptor {
3982 conversion_calls: conversion_calls.clone(),
3983 })],
3984 };
3985
3986 client
3987 .start_workflow(
3988 TestWorkflow::run,
3989 vec!["initial".to_owned()],
3990 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
3991 )
3992 .await
3993 .unwrap();
3994
3995 assert_eq!(conversion_calls.load(Ordering::SeqCst), 1);
3996 assert_eq!(encode_calls.load(Ordering::SeqCst), 1);
3997 assert_eq!(recorded.lock().calls, 1);
3998 }
3999
4000 #[tokio::test]
4001 async fn custom_client_defaults_to_empty_chain() {
4002 let recorded = Arc::new(Mutex::new(RecordedStart::default()));
4003 let client = MockStartWorkflowClient {
4004 recorded: recorded.clone(),
4005 data_converter: DataConverter::default(),
4006 };
4007 assert!(client.client_interceptors().is_empty());
4008
4009 client
4010 .start_workflow(
4011 TestWorkflow::run,
4012 vec!["initial".to_owned()],
4013 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
4014 )
4015 .await
4016 .unwrap();
4017 assert_eq!(recorded.lock().calls, 1);
4018 }
4019
4020 #[tokio::test]
4021 async fn rpc_options_reach_the_request() {
4022 let (client, recorded) = mock_client(Vec::new(), Arc::new(AtomicUsize::new(0)));
4023 let mut metadata = RpcMetadata::new();
4024 metadata.insert("call-meta", "call-value").unwrap();
4025 metadata
4026 .insert_binary("call-meta-bin", vec![0, 255])
4027 .unwrap();
4028 let rpc_options = RpcOptions::builder()
4029 .metadata(metadata)
4030 .timeout(Duration::from_millis(250))
4031 .retry_options(RetryOptions::no_retries())
4032 .build();
4033 let mut options = WorkflowStartOptions::new("task-queue", "workflow-id").build();
4034 options.rpc_options = rpc_options.clone();
4035
4036 client
4037 .start_workflow(TestWorkflow::run, vec!["initial".to_owned()], options)
4038 .await
4039 .unwrap();
4040
4041 {
4042 let recorded = recorded.lock();
4043 assert_eq!(recorded.ascii_metadata.as_deref(), Some("call-value"));
4044 assert_eq!(recorded.binary_metadata.as_deref(), Some(&[0, 255][..]));
4045 assert_eq!(recorded.grpc_timeout.as_deref(), Some("250000u"));
4046 assert_eq!(recorded.retry_options, Some(RetryOptions::no_retries()));
4047 }
4048
4049 let mut options = WorkflowStartOptions::new("task-queue", "signal-workflow-id").build();
4050 options.rpc_options = rpc_options;
4051 let handle = client
4052 .signal_with_start_workflow(
4053 TestWorkflow::run,
4054 vec!["initial".to_owned()],
4055 TestWorkflow::test_signal,
4056 vec!["signal".to_owned()],
4057 options,
4058 )
4059 .await
4060 .unwrap();
4061
4062 let recorded = recorded.lock();
4063 assert_eq!(recorded.calls, 2);
4064 assert_eq!(recorded.ascii_metadata.as_deref(), Some("call-value"));
4065 assert_eq!(recorded.binary_metadata.as_deref(), Some(&[0, 255][..]));
4066 assert_eq!(recorded.grpc_timeout.as_deref(), Some("250000u"));
4067 assert_eq!(recorded.retry_options, Some(RetryOptions::no_retries()));
4068 assert_eq!(recorded.signal_name, "test_signal");
4069 assert_eq!(recorded.signal_payloads.len(), 1);
4070 assert_eq!(handle.run_id(), Some("signal-server-run-id"));
4071 }
4072
4073 #[tokio::test]
4074 async fn signal_with_start_interceptor_can_replace_both_argument_sets() {
4075 let (client, recorded) = mock_client(
4076 vec![Arc::new(ReplacingSignalWithStartInterceptor)],
4077 Arc::new(AtomicUsize::new(0)),
4078 );
4079
4080 client
4081 .signal_with_start_workflow(
4082 TestWorkflow::run,
4083 vec!["workflow".to_owned()],
4084 TestWorkflow::test_signal,
4085 vec!["signal".to_owned()],
4086 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
4087 )
4088 .await
4089 .unwrap();
4090
4091 let data_converter = DataConverter::default();
4092 let (workflow_payloads, signal_payloads) = {
4093 let recorded = recorded.lock();
4094 (recorded.payloads.clone(), recorded.signal_payloads.clone())
4095 };
4096 assert_eq!(
4097 data_converter
4098 .from_payloads::<Vec<String>>(
4099 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4100 workflow_payloads,
4101 )
4102 .await
4103 .unwrap(),
4104 vec!["replaced-workflow".to_owned()]
4105 );
4106 assert_eq!(
4107 data_converter
4108 .from_payloads::<Vec<String>>(
4109 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4110 signal_payloads,
4111 )
4112 .await
4113 .unwrap(),
4114 vec!["replaced-signal".to_owned()]
4115 );
4116 }
4117
4118 #[tokio::test]
4119 async fn signal_with_start_payload_conversion_failure_does_not_call_service() {
4120 let (client, recorded) = mock_client(Vec::new(), Arc::new(AtomicUsize::new(0)));
4121
4122 let result = client
4123 .signal_with_start_workflow(
4124 TestWorkflow::run,
4125 vec!["workflow".to_owned()],
4126 FailingSignal,
4127 FailingSignalInput,
4128 WorkflowStartOptions::new("task-queue", "workflow-id").build(),
4129 )
4130 .await;
4131
4132 assert!(matches!(
4133 result,
4134 Err(WorkflowStartError::PayloadConversion(_))
4135 ));
4136 assert_eq!(recorded.lock().calls, 0);
4137 }
4138
4139 #[test]
4140 fn rpc_metadata_combines_with_and_overrides_connection_defaults() {
4141 let headers = Arc::new(RwLock::new(ClientHeaders {
4142 user_headers: HashMap::from([
4143 (
4144 "shared-meta".parse().unwrap(),
4145 "connection-value".parse().unwrap(),
4146 ),
4147 (
4148 "connection-meta".parse().unwrap(),
4149 "connection-only".parse().unwrap(),
4150 ),
4151 ]),
4152 user_binary_headers: HashMap::from([
4153 (
4154 "shared-meta-bin".parse().unwrap(),
4155 BinaryMetadataValue::from_bytes(&[1]),
4156 ),
4157 (
4158 "connection-meta-bin".parse().unwrap(),
4159 BinaryMetadataValue::from_bytes(&[2]),
4160 ),
4161 ]),
4162 api_key: None,
4163 }));
4164 let mut service_interceptor = ServiceCallInterceptor {
4165 client_name: "test-client".to_owned(),
4166 client_version: "test-version".to_owned(),
4167 headers,
4168 };
4169 let mut rpc_options = RpcOptions::default();
4170 rpc_options
4171 .metadata
4172 .insert("shared-meta", "call-value")
4173 .unwrap();
4174 rpc_options
4175 .metadata
4176 .insert("call-meta", "call-only")
4177 .unwrap();
4178 rpc_options
4179 .metadata
4180 .insert_binary("shared-meta-bin", vec![3])
4181 .unwrap();
4182 rpc_options
4183 .metadata
4184 .insert_binary("call-meta-bin", vec![4])
4185 .unwrap();
4186 let mut request = Request::new(());
4187 rpc_options.apply_to(&mut request);
4188
4189 let request = service_interceptor.call(request).unwrap();
4190 assert_eq!(request.metadata().get("shared-meta").unwrap(), "call-value");
4191 assert_eq!(request.metadata().get("call-meta").unwrap(), "call-only");
4192 assert_eq!(
4193 request.metadata().get("connection-meta").unwrap(),
4194 "connection-only"
4195 );
4196 assert_eq!(
4197 request.metadata().get_bin("shared-meta-bin").unwrap(),
4198 &[3][..]
4199 );
4200 assert_eq!(
4201 request.metadata().get_bin("call-meta-bin").unwrap(),
4202 &[4][..]
4203 );
4204 assert_eq!(
4205 request.metadata().get_bin("connection-meta-bin").unwrap(),
4206 &[2][..]
4207 );
4208 }
4209 }
4210
4211 mod update_with_start_tests {
4212 use super::*;
4213 use assert_matches::assert_matches;
4214 use parking_lot::Mutex;
4215 use std::collections::VecDeque;
4216 use temporalio_common::{
4217 UpdateDefinition, WorkerDeploymentVersion, WorkflowDefinition,
4218 data_converters::{GenericPayloadConverter, PayloadConverter},
4219 protos::temporal::api::{
4220 common::v1::{
4221 Header, Payload, Payloads, WorkflowExecution as ProtoWorkflowExecution,
4222 },
4223 enums::v1::{
4224 UpdateWorkflowExecutionLifecycleStage,
4225 WorkflowIdConflictPolicy as ProtoWorkflowIdConflictPolicy,
4226 },
4227 update::v1::{
4228 Input as UpdateInput, Meta as UpdateMeta, Outcome, Request as UpdateRequest,
4229 UpdateRef, WaitPolicy, outcome,
4230 },
4231 },
4232 };
4233 use tonic::{Request, Response};
4234
4235 struct TestWorkflow;
4236
4237 impl WorkflowDefinition for TestWorkflow {
4238 type Input = String;
4239 type Output = ();
4240
4241 fn name(&self) -> &str {
4242 "test-workflow"
4243 }
4244 }
4245
4246 impl HasWorkflowDefinition for TestWorkflow {
4247 type Run = Self;
4248 }
4249
4250 struct TestUpdate;
4251
4252 impl UpdateDefinition for TestUpdate {
4253 type Workflow = TestWorkflow;
4254 type Input = String;
4255 type Output = String;
4256
4257 fn name(&self) -> &str {
4258 "test-update"
4259 }
4260 }
4261
4262 fn successful_multi_operation_response(
4263 stage: UpdateWorkflowExecutionLifecycleStage,
4264 ) -> ExecuteMultiOperationResponse {
4265 let outcome = (stage == UpdateWorkflowExecutionLifecycleStage::Completed).then(|| {
4266 let payload_converter = PayloadConverter::default();
4267 let result_payloads =
4268 payload_converter
4269 .to_payloads(
4270 &SerializationContext::new(
4271 &SerializationContextData::Workflow(
4272 WorkflowSerializationContext::new(),
4273 ),
4274 &payload_converter,
4275 ),
4276 &"update-result".to_owned(),
4277 )
4278 .unwrap();
4279 Outcome {
4280 value: Some(outcome::Value::Success(Payloads {
4281 payloads: result_payloads,
4282 })),
4283 }
4284 });
4285 ExecuteMultiOperationResponse {
4286 responses: vec![
4287 execute_multi_operation_response::Response {
4288 response: Some(MultiOperationResponse::StartWorkflow(
4289 StartWorkflowExecutionResponse {
4290 run_id: "started-run-id".to_owned(),
4291 first_execution_run_id: "first-run-id".to_owned(),
4292 started: true,
4293 ..Default::default()
4294 },
4295 )),
4296 },
4297 execute_multi_operation_response::Response {
4298 response: Some(MultiOperationResponse::UpdateWorkflow(
4299 UpdateWorkflowExecutionResponse {
4300 update_ref: Some(UpdateRef {
4301 workflow_execution: Some(ProtoWorkflowExecution {
4302 workflow_id: "workflow-id".to_owned(),
4303 run_id: "update-run-id".to_owned(),
4304 }),
4305 update_id: "server-update-id".to_owned(),
4306 }),
4307 outcome,
4308 stage: stage as i32,
4309 ..Default::default()
4310 },
4311 )),
4312 },
4313 ],
4314 }
4315 }
4316
4317 #[derive(Clone)]
4318 struct MockMultiOperationClient {
4319 recorded: Arc<Mutex<Option<ExecuteMultiOperationRequest>>>,
4320 responses: Arc<Mutex<VecDeque<ExecuteMultiOperationResponse>>>,
4321 call_count: Arc<Mutex<usize>>,
4322 interceptors: Vec<Arc<dyn ClientInterceptor>>,
4323 }
4324
4325 impl MockMultiOperationClient {
4326 fn new(
4327 interceptors: Vec<Arc<dyn ClientInterceptor>>,
4328 responses: impl IntoIterator<Item = ExecuteMultiOperationResponse>,
4329 ) -> Self {
4330 Self {
4331 recorded: Arc::new(Mutex::new(None)),
4332 responses: Arc::new(Mutex::new(responses.into_iter().collect())),
4333 call_count: Arc::new(Mutex::new(0)),
4334 interceptors,
4335 }
4336 }
4337 }
4338
4339 impl NamespacedClient for MockMultiOperationClient {
4340 fn namespace(&self) -> String {
4341 "test-namespace".to_owned()
4342 }
4343
4344 fn identity(&self) -> String {
4345 "test-identity".to_owned()
4346 }
4347
4348 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
4349 &self.interceptors
4350 }
4351 }
4352
4353 impl WorkflowService for MockMultiOperationClient {
4354 fn execute_multi_operation(
4355 &mut self,
4356 request: Request<ExecuteMultiOperationRequest>,
4357 ) -> futures_util::future::BoxFuture<
4358 '_,
4359 Result<Response<ExecuteMultiOperationResponse>, tonic::Status>,
4360 > {
4361 *self.recorded.lock() = Some(request.into_inner());
4362 *self.call_count.lock() += 1;
4363 let response = self.responses.lock().pop_front().unwrap_or_else(|| {
4364 successful_multi_operation_response(
4365 UpdateWorkflowExecutionLifecycleStage::Completed,
4366 )
4367 });
4368 Box::pin(async { Ok(Response::new(response)) })
4369 }
4370 }
4371
4372 fn update_with_start_options(
4373 conflict_policy: WorkflowIdConflictPolicy,
4374 ) -> WorkflowUpdateWithStartOptions {
4375 WorkflowUpdateWithStartOptions::new("task-queue", "workflow-id", conflict_policy)
4376 .build()
4377 }
4378
4379 #[rstest::rstest]
4380 #[case::unset(None)]
4381 #[case::pinned(Some(VersioningOverride::Pinned(
4382 WorkerDeploymentVersion::builder().deployment_name("deployment").build_id("build").build()
4383 )))]
4384 #[case::auto_upgrade(Some(VersioningOverride::AutoUpgrade))]
4385 #[case::one_time(Some(VersioningOverride::OneTime(
4386 WorkerDeploymentVersion::builder().deployment_name("deployment").build_id("build").build()
4387 )))]
4388 #[tokio::test]
4389 async fn update_with_start_builds_multi_operation_request(
4390 #[case] versioning_override: Option<VersioningOverride>,
4391 ) {
4392 let client = MockMultiOperationClient::new(Vec::new(), []);
4393 let recorded = client.recorded.clone();
4394
4395 let start_header = Header {
4396 fields: HashMap::from([("start-header".to_owned(), Payload::default())]),
4397 };
4398 let update_header = Header {
4399 fields: HashMap::from([("update-header".to_owned(), Payload::default())]),
4400 };
4401 let update_handle = client
4402 .start_update_with_start_workflow(
4403 TestWorkflow,
4404 "workflow-input".to_owned(),
4405 TestUpdate,
4406 "update-input".to_owned(),
4407 WorkflowUpdateWithStartOptions::new(
4408 "task-queue",
4409 "workflow-id",
4410 WorkflowIdConflictPolicy::UseExisting,
4411 )
4412 .update_id("my-update-id".to_owned())
4413 .start_header(start_header.clone())
4414 .update_header(update_header.clone())
4415 .maybe_versioning_override(versioning_override.clone())
4416 .build(),
4417 )
4418 .await
4419 .unwrap();
4420
4421 let payload_converter = PayloadConverter::default();
4422 let context_data =
4423 SerializationContextData::Workflow(WorkflowSerializationContext::new());
4424 let context = SerializationContext::new(&context_data, &payload_converter);
4425 let workflow_payloads = payload_converter
4426 .to_payloads(&context, &"workflow-input".to_owned())
4427 .unwrap();
4428 let update_payloads = payload_converter
4429 .to_payloads(&context, &"update-input".to_owned())
4430 .unwrap();
4431
4432 let request = recorded.lock().take().unwrap();
4433 let request_id = assert_matches!(
4434 &request.operations[0].operation,
4435 Some(execute_multi_operation_request::operation::Operation::StartWorkflow(r)) => r
4436 )
4437 .request_id
4438 .clone();
4439 assert_eq!(
4440 request,
4441 ExecuteMultiOperationRequest {
4442 namespace: "test-namespace".to_owned(),
4443 operations: vec![
4444 execute_multi_operation_request::Operation {
4445 operation: Some(MultiOperationRequest::StartWorkflow(
4446 StartWorkflowExecutionRequest {
4447 namespace: "test-namespace".to_owned(),
4448 workflow_id: "workflow-id".to_owned(),
4449 workflow_type: Some(WorkflowType {
4450 name: "test-workflow".to_owned(),
4451 }),
4452 task_queue: Some(TaskQueue {
4453 name: "task-queue".to_owned(),
4454 ..Default::default()
4455 }),
4456 input: Some(Payloads {
4457 payloads: workflow_payloads,
4458 }),
4459 request_id,
4460 identity: "test-identity".to_owned(),
4461 workflow_id_conflict_policy:
4462 ProtoWorkflowIdConflictPolicy::UseExisting as i32,
4463 header: Some(start_header),
4464 priority: Some(Default::default()),
4465 versioning_override: versioning_override.map(Into::into),
4466 ..Default::default()
4467 },
4468 )),
4469 },
4470 execute_multi_operation_request::Operation {
4471 operation: Some(MultiOperationRequest::UpdateWorkflow(
4472 UpdateWorkflowExecutionRequest {
4473 namespace: "test-namespace".to_owned(),
4474 workflow_execution: Some(ProtoWorkflowExecution {
4475 workflow_id: "workflow-id".to_owned(),
4476 run_id: String::new(),
4477 }),
4478 wait_policy: Some(WaitPolicy {
4479 lifecycle_stage:
4480 UpdateWorkflowExecutionLifecycleStage::Accepted as i32,
4481 }),
4482 request: Some(UpdateRequest {
4483 meta: Some(UpdateMeta {
4484 update_id: "my-update-id".to_owned(),
4485 identity: "test-identity".to_owned(),
4486 }),
4487 input: Some(UpdateInput {
4488 header: Some(update_header),
4489 name: "test-update".to_owned(),
4490 args: Some(Payloads {
4491 payloads: update_payloads,
4492 }),
4493 }),
4494 ..Default::default()
4495 }),
4496 ..Default::default()
4497 },
4498 )),
4499 },
4500 ],
4501 resource_id: "workflow-id".to_owned(),
4502 }
4503 );
4504
4505 assert_eq!(update_handle.id(), "my-update-id");
4506 assert_eq!(update_handle.workflow_run_id(), Some("update-run-id"));
4507 let result: String = update_handle
4510 .get_result(RpcOptions::default())
4511 .await
4512 .unwrap();
4513 assert_eq!(result, "update-result");
4514 }
4515
4516 #[tokio::test]
4517 async fn update_with_start_retries_until_update_is_accepted() {
4518 let client = MockMultiOperationClient::new(
4519 Vec::new(),
4520 [
4521 successful_multi_operation_response(
4522 UpdateWorkflowExecutionLifecycleStage::Unspecified,
4523 ),
4524 successful_multi_operation_response(
4525 UpdateWorkflowExecutionLifecycleStage::Accepted,
4526 ),
4527 ],
4528 );
4529 let call_count = client.call_count.clone();
4530
4531 let update_handle = client
4532 .start_update_with_start_workflow(
4533 TestWorkflow,
4534 "workflow-input".to_owned(),
4535 TestUpdate,
4536 "update-input".to_owned(),
4537 update_with_start_options(WorkflowIdConflictPolicy::Fail),
4538 )
4539 .await
4540 .unwrap();
4541
4542 assert_eq!(*call_count.lock(), 2);
4543 assert_eq!(update_handle.workflow_run_id(), Some("update-run-id"));
4544 }
4545
4546 #[tokio::test]
4547 async fn update_with_start_rejects_malformed_operation_responses() {
4548 let mut missing_response = successful_multi_operation_response(
4549 UpdateWorkflowExecutionLifecycleStage::Accepted,
4550 );
4551 missing_response.responses[0] = execute_multi_operation_response::Response::default();
4552 let mut extra_response = successful_multi_operation_response(
4553 UpdateWorkflowExecutionLifecycleStage::Accepted,
4554 );
4555 extra_response
4556 .responses
4557 .push(execute_multi_operation_response::Response::default());
4558 let mut wrong_order = successful_multi_operation_response(
4559 UpdateWorkflowExecutionLifecycleStage::Accepted,
4560 );
4561 wrong_order.responses.swap(0, 1);
4562
4563 for response in [missing_response, extra_response, wrong_order] {
4564 let client = MockMultiOperationClient::new(Vec::new(), [response]);
4565 let result = client
4566 .start_update_with_start_workflow(
4567 TestWorkflow,
4568 "workflow-input".to_owned(),
4569 TestUpdate,
4570 "update-input".to_owned(),
4571 update_with_start_options(WorkflowIdConflictPolicy::Fail),
4572 )
4573 .await;
4574 assert!(matches!(
4575 result,
4576 Err(WorkflowUpdateWithStartError::Other(_))
4577 ));
4578 }
4579 }
4580
4581 #[tokio::test]
4582 async fn update_with_start_interceptor_can_mutate_args() {
4583 struct ReplaceArgsInterceptor;
4584
4585 impl ClientInterceptor for ReplaceArgsInterceptor {
4586 fn update_with_start_workflow<'a>(
4587 &'a self,
4588 mut input: UpdateWithStartWorkflowInput,
4589 next: Next<
4590 'a,
4591 UpdateWithStartWorkflowInput,
4592 BoxFuture<
4593 'a,
4594 Result<UpdateWithStartWorkflowOutput, WorkflowUpdateWithStartError>,
4595 >,
4596 >,
4597 ) -> BoxFuture<
4598 'a,
4599 Result<UpdateWithStartWorkflowOutput, WorkflowUpdateWithStartError>,
4600 > {
4601 assert_eq!(
4602 input.workflow_args_ref::<String>().unwrap(),
4603 "workflow-input"
4604 );
4605 input.replace_workflow_args("replaced-workflow-input".to_owned());
4606 *input.update_args_mut::<String>().unwrap() =
4607 "replaced-update-input".to_owned();
4608 next.run(input)
4609 }
4610 }
4611
4612 let client = MockMultiOperationClient::new(vec![Arc::new(ReplaceArgsInterceptor)], []);
4613 let recorded = client.recorded.clone();
4614
4615 client
4616 .start_update_with_start_workflow(
4617 TestWorkflow,
4618 "workflow-input".to_owned(),
4619 TestUpdate,
4620 "update-input".to_owned(),
4621 update_with_start_options(WorkflowIdConflictPolicy::Fail),
4622 )
4623 .await
4624 .unwrap();
4625
4626 let request = recorded.lock().take().unwrap();
4627 let start_request = assert_matches!(
4628 &request.operations[0].operation,
4629 Some(execute_multi_operation_request::operation::Operation::StartWorkflow(r)) => r
4630 );
4631 let workflow_input: String = client
4632 .data_converter()
4633 .from_payloads(
4634 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4635 start_request.input.clone().unwrap().payloads,
4636 )
4637 .await
4638 .unwrap();
4639 assert_eq!(workflow_input, "replaced-workflow-input");
4640 let update_request = assert_matches!(
4641 &request.operations[1].operation,
4642 Some(execute_multi_operation_request::operation::Operation::UpdateWorkflow(r)) => r
4643 );
4644 let update_input: String = client
4645 .data_converter()
4646 .from_payloads(
4647 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4648 update_request
4649 .request
4650 .clone()
4651 .unwrap()
4652 .input
4653 .unwrap()
4654 .args
4655 .unwrap()
4656 .payloads,
4657 )
4658 .await
4659 .unwrap();
4660 assert_eq!(update_input, "replaced-update-input");
4661 }
4662 }
4663
4664 mod list_workflows_tests {
4665 use super::*;
4666 use crate::test_helpers::{FailingCodec, XorCodec};
4667 use futures_util::{FutureExt, StreamExt};
4668 use std::sync::atomic::{AtomicUsize, Ordering};
4669 use temporalio_common::{
4670 data_converters::{DefaultFailureConverter, PayloadConverter},
4671 protos::temporal::api::common::v1::{
4672 Memo as ProtoMemo, Payload, WorkflowExecution as ProtoWorkflowExecution,
4673 },
4674 };
4675 use tonic::{Request, Response};
4676
4677 #[derive(Clone)]
4678 struct MockListWorkflowsClient {
4679 call_count: Arc<AtomicUsize>,
4680 page_size: usize,
4682 total_workflows: usize,
4684 data_converter: DataConverter,
4685 memo_payload: Option<Payload>,
4686 interceptors: Vec<Arc<dyn ClientInterceptor>>,
4687 }
4688
4689 impl NamespacedClient for MockListWorkflowsClient {
4690 fn namespace(&self) -> String {
4691 "test-namespace".to_string()
4692 }
4693 fn identity(&self) -> String {
4694 "test-identity".to_string()
4695 }
4696 fn data_converter(&self) -> &DataConverter {
4697 &self.data_converter
4698 }
4699 fn client_interceptors(&self) -> &[Arc<dyn ClientInterceptor>] {
4700 &self.interceptors
4701 }
4702 }
4703
4704 struct CountingListInterceptor {
4705 calls: Arc<AtomicUsize>,
4706 }
4707
4708 impl ClientInterceptor for CountingListInterceptor {
4709 fn list_workflows_page<'a>(
4710 &'a self,
4711 input: ListWorkflowsPageInput,
4712 next: Next<
4713 'a,
4714 ListWorkflowsPageInput,
4715 BoxFuture<'a, Result<ListWorkflowsPageOutput, ClientError>>,
4716 >,
4717 ) -> BoxFuture<'a, Result<ListWorkflowsPageOutput, ClientError>> {
4718 self.calls.fetch_add(1, Ordering::SeqCst);
4719 next.run(input)
4720 }
4721 }
4722
4723 impl WorkflowService for MockListWorkflowsClient {
4724 fn list_workflow_executions(
4725 &mut self,
4726 request: Request<ListWorkflowExecutionsRequest>,
4727 ) -> futures_util::future::BoxFuture<
4728 '_,
4729 Result<Response<ListWorkflowExecutionsResponse>, tonic::Status>,
4730 > {
4731 self.call_count.fetch_add(1, Ordering::SeqCst);
4732 let req = request.into_inner();
4733
4734 let offset: usize = if req.next_page_token.is_empty() {
4736 0
4737 } else {
4738 String::from_utf8(req.next_page_token)
4739 .unwrap()
4740 .parse()
4741 .unwrap()
4742 };
4743
4744 let remaining = self.total_workflows.saturating_sub(offset);
4745 let count = remaining.min(self.page_size);
4746 let new_offset = offset + count;
4747
4748 let executions: Vec<_> = (offset..offset + count)
4749 .map(|i| workflow::WorkflowExecutionInfo {
4750 execution: Some(ProtoWorkflowExecution {
4751 workflow_id: format!("wf-{i}"),
4752 run_id: format!("run-{i}"),
4753 }),
4754 r#type: Some(WorkflowType {
4755 name: "TestWorkflow".to_string(),
4756 }),
4757 task_queue: "test-queue".to_string(),
4758 memo: self.memo_payload.clone().map(|payload| ProtoMemo {
4759 fields: HashMap::from([("memo-key".to_owned(), payload)]),
4760 }),
4761 ..Default::default()
4762 })
4763 .collect();
4764
4765 let next_page_token = if new_offset < self.total_workflows {
4766 new_offset.to_string().into_bytes()
4767 } else {
4768 vec![]
4769 };
4770
4771 async move {
4772 Ok(Response::new(ListWorkflowExecutionsResponse {
4773 executions,
4774 next_page_token,
4775 }))
4776 }
4777 .boxed()
4778 }
4779 }
4780
4781 #[tokio::test]
4782 async fn list_workflows_paginates_through_all_results() {
4783 let call_count = Arc::new(AtomicUsize::new(0));
4784 let interceptor_calls = Arc::new(AtomicUsize::new(0));
4785 let client = MockListWorkflowsClient {
4786 call_count: call_count.clone(),
4787 page_size: 3,
4788 total_workflows: 10,
4789 data_converter: DataConverter::default(),
4790 memo_payload: None,
4791 interceptors: vec![Arc::new(CountingListInterceptor {
4792 calls: interceptor_calls.clone(),
4793 })],
4794 };
4795
4796 let stream = client.list_workflows("", WorkflowListOptions::default());
4797 let results: Vec<_> = stream.collect().await;
4798
4799 assert_eq!(results.len(), 10);
4800 for (i, result) in results.iter().enumerate() {
4801 let wf = result.as_ref().unwrap();
4802 assert_eq!(wf.id(), format!("wf-{i}"));
4803 assert_eq!(wf.run_id(), format!("run-{i}"));
4804 }
4805 assert_eq!(call_count.load(Ordering::SeqCst), 4);
4807 assert_eq!(interceptor_calls.load(Ordering::SeqCst), 4);
4808 }
4809
4810 #[tokio::test]
4811 async fn list_workflows_respects_limit() {
4812 let call_count = Arc::new(AtomicUsize::new(0));
4813 let client = MockListWorkflowsClient {
4814 call_count: call_count.clone(),
4815 page_size: 3,
4816 total_workflows: 10,
4817 data_converter: DataConverter::default(),
4818 memo_payload: None,
4819 interceptors: Vec::new(),
4820 };
4821
4822 let opts = WorkflowListOptions::builder().limit(5).build();
4823 let stream = client.list_workflows("", opts);
4824 let results: Vec<_> = stream.collect().await;
4825
4826 assert_eq!(results.len(), 5);
4827 for (i, result) in results.iter().enumerate() {
4828 let wf = result.as_ref().unwrap();
4829 assert_eq!(wf.id(), format!("wf-{i}"));
4830 }
4831 assert_eq!(call_count.load(Ordering::SeqCst), 2);
4833 }
4834
4835 #[tokio::test]
4836 async fn list_workflows_limit_less_than_page_size() {
4837 let call_count = Arc::new(AtomicUsize::new(0));
4838 let client = MockListWorkflowsClient {
4839 call_count: call_count.clone(),
4840 page_size: 10,
4841 total_workflows: 100,
4842 data_converter: DataConverter::default(),
4843 memo_payload: None,
4844 interceptors: Vec::new(),
4845 };
4846
4847 let opts = WorkflowListOptions::builder().limit(3).build();
4848 let stream = client.list_workflows("", opts);
4849 let results: Vec<_> = stream.collect().await;
4850
4851 assert_eq!(results.len(), 3);
4852 assert_eq!(call_count.load(Ordering::SeqCst), 1);
4854 }
4855
4856 #[tokio::test]
4857 async fn list_workflows_empty_results() {
4858 let call_count = Arc::new(AtomicUsize::new(0));
4859 let client = MockListWorkflowsClient {
4860 call_count: call_count.clone(),
4861 page_size: 10,
4862 total_workflows: 0,
4863 data_converter: DataConverter::default(),
4864 memo_payload: None,
4865 interceptors: Vec::new(),
4866 };
4867
4868 let stream = client.list_workflows("", WorkflowListOptions::default());
4869 let results: Vec<_> = stream.collect().await;
4870
4871 assert_eq!(results.len(), 0);
4872 assert_eq!(call_count.load(Ordering::SeqCst), 1);
4873 }
4874
4875 #[tokio::test]
4876 async fn list_workflows_exposes_typed_memo() {
4877 let data_converter = DataConverter::new(
4878 PayloadConverter::default(),
4879 DefaultFailureConverter::default(),
4880 XorCodec,
4881 );
4882 let memo_payload = data_converter
4883 .to_payload(
4884 &SerializationContextData::Workflow(WorkflowSerializationContext::new()),
4885 &"memo-value".to_owned(),
4886 )
4887 .await
4888 .unwrap();
4889 let client = MockListWorkflowsClient {
4890 call_count: Arc::new(AtomicUsize::new(0)),
4891 page_size: 1,
4892 total_workflows: 1,
4893 data_converter,
4894 memo_payload: Some(memo_payload),
4895 interceptors: Vec::new(),
4896 };
4897
4898 let workflow = client
4899 .list_workflows("", WorkflowListOptions::default())
4900 .next()
4901 .await
4902 .unwrap()
4903 .unwrap();
4904
4905 assert_eq!(
4906 workflow.memo().get::<String>("memo-key").unwrap(),
4907 Some("memo-value".to_owned())
4908 );
4909 }
4910
4911 #[tokio::test]
4912 async fn list_workflows_yields_codec_error_then_ends() {
4913 let client = MockListWorkflowsClient {
4914 call_count: Arc::new(AtomicUsize::new(0)),
4915 page_size: 1,
4916 total_workflows: 1,
4917 data_converter: DataConverter::new(
4918 PayloadConverter::default(),
4919 DefaultFailureConverter::default(),
4920 FailingCodec,
4921 ),
4922 memo_payload: Some(Payload::default()),
4923 interceptors: Vec::new(),
4924 };
4925 let mut stream = client.list_workflows("", WorkflowListOptions::default());
4926
4927 let err = stream.next().await.unwrap().unwrap_err();
4928
4929 assert!(matches!(err, ClientError::PayloadConversion(_)));
4930 assert!(stream.next().await.is_none());
4931 }
4932 }
4933}