Skip to main content

temporalio_client/
lib.rs

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