1use 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#[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 #[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 #[must_use]
41 pub fn with_policy(mut self, policy: Policy) -> Self {
42 self.policy = Some(policy);
43 self
44 }
45
46 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
118pub 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}