1use 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
14pub const ECHO_KIND: &str = "echo";
16
17#[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
62pub 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}