1use std::any::TypeId;
2use std::collections::HashMap;
3use std::sync::Arc;
4
5use async_trait::async_trait;
6
7use crate::secrets::Secret;
8
9pub trait CredentialTarget: Send + Sync + 'static {}
11
12#[async_trait]
14pub trait ProvideBearerToken<Target: CredentialTarget>: Send + Sync + 'static {
15 async fn resolve_bearer_token(&self, target: &Target) -> Result<Secret<String>, String>;
16}
17
18#[async_trait]
20pub trait ProvideApiKey<Target: CredentialTarget>: Send + Sync + 'static {
21 async fn resolve_api_key(&self, target: &Target) -> Result<Secret<String>, String>;
22}
23
24#[async_trait]
26pub trait ProvideBasicAuth<Target: CredentialTarget>: Send + Sync + 'static {
27 async fn resolve_basic_auth(&self, target: &Target)
28 -> Result<(String, Secret<String>), String>;
29}
30
31#[async_trait]
32trait ErasedBearerResolver: Send + Sync + 'static {
33 async fn resolve_erased(
34 &self,
35 target: &(dyn std::any::Any + Send + Sync),
36 ) -> Result<Secret<String>, String>;
37}
38
39struct TypedBearerResolver<Target, P> {
40 provider: Arc<P>,
41 _marker: std::marker::PhantomData<Target>,
42}
43
44#[async_trait]
45impl<Target: CredentialTarget, P: ProvideBearerToken<Target>> ErasedBearerResolver
46 for TypedBearerResolver<Target, P>
47{
48 async fn resolve_erased(
49 &self,
50 target: &(dyn std::any::Any + Send + Sync),
51 ) -> Result<Secret<String>, String> {
52 let typed_target = target.downcast_ref::<Target>().ok_or_else(|| {
53 format!(
54 "Invalid target type: expected {}",
55 std::any::type_name::<Target>()
56 )
57 })?;
58 self.provider.resolve_bearer_token(typed_target).await
59 }
60}
61
62#[async_trait]
63trait ErasedApiKeyResolver: Send + Sync + 'static {
64 async fn resolve_erased(
65 &self,
66 target: &(dyn std::any::Any + Send + Sync),
67 ) -> Result<Secret<String>, String>;
68}
69
70struct TypedApiKeyResolver<Target, P> {
71 provider: Arc<P>,
72 _marker: std::marker::PhantomData<Target>,
73}
74
75#[async_trait]
76impl<Target: CredentialTarget, P: ProvideApiKey<Target>> ErasedApiKeyResolver
77 for TypedApiKeyResolver<Target, P>
78{
79 async fn resolve_erased(
80 &self,
81 target: &(dyn std::any::Any + Send + Sync),
82 ) -> Result<Secret<String>, String> {
83 let typed_target = target.downcast_ref::<Target>().ok_or_else(|| {
84 format!(
85 "Invalid target type: expected {}",
86 std::any::type_name::<Target>()
87 )
88 })?;
89 self.provider.resolve_api_key(typed_target).await
90 }
91}
92
93#[async_trait]
94trait ErasedBasicAuthResolver: Send + Sync + 'static {
95 async fn resolve_erased(
96 &self,
97 target: &(dyn std::any::Any + Send + Sync),
98 ) -> Result<(String, Secret<String>), String>;
99}
100
101struct TypedBasicAuthResolver<Target, P> {
102 provider: Arc<P>,
103 _marker: std::marker::PhantomData<Target>,
104}
105
106#[async_trait]
107impl<Target: CredentialTarget, P: ProvideBasicAuth<Target>> ErasedBasicAuthResolver
108 for TypedBasicAuthResolver<Target, P>
109{
110 async fn resolve_erased(
111 &self,
112 target: &(dyn std::any::Any + Send + Sync),
113 ) -> Result<(String, Secret<String>), String> {
114 let typed_target = target.downcast_ref::<Target>().ok_or_else(|| {
115 format!(
116 "Invalid target type: expected {}",
117 std::any::type_name::<Target>()
118 )
119 })?;
120 self.provider.resolve_basic_auth(typed_target).await
121 }
122}
123
124#[derive(Default)]
126pub struct CredentialsBuilder {
127 bearer_resolvers: HashMap<TypeId, Arc<dyn ErasedBearerResolver>>,
128 api_key_resolvers: HashMap<TypeId, Arc<dyn ErasedApiKeyResolver>>,
129 basic_auth_resolvers: HashMap<TypeId, Arc<dyn ErasedBasicAuthResolver>>,
130}
131
132impl CredentialsBuilder {
133 pub fn register_bearer<Target: CredentialTarget, P: ProvideBearerToken<Target>>(
134 &mut self,
135 provider: Arc<P>,
136 ) -> Result<(), String> {
137 let type_id = TypeId::of::<Target>();
138 if self.bearer_resolvers.contains_key(&type_id) {
139 return Err(format!(
140 "Credentials resolver already registered for target: {}",
141 std::any::type_name::<Target>()
142 ));
143 }
144 let resolver = TypedBearerResolver {
145 provider,
146 _marker: std::marker::PhantomData,
147 };
148 self.bearer_resolvers.insert(type_id, Arc::new(resolver));
149 Ok(())
150 }
151
152 pub fn register_api_key<Target: CredentialTarget, P: ProvideApiKey<Target>>(
153 &mut self,
154 provider: Arc<P>,
155 ) -> Result<(), String> {
156 let type_id = TypeId::of::<Target>();
157 if self.api_key_resolvers.contains_key(&type_id) {
158 return Err(format!(
159 "Credentials resolver already registered for target: {}",
160 std::any::type_name::<Target>()
161 ));
162 }
163 let resolver = TypedApiKeyResolver {
164 provider,
165 _marker: std::marker::PhantomData,
166 };
167 self.api_key_resolvers.insert(type_id, Arc::new(resolver));
168 Ok(())
169 }
170
171 pub fn register_basic_auth<Target: CredentialTarget, P: ProvideBasicAuth<Target>>(
172 &mut self,
173 provider: Arc<P>,
174 ) -> Result<(), String> {
175 let type_id = TypeId::of::<Target>();
176 if self.basic_auth_resolvers.contains_key(&type_id) {
177 return Err(format!(
178 "Credentials resolver already registered for target: {}",
179 std::any::type_name::<Target>()
180 ));
181 }
182 let resolver = TypedBasicAuthResolver {
183 provider,
184 _marker: std::marker::PhantomData,
185 };
186 self.basic_auth_resolvers
187 .insert(type_id, Arc::new(resolver));
188 Ok(())
189 }
190
191 pub fn build(self) -> CredentialsHandle {
192 CredentialsHandle {
193 bearer_resolvers: Arc::new(self.bearer_resolvers),
194 api_key_resolvers: Arc::new(self.api_key_resolvers),
195 basic_auth_resolvers: Arc::new(self.basic_auth_resolvers),
196 }
197 }
198}
199
200#[derive(Clone, Default)]
202pub struct CredentialsHandle {
203 bearer_resolvers: Arc<HashMap<TypeId, Arc<dyn ErasedBearerResolver>>>,
204 api_key_resolvers: Arc<HashMap<TypeId, Arc<dyn ErasedApiKeyResolver>>>,
205 basic_auth_resolvers: Arc<HashMap<TypeId, Arc<dyn ErasedBasicAuthResolver>>>,
206}
207
208impl CredentialsHandle {
209 pub async fn resolve_bearer_token<Target: CredentialTarget>(
211 &self,
212 target: &Target,
213 ) -> Result<Secret<String>, String> {
214 let type_id = TypeId::of::<Target>();
215 if let Some(resolver) = self.bearer_resolvers.get(&type_id) {
216 resolver.resolve_erased(target).await
217 } else {
218 Err(format!(
219 "No registered credentials provider implements ProvideBearerToken<{}>",
220 std::any::type_name::<Target>()
221 ))
222 }
223 }
224
225 pub async fn resolve_api_key<Target: CredentialTarget>(
227 &self,
228 target: &Target,
229 ) -> Result<Secret<String>, String> {
230 let type_id = TypeId::of::<Target>();
231 if let Some(resolver) = self.api_key_resolvers.get(&type_id) {
232 resolver.resolve_erased(target).await
233 } else {
234 Err(format!(
235 "No registered credentials provider implements ProvideApiKey<{}>",
236 std::any::type_name::<Target>()
237 ))
238 }
239 }
240
241 pub async fn resolve_basic_auth<Target: CredentialTarget>(
243 &self,
244 target: &Target,
245 ) -> Result<(String, Secret<String>), String> {
246 let type_id = TypeId::of::<Target>();
247 if let Some(resolver) = self.basic_auth_resolvers.get(&type_id) {
248 resolver.resolve_erased(target).await
249 } else {
250 Err(format!(
251 "No registered credentials provider implements ProvideBasicAuth<{}>",
252 std::any::type_name::<Target>()
253 ))
254 }
255 }
256}
257
258pub trait PluggableCredentialsProvider: Send + Sync + 'static {
260 type Config: serde::de::DeserializeOwned + Default + Send + Sync;
261
262 fn register_provider(
263 config: Self::Config,
264 builder: &mut CredentialsBuilder,
265 ) -> Result<(), String>;
266}