Skip to main content

cpd_semantic/embed/
mod.rs

1// embed — the embedding side of `--semantic` (experimental).
2//
3// The search lives in `crate::search`; this module turns function texts
4// into vectors with one of two providers behind one on-disk vector cache:
5//
6// - `local` (the default): a model downloaded once with
7//   `--semantic-download` and run in-process on the CPU, so a scan makes no
8//   network call;
9// - `http`: any OpenAI-compatible embeddings API — Ollama, LM Studio,
10//   llama.cpp's llama-server, text-embeddings-inference, hosted APIs. The
11//   API key, when one is needed, comes from the environment only, never
12//   from a flag or a config file that could be committed.
13
14mod bert;
15mod cache;
16pub mod catalog;
17mod http;
18mod jina_bert;
19mod local;
20pub mod models;
21mod nomic_bert;
22
23use crate::search::{Embedder, SemanticScope};
24use serde::Serialize;
25use serde_json::{Map, Value};
26use std::io::IsTerminal;
27use std::path::PathBuf;
28use std::sync::Arc;
29
30pub const DEFAULT_URL: &str = "http://localhost:11434/v1";
31/// The default model of the local provider, by its name in
32/// [`catalog::KNOWN_MODELS`].
33pub const DEFAULT_LOCAL_MODEL: &str = "CodeRankEmbed";
34/// The default model of an embeddings API: jina-embeddings-v2-base-code
35/// under the name Ollama serves it by. Ollama's library has no copy of the
36/// local default.
37pub const DEFAULT_HTTP_MODEL: &str = "unclemusclez/jina-embeddings-v2-base-code";
38pub const API_KEY_ENV: &str = "JSCPD_SEMANTIC_API_KEY";
39
40/// Longest function text embedded, in bytes; a longer function is embedded
41/// by its head, which is where its signature and main logic are.
42const MAX_TEXT_BYTES: usize = 8_000;
43
44/// Where vectors come from.
45#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
46#[serde(rename_all = "lowercase")]
47pub enum Provider {
48    /// A downloaded model run in-process.
49    Local,
50    /// An OpenAI-compatible embeddings API.
51    Http,
52}
53
54impl std::str::FromStr for Provider {
55    type Err = String;
56
57    fn from_str(s: &str) -> Result<Self, Self::Err> {
58        match s.trim().to_ascii_lowercase().as_str() {
59            "local" => Ok(Provider::Local),
60            "http" => Ok(Provider::Http),
61            other => Err(format!("unknown provider '{other}': must be local or http")),
62        }
63    }
64}
65
66/// What `--semantic` runs with, after merging flags and the config file.
67#[derive(Debug, Clone, PartialEq, Serialize)]
68pub struct SemanticOptions {
69    pub provider: Provider,
70    /// Lowest cosine similarity of a pair across languages.
71    pub threshold: f32,
72    /// Lowest cosine similarity of a pair within one language; `None`
73    /// keeps the model's gap above `threshold` (see [`catalog::thresholds`]).
74    #[serde(skip_serializing_if = "Option::is_none")]
75    pub same_threshold: Option<f32>,
76    #[serde(serialize_with = "scope_name")]
77    pub scope: SemanticScope,
78    pub model: String,
79    /// The embeddings API; only the `http` provider uses it.
80    pub url: String,
81    #[serde(skip_serializing_if = "Option::is_none")]
82    pub dimensions: Option<u32>,
83    /// Extra request fields for an API, e.g. `{"task": "code2code.query"}`.
84    #[serde(skip_serializing_if = "Map::is_empty")]
85    pub params: Map<String, Value>,
86    /// Put before every function text; `None` is the prefix of a model in
87    /// [`catalog`], or none.
88    #[serde(skip_serializing_if = "Option::is_none")]
89    pub prefix: Option<String>,
90    pub cache: bool,
91    /// Whether `url` came from a config file rather than the command line.
92    /// A config file is shared, and can arrive with the code being scanned,
93    /// so such a URL, unless it is on this machine, gets no API key, and
94    /// gets code only when `--semantic` itself was typed.
95    #[serde(skip)]
96    pub url_from_config: bool,
97    /// Whether `--semantic` was given on the command line, rather than
98    /// turned on by a config file.
99    #[serde(skip)]
100    pub on_command_line: bool,
101    /// `--semantic-rebuild-cache`: embed every function again and replace
102    /// the cached vectors of this model and request shape.
103    #[serde(skip)]
104    pub rebuild_cache: bool,
105}
106
107fn scope_name<S: serde::Serializer>(scope: &SemanticScope, s: S) -> Result<S::Ok, S::Error> {
108    s.serialize_str(scope.as_str())
109}
110
111impl Default for SemanticOptions {
112    fn default() -> Self {
113        Self {
114            provider: Provider::Local,
115            threshold: catalog::thresholds(DEFAULT_LOCAL_MODEL, None, None).across,
116            same_threshold: None,
117            scope: SemanticScope::All,
118            model: DEFAULT_LOCAL_MODEL.to_string(),
119            url: DEFAULT_URL.to_string(),
120            dimensions: None,
121            params: Map::new(),
122            prefix: None,
123            cache: true,
124            url_from_config: false,
125            on_command_line: false,
126            rebuild_cache: false,
127        }
128    }
129}
130
131impl SemanticOptions {
132    /// The thresholds of the rules: the model's (see [`catalog`]), with
133    /// [`Self::threshold`] and [`Self::same_threshold`] put in.
134    pub fn thresholds(&self) -> crate::search::Thresholds {
135        catalog::thresholds(&self.model, Some(self.threshold), self.same_threshold)
136    }
137}
138
139/// One way of turning texts into vectors.
140trait Backend: Send + Sync {
141    /// The model and where it runs, for the progress line.
142    fn label(&self) -> String;
143    /// The model's name, for the cache file name.
144    fn model_name(&self) -> &str;
145    /// Everything that changes the vectors of a text, for the cache key.
146    fn cache_identity(&self) -> Value;
147    /// One vector per text, in order; `progress` gets the number done.
148    fn embed(&self, texts: &[&str], progress: &dyn Fn(usize)) -> Result<Vec<Vec<f32>>, String>;
149}
150
151/// The embedder `options` asks for, keeping its vectors in the cache folder
152/// of the `scanned` paths. The local provider's model must be downloaded
153/// already (see [`download`]); this fails before any scanning starts when it
154/// is not.
155pub fn embedder(
156    options: &SemanticOptions,
157    scanned: &[PathBuf],
158    quiet: bool,
159) -> Result<Arc<dyn Embedder>, String> {
160    let root = cache::root();
161    let backend: Box<dyn Backend> = match options.provider {
162        Provider::Http => Box::new(http::HttpBackend::new(options)?),
163        Provider::Local => {
164            let (known, model) = local_model(options)?;
165            let root = root.clone().ok_or_else(no_cache_dir)?;
166            let dir = model.dir(&root);
167            if !model.is_downloaded(&dir) {
168                let download = match known.name == DEFAULT_LOCAL_MODEL {
169                    true => "jscpd --semantic-download".to_string(),
170                    false => format!("jscpd --semantic-download {}", known.name),
171                };
172                return Err(format!(
173                    "the model {} is not downloaded yet. Run `{download}` once ({:.0} MB into {}), or use an embeddings API with --semantic-url",
174                    model.id,
175                    model.size() as f64 / 1e6,
176                    dir.display()
177                ));
178            }
179            Box::new(local::LocalBackend::new(model, dir))
180        }
181    };
182    let prefix = match &options.prefix {
183        Some(prefix) => prefix.clone(),
184        None => catalog::find(&options.model)
185            .map_or("", |m| m.prefix)
186            .to_string(),
187    };
188    let cache_file = match options.cache {
189        true => root.map(|r| {
190            cache::project_dir(&r, scanned).join(cache::file_name(
191                backend.model_name(),
192                &cache_identity(backend.as_ref(), &prefix),
193            ))
194        }),
195        false => None,
196    };
197    Ok(Arc::new(Cached {
198        backend,
199        prefix,
200        cache_file,
201        rebuild: options.rebuild_cache,
202        quiet,
203    }))
204}
205
206/// Everything that changes the vectors of a text: what the backend says,
207/// and the prefix when there is one.
208fn cache_identity(backend: &dyn Backend, prefix: &str) -> Value {
209    let mut identity = backend.cache_identity();
210    if let (false, Value::Object(fields)) = (prefix.is_empty(), &mut identity) {
211        fields.insert("prefix".into(), Value::from(prefix));
212    }
213    identity
214}
215
216/// The local model `options` needs and has not downloaded yet: its id and
217/// its size in bytes, so a front end can ask before fetching it. `None` for
218/// an embeddings API, and once the model is there.
219pub fn missing_model(options: &SemanticOptions) -> Option<(String, u64)> {
220    if options.provider != Provider::Local {
221        return None;
222    }
223    let (_, model) = local_model(options).ok()?;
224    let dir = model.dir(&cache::root()?);
225    (!model.is_downloaded(&dir)).then(|| (model.id.to_string(), model.size()))
226}
227
228/// `--semantic-download`: fetch the local model's files that are missing,
229/// verified against their pinned checksums. Returns their directory.
230pub fn download(options: &SemanticOptions, quiet: bool) -> Result<PathBuf, String> {
231    if options.provider == Provider::Http {
232        return Err(
233            "--semantic-download fetches a model for the local provider; an embeddings API needs none"
234                .to_string(),
235        );
236    }
237    let (_, model) = local_model(options)?;
238    let dir = model.dir(&cache::root().ok_or_else(no_cache_dir)?);
239    if model.is_downloaded(&dir) {
240        if !quiet {
241            eprintln!("{} is already downloaded: {}", model.id, dir.display());
242        }
243        return Ok(dir);
244    }
245    model.download(&dir, &http::agent(), quiet)?;
246    Ok(dir)
247}
248
249/// `--semantic-models`: the models jscpd has thresholds for, as a table,
250/// with where each one runs and what the columns mean.
251pub fn model_list() -> String {
252    let root = cache::root();
253    let mut rows = vec![["MODEL", "CROSS", "SAME", "LICENSE", "RUNS"].map(String::from)];
254    for model in catalog::KNOWN_MODELS {
255        let name = match model.name == DEFAULT_LOCAL_MODEL {
256            true => format!("{} (default)", model.name),
257            false => model.name.to_string(),
258        };
259        let runs = match (model.local, model.ollama.first()) {
260            (Some(local), _) => {
261                let downloaded = root
262                    .as_ref()
263                    .is_some_and(|r| local.is_stamped(&local.dir(r)));
264                format!(
265                    "in jscpd, {:.0} MB{}",
266                    local.size() as f64 / 1e6,
267                    if downloaded { ", downloaded" } else { "" }
268                )
269            }
270            (None, Some(ollama)) => format!("API (Ollama: {ollama})"),
271            (None, None) => "API".to_string(),
272        };
273        rows.push([
274            name,
275            model.threshold.to_string(),
276            model.same_threshold.to_string(),
277            model.license.to_string(),
278            runs,
279        ]);
280    }
281    let widths: Vec<usize> = (0..5)
282        .map(|c| rows.iter().map(|r| r[c].len()).max().unwrap_or(0))
283        .collect();
284    let mut out = String::new();
285    for row in &rows {
286        let cells: Vec<String> = row
287            .iter()
288            .zip(&widths)
289            .map(|(cell, &width)| format!("{cell:width$}"))
290            .collect();
291        out.push_str(cells.join("  ").trim_end());
292        out.push('\n');
293    }
294    out.push_str(
295        "\nCROSS is the default --semantic-threshold, for pairs across languages, and SAME\n\
296         the default --semantic-same-threshold, for pairs within one language.\n\
297         --semantic-model takes the name in the MODEL column, in any letter case, or\n\
298         the model's Hugging Face id.\n\
299         jscpd runs the models marked \"in jscpd\" on this machine once\n\
300         `jscpd --semantic-download <model>` has fetched them; the\n\
301         others need an embeddings API that serves them, given with --semantic-url.\n",
302    );
303    out
304}
305
306/// The model the local provider runs for `options`, and its files.
307fn local_model(
308    options: &SemanticOptions,
309) -> Result<(&'static catalog::KnownModel, &'static models::LocalModel), String> {
310    match catalog::find(&options.model) {
311        Some(
312            known @ catalog::KnownModel {
313                local: Some(model), ..
314            },
315        ) => Ok((known, model)),
316        Some(known) => Err(format!(
317            "jscpd does not run {} itself (it runs {}); serve it with an embeddings API and pass --semantic-url",
318            known.name,
319            catalog::local_names()
320        )),
321        None => Err(format!(
322            "jscpd runs {} itself, not '{}' (`jscpd --semantic-models` lists the models it knows); for another model use an embeddings API with --semantic-url",
323            catalog::local_names(),
324            options.model
325        )),
326    }
327}
328
329fn no_cache_dir() -> String {
330    format!(
331        "no cache directory to keep the model in: set {}",
332        cache::CACHE_DIR_ENV
333    )
334}
335
336/// A backend behind the vector cache: identical texts and texts embedded by
337/// an earlier run are never embedded again.
338struct Cached {
339    backend: Box<dyn Backend>,
340    /// Put before every text the backend embeds; see [`catalog`].
341    prefix: String,
342    cache_file: Option<PathBuf>,
343    /// Read nothing from `cache_file`, and replace it with this run's
344    /// vectors.
345    rebuild: bool,
346    quiet: bool,
347}
348
349impl Cached {
350    fn note(&self, message: &str) {
351        if !self.quiet {
352            eprintln!("{message}");
353        }
354    }
355}
356
357impl Embedder for Cached {
358    fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, String> {
359        use xxhash_rust::xxh3::xxh3_128;
360        let texts: Vec<&str> = texts.iter().map(|t| clip(t, MAX_TEXT_BYTES)).collect();
361        let keys: Vec<u128> = texts.iter().map(|t| xxh3_128(t.as_bytes())).collect();
362        let mut cache = match &self.cache_file {
363            Some(_) if self.rebuild => cache::Cache::replacing(),
364            Some(path) => cache::load(path),
365            None => cache::Cache::default(),
366        };
367        let mut announced = false;
368        loop {
369            let mut seen = std::collections::HashSet::new();
370            let missing: Vec<usize> = (0..keys.len())
371                .filter(|&i| !cache.vectors.contains_key(&keys[i]) && seen.insert(keys[i]))
372                .collect();
373            if !announced {
374                let distinct = keys.iter().collect::<std::collections::HashSet<_>>().len();
375                let mut line = announcement(
376                    texts.len(),
377                    distinct - missing.len(),
378                    missing.len(),
379                    &self.backend.label(),
380                    self.backend.model_name(),
381                );
382                if self.rebuild && self.cache_file.is_some() {
383                    line.push_str(", rebuilding the cache");
384                }
385                self.note(&line);
386                announced = true;
387            }
388            if missing.is_empty() {
389                return Ok(keys.iter().map(|k| cache.vectors[k].clone()).collect());
390            }
391            let prefixed: Vec<String> = missing
392                .iter()
393                .map(|&i| format!("{}{}", self.prefix, texts[i]))
394                .collect();
395            let batch: Vec<&str> = prefixed.iter().map(String::as_str).collect();
396            let live = !self.quiet && std::io::stderr().is_terminal();
397            let progress = |done: usize| {
398                if live {
399                    eprint!("\r  embedded {done} of {}", batch.len());
400                }
401            };
402            let vectors = self.backend.embed(&batch, &progress)?;
403            if live {
404                eprintln!();
405            }
406            let dims = vectors.first().map_or(0, Vec::len);
407            if cache.dims != 0 && cache.dims != dims {
408                // Same name, other vectors: the model changed. Nothing
409                // cached can be mixed with the new ones.
410                cache = cache::Cache::default();
411                continue;
412            }
413            cache.dims = dims;
414            let fresh: Vec<(u128, Vec<f32>)> =
415                missing.iter().map(|&i| keys[i]).zip(vectors).collect();
416            if let Some(path) = &self.cache_file
417                && let Err(e) = cache::save(path, &mut cache, &fresh, &keys)
418            {
419                self.note(&format!(
420                    "Warning: --semantic: cache {} not written: {e}",
421                    path.display()
422                ));
423            }
424            cache.vectors.extend(fresh);
425        }
426    }
427}
428
429/// The line announcing a run's embedding: of `total` functions, the
430/// distinct texts found in the cache and the ones to embed. Identical texts
431/// are embedded once, which is not caching.
432fn announcement(total: usize, cached: usize, todo: usize, label: &str, model: &str) -> String {
433    match (todo, cached) {
434        (0, _) => {
435            format!(
436                "Semantic clones (experimental): {total} functions, all embeddings cached ({model})"
437            )
438        }
439        (_, 0) => {
440            format!("Semantic clones (experimental): embedding {total} functions with {label}")
441        }
442        _ => format!(
443            "Semantic clones (experimental): embedding {todo} of {total} functions with {label}, the rest cached"
444        ),
445    }
446}
447
448/// The longest prefix of `text` of at most `max` bytes, cut at a char
449/// boundary.
450fn clip(text: &str, max: usize) -> &str {
451    if text.len() <= max {
452        return text;
453    }
454    let mut end = max;
455    while !text.is_char_boundary(end) {
456        end -= 1;
457    }
458    &text[..end]
459}
460
461/// A fresh directory in the system temp dir for one test of this module.
462#[cfg(test)]
463pub(crate) fn test_dir(name: &str) -> std::path::PathBuf {
464    let dir = std::env::temp_dir().join(format!("jscpd-semantic-{name}-{}", std::process::id()));
465    let _ = std::fs::remove_dir_all(&dir);
466    std::fs::create_dir_all(&dir).unwrap();
467    dir
468}
469#[cfg(test)]
470mod tests {
471    use super::*;
472
473    #[test]
474    fn clip_cuts_at_a_char_boundary() {
475        assert_eq!(clip("abc", 10), "abc");
476        assert_eq!(clip("aé", 2), "a");
477        assert_eq!(clip("aéb", 3), "aé");
478    }
479
480    #[test]
481    fn the_announcement_counts_duplicates_as_embedded_not_cached() {
482        let line = |cached, todo| announcement(897, cached, todo, "m here", "m");
483        assert!(line(0, 890).ends_with("embedding 897 functions with m here"));
484        assert!(
485            line(10, 880).ends_with("embedding 880 of 897 functions with m here, the rest cached")
486        );
487        assert!(line(890, 0).ends_with("897 functions, all embeddings cached (m)"));
488    }
489
490    /// A backend that remembers the texts it was given.
491    struct Recorder(std::sync::Arc<std::sync::Mutex<Vec<String>>>);
492
493    impl Backend for Recorder {
494        fn label(&self) -> String {
495            "recorder".into()
496        }
497
498        fn model_name(&self) -> &str {
499            "recorder"
500        }
501
502        fn cache_identity(&self) -> Value {
503            serde_json::json!({"model": "recorder"})
504        }
505
506        fn embed(&self, texts: &[&str], _: &dyn Fn(usize)) -> Result<Vec<Vec<f32>>, String> {
507            let mut seen = self.0.lock().unwrap();
508            seen.extend(texts.iter().map(|t| t.to_string()));
509            Ok(texts.iter().map(|t| vec![t.len() as f32, 1.0]).collect())
510        }
511    }
512
513    #[test]
514    fn the_prefix_goes_before_every_text_and_into_the_cache_key() {
515        let seen = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
516        let cached = Cached {
517            backend: Box::new(Recorder(seen.clone())),
518            prefix: "Code: ".into(),
519            cache_file: None,
520            rebuild: false,
521            quiet: true,
522        };
523        assert_eq!(cached.embed(&["a", "bb", "a"]).unwrap().len(), 3);
524        assert_eq!(*seen.lock().unwrap(), ["Code: a", "Code: bb"]);
525
526        let backend = Recorder(seen);
527        assert_eq!(
528            cache_identity(&backend, ""),
529            backend.cache_identity(),
530            "the caches of a model without a prefix stay valid"
531        );
532        assert_eq!(cache_identity(&backend, "Code: ")["prefix"], "Code: ");
533    }
534
535    #[test]
536    fn provider_names() {
537        assert_eq!("LOCAL".parse::<Provider>(), Ok(Provider::Local));
538        assert_eq!("http".parse::<Provider>(), Ok(Provider::Http));
539        assert!(
540            "grpc"
541                .parse::<Provider>()
542                .unwrap_err()
543                .contains("local or http")
544        );
545    }
546
547    #[test]
548    fn an_unknown_local_model_names_the_known_ones() {
549        let options = SemanticOptions {
550            model: "some/other-model".into(),
551            ..SemanticOptions::default()
552        };
553        let err = embedder(&options, &[], true).err().unwrap();
554        assert!(
555            err.contains("jscpd runs CodeRankEmbed, jina-embeddings-v2-base-code itself, not 'some/other-model'"),
556            "{err}"
557        );
558        assert!(err.contains("--semantic-url"), "{err}");
559        let api_only = SemanticOptions {
560            model: "qwen3-embedding:0.6b".into(),
561            ..SemanticOptions::default()
562        };
563        let err = embedder(&api_only, &[], true).err().unwrap();
564        assert!(
565            err.contains("jscpd does not run Qwen3-Embedding-0.6B itself"),
566            "{err}"
567        );
568        let http = SemanticOptions {
569            provider: Provider::Http,
570            ..SemanticOptions::default()
571        };
572        assert!(download(&http, true).unwrap_err().contains("needs none"));
573    }
574}