Skip to main content

rskit_llm_common/
runner.rs

1//! Shared non-streaming chat adapter execution mechanics.
2
3use std::future::Future;
4use std::sync::Arc;
5use std::sync::atomic::{AtomicU64, Ordering};
6use std::time::{SystemTime, UNIX_EPOCH};
7
8use rskit_ai::semconv;
9use rskit_errors::{AppError, AppResult, ErrorCode};
10use rskit_httpclient::{HttpClient, Request};
11use rskit_llm::types::{CompletionRequest, CompletionResponse};
12use rskit_observability::{
13    record_current_span_attribute, record_span_attribute, set_span_attribute,
14};
15use rskit_resilience::Policy;
16use tracing::Instrument;
17
18/// Shared runner for provider adapters that differ only by wire dialect.
19#[derive(Clone)]
20pub struct ChatRunner {
21    system: &'static str,
22    default_model: String,
23    policy: Option<Policy>,
24    last_call_at: Arc<AtomicU64>,
25}
26
27impl ChatRunner {
28    /// Create a runner with the provider system name and default model.
29    #[must_use]
30    pub fn new(system: &'static str, default_model: impl Into<String>) -> Self {
31        Self {
32            system,
33            default_model: default_model.into(),
34            policy: None,
35            last_call_at: Arc::new(AtomicU64::new(0)),
36        }
37    }
38
39    /// Inject a resilience policy for outbound completion requests.
40    #[must_use]
41    pub fn with_policy(mut self, policy: Policy) -> Self {
42        self.policy = Some(policy);
43        self
44    }
45
46    /// Complete a request using provider-specific wire conversion.
47    pub async fn complete<F, Fut>(
48        &self,
49        mut req: CompletionRequest,
50        complete_once: F,
51    ) -> AppResult<CompletionResponse>
52    where
53        F: Fn(CompletionRequest) -> Fut + Send + Sync,
54        Fut: Future<Output = AppResult<CompletionResponse>> + Send,
55    {
56        if req.model.is_empty() {
57            req.model.clone_from(&self.default_model);
58        }
59
60        let span = tracing::info_span!(
61            "llm.complete",
62            "gen_ai.system" = self.system,
63            "gen_ai.operation.name" = semconv::Operation::Chat.as_str(),
64            "gen_ai.request.model" = req.model.as_str(),
65            "gen_ai.request.max_tokens" = tracing::field::Empty,
66            "gen_ai.request.temperature" = tracing::field::Empty,
67            "gen_ai.usage.input_tokens" = tracing::field::Empty,
68            "gen_ai.usage.output_tokens" = tracing::field::Empty,
69            "gen_ai.response.model" = tracing::field::Empty,
70            "gen_ai.response.finish_reason" = tracing::field::Empty,
71        );
72        set_span_attribute(&span, semconv::SYSTEM, self.system);
73        set_span_attribute(
74            &span,
75            semconv::OPERATION_NAME,
76            semconv::Operation::Chat.as_str(),
77        );
78        set_span_attribute(&span, semconv::REQUEST_MODEL, req.model.clone());
79        if let Some(max) = req.max_tokens {
80            record_span_attribute(&span, semconv::REQUEST_MAX_TOKENS, i64::from(max));
81        }
82        if let Some(temp) = req.temperature {
83            record_span_attribute(&span, semconv::REQUEST_TEMPERATURE, f64::from(temp));
84        }
85
86        let policy = self.policy.clone();
87        async {
88            let response = if let Some(policy) = policy {
89                let req = req.clone();
90                policy
91                    .execute(|| {
92                        let req = req.clone();
93                        complete_once(req)
94                    })
95                    .await?
96            } else {
97                complete_once(req).await?
98            };
99
100            self.record_call();
101            annotate_response(&response);
102            Ok(response)
103        }
104        .instrument(span)
105        .await
106    }
107
108    fn record_call(&self) {
109        let now_ms = SystemTime::now()
110            .duration_since(UNIX_EPOCH)
111            .map_or(0, |duration| {
112                u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
113            });
114        self.last_call_at.store(now_ms, Ordering::Relaxed);
115    }
116}
117
118/// Send a provider request, map provider errors, and return the response body text.
119pub async fn send_text(
120    client: &HttpClient,
121    request: Request,
122    provider: &'static str,
123    parse_error: impl FnOnce(u16, &str) -> AppError,
124) -> AppResult<String> {
125    let response = client.send(request).await?;
126
127    if !response.is_success() {
128        let status = response.status_u16();
129        let text = response.text_or_diagnostic();
130        return Err(parse_error(status, &text));
131    }
132
133    response.text().map_err(|error| {
134        AppError::new(
135            ErrorCode::ExternalService,
136            format!("failed to read {provider} response: {error}"),
137        )
138    })
139}
140
141fn annotate_response(response: &CompletionResponse) {
142    record_current_span_attribute(
143        semconv::USAGE_INPUT_TOKENS,
144        i64::try_from(response.usage.input_tokens).unwrap_or(i64::MAX),
145    );
146    record_current_span_attribute(
147        semconv::USAGE_OUTPUT_TOKENS,
148        i64::try_from(response.usage.output_tokens).unwrap_or(i64::MAX),
149    );
150    record_current_span_attribute(semconv::RESPONSE_MODEL, response.model.clone());
151    if let Some(reason) = response.stop_reason.as_ref() {
152        record_current_span_attribute(semconv::RESPONSE_FINISH_REASON, format!("{reason:?}"));
153    }
154}
155
156#[cfg(test)]
157mod tests {
158    use super::*;
159    use rskit_ai::{ContentPart, FinishReason, Usage};
160    use rskit_llm::types::{AssistantMessage, CompletionRequest, Message};
161
162    fn request(model: &str) -> CompletionRequest {
163        CompletionRequest {
164            model: model.to_owned(),
165            messages: vec![rskit_llm::types::user("hello")],
166            max_tokens: Some(8),
167            temperature: Some(0.1),
168            stream: false,
169            tools: None,
170            tool_choice: None,
171        }
172    }
173
174    fn response(model: &str) -> CompletionResponse {
175        CompletionResponse {
176            message: AssistantMessage {
177                content: vec![ContentPart::Text {
178                    text: "ok".to_owned(),
179                }],
180                tool_calls: Vec::new(),
181                usage: None,
182            },
183            model: model.to_owned(),
184            usage: Usage {
185                input_tokens: 1,
186                output_tokens: 2,
187                cached_tokens: 3,
188                reasoning_tokens: 4,
189            },
190            stop_reason: Some(FinishReason::Stop),
191        }
192    }
193
194    #[tokio::test]
195    async fn complete_fills_default_model_and_records_success() {
196        let runner = ChatRunner::new("test", "default-model");
197
198        let completed = runner
199            .complete(request(""), |req| async move {
200                assert_eq!(req.model, "default-model");
201                assert!(matches!(req.messages.first(), Some(Message::User(_))));
202                Ok(response(&req.model))
203            })
204            .await
205            .unwrap();
206
207        assert_eq!(completed.model, "default-model");
208        assert_eq!(completed.usage.output_tokens, 2);
209    }
210
211    #[tokio::test]
212    async fn complete_with_policy_runs_closure() {
213        let runner = ChatRunner::new("test", "fallback").with_policy(Policy::new());
214
215        let completed = runner
216            .complete(request("explicit"), |req| async move {
217                Ok(response(&req.model))
218            })
219            .await
220            .unwrap();
221
222        assert_eq!(completed.model, "explicit");
223    }
224
225    #[tokio::test]
226    async fn complete_propagates_adapter_error() {
227        let runner = ChatRunner::new("test", "fallback");
228
229        let err = runner
230            .complete(request("explicit"), |_req| async move {
231                Err(AppError::new(ErrorCode::ExternalService, "provider failed"))
232            })
233            .await
234            .unwrap_err();
235
236        assert_eq!(err.code(), ErrorCode::ExternalService);
237    }
238}