diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-05 09:02:23 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-05 09:02:23 +0800 |
| commit | 1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch) | |
| tree | 6f68a860a39b0fbe87784de925b5a31e00e74e33 | |
| parent | 71f2e94579511d7524f5534377001ab3f02a9597 (diff) | |
| download | aichat-1ec6abfaee2fdc189b348b7e3a8145bd9a84da74.tar.gz | |
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)
34 files changed, 2616 insertions, 304 deletions
@@ -18,6 +18,27 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -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", @@ -86,6 +111,12 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -101,6 +132,24 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -471,6 +520,16 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -487,6 +546,16 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -505,6 +574,31 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -578,6 +672,12 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -637,6 +737,104 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -659,6 +857,15 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -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" @@ -937,6 +1148,31 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -992,6 +1228,12 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1219,6 +1461,12 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1257,6 +1505,32 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1266,12 +1540,27 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1315,6 +1604,23 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1325,6 +1631,19 @@ dependencies = [ [[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" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ab2156c4fce2f8df6c499cc1c763e4394b7482525bf2a9701c9d79d215f519e4" @@ -1601,6 +1920,39 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1669,6 +2021,18 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1681,6 +2045,16 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1738,6 +2112,26 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -2291,6 +2685,20 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -2608,6 +3016,15 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -2931,6 +3348,18 @@ dependencies = [ ] [[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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -2963,6 +3392,15 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9252e5725dbed82865af151df558e754e4a3c2c30818359eb17465f1346a1b49" @@ -3157,7 +3595,7 @@ dependencies = [ "derive-new", "libc", "log", - "nix", + "nix 0.28.0", "os_pipe", "tempfile", "thiserror", @@ -3186,6 +3624,32 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -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] @@ -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: <your_api_key_here> -β¨ Saved config file to <config-dir>/aichat/config.yaml +β¨ Saved config file to '<user-config-dir>/aichat/config.yaml' ``` Feel free to adjust the configuration according to your needs. @@ -340,12 +341,27 @@ Usage: .file <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 '<user-config-dir>/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> + __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: # <regex>: # 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: <path-to/gcloud/application_default_credentials.json> 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<String>, pub api_base: Option<String>, pub api_key: Option<String>, + #[serde(default)] pub models: Vec<ModelData>, pub patches: Option<ModelPatches>, pub extra: Option<ExtraConfig>, @@ -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<RequestBuilder> { + 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<RequestBuilder> { + 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<ChatCompletionsOutput> { let res = builder.send().await?; @@ -100,6 +129,24 @@ async fn chat_completions_streaming( Ok(()) } +async fn embeddings( + builder: RequestBuilder, +) -> Result<EmbeddingsOutput> { + 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<Vec<f32>>, +} + fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> { 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<BuiltinModels> = - serde_yaml::from_str(MODELS_YAML).unwrap(); + pub static ref ALL_MODELS: Vec<BuiltinModels> = serde_yaml::from_str(MODELS_YAML).unwrap(); static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(?<!\\)/").unwrap(); } @@ -92,10 +94,10 @@ macro_rules! register_client { pub fn list_models(local_config: &$config) -> Vec<Model> { 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<Vec<$crate::client::Model>> = None; + static mut ALL_CLIENT_MODELS: Option<Vec<$crate::client::Model>> = 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<Model> { - 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<Vec<Vec<f32>>> { + 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<Model>; - 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<ReqwestClient> { 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<Vec<Vec<f32>>> { + 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<Vec<Vec<f32>>> { + bail!("No embeddings api") + } } impl Default for ClientConfig { @@ -375,7 +411,7 @@ pub type ModelPatches = IndexMap<String, ModelPatch>; #[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<String>, + pub query: bool, +} + +impl EmbeddingsData { + pub fn new(texts: Vec<String>, query: bool) -> Self { + Self { texts, query } + } +} + +pub type EmbeddingsOutput = Vec<Vec<f32>>; + 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<Option<(St "name": name, "api_base": api_base, }); - let prompts = if ALL_CLIENT_MODELS.iter().any(|v| &v.platform == name) { + let prompts = if ALL_MODELS.iter().any(|v| &v.platform == name) { vec![("api_key", "API Key:", false, PromptKind::String)] } else { vec![ diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 097ee68..77f4741 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,8 +1,5 @@ -use super::{ - access_token::*, maybe_catch_error, patch_system_message, sse_stream, ChatCompletionsData, - ChatCompletionsOutput, Client, ErnieClient, ExtraConfig, Model, ModelData, ModelPatches, - PromptAction, PromptKind, SseHandler, SseMmessage, -}; +use super::*; +use super::access_token::*; use anyhow::{anyhow, Context, Result}; use async_trait::async_trait; @@ -37,7 +34,7 @@ impl ErnieClient { data: ChatCompletionsData, ) -> Result<RequestBuilder> { 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<RequestBuilder> { + 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<EmbeddingsOutput> { + 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<f32>, +} 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<isize> { 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<usize>, + pub input_price: Option<f64>, + pub output_price: Option<f64>, + + // chat-only properties pub max_output_tokens: Option<isize>, #[serde(default)] pub pass_max_tokens: bool, - pub input_price: Option<f64>, - pub output_price: Option<f64>, #[serde(default)] pub supports_vision: bool, #[serde(default)] pub supports_function_calling: bool, + + // embedding-only properties + pub default_chunk_size: Option<usize>, + pub max_concurrent_chunks: Option<usize>, } impl ModelData { @@ -222,3 +252,7 @@ pub struct BuiltinModels { pub platform: String, pub models: Vec<ModelData>, } + +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<String>, pub api_base: Option<String>, pub api_auth: Option<String>, - pub chat_endpoint: Option<String>, + #[serde(default)] pub models: Vec<ModelData>, pub patches: Option<ModelPatches>, pub extra: Option<ExtraConfig>, @@ -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<RequestBuilder> { + 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<ChatCompletionsOutput> { let res = builder.send().await?; @@ -109,6 +133,25 @@ async fn chat_completions_streaming( Ok(()) } +async fn embeddings( + builder: RequestBuilder, +) -> Result<EmbeddingsOutput> { + 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<f32>, +} + fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> { 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<RequestBuilder> { + 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<ChatCompletionsOutput> { @@ -125,6 +140,30 @@ pub async fn openai_chat_completions_streaming( sse_stream(builder, handle).await } +pub async fn openai_embeddings( + builder: RequestBuilder, +) -> Result<EmbeddingsOutput> { + 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<EmbeddingsResBodyEmbedding>, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyEmbedding { + embedding: Vec<f32>, +} + 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<ChatCompletionsOutput> { let text = data["choices"][0]["message"]["content"] .as_str() @@ -244,5 +292,6 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOu impl_client_trait!( OpenAIClient, openai_chat_completions, - openai_chat_completions_streaming + openai_chat_completions_streaming, + openai_embeddings ); diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index 74cd954..af7cd0e 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -1,7 +1,5 @@ -use super::{ - openai::*, ChatCompletionsData, Client, ExtraConfig, Model, ModelData, ModelPatches, - OpenAICompatibleClient, PromptAction, PromptKind, OPENAI_COMPATIBLE_PLATFORMS, -}; +use super::*; +use super::openai::*; use anyhow::Result; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -41,27 +39,11 @@ impl OpenAICompatibleClient { client: &ReqwestClient, data: ChatCompletionsData, ) -> Result<RequestBuilder> { - 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<RequestBuilder> { + 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<String> { + 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<String>, @@ -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<RequestBuilder> { + 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<Vec<Vec<f32>>> { + let builder = self.embeddings_builder(client, data)?; + embeddings(builder).await + } } async fn chat_completions(builder: RequestBuilder, model: &Model) -> Result<ChatCompletionsOutput> { @@ -210,6 +249,31 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu Ok((body, has_upload)) } +async fn embeddings( + builder: RequestBuilder, +) -> Result<EmbeddingsOutput> { + 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<EmbeddingsResBodyOutputEmbedding>, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyOutputEmbedding { + embedding: Vec<f32>, +} + fn extract_chat_completions_text(data: &Value, model: &Model) -> Result<ChatCompletionsOutput> { 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<RequestBuilder> { 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<RequestBuilder> { + 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<Vec<Vec<f32>>> { + 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<ChatCompletionsOutput> { @@ -138,6 +173,34 @@ pub async fn gemini_chat_completions_streaming( Ok(()) } +async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> { + 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<EmbeddingsResBodyPrediction>, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyPrediction { + embeddings: EmbeddingsResBodyPredictionEmbeddings, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyPredictionEmbeddings { + values: Vec<f32> +} + fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsOutput> { let text = data["candidates"][0]["content"]["parts"][0]["text"] .as_str() @@ -179,7 +242,7 @@ fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsO Ok(output) } -pub(crate) fn gemini_build_chat_completions_body( +pub fn gemini_build_chat_completions_body( data: ChatCompletionsData, model: &Model, ) -> Result<Value> { 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<String>, medias: Vec<String>, data_urls: HashMap<String, String>, tool_call: Option<ToolResults>, + rag: Option<String>, 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<String> = 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<MessageContentPart> = 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> +__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<String>, pub buffer_editor: Option<String>, + pub embedding_model: Option<String>, + pub rag_top_k: usize, + pub rag_template: Option<String>, pub function_calling: bool, pub compress_threshold: usize, pub summarize_prompt: Option<String>, @@ -85,6 +99,8 @@ pub struct Config { #[serde(skip)] pub session: Option<Session>, #[serde(skip)] + pub rag: Option<Arc<Rag>>, + #[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<PathBuf> { + Self::local_path(RAGS_DIR_NAME) + } + pub fn functions_dir() -> Result<PathBuf> { 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<PathBuf> { + 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<String> { + if let Some(rag) = &self.rag { + rag.export() + } else { + bail!("No rag") + } + } + pub fn info(&self) -> Result<String> { 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<String> { + 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<RenderOptions> { 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<Vec<RagDocument>> { + 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<Vec<RagDocument>> { + let contents = fs::read_to_string(path).await?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pdf(path: &str) -> Result<Vec<RagDocument>> { + let contents = pdf_extract::extract_text(path)?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pandoc(path: &str) -> Result<Vec<RagDocument>> { + 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<String>)> { + 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::<Vec<String>>() + } 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<String>, + entry_path: &Path, + suffixes: Option<&Vec<String>>, +) -> 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<String>, suffixes: Option<&Vec<String>>, path: &Path) { + if is_valid_extension(suffixes, path) { + files.push(path.display().to_string()); + } +} + +fn is_valid_extension(suffixes: Option<&Vec<String>>, 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<dyn Client>, + 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<Self> { + 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<Self> { + 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<Self> { + 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<String> { + 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<String> { + 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<T: AsRef<Path>>( + &mut self, + paths: &[T], + progress_tx: Option<mpsc::UnboundedSender<String>>, + ) -> 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<Vec<String>> { + 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::<Vec<_>>() + }) + .collect(); + Ok(output) + } + + async fn create_embeddings( + &self, + data: EmbeddingsData, + progress_tx: Option<mpsc::UnboundedSender<String>>, + ) -> Result<EmbeddingsOutput> { + 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<RagFile>, + pub vectors: IndexMap<VectorID, Vec<f32>>, +} + +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<RagFile>, + vector_ids: Vec<VectorID>, + 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<RagDocument>, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagDocument { + pub page_content: String, + pub metadata: RagMetadata, +} + +impl RagDocument { + pub fn new<S: Into<String>>(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<String, String>; + +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<Model> { + 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<Model> { + 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<usize> { + let value = Text::new("Set chunk size:") + .with_default(&chunk_size.to_string()) + .with_validator(move |text: &str| { + let out = match text.parse::<usize>() { + 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<Vec<String>> { + 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<mpsc::UnboundedSender<String>>, 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 + "<body>", "<div>", "<p>", "<br>", "<li>", "<h1>", "<h2>", "<h3>", "<h4>", "<h5>", "<h6>", + "<span>", "<table>", "<tr>", "<td>", "<th>", "<ul>", "<ol>", "<header>", "<footer>", "<nav>", + // Head + "<head>", "<style>", "<script>", "<meta>", "<title>", // Normal type of lines + " ", "", +]; +pub const MARKDOWN_SEPARATES: [&str; 13] = [ + // First, try to split along Markdown headings (starting with level 2) + "\n## ", + "\n### ", + "\n#### ", + "\n##### ", + "\n###### ", + // Note the alternative syntax for headings (below) is not handled here + // Heading level 2 + // --------------- + // End of code block + "```\n\n", + // Horizontal lines + "\n\n***\n\n", + "\n\n---\n\n", + "\n\n___\n\n", + // Note that this splitter doesn't handle horizontal lines defined + // by *three or more* of ***, ---, or ___, but this is not handled + "\n\n", + "\n", + " ", + "", +]; +pub const LATEX_SEPARATES: [&str; 19] = [ + // First, try to split along Latex sections + "\n\\chapter{", + "\n\\section{", + "\n\\subsection{", + "\n\\subsubsection{", + // Now split by environments + "\n\\begin{enumerate}", + "\n\\begin{itemize}", + "\n\\begin{description}", + "\n\\begin{list}", + "\n\\begin{quote}", + "\n\\begin{quotation}", + "\n\\begin{verse}", + "\n\\begin{verbatim}", + // Now split by math environments + "\n\\begin{align}", + "$$", + "$", + // Now split by the normal type of lines + "\n\n", + "\n", + " ", + "", +]; + +pub fn autodetect_separator(extension: &str) -> &[&'static str] { + match extension { + "md" | "mkd" => &MARKDOWN_SEPARATES, + "htm" | "html" => &HTML_SEPARATES, + "tex" => &LATEX_SEPARATES, + _ => &DEFAULT_SEPARATES, + } +} + +pub struct Splitter { + pub chunk_size: usize, + pub chunk_overlap: usize, + pub separators: Vec<String>, + pub length_function: Box<dyn Fn(&str) -> usize + Send + Sync>, +} + +impl Default for Splitter { + fn default() -> Self { + Self { + chunk_size: 1000, + chunk_overlap: 20, + separators: DEFAULT_SEPARATES.iter().map(|v| v.to_string()).collect(), + length_function: Box::new(|text| text.len()), + } + } +} + +// Builder pattern for Options struct +impl Splitter { + pub fn new(chunk_size: usize, chunk_overlap: usize, separators: &[&str]) -> Self { + Self::default() + .with_chunk_size(chunk_size) + .with_chunk_overlap(chunk_overlap) + .with_separators(separators) + } + + pub fn with_chunk_size(mut self, chunk_size: usize) -> Self { + self.chunk_size = chunk_size; + self + } + + pub fn with_chunk_overlap(mut self, chunk_overlap: usize) -> Self { + self.chunk_overlap = chunk_overlap; + self + } + + pub fn with_separators(mut self, separators: &[&str]) -> Self { + self.separators = separators.iter().map(|v| v.to_string()).collect(); + self + } + + #[allow(unused)] + pub fn with_length_function<F>(mut self, length_function: F) -> Self + where + F: Fn(&str) -> usize + Send + Sync + 'static, + { + self.length_function = Box::new(length_function); + self + } + + pub fn split_documents( + &self, + documents: &[RagDocument], + chunk_header_options: &SplitterChunkHeaderOptions, + ) -> Vec<RagDocument> { + let mut texts: Vec<String> = Vec::new(); + let mut metadatas: Vec<RagMetadata> = Vec::new(); + documents.iter().for_each(|d| { + if !d.page_content.is_empty() { + texts.push(d.page_content.clone()); + metadatas.push(d.metadata.clone()); + } + }); + + self.create_documents(&texts, &metadatas, chunk_header_options) + } + + pub fn create_documents( + &self, + texts: &[String], + metadatas: &[RagMetadata], + chunk_header_options: &SplitterChunkHeaderOptions, + ) -> Vec<RagDocument> { + let SplitterChunkHeaderOptions { + chunk_header, + chunk_overlap_header, + append_chunk_overlap_header, + } = chunk_header_options; + + let mut documents = Vec::new(); + for (i, text) in texts.iter().enumerate() { + let mut line_counter_index = 1; + let mut prev_chunk = None; + let mut index_prev_chunk = None; + + for chunk in self.split_text(text) { + let mut page_content = chunk_header.clone(); + + let index_chunk = { + let idx = match index_prev_chunk { + Some(v) => v + 1, + None => 0, + }; + text[idx..].find(&chunk).map(|i| i + idx).unwrap_or(0) + }; + if prev_chunk.is_none() { + line_counter_index += self.number_of_newlines(text, 0, index_chunk); + } else { + let index_end_prev_chunk: usize = index_prev_chunk.unwrap_or_default() + + (self.length_function)(prev_chunk.as_deref().unwrap_or_default()); + + match index_end_prev_chunk.cmp(&index_chunk) { + Ordering::Less => { + line_counter_index += + self.number_of_newlines(text, index_end_prev_chunk, index_chunk); + } + Ordering::Greater => { + let number = + self.number_of_newlines(text, index_chunk, index_end_prev_chunk); + line_counter_index = line_counter_index.saturating_sub(number); + } + Ordering::Equal => {} + } + + if *append_chunk_overlap_header { + page_content += chunk_overlap_header; + } + } + + let newlines_count = self.number_of_newlines(&chunk, 0, chunk.len()); + + let mut metadata = metadatas[i].clone(); + metadata.insert( + "loc".into(), + format!( + "{}:{}", + line_counter_index, + line_counter_index + newlines_count + ), + ); + page_content += &chunk; + documents.push(RagDocument { + page_content, + metadata, + }); + + line_counter_index += newlines_count; + prev_chunk = Some(chunk); + index_prev_chunk = Some(index_chunk); + } + } + + documents + } + + fn number_of_newlines(&self, text: &str, start: usize, end: usize) -> usize { + text[start..end].matches('\n').count() + } + + pub fn split_text(&self, text: &str) -> Vec<String> { + let keep_separator = self + .separators + .iter() + .any(|v| v.chars().any(|v| !v.is_whitespace())); + self.split_text_impl(text, &self.separators, keep_separator) + } + + fn split_text_impl( + &self, + text: &str, + separators: &[String], + keep_separator: bool, + ) -> Vec<String> { + let mut final_chunks = Vec::new(); + + let mut separator: String = separators.last().cloned().unwrap_or_default(); + let mut new_separators: Vec<String> = vec![]; + for (i, s) in separators.iter().enumerate() { + if s.is_empty() { + separator.clone_from(s); + break; + } + if text.contains(s) { + separator.clone_from(s); + new_separators = separators[i + 1..].to_vec(); + break; + } + } + + // Now that we have the separator, split the text + let splits = split_on_separator(text, &separator, keep_separator); + + // Now go merging things, recursively splitting longer texts. + let mut good_splits = Vec::new(); + let _separator = if keep_separator { "" } else { &separator }; + for s in splits { + if (self.length_function)(s) < self.chunk_size { + good_splits.push(s.to_string()); + } else { + if !good_splits.is_empty() { + let merged_text = self.merge_splits(&good_splits, _separator); + final_chunks.extend(merged_text); + good_splits.clear(); + } + if new_separators.is_empty() { + final_chunks.push(s.to_string()); + } else { + let other_info = self.split_text_impl(s, &new_separators, keep_separator); + final_chunks.extend(other_info); + } + } + } + if !good_splits.is_empty() { + let merged_text = self.merge_splits(&good_splits, _separator); + final_chunks.extend(merged_text); + } + final_chunks + } + + fn merge_splits(&self, splits: &[String], separator: &str) -> Vec<String> { + let mut docs = Vec::new(); + let mut current_doc = Vec::new(); + let mut total = 0; + for d in splits { + let _len = (self.length_function)(d); + if total + _len + current_doc.len() * separator.len() > self.chunk_size { + if total > self.chunk_size { + // warn!("Warning: Created a chunk of size {}, which is longer than the specified {}", total, self.chunk_size); + } + if !current_doc.is_empty() { + let doc = self.join_docs(¤t_doc, separator); + if let Some(doc) = doc { + docs.push(doc); + } + // Keep on popping if: + // - we have a larger chunk than in the chunk overlap + // - or if we still have any chunks and the length is long + while total > self.chunk_overlap + || (total + _len + current_doc.len() * separator.len() > self.chunk_size + && total > 0) + { + total -= (self.length_function)(¤t_doc[0]); + current_doc.remove(0); + } + } + } + current_doc.push(d.to_string()); + total += _len; + } + let doc = self.join_docs(¤t_doc, separator); + if let Some(doc) = doc { + docs.push(doc); + } + docs + } + + fn join_docs(&self, docs: &[String], separator: &str) -> Option<String> { + let text = docs.join(separator).trim().to_string(); + if text.is_empty() { + None + } else { + Some(text) + } + } +} + +pub struct SplitterChunkHeaderOptions { + pub chunk_header: String, + pub chunk_overlap_header: String, + pub append_chunk_overlap_header: bool, +} + +impl Default for SplitterChunkHeaderOptions { + fn default() -> Self { + Self { + chunk_header: "".into(), + chunk_overlap_header: "(cont'd) ".into(), + append_chunk_overlap_header: false, + } + } +} + +impl SplitterChunkHeaderOptions { + // Set the value of chunk_header + #[allow(unused)] + pub fn with_chunk_header(mut self, header: &str) -> Self { + self.chunk_header = header.to_string(); + self + } + + // Set the value of chunk_overlap_header + #[allow(unused)] + pub fn with_chunk_overlap_header(mut self, overlap_header: &str) -> Self { + self.chunk_overlap_header = overlap_header.to_string(); + self + } + + // Set the value of append_chunk_overlap_header + #[allow(unused)] + pub fn with_append_chunk_overlap_header(mut self, value: bool) -> Self { + self.append_chunk_overlap_header = value; + self + } +} + +fn split_on_separator<'a>(text: &'a str, separator: &str, keep_separator: bool) -> Vec<&'a str> { + let splits: Vec<&str> = if !separator.is_empty() { + if keep_separator { + let mut splits = Vec::new(); + let mut prev_idx = 0; + let sep_len = separator.len(); + + while let Some(idx) = text[prev_idx..].find(separator) { + splits.push(&text[prev_idx.saturating_sub(sep_len)..prev_idx + idx]); + prev_idx += idx + sep_len; + } + + if prev_idx < text.len() { + splits.push(&text[prev_idx.saturating_sub(sep_len)..]); + } + + splits + } else { + text.split(separator).collect() + } + } else { + text.split("").collect() + }; + splits.into_iter().filter(|s| !s.is_empty()).collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use indexmap::IndexMap; + use pretty_assertions::assert_eq; + use serde_json::{json, Value}; + + fn build_metadata(source: &str, loc_from_line: usize, loc_to_line: usize) -> Value { + json!({ + "source": source, + "loc": format!("{loc_from_line}:{loc_to_line}"), + }) + } + #[test] + fn test_split_text() { + let splitter = Splitter { + chunk_size: 7, + chunk_overlap: 3, + separators: vec![" ".into()], + ..Default::default() + }; + let output = splitter.split_text("foo bar baz 123"); + assert_eq!(output, vec!["foo bar", "bar baz", "baz 123"]); + } + + #[test] + fn test_create_document() { + let splitter = Splitter::new(3, 0, &[" "]); + let chunk_header_options = SplitterChunkHeaderOptions::default(); + let mut metadata1 = IndexMap::new(); + metadata1.insert("source".into(), "1".into()); + let mut metadata2 = IndexMap::new(); + metadata2.insert("source".into(), "2".into()); + let output = splitter.create_documents( + &["foo bar".into(), "baz".into()], + &[metadata1, metadata2], + &chunk_header_options, + ); + let output = json!(output); + assert_eq!( + output, + json!([ + { + "page_content": "foo", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "bar", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "baz", + "metadata": build_metadata("2", 1, 1), + }, + ]) + ); + } + + #[test] + fn test_chunk_header() { + let splitter = Splitter::new(3, 0, &[" "]); + let chunk_header_options = SplitterChunkHeaderOptions::default() + .with_chunk_header("SOURCE NAME: testing\n-----\n") + .with_append_chunk_overlap_header(true); + let mut metadata1 = IndexMap::new(); + metadata1.insert("source".into(), "1".into()); + let mut metadata2 = IndexMap::new(); + metadata2.insert("source".into(), "2".into()); + let output = splitter.create_documents( + &["foo bar".into(), "baz".into()], + &[metadata1, metadata2], + &chunk_header_options, + ); + let output = json!(output); + assert_eq!( + output, + json!([ + { + "page_content": "SOURCE NAME: testing\n-----\nfoo", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "SOURCE NAME: testing\n-----\n(cont'd) bar", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "SOURCE NAME: testing\n-----\nbaz", + "metadata": build_metadata("2", 1, 1), + }, + ]) + ); + } + + #[test] + fn test_markdown_splitter() { + let text = r#"# π¦οΈπ LangChain + +β‘ Building applications with LLMs through composability β‘ + +## Quick Install + +```bash +# Hopefully this code block isn't split +pip install langchain +``` + +As an open source project in a rapidly developing field, we are extremely open to contributions."#; + let splitter = Splitter::new(100, 0, &MARKDOWN_SEPARATES); + let output = splitter.split_text(text); + let expected_output = vec![ + "# π¦οΈπ LangChain\n\nβ‘ Building applications with LLMs through composability β‘", + "## Quick Install\n\n```bash\n# Hopefully this code block isn't split\npip install langchain", + "```", + "As an open source project in a rapidly developing field, we are extremely open to contributions.", + ]; + assert_eq!(output, expected_output); + } + + #[test] + fn test_html_splitter() { + let text = r#"<!DOCTYPE html> +<html> + <head> + <title>π¦οΈπ LangChain</title> + <style> + body { + font-family: Arial, sans-serif; + } + h1 { + color: darkblue; + } + </style> + </head> + <body> + <div> + <h1>π¦οΈπ LangChain</h1> + <p>β‘ Building applications with LLMs through composability β‘</p> + </div> + <div> + As an open source project in a rapidly developing field, we are extremely open to contributions. + </div> + </body> +</html>"#; + let splitter = Splitter::new(175, 20, &HTML_SEPARATES); + let output = splitter.split_text(text); + let expected_output = vec![ + "<!DOCTYPE html>\n<html>", + "<head>\n <title>π¦οΈπ LangChain</title>", + r#"<style> + body { + font-family: Arial, sans-serif; + } + h1 { + color: darkblue; + } + </style> + </head>"#, + r#"<body> + <div> + <h1>π¦οΈπ LangChain</h1> + <p>β‘ Building applications with LLMs through composability β‘</p> + </div>"#, + r#"<div> + As an open source project in a rapidly developing field, we are extremely open to contributions. + </div> + </body> +</html>"#, + ]; + assert_eq!(output, expected_output); + } +} diff --git a/src/render/stream.rs b/src/render/stream.rs index f35831c..0c70bce 100644 --- a/src/render/stream.rs +++ b/src/render/stream.rs @@ -14,7 +14,7 @@ use std::{ time::Duration, }; use textwrap::core::display_width; -use tokio::sync::{mpsc::UnboundedReceiver, oneshot}; +use tokio::sync::mpsc::UnboundedReceiver; pub async fn markdown_stream( rx: UnboundedReceiver<SseEvent>, @@ -62,17 +62,16 @@ async fn markdown_stream_inner( let columns = terminal::size()?.0; - let (spinner_tx, spinner_rx) = oneshot::channel(); - let mut spinner_tx = Some(spinner_tx); - tokio::spawn(run_spinner(" Generating", spinner_rx)); + let (stop_spinner_tx, _) = run_spinner("Generating").await; + let mut stop_spinner_tx = Some(stop_spinner_tx); 'outer: loop { if abort.aborted() { return Ok(()); } for reply_event in gather_events(&mut rx).await { - if let Some(spinner_tx) = spinner_tx.take() { - let _ = spinner_tx.send(()); + if let Some(stop_spinner_tx) = stop_spinner_tx.take() { + let _ = stop_spinner_tx.send(()); } match reply_event { @@ -150,8 +149,8 @@ async fn markdown_stream_inner( } } - if let Some(spinner_tx) = spinner_tx.take() { - let _ = spinner_tx.send(()); + if let Some(stop_spinner_tx) = stop_spinner_tx.take() { + let _ = stop_spinner_tx.send(()); } Ok(()) } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index ba506dd..291bcd2 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -7,7 +7,7 @@ use self::highlighter::ReplHighlighter; use self::prompt::ReplPrompt; use crate::client::send_stream; -use crate::config::{AssertState, GlobalConfig, Input, InputContext, StateFlags}; +use crate::config::{AssertState, Config, GlobalConfig, Input, InputContext, StateFlags}; use crate::function::need_send_call_results; use crate::render::render_error; use crate::utils::{create_abort_signal, set_text, AbortSignal}; @@ -33,7 +33,7 @@ lazy_static! { const MENU_NAME: &str = "completion_menu"; lazy_static! { - static ref REPL_COMMANDS: [ReplCommand; 16] = [ + static ref REPL_COMMANDS: [ReplCommand; 19] = [ ReplCommand::new(".help", "Show this help message", AssertState::any()), ReplCommand::new(".info", "View system info", AssertState::any()), ReplCommand::new(".model", "Change the current LLM", AssertState::any()), @@ -82,6 +82,17 @@ lazy_static! { "End the current session", AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION) ), + ReplCommand::new(".rag", "Init or use a rag", AssertState::any()), + ReplCommand::new( + ".info rag", + "View rag info", + AssertState::True(StateFlags::RAG), + ), + ReplCommand::new( + ".exit rag", + "Leave the rag", + AssertState::True(StateFlags::RAG) + ), ReplCommand::new( ".file", "Include files with the message", @@ -99,7 +110,7 @@ pub struct Repl { config: GlobalConfig, editor: Reedline, prompt: ReplPrompt, - abort: AbortSignal, + abort_signal: AbortSignal, } impl Repl { @@ -114,7 +125,7 @@ impl Repl { config: config.clone(), editor, prompt, - abort, + abort_signal: abort, }) } @@ -122,13 +133,13 @@ impl Repl { self.banner(); loop { - if self.abort.aborted_ctrld() { + if self.abort_signal.aborted_ctrld() { break; } let sig = self.editor.read_line(&self.prompt); match sig { Ok(Signal::Success(line)) => { - self.abort.reset(); + self.abort_signal.reset(); match self.handle(&line).await { Ok(exit) => { if exit { @@ -142,11 +153,11 @@ impl Repl { } } Ok(Signal::CtrlC) => { - self.abort.set_ctrlc(); + self.abort_signal.set_ctrlc(); println!("(To exit, press Ctrl+D or enter \".exit\")\n"); } Ok(Signal::CtrlD) => { - self.abort.set_ctrld(); + self.abort_signal.set_ctrld(); break; } _ => {} @@ -176,6 +187,10 @@ impl Repl { let info = self.config.read().session_info()?; println!("{}", info); } + Some("rag") => { + let info = self.config.read().rag_info()?; + println!("{}", info); + } Some(_) => unknown_command()?, None => { let output = self.config.read().system_info()?; @@ -193,7 +208,7 @@ impl Repl { }, ".prompt" => match args { Some(text) => { - self.config.write().set_prompt(text)?; + self.config.write().use_prompt(text)?; } None => println!("Usage: .prompt <text>..."), }, @@ -206,16 +221,19 @@ impl Repl { text.trim(), Some(InputContext::role(role)), ); - ask(&self.config, self.abort.clone(), input).await?; + ask(&self.config, self.abort_signal.clone(), input).await?; } None => { - self.config.write().set_role(args)?; + self.config.write().use_role(args)?; } }, None => println!(r#"Usage: .role <name> [text]..."#), }, ".session" => { - self.config.write().start_session(args)?; + self.config.write().use_session(args)?; + } + ".rag" => { + Config::use_rag(&self.config, args, self.abort_signal.clone()).await?; } ".save" => { match args.map(|v| match v.split_once(' ') { @@ -248,16 +266,19 @@ impl Repl { let (files, text) = split_files_text(args); let files = shell_words::split(files).with_context(|| "Invalid args")?; let input = Input::new(&self.config, text, files, None)?; - ask(&self.config, self.abort.clone(), input).await?; + ask(&self.config, self.abort_signal.clone(), input).await?; } None => println!("Usage: .file <files>... [-- <text>...]"), }, ".exit" => match args { Some("role") => { - self.config.write().clear_role()?; + self.config.write().exit_role()?; } Some("session") => { - self.config.write().end_session()?; + self.config.write().exit_session()?; + } + Some("rag") => { + self.config.write().exit_rag()?; } Some(_) => unknown_command()?, None => { @@ -273,8 +294,9 @@ impl Repl { _ => unknown_command()?, }, None => { - let input = Input::from_str(&self.config, line, None); - ask(&self.config, self.abort.clone(), input).await?; + let mut input = Input::from_str(&self.config, line, None); + input.maybe_embeddings(self.abort_signal.clone()).await?; + ask(&self.config, self.abort_signal.clone(), input).await?; } } @@ -407,7 +429,7 @@ impl Validator for ReplValidator { } #[async_recursion] -async fn ask(config: &GlobalConfig, abort: AbortSignal, input: Input) -> Result<()> { +async fn ask(config: &GlobalConfig, abort: AbortSignal, mut input: Input) -> Result<()> { if input.is_empty() { return Ok(()); } @@ -417,9 +439,10 @@ async fn ask(config: &GlobalConfig, abort: AbortSignal, input: Input) -> Result< let client = input.create_client()?; let (output, tool_call_results) = send_stream(&input, client.as_ref(), config, abort.clone()).await?; + config .write() - .save_message(&input, &output, &tool_call_results)?; + .save_message(&mut input, &output, &tool_call_results)?; config.read().maybe_copy(&output); if config.write().should_compress_session() { let config = config.clone(); diff --git a/src/serve.rs b/src/serve.rs index 2d375f8..116942a 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -1,11 +1,4 @@ -use crate::{ - client::{ - init_client, list_models, ChatCompletionsData, ChatCompletionsOutput, ClientConfig, - Message, Model, ModelData, SseEvent, SseHandler, - }, - config::{Config, GlobalConfig, Role}, - utils::create_abort_signal, -}; +use crate::{client::*, config::*, utils::*}; use anyhow::{anyhow, bail, Result}; use bytes::Bytes; @@ -76,7 +69,7 @@ impl Server { let clients = config.clients.clone(); let model = config.model.clone(); let roles = config.roles.clone(); - let mut models = list_models(&config); + let mut models = list_chat_models(&config); let mut default_model = model.clone(); default_model.data_mut().name = DEFAULT_MODEL_NAME.into(); models.insert(0, &default_model); diff --git a/src/utils/abort_signal.rs b/src/utils/abort_signal.rs index af58b35..ac93653 100644 --- a/src/utils/abort_signal.rs +++ b/src/utils/abort_signal.rs @@ -53,3 +53,12 @@ impl AbortSignalInner { self.ctrld.store(true, Ordering::SeqCst); } } + +pub async fn watch_abort_signal(abort: AbortSignal) { + loop { + if abort.aborted() { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + } +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index fa67c63..95c6725 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -6,7 +6,7 @@ mod prompt_input; mod render_prompt; mod spinner; -pub use self::abort_signal::{create_abort_signal, AbortSignal}; +pub use self::abort_signal::*; pub use self::clipboard::set_text; pub use self::command::*; pub use self::crypto::*; diff --git a/src/utils/spinner.rs b/src/utils/spinner.rs index 0dd4d01..c746469 100644 --- a/src/utils/spinner.rs +++ b/src/utils/spinner.rs @@ -4,7 +4,10 @@ use std::{ io::{stdout, Stdout, Write}, time::Duration, }; -use tokio::{sync::oneshot, time::interval}; +use tokio::{ + sync::{mpsc, oneshot}, + time::interval, +}; pub struct Spinner { index: usize, @@ -23,6 +26,10 @@ impl Spinner { } } + pub fn set_message(&mut self, message: &str) { + self.message = format!(" {message}"); + } + pub fn step(&mut self, writer: &mut Stdout) -> Result<()> { if self.stopped { return Ok(()); @@ -55,18 +62,38 @@ impl Spinner { } } -pub async fn run_spinner(message: &str, rx: oneshot::Receiver<()>) -> Result<()> { +pub async fn run_spinner(message: &str) -> (oneshot::Sender<()>, mpsc::UnboundedSender<String>) { + let message = format!(" {message}"); + let (stop_tx, stop_rx) = oneshot::channel(); + let (message_tx, message_rx) = mpsc::unbounded_channel(); + tokio::spawn(run_spinner_inner(message, stop_rx, message_rx)); + (stop_tx, message_tx) +} + +async fn run_spinner_inner( + message: String, + stop_rx: oneshot::Receiver<()>, + mut message_rx: mpsc::UnboundedReceiver<String>, +) -> Result<()> { let mut writer = stdout(); - let mut spinner = Spinner::new(message); + let mut spinner = Spinner::new(&message); let mut interval = interval(Duration::from_millis(50)); tokio::select! { _ = async { loop { - interval.tick().await; - let _ = spinner.step(&mut writer); + tokio::select! { + _ = interval.tick() => { + let _ = spinner.step(&mut writer); + } + message = message_rx.recv() => { + if let Some(message) = message { + spinner.set_message(&message); + } + } + } } } => {} - _ = rx => { + _ = stop_rx => { spinner.stop(&mut writer)?; } } |
