Skip to repository content487 lines · 15.4 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T04:47:08.698Z Public web read
NIP-34 coordinate
30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omegaMaintainersHidden in public view
References2 branches · 1 tag
Read-only clone
git clone https://openagents.com/git/tenant.openagents/omega.gitBrowse files
x_ai.rs
1use anyhow::Result;
2use collections::BTreeMap;
3use credentials_provider::CredentialsProvider;
4use futures::{FutureExt, StreamExt, future::BoxFuture};
5use gpui::{App, AppContext, AsyncApp, Context, Entity, SharedString, Task};
6use http_client::{CustomHeaders, HttpClient};
7use language_model::{
8 ApiKeyConfiguration, ApiKeyState, AuthenticateError, EnvVar, IconOrSvg, LanguageModel,
9 LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelEffortLevel,
10 LanguageModelId, LanguageModelName, LanguageModelProvider, LanguageModelProviderId,
11 LanguageModelProviderName, LanguageModelProviderState, LanguageModelRequest,
12 LanguageModelToolChoice, LanguageModelToolSchemaFormat, ProviderSettingsView, RateLimiter,
13 env_var,
14};
15use open_ai::ResponseStreamEvent;
16pub use settings::XaiAvailableModel as AvailableModel;
17use settings::{Settings, SettingsStore};
18use std::sync::{Arc, LazyLock};
19use strum::IntoEnumIterator;
20use ui::IconName;
21use x_ai::XAI_API_URL;
22
23const PROVIDER_ID: LanguageModelProviderId = LanguageModelProviderId::new("x_ai");
24const PROVIDER_NAME: LanguageModelProviderName = LanguageModelProviderName::new("xAI");
25
26const API_KEY_ENV_VAR_NAME: &str = "XAI_API_KEY";
27static API_KEY_ENV_VAR: LazyLock<EnvVar> = env_var!(API_KEY_ENV_VAR_NAME);
28
29#[derive(Default, Clone, Debug, PartialEq)]
30pub struct XAiSettings {
31 pub api_url: String,
32 pub available_models: Vec<AvailableModel>,
33 pub custom_headers: CustomHeaders,
34}
35
36pub struct XAiLanguageModelProvider {
37 http_client: Arc<dyn HttpClient>,
38 state: Entity<State>,
39}
40
41pub struct State {
42 api_key_state: ApiKeyState,
43 credentials_provider: Arc<dyn CredentialsProvider>,
44}
45
46impl State {
47 fn is_authenticated(&self) -> bool {
48 self.api_key_state.has_key()
49 }
50
51 fn set_api_key(&mut self, api_key: Option<String>, cx: &mut Context<Self>) -> Task<Result<()>> {
52 let credentials_provider = self.credentials_provider.clone();
53 let api_url = XAiLanguageModelProvider::api_url(cx);
54 self.api_key_state.store(
55 api_url,
56 api_key,
57 |this| &mut this.api_key_state,
58 credentials_provider,
59 cx,
60 )
61 }
62
63 fn authenticate(&mut self, cx: &mut Context<Self>) -> Task<Result<(), AuthenticateError>> {
64 let credentials_provider = self.credentials_provider.clone();
65 let api_url = XAiLanguageModelProvider::api_url(cx);
66 self.api_key_state.load_if_needed(
67 api_url,
68 |this| &mut this.api_key_state,
69 credentials_provider,
70 cx,
71 )
72 }
73}
74
75impl XAiLanguageModelProvider {
76 pub fn new(
77 http_client: Arc<dyn HttpClient>,
78 credentials_provider: Arc<dyn CredentialsProvider>,
79 cx: &mut App,
80 ) -> Self {
81 let state = cx.new(|cx| {
82 cx.observe_global::<SettingsStore>(|this: &mut State, cx| {
83 let credentials_provider = this.credentials_provider.clone();
84 let api_url = Self::api_url(cx);
85 this.api_key_state.handle_url_change(
86 api_url,
87 |this| &mut this.api_key_state,
88 credentials_provider,
89 cx,
90 );
91 cx.notify();
92 })
93 .detach();
94 State {
95 api_key_state: ApiKeyState::new(Self::api_url(cx), (*API_KEY_ENV_VAR).clone()),
96 credentials_provider,
97 }
98 });
99
100 Self { http_client, state }
101 }
102
103 fn create_language_model(&self, model: x_ai::Model) -> Arc<dyn LanguageModel> {
104 Arc::new(XAiLanguageModel {
105 id: LanguageModelId::from(model.id().to_string()),
106 model,
107 state: self.state.clone(),
108 http_client: self.http_client.clone(),
109 request_limiter: RateLimiter::new(4),
110 })
111 }
112
113 fn settings(cx: &App) -> &XAiSettings {
114 &crate::AllLanguageModelSettings::get_global(cx).x_ai
115 }
116
117 fn api_url(cx: &App) -> SharedString {
118 let api_url = &Self::settings(cx).api_url;
119 if api_url.is_empty() {
120 XAI_API_URL.into()
121 } else {
122 SharedString::new(api_url.as_str())
123 }
124 }
125}
126
127impl LanguageModelProviderState for XAiLanguageModelProvider {
128 type ObservableEntity = State;
129
130 fn observable_entity(&self) -> Option<Entity<Self::ObservableEntity>> {
131 Some(self.state.clone())
132 }
133}
134
135impl LanguageModelProvider for XAiLanguageModelProvider {
136 fn id(&self) -> LanguageModelProviderId {
137 PROVIDER_ID
138 }
139
140 fn name(&self) -> LanguageModelProviderName {
141 PROVIDER_NAME
142 }
143
144 fn icon(&self) -> IconOrSvg {
145 IconOrSvg::Icon(IconName::AiXAi)
146 }
147
148 fn default_model(&self, _cx: &App) -> Option<Arc<dyn LanguageModel>> {
149 Some(self.create_language_model(x_ai::Model::default()))
150 }
151
152 fn default_fast_model(&self, _cx: &App) -> Option<Arc<dyn LanguageModel>> {
153 Some(self.create_language_model(x_ai::Model::default_fast()))
154 }
155
156 fn provided_models(&self, cx: &App) -> Vec<Arc<dyn LanguageModel>> {
157 let mut models = BTreeMap::default();
158
159 for model in x_ai::Model::iter() {
160 if !matches!(model, x_ai::Model::Custom { .. }) {
161 models.insert(model.id().to_string(), model);
162 }
163 }
164
165 for model in &Self::settings(cx).available_models {
166 models.insert(
167 model.name.clone(),
168 x_ai::Model::Custom {
169 name: model.name.clone(),
170 display_name: model.display_name.clone(),
171 max_tokens: model.max_tokens,
172 max_output_tokens: model.max_output_tokens,
173 max_completion_tokens: model.max_completion_tokens,
174 supports_images: model.supports_images,
175 supports_tools: model.supports_tools,
176 parallel_tool_calls: model.parallel_tool_calls,
177 },
178 );
179 }
180
181 models
182 .into_values()
183 .map(|model| self.create_language_model(model))
184 .collect()
185 }
186
187 fn is_authenticated(&self, cx: &App) -> bool {
188 self.state.read(cx).is_authenticated()
189 }
190
191 fn authenticate(&self, cx: &mut App) -> Task<Result<(), AuthenticateError>> {
192 self.state.update(cx, |state, cx| state.authenticate(cx))
193 }
194
195 fn settings_view(&self, cx: &mut App) -> Option<ProviderSettingsView> {
196 let state = self.state.read(cx);
197 Some(ProviderSettingsView::ApiKey(ApiKeyConfiguration::new(
198 state.api_key_state.has_key(),
199 state.api_key_state.is_from_env_var(),
200 state.api_key_state.env_var_name().clone(),
201 "https://console.x.ai/team/default/api-keys".into(),
202 )))
203 }
204
205 fn set_api_key(&self, api_key: Option<String>, cx: &mut App) -> Task<Result<()>> {
206 self.state
207 .update(cx, |state, cx| state.set_api_key(api_key, cx))
208 }
209}
210
211pub struct XAiLanguageModel {
212 id: LanguageModelId,
213 model: x_ai::Model,
214 state: Entity<State>,
215 http_client: Arc<dyn HttpClient>,
216 request_limiter: RateLimiter,
217}
218
219impl XAiLanguageModel {
220 fn stream_completion(
221 &self,
222 request: open_ai::Request,
223 cx: &AsyncApp,
224 ) -> BoxFuture<
225 'static,
226 Result<
227 futures::stream::BoxStream<'static, Result<ResponseStreamEvent>>,
228 LanguageModelCompletionError,
229 >,
230 > {
231 let http_client = self.http_client.clone();
232
233 let (api_key, api_url, extra_headers) = self.state.read_with(cx, |state, cx| {
234 let api_url = XAiLanguageModelProvider::api_url(cx);
235 let extra_headers = XAiLanguageModelProvider::settings(cx)
236 .custom_headers
237 .clone();
238 (state.api_key_state.key(&api_url), api_url, extra_headers)
239 });
240
241 let future = self.request_limiter.stream(async move {
242 let provider = PROVIDER_NAME;
243 let Some(api_key) = api_key else {
244 return Err(LanguageModelCompletionError::NoApiKey { provider });
245 };
246 let request = open_ai::stream_completion(
247 http_client.as_ref(),
248 provider.0.as_str(),
249 &api_url,
250 &api_key,
251 request,
252 &extra_headers,
253 );
254 let response = request.await?;
255 Ok(response)
256 });
257
258 async move { Ok(future.await?.boxed()) }.boxed()
259 }
260}
261
262fn x_ai_reasoning_efforts(model: &x_ai::Model) -> &'static [open_ai::ReasoningEffort] {
263 if model.supports_reasoning_effort() {
264 &[
265 open_ai::ReasoningEffort::None,
266 open_ai::ReasoningEffort::Low,
267 open_ai::ReasoningEffort::Medium,
268 open_ai::ReasoningEffort::High,
269 ]
270 } else {
271 &[]
272 }
273}
274
275fn default_thinking_reasoning_effort(model: &x_ai::Model) -> Option<open_ai::ReasoningEffort> {
276 if model.supports_reasoning_effort() {
277 Some(open_ai::ReasoningEffort::Low)
278 } else {
279 None
280 }
281}
282
283fn reasoning_effort_for_request(
284 request: &LanguageModelRequest,
285 model: &x_ai::Model,
286) -> Option<open_ai::ReasoningEffort> {
287 let supported_efforts = x_ai_reasoning_efforts(model);
288 if supported_efforts.is_empty() {
289 return None;
290 }
291
292 if request.thinking_allowed {
293 request
294 .thinking_effort
295 .as_deref()
296 .and_then(|effort| effort.parse::<open_ai::ReasoningEffort>().ok())
297 .filter(|effort| supported_efforts.contains(effort))
298 .filter(|effort| *effort != open_ai::ReasoningEffort::None)
299 .or_else(|| default_thinking_reasoning_effort(model))
300 } else if supported_efforts.contains(&open_ai::ReasoningEffort::None) {
301 Some(open_ai::ReasoningEffort::None)
302 } else {
303 None
304 }
305}
306
307fn supported_thinking_effort_levels(model: &x_ai::Model) -> Vec<LanguageModelEffortLevel> {
308 let default_effort = default_thinking_reasoning_effort(model);
309 x_ai_reasoning_efforts(model)
310 .iter()
311 .copied()
312 .filter_map(|effort| {
313 let (name, value) = match effort {
314 open_ai::ReasoningEffort::None => return None,
315 open_ai::ReasoningEffort::Minimal => ("Minimal", "minimal"),
316 open_ai::ReasoningEffort::Low => ("Low", "low"),
317 open_ai::ReasoningEffort::Medium => ("Medium", "medium"),
318 open_ai::ReasoningEffort::High => ("High", "high"),
319 open_ai::ReasoningEffort::XHigh => ("Extra High", "xhigh"),
320 open_ai::ReasoningEffort::Max => return None, // Not supported by any xAI models
321 };
322
323 Some(LanguageModelEffortLevel {
324 name: name.into(),
325 value: value.into(),
326 is_default: Some(effort) == default_effort,
327 })
328 })
329 .collect()
330}
331
332impl LanguageModel for XAiLanguageModel {
333 fn id(&self) -> LanguageModelId {
334 self.id.clone()
335 }
336
337 fn name(&self) -> LanguageModelName {
338 LanguageModelName::from(self.model.display_name().to_string())
339 }
340
341 fn provider_id(&self) -> LanguageModelProviderId {
342 PROVIDER_ID
343 }
344
345 fn provider_name(&self) -> LanguageModelProviderName {
346 PROVIDER_NAME
347 }
348
349 fn supports_tools(&self) -> bool {
350 self.model.supports_tool()
351 }
352
353 fn supports_images(&self) -> bool {
354 self.model.supports_images()
355 }
356
357 fn supports_streaming_tools(&self) -> bool {
358 true
359 }
360
361 fn supports_tool_choice(&self, choice: LanguageModelToolChoice) -> bool {
362 match choice {
363 LanguageModelToolChoice::Auto
364 | LanguageModelToolChoice::Any
365 | LanguageModelToolChoice::None => true,
366 }
367 }
368
369 fn supports_thinking(&self) -> bool {
370 self.model.supports_reasoning_effort()
371 }
372
373 fn supported_effort_levels(&self) -> Vec<LanguageModelEffortLevel> {
374 supported_thinking_effort_levels(&self.model)
375 }
376
377 fn tool_input_format(&self) -> LanguageModelToolSchemaFormat {
378 if self.model.requires_json_schema_subset() {
379 LanguageModelToolSchemaFormat::JsonSchemaSubset
380 } else {
381 LanguageModelToolSchemaFormat::JsonSchema
382 }
383 }
384
385 fn telemetry_id(&self) -> String {
386 format!("x_ai/{}", self.model.id())
387 }
388
389 fn max_token_count(&self) -> u64 {
390 self.model.max_token_count()
391 }
392
393 fn max_output_tokens(&self) -> Option<u64> {
394 self.model.max_output_tokens()
395 }
396
397 fn supports_split_token_display(&self) -> bool {
398 true
399 }
400
401 fn stream_completion(
402 &self,
403 request: LanguageModelRequest,
404 cx: &AsyncApp,
405 ) -> BoxFuture<
406 'static,
407 Result<
408 futures::stream::BoxStream<
409 'static,
410 Result<LanguageModelCompletionEvent, LanguageModelCompletionError>,
411 >,
412 LanguageModelCompletionError,
413 >,
414 > {
415 let reasoning_effort = reasoning_effort_for_request(&request, &self.model);
416 let request = match crate::provider::open_ai::into_open_ai(
417 request,
418 self.model.id(),
419 self.model.supports_parallel_tool_calls(),
420 self.model.supports_prompt_cache_key(),
421 self.max_output_tokens(),
422 crate::provider::open_ai::ChatCompletionMaxTokensParameter::MaxCompletionTokens,
423 reasoning_effort,
424 false,
425 ) {
426 Ok(request) => request,
427 Err(error) => return async move { Err(error.into()) }.boxed(),
428 };
429 let completions = self.stream_completion(request, cx);
430 async move {
431 let mapper = crate::provider::open_ai::OpenAiEventMapper::new();
432 Ok(mapper.map_stream(completions.await?).boxed())
433 }
434 .boxed()
435 }
436}
437
438#[cfg(test)]
439mod tests {
440 use super::*;
441
442 #[test]
443 fn grok_43_supports_selectable_thinking_effort_levels() {
444 let effort_levels = supported_thinking_effort_levels(&x_ai::Model::Grok43);
445 let values = effort_levels
446 .iter()
447 .map(|level| level.value.as_ref())
448 .collect::<Vec<_>>();
449
450 assert_eq!(values, ["low", "medium", "high"]);
451 assert_eq!(
452 effort_levels
453 .iter()
454 .find(|level| level.is_default)
455 .map(|level| level.value.as_ref()),
456 Some("low")
457 );
458 }
459
460 #[test]
461 fn grok_43_request_uses_selected_reasoning_effort() {
462 let request = LanguageModelRequest {
463 thinking_allowed: true,
464 thinking_effort: Some("high".to_string()),
465 ..Default::default()
466 };
467
468 assert_eq!(
469 reasoning_effort_for_request(&request, &x_ai::Model::Grok43),
470 Some(open_ai::ReasoningEffort::High)
471 );
472 }
473
474 #[test]
475 fn grok_43_request_uses_none_when_thinking_is_disabled() {
476 let request = LanguageModelRequest {
477 thinking_allowed: false,
478 ..Default::default()
479 };
480
481 assert_eq!(
482 reasoning_effort_for_request(&request, &x_ai::Model::Grok43),
483 Some(open_ai::ReasoningEffort::None)
484 );
485 }
486}
487