1use crate::telemetry::{CompletionOperation, CompletionSpanBuilder, SpanCombinator};
6use bytes::Bytes;
7use serde::{Deserialize, Serialize};
8use serde_json::Value;
9use tracing::{Instrument, Level, enabled};
10
11use super::api::{ApiResponse, Message, ToolDefinition};
12use super::client::Client;
13use crate::OneOrMany;
14use crate::completion::{self, CompletionError, CompletionRequest, GetTokenUsage};
15use crate::http_client::HttpClientExt;
16use crate::providers::openai::responses_api::ToolChoice;
17use crate::providers::openai::responses_api::streaming::StreamingCompletionResponse;
18use crate::providers::openai::responses_api::{Output, ResponsesUsage};
19use crate::streaming::StreamingCompletionResponse as BaseStreamingCompletionResponse;
20
21pub const GROK_2_1212: &str = "grok-2-1212";
23pub const GROK_2_VISION_1212: &str = "grok-2-vision-1212";
24pub const GROK_3: &str = "grok-3";
25pub const GROK_3_FAST: &str = "grok-3-fast";
26pub const GROK_3_MINI: &str = "grok-3-mini";
27pub const GROK_3_MINI_FAST: &str = "grok-3-mini-fast";
28pub const GROK_2_IMAGE_1212: &str = "grok-2-image-1212";
29pub const GROK_4: &str = "grok-4-0709";
30
31#[derive(Debug, Serialize, Deserialize)]
36pub(super) struct XAICompletionRequest {
37 pub(super) model: String,
38 pub input: Vec<Message>,
39 #[serde(skip_serializing_if = "Option::is_none")]
40 temperature: Option<f64>,
41 #[serde(skip_serializing_if = "Option::is_none")]
42 max_output_tokens: Option<u64>,
43 #[serde(skip_serializing_if = "Vec::is_empty")]
44 tools: Vec<Value>,
45 #[serde(skip_serializing_if = "Option::is_none")]
46 tool_choice: Option<ToolChoice>,
47 #[serde(flatten, skip_serializing_if = "Option::is_none")]
48 pub additional_params: Option<serde_json::Value>,
49}
50
51impl TryFrom<(&str, CompletionRequest)> for XAICompletionRequest {
52 type Error = CompletionError;
53
54 fn try_from((model, req): (&str, CompletionRequest)) -> Result<Self, Self::Error> {
55 let chat_history = req.chat_history_with_documents();
56 if req.output_schema.is_some() {
57 tracing::warn!("Structured outputs currently not supported for xAI");
58 }
59 let model = req.model.clone().unwrap_or_else(|| model.to_string());
60 let mut input: Vec<Message> = req
61 .preamble
62 .as_ref()
63 .map_or_else(Vec::new, |p| vec![Message::system(p)]);
64
65 let mut additional_params_payload = req.additional_params.unwrap_or(Value::Null);
66
67 for msg in chat_history {
68 let msg: Vec<Message> = msg.try_into()?;
69 input.extend(msg);
70 }
71
72 let tool_choice = req.tool_choice.map(ToolChoice::try_from).transpose()?;
73 let mut additional_tools =
74 extract_tools_from_additional_params(&mut additional_params_payload)?;
75 let mut tools = req
76 .tools
77 .into_iter()
78 .map(ToolDefinition::from)
79 .map(serde_json::to_value)
80 .collect::<Result<Vec<_>, _>>()?;
81 tools.append(&mut additional_tools);
82 let additional_params = if additional_params_payload.is_null() {
83 None
84 } else {
85 Some(additional_params_payload)
86 };
87
88 Ok(Self {
89 model: model.to_string(),
90 input,
91 temperature: req.temperature,
92 max_output_tokens: req.max_tokens,
93 tools,
94 tool_choice,
95 additional_params,
96 })
97 }
98}
99
100fn extract_tools_from_additional_params(
101 additional_params: &mut Value,
102) -> Result<Vec<Value>, CompletionError> {
103 if let Some(map) = additional_params.as_object_mut()
104 && let Some(raw_tools) = map.remove("tools")
105 {
106 return serde_json::from_value::<Vec<Value>>(raw_tools).map_err(|err| {
107 CompletionError::RequestError(
108 format!("Invalid xAI `additional_params.tools` payload: {err}").into(),
109 )
110 });
111 }
112
113 Ok(Vec::new())
114}
115
116#[derive(Debug, Deserialize, Serialize)]
121pub struct CompletionResponse {
122 pub id: String,
123 pub model: String,
124 pub output: Vec<Output>,
125 #[serde(default)]
126 pub created: i64,
127 #[serde(default)]
128 pub object: String,
129 #[serde(default)]
130 pub status: Option<String>,
131 pub usage: Option<ResponsesUsage>,
132}
133
134impl TryFrom<CompletionResponse> for completion::CompletionResponse<CompletionResponse> {
135 type Error = CompletionError;
136
137 fn try_from(response: CompletionResponse) -> Result<Self, Self::Error> {
138 let content: Vec<completion::AssistantContent> = response
139 .output
140 .iter()
141 .cloned()
142 .flat_map(<Vec<completion::AssistantContent>>::from)
143 .collect();
144
145 let choice = OneOrMany::many(content).map_err(|_| {
146 CompletionError::ResponseError("Response contained no output".to_owned())
147 })?;
148
149 let usage = response
150 .usage
151 .as_ref()
152 .map(GetTokenUsage::token_usage)
153 .unwrap_or_default();
154 let message_id = response.output.iter().find_map(|item| match item {
155 Output::Message(message) => Some(message.id.clone()),
156 _ => None,
157 });
158
159 Ok(completion::CompletionResponse {
160 choice,
161 usage,
162 raw_response: response,
163 message_id,
164 })
165 }
166}
167
168#[derive(Clone)]
173pub struct CompletionModel<T = reqwest::Client> {
174 pub(crate) client: Client<T>,
175 pub model: String,
176}
177
178impl<T> CompletionModel<T> {
179 pub fn new(client: Client<T>, model: impl Into<String>) -> Self {
180 Self {
181 client,
182 model: model.into(),
183 }
184 }
185}
186
187impl<T> completion::CompletionModel for CompletionModel<T>
188where
189 T: HttpClientExt + Clone + Default + std::fmt::Debug + Send + 'static,
190{
191 type Response = CompletionResponse;
192 type StreamingResponse = StreamingCompletionResponse;
193
194 type Client = Client<T>;
195
196 fn make(client: &Self::Client, model: impl Into<String>) -> Self {
197 Self::new(client.clone(), model)
198 }
199
200 async fn completion(
201 &self,
202 completion_request: completion::CompletionRequest,
203 ) -> Result<completion::CompletionResponse<CompletionResponse>, CompletionError> {
204 let system_instructions = completion_request.preamble.clone();
205 let record_telemetry_content = completion_request.record_telemetry_content;
206 let request =
207 XAICompletionRequest::try_from((self.model.to_string().as_ref(), completion_request))?;
208 let span = CompletionSpanBuilder::new("xai", &request.model, CompletionOperation::Chat)
209 .system_instructions(system_instructions.as_deref(), record_telemetry_content)
210 .build();
211
212 if enabled!(Level::TRACE) {
213 tracing::trace!(target: "rig::completions",
214 "xAI completion request: {}",
215 serde_json::to_string_pretty(&request)?
216 );
217 }
218
219 let body = serde_json::to_vec(&request)?;
220 let req = self
221 .client
222 .post("/v1/responses")?
223 .body(body)
224 .map_err(|e| CompletionError::HttpError(e.into()))?;
225
226 async move {
227 let response = self.client.send::<_, Bytes>(req).await?;
228 let status = response.status();
229 let response_body = response.into_body().into_future().await?.to_vec();
230
231 if status.is_success() {
232 match serde_json::from_slice::<ApiResponse<CompletionResponse>>(&response_body)? {
233 ApiResponse::Ok(response) => {
234 let span = tracing::Span::current();
235 span.record("gen_ai.response.id", response.id.as_str());
236 span.record("gen_ai.response.model", response.model.as_str());
237 if let Some(usage) = &response.usage {
238 span.record_token_usage(usage);
239 }
240
241 if enabled!(Level::TRACE) {
242 tracing::trace!(target: "rig::completions",
243 "xAI completion response: {}",
244 serde_json::to_string_pretty(&response)?
245 );
246 }
247
248 response.try_into()
249 }
250 ApiResponse::Error(error) => {
251 tracing::warn!(message = %error.message(), "provider returned an error response");
252 Err(CompletionError::from_http_response(
253 status,
254 String::from_utf8_lossy(&response_body),
255 ))
256 }
257 }
258 } else {
259 Err(CompletionError::from_http_response(
260 status,
261 String::from_utf8_lossy(&response_body),
262 ))
263 }
264 }
265 .instrument(span)
266 .await
267 }
268
269 async fn stream(
270 &self,
271 request: CompletionRequest,
272 ) -> Result<BaseStreamingCompletionResponse<Self::StreamingResponse>, CompletionError> {
273 self.stream(request).await
274 }
275}
276
277#[cfg(test)]
278mod tests {
279 use super::XAICompletionRequest;
280 use crate::OneOrMany;
281 use crate::completion::request::Document;
282 use crate::completion::{CompletionRequest, CompletionRequestBuilder, Message, ToolDefinition};
283 use crate::message::ToolChoice;
284 use crate::test_utils::MockCompletionModel;
285
286 #[test]
287 fn xai_request_includes_normalized_documents() {
288 let request =
289 CompletionRequestBuilder::new(MockCompletionModel::default(), "What is glarb-glarb?")
290 .message(Message::system("Use the provided context."))
291 .document(Document {
292 id: "doc_1".to_string(),
293 text: "Definition of glarb-glarb: an ancient tool.".to_string(),
294 additional_props: Default::default(),
295 })
296 .build();
297
298 let xai_request = XAICompletionRequest::try_from(("grok-4-0709", request))
299 .expect("request conversion should succeed");
300 let serialized = serde_json::to_value(xai_request).expect("serialization should succeed");
301 let input = serialized["input"]
302 .as_array()
303 .expect("xAI request input should be an array");
304
305 assert!(
306 input
307 .iter()
308 .any(|message| message.to_string().contains("glarb-glarb")),
309 "normalized documents should be forwarded into xAI input"
310 );
311 }
312
313 #[test]
314 fn xai_direct_request_keeps_documents_after_system_messages() {
315 let request = CompletionRequest {
316 model: None,
317 preamble: None,
318 chat_history: OneOrMany::many(vec![
319 Message::system("System prompt"),
320 Message::assistant("Earlier assistant turn"),
321 Message::system("Mid-conversation instruction"),
322 Message::user("What is glarb-glarb?"),
323 ])
324 .unwrap(),
325 documents: vec![Document {
326 id: "doc_1".to_string(),
327 text: "Definition of glarb-glarb: an ancient tool.".to_string(),
328 additional_props: Default::default(),
329 }],
330 tools: vec![],
331 temperature: None,
332 max_tokens: None,
333 tool_choice: None,
334 additional_params: None,
335 output_schema: None,
336 record_telemetry_content: false,
337 };
338
339 let xai_request = XAICompletionRequest::try_from(("grok-4-0709", request))
340 .expect("request conversion should succeed");
341 let serialized = serde_json::to_value(xai_request).expect("serialization should succeed");
342 let input = serialized["input"]
343 .as_array()
344 .expect("xAI request input should be an array");
345
346 assert_eq!(input.len(), 5);
347 assert_eq!(input[0]["role"], "system");
348 assert_eq!(input[1]["role"], "user");
349 assert!(
350 input[1].to_string().contains("<file id: doc_1>"),
351 "document input should follow leading system input: {input:?}"
352 );
353 assert_eq!(input[2]["role"], "assistant");
354 assert_eq!(input[3]["role"], "system");
355 assert_eq!(input[4]["role"], "user");
356 assert_eq!(
357 input
358 .iter()
359 .filter(|message| message.to_string().contains("<file id: doc_1>"))
360 .count(),
361 1,
362 "document input should appear exactly once: {input:?}"
363 );
364 }
365
366 #[test]
367 fn xai_request_uses_responses_tool_choice_for_specific_tool() {
368 let request = CompletionRequestBuilder::new(MockCompletionModel::default(), "Use a tool.")
369 .tool(ToolDefinition {
370 name: "alpha".to_string(),
371 description: "Alpha tool".to_string(),
372 parameters: serde_json::json!({
373 "type": "object",
374 "properties": {},
375 "required": []
376 }),
377 })
378 .tool(ToolDefinition {
379 name: "beta".to_string(),
380 description: "Beta tool".to_string(),
381 parameters: serde_json::json!({
382 "type": "object",
383 "properties": {},
384 "required": []
385 }),
386 })
387 .tool_choice(ToolChoice::Specific {
388 function_names: vec!["beta".to_string()],
389 })
390 .build();
391
392 let xai_request = XAICompletionRequest::try_from(("grok-4.3", request))
393 .expect("xAI Responses API should support specific tool choice");
394 let serialized = serde_json::to_value(xai_request).expect("serialization should succeed");
395
396 assert_eq!(
397 serialized["tool_choice"],
398 serde_json::json!({"type": "function", "name": "beta"})
399 );
400 }
401
402 #[test]
403 fn xai_response_preserves_message_id_and_reasoning_token_usage() {
404 let raw: super::CompletionResponse = serde_json::from_value(serde_json::json!({
405 "id": "resp_123",
406 "model": "grok-4.3",
407 "output": [
408 {
409 "type": "reasoning",
410 "id": "rs_123",
411 "summary": [{ "type": "summary_text", "text": "thinking" }],
412 "status": "completed"
413 },
414 {
415 "type": "message",
416 "id": "msg_123",
417 "role": "assistant",
418 "status": "completed",
419 "content": [
420 { "type": "output_text", "text": "done", "annotations": [] }
421 ]
422 }
423 ],
424 "usage": {
425 "input_tokens": 10,
426 "input_tokens_details": { "cached_tokens": 3 },
427 "output_tokens": 8,
428 "output_tokens_details": { "reasoning_tokens": 5 },
429 "total_tokens": 18
430 }
431 }))
432 .expect("fixture should deserialize");
433
434 let converted = crate::completion::CompletionResponse::try_from(raw)
435 .expect("xAI response should convert");
436
437 assert_eq!(converted.message_id.as_deref(), Some("msg_123"));
438 assert_eq!(converted.usage.input_tokens, 10);
439 assert_eq!(converted.usage.cached_input_tokens, 3);
440 assert_eq!(converted.usage.output_tokens, 8);
441 assert_eq!(converted.usage.reasoning_tokens, 5);
442 }
443
444 #[tokio::test]
445 async fn completion_non_success_preserves_status_and_body() {
446 use crate::client::CompletionClient;
447 use crate::completion::{CompletionError, CompletionModel as _};
448 use crate::test_utils::RecordingHttpClient;
449
450 let body = r#"{"error":"boom","code":"503"}"#;
451 let http_client =
452 RecordingHttpClient::with_error_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
453 let client = crate::providers::xai::Client::builder()
454 .api_key("test-key")
455 .http_client(http_client)
456 .build()
457 .expect("build client");
458 let model = client.completion_model(crate::providers::xai::completion::GROK_4);
459 let request = model.completion_request("hello").build();
460
461 let error = model
462 .completion(request)
463 .await
464 .expect_err("should fail with non-success status");
465
466 assert!(matches!(error, CompletionError::HttpError(_)));
467 assert_eq!(
468 error.provider_response_status(),
469 Some(http::StatusCode::SERVICE_UNAVAILABLE)
470 );
471 assert_eq!(error.provider_response_body(), Some(body));
472 }
473
474 #[tokio::test]
475 async fn completion_2xx_error_envelope_preserves_status_and_body() {
476 use crate::client::CompletionClient;
477 use crate::completion::{CompletionError, CompletionModel as _};
478 use crate::test_utils::RecordingHttpClient;
479
480 let body = r#"{"error":"boom","code":"503"}"#;
482 let http_client = RecordingHttpClient::new(body);
483 let client = crate::providers::xai::Client::builder()
484 .api_key("test-key")
485 .http_client(http_client)
486 .build()
487 .expect("build client");
488 let model = client.completion_model(crate::providers::xai::completion::GROK_4);
489 let request = model.completion_request("hello").build();
490
491 let error = model
492 .completion(request)
493 .await
494 .expect_err("should fail with provider error envelope");
495
496 match &error {
497 CompletionError::ProviderResponse(stored) => {
498 assert_eq!(stored.body, body);
499 assert_eq!(stored.status, Some(http::StatusCode::OK));
500 }
501 other => panic!("expected ProviderResponse, got {other:?}"),
502 }
503 }
504}