1use super::HttpConnectorError;
35use crate::env_ref::parse_env_ref;
40use async_trait::async_trait;
41use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
42use serde::{Deserialize, Serialize};
43use std::collections::HashMap;
44use std::sync::Arc;
45
46fn default_true() -> bool {
48 true
49}
50
51fn default_auth_header() -> String {
53 "Authorization".to_string()
54}
55
56#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize, Default)]
64#[serde(tag = "type", rename_all = "snake_case")]
65pub enum AuthConfig {
66 #[default]
68 None,
69
70 ApiKey {
72 #[serde(default)]
74 query_params: HashMap<String, String>,
75 #[serde(default)]
77 headers: HashMap<String, String>,
78 #[serde(default = "default_true")]
80 required: bool,
81 },
82
83 Bearer {
85 token: String,
89 #[serde(default = "default_true")]
91 required: bool,
92 },
93
94 Basic {
96 username: String,
99 password: String,
103 #[serde(default = "default_true")]
105 required: bool,
106 },
107
108 #[serde(alias = "oauth2_client_credentials")]
115 OAuth2ClientCredentials {
116 token_url: String,
118 client_id: String,
121 client_secret: String,
125 #[serde(default)]
127 scopes: Vec<String>,
128 #[serde(default = "default_true")]
130 required: bool,
131 },
132
133 #[serde(alias = "oauth_passthrough")]
140 OAuthPassthrough {
141 #[serde(default = "default_auth_header")]
143 target_header: String,
144 #[serde(default = "default_true")]
146 required: bool,
147 },
148}
149
150impl AuthConfig {
151 #[must_use]
200 pub fn malformed_env_ref_field(&self) -> Option<String> {
201 fn scalar(label: &str, value: &str) -> Option<String> {
204 (parse_env_ref(value) == Some("")).then(|| label.to_string())
205 }
206 fn entry(label: &str, map: &HashMap<String, String>) -> Option<String> {
207 map.iter()
208 .filter(|(_, v)| parse_env_ref(v) == Some(""))
209 .map(|(k, _)| k)
210 .min()
211 .map(|k| format!("{label}.{k}"))
212 }
213
214 match self {
215 Self::None | Self::OAuthPassthrough { .. } => None,
217 Self::ApiKey {
218 query_params,
219 headers,
220 ..
221 } => entry("query_params", query_params).or_else(|| entry("headers", headers)),
222 Self::Bearer { token, .. } => scalar("token", token),
223 Self::Basic {
224 username, password, ..
225 } => scalar("username", username).or_else(|| scalar("password", password)),
226 Self::OAuth2ClientCredentials {
230 client_id,
231 client_secret,
232 ..
233 } => scalar("client_id", client_id).or_else(|| scalar("client_secret", client_secret)),
234 }
235 }
236
237 #[must_use]
239 pub fn is_required(&self) -> bool {
240 match self {
241 Self::None => false,
242 Self::ApiKey { required, .. }
243 | Self::Bearer { required, .. }
244 | Self::Basic { required, .. }
245 | Self::OAuth2ClientCredentials { required, .. }
246 | Self::OAuthPassthrough { required, .. } => *required,
247 }
248 }
249}
250
251#[async_trait]
259pub trait HttpAuthProvider: Send + Sync + 'static {
260 async fn apply(
271 &self,
272 headers: &mut HeaderMap,
273 query: &mut HashMap<String, String>,
274 inbound_token: Option<&str>,
275 ) -> Result<(), HttpConnectorError>;
276}
277
278pub struct NoAuth;
280
281#[async_trait]
282impl HttpAuthProvider for NoAuth {
283 async fn apply(
284 &self,
285 _headers: &mut HeaderMap,
286 _query: &mut HashMap<String, String>,
287 _inbound_token: Option<&str>,
288 ) -> Result<(), HttpConnectorError> {
289 Ok(())
290 }
291}
292
293pub struct MissingTokenAuth;
295
296#[async_trait]
297impl HttpAuthProvider for MissingTokenAuth {
298 async fn apply(
299 &self,
300 _headers: &mut HeaderMap,
301 _query: &mut HashMap<String, String>,
302 inbound_token: Option<&str>,
303 ) -> Result<(), HttpConnectorError> {
304 if inbound_token.map(str::is_empty) == Some(false) {
307 return Ok(());
308 }
309 Err(HttpConnectorError::Auth(
310 "authentication required but no inbound token was provided".to_string(),
311 ))
312 }
313}
314
315pub struct ApiKeyAuth {
317 query_params: HashMap<String, String>,
318 headers: HashMap<String, String>,
319}
320
321#[async_trait]
322impl HttpAuthProvider for ApiKeyAuth {
323 async fn apply(
324 &self,
325 headers: &mut HeaderMap,
326 query: &mut HashMap<String, String>,
327 _inbound_token: Option<&str>,
328 ) -> Result<(), HttpConnectorError> {
329 for (key, value) in &self.query_params {
330 query.insert(key.clone(), value.clone());
331 }
332 for (key, value) in &self.headers {
333 let name = HeaderName::try_from(key.as_str()).map_err(|_| {
334 HttpConnectorError::InvalidHeader("invalid header name".to_string())
335 })?;
336 let val = HeaderValue::try_from(value.as_str()).map_err(|_| {
337 HttpConnectorError::InvalidHeader("invalid header value".to_string())
338 })?;
339 headers.insert(name, val);
340 }
341 Ok(())
342 }
343}
344
345pub struct BearerAuth {
347 token: String,
348}
349
350#[async_trait]
351impl HttpAuthProvider for BearerAuth {
352 async fn apply(
353 &self,
354 headers: &mut HeaderMap,
355 _query: &mut HashMap<String, String>,
356 _inbound_token: Option<&str>,
357 ) -> Result<(), HttpConnectorError> {
358 let value = format!("Bearer {}", self.token);
359 let header_value = HeaderValue::try_from(value)
360 .map_err(|_| HttpConnectorError::InvalidHeader("invalid bearer token".to_string()))?;
361 headers.insert(reqwest::header::AUTHORIZATION, header_value);
362 Ok(())
363 }
364}
365
366pub struct BasicAuth {
368 username: String,
369 password: String,
370}
371
372#[async_trait]
373impl HttpAuthProvider for BasicAuth {
374 async fn apply(
375 &self,
376 headers: &mut HeaderMap,
377 _query: &mut HashMap<String, String>,
378 _inbound_token: Option<&str>,
379 ) -> Result<(), HttpConnectorError> {
380 use base64::Engine;
381 let credentials = format!("{}:{}", self.username, self.password);
382 let encoded = base64::engine::general_purpose::STANDARD.encode(credentials.as_bytes());
383 let value = format!("Basic {encoded}");
384 let header_value = HeaderValue::try_from(value).map_err(|_| {
385 HttpConnectorError::InvalidHeader("invalid basic credentials".to_string())
386 })?;
387 headers.insert(reqwest::header::AUTHORIZATION, header_value);
388 Ok(())
389 }
390}
391
392pub struct OAuth2ClientCredentialsAuth {
398 token_url: String,
399 client_id: String,
400 client_secret: String,
401 scopes: Vec<String>,
402 cached: tokio::sync::RwLock<Option<String>>,
403}
404
405impl OAuth2ClientCredentialsAuth {
406 #[must_use]
408 pub fn new(
409 token_url: String,
410 client_id: String,
411 client_secret: String,
412 scopes: Vec<String>,
413 ) -> Self {
414 Self {
415 token_url,
416 client_id,
417 client_secret,
418 scopes,
419 cached: tokio::sync::RwLock::new(None),
420 }
421 }
422
423 async fn fetch_token(&self) -> Result<String, HttpConnectorError> {
424 let client = reqwest::Client::new();
425 let mut params = vec![
426 ("grant_type", "client_credentials".to_string()),
427 ("client_id", self.client_id.clone()),
428 ("client_secret", self.client_secret.clone()),
429 ];
430 if !self.scopes.is_empty() {
431 params.push(("scope", self.scopes.join(" ")));
432 }
433 let response = client
434 .post(&self.token_url)
435 .form(¶ms)
436 .send()
437 .await
438 .map_err(|_| HttpConnectorError::Auth("oauth2 token request failed".to_string()))?;
439 if !response.status().is_success() {
440 return Err(HttpConnectorError::Auth(format!(
441 "oauth2 token endpoint returned status {}",
442 response.status().as_u16()
443 )));
444 }
445 #[derive(Deserialize)]
446 struct TokenResponse {
447 access_token: String,
448 }
449 let token: TokenResponse = response.json().await.map_err(|_| {
450 HttpConnectorError::Auth("oauth2 token response unparseable".to_string())
451 })?;
452 Ok(token.access_token)
453 }
454}
455
456#[async_trait]
457impl HttpAuthProvider for OAuth2ClientCredentialsAuth {
458 async fn apply(
459 &self,
460 headers: &mut HeaderMap,
461 _query: &mut HashMap<String, String>,
462 _inbound_token: Option<&str>,
463 ) -> Result<(), HttpConnectorError> {
464 {
465 let cached = self.cached.read().await;
466 if cached.is_none() {
467 drop(cached);
468 let fetched = self.fetch_token().await?;
469 *self.cached.write().await = Some(fetched);
470 }
471 }
472 let cached = self.cached.read().await;
473 if let Some(access_token) = cached.as_ref() {
474 let value = format!("Bearer {access_token}");
475 let header_value = HeaderValue::try_from(value).map_err(|_| {
476 HttpConnectorError::InvalidHeader("invalid oauth2 access token".to_string())
477 })?;
478 headers.insert(reqwest::header::AUTHORIZATION, header_value);
479 }
480 Ok(())
481 }
482}
483
484pub struct OAuthPassthroughAuth {
512 target_header: String,
513 incoming_token: Option<String>,
514 required: bool,
515}
516
517#[async_trait]
518impl HttpAuthProvider for OAuthPassthroughAuth {
519 async fn apply(
520 &self,
521 headers: &mut HeaderMap,
522 _query: &mut HashMap<String, String>,
523 inbound_token: Option<&str>,
524 ) -> Result<(), HttpConnectorError> {
525 let token: Option<&str> = inbound_token
527 .filter(|t| !t.is_empty())
528 .or_else(|| self.incoming_token.as_deref().filter(|t| !t.is_empty()));
529
530 match token {
531 Some(tok) => {
532 let header_name =
533 HeaderName::try_from(self.target_header.as_str()).map_err(|_| {
534 HttpConnectorError::InvalidHeader(
535 "invalid passthrough target header".to_string(),
536 )
537 })?;
538 let value = if tok.starts_with("Bearer ") || tok.starts_with("Basic ") {
541 tok.to_string()
542 } else {
543 format!("Bearer {tok}")
544 };
545 let header_value = HeaderValue::try_from(value).map_err(|_| {
546 HttpConnectorError::InvalidHeader("invalid passthrough token value".to_string())
547 })?;
548 headers.insert(header_name, header_value);
558 Ok(())
559 },
560 None if self.required => Err(HttpConnectorError::Auth(
561 "passthrough authentication required but no inbound token was provided".to_string(),
562 )),
563 None => Ok(()),
564 }
565 }
566}
567
568fn resolve_secret_ref(raw: &str) -> String {
607 match parse_env_ref(raw) {
608 None => raw.to_string(),
610 Some("") => String::new(),
615 Some(name) => std::env::var(name)
616 .ok()
617 .filter(|v| !v.trim().is_empty())
618 .unwrap_or_default(),
619 }
620}
621
622fn expand_api_key_map(map: &HashMap<String, String>) -> HashMap<String, String> {
625 map.iter()
626 .filter_map(|(k, v)| {
627 let resolved = resolve_secret_ref(v);
628 (!resolved.is_empty()).then(|| (k.clone(), resolved))
629 })
630 .collect()
631}
632
633pub fn create_auth_provider(
650 cfg: &AuthConfig,
651) -> Result<Arc<dyn HttpAuthProvider>, HttpConnectorError> {
652 if let Some(field) = cfg.malformed_env_ref_field() {
658 return Err(HttpConnectorError::Auth(format!(
659 "[backend.auth].{field} is a malformed environment reference; a reference must be \
660 exactly one `${{VAR}}` (name matching [A-Za-z0-9_]+) or `env:VAR` naming a single \
661 variable. The value is not echoed here because a malformed reference is often a \
662 mistyped literal secret."
663 )));
664 }
665 let provider: Arc<dyn HttpAuthProvider> = match cfg {
666 AuthConfig::None => Arc::new(NoAuth),
667 AuthConfig::ApiKey {
668 query_params,
669 headers,
670 ..
671 } => {
672 let query_params = expand_api_key_map(query_params);
676 let headers = expand_api_key_map(headers);
677 let has_values = query_params.values().any(|v| !v.is_empty())
678 || headers.values().any(|v| !v.is_empty());
679 if has_values {
680 Arc::new(ApiKeyAuth {
681 query_params,
682 headers,
683 })
684 } else {
685 Arc::new(NoAuth)
686 }
687 },
688 AuthConfig::Bearer { token, .. } => {
689 let token = resolve_secret_ref(token);
693 if token.is_empty() {
694 Arc::new(NoAuth)
695 } else {
696 Arc::new(BearerAuth { token })
697 }
698 },
699 AuthConfig::Basic {
700 username, password, ..
701 } => {
702 let username = resolve_secret_ref(username);
705 let password = resolve_secret_ref(password);
706 if username.is_empty() && password.is_empty() {
707 Arc::new(NoAuth)
708 } else {
709 Arc::new(BasicAuth { username, password })
710 }
711 },
712 AuthConfig::OAuth2ClientCredentials {
713 token_url,
714 client_id,
715 client_secret,
716 scopes,
717 ..
718 } => {
719 let client_id = resolve_secret_ref(client_id);
723 let client_secret = resolve_secret_ref(client_secret);
724 if client_id.is_empty() || client_secret.is_empty() {
725 Arc::new(NoAuth)
726 } else {
727 Arc::new(OAuth2ClientCredentialsAuth::new(
728 token_url.clone(),
729 client_id,
730 client_secret,
731 scopes.clone(),
732 ))
733 }
734 },
735 AuthConfig::OAuthPassthrough { required, .. } => {
736 if *required {
737 Arc::new(MissingTokenAuth)
738 } else {
739 Arc::new(NoAuth)
740 }
741 },
742 };
743 Ok(provider)
744}
745
746pub fn create_passthrough_auth_provider(
758 cfg: &AuthConfig,
759 incoming_token: Option<String>,
760) -> Result<Arc<dyn HttpAuthProvider>, HttpConnectorError> {
761 match cfg {
762 AuthConfig::OAuthPassthrough {
763 target_header,
764 required,
765 } => Ok(Arc::new(OAuthPassthroughAuth {
766 target_header: target_header.clone(),
767 incoming_token: incoming_token.filter(|t| !t.is_empty()),
768 required: *required,
769 })),
770 other => create_auth_provider(other),
771 }
772}
773
774#[cfg(test)]
775mod tests {
776 use super::*;
777
778 #[tokio::test]
779 async fn test_no_auth() {
780 let auth = create_auth_provider(&AuthConfig::None).unwrap();
781 let mut headers = HeaderMap::new();
782 let mut query = HashMap::new();
783 auth.apply(&mut headers, &mut query, None).await.unwrap();
784 assert!(headers.is_empty());
785 assert!(query.is_empty());
786 }
787
788 #[tokio::test]
789 async fn test_bearer_auth() {
790 let cfg = AuthConfig::Bearer {
791 token: "my_token".to_string(),
792 required: true,
793 };
794 let auth = create_auth_provider(&cfg).unwrap();
795 let mut headers = HeaderMap::new();
796 let mut query = HashMap::new();
797 auth.apply(&mut headers, &mut query, Some("client-tok"))
799 .await
800 .unwrap();
801 assert_eq!(
802 headers.get(reqwest::header::AUTHORIZATION).unwrap(),
803 "Bearer my_token"
804 );
805 assert!(query.is_empty());
806 }
807
808 #[tokio::test]
809 async fn test_basic_auth() {
810 let cfg = AuthConfig::Basic {
811 username: "user".to_string(),
812 password: "pass".to_string(),
813 required: true,
814 };
815 let auth = create_auth_provider(&cfg).unwrap();
816 let mut headers = HeaderMap::new();
817 let mut query = HashMap::new();
818 auth.apply(&mut headers, &mut query, None).await.unwrap();
819 assert_eq!(
821 headers.get(reqwest::header::AUTHORIZATION).unwrap(),
822 "Basic dXNlcjpwYXNz"
823 );
824 }
825
826 #[tokio::test]
827 async fn test_api_key_query_param() {
828 let cfg = AuthConfig::ApiKey {
830 query_params: [("app_key".to_string(), "secret123".to_string())]
831 .into_iter()
832 .collect(),
833 headers: HashMap::new(),
834 required: true,
835 };
836 let auth = create_auth_provider(&cfg).unwrap();
837 let mut headers = HeaderMap::new();
838 let mut query = HashMap::new();
839 auth.apply(&mut headers, &mut query, None).await.unwrap();
840 assert_eq!(query.get("app_key"), Some(&"secret123".to_string()));
841 assert!(
842 headers.is_empty(),
843 "api-key-in-query must not touch headers"
844 );
845 }
846
847 #[tokio::test]
848 async fn test_api_key_query_param_expands_braced_env_ref() {
849 let var = "PMCP_TEST_TFL_APP_KEY_BRACED";
851 std::env::set_var(var, "dummy");
852 let cfg = AuthConfig::ApiKey {
853 query_params: [("app_key".to_string(), format!("${{{var}}}"))]
854 .into_iter()
855 .collect(),
856 headers: HashMap::new(),
857 required: false,
858 };
859 let auth = create_auth_provider(&cfg).unwrap();
860 let mut headers = HeaderMap::new();
861 let mut query = HashMap::new();
862 auth.apply(&mut headers, &mut query, None).await.unwrap();
863 assert_eq!(
864 query.get("app_key"),
865 Some(&"dummy".to_string()),
866 "resolved env value lands on the query, not the literal ${{...}}"
867 );
868 std::env::remove_var(var);
869 }
870
871 #[tokio::test]
872 async fn test_api_key_query_param_unset_ref_is_omitted() {
873 let var = "PMCP_TEST_TFL_APP_KEY_UNSET";
876 std::env::remove_var(var);
877 let cfg = AuthConfig::ApiKey {
878 query_params: [("app_key".to_string(), format!("${{{var}}}"))]
879 .into_iter()
880 .collect(),
881 headers: HashMap::new(),
882 required: false,
883 };
884 let auth = create_auth_provider(&cfg).unwrap();
885 let mut headers = HeaderMap::new();
886 let mut query = HashMap::new();
887 auth.apply(&mut headers, &mut query, None).await.unwrap();
888 assert!(
889 !query.contains_key("app_key"),
890 "an unset required=false api_key ref is omitted, not sent empty/literal"
891 );
892 }
893
894 #[test]
895 fn test_resolve_api_key_value_forms() {
896 let var = "PMCP_TEST_RESOLVE_API_KEY_FORM";
898 std::env::set_var(var, "resolved");
899 assert_eq!(resolve_secret_ref(&format!("${{{var}}}")), "resolved");
900 assert_eq!(resolve_secret_ref(&format!("env:{var}")), "resolved");
901 assert_eq!(resolve_secret_ref("plain-literal"), "plain-literal");
902 std::env::remove_var(var);
903 assert_eq!(resolve_secret_ref(&format!("${{{var}}}")), "");
904 assert_eq!(resolve_secret_ref("${}"), "");
905 }
906
907 #[tokio::test]
908 async fn test_passthrough_forwards_inbound_token() {
909 let cfg = AuthConfig::OAuthPassthrough {
911 target_header: "Authorization".to_string(),
912 required: true,
913 };
914 let auth = create_passthrough_auth_provider(&cfg, None).unwrap();
915 let mut headers = HeaderMap::new();
916 let mut query = HashMap::new();
917 auth.apply(&mut headers, &mut query, Some("client-tok"))
918 .await
919 .unwrap();
920 assert_eq!(
921 headers.get(reqwest::header::AUTHORIZATION).unwrap(),
922 "Bearer client-tok"
923 );
924 }
925
926 #[tokio::test]
927 async fn test_passthrough_custom_target_header() {
928 let cfg = AuthConfig::OAuthPassthrough {
929 target_header: "X-Forwarded-Token".to_string(),
930 required: true,
931 };
932 let auth = create_passthrough_auth_provider(&cfg, None).unwrap();
933 let mut headers = HeaderMap::new();
934 let mut query = HashMap::new();
935 auth.apply(&mut headers, &mut query, Some("client-tok"))
936 .await
937 .unwrap();
938 assert_eq!(
939 headers.get("X-Forwarded-Token").unwrap(),
940 "Bearer client-tok"
941 );
942 }
943
944 #[tokio::test]
945 async fn test_passthrough_uses_construction_time_token() {
946 let cfg = AuthConfig::OAuthPassthrough {
948 target_header: "Authorization".to_string(),
949 required: true,
950 };
951 let auth =
952 create_passthrough_auth_provider(&cfg, Some("captured-tok".to_string())).unwrap();
953 let mut headers = HeaderMap::new();
954 let mut query = HashMap::new();
955 auth.apply(&mut headers, &mut query, None).await.unwrap();
956 assert_eq!(
957 headers.get(reqwest::header::AUTHORIZATION).unwrap(),
958 "Bearer captured-tok"
959 );
960 }
961
962 #[tokio::test]
963 async fn test_passthrough_required_missing_token_errors() {
964 let cfg = AuthConfig::OAuthPassthrough {
965 target_header: "Authorization".to_string(),
966 required: true,
967 };
968 let auth = create_passthrough_auth_provider(&cfg, None).unwrap();
969 let mut headers = HeaderMap::new();
970 let mut query = HashMap::new();
971 let err = auth
972 .apply(&mut headers, &mut query, None)
973 .await
974 .unwrap_err();
975 assert!(matches!(err, HttpConnectorError::Auth(_)));
976 }
977
978 #[test]
979 fn test_oauth_passthrough_documented_tag_deserializes() {
980 let cfg: AuthConfig = toml::from_str(r#"type = "oauth_passthrough""#)
984 .expect("documented oauth_passthrough tag must deserialize via the serde alias");
985 assert!(matches!(cfg, AuthConfig::OAuthPassthrough { .. }));
986 }
987
988 #[test]
989 fn test_oauth2_client_credentials_documented_tag_deserializes() {
990 let cfg: AuthConfig = toml::from_str(
991 r#"
992 type = "oauth2_client_credentials"
993 token_url = "https://example.test/token"
994 client_id = "${CID}"
995 client_secret = "${CSECRET}"
996 "#,
997 )
998 .expect("documented oauth2_client_credentials tag must deserialize via the serde alias");
999 assert!(matches!(cfg, AuthConfig::OAuth2ClientCredentials { .. }));
1000 }
1001
1002 #[test]
1003 fn test_snake_case_tag_still_deserializes_after_alias() {
1004 let cfg: AuthConfig = toml::from_str(r#"type = "o_auth_passthrough""#)
1007 .expect("canonical snake_case tag must still deserialize");
1008 assert!(matches!(cfg, AuthConfig::OAuthPassthrough { .. }));
1009 }
1010
1011 #[tokio::test]
1012 async fn test_static_provider_ignores_inbound_token() {
1013 let bearer = create_auth_provider(&AuthConfig::Bearer {
1016 token: "static-tok".to_string(),
1017 required: true,
1018 })
1019 .unwrap();
1020 let mut headers = HeaderMap::new();
1021 let mut query = HashMap::new();
1022 bearer
1023 .apply(&mut headers, &mut query, Some("client-tok"))
1024 .await
1025 .unwrap();
1026 let rendered = headers
1027 .get(reqwest::header::AUTHORIZATION)
1028 .unwrap()
1029 .to_str()
1030 .unwrap();
1031 assert_eq!(rendered, "Bearer static-tok");
1032 assert!(
1033 !rendered.contains("client-tok"),
1034 "static provider must not forward the inbound token"
1035 );
1036
1037 let apikey = create_auth_provider(&AuthConfig::ApiKey {
1039 query_params: [("app_key".to_string(), "kkk".to_string())]
1040 .into_iter()
1041 .collect(),
1042 headers: HashMap::new(),
1043 required: true,
1044 })
1045 .unwrap();
1046 let mut headers2 = HeaderMap::new();
1047 let mut query2 = HashMap::new();
1048 apikey
1049 .apply(&mut headers2, &mut query2, Some("client-tok"))
1050 .await
1051 .unwrap();
1052 assert_eq!(query2.get("app_key"), Some(&"kkk".to_string()));
1053 assert!(
1054 !query2.values().any(|v| v.contains("client-tok")),
1055 "static api-key provider must not forward the inbound token"
1056 );
1057 assert!(headers2.is_empty());
1058 }
1059
1060 #[tokio::test]
1061 async fn test_auth_error_display_no_secret() {
1062 let cfg = AuthConfig::OAuthPassthrough {
1064 target_header: "Authorization".to_string(),
1065 required: true,
1066 };
1067 let auth = create_passthrough_auth_provider(&cfg, None).unwrap();
1068 let mut headers = HeaderMap::new();
1069 let mut query = HashMap::new();
1070 let err = auth
1071 .apply(&mut headers, &mut query, None)
1072 .await
1073 .unwrap_err();
1074 let rendered = err.to_string();
1075 for forbidden in ["Bearer", "client-tok", "app_key", "https://"] {
1076 assert!(
1077 !rendered.contains(forbidden),
1078 "auth error Display must not echo {forbidden:?}; got {rendered:?}"
1079 );
1080 }
1081 }
1082
1083 #[test]
1084 fn test_auth_config_deserializes_snake_case_tag() {
1085 let toml_src = r#"type = "bearer"
1086token = "abc"
1087"#;
1088 let cfg: AuthConfig = toml::from_str(toml_src).unwrap();
1089 assert!(matches!(cfg, AuthConfig::Bearer { .. }));
1090 assert!(cfg.is_required());
1091 }
1092
1093 #[test]
1094 fn test_auth_config_default_is_none() {
1095 assert!(matches!(AuthConfig::default(), AuthConfig::None));
1096 assert!(!AuthConfig::None.is_required());
1097 }
1098
1099 #[test]
1104 fn test_resolve_secret_ref_forms() {
1105 let var = "PMCP_TEST_RESOLVE_SECRET_REF_FORM";
1106 std::env::set_var(var, "secret");
1107 assert_eq!(resolve_secret_ref(&format!("${{{var}}}")), "secret");
1108 assert_eq!(resolve_secret_ref(&format!("env:{var}")), "secret");
1109 assert_eq!(resolve_secret_ref("plain-literal"), "plain-literal");
1110 std::env::remove_var(var);
1111 assert_eq!(resolve_secret_ref(&format!("${{{var}}}")), "");
1113 assert_eq!(resolve_secret_ref("${}"), "");
1114 }
1115
1116 #[test]
1135 fn create_auth_provider_refuses_malformed_bearer_token_ref() {
1136 let cfg = AuthConfig::Bearer {
1137 token: "${TFL-APP-KEY}".to_string(),
1138 required: true,
1139 };
1140 let Err(err) = create_auth_provider(&cfg) else {
1142 panic!("a malformed ref must not build a provider");
1143 };
1144 let msg = err.to_string();
1145 assert!(msg.contains("token"), "the error names the field: {msg}");
1146 assert!(
1147 !msg.contains("TFL-APP-KEY"),
1148 "the error must not echo the configured value: {msg}"
1149 );
1150 }
1151
1152 #[test]
1153 fn create_auth_provider_refuses_malformed_api_key_map_value() {
1154 let cfg = AuthConfig::ApiKey {
1158 query_params: [("app_key".to_string(), "${TFL-APP-KEY}".to_string())]
1159 .into_iter()
1160 .collect(),
1161 headers: HashMap::new(),
1162 required: true,
1163 };
1164 let Err(err) = create_auth_provider(&cfg) else {
1165 panic!("a malformed ref must not build a provider");
1166 };
1167 let msg = err.to_string();
1168 assert!(
1169 msg.contains("query_params.app_key"),
1170 "the error names the offending map entry: {msg}"
1171 );
1172 }
1173
1174 #[test]
1175 fn create_auth_provider_refuses_malformed_composition_in_every_variant() {
1176 let composed = "${A}://${B}".to_string();
1178 let variants = [
1179 AuthConfig::Bearer {
1180 token: composed.clone(),
1181 required: false,
1182 },
1183 AuthConfig::Basic {
1184 username: "user".to_string(),
1185 password: composed.clone(),
1186 required: false,
1187 },
1188 AuthConfig::OAuth2ClientCredentials {
1189 token_url: "https://example.com/token".to_string(),
1190 client_id: "id".to_string(),
1191 client_secret: composed.clone(),
1192 scopes: vec![],
1193 required: false,
1194 },
1195 AuthConfig::ApiKey {
1196 query_params: HashMap::new(),
1197 headers: [("X-Api-Key".to_string(), composed)].into_iter().collect(),
1198 required: false,
1199 },
1200 ];
1201 for cfg in variants {
1202 assert!(
1203 create_auth_provider(&cfg).is_err(),
1204 "every credential-bearing variant refuses a malformed ref: {cfg:?}"
1205 );
1206 }
1207 }
1208
1209 #[test]
1210 fn malformed_api_key_report_is_deterministic_across_map_orderings() {
1211 for _ in 0..20 {
1215 let cfg = AuthConfig::ApiKey {
1216 query_params: [
1217 ("zeta".to_string(), "${}".to_string()),
1218 ("alpha".to_string(), "${}".to_string()),
1219 ("mid".to_string(), "${}".to_string()),
1220 ]
1221 .into_iter()
1222 .collect(),
1223 headers: HashMap::new(),
1224 required: true,
1225 };
1226 let Err(err) = create_auth_provider(&cfg) else {
1227 panic!("malformed refs are refused");
1228 };
1229 let msg = err.to_string();
1230 assert!(
1231 msg.contains("query_params.alpha"),
1232 "the lexicographically first offender is reported every time: {msg}"
1233 );
1234 }
1235 }
1236
1237 mod malformed_ref_properties {
1238 use super::super::{create_auth_provider, AuthConfig};
1239 use crate::env_ref::is_valid_env_var_name;
1240 use proptest::prelude::*;
1241
1242 proptest! {
1243 #[test]
1251 fn braced_credential_is_refused_exactly_when_its_name_is_unsettable(
1252 interior in "\\PC{0,64}"
1253 ) {
1254 let cfg = AuthConfig::Bearer {
1255 token: format!("${{{interior}}}"),
1256 required: true,
1257 };
1258 let settable = is_valid_env_var_name(&interior);
1259 prop_assert_eq!(cfg.malformed_env_ref_field().is_some(), !settable);
1260 prop_assert_eq!(create_auth_provider(&cfg).is_err(), !settable);
1261 }
1262
1263 #[test]
1267 fn plain_literal_credentials_are_never_refused(raw in "\\PC{0,64}") {
1268 prop_assume!(!raw.starts_with("env:") && !raw.starts_with("${"));
1269 let cfg = AuthConfig::Basic {
1270 username: "svc".to_string(),
1271 password: raw,
1272 required: true,
1273 };
1274 prop_assert!(cfg.malformed_env_ref_field().is_none());
1275 prop_assert!(create_auth_provider(&cfg).is_ok());
1276 }
1277
1278 #[test]
1282 fn guard_never_panics_on_arbitrary_credential_text(raw in "\\PC{0,256}") {
1283 let cfg = AuthConfig::Bearer { token: raw, required: true };
1284 let _ = cfg.malformed_env_ref_field();
1285 let _ = create_auth_provider(&cfg);
1286 }
1287 }
1288 }
1289
1290 #[tokio::test]
1291 async fn unset_wellformed_ref_still_collapses_to_no_auth() {
1292 let var = "PMCP_TEST_UNSET_POLICY_UNCHANGED";
1296 std::env::remove_var(var);
1297 let cfg = AuthConfig::Bearer {
1298 token: format!("${{{var}}}"),
1299 required: false,
1300 };
1301 let auth = create_auth_provider(&cfg).expect("a well-formed ref still builds a provider");
1302 let mut headers = HeaderMap::new();
1303 let mut query = HashMap::new();
1304 auth.apply(&mut headers, &mut query, None).await.unwrap();
1305 assert!(
1306 !headers.contains_key(reqwest::header::AUTHORIZATION),
1307 "an unset optional credential is omitted, not refused"
1308 );
1309 }
1310
1311 #[tokio::test]
1312 async fn test_bearer_resolves_braced_env_ref() {
1313 let var = "PMCP_TEST_BEARER_BRACED_PAT";
1314 std::env::set_var(var, "ghp_abc");
1315 let cfg = AuthConfig::Bearer {
1316 token: format!("${{{var}}}"),
1317 required: true,
1318 };
1319 let auth = create_auth_provider(&cfg).unwrap();
1320 let mut headers = HeaderMap::new();
1321 let mut query = HashMap::new();
1322 auth.apply(&mut headers, &mut query, None).await.unwrap();
1323 let rendered = headers
1324 .get(reqwest::header::AUTHORIZATION)
1325 .unwrap()
1326 .to_str()
1327 .unwrap();
1328 assert_eq!(rendered, "Bearer ghp_abc");
1329 assert!(
1330 !rendered.contains("${"),
1331 "the literal ${{...}} must never reach the Authorization header"
1332 );
1333 std::env::remove_var(var);
1334 }
1335
1336 #[tokio::test]
1337 async fn test_bearer_resolves_env_prefix_ref() {
1338 let var = "PMCP_TEST_BEARER_ENV_PAT";
1339 std::env::set_var(var, "ghp_xyz");
1340 let cfg = AuthConfig::Bearer {
1341 token: format!("env:{var}"),
1342 required: true,
1343 };
1344 let auth = create_auth_provider(&cfg).unwrap();
1345 let mut headers = HeaderMap::new();
1346 let mut query = HashMap::new();
1347 auth.apply(&mut headers, &mut query, None).await.unwrap();
1348 assert_eq!(
1349 headers.get(reqwest::header::AUTHORIZATION).unwrap(),
1350 "Bearer ghp_xyz"
1351 );
1352 std::env::remove_var(var);
1353 }
1354
1355 #[tokio::test]
1356 async fn test_bearer_unset_ref_collapses_to_no_auth() {
1357 let var = "PMCP_TEST_BEARER_UNSET_PAT";
1358 std::env::remove_var(var);
1359 let cfg = AuthConfig::Bearer {
1360 token: format!("${{{var}}}"),
1361 required: true,
1362 };
1363 let auth = create_auth_provider(&cfg).unwrap();
1364 let mut headers = HeaderMap::new();
1365 let mut query = HashMap::new();
1366 auth.apply(&mut headers, &mut query, None).await.unwrap();
1367 assert!(headers.is_empty());
1369 assert!(query.is_empty());
1370 }
1371
1372 #[tokio::test]
1373 async fn test_basic_resolves_password_braced_env_ref() {
1374 use base64::Engine;
1375 let var = "PMCP_TEST_BASIC_BRACED_PW";
1376 std::env::set_var(var, "s3cr3t");
1377 let cfg = AuthConfig::Basic {
1378 username: "u".to_string(),
1379 password: format!("${{{var}}}"),
1380 required: true,
1381 };
1382 let auth = create_auth_provider(&cfg).unwrap();
1383 let mut headers = HeaderMap::new();
1384 let mut query = HashMap::new();
1385 auth.apply(&mut headers, &mut query, None).await.unwrap();
1386 let rendered = headers
1387 .get(reqwest::header::AUTHORIZATION)
1388 .unwrap()
1389 .to_str()
1390 .unwrap();
1391 let expected = format!(
1392 "Basic {}",
1393 base64::engine::general_purpose::STANDARD.encode("u:s3cr3t")
1394 );
1395 assert_eq!(rendered, expected);
1396 assert!(
1397 !rendered.contains("${"),
1398 "the literal ${{...}} must never reach the Basic credential"
1399 );
1400 std::env::remove_var(var);
1401 }
1402
1403 #[tokio::test]
1404 async fn test_basic_resolves_password_env_prefix_ref() {
1405 use base64::Engine;
1406 let var = "PMCP_TEST_BASIC_ENV_PW";
1407 std::env::set_var(var, "p4ss");
1408 let cfg = AuthConfig::Basic {
1409 username: "user".to_string(),
1410 password: format!("env:{var}"),
1411 required: true,
1412 };
1413 let auth = create_auth_provider(&cfg).unwrap();
1414 let mut headers = HeaderMap::new();
1415 let mut query = HashMap::new();
1416 auth.apply(&mut headers, &mut query, None).await.unwrap();
1417 let expected = format!(
1418 "Basic {}",
1419 base64::engine::general_purpose::STANDARD.encode("user:p4ss")
1420 );
1421 assert_eq!(
1422 headers.get(reqwest::header::AUTHORIZATION).unwrap(),
1423 expected.as_str()
1424 );
1425 std::env::remove_var(var);
1426 }
1427
1428 #[tokio::test]
1429 async fn test_oauth2_resolves_client_secret_via_token_endpoint() {
1430 use wiremock::matchers::{body_string_contains, method, path};
1433 use wiremock::{Mock, MockServer, ResponseTemplate};
1434
1435 let var = "PMCP_TEST_OAUTH2_BRACED_CS";
1436 std::env::set_var(var, "xyz");
1437
1438 let server = MockServer::start().await;
1439 Mock::given(method("POST"))
1440 .and(path("/token"))
1441 .and(body_string_contains("client_secret=xyz"))
1442 .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
1443 "access_token": "issued-token"
1444 })))
1445 .mount(&server)
1446 .await;
1447
1448 let cfg = AuthConfig::OAuth2ClientCredentials {
1449 token_url: format!("{}/token", server.uri()),
1450 client_id: "cid".to_string(),
1451 client_secret: format!("${{{var}}}"),
1452 scopes: vec![],
1453 required: true,
1454 };
1455 let auth = create_auth_provider(&cfg).unwrap();
1456 let mut headers = HeaderMap::new();
1457 let mut query = HashMap::new();
1458 auth.apply(&mut headers, &mut query, None).await.unwrap();
1462 assert_eq!(
1463 headers.get(reqwest::header::AUTHORIZATION).unwrap(),
1464 "Bearer issued-token"
1465 );
1466 std::env::remove_var(var);
1467 }
1468
1469 #[tokio::test]
1470 async fn test_oauth2_unset_secret_collapses_to_no_auth() {
1471 let var = "PMCP_TEST_OAUTH2_UNSET_CS";
1472 std::env::remove_var(var);
1473 let cfg = AuthConfig::OAuth2ClientCredentials {
1474 token_url: "http://127.0.0.1:1/token".to_string(),
1475 client_id: "cid".to_string(),
1476 client_secret: format!("${{{var}}}"),
1477 scopes: vec![],
1478 required: true,
1479 };
1480 let auth = create_auth_provider(&cfg).unwrap();
1481 let mut headers = HeaderMap::new();
1482 let mut query = HashMap::new();
1483 auth.apply(&mut headers, &mut query, None).await.unwrap();
1485 assert!(headers.is_empty());
1486 assert!(query.is_empty());
1487 }
1488}