Skip to main content

temporalio_client/
lib.rs

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