From 1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 5 Jun 2024 09:02:23 +0800 Subject: feat: support RAG (#560) * feat: support RAG * support more embeddings models and implement concurrent embedding api * show the progress of addings paths * ignore embedding context when saving message * embedding model max_chunk_size => default_chunk_size * support pdf and pandoc formats (docx, epub, ipynb) --- Cargo.lock | 468 ++++++++++++++++++++++++++++++++- Cargo.toml | 4 + README.md | 34 ++- config.example.yaml | 37 ++- models.yaml | 43 +++ src/client/azure_openai.rs | 31 ++- src/client/bedrock.rs | 10 +- src/client/claude.rs | 9 +- src/client/cloudflare.rs | 7 +- src/client/cohere.rs | 69 ++++- src/client/common.rs | 140 ++++++---- src/client/ernie.rs | 9 +- src/client/gemini.rs | 73 +++++- src/client/model.rs | 44 +++- src/client/ollama.rs | 67 ++++- src/client/openai.rs | 67 ++++- src/client/openai_compatible.rs | 73 ++++-- src/client/qianwen.rs | 86 +++++- src/client/replicate.rs | 9 +- src/client/vertexai.rs | 79 +++++- src/client/vertexai_claude.rs | 11 +- src/config/input.rs | 59 +++-- src/config/mod.rs | 185 ++++++++++--- src/config/session.rs | 2 +- src/main.rs | 32 ++- src/rag/loader.rs | 146 +++++++++++ src/rag/mod.rs | 425 ++++++++++++++++++++++++++++++ src/rag/splitter.rs | 564 ++++++++++++++++++++++++++++++++++++++++ src/render/stream.rs | 15 +- src/repl/mod.rs | 61 +++-- src/serve.rs | 11 +- src/utils/abort_signal.rs | 9 + src/utils/mod.rs | 2 +- src/utils/spinner.rs | 39 ++- 34 files changed, 2616 insertions(+), 304 deletions(-) create mode 100644 src/rag/loader.rs create mode 100644 src/rag/mod.rs create mode 100644 src/rag/splitter.rs diff --git a/Cargo.lock b/Cargo.lock index 266c2df..e4b7137 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -17,6 +17,27 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f26201604c87b1e01bd3d98f8d5d9a8fcbb815e8cedb41ffccbeb4bf593a35fe" +[[package]] +name = "adobe-cmap-parser" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "261a937a307ddc70a1605dec925987d7256adb00232161a1855ef9cc820bd8d5" +dependencies = [ + "pom", +] + +[[package]] +name = "ahash" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e89da841a80418a9b391ebaea17f5c112ffaaa96f621d2c285b5174da76b9011" +dependencies = [ + "cfg-if", + "once_cell", + "version_check", + "zerocopy", +] + [[package]] name = "aho-corasick" version = "1.1.3" @@ -48,6 +69,7 @@ dependencies = [ "fancy-regex", "futures-util", "hmac", + "hnsw_rs", "http", "http-body-util", "hyper", @@ -62,6 +84,9 @@ dependencies = [ "nu-ansi-term 0.50.0", "num_cpus", "parking_lot", + "path-absolutize", + "pdf-extract", + "pretty_assertions", "rand", "reedline", "reqwest", @@ -85,6 +110,12 @@ dependencies = [ "urlencoding", ] +[[package]] +name = "allocator-api2" +version = "0.2.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c6cb57a04249c6480766f7f7cef5467412af1490f8d1e243141daddada3264f" + [[package]] name = "android-tzdata" version = "0.1.1" @@ -100,6 +131,24 @@ dependencies = [ "libc", ] +[[package]] +name = "anndists" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4747593401c8d692fb589ac2a208a27ef968b95f9392af837728933348fc199c" +dependencies = [ + "anyhow", + "cfg-if", + "cpu-time", + "env_logger", + "lazy_static", + "log", + "num-traits", + "num_cpus", + "rand", + "rayon", +] + [[package]] name = "ansi_colours" version = "1.2.2" @@ -470,6 +519,16 @@ version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0b6a852b24ab71dffc585bcb46eaf7959d175cb865a7152e35b348d1b2960422" +[[package]] +name = "combine" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba5a308b75df32fe02788e748662718f03fde005016435c444eea572398219fd" +dependencies = [ + "bytes", + "memchr", +] + [[package]] name = "core-foundation" version = "0.9.4" @@ -486,6 +545,16 @@ version = "0.8.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06ea2b9bc92be3c2baa9334a323ebca2d6f074ff852cd1d7b11064035cd3868f" +[[package]] +name = "cpu-time" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e9e393a7668fe1fad3075085b86c781883000b4ede868f43627b34a87c8b7ded" +dependencies = [ + "libc", + "winapi", +] + [[package]] name = "cpufeatures" version = "0.2.12" @@ -504,6 +573,31 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "crossbeam-deque" +version = "0.8.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "613f8cc01fe9cf1a3eb3d7f488fd2fa8388403e97039e2f73692932e291a770d" +dependencies = [ + "crossbeam-epoch", + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-epoch" +version = "0.9.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5b82ac4a3c2ca9c3460964f020e1402edd5753411d7737aa39c3714ad1b5420e" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22ec99545bb0ed0ea7bb9b8e1e9122ea386ff8a48c0922e43f36d45ab09e0e80" + [[package]] name = "crossterm" version = "0.25.0" @@ -577,6 +671,12 @@ dependencies = [ "syn", ] +[[package]] +name = "diff" +version = "0.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "56254986775e3233ffa9c4d7d3faaf6d36a2c09d30b20687e9f88bc8bafc16c8" + [[package]] name = "digest" version = "0.10.7" @@ -636,6 +736,104 @@ version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3dca9240753cf90908d7e4aac30f630662b02aebaa1b58a3cadabdb23385b58b" +[[package]] +name = "encoding" +version = "0.2.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b0d943856b990d12d3b55b359144ff341533e516d94098b1d3fc1ac666d36ec" +dependencies = [ + "encoding-index-japanese", + "encoding-index-korean", + "encoding-index-simpchinese", + "encoding-index-singlebyte", + "encoding-index-tradchinese", +] + +[[package]] +name = "encoding-index-japanese" +version = "1.20141219.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "04e8b2ff42e9a05335dbf8b5c6f7567e5591d0d916ccef4e0b1710d32a0d0c91" +dependencies = [ + "encoding_index_tests", +] + +[[package]] +name = "encoding-index-korean" +version = "1.20141219.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4dc33fb8e6bcba213fe2f14275f0963fd16f0a02c878e3095ecfdf5bee529d81" +dependencies = [ + "encoding_index_tests", +] + +[[package]] +name = "encoding-index-simpchinese" +version = "1.20141219.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d87a7194909b9118fc707194baa434a4e3b0fb6a5a757c73c3adb07aa25031f7" +dependencies = [ + "encoding_index_tests", +] + +[[package]] +name = "encoding-index-singlebyte" +version = "1.20141219.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3351d5acffb224af9ca265f435b859c7c01537c0849754d3db3fdf2bfe2ae84a" +dependencies = [ + "encoding_index_tests", +] + +[[package]] +name = "encoding-index-tradchinese" +version = "1.20141219.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fd0e20d5688ce3cab59eb3ef3a2083a5c77bf496cb798dc6fcdb75f323890c18" +dependencies = [ + "encoding_index_tests", +] + +[[package]] +name = "encoding_index_tests" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a246d82be1c9d791c5dfde9a2bd045fc3cbba3fa2b11ad558f27d01712f00569" + +[[package]] +name = "encoding_rs" +version = "0.8.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b45de904aa0b010bce2ab45264d0631681847fa7b6f2eaa7dab7619943bc4f59" +dependencies = [ + "cfg-if", +] + +[[package]] +name = "enum-as-inner" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ffccbb6966c05b32ef8fbac435df276c4ae4d3dc55a8cd0eb9745e6c12f546a" +dependencies = [ + "heck 0.4.1", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "env_logger" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cd405aab171cb85d6735e5c8d9db038c17d3ca007a4d2c25f337935c3d90580" +dependencies = [ + "humantime", + "is-terminal", + "log", + "regex", + "termcolor", +] + [[package]] name = "equivalent" version = "1.0.1" @@ -658,6 +856,15 @@ version = "3.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a0474425d51df81997e2f90a21591180b38eccf27292d755f3e30750225c175b" +[[package]] +name = "euclid" +version = "0.20.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2bb7ef65b3777a325d1eeefefab5b6d4959da54747e33bd6258e789640f307ad" +dependencies = [ + "num-traits", +] + [[package]] name = "eventsource-stream" version = "0.2.3" @@ -844,7 +1051,7 @@ dependencies = [ "libc", "log", "rustversion", - "windows", + "windows 0.54.0", ] [[package]] @@ -908,6 +1115,10 @@ name = "hashbrown" version = "0.14.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" +dependencies = [ + "ahash", + "allocator-api2", +] [[package]] name = "heck" @@ -936,6 +1147,31 @@ dependencies = [ "digest", ] +[[package]] +name = "hnsw_rs" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e4594a6fe509cc5d6549fd56af7e6435c476e59e3bc498efa4aafcd4311b6d66" +dependencies = [ + "anndists", + "anyhow", + "bincode", + "cfg-if", + "cpu-time", + "env_logger", + "hashbrown", + "indexmap", + "lazy_static", + "log", + "mmap-rs", + "num-traits", + "num_cpus", + "parking_lot", + "rand", + "rayon", + "serde", +] + [[package]] name = "home" version = "0.5.9" @@ -991,6 +1227,12 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" +[[package]] +name = "humantime" +version = "2.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a3a5bfb195931eeb336b2a7b4d761daec841b97f947d34394601737a7bba5e4" + [[package]] name = "hyper" version = "1.3.1" @@ -1218,6 +1460,12 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dd1bc4d24ad230d21fb898d1116b1801d7adfc449d42026475862ab48b11e70e" +[[package]] +name = "linked-hash-map" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0717cef1bc8b636c6e1c1bbdefc09e6322da8a9321966e8928ef80d20f7f770f" + [[package]] name = "linux-raw-sys" version = "0.4.14" @@ -1256,6 +1504,32 @@ dependencies = [ "tracing-subscriber", ] +[[package]] +name = "lopdf" +version = "0.32.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e775e4ee264e8a87d50a9efef7b67b4aa988cf94e75630859875fc347e6c872b" +dependencies = [ + "encoding_rs", + "flate2", + "itoa", + "linked-hash-map", + "log", + "md5", + "nom", + "time", + "weezl", +] + +[[package]] +name = "mach2" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19b955cdeb2a02b9117f121ce63aa52d08ade45de53e48fe6a38b39c10f6f709" +dependencies = [ + "libc", +] + [[package]] name = "matchers" version = "0.1.0" @@ -1265,12 +1539,27 @@ dependencies = [ "regex-automata 0.1.10", ] +[[package]] +name = "md5" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "490cc448043f947bae3cbee9c203358d62dbee0db12107a74be5c30ccfd09771" + [[package]] name = "memchr" version = "2.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6c8640c5d730cb13ebd907d8d04b52f55ac9a2eec55b440c8892f40d56c76c1d" +[[package]] +name = "memoffset" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5de893c32cde5f383baa4c04c5d6dbdd735cfd4a794b0debdb2bb1b421da5ff4" +dependencies = [ + "autocfg", +] + [[package]] name = "mime" version = "0.3.17" @@ -1314,6 +1603,23 @@ dependencies = [ "windows-sys 0.48.0", ] +[[package]] +name = "mmap-rs" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "86968d85441db75203c34deefd0c88032f275aaa85cee19a1dcfff6ae9df56da" +dependencies = [ + "bitflags 1.3.2", + "combine", + "libc", + "mach2", + "nix 0.26.4", + "sysctl", + "thiserror", + "widestring", + "windows 0.48.0", +] + [[package]] name = "newline-converter" version = "0.3.0" @@ -1323,6 +1629,19 @@ dependencies = [ "unicode-segmentation", ] +[[package]] +name = "nix" +version = "0.26.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "598beaf3cc6fdd9a5dfb1630c2800c7acd31df7aaf0f565796fba2b53ca1af1b" +dependencies = [ + "bitflags 1.3.2", + "cfg-if", + "libc", + "memoffset", + "pin-utils", +] + [[package]] name = "nix" version = "0.28.0" @@ -1600,6 +1919,39 @@ dependencies = [ "windows-targets 0.52.5", ] +[[package]] +name = "path-absolutize" +version = "3.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e4af381fe79fa195b4909485d99f73a80792331df0625188e707854f0b3383f5" +dependencies = [ + "path-dedot", +] + +[[package]] +name = "path-dedot" +version = "3.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07ba0ad7e047712414213ff67533e6dd477af0a4e1d14fb52343e53d30ea9397" +dependencies = [ + "once_cell", +] + +[[package]] +name = "pdf-extract" +version = "0.7.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3423481005e61b95855d53d7b6c0bcc514b4fbab45d165776d6e42c0a1642b22" +dependencies = [ + "adobe-cmap-parser", + "encoding", + "euclid", + "lopdf", + "postscript", + "type1-encoding-parser", + "unicode-normalization", +] + [[package]] name = "percent-encoding" version = "2.3.1" @@ -1668,6 +2020,18 @@ dependencies = [ "time", ] +[[package]] +name = "pom" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "60f6ce597ecdcc9a098e7fddacb1065093a3d66446fa16c675e7e71d1b5c28e6" + +[[package]] +name = "postscript" +version = "0.14.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78451badbdaebaf17f053fd9152b3ffb33b516104eacb45e7864aaa9c712f306" + [[package]] name = "powerfmt" version = "0.2.0" @@ -1680,6 +2044,16 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5b40af805b3121feab8a3c29f04d8ad262fa8e0561883e7653e024ae4479e6de" +[[package]] +name = "pretty_assertions" +version = "1.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af7cee1a6c8a5b9208b3cb1061f10c0cb689087b3d8ce85fb9d2dd7a29b6ba66" +dependencies = [ + "diff", + "yansi", +] + [[package]] name = "proc-macro2" version = "1.0.84" @@ -1737,6 +2111,26 @@ dependencies = [ "getrandom", ] +[[package]] +name = "rayon" +version = "1.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b418a60154510ca1a002a752ca9714984e21e4241e804d32555251faf8b78ffa" +dependencies = [ + "either", + "rayon-core", +] + +[[package]] +name = "rayon-core" +version = "1.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1465873a3dfdaa8ae7cb14b4383657caab0b3e8a0aa9ae8e04b044854c8dfce2" +dependencies = [ + "crossbeam-deque", + "crossbeam-utils", +] + [[package]] name = "redox_syscall" version = "0.5.1" @@ -2290,6 +2684,20 @@ dependencies = [ "walkdir", ] +[[package]] +name = "sysctl" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec7dddc5f0fee506baf8b9fdb989e242f17e4b11c61dfbb0635b705217199eea" +dependencies = [ + "bitflags 2.5.0", + "byteorder", + "enum-as-inner", + "libc", + "thiserror", + "walkdir", +] + [[package]] name = "tempfile" version = "3.10.1" @@ -2607,6 +3015,15 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "type1-encoding-parser" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3d6cc09e1a99c7e01f2afe4953789311a1c50baebbdac5b477ecf78e2e92a5b" +dependencies = [ + "pom", +] + [[package]] name = "typenum" version = "1.17.0" @@ -2930,6 +3347,18 @@ dependencies = [ "rustls-pki-types", ] +[[package]] +name = "weezl" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53a85b86a771b1c87058196170769dd264f66c0782acf1ae6cc51bfd64b39082" + +[[package]] +name = "widestring" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7219d36b6eac893fa81e84ebe06485e7dcbb616177469b142df14f1f4deb1311" + [[package]] name = "winapi" version = "0.3.9" @@ -2961,6 +3390,15 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" +[[package]] +name = "windows" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e686886bc078bc1b0b600cac0147aadb815089b6e4da64016cbd754b6342700f" +dependencies = [ + "windows-targets 0.48.5", +] + [[package]] name = "windows" version = "0.54.0" @@ -3157,7 +3595,7 @@ dependencies = [ "derive-new", "libc", "log", - "nix", + "nix 0.28.0", "os_pipe", "tempfile", "thiserror", @@ -3185,6 +3623,32 @@ version = "0.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ec107c4503ea0b4a98ef47356329af139c0a4f7750e621cf2973cd3385ebcb3d" +[[package]] +name = "yansi" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09041cd90cf85f7f8b2df60c646f853b7f535ce68f85244eb6731cf89fa498ec" + +[[package]] +name = "zerocopy" +version = "0.7.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae87e3fcd617500e5d106f0380cf7b77f3c6092aae37191433159dda23cfb087" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.7.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "15e934569e47891f7d9411f1a451d947a60e000ab3bd24fbb970f000387d1b3b" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "zeroize" version = "1.8.1" diff --git a/Cargo.toml b/Cargo.toml index 44de068..7363323 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -60,6 +60,9 @@ num_cpus = "1.16.0" threadpool = "1.8.1" json-patch = { version = "2.0.0", default-features = false } bitflags = "2.5.0" +path-absolutize = "3.1.1" +hnsw_rs = "0.3.0" +pdf-extract = "0.7.7" [dependencies.reqwest] version = "0.12.0" @@ -81,6 +84,7 @@ arboard = { version = "3.3.0", default-features = false, features = ["wayland-da arboard = { version = "3.3.0", default-features = false } [dev-dependencies] +pretty_assertions = "1.4.0" rand = "0.8.5" [profile.release] diff --git a/README.md b/README.md index 09b8c26..7da7b85 100644 --- a/README.md +++ b/README.md @@ -16,6 +16,7 @@ AIChat is an all-in-one AI CLI tool that accesses 100+ LLMs across 20+ AI platfo - **Chat-REPL**: Powerful and feature-rich interactive chat interface. - **Custom Roles**: Tailor LLM behavior with customizable roles. - **Unlimited Sessions**: Automatic message compression for endless conversations. +- **RAG Retrieval**: Get answer enhanced by your knowledge base.. - **Function Calling**: Connect LLMs to external tools seamlessly. - **Execute Commands**: Use natural language to run shell commands. - **Shell Auto-Completion**: AI-based auto-completion for shell commands. @@ -25,22 +26,22 @@ AIChat is an all-in-one AI CLI tool that accesses 100+ LLMs across 20+ AI platfo ## Supported AI Platforms -- OpenAI GPT-3.5/GPT-4 (paid, vision, function-calling) -- Gemini: Gemini-1.0/Gemini-1.5 (free, paid, vision, function-calling) +- OpenAI GPT-3.5/GPT-4 (paid, vision, embedding, function-calling) +- Gemini: Gemini-1.0/Gemini-1.5 (free, paid, vision, embedding, function-calling) - Claude: Claude-3 (vision, paid, function-calling) -- Mistral (paid, function-calling) -- Cohere: Command-R/Command-R+ (paid, function-calling) +- Mistral (paid, embedding, function-calling) +- Cohere: Command-R/Command-R+ (paid, embedding, function-calling) - Perplexity: Llama-3/Mixtral (paid) - Groq: Llama-3/Mixtral/Gemma (free) -- Ollama (free, local) -- Azure OpenAI (paid) -- VertexAI: Gemini-1.0/Gemini-1.5 (paid, vision, function-calling) +- Ollama (free, local, embedding) +- Azure OpenAI (paid, vision, embedding, function-calling) +- VertexAI: Gemini-1.0/Gemini-1.5 (paid, vision, embedding, function-calling) - VertexAI-Claude: Claude-3 (paid, vision) - Bedrock: Llama-3/Claude-3/Mistral (paid, vision) - Cloudflare (free, paid, vision) - Replicate (paid) - Ernie (paid) -- Qianwen (paid, vision) +- Qianwen (paid, vision, embedding) - Moonshot (paid) - ZhipuAI: GLM-3.5/GLM-4 (paid, vision) - Deepseek (paid) @@ -68,7 +69,7 @@ Upon first launch, AIChat will guide you through the configuration process. > No config file, create a new one? Yes > AI Platform: openai > API Key: -✨ Saved config file to /aichat/config.yaml +✨ Saved config file to '/aichat/config.yaml' ``` Feel free to adjust the configuration according to your needs. @@ -340,12 +341,27 @@ Usage: .file ... [-- text...] > The capability to process images through `.file` command depends on the current model’s vision support. +### `.rag` - chat with your documents and knowledge bases. + +``` +> .rag test1 +> Select embedding model: openai:text-embedding-3-small +> Set chunk size: 2000 +> Add document paths: tmp/files/paul_graham_essay.txt +✨ Saved rag to '/aichat/rags/test5.bin' + +#test1> What did the author do growing up? +The author mainly focused on writing and programming growing up ... +``` + ### `.set` - adjust settings (non-persistent) ``` .set max_output_tokens 4096 .set temperature 1.2 .set top_p 0.8 +.set rag_top_k 4 +.set function_calling true .set compress_threshold 1000 .set dry_run true ``` diff --git a/config.example.yaml b/config.example.yaml index b6ec6fe..86fe221 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -15,9 +15,24 @@ prelude: null # Set a default role or session to start with ( # if unset fallback to $EDITOR and $VISUAL buffer_editor: null -# Controls the function calling feature. For setup instructions, visit https://github.com/sigoden/llm-functions. +# Controls the function calling feature. For setup instructions, visit https://github.com/sigoden/llm-functions function_calling: false +# Specifies the embedding model to use +embedding_model: null + +# Determines how many relevant documents are retrieved +rag_top_k: 4 + +# Defines the query structure using variables like __CONTEXT__ and __INPUT__ to tailor searches to specific needs +rag_template: | + Answer the following question based only on the provided context: + + __CONTEXT__ + + + Question: __INPUT__ + # Compress session when token count reaches or exceeds this threshold (must be at least 1000) compress_threshold: 4000 # Text prompt used for creating a concise summary of session message @@ -26,7 +41,7 @@ summarize_prompt: 'Summarize the discussion briefly in 200 words or less to use summary_prompt: 'This is a summary of the chat history as a recap: ' # Custom REPL prompt, see https://github.com/sigoden/aichat/wiki/Custom-REPL-Prompt -left_prompt: '{color.green}{?session {session}{?role /}}{role}{color.cyan}{?session )}{!session >}{color.reset} ' +left_prompt: '{color.green}{?session {session}{?role /}}{role}{?rag #{rag}{color.cyan}{?session )}{!session >}{color.reset} ' right_prompt: '{color.purple}{?session {?consume_tokens {consume_tokens}({consume_percent}%)}{!consume_tokens {consume_tokens}}}{color.reset}' clients: @@ -34,13 +49,19 @@ clients: # - type: xxxx # name: xxxx # Only use it to distinguish clients with the same client type. Optional # models: - # - name: xxxx # The model name + # - name: xxxx + # mode: chat # Chat model # max_input_tokens: 100000 # supports_vision: true # supports_function_calling: true + # - name: xxxx + # mode: embedding # Embedding model + # max_input_tokens: 2048 + # default_chunk_size: 2000 + # max_concurrent_chunks: 100 # patches: # : # The regex to match model names, e.g. '.*' 'gpt-4o' 'gpt-4o|gpt-4-.*' - # request_body: # The JSON to be merged with the request body. + # chat_completions_body: # The JSON to be merged with the chat completions request body. # extra: # proxy: socks5://127.0.0.1:1080 # Set https/socks5 proxy. ENV: HTTPS_PROXY/https_proxy/ALL_PROXY/all_proxy # connect_timeout: 10 # Set timeout in seconds for connect to api @@ -66,7 +87,7 @@ clients: api_key: xxx # ENV: {client}_API_KEY patches: '.*': - request_body: # Override safetySettings for all models + chat_completions_body: # Override safetySettings for all models safetySettings: - category: HARM_CATEGORY_HARASSMENT threshold: BLOCK_NONE @@ -107,10 +128,12 @@ clients: - type: ollama api_base: http://localhost:11434 # ENV: {client}_API_BASE api_auth: Basic xxx # ENV: {client}_API_AUTH - chat_endpoint: /api/chat # Optional models: # Required - name: llama3 max_input_tokens: 8192 + - name: all-minilm:l6-v2 + mode: embedding + max_chunk_size: 1000 # See https://learn.microsoft.com/en-us/azure/ai-services/openai/chatgpt-quickstart - type: azure-openai @@ -130,7 +153,7 @@ clients: adc_file: patches: 'gemini-.*': - request_body: # Override safetySettings for all gemini models + chat_completions_body: # Override safetySettings for all gemini models safetySettings: - category: HARM_CATEGORY_HARASSMENT threshold: BLOCK_ONLY_HIGH diff --git a/models.yaml b/models.yaml index c306f1b..f4c980b 100644 --- a/models.yaml +++ b/models.yaml @@ -65,6 +65,16 @@ max_output_tokens: 4096 input_price: 60 output_price: 120 + - name: text-embedding-3-large + mode: embedding + max_input_tokens: 8191 + default_chunk_size: 8000 + max_concurrent_chunks: 100 + - name: text-embedding-3-small + mode: embedding + max_input_tokens: 8191 + default_chunk_size: 8000 + max_concurrent_chunks: 100 - platform: gemini # docs: @@ -100,6 +110,10 @@ output_price: 10.5 supports_vision: true supports_function_calling: true + - name: text-embedding-004 + mode: embedding + max_input_tokens: 2048 + default_chunk_size: 2000 - platform: claude # docs: @@ -162,6 +176,10 @@ input_price: 8 output_price: 24 supports_function_calling: true + - name: mistral-embed + mode: embedding + max_input_tokens: 8092 + default_chunk_size: 8000 - platform: cohere # docs: @@ -183,6 +201,16 @@ input_price: 3 output_price: 15 supports_function_calling: true + - name: embed-english-v3.0 + mode: embedding + max_input_tokens: 512 + default_chunk_size: 1000 + max_concurrent_chunks: 96 + - name: embed-multilingual-v3.0 + mode: embedding + max_input_tokens: 512 + default_chunk_size: 1000 + max_concurrent_chunks: 96 - platform: perplexity # docs: @@ -281,6 +309,16 @@ input_price: 1.25 output_price: 3.75 supports_vision: true + - name: text-embedding-004 + mode: embedding + max_input_tokens: 3072 + default_chunk_size: 3000 + max_concurrent_chunks: 5 + - name: text-multilingual-embedding-002 + mode: embedding + max_input_tokens: 3072 + default_chunk_size: 3000 + max_concurrent_chunks: 5 - platform: vertexai-claude # docs: @@ -509,6 +547,11 @@ input_price: 2.8 output_price: 2.8 supports_vision: true + - name: text-embedding-v2 + mode: embedding + max_input_tokens: 2048 + default_chunk_size: 2000 + max_concurrent_chunks: 5 - platform: moonshot # docs: diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index 52c8a34..19d234a 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,7 +1,5 @@ -use super::{ - openai::*, AzureOpenAIClient, ChatCompletionsData, Client, ExtraConfig, Model, ModelData, - ModelPatches, PromptAction, PromptKind, -}; +use super::*; +use super::openai::*; use anyhow::Result; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -12,6 +10,7 @@ pub struct AzureOpenAIConfig { pub name: Option, pub api_base: Option, pub api_key: Option, + #[serde(default)] pub models: Vec, pub patches: Option, pub extra: Option, @@ -42,7 +41,7 @@ impl AzureOpenAIClient { let api_key = self.get_api_key()?; let mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!( "{}/openai/deployments/{}/chat/completions?api-version=2024-02-01", @@ -56,10 +55,28 @@ impl AzureOpenAIClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_base = self.get_api_base()?; + let api_key = self.get_api_key()?; + + let body = openai_build_embeddings_body(data, &self.model); + + let url = format!("{api_base}/embeddings"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } } impl_client_trait!( AzureOpenAIClient, - crate::client::openai::openai_chat_completions, - crate::client::openai::openai_chat_completions_streaming + openai_chat_completions, + openai_chat_completions_streaming, + openai_embeddings ); diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 3dfa977..981d1cb 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -1,8 +1,6 @@ -use super::{ - catch_error, claude::*, prompt_format::*, BedrockClient, ChatCompletionsData, - ChatCompletionsOutput, Client, ExtraConfig, Model, ModelData, ModelPatches, PromptAction, - PromptKind, SseHandler, -}; +use super::*; +use super::claude::*; +use super::prompt_format::*; use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256}; @@ -102,7 +100,7 @@ impl BedrockClient { let headers = IndexMap::new(); let mut body = build_chat_completions_body(data, &self.model, model_category)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let builder = aws_fetch( client, diff --git a/src/client/claude.rs b/src/client/claude.rs index 6533a9a..a16722b 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,9 +1,4 @@ -use super::{ - catch_error, extract_system_message, message::*, sse_stream, ChatCompletionsData, - ChatCompletionsOutput, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent, - MessageContentPart, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, - SseMmessage, ToolCall, -}; +use super::*; use anyhow::{bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -36,7 +31,7 @@ impl ClaudeClient { let api_key = self.get_api_key().ok(); let mut body = claude_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = API_BASE; diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs index 966cee4..965f20a 100644 --- a/src/client/cloudflare.rs +++ b/src/client/cloudflare.rs @@ -1,7 +1,4 @@ -use super::{ - catch_error, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, CloudflareClient, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, SseMmessage, -}; +use super::*; use anyhow::{anyhow, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -39,7 +36,7 @@ impl CloudflareClient { let api_key = self.get_api_key()?; let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!( "{API_BASE}/accounts/{account_id}/ai/run/{}", diff --git a/src/client/cohere.rs b/src/client/cohere.rs index e0a5eec..69c343b 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,15 +1,12 @@ -use super::{ - catch_error, extract_system_message, json_stream, message::*, ChatCompletionsData, - ChatCompletionsOutput, Client, CohereClient, ExtraConfig, Model, ModelData, ModelPatches, - PromptAction, PromptKind, SseHandler, ToolCall, -}; +use super::*; -use anyhow::{bail, Result}; +use anyhow::{bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; -const API_URL: &str = "https://api.cohere.ai/v1/chat"; +const CHAT_COMPLETIONS_API_URL: &str = "https://api.cohere.ai/v1/chat"; +const EMBEDDINGS_API_URL: &str = "https://api.cohere.ai/v1/embed"; #[derive(Debug, Clone, Deserialize, Default)] pub struct CohereConfig { @@ -35,11 +32,38 @@ impl CohereClient { let api_key = self.get_api_key()?; let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - let url = API_URL; + let url = CHAT_COMPLETIONS_API_URL; - debug!("Cohere Request: {url} {body}"); + debug!("Cohere Chat Completions Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + + let input_type = match data.query { + true => "search_query", + false => "search_document", + }; + + let body = json!({ + "model": self.model.name(), + "texts": data.texts, + "input_type": input_type, + }); + + let url = EMBEDDINGS_API_URL; + + debug!("Cohere Embeddings Request: {url} {body}"); let builder = client.post(url).bearer_auth(api_key).json(&body); @@ -47,7 +71,12 @@ impl CohereClient { } } -impl_client_trait!(CohereClient, chat_completions, chat_completions_streaming); +impl_client_trait!( + CohereClient, + chat_completions, + chat_completions_streaming, + embeddings +); async fn chat_completions(builder: RequestBuilder) -> Result { let res = builder.send().await?; @@ -100,6 +129,24 @@ async fn chat_completions_streaming( Ok(()) } +async fn embeddings( + builder: RequestBuilder, +) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?; + Ok(res_body.embeddings) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + embeddings: Vec>, +} + fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result { let ChatCompletionsData { mut messages, diff --git a/src/client/common.rs b/src/client/common.rs index 96ec90b..1055b84 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,10 +1,13 @@ -use super::{openai::OpenAIConfig, BuiltinModels, ClientConfig, Message, Model, SseHandler}; +use super::*; use crate::{ config::{GlobalConfig, Input}, function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolCallResult}, render::{render_error, render_stream}, - utils::{prompt_input_integer, prompt_input_string, tokenize, AbortSignal, PromptKind}, + utils::{ + prompt_input_integer, prompt_input_string, tokenize, watch_abort_signal, AbortSignal, + PromptKind, + }, }; use anyhow::{bail, Context, Result}; @@ -16,13 +19,12 @@ use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; use std::{env, future::Future, time::Duration}; -use tokio::{sync::mpsc::unbounded_channel, time::sleep}; +use tokio::sync::mpsc::unbounded_channel; const MODELS_YAML: &str = include_str!("../../models.yaml"); lazy_static! { - pub static ref ALL_CLIENT_MODELS: Vec = - serde_yaml::from_str(MODELS_YAML).unwrap(); + pub static ref ALL_MODELS: Vec = serde_yaml::from_str(MODELS_YAML).unwrap(); static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(? Vec { let client_name = Self::name(local_config); if local_config.models.is_empty() { - if let Some(client_models) = $crate::client::ALL_CLIENT_MODELS.iter().find(|v| { + if let Some(models) = $crate::client::ALL_MODELS.iter().find(|v| { v.platform == $name || ($name == "openai-compatible" && local_config.name.as_deref() == Some(&v.platform)) }) { - return Model::from_config(client_name, &client_models.models); + return Model::from_config(client_name, &models.models); } vec![] } else { @@ -137,10 +139,10 @@ macro_rules! register_client { anyhow::bail!("Unknown client '{}'", client) } - static mut ALL_CLIENTS: Option> = None; + static mut ALL_CLIENT_MODELS: Option> = None; - pub fn list_models(config: &$crate::config::Config) -> Vec<&$crate::client::Model> { - if unsafe { ALL_CLIENTS.is_none() } { + pub fn list_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + if unsafe { ALL_CLIENT_MODELS.is_none() } { let models: Vec<_> = config .clients .iter() @@ -149,9 +151,17 @@ macro_rules! register_client { ClientConfig::Unknown => vec![], }) .collect(); - unsafe { ALL_CLIENTS = Some(models) }; + unsafe { ALL_CLIENT_MODELS = Some(models) }; } - unsafe { ALL_CLIENTS.as_ref().unwrap().iter().collect() } + unsafe { ALL_CLIENT_MODELS.as_ref().unwrap().iter().collect() } + } + + pub fn list_chat_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + list_models(config).into_iter().filter(|v| v.mode() == "chat").collect() + } + + pub fn list_embedding_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + list_models(config).into_iter().filter(|v| v.mode() == "embedding").collect() } }; } @@ -171,10 +181,6 @@ macro_rules! client_common_fns { self.config.patches.as_ref() } - fn list_models(&self) -> Vec { - Self::list_models(&self.config) - } - fn name(&self) -> &str { Self::name(&self.config) } @@ -186,10 +192,6 @@ macro_rules! client_common_fns { fn model_mut(&mut self) -> &mut Model { &mut self.model } - - fn set_model(&mut self, model: Model) { - self.model = model; - } }; } @@ -220,6 +222,40 @@ macro_rules! impl_client_trait { } } }; + ($client:ident, $chat_completions:path, $chat_completions_streaming:path, $embeddings:path) => { + #[async_trait::async_trait] + impl $crate::client::Client for $crate::client::$client { + client_common_fns!(); + + async fn chat_completions_inner( + &self, + client: &reqwest::Client, + data: $crate::client::ChatCompletionsData, + ) -> anyhow::Result<$crate::client::ChatCompletionsOutput> { + let builder = self.chat_completions_builder(client, data)?; + $chat_completions(builder).await + } + + async fn chat_completions_streaming_inner( + &self, + client: &reqwest::Client, + handler: &mut $crate::client::SseHandler, + data: $crate::client::ChatCompletionsData, + ) -> Result<()> { + let builder = self.chat_completions_builder(client, data)?; + $chat_completions_streaming(builder, handler).await + } + + async fn embeddings_inner( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result>> { + let builder = self.embeddings_builder(client, data)?; + $embeddings(builder).await + } + } + }; } #[macro_export] @@ -256,19 +292,12 @@ pub trait Client: Sync + Send { fn patches_config(&self) -> Option<&ModelPatches>; - #[allow(unused)] fn name(&self) -> &str; - #[allow(unused)] - fn list_models(&self) -> Vec; - fn model(&self) -> &Model; fn model_mut(&mut self) -> &mut Model; - #[allow(unused)] - fn set_model(&mut self, model: Model); - fn build_client(&self) -> Result { let mut builder = ReqwestClient::builder(); let extra = self.extra_config(); @@ -288,11 +317,10 @@ pub trait Client: Sync + Send { return Ok(ChatCompletionsOutput::new(&content)); } let client = self.build_client()?; - let data = input.prepare_completion_data(self.model(), false)?; self.chat_completions_inner(&client, data) .await - .with_context(|| "Failed to get answer") + .with_context(|| "Failed to get chat completions") } async fn chat_completions_streaming( @@ -300,15 +328,7 @@ pub trait Client: Sync + Send { input: &Input, handler: &mut SseHandler, ) -> Result<()> { - async fn watch_abort(abort: AbortSignal) { - loop { - if abort.aborted() { - break; - } - sleep(Duration::from_millis(100)).await; - } - } - let abort = handler.get_abort(); + let abort_signal = handler.get_abort(); let input = input.clone(); tokio::select! { ret = async { @@ -326,20 +346,28 @@ pub trait Client: Sync + Send { self.chat_completions_streaming_inner(&client, handler, data).await } => { handler.done()?; - ret.with_context(|| "Failed to get answer") + ret.with_context(|| "Failed to get chat completions") } - _ = watch_abort(abort.clone()) => { + _ = watch_abort_signal(abort_signal) => { handler.done()?; Ok(()) }, } } - fn patch_request_body(&self, body: &mut Value) { + async fn embeddings(&self, data: EmbeddingsData) -> Result>> { + let client = self.build_client()?; + self.model().guard_max_concurrent_chunks(&data)?; + self.embeddings_inner(&client, data) + .await + .with_context(|| "Failed to get embeddings") + } + + fn patch_chat_completions_body(&self, body: &mut Value) { let model_name = self.model().name(); if let Some(patch_data) = select_model_patch(self.patches_config(), model_name) { - if body.is_object() && patch_data.request_body.is_object() { - json_patch::merge(body, &patch_data.request_body) + if body.is_object() && patch_data.chat_completions_body.is_object() { + json_patch::merge(body, &patch_data.chat_completions_body) } } } @@ -356,6 +384,14 @@ pub trait Client: Sync + Send { handler: &mut SseHandler, data: ChatCompletionsData, ) -> Result<()>; + + async fn embeddings_inner( + &self, + _client: &ReqwestClient, + _data: EmbeddingsData, + ) -> Result>> { + bail!("No embeddings api") + } } impl Default for ClientConfig { @@ -375,7 +411,7 @@ pub type ModelPatches = IndexMap; #[derive(Debug, Clone, Deserialize)] pub struct ModelPatch { #[serde(default)] - pub request_body: Value, + pub chat_completions_body: Value, } pub fn select_model_patch<'a>( @@ -421,6 +457,20 @@ impl ChatCompletionsOutput { } } +#[derive(Debug)] +pub struct EmbeddingsData { + pub texts: Vec, + pub query: bool, +} + +impl EmbeddingsData { + pub fn new(texts: Vec, query: bool) -> Self { + Self { texts, query } + } +} + +pub type EmbeddingsOutput = Vec>; + pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind); pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> { @@ -445,7 +495,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result Result { let mut body = build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let access_token = get_access_token(self.name())?; diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 5cc45c5..03eef7a 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -1,11 +1,10 @@ -use super::{ - vertexai::*, ChatCompletionsData, Client, ExtraConfig, GeminiClient, Model, ModelData, - ModelPatches, PromptAction, PromptKind, -}; +use super::vertexai::*; +use super::*; -use anyhow::Result; +use anyhow::{Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; +use serde_json::{json, Value}; const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/"; @@ -38,13 +37,41 @@ impl GeminiClient { }; let mut body = gemini_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - let model = &self.model.name(); + let url = format!("{API_BASE}{}:{}?key={}", &self.model.name(), func, api_key); - let url = format!("{API_BASE}{}:{}?key={}", model, func, api_key); + debug!("Gemini Chat Completions Request: {url} {body}"); - debug!("Gemini Request: {url} {body}"); + let builder = client.post(url).json(&body); + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + + let body = json!({ + "content": { + "parts": [ + { + "text": data.texts[0], + } + ] + } + }); + + let url = format!( + "{API_BASE}{}:embedContent?key={}", + &self.model.name(), + api_key + ); + + debug!("Gemini Embeddings Request: {url} {body}"); let builder = client.post(url).json(&body); @@ -54,6 +81,30 @@ impl GeminiClient { impl_client_trait!( GeminiClient, - crate::client::vertexai::gemini_chat_completions, - crate::client::vertexai::gemini_chat_completions_streaming + gemini_chat_completions, + gemini_chat_completions_streaming, + gemini_embeddings ); + +async fn gemini_embeddings(builder: RequestBuilder) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + let res_body: EmbeddingsResBody = + serde_json::from_value(data).context("Invalid request data")?; + let output = vec![res_body.embedding.values]; + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + embedding: EmbeddingsResBodyEmbedding, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyEmbedding { + values: Vec, +} diff --git a/src/client/model.rs b/src/client/model.rs index 65e4143..e16cb4e 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,4 +1,7 @@ -use super::message::{Message, MessageContent}; +use super::{ + message::{Message, MessageContent}, + EmbeddingsData, +}; use crate::utils::{estimate_token_length, format_option_value}; @@ -81,6 +84,10 @@ impl Model { &self.data.name } + pub fn mode(&self) -> &str { + &self.data.mode + } + pub fn data(&self) -> &ModelData { &self.data } @@ -137,6 +144,14 @@ impl Model { self.data.supports_function_calling } + pub fn default_chunk_size(&self) -> usize { + self.data.default_chunk_size.unwrap_or(1000) + } + + pub fn max_concurrent_chunks(&self) -> usize { + self.data.max_concurrent_chunks.unwrap_or(1) + } + pub fn max_tokens_param(&self) -> Option { if self.data.pass_max_tokens { self.data.max_output_tokens @@ -182,30 +197,45 @@ impl Model { } } - pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> { + pub fn guard_max_input_tokens(&self, messages: &[Message]) -> Result<()> { let total_tokens = self.total_tokens(messages) + BASIS_TOKENS; if let Some(max_input_tokens) = self.data.max_input_tokens { if total_tokens >= max_input_tokens { - bail!("Exceed max input tokens limit") + bail!("Exceed max_input_tokens limit") } } Ok(()) } + + pub fn guard_max_concurrent_chunks(&self, data: &EmbeddingsData) -> Result<()> { + if data.texts.len() > self.max_concurrent_chunks() { + bail!("Exceed max_concurrent_chunks limit"); + } + Ok(()) + } } #[derive(Debug, Clone, Default, Deserialize)] pub struct ModelData { pub name: String, + #[serde(default = "default_model_mode")] + pub mode: String, pub max_input_tokens: Option, + pub input_price: Option, + pub output_price: Option, + + // chat-only properties pub max_output_tokens: Option, #[serde(default)] pub pass_max_tokens: bool, - pub input_price: Option, - pub output_price: Option, #[serde(default)] pub supports_vision: bool, #[serde(default)] pub supports_function_calling: bool, + + // embedding-only properties + pub default_chunk_size: Option, + pub max_concurrent_chunks: Option, } impl ModelData { @@ -222,3 +252,7 @@ pub struct BuiltinModels { pub platform: String, pub models: Vec, } + +fn default_model_mode() -> String { + "chat".into() +} diff --git a/src/client/ollama.rs b/src/client/ollama.rs index beba8a1..f9bf8d5 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -1,10 +1,6 @@ -use super::{ - catch_error, json_stream, message::*, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, OllamaClient, PromptAction, PromptKind, - SseHandler, -}; +use super::*; -use anyhow::{anyhow, bail, Result}; +use anyhow::{anyhow, bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -14,7 +10,7 @@ pub struct OllamaConfig { pub name: Option, pub api_base: Option, pub api_auth: Option, - pub chat_endpoint: Option, + #[serde(default)] pub models: Vec, pub patches: Option, pub extra: Option, @@ -45,13 +41,36 @@ impl OllamaClient { let api_auth = self.get_api_auth().ok(); let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - let chat_endpoint = self.config.chat_endpoint.as_deref().unwrap_or("/api/chat"); + let url = format!("{api_base}/api/chat"); - let url = format!("{api_base}{chat_endpoint}"); + debug!("Ollama Chat Completions Request: {url} {body}"); - debug!("Ollama Request: {url} {body}"); + let mut builder = client.post(url).json(&body); + if let Some(api_auth) = api_auth { + builder = builder.header("Authorization", api_auth) + } + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_base = self.get_api_base()?; + let api_auth = self.get_api_auth().ok(); + + let body = json!({ + "model": self.model.name(), + "prompt": data.texts[0], + }); + + let url = format!("{api_base}/api/embeddings"); + + debug!("Ollama Embeddings Request: {url} {body}"); let mut builder = client.post(url).json(&body); if let Some(api_auth) = api_auth { @@ -62,7 +81,12 @@ impl OllamaClient { } } -impl_client_trait!(OllamaClient, chat_completions, chat_completions_streaming); +impl_client_trait!( + OllamaClient, + chat_completions, + chat_completions_streaming, + embeddings +); async fn chat_completions(builder: RequestBuilder) -> Result { let res = builder.send().await?; @@ -109,6 +133,25 @@ async fn chat_completions_streaming( Ok(()) } +async fn embeddings( + builder: RequestBuilder, +) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?; + let output = vec![res_body.embedding]; + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + embedding: Vec, +} + fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result { let ChatCompletionsData { messages, diff --git a/src/client/openai.rs b/src/client/openai.rs index 3cdea24..0da8166 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,10 +1,6 @@ -use super::{ - catch_error, message::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, OpenAIClient, PromptAction, PromptKind, - SseHandler, SseMmessage, ToolCall, -}; +use super::*; -use anyhow::{bail, Result}; +use anyhow::{bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -39,11 +35,11 @@ impl OpenAIClient { let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string()); let mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!("{api_base}/chat/completions"); - debug!("OpenAI Request: {url} {body}"); + debug!("OpenAI Chat Completions Request: {url} {body}"); let mut builder = client.post(url).bearer_auth(api_key).json(&body); @@ -53,6 +49,25 @@ impl OpenAIClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string()); + + let body = openai_build_embeddings_body(data, &self.model); + + let url = format!("{api_base}/embeddings"); + + debug!("OpenAI Embeddings Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } } pub async fn openai_chat_completions(builder: RequestBuilder) -> Result { @@ -125,6 +140,30 @@ pub async fn openai_chat_completions_streaming( sse_stream(builder, handle).await } +pub async fn openai_embeddings( + builder: RequestBuilder, +) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?; + let output = res_body.data.into_iter().map(|v| v.embedding).collect(); + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + data: Vec, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyEmbedding { + embedding: Vec, +} + pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value { let ChatCompletionsData { messages, @@ -201,6 +240,15 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod body } + +pub fn openai_build_embeddings_body(data: EmbeddingsData, model: &Model) -> Value { + json!({ + "input": data.texts, + "model": model.name(), + "encoding_format": "float", + }) +} + pub fn openai_extract_chat_completions(data: &Value) -> Result { let text = data["choices"][0]["message"]["content"] .as_str() @@ -244,5 +292,6 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result Result { - let api_base = match self.get_api_base() { - Ok(v) => v, - Err(err) => { - match OPENAI_COMPATIBLE_PLATFORMS - .into_iter() - .find_map(|(name, api_base)| { - if name == self.model.client_name() { - Some(api_base.to_string()) - } else { - None - } - }) { - Some(v) => v, - None => return Err(err), - } - } - }; let api_key = self.get_api_key().ok(); + let api_base = self.get_api_base_ext()?; let mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let chat_endpoint = self .config @@ -71,7 +53,7 @@ impl OpenAICompatibleClient { let url = format!("{api_base}{chat_endpoint}"); - debug!("OpenAICompatible Request: {url} {body}"); + debug!("OpenAICompatible Chat Completions Request: {url} {body}"); let mut builder = client.post(url).json(&body); if let Some(api_key) = api_key { @@ -80,10 +62,51 @@ impl OpenAICompatibleClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + let api_base = self.get_api_base_ext()?; + + let body = openai_build_embeddings_body(data, &self.model); + + let url = format!("{api_base}/embeddings"); + + debug!("OpenAICompatible Embeddings Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } + + fn get_api_base_ext(&self) -> Result { + let api_base = match self.get_api_base() { + Ok(v) => v, + Err(err) => { + match OPENAI_COMPATIBLE_PLATFORMS + .into_iter() + .find_map(|(name, api_base)| { + if name == self.model.client_name() { + Some(api_base.to_string()) + } else { + None + } + }) { + Some(v) => v, + None => return Err(err), + } + } + }; + Ok(api_base) + } } impl_client_trait!( OpenAICompatibleClient, - crate::client::openai::openai_chat_completions, - crate::client::openai::openai_chat_completions_streaming + openai_chat_completions, + openai_chat_completions_streaming, + openai_embeddings ); diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index 0230e21..c34e409 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -1,8 +1,4 @@ -use super::{ - maybe_catch_error, message::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, QianwenClient, - SseHandler, SseMmessage, -}; +use super::*; use crate::utils::{base64_decode, sha256}; @@ -16,12 +12,15 @@ use serde::Deserialize; use serde_json::{json, Value}; use std::borrow::BorrowMut; -const API_URL: &str = +const CHAT_COMPLETIONS_API_URL: &str = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation"; -const API_URL_VL: &str = +const CHAT_COMPLETIONS_API_URL_VL: &str = "https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation"; +const EMBEDDINGS_API_URL: &str = + "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding"; + #[derive(Debug, Clone, Deserialize, Default)] pub struct QianwenConfig { pub name: Option, @@ -48,13 +47,13 @@ impl QianwenClient { let stream = data.stream; let url = match self.model.supports_vision() { - true => API_URL_VL, - false => API_URL, + true => CHAT_COMPLETIONS_API_URL_VL, + false => CHAT_COMPLETIONS_API_URL, }; let (mut body, has_upload) = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - debug!("Qianwen Request: {url} {body}"); + debug!("Qianwen Chat Completions Request: {url} {body}"); let mut builder = client.post(url).bearer_auth(api_key).json(&body); if stream { @@ -66,6 +65,37 @@ impl QianwenClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + + let text_type = match data.query { + true => "query", + false => "document", + }; + + let body = json!({ + "model": self.model.name(), + "input": { + "texts": data.texts, + }, + "parameters": { + "text_type": text_type, + } + }); + + let url = EMBEDDINGS_API_URL; + + debug!("Qianwen Embeddings Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } } #[async_trait] @@ -94,6 +124,15 @@ impl Client for QianwenClient { let builder = self.chat_completions_builder(client, data)?; chat_completions_streaming(builder, handler, &self.model).await } + + async fn embeddings_inner( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result>> { + let builder = self.embeddings_builder(client, data)?; + embeddings(builder).await + } } async fn chat_completions(builder: RequestBuilder, model: &Model) -> Result { @@ -210,6 +249,31 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu Ok((body, has_upload)) } +async fn embeddings( + builder: RequestBuilder, +) -> Result { + let data: Value = builder.send().await?.json().await?; + maybe_catch_error(&data)?; + let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?; + let output = res_body.output.embeddings.into_iter().map(|v| v.embedding).collect(); + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + output: EmbeddingsResBodyOutput, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyOutput { + embeddings: Vec, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyOutputEmbedding { + embedding: Vec, +} + fn extract_chat_completions_text(data: &Value, model: &Model) -> Result { let err = || anyhow!("Invalid response data: {data}"); let text = if model.name() == "qwen-long" { diff --git a/src/client/replicate.rs b/src/client/replicate.rs index 92c7e18..e96ed64 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -1,8 +1,5 @@ -use super::{ - catch_error, prompt_format::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, ReplicateClient, - SseHandler, SseMmessage, -}; +use super::*; +use super::prompt_format::*; use anyhow::{anyhow, Result}; use async_trait::async_trait; @@ -36,7 +33,7 @@ impl ReplicateClient { api_key: &str, ) -> Result { let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!("{API_BASE}/models/{}/predictions", self.model.name()); diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index b40247d..a9f84f8 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,8 +1,5 @@ -use super::{ - access_token::*, catch_error, json_stream, message::*, patch_system_message, - ChatCompletionsData, ChatCompletionsOutput, Client, ExtraConfig, Model, ModelData, - ModelPatches, PromptAction, PromptKind, SseHandler, ToolCall, VertexAIClient, -}; +use super::*; +use super::access_token::*; use anyhow::{anyhow, bail, Context, Result}; use async_trait::async_trait; @@ -51,9 +48,37 @@ impl VertexAIClient { let url = format!("{base_url}/google/models/{}:{func}", self.model.name()); let mut body = gemini_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - debug!("VertexAI Request: {url} {body}"); + debug!("VertexAI Chat Completions Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(access_token).json(&body); + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let project_id = self.get_project_id()?; + let location = self.get_location()?; + let access_token = get_access_token(self.name())?; + + let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers"); + let url = format!("{base_url}/google/models/{}:predict", self.model.name()); + + let task_type = match data.query { + true => "RETRIEVAL_DOCUMENT", + false => "QUESTION_ANSWERING", + }; + let instances: Vec<_> = data.texts.into_iter().map(|v| json!({"task_type": task_type, "content": v})).collect(); + let body = json!({ + "instances": instances, + }); + + debug!("VertexAI Embeddings Request: {url} {body}"); let builder = client.post(url).bearer_auth(access_token).json(&body); @@ -85,6 +110,16 @@ impl Client for VertexAIClient { let builder = self.chat_completions_builder(client, data)?; gemini_chat_completions_streaming(builder, handler).await } + + async fn embeddings_inner( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result>> { + prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?; + let builder = self.embeddings_builder(client, data)?; + embeddings(builder).await + } } pub async fn gemini_chat_completions(builder: RequestBuilder) -> Result { @@ -138,6 +173,34 @@ pub async fn gemini_chat_completions_streaming( Ok(()) } +async fn embeddings(builder: RequestBuilder) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + let res_body: EmbeddingsResBody = + serde_json::from_value(data).context("Invalid request data")?; + let output = res_body.predictions.into_iter().map(|v| v.embeddings.values).collect(); + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + predictions: Vec, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyPrediction { + embeddings: EmbeddingsResBodyPredictionEmbeddings, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyPredictionEmbeddings { + values: Vec +} + fn gemini_extract_chat_completions_text(data: &Value) -> Result { let text = data["candidates"][0]["content"]["parts"][0]["text"] .as_str() @@ -179,7 +242,7 @@ fn gemini_extract_chat_completions_text(data: &Value) -> Result Result { diff --git a/src/client/vertexai_claude.rs b/src/client/vertexai_claude.rs index bdce7d8..3993078 100644 --- a/src/client/vertexai_claude.rs +++ b/src/client/vertexai_claude.rs @@ -1,8 +1,7 @@ -use super::{ - access_token::*, claude::*, vertexai::*, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, - VertexAIClaudeClient, -}; +use super::*; +use super::access_token::*; +use super::claude::*; +use super::vertexai::*; use anyhow::Result; use async_trait::async_trait; @@ -46,7 +45,7 @@ impl VertexAIClaudeClient { ); let mut body = claude_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); if let Some(body_obj) = body.as_object_mut() { body_obj.remove("model"); } diff --git a/src/config/input.rs b/src/config/input.rs index ae94799..56ae5ed 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,11 +1,11 @@ use super::{role::Role, session::Session, GlobalConfig}; use crate::client::{ - init_client, list_models, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, + init_client, list_chat_models, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, MessageContentPart, MessageRole, Model, }; use crate::function::{ToolCallResult, ToolResults}; -use crate::utils::{base64_encode, sha256}; +use crate::utils::{base64_encode, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; use fancy_regex::Regex; @@ -29,9 +29,11 @@ lazy_static! { pub struct Input { config: GlobalConfig, text: String, + patch_text: Option, medias: Vec, data_urls: HashMap, tool_call: Option, + rag: Option, context: InputContext, } @@ -40,9 +42,11 @@ impl Input { Self { config: config.clone(), text: text.to_string(), + patch_text: None, medias: Default::default(), data_urls: Default::default(), tool_call: None, + rag: None, context: context.unwrap_or_else(|| InputContext::from_config(config)), } } @@ -92,9 +96,11 @@ impl Input { Ok(Self { config: config.clone(), text: texts.join("\n"), + patch_text: None, medias, data_urls, tool_call: Default::default(), + rag: None, context: context.unwrap_or_else(|| InputContext::from_config(config)), }) } @@ -108,13 +114,41 @@ impl Input { } pub fn text(&self) -> String { - self.text.clone() + match self.patch_text.clone() { + Some(text) => text, + None => self.text.clone(), + } } pub fn set_text(&mut self, text: String) { self.text = text; } + pub async fn maybe_embeddings(&mut self, abort_signal: AbortSignal) -> Result<()> { + if self.text.is_empty() { + return Ok(()); + } + if !self.text.is_empty() { + let rag = self.config.read().rag.clone(); + if let Some(rag) = rag { + let top_k = self.config.read().rag_top_k; + let embeddings = rag.search(&self.text, top_k, abort_signal).await?; + let text = self.config.read().rag_template(&embeddings, &self.text); + self.patch_text = Some(text); + self.rag = Some(rag.name().to_string()); + } + } + Ok(()) + } + + pub fn rag(&self) -> Option<&str> { + self.rag.as_deref() + } + + pub fn clear_patch_text(&mut self) { + self.patch_text.take(); + } + pub fn merge_tool_call( mut self, output: String, @@ -134,7 +168,7 @@ impl Input { let model = self.config.read().model.clone(); if let Some(model_id) = self.role().and_then(|v| v.model_id.clone()) { if model.id() != model_id { - if let Some(model) = list_models(&self.config.read()) + if let Some(model) = list_chat_models(&self.config.read()) .into_iter() .find(|v| v.id() == model_id) { @@ -158,7 +192,7 @@ impl Input { bail!("The current model does not support vision."); } let messages = self.build_messages()?; - self.config.read().model.max_input_tokens_limit(&messages)?; + self.config.read().model.guard_max_input_tokens(&messages)?; let (temperature, top_p) = if let Some(session) = self.session(&self.config.read().session) { (session.temperature(), session.top_p()) @@ -262,12 +296,12 @@ impl Input { pub fn render(&self) -> String { if self.medias.is_empty() { - return self.text.clone(); + return self.text(); } let text = if self.text.is_empty() { - self.text.to_string() + String::new() } else { - format!(" -- {}", self.text) + format!(" -- {}", self.text()) }; let files: Vec = self .medias @@ -280,7 +314,7 @@ impl Input { pub fn message_content(&self) -> MessageContent { if self.medias.is_empty() { - MessageContent::Text(self.text.clone()) + MessageContent::Text(self.text()) } else { let mut list: Vec = self .medias @@ -291,12 +325,7 @@ impl Input { }) .collect(); if !self.text.is_empty() { - list.insert( - 0, - MessageContentPart::Text { - text: self.text.clone(), - }, - ); + list.insert(0, MessageContentPart::Text { text: self.text() }); } MessageContent::Array(list) } diff --git a/src/config/mod.rs b/src/config/mod.rs index 5ae6bec..fc2df17 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -7,14 +7,15 @@ pub use self::role::{Role, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ - create_client_config, list_client_types, list_models, ClientConfig, Model, + create_client_config, list_chat_models, list_client_types, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{Function, ToolCallResult}; +use crate::rag::{Rag, TEMP_RAG_NAME}; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{ format_option_value, fuzzy_match, get_env_name, light_theme_from_colorfgbg, now, render_prompt, - set_text, + set_text, AbortSignal, }; use anyhow::{anyhow, bail, Context, Result}; @@ -42,6 +43,7 @@ const CONFIG_FILE_NAME: &str = "config.yaml"; const ROLES_FILE_NAME: &str = "roles.yaml"; const MESSAGES_FILE_NAME: &str = "messages.md"; const SESSIONS_DIR_NAME: &str = "sessions"; +const RAGS_DIR_NAME: &str = "rags"; const FUNCTIONS_DIR_NAME: &str = "functions"; const CLIENTS_FIELD: &str = "clients"; @@ -49,7 +51,16 @@ const CLIENTS_FIELD: &str = "clients"; const SUMMARIZE_PROMPT: &str = "Summarize the discussion briefly in 200 words or less to use as a prompt for future context."; const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: "; -const LEFT_PROMPT: &str = "{color.green}{?session {session}{?role /}}{role}{color.cyan}{?session )}{!session >}{color.reset} "; + +const RAG_TEMPLATE: &str = r#"Answer the following question based only on the provided context: + +__CONTEXT__ + + +Question: __INPUT__ +"#; + +const LEFT_PROMPT: &str = "{color.green}{?session {session}{?role /}}{role}{?rag #{rag}}{color.cyan}{?session )}{!session >}{color.reset} "; const RIGHT_PROMPT: &str = "{color.purple}{?session {?consume_tokens {consume_tokens}({consume_percent}%)}{!consume_tokens {consume_tokens}}}{color.reset}"; #[derive(Debug, Clone, Deserialize)] @@ -71,6 +82,9 @@ pub struct Config { pub keybindings: Keybindings, pub prelude: Option, pub buffer_editor: Option, + pub embedding_model: Option, + pub rag_top_k: usize, + pub rag_template: Option, pub function_calling: bool, pub compress_threshold: usize, pub summarize_prompt: Option, @@ -85,6 +99,8 @@ pub struct Config { #[serde(skip)] pub session: Option, #[serde(skip)] + pub rag: Option>, + #[serde(skip)] pub model: Model, #[serde(skip)] pub function: Function, @@ -111,6 +127,9 @@ impl Default for Config { keybindings: Default::default(), prelude: None, buffer_editor: None, + embedding_model: None, + rag_top_k: 4, + rag_template: None, function_calling: false, compress_threshold: 4000, summarize_prompt: None, @@ -121,6 +140,7 @@ impl Default for Config { roles: vec![], role: None, session: None, + rag: None, model: Default::default(), function: Default::default(), working_mode: WorkingMode::Command, @@ -170,12 +190,12 @@ impl Config { match prelude.split_once(':') { Some(("role", name)) => { if self.role.is_none() && self.session.is_none() { - self.set_role(name).with_context(err_msg)?; + self.use_role(name).with_context(err_msg)?; } } Some(("session", name)) => { if self.session.is_none() { - self.start_session(Some(name)).with_context(err_msg)?; + self.use_session(Some(name)).with_context(err_msg)?; } } _ => { @@ -223,10 +243,11 @@ impl Config { pub fn save_message( &mut self, - input: &Input, + input: &mut Input, output: &str, tool_call_results: &[ToolCallResult], ) -> Result<()> { + input.clear_patch_text(); self.last_message = Some((input.clone(), output.to_string())); if self.dry_run || output.is_empty() || !tool_call_results.is_empty() { @@ -248,17 +269,13 @@ impl Config { let timestamp = now(); let summary = input.summary(); let input_markdown = input.render(); - let output = match input.role() { - None => { - format!("# CHAT: {summary} [{timestamp}]\n{input_markdown}\n--------\n{output}\n--------\n\n",) - } - Some(v) => { - format!( - "# CHAT: {summary} [{timestamp}] ({})\n{input_markdown}\n--------\n{output}\n--------\n\n", - v.name, - ) - } + let scope = match (input.role().map(|v| v.name.as_str()), input.rag()) { + (Some(role), Some(rag)) => format!(" ({role}#{rag})"), + (Some(role), _) => format!(" ({role})"), + (None, Some(rag)) => format!(" (#{rag})"), + _ => String::new(), }; + let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",); file.write_all(output.as_bytes()) .with_context(|| "Failed to save message") } @@ -289,6 +306,10 @@ impl Config { Self::local_path(SESSIONS_DIR_NAME) } + pub fn rags_dir() -> Result { + Self::local_path(RAGS_DIR_NAME) + } + pub fn functions_dir() -> Result { Self::local_path(FUNCTIONS_DIR_NAME) } @@ -299,17 +320,23 @@ impl Config { Ok(path) } - pub fn set_prompt(&mut self, prompt: &str) -> Result<()> { + pub fn rag_file(name: &str) -> Result { + let mut path = Self::rags_dir()?; + path.push(&format!("{name}.bin")); + Ok(path) + } + + pub fn use_prompt(&mut self, prompt: &str) -> Result<()> { let role = Role::temp(prompt); - self.set_role_obj(role) + self.use_role_obj(role) } - pub fn set_role(&mut self, name: &str) -> Result<()> { + pub fn use_role(&mut self, name: &str) -> Result<()> { let role = self.retrieve_role(name)?; - self.set_role_obj(role) + self.use_role_obj(role) } - pub fn set_role_obj(&mut self, role: Role) -> Result<()> { + pub fn use_role_obj(&mut self, role: Role) -> Result<()> { if let Some(session) = self.session.as_mut() { session.guard_empty()?; session.set_role_properties(&role); @@ -321,7 +348,7 @@ impl Config { Ok(()) } - pub fn clear_role(&mut self) -> Result<()> { + pub fn exit_role(&mut self) -> Result<()> { self.role = None; self.restore_model()?; Ok(()) @@ -337,7 +364,10 @@ impl Config { } } if self.role.is_some() { - flags |= StateFlags::ROLE + flags |= StateFlags::ROLE; + } + if self.rag.is_some() { + flags |= StateFlags::RAG; } flags } @@ -393,7 +423,7 @@ impl Config { } pub fn set_model(&mut self, value: &str) -> Result<()> { - let models = list_models(self); + let models = list_chat_models(self); let model = Model::find(&models, value); match model { None => bail!("No model '{}'", value), @@ -442,6 +472,7 @@ impl Config { ), ("temperature", format_option_value(&temperature)), ("top_p", format_option_value(&top_p)), + ("rag_top_k", self.rag_top_k.to_string()), ("function_calling", self.function_calling.to_string()), ("compress_threshold", self.compress_threshold.to_string()), ("dry_run", self.dry_run.to_string()), @@ -458,6 +489,7 @@ impl Config { ("roles_file", display_path(&Self::roles_file()?)), ("messages_file", display_path(&Self::messages_file()?)), ("sessions_dir", display_path(&Self::sessions_dir()?)), + ("rags_dir", display_path(&Self::rags_dir()?)), ("functions_dir", display_path(&Self::functions_dir()?)), ]; let output = items @@ -486,11 +518,21 @@ impl Config { } } + pub fn rag_info(&self) -> Result { + if let Some(rag) = &self.rag { + rag.export() + } else { + bail!("No rag") + } + } + pub fn info(&self) -> Result { if let Some(session) = &self.session { session.export() } else if let Some(role) = &self.role { role.export() + } else if let Some(rag) = &self.rag { + rag.export() } else { self.system_info() } @@ -511,7 +553,7 @@ impl Config { .iter() .map(|v| (v.name.clone(), String::new())) .collect(), - ".model" => list_models(self) + ".model" => list_chat_models(self) .into_iter() .map(|v| (v.id(), v.description())) .collect(), @@ -520,10 +562,16 @@ impl Config { .into_iter() .map(|v| (v.clone(), String::new())) .collect(), + ".rag" => self + .list_rags() + .into_iter() + .map(|v| (v.clone(), String::new())) + .collect(), ".set" => vec![ "max_output_tokens", "temperature", "top_p", + "rag_top_k", "function_calling", "compress_threshold", "save", @@ -592,6 +640,11 @@ impl Config { let value = parse_value(value)?; self.set_top_p(value); } + "rag_top_k" => { + if let Some(value) = parse_value(value)? { + self.rag_top_k = value; + } + } "function_calling" => { let value = value.parse().with_context(|| "Invalid value")?; self.function_calling = value; @@ -625,7 +678,7 @@ impl Config { Ok(()) } - pub fn start_session(&mut self, session: Option<&str>) -> Result<()> { + pub fn use_session(&mut self, session: Option<&str>) -> Result<()> { if self.session.is_some() { bail!( "Already in a session, please run '.exit session' first to exit the current session." @@ -671,7 +724,7 @@ impl Config { Ok(()) } - pub fn end_session(&mut self) -> Result<()> { + pub fn exit_session(&mut self) -> Result<()> { if let Some(mut session) = self.session.take() { self.last_message = None; let save_session = session.save_session(); @@ -767,6 +820,74 @@ impl Config { } } + pub async fn use_rag( + config: &GlobalConfig, + rag: Option<&str>, + abort_signal: AbortSignal, + ) -> Result<()> { + if config.read().rag.is_some() { + bail!("Already in a rag, please run '.exit rag' first to exit the current rag."); + } + let rag = match rag { + None => { + let rag_path = Self::rag_file(TEMP_RAG_NAME)?; + if rag_path.exists() { + remove_file(&rag_path).with_context(|| { + format!("Failed to cleanup previous '{TEMP_RAG_NAME}' rag") + })?; + } + Rag::init(config, TEMP_RAG_NAME, &rag_path, abort_signal).await? + } + Some(name) => { + let rag_path = Self::rag_file(name)?; + if !rag_path.exists() { + Rag::init(config, name, &rag_path, abort_signal).await? + } else { + Rag::load(config, name, &rag_path)? + } + } + }; + config.write().rag = Some(Arc::new(rag)); + Ok(()) + } + + pub fn exit_rag(&mut self) -> Result<()> { + self.rag.take(); + Ok(()) + } + + pub fn list_rags(&self) -> Vec { + let rags_dir = match Self::rags_dir() { + Ok(dir) => dir, + Err(_) => return vec![], + }; + match read_dir(rags_dir) { + Ok(rd) => { + let mut names = vec![]; + for entry in rd.flatten() { + let name = entry.file_name(); + if let Some(name) = name.to_string_lossy().strip_suffix(".bin") { + names.push(name.to_string()); + } + } + names.sort_unstable(); + names + } + Err(_) => vec![], + } + } + + pub fn rag_template(&self, embeddings: &str, text: &str) -> String { + if embeddings.is_empty() { + return text.to_string(); + } + self.rag_template + .as_deref() + .unwrap_or(RAG_TEMPLATE) + .replace("__CONTEXT__", embeddings) + .replace("__INPUT__", text) + } + pub fn get_render_options(&self) -> Result { let theme = if self.highlight { let theme_mode = if self.light_theme { "light" } else { "dark" }; @@ -858,6 +979,9 @@ impl Config { output.insert("consume_percent", percent.to_string()); output.insert("user_messages_len", session.user_messages_len().to_string()); } + if let Some(rag) = &self.rag { + output.insert("rag", rag.name().to_string()); + } if self.highlight { output.insert("color.reset", "\u{1b}[0m".to_string()); @@ -974,7 +1098,7 @@ impl Config { fn setup_model(&mut self) -> Result<()> { let model_id = if self.model_id.is_empty() { - let models = list_models(self); + let models = list_chat_models(self); if models.is_empty() { bail!("No available model"); } @@ -1049,6 +1173,7 @@ bitflags::bitflags! { const ROLE = 1 << 0; const SESSION_EMPTY = 1 << 1; const SESSION = 1 << 2; + const RAG = 1 << 3; } } @@ -1090,12 +1215,12 @@ fn create_config_file(config_path: &Path) -> Result<()> { std::fs::set_permissions(config_path, perms)?; } - println!("✨ Saved config file to {}\n", config_path.display()); + println!("✨ Saved config file to '{}'\n", config_path.display()); Ok(()) } -fn ensure_parent_exists(path: &Path) -> Result<()> { +pub(crate) fn ensure_parent_exists(path: &Path) -> Result<()> { if path.exists() { return Ok(()); } diff --git a/src/config/session.rs b/src/config/session.rs index 8ac5ad9..14e0731 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -160,7 +160,7 @@ impl Session { data["messages"] = json!(self.messages); let output = serde_yaml::to_string(&data) - .with_context(|| format!("Unable to show info about session {}", &self.name))?; + .with_context(|| format!("Unable to show info about session '{}'", &self.name))?; Ok(output) } diff --git a/src/main.rs b/src/main.rs index 0a5e404..3222285 100644 --- a/src/main.rs +++ b/src/main.rs @@ -3,6 +3,7 @@ mod client; mod config; mod function; mod logger; +mod rag; mod render; mod repl; mod serve; @@ -13,12 +14,12 @@ mod utils; extern crate log; use crate::cli::Cli; -use crate::client::{list_models, send_stream, ChatCompletionsOutput}; +use crate::client::{list_chat_models, send_stream, ChatCompletionsOutput}; use crate::config::{ Config, GlobalConfig, Input, InputContext, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, }; -use crate::function::eval_tool_calls; +use crate::function::{eval_tool_calls, need_send_call_results}; use crate::render::{render_error, MarkdownRender}; use crate::repl::Repl; use crate::utils::{ @@ -29,14 +30,12 @@ use crate::utils::{ use anyhow::{bail, Result}; use async_recursion::async_recursion; use clap::Parser; -use function::need_send_call_results; use inquire::{Select, Text}; use is_terminal::IsTerminal; use parking_lot::RwLock; use std::io::{stderr, stdin, stdout, Read}; use std::process; use std::sync::Arc; -use tokio::sync::oneshot; #[tokio::main] async fn main() -> Result<()> { @@ -67,7 +66,7 @@ async fn main() -> Result<()> { return Ok(()); } if cli.list_models { - for model in list_models(&config.read()) { + for model in list_chat_models(&config.read()) { println!("{}", model.id()); } return Ok(()); @@ -87,18 +86,18 @@ async fn main() -> Result<()> { config.write().dry_run = true; } if let Some(prompt) = &cli.prompt { - config.write().set_prompt(prompt)?; + config.write().use_prompt(prompt)?; } else if let Some(name) = &cli.role { - config.write().set_role(name)?; + config.write().use_role(name)?; } else if cli.execute { - config.write().set_role(SHELL_ROLE)?; + config.write().use_role(SHELL_ROLE)?; } else if cli.code { - config.write().set_role(CODE_ROLE)?; + config.write().use_role(CODE_ROLE)?; } if let Some(session) = &cli.session { config .write() - .start_session(session.as_ref().map(|v| v.as_str()))?; + .use_session(session.as_ref().map(|v| v.as_str()))?; } if let Some(model) = &cli.model { config.write().set_model(model)?; @@ -142,7 +141,7 @@ async fn main() -> Result<()> { #[async_recursion] async fn start_directive( config: &GlobalConfig, - input: Input, + mut input: Input, no_stream: bool, code_mode: bool, ) -> Result<()> { @@ -176,8 +175,8 @@ async fn start_directive( }; config .write() - .save_message(&input, &output, &tool_call_results)?; - config.write().end_session()?; + .save_message(&mut input, &output, &tool_call_results)?; + config.write().exit_session()?; if need_send_call_results(&tool_call_results) { start_directive( config, @@ -201,10 +200,9 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - let client = input.create_client()?; let is_terminal_stdout = stdout().is_terminal(); let ret = if is_terminal_stdout { - let (spinner_tx, spinner_rx) = oneshot::channel(); - tokio::spawn(run_spinner(" Generating", spinner_rx)); + let (stop_spinner_tx, _) = run_spinner("Generating").await; let ret = client.chat_completions(input.clone()).await; - let _ = spinner_tx.send(()); + let _ = stop_spinner_tx.send(()); ret } else { client.chat_completions(input.clone()).await @@ -213,7 +211,7 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } - config.write().save_message(&input, &eval_str, &[])?; + config.write().save_message(&mut input, &eval_str, &[])?; config.read().maybe_copy(&eval_str); let render_options = config.read().get_render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; diff --git a/src/rag/loader.rs b/src/rag/loader.rs new file mode 100644 index 0000000..106802a --- /dev/null +++ b/src/rag/loader.rs @@ -0,0 +1,146 @@ +use super::RagDocument; + +use anyhow::{bail, Context, Result}; +use async_recursion::async_recursion; +use std::{path::Path, process::Command}; +use tokio::fs; + +pub async fn load(path: &str, extension: &str) -> Result> { + match extension { + "docx" | "epub" | "ipynb" => load_pandoc(path) + .await + .context("Failed to load with pandoc"), + "pdf" => load_pdf(path).await, + _ => load_plain(path).await, + } +} + +async fn load_plain(path: &str) -> Result> { + let contents = fs::read_to_string(path).await?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pdf(path: &str) -> Result> { + let contents = pdf_extract::extract_text(path)?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pandoc(path: &str) -> Result> { + let output = Command::new("pandoc") + .arg("--to") + .arg("plain") + .arg(path) + .output()?; + + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + bail!( + "Pandoc conversion failed with exit code {:?}: {}", + output.status.code(), + stderr + ); + } + + let contents = std::str::from_utf8(&output.stdout)?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +pub fn parse_glob(path_str: &str) -> Result<(String, Vec)> { + if let Some(start) = path_str.find("/**/*.").or_else(|| path_str.find(r"\**\*.")) { + let base_path = path_str[..start].to_string(); + if let Some(curly_brace_end) = path_str[start..].find('}') { + let end = start + curly_brace_end; + let extensions_str = &path_str[start + 6..end + 1]; + let extensions = if extensions_str.starts_with('{') && extensions_str.ends_with('}') { + extensions_str[1..extensions_str.len() - 1] + .split(',') + .map(|s| s.to_string()) + .collect::>() + } else { + bail!("Invalid path '{path_str}'"); + }; + Ok((base_path, extensions)) + } else { + let extensions_str = &path_str[start + 6..]; + let extensions = vec![extensions_str.to_string()]; + Ok((base_path, extensions)) + } + } else { + Ok((path_str.to_string(), vec![])) + } +} + +#[async_recursion] +pub async fn list_files( + files: &mut Vec, + entry_path: &Path, + suffixes: Option<&Vec>, +) -> Result<()> { + if !entry_path.exists() { + bail!("Not found: {:?}", entry_path); + } + if entry_path.is_file() { + add_file(files, suffixes, entry_path); + return Ok(()); + } + if !entry_path.is_dir() { + bail!("Not a directory: {:?}", entry_path); + } + let mut reader = fs::read_dir(entry_path).await?; + while let Some(entry) = reader.next_entry().await? { + let path = entry.path(); + if path.is_file() { + add_file(files, suffixes, &path); + } else if path.is_dir() { + list_files(files, &path, suffixes).await?; + } + } + Ok(()) +} + +fn add_file(files: &mut Vec, suffixes: Option<&Vec>, path: &Path) { + if is_valid_extension(suffixes, path) { + files.push(path.display().to_string()); + } +} + +fn is_valid_extension(suffixes: Option<&Vec>, path: &Path) -> bool { + if let Some(suffixes) = suffixes { + if !suffixes.is_empty() { + if let Some(extension) = path.extension().map(|v| v.to_string_lossy().to_string()) { + return suffixes.contains(&extension); + } + return false; + } + } + true +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_glob() { + assert_eq!(parse_glob("dir").unwrap(), ("dir".into(), vec![])); + assert_eq!( + parse_glob("dir/file.md").unwrap(), + ("dir/file.md".into(), vec![]) + ); + assert_eq!( + parse_glob("dir/**/*.md").unwrap(), + ("dir".into(), vec!["md".into()]) + ); + assert_eq!( + parse_glob("dir/**/*.{md,txt}").unwrap(), + ("dir".into(), vec!["md".into(), "txt".into()]) + ); + assert_eq!( + parse_glob("C:\\dir\\**\\*.{md,txt}").unwrap(), + ("C:\\dir".into(), vec!["md".into(), "txt".into()]) + ); + } +} diff --git a/src/rag/mod.rs b/src/rag/mod.rs new file mode 100644 index 0000000..387d3d9 --- /dev/null +++ b/src/rag/mod.rs @@ -0,0 +1,425 @@ +use self::loader::*; +use self::splitter::*; + +use crate::client::*; +use crate::config::*; +use crate::utils::*; + +mod loader; +mod splitter; + +use anyhow::bail; +use anyhow::{anyhow, Context, Result}; +use hnsw_rs::prelude::*; +use indexmap::IndexMap; +use inquire::{required, validator::Validation, Select, Text}; +use path_absolutize::Absolutize; +use serde::{Deserialize, Serialize}; +use serde_json::json; +use std::fmt::Debug; +use std::{io::BufReader, path::Path}; +use tokio::sync::mpsc; + +pub const TEMP_RAG_NAME: &str = "temp"; +pub const CHUNK_OVERLAP: usize = 20; +pub const SIMILARITY_THRESHOLD: f32 = 0.25; + +pub struct Rag { + client: Box, + name: String, + path: String, + model: Model, + hnsw: Hnsw<'static, f32, DistCosine>, + data: RagData, +} + +impl Debug for Rag { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Rag") + .field("name", &self.name) + .field("path", &self.path) + .field("model", &self.model) + .field("data", &self.data) + .finish() + } +} + +impl Rag { + pub async fn init( + config: &GlobalConfig, + name: &str, + path: &Path, + abort_signal: AbortSignal, + ) -> Result { + debug!("init rag: {name}"); + let model = select_embedding_model(config)?; + let chunk_size = model.default_chunk_size(); + let chunk_size = set_chunk_size(chunk_size)?; + let data = RagData::new(&model.id(), chunk_size); + let mut rag = Self::create(config, name, path, data)?; + let paths = add_document_paths()?; + debug!("document paths: {paths:?}"); + let (stop_spinner_tx, set_spinner_message_tx) = run_spinner("Starting").await; + tokio::select! { + ret = rag.add_paths(&paths, Some(set_spinner_message_tx)) => { + let _ = stop_spinner_tx.send(()); + ret?; + } + _ = watch_abort_signal(abort_signal) => { + let _ = stop_spinner_tx.send(()); + bail!("Aborted!") + }, + }; + if !rag.is_temp() { + rag.save(path)?; + println!("✨ Saved rag to '{}'", path.display()); + } + Ok(rag) + } + + pub fn load(config: &GlobalConfig, name: &str, path: &Path) -> Result { + let err = || format!("Failed to load rag '{name}'"); + let file = std::fs::File::open(path).with_context(err)?; + let reader = BufReader::new(file); + let data: RagData = bincode::deserialize_from(reader).with_context(err)?; + Self::create(config, name, path, data) + } + + pub fn create(config: &GlobalConfig, name: &str, path: &Path, data: RagData) -> Result { + let hnsw = data.build_hnsw(); + let model = retrieve_embedding_model(&config.read(), &data.model)?; + let client = init_client(config, Some(model.clone()))?; + let rag = Rag { + client, + name: name.to_string(), + path: path.display().to_string(), + data, + model, + hnsw, + }; + Ok(rag) + } + + pub fn save(&self, path: &Path) -> Result<()> { + ensure_parent_exists(path)?; + let mut file = std::fs::File::create(path)?; + bincode::serialize_into(&mut file, &self.data) + .with_context(|| format!("Failed to save rag '{}'", self.name))?; + Ok(()) + } + + pub fn export(&self) -> Result { + let files: Vec<_> = self.data.files.iter().map(|v| &v.path).collect(); + let data = json!({ + "path": self.path, + "model": self.model.id(), + "chunk_size": self.data.chunk_size, + "files": files, + }); + let output = serde_yaml::to_string(&data) + .with_context(|| format!("Unable to show info about rag '{}'", self.name))?; + Ok(output) + } + + pub fn name(&self) -> &str { + &self.name + } + + pub fn is_temp(&self) -> bool { + self.name == TEMP_RAG_NAME + } + + pub async fn search( + &self, + text: &str, + top_k: usize, + abort_signal: AbortSignal, + ) -> Result { + let (stop_spinner_tx, _) = run_spinner("Embedding").await; + let ret = tokio::select! { + ret = self.search_impl(text, top_k) => { + ret + } + _ = watch_abort_signal(abort_signal) => { + bail!("Aborted!") + }, + }; + let _ = stop_spinner_tx.send(()); + let output = ret?.join("\n\n"); + Ok(output) + } + + pub async fn add_paths>( + &mut self, + paths: &[T], + progress_tx: Option>, + ) -> Result<()> { + // List files + let mut file_paths = vec![]; + progress(&progress_tx, "Listing paths".into()); + for path in paths { + let path = path + .as_ref() + .absolutize() + .with_context(|| anyhow!("Invalid path '{}'", path.as_ref().display()))?; + let path_str = path.display().to_string(); + if self.data.files.iter().any(|v| v.path == path_str) { + continue; + } + let (path_str, suffixes) = parse_glob(&path_str)?; + let suffixes = if suffixes.is_empty() { + None + } else { + Some(&suffixes) + }; + list_files(&mut file_paths, Path::new(&path_str), suffixes).await?; + } + + // Load files + let mut rag_files = vec![]; + let file_paths_len = file_paths.len(); + progress(&progress_tx, format!("Loading files [1/{file_paths_len}]")); + for path in file_paths { + let extension = Path::new(&path) + .extension() + .map(|v| v.to_string_lossy().to_lowercase()) + .unwrap_or_default(); + let separator = autodetect_separator(&extension); + let splitter = Splitter::new(self.data.chunk_size, CHUNK_OVERLAP, separator); + let documents = load(&path, &extension) + .await + .with_context(|| format!("Failed to load text at '{path}'"))?; + let documents = + splitter.split_documents(&documents, &SplitterChunkHeaderOptions::default()); + rag_files.push(RagFile { path, documents }); + progress( + &progress_tx, + format!("Loading files [{}/{file_paths_len}]", rag_files.len()), + ); + } + + if rag_files.is_empty() { + return Ok(()); + } + + // Convert vectors + let mut vector_ids = vec![]; + let mut texts = vec![]; + for (file_index, file) in rag_files.iter().enumerate() { + for (document_index, doc) in file.documents.iter().enumerate() { + vector_ids.push(combine_vector_id(file_index, document_index)); + texts.push(doc.page_content.clone()) + } + } + + let embeddings_data = EmbeddingsData::new(texts, false); + let embeddings = self + .create_embeddings(embeddings_data, progress_tx.clone()) + .await?; + + self.data.add(rag_files, vector_ids, embeddings); + progress(&progress_tx, "Building vector store".into()); + self.hnsw = self.data.build_hnsw(); + + Ok(()) + } + + async fn search_impl(&self, text: &str, top_k: usize) -> Result> { + let splitter = Splitter::new(self.data.chunk_size, CHUNK_OVERLAP, &DEFAULT_SEPARATES); + let texts = splitter.split_text(text); + let embeddings_data = EmbeddingsData::new(texts, true); + let embeddings = self.create_embeddings(embeddings_data, None).await?; + let output = self + .hnsw + .parallel_search(&embeddings, top_k, 30) + .into_iter() + .flat_map(|list| { + list.into_iter() + .filter_map(|v| { + if v.distance < SIMILARITY_THRESHOLD { + return None; + } + let (file_index, document_index) = split_vector_id(v.d_id); + let text = self.data.files[file_index].documents[document_index] + .page_content + .clone(); + Some(text) + }) + .collect::>() + }) + .collect(); + Ok(output) + } + + async fn create_embeddings( + &self, + data: EmbeddingsData, + progress_tx: Option>, + ) -> Result { + let EmbeddingsData { texts, query } = data; + let mut output = vec![]; + let chunks = texts.chunks(self.model.max_concurrent_chunks()); + let chunks_len = chunks.len(); + progress( + &progress_tx, + format!("Creating embeddings [1/{chunks_len}]"), + ); + for (index, texts) in chunks.enumerate() { + let chunk_data = EmbeddingsData { + texts: texts.to_vec(), + query, + }; + let chunk_output = self + .client + .embeddings(chunk_data) + .await + .context("Failed to create embedding")?; + output.extend(chunk_output); + progress( + &progress_tx, + format!("Creating embeddings [{}/{chunks_len}]", index + 1), + ); + } + Ok(output) + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagData { + pub model: String, + pub chunk_size: usize, + pub files: Vec, + pub vectors: IndexMap>, +} + +impl RagData { + pub fn new(model: &str, chunk_size: usize) -> Self { + Self { + model: model.to_string(), + chunk_size, + files: Default::default(), + vectors: Default::default(), + } + } + + pub fn add( + &mut self, + files: Vec, + vector_ids: Vec, + embeddings: EmbeddingsOutput, + ) { + self.files.extend(files); + self.vectors.extend(vector_ids.into_iter().zip(embeddings)); + } + + pub fn build_hnsw(&self) -> Hnsw<'static, f32, DistCosine> { + let hnsw = Hnsw::new(32, self.vectors.len(), 16, 200, DistCosine {}); + let list: Vec<_> = self.vectors.iter().map(|(k, v)| (v, *k)).collect(); + hnsw.parallel_insert(&list); + hnsw + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagFile { + path: String, + documents: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagDocument { + pub page_content: String, + pub metadata: RagMetadata, +} + +impl RagDocument { + pub fn new>(page_content: S) -> Self { + RagDocument { + page_content: page_content.into(), + metadata: IndexMap::new(), + } + } + + #[allow(unused)] + pub fn with_metadata(mut self, metadata: RagMetadata) -> Self { + self.metadata = metadata; + self + } +} + +impl Default for RagDocument { + fn default() -> Self { + RagDocument { + page_content: "".to_string(), + metadata: IndexMap::new(), + } + } +} + +pub type RagMetadata = IndexMap; + +pub type VectorID = usize; + +pub fn combine_vector_id(file_index: usize, document_index: usize) -> VectorID { + file_index << (usize::BITS / 2) | document_index +} + +pub fn split_vector_id(value: VectorID) -> (usize, usize) { + let low_mask = (1 << (usize::BITS / 2)) - 1; + let low = value & low_mask; + let high = value >> (usize::BITS / 2); + (high, low) +} + +fn retrieve_embedding_model(config: &Config, model_id: &str) -> Result { + let models = list_embedding_models(config); + let model = + Model::find(&models, model_id).ok_or_else(|| anyhow!("No embedding model '{model_id}'"))?; + Ok(model) +} + +fn select_embedding_model(config: &GlobalConfig) -> Result { + let config = config.read(); + let model = match config.embedding_model.clone() { + Some(model_id) => retrieve_embedding_model(&config, &model_id)?, + None => { + let models = list_embedding_models(&config); + if models.is_empty() { + bail!("No embedding model"); + } + let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect(); + let model_id = Select::new("Select embedding model:", model_ids).prompt()?; + retrieve_embedding_model(&config, &model_id)? + } + }; + Ok(model) +} + +fn set_chunk_size(chunk_size: usize) -> Result { + let value = Text::new("Set chunk size:") + .with_default(&chunk_size.to_string()) + .with_validator(move |text: &str| { + let out = match text.parse::() { + Ok(_) => Validation::Valid, + Err(_) => Validation::Invalid("Must be a integer".into()), + }; + Ok(out) + }) + .prompt()?; + value.parse().map_err(|_| anyhow!("Invalid chunk_size")) +} + +fn add_document_paths() -> Result> { + let text = Text::new("Add document paths:") + .with_validator(required!("This field is required")) + .with_help_message("e.g. file1;dir2/;dir3/**/*.md") + .prompt()?; + let paths = text.split(';').map(|v| v.to_string()).collect(); + Ok(paths) +} + +fn progress(spinner_message_tx: &Option>, message: String) { + if let Some(tx) = spinner_message_tx { + let _ = tx.send(message); + } +} diff --git a/src/rag/splitter.rs b/src/rag/splitter.rs new file mode 100644 index 0000000..5fdacee --- /dev/null +++ b/src/rag/splitter.rs @@ -0,0 +1,564 @@ +use super::{RagDocument, RagMetadata}; + +use std::cmp::Ordering; + +pub const DEFAULT_SEPARATES: [&str; 4] = ["\n\n", "\n", " ", ""]; +pub const HTML_SEPARATES: [&str; 28] = [ + // First, try to split along HTML tags + "", "
", "

", "
", "

  • ", "

    ", "

    ", "

    ", "

    ", "

    ", "
    ", + "", "", "", "
    ", "", "
      ", "
        ", "
        ", "