use std::time::Duration;
use serde::{Deserialize, Serialize};
pub const VISION_MODEL_PATTERNS: &[&str] = &[
"llama3.2-vision",
"llava",
"bakllava",
"qwen2-vl",
"qwen3-vl",
"gemma3-vl",
"pixtral",
"minicpm-v",
"moondream",
];
pub const DEFAULT_VLM_AUTO_PULL: &str = "llama3.2-vision:11b";
pub const DEFAULT_WHISPER_MODEL: &str = "base";
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct Capabilities {
pub vlm: Option<VlmBackend>,
pub stt: Option<SttBackend>,
pub ocr: Option<OcrBackend>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VlmBackend {
pub endpoint: String,
pub model: String,
pub source: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SttBackend {
pub kind: SttKind,
pub endpoint: String,
pub model: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SttKind {
WhisperCli,
OpenAIApi,
LocalServer,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OcrBackend {
pub binary: String,
}
pub async fn probe() -> Capabilities {
let (vlm, stt, ocr) = tokio::join!(probe_vlm(), probe_stt(), probe_ocr());
Capabilities { vlm, stt, ocr }
}
async fn probe_vlm() -> Option<VlmBackend> {
let endpoint = "http://localhost:11434";
let client = match reqwest::Client::builder()
.timeout(Duration::from_millis(800))
.build()
{
Ok(c) => c,
Err(_) => return None,
};
let resp = client
.get(format!("{endpoint}/api/tags"))
.send()
.await
.ok()?;
if !resp.status().is_success() {
return None;
}
#[derive(Deserialize)]
struct OllamaTags {
models: Vec<OllamaModel>,
}
#[derive(Deserialize)]
struct OllamaModel {
name: String,
}
let body: OllamaTags = resp.json().await.ok()?;
let model = pick_best_vision_model(
&body
.models
.iter()
.map(|m| m.name.as_str())
.collect::<Vec<_>>(),
)?;
Some(VlmBackend {
endpoint: endpoint.to_string(),
model: model.to_string(),
source: "ollama".to_string(),
})
}
pub fn pick_best_vision_model<'a>(model_names: &[&'a str]) -> Option<&'a str> {
for pattern in VISION_MODEL_PATTERNS {
for name in model_names {
if name.contains(pattern) {
return Some(name);
}
}
}
None
}
async fn probe_stt() -> Option<SttBackend> {
if let Some(path) = which("whisper").or_else(|| which("whisper-cli")) {
return Some(SttBackend {
kind: SttKind::WhisperCli,
endpoint: path,
model: Some(DEFAULT_WHISPER_MODEL.to_string()),
});
}
if std::env::var("OPENAI_API_KEY").is_ok() {
return Some(SttBackend {
kind: SttKind::OpenAIApi,
endpoint: "https://api.openai.com/v1/audio/transcriptions".to_string(),
model: Some("whisper-1".to_string()),
});
}
let client = reqwest::Client::builder()
.timeout(Duration::from_millis(300))
.build()
.ok()?;
if client
.get("http://localhost:9000/")
.send()
.await
.map(|r| r.status().is_success() || r.status().as_u16() == 405)
.unwrap_or(false)
{
return Some(SttBackend {
kind: SttKind::LocalServer,
endpoint: "http://localhost:9000/asr".to_string(),
model: None,
});
}
None
}
async fn probe_ocr() -> Option<OcrBackend> {
which("tesseract").map(|binary| OcrBackend { binary })
}
pub fn which(binary: &str) -> Option<String> {
let path = std::env::var_os("PATH")?;
for dir in std::env::split_paths(&path) {
let candidate = dir.join(binary);
if candidate.is_file() {
return Some(candidate.to_string_lossy().into_owned());
}
}
None
}
impl Capabilities {
pub fn any(&self) -> bool {
self.vlm.is_some() || self.stt.is_some() || self.ocr.is_some()
}
pub fn summary(&self) -> String {
let mut parts = Vec::new();
if let Some(v) = &self.vlm {
parts.push(format!("vlm:{} ({})", v.model, v.source));
}
if let Some(s) = &self.stt {
parts.push(format!("stt:{:?}", s.kind));
}
if self.ocr.is_some() {
parts.push("ocr:tesseract".to_string());
}
if parts.is_empty() {
"no backends detected".to_string()
} else {
parts.join(", ")
}
}
pub fn install_hints(&self) -> Vec<String> {
let mut out = Vec::new();
if self.vlm.is_none() {
out.push(
"vlm: install ollama (https://ollama.com/download) and run \
`ollama pull llama3.2-vision:11b` — captchaforge will then \
auto-detect it. Solves canvas / SVG / image-grid CAPTCHAs."
.to_string(),
);
}
if self.stt.is_none() {
out.push(
"stt: install OpenAI Whisper CLI (`pip install -U openai-whisper`) \
OR set OPENAI_API_KEY for the hosted Whisper API. Solves \
reCAPTCHA v2 audio + generic spoken-digit CAPTCHAs."
.to_string(),
);
}
if self.ocr.is_none() {
out.push(
"ocr: install tesseract (`apt install tesseract-ocr` / \
`brew install tesseract`). Used as a fallback when no VLM \
is available."
.to_string(),
);
}
out
}
pub async fn ensure_vlm_model(&mut self) -> anyhow::Result<bool> {
if self.vlm.is_some() {
return Ok(true);
}
let ollama = which("ollama");
let Some(ollama_bin) = ollama else {
return Ok(false);
};
let client = reqwest::Client::builder()
.timeout(Duration::from_millis(800))
.build()?;
if client
.get("http://localhost:11434/api/tags")
.send()
.await
.map(|r| !r.status().is_success())
.unwrap_or(true)
{
return Ok(false);
}
tracing::info!(
model = DEFAULT_VLM_AUTO_PULL,
"auto-pulling vision model via Ollama (first-run install)"
);
let status = std::process::Command::new(&ollama_bin)
.arg("pull")
.arg(DEFAULT_VLM_AUTO_PULL)
.status()?;
if !status.success() {
anyhow::bail!("ollama pull {DEFAULT_VLM_AUTO_PULL} exited {status}");
}
self.vlm = probe_vlm().await;
Ok(self.vlm.is_some())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn vision_model_patterns_are_lowercase() {
for p in VISION_MODEL_PATTERNS {
assert_eq!(*p, p.to_lowercase(), "pattern {p} not lowercase");
}
}
#[test]
fn pick_best_vision_model_returns_none_for_text_only() {
let names = ["mistral:7b", "qwen2.5:14b", "llama3.1:8b"];
assert_eq!(pick_best_vision_model(&names), None);
}
#[test]
fn pick_best_vision_model_finds_llava() {
let names = ["llama3.1:8b", "llava:13b"];
assert_eq!(pick_best_vision_model(&names), Some("llava:13b"));
}
#[test]
fn pick_best_vision_model_priority_order_respected() {
let names = ["llava:13b", "llama3.2-vision:11b"];
assert_eq!(pick_best_vision_model(&names), Some("llama3.2-vision:11b"));
}
#[test]
fn pick_best_vision_model_handles_empty_list() {
assert_eq!(pick_best_vision_model(&[]), None);
}
#[test]
fn capabilities_summary_no_backends() {
let caps = Capabilities::default();
assert_eq!(caps.summary(), "no backends detected");
assert!(!caps.any());
}
#[test]
fn capabilities_summary_lists_each_backend() {
let caps = Capabilities {
vlm: Some(VlmBackend {
endpoint: "http://localhost:11434".to_string(),
model: "llava:13b".to_string(),
source: "ollama".to_string(),
}),
stt: Some(SttBackend {
kind: SttKind::WhisperCli,
endpoint: "/usr/local/bin/whisper".to_string(),
model: Some("tiny".to_string()),
}),
ocr: Some(OcrBackend {
binary: "/usr/bin/tesseract".to_string(),
}),
};
let s = caps.summary();
assert!(s.contains("vlm:llava:13b"));
assert!(s.contains("stt:WhisperCli"));
assert!(s.contains("ocr:tesseract"));
assert!(caps.any());
}
#[test]
fn default_vlm_pull_target_is_a_known_pattern() {
assert!(
VISION_MODEL_PATTERNS
.iter()
.any(|p| DEFAULT_VLM_AUTO_PULL.contains(p)),
"DEFAULT_VLM_AUTO_PULL must match a VISION_MODEL_PATTERNS entry; \
otherwise the post-pull probe won't find the new model",
);
}
#[test]
fn which_finds_existing_binary_unix() {
if cfg!(unix) {
assert!(which("sh").is_some(), "sh should be on PATH");
}
}
#[test]
fn which_returns_none_for_missing_binary() {
assert!(which("definitely-not-a-real-binary-xyzzy12345").is_none());
}
#[tokio::test]
async fn probe_returns_some_capabilities_field_default() {
let caps = probe().await;
let _ = caps.summary();
}
#[test]
fn stt_kind_serializes_snake_case() {
let json = serde_json::to_string(&SttKind::WhisperCli).unwrap();
assert_eq!(json, r#""whisper_cli""#);
let json = serde_json::to_string(&SttKind::OpenAIApi).unwrap();
assert_eq!(json, r#""open_a_i_api""#);
let json = serde_json::to_string(&SttKind::LocalServer).unwrap();
assert_eq!(json, r#""local_server""#);
}
}