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