Skip to repository content903 lines · 29.9 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T02:55:10.274Z 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
language_model_selector.rs
1use std::{cmp::Reverse, sync::Arc};
2
3use agent_settings::AgentSettings;
4use collections::{HashMap, HashSet, IndexMap};
5use fuzzy::{StringMatch, StringMatchCandidate, match_strings};
6use gpui::{
7 Action, AnyElement, App, BackgroundExecutor, DismissEvent, FocusHandle, ForegroundExecutor,
8 Subscription, Task,
9};
10use language_model::{
11 ConfiguredModel, IconOrSvg, LanguageModel, LanguageModelId, LanguageModelProvider,
12 LanguageModelProviderId, LanguageModelRegistry,
13};
14use ordered_float::OrderedFloat;
15use picker::{Picker, PickerDelegate};
16use settings::Settings;
17use ui::prelude::*;
18use zed_actions::agent::OpenSettings;
19
20use crate::ui::{ModelSelectorFooter, ModelSelectorHeader, ModelSelectorListItem};
21
22type OnModelChanged = Arc<dyn Fn(Arc<dyn LanguageModel>, &mut App) + 'static>;
23type GetActiveModel = Arc<dyn Fn(&App) -> Option<ConfiguredModel> + 'static>;
24type OnToggleFavorite = Arc<dyn Fn(Arc<dyn LanguageModel>, bool, &mut App) + 'static>;
25
26pub type LanguageModelSelector = Picker<LanguageModelPickerDelegate>;
27
28pub fn language_model_selector(
29 get_active_model: impl Fn(&App) -> Option<ConfiguredModel> + 'static,
30 on_model_changed: impl Fn(Arc<dyn LanguageModel>, &mut App) + 'static,
31 on_toggle_favorite: impl Fn(Arc<dyn LanguageModel>, bool, &mut App) + 'static,
32 popover_styles: bool,
33 focus_handle: FocusHandle,
34 window: &mut Window,
35 cx: &mut Context<LanguageModelSelector>,
36) -> LanguageModelSelector {
37 let delegate = LanguageModelPickerDelegate::new(
38 get_active_model,
39 on_model_changed,
40 on_toggle_favorite,
41 popover_styles,
42 focus_handle,
43 window,
44 cx,
45 );
46
47 if popover_styles {
48 Picker::list(delegate, window, cx)
49 .show_scrollbar(true)
50 .initial_width(rems(20.))
51 .popover()
52 } else {
53 Picker::list(delegate, window, cx)
54 .show_scrollbar(true)
55 .embedded()
56 }
57}
58
59fn all_models(cx: &App) -> GroupedModels {
60 let lm_registry = LanguageModelRegistry::global(cx).read(cx);
61 let providers = lm_registry.visible_providers();
62
63 let mut favorites_index = FavoritesIndex::default();
64
65 for sel in &AgentSettings::get_global(cx).favorite_models {
66 favorites_index
67 .entry(sel.provider.0.clone().into())
68 .or_default()
69 .insert(sel.model.clone().into());
70 }
71
72 let recommended = providers
73 .iter()
74 .flat_map(|provider| {
75 provider
76 .recommended_models(cx)
77 .into_iter()
78 .map(|model| ModelInfo::new(&**provider, model, &favorites_index))
79 })
80 .collect();
81
82 let all = providers
83 .iter()
84 .flat_map(|provider| {
85 provider
86 .provided_models(cx)
87 .into_iter()
88 .map(|model| ModelInfo::new(&**provider, model, &favorites_index))
89 })
90 .collect();
91
92 GroupedModels::new(all, recommended)
93}
94
95type FavoritesIndex = HashMap<LanguageModelProviderId, HashSet<LanguageModelId>>;
96
97#[derive(Clone)]
98struct ModelInfo {
99 model: Arc<dyn LanguageModel>,
100 icon: IconOrSvg,
101 is_favorite: bool,
102}
103
104impl ModelInfo {
105 fn new(
106 provider: &dyn LanguageModelProvider,
107 model: Arc<dyn LanguageModel>,
108 favorites_index: &FavoritesIndex,
109 ) -> Self {
110 let is_favorite = favorites_index
111 .get(&provider.id())
112 .map_or(false, |set| set.contains(&model.id()));
113
114 Self {
115 model,
116 icon: provider.icon(),
117 is_favorite,
118 }
119 }
120}
121
122pub struct LanguageModelPickerDelegate {
123 on_model_changed: OnModelChanged,
124 get_active_model: GetActiveModel,
125 on_toggle_favorite: OnToggleFavorite,
126 all_models: Arc<GroupedModels>,
127 filtered_entries: Vec<LanguageModelPickerEntry>,
128 selected_index: usize,
129 _subscriptions: Vec<Subscription>,
130 popover_styles: bool,
131 focus_handle: FocusHandle,
132}
133
134impl LanguageModelPickerDelegate {
135 fn new(
136 get_active_model: impl Fn(&App) -> Option<ConfiguredModel> + 'static,
137 on_model_changed: impl Fn(Arc<dyn LanguageModel>, &mut App) + 'static,
138 on_toggle_favorite: impl Fn(Arc<dyn LanguageModel>, bool, &mut App) + 'static,
139 popover_styles: bool,
140 focus_handle: FocusHandle,
141 window: &mut Window,
142 cx: &mut Context<Picker<Self>>,
143 ) -> Self {
144 let on_model_changed = Arc::new(on_model_changed);
145 let models = all_models(cx);
146 let entries = models.entries();
147
148 Self {
149 on_model_changed,
150 all_models: Arc::new(models),
151 selected_index: Self::get_active_model_index(&entries, get_active_model(cx)),
152 filtered_entries: entries,
153 get_active_model: Arc::new(get_active_model),
154 on_toggle_favorite: Arc::new(on_toggle_favorite),
155 _subscriptions: vec![cx.subscribe_in(
156 &LanguageModelRegistry::global(cx),
157 window,
158 |picker, _, event, window, cx| {
159 match event {
160 language_model::Event::ProviderStateChanged(_)
161 | language_model::Event::AddedProvider(_)
162 | language_model::Event::RemovedProvider(_) => {
163 let query = picker.query(cx);
164 picker.delegate.all_models = Arc::new(all_models(cx));
165 // Update matches will automatically drop the previous task
166 // if we get a provider event again
167 picker.update_matches(query, window, cx)
168 }
169 _ => {}
170 }
171 },
172 )],
173 popover_styles,
174 focus_handle,
175 }
176 }
177
178 fn get_active_model_index(
179 entries: &[LanguageModelPickerEntry],
180 active_model: Option<ConfiguredModel>,
181 ) -> usize {
182 entries
183 .iter()
184 .position(|entry| {
185 if let LanguageModelPickerEntry::Model(model) = entry {
186 active_model
187 .as_ref()
188 .map(|active_model| {
189 active_model.model.id() == model.model.id()
190 && active_model.provider.id() == model.model.provider_id()
191 })
192 .unwrap_or_default()
193 } else {
194 false
195 }
196 })
197 .unwrap_or(0)
198 }
199
200 pub fn active_model(&self, cx: &App) -> Option<ConfiguredModel> {
201 (self.get_active_model)(cx)
202 }
203
204 pub fn favorites_count(&self) -> usize {
205 self.all_models.favorites.len()
206 }
207
208 pub fn cycle_favorite_models(&mut self, window: &mut Window, cx: &mut Context<Picker<Self>>) {
209 if self.all_models.favorites.is_empty() {
210 return;
211 }
212
213 let active_model = (self.get_active_model)(cx);
214 let active_provider_id = active_model.as_ref().map(|m| m.provider.id());
215 let active_model_id = active_model.as_ref().map(|m| m.model.id());
216
217 let current_index = self
218 .all_models
219 .favorites
220 .iter()
221 .position(|info| {
222 Some(info.model.provider_id()) == active_provider_id
223 && Some(info.model.id()) == active_model_id
224 })
225 .unwrap_or(usize::MAX);
226
227 let next_index = if current_index == usize::MAX {
228 0
229 } else {
230 (current_index + 1) % self.all_models.favorites.len()
231 };
232
233 let next_model = self.all_models.favorites[next_index].model.clone();
234
235 (self.on_model_changed)(next_model, cx);
236
237 // Align the picker selection with the newly-active model
238 let new_index =
239 Self::get_active_model_index(&self.filtered_entries, (self.get_active_model)(cx));
240 self.set_selected_index(new_index, window, cx);
241 }
242}
243
244struct GroupedModels {
245 favorites: Vec<ModelInfo>,
246 recommended: Vec<ModelInfo>,
247 all: IndexMap<LanguageModelProviderId, Vec<ModelInfo>>,
248}
249
250impl GroupedModels {
251 pub fn new(all: Vec<ModelInfo>, recommended: Vec<ModelInfo>) -> Self {
252 let favorites = all
253 .iter()
254 .filter(|info| info.is_favorite)
255 .cloned()
256 .collect();
257
258 let mut all_by_provider: IndexMap<_, Vec<ModelInfo>> = IndexMap::default();
259 for model in all {
260 let provider = model.model.provider_id();
261 if let Some(models) = all_by_provider.get_mut(&provider) {
262 models.push(model);
263 } else {
264 all_by_provider.insert(provider, vec![model]);
265 }
266 }
267
268 Self {
269 favorites,
270 recommended,
271 all: all_by_provider,
272 }
273 }
274
275 fn entries(&self) -> Vec<LanguageModelPickerEntry> {
276 let mut entries = Vec::new();
277
278 if !self.favorites.is_empty() {
279 entries.push(LanguageModelPickerEntry::Separator("Favorite".into()));
280 for info in &self.favorites {
281 entries.push(LanguageModelPickerEntry::Model(info.clone()));
282 }
283 }
284
285 if !self.recommended.is_empty() {
286 entries.push(LanguageModelPickerEntry::Separator("Recommended".into()));
287 for info in &self.recommended {
288 entries.push(LanguageModelPickerEntry::Model(info.clone()));
289 }
290 }
291
292 for models in self.all.values() {
293 if models.is_empty() {
294 continue;
295 }
296 entries.push(LanguageModelPickerEntry::Separator(
297 models[0].model.provider_name().0,
298 ));
299 for info in models {
300 entries.push(LanguageModelPickerEntry::Model(info.clone()));
301 }
302 }
303
304 entries
305 }
306}
307
308enum LanguageModelPickerEntry {
309 Model(ModelInfo),
310 Separator(SharedString),
311}
312
313struct ModelMatcher {
314 models: Vec<ModelInfo>,
315 fg_executor: ForegroundExecutor,
316 bg_executor: BackgroundExecutor,
317 candidates: Vec<StringMatchCandidate>,
318}
319
320impl ModelMatcher {
321 fn new(
322 models: Vec<ModelInfo>,
323 fg_executor: ForegroundExecutor,
324 bg_executor: BackgroundExecutor,
325 ) -> ModelMatcher {
326 let candidates = Self::make_match_candidates(&models);
327 Self {
328 models,
329 fg_executor,
330 bg_executor,
331 candidates,
332 }
333 }
334
335 pub fn fuzzy_search(&self, query: &str) -> Vec<ModelInfo> {
336 let mut matches = self.fg_executor.block_on(match_strings(
337 &self.candidates,
338 query,
339 false,
340 true,
341 100,
342 &Default::default(),
343 self.bg_executor.clone(),
344 ));
345
346 let sorting_key = |mat: &StringMatch| {
347 let candidate = &self.candidates[mat.candidate_id];
348 (Reverse(OrderedFloat(mat.score)), candidate.id)
349 };
350 matches.sort_unstable_by_key(sorting_key);
351
352 let matched_models: Vec<_> = matches
353 .into_iter()
354 .map(|mat| self.models[mat.candidate_id].clone())
355 .collect();
356
357 matched_models
358 }
359
360 pub fn exact_search(&self, query: &str) -> Vec<ModelInfo> {
361 self.models
362 .iter()
363 .filter(|m| {
364 m.model
365 .name()
366 .0
367 .to_lowercase()
368 .contains(&query.to_lowercase())
369 })
370 .cloned()
371 .collect::<Vec<_>>()
372 }
373
374 fn make_match_candidates(model_infos: &Vec<ModelInfo>) -> Vec<StringMatchCandidate> {
375 model_infos
376 .iter()
377 .enumerate()
378 .map(|(index, model)| {
379 StringMatchCandidate::new(
380 index,
381 &format!(
382 "{}/{}",
383 &model.model.provider_name().0,
384 &model.model.name().0
385 ),
386 )
387 })
388 .collect::<Vec<_>>()
389 }
390}
391
392impl PickerDelegate for LanguageModelPickerDelegate {
393 type ListItem = AnyElement;
394
395 fn name() -> &'static str {
396 "language model selector"
397 }
398
399 fn match_count(&self) -> usize {
400 self.filtered_entries.len()
401 }
402
403 fn selected_index(&self) -> usize {
404 self.selected_index
405 }
406
407 fn set_selected_index(&mut self, ix: usize, _: &mut Window, cx: &mut Context<Picker<Self>>) {
408 self.selected_index = ix.min(self.filtered_entries.len().saturating_sub(1));
409 cx.notify();
410 }
411
412 fn can_select(&self, ix: usize, _window: &mut Window, _cx: &mut Context<Picker<Self>>) -> bool {
413 match self.filtered_entries.get(ix) {
414 Some(LanguageModelPickerEntry::Model(_)) => true,
415 Some(LanguageModelPickerEntry::Separator(_)) | None => false,
416 }
417 }
418
419 fn placeholder_text(&self, _window: &mut Window, _cx: &mut App) -> Arc<str> {
420 "Select a model…".into()
421 }
422
423 fn update_matches(
424 &mut self,
425 query: String,
426 window: &mut Window,
427 cx: &mut Context<Picker<Self>>,
428 ) -> Task<()> {
429 let all_models = self.all_models.clone();
430 let active_model = (self.get_active_model)(cx);
431 let fg_executor = cx.foreground_executor();
432 let bg_executor = cx.background_executor();
433
434 let language_model_registry = LanguageModelRegistry::global(cx);
435
436 let configured_providers = language_model_registry
437 .read(cx)
438 .visible_providers()
439 .into_iter()
440 .filter(|provider| provider.is_authenticated(cx))
441 .collect::<Vec<_>>();
442
443 let configured_provider_ids = configured_providers
444 .iter()
445 .map(|provider| provider.id())
446 .collect::<Vec<_>>();
447
448 let recommended_models = all_models
449 .recommended
450 .iter()
451 .filter(|m| configured_provider_ids.contains(&m.model.provider_id()))
452 .cloned()
453 .collect::<Vec<_>>();
454
455 let available_models = all_models
456 .all
457 .values()
458 .flat_map(|models| models.iter())
459 .filter(|m| configured_provider_ids.contains(&m.model.provider_id()))
460 .cloned()
461 .collect::<Vec<_>>();
462
463 let matcher_rec =
464 ModelMatcher::new(recommended_models, fg_executor.clone(), bg_executor.clone());
465 let matcher_all =
466 ModelMatcher::new(available_models, fg_executor.clone(), bg_executor.clone());
467
468 let recommended = matcher_rec.exact_search(&query);
469 let all = matcher_all.fuzzy_search(&query);
470
471 let filtered_models = GroupedModels::new(all, recommended);
472
473 cx.spawn_in(window, async move |this, cx| {
474 this.update_in(cx, |this, window, cx| {
475 this.delegate.filtered_entries = filtered_models.entries();
476 // Finds the currently selected model in the list
477 let new_index =
478 Self::get_active_model_index(&this.delegate.filtered_entries, active_model);
479 this.set_selected_index(new_index, Some(picker::Direction::Down), true, window, cx);
480 cx.notify();
481 })
482 .ok();
483 })
484 }
485
486 fn confirm(&mut self, _secondary: bool, window: &mut Window, cx: &mut Context<Picker<Self>>) {
487 if let Some(LanguageModelPickerEntry::Model(model_info)) =
488 self.filtered_entries.get(self.selected_index)
489 {
490 let model = model_info.model.clone();
491 (self.on_model_changed)(model.clone(), cx);
492
493 let current_index = self.selected_index;
494 self.set_selected_index(current_index, window, cx);
495
496 cx.emit(DismissEvent);
497 }
498 }
499
500 fn dismissed(&mut self, _: &mut Window, cx: &mut Context<Picker<Self>>) {
501 cx.emit(DismissEvent);
502 }
503
504 fn render_match(
505 &self,
506 ix: usize,
507 selected: bool,
508 _: &mut Window,
509 cx: &mut Context<Picker<Self>>,
510 ) -> Option<Self::ListItem> {
511 match self.filtered_entries.get(ix)? {
512 LanguageModelPickerEntry::Separator(title) => {
513 Some(ModelSelectorHeader::new(title, ix > 1).into_any_element())
514 }
515 LanguageModelPickerEntry::Model(model_info) => {
516 let active_model = (self.get_active_model)(cx);
517 let active_provider_id = active_model.as_ref().map(|m| m.provider.id());
518 let active_model_id = active_model.map(|m| m.model.id());
519
520 let is_selected = Some(model_info.model.provider_id()) == active_provider_id
521 && Some(model_info.model.id()) == active_model_id;
522
523 let model_cost = model_info
524 .model
525 .model_cost_info()
526 .map(|cost| cost.to_shared_string());
527
528 let is_favorite = model_info.is_favorite;
529 let handle_action_click = {
530 let model = model_info.model.clone();
531 let on_toggle_favorite = self.on_toggle_favorite.clone();
532 cx.listener(move |picker, _, window, cx| {
533 on_toggle_favorite(model.clone(), !is_favorite, cx);
534 picker.refresh(window, cx);
535 })
536 };
537
538 Some(
539 ModelSelectorListItem::new(ix, model_info.model.name().0)
540 .map(|this| match &model_info.icon {
541 IconOrSvg::Icon(icon_name) => this.icon(*icon_name),
542 IconOrSvg::Svg(icon_path) => this.icon_path(icon_path.clone()),
543 })
544 .is_selected(is_selected)
545 .is_focused(selected)
546 .is_latest(model_info.model.is_latest())
547 .is_favorite(is_favorite)
548 .cost_info(model_cost)
549 .on_toggle_favorite(handle_action_click)
550 .into_any_element(),
551 )
552 }
553 }
554 }
555
556 fn render_footer(
557 &self,
558 _window: &mut Window,
559 _cx: &mut Context<Picker<Self>>,
560 ) -> Option<gpui::AnyElement> {
561 let focus_handle = self.focus_handle.clone();
562
563 if !self.popover_styles {
564 return None;
565 }
566
567 Some(ModelSelectorFooter::new(OpenSettings.boxed_clone(), focus_handle).into_any_element())
568 }
569}
570
571#[cfg(test)]
572mod tests {
573 use super::*;
574 use futures::{future::BoxFuture, stream::BoxStream};
575 use gpui::{AsyncApp, TestAppContext};
576 use language_model::{
577 LanguageModelCompletionError, LanguageModelCompletionEvent, LanguageModelId,
578 LanguageModelName, LanguageModelProviderId, LanguageModelProviderName,
579 LanguageModelRequest, LanguageModelToolChoice,
580 };
581 use ui::IconName;
582
583 #[derive(Clone)]
584 struct TestLanguageModel {
585 name: LanguageModelName,
586 id: LanguageModelId,
587 provider_id: LanguageModelProviderId,
588 provider_name: LanguageModelProviderName,
589 }
590
591 impl TestLanguageModel {
592 fn new(name: &str, provider: &str) -> Self {
593 Self {
594 name: LanguageModelName::from(name.to_string()),
595 id: LanguageModelId::from(name.to_string()),
596 provider_id: LanguageModelProviderId::from(provider.to_string()),
597 provider_name: LanguageModelProviderName::from(provider.to_string()),
598 }
599 }
600 }
601
602 impl LanguageModel for TestLanguageModel {
603 fn id(&self) -> LanguageModelId {
604 self.id.clone()
605 }
606
607 fn name(&self) -> LanguageModelName {
608 self.name.clone()
609 }
610
611 fn provider_id(&self) -> LanguageModelProviderId {
612 self.provider_id.clone()
613 }
614
615 fn provider_name(&self) -> LanguageModelProviderName {
616 self.provider_name.clone()
617 }
618
619 fn supports_tools(&self) -> bool {
620 false
621 }
622
623 fn supports_tool_choice(&self, _choice: LanguageModelToolChoice) -> bool {
624 false
625 }
626
627 fn supports_images(&self) -> bool {
628 false
629 }
630
631 fn telemetry_id(&self) -> String {
632 format!("{}/{}", self.provider_id.0, self.name.0)
633 }
634
635 fn max_token_count(&self) -> u64 {
636 1000
637 }
638
639 fn stream_completion(
640 &self,
641 _: LanguageModelRequest,
642 _: &AsyncApp,
643 ) -> BoxFuture<
644 'static,
645 Result<
646 BoxStream<
647 'static,
648 Result<LanguageModelCompletionEvent, LanguageModelCompletionError>,
649 >,
650 LanguageModelCompletionError,
651 >,
652 > {
653 unimplemented!()
654 }
655 }
656
657 fn create_models(model_specs: Vec<(&str, &str)>) -> Vec<ModelInfo> {
658 create_models_with_favorites(model_specs, vec![])
659 }
660
661 fn create_models_with_favorites(
662 model_specs: Vec<(&str, &str)>,
663 favorites: Vec<(&str, &str)>,
664 ) -> Vec<ModelInfo> {
665 model_specs
666 .into_iter()
667 .map(|(provider, name)| {
668 let is_favorite = favorites
669 .iter()
670 .any(|(fav_provider, fav_name)| *fav_provider == provider && *fav_name == name);
671 ModelInfo {
672 model: Arc::new(TestLanguageModel::new(name, provider)),
673 icon: IconOrSvg::Icon(IconName::OmegaAgent),
674 is_favorite,
675 }
676 })
677 .collect()
678 }
679
680 fn assert_models_eq(result: Vec<ModelInfo>, expected: Vec<&str>) {
681 assert_eq!(
682 result.len(),
683 expected.len(),
684 "Number of models doesn't match"
685 );
686
687 for (i, expected_name) in expected.iter().enumerate() {
688 assert_eq!(
689 result[i].model.telemetry_id(),
690 *expected_name,
691 "Model at position {} doesn't match expected model",
692 i
693 );
694 }
695 }
696
697 #[gpui::test]
698 fn test_exact_match(cx: &mut TestAppContext) {
699 let models = create_models(vec![
700 ("zed", "Claude 3.7 Sonnet"),
701 ("zed", "Claude 3.7 Sonnet Thinking"),
702 ("zed", "gpt-5"),
703 ("zed", "gpt-5-mini"),
704 ("openai", "gpt-3.5-turbo"),
705 ("openai", "gpt-5"),
706 ("openai", "gpt-5-mini"),
707 ("ollama", "mistral"),
708 ("ollama", "deepseek"),
709 ]);
710 let matcher = ModelMatcher::new(
711 models,
712 cx.foreground_executor().clone(),
713 cx.background_executor.clone(),
714 );
715
716 // The order of models should be maintained, case doesn't matter
717 let results = matcher.exact_search("GPT-5");
718 assert_models_eq(
719 results,
720 vec![
721 "zed/gpt-5",
722 "zed/gpt-5-mini",
723 "openai/gpt-5",
724 "openai/gpt-5-mini",
725 ],
726 );
727 }
728
729 #[gpui::test]
730 fn test_fuzzy_match(cx: &mut TestAppContext) {
731 let models = create_models(vec![
732 ("zed", "Claude 3.7 Sonnet"),
733 ("zed", "Claude 3.7 Sonnet Thinking"),
734 ("zed", "gpt-5"),
735 ("zed", "gpt-5-mini"),
736 ("openai", "gpt-3.5-turbo"),
737 ("openai", "gpt-5"),
738 ("openai", "gpt-5-mini"),
739 ("ollama", "mistral"),
740 ("ollama", "deepseek"),
741 ]);
742 let matcher = ModelMatcher::new(
743 models,
744 cx.foreground_executor().clone(),
745 cx.background_executor.clone(),
746 );
747
748 // Results should preserve models order whenever possible.
749 // In the case below, `zed/gpt-5-mini` and `openai/gpt-5-mini` have identical
750 // similarity scores, but `zed/gpt-5-mini` was higher in the models list,
751 // so it should appear first in the results.
752 let results = matcher.fuzzy_search("mini");
753 assert_models_eq(results, vec!["zed/gpt-5-mini", "openai/gpt-5-mini"]);
754
755 // Model provider should be searchable as well
756 let results = matcher.fuzzy_search("ol"); // meaning "ollama"
757 assert_models_eq(results, vec!["ollama/mistral", "ollama/deepseek"]);
758
759 // Fuzzy search - search for Claude to get the Thinking variant
760 let results = matcher.fuzzy_search("thinking");
761 assert_models_eq(results, vec!["zed/Claude 3.7 Sonnet Thinking"]);
762 }
763
764 #[gpui::test]
765 fn test_recommended_models_also_appear_in_other(_cx: &mut TestAppContext) {
766 let recommended_models = create_models(vec![("zed", "claude")]);
767 let all_models = create_models(vec![
768 ("zed", "claude"), // Should also appear in "other"
769 ("zed", "gemini"),
770 ("copilot", "o3"),
771 ]);
772
773 let grouped_models = GroupedModels::new(all_models, recommended_models);
774
775 let actual_all_models = grouped_models
776 .all
777 .values()
778 .flatten()
779 .cloned()
780 .collect::<Vec<_>>();
781
782 // Recommended models should also appear in "all"
783 assert_models_eq(
784 actual_all_models,
785 vec!["zed/claude", "zed/gemini", "copilot/o3"],
786 );
787 }
788
789 #[gpui::test]
790 fn test_models_from_different_providers(_cx: &mut TestAppContext) {
791 let recommended_models = create_models(vec![("zed", "claude")]);
792 let all_models = create_models(vec![
793 ("zed", "claude"), // Should also appear in "other"
794 ("zed", "gemini"),
795 ("copilot", "claude"), // Different provider, should appear in "other"
796 ]);
797
798 let grouped_models = GroupedModels::new(all_models, recommended_models);
799
800 let actual_all_models = grouped_models
801 .all
802 .values()
803 .flatten()
804 .cloned()
805 .collect::<Vec<_>>();
806
807 // All models should appear in "all" regardless of recommended status
808 assert_models_eq(
809 actual_all_models,
810 vec!["zed/claude", "zed/gemini", "copilot/claude"],
811 );
812 }
813
814 #[gpui::test]
815 fn test_favorites_section_appears_when_favorites_exist(_cx: &mut TestAppContext) {
816 let recommended_models = create_models(vec![("zed", "claude")]);
817 let all_models = create_models_with_favorites(
818 vec![("zed", "claude"), ("zed", "gemini"), ("openai", "gpt-4")],
819 vec![("zed", "gemini")],
820 );
821
822 let grouped_models = GroupedModels::new(all_models, recommended_models);
823 let entries = grouped_models.entries();
824
825 assert!(matches!(
826 entries.first(),
827 Some(LanguageModelPickerEntry::Separator(s)) if s == "Favorite"
828 ));
829
830 assert_models_eq(grouped_models.favorites, vec!["zed/gemini"]);
831 }
832
833 #[gpui::test]
834 fn test_no_favorites_section_when_no_favorites(_cx: &mut TestAppContext) {
835 let recommended_models = create_models(vec![("zed", "claude")]);
836 let all_models = create_models(vec![("zed", "claude"), ("zed", "gemini")]);
837
838 let grouped_models = GroupedModels::new(all_models, recommended_models);
839 let entries = grouped_models.entries();
840
841 assert!(matches!(
842 entries.first(),
843 Some(LanguageModelPickerEntry::Separator(s)) if s == "Recommended"
844 ));
845
846 assert!(grouped_models.favorites.is_empty());
847 }
848
849 #[gpui::test]
850 fn test_models_have_correct_actions(_cx: &mut TestAppContext) {
851 let recommended_models =
852 create_models_with_favorites(vec![("zed", "claude")], vec![("zed", "claude")]);
853 let all_models = create_models_with_favorites(
854 vec![("zed", "claude"), ("zed", "gemini"), ("openai", "gpt-4")],
855 vec![("zed", "claude")],
856 );
857
858 let grouped_models = GroupedModels::new(all_models, recommended_models);
859 let entries = grouped_models.entries();
860
861 for entry in &entries {
862 if let LanguageModelPickerEntry::Model(info) = entry {
863 if info.model.telemetry_id() == "zed/claude" {
864 assert!(info.is_favorite, "zed/claude should be a favorite");
865 } else {
866 assert!(
867 !info.is_favorite,
868 "{} should not be a favorite",
869 info.model.telemetry_id()
870 );
871 }
872 }
873 }
874 }
875
876 #[gpui::test]
877 fn test_favorites_appear_in_other_sections(_cx: &mut TestAppContext) {
878 let favorites = vec![("zed", "gemini"), ("openai", "gpt-4")];
879
880 let recommended_models =
881 create_models_with_favorites(vec![("zed", "claude")], favorites.clone());
882
883 let all_models = create_models_with_favorites(
884 vec![
885 ("zed", "claude"),
886 ("zed", "gemini"),
887 ("openai", "gpt-4"),
888 ("openai", "gpt-3.5"),
889 ],
890 favorites,
891 );
892
893 let grouped_models = GroupedModels::new(all_models, recommended_models);
894
895 assert_models_eq(grouped_models.favorites, vec!["zed/gemini", "openai/gpt-4"]);
896 assert_models_eq(grouped_models.recommended, vec!["zed/claude"]);
897 assert_models_eq(
898 grouped_models.all.values().flatten().cloned().collect(),
899 vec!["zed/claude", "zed/gemini", "openai/gpt-4", "openai/gpt-3.5"],
900 );
901 }
902}
903