Skip to main content

openkind_engine/
dispatch.rs

1//! Request dispatching, metrics tracking, and token estimation.
2
3use openkind_core::{Answer, ResponseContract, SystemRequest, SystemResponse};
4use tracing::instrument;
5
6use crate::error::{EngineError, EngineResult};
7use crate::registry::EngineRegistry;
8
9/// Validate + dispatch. The HTTP and gRPC layers both call this — it
10/// contains the cross-cutting logic (validation, telemetry, routing).
11#[instrument(skip(req, registry), fields(model = %req.model, n_questions = req.questions.len()))]
12pub async fn dispatch(
13    req: SystemRequest,
14    registry: &EngineRegistry,
15) -> EngineResult<SystemResponse> {
16    metrics::counter!("openkind_requests_total").increment(1);
17
18    let mut observation = DispatchObservation {
19        start: std::time::Instant::now(),
20        outcome: "cancelled",
21    };
22    let result = dispatch_inner(req, registry).await;
23    observation.outcome = match &result {
24        Ok(_) => "success",
25        Err(EngineError::Invalid(_)) => "invalid_request",
26        Err(EngineError::UnknownModel(_)) => "unknown_model",
27        Err(EngineError::Unsupported { .. }) => "unsupported",
28        Err(EngineError::Overloaded { .. }) => "overloaded",
29        Err(EngineError::DeadlineExceeded { .. }) => "deadline_exceeded",
30        Err(EngineError::BackendValidation { .. }) => "backend_validation",
31        Err(EngineError::Backend { .. }) => "backend_error",
32    };
33    result
34}
35
36struct DispatchObservation {
37    start: std::time::Instant,
38    outcome: &'static str,
39}
40
41impl Drop for DispatchObservation {
42    fn drop(&mut self) {
43        // Drop also observes callers abandoning an in-flight dispatch.
44        metrics::histogram!("openkind_request_duration_ms")
45            .record(duration_millis(self.start.elapsed()));
46        metrics::counter!("openkind_request_outcomes_total", "outcome" => self.outcome)
47            .increment(1);
48    }
49}
50
51async fn dispatch_inner(
52    req: SystemRequest,
53    registry: &EngineRegistry,
54) -> EngineResult<SystemResponse> {
55    let requested_model = req.model.clone();
56    let engine = registry
57        .get(&req.model)
58        .ok_or_else(|| EngineError::UnknownModel(req.model.clone()))?;
59
60    let contract = ResponseContract::from_request(&req)?;
61    let input_tokens = engine.estimate_input_tokens(&req);
62
63    let mut resp = engine.evaluate(req).await?;
64
65    if resp.model != requested_model {
66        return Err(EngineError::Backend {
67            backend: engine.backend_id().to_string(),
68            message: format!(
69                "backend returned model `{}`, expected registered alias `{requested_model}`",
70                resp.model
71            ),
72        });
73    }
74
75    // Never forward a contract-violating engine response to the client:
76    // a bad answer shape is a backend fault, so it maps to BackendValidation (500),
77    // not to a client-facing 422.
78    if let Err(validation) = contract.validate(&resp) {
79        return Err(EngineError::BackendValidation {
80            backend: engine.backend_id().to_string(),
81            source: validation,
82        });
83    }
84
85    // If the backend didn't fill in usage, do it from the estimator.
86    // Real backends will fill it precisely.
87    if resp.usage.input_tokens == 0 {
88        resp.usage.input_tokens = input_tokens;
89    }
90    if resp.usage.output_tokens == 0 {
91        resp.usage.output_tokens = estimate_output_tokens(&resp);
92    }
93
94    metrics::counter!("openkind_responses_total").increment(1);
95
96    Ok(resp)
97}
98
99fn estimate_output_tokens(resp: &SystemResponse) -> u32 {
100    // Noul = 1 token. Choice = 1 (just the picked label).
101    // Score = ~ level descriptions worth of tokens.
102    let sum: usize = resp
103        .answers
104        .values()
105        .map(|a| match a {
106            Answer::Noul(_) => 1,
107            Answer::Choice(_) => 1,
108            Answer::Score(_) => 4,
109        })
110        .sum();
111    u32::try_from(sum).unwrap_or(u32::MAX)
112}
113
114fn duration_millis(duration: std::time::Duration) -> f64 {
115    duration.as_secs_f64() * 1000.0
116}
117
118#[cfg(test)]
119mod telemetry_tests {
120    use super::*;
121    use crate::{DecisionEngine, MockEngine};
122    use async_trait::async_trait;
123    use metrics_util::debugging::{DebugValue, DebuggingRecorder};
124    use openkind_core::{NoulQuestion, Question, State};
125    use std::{
126        future::Future,
127        sync::Arc,
128        task::{Context, Poll, Waker},
129        time::Duration,
130    };
131
132    struct OutcomeEngine(&'static str);
133    #[async_trait]
134    impl DecisionEngine for OutcomeEngine {
135        fn backend_id(&self) -> &str {
136            "outcome-test"
137        }
138        async fn evaluate(&self, req: SystemRequest) -> EngineResult<SystemResponse> {
139            match self.0 {
140                "unsupported" => Err(EngineError::Unsupported {
141                    backend: "test".into(),
142                    message: "unsupported".into(),
143                }),
144                "overloaded" => Err(EngineError::Overloaded {
145                    backend: "test".into(),
146                    retry_after_ms: 1,
147                }),
148                "deadline_exceeded" => Err(EngineError::DeadlineExceeded {
149                    backend: "test".into(),
150                    timeout_ms: 1,
151                }),
152                "backend_error" => Err(EngineError::Backend {
153                    backend: "test".into(),
154                    message: "error".into(),
155                }),
156                "backend_validation" => {
157                    let mut response = MockEngine::new().evaluate(req).await?;
158                    response.answers.clear();
159                    Ok(response)
160                }
161                "cancelled" => std::future::pending().await,
162                _ => unreachable!(),
163            }
164        }
165    }
166
167    fn request() -> SystemRequest {
168        SystemRequest {
169            model: "mock".into(),
170            state: State::Text("private text".into()),
171            questions: [(
172                "q".into(),
173                Question::Noul(NoulQuestion {
174                    instructions: serde_json::json!("?"),
175                    criteria: None,
176                }),
177            )]
178            .into_iter()
179            .collect(),
180        }
181    }
182
183    #[test]
184    fn duration_keeps_sub_millisecond_precision() {
185        assert_eq!(duration_millis(Duration::from_micros(125)), 0.125);
186    }
187
188    #[tokio::test(flavor = "current_thread")]
189    async fn every_dispatch_observes_duration_and_fixed_outcome_including_cancellation() {
190        let recorder = DebuggingRecorder::new();
191        let _guard = metrics::set_default_local_recorder(&recorder);
192        let mut registry = EngineRegistry::new();
193        registry.register("mock", Arc::new(MockEngine::new()));
194        dispatch(request(), &registry).await.unwrap();
195        let mut invalid = request();
196        invalid.questions.clear();
197        assert!(dispatch(invalid, &registry).await.is_err());
198        assert!(dispatch(request(), &EngineRegistry::new()).await.is_err());
199        for outcome in [
200            "unsupported",
201            "overloaded",
202            "deadline_exceeded",
203            "backend_error",
204            "backend_validation",
205        ] {
206            registry.register("mock", Arc::new(OutcomeEngine(outcome)));
207            assert!(dispatch(request(), &registry).await.is_err());
208        }
209        registry.register("mock", Arc::new(OutcomeEngine("cancelled")));
210        let mut future = Box::pin(dispatch(request(), &registry));
211        assert!(matches!(
212            future
213                .as_mut()
214                .poll(&mut Context::from_waker(Waker::noop())),
215            Poll::Pending
216        ));
217        drop(future);
218
219        let snapshot = recorder.snapshotter().snapshot().into_vec();
220        let mut outcomes = std::collections::BTreeSet::new();
221        for (key, _, _, value) in snapshot {
222            match key.key().name() {
223                "openkind_requests_total" => assert_eq!(value, DebugValue::Counter(9)),
224                "openkind_responses_total" => assert_eq!(value, DebugValue::Counter(1)),
225                "openkind_request_duration_ms" => {
226                    let DebugValue::Histogram(values) = value else {
227                        panic!("expected histogram")
228                    };
229                    assert_eq!(values.len(), 9);
230                    assert!(values.iter().all(|value| value.0 > 0.0));
231                    assert!(values.iter().any(|value| value.0.fract() > 0.0));
232                }
233                "openkind_request_outcomes_total" => {
234                    assert_eq!(value, DebugValue::Counter(1));
235                    let labels: Vec<_> = key.key().labels().collect();
236                    assert_eq!(labels.len(), 1);
237                    assert_eq!(labels[0].key(), "outcome");
238                    outcomes.insert(labels[0].value().to_owned());
239                }
240                _ => panic!("unexpected metric"),
241            }
242        }
243        assert_eq!(
244            outcomes,
245            [
246                "success",
247                "invalid_request",
248                "unknown_model",
249                "unsupported",
250                "overloaded",
251                "deadline_exceeded",
252                "backend_error",
253                "backend_validation",
254                "cancelled"
255            ]
256            .into_iter()
257            .map(str::to_owned)
258            .collect()
259        );
260    }
261}