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]
72 pub fn contains(&self, name: &str) -> bool {
73 self.factories.contains_key(name.trim())
74 }
75
76 #[must_use]
78 pub fn len(&self) -> usize {
79 self.factories.len()
80 }
81
82 #[must_use]
84 pub fn is_empty(&self) -> bool {
85 self.factories.is_empty()
86 }
87
88 #[must_use]
90 pub fn names(&self) -> Vec<String> {
91 self.factories.keys().cloned().collect()
92 }
93
94 pub async fn build(&self, config: &WorkloadConfig) -> AppResult<Arc<dyn Manager>> {
102 config.validate()?;
103 let provider = config.provider.trim();
104 self.factories
105 .get(provider)
106 .ok_or_else(|| {
107 AppError::new(
108 ErrorCode::NotFound,
109 format!("workload provider '{provider}' is not registered"),
110 )
111 })?
112 .create(config)
113 .await
114 }
115}
116
117#[cfg(test)]
118mod tests {
119 use std::sync::atomic::{AtomicUsize, Ordering};
120
121 use super::*;
122 use crate::test_support::FakeManager;
123
124 struct CountingFactory {
125 calls: Arc<AtomicUsize>,
126 }
127
128 #[async_trait]
129 impl ManagerFactory for CountingFactory {
130 async fn create(&self, _config: &WorkloadConfig) -> AppResult<Arc<dyn Manager>> {
131 self.calls.fetch_add(1, Ordering::SeqCst);
132 Ok(Arc::new(FakeManager))
133 }
134 }
135
136 fn factory(calls: &Arc<AtomicUsize>) -> Arc<CountingFactory> {
137 Arc::new(CountingFactory {
138 calls: Arc::clone(calls),
139 })
140 }
141
142 #[tokio::test]
143 async fn registers_and_builds_selected_provider() {
144 let calls = Arc::new(AtomicUsize::new(0));
145 let mut registry = WorkloadRegistry::new();
146 assert!(registry.is_empty());
147
148 registry.register(" docker ", factory(&calls)).unwrap();
149 assert!(registry.contains("docker"));
150 assert!(registry.contains(" docker "));
151 assert_eq!(registry.len(), 1);
152 assert_eq!(registry.names(), vec!["docker".to_string()]);
153
154 let config = WorkloadConfig {
155 provider: "docker".to_string(),
156 ..Default::default()
157 };
158 registry.build(&config).await.unwrap();
159 assert_eq!(calls.load(Ordering::SeqCst), 1);
160 }
161
162 #[tokio::test]
163 async fn rejects_empty_and_duplicate_names() {
164 let calls = Arc::new(AtomicUsize::new(0));
165 let mut registry = WorkloadRegistry::new();
166 assert_eq!(
167 registry.register(" ", factory(&calls)).unwrap_err().code(),
168 ErrorCode::InvalidInput
169 );
170 registry.register("docker", factory(&calls)).unwrap();
171 assert_eq!(
172 registry
173 .register("docker", factory(&calls))
174 .unwrap_err()
175 .code(),
176 ErrorCode::AlreadyExists
177 );
178 }
179
180 #[tokio::test]
181 async fn build_reports_missing_and_unregistered_providers() {
182 let registry = WorkloadRegistry::new();
183 let empty = WorkloadConfig {
184 provider: String::new(),
185 ..Default::default()
186 };
187 assert_eq!(
188 registry.build(&empty).await.map(|_| ()).unwrap_err().code(),
189 ErrorCode::MissingField
190 );
191
192 let unknown = WorkloadConfig {
193 provider: "podman".to_string(),
194 ..Default::default()
195 };
196 assert_eq!(
197 registry
198 .build(&unknown)
199 .await
200 .map(|_| ())
201 .unwrap_err()
202 .code(),
203 ErrorCode::NotFound
204 );
205 }
206}