Skip to repository content

tenant.openagents/omega

No repository description is available.

OpenAgents Git authority 2026-07-28T02:55:10.274Z Public web read
NIP-34 coordinate30617:7649603503856e5148d571eac2766b288a8ff1e9e35d380337a1d2b0015b4f92:omega
MaintainersHidden in public view
References2 branches · 1 tag
Read-only clonegit clone https://openagents.com/git/tenant.openagents/omega.git
Browse files

language_model_selector.rs

903 lines · 29.9 KB · rust
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
Served at tenant.openagents/omega Member data and write actions are omitted.