Skip to main content

vtcode_llm/providers/ollama/
pull.rs

1use hashbrown::HashMap;
2use ratatui::crossterm::{
3    cursor::MoveToColumn,
4    execute,
5    terminal::{Clear, ClearType},
6};
7use std::io;
8use std::io::Write;
9
10/// Ollama model pull functionality with progress reporting.
11/// Adapted from OpenAI Codex's codex-ollama/src/pull.rs
12/// Events emitted while pulling a model from Ollama.
13#[derive(Debug, Clone)]
14pub enum OllamaPullEvent {
15    /// A human-readable status message (e.g., "verifying", "writing").
16    Status(String),
17    /// Byte-level progress update for a specific layer digest.
18    ChunkProgress {
19        digest: String,
20        total: Option<u64>,
21        completed: Option<u64>,
22    },
23    /// The pull finished successfully.
24    Success,
25    /// Error event with a message.
26    Error(String),
27}
28
29/// A progress reporter for pull operations. Implementations decide how to render progress
30/// (CLI, TUI, logs, etc.).
31///
32/// `Send` is required so pull progress can run inside the `Send` provider futures
33/// (e.g. `ensure_oss_ready` invoked from `OllamaProvider::generate`).
34pub trait OllamaPullProgressReporter: Send {
35    fn on_event(&mut self, event: &OllamaPullEvent) -> io::Result<()>;
36}
37
38/// A minimal CLI reporter that writes inline progress to stderr.
39pub 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                // Avoid noisy manifest messages; otherwise show status inline.
71                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                // This will be handled by the caller, so we don't do anything
121                // here or the error will be printed twice.
122                Ok(())
123            }
124            OllamaPullEvent::Success => {
125                out.write_all(b"\n")?;
126                out.flush()
127            }
128        }
129    }
130}
131
132/// For now the TUI reporter delegates to the CLI reporter. This keeps UI and
133/// CLI behavior aligned until a dedicated TUI integration is implemented.
134#[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), // 1 GB
161            completed: Some(500_000_000),
162        };
163        let result = reporter.on_event(&event);
164        result.unwrap();
165    }
166}