vtcode_llm/providers/ollama/
pull.rs1use hashbrown::HashMap;
2use ratatui::crossterm::{
3 cursor::MoveToColumn,
4 execute,
5 terminal::{Clear, ClearType},
6};
7use std::io;
8use std::io::Write;
9
10#[derive(Debug, Clone)]
14pub enum OllamaPullEvent {
15 Status(String),
17 ChunkProgress {
19 digest: String,
20 total: Option<u64>,
21 completed: Option<u64>,
22 },
23 Success,
25 Error(String),
27}
28
29pub trait OllamaPullProgressReporter: Send {
35 fn on_event(&mut self, event: &OllamaPullEvent) -> io::Result<()>;
36}
37
38pub struct CliPullProgressReporter {
40 printed_header: bool,
41 last_line_len: usize,
42 last_completed_sum: u64,
43 last_instant: std::time::Instant,
44 totals_by_digest: HashMap<String, (u64, u64)>,
45}
46
47impl Default for CliPullProgressReporter {
48 fn default() -> Self {
49 Self::new()
50 }
51}
52
53impl CliPullProgressReporter {
54 pub(crate) fn new() -> Self {
55 Self {
56 printed_header: false,
57 last_line_len: 0,
58 last_completed_sum: 0,
59 last_instant: std::time::Instant::now(),
60 totals_by_digest: HashMap::new(),
61 }
62 }
63}
64
65impl OllamaPullProgressReporter for CliPullProgressReporter {
66 fn on_event(&mut self, event: &OllamaPullEvent) -> io::Result<()> {
67 let mut out = io::stderr();
68 match event {
69 OllamaPullEvent::Status(status) => {
70 if status.eq_ignore_ascii_case("pulling manifest") {
72 return Ok(());
73 }
74 let pad = self.last_line_len.saturating_sub(status.len());
75 let line = format!("\r{status}{}", " ".repeat(pad));
76 self.last_line_len = status.len();
77 out.write_all(line.as_bytes())?;
78 out.flush()
79 }
80 OllamaPullEvent::ChunkProgress { digest, total, completed } => {
81 if let Some(t) = total {
82 self.totals_by_digest.entry(digest.clone()).or_insert((0, 0)).0 = *t;
83 }
84 if let Some(c) = completed {
85 self.totals_by_digest.entry(digest.clone()).or_insert((0, 0)).1 = *c;
86 }
87 let (sum_total, sum_completed) = self
88 .totals_by_digest
89 .values()
90 .fold((0u64, 0u64), |acc, (t, c)| (acc.0 + t, acc.1 + c));
91
92 if sum_total > 0 {
93 if !self.printed_header {
94 let gb = (sum_total as f64) / (1024.0 * 1024.0 * 1024.0);
95 let header = format!("Downloading model: total {gb:.2} GB\n");
96 execute!(out, MoveToColumn(0), Clear(ClearType::CurrentLine))?;
97 out.write_all(header.as_bytes())?;
98 self.printed_header = true;
99 }
100 let now = std::time::Instant::now();
101 let dt = now.duration_since(self.last_instant).as_secs_f64().max(0.001);
102 let dbytes = sum_completed.saturating_sub(self.last_completed_sum) as f64;
103 let speed_mb_s = dbytes / (1024.0 * 1024.0) / dt;
104 self.last_completed_sum = sum_completed;
105 self.last_instant = now;
106 let done_gb = (sum_completed as f64) / (1024.0 * 1024.0 * 1024.0);
107 let total_gb = (sum_total as f64) / (1024.0 * 1024.0 * 1024.0);
108 let pct = (sum_completed as f64) * 100.0 / (sum_total as f64);
109 let text = format!("{done_gb:.2}/{total_gb:.2} GB ({pct:.1}%) {speed_mb_s:.1} MB/s");
110 let pad = self.last_line_len.saturating_sub(text.len());
111 let line = format!("\r{text}{}", " ".repeat(pad));
112 self.last_line_len = text.len();
113 out.write_all(line.as_bytes())?;
114 out.flush()
115 } else {
116 Ok(())
117 }
118 }
119 OllamaPullEvent::Error(_) => {
120 Ok(())
123 }
124 OllamaPullEvent::Success => {
125 out.write_all(b"\n")?;
126 out.flush()
127 }
128 }
129 }
130}
131
132#[derive(Default)]
135pub struct TuiPullProgressReporter(CliPullProgressReporter);
136
137impl OllamaPullProgressReporter for TuiPullProgressReporter {
138 fn on_event(&mut self, event: &OllamaPullEvent) -> io::Result<()> {
139 self.0.on_event(event)
140 }
141}
142
143#[cfg(test)]
144mod tests {
145 use super::*;
146
147 #[test]
148 fn cli_reporter_formats_status_messages() {
149 let mut reporter = CliPullProgressReporter::new();
150 let event = OllamaPullEvent::Status("verifying".to_string());
151 let result = reporter.on_event(&event);
152 result.unwrap();
153 }
154
155 #[test]
156 fn cli_reporter_tracks_download_progress() {
157 let mut reporter = CliPullProgressReporter::new();
158 let event = OllamaPullEvent::ChunkProgress {
159 digest: "sha256:abc".to_string(),
160 total: Some(1_000_000_000), completed: Some(500_000_000),
162 };
163 let result = reporter.on_event(&event);
164 result.unwrap();
165 }
166}