Skip to main content

chatty_rs/app/ui/
models.rs

1use std::collections::{BTreeMap, HashMap};
2
3use crate::{
4    config::{self},
5    info_event,
6    models::{Event, Model},
7};
8use ratatui::{
9    Frame,
10    layout::{Alignment, Rect},
11    style::{Color, Modifier, Style, Stylize},
12    text::{Line, Span, Text},
13    widgets::{Block, BorderType, Borders, Clear, List, ListItem, ListState, Padding},
14};
15use ratatui_macros::span;
16use tokio::sync::mpsc;
17use tui_textarea::Key;
18
19use super::{
20    input_box::{self, InputBox},
21    utils,
22};
23
24pub struct ModelsScreen<'a> {
25    event_tx: mpsc::UnboundedSender<Event>,
26
27    showing: bool,
28    models: Vec<Model>,
29    idx_map: HashMap<usize, String>,
30
31    current_model: String,
32    state: ListState,
33    items: Vec<ListItem<'a>>,
34
35    last_known_width: usize,
36
37    search: InputBox<'a>,
38    current_search: String,
39}
40
41impl<'a> ModelsScreen<'a> {
42    pub fn new(models: Vec<Model>, event_tx: mpsc::UnboundedSender<Event>) -> ModelsScreen<'a> {
43        let want_model = config::instance()
44            .backend
45            .default_model
46            .as_deref()
47            .unwrap_or_default();
48
49        let default_model = models
50            .iter()
51            .find(|model| want_model == model.id())
52            .unwrap_or_else(|| &models[0])
53            .id()
54            .to_string();
55
56        ModelsScreen {
57            event_tx,
58            showing: false,
59            models,
60            current_model: default_model,
61            search: InputBox::default().with_title(" Search "),
62            current_search: String::new(),
63            last_known_width: 0,
64
65            state: ListState::default(),
66            idx_map: HashMap::new(),
67            items: vec![],
68        }
69    }
70
71    pub fn current_model(&self) -> &str {
72        &self.current_model
73    }
74
75    pub fn set_current_model(&mut self, model: &str) {
76        if self.current_model == model {
77            return;
78        }
79        self.current_model = model.to_string();
80
81        let _ = self
82            .event_tx
83            .send(info_event!(format!("Model changed to \"{}\"", model)));
84
85        self.build_items();
86        self.set_cursor_to_selected();
87    }
88
89    pub fn showing(&self) -> bool {
90        self.showing
91    }
92
93    pub fn toggle_showing(&mut self) {
94        self.showing = !self.showing;
95    }
96
97    fn next_row(&mut self) {
98        if self.models.is_empty() {
99            self.state.select(None);
100            return;
101        }
102
103        let i = match self.state.selected() {
104            Some(i) => (i + 1).min(self.items.len() - 1),
105            None => 0,
106        };
107        self.state.select(Some(i));
108    }
109
110    fn prev_row(&mut self) {
111        if self.models.is_empty() {
112            self.state.select(None);
113            return;
114        }
115
116        let i = match self.state.selected() {
117            Some(i) => (i as isize - 1).max(0) as usize,
118            None => 0,
119        };
120        self.state.select(Some(i));
121    }
122
123    fn first(&mut self) {
124        if self.models.is_empty() {
125            self.state.select(None);
126            return;
127        }
128        self.state.select(Some(0));
129        // if the first item is a group header, we need to select the next item
130        self.next_row();
131    }
132
133    fn last(&mut self) {
134        if self.models.is_empty() {
135            self.state.select(None);
136            return;
137        }
138        self.state.select(Some(self.items.len() - 1));
139    }
140
141    fn request_change_model(&mut self) -> bool {
142        let index = self.state.selected().unwrap_or(0);
143        if index >= self.models.len() {
144            return false;
145        }
146
147        let model = match self.idx_map.get(&index) {
148            Some(idx) => idx,
149            None => return false,
150        };
151
152        if self.current_model == *model {
153            return false;
154        }
155
156        let model = model.to_string();
157        self.set_current_model(&model);
158
159        true
160    }
161
162    fn set_cursor_to_selected(&mut self) {
163        if let Some(item) = self
164            .idx_map
165            .iter()
166            .find(|(_, model)| **model == self.current_model)
167        {
168            self.state.select(Some(*item.0));
169        }
170    }
171
172    pub fn render(&mut self, f: &mut Frame, area: Rect) {
173        if !self.showing {
174            return;
175        }
176
177        let instructions = vec![
178            " ".into(),
179            span!("q").green().bold(),
180            span!(" to close, ").white(),
181            span!("Enter").green().bold(),
182            span!(" to select, ").white(),
183            span!("/").green().bold(),
184            span!(" to search ").white(),
185        ];
186
187        let block = Block::default()
188            .borders(Borders::ALL)
189            .border_type(BorderType::Rounded)
190            .border_style(Style::default().fg(Color::LightBlue))
191            .padding(Padding::symmetric(1, 0))
192            .title(Line::from(" Models ").bold())
193            .title_alignment(Alignment::Center)
194            .title_bottom(Line::from(instructions))
195            .style(Style::default());
196        f.render_widget(Clear, area);
197
198        let inner = block.inner(area);
199
200        if self.last_known_width != inner.width as usize {
201            let prev_width = self.last_known_width;
202
203            self.last_known_width = inner.width as usize;
204            self.build_items();
205            // Only set the cursor to the selected item at the first time
206            if prev_width == 0 {
207                self.set_cursor_to_selected();
208            }
209        }
210
211        let list = List::new(self.items.clone())
212            .block(block)
213            .highlight_style(Style::default().add_modifier(Modifier::REVERSED));
214        f.render_stateful_widget(list, area, &mut self.state);
215
216        let search_area = input_box::build_area(inner, ((inner.width as f32 * 0.9).ceil()) as u16);
217        self.search.render(f, search_area);
218    }
219
220    pub async fn handle_key_event(&mut self, event: &Event) -> bool {
221        if self.search.showing() {
222            self.handle_search_popup(event).await;
223            return false;
224        }
225
226        match event {
227            Event::KeyboardCtrlL => {
228                self.showing = !self.showing;
229            }
230
231            Event::Quit => {
232                self.showing = false;
233                return true;
234            }
235
236            Event::KeyboardEnter => {
237                self.showing = !self.request_change_model();
238            }
239
240            Event::KeyboardCharInput(input) => match input.key {
241                Key::Char('j') => self.next_row(),
242                Key::Char('k') => self.prev_row(),
243                Key::Char('g') => self.first(),
244                Key::Char('G') => self.last(),
245                Key::Char('/') => self.search.open(&self.current_search),
246                Key::Char('q') => {
247                    self.showing = false;
248                }
249                _ => {}
250            },
251
252            Event::UiScrollDown => self.next_row(),
253            Event::UiScrollUp => self.prev_row(),
254            _ => {}
255        }
256
257        false
258    }
259
260    async fn handle_search_popup(&mut self, event: &Event) {
261        match event {
262            Event::KeyboardEsc | Event::KeyboardCtrlC => {
263                self.search.close();
264            }
265            Event::KeyboardEnter => {
266                self.current_search = self.search.close().unwrap_or_default();
267                self.build_items();
268                if !self.items.is_empty() {
269                    self.state.select(Some(0));
270                }
271            }
272            _ => self.search.handle_key_event(event),
273        }
274    }
275
276    fn build_items(&mut self) {
277        self.idx_map.clear();
278        self.items.clear();
279
280        let mut models: BTreeMap<String, Vec<String>> = BTreeMap::new();
281
282        self.models
283            .iter()
284            .filter(|model| {
285                if self.current_search.is_empty() {
286                    return true;
287                }
288                model
289                    .id()
290                    .to_lowercase()
291                    .contains(&self.current_search.to_lowercase())
292            })
293            .for_each(|m| {
294                let model = m.id().to_string();
295                let alias = m.provider().to_string();
296                models.entry(alias).or_default().push(model);
297            });
298
299        for (provider, models) in models {
300            self.items.push(header_item(provider));
301
302            for model in models {
303                let mut spans = vec![span!(model)];
304                if self.current_model == model {
305                    spans.push(Span::styled(" ", Style::default()));
306                    spans.push(Span::styled("[*]", Style::default().fg(Color::LightRed)))
307                }
308
309                let lines = utils::split_to_lines(spans, self.last_known_width - 2);
310                self.items.push(ListItem::new(Text::from(lines)));
311                self.idx_map.insert(self.items.len() - 1, model);
312            }
313        }
314    }
315}
316
317fn header_item<'a>(value: String) -> ListItem<'a> {
318    ListItem::new(Text::from(value).alignment(Alignment::Center).bold())
319        .style(
320            Style::default()
321                .fg(Color::Yellow)
322                .bg(Color::Rgb(26, 35, 126)),
323        )
324        .add_modifier(Modifier::BOLD)
325}