Skip to main content

rskit_inference/
echo.rs

1//! Deterministic echo inference adapter.
2
3use async_trait::async_trait;
4use rskit_ai::{Capabilities, Model, Provider as ModelProvider, Usage};
5use rskit_component::{Component, Health};
6use rskit_errors::AppResult;
7use rskit_tool::Envelope;
8
9use crate::{
10    Factory, Inference, InferenceDescriptor, InferenceError, PredictRequest, PredictResponse,
11    PredictStatus, Registry, RegistryError, ServingProtocol,
12};
13
14/// Registry kind for the echo adapter.
15pub const ECHO_KIND: &str = "echo";
16
17/// Echo inference adapter for tests and local composition.
18#[derive(Debug, Clone, Default)]
19pub struct Echo;
20
21#[async_trait]
22impl rskit_provider::Provider for Echo {
23    fn name(&self) -> &'static str {
24        ECHO_KIND
25    }
26}
27
28#[async_trait]
29impl rskit_provider::RequestResponse<PredictRequest, PredictResponse> for Echo {
30    async fn execute(&self, input: PredictRequest) -> AppResult<PredictResponse> {
31        self.predict(input).await.map_err(Into::into)
32    }
33}
34
35#[async_trait]
36impl Inference for Echo {
37    async fn predict(&self, request: PredictRequest) -> Result<PredictResponse, InferenceError> {
38        Ok(PredictResponse {
39            outputs: request.inputs,
40            usage: Usage::default(),
41            model: Model {
42                name: request.model_name,
43                provider: ModelProvider::Custom("echo".to_string()),
44                version: request.model_version,
45                capabilities: Capabilities::default(),
46            },
47            status: PredictStatus::Success,
48            metadata: Default::default(),
49        })
50    }
51
52    fn descriptor(&self) -> InferenceDescriptor {
53        InferenceDescriptor {
54            name: ECHO_KIND.to_string(),
55            description: "Echo inputs unchanged for tests".to_string(),
56            serving_protocol: ServingProtocol::Custom,
57            envelope: Envelope::default(),
58        }
59    }
60}
61
62/// Explicitly register the echo adapter.
63pub fn register(registry: &mut Registry) -> Result<(), RegistryError> {
64    let factory: Factory = std::sync::Arc::new(|| Ok(std::sync::Arc::new(Echo)));
65    registry.register(ECHO_KIND, factory)
66}
67
68#[async_trait]
69impl Component for Echo {
70    fn name(&self) -> &str {
71        "rskit-inference.echo"
72    }
73
74    async fn start(&self) -> rskit_errors::AppResult<()> {
75        Ok(())
76    }
77
78    async fn stop(&self) -> rskit_errors::AppResult<()> {
79        Ok(())
80    }
81
82    fn health(&self) -> Health {
83        Health::healthy(self.name())
84    }
85}
86
87#[cfg(test)]
88mod tests {
89    use super::*;
90    use crate::Value;
91    use std::collections::HashMap;
92
93    #[tokio::test]
94    async fn echo_returns_inputs_unchanged() {
95        let adapter = Echo;
96        let inputs = HashMap::from([(
97            "text".to_string(),
98            Value::Text {
99                text: "hello".to_string(),
100            },
101        )]);
102        let response = adapter
103            .predict(PredictRequest {
104                model_name: "echo-model".to_string(),
105                inputs: inputs.clone(),
106                ..PredictRequest::default()
107            })
108            .await
109            .expect("predict");
110        assert_eq!(response.outputs, inputs);
111        assert_eq!(response.usage, Usage::default());
112        assert_eq!(response.model.name, "echo-model");
113        assert_eq!(response.status, PredictStatus::Success);
114    }
115
116    #[test]
117    fn register_adds_echo_kind() {
118        let mut registry = Registry::new();
119        register(&mut registry).expect("register echo");
120        assert_eq!(registry.kinds(), vec![ECHO_KIND.to_string()]);
121    }
122}