Skip to repository content896 lines · 31.3 KB · rust
tenant.openagents/omega
No repository description is available.
OpenAgents Git authority 2026-07-28T02:57:06.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
model_selector.rs
1use std::{cmp::Reverse, rc::Rc, sync::Arc};
2
3use acp_thread::{
4 AgentModelIcon, AgentModelId, AgentModelInfo, AgentModelList, AgentModelSelector,
5};
6
7use anyhow::Result;
8use collections::{HashSet, IndexMap};
9use futures::FutureExt;
10use fuzzy::{StringMatchCandidate, match_strings};
11use gpui::{
12 Action, AsyncWindowContext, BackgroundExecutor, DismissEvent, FocusHandle, Subscription, Task,
13 TaskExt, WeakEntity,
14};
15use itertools::Itertools;
16use ordered_float::OrderedFloat;
17use picker::{Picker, PickerDelegate};
18use settings::SettingsStore;
19use ui::{DocumentationAside, IntoElement, prelude::*};
20use util::ResultExt;
21use zed_actions::agent::OpenSettings;
22
23use crate::ui::{
24 ModelSelectorFooter, ModelSelectorHeader, ModelSelectorListItem, documentation_aside_side,
25};
26
27pub type ModelSelector = Picker<ModelPickerDelegate>;
28
29pub fn acp_model_selector(
30 selector: Rc<dyn AgentModelSelector>,
31 focus_handle: FocusHandle,
32 window: &mut Window,
33 cx: &mut Context<ModelSelector>,
34) -> ModelSelector {
35 let delegate = ModelPickerDelegate::new(selector, focus_handle, window, cx);
36 Picker::list(delegate, window, cx)
37 .show_scrollbar(true)
38 .initial_width(rems(20.))
39}
40
41enum ModelPickerEntry {
42 Separator(SharedString),
43 Model(AgentModelInfo, bool),
44}
45
46pub struct ModelPickerDelegate {
47 selector: Rc<dyn AgentModelSelector>,
48 filtered_entries: Vec<ModelPickerEntry>,
49 models: Option<AgentModelList>,
50 selected_index: usize,
51 selected_description: Option<(usize, SharedString)>,
52 selected_model: Option<AgentModelInfo>,
53 favorites: HashSet<AgentModelId>,
54 _refresh_models_task: Task<()>,
55 _settings_subscription: Subscription,
56 focus_handle: FocusHandle,
57}
58
59impl ModelPickerDelegate {
60 fn new(
61 selector: Rc<dyn AgentModelSelector>,
62 focus_handle: FocusHandle,
63 window: &mut Window,
64 cx: &mut Context<ModelSelector>,
65 ) -> Self {
66 let rx = selector.watch(cx);
67 let refresh_models_task = {
68 cx.spawn_in(window, {
69 async move |this, cx| {
70 async fn refresh(
71 this: &WeakEntity<Picker<ModelPickerDelegate>>,
72 cx: &mut AsyncWindowContext,
73 ) -> Result<()> {
74 let (models_task, selected_model_task) = this.update(cx, |this, cx| {
75 (
76 this.delegate.selector.list_models(cx),
77 this.delegate.selector.selected_model(cx),
78 )
79 })?;
80
81 let (models, selected_model) =
82 futures::join!(models_task, selected_model_task);
83
84 this.update_in(cx, |this, window, cx| {
85 this.delegate.models = models.ok();
86 this.delegate.selected_model = selected_model.ok();
87 this.refresh(window, cx)
88 })
89 }
90
91 refresh(&this, cx).await.log_err();
92 if let Some(mut rx) = rx {
93 while let Ok(()) = rx.recv().await {
94 refresh(&this, cx).await.log_err();
95 }
96 }
97 }
98 })
99 };
100
101 let selector_for_subscription = selector.clone();
102 let settings_subscription =
103 cx.observe_global_in::<SettingsStore>(window, move |picker, window, cx| {
104 // Only refresh if the favorites actually changed to avoid redundant work
105 // when other settings are modified (e.g., user editing settings.json)
106 let new_favorites = selector_for_subscription.favorite_model_ids(cx);
107 if new_favorites != picker.delegate.favorites {
108 picker.delegate.favorites = new_favorites;
109 picker.refresh(window, cx);
110 }
111 });
112 let favorites = selector.favorite_model_ids(cx);
113
114 Self {
115 selector,
116 filtered_entries: Vec::new(),
117 models: None,
118 selected_model: None,
119 selected_index: 0,
120 selected_description: None,
121 favorites,
122 _refresh_models_task: refresh_models_task,
123 _settings_subscription: settings_subscription,
124 focus_handle,
125 }
126 }
127
128 pub fn active_model(&self) -> Option<&AgentModelInfo> {
129 self.selected_model.as_ref()
130 }
131
132 pub fn favorites_count(&self) -> usize {
133 self.favorites.len()
134 }
135
136 pub fn cycle_favorite_models(&mut self, window: &mut Window, cx: &mut Context<Picker<Self>>) {
137 if self.favorites.is_empty() {
138 return;
139 }
140
141 let Some(models) = &self.models else {
142 return;
143 };
144
145 let all_models: Vec<&AgentModelInfo> = match models {
146 AgentModelList::Flat(list) => list.iter().collect(),
147 AgentModelList::Grouped(index_map) => index_map.values().flatten().collect(),
148 };
149
150 let favorite_models: Vec<_> = all_models
151 .into_iter()
152 .filter(|model| self.favorites.contains(&model.id))
153 .unique_by(|model| &model.id)
154 .collect();
155
156 if favorite_models.is_empty() {
157 return;
158 }
159
160 let current_id = self.selected_model.as_ref().map(|m| &m.id);
161
162 let current_index_in_favorites = current_id
163 .and_then(|id| favorite_models.iter().position(|m| &m.id == id))
164 .unwrap_or(usize::MAX);
165
166 let next_index = if current_index_in_favorites == usize::MAX {
167 0
168 } else {
169 (current_index_in_favorites + 1) % favorite_models.len()
170 };
171
172 let next_model = favorite_models[next_index].clone();
173
174 self.selector
175 .select_model(next_model.id.clone(), cx)
176 .detach_and_log_err(cx);
177
178 self.selected_model = Some(next_model);
179
180 // Keep the picker selection aligned with the newly-selected model
181 if let Some(new_index) = self.filtered_entries.iter().position(|entry| {
182 matches!(entry, ModelPickerEntry::Model(model_info, _) if self.selected_model.as_ref().is_some_and(|selected| model_info.id == selected.id))
183 }) {
184 self.set_selected_index(new_index, window, cx);
185 } else {
186 cx.notify();
187 }
188 }
189}
190
191impl PickerDelegate for ModelPickerDelegate {
192 type ListItem = AnyElement;
193
194 fn name() -> &'static str {
195 "model selector"
196 }
197
198 fn match_count(&self) -> usize {
199 self.filtered_entries.len()
200 }
201
202 fn selected_index(&self) -> usize {
203 self.selected_index
204 }
205
206 fn set_selected_index(&mut self, ix: usize, _: &mut Window, cx: &mut Context<Picker<Self>>) {
207 self.selected_index = ix.min(self.filtered_entries.len().saturating_sub(1));
208 cx.notify();
209 }
210
211 fn can_select(&self, ix: usize, _window: &mut Window, _cx: &mut Context<Picker<Self>>) -> bool {
212 match self.filtered_entries.get(ix) {
213 Some(ModelPickerEntry::Model(_, _)) => true,
214 Some(ModelPickerEntry::Separator(_)) | None => false,
215 }
216 }
217
218 fn placeholder_text(&self, _window: &mut Window, _cx: &mut App) -> Arc<str> {
219 "Select a model…".into()
220 }
221
222 fn update_matches(
223 &mut self,
224 query: String,
225 window: &mut Window,
226 cx: &mut Context<Picker<Self>>,
227 ) -> Task<()> {
228 let favorites = self.favorites.clone();
229
230 cx.spawn_in(window, async move |this, cx| {
231 let filtered_models = match this
232 .read_with(cx, |this, cx| {
233 this.delegate.models.clone().map(move |models| {
234 fuzzy_search(models, query, cx.background_executor().clone())
235 })
236 })
237 .ok()
238 .flatten()
239 {
240 Some(task) => task.await,
241 None => AgentModelList::Flat(vec![]),
242 };
243
244 this.update_in(cx, |this, window, cx| {
245 this.delegate.filtered_entries =
246 info_list_to_picker_entries(filtered_models, &favorites);
247 // Finds the currently selected model in the list
248 let new_index = this
249 .delegate
250 .selected_model
251 .as_ref()
252 .and_then(|selected| {
253 this.delegate.filtered_entries.iter().position(|entry| {
254 if let ModelPickerEntry::Model(model_info, _) = entry {
255 model_info.id == selected.id
256 } else {
257 false
258 }
259 })
260 })
261 .unwrap_or(0);
262 this.set_selected_index(new_index, Some(picker::Direction::Down), true, window, cx);
263 cx.notify();
264 })
265 .ok();
266 })
267 }
268
269 fn confirm(&mut self, _secondary: bool, window: &mut Window, cx: &mut Context<Picker<Self>>) {
270 if let Some(ModelPickerEntry::Model(model_info, _)) =
271 self.filtered_entries.get(self.selected_index)
272 && model_info.disabled.is_none()
273 {
274 self.selector
275 .select_model(model_info.id.clone(), cx)
276 .detach_and_log_err(cx);
277 self.selected_model = Some(model_info.clone());
278 let current_index = self.selected_index;
279 self.set_selected_index(current_index, window, cx);
280
281 cx.emit(DismissEvent);
282 }
283 }
284
285 fn dismissed(&mut self, window: &mut Window, cx: &mut Context<Picker<Self>>) {
286 cx.defer_in(window, |picker, window, cx| {
287 picker.set_query("", window, cx);
288 });
289 }
290
291 fn render_match(
292 &self,
293 ix: usize,
294 selected: bool,
295 _: &mut Window,
296 cx: &mut Context<Picker<Self>>,
297 ) -> Option<Self::ListItem> {
298 match self.filtered_entries.get(ix)? {
299 ModelPickerEntry::Separator(title) => {
300 Some(ModelSelectorHeader::new(title, ix > 1).into_any_element())
301 }
302 ModelPickerEntry::Model(model_info, is_favorite) => {
303 let is_selected = Some(model_info) == self.selected_model.as_ref();
304
305 let is_favorite = *is_favorite;
306 let handle_action_click = {
307 let model_id = model_info.id.clone();
308 let selector = self.selector.clone();
309
310 cx.listener(move |_, _, _, cx| {
311 selector.toggle_favorite_model(model_id.clone(), !is_favorite, cx);
312 })
313 };
314
315 let model_cost = model_info.cost.clone();
316
317 Some(
318 div()
319 .id(("model-picker-menu-child", ix))
320 .when_some(model_info.description.clone(), |this, description| {
321 this.on_hover(cx.listener(move |menu, hovered, _, cx| {
322 if *hovered {
323 menu.delegate.selected_description =
324 Some((ix, description.clone()));
325 } else if matches!(menu.delegate.selected_description, Some((id, _)) if id == ix) {
326 menu.delegate.selected_description = None;
327 }
328 cx.notify();
329 }))
330 })
331 .child(
332 ModelSelectorListItem::new(ix, model_info.name.clone())
333 .map(|this| match &model_info.icon {
334 Some(AgentModelIcon::Path(path)) => this.icon_path(path.clone()),
335 Some(AgentModelIcon::Named(icon)) => this.icon(*icon),
336 None => this,
337 })
338 .disabled(model_info.disabled.clone())
339 .is_selected(is_selected)
340 .is_focused(selected)
341 .is_latest(model_info.is_latest)
342 .is_favorite(is_favorite)
343 .on_toggle_favorite(handle_action_click)
344 .cost_info(model_cost)
345 )
346 .into_any_element(),
347 )
348 }
349 }
350 }
351
352 fn documentation_aside(
353 &self,
354 _window: &mut Window,
355 cx: &mut Context<Picker<Self>>,
356 ) -> Option<ui::DocumentationAside> {
357 self.selected_description.as_ref().map(|(_, description)| {
358 let description = description.clone();
359
360 let side = documentation_aside_side(cx);
361
362 DocumentationAside::new(
363 side,
364 Rc::new(move |_| Label::new(description.clone()).into_any_element()),
365 )
366 })
367 }
368
369 fn documentation_aside_index(&self) -> Option<usize> {
370 self.selected_description.as_ref().map(|(ix, _)| *ix)
371 }
372
373 fn render_footer(
374 &self,
375 _window: &mut Window,
376 _cx: &mut Context<Picker<Self>>,
377 ) -> Option<AnyElement> {
378 let focus_handle = self.focus_handle.clone();
379
380 if !self.selector.should_render_footer() {
381 return None;
382 }
383
384 Some(ModelSelectorFooter::new(OpenSettings.boxed_clone(), focus_handle).into_any_element())
385 }
386}
387
388fn info_list_to_picker_entries(
389 model_list: AgentModelList,
390 favorites: &HashSet<AgentModelId>,
391) -> Vec<ModelPickerEntry> {
392 let mut entries = Vec::new();
393
394 let all_models: Vec<_> = match &model_list {
395 AgentModelList::Flat(list) => list.iter().collect(),
396 AgentModelList::Grouped(index_map) => index_map.values().flatten().collect(),
397 };
398
399 let favorite_models: Vec<_> = all_models
400 .iter()
401 .filter(|m| favorites.contains(&m.id))
402 .unique_by(|m| &m.id)
403 .collect();
404
405 let has_favorites = !favorite_models.is_empty();
406 if has_favorites {
407 entries.push(ModelPickerEntry::Separator("Favorite".into()));
408 for model in favorite_models {
409 entries.push(ModelPickerEntry::Model((*model).clone(), true));
410 }
411 }
412
413 match model_list {
414 AgentModelList::Flat(list) => {
415 if has_favorites {
416 entries.push(ModelPickerEntry::Separator("All".into()));
417 }
418 for model in list {
419 let is_favorite = favorites.contains(&model.id);
420 entries.push(ModelPickerEntry::Model(model, is_favorite));
421 }
422 }
423 AgentModelList::Grouped(index_map) => {
424 for (group_name, models) in index_map {
425 entries.push(ModelPickerEntry::Separator(group_name.0));
426 for model in models {
427 let is_favorite = favorites.contains(&model.id);
428 entries.push(ModelPickerEntry::Model(model, is_favorite));
429 }
430 }
431 }
432 }
433
434 entries
435}
436
437async fn fuzzy_search(
438 model_list: AgentModelList,
439 query: String,
440 executor: BackgroundExecutor,
441) -> AgentModelList {
442 async fn fuzzy_search_list(
443 model_list: Vec<AgentModelInfo>,
444 query: &str,
445 executor: BackgroundExecutor,
446 ) -> Vec<AgentModelInfo> {
447 let candidates = model_list
448 .iter()
449 .enumerate()
450 .map(|(ix, model)| StringMatchCandidate::new(ix, model.name.as_ref()))
451 .collect::<Vec<_>>();
452 let mut matches = match_strings(
453 &candidates,
454 query,
455 false,
456 true,
457 100,
458 &Default::default(),
459 executor,
460 )
461 .await;
462
463 matches.sort_unstable_by_key(|mat| {
464 let candidate = &candidates[mat.candidate_id];
465 (Reverse(OrderedFloat(mat.score)), candidate.id)
466 });
467
468 matches
469 .into_iter()
470 .map(|mat| model_list[mat.candidate_id].clone())
471 .collect()
472 }
473
474 match model_list {
475 AgentModelList::Flat(model_list) => {
476 AgentModelList::Flat(fuzzy_search_list(model_list, &query, executor).await)
477 }
478 AgentModelList::Grouped(index_map) => {
479 let groups =
480 futures::future::join_all(index_map.into_iter().map(|(group_name, models)| {
481 fuzzy_search_list(models, &query, executor.clone())
482 .map(|results| (group_name, results))
483 }))
484 .await;
485 AgentModelList::Grouped(IndexMap::from_iter(
486 groups
487 .into_iter()
488 .filter(|(_, results)| !results.is_empty()),
489 ))
490 }
491 }
492}
493
494#[cfg(test)]
495mod tests {
496 use gpui::{App, TestAppContext, VisualTestContext};
497 use std::cell::RefCell;
498
499 use super::*;
500
501 fn create_model_list(grouped_models: Vec<(&str, Vec<&str>)>) -> AgentModelList {
502 AgentModelList::Grouped(IndexMap::from_iter(grouped_models.into_iter().map(
503 |(group, models)| {
504 (
505 acp_thread::AgentModelGroupName(group.to_string().into()),
506 models
507 .into_iter()
508 .map(|model| acp_thread::AgentModelInfo {
509 id: AgentModelId::new(model),
510 name: model.to_string().into(),
511 description: None,
512 icon: None,
513 is_latest: false,
514 disabled: None,
515 cost: None,
516 })
517 .collect::<Vec<_>>(),
518 )
519 },
520 )))
521 }
522
523 fn assert_models_eq(result: AgentModelList, expected: Vec<(&str, Vec<&str>)>) {
524 let AgentModelList::Grouped(groups) = result else {
525 panic!("Expected LanguageModelInfoList::Grouped, got {:?}", result);
526 };
527
528 assert_eq!(
529 groups.len(),
530 expected.len(),
531 "Number of groups doesn't match"
532 );
533
534 for (i, (expected_group, expected_models)) in expected.iter().enumerate() {
535 let (actual_group, actual_models) = groups.get_index(i).unwrap();
536 assert_eq!(
537 actual_group.0.as_ref(),
538 *expected_group,
539 "Group at position {} doesn't match expected group",
540 i
541 );
542 assert_eq!(
543 actual_models.len(),
544 expected_models.len(),
545 "Number of models in group {} doesn't match",
546 expected_group
547 );
548
549 for (j, expected_model_name) in expected_models.iter().enumerate() {
550 assert_eq!(
551 actual_models[j].name, *expected_model_name,
552 "Model at position {} in group {} doesn't match expected model",
553 j, expected_group
554 );
555 }
556 }
557 }
558
559 fn create_favorites(models: Vec<&str>) -> HashSet<AgentModelId> {
560 models.into_iter().map(AgentModelId::new).collect()
561 }
562
563 fn get_entry_model_ids(entries: &[ModelPickerEntry]) -> Vec<&str> {
564 entries
565 .iter()
566 .filter_map(|entry| match entry {
567 ModelPickerEntry::Model(info, _) => Some(info.id.as_ref()),
568 _ => None,
569 })
570 .collect()
571 }
572
573 #[gpui::test]
574 fn confirming_model_selects_model(cx: &mut TestAppContext) {
575 cx.update(|cx| {
576 let settings_store = settings::SettingsStore::test(cx);
577 cx.set_global(settings_store);
578 theme_settings::init(theme::LoadThemes::JustBase, cx);
579 editor::init(cx);
580 });
581
582 let model_selector = Rc::new(TestModelSelector::new());
583
584 let window_handle = cx.add_window({
585 let model_selector = model_selector.clone();
586 move |window, cx| {
587 let selector: Rc<dyn AgentModelSelector> = model_selector;
588 acp_model_selector(selector, cx.focus_handle(), window, cx)
589 }
590 });
591 cx.run_until_parked();
592
593 let mut cx = VisualTestContext::from_window(window_handle.into(), cx);
594 window_handle
595 .update(&mut cx, |picker, window, cx| {
596 picker.delegate.set_selected_index(1, window, cx);
597 picker.delegate.confirm(false, window, cx);
598 })
599 .unwrap();
600
601 assert_eq!(
602 model_selector.selected_models.borrow().as_slice(),
603 &[AgentModelId::new("manual")]
604 );
605 }
606
607 struct TestModelSelector {
608 models: Vec<AgentModelInfo>,
609 selected_model: RefCell<AgentModelInfo>,
610 selected_models: RefCell<Vec<AgentModelId>>,
611 }
612
613 impl TestModelSelector {
614 fn new() -> Self {
615 let models = vec![
616 AgentModelInfo {
617 id: AgentModelId::new("auto"),
618 name: "Auto".into(),
619 description: None,
620 icon: None,
621 is_latest: false,
622 disabled: None,
623 cost: None,
624 },
625 AgentModelInfo {
626 id: AgentModelId::new("manual"),
627 name: "Manual".into(),
628 description: None,
629 icon: None,
630 is_latest: false,
631 disabled: None,
632 cost: None,
633 },
634 ];
635
636 Self {
637 selected_model: RefCell::new(models[0].clone()),
638 models,
639 selected_models: RefCell::new(Vec::new()),
640 }
641 }
642 }
643
644 impl AgentModelSelector for TestModelSelector {
645 fn list_models(&self, _cx: &mut App) -> Task<Result<AgentModelList>> {
646 Task::ready(Ok(AgentModelList::Flat(self.models.clone())))
647 }
648
649 fn select_model(&self, model_id: AgentModelId, _cx: &mut App) -> Task<Result<()>> {
650 self.selected_models.borrow_mut().push(model_id.clone());
651 if let Some(model) = self.models.iter().find(|model| model.id == model_id) {
652 *self.selected_model.borrow_mut() = model.clone();
653 }
654 Task::ready(Ok(()))
655 }
656
657 fn selected_model(&self, _cx: &mut App) -> Task<Result<AgentModelInfo>> {
658 Task::ready(Ok(self.selected_model.borrow().clone()))
659 }
660 }
661
662 fn get_entry_labels(entries: &[ModelPickerEntry]) -> Vec<&str> {
663 entries
664 .iter()
665 .map(|entry| match entry {
666 ModelPickerEntry::Model(info, _) => info.id.as_ref(),
667 ModelPickerEntry::Separator(s) => &s,
668 })
669 .collect()
670 }
671
672 #[gpui::test]
673 async fn test_fuzzy_match(cx: &mut TestAppContext) {
674 let models = create_model_list(vec![
675 (
676 "zed",
677 vec![
678 "Claude 3.7 Sonnet",
679 "Claude 3.7 Sonnet Thinking",
680 "gpt-5",
681 "gpt-5-mini",
682 ],
683 ),
684 ("openai", vec!["gpt-3.5-turbo", "gpt-5", "gpt-5-mini"]),
685 ("ollama", vec!["mistral", "deepseek"]),
686 ]);
687
688 // Results should preserve models order whenever possible.
689 // In the case below, `zed/gpt-5-mini` and `openai/gpt-5-mini` have identical
690 // similarity scores, but `zed/gpt-5-mini` was higher in the models list,
691 // so it should appear first in the results.
692 let results = fuzzy_search(models.clone(), "mini".into(), cx.executor()).await;
693 assert_models_eq(
694 results,
695 vec![("zed", vec!["gpt-5-mini"]), ("openai", vec!["gpt-5-mini"])],
696 );
697
698 // Fuzzy search - test with specific model name
699 let results = fuzzy_search(models.clone(), "mistral".into(), cx.executor()).await;
700 assert_models_eq(results, vec![("ollama", vec!["mistral"])]);
701 }
702
703 #[gpui::test]
704 fn test_favorites_section_appears_when_favorites_exist(_cx: &mut TestAppContext) {
705 let models = create_model_list(vec![
706 ("zed", vec!["zed/claude", "zed/gemini"]),
707 ("openai", vec!["openai/gpt-5"]),
708 ]);
709 let favorites = create_favorites(vec!["zed/gemini"]);
710
711 let entries = info_list_to_picker_entries(models, &favorites);
712
713 assert!(matches!(
714 entries.first(),
715 Some(ModelPickerEntry::Separator(s)) if s == "Favorite"
716 ));
717
718 let model_ids = get_entry_model_ids(&entries);
719 assert_eq!(model_ids[0], "zed/gemini");
720 }
721
722 #[gpui::test]
723 fn test_no_favorites_section_when_no_favorites(_cx: &mut TestAppContext) {
724 let models = create_model_list(vec![("zed", vec!["zed/claude", "zed/gemini"])]);
725 let favorites = create_favorites(vec![]);
726
727 let entries = info_list_to_picker_entries(models, &favorites);
728
729 assert!(matches!(
730 entries.first(),
731 Some(ModelPickerEntry::Separator(s)) if s == "zed"
732 ));
733 }
734
735 #[gpui::test]
736 fn test_models_have_correct_actions(_cx: &mut TestAppContext) {
737 let models = create_model_list(vec![
738 ("zed", vec!["zed/claude", "zed/gemini"]),
739 ("openai", vec!["openai/gpt-5"]),
740 ]);
741 let favorites = create_favorites(vec!["zed/claude"]);
742
743 let entries = info_list_to_picker_entries(models, &favorites);
744
745 for entry in &entries {
746 if let ModelPickerEntry::Model(info, is_favorite) = entry {
747 if info.id.as_ref() == "zed/claude" {
748 assert!(is_favorite, "zed/claude should be a favorite");
749 } else {
750 assert!(!is_favorite, "{} should not be a favorite", info.id);
751 }
752 }
753 }
754 }
755
756 #[gpui::test]
757 fn test_favorites_appear_in_both_sections(_cx: &mut TestAppContext) {
758 let models = create_model_list(vec![
759 ("zed", vec!["zed/claude", "zed/gemini"]),
760 ("openai", vec!["openai/gpt-5", "openai/gpt-4"]),
761 ]);
762 let favorites = create_favorites(vec!["zed/gemini", "openai/gpt-5"]);
763
764 let entries = info_list_to_picker_entries(models, &favorites);
765 let model_ids = get_entry_model_ids(&entries);
766
767 assert_eq!(model_ids[0], "zed/gemini");
768 assert_eq!(model_ids[1], "openai/gpt-5");
769
770 assert!(model_ids[2..].contains(&"zed/gemini"));
771 assert!(model_ids[2..].contains(&"openai/gpt-5"));
772 }
773
774 #[gpui::test]
775 fn test_favorites_are_not_duplicated_when_repeated_in_other_sections(_cx: &mut TestAppContext) {
776 let models = create_model_list(vec![
777 ("Recommended", vec!["zed/claude", "anthropic/claude"]),
778 ("Zed", vec!["zed/claude", "zed/gpt-5"]),
779 ("Antropic", vec!["anthropic/claude"]),
780 ("OpenAI", vec!["openai/gpt-5"]),
781 ]);
782
783 let favorites = create_favorites(vec!["zed/claude"]);
784
785 let entries = info_list_to_picker_entries(models, &favorites);
786 let labels = get_entry_labels(&entries);
787
788 assert_eq!(
789 labels,
790 vec![
791 "Favorite",
792 "zed/claude",
793 "Recommended",
794 "zed/claude",
795 "anthropic/claude",
796 "Zed",
797 "zed/claude",
798 "zed/gpt-5",
799 "Antropic",
800 "anthropic/claude",
801 "OpenAI",
802 "openai/gpt-5"
803 ]
804 );
805 }
806
807 #[gpui::test]
808 fn test_flat_model_list_with_favorites(_cx: &mut TestAppContext) {
809 let models = AgentModelList::Flat(vec![
810 acp_thread::AgentModelInfo {
811 id: AgentModelId::new("zed/claude"),
812 name: "Claude".into(),
813 description: None,
814 icon: None,
815 is_latest: false,
816 disabled: None,
817 cost: None,
818 },
819 acp_thread::AgentModelInfo {
820 id: AgentModelId::new("zed/gemini"),
821 name: "Gemini".into(),
822 description: None,
823 icon: None,
824 is_latest: false,
825 disabled: None,
826 cost: None,
827 },
828 ]);
829 let favorites = create_favorites(vec!["zed/gemini"]);
830
831 let entries = info_list_to_picker_entries(models, &favorites);
832
833 assert!(matches!(
834 entries.first(),
835 Some(ModelPickerEntry::Separator(s)) if s == "Favorite"
836 ));
837
838 assert!(entries.iter().any(|e| matches!(
839 e,
840 ModelPickerEntry::Separator(s) if s == "All"
841 )));
842 }
843
844 #[gpui::test]
845 fn test_favorites_count_returns_correct_count(_cx: &mut TestAppContext) {
846 let empty_favorites: HashSet<AgentModelId> = HashSet::default();
847 assert_eq!(empty_favorites.len(), 0);
848
849 let one_favorite = create_favorites(vec!["model-a"]);
850 assert_eq!(one_favorite.len(), 1);
851
852 let multiple_favorites = create_favorites(vec!["model-a", "model-b", "model-c"]);
853 assert_eq!(multiple_favorites.len(), 3);
854
855 let with_duplicates = create_favorites(vec!["model-a", "model-a", "model-b"]);
856 assert_eq!(with_duplicates.len(), 2);
857 }
858
859 #[gpui::test]
860 fn test_is_favorite_flag_set_correctly_in_entries(_cx: &mut TestAppContext) {
861 let models = AgentModelList::Flat(vec![
862 acp_thread::AgentModelInfo {
863 id: AgentModelId::new("favorite-model"),
864 name: "Favorite".into(),
865 description: None,
866 icon: None,
867 is_latest: false,
868 disabled: None,
869 cost: None,
870 },
871 acp_thread::AgentModelInfo {
872 id: AgentModelId::new("regular-model"),
873 name: "Regular".into(),
874 description: None,
875 icon: None,
876 is_latest: false,
877 disabled: None,
878 cost: None,
879 },
880 ]);
881 let favorites = create_favorites(vec!["favorite-model"]);
882
883 let entries = info_list_to_picker_entries(models, &favorites);
884
885 for entry in &entries {
886 if let ModelPickerEntry::Model(info, is_favorite) = entry {
887 if info.id.as_ref() == "favorite-model" {
888 assert!(*is_favorite, "favorite-model should have is_favorite=true");
889 } else if info.id.as_ref() == "regular-model" {
890 assert!(!*is_favorite, "regular-model should have is_favorite=false");
891 }
892 }
893 }
894 }
895}
896