1use crate::handle::{EmbeddingHandle, NetworkHandle};
11use std::collections::HashMap;
12
13#[derive(Debug, Clone)]
17#[non_exhaustive]
18pub struct NetworkConfig {
19 pub name: String,
21
22 pub batching: bool,
25
26 pub k: Option<usize>,
29
30 pub det: bool,
33
34 pub cache_enabled: bool,
37
38 pub cache_size: usize,
40
41 pub arity: Option<usize>,
48
49 pub arg_sorts: Option<Vec<i64>>,
56
57 pub artifact_hash: Option<String>,
62}
63
64impl NetworkConfig {
65 pub fn default(name: &str) -> Self {
74 Self {
75 name: name.to_string(),
76 batching: true,
77 k: None,
78 det: false,
79 cache_enabled: true,
80 cache_size: 10000,
81 arity: None,
82 arg_sorts: None,
83 artifact_hash: None,
84 }
85 }
86
87 pub fn deterministic(name: &str) -> Self {
89 Self {
90 name: name.to_string(),
91 batching: true,
92 k: None,
93 det: true,
94 cache_enabled: true,
95 cache_size: 10000,
96 arity: None,
97 arg_sorts: None,
98 artifact_hash: None,
99 }
100 }
101
102 pub fn with_top_k(name: &str, k: usize) -> Self {
104 Self {
105 name: name.to_string(),
106 batching: true,
107 k: Some(k),
108 det: false,
109 cache_enabled: true,
110 cache_size: 10000,
111 arity: None,
112 arg_sorts: None,
113 artifact_hash: None,
114 }
115 }
116
117 pub fn batching(mut self, enabled: bool) -> Self {
119 self.batching = enabled;
120 self
121 }
122
123 pub fn k(mut self, k: Option<usize>) -> Self {
125 self.k = k;
126 self
127 }
128
129 pub fn det(mut self, det: bool) -> Self {
131 self.det = det;
132 self
133 }
134
135 pub fn cache(mut self, enabled: bool, size: usize) -> Self {
137 self.cache_enabled = enabled;
138 self.cache_size = size;
139 self
140 }
141}
142
143pub struct NetworkRegistry {
150 networks: HashMap<String, NetworkHandle>,
152 embeddings: HashMap<String, EmbeddingHandle>,
154}
155
156impl NetworkRegistry {
157 pub fn new() -> Self {
159 Self {
160 networks: HashMap::new(),
161 embeddings: HashMap::new(),
162 }
163 }
164
165 pub fn register(&mut self, config: NetworkConfig) {
169 let handle = NetworkHandle::from_config(&config);
170 self.networks.insert(config.name, handle);
171 }
172
173 pub fn get(&self, name: &str) -> Option<&NetworkHandle> {
175 self.networks.get(name)
176 }
177
178 pub fn get_mut(&mut self, name: &str) -> Option<&mut NetworkHandle> {
180 self.networks.get_mut(name)
181 }
182
183 pub fn contains(&self, name: &str) -> bool {
185 self.networks.contains_key(name)
186 }
187
188 pub fn unregister(&mut self, name: &str) -> Option<NetworkHandle> {
190 self.networks.remove(name)
191 }
192
193 pub fn set_train_mode(&mut self, train: bool) {
198 for handle in self.networks.values_mut() {
199 handle.train_mode = train;
200 }
201 }
202
203 pub fn names(&self) -> Vec<&str> {
205 self.networks.keys().map(|s| s.as_str()).collect()
206 }
207
208 pub fn len(&self) -> usize {
210 self.networks.len()
211 }
212
213 pub fn is_empty(&self) -> bool {
215 self.networks.is_empty()
216 }
217
218 pub fn clear(&mut self) {
220 self.networks.clear();
221 }
222
223 pub fn iter(&self) -> impl Iterator<Item = (&str, &NetworkHandle)> {
225 self.networks.iter().map(|(k, v)| (k.as_str(), v))
226 }
227
228 pub fn iter_mut(&mut self) -> impl Iterator<Item = (&str, &mut NetworkHandle)> {
230 self.networks.iter_mut().map(|(k, v)| (k.as_str(), v))
231 }
232
233 pub fn register_embedding(&mut self, handle: EmbeddingHandle) {
235 self.embeddings.insert(handle.name.clone(), handle);
236 }
237
238 pub fn get_embedding(&self, name: &str) -> Option<&EmbeddingHandle> {
240 self.embeddings.get(name)
241 }
242
243 pub fn get_embedding_mut(&mut self, name: &str) -> Option<&mut EmbeddingHandle> {
245 self.embeddings.get_mut(name)
246 }
247
248 pub fn contains_embedding(&self, name: &str) -> bool {
250 self.embeddings.contains_key(name)
251 }
252}
253
254impl Default for NetworkRegistry {
255 fn default() -> Self {
256 Self::new()
257 }
258}
259
260#[cfg(test)]
261mod tests {
262 use super::*;
263
264 #[test]
265 fn test_config_default() {
266 let config = NetworkConfig::default("test");
267 assert_eq!(config.name, "test");
268 assert!(config.batching);
269 assert!(config.k.is_none());
270 assert!(!config.det);
271 assert!(config.cache_enabled);
272 assert_eq!(config.cache_size, 10000);
273 assert!(config.arity.is_none());
274 assert!(config.arg_sorts.is_none());
275 assert!(config.artifact_hash.is_none());
276 }
277
278 #[test]
279 fn test_config_registration_metadata_none_across_constructors() {
280 let det = NetworkConfig::deterministic("det_test");
284 assert!(det.arity.is_none());
285 assert!(det.arg_sorts.is_none());
286 assert!(det.artifact_hash.is_none());
287
288 let top_k = NetworkConfig::with_top_k("top_k_test", 5);
289 assert!(top_k.arity.is_none());
290 assert!(top_k.arg_sorts.is_none());
291 assert!(top_k.artifact_hash.is_none());
292 }
293
294 #[test]
295 fn test_config_registration_metadata_survives_clone() {
296 let mut config = NetworkConfig::default("meta_test");
297 config.arity = Some(2);
298 config.arg_sorts = Some(vec![0, 1]);
299 config.artifact_hash = Some("deadbeef".to_string());
300
301 let cloned = config.clone();
302 assert_eq!(cloned.arity, Some(2));
303 assert_eq!(cloned.arg_sorts, Some(vec![0, 1]));
304 assert_eq!(cloned.artifact_hash, Some("deadbeef".to_string()));
305 }
306
307 #[test]
308 fn test_config_deterministic() {
309 let config = NetworkConfig::deterministic("det_test");
310 assert!(config.det);
311 }
312
313 #[test]
314 fn test_config_with_top_k() {
315 let config = NetworkConfig::with_top_k("top_k_test", 5);
316 assert_eq!(config.k, Some(5));
317 }
318
319 #[test]
320 fn test_config_builder() {
321 let config = NetworkConfig::default("builder_test")
322 .batching(false)
323 .k(Some(3))
324 .det(true)
325 .cache(false, 0);
326
327 assert!(!config.batching);
328 assert_eq!(config.k, Some(3));
329 assert!(config.det);
330 assert!(!config.cache_enabled);
331 assert_eq!(config.cache_size, 0);
332 }
333
334 #[test]
335 fn test_registry_new() {
336 let registry = NetworkRegistry::new();
337 assert!(registry.is_empty());
338 assert_eq!(registry.len(), 0);
339 }
340
341 #[test]
342 fn test_registry_register_get() {
343 let mut registry = NetworkRegistry::new();
344 registry.register(NetworkConfig::default("net1"));
345
346 assert!(registry.contains("net1"));
347 assert!(registry.get("net1").is_some());
348 assert!(registry.get("nonexistent").is_none());
349 }
350
351 #[test]
352 fn test_registry_iter() {
353 let mut registry = NetworkRegistry::new();
354 registry.register(NetworkConfig::default("a"));
355 registry.register(NetworkConfig::default("b"));
356
357 let names: Vec<&str> = registry.iter().map(|(name, _)| name).collect();
358 assert_eq!(names.len(), 2);
359 }
360
361 use crate::handle::EmbeddingHandle;
362
363 #[test]
364 fn test_registry_embedding_register_get() {
365 let mut registry = NetworkRegistry::new();
366 let handle = EmbeddingHandle::new("embed1".to_string(), true, 64, 100);
367 registry.register_embedding(handle);
368
369 assert!(registry.contains_embedding("embed1"));
370 assert!(!registry.contains_embedding("nonexistent"));
371
372 let h = registry.get_embedding("embed1").unwrap();
373 assert_eq!(h.dim, 64);
374 assert_eq!(h.vocab_size, 100);
375 }
376
377 #[test]
378 fn test_registry_embedding_get_mut() {
379 let mut registry = NetworkRegistry::new();
380 let handle = EmbeddingHandle::new("embed1".to_string(), true, 64, 100);
381 registry.register_embedding(handle);
382
383 let h = registry.get_embedding_mut("embed1").unwrap();
384 h.trainable = false;
385 assert!(!registry.get_embedding("embed1").unwrap().trainable);
386 }
387}