summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--assets/playground.html41
-rw-r--r--src/config/input.rs33
-rw-r--r--src/config/mod.rs45
-rw-r--r--src/main.rs2
-rw-r--r--src/serve.rs92
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>,