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