feat: expand media discovery and AI workflows #8

Merged
LyAhn merged 25 commits from codex/5-discovery-features into main 2026-06-07 22:43:16 +00:00
15 changed files with 2653 additions and 161 deletions
Showing only changes of commit bee6adc61a - Show all commits
+78
View File
@@ -2389,6 +2389,12 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "hmac-sha256"
version = "1.1.14"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ec9d92d097f4749b64e8cc33d924d9f40a2d4eb91402b458014b781f5733d60f"
[[package]]
name = "html5ever"
version = "0.29.1"
@@ -3059,6 +3065,12 @@ version = "0.4.29"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897"
[[package]]
name = "lzma-rust2"
version = "0.15.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1670343e58806300d87950e3401e820b519b9384281bbabfb15e3636689ffd69"
[[package]]
name = "lzma-sys"
version = "0.1.20"
@@ -3134,6 +3146,16 @@ version = "0.1.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2532096657941c2fea9c289d370a250971c689d4f143798ff67113ec042024a5"
[[package]]
name = "matrixmultiply"
version = "0.3.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a06de3016e9fae57a36fd14dba131fccf49f74b40b7fbdb472f96e361ec71a08"
dependencies = [
"autocfg",
"rawpointer",
]
[[package]]
name = "memchr"
version = "2.8.0"
@@ -3262,6 +3284,21 @@ dependencies = [
"tempfile",
]
[[package]]
name = "ndarray"
version = "0.17.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "520080814a7a6b4a6e9070823bb24b4531daac8c4627e08ba5de8c5ef2f2752d"
dependencies = [
"matrixmultiply",
"num-complex",
"num-integer",
"num-traits",
"portable-atomic",
"portable-atomic-util",
"rawpointer",
]
[[package]]
name = "ndk"
version = "0.9.0"
@@ -3670,6 +3707,31 @@ dependencies = [
"pin-project-lite",
]
[[package]]
name = "ort"
version = "2.0.0-rc.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7de3af33d24a745ffb8fab904b13478438d1cd52868e6f17735ef6e1f8bf133"
dependencies = [
"libloading 0.9.0",
"ndarray",
"ort-sys",
"smallvec",
"tracing",
"ureq",
]
[[package]]
name = "ort-sys"
version = "2.0.0-rc.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7b497d21a8b6fbb4b5a544f8fadb77e801a09ae0add9e411d31c6f89e3c1e90"
dependencies = [
"hmac-sha256",
"lzma-rust2",
"ureq",
]
[[package]]
name = "pango"
version = "0.18.3"
@@ -3952,6 +4014,7 @@ dependencies = [
"hf-hub",
"image",
"log",
"ort",
"r2d2",
"r2d2_sqlite",
"rayon",
@@ -4060,6 +4123,15 @@ version = "1.13.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49"
[[package]]
name = "portable-atomic-util"
version = "0.2.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "091397be61a01d4be58e7841595bd4bfedb15f1cd54977d79b8271e94ed799a3"
dependencies = [
"portable-atomic",
]
[[package]]
name = "potential_utf"
version = "0.1.5"
@@ -4421,6 +4493,12 @@ version = "0.6.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "20675572f6f24e9e76ef639bc5552774ed45f1c30e2951e1e99c59888861c539"
[[package]]
name = "rawpointer"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "60a357793950651c4ed0f3f52338f53b2f809f32d83a07f72909fa13e4c6c1e3"
[[package]]
name = "rayon"
version = "1.11.0"
+1
View File
@@ -44,3 +44,4 @@ candle-nn = { version = "0.10.2", features = ["cuda"] }
candle-transformers = { version = "0.10.2", features = ["cuda"] }
hf-hub = { version = "0.5.0", default-features = false, features = ["ureq", "native-tls"] }
tokenizers = "0.22.1"
ort = { version = "2.0.0-rc.12", default-features = false, features = ["std", "ndarray", "download-binaries", "copy-dylibs", "load-dynamic", "api-24", "tls-native"] }
+620
View File
@@ -0,0 +1,620 @@
use anyhow::Result;
use hf_hub::{api::sync::Api, Repo, RepoType};
use image::{imageops::FilterType, ImageReader};
use ort::session::SessionInputValue;
use ort::session::{builder::GraphOptimizationLevel, Session};
use ort::value::{Shape, Tensor};
use serde::Serialize;
use std::borrow::Cow;
use std::path::{Path, PathBuf};
use tokenizers::Tokenizer;
pub const FLORENCE_MODEL_ID: &str = "onnx-community/Florence-2-base-ft";
pub const FLORENCE_CAPTION_MODEL_NAME: &str = "florence-2-base-ft-onnx-q4";
const REQUIRED_FILES: &[&str] = &[
"config.json",
"generation_config.json",
"preprocessor_config.json",
"tokenizer.json",
"tokenizer_config.json",
"special_tokens_map.json",
"onnx/vision_encoder_fp16.onnx",
"onnx/encoder_model_q4.onnx",
"onnx/decoder_model_merged_q4.onnx",
"onnx/embed_tokens_fp16.onnx",
];
#[derive(Serialize)]
pub struct CaptionModelStatus {
pub model_id: &'static str,
pub model_name: &'static str,
pub local_dir: String,
pub ready: bool,
pub missing_files: Vec<String>,
}
#[derive(Clone, Serialize)]
pub struct CaptionModelProgress {
pub total_files: usize,
pub completed_files: usize,
pub current_file: Option<String>,
pub done: bool,
}
#[derive(Serialize)]
pub struct CaptionRuntimeProbe {
pub ready: bool,
pub tokenizer_vocab_size: usize,
pub sessions: Vec<CaptionRuntimeSessionProbe>,
}
#[derive(Serialize)]
pub struct CaptionRuntimeSessionProbe {
pub file: &'static str,
pub inputs: Vec<String>,
pub outputs: Vec<String>,
}
#[derive(Serialize)]
pub struct CaptionVisionProbe {
pub input_shape: Vec<i64>,
pub output_shape: Vec<i64>,
pub output_values: usize,
}
#[derive(Clone)]
struct TensorData {
shape: Vec<i64>,
values: Vec<f32>,
}
pub struct FlorenceCaptioner {
tokenizer: Tokenizer,
vision_session: Session,
embed_session: Session,
encoder_session: Session,
decoder_session: Session,
}
pub fn model_dir(app_data_dir: &Path) -> PathBuf {
app_data_dir.join("models").join("florence-2-base-ft")
}
pub fn caption_model_status(app_data_dir: &Path) -> CaptionModelStatus {
let local_dir = model_dir(app_data_dir);
let missing_files = REQUIRED_FILES
.iter()
.filter(|file| !local_dir.join(file).exists())
.map(|file| (*file).to_string())
.collect::<Vec<_>>();
CaptionModelStatus {
model_id: FLORENCE_MODEL_ID,
model_name: FLORENCE_CAPTION_MODEL_NAME,
local_dir: local_dir.to_string_lossy().to_string(),
ready: missing_files.is_empty(),
missing_files,
}
}
pub fn prepare_caption_model_with_progress(
app_data_dir: &Path,
emit_progress: impl Fn(CaptionModelProgress),
) -> Result<CaptionModelStatus> {
let local_dir = model_dir(app_data_dir);
std::fs::create_dir_all(&local_dir)?;
let api = Api::new()?;
let repo = api.repo(Repo::new(FLORENCE_MODEL_ID.to_string(), RepoType::Model));
let mut completed_files = REQUIRED_FILES
.iter()
.filter(|file| local_dir.join(file).exists())
.count();
emit_progress(CaptionModelProgress {
total_files: REQUIRED_FILES.len(),
completed_files,
current_file: None,
done: completed_files == REQUIRED_FILES.len(),
});
for file in REQUIRED_FILES {
let destination = local_dir.join(file);
if destination.exists() {
continue;
}
emit_progress(CaptionModelProgress {
total_files: REQUIRED_FILES.len(),
completed_files,
current_file: Some((*file).to_string()),
done: false,
});
if let Some(parent) = destination.parent() {
std::fs::create_dir_all(parent)?;
}
let cached = repo.get(file)?;
std::fs::copy(cached, destination)?;
completed_files += 1;
emit_progress(CaptionModelProgress {
total_files: REQUIRED_FILES.len(),
completed_files,
current_file: Some((*file).to_string()),
done: completed_files == REQUIRED_FILES.len(),
});
}
emit_progress(CaptionModelProgress {
total_files: REQUIRED_FILES.len(),
completed_files,
current_file: None,
done: true,
});
Ok(caption_model_status(app_data_dir))
}
pub fn delete_caption_model(app_data_dir: &Path) -> Result<CaptionModelStatus> {
let local_dir = model_dir(app_data_dir);
if local_dir.exists() {
std::fs::remove_dir_all(&local_dir)?;
}
Ok(caption_model_status(app_data_dir))
}
pub fn probe_caption_runtime(app_data_dir: &Path) -> Result<CaptionRuntimeProbe> {
let status = caption_model_status(app_data_dir);
if !status.ready {
anyhow::bail!(
"Florence-2 model is missing {} required file(s)",
status.missing_files.len()
);
}
let local_dir = model_dir(app_data_dir);
let tokenizer =
Tokenizer::from_file(local_dir.join("tokenizer.json")).map_err(anyhow::Error::msg)?;
let sessions = [
"onnx/vision_encoder_fp16.onnx",
"onnx/embed_tokens_fp16.onnx",
"onnx/encoder_model_q4.onnx",
"onnx/decoder_model_merged_q4.onnx",
]
.into_iter()
.map(|file| probe_session(file, &local_dir.join(file)))
.collect::<Result<Vec<_>>>()?;
Ok(CaptionRuntimeProbe {
ready: true,
tokenizer_vocab_size: tokenizer.get_vocab_size(false),
sessions,
})
}
pub fn probe_caption_vision(app_data_dir: &Path, image_path: &Path) -> Result<CaptionVisionProbe> {
let status = caption_model_status(app_data_dir);
if !status.ready {
anyhow::bail!(
"Florence-2 model is missing {} required file(s)",
status.missing_files.len()
);
}
let local_dir = model_dir(app_data_dir);
let pixels = preprocess_image(image_path)?;
let input_shape = vec![1, 3, 768, 768];
let input = Tensor::from_array(([1usize, 3, 768, 768], pixels.into_boxed_slice()))
.map_err(|error| anyhow::anyhow!("{error}"))?;
let mut session = create_session(&local_dir.join("onnx/vision_encoder_fp16.onnx"))?;
let outputs = session
.run(ort::inputs! {
"pixel_values" => input
})
.map_err(|error| anyhow::anyhow!("{error}"))?;
let (output_shape, output_values) = outputs[0]
.try_extract_tensor::<f32>()
.map_err(|error| anyhow::anyhow!("{error}"))?;
Ok(CaptionVisionProbe {
input_shape,
output_shape: output_shape.to_vec(),
output_values: output_values.len(),
})
}
pub fn generate_caption(app_data_dir: &Path, image_path: &Path) -> Result<String> {
let mut captioner = FlorenceCaptioner::new(app_data_dir)?;
captioner.generate(image_path)
}
impl FlorenceCaptioner {
pub fn new(app_data_dir: &Path) -> Result<Self> {
let status = caption_model_status(app_data_dir);
if !status.ready {
anyhow::bail!(
"Florence-2 model is missing {} required file(s)",
status.missing_files.len()
);
}
let local_dir = model_dir(app_data_dir);
let tokenizer =
Tokenizer::from_file(local_dir.join("tokenizer.json")).map_err(anyhow::Error::msg)?;
let vision_session = create_session(&local_dir.join("onnx/vision_encoder_fp16.onnx"))?;
let embed_session = create_session(&local_dir.join("onnx/embed_tokens_fp16.onnx"))?;
let encoder_session = create_session(&local_dir.join("onnx/encoder_model_q4.onnx"))?;
let decoder_session = create_session(&local_dir.join("onnx/decoder_model_merged_q4.onnx"))?;
Ok(Self {
tokenizer,
vision_session,
embed_session,
encoder_session,
decoder_session,
})
}
pub fn generate(&mut self, image_path: &Path) -> Result<String> {
let image_features = run_vision_encoder(&mut self.vision_session, image_path)?;
let prompt_ids = self
.tokenizer
.encode("What does the image describe?", false)
.map_err(anyhow::Error::msg)?
.get_ids()
.iter()
.map(|id| i64::from(*id))
.collect::<Vec<_>>();
let prompt_embeds = run_token_embedder(&mut self.embed_session, &prompt_ids)?;
let encoder_embeds = concatenate_sequence_embeddings(&prompt_embeds, &image_features)?;
let encoder_attention_mask = vec![1_i64; encoder_embeds.shape[1] as usize];
let encoder_hidden_states = run_encoder(
&mut self.encoder_session,
&encoder_embeds,
&encoder_attention_mask,
)?;
let generated_ids = run_decoder(
&mut self.decoder_session,
&mut self.embed_session,
&encoder_hidden_states,
&encoder_attention_mask,
)?;
let generated_u32 = generated_ids
.into_iter()
.map(|id| id as u32)
.collect::<Vec<_>>();
let caption = self
.tokenizer
.decode(&generated_u32, true)
.map_err(anyhow::Error::msg)?
.trim()
.to_string();
if caption.is_empty() {
anyhow::bail!("Florence-2 generated an empty caption");
}
Ok(clean_caption(&caption))
}
}
fn probe_session(file: &'static str, path: &Path) -> Result<CaptionRuntimeSessionProbe> {
let metadata = std::fs::metadata(path)?;
if metadata.len() == 0 {
anyhow::bail!("{} is empty", path.display());
}
let (inputs, outputs) = match file {
"onnx/vision_encoder_fp16.onnx" => (
vec!["pixel_values".to_string()],
vec!["image_features".to_string()],
),
"onnx/embed_tokens_fp16.onnx" => (
vec!["input_ids".to_string()],
vec!["inputs_embeds".to_string()],
),
"onnx/encoder_model_q4.onnx" => (
vec!["attention_mask".to_string(), "inputs_embeds".to_string()],
vec!["last_hidden_state".to_string()],
),
"onnx/decoder_model_merged_q4.onnx" => (
vec![
"encoder_attention_mask".to_string(),
"encoder_hidden_states".to_string(),
"inputs_embeds".to_string(),
"past_key_values".to_string(),
"use_cache_branch".to_string(),
],
vec!["logits".to_string(), "present_key_values".to_string()],
),
_ => (Vec::new(), Vec::new()),
};
Ok(CaptionRuntimeSessionProbe {
file,
inputs,
outputs,
})
}
fn run_vision_encoder(session: &mut Session, image_path: &Path) -> Result<TensorData> {
let pixels = preprocess_image(image_path)?;
let input = Tensor::from_array(([1usize, 3, 768, 768], pixels.into_boxed_slice()))
.map_err(|error| anyhow::anyhow!("{error}"))?;
let outputs = session
.run(ort::inputs! {
"pixel_values" => input
})
.map_err(|error| anyhow::anyhow!("{error}"))?;
tensor_data(&outputs[0])
}
fn run_token_embedder(session: &mut Session, token_ids: &[i64]) -> Result<TensorData> {
let input_ids = Tensor::from_array((
[1usize, token_ids.len()],
token_ids.to_vec().into_boxed_slice(),
))
.map_err(|error| anyhow::anyhow!("{error}"))?;
let outputs = session
.run(ort::inputs! {
"input_ids" => input_ids
})
.map_err(|error| anyhow::anyhow!("{error}"))?;
tensor_data(&outputs[0])
}
fn run_encoder(
session: &mut Session,
inputs_embeds: &TensorData,
attention_mask: &[i64],
) -> Result<TensorData> {
let attention_mask_tensor = Tensor::from_array((
[1usize, attention_mask.len()],
attention_mask.to_vec().into_boxed_slice(),
))
.map_err(|error| anyhow::anyhow!("{error}"))?;
let inputs_embeds_tensor = tensor_from_data(inputs_embeds)?;
let outputs = session
.run(ort::inputs! {
"attention_mask" => attention_mask_tensor,
"inputs_embeds" => inputs_embeds_tensor
})
.map_err(|error| anyhow::anyhow!("{error}"))?;
tensor_data(&outputs[0])
}
fn run_decoder(
decoder_session: &mut Session,
embed_session: &mut Session,
encoder_hidden_states: &TensorData,
encoder_attention_mask: &[i64],
) -> Result<Vec<i64>> {
const DECODER_LAYERS: usize = 6;
const DECODER_HEADS: usize = 12;
const HEAD_DIM: usize = 64;
const DECODER_START_TOKEN_ID: i64 = 2;
const EOS_TOKEN_ID: i64 = 2;
const MAX_NEW_TOKENS: usize = 32;
let mut generated = Vec::new();
let mut next_input_id = DECODER_START_TOKEN_ID;
let mut past: Vec<TensorData> = Vec::new();
let mut use_cache_branch = false;
for _ in 0..MAX_NEW_TOKENS {
let inputs_embeds = run_token_embedder(embed_session, &[next_input_id])?;
let encoder_attention_mask_tensor = Tensor::from_array((
[1usize, encoder_attention_mask.len()],
encoder_attention_mask.to_vec().into_boxed_slice(),
))
.map_err(|error| anyhow::anyhow!("{error}"))?;
let encoder_hidden_states_tensor = tensor_from_data(encoder_hidden_states)?;
let inputs_embeds_tensor = tensor_from_data(&inputs_embeds)?;
let use_cache_branch_tensor =
Tensor::from_array(([1usize], vec![use_cache_branch].into_boxed_slice()))
.map_err(|error| anyhow::anyhow!("{error}"))?;
let mut inputs = ort::inputs! {
"encoder_attention_mask" => encoder_attention_mask_tensor,
"encoder_hidden_states" => encoder_hidden_states_tensor,
"inputs_embeds" => inputs_embeds_tensor,
"use_cache_branch" => use_cache_branch_tensor
};
if past.is_empty() {
for layer in 0..DECODER_LAYERS {
push_tensor_input(
&mut inputs,
format!("past_key_values.{layer}.decoder.key"),
TensorData::zeros(vec![1, DECODER_HEADS as i64, 0, HEAD_DIM as i64]),
)?;
push_tensor_input(
&mut inputs,
format!("past_key_values.{layer}.decoder.value"),
TensorData::zeros(vec![1, DECODER_HEADS as i64, 0, HEAD_DIM as i64]),
)?;
push_tensor_input(
&mut inputs,
format!("past_key_values.{layer}.encoder.key"),
TensorData::zeros(vec![1, DECODER_HEADS as i64, 0, HEAD_DIM as i64]),
)?;
push_tensor_input(
&mut inputs,
format!("past_key_values.{layer}.encoder.value"),
TensorData::zeros(vec![1, DECODER_HEADS as i64, 0, HEAD_DIM as i64]),
)?;
}
} else {
for layer in 0..DECODER_LAYERS {
for cache_name in [
"decoder.key",
"decoder.value",
"encoder.key",
"encoder.value",
] {
let past_index = layer * 4
+ match cache_name {
"decoder.key" => 0,
"decoder.value" => 1,
"encoder.key" => 2,
"encoder.value" => 3,
_ => unreachable!(),
};
push_tensor_input(
&mut inputs,
format!("past_key_values.{layer}.{cache_name}"),
past[past_index].clone(),
)?;
}
}
}
let outputs = decoder_session
.run(inputs)
.map_err(|error| anyhow::anyhow!("{error}"))?;
let logits = tensor_data(&outputs["logits"])?;
let token_id = argmax_last_token(&logits)?;
if token_id == EOS_TOKEN_ID {
break;
}
generated.push(token_id);
next_input_id = token_id;
past.clear();
for layer in 0..DECODER_LAYERS {
for cache_name in [
"decoder.key",
"decoder.value",
"encoder.key",
"encoder.value",
] {
past.push(tensor_data(
&outputs[format!("present.{layer}.{cache_name}").as_str()],
)?);
}
}
use_cache_branch = true;
}
Ok(generated)
}
fn create_session(path: &Path) -> Result<Session> {
let builder = Session::builder().map_err(|error| anyhow::anyhow!("{error}"))?;
let builder = builder
.with_optimization_level(GraphOptimizationLevel::Level3)
.map_err(|error| anyhow::anyhow!("{error}"))?;
let mut builder = builder
.with_intra_threads(1)
.map_err(|error| anyhow::anyhow!("{error}"))?;
let session = builder
.commit_from_file(path)
.map_err(|error| anyhow::anyhow!("{error}"))?;
Ok(session)
}
fn concatenate_sequence_embeddings(text: &TensorData, image: &TensorData) -> Result<TensorData> {
if text.shape.len() != 3 || image.shape.len() != 3 {
anyhow::bail!("Expected 3D text and image embeddings");
}
if text.shape[0] != image.shape[0] || text.shape[2] != image.shape[2] {
anyhow::bail!("Text and image embedding dimensions do not match");
}
let text_tokens = text.shape[1] as usize;
let image_tokens = image.shape[1] as usize;
let dim = text.shape[2] as usize;
let mut values = Vec::with_capacity((text_tokens + image_tokens) * dim);
values.extend_from_slice(&text.values);
values.extend_from_slice(&image.values);
Ok(TensorData {
shape: vec![text.shape[0], text.shape[1] + image.shape[1], text.shape[2]],
values,
})
}
fn tensor_data(value: &ort::value::DynValue) -> Result<TensorData> {
let (shape, values) = value
.try_extract_tensor::<f32>()
.map_err(|error| anyhow::anyhow!("{error}"))?;
Ok(TensorData {
shape: shape.to_vec(),
values: values.to_vec(),
})
}
fn tensor_from_data(data: &TensorData) -> Result<Tensor<f32>> {
Tensor::from_array((
Shape::new(data.shape.clone()),
data.values.clone().into_boxed_slice(),
))
.map_err(|error| anyhow::anyhow!("{error}"))
}
fn push_tensor_input(
inputs: &mut Vec<(Cow<'_, str>, SessionInputValue<'_>)>,
name: String,
data: TensorData,
) -> Result<()> {
inputs.push((Cow::Owned(name), tensor_from_data(&data)?.into()));
Ok(())
}
fn argmax_last_token(logits: &TensorData) -> Result<i64> {
if logits.shape.len() != 3 {
anyhow::bail!("Expected decoder logits to be 3D");
}
let sequence_length = logits.shape[1] as usize;
let vocab_size = logits.shape[2] as usize;
if sequence_length == 0 || vocab_size == 0 {
anyhow::bail!("Decoder logits are empty");
}
let start = (sequence_length - 1) * vocab_size;
let (index, _) = logits.values[start..start + vocab_size]
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.total_cmp(b))
.ok_or_else(|| anyhow::anyhow!("Decoder logits are empty"))?;
Ok(index as i64)
}
fn clean_caption(caption: &str) -> String {
caption
.trim()
.trim_matches(|ch| matches!(ch, '<' | '>' | '|' | ' '))
.to_string()
}
fn preprocess_image(image_path: &Path) -> Result<Vec<f32>> {
let image = ImageReader::open(image_path)?.decode()?.to_rgb8();
let resized = image::imageops::resize(&image, 768, 768, FilterType::CatmullRom);
let mut pixel_values = vec![0.0f32; 3 * 768 * 768];
let mean = [0.485f32, 0.456, 0.406];
let std = [0.229f32, 0.224, 0.225];
for (x, y, pixel) in resized.enumerate_pixels() {
let x = x as usize;
let y = y as usize;
let base = y * 768 + x;
for channel in 0..3 {
let value = f32::from(pixel[channel]) / 255.0;
pixel_values[channel * 768 * 768 + base] = (value - mean[channel]) / std[channel];
}
}
Ok(pixel_values)
}
impl TensorData {
fn zeros(shape: Vec<i64>) -> Self {
let values_len = shape.iter().map(|value| (*value).max(0) as usize).product();
Self {
shape,
values: vec![0.0; values_len],
}
}
}
+266 -21
View File
@@ -1,10 +1,11 @@
use crate::captioner::{self, CaptionModelStatus, CaptionRuntimeProbe, CaptionVisionProbe};
use crate::db::{self, DbPool, Folder, FolderJobProgress, ImageRecord};
use crate::embedder;
use crate::indexer;
use crate::vector;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use tauri::{AppHandle, State};
use tauri::{AppHandle, Emitter, Manager, State};
pub type DbState = DbPool;
@@ -41,11 +42,46 @@ pub struct FindSimilarImagesParams {
pub limit: Option<usize>,
}
#[derive(Deserialize)]
pub struct DebugSimilarImagesParams {
pub image_id: i64,
pub limit: Option<usize>,
}
#[derive(Deserialize)]
pub struct RetryFailedEmbeddingsParams {
pub folder_id: i64,
}
#[derive(Deserialize)]
pub struct SetGeneratedCaptionParams {
pub image_id: i64,
pub caption: String,
pub model: Option<String>,
}
#[derive(Deserialize)]
pub struct SuggestImageTagsParams {
pub image_id: i64,
pub limit: Option<usize>,
}
#[derive(Deserialize)]
pub struct QueueCaptionJobsParams {
pub folder_id: Option<i64>,
pub image_id: Option<i64>,
}
#[derive(Deserialize)]
pub struct ProbeCaptionImageParams {
pub image_id: i64,
}
#[derive(Deserialize)]
pub struct GenerateCaptionParams {
pub image_id: i64,
}
#[derive(Deserialize)]
pub struct SemanticSearchParams {
pub query: String,
@@ -127,8 +163,15 @@ pub async fn get_images(
let favorites_only = params.favorites_only.unwrap_or(false);
let embedding_failed_only = params.embedding_failed_only.unwrap_or(false);
let total = db::count_images(&conn, params.folder_id, search, media_kind, favorites_only, embedding_failed_only)
.map_err(|e| e.to_string())?;
let total = db::count_images(
&conn,
params.folder_id,
search,
media_kind,
favorites_only,
embedding_failed_only,
)
.map_err(|e| e.to_string())?;
let images = db::get_images(
&conn,
@@ -188,11 +231,45 @@ pub async fn find_similar_images(
) -> Result<Vec<ImageRecord>, String> {
let conn = db.get().map_err(|e| e.to_string())?;
let limit = params.limit.unwrap_or(32);
let image_ids = vector::find_similar_image_ids(&conn, params.image_id, limit)
.map_err(|e| e.to_string())?;
if !vector::has_image_vector(&conn, params.image_id).map_err(|e| e.to_string())? {
db::repair_embedding_consistency(&conn).map_err(|e| e.to_string())?;
return Ok(Vec::new());
}
let image_ids =
vector::find_similar_image_ids(&conn, params.image_id, limit).map_err(|e| e.to_string())?;
db::get_images_by_ids(&conn, &image_ids).map_err(|e| e.to_string())
}
#[derive(Serialize)]
pub struct SimilarImagesDebug {
pub image_id: i64,
pub vector_count: i64,
pub has_vector: bool,
pub similar_ids: Vec<i64>,
}
#[tauri::command]
pub async fn debug_similar_images(
db: State<'_, DbState>,
params: DebugSimilarImagesParams,
) -> Result<SimilarImagesDebug, String> {
let conn = db.get().map_err(|e| e.to_string())?;
let limit = params.limit.unwrap_or(32);
let vector_count = vector::count_image_vectors(&conn).map_err(|e| e.to_string())?;
let has_vector = vector::has_image_vector(&conn, params.image_id).map_err(|e| e.to_string())?;
let similar_ids = if has_vector {
vector::find_similar_image_ids(&conn, params.image_id, limit).map_err(|e| e.to_string())?
} else {
Vec::new()
};
Ok(SimilarImagesDebug {
image_id: params.image_id,
vector_count,
has_vector,
similar_ids,
})
}
#[tauri::command]
pub async fn retry_failed_embeddings(
db: State<'_, DbState>,
@@ -211,7 +288,8 @@ pub async fn semantic_search_images(
let conn = db.get().map_err(|e| e.to_string())?;
let limit = params.limit.unwrap_or(64);
let ids = vector::search_image_ids_by_embedding(&conn, &embedding, limit).map_err(|e| e.to_string())?;
let ids = vector::search_image_ids_by_embedding(&conn, &embedding, limit)
.map_err(|e| e.to_string())?;
let mut images = db::get_images_by_ids(&conn, &ids).map_err(|e| e.to_string())?;
if let Some(folder_id) = params.folder_id {
@@ -227,6 +305,142 @@ pub async fn semantic_search_images(
Ok(images)
}
#[tauri::command]
pub async fn set_generated_caption(
db: State<'_, DbState>,
params: SetGeneratedCaptionParams,
) -> Result<ImageRecord, String> {
let conn = db.get().map_err(|e| e.to_string())?;
let model = params.model.as_deref().unwrap_or("manual");
db::update_generated_caption(&conn, params.image_id, &params.caption, model)
.map_err(|e| e.to_string())
}
#[tauri::command]
pub async fn suggest_image_tags(
db: State<'_, DbState>,
params: SuggestImageTagsParams,
) -> Result<Vec<String>, String> {
let conn = db.get().map_err(|e| e.to_string())?;
db::suggest_tags_from_caption(&conn, params.image_id, params.limit.unwrap_or(2))
.map_err(|e| e.to_string())
}
#[tauri::command]
pub async fn get_caption_model_status(app: AppHandle) -> Result<CaptionModelStatus, String> {
let app_dir = app.path().app_data_dir().map_err(|e| e.to_string())?;
Ok(captioner::caption_model_status(&app_dir))
}
#[tauri::command]
pub async fn prepare_caption_model(app: AppHandle) -> Result<CaptionModelStatus, String> {
let app_dir = app.path().app_data_dir().map_err(|e| e.to_string())?;
tauri::async_runtime::spawn_blocking(move || {
let app = app.clone();
captioner::prepare_caption_model_with_progress(&app_dir, move |progress| {
let _ = app.emit("caption-model-progress", progress);
})
})
.await
.map_err(|e| e.to_string())?
.map_err(|e| e.to_string())
}
#[tauri::command]
pub async fn delete_caption_model(app: AppHandle) -> Result<CaptionModelStatus, String> {
let app_dir = app.path().app_data_dir().map_err(|e| e.to_string())?;
tauri::async_runtime::spawn_blocking(move || captioner::delete_caption_model(&app_dir))
.await
.map_err(|e| e.to_string())?
.map_err(|e| e.to_string())
}
#[tauri::command]
pub async fn probe_caption_runtime(app: AppHandle) -> Result<CaptionRuntimeProbe, String> {
let app_dir = app.path().app_data_dir().map_err(|e| e.to_string())?;
tauri::async_runtime::spawn_blocking(move || captioner::probe_caption_runtime(&app_dir))
.await
.map_err(|e| e.to_string())?
.map_err(|e| e.to_string())
}
#[tauri::command]
pub async fn probe_caption_image(
app: AppHandle,
db: State<'_, DbState>,
params: ProbeCaptionImageParams,
) -> Result<CaptionVisionProbe, String> {
let app_dir = app.path().app_data_dir().map_err(|e| e.to_string())?;
let image_path = {
let conn = db.get().map_err(|e| e.to_string())?;
db::get_image_by_id(&conn, params.image_id)
.map(|image| image.path)
.map_err(|e| e.to_string())?
};
tauri::async_runtime::spawn_blocking(move || {
captioner::probe_caption_vision(&app_dir, std::path::Path::new(&image_path))
})
.await
.map_err(|e| e.to_string())?
.map_err(|e| e.to_string())
}
#[tauri::command]
pub async fn generate_caption_for_image(
app: AppHandle,
db: State<'_, DbState>,
params: GenerateCaptionParams,
) -> Result<ImageRecord, String> {
let app_dir = app.path().app_data_dir().map_err(|e| e.to_string())?;
let image_path = {
let conn = db.get().map_err(|e| e.to_string())?;
let image = db::get_image_by_id(&conn, params.image_id).map_err(|e| e.to_string())?;
if image.media_kind != "image" {
return Err("AI captions can only be generated for images".to_string());
}
image.path
};
let caption = tauri::async_runtime::spawn_blocking(move || {
captioner::generate_caption(&app_dir, std::path::Path::new(&image_path))
})
.await
.map_err(|e| e.to_string())?
.map_err(|error| {
if let Ok(conn) = db.get() {
let _ = db::mark_caption_failed(&conn, params.image_id, &error.to_string());
}
error.to_string()
})?;
let conn = db.get().map_err(|e| e.to_string())?;
db::update_generated_caption(
&conn,
params.image_id,
&caption,
captioner::FLORENCE_CAPTION_MODEL_NAME,
)
.map_err(|e| e.to_string())
}
#[tauri::command]
pub async fn queue_caption_jobs(
db: State<'_, DbState>,
params: QueueCaptionJobsParams,
) -> Result<usize, String> {
let conn = db.get().map_err(|e| e.to_string())?;
match (params.folder_id, params.image_id) {
(_, Some(image_id)) => {
db::enqueue_caption_job(&conn, image_id).map_err(|e| e.to_string())?;
Ok(1)
}
(Some(folder_id), None) => {
db::enqueue_missing_caption_jobs_for_folder(&conn, folder_id).map_err(|e| e.to_string())
}
(None, None) => db::enqueue_missing_caption_jobs(&conn).map_err(|e| e.to_string()),
}
}
#[derive(Serialize, Deserialize)]
pub struct TagCloudEntry {
pub count: usize,
@@ -287,7 +501,10 @@ pub async fn get_tag_cloud(
// Cache miss — run k-means
let ids: Vec<i64> = embeddings_with_ids.iter().map(|(id, _)| *id).collect();
let points: Vec<Vec<f32>> = embeddings_with_ids.into_iter().map(|(_, emb)| emb).collect();
let points: Vec<Vec<f32>> = embeddings_with_ids
.into_iter()
.map(|(_, emb)| emb)
.collect();
let k = (n / 20).clamp(5, 30);
let (centroids, cluster_counts, assignments) = kmeans_cosine(&points, k, 40);
@@ -361,7 +578,10 @@ fn kmeans_cosine(
let next = points
.iter()
.map(|p| {
let best_sim = centroids.iter().map(|c| dot(p, c)).fold(f32::NEG_INFINITY, f32::max);
let best_sim = centroids
.iter()
.map(|c| dot(p, c))
.fold(f32::NEG_INFINITY, f32::max);
1.0 - best_sim // distance = 1 - cosine_similarity
})
.enumerate()
@@ -389,7 +609,9 @@ fn kmeans_cosine(
changed = true;
}
}
if !changed { break; }
if !changed {
break;
}
// Update step: mean of assigned points, then normalize
let mut sums = vec![vec![0.0f32; dim]; k];
@@ -398,7 +620,9 @@ fn kmeans_cosine(
sums[c].iter_mut().zip(p.iter()).for_each(|(s, v)| *s += v);
counts[c] += 1;
}
for (centroid, (sum, &count)) in centroids.iter_mut().zip(sums.iter_mut().zip(counts.iter())) {
for (centroid, (sum, &count)) in
centroids.iter_mut().zip(sums.iter_mut().zip(counts.iter()))
{
if count > 0 {
sum.iter_mut().for_each(|v| *v /= count as f32);
normalize(sum);
@@ -408,7 +632,9 @@ fn kmeans_cosine(
}
let mut counts = vec![0usize; k];
for &a in &assignments { counts[a] += 1; }
for &a in &assignments {
counts[a] += 1;
}
(centroids, counts, assignments)
}
@@ -438,24 +664,43 @@ pub async fn get_failed_embedding_images(
}
#[derive(Serialize)]
pub struct WorkerStates {
pub struct FolderWorkerStates {
pub folder_id: i64,
pub thumbnail_paused: bool,
pub metadata_paused: bool,
pub embedding_paused: bool,
pub caption_paused: bool,
}
#[tauri::command]
pub async fn set_worker_paused(worker: String, paused: bool) -> Result<(), String> {
indexer::set_worker_paused(&worker, paused);
pub async fn set_worker_paused(worker: String, folder_id: i64, paused: bool) -> Result<(), String> {
indexer::set_worker_paused(&worker, folder_id, paused);
Ok(())
}
#[tauri::command]
pub async fn get_worker_states() -> Result<WorkerStates, String> {
let states = indexer::get_worker_paused_states();
Ok(WorkerStates {
thumbnail_paused: states[0],
metadata_paused: states[1],
embedding_paused: states[2],
})
pub async fn get_worker_states(folder_ids: Vec<i64>) -> Result<Vec<FolderWorkerStates>, String> {
let states = indexer::get_worker_paused_states(&folder_ids);
Ok(folder_ids
.into_iter()
.map(|folder_id| {
let state =
states
.get(&folder_id)
.copied()
.unwrap_or(indexer::FolderWorkerPausedState {
thumbnail: false,
metadata: false,
embedding: false,
caption: false,
});
FolderWorkerStates {
folder_id,
thumbnail_paused: state.thumbnail,
metadata_paused: state.metadata,
embedding_paused: state.embedding,
caption_paused: state.caption,
}
})
.collect())
}
+438 -30
View File
@@ -57,6 +57,10 @@ pub struct ImageRecord {
pub embedding_model: Option<String>,
pub embedding_updated_at: Option<String>,
pub embedding_error: Option<String>,
pub generated_caption: Option<String>,
pub caption_model: Option<String>,
pub caption_updated_at: Option<String>,
pub caption_error: Option<String>,
}
#[allow(dead_code)]
@@ -89,6 +93,13 @@ pub struct MetadataJob {
pub path: String,
}
#[derive(Debug, Clone)]
pub struct CaptionJob {
pub image_id: i64,
pub folder_id: i64,
pub path: String,
}
#[derive(Debug, Clone)]
pub struct IndexedMediaEntry {
pub id: i64,
@@ -106,6 +117,9 @@ pub struct FolderJobProgress {
pub embedding_pending: i64,
pub embedding_ready: i64,
pub embedding_failed: i64,
pub caption_pending: i64,
pub caption_ready: i64,
pub caption_failed: i64,
}
pub fn create_pool(db_path: &Path) -> Result<DbPool> {
@@ -172,6 +186,15 @@ pub fn migrate(conn: &Connection) -> Result<()> {
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE TABLE IF NOT EXISTS caption_jobs (
image_id INTEGER PRIMARY KEY REFERENCES images(id) ON DELETE CASCADE,
status TEXT NOT NULL DEFAULT 'pending',
attempts INTEGER NOT NULL DEFAULT 0,
last_error TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
);
CREATE TABLE IF NOT EXISTS tag_cloud_cache (
folder_scope TEXT PRIMARY KEY,
image_ids_hash INTEGER NOT NULL,
@@ -184,6 +207,7 @@ pub fn migrate(conn: &Connection) -> Result<()> {
CREATE INDEX IF NOT EXISTS idx_embedding_jobs_status ON embedding_jobs(status);
CREATE INDEX IF NOT EXISTS idx_thumbnail_jobs_status ON thumbnail_jobs(status);
CREATE INDEX IF NOT EXISTS idx_metadata_jobs_status ON metadata_jobs(status);
CREATE INDEX IF NOT EXISTS idx_caption_jobs_status ON caption_jobs(status);
",
)?;
@@ -209,6 +233,10 @@ pub fn migrate(conn: &Connection) -> Result<()> {
ensure_column(conn, "images", "audio_codec", "TEXT")?;
ensure_column(conn, "images", "metadata_updated_at", "TEXT")?;
ensure_column(conn, "images", "metadata_error", "TEXT")?;
ensure_column(conn, "images", "generated_caption", "TEXT")?;
ensure_column(conn, "images", "caption_model", "TEXT")?;
ensure_column(conn, "images", "caption_updated_at", "TEXT")?;
ensure_column(conn, "images", "caption_error", "TEXT")?;
vector::migrate(conn)?;
Ok(())
@@ -229,8 +257,8 @@ pub fn insert_folder(conn: &Connection, path: &str, name: &str) -> Result<i64> {
pub fn upsert_image(conn: &Connection, img: &ImageRecord) -> Result<i64> {
let id = conn.query_row(
"INSERT INTO images (folder_id, path, filename, thumbnail_path, width, height, file_size, created_at, modified_at, mime_type, media_kind, duration_ms, video_codec, audio_codec, metadata_updated_at, metadata_error, favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21, ?22)
"INSERT INTO images (folder_id, path, filename, thumbnail_path, width, height, file_size, created_at, modified_at, mime_type, media_kind, duration_ms, video_codec, audio_codec, metadata_updated_at, metadata_error, favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error, generated_caption, caption_model, caption_updated_at, caption_error)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21, ?22, ?23, ?24, ?25, ?26)
ON CONFLICT(path) DO UPDATE SET
folder_id = excluded.folder_id,
filename = excluded.filename,
@@ -250,7 +278,11 @@ pub fn upsert_image(conn: &Connection, img: &ImageRecord) -> Result<i64> {
embedding_status = excluded.embedding_status,
embedding_model = excluded.embedding_model,
embedding_updated_at = excluded.embedding_updated_at,
embedding_error = excluded.embedding_error
embedding_error = excluded.embedding_error,
generated_caption = excluded.generated_caption,
caption_model = excluded.caption_model,
caption_updated_at = excluded.caption_updated_at,
caption_error = excluded.caption_error
RETURNING id",
params![
img.folder_id,
@@ -275,6 +307,10 @@ pub fn upsert_image(conn: &Connection, img: &ImageRecord) -> Result<i64> {
img.embedding_model,
img.embedding_updated_at,
img.embedding_error,
img.generated_caption,
img.caption_model,
img.caption_updated_at,
img.caption_error,
],
|row| row.get(0),
)?;
@@ -307,6 +343,51 @@ pub fn backfill_embedding_jobs(conn: &Connection) -> Result<usize> {
Ok(inserted)
}
pub fn repair_embedding_consistency(conn: &Connection) -> Result<(usize, usize)> {
let orphaned_vectors = vector::delete_orphaned_embeddings(conn)?;
let ready_ids = {
let mut stmt = conn.prepare(
"SELECT id
FROM images
WHERE embedding_status = 'ready'",
)?;
let rows = stmt
.query_map([], |row| row.get::<_, i64>(0))?
.collect::<rusqlite::Result<Vec<_>>>()?;
rows
};
let mut missing_vector_ids = Vec::new();
for image_id in ready_ids {
if !vector::has_image_vector(conn, image_id)? {
missing_vector_ids.push(image_id);
}
}
let tx = conn.unchecked_transaction()?;
for image_id in &missing_vector_ids {
tx.execute(
"INSERT INTO embedding_jobs (image_id, status, attempts, last_error, created_at, updated_at)
VALUES (?1, 'pending', 0, NULL, datetime('now'), datetime('now'))
ON CONFLICT(image_id) DO UPDATE SET
status = 'pending',
last_error = NULL,
updated_at = datetime('now')",
[image_id],
)?;
tx.execute(
"UPDATE images
SET embedding_status = 'pending',
embedding_error = NULL
WHERE id = ?1",
[image_id],
)?;
}
tx.commit()?;
Ok((orphaned_vectors, missing_vector_ids.len()))
}
pub fn retry_failed_embedding_jobs(conn: &Connection, folder_id: i64) -> Result<usize> {
// Only re-queue images that are actually embeddable right now.
// Videos without a thumbnail would just fail again immediately, so skip them —
@@ -348,6 +429,10 @@ pub fn reset_inflight_jobs(conn: &Connection) -> Result<()> {
"UPDATE embedding_jobs SET status = 'pending' WHERE status = 'processing'",
[],
)?;
conn.execute(
"UPDATE caption_jobs SET status = 'pending' WHERE status = 'processing'",
[],
)?;
Ok(())
}
@@ -377,16 +462,102 @@ pub fn enqueue_metadata_job(conn: &Connection, image_id: i64) -> Result<()> {
Ok(())
}
pub fn get_pending_embedding_jobs(conn: &Connection, limit: usize) -> Result<Vec<EmbeddingJob>> {
let mut stmt = conn.prepare(
pub fn enqueue_caption_job(conn: &Connection, image_id: i64) -> Result<()> {
conn.execute(
"INSERT INTO caption_jobs (image_id, status, attempts, last_error, created_at, updated_at)
VALUES (?1, 'pending', 0, NULL, datetime('now'), datetime('now'))
ON CONFLICT(image_id) DO UPDATE SET
status = 'pending',
last_error = NULL,
updated_at = datetime('now')",
[image_id],
)?;
conn.execute(
"UPDATE images
SET caption_error = NULL
WHERE id = ?1",
[image_id],
)?;
Ok(())
}
pub fn requeue_caption_jobs(conn: &Connection, image_ids: &[i64]) -> Result<()> {
for image_id in image_ids {
conn.execute(
"UPDATE caption_jobs
SET status = 'pending', updated_at = datetime('now')
WHERE image_id = ?1 AND status = 'processing'",
[image_id],
)?;
}
Ok(())
}
pub fn enqueue_missing_caption_jobs_for_folder(conn: &Connection, folder_id: i64) -> Result<usize> {
let inserted = conn.execute(
"INSERT INTO caption_jobs (image_id, status, attempts, last_error, created_at, updated_at)
SELECT id, 'pending', 0, NULL, datetime('now'), datetime('now')
FROM images
WHERE folder_id = ?1
AND media_kind = 'image'
AND generated_caption IS NULL
ON CONFLICT(image_id) DO UPDATE SET
status = 'pending',
last_error = NULL,
updated_at = datetime('now')",
[folder_id],
)?;
conn.execute(
"UPDATE images
SET caption_error = NULL
WHERE folder_id = ?1
AND media_kind = 'image'
AND generated_caption IS NULL",
[folder_id],
)?;
Ok(inserted)
}
pub fn enqueue_missing_caption_jobs(conn: &Connection) -> Result<usize> {
let inserted = conn.execute(
"INSERT INTO caption_jobs (image_id, status, attempts, last_error, created_at, updated_at)
SELECT id, 'pending', 0, NULL, datetime('now'), datetime('now')
FROM images
WHERE media_kind = 'image'
AND generated_caption IS NULL
ON CONFLICT(image_id) DO UPDATE SET
status = 'pending',
last_error = NULL,
updated_at = datetime('now')",
[],
)?;
conn.execute(
"UPDATE images
SET caption_error = NULL
WHERE media_kind = 'image'
AND generated_caption IS NULL",
[],
)?;
Ok(inserted)
}
pub fn get_pending_embedding_jobs(
conn: &Connection,
excluded_folder_ids: &std::collections::HashSet<i64>,
limit: usize,
) -> Result<Vec<EmbeddingJob>> {
let sql = format!(
"SELECT j.image_id, i.folder_id, i.path, i.thumbnail_path, i.media_kind,
j.status, j.attempts, j.last_error, j.created_at, j.updated_at
FROM embedding_jobs j
JOIN images i ON i.id = j.image_id
WHERE status = 'pending'
{}
ORDER BY j.updated_at, j.image_id
LIMIT ?1",
)?;
folder_exclusion_clause("i", excluded_folder_ids)
);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map([limit as i64], |row| {
Ok(EmbeddingJob {
image_id: row.get(0)?,
@@ -404,9 +575,13 @@ pub fn get_pending_embedding_jobs(conn: &Connection, limit: usize) -> Result<Vec
Ok(rows.collect::<rusqlite::Result<Vec<_>>>()?)
}
pub fn claim_embedding_jobs(conn: &mut Connection, limit: usize) -> Result<Vec<EmbeddingJob>> {
pub fn claim_embedding_jobs(
conn: &mut Connection,
paused_folder_ids: &std::collections::HashSet<i64>,
limit: usize,
) -> Result<Vec<EmbeddingJob>> {
let tx = conn.transaction()?;
let candidates = get_pending_embedding_jobs(&tx, limit * 2)?;
let candidates = get_pending_embedding_jobs(&tx, paused_folder_ids, limit * 2)?;
let mut claimed = Vec::with_capacity(limit);
for job in candidates {
@@ -430,6 +605,64 @@ pub fn claim_embedding_jobs(conn: &mut Connection, limit: usize) -> Result<Vec<E
Ok(claimed)
}
fn get_pending_caption_jobs_excluding(
conn: &Connection,
excluded_folder_ids: &std::collections::HashSet<i64>,
limit: usize,
) -> Result<Vec<CaptionJob>> {
let sql = format!(
"SELECT j.image_id, i.folder_id, i.path
FROM caption_jobs j
JOIN images i ON i.id = j.image_id
WHERE j.status = 'pending'
AND i.media_kind = 'image'
AND i.generated_caption IS NULL
{}
ORDER BY j.updated_at, j.image_id
LIMIT ?1",
folder_exclusion_clause("i", excluded_folder_ids)
);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map([limit as i64], |row| {
Ok(CaptionJob {
image_id: row.get(0)?,
folder_id: row.get(1)?,
path: row.get(2)?,
})
})?;
Ok(rows.collect::<rusqlite::Result<Vec<_>>>()?)
}
pub fn claim_caption_jobs(
conn: &mut Connection,
paused_folder_ids: &std::collections::HashSet<i64>,
limit: usize,
) -> Result<Vec<CaptionJob>> {
let tx = conn.transaction()?;
let candidates = get_pending_caption_jobs_excluding(&tx, paused_folder_ids, limit * 2)?;
let mut claimed = Vec::with_capacity(limit);
for job in candidates {
let updated = tx.execute(
"UPDATE caption_jobs
SET status = 'processing', attempts = attempts + 1, updated_at = datetime('now')
WHERE image_id = ?1 AND status = 'pending'",
[job.image_id],
)?;
if updated == 1 {
claimed.push(job);
}
if claimed.len() >= limit {
break;
}
}
tx.commit()?;
Ok(claimed)
}
#[allow(dead_code)]
pub fn mark_embedding_ready(conn: &Connection, image_id: i64, model: &str) -> Result<()> {
conn.execute(
@@ -489,6 +722,7 @@ pub fn delete_images_by_ids(conn: &Connection, image_ids: &[i64]) -> Result<()>
let tx = conn.unchecked_transaction()?;
for image_id in image_ids {
vector::delete_embedding(&tx, *image_id)?;
vector::delete_caption_embedding(&tx, *image_id)?;
tx.execute("DELETE FROM images WHERE id = ?1", [image_id])?;
}
tx.commit()?;
@@ -539,6 +773,31 @@ pub fn get_folder_job_progress(conn: &Connection, folder_id: i64) -> Result<Fold
|row| row.get(0),
)?;
let caption_pending = conn.query_row(
"SELECT COUNT(*)
FROM caption_jobs j
JOIN images i ON i.id = j.image_id
WHERE i.folder_id = ?1 AND j.status IN ('pending', 'processing')",
[folder_id],
|row| row.get(0),
)?;
let caption_ready = conn.query_row(
"SELECT COUNT(*)
FROM images
WHERE folder_id = ?1 AND generated_caption IS NOT NULL",
[folder_id],
|row| row.get(0),
)?;
let caption_failed = conn.query_row(
"SELECT COUNT(*)
FROM images
WHERE folder_id = ?1 AND caption_error IS NOT NULL",
[folder_id],
|row| row.get(0),
)?;
Ok(FolderJobProgress {
folder_id,
thumbnail_pending,
@@ -546,6 +805,9 @@ pub fn get_folder_job_progress(conn: &Connection, folder_id: i64) -> Result<Fold
embedding_pending,
embedding_ready,
embedding_failed,
caption_pending,
caption_ready,
caption_failed,
})
}
@@ -562,15 +824,22 @@ pub fn get_all_folder_job_progress(conn: &Connection) -> Result<Vec<FolderJobPro
Ok(progress)
}
pub fn get_pending_thumbnail_jobs(conn: &Connection, limit: usize) -> Result<Vec<ThumbnailJob>> {
let mut stmt = conn.prepare(
fn get_pending_thumbnail_jobs_excluding(
conn: &Connection,
excluded_folder_ids: &std::collections::HashSet<i64>,
limit: usize,
) -> Result<Vec<ThumbnailJob>> {
let sql = format!(
"SELECT j.image_id, i.folder_id, i.path, i.media_kind
FROM thumbnail_jobs j
JOIN images i ON i.id = j.image_id
WHERE j.status = 'pending'
{}
ORDER BY j.updated_at, j.image_id
LIMIT ?1",
)?;
folder_exclusion_clause("i", excluded_folder_ids)
);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map([limit as i64], |row| {
Ok(ThumbnailJob {
image_id: row.get(0)?,
@@ -585,18 +854,20 @@ pub fn get_pending_thumbnail_jobs(conn: &Connection, limit: usize) -> Result<Vec
pub fn claim_thumbnail_jobs(
conn: &mut Connection,
active_folder_ids: &std::collections::HashSet<i64>,
paused_folder_ids: &std::collections::HashSet<i64>,
fetch_limit: usize,
claim_limit: usize,
) -> Result<Vec<ThumbnailJob>> {
let tx = conn.transaction()?;
let candidates = get_pending_thumbnail_jobs(&tx, fetch_limit)?;
let excluded_folder_ids = active_folder_ids
.union(paused_folder_ids)
.copied()
.collect::<std::collections::HashSet<_>>();
let candidates = get_pending_thumbnail_jobs_excluding(&tx, &excluded_folder_ids, fetch_limit)?;
let mut claimed = Vec::with_capacity(claim_limit);
for job in candidates {
if active_folder_ids.contains(&job.folder_id) {
continue;
}
debug_assert!(!excluded_folder_ids.contains(&job.folder_id));
let updated = tx.execute(
"UPDATE thumbnail_jobs
SET status = 'processing', attempts = attempts + 1, updated_at = datetime('now')
@@ -617,15 +888,22 @@ pub fn claim_thumbnail_jobs(
Ok(claimed)
}
pub fn get_pending_metadata_jobs(conn: &Connection, limit: usize) -> Result<Vec<MetadataJob>> {
let mut stmt = conn.prepare(
fn get_pending_metadata_jobs_excluding(
conn: &Connection,
excluded_folder_ids: &std::collections::HashSet<i64>,
limit: usize,
) -> Result<Vec<MetadataJob>> {
let sql = format!(
"SELECT j.image_id, i.folder_id, i.path
FROM metadata_jobs j
JOIN images i ON i.id = j.image_id
WHERE j.status = 'pending' AND i.media_kind = 'video'
{}
ORDER BY j.updated_at, j.image_id
LIMIT ?1",
)?;
folder_exclusion_clause("i", excluded_folder_ids)
);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map([limit as i64], |row| {
Ok(MetadataJob {
image_id: row.get(0)?,
@@ -639,18 +917,20 @@ pub fn get_pending_metadata_jobs(conn: &Connection, limit: usize) -> Result<Vec<
pub fn claim_metadata_jobs(
conn: &mut Connection,
active_folder_ids: &std::collections::HashSet<i64>,
paused_folder_ids: &std::collections::HashSet<i64>,
fetch_limit: usize,
claim_limit: usize,
) -> Result<Vec<MetadataJob>> {
let tx = conn.transaction()?;
let candidates = get_pending_metadata_jobs(&tx, fetch_limit)?;
let excluded_folder_ids = active_folder_ids
.union(paused_folder_ids)
.copied()
.collect::<std::collections::HashSet<_>>();
let candidates = get_pending_metadata_jobs_excluding(&tx, &excluded_folder_ids, fetch_limit)?;
let mut claimed = Vec::with_capacity(claim_limit);
for job in candidates {
if active_folder_ids.contains(&job.folder_id) {
continue;
}
debug_assert!(!excluded_folder_ids.contains(&job.folder_id));
let updated = tx.execute(
"UPDATE metadata_jobs
SET status = 'processing', attempts = attempts + 1, updated_at = datetime('now')
@@ -765,7 +1045,8 @@ pub fn update_image_details(
conn.query_row(
"SELECT id, folder_id, path, filename, thumbnail_path, width, height, file_size, created_at, modified_at, mime_type,
media_kind, duration_ms, video_codec, audio_codec, metadata_updated_at, metadata_error,
favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error
favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error,
generated_caption, caption_model, caption_updated_at, caption_error
FROM images
WHERE id = ?1",
[image_id],
@@ -778,7 +1059,8 @@ pub fn get_image_by_id(conn: &Connection, image_id: i64) -> Result<ImageRecord>
conn.query_row(
"SELECT id, folder_id, path, filename, thumbnail_path, width, height, file_size, created_at, modified_at, mime_type,
media_kind, duration_ms, video_codec, audio_codec, metadata_updated_at, metadata_error,
favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error
favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error,
generated_caption, caption_model, caption_updated_at, caption_error
FROM images
WHERE id = ?1",
[image_id],
@@ -790,7 +1072,11 @@ pub fn get_image_by_id(conn: &Connection, image_id: i64) -> Result<ImageRecord>
pub fn get_images_by_ids(conn: &Connection, image_ids: &[i64]) -> Result<Vec<ImageRecord>> {
let mut images = Vec::with_capacity(image_ids.len());
for image_id in image_ids {
images.push(get_image_by_id(conn, *image_id)?);
match get_image_by_id(conn, *image_id) {
Ok(image) => images.push(image),
Err(error) if is_query_returned_no_rows(&error) => {}
Err(error) => return Err(error),
}
}
Ok(images)
}
@@ -839,7 +1125,8 @@ pub fn get_images(
let sql = format!(
"SELECT id, folder_id, path, filename, thumbnail_path, width, height, file_size, created_at, modified_at, mime_type,
media_kind, duration_ms, video_codec, audio_codec, metadata_updated_at, metadata_error,
favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error
favorite, rating, embedding_status, embedding_model, embedding_updated_at, embedding_error,
generated_caption, caption_model, caption_updated_at, caption_error
FROM images
WHERE (?1 IS NULL OR folder_id = ?1)
AND (?2 IS NULL OR filename LIKE ?2)
@@ -918,11 +1205,104 @@ pub fn get_failed_embedding_images(
Ok(rows.collect::<rusqlite::Result<Vec<_>>>()?)
}
pub fn delete_folder(conn: &Connection, folder_id: i64) -> Result<()> {
conn.execute("DELETE FROM folders WHERE id = ?1", params![folder_id])?;
pub fn update_generated_caption(
conn: &Connection,
image_id: i64,
caption: &str,
model: &str,
) -> Result<ImageRecord> {
conn.execute(
"UPDATE images
SET generated_caption = ?2,
caption_model = ?3,
caption_updated_at = datetime('now'),
caption_error = NULL
WHERE id = ?1",
params![image_id, caption, model],
)?;
conn.execute("DELETE FROM caption_jobs WHERE image_id = ?1", [image_id])?;
get_image_by_id(conn, image_id)
}
pub fn mark_caption_failed(conn: &Connection, image_id: i64, error: &str) -> Result<()> {
conn.execute(
"UPDATE images
SET caption_error = ?2
WHERE id = ?1",
params![image_id, error],
)?;
conn.execute(
"UPDATE caption_jobs
SET status = 'failed', last_error = ?2, updated_at = datetime('now')
WHERE image_id = ?1",
params![image_id, error],
)?;
Ok(())
}
pub fn suggest_tags_from_caption(
conn: &Connection,
image_id: i64,
limit: usize,
) -> Result<Vec<String>> {
let caption = conn.query_row(
"SELECT generated_caption FROM images WHERE id = ?1",
[image_id],
|row| row.get::<_, Option<String>>(0),
)?;
Ok(caption
.as_deref()
.map(|caption| derive_caption_tags(caption, limit))
.unwrap_or_default())
}
pub fn delete_folder(conn: &Connection, folder_id: i64) -> Result<()> {
let image_ids = {
let mut stmt = conn.prepare("SELECT id FROM images WHERE folder_id = ?1")?;
let rows = stmt
.query_map([folder_id], |row| row.get::<_, i64>(0))?
.collect::<rusqlite::Result<Vec<_>>>()?;
rows
};
let tx = conn.unchecked_transaction()?;
for image_id in image_ids {
vector::delete_embedding(&tx, image_id)?;
vector::delete_caption_embedding(&tx, image_id)?;
}
tx.execute("DELETE FROM folders WHERE id = ?1", params![folder_id])?;
tx.commit()?;
Ok(())
}
fn derive_caption_tags(caption: &str, limit: usize) -> Vec<String> {
const STOPWORDS: &[&str] = &[
"a", "an", "and", "are", "as", "at", "be", "by", "for", "from", "in", "into", "is", "it",
"near", "of", "on", "or", "the", "to", "with", "without", "image", "photo", "picture",
"showing", "shows", "there", "this", "that", "over", "under",
];
let stopwords = STOPWORDS
.iter()
.copied()
.collect::<std::collections::HashSet<_>>();
let mut tags = Vec::new();
for word in caption
.split(|ch: char| !ch.is_alphanumeric())
.map(|word| word.trim().to_lowercase())
.filter(|word| word.len() >= 3 && !stopwords.contains(word.as_str()))
{
if !tags.contains(&word) {
tags.push(word);
}
if tags.len() >= limit {
break;
}
}
tags
}
fn map_image_row(row: &Row<'_>) -> rusqlite::Result<ImageRecord> {
Ok(ImageRecord {
id: row.get(0)?,
@@ -948,6 +1328,10 @@ fn map_image_row(row: &Row<'_>) -> rusqlite::Result<ImageRecord> {
embedding_model: row.get(20)?,
embedding_updated_at: row.get(21)?,
embedding_error: row.get(22)?,
generated_caption: row.get(23)?,
caption_model: row.get(24)?,
caption_updated_at: row.get(25)?,
caption_error: row.get(26)?,
})
}
@@ -1005,3 +1389,27 @@ fn ensure_column(conn: &Connection, table: &str, column: &str, definition: &str)
)?;
Ok(())
}
fn is_query_returned_no_rows(error: &anyhow::Error) -> bool {
error
.downcast_ref::<rusqlite::Error>()
.is_some_and(|error| matches!(error, rusqlite::Error::QueryReturnedNoRows))
}
fn folder_exclusion_clause(
image_alias: &str,
excluded_folder_ids: &std::collections::HashSet<i64>,
) -> String {
if excluded_folder_ids.is_empty() {
return String::new();
}
let mut ids = excluded_folder_ids.iter().copied().collect::<Vec<_>>();
ids.sort_unstable();
let id_list = ids
.into_iter()
.map(|id| id.to_string())
.collect::<Vec<_>>()
.join(",");
format!("AND {}.folder_id NOT IN ({})", image_alias, id_list)
}
+218 -44
View File
@@ -1,3 +1,4 @@
use crate::captioner::{self, FlorenceCaptioner};
use crate::db::{self, DbPool, EmbeddingJob, FolderJobProgress, ImageRecord, IndexedMediaEntry};
use crate::embedder::{embedding_source_path, ClipImageEmbedder};
use crate::media::{probe_video_metadata, MediaTools};
@@ -9,7 +10,6 @@ use rayon::prelude::*;
use serde::Serialize;
use std::collections::{HashMap, HashSet};
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Mutex, OnceLock};
use std::time::{Duration, Instant};
use tauri::{AppHandle, Emitter};
@@ -22,29 +22,91 @@ const IMAGE_EXTENSIONS: &[&str] = &[
const VIDEO_EXTENSIONS: &[&str] = &["mp4", "mov", "m4v", "webm"];
const JOB_PROGRESS_EMIT_INTERVAL: Duration = Duration::from_millis(750);
const CAPTION_BATCH_SIZE: usize = 1;
static LAST_JOB_PROGRESS_EMIT: OnceLock<Mutex<HashMap<i64, Instant>>> = OnceLock::new();
static ACTIVE_INDEXING_FOLDERS: OnceLock<Mutex<HashSet<i64>>> = OnceLock::new();
static THUMBNAIL_WORKER_PAUSED: AtomicBool = AtomicBool::new(false);
static METADATA_WORKER_PAUSED: AtomicBool = AtomicBool::new(false);
static EMBEDDING_WORKER_PAUSED: AtomicBool = AtomicBool::new(false);
static PAUSED_WORKER_FOLDERS: OnceLock<Mutex<PausedWorkerFolders>> = OnceLock::new();
pub fn set_worker_paused(worker: &str, paused: bool) {
match worker {
"thumbnail" => THUMBNAIL_WORKER_PAUSED.store(paused, Ordering::Relaxed),
"metadata" => METADATA_WORKER_PAUSED.store(paused, Ordering::Relaxed),
"embedding" => EMBEDDING_WORKER_PAUSED.store(paused, Ordering::Relaxed),
_ => {}
#[derive(Default)]
struct PausedWorkerFolders {
thumbnail: HashSet<i64>,
metadata: HashSet<i64>,
embedding: HashSet<i64>,
caption: HashSet<i64>,
}
#[derive(Clone, Copy)]
pub struct FolderWorkerPausedState {
pub thumbnail: bool,
pub metadata: bool,
pub embedding: bool,
pub caption: bool,
}
pub fn set_worker_paused(worker: &str, folder_id: i64, paused: bool) {
if let Ok(mut paused_folders) = PAUSED_WORKER_FOLDERS
.get_or_init(|| Mutex::new(PausedWorkerFolders::default()))
.lock()
{
let folder_set = match worker {
"thumbnail" => Some(&mut paused_folders.thumbnail),
"metadata" => Some(&mut paused_folders.metadata),
"embedding" => Some(&mut paused_folders.embedding),
"caption" => Some(&mut paused_folders.caption),
_ => None,
};
if let Some(folder_set) = folder_set {
if paused {
folder_set.insert(folder_id);
} else {
folder_set.remove(&folder_id);
}
}
}
}
pub fn get_worker_paused_states() -> [bool; 3] {
[
THUMBNAIL_WORKER_PAUSED.load(Ordering::Relaxed),
METADATA_WORKER_PAUSED.load(Ordering::Relaxed),
EMBEDDING_WORKER_PAUSED.load(Ordering::Relaxed),
]
pub fn get_worker_paused_states(folder_ids: &[i64]) -> HashMap<i64, FolderWorkerPausedState> {
let Ok(paused_folders) = PAUSED_WORKER_FOLDERS
.get_or_init(|| Mutex::new(PausedWorkerFolders::default()))
.lock()
else {
return HashMap::new();
};
folder_ids
.iter()
.copied()
.map(|folder_id| {
(
folder_id,
FolderWorkerPausedState {
thumbnail: paused_folders.thumbnail.contains(&folder_id),
metadata: paused_folders.metadata.contains(&folder_id),
embedding: paused_folders.embedding.contains(&folder_id),
caption: paused_folders.caption.contains(&folder_id),
},
)
})
.collect()
}
fn paused_folder_ids(worker: &str) -> HashSet<i64> {
let Ok(paused_folders) = PAUSED_WORKER_FOLDERS
.get_or_init(|| Mutex::new(PausedWorkerFolders::default()))
.lock()
else {
return HashSet::new();
};
match worker {
"thumbnail" => paused_folders.thumbnail.clone(),
"metadata" => paused_folders.metadata.clone(),
"embedding" => paused_folders.embedding.clone(),
_ => HashSet::new(),
}
}
static FOLDER_STORAGE_PROFILES: OnceLock<Mutex<HashMap<i64, RuntimeAdaptiveProfile>>> =
OnceLock::new();
@@ -95,10 +157,6 @@ pub fn start_thumbnail_worker(
cache_dir: PathBuf,
) {
std::thread::spawn(move || loop {
if THUMBNAIL_WORKER_PAUSED.load(Ordering::Relaxed) {
std::thread::sleep(std::time::Duration::from_millis(500));
continue;
}
if let Err(error) = process_thumbnail_batch(&app, &pool, &media_tools, &cache_dir) {
eprintln!("Thumbnail worker error: {}", error);
}
@@ -108,10 +166,6 @@ pub fn start_thumbnail_worker(
pub fn start_metadata_worker(app: AppHandle, pool: DbPool, media_tools: MediaTools) {
std::thread::spawn(move || loop {
if METADATA_WORKER_PAUSED.load(Ordering::Relaxed) {
std::thread::sleep(std::time::Duration::from_millis(500));
continue;
}
if let Err(error) = process_metadata_batch(&app, &pool, &media_tools) {
eprintln!("Metadata worker error: {}", error);
}
@@ -124,10 +178,6 @@ pub fn start_embedding_worker(app: AppHandle, pool: DbPool) {
let mut embedder: Option<ClipImageEmbedder> = None;
println!("Embedding worker started.");
loop {
if EMBEDDING_WORKER_PAUSED.load(Ordering::Relaxed) {
std::thread::sleep(std::time::Duration::from_millis(500));
continue;
}
if let Err(error) = process_embedding_batch(&app, &pool, &mut embedder) {
eprintln!("Embedding worker error: {}", error);
}
@@ -136,6 +186,20 @@ pub fn start_embedding_worker(app: AppHandle, pool: DbPool) {
});
}
pub fn start_caption_worker(app: AppHandle, pool: DbPool, app_data_dir: PathBuf) {
std::thread::spawn(move || {
let mut captioner: Option<FlorenceCaptioner> = None;
println!("Caption worker started.");
loop {
if let Err(error) = process_caption_batch(&app, &pool, &app_data_dir, &mut captioner) {
eprintln!("Caption worker error: {}", error);
captioner = None;
}
std::thread::sleep(std::time::Duration::from_millis(750));
}
});
}
fn do_index(app: AppHandle, pool: DbPool, folder_id: i64, folder_path: PathBuf) -> Result<()> {
let existing_entries = {
let conn = pool.get()?;
@@ -197,7 +261,7 @@ fn do_index(app: AppHandle, pool: DbPool, folder_id: i64, folder_path: PathBuf)
images: committed,
},
);
emit_folder_job_progress(&app, &pool, &[folder_id]);
emit_folder_job_progress(&app, &pool, &[folder_id], false);
}
processed += path_chunk.len();
@@ -245,7 +309,7 @@ fn do_index(app: AppHandle, pool: DbPool, folder_id: i64, folder_path: PathBuf)
done: true,
},
);
emit_folder_job_progress(&app, &pool, &[folder_id]);
emit_folder_job_progress(&app, &pool, &[folder_id], true);
Ok(())
}
@@ -302,6 +366,10 @@ fn build_record(
embedding_model: Some(vector::CLIP_MODEL_NAME.to_string()),
embedding_updated_at: None,
embedding_error: None,
generated_caption: None,
caption_model: None,
caption_updated_at: None,
caption_error: None,
})
}
@@ -335,11 +403,13 @@ fn process_thumbnail_batch(
with_db_write_lock(|| {
let mut conn = pool.get()?;
let active_folders = active_indexing_folders();
let paused_folders = paused_folder_ids("thumbnail");
let worker_batch_size = max_worker_batch_size(&active_folders);
let worker_fetch_size = max_worker_fetch_size(&active_folders);
db::claim_thumbnail_jobs(
&mut conn,
&active_folders,
&paused_folders,
worker_fetch_size,
worker_batch_size,
)
@@ -428,7 +498,7 @@ fn process_thumbnail_batch(
images: updated_images,
},
);
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>());
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>(), true);
}
Ok(())
@@ -439,11 +509,13 @@ fn process_metadata_batch(app: &AppHandle, pool: &DbPool, media_tools: &MediaToo
with_db_write_lock(|| {
let mut conn = pool.get()?;
let active_folders = active_indexing_folders();
let paused_folders = paused_folder_ids("metadata");
let worker_batch_size = max_worker_batch_size(&active_folders);
let worker_fetch_size = max_worker_fetch_size(&active_folders);
db::claim_metadata_jobs(
&mut conn,
&active_folders,
&paused_folders,
worker_fetch_size,
worker_batch_size,
)
@@ -506,7 +578,7 @@ fn process_metadata_batch(app: &AppHandle, pool: &DbPool, media_tools: &MediaToo
images: updated_images,
},
);
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>());
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>(), true);
}
Ok(())
@@ -518,14 +590,11 @@ fn process_embedding_batch(
embedder: &mut Option<ClipImageEmbedder>,
) -> Result<()> {
let batch_started_at = Instant::now();
if embedder.is_none() {
*embedder = Some(ClipImageEmbedder::new()?);
}
let claim_started_at = Instant::now();
let paused_folders = paused_folder_ids("embedding");
let jobs = with_db_write_lock(|| {
let mut conn = pool.get()?;
db::claim_embedding_jobs(&mut conn, EMBEDDING_BATCH_SIZE)
db::claim_embedding_jobs(&mut conn, &paused_folders, EMBEDDING_BATCH_SIZE)
})?;
let claim_elapsed = claim_started_at.elapsed();
@@ -533,9 +602,18 @@ fn process_embedding_batch(
return Ok(());
}
if embedder.is_none() {
*embedder = Some(ClipImageEmbedder::new()?);
}
println!("Embedding batch claimed: {} items", jobs.len());
let folder_ids = jobs.iter().map(|job| job.folder_id).collect::<HashSet<_>>();
emit_folder_job_progress(app, pool, &folder_ids.iter().copied().collect::<Vec<_>>());
emit_folder_job_progress(
app,
pool,
&folder_ids.iter().copied().collect::<Vec<_>>(),
false,
);
let embedder = embedder.as_ref().expect("embedder should be initialized");
let infer_started_at = Instant::now();
@@ -639,7 +717,7 @@ fn process_embedding_batch(
images: updated_images,
},
);
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>());
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>(), true);
}
let write_elapsed = write_started_at.elapsed();
@@ -652,6 +730,101 @@ fn process_embedding_batch(
Ok(())
}
fn process_caption_batch(
app: &AppHandle,
pool: &DbPool,
app_data_dir: &Path,
captioner: &mut Option<FlorenceCaptioner>,
) -> Result<()> {
if !captioner::caption_model_status(app_data_dir).ready {
return Ok(());
}
let paused_folders = paused_folder_ids("caption");
let jobs = with_db_write_lock(|| {
let mut conn = pool.get()?;
db::claim_caption_jobs(&mut conn, &paused_folders, CAPTION_BATCH_SIZE)
})?;
if jobs.is_empty() {
return Ok(());
}
if captioner.is_none() {
match FlorenceCaptioner::new(app_data_dir) {
Ok(model) => *captioner = Some(model),
Err(error) => {
with_db_write_lock(|| {
let conn = pool.get()?;
db::requeue_caption_jobs(
&conn,
&jobs.iter().map(|job| job.image_id).collect::<Vec<_>>(),
)
})?;
return Err(error);
}
}
}
let folder_ids = jobs.iter().map(|job| job.folder_id).collect::<HashSet<_>>();
emit_folder_job_progress(
app,
pool,
&folder_ids.iter().copied().collect::<Vec<_>>(),
false,
);
let captioner = captioner
.as_mut()
.expect("captioner should be initialized before caption batch processing");
let caption_results = jobs
.iter()
.map(|job| (job.clone(), captioner.generate(Path::new(&job.path))))
.collect::<Vec<_>>();
let updated_images = with_db_write_lock(|| {
let mut conn = pool.get()?;
let tx = conn.transaction()?;
let mut updated_images = Vec::with_capacity(caption_results.len());
for (job, caption_result) in &caption_results {
match caption_result {
Ok(caption) => {
updated_images.push(db::update_generated_caption(
&tx,
job.image_id,
caption,
captioner::FLORENCE_CAPTION_MODEL_NAME,
)?);
}
Err(error) => {
db::mark_caption_failed(&tx, job.image_id, &error.to_string())?;
updated_images.push(db::get_image_by_id(&tx, job.image_id)?);
}
}
}
tx.commit()?;
Ok(updated_images)
})?;
if !updated_images.is_empty() {
let folder_ids = updated_images
.iter()
.map(|image| image.folder_id)
.collect::<HashSet<_>>();
emit_media_updates(
app,
&MediaUpdateBatch {
images: updated_images,
},
);
emit_folder_job_progress(app, pool, &folder_ids.into_iter().collect::<Vec<_>>(), true);
}
Ok(())
}
fn active_indexing_folders() -> HashSet<i64> {
ACTIVE_INDEXING_FOLDERS
.get_or_init(|| Mutex::new(HashSet::new()))
@@ -737,7 +910,7 @@ fn emit_media_updates(app: &AppHandle, batch: &MediaUpdateBatch) {
let _ = app.emit("media-updated", batch);
}
fn emit_folder_job_progress(app: &AppHandle, pool: &DbPool, folder_ids: &[i64]) {
fn emit_folder_job_progress(app: &AppHandle, pool: &DbPool, folder_ids: &[i64], force: bool) {
let mut unique_folder_ids = folder_ids.iter().copied().collect::<Vec<_>>();
unique_folder_ids.sort_unstable();
unique_folder_ids.dedup();
@@ -749,10 +922,11 @@ fn emit_folder_job_progress(app: &AppHandle, pool: &DbPool, folder_ids: &[i64])
Err(_) => return,
};
unique_folder_ids.retain(|folder_id| {
let should_emit = tracker
.get(folder_id)
.map(|last_emit| now.duration_since(*last_emit) >= JOB_PROGRESS_EMIT_INTERVAL)
.unwrap_or(true);
let should_emit = force
|| tracker
.get(folder_id)
.map(|last_emit| now.duration_since(*last_emit) >= JOB_PROGRESS_EMIT_INTERVAL)
.unwrap_or(true);
if should_emit {
tracker.insert(*folder_id, now);
}
+23 -3
View File
@@ -1,3 +1,4 @@
mod captioner;
mod commands;
mod db;
mod embedder;
@@ -7,8 +8,8 @@ mod storage;
mod thumbnail;
mod vector;
use tauri::Manager;
use crate::storage::StorageProfile;
use tauri::Manager;
#[cfg_attr(mobile, tauri::mobile_entry_point)]
pub fn run() {
@@ -34,11 +35,19 @@ pub fn run() {
let conn = pool.get().expect("Failed to get connection for migration");
db::migrate(&conn).expect("Failed to run migrations");
db::reset_inflight_jobs(&conn).expect("Failed to reset inflight jobs");
let backfilled = db::backfill_embedding_jobs(&conn)
.expect("Failed to backfill embedding jobs");
let backfilled =
db::backfill_embedding_jobs(&conn).expect("Failed to backfill embedding jobs");
if backfilled > 0 {
println!("Backfilled {} embedding jobs.", backfilled);
}
let (orphaned_vectors, missing_vectors) = db::repair_embedding_consistency(&conn)
.expect("Failed to repair embedding consistency");
if orphaned_vectors > 0 || missing_vectors > 0 {
println!(
"Repaired embedding consistency: removed {} orphaned vectors, requeued {} missing vectors.",
orphaned_vectors, missing_vectors
);
}
}
let thumb_dir = app_dir.join("thumbnails");
@@ -58,6 +67,7 @@ pub fn run() {
}
indexer::start_metadata_worker(app.handle().clone(), pool.clone(), media_tools.clone());
indexer::start_embedding_worker(app.handle().clone(), pool.clone());
indexer::start_caption_worker(app.handle().clone(), pool.clone(), app_dir.clone());
app.manage(pool);
app.manage(media_tools);
@@ -73,8 +83,18 @@ pub fn run() {
commands::reindex_folder,
commands::update_image_details,
commands::find_similar_images,
commands::debug_similar_images,
commands::retry_failed_embeddings,
commands::semantic_search_images,
commands::get_caption_model_status,
commands::prepare_caption_model,
commands::delete_caption_model,
commands::probe_caption_runtime,
commands::probe_caption_image,
commands::generate_caption_for_image,
commands::queue_caption_jobs,
commands::set_generated_caption,
commands::suggest_image_tags,
commands::set_worker_paused,
commands::get_worker_states,
commands::get_tag_cloud,
+147 -4
View File
@@ -1,5 +1,5 @@
use anyhow::{anyhow, Result};
use rusqlite::{ffi::sqlite3_auto_extension, Connection};
use rusqlite::{ffi::sqlite3_auto_extension, Connection, Error as SqliteError};
use sqlite_vec::sqlite3_vec_init;
use std::sync::Once;
@@ -19,8 +19,13 @@ pub fn migrate(conn: &Connection) -> Result<()> {
"CREATE VIRTUAL TABLE IF NOT EXISTS image_vec USING vec0(
image_id INTEGER PRIMARY KEY,
embedding FLOAT[{}] distance_metric=cosine
);
CREATE VIRTUAL TABLE IF NOT EXISTS caption_vec USING vec0(
image_id INTEGER PRIMARY KEY,
embedding FLOAT[{}] distance_metric=cosine
);",
CLIP_VECTOR_DIM
CLIP_VECTOR_DIM, CLIP_VECTOR_DIM
))?;
Ok(())
}
@@ -31,6 +36,12 @@ pub fn delete_embedding(conn: &Connection, image_id: i64) -> Result<()> {
Ok(())
}
#[allow(dead_code)]
pub fn delete_caption_embedding(conn: &Connection, image_id: i64) -> Result<()> {
conn.execute("DELETE FROM caption_vec WHERE image_id = ?1", [image_id])?;
Ok(())
}
#[allow(dead_code)]
pub fn upsert_embedding(conn: &Connection, image_id: i64, embedding: &[f32]) -> Result<()> {
if embedding.len() != CLIP_VECTOR_DIM {
@@ -50,12 +61,35 @@ pub fn upsert_embedding(conn: &Connection, image_id: i64, embedding: &[f32]) ->
Ok(())
}
#[allow(dead_code)]
pub fn upsert_caption_embedding(conn: &Connection, image_id: i64, embedding: &[f32]) -> Result<()> {
if embedding.len() != CLIP_VECTOR_DIM {
return Err(anyhow!(
"expected {}-dimensional embedding, got {}",
CLIP_VECTOR_DIM,
embedding.len()
));
}
let packed = pack_f32(embedding);
conn.execute("DELETE FROM caption_vec WHERE image_id = ?1", [image_id])?;
conn.execute(
"INSERT INTO caption_vec (image_id, embedding) VALUES (?1, ?2)",
(&image_id, &packed),
)?;
Ok(())
}
pub fn find_similar_image_ids(conn: &Connection, image_id: i64, limit: usize) -> Result<Vec<i64>> {
let embedding: Vec<u8> = conn.query_row(
let embedding: Vec<u8> = match conn.query_row(
"SELECT embedding FROM image_vec WHERE image_id = ?1",
[image_id],
|row| row.get(0),
)?;
) {
Ok(embedding) => embedding,
Err(SqliteError::QueryReturnedNoRows) => return Ok(Vec::new()),
Err(error) => return Err(error.into()),
};
let mut stmt = conn.prepare(
"SELECT image_id
@@ -155,6 +189,115 @@ pub fn search_image_ids_by_embedding(
Ok(ids)
}
#[allow(dead_code)]
pub fn search_caption_ids_by_embedding(
conn: &Connection,
embedding: &[f32],
limit: usize,
) -> Result<Vec<i64>> {
if embedding.len() != CLIP_VECTOR_DIM {
return Err(anyhow!(
"expected {}-dimensional embedding, got {}",
CLIP_VECTOR_DIM,
embedding.len()
));
}
let packed = pack_f32(embedding);
let mut stmt = conn.prepare(
"SELECT image_id
FROM caption_vec
WHERE embedding MATCH vec_f32(?1)
AND k = ?2",
)?;
let rows = stmt.query_map((&packed, limit as i64), |row| row.get::<_, i64>(0))?;
let mut ids = Vec::new();
for row in rows {
ids.push(row?);
if ids.len() >= limit {
break;
}
}
Ok(ids)
}
pub fn count_image_vectors(conn: &Connection) -> Result<i64> {
conn.query_row("SELECT COUNT(*) FROM image_vec", [], |row| row.get(0))
.map_err(Into::into)
}
#[allow(dead_code)]
pub fn count_caption_vectors(conn: &Connection) -> Result<i64> {
conn.query_row("SELECT COUNT(*) FROM caption_vec", [], |row| row.get(0))
.map_err(Into::into)
}
pub fn delete_orphaned_embeddings(conn: &Connection) -> Result<usize> {
let image_ids = {
let mut stmt = conn.prepare("SELECT id FROM images")?;
let rows = stmt
.query_map([], |row| row.get::<_, i64>(0))?
.collect::<rusqlite::Result<std::collections::HashSet<_>>>()?;
rows
};
let vector_ids = {
let mut stmt = conn.prepare("SELECT image_id FROM image_vec")?;
let rows = stmt
.query_map([], |row| row.get::<_, i64>(0))?
.collect::<rusqlite::Result<Vec<_>>>()?;
rows
};
let orphaned_ids = vector_ids
.into_iter()
.filter(|image_id| !image_ids.contains(image_id))
.collect::<Vec<_>>();
for image_id in &orphaned_ids {
delete_embedding(conn, *image_id)?;
}
Ok(orphaned_ids.len())
}
#[allow(dead_code)]
pub fn delete_orphaned_caption_embeddings(conn: &Connection) -> Result<usize> {
let image_ids = {
let mut stmt = conn.prepare("SELECT id FROM images")?;
let rows = stmt
.query_map([], |row| row.get::<_, i64>(0))?
.collect::<rusqlite::Result<std::collections::HashSet<_>>>()?;
rows
};
let vector_ids = {
let mut stmt = conn.prepare("SELECT image_id FROM caption_vec")?;
let rows = stmt
.query_map([], |row| row.get::<_, i64>(0))?
.collect::<rusqlite::Result<Vec<_>>>()?;
rows
};
let orphaned_ids = vector_ids
.into_iter()
.filter(|image_id| !image_ids.contains(image_id))
.collect::<Vec<_>>();
for image_id in &orphaned_ids {
delete_caption_embedding(conn, *image_id)?;
}
Ok(orphaned_ids.len())
}
pub fn has_image_vector(conn: &Connection, image_id: i64) -> Result<bool> {
conn.query_row(
"SELECT EXISTS(SELECT 1 FROM image_vec WHERE image_id = ?1)",
[image_id],
|row| row.get::<_, i64>(0),
)
.map(|value| value != 0)
.map_err(Into::into)
}
#[allow(dead_code)]
fn pack_f32(values: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(values.len() * std::mem::size_of::<f32>());
+4
View File
@@ -7,17 +7,20 @@ import { Gallery } from "./components/Gallery";
import { Lightbox } from "./components/Lightbox";
import { TagCloud } from "./components/TagCloud";
import { TitleBar } from "./components/TitleBar";
import { SettingsModal } from "./components/SettingsModal";
export default function App() {
const loadFolders = useGalleryStore((state) => state.loadFolders);
const loadBackgroundJobProgress = useGalleryStore((state) => state.loadBackgroundJobProgress);
const loadImages = useGalleryStore((state) => state.loadImages);
const loadCaptionModelStatus = useGalleryStore((state) => state.loadCaptionModelStatus);
const subscribeToProgress = useGalleryStore((state) => state.subscribeToProgress);
const activeView = useGalleryStore((state) => state.activeView);
useEffect(() => {
loadFolders().then(() => {
void loadBackgroundJobProgress();
void loadCaptionModelStatus();
return loadImages(true);
});
let unlisten: (() => void) | undefined;
@@ -54,6 +57,7 @@ export default function App() {
</div>
<Lightbox />
<SettingsModal />
</div>
);
}
+86 -28
View File
@@ -2,12 +2,13 @@ import { useEffect, useMemo, useState } from "react";
import { invoke } from "@tauri-apps/api/core";
import { useGalleryStore } from "../store";
type WorkerKey = "thumbnail" | "metadata" | "embedding";
type WorkerKey = "thumbnail" | "metadata" | "embedding" | "caption";
const WORKER_FOR_STAGE: Record<string, WorkerKey> = {
Thumbnails: "thumbnail",
Metadata: "metadata",
Embeddings: "embedding",
Captions: "caption",
};
interface TaskStage {
@@ -22,6 +23,7 @@ interface Task {
name: string;
stages: TaskStage[];
hasFailedEmbeddings: boolean;
hasFailedCaptions: boolean;
pendingMediaWork: number;
embeddingProcessed: number;
embeddingTotal: number;
@@ -35,6 +37,21 @@ interface FailedEmbeddingItem {
error: string | null;
}
interface FolderWorkerStates {
folder_id: number;
thumbnail_paused: boolean;
metadata_paused: boolean;
embedding_paused: boolean;
caption_paused: boolean;
}
const DEFAULT_PAUSED_STATE: Record<WorkerKey, boolean> = {
thumbnail: false,
metadata: false,
embedding: false,
caption: false,
};
export function BackgroundTasks() {
const folders = useGalleryStore((state) => state.folders);
const indexingProgress = useGalleryStore((state) => state.indexingProgress);
@@ -42,24 +59,32 @@ export function BackgroundTasks() {
const retryFailedEmbeddings = useGalleryStore((state) => state.retryFailedEmbeddings);
const [expanded, setExpanded] = useState(false);
const [dismissed, setDismissed] = useState<Record<number, string>>({});
const [paused, setPaused] = useState<Record<WorkerKey, boolean>>({
thumbnail: false,
metadata: false,
embedding: false,
});
const [paused, setPaused] = useState<Record<number, Record<WorkerKey, boolean>>>({});
const [failedItems, setFailedItems] = useState<Record<number, FailedEmbeddingItem[]>>({});
useEffect(() => {
invoke<{ thumbnail_paused: boolean; metadata_paused: boolean; embedding_paused: boolean }>(
"get_worker_states",
).then((states) => {
setPaused({
thumbnail: states.thumbnail_paused,
metadata: states.metadata_paused,
embedding: states.embedding_paused,
});
const folderIds = folders.map((folder) => folder.id);
if (folderIds.length === 0) {
setPaused({});
return;
}
invoke<FolderWorkerStates[]>("get_worker_states", { folderIds }).then((states) => {
setPaused(
Object.fromEntries(
states.map((state) => [
state.folder_id,
{
thumbnail: state.thumbnail_paused,
metadata: state.metadata_paused,
embedding: state.embedding_paused,
caption: state.caption_paused,
},
]),
),
);
});
}, []);
}, [folders]);
// Fetch failed embedding filenames whenever the expanded panel opens or failure counts change.
const failedCounts = useMemo(
@@ -83,10 +108,21 @@ export function BackgroundTasks() {
}
}, [expanded, failedCounts]);
const toggleWorker = (worker: WorkerKey) => {
const next = !paused[worker];
setPaused((prev) => ({ ...prev, [worker]: next }));
void invoke("set_worker_paused", { worker, paused: next });
const isWorkerPaused = (folderId: number, worker: WorkerKey) => {
return paused[folderId]?.[worker] ?? DEFAULT_PAUSED_STATE[worker];
};
const toggleWorker = (folderId: number, worker: WorkerKey) => {
const next = !isWorkerPaused(folderId, worker);
setPaused((prev) => ({
...prev,
[folderId]: {
...DEFAULT_PAUSED_STATE,
...prev[folderId],
[worker]: next,
},
}));
void invoke("set_worker_paused", { worker, folderId, paused: next });
};
const dismissTask = (id: number, snapshot: string) => {
@@ -105,13 +141,16 @@ export function BackgroundTasks() {
const embeddingPending = jobs?.embedding_pending ?? 0;
const embeddingReady = jobs?.embedding_ready ?? 0;
const embeddingFailed = jobs?.embedding_failed ?? 0;
const captionPending = jobs?.caption_pending ?? 0;
const captionFailed = jobs?.caption_failed ?? 0;
const pendingMediaWork = thumbnailPending + metadataPending + embeddingPending;
const pendingMediaWork = thumbnailPending + metadataPending + embeddingPending + captionPending;
const embeddingProcessed = embeddingReady + embeddingFailed;
const embeddingTotal = embeddingProcessed + embeddingPending;
const hasFailedEmbeddings = embeddingFailed > 0;
const hasFailedCaptions = captionFailed > 0;
if (!index && pendingMediaWork === 0 && !hasFailedEmbeddings) return null;
if (!index && pendingMediaWork === 0 && !hasFailedEmbeddings && !hasFailedCaptions) return null;
const stages: TaskStage[] = [];
@@ -153,6 +192,15 @@ export function BackgroundTasks() {
});
}
if (captionPending > 0) {
stages.push({
label: "Captions",
detail: captionPending.toLocaleString(),
progress: null,
failed: false,
});
}
if (hasFailedEmbeddings && pendingMediaWork === 0) {
stages.push({
label: "Failed",
@@ -162,13 +210,23 @@ export function BackgroundTasks() {
});
}
const snapshot = `${pendingMediaWork}:${embeddingFailed}`;
if (hasFailedCaptions && pendingMediaWork === 0) {
stages.push({
label: "Failed",
detail: `${captionFailed.toLocaleString()} captions`,
progress: null,
failed: true,
});
}
const snapshot = `${pendingMediaWork}:${embeddingFailed}:${captionFailed}`;
return {
id: folder.id,
name: folder.name,
stages,
hasFailedEmbeddings,
hasFailedCaptions,
pendingMediaWork,
embeddingProcessed,
embeddingTotal,
@@ -184,7 +242,7 @@ export function BackgroundTasks() {
const primary = tasks[0];
const extraCount = tasks.length - 1;
const hasFailed = tasks.some((t) => t.hasFailedEmbeddings && t.pendingMediaWork === 0);
const hasFailed = tasks.some((t) => (t.hasFailedEmbeddings || t.hasFailedCaptions) && t.pendingMediaWork === 0);
// Best progress bar value: use embedding progress if available (most informative),
// otherwise fall back to scanning progress, otherwise indeterminate.
@@ -214,7 +272,7 @@ export function BackgroundTasks() {
<div className="flex items-center gap-1.5 flex-1 min-w-0 overflow-hidden">
{primary.stages.map((stage) => {
const workerKey = WORKER_FOR_STAGE[stage.label];
const isPaused = workerKey ? paused[workerKey] : false;
const isPaused = workerKey ? isWorkerPaused(primary.id, workerKey) : false;
return (
<span
key={stage.label}
@@ -234,7 +292,7 @@ export function BackgroundTasks() {
<button
className="ml-0.5 opacity-0 group-hover:opacity-100 hover:text-white transition-opacity"
title={isPaused ? `Resume ${stage.label}` : `Pause ${stage.label}`}
onClick={(e) => { e.stopPropagation(); toggleWorker(workerKey); }}
onClick={(e) => { e.stopPropagation(); toggleWorker(primary.id, workerKey); }}
>
{isPaused ? (
<svg className="h-2.5 w-2.5" fill="currentColor" viewBox="0 0 24 24">
@@ -312,7 +370,7 @@ export function BackgroundTasks() {
const taskEmbeddingStage = task.stages.find((s) => s.label === "Embeddings");
const taskScanningStage = task.stages.find((s) => s.label === "Scanning");
const taskBarProgress = taskEmbeddingStage?.progress ?? taskScanningStage?.progress ?? null;
const taskHasFailed = task.hasFailedEmbeddings && task.pendingMediaWork === 0;
const taskHasFailed = (task.hasFailedEmbeddings || task.hasFailedCaptions) && task.pendingMediaWork === 0;
return (
<div key={task.id}>
@@ -322,7 +380,7 @@ export function BackgroundTasks() {
<div className="flex items-center gap-1.5 flex-1 min-w-0 overflow-hidden">
{task.stages.map((stage) => {
const workerKey = WORKER_FOR_STAGE[stage.label];
const isPaused = workerKey ? paused[workerKey] : false;
const isPaused = workerKey ? isWorkerPaused(task.id, workerKey) : false;
return (
<span
key={stage.label}
@@ -347,7 +405,7 @@ export function BackgroundTasks() {
<button
className="ml-0.5 text-gray-600 hover:text-white transition-colors"
title={isPaused ? `Resume ${stage.label}` : `Pause ${stage.label}`}
onClick={() => toggleWorker(workerKey)}
onClick={() => toggleWorker(task.id, workerKey)}
>
{isPaused ? (
<svg className="h-2.5 w-2.5" fill="currentColor" viewBox="0 0 24 24">
+24 -4
View File
@@ -251,6 +251,10 @@ export function Gallery() {
const zoomPreset = useGalleryStore((state) => state.zoomPreset);
const search = useGalleryStore((state) => state.search);
const searchMode = useGalleryStore((state) => state.searchMode);
const collectionTitle = useGalleryStore((state) => state.collectionTitle);
const imageLoadError = useGalleryStore((state) => state.imageLoadError);
const galleryScrollResetKey = useGalleryStore((state) => state.galleryScrollResetKey);
const isSimilarResults = collectionTitle === "Similar Images";
const parentRef = useRef<HTMLDivElement>(null);
const [contextMenu, setContextMenu] = useState<{ x: number; y: number; image: ImageRecord } | null>(null);
@@ -271,6 +275,10 @@ export function Gallery() {
return () => element.removeEventListener("scroll", handleScroll);
}, [handleScroll]);
useEffect(() => {
parentRef.current?.scrollTo({ top: 0, left: 0 });
}, [galleryScrollResetKey]);
useEffect(() => {
const close = (event: PointerEvent) => {
if ((event.target as HTMLElement | null)?.closest("[data-gallery-context-menu]")) return;
@@ -293,12 +301,16 @@ export function Gallery() {
<div className="rounded-2xl border border-white/[0.06] bg-white/[0.02] p-8 min-w-72">
<div className="h-5 w-5 mx-auto rounded-full border-2 border-white/20 border-t-white/60 animate-spin" />
<p className="mt-4 text-sm text-white/40 font-medium">
{searchMode === "semantic" && search.trim().length > 0
{isSimilarResults
? "Finding similar images"
: searchMode === "semantic" && search.trim().length > 0
? `Searching for matches to "${search}"`
: "Loading media"}
</p>
<p className="text-xs text-white/20 mt-1">
{searchMode === "semantic" && search.trim().length > 0
{isSimilarResults
? "Comparing visual embeddings"
: searchMode === "semantic" && search.trim().length > 0
? "Semantic search can take a little longer than filename search"
: "Fetching results"}
</p>
@@ -316,12 +328,20 @@ export function Gallery() {
d="M4 16l4.586-4.586a2 2 0 012.828 0L16 16m-2-2l1.586-1.586a2 2 0 012.828 0L20 14m-6-6h.01M6 20h12a2 2 0 002-2V6a2 2 0 00-2-2H6a2 2 0 00-2 2v12a2 2 0 002 2z" />
</svg>
<p className="text-sm text-white/30 font-medium">
{searchMode === "semantic" && search.trim().length > 0
{imageLoadError
? "Could not load results"
: isSimilarResults
? "No similar images found"
: searchMode === "semantic" && search.trim().length > 0
? "No semantic matches found"
: "No media found"}
</p>
<p className="text-xs text-white/15 mt-1">
{searchMode === "semantic" && search.trim().length > 0
{imageLoadError
? imageLoadError
: isSimilarResults
? "This item may be visually isolated, or more embeddings may need to finish processing"
: searchMode === "semantic" && search.trim().length > 0
? "Try a broader phrase, or wait for more embeddings to finish processing"
: "Try adjusting your filters or add a new folder"}
</p>
+137
View File
@@ -47,12 +47,27 @@ function embeddingLabel(status: string, model: string | null): string {
export function Lightbox() {
const selectedImage = useGalleryStore((state) => state.selectedImage);
const selectedFolderId = useGalleryStore((state) => state.selectedFolderId);
const closeImage = useGalleryStore((state) => state.closeImage);
const images = useGalleryStore((state) => state.images);
const openImage = useGalleryStore((state) => state.openImage);
const loadSimilarImages = useGalleryStore((state) => state.loadSimilarImages);
const updateImageDetails = useGalleryStore((state) => state.updateImageDetails);
const suggestImageTags = useGalleryStore((state) => state.suggestImageTags);
const captionModelStatus = useGalleryStore((state) => state.captionModelStatus);
const captionModelPreparing = useGalleryStore((state) => state.captionModelPreparing);
const captionModelError = useGalleryStore((state) => state.captionModelError);
const captionModelProgress = useGalleryStore((state) => state.captionModelProgress);
const aiCaptionsEnabled = useGalleryStore((state) => state.aiCaptionsEnabled);
const setAiCaptionsEnabled = useGalleryStore((state) => state.setAiCaptionsEnabled);
const setSettingsOpen = useGalleryStore((state) => state.setSettingsOpen);
const prepareCaptionModel = useGalleryStore((state) => state.prepareCaptionModel);
const queueCaptionForImage = useGalleryStore((state) => state.queueCaptionForImage);
const queueCaptionJobs = useGalleryStore((state) => state.queueCaptionJobs);
const [zoom, setZoom] = useState(1);
const [suggestedTags, setSuggestedTags] = useState<string[]>([]);
const [captionQueueing, setCaptionQueueing] = useState(false);
const [captionQueueStatus, setCaptionQueueStatus] = useState<string | null>(null);
const imageViewportRef = useRef<HTMLDivElement>(null);
const currentIndex = selectedImage ? images.findIndex((image) => image.id === selectedImage.id) : -1;
@@ -68,8 +83,17 @@ export function Lightbox() {
useEffect(() => {
setZoom(1);
setSuggestedTags([]);
setCaptionQueueStatus(null);
}, [selectedImage?.id]);
useEffect(() => {
if (!selectedImage?.generated_caption) return;
void suggestImageTags(selectedImage.id)
.then(setSuggestedTags)
.catch(() => setSuggestedTags([]));
}, [selectedImage?.id, selectedImage?.generated_caption, suggestImageTags]);
useEffect(() => {
const viewport = imageViewportRef.current;
if (!viewport || !selectedImage || selectedImage.media_kind !== "image") return;
@@ -316,6 +340,119 @@ export function Lightbox() {
) : null}
</div>
<div>
<p className="mb-1 text-xs uppercase tracking-wider text-gray-500">AI Caption</p>
{selectedImage.generated_caption ? (
<>
<p className="text-white">{selectedImage.generated_caption}</p>
{suggestedTags.length > 0 ? (
<div className="mt-2 flex flex-wrap gap-1.5">
{suggestedTags.map((tag) => (
<span
key={tag}
className="rounded-md border border-white/10 bg-white/5 px-2 py-0.5 text-xs text-gray-300"
>
{tag}
</span>
))}
</div>
) : null}
{selectedImage.caption_model ? (
<p className="mt-1 text-xs text-gray-600">{selectedImage.caption_model}</p>
) : null}
</>
) : (
<>
<p className="text-gray-600">
{selectedImage.caption_error ?? "Not generated"}
</p>
{captionModelStatus?.ready ? (
<>
<p className={`mt-2 text-xs ${aiCaptionsEnabled ? "text-emerald-400/70" : "text-gray-600"}`}>
{aiCaptionsEnabled ? "Florence-2 enabled locally" : "Florence-2 downloaded but disabled"}
</p>
<div className="mt-2 flex flex-wrap gap-2">
<button
className="rounded-md border border-white/10 bg-white/5 px-3 py-1.5 text-xs text-gray-300 transition-colors hover:bg-white/10 hover:text-white"
onClick={() => setAiCaptionsEnabled(!aiCaptionsEnabled)}
>
{aiCaptionsEnabled ? "Disable captions" : "Enable captions"}
</button>
{selectedImage.media_kind === "image" && aiCaptionsEnabled ? (
<button
className="rounded-md border border-white/10 bg-white/5 px-3 py-1.5 text-xs text-gray-300 transition-colors hover:bg-white/10 hover:text-white disabled:cursor-not-allowed disabled:opacity-50"
onClick={() => {
setCaptionQueueing(true);
setCaptionQueueStatus(null);
void queueCaptionForImage(selectedImage.id)
.then(() => setCaptionQueueStatus("Queued for background captions."))
.catch((error) => setCaptionQueueStatus(String(error)))
.finally(() => setCaptionQueueing(false));
}}
disabled={captionQueueing}
>
{captionQueueing ? "Queueing..." : "Queue caption"}
</button>
) : null}
{aiCaptionsEnabled ? (
<button
className="rounded-md border border-white/10 bg-white/5 px-3 py-1.5 text-xs text-gray-300 transition-colors hover:bg-white/10 hover:text-white disabled:cursor-not-allowed disabled:opacity-50"
onClick={() => {
setCaptionQueueing(true);
setCaptionQueueStatus(null);
void queueCaptionJobs(selectedFolderId)
.then((queued) =>
setCaptionQueueStatus(
queued === 0
? "No missing captions found."
: `Queued ${queued.toLocaleString()} image${queued === 1 ? "" : "s"} for background captions.`,
),
)
.catch((error) => setCaptionQueueStatus(String(error)))
.finally(() => setCaptionQueueing(false));
}}
disabled={captionQueueing}
>
{selectedFolderId === null ? "Queue all" : "Queue folder"}
</button>
) : null}
</div>
{captionQueueStatus ? (
<p className="mt-2 break-all text-xs text-gray-500">{captionQueueStatus}</p>
) : null}
</>
) : (
<button
className="mt-2 rounded-md border border-white/10 bg-white/5 px-3 py-1.5 text-xs text-gray-300 transition-colors hover:bg-white/10 hover:text-white disabled:cursor-not-allowed disabled:opacity-50"
onClick={() => void prepareCaptionModel()}
disabled={captionModelPreparing}
>
{captionModelProgress
? `Downloading ${captionModelProgress.completed_files}/${captionModelProgress.total_files}`
: captionModelPreparing
? "Preparing Florence-2..."
: "Enable local captions"}
</button>
)}
{captionModelProgress?.current_file ? (
<p className="mt-2 break-all text-xs text-gray-600">{captionModelProgress.current_file}</p>
) : null}
{captionModelError ? (
<p className="mt-2 text-xs text-amber-300">{captionModelError}</p>
) : null}
{!captionModelStatus?.ready && !captionModelError ? (
<p className="mt-2 text-xs text-gray-600">Downloads Florence-2 on demand for offline captions.</p>
) : null}
<button
className="mt-2 block text-xs text-gray-500 transition-colors hover:text-gray-300"
onClick={() => setSettingsOpen(true)}
>
Open settings
</button>
</>
)}
</div>
<div>
<p className="mb-1 text-xs uppercase tracking-wider text-gray-500">Path</p>
<p className="break-all text-xs text-gray-400">{selectedImage.path}</p>
+354
View File
@@ -0,0 +1,354 @@
import { useEffect, useState } from "react";
import { useGalleryStore } from "../store";
type SettingsSection = "ai" | "library" | "display" | "storage";
const SECTIONS: { id: SettingsSection; label: string; detail: string }[] = [
{ id: "ai", label: "Local AI", detail: "Captions and suggestions" },
{ id: "library", label: "Library", detail: "Indexing and scanning" },
{ id: "display", label: "Display", detail: "Gallery preferences" },
{ id: "storage", label: "Storage", detail: "Cache and model files" },
];
function ToggleSwitch({
checked,
disabled,
onChange,
label,
}: {
checked: boolean;
disabled?: boolean;
onChange: (checked: boolean) => void;
label: string;
}) {
return (
<button
type="button"
role="switch"
aria-checked={checked}
aria-label={label}
disabled={disabled}
onClick={() => onChange(!checked)}
className={`relative h-7 w-12 shrink-0 rounded-full border transition-colors duration-150 focus:outline-none focus:ring-2 focus:ring-emerald-400/30 disabled:cursor-not-allowed disabled:opacity-45 ${
checked
? "border-emerald-400/50 bg-emerald-500/35"
: "border-white/12 bg-white/[0.06]"
}`}
>
<span
className={`absolute left-0.5 top-0.5 h-6 w-6 rounded-full bg-white shadow-lg shadow-black/40 transition-transform duration-150 ${
checked ? "translate-x-5" : "translate-x-0"
}`}
/>
</button>
);
}
function StatusPill({ children, tone }: { children: React.ReactNode; tone: "ready" | "muted" | "busy" }) {
const className =
tone === "ready"
? "border-emerald-400/25 bg-emerald-500/10 text-emerald-300"
: tone === "busy"
? "border-sky-400/25 bg-sky-500/10 text-sky-300"
: "border-white/10 bg-white/[0.04] text-gray-500";
return (
<span className={`inline-flex rounded-md border px-2 py-1 text-[11px] font-medium ${className}`}>
{children}
</span>
);
}
function SettingsRow({
title,
description,
children,
}: {
title: string;
description: string;
children: React.ReactNode;
}) {
return (
<div className="flex items-start justify-between gap-5 border-b border-white/[0.07] py-4 last:border-b-0">
<div className="min-w-0">
<p className="text-sm font-medium text-white">{title}</p>
<p className="mt-1 max-w-md text-xs leading-relaxed text-gray-500">{description}</p>
</div>
<div className="shrink-0">{children}</div>
</div>
);
}
function SectionShell({
eyebrow,
title,
children,
}: {
eyebrow: string;
title: string;
children: React.ReactNode;
}) {
return (
<div>
<p className="text-[10px] font-semibold uppercase tracking-[0.18em] text-gray-600">{eyebrow}</p>
<h3 className="mt-1 text-lg font-semibold text-white">{title}</h3>
<div className="mt-5 border-t border-white/[0.08]">{children}</div>
</div>
);
}
export function SettingsModal() {
const [activeSection, setActiveSection] = useState<SettingsSection>("ai");
const [captionQueueStatus, setCaptionQueueStatus] = useState<string | null>(null);
const [captionQueueing, setCaptionQueueing] = useState(false);
const settingsOpen = useGalleryStore((state) => state.settingsOpen);
const setSettingsOpen = useGalleryStore((state) => state.setSettingsOpen);
const folders = useGalleryStore((state) => state.folders);
const selectedFolderId = useGalleryStore((state) => state.selectedFolderId);
const captionModelStatus = useGalleryStore((state) => state.captionModelStatus);
const captionModelPreparing = useGalleryStore((state) => state.captionModelPreparing);
const captionModelProgress = useGalleryStore((state) => state.captionModelProgress);
const captionModelError = useGalleryStore((state) => state.captionModelError);
const captionRuntimeProbe = useGalleryStore((state) => state.captionRuntimeProbe);
const captionRuntimeChecking = useGalleryStore((state) => state.captionRuntimeChecking);
const aiCaptionsEnabled = useGalleryStore((state) => state.aiCaptionsEnabled);
const setAiCaptionsEnabled = useGalleryStore((state) => state.setAiCaptionsEnabled);
const prepareCaptionModel = useGalleryStore((state) => state.prepareCaptionModel);
const deleteCaptionModel = useGalleryStore((state) => state.deleteCaptionModel);
const probeCaptionRuntime = useGalleryStore((state) => state.probeCaptionRuntime);
const queueCaptionJobs = useGalleryStore((state) => state.queueCaptionJobs);
useEffect(() => {
if (!settingsOpen) return;
const handleKeyDown = (event: KeyboardEvent) => {
if (event.key === "Escape") setSettingsOpen(false);
};
window.addEventListener("keydown", handleKeyDown);
return () => window.removeEventListener("keydown", handleKeyDown);
}, [settingsOpen, setSettingsOpen]);
if (!settingsOpen) return null;
const modelReady = captionModelStatus?.ready ?? false;
const downloadLabel = captionModelProgress
? `Downloading ${captionModelProgress.completed_files}/${captionModelProgress.total_files}`
: captionModelPreparing
? "Preparing Florence-2..."
: modelReady
? "Downloaded"
: "Download Florence-2";
const downloadPercent = captionModelProgress
? Math.round((captionModelProgress.completed_files / Math.max(captionModelProgress.total_files, 1)) * 100)
: 0;
const selectedFolder = folders.find((folder) => folder.id === selectedFolderId);
const captionScopeLabel = selectedFolder ? selectedFolder.name : "all libraries";
return (
<div
className="fixed inset-0 z-[60] flex items-center justify-center bg-black/65 px-6 backdrop-blur-sm"
onClick={() => setSettingsOpen(false)}
>
<div
className="flex h-[min(680px,calc(100vh-56px))] w-full max-w-4xl overflow-hidden rounded-lg border border-white/10 bg-[#07080f] shadow-2xl shadow-black/60"
onClick={(event) => event.stopPropagation()}
>
<aside className="flex w-56 shrink-0 flex-col border-r border-white/[0.07] bg-white/[0.025]">
<div className="border-b border-white/[0.07] px-5 py-5">
<p className="text-base font-semibold text-white">Settings</p>
<p className="mt-1 text-xs text-gray-600">Phokus preferences</p>
</div>
<div className="flex-1 space-y-1 overflow-y-auto p-2">
{SECTIONS.map((section) => (
<button
key={section.id}
className={`w-full rounded-md px-3 py-2.5 text-left transition-colors ${
activeSection === section.id
? "bg-white/10 text-white"
: "text-gray-500 hover:bg-white/[0.055] hover:text-gray-200"
}`}
onClick={() => setActiveSection(section.id)}
>
<span className="block text-[13px] font-medium">{section.label}</span>
<span className="mt-0.5 block text-[11px] text-gray-600">{section.detail}</span>
</button>
))}
</div>
</aside>
<main className="flex min-w-0 flex-1 flex-col">
<div className="flex h-14 shrink-0 items-center justify-between border-b border-white/[0.07] px-6">
<div className="flex items-center gap-2">
{modelReady ? <StatusPill tone="ready">Florence-2 ready</StatusPill> : <StatusPill tone="muted">Optional</StatusPill>}
{captionModelPreparing ? <StatusPill tone="busy">Working</StatusPill> : null}
</div>
<button
className="rounded-md p-1.5 text-gray-500 transition-colors hover:bg-white/[0.06] hover:text-white"
onClick={() => setSettingsOpen(false)}
title="Close settings"
>
<svg className="h-4 w-4" fill="none" viewBox="0 0 24 24" stroke="currentColor">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M6 18L18 6M6 6l12 12" />
</svg>
</button>
</div>
<div className="flex-1 overflow-y-auto px-7 py-6">
{activeSection === "ai" ? (
<SectionShell eyebrow="Local AI" title="Captions and suggested tags">
<SettingsRow
title="AI captions"
description="Generate captions and suggested tags with the local Florence-2 model."
>
<ToggleSwitch
checked={aiCaptionsEnabled && modelReady}
disabled={!modelReady || captionModelPreparing}
onChange={setAiCaptionsEnabled}
label="AI captions"
/>
</SettingsRow>
<SettingsRow
title="Florence-2 model"
description={modelReady ? "Stored locally and available offline." : "Download the model before enabling local captions."}
>
<div className="flex flex-col items-end gap-2">
<button
className="relative overflow-hidden rounded-md border border-white/10 bg-white/[0.055] px-3 py-1.5 text-xs text-gray-300 transition-colors hover:bg-white/10 hover:text-white disabled:cursor-not-allowed disabled:opacity-45"
onClick={() => void prepareCaptionModel()}
disabled={captionModelPreparing || modelReady}
>
{captionModelProgress ? (
<span
className="absolute inset-y-0 left-0 bg-emerald-400/15 transition-[width] duration-200"
style={{ width: `${downloadPercent}%` }}
/>
) : null}
<span className="relative">{downloadLabel}</span>
</button>
{modelReady ? (
<>
<button
className="rounded-md border border-white/10 bg-white/[0.055] px-3 py-1.5 text-xs text-gray-300 transition-colors hover:bg-white/10 hover:text-white disabled:cursor-not-allowed disabled:opacity-45"
onClick={() => void probeCaptionRuntime()}
disabled={captionRuntimeChecking}
>
{captionRuntimeChecking ? "Checking runtime..." : "Check runtime"}
</button>
<button
className="rounded-md border border-red-400/20 bg-red-500/10 px-3 py-1.5 text-xs text-red-200 transition-colors hover:bg-red-500/15 disabled:cursor-not-allowed disabled:opacity-45"
onClick={() => void deleteCaptionModel()}
disabled={captionModelPreparing || captionRuntimeChecking}
>
Delete model files
</button>
</>
) : null}
</div>
</SettingsRow>
<SettingsRow
title="Caption queue"
description={`Generate missing captions in ${captionScopeLabel}. Captions update in the gallery as the local worker finishes each image.`}
>
<button
className="rounded-md border border-white/10 bg-white/[0.055] px-3 py-1.5 text-xs text-gray-300 transition-colors hover:bg-white/10 hover:text-white disabled:cursor-not-allowed disabled:opacity-45"
onClick={() => {
setCaptionQueueing(true);
setCaptionQueueStatus(null);
void queueCaptionJobs(selectedFolderId)
.then((queued) => {
setCaptionQueueStatus(
queued === 0
? "No missing captions found."
: `Queued ${queued.toLocaleString()} image${queued === 1 ? "" : "s"}.`,
);
})
.catch((error) => setCaptionQueueStatus(String(error)))
.finally(() => setCaptionQueueing(false));
}}
disabled={!modelReady || !aiCaptionsEnabled || captionQueueing}
>
{captionQueueing ? "Queueing..." : selectedFolder ? "Caption this library" : "Caption all libraries"}
</button>
</SettingsRow>
<div className="py-4">
<p className="text-xs font-medium text-gray-400">Model location</p>
<p className="mt-2 break-all rounded-md border border-white/[0.07] bg-black/20 px-3 py-2 text-xs text-gray-600">
{modelReady ? captionModelStatus?.local_dir : "Not downloaded"}
</p>
{captionModelProgress?.current_file ? (
<p className="mt-3 break-all text-xs text-gray-500">{captionModelProgress.current_file}</p>
) : null}
{captionModelError ? (
<p className="mt-3 text-xs text-amber-300">{captionModelError}</p>
) : null}
{captionQueueStatus ? (
<p className="mt-3 text-xs text-gray-500">{captionQueueStatus}</p>
) : null}
{captionRuntimeProbe ? (
<div className="mt-4 border-t border-white/[0.07] pt-4">
<div className="flex items-center justify-between gap-3">
<p className="text-xs font-medium text-gray-400">Runtime check</p>
<StatusPill tone="ready">Ready</StatusPill>
</div>
<p className="mt-2 text-xs text-gray-600">
Tokenizer vocabulary: {captionRuntimeProbe.tokenizer_vocab_size.toLocaleString()}
</p>
<div className="mt-3 space-y-2">
{captionRuntimeProbe.sessions.map((session) => (
<div key={session.file} className="rounded-md border border-white/[0.07] bg-black/20 px-3 py-2">
<p className="break-all text-xs text-gray-400">{session.file}</p>
<p className="mt-1 text-[11px] text-gray-600">
{session.inputs.length} input{session.inputs.length === 1 ? "" : "s"} · {session.outputs.length} output{session.outputs.length === 1 ? "" : "s"}
</p>
</div>
))}
</div>
</div>
) : null}
</div>
</SectionShell>
) : null}
{activeSection === "library" ? (
<SectionShell eyebrow="Library" title="Indexing and scanning">
<SettingsRow title="Background workers" description="Folder-level pause controls remain in the background tasks panel.">
<StatusPill tone="muted">Managed per folder</StatusPill>
</SettingsRow>
<SettingsRow title="Reindexing" description="Use the library sidebar to rescan a folder when files change.">
<StatusPill tone="muted">Available</StatusPill>
</SettingsRow>
</SectionShell>
) : null}
{activeSection === "display" ? (
<SectionShell eyebrow="Display" title="Gallery preferences">
<SettingsRow title="Grid density" description="Use the toolbar size control to change thumbnail density.">
<StatusPill tone="muted">Toolbar</StatusPill>
</SettingsRow>
<SettingsRow title="Result view" description="Similar results reset to the top on each new search.">
<StatusPill tone="ready">Enabled</StatusPill>
</SettingsRow>
</SectionShell>
) : null}
{activeSection === "storage" ? (
<SectionShell eyebrow="Storage" title="Local files">
<SettingsRow title="Florence-2" description="Remove the local model without changing the rest of the library.">
<button
className="rounded-md border border-red-400/20 bg-red-500/10 px-3 py-1.5 text-xs text-red-200 transition-colors hover:bg-red-500/15 disabled:cursor-not-allowed disabled:opacity-45"
onClick={() => void deleteCaptionModel()}
disabled={captionModelPreparing || !modelReady}
>
Delete model files
</button>
</SettingsRow>
</SectionShell>
) : null}
</div>
</main>
</div>
</div>
);
}
+13
View File
@@ -1,5 +1,6 @@
import { useState, useEffect } from "react";
import { getCurrentWindow } from "@tauri-apps/api/window";
import { useGalleryStore } from "../store";
// SVG icons for window controls
function MinimizeIcon() {
@@ -37,6 +38,7 @@ function CloseIcon() {
export function TitleBar() {
const [isMaximized, setIsMaximized] = useState(false);
const setSettingsOpen = useGalleryStore((state) => state.setSettingsOpen);
const appWindow = getCurrentWindow();
useEffect(() => {
@@ -85,6 +87,17 @@ export function TitleBar() {
className="flex items-stretch h-full"
style={{ WebkitAppRegion: "no-drag" } as React.CSSProperties}
>
<button
onClick={() => setSettingsOpen(true)}
title="Settings"
className="group flex h-full w-10 items-center justify-center text-gray-600 transition-colors duration-100 hover:bg-white/6 hover:text-gray-300"
>
<svg className="h-3.5 w-3.5" fill="none" viewBox="0 0 24 24" stroke="currentColor">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={1.8} d="M10.325 4.317c.426-1.756 2.924-1.756 3.35 0a1.724 1.724 0 002.573 1.066c1.543-.94 3.31.826 2.37 2.37a1.724 1.724 0 001.065 2.572c1.756.426 1.756 2.924 0 3.35a1.724 1.724 0 00-1.066 2.573c.94 1.543-.826 3.31-2.37 2.37a1.724 1.724 0 00-2.572 1.065c-.426 1.756-2.924 1.756-3.35 0a1.724 1.724 0 00-2.573-1.066c-1.543.94-3.31-.826-2.37-2.37a1.724 1.724 0 00-1.065-2.572c-1.756-.426-1.756-2.924 0-3.35a1.724 1.724 0 001.066-2.573c-.94-1.543.826-3.31 2.37-2.37.996.607 2.296.07 2.572-1.065z" />
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={1.8} d="M15 12a3 3 0 11-6 0 3 3 0 016 0z" />
</svg>
</button>
{/* Minimize */}
<button
onClick={handleMinimize}
+244 -27
View File
@@ -40,6 +40,10 @@ export interface ImageRecord {
embedding_model: string | null;
embedding_updated_at: string | null;
embedding_error: string | null;
generated_caption: string | null;
caption_model: string | null;
caption_updated_at: string | null;
caption_error: string | null;
}
export interface IndexProgress {
@@ -57,6 +61,9 @@ export interface FolderJobProgress {
embedding_pending: number;
embedding_ready: number;
embedding_failed: number;
caption_pending: number;
caption_ready: number;
caption_failed: number;
}
export interface MediaJobProgressEvent {
@@ -80,6 +87,39 @@ export interface TagCloudEntry {
thumbnail_path: string | null;
}
export interface CaptionModelStatus {
model_id: string;
model_name: string;
local_dir: string;
ready: boolean;
missing_files: string[];
}
export interface CaptionModelProgress {
total_files: number;
completed_files: number;
current_file: string | null;
done: boolean;
}
export interface CaptionRuntimeSessionProbe {
file: string;
inputs: string[];
outputs: string[];
}
export interface CaptionRuntimeProbe {
ready: boolean;
tokenizer_vocab_size: number;
sessions: CaptionRuntimeSessionProbe[];
}
export interface CaptionVisionProbe {
input_shape: number[];
output_shape: number[];
output_values: number;
}
export type SortOrder =
| "date_desc"
| "date_asc"
@@ -97,6 +137,7 @@ interface GalleryState {
totalImages: number;
loadedCount: number;
loadingImages: boolean;
imageLoadError: string | null;
search: string;
searchMode: SearchMode;
sort: SortOrder;
@@ -106,6 +147,9 @@ interface GalleryState {
zoomPreset: ZoomPreset;
selectedImage: ImageRecord | null;
collectionTitle: string | null;
similarSourceImageId: number | null;
similarHasMore: boolean;
galleryScrollResetKey: number;
activeView: ActiveView;
tagCloudEntries: TagCloudEntry[];
tagCloudLoading: boolean;
@@ -113,6 +157,14 @@ interface GalleryState {
indexingProgress: Record<number, IndexProgress>;
mediaJobProgress: Record<number, FolderJobProgress>;
cacheDir: string;
captionModelStatus: CaptionModelStatus | null;
captionModelPreparing: boolean;
captionModelError: string | null;
captionModelProgress: CaptionModelProgress | null;
captionRuntimeProbe: CaptionRuntimeProbe | null;
captionRuntimeChecking: boolean;
aiCaptionsEnabled: boolean;
settingsOpen: boolean;
loadFolders: () => Promise<void>;
loadBackgroundJobProgress: () => Promise<void>;
@@ -136,7 +188,18 @@ interface GalleryState {
setView: (view: ActiveView) => void;
loadTagCloud: () => Promise<void>;
searchByTag: (imageId: number) => void;
loadSimilarImages: (imageId: number) => Promise<void>;
loadSimilarImages: (imageId: number, folderId?: number | null, reset?: boolean) => Promise<void>;
suggestImageTags: (imageId: number) => Promise<string[]>;
loadCaptionModelStatus: () => Promise<void>;
prepareCaptionModel: () => Promise<void>;
deleteCaptionModel: () => Promise<void>;
probeCaptionRuntime: () => Promise<void>;
probeCaptionImage: (imageId: number) => Promise<CaptionVisionProbe>;
generateCaptionForImage: (imageId: number) => Promise<ImageRecord>;
queueCaptionJobs: (folderId?: number | null) => Promise<number>;
queueCaptionForImage: (imageId: number) => Promise<number>;
setAiCaptionsEnabled: (enabled: boolean) => void;
setSettingsOpen: (open: boolean) => void;
retryFailedEmbeddings: (folderId: number) => Promise<void>;
updateImageDetails: (imageId: number, updates: { favorite?: boolean; rating?: number }) => Promise<void>;
setCacheDir: (dir: string) => void;
@@ -144,6 +207,12 @@ interface GalleryState {
}
const PAGE_SIZE = 200;
const AI_CAPTIONS_ENABLED_KEY = "phokus.aiCaptionsEnabled";
function initialAiCaptionsEnabled(): boolean {
if (typeof window === "undefined") return false;
return window.localStorage.getItem(AI_CAPTIONS_ENABLED_KEY) === "true";
}
function mergeIntoVisibleWindow(
currentImages: ImageRecord[],
@@ -266,6 +335,7 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
totalImages: 0,
loadedCount: 0,
loadingImages: false,
imageLoadError: null,
search: "",
searchMode: "filename",
sort: "date_desc",
@@ -275,6 +345,9 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
zoomPreset: "comfortable",
selectedImage: null,
collectionTitle: null,
similarSourceImageId: null,
similarHasMore: false,
galleryScrollResetKey: 0,
activeView: "gallery",
tagCloudEntries: [],
tagCloudLoading: false,
@@ -282,6 +355,14 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
indexingProgress: {},
mediaJobProgress: {},
cacheDir: "",
captionModelStatus: null,
captionModelPreparing: false,
captionModelError: null,
captionModelProgress: null,
captionRuntimeProbe: null,
captionRuntimeChecking: false,
aiCaptionsEnabled: initialAiCaptionsEnabled(),
settingsOpen: false,
setCacheDir: (cacheDir) => set({ cacheDir }),
@@ -327,13 +408,13 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
},
selectFolder: (folderId) => {
set({ selectedFolderId: folderId, images: [], loadedCount: 0, collectionTitle: null, activeView: "gallery", failedEmbeddingsOnly: false });
set({ selectedFolderId: folderId, images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, activeView: "gallery", failedEmbeddingsOnly: false, imageLoadError: null });
void get().loadImages(true);
},
loadImages: async (reset = false) => {
const { selectedFolderId, search, searchMode, sort, loadedCount, mediaFilter, favoritesOnly, failedEmbeddingsOnly } = get();
set({ loadingImages: true });
set({ loadingImages: true, imageLoadError: null });
try {
if (searchMode === "semantic" && search.trim()) {
@@ -353,6 +434,8 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
loadedCount: images.length,
loadingImages: false,
collectionTitle: `Semantic search: ${search}`,
similarSourceImageId: null,
similarHasMore: false,
});
return;
}
@@ -382,56 +465,63 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
loadedCount: reset ? result.images.length : state.loadedCount + result.images.length,
loadingImages: false,
collectionTitle: reset ? null : state.collectionTitle,
similarSourceImageId: null,
similarHasMore: false,
}));
} catch (error) {
console.error("Failed to load media:", error);
set({ loadingImages: false });
set({ loadingImages: false, imageLoadError: String(error) });
}
},
loadMoreImages: async () => {
const { loadedCount, totalImages, loadingImages } = get();
const { loadedCount, totalImages, loadingImages, collectionTitle, similarSourceImageId, similarHasMore, selectedFolderId } = get();
if (loadingImages || loadedCount >= totalImages) return;
if (collectionTitle === "Similar Images" && similarSourceImageId !== null) {
if (!similarHasMore) return;
await get().loadSimilarImages(similarSourceImageId, selectedFolderId, false);
return;
}
await get().loadImages(false);
},
setSearch: (search) => {
set({ search, images: [], loadedCount: 0, collectionTitle: null });
set({ search, images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
clearSearch: () => {
set({ search: "", images: [], loadedCount: 0, collectionTitle: null });
set({ search: "", images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
resetSearch: () => {
set({ search: "", searchMode: "filename", images: [], loadedCount: 0, collectionTitle: null });
set({ search: "", searchMode: "filename", images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
setSearchMode: (searchMode) => {
set({ searchMode, images: [], loadedCount: 0, collectionTitle: null });
set({ searchMode, images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
setSort: (sort) => {
set({ sort, images: [], loadedCount: 0, collectionTitle: null });
set({ sort, images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
setMediaFilter: (mediaFilter) => {
set({ mediaFilter, images: [], loadedCount: 0, collectionTitle: null });
set({ mediaFilter, images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
setFavoritesOnly: (favoritesOnly) => {
set({ favoritesOnly, images: [], loadedCount: 0, collectionTitle: null });
set({ favoritesOnly, images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
setFailedEmbeddingsOnly: (failedEmbeddingsOnly) => {
set({ failedEmbeddingsOnly, images: [], loadedCount: 0, collectionTitle: null });
set({ failedEmbeddingsOnly, images: [], loadedCount: 0, collectionTitle: null, similarSourceImageId: null, similarHasMore: false, imageLoadError: null });
void get().loadImages(true);
},
@@ -461,26 +551,145 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
},
searchByTag: (imageId) => {
set({ activeView: "gallery", images: [], loadedCount: 0, loadingImages: true, collectionTitle: "Similar Images" });
void get().loadSimilarImages(imageId);
const { selectedFolderId } = get();
set((state) => ({ activeView: "gallery", images: [], loadedCount: 0, loadingImages: true, collectionTitle: "Similar Images", imageLoadError: null, galleryScrollResetKey: state.galleryScrollResetKey + 1 }));
void get().loadSimilarImages(imageId, selectedFolderId);
},
loadSimilarImages: async (imageId) => {
set({ images: [], loadedCount: 0, loadingImages: true, collectionTitle: "Similar Images" });
const images = await invoke<ImageRecord[]>("find_similar_images", {
params: { image_id: imageId, limit: PAGE_SIZE },
});
set({
images,
totalImages: images.length,
loadedCount: images.length,
loadingImages: false,
loadSimilarImages: async (imageId, folderId = get().selectedFolderId, reset = true) => {
const requestedLimit = reset ? PAGE_SIZE : get().loadedCount + PAGE_SIZE;
set((state) => ({
images: reset ? [] : get().images,
loadedCount: reset ? 0 : get().loadedCount,
loadingImages: true,
collectionTitle: "Similar Images",
selectedFolderId: null,
selectedImage: null,
imageLoadError: null,
similarSourceImageId: imageId,
galleryScrollResetKey: reset ? state.galleryScrollResetKey + 1 : state.galleryScrollResetKey,
}));
try {
const images = await invoke<ImageRecord[]>("find_similar_images", {
params: { image_id: imageId, limit: requestedLimit },
});
const hasMore = images.length >= requestedLimit;
set({
images,
totalImages: hasMore ? images.length + PAGE_SIZE : images.length,
loadedCount: images.length,
loadingImages: false,
imageLoadError: null,
collectionTitle: "Similar Images",
similarSourceImageId: imageId,
similarHasMore: hasMore,
selectedFolderId: folderId ?? null,
selectedImage: reset ? null : get().selectedImage,
});
} catch (error) {
console.error("Failed to load similar images:", error);
set({
images: [],
totalImages: 0,
loadedCount: 0,
loadingImages: false,
imageLoadError: String(error),
collectionTitle: "Similar Images",
similarSourceImageId: imageId,
similarHasMore: false,
selectedFolderId: folderId ?? null,
selectedImage: null,
});
}
},
suggestImageTags: async (imageId) => {
return invoke<string[]>("suggest_image_tags", {
params: { image_id: imageId, limit: 2 },
});
},
loadCaptionModelStatus: async () => {
try {
const captionModelStatus = await invoke<CaptionModelStatus>("get_caption_model_status");
set({ captionModelStatus, captionModelError: null });
} catch (error) {
set({ captionModelError: String(error) });
}
},
prepareCaptionModel: async () => {
set({ captionModelPreparing: true, captionModelError: null, captionModelProgress: null });
try {
const captionModelStatus = await invoke<CaptionModelStatus>("prepare_caption_model");
window.localStorage.setItem(AI_CAPTIONS_ENABLED_KEY, String(captionModelStatus.ready));
set({ captionModelStatus, captionModelPreparing: false, captionModelError: null, captionModelProgress: null, aiCaptionsEnabled: captionModelStatus.ready });
} catch (error) {
set({ captionModelPreparing: false, captionModelError: String(error), captionModelProgress: null });
}
},
deleteCaptionModel: async () => {
set({ captionModelPreparing: true, captionModelError: null, captionModelProgress: null });
try {
const captionModelStatus = await invoke<CaptionModelStatus>("delete_caption_model");
window.localStorage.setItem(AI_CAPTIONS_ENABLED_KEY, "false");
set({ captionModelStatus, captionModelPreparing: false, captionModelError: null, captionModelProgress: null, captionRuntimeProbe: null, aiCaptionsEnabled: false });
} catch (error) {
set({ captionModelPreparing: false, captionModelError: String(error), captionModelProgress: null });
}
},
probeCaptionRuntime: async () => {
set({ captionRuntimeChecking: true, captionModelError: null });
try {
const captionRuntimeProbe = await invoke<CaptionRuntimeProbe>("probe_caption_runtime");
set({ captionRuntimeProbe, captionRuntimeChecking: false, captionModelError: null });
} catch (error) {
set({ captionRuntimeChecking: false, captionModelError: String(error), captionRuntimeProbe: null });
}
},
probeCaptionImage: async (imageId) => {
return invoke<CaptionVisionProbe>("probe_caption_image", {
params: { image_id: imageId },
});
},
generateCaptionForImage: async (imageId) => {
const updatedImage = await invoke<ImageRecord>("generate_caption_for_image", {
params: { image_id: imageId },
});
set((state) => ({
images: replaceImage(state.images, updatedImage, state.sort),
selectedImage: state.selectedImage?.id === updatedImage.id ? updatedImage : state.selectedImage,
}));
return updatedImage;
},
queueCaptionJobs: async (folderId = get().selectedFolderId) => {
const queued = await invoke<number>("queue_caption_jobs", {
params: { folder_id: folderId ?? null, image_id: null },
});
await get().loadBackgroundJobProgress();
return queued;
},
queueCaptionForImage: async (imageId) => {
const queued = await invoke<number>("queue_caption_jobs", {
params: { folder_id: null, image_id: imageId },
});
await get().loadBackgroundJobProgress();
return queued;
},
setAiCaptionsEnabled: (aiCaptionsEnabled) => {
window.localStorage.setItem(AI_CAPTIONS_ENABLED_KEY, String(aiCaptionsEnabled));
set({ aiCaptionsEnabled });
},
setSettingsOpen: (settingsOpen) => set({ settingsOpen }),
retryFailedEmbeddings: async (folderId) => {
await invoke("retry_failed_embeddings", { params: { folder_id: folderId } });
await get().loadBackgroundJobProgress();
@@ -536,6 +745,13 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
});
});
const unlistenCaptionModelProgress = await listen<CaptionModelProgress>("caption-model-progress", (event) => {
set({
captionModelProgress: event.payload.done ? null : event.payload,
captionModelPreparing: !event.payload.done,
});
});
const unlistenImages = await listen<IndexedImagesBatch>("indexed-images", (event) => {
const batch = event.payload;
@@ -600,6 +816,7 @@ export const useGalleryStore = create<GalleryState>((set, get) => ({
return () => {
unlistenProgress();
unlistenMediaJobs();
unlistenCaptionModelProgress();
unlistenImages();
unlistenThumbnails();
};