1use crate::config::{GrpcAuth, GrpcStreamConfig, MetadataEntry, RpcKind};
4use async_trait::async_trait;
5use base64::Engine as _;
6use faucet_core::{AuthSpec, Credential, FaucetError, SharedAuthProvider, StreamPage};
7use futures_core::Stream;
8use prost::Message;
9use prost::bytes::Bytes;
10use prost_reflect::{DescriptorPool, DynamicMessage, MessageDescriptor, SerializeOptions};
11use serde_json::Value;
12use std::pin::Pin;
13use std::time::Duration;
14use tonic::codec::{Codec, DecodeBuf, Decoder, EncodeBuf, Encoder};
15use tonic::transport::Channel;
16
17pub struct GrpcStream {
20 config: GrpcStreamConfig,
21 pool: DescriptorPool,
22 auth_provider: Option<SharedAuthProvider>,
25}
26
27impl GrpcStream {
28 pub fn new(config: GrpcStreamConfig) -> Result<Self, FaucetError> {
30 if config.reconnect_initial_backoff.is_zero() {
33 return Err(FaucetError::Config(
34 "grpc reconnect_initial_backoff must be > 0 (a zero backoff busy-spins reconnects)"
35 .into(),
36 ));
37 }
38 let descriptor_bytes = std::fs::read(&config.descriptor_set_path).map_err(|e| {
39 FaucetError::Config(format!(
40 "failed to read descriptor set at {}: {e}",
41 config.descriptor_set_path.display()
42 ))
43 })?;
44
45 let pool = DescriptorPool::decode(Bytes::from(descriptor_bytes))
46 .map_err(|e| FaucetError::Config(format!("failed to parse FileDescriptorSet: {e}")))?;
47
48 Ok(Self {
49 config,
50 pool,
51 auth_provider: None,
52 })
53 }
54
55 pub fn with_auth_provider(mut self, provider: SharedAuthProvider) -> Self {
62 self.auth_provider = Some(provider);
63 self
64 }
65
66 pub async fn fetch_all(&self) -> Result<Vec<Value>, FaucetError> {
68 self.fetch_resolved(
69 &self.config.endpoint,
70 &self.config.service_name,
71 &self.config.method_name,
72 &self.config.request,
73 )
74 .await
75 }
76
77 async fn fetch_resolved(
79 &self,
80 endpoint: &str,
81 service_name: &str,
82 method_name: &str,
83 request: &Value,
84 ) -> Result<Vec<Value>, FaucetError> {
85 let (output_desc, request_bytes) = self.prepare_call(service_name, method_name, request)?;
86 let path = parse_method_path(service_name, method_name)?;
87
88 match self.config.rpc_kind {
89 RpcKind::Unary => {
90 let channel = self.connect_channel(endpoint).await?;
91 let mut grpc_client = self.configure_client(tonic::client::Grpc::new(channel));
92 grpc_client
93 .ready()
94 .await
95 .map_err(|e| FaucetError::Config(format!("gRPC channel not ready: {e}")))?;
96
97 let codec = DynamicCodec::new(output_desc);
98 let request = self.build_grpc_request(request_bytes).await?;
99
100 let response: tonic::Response<DynamicMessage> = grpc_client
101 .unary(request, path, codec)
102 .await
103 .map_err(|e| FaucetError::Source(format!("gRPC unary call failed: {e}")))?;
104
105 let resp_msg = response.into_inner();
106 let records =
107 serialize_and_extract(&resp_msg, self.config.records_path.as_deref())?;
108 tracing::info!(records = records.len(), "gRPC unary fetch complete");
109 Ok(records)
110 }
111 RpcKind::ServerStreaming => {
112 self.fetch_server_streaming_collect(endpoint, path, output_desc, request_bytes)
113 .await
114 }
115 }
116 }
117
118 async fn fetch_server_streaming_collect(
124 &self,
125 endpoint: &str,
126 path: tonic::codegen::http::uri::PathAndQuery,
127 output_desc: MessageDescriptor,
128 request_bytes: Vec<u8>,
129 ) -> Result<Vec<Value>, FaucetError> {
130 let max_messages = self.config.max_messages.unwrap_or(usize::MAX);
131 let mut all: Vec<Value> = Vec::new();
132 let mut messages_seen: usize = 0;
133 let mut attempt: u32 = 0;
134 let mut backoff = self.config.reconnect_initial_backoff;
135 let max_backoff = self.config.reconnect_max_backoff;
136
137 loop {
138 match self
139 .drive_server_streaming_once(
140 endpoint,
141 path.clone(),
142 output_desc.clone(),
143 request_bytes.clone(),
144 max_messages,
145 messages_seen,
146 |records| {
147 all.extend(records);
148 Ok(())
149 },
150 )
151 .await
152 {
153 Ok(consumed) => {
154 tracing::info!(
155 records = all.len(),
156 messages = messages_seen + consumed,
157 "gRPC server-streaming fetch complete"
158 );
159 return Ok(all);
160 }
161 Err(StreamOutcome::Done(consumed)) => {
162 messages_seen += consumed;
163 tracing::info!(
164 records = all.len(),
165 messages = messages_seen,
166 "gRPC server-streaming fetch complete (max_messages reached)"
167 );
168 return Ok(all);
169 }
170 Err(StreamOutcome::Transient { consumed, error }) => {
171 messages_seen += consumed;
172 if self.config.terminate_on_error {
173 return Err(error);
174 }
175 if consumed > 0 {
183 attempt = 0;
184 backoff = self.config.reconnect_initial_backoff;
185 }
186 if let Some(max_attempts) = self.config.reconnect_max_attempts
187 && attempt >= max_attempts
188 {
189 return Err(FaucetError::Source(format!(
190 "gRPC server-streaming exceeded reconnect_max_attempts={max_attempts}: {error}"
191 )));
192 }
193 attempt += 1;
194 tracing::warn!(
195 attempt,
196 backoff_ms = backoff.as_millis() as u64,
197 error = %error,
198 "gRPC server-streaming transient error, reconnecting"
199 );
200 tokio::time::sleep(backoff).await;
201 backoff = next_backoff(backoff, max_backoff);
202 }
203 }
204 }
205 }
206
207 #[allow(clippy::too_many_arguments)]
215 async fn drive_server_streaming_once<F>(
216 &self,
217 endpoint: &str,
218 path: tonic::codegen::http::uri::PathAndQuery,
219 output_desc: MessageDescriptor,
220 request_bytes: Vec<u8>,
221 max_messages: usize,
222 already_seen: usize,
223 mut on_records: F,
224 ) -> Result<usize, StreamOutcome>
225 where
226 F: FnMut(Vec<Value>) -> Result<(), FaucetError>,
227 {
228 let channel = match self.connect_channel(endpoint).await {
229 Ok(c) => c,
230 Err(e) => {
231 return Err(StreamOutcome::Transient {
232 consumed: 0,
233 error: e,
234 });
235 }
236 };
237
238 let mut grpc_client = self.configure_client(tonic::client::Grpc::new(channel));
239 if let Err(e) = grpc_client.ready().await {
240 return Err(StreamOutcome::Transient {
241 consumed: 0,
242 error: FaucetError::Source(format!("gRPC channel not ready: {e}")),
243 });
244 }
245
246 let codec = DynamicCodec::new(output_desc);
247 let request = match self.build_grpc_request(request_bytes).await {
248 Ok(r) => r,
249 Err(e) => {
250 return Err(StreamOutcome::Transient {
252 consumed: 0,
253 error: e,
254 });
255 }
256 };
257
258 let response = match grpc_client.server_streaming(request, path, codec).await {
259 Ok(r) => r,
260 Err(status) => {
261 return Err(StreamOutcome::Transient {
262 consumed: 0,
263 error: FaucetError::Source(format!(
264 "gRPC server-streaming start failed: {status}"
265 )),
266 });
267 }
268 };
269
270 let mut streaming = response.into_inner();
271 let records_path = self.config.records_path.as_deref();
272 let skip = if self.config.reconnect_replay_from_start {
277 already_seen
278 } else {
279 0
280 };
281 let mut position: usize = 0; let mut emitted: usize = 0; loop {
285 if already_seen + emitted >= max_messages {
286 return Err(StreamOutcome::Done(emitted));
287 }
288 match streaming.message().await {
289 Ok(Some(msg)) => {
290 position += 1;
291 if position <= skip {
292 continue;
294 }
295 let records = match serialize_and_extract(&msg, records_path) {
296 Ok(r) => r,
297 Err(e) => {
298 return Err(StreamOutcome::Transient {
299 consumed: emitted,
300 error: e,
301 });
302 }
303 };
304 if let Err(e) = on_records(records) {
305 return Err(StreamOutcome::Transient {
306 consumed: emitted,
307 error: e,
308 });
309 }
310 emitted += 1;
311 }
312 Ok(None) => {
313 return Ok(emitted);
314 }
315 Err(status) => {
316 return Err(StreamOutcome::Transient {
317 consumed: emitted,
318 error: FaucetError::Source(format!(
319 "gRPC server-streaming recv failed: {status}"
320 )),
321 });
322 }
323 }
324 }
325 }
326
327 fn prepare_call(
329 &self,
330 service_name: &str,
331 method_name: &str,
332 request: &Value,
333 ) -> Result<(MessageDescriptor, Vec<u8>), FaucetError> {
334 let service = self.pool.get_service_by_name(service_name).ok_or_else(|| {
335 FaucetError::Config(format!(
336 "service '{service_name}' not found in descriptor set",
337 ))
338 })?;
339
340 let method = service
341 .methods()
342 .find(|m| m.name() == method_name)
343 .ok_or_else(|| {
344 FaucetError::Config(format!(
345 "method '{method_name}' not found in service '{service_name}'",
346 ))
347 })?;
348
349 let input_desc = method.input();
350 let request_msg = DynamicMessage::deserialize(input_desc, request)
351 .map_err(|e| FaucetError::Config(format!("failed to build request message: {e}")))?;
352 let request_bytes = request_msg.encode_to_vec();
353
354 Ok((method.output(), request_bytes))
355 }
356
357 async fn connect_channel(&self, endpoint: &str) -> Result<Channel, FaucetError> {
359 let use_tls = self
360 .config
361 .tls
362 .unwrap_or_else(|| endpoint.starts_with("https"));
363
364 let channel_endpoint = Channel::from_shared(endpoint.to_string())
365 .map_err(|e| FaucetError::Url(format!("invalid gRPC endpoint: {e}")))?;
366
367 let channel = if use_tls {
368 channel_endpoint
369 .tls_config(tonic::transport::ClientTlsConfig::new())
370 .map_err(|e| FaucetError::Config(format!("TLS config failed: {e}")))?
371 .connect()
372 .await
373 .map_err(|e| FaucetError::Source(format!("gRPC connect failed: {e}")))?
374 } else {
375 channel_endpoint
376 .connect()
377 .await
378 .map_err(|e| FaucetError::Source(format!("gRPC connect failed: {e}")))?
379 };
380
381 Ok(channel)
382 }
383
384 fn configure_client(
387 &self,
388 mut client: tonic::client::Grpc<Channel>,
389 ) -> tonic::client::Grpc<Channel> {
390 if let Some(n) = self.config.max_decoding_message_size {
391 client = client.max_decoding_message_size(n);
392 }
393 if let Some(n) = self.config.max_encoding_message_size {
394 client = client.max_encoding_message_size(n);
395 }
396 client
397 }
398
399 async fn build_grpc_request(
406 &self,
407 request_bytes: Vec<u8>,
408 ) -> Result<tonic::Request<Vec<u8>>, FaucetError> {
409 let effective = if let Some(provider) = &self.auth_provider {
410 credential_to_auth(provider.credential().await?)
411 } else {
412 match &self.config.auth {
413 AuthSpec::Inline(a) => a.clone(),
414 AuthSpec::Reference(r) => {
415 return Err(FaucetError::Auth(format!(
416 "auth references provider '{}' but no provider was supplied; \
417 set one via the CLI `auth:` catalog or `with_auth_provider`",
418 r.name
419 )));
420 }
421 }
422 };
423
424 let mut request = tonic::Request::new(request_bytes);
425 apply_grpc_auth(&effective, &mut request)?;
426 Ok(request)
427 }
428}
429
430#[async_trait]
431impl faucet_core::Source for GrpcStream {
432 async fn fetch_with_context(
433 &self,
434 context: &std::collections::HashMap<String, serde_json::Value>,
435 ) -> Result<Vec<Value>, FaucetError> {
436 if context.is_empty() {
437 return GrpcStream::fetch_all(self).await;
438 }
439
440 let endpoint = faucet_core::util::substitute_context(&self.config.endpoint, context);
441 let service_name =
442 faucet_core::util::substitute_context(&self.config.service_name, context);
443 let method_name = faucet_core::util::substitute_context(&self.config.method_name, context);
444
445 let request = {
446 let s = serde_json::to_string(&self.config.request)
447 .map_err(|e| FaucetError::Config(format!("failed to serialize request: {e}")))?;
448 let s = faucet_core::util::substitute_context_json(&s, context);
449 serde_json::from_str(&s).map_err(|e| {
450 FaucetError::Config(format!("failed to parse substituted request: {e}"))
451 })?
452 };
453
454 self.fetch_resolved(&endpoint, &service_name, &method_name, &request)
455 .await
456 }
457
458 fn stream_pages<'a>(
476 &'a self,
477 context: &'a std::collections::HashMap<String, Value>,
478 batch_size: usize,
479 ) -> Pin<Box<dyn Stream<Item = Result<StreamPage, FaucetError>> + Send + 'a>> {
480 match self.config.rpc_kind {
481 RpcKind::Unary => {
482 self.default_stream_pages(context, batch_size)
485 }
486 RpcKind::ServerStreaming => self.server_streaming_pages(context),
487 }
488 }
489
490 fn connector_name(&self) -> &'static str {
491 "grpc"
492 }
493
494 fn config_schema(&self) -> serde_json::Value {
495 serde_json::to_value(faucet_core::schema_for!(GrpcStreamConfig))
496 .expect("schema serialization")
497 }
498
499 fn dataset_uri(&self) -> String {
500 format!(
501 "{}/{}/{}",
502 faucet_core::redact_uri_credentials(&self.config.endpoint),
503 self.config.service_name,
504 self.config.method_name
505 )
506 }
507}
508
509impl GrpcStream {
510 fn default_stream_pages<'a>(
514 &'a self,
515 context: &'a std::collections::HashMap<String, Value>,
516 batch_size: usize,
517 ) -> Pin<Box<dyn Stream<Item = Result<StreamPage, FaucetError>> + Send + 'a>> {
518 use faucet_core::Source;
519 Box::pin(async_stream::try_stream! {
520 let (records, bookmark) = Source::fetch_with_context_incremental(self, context).await?;
521 let total = records.len();
522 let chunk = if batch_size == 0 { usize::MAX } else { batch_size };
523
524 if total == 0 {
525 if bookmark.is_some() {
526 yield StreamPage { records: Vec::new(), bookmark };
527 }
528 return;
529 }
530
531 let mut iter = records.into_iter();
532 let mut consumed = 0usize;
533 loop {
534 let batch: Vec<Value> = iter.by_ref().take(chunk).collect();
535 if batch.is_empty() {
536 break;
537 }
538 consumed += batch.len();
539 let page_bookmark = if consumed >= total { bookmark.clone() } else { None };
540 yield StreamPage { records: batch, bookmark: page_bookmark };
541 }
542 })
543 }
544
545 fn server_streaming_pages<'a>(
549 &'a self,
550 context: &'a std::collections::HashMap<String, Value>,
551 ) -> Pin<Box<dyn Stream<Item = Result<StreamPage, FaucetError>> + Send + 'a>> {
552 let batch_size = self.config.batch_size;
553 let page_chunk = if batch_size == 0 {
554 usize::MAX
555 } else {
556 batch_size
557 };
558 let initial_capacity = if batch_size == 0 { 1024 } else { batch_size };
559 let max_messages = self.config.max_messages.unwrap_or(usize::MAX);
560 let terminate_on_error = self.config.terminate_on_error;
561 let reconnect_max_attempts = self.config.reconnect_max_attempts;
562 let reconnect_initial_backoff = self.config.reconnect_initial_backoff;
563 let mut backoff = reconnect_initial_backoff;
564 let max_backoff = self.config.reconnect_max_backoff;
565
566 Box::pin(async_stream::try_stream! {
567 let endpoint = if context.is_empty() {
571 self.config.endpoint.clone()
572 } else {
573 faucet_core::util::substitute_context(&self.config.endpoint, context)
574 };
575 let service_name = if context.is_empty() {
576 self.config.service_name.clone()
577 } else {
578 faucet_core::util::substitute_context(&self.config.service_name, context)
579 };
580 let method_name = if context.is_empty() {
581 self.config.method_name.clone()
582 } else {
583 faucet_core::util::substitute_context(&self.config.method_name, context)
584 };
585 let request: Value = if context.is_empty() {
586 self.config.request.clone()
587 } else {
588 let s = serde_json::to_string(&self.config.request)
589 .map_err(|e| FaucetError::Config(format!("failed to serialize request: {e}")))?;
590 let s = faucet_core::util::substitute_context_json(&s, context);
591 serde_json::from_str(&s).map_err(|e| FaucetError::Config(format!(
592 "failed to parse substituted request: {e}"
593 )))?
594 };
595
596 let (output_desc, request_bytes) =
597 self.prepare_call(&service_name, &method_name, &request)?;
598 let path = parse_method_path(&service_name, &method_name)?;
599
600 let mut buffer: Vec<Value> = Vec::with_capacity(initial_capacity);
601 let mut messages_seen: usize = 0;
602 let mut attempt: u32 = 0;
603
604 'reconnect: loop {
605 let outcome = self.drive_server_streaming_once(
609 &endpoint,
610 path.clone(),
611 output_desc.clone(),
612 request_bytes.clone(),
613 max_messages,
614 messages_seen,
615 |records| {
616 buffer.extend(records);
617 Ok(())
618 },
619 ).await;
620
621 while buffer.len() >= page_chunk {
625 let drained: Vec<Value> = buffer.drain(..page_chunk).collect();
626 yield StreamPage { records: drained, bookmark: None };
627 }
628
629 match outcome {
630 Ok(consumed) => {
631 messages_seen += consumed;
632 if !buffer.is_empty() {
635 let final_records = std::mem::take(&mut buffer);
636 yield StreamPage { records: final_records, bookmark: None };
637 }
638 tracing::info!(
639 messages = messages_seen,
640 "gRPC server-streaming complete"
641 );
642 break 'reconnect;
643 }
644 Err(StreamOutcome::Done(consumed)) => {
645 messages_seen += consumed;
646 if !buffer.is_empty() {
647 let final_records = std::mem::take(&mut buffer);
648 yield StreamPage { records: final_records, bookmark: None };
649 }
650 tracing::info!(
651 messages = messages_seen,
652 "gRPC server-streaming complete (max_messages reached)"
653 );
654 break 'reconnect;
655 }
656 Err(StreamOutcome::Transient { consumed, error }) => {
657 messages_seen += consumed;
658 if terminate_on_error {
659 Err(error)?;
665 return;
666 }
667 if consumed > 0 {
672 attempt = 0;
673 backoff = reconnect_initial_backoff;
674 }
675 if let Some(max_attempts) = reconnect_max_attempts
676 && attempt >= max_attempts
677 {
678 let final_err = FaucetError::Source(format!(
679 "gRPC server-streaming exceeded reconnect_max_attempts={max_attempts}: {error}"
680 ));
681 Err(final_err)?;
682 return;
683 }
684 attempt += 1;
685 tracing::warn!(
686 attempt,
687 backoff_ms = backoff.as_millis() as u64,
688 error = %error,
689 "gRPC server-streaming transient error, reconnecting"
690 );
691 tokio::time::sleep(backoff).await;
692 backoff = next_backoff(backoff, max_backoff);
693 }
694 }
695 }
696 })
697 }
698}
699
700fn credential_to_auth(cred: Credential) -> GrpcAuth {
708 match cred {
709 Credential::Bearer(token) => GrpcAuth::Bearer { token },
710 Credential::Token(token) => GrpcAuth::Metadata {
711 entries: vec![MetadataEntry {
712 key: "authorization".into(),
713 value: token,
714 }],
715 },
716 Credential::Header { name, value } => GrpcAuth::Metadata {
717 entries: vec![MetadataEntry { key: name, value }],
718 },
719 Credential::Basic { username, password } => GrpcAuth::Metadata {
720 entries: vec![MetadataEntry {
721 key: "authorization".into(),
722 value: format!(
723 "Basic {}",
724 base64::engine::general_purpose::STANDARD
725 .encode(format!("{username}:{password}"))
726 ),
727 }],
728 },
729 }
730}
731
732fn apply_grpc_auth(
734 auth: &GrpcAuth,
735 request: &mut tonic::Request<Vec<u8>>,
736) -> Result<(), FaucetError> {
737 match auth {
738 GrpcAuth::None => {}
739 GrpcAuth::Bearer { token } => {
740 let val: tonic::metadata::MetadataValue<tonic::metadata::Ascii> =
741 format!("Bearer {token}")
742 .parse()
743 .map_err(|e| FaucetError::Auth(format!("invalid bearer token: {e}")))?;
744 request.metadata_mut().insert("authorization", val);
745 }
746 GrpcAuth::Metadata { entries } => {
747 for entry in entries {
748 let val: tonic::metadata::MetadataValue<tonic::metadata::Ascii> = entry
749 .value
750 .parse()
751 .map_err(|e| FaucetError::Auth(format!("invalid metadata value: {e}")))?;
752 let key: tonic::metadata::MetadataKey<tonic::metadata::Ascii> =
753 entry
754 .key
755 .parse()
756 .map_err(|e| FaucetError::Auth(format!("invalid metadata key: {e}")))?;
757 request.metadata_mut().insert(key, val);
758 }
759 }
760 }
761 Ok(())
762}
763
764fn parse_method_path(
765 service_name: &str,
766 method_name: &str,
767) -> Result<tonic::codegen::http::uri::PathAndQuery, FaucetError> {
768 let full_method = format!("/{service_name}/{method_name}");
769 tonic::codegen::http::uri::PathAndQuery::from_maybe_shared(full_method)
770 .map_err(|e| FaucetError::Url(format!("invalid method path: {e}")))
771}
772
773fn serialize_and_extract(
774 msg: &DynamicMessage,
775 records_path: Option<&str>,
776) -> Result<Vec<Value>, FaucetError> {
777 let serialize_opts = SerializeOptions::new().stringify_64_bit_integers(false);
778 let json_value = msg
779 .serialize_with_options(serde_json::value::Serializer, &serialize_opts)
780 .map_err(|e| {
781 FaucetError::Transform(format!("failed to serialize gRPC response to JSON: {e}"))
782 })?;
783 faucet_core::util::extract_records(&json_value, records_path)
784}
785
786fn next_backoff(current: Duration, cap: Duration) -> Duration {
787 current.saturating_mul(2).min(cap)
788}
789
790enum StreamOutcome {
794 Done(usize),
795 Transient { consumed: usize, error: FaucetError },
796}
797
798struct DynamicCodec {
802 output_desc: prost_reflect::MessageDescriptor,
803}
804
805impl DynamicCodec {
806 fn new(output_desc: prost_reflect::MessageDescriptor) -> Self {
807 Self { output_desc }
808 }
809}
810
811impl Codec for DynamicCodec {
812 type Encode = Vec<u8>;
813 type Decode = DynamicMessage;
814 type Encoder = RawEncoder;
815 type Decoder = DynamicDecoder;
816
817 fn encoder(&mut self) -> Self::Encoder {
818 RawEncoder
819 }
820
821 fn decoder(&mut self) -> Self::Decoder {
822 DynamicDecoder {
823 desc: self.output_desc.clone(),
824 }
825 }
826}
827
828struct RawEncoder;
829
830impl Encoder for RawEncoder {
831 type Item = Vec<u8>;
832 type Error = tonic::Status;
833
834 fn encode(&mut self, item: Self::Item, buf: &mut EncodeBuf<'_>) -> Result<(), Self::Error> {
835 use prost::bytes::BufMut;
836 buf.put_slice(&item);
837 Ok(())
838 }
839}
840
841struct DynamicDecoder {
842 desc: prost_reflect::MessageDescriptor,
843}
844
845impl Decoder for DynamicDecoder {
846 type Item = DynamicMessage;
847 type Error = tonic::Status;
848
849 fn decode(&mut self, buf: &mut DecodeBuf<'_>) -> Result<Option<Self::Item>, Self::Error> {
850 use prost::bytes::Buf;
851 if !buf.has_remaining() {
852 return Ok(None);
853 }
854 let bytes = buf.copy_to_bytes(buf.remaining());
855 let msg = DynamicMessage::decode(self.desc.clone(), bytes)
856 .map_err(|e| tonic::Status::internal(format!("protobuf decode error: {e}")))?;
857 Ok(Some(msg))
858 }
859}
860
861#[cfg(test)]
862mod tests {
863 use super::*;
864 use faucet_core::{AuthProvider, AuthReference, AuthSpec, Credential, FaucetError};
865 use std::sync::Arc;
866
867 #[test]
870 fn next_backoff_doubles_up_to_cap() {
871 let cap = Duration::from_secs(30);
872 let a = Duration::from_secs(1);
873 let b = next_backoff(a, cap);
874 let c = next_backoff(b, cap);
875 let d = next_backoff(c, cap);
876 let e = next_backoff(d, cap);
877 let f = next_backoff(e, cap);
878 let g = next_backoff(f, cap);
879 assert_eq!(b, Duration::from_secs(2));
880 assert_eq!(c, Duration::from_secs(4));
881 assert_eq!(d, Duration::from_secs(8));
882 assert_eq!(e, Duration::from_secs(16));
883 assert_eq!(f, Duration::from_secs(30));
884 assert_eq!(g, Duration::from_secs(30));
885 }
886
887 #[test]
890 fn credential_bearer_maps_to_grpc_bearer() {
891 let auth = credential_to_auth(Credential::Bearer("tok".into()));
892 assert!(matches!(auth, GrpcAuth::Bearer { token } if token == "tok"));
893 }
894
895 #[test]
896 fn credential_token_maps_to_metadata_authorization() {
897 let auth = credential_to_auth(Credential::Token("Custom xyz".into()));
898 match auth {
899 GrpcAuth::Metadata { entries } => {
900 assert_eq!(entries.len(), 1);
901 assert_eq!(entries[0].key, "authorization");
902 assert_eq!(entries[0].value, "Custom xyz");
903 }
904 other => panic!("expected Metadata, got {other:?}"),
905 }
906 }
907
908 #[test]
909 fn credential_header_maps_to_metadata_with_given_name() {
910 let auth = credential_to_auth(Credential::Header {
911 name: "x-api-key".into(),
912 value: "secret".into(),
913 });
914 match auth {
915 GrpcAuth::Metadata { entries } => {
916 assert_eq!(entries.len(), 1);
917 assert_eq!(entries[0].key, "x-api-key");
918 assert_eq!(entries[0].value, "secret");
919 }
920 other => panic!("expected Metadata, got {other:?}"),
921 }
922 }
923
924 #[test]
925 fn credential_basic_maps_to_base64_authorization_metadata() {
926 let auth = credential_to_auth(Credential::Basic {
927 username: "alice".into(),
928 password: "p@ss".into(),
929 });
930 match auth {
931 GrpcAuth::Metadata { entries } => {
932 assert_eq!(entries.len(), 1);
933 assert_eq!(entries[0].key, "authorization");
934 let expected = format!(
935 "Basic {}",
936 base64::engine::general_purpose::STANDARD.encode("alice:p@ss")
937 );
938 assert_eq!(entries[0].value, expected);
939 }
940 other => panic!("expected Metadata, got {other:?}"),
941 }
942 }
943
944 fn make_dummy_stream() -> GrpcStream {
949 use prost::Message;
950 let fds_set = prost_types::FileDescriptorSet {
951 file: vec![prost_types::FileDescriptorProto {
952 name: Some("dummy.proto".into()),
953 syntax: Some("proto3".into()),
954 ..Default::default()
955 }],
956 };
957 let bytes = fds_set.encode_to_vec();
958 let tmp = tempfile::NamedTempFile::new().expect("tempfile");
959 std::fs::write(tmp.path(), &bytes).expect("write descriptor");
960 let config = GrpcStreamConfig::new(
961 "http://localhost:50051",
962 "dummy.Svc",
963 "Call",
964 tmp.path().to_str().unwrap(),
965 );
966 GrpcStream::new(config).expect("new from in-memory descriptor")
968 }
969
970 #[test]
971 fn rejects_zero_reconnect_initial_backoff() {
972 use prost::Message;
973 let fds_set = prost_types::FileDescriptorSet {
976 file: vec![prost_types::FileDescriptorProto {
977 name: Some("dummy.proto".into()),
978 syntax: Some("proto3".into()),
979 ..Default::default()
980 }],
981 };
982 let tmp = tempfile::NamedTempFile::new().expect("tempfile");
983 std::fs::write(tmp.path(), fds_set.encode_to_vec()).expect("write descriptor");
984 let mut config = GrpcStreamConfig::new(
985 "http://localhost:50051",
986 "dummy.Svc",
987 "Call",
988 tmp.path().to_str().unwrap(),
989 );
990 config.reconnect_initial_backoff = std::time::Duration::ZERO;
991 let Err(err) = GrpcStream::new(config) else {
992 panic!("a zero reconnect_initial_backoff must be rejected (it busy-spins)");
993 };
994 assert!(matches!(err, FaucetError::Config(_)), "{err:?}");
995 assert!(
996 err.to_string().contains("reconnect_initial_backoff"),
997 "{err}"
998 );
999 }
1000
1001 #[tokio::test]
1004 async fn unresolved_auth_reference_errors_at_request_time() {
1005 let mut stream = make_dummy_stream();
1006 stream.config.auth = AuthSpec::Reference(AuthReference {
1007 name: "missing-provider".into(),
1008 });
1009 let err = stream.build_grpc_request(vec![]).await.unwrap_err();
1010 assert!(
1011 matches!(err, FaucetError::Auth(_)),
1012 "expected Auth error, got {err:?}"
1013 );
1014 let msg = err.to_string();
1015 assert!(
1016 msg.contains("missing-provider"),
1017 "error message should name the provider: {msg}"
1018 );
1019 }
1020
1021 #[derive(Debug)]
1024 struct FixedBearer(&'static str);
1025
1026 #[async_trait::async_trait]
1027 impl AuthProvider for FixedBearer {
1028 async fn credential(&self) -> Result<Credential, FaucetError> {
1029 Ok(Credential::Bearer(self.0.to_string()))
1030 }
1031 fn provider_name(&self) -> &'static str {
1032 "fixed-bearer"
1033 }
1034 }
1035
1036 #[tokio::test]
1037 async fn injected_provider_overrides_inline_none() {
1038 let provider: SharedAuthProvider = Arc::new(FixedBearer("MYTOKEN"));
1039 let stream = make_dummy_stream().with_auth_provider(provider);
1040 let req = stream
1042 .build_grpc_request(vec![])
1043 .await
1044 .expect("build request");
1045 let auth_header = req
1046 .metadata()
1047 .get("authorization")
1048 .expect("authorization metadata must be present")
1049 .to_str()
1050 .expect("ascii");
1051 assert_eq!(auth_header, "Bearer MYTOKEN");
1052 }
1053
1054 #[test]
1055 fn dataset_uri_combines_endpoint_service_method() {
1056 use faucet_core::Source;
1057 let stream = make_dummy_stream();
1058 assert_eq!(
1059 stream.dataset_uri(),
1060 "http://localhost:50051/dummy.Svc/Call"
1061 );
1062 }
1063}