Skip to main content

rig_core/providers/xai/
completion.rs

1//! xAI Completion Integration
2//!
3//! Uses the xAI Responses API: <https://docs.x.ai/docs/guides/chat>
4
5use 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
21/// xAI completion models as of 2025-06-04
22pub 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// ================================================================
32// Request Types
33// ================================================================
34
35#[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// ================================================================
117// Response Types
118// ================================================================
119
120#[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// ================================================================
169// Completion Model
170// ================================================================
171
172#[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        // Deserializes to `ApiResponse::Error(ApiError { error, code })` on a 200 OK.
481        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}