Skip to main content

camel_function/provider/
mod.rs

1use crate::pool::{RunnerHandle, RunnerPoolKey};
2use camel_api::{Exchange, function::*};
3use std::time::Duration;
4
5mod sealed {
6    pub trait Sealed {}
7}
8
9/// Health status of a function provider instance.
10/// Named `FunctionHealthStatus` to avoid collision with `camel_api::FunctionHealthStatus`.
11#[derive(Debug, Clone)]
12pub enum FunctionHealthStatus {
13    Healthy,
14    Unhealthy(String),
15}
16
17#[derive(Debug, thiserror::Error)]
18pub enum ProviderError {
19    #[error("spawn failed: {0}")]
20    SpawnFailed(String),
21    #[error("health check failed: {0}")]
22    HealthFailed(String),
23    #[error("register failed: {0}")]
24    RegisterFailed(String),
25    #[error("unregister failed: {0}")]
26    UnregisterFailed(String),
27    #[error("invoke failed: {0}")]
28    InvokeFailed(String),
29    #[error("shutdown failed: {0}")]
30    ShutdownFailed(String),
31    #[error("invalid provider config: {0}")]
32    InvalidConfig(String),
33    #[error("boot timeout")]
34    BootTimeout,
35}
36
37#[async_trait::async_trait]
38pub(crate) trait FunctionProvider: Send + Sync + sealed::Sealed {
39    async fn spawn(&self, key: &RunnerPoolKey) -> Result<RunnerHandle, ProviderError>;
40    async fn shutdown(&self, handle: RunnerHandle) -> Result<(), ProviderError>;
41    async fn health(&self, handle: &RunnerHandle) -> Result<FunctionHealthStatus, ProviderError>;
42    async fn register(
43        &self,
44        handle: &RunnerHandle,
45        def: &FunctionDefinition,
46    ) -> Result<(), ProviderError>;
47    async fn unregister(&self, handle: &RunnerHandle, id: &FunctionId)
48    -> Result<(), ProviderError>;
49    async fn invoke(
50        &self,
51        handle: &RunnerHandle,
52        id: &FunctionId,
53        ex: &Exchange,
54        timeout: Duration,
55    ) -> Result<ExchangePatch, ProviderError>;
56}
57
58pub mod container;
59pub mod fake {
60    use super::*;
61    use std::collections::{HashMap, HashSet};
62    use std::sync::atomic::{AtomicUsize, Ordering};
63    use std::sync::{Arc, Mutex};
64    use tokio_util::sync::CancellationToken;
65
66    #[derive(Debug, Clone, Default)]
67    pub struct FakeProviderConfig {
68        pub fail_on_spawn: bool,
69        pub fail_on_register: usize,
70        pub fail_on_health: bool,
71        pub fail_on_shutdown: bool,
72        pub invoke_response: Option<ExchangePatch>,
73        pub invoke_delay: Option<std::time::Duration>,
74    }
75
76    #[derive(Debug, Clone)]
77    pub enum FakeCall {
78        Spawn(RunnerPoolKey),
79        Shutdown(RunnerPoolKey),
80        Health(String),
81        Register(String, FunctionId),
82        Unregister(String, FunctionId),
83        Invoke(String, FunctionId),
84    }
85
86    pub struct FakeProvider {
87        pub config: Arc<Mutex<FakeProviderConfig>>,
88        pub calls: Arc<Mutex<Vec<FakeCall>>>,
89        pub registered: Arc<Mutex<HashMap<String, HashSet<FunctionId>>>>,
90        pub spawned: Arc<Mutex<Vec<RunnerPoolKey>>>,
91        pub shutdowns: Arc<Mutex<Vec<RunnerPoolKey>>>,
92        register_ok_count: Arc<Mutex<usize>>,
93        spawn_count: AtomicUsize,
94    }
95
96    impl FakeProvider {
97        pub fn new(config: FakeProviderConfig) -> Self {
98            Self {
99                config: Arc::new(Mutex::new(config)),
100                calls: Arc::new(Mutex::new(Vec::new())),
101                registered: Arc::new(Mutex::new(HashMap::new())),
102                spawned: Arc::new(Mutex::new(Vec::new())),
103                shutdowns: Arc::new(Mutex::new(Vec::new())),
104                register_ok_count: Arc::new(Mutex::new(0)),
105                spawn_count: AtomicUsize::new(0),
106            }
107        }
108
109        pub fn spawn_count(&self) -> usize {
110            self.spawn_count.load(Ordering::SeqCst)
111        }
112    }
113
114    impl super::sealed::Sealed for FakeProvider {}
115
116    #[async_trait::async_trait]
117    impl FunctionProvider for FakeProvider {
118        async fn spawn(&self, key: &RunnerPoolKey) -> Result<RunnerHandle, ProviderError> {
119            self.spawn_count.fetch_add(1, Ordering::SeqCst);
120            self.calls
121                .lock()
122                .expect("calls") // allow-unwrap
123                .push(FakeCall::Spawn(key.clone()));
124            self.spawned.lock().expect("spawned").push(key.clone()); // allow-unwrap
125            if self.config.lock().expect("config").fail_on_spawn {
126                // allow-unwrap
127                return Err(ProviderError::SpawnFailed("configured".into()));
128            }
129            Ok(RunnerHandle {
130                id: format!("fake-{}", key.runtime),
131                state: Arc::new(Mutex::new(crate::pool::RunnerState::Booting)),
132                cancel: CancellationToken::new(),
133            })
134        }
135
136        async fn shutdown(&self, handle: RunnerHandle) -> Result<(), ProviderError> {
137            self.calls
138                .lock()
139                .expect("calls") // allow-unwrap
140                .push(FakeCall::Shutdown(RunnerPoolKey {
141                    runtime: handle.id.replace("fake-", ""),
142                }));
143            self.shutdowns
144                .lock()
145                .expect("shutdowns") // allow-unwrap
146                .push(RunnerPoolKey {
147                    runtime: handle.id.replace("fake-", ""),
148                });
149            if self.config.lock().expect("config").fail_on_shutdown {
150                // allow-unwrap
151                return Err(ProviderError::ShutdownFailed(
152                    "configured shutdown failure".into(),
153                ));
154            }
155            Ok(())
156        }
157
158        async fn health(
159            &self,
160            handle: &RunnerHandle,
161        ) -> Result<FunctionHealthStatus, ProviderError> {
162            self.calls
163                .lock()
164                .expect("calls") // allow-unwrap
165                .push(FakeCall::Health(handle.id.clone()));
166            if self.config.lock().expect("config").fail_on_health {
167                // allow-unwrap
168                return Ok(FunctionHealthStatus::Unhealthy("configured".into()));
169            }
170            Ok(FunctionHealthStatus::Healthy)
171        }
172
173        async fn register(
174            &self,
175            handle: &RunnerHandle,
176            def: &FunctionDefinition,
177        ) -> Result<(), ProviderError> {
178            self.calls
179                .lock()
180                .expect("calls") // allow-unwrap
181                .push(FakeCall::Register(handle.id.clone(), def.id.clone()));
182            let mut count = self.register_ok_count.lock().expect("count"); // allow-unwrap
183            let cfg = self.config.lock().expect("config").clone(); // allow-unwrap
184            if cfg.fail_on_register > 0 && *count >= cfg.fail_on_register {
185                return Err(ProviderError::RegisterFailed("configured".into()));
186            }
187            *count += 1;
188            self.registered
189                .lock()
190                .expect("registered") // allow-unwrap
191                .entry(handle.id.clone())
192                .or_default()
193                .insert(def.id.clone());
194            Ok(())
195        }
196
197        async fn unregister(
198            &self,
199            handle: &RunnerHandle,
200            id: &FunctionId,
201        ) -> Result<(), ProviderError> {
202            self.calls
203                .lock()
204                .expect("calls") // allow-unwrap
205                .push(FakeCall::Unregister(handle.id.clone(), id.clone()));
206            if let Some(set) = self
207                .registered
208                .lock()
209                .expect("registered") // allow-unwrap
210                .get_mut(&handle.id)
211            {
212                set.remove(id);
213            }
214            Ok(())
215        }
216
217        async fn invoke(
218            &self,
219            handle: &RunnerHandle,
220            id: &FunctionId,
221            _ex: &Exchange,
222            _timeout: Duration,
223        ) -> Result<ExchangePatch, ProviderError> {
224            self.calls
225                .lock()
226                .expect("calls") // allow-unwrap
227                .push(FakeCall::Invoke(handle.id.clone(), id.clone()));
228            let exists = self
229                .registered
230                .lock()
231                .expect("registered") // allow-unwrap
232                .get(&handle.id)
233                .map(|s| s.contains(id))
234                .unwrap_or(false);
235            if !exists {
236                return Err(ProviderError::InvokeFailed("not registered".into()));
237            }
238            let cfg = self.config.lock().expect("config").clone(); // allow-unwrap
239            if let Some(delay) = cfg.invoke_delay {
240                tokio::time::sleep(delay).await;
241            }
242            Ok(cfg.invoke_response.unwrap_or_default())
243        }
244    }
245}