llm_kernel/llm/
middleware.rs1use std::time::Duration;
21
22use async_trait::async_trait;
23
24use crate::error::{KernelError, Result};
25use crate::llm::client::LLMClient;
26use crate::llm::types::{LLMRequest, LLMResponse, LLMStream};
27
28#[async_trait]
36pub trait LLMClientMiddleware: Send + Sync {
37 async fn on_request(&self, _request: &LLMRequest) {}
41
42 async fn on_response(
48 &self,
49 _request: &LLMRequest,
50 _response: &LLMResponse,
51 _elapsed: Duration,
52 ) {
53 }
54
55 async fn on_error(&self, _request: &LLMRequest, _error: &KernelError, _elapsed: Duration) {}
60}
61
62pub struct NoopMiddleware;
64
65#[async_trait]
66impl LLMClientMiddleware for NoopMiddleware {
67 }
69
70pub struct MiddlewareClient<C, M> {
75 inner: C,
76 middleware: M,
77}
78
79impl<C, M> MiddlewareClient<C, M> {
80 pub fn new(inner: C, middleware: M) -> Self {
82 Self { inner, middleware }
83 }
84
85 pub fn inner(&self) -> &C {
87 &self.inner
88 }
89}
90
91#[async_trait]
92impl<C: LLMClient, M: LLMClientMiddleware> LLMClient for MiddlewareClient<C, M> {
93 async fn complete(&self, request: LLMRequest) -> Result<LLMResponse> {
94 self.middleware.on_request(&request).await;
95 let started = std::time::Instant::now();
96 match self.inner.complete(request.clone()).await {
97 Ok(response) => {
98 self.middleware
99 .on_response(&request, &response, started.elapsed())
100 .await;
101 Ok(response)
102 }
103 Err(err) => {
104 self.middleware
105 .on_error(&request, &err, started.elapsed())
106 .await;
107 Err(err)
108 }
109 }
110 }
111
112 fn model_name(&self) -> &str {
113 self.inner.model_name()
114 }
115
116 async fn stream_complete(&self, request: LLMRequest) -> Result<LLMStream> {
117 self.inner.stream_complete(request).await
120 }
121}
122
123#[cfg(test)]
124mod tests {
125 use super::*;
126 use std::sync::{Arc, Mutex};
127
128 #[derive(Default)]
130 struct RecordingMiddleware {
131 on_request_called: Arc<Mutex<bool>>,
132 on_response_called: Arc<Mutex<bool>>,
133 on_error_called: Arc<Mutex<bool>>,
134 response_elapsed: Arc<Mutex<Option<Duration>>>,
135 error_elapsed: Arc<Mutex<Option<Duration>>>,
136 }
137
138 #[async_trait]
139 impl LLMClientMiddleware for RecordingMiddleware {
140 async fn on_request(&self, _request: &LLMRequest) {
141 *self.on_request_called.lock().unwrap() = true;
142 }
143
144 async fn on_response(
145 &self,
146 _request: &LLMRequest,
147 _response: &LLMResponse,
148 elapsed: Duration,
149 ) {
150 *self.on_response_called.lock().unwrap() = true;
151 *self.response_elapsed.lock().unwrap() = Some(elapsed);
152 }
153
154 async fn on_error(&self, _request: &LLMRequest, _error: &KernelError, elapsed: Duration) {
155 *self.on_error_called.lock().unwrap() = true;
156 *self.error_elapsed.lock().unwrap() = Some(elapsed);
157 }
158 }
159
160 struct MockClient {
162 response: std::sync::Mutex<Option<Result<LLMResponse>>>,
163 }
164
165 impl MockClient {
166 fn ok() -> Self {
167 Self {
168 response: std::sync::Mutex::new(Some(Ok(LLMResponse {
169 content: "hello".into(),
170 model: "mock".into(),
171 ..Default::default()
172 }))),
173 }
174 }
175
176 fn err() -> Self {
177 Self {
178 response: std::sync::Mutex::new(Some(Err(KernelError::LlmApi("fail".into())))),
179 }
180 }
181 }
182
183 #[async_trait]
184 impl LLMClient for MockClient {
185 async fn complete(&self, _request: LLMRequest) -> Result<LLMResponse> {
186 self.response.lock().unwrap().take().unwrap()
187 }
188
189 fn model_name(&self) -> &str {
190 "mock"
191 }
192
193 async fn stream_complete(&self, _request: LLMRequest) -> Result<LLMStream> {
194 unimplemented!()
195 }
196 }
197
198 #[tokio::test]
199 async fn middleware_calls_on_request_and_on_response_on_success() {
200 let mid = RecordingMiddleware::default();
201 let req_called = mid.on_request_called.clone();
202 let res_called = mid.on_response_called.clone();
203 let elapsed = mid.response_elapsed.clone();
204
205 let client = MiddlewareClient::new(MockClient::ok(), mid);
206 let result = client.complete(LLMRequest::builder().build()).await;
207
208 assert!(result.is_ok());
209 assert!(*req_called.lock().unwrap());
210 assert!(*res_called.lock().unwrap());
211 assert!(elapsed.lock().unwrap().is_some());
213 }
214
215 #[tokio::test]
216 async fn middleware_calls_on_error_on_failure() {
217 let mid = RecordingMiddleware::default();
218 let err_called = mid.on_error_called.clone();
219 let elapsed = mid.error_elapsed.clone();
220
221 let client = MiddlewareClient::new(MockClient::err(), mid);
222 let result = client.complete(LLMRequest::builder().build()).await;
223
224 assert!(result.is_err());
225 assert!(*err_called.lock().unwrap());
226 assert!(elapsed.lock().unwrap().is_some());
227 }
228
229 #[tokio::test]
230 async fn middleware_delegates_model_name() {
231 let client = MiddlewareClient::new(MockClient::ok(), NoopMiddleware);
232 assert_eq!(client.model_name(), "mock");
233 }
234
235 #[tokio::test]
236 async fn noop_middleware_compiles_and_works() {
237 let client = MiddlewareClient::new(MockClient::ok(), NoopMiddleware);
238 let result = client.complete(LLMRequest::builder().build()).await;
239 assert!(result.is_ok());
240 }
241}