Skip to main content

llm_kernel/llm/
middleware.rs

1//! Middleware hooks for [`LLMClient`] request/response lifecycle.
2//!
3//! Provides an [`LLMClientMiddleware`] trait with hooks that fire before
4//! each request, after each successful response, and on each error.
5//!
6//! The [`MiddlewareClient`] wrapper composes with any [`LLMClient`],
7//! including [`RetryClient`](crate::llm::retry::RetryClient).
8//!
9//! # Example
10//!
11//! ```ignore
12//! use llm_kernel::llm::{LLMClient, MiddlewareClient, NoopMiddleware};
13//!
14//! let client = OpenAIClient::from_key("gpt-4o", "sk-...")?;
15//! let middleware = NoopMiddleware;
16//! let wrapped = MiddlewareClient::new(client, middleware);
17//! let response = wrapped.complete(request).await?;
18//! ```
19
20use 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/// Hook trait for intercepting [`LLMClient`] request/response cycles.
29///
30/// All methods have default no-op implementations. Override only the hooks
31/// you need. Each hook receives an immutable reference, so middleware cannot
32/// mutate requests or responses — only observe.
33///
34/// The trait is object-safe: `Box<dyn LLMClientMiddleware>` is usable.
35#[async_trait]
36pub trait LLMClientMiddleware: Send + Sync {
37    /// Called before each `complete` request is sent.
38    ///
39    /// Use for logging, metrics, or request tracing.
40    async fn on_request(&self, _request: &LLMRequest) {}
41
42    /// Called after a successful `complete` response.
43    ///
44    /// `elapsed` is the wall time spent in the inner client (including
45    /// retries when the middleware wraps a `RetryClient`). Use for
46    /// logging, metrics, or response tracing.
47    async fn on_response(
48        &self,
49        _request: &LLMRequest,
50        _response: &LLMResponse,
51        _elapsed: Duration,
52    ) {
53    }
54
55    /// Called when `complete` returns an error.
56    ///
57    /// `elapsed` is the wall time spent in the inner client before the
58    /// error surfaced. Use for error logging, alerting, or metrics.
59    async fn on_error(&self, _request: &LLMRequest, _error: &KernelError, _elapsed: Duration) {}
60}
61
62/// A default no-op middleware. Useful as a type parameter default.
63pub struct NoopMiddleware;
64
65#[async_trait]
66impl LLMClientMiddleware for NoopMiddleware {
67    // All methods use default no-op implementations.
68}
69
70/// An [`LLMClient`] wrapper that calls [`LLMClientMiddleware`] hooks around
71/// each request.
72///
73/// Composable: `MiddlewareClient<RetryClient<OpenAIClient>, LogMiddleware>`.
74pub struct MiddlewareClient<C, M> {
75    inner: C,
76    middleware: M,
77}
78
79impl<C, M> MiddlewareClient<C, M> {
80    /// Create a new middleware-wrapped client.
81    pub fn new(inner: C, middleware: M) -> Self {
82        Self { inner, middleware }
83    }
84
85    /// Access the underlying client.
86    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        // Streaming does not invoke middleware hooks — the stream is opaque
118        // and errors arrive asynchronously in the stream events.
119        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    /// A middleware that records hook invocations.
129    #[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    /// A mock client that returns a fixed response or error.
161    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        // Timing is delivered, not just set.
212        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}