From 85ad276a29b75b24903ab925956118d537afd52d Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 7 May 2024 09:12:18 +0800 Subject: feat: support playground/arena webui (#487) --- src/client/model.rs | 4 +++ src/config/mod.rs | 3 +- src/serve.rs | 99 +++++++++++++++++++++++++++++++++++++++++++++++------ 3 files changed, 93 insertions(+), 13 deletions(-) (limited to 'src') diff --git a/src/client/model.rs b/src/client/model.rs index 90e46e4..b24a546 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -111,6 +111,10 @@ impl Model { ) } + pub fn supports_vision(&self) -> bool { + self.capabilities.contains(ModelCapabilities::Vision) + } + pub fn show_max_output_tokens(&self) -> Option { self.max_output_tokens.or(self.ref_max_output_tokens) } diff --git a/src/config/mod.rs b/src/config/mod.rs index 3dbd8ac..f33ee73 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -3,8 +3,7 @@ mod role; mod session; pub use self::input::{Input, InputContext}; -use self::role::Role; -pub use self::role::{CODE_ROLE, EXPLAIN_ROLE, SHELL_ROLE}; +pub use self::role::{Role, CODE_ROLE, EXPLAIN_ROLE, SHELL_ROLE}; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ diff --git a/src/serve.rs b/src/serve.rs index 714e8f8..5f748e8 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -1,9 +1,9 @@ use crate::{ client::{ - init_client, ClientConfig, CompletionDetails, Message, Model, SendData, SseEvent, - SseHandler, + init_client, list_models, ClientConfig, CompletionDetails, Message, Model, SendData, + SseEvent, SseHandler, }, - config::{Config, GlobalConfig}, + config::{Config, GlobalConfig, Role}, utils::create_abort_signal, }; @@ -34,6 +34,8 @@ use tokio_stream::wrappers::UnboundedReceiverStream; const DEFAULT_ADDRESS: &str = "127.0.0.1:8000"; const DEFAULT_MODEL_NAME: &str = "default"; +const PLAYGROUND_HTML: &[u8] = include_bytes!("../assets/playground.html"); +const ARENA_HTML: &[u8] = include_bytes!("../assets/arena.html"); type AppResponse = Response>; @@ -50,12 +52,12 @@ pub async fn run(config: GlobalConfig, addr: Option) -> Result<()> { } None => DEFAULT_ADDRESS.to_string(), }; - let clients = config.read().clients.clone(); - let model = config.read().model.clone(); + let server = Arc::new(Server::new(&config)); let listener = TcpListener::bind(&addr).await?; - let server = Arc::new(Server { clients, model }); let stop_server = server.run(listener).await?; - println!("Access the chat completion API at: http://{addr}/v1/chat/completions"); + println!("Chat Completions API: http://{addr}/v1/chat/completions"); + println!("LLM Playground: http://{addr}/playground"); + println!("LLM ARENA: http://{addr}/arena"); shutdown_signal().await; let _ = stop_server.send(()); Ok(()) @@ -64,9 +66,47 @@ pub async fn run(config: GlobalConfig, addr: Option) -> Result<()> { struct Server { clients: Vec, model: Model, + models: Vec, + roles: Vec, } impl Server { + fn new(config: &GlobalConfig) -> Self { + let config = config.read(); + let clients = config.clients.clone(); + let model = config.model.clone(); + let roles = config.roles.clone(); + let mut models = list_models(&config); + let mut default_model = model.clone(); + default_model.name = DEFAULT_MODEL_NAME.into(); + models.insert(0, &default_model); + let models: Vec = models + .into_iter() + .enumerate() + .map(|(i, model)| { + let id = if i == 0 { + DEFAULT_MODEL_NAME.into() + } else { + model.id() + }; + json!({ + "id": id, + "max_input_tokens": model.max_input_tokens, + "max_output_tokens": model.max_output_tokens, + "max_output_tokens?": model.ref_max_output_tokens, + "input_price": model.input_price, + "output_price": model.output_price, + "supports_vision": model.supports_vision(), + }) + }) + .collect(); + Self { + clients, + model, + roles, + models, + } + } async fn run(self: Arc, listener: TcpListener) -> Result> { let (tx, rx) = oneshot::channel(); tokio::spawn(async move { @@ -106,12 +146,24 @@ impl Server { ) -> std::result::Result { let method = req.method().clone(); let uri = req.uri().clone(); + let path = uri.path(); + + if method == Method::OPTIONS { + let mut res = Response::default(); + *res.status_mut() = StatusCode::NO_CONTENT; + set_cors_header(&mut res); + return Ok(res); + } + let mut status = StatusCode::OK; - let res = if method == Method::POST && uri == "/v1/chat/completions" { + let res = if path == "/v1/chat/completions" { self.chat_completion(req).await - } else if method == Method::OPTIONS && uri == "/v1/chat/completions" { - status = StatusCode::NO_CONTENT; - Ok(Response::default()) + } else if path == "/playground" || path == "/playground.html" { + self.playground_page() + } else if path == "/arena" || path == "/arena.html" { + self.arena_page() + } else if path == "/data.json" { + self.data_json() } else { status = StatusCode::NOT_FOUND; Err(anyhow!("The requested endpoint was not found.")) @@ -132,6 +184,31 @@ impl Server { Ok(res) } + fn playground_page(&self) -> Result { + let res = Response::builder() + .header("Content-Type", "text/html; charset=utf-8") + .body(Full::new(Bytes::from(PLAYGROUND_HTML)).boxed())?; + Ok(res) + } + + fn arena_page(&self) -> Result { + let res = Response::builder() + .header("Content-Type", "text/html; charset=utf-8") + .body(Full::new(Bytes::from(ARENA_HTML)).boxed())?; + Ok(res) + } + + fn data_json(&self) -> Result { + let data = json!({ + "models": self.models, + "roles": self.roles, + }); + 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_completion(&self, req: hyper::Request) -> Result { let req_body = req.collect().await?.to_bytes(); let req_body: ChatCompletionReqBody = serde_json::from_slice(&req_body) -- cgit v1.2.3