rskit_workload/
registry.rs1use std::collections::BTreeMap;
8use std::sync::Arc;
9
10use async_trait::async_trait;
11use rskit_errors::{AppError, AppResult, ErrorCode};
12
13use crate::config::WorkloadConfig;
14use crate::manager::Manager;
15
16#[async_trait]
21pub trait ManagerFactory: Send + Sync {
22 async fn create(&self, config: &WorkloadConfig) -> AppResult<Arc<dyn Manager>>;
24}
25
26#[derive(Default)]
28pub struct WorkloadRegistry {
29 factories: BTreeMap<String, Arc<dyn ManagerFactory>>,
30}
31
32impl WorkloadRegistry {
33 #[must_use]
35 pub fn new() -> Self {
36 Self::default()
37 }
38
39 pub fn register(
46 &mut self,
47 name: impl Into<String>,
48 factory: Arc<dyn ManagerFactory>,
49 ) -> AppResult<()> {
50 let name = name.into().trim().to_owned();
51 if name.is_empty() {
52 return Err(AppError::new(
53 ErrorCode::InvalidInput,
54 "workload provider name is required",
55 ));
56 }
57 if self.factories.contains_key(&name) {
58 return Err(AppError::new(
59 ErrorCode::AlreadyExists,
60 format!("workload provider '{name}' is already registered"),
61 ));
62 }
63 self.factories.insert(name, factory);
64 Ok(())
65 }
66
67 #[must_use]
71 pub fn contains(&self, name: &str) -> bool {
72 self.factories.contains_key(name.trim())
73 }
74
75 #[must_use]
77 pub fn len(&self) -> usize {
78 self.factories.len()
79 }
80
81 #[must_use]
83 pub fn is_empty(&self) -> bool {
84 self.factories.is_empty()
85 }
86
87 #[must_use]
89 pub fn names(&self) -> Vec<String> {
90 self.factories.keys().cloned().collect()
91 }
92
93 pub async fn build(&self, config: &WorkloadConfig) -> AppResult<Arc<dyn Manager>> {
101 config.validate()?;
102 let provider = config.provider.trim();
103 self.factories
104 .get(provider)
105 .ok_or_else(|| {
106 AppError::new(
107 ErrorCode::NotFound,
108 format!("workload provider '{provider}' is not registered"),
109 )
110 })?
111 .create(config)
112 .await
113 }
114}
115
116#[cfg(test)]
117mod tests {
118 use std::sync::atomic::{AtomicUsize, Ordering};
119
120 use super::*;
121 use crate::test_support::FakeManager;
122
123 struct CountingFactory {
124 calls: Arc<AtomicUsize>,
125 }
126
127 #[async_trait]
128 impl ManagerFactory for CountingFactory {
129 async fn create(&self, _config: &WorkloadConfig) -> AppResult<Arc<dyn Manager>> {
130 self.calls.fetch_add(1, Ordering::SeqCst);
131 Ok(Arc::new(FakeManager))
132 }
133 }
134
135 fn factory(calls: &Arc<AtomicUsize>) -> Arc<CountingFactory> {
136 Arc::new(CountingFactory {
137 calls: Arc::clone(calls),
138 })
139 }
140
141 #[tokio::test]
142 async fn registers_and_builds_selected_provider() {
143 let calls = Arc::new(AtomicUsize::new(0));
144 let mut registry = WorkloadRegistry::new();
145 assert!(registry.is_empty());
146
147 registry.register(" docker ", factory(&calls)).unwrap();
148 assert!(registry.contains("docker"));
149 assert!(registry.contains(" docker "));
150 assert_eq!(registry.len(), 1);
151 assert_eq!(registry.names(), vec!["docker".to_string()]);
152
153 let config = WorkloadConfig {
154 provider: "docker".to_string(),
155 ..Default::default()
156 };
157 registry.build(&config).await.unwrap();
158 assert_eq!(calls.load(Ordering::SeqCst), 1);
159 }
160
161 #[tokio::test]
162 async fn rejects_empty_and_duplicate_names() {
163 let calls = Arc::new(AtomicUsize::new(0));
164 let mut registry = WorkloadRegistry::new();
165 assert_eq!(
166 registry.register(" ", factory(&calls)).unwrap_err().code(),
167 ErrorCode::InvalidInput
168 );
169 registry.register("docker", factory(&calls)).unwrap();
170 assert_eq!(
171 registry
172 .register("docker", factory(&calls))
173 .unwrap_err()
174 .code(),
175 ErrorCode::AlreadyExists
176 );
177 }
178
179 #[tokio::test]
180 async fn build_reports_missing_and_unregistered_providers() {
181 let registry = WorkloadRegistry::new();
182 let empty = WorkloadConfig {
183 provider: String::new(),
184 ..Default::default()
185 };
186 assert_eq!(
187 registry.build(&empty).await.map(|_| ()).unwrap_err().code(),
188 ErrorCode::MissingField
189 );
190
191 let unknown = WorkloadConfig {
192 provider: "podman".to_string(),
193 ..Default::default()
194 };
195 assert_eq!(
196 registry
197 .build(&unknown)
198 .await
199 .map(|_| ())
200 .unwrap_err()
201 .code(),
202 ErrorCode::NotFound
203 );
204 }
205}