summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--Cargo.lock1
-rw-r--r--Cargo.toml1
-rw-r--r--src/config/mod.rs11
-rw-r--r--src/repl/completer.rs2
-rw-r--r--src/utils/mod.rs25
5 files changed, 14 insertions, 26 deletions
diff --git a/Cargo.lock b/Cargo.lock
index 3a64cd0..ecfd9ae 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -59,6 +59,7 @@ dependencies = [
"dirs",
"fancy-regex",
"futures-util",
+ "fuzzy-matcher",
"hmac",
"hnsw_rs",
"html_to_markdown",
diff --git a/Cargo.toml b/Cargo.toml
index 1af8aa3..7add0be 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -66,6 +66,7 @@ rust-embed = "8.5.0"
os_info = { version = "3.8.2", default-features = false }
bm25 = { version = "2.0.1", features = ["parallelism"] }
which = "7.0.1"
+fuzzy-matcher = "0.3.7"
[dependencies.reqwest]
version = "0.12.0"
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 98bcab3..ec6a930 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1850,10 +1850,15 @@ impl Config {
}
values.extend(complete_agent_variables(args[0]));
};
- values
+ let mut values_with_score: Vec<_> = values
.into_iter()
- .filter(|(value, _)| fuzzy_match(value, filter))
- .collect()
+ .filter_map(|v| {
+ let score = fuzzy_match(&v.0, filter)?;
+ Some((v, score))
+ })
+ .collect();
+ values_with_score.sort_unstable_by(|a, b| b.1.cmp(&a.1));
+ values_with_score.into_iter().map(|v| v.0).collect()
}
pub fn sync_models_url(&self) -> String {
diff --git a/src/repl/completer.rs b/src/repl/completer.rs
index 12e871a..bcce694 100644
--- a/src/repl/completer.rs
+++ b/src/repl/completer.rs
@@ -45,7 +45,7 @@ impl Completer for ReplCompleter {
if line == "." {
return true;
}
- line.starts_with(&cmd.name[..2]) && fuzzy_match(cmd.name, &line)
+ line.starts_with(&cmd.name[..2]) && fuzzy_match(cmd.name, &line).is_some()
})
.collect();
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index 8d66470..f4b98c1 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -24,6 +24,7 @@ pub use self::variables::*;
use anyhow::{Context, Result};
use fancy_regex::Regex;
+use fuzzy_matcher::{skim::SkimMatcherV2, FuzzyMatcher};
use is_terminal::IsTerminal;
use std::{env, path::PathBuf, process};
use unicode_segmentation::UnicodeSegmentation;
@@ -118,21 +119,8 @@ pub fn convert_option_string(value: &str) -> Option<String> {
}
}
-pub fn fuzzy_match(text: &str, pattern: &str) -> bool {
- let text_chars: Vec<char> = text.chars().collect();
- let pattern_chars: Vec<char> = pattern.chars().collect();
-
- let mut pattern_index = 0;
- let mut text_index = 0;
-
- while pattern_index < pattern_chars.len() && text_index < text_chars.len() {
- if pattern_chars[pattern_index] == text_chars[text_index] {
- pattern_index += 1;
- }
- text_index += 1;
- }
-
- pattern_index == pattern_chars.len()
+pub fn fuzzy_match(choice: &str, pattern: &str) -> Option<i64> {
+ SkimMatcherV2::default().fuzzy_match(choice, pattern)
}
pub fn pretty_error(err: &anyhow::Error) -> String {
@@ -235,13 +223,6 @@ mod tests {
use super::*;
#[test]
- fn test_fuzzy_match() {
- assert!(fuzzy_match("openai:gpt-4-turbo", "gpt4"));
- assert!(fuzzy_match("openai:gpt-4-turbo", "oai4"));
- assert!(!fuzzy_match("openai:gpt-4-turbo", "4gpt"));
- }
-
- #[test]
#[cfg(not(target_os = "windows"))]
fn test_safe_join_path() {
assert_eq!(