Skip to main content

studio_worker/engine/
llama_subprocess.rs

1//! Windows LLM via a subprocess `llama-cli` (mirrors the `sd-cli`
2//! pattern) so Windows workers reach in-process-llama parity.
3//!
4//! `llama-cpp-2` (the in-process backend) doesn't link on Windows MSVC
5//! (a static-vs-dynamic CRT clash, documented in `Cargo.toml`), so a
6//! Windows release worker would otherwise fall back to the synthetic
7//! LLM.  Instead we auto-provision the official `llama.cpp` Windows
8//! Vulkan release binary into `<models_root>/bin/` on first use and run
9//! `llama-cli` per request, exactly like the image engine runs
10//! `sd-cli`.
11//!
12//! The module is always compiled (so its pure argv / response logic and
13//! the provisioner are unit-tested on every platform), but only
14//! *registered* on Windows — Linux/macOS keep the faster in-process
15//! `llama-cpp-2` backend.
16
17use std::path::{Path, PathBuf};
18
19use anyhow::{bail, Context, Result};
20use tracing::info;
21
22use super::download;
23use crate::types::LlmParams;
24
25// The engine struct + dispatch are only *registered* on Windows, but we
26// also compile them in `test` builds on every platform so the dispatch
27// path is genuinely type-checked (not just the pure helpers).  The
28// pure argv/response helpers above them need none of these.
29#[cfg(any(target_os = "windows", test))]
30use crate::types::{ModelFileRole, ModelSource, Task, TaskKind, TaskResult};
31#[cfg(any(target_os = "windows", test))]
32use tracing::warn;
33
34const TRACE_TARGET: &str = "studio_worker::engine::llama_subprocess";
35
36/// Pinned llama.cpp release build.  Bumping this must keep the asset
37/// naming (`llama-<build>-bin-win-vulkan-x64.zip`) in step; the
38/// `pinned-assets-drift` CI job HEADs it.
39const DEFAULT_BUILD: &str = "b6414";
40const BUILD_ENV: &str = "STUDIO_WORKER_LLAMA_BUILD";
41const URL_ENV: &str = "STUDIO_WORKER_LLAMA_URL";
42
43/// The `llama-cli` executable name for this platform.
44fn binary_name() -> &'static str {
45    if cfg!(target_os = "windows") {
46        "llama-cli.exe"
47    } else {
48        "llama-cli"
49    }
50}
51
52/// Resolve the release build tag: `STUDIO_WORKER_LLAMA_BUILD` wins, else
53/// the pinned default.
54pub fn select_build(env_value: Option<&str>) -> String {
55    match env_value.map(str::trim).filter(|s| !s.is_empty()) {
56        Some(b) => b.to_string(),
57        None => DEFAULT_BUILD.to_string(),
58    }
59}
60
61/// The Windows-Vulkan release asset name for `build`.
62pub fn asset_name(build: &str) -> String {
63    format!("llama-{build}-bin-win-vulkan-x64.zip")
64}
65
66/// The full download URL: a `STUDIO_WORKER_LLAMA_URL` override wins
67/// (air-gapped mirror / tests), else the pinned GitHub release asset.
68pub fn resolve_url(build_env: Option<&str>, url_env: Option<&str>) -> String {
69    if let Some(url) = url_env.map(str::trim).filter(|s| !s.is_empty()) {
70        return url.to_string();
71    }
72    let build = select_build(build_env);
73    format!(
74        "https://github.com/ggml-org/llama.cpp/releases/download/{build}/{}",
75        asset_name(&build)
76    )
77}
78
79/// Build the `llama-cli` argument vector for a chat completion.  Pure so
80/// the flag set is unit-tested without a binary.  `--no-display-prompt`
81/// keeps the echoed prompt out of stdout so only the completion is
82/// captured; `-st` (single-turn) + `-no-cnv` runs one non-interactive
83/// turn and exits.
84pub fn build_argv(model_path: &Path, prompt: &str, params: &LlmParams) -> Vec<String> {
85    let mut argv = vec![
86        "-m".to_string(),
87        model_path.display().to_string(),
88        "-p".to_string(),
89        prompt.to_string(),
90        "-n".to_string(),
91        params.max_tokens.to_string(),
92        "--temp".to_string(),
93        params.temperature.to_string(),
94        "--no-display-prompt".to_string(),
95        "-no-cnv".to_string(),
96    ];
97    if let Some(top_p) = params.top_p {
98        argv.push("--top-p".to_string());
99        argv.push(top_p.to_string());
100    }
101    argv
102}
103
104/// Wrap raw `llama-cli` stdout in an OpenAI `chat.completion` object so
105/// the local API / studio consumers parse it uniformly with the
106/// synthetic + in-process engines.  Pure.
107pub fn wrap_response(stdout: &str, model_id: &str, prompt: &str) -> serde_json::Value {
108    let content = stdout.trim();
109    serde_json::json!({
110        "object": "chat.completion",
111        "model": model_id,
112        "choices": [{
113            "index": 0,
114            "message": { "role": "assistant", "content": content },
115            "finish_reason": "stop",
116        }],
117        "usage": {
118            "prompt_tokens": prompt.split_whitespace().count(),
119            "completion_tokens": content.split_whitespace().count(),
120            "total_tokens": prompt.split_whitespace().count() + content.split_whitespace().count(),
121        },
122    })
123}
124
125/// The last non-empty line of `prompt_for`-style chat input: the user
126/// turn we feed llama-cli.  (llama-cli takes a single `-p` string; a
127/// full chat template is a future enhancement.)
128pub fn prompt_from_params(params: &LlmParams) -> String {
129    params
130        .messages
131        .last()
132        .map(|m| m.content.clone())
133        .unwrap_or_default()
134}
135
136/// Extract every file from the release zip, flattened into `dest_dir`
137/// (defusing zip-slip by keeping only base file names).  Mirrors the
138/// sd-cli provisioner's extractor.
139fn extract_zip(zip_path: &Path, dest_dir: &Path) -> Result<usize> {
140    let file =
141        std::fs::File::open(zip_path).with_context(|| format!("opening {}", zip_path.display()))?;
142    let mut archive = zip::ZipArchive::new(file)
143        .with_context(|| format!("reading zip {}", zip_path.display()))?;
144    std::fs::create_dir_all(dest_dir)
145        .with_context(|| format!("creating {}", dest_dir.display()))?;
146    let mut written = 0usize;
147    for i in 0..archive.len() {
148        let mut entry = archive.by_index(i)?;
149        if entry.is_dir() {
150            continue;
151        }
152        let Some(name) = Path::new(entry.name()).file_name().map(|n| n.to_owned()) else {
153            continue;
154        };
155        let out = dest_dir.join(&name);
156        let mut writer =
157            std::fs::File::create(&out).with_context(|| format!("creating {}", out.display()))?;
158        std::io::copy(&mut entry, &mut writer)
159            .with_context(|| format!("writing {}", out.display()))?;
160        written += 1;
161    }
162    Ok(written)
163}
164
165/// Ensure `llama-cli` is present under `<models_root>/bin/`,
166/// downloading + extracting the pinned release on first use.  Excluded
167/// from coverage: the happy path needs a real multi-hundred-MB download;
168/// the URL/asset/argv/response logic is unit-tested and the extractor is
169/// fixture-tested.
170#[cfg_attr(coverage_nightly, coverage(off))]
171pub fn provision(models_root: &Path) -> Result<PathBuf> {
172    let bin_dir = models_root.join("bin");
173    let binary = bin_dir.join(binary_name());
174    if binary.is_file() {
175        return Ok(binary);
176    }
177    let url = resolve_url(
178        std::env::var(BUILD_ENV).ok().as_deref(),
179        std::env::var(URL_ENV).ok().as_deref(),
180    );
181    info!(
182        target: TRACE_TARGET,
183        op = "provision",
184        url = %url,
185        dest = %bin_dir.display(),
186        "llama-cli not found; provisioning llama.cpp"
187    );
188    std::fs::create_dir_all(models_root)
189        .with_context(|| format!("creating {}", models_root.display()))?;
190    let zip_path = models_root.join(format!(".llama-cli-{}.zip", std::process::id()));
191    let result = (|| -> Result<PathBuf> {
192        download::download_file(&url, &zip_path)?;
193        extract_zip(&zip_path, &bin_dir)?;
194        if !binary.is_file() {
195            bail!(
196                "llama.cpp release {url} did not contain {} after extraction",
197                binary_name()
198            );
199        }
200        Ok(binary.clone())
201    })();
202    download::remove_temp_file(&zip_path);
203    result
204}
205
206/// Subprocess-`llama-cli` LLM engine (Windows).
207#[cfg(any(target_os = "windows", test))]
208pub struct LlamaSubprocessEngine {
209    models_root: PathBuf,
210}
211
212#[cfg(any(target_os = "windows", test))]
213#[allow(dead_code)] // run_chat / ensure_model are the live Windows path
214impl LlamaSubprocessEngine {
215    pub fn new(models_root: PathBuf) -> Self {
216        Self { models_root }
217    }
218
219    /// Resolve the GGUF weights from the offer's `ModelSource` (role
220    /// `model` or `diffusion-model` fallback), downloading on first use.
221    #[cfg_attr(coverage_nightly, coverage(off))]
222    fn ensure_model(&self, model: &str, source: &ModelSource) -> Result<PathBuf> {
223        let file = source
224            .files
225            .iter()
226            .find(|f| matches!(f.role, ModelFileRole::Model))
227            .ok_or_else(|| {
228                anyhow::anyhow!("llama modelSource has no `model` file (the .gguf weights)")
229            })?;
230        download::ensure_file_for_model(&self.models_root, model, file)
231    }
232
233    #[cfg_attr(coverage_nightly, coverage(off))]
234    fn run_chat(&self, model: &str, params: LlmParams, source: &ModelSource) -> Result<TaskResult> {
235        let binary = provision(&self.models_root)?;
236        let model_path = self.ensure_model(model, source)?;
237        let prompt = prompt_from_params(&params);
238        let argv = build_argv(&model_path, &prompt, &params);
239        let output = std::process::Command::new(&binary)
240            .args(&argv)
241            .output()
242            .with_context(|| format!("spawning {}", binary.display()))?;
243        if !output.status.success() {
244            let stderr = String::from_utf8_lossy(&output.stderr);
245            let last = stderr.lines().last().unwrap_or("");
246            warn!(
247                target: TRACE_TARGET,
248                op = "dispatch",
249                model,
250                code = ?output.status.code(),
251                "llama-cli failed: {last}"
252            );
253            bail!("llama-cli exited {:?}: {last}", output.status.code());
254        }
255        let stdout = String::from_utf8_lossy(&output.stdout);
256        Ok(TaskResult::Llm {
257            json: wrap_response(&stdout, model, &prompt),
258        })
259    }
260}
261
262#[cfg(any(target_os = "windows", test))]
263impl super::Engine for LlamaSubprocessEngine {
264    fn name(&self) -> &'static str {
265        "llama-subprocess"
266    }
267
268    fn capabilities(&self) -> super::EngineCapabilities {
269        let mut per_kind = std::collections::BTreeMap::new();
270        // Kind-based selection on the studio side; the sentinel model
271        // name is informational (mirrors sdcpp's `sd-cpp:*`).
272        per_kind.insert(TaskKind::Llm, vec!["llama-cpp:*".to_string()]);
273        super::EngineCapabilities {
274            supported_models_per_kind: per_kind,
275        }
276    }
277
278    fn dispatch(&self, _model: &str, _task: Task) -> Result<TaskResult> {
279        bail!("llama-subprocess requires a ModelSource on the offer")
280    }
281
282    #[cfg_attr(coverage_nightly, coverage(off))]
283    fn dispatch_with_source(
284        &self,
285        model: &str,
286        task: Task,
287        source: &ModelSource,
288    ) -> Result<TaskResult> {
289        match task {
290            Task::Llm(p) => self.run_chat(model, p, source),
291            other => Err(super::UnsupportedTask::new("llama-subprocess", other.kind()).into()),
292        }
293    }
294}
295
296#[cfg(test)]
297mod tests {
298    use super::*;
299    use crate::types::ChatMessage;
300
301    #[test]
302    fn select_build_prefers_env_then_default() {
303        assert_eq!(select_build(Some("b9999")), "b9999");
304        assert_eq!(select_build(Some("  ")), DEFAULT_BUILD);
305        assert_eq!(select_build(None), DEFAULT_BUILD);
306    }
307
308    #[test]
309    fn asset_and_url_follow_the_pinned_naming() {
310        assert_eq!(asset_name("b6414"), "llama-b6414-bin-win-vulkan-x64.zip");
311        let url = resolve_url(None, None);
312        assert!(url.contains("ggml-org/llama.cpp/releases/download/"));
313        assert!(url.ends_with("llama-b6414-bin-win-vulkan-x64.zip"));
314        // Env URL override wins verbatim.
315        assert_eq!(
316            resolve_url(None, Some("http://127.0.0.1/x.zip")),
317            "http://127.0.0.1/x.zip"
318        );
319        // Build override flows into the URL.
320        assert!(resolve_url(Some("b7000"), None).ends_with("llama-b7000-bin-win-vulkan-x64.zip"));
321    }
322
323    #[test]
324    fn build_argv_carries_prompt_tokens_and_temp() {
325        let params = LlmParams {
326            max_tokens: 128,
327            temperature: 0.3,
328            top_p: Some(0.9),
329            ..Default::default()
330        };
331        let argv = build_argv(Path::new("/m/model.gguf"), "hello world", &params);
332        // Model + prompt + token budget + temp + top-p are all present.
333        assert!(argv.windows(2).any(|w| w == ["-m", "/m/model.gguf"]));
334        assert!(argv.windows(2).any(|w| w == ["-p", "hello world"]));
335        assert!(argv.windows(2).any(|w| w == ["-n", "128"]));
336        assert!(argv.windows(2).any(|w| w == ["--temp", "0.3"]));
337        assert!(argv.windows(2).any(|w| w == ["--top-p", "0.9"]));
338        // Single-turn, no prompt echo.
339        assert!(argv.iter().any(|a| a == "--no-display-prompt"));
340        assert!(argv.iter().any(|a| a == "-no-cnv"));
341    }
342
343    #[test]
344    fn build_argv_omits_top_p_when_unset() {
345        let argv = build_argv(Path::new("/m.gguf"), "x", &LlmParams::default());
346        assert!(!argv.iter().any(|a| a == "--top-p"));
347    }
348
349    #[test]
350    fn wrap_response_produces_a_chat_completion() {
351        let json = wrap_response("  the answer is 42  ", "my-llm", "what is it");
352        assert_eq!(json["object"], "chat.completion");
353        assert_eq!(json["model"], "my-llm");
354        assert_eq!(json["choices"][0]["message"]["role"], "assistant");
355        // stdout is trimmed into the content.
356        assert_eq!(json["choices"][0]["message"]["content"], "the answer is 42");
357        assert_eq!(json["choices"][0]["finish_reason"], "stop");
358        assert_eq!(json["usage"]["completion_tokens"], 4);
359    }
360
361    #[test]
362    fn prompt_from_params_takes_the_last_message() {
363        let params = LlmParams {
364            messages: vec![
365                ChatMessage {
366                    role: "system".into(),
367                    content: "be terse".into(),
368                },
369                ChatMessage {
370                    role: "user".into(),
371                    content: "ping".into(),
372                },
373            ],
374            ..Default::default()
375        };
376        assert_eq!(prompt_from_params(&params), "ping");
377        assert_eq!(prompt_from_params(&LlmParams::default()), "");
378    }
379
380    #[test]
381    fn extract_zip_flattens_entries_into_the_bin_dir() {
382        // Build a fixture zip with a nested path; extraction must flatten
383        // it to the base name (zip-slip defence) next to a sibling.
384        let dir = tempfile::tempdir().unwrap();
385        let zip_path = dir.path().join("llama.zip");
386        {
387            use std::io::Write as _;
388            let f = std::fs::File::create(&zip_path).unwrap();
389            let mut zw = zip::ZipWriter::new(f);
390            let opts: zip::write::FileOptions<'_, ()> = zip::write::FileOptions::default();
391            zw.start_file("build/bin/llama-cli.exe", opts).unwrap();
392            zw.write_all(b"MZ fake binary").unwrap();
393            zw.start_file("build/bin/ggml.dll", opts).unwrap();
394            zw.write_all(b"dll").unwrap();
395            zw.finish().unwrap();
396        }
397        let bin = dir.path().join("bin");
398        let n = extract_zip(&zip_path, &bin).unwrap();
399        assert_eq!(n, 2);
400        assert!(bin.join("llama-cli.exe").is_file());
401        assert!(bin.join("ggml.dll").is_file());
402        // Nothing escaped into a nested subdir.
403        assert!(!bin.join("build").exists());
404    }
405
406    #[test]
407    fn engine_advertises_only_the_llm_kind() {
408        use super::super::Engine as _;
409        let engine = LlamaSubprocessEngine::new(PathBuf::from("/models"));
410        let caps = engine.capabilities();
411        assert_eq!(caps.kinds(), vec![TaskKind::Llm]);
412        assert_eq!(engine.name(), "llama-subprocess");
413        // Non-LLM tasks are rejected as unsupported.
414        let err = engine
415            .dispatch_with_source(
416                "m",
417                Task::Image(crate::types::ImageParams::default()),
418                &ModelSource {
419                    engine: crate::types::ModelEngine::Synthetic,
420                    files: vec![],
421                    cli_defaults: crate::types::ModelCliDefaults::default(),
422                },
423            )
424            .unwrap_err();
425        assert!(err.to_string().contains("cannot serve"));
426    }
427}