diff options
| -rw-r--r-- | assets/playground.html | 41 | ||||
| -rw-r--r-- | src/config/input.rs | 33 | ||||
| -rw-r--r-- | src/config/mod.rs | 45 | ||||
| -rw-r--r-- | src/main.rs | 2 | ||||
| -rw-r--r-- | src/serve.rs | 92 |
5 files changed, 152 insertions, 61 deletions
diff --git a/assets/playground.html b/assets/playground.html index da2e32c..8ae0ffa 100644 --- a/assets/playground.html +++ b/assets/playground.html @@ -502,6 +502,15 @@ </div> <div class="settings"> <div class="control"> + <label for="role">RAG</label> + <select id="role" x-model="settings.rag" :disabled="sessionMode"> + <template x-for="rag in rags"> + <option :value="rag" :selected="rag == settings.rag" x-text="rag"></option> + </template> + </select> + </div> + + <div class="control"> <label for="role">Role</label> <select id="role" x-model="settings.role" :disabled="sessionMode"> <template x-for="role in roles"> @@ -700,6 +709,8 @@ const CHAT_COMPLETIONS_URL = API_BASE + "/chat/completions"; const MODELS_API = API_BASE + "/models"; const ROLES_API = API_BASE + "/roles"; + const RAGS_API = API_BASE + "/rags"; + const SEARCH_RAG_API = API_BASE + "/rags/search"; document.addEventListener("alpine:init", () => { setupMarked(); @@ -710,6 +721,7 @@ let msgIdx = 0; let defaultSettings = { model: QUERY.model || "default", + rag: QUERY.rag || "", role: QUERY.role || "", prompt: "", max_output_tokens: parseInt(QUERY.max_output_tokens) || null, @@ -719,6 +731,7 @@ Alpine.data("app", () => ({ models: [], + rags: [""], roles: [{ name: "", prompt: "" }], modelData: {}, messages: [], @@ -741,6 +754,9 @@ toast("No model available"); console.error("Failed to load models", err); }), + fetchJSON(RAGS_API).then(rags => { + this.rags.push(...rags); + }).catch(() => { }), fetchJSON(ROLES_API).then(roles => { this.roles.push(...roles.filter(v => !!v.prompt)); }).catch(() => { }), @@ -753,6 +769,9 @@ } else { this.settings.model = "default"; } + if (!this.rags.find(rag => rag === this.settings.rag)) { + this.settings.rag = ""; + } this.$watch("settings.role", () => this.handleRoleChange()) if (this.roles.find(role => role.name === this.settings.role)) { this.handleRoleChange(); @@ -922,7 +941,7 @@ updateUrl() { const newUrl = new URL(location.href); - ["model", "role", "max_output_tokens", "temperature", "top_p"].forEach(key => { + ["model", "rag", "role", "max_output_tokens", "temperature", "top_p"].forEach(key => { if (this.settings[key] || typeof this.settings[key] === "number") { newUrl.searchParams.set(key, this.settings[key]); } else { @@ -957,6 +976,12 @@ const body = this.buildBody(); let succeed = false; try { + if (this.settings.rag) { + const message = body.messages[body.messages.length - 1]; + if (message.role === "user" && typeof message.content === "string") { + message.content = await this.searchRag(this.settings.rag, message.content); + } + } const stream = await fetchChatCompletions(CHAT_COMPLETIONS_URL, body, this.askAbortController.signal) for await (const chunk of stream) { lastMessage.state = "streaming"; @@ -983,6 +1008,20 @@ this.asking = false; }, + async searchRag(name, input) { + const res = await fetch(SEARCH_RAG_API, { + method: "POST", + headers: getHeaders(), + signal: this.askAbortController.signal, + body: JSON.stringify({ + name, + input + }) + }); + const data = await res.json(); + return data.data; + }, + buildBody() { let messages = []; for ([userMessage, assistantMessage] of chunkArray(this.messages, 2)) { diff --git a/src/config/input.rs b/src/config/input.rs index 9c7a666..829d9ed 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -145,36 +145,9 @@ impl Input { if !self.text.is_empty() { let rag = self.config.read().rag.clone(); if let Some(rag) = rag { - let (top_k, min_score_vector_search, min_score_keyword_search) = { - let config = self.config.read(); - ( - config.rag_top_k, - config.rag_min_score_vector_search, - config.rag_min_score_keyword_search, - ) - }; - let rerank = match self.config.read().rag_reranker_model.clone() { - Some(reranker_model_id) => { - let min_score = self.config.read().rag_min_score_rerank; - let rerank_model = - Model::retrieve_reranker(&self.config.read(), &reranker_model_id)?; - let rerank_client = init_client(&self.config, Some(rerank_model))?; - Some((rerank_client, min_score)) - } - None => None, - }; - let embeddings = rag - .search( - &self.text, - top_k, - min_score_vector_search, - min_score_keyword_search, - rerank, - abort_signal, - ) - .await?; - let text = self.config.read().rag_template(&embeddings, &self.text); - self.patched_text = Some(text); + let result = + Config::search_rag(&self.config, &rag, &self.text, abort_signal).await?; + self.patched_text = Some(result); self.rag_name = Some(rag.name().to_string()); } } diff --git a/src/config/mod.rs b/src/config/mod.rs index 1a0b1d6..5d1172b 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -9,8 +9,8 @@ pub use self::role::{Role, RoleLike, BUILTIN_ROLES, CODE_ROLE, EXPLAIN_SHELL_ROL use self::session::Session; use crate::client::{ - create_client_config, list_chat_models, list_client_types, list_reranker_models, ClientConfig, - Model, OPENAI_COMPATIBLE_PLATFORMS, + create_client_config, init_client, list_chat_models, list_client_types, list_reranker_models, + ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{FunctionDeclaration, Functions, ToolResult}; use crate::rag::Rag; @@ -1095,7 +1095,44 @@ impl Config { Ok(()) } - pub fn list_rags(&self) -> Vec<String> { + pub async fn search_rag( + config: &GlobalConfig, + rag: &Rag, + text: &str, + abort_signal: AbortSignal, + ) -> Result<String> { + let (top_k, min_score_vector_search, min_score_keyword_search) = { + let config = config.read(); + ( + config.rag_top_k, + config.rag_min_score_vector_search, + config.rag_min_score_keyword_search, + ) + }; + let rerank = match config.read().rag_reranker_model.clone() { + Some(reranker_model_id) => { + let min_score = config.read().rag_min_score_rerank; + let rerank_model = Model::retrieve_reranker(&config.read(), &reranker_model_id)?; + let rerank_client = init_client(config, Some(rerank_model))?; + Some((rerank_client, min_score)) + } + None => None, + }; + let embeddings = rag + .search( + text, + top_k, + min_score_vector_search, + min_score_keyword_search, + rerank, + abort_signal, + ) + .await?; + let text = config.read().rag_template(&embeddings, text); + Ok(text) + } + + pub fn list_rags() -> Vec<String> { let rags_dir = match Self::rags_dir() { Ok(dir) => dir, Err(_) => return vec![], @@ -1327,7 +1364,7 @@ impl Config { .into_iter() .map(|v| (v, None)) .collect(), - ".rag" => self.list_rags().into_iter().map(|v| (v, None)).collect(), + ".rag" => Self::list_rags().into_iter().map(|v| (v, None)).collect(), ".agent" => list_agents().into_iter().map(|v| (v, None)).collect(), ".starter" => match &self.agent { Some(agent) => agent diff --git a/src/main.rs b/src/main.rs index 83d2536..780678d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -88,7 +88,7 @@ async fn run(config: GlobalConfig, cli: Cli, text: Option<String>) -> Result<()> return Ok(()); } if cli.list_rags { - let rags = config.read().list_rags().join("\n"); + let rags = Config::list_rags().join("\n"); println!("{rags}"); return Ok(()); } diff --git a/src/serve.rs b/src/serve.rs index 65ac4cf..95b1402 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -1,4 +1,4 @@ -use crate::{client::*, config::*, utils::*}; +use crate::{client::*, config::*, function::*, rag::*, utils::*}; use anyhow::{anyhow, bail, Result}; use bytes::Bytes; @@ -65,20 +65,18 @@ pub async fn run(config: GlobalConfig, addr: Option<String>) -> Result<()> { } struct Server { - clients: Vec<ClientConfig>, - model: Model, + config: Config, models: Vec<Value>, roles: Vec<Role>, + rags: Vec<String>, } impl Server { fn new(config: &GlobalConfig) -> Self { - let config = config.read(); - let clients = config.clients.clone(); - let model = config.model.clone(); - let roles = Config::all_roles(); + let mut config = config.read().clone(); + config.functions = Functions::default(); let mut models = list_models(&config); - let mut default_model = model.clone(); + let mut default_model = config.model.clone(); default_model.data_mut().name = DEFAULT_MODEL_NAME.into(); models.insert(0, &default_model); let models: Vec<Value> = models @@ -101,12 +99,13 @@ impl Server { }) .collect(); Self { - clients, - model, - roles, + config, models, + roles: Config::all_roles(), + rags: Config::list_rags(), } } + async fn run(self: Arc<Self>, listener: TcpListener) -> Result<oneshot::Sender<()>> { let (tx, rx) = oneshot::channel(); tokio::spawn(async move { @@ -164,6 +163,10 @@ impl Server { self.list_models() } else if path == "/v1/roles" { self.list_roles() + } else if path == "/v1/rags" { + self.list_rags() + } else if path == "/v1/rags/search" { + self.search_rag(req).await } else if path == "/playground" || path == "/playground.html" { self.playground_page() } else if path == "/arena" || path == "/arena.html" { @@ -220,6 +223,43 @@ impl Server { Ok(res) } + fn list_rags(&self) -> Result<AppResponse> { + let data = json!({ "data": self.rags }); + let res = Response::builder() + .header("Content-Type", "application/json; charset=utf-8") + .body(Full::new(Bytes::from(data.to_string())).boxed())?; + Ok(res) + } + + async fn search_rag(&self, req: hyper::Request<Incoming>) -> Result<AppResponse> { + let req_body = req.collect().await?.to_bytes(); + let req_body: Value = serde_json::from_slice(&req_body) + .map_err(|err| anyhow!("Invalid request json, {err}"))?; + + debug!("search rag request: {req_body}"); + let SearchRagReqBody { name, input } = serde_json::from_value(req_body) + .map_err(|err| anyhow!("Invalid request body, {err}"))?; + + let config = Arc::new(RwLock::new(self.config.clone())); + + let abort_signal = create_abort_signal(); + + let rag = config + .read() + .rag_file(&name) + .ok() + .and_then(|rag_path| Rag::load(&config, &name, &rag_path).ok()) + .ok_or_else(|| anyhow!("Invalid rag"))?; + + let rag_result = Config::search_rag(&config, &rag, &input, abort_signal).await?; + + let data = json!({ "data": rag_result }); + let res = Response::builder() + .header("Content-Type", "application/json; charset=utf-8") + .body(Full::new(Bytes::from(data.to_string())).boxed())?; + Ok(res) + } + async fn chat_completions(&self, req: hyper::Request<Incoming>) -> Result<AppResponse> { let req_body = req.collect().await?.to_bytes(); let req_body: Value = serde_json::from_slice(&req_body) @@ -238,16 +278,15 @@ impl Server { stream, } = req_body; - let config = Config { - clients: self.clients.to_vec(), - model: self.model.clone(), - ..Default::default() - }; + let config = self.config.clone(); + + let default_model = config.model.clone(); + let config = Arc::new(RwLock::new(config)); let (model_name, change) = if model == DEFAULT_MODEL_NAME { - (self.model.id(), true) - } else if self.model.id() == model { + (default_model.id(), true) + } else if default_model.id() == model { (model, false) } else { (model, true) @@ -261,7 +300,7 @@ impl Server { if max_tokens.is_some() { client.model_mut().set_max_tokens(max_tokens, true); } - let abort = create_abort_signal(); + let abort_signal = create_abort_signal(); let http_client = client.build_client()?; let completion_id = generate_completion_id(); @@ -280,7 +319,7 @@ impl Server { tokio::spawn(async move { let is_first = Arc::new(AtomicBool::new(true)); let (sse_tx, sse_rx) = unbounded_channel(); - let mut handler = SseHandler::new(sse_tx, abort); + let mut handler = SseHandler::new(sse_tx, abort_signal); async fn map_event( mut sse_rx: UnboundedReceiver<SseEvent>, tx: &UnboundedSender<ResEvent>, @@ -399,11 +438,8 @@ impl Server { model: embedding_model_id, } = req_body; - let config = Config { - clients: self.clients.to_vec(), - ..Default::default() - }; - let config = Arc::new(RwLock::new(config)); + let config = Arc::new(RwLock::new(self.config.clone())); + let embedding_model = Model::retrieve_embedding(&config.read(), &embedding_model_id)?; let texts = match input { @@ -445,6 +481,12 @@ impl Server { } #[derive(Debug, Deserialize)] +struct SearchRagReqBody { + name: String, + input: String, +} + +#[derive(Debug, Deserialize)] struct ChatCompletionsReqBody { model: String, messages: Vec<Message>, |
