summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-24 07:16:56 +0800
committerGitHub <noreply@github.com>2024-04-24 07:16:56 +0800
commit0a4c0413ef0154cde2e3485ec0415e6069596a23 (patch)
tree02527d5ee4b23b43a34bf90eeb817e1cc066d97c /src
parent9c6c9f10a27d0993636b453f39d8934c95c5c2b2 (diff)
downloadaichat-0a4c0413ef0154cde2e3485ec0415e6069596a23.tar.gz
feat: serve all LLMs as OpenAI-compatible API (#431)
Diffstat (limited to 'src')
-rw-r--r--src/cli.rs3
-rw-r--r--src/client/reply_handler.rs4
-rw-r--r--src/config/mod.rs44
-rw-r--r--src/logger.rs40
-rw-r--r--src/main.rs61
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/serve.rs364
7 files changed, 461 insertions, 57 deletions
diff --git a/src/cli.rs b/src/cli.rs
index 7cfd09a..aff608a 100644
--- a/src/cli.rs
+++ b/src/cli.rs
@@ -15,6 +15,9 @@ pub struct Cli {
/// Forces the session to be saved
#[clap(long)]
pub save_session: bool,
+ /// Serve all LLMs as OpenAI-compatible API
+ #[clap(long, value_name = "ADDRESS")]
+ pub serve: Option<Option<String>>,
/// Execute commands in natural language
#[clap(short = 'e', long)]
pub execute: bool,
diff --git a/src/client/reply_handler.rs b/src/client/reply_handler.rs
index e11ea1d..e024685 100644
--- a/src/client/reply_handler.rs
+++ b/src/client/reply_handler.rs
@@ -19,7 +19,7 @@ impl ReplyHandler {
}
pub fn text(&mut self, text: &str) -> Result<()> {
- debug!("ReplyText: {}", text);
+ // debug!("ReplyText: {}", text);
if text.is_empty() {
return Ok(());
}
@@ -33,7 +33,7 @@ impl ReplyHandler {
}
pub fn done(&mut self) -> Result<()> {
- debug!("ReplyDone");
+ // debug!("ReplyDone");
let ret = self
.sender
.send(ReplyEvent::Done)
diff --git a/src/config/mod.rs b/src/config/mod.rs
index f86cda1..e1cec4d 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -79,9 +79,9 @@ pub struct Config {
#[serde(skip)]
pub model: Model,
#[serde(skip)]
- pub last_message: Option<(Input, String)>,
+ pub working_mode: WorkingMode,
#[serde(skip)]
- pub in_repl: bool,
+ pub last_message: Option<(Input, String)>,
}
impl Default for Config {
@@ -110,8 +110,8 @@ impl Default for Config {
role: None,
session: None,
model: Default::default(),
+ working_mode: WorkingMode::Command,
last_message: None,
- in_repl: false,
}
}
}
@@ -119,13 +119,13 @@ impl Default for Config {
pub type GlobalConfig = Arc<RwLock<Config>>;
impl Config {
- pub fn init(is_interactive: bool) -> Result<Self> {
+ pub fn init(working_mode: WorkingMode) -> Result<Self> {
let config_path = Self::config_file()?;
let api_key = env::var("OPENAI_API_KEY").ok();
let exist_config_path = config_path.exists();
- if is_interactive && api_key.is_none() && !exist_config_path {
+ if working_mode != WorkingMode::Command && api_key.is_none() && !exist_config_path {
create_config_file(&config_path)?;
}
let mut config = if api_key.is_some() && !exist_config_path {
@@ -143,14 +143,13 @@ impl Config {
config.set_wrap(&wrap)?;
}
+ config.working_mode = working_mode;
config.load_roles()?;
config.setup_model()?;
config.setup_highlight();
config.setup_light_theme()?;
- setup_logger()?;
-
Ok(config)
}
@@ -611,7 +610,7 @@ impl Config {
let save_session = session.save_session();
if session.dirty && save_session != Some(false) {
if save_session.is_none() || session.is_temp() {
- if !self.in_repl {
+ if self.working_mode != WorkingMode::Repl {
return Ok(());
}
let ans = Confirm::new("Save session?").with_default(false).prompt()?;
@@ -999,6 +998,13 @@ impl Keybindings {
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
+pub enum WorkingMode {
+ Command,
+ Repl,
+ Serve,
+}
+
+#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum State {
Normal,
Role,
@@ -1145,25 +1151,3 @@ fn complete_option_bool(value: Option<bool>) -> Vec<String> {
None => vec!["true".to_string(), "false".to_string()],
}
}
-
-#[cfg(debug_assertions)]
-fn setup_logger() -> Result<()> {
- use simplelog::{LevelFilter, WriteLogger};
- let file = std::fs::File::create(Config::local_path("debug.log")?)?;
- let log_filter = match std::env::var("AICHAT_LOG_FILTER") {
- Ok(v) => v,
- Err(_) => "aichat".into(),
- };
- let config = simplelog::ConfigBuilder::new()
- .add_filter_allow(log_filter)
- .set_thread_level(LevelFilter::Off)
- .set_time_level(LevelFilter::Off)
- .build();
- WriteLogger::init(log::LevelFilter::Debug, config, file)?;
- Ok(())
-}
-
-#[cfg(not(debug_assertions))]
-fn setup_logger() -> Result<()> {
- Ok(())
-}
diff --git a/src/logger.rs b/src/logger.rs
new file mode 100644
index 0000000..f7ef2f5
--- /dev/null
+++ b/src/logger.rs
@@ -0,0 +1,40 @@
+use crate::config::WorkingMode;
+
+use anyhow::Result;
+use log::LevelFilter;
+use simplelog::{format_description, Config as LogConfig, ConfigBuilder};
+
+#[cfg(debug_assertions)]
+pub fn setup_logger(working_mode: WorkingMode) -> Result<()> {
+ let config = build_config();
+ if working_mode == WorkingMode::Serve {
+ simplelog::SimpleLogger::init(LevelFilter::Debug, config)?;
+ } else {
+ let file = std::fs::File::create(crate::config::Config::local_path("debug.log")?)?;
+ simplelog::WriteLogger::init(LevelFilter::Debug, config, file)?;
+ }
+ Ok(())
+}
+
+#[cfg(not(debug_assertions))]
+pub fn setup_logger(working_mode: WorkingMode) -> Result<()> {
+ let config = build_config();
+ if working_mode == WorkingMode::Serve {
+ simplelog::SimpleLogger::init(log::LevelFilter::Info, config)?;
+ }
+ Ok(())
+}
+
+fn build_config() -> LogConfig {
+ let log_filter = match std::env::var("AICHAT_LOG_FILTER") {
+ Ok(v) => v,
+ Err(_) => "aichat".into(),
+ };
+ ConfigBuilder::new()
+ .add_filter_allow(log_filter)
+ .set_time_format_custom(format_description!(
+ "[year]-[month]-[day]T[hour]:[minute]:[second].[subsecond digits:3]Z"
+ ))
+ .set_thread_level(LevelFilter::Off)
+ .build()
+}
diff --git a/src/main.rs b/src/main.rs
index bc6b8f1..b396917 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -1,17 +1,21 @@
mod cli;
mod client;
mod config;
+mod logger;
mod render;
mod repl;
+mod serve;
+#[macro_use]
+mod utils;
#[macro_use]
extern crate log;
-#[macro_use]
-mod utils;
use crate::cli::Cli;
use crate::client::{ensure_model_capabilities, init_client, list_models, send_stream};
-use crate::config::{Config, GlobalConfig, Input, CODE_ROLE, EXPLAIN_ROLE, SHELL_ROLE};
+use crate::config::{
+ Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_ROLE, SHELL_ROLE,
+};
use crate::render::{render_error, MarkdownRender};
use crate::repl::Repl;
use crate::utils::{
@@ -33,7 +37,21 @@ use tokio::sync::oneshot;
async fn main() -> Result<()> {
let cli = Cli::parse();
let text = cli.text();
- let config = Arc::new(RwLock::new(Config::init(text.is_none())?));
+ let file = &cli.file;
+ let no_input = text.is_none() && file.is_empty();
+ let working_mode = if cli.serve.is_some() {
+ WorkingMode::Serve
+ } else if no_input {
+ WorkingMode::Repl
+ } else {
+ WorkingMode::Command
+ };
+ crate::logger::setup_logger(working_mode)?;
+ let config = Arc::new(RwLock::new(Config::init(working_mode)?));
+
+ if let Some(addr) = cli.serve {
+ return serve::run(config, addr).await;
+ }
if cli.list_roles {
config
.read()
@@ -89,20 +107,21 @@ async fn main() -> Result<()> {
return Ok(());
}
let text = aggregate_text(text)?;
- let input = create_input(&config, text, &cli.file)?;
if cli.execute {
- match input {
- Some(input) => {
- execute(&config, input).await?;
- return Ok(());
- }
- None => bail!("No input text"),
+ if no_input {
+ bail!("No input");
}
+ let input = create_input(&config, text, file)?;
+ execute(&config, input).await?;
+ return Ok(());
}
config.write().apply_prelude()?;
- if let Err(err) = match input {
- Some(input) => start_directive(&config, input, cli.no_stream, cli.code).await,
- None => start_interactive(&config).await,
+ if let Err(err) = match no_input {
+ false => {
+ let input = create_input(&config, text, file)?;
+ start_directive(&config, input, cli.no_stream, cli.code).await
+ }
+ true => start_interactive(&config).await,
} {
let highlight = stderr().is_terminal() && config.read().highlight;
render_error(err, highlight)
@@ -232,19 +251,15 @@ fn aggregate_text(text: Option<String>) -> Result<Option<String>> {
Ok(text)
}
-fn create_input(
- config: &GlobalConfig,
- text: Option<String>,
- file: &[String],
-) -> Result<Option<Input>> {
- if text.is_none() && file.is_empty() {
- return Ok(None);
- }
+fn create_input(config: &GlobalConfig, text: Option<String>, file: &[String]) -> Result<Input> {
let input_context = config.read().input_context();
let input = if file.is_empty() {
Input::from_str(&text.unwrap_or_default(), input_context)
} else {
Input::new(&text.unwrap_or_default(), file.to_vec(), input_context)?
};
- Ok(Some(input))
+ if input.is_empty() {
+ bail!("No input");
+ }
+ Ok(input)
}
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index 3c0b98b..dd1c735 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -77,8 +77,6 @@ pub struct Repl {
impl Repl {
pub fn init(config: &GlobalConfig) -> Result<Self> {
- config.write().in_repl = true;
-
let editor = Self::create_editor(config)?;
let prompt = ReplPrompt::new(config);
diff --git a/src/serve.rs b/src/serve.rs
new file mode 100644
index 0000000..2286b61
--- /dev/null
+++ b/src/serve.rs
@@ -0,0 +1,364 @@
+use crate::{
+ client::{init_client, ClientConfig, Message, Model, ReplyEvent, ReplyHandler, SendData},
+ config::{Config, GlobalConfig},
+ utils::create_abort_signal,
+};
+
+use anyhow::{anyhow, bail, Result};
+use bytes::Bytes;
+use chrono::{Timelike, Utc};
+use futures_util::StreamExt;
+use http::{Method, Response, StatusCode};
+use http_body_util::{combinators::BoxBody, BodyExt, Full, StreamBody};
+use hyper::{
+ body::{Frame, Incoming},
+ service::service_fn,
+};
+use hyper_util::rt::{TokioExecutor, TokioIo};
+use parking_lot::RwLock;
+use serde::Deserialize;
+use serde_json::{json, Value};
+use std::{convert::Infallible, sync::Arc};
+use tokio::{
+ net::TcpListener,
+ sync::{
+ mpsc::{unbounded_channel, UnboundedReceiver, UnboundedSender},
+ oneshot,
+ },
+};
+use tokio_graceful::Shutdown;
+use tokio_stream::wrappers::UnboundedReceiverStream;
+
+const DEFAULT_ADDRESS: &str = "0.0.0.0:8080";
+
+type AppResponse = Response<BoxBody<Bytes, Infallible>>;
+
+pub async fn run(config: GlobalConfig, addr: Option<String>) -> Result<()> {
+ let addr = match addr {
+ Some(addr) => {
+ if let Ok(port) = addr.parse::<u16>() {
+ format!("0.0.0.0:{port}")
+ } else {
+ addr
+ }
+ }
+ None => DEFAULT_ADDRESS.to_string(),
+ };
+ let clients = config.read().clients.clone();
+ let model = config.read().model.clone();
+ 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");
+ shutdown_signal().await;
+ let _ = stop_server.send(());
+ Ok(())
+}
+
+struct Server {
+ clients: Vec<ClientConfig>,
+ model: Model,
+}
+
+impl Server {
+ async fn run(self: Arc<Self>, listener: TcpListener) -> Result<oneshot::Sender<()>> {
+ let (tx, rx) = oneshot::channel();
+ tokio::spawn(async move {
+ let shutdown = Shutdown::new(async { rx.await.unwrap_or_default() });
+ let guard = shutdown.guard_weak();
+
+ loop {
+ tokio::select! {
+ res = listener.accept() => {
+ let Ok((cnx, _)) = res else {
+ continue;
+ };
+
+ let stream = TokioIo::new(cnx);
+ let server = self.clone();
+ shutdown.spawn_task(async move {
+ let hyper_service = service_fn(move |request: hyper::Request<Incoming>| {
+ server.clone().handle(request)
+ });
+ let _ = hyper_util::server::conn::auto::Builder::new(TokioExecutor::new())
+ .serve_connection_with_upgrades(stream, hyper_service)
+ .await;
+ });
+ }
+ _ = guard.cancelled() => {
+ break;
+ }
+ }
+ }
+ });
+ Ok(tx)
+ }
+
+ async fn handle(
+ self: Arc<Self>,
+ req: hyper::Request<Incoming>,
+ ) -> std::result::Result<AppResponse, hyper::Error> {
+ let method = req.method().clone();
+ let uri = req.uri().clone();
+ let mut status = StatusCode::OK;
+ let res = if method == Method::POST && uri == "/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 {
+ status = StatusCode::NOT_FOUND;
+ Err(anyhow!("The requested endpoint was not found."))
+ };
+ let mut res = match res {
+ Ok(res) => {
+ info!("{method} {uri} {}", status.as_u16());
+ res
+ }
+ Err(err) => {
+ error!("{method} {uri} {} {err}", status.as_u16());
+ ret_err(err)
+ }
+ };
+ *res.status_mut() = status;
+ set_cors_header(&mut res);
+ Ok(res)
+ }
+
+ async fn chat_completion(&self, req: hyper::Request<Incoming>) -> Result<AppResponse> {
+ let req_body = req.collect().await?.to_bytes();
+ let req_body: ChatCompletionReqBody = serde_json::from_slice(&req_body)
+ .map_err(|err| anyhow!("Invalid request body, {err}"))?;
+
+ let ChatCompletionReqBody {
+ model,
+ messages,
+ temperature,
+ max_tokens,
+ stream,
+ } = req_body;
+
+ let config = Config {
+ clients: self.clients.to_vec(),
+ model: self.model.clone(),
+ ..Default::default()
+ };
+ let config = Arc::new(RwLock::new(config));
+ if model != "default" && model != self.model.id() {
+ config.write().set_model(&model)?;
+ }
+
+ let mut client = init_client(&config)?;
+ if max_tokens.is_some() {
+ client.set_model(client.model().clone().set_max_output_tokens(max_tokens));
+ }
+ let abort = create_abort_signal();
+ let http_client = client.build_client()?;
+
+ let completion_id = generate_completion_id();
+ let created = Utc::now().timestamp();
+
+ let send_data: SendData = SendData {
+ messages,
+ temperature,
+ stream,
+ };
+
+ if stream {
+ let (tx, mut rx) = unbounded_channel();
+ tokio::spawn(async move {
+ let mut is_first = true;
+ let (tx2, rx2) = unbounded_channel();
+ let mut handler = ReplyHandler::new(tx2, abort);
+ async fn map_event(
+ mut rx: UnboundedReceiver<ReplyEvent>,
+ tx: &UnboundedSender<ResEvent>,
+ is_first: &mut bool,
+ ) {
+ while let Some(reply_event) = rx.recv().await {
+ if *is_first {
+ let _ = tx.send(ResEvent::First(None));
+ *is_first = false;
+ }
+ match reply_event {
+ ReplyEvent::Text(text) => {
+ let _ = tx.send(ResEvent::Text(text));
+ }
+ ReplyEvent::Done => {
+ let _ = tx.send(ResEvent::Done);
+ }
+ }
+ }
+ }
+ tokio::select! {
+ _ = map_event(rx2, &tx, &mut is_first) => {}
+ ret = client.send_message_streaming_inner(&http_client, &mut handler, send_data) => {
+ if let Err(err) = ret {
+ send_first_event(&tx, Some(format!("{err:?}")), &mut is_first)
+ }
+ }
+ }
+ });
+
+ let first_event = rx.recv().await;
+
+ if let Some(ResEvent::First(Some(err))) = first_event {
+ bail!("{err}");
+ }
+
+ let shared: Arc<(String, i64)> = Arc::new((completion_id, created));
+ let stream = UnboundedReceiverStream::new(rx);
+ let stream = stream.filter_map(move |res_event| {
+ let shared = shared.clone();
+ async move {
+ match res_event {
+ ResEvent::Text(text) => {
+ Some(Ok(create_frame(&shared.0, shared.1, &text, false)))
+ }
+ ResEvent::Done => Some(Ok(create_frame(&shared.0, shared.1, "", true))),
+ _ => None,
+ }
+ }
+ });
+ let res = Response::builder()
+ .status(StatusCode::OK)
+ .header("Content-Type", "text/event-stream")
+ .header("Cache-Control", "no-cache")
+ .header("Connection", "keep-alive")
+ .body(BodyExt::boxed(StreamBody::new(stream)))?;
+ Ok(res)
+ } else {
+ let content = client.send_message_inner(&http_client, send_data).await?;
+ let res = Response::builder()
+ .header("Content-Type", "application/json")
+ .body(Full::new(ret_non_stream(&completion_id, created, &content)).boxed())?;
+ Ok(res)
+ }
+ }
+}
+
+#[derive(Debug, Deserialize)]
+struct ChatCompletionReqBody {
+ model: String,
+ messages: Vec<Message>,
+ temperature: Option<f64>,
+ max_tokens: Option<isize>,
+ #[serde(default)]
+ stream: bool,
+}
+
+#[derive(Debug)]
+enum ResEvent {
+ First(Option<String>),
+ Text(String),
+ Done,
+}
+
+fn send_first_event(tx: &UnboundedSender<ResEvent>, data: Option<String>, is_first: &mut bool) {
+ if *is_first {
+ let _ = tx.send(ResEvent::First(data));
+ *is_first = false;
+ }
+}
+
+async fn shutdown_signal() {
+ tokio::signal::ctrl_c()
+ .await
+ .expect("Failed to install CTRL+C signal handler")
+}
+
+fn generate_completion_id() -> String {
+ let random_id = chrono::Utc::now().nanosecond();
+ format!("chatcmpl-{}", random_id)
+}
+
+fn set_cors_header(res: &mut AppResponse) {
+ res.headers_mut().insert(
+ hyper::header::ACCESS_CONTROL_ALLOW_ORIGIN,
+ hyper::header::HeaderValue::from_static("*"),
+ );
+ res.headers_mut().insert(
+ hyper::header::ACCESS_CONTROL_ALLOW_METHODS,
+ hyper::header::HeaderValue::from_static("GET,POST,PUT,PATCH,DELETE"),
+ );
+ res.headers_mut().insert(
+ hyper::header::ACCESS_CONTROL_ALLOW_HEADERS,
+ hyper::header::HeaderValue::from_static("Content-Type,Authorization"),
+ );
+}
+
+fn create_frame(id: &str, created: i64, content: &str, done: bool) -> Frame<Bytes> {
+ let (delta, finish_reason) = if done {
+ (json!({}), "stop".into())
+ } else {
+ let delta = if content.is_empty() {
+ json!({ "role": "assistant", "content": content })
+ } else {
+ json!({ "content": content })
+ };
+ (delta, Value::Null)
+ };
+ let mut value = json!({
+ "id": id,
+ "object": "chat.completion.chunk",
+ "created": created,
+ "model": "gpt-3.5-turbo",
+ "choices": [
+ {
+ "index": 0,
+ "delta": delta,
+ "finish_reason": finish_reason,
+ },
+ ],
+ });
+ let output = if done {
+ value["usage"] = json!({
+ "prompt_tokens": 0,
+ "completion_tokens": 0,
+ "total_tokens": 0,
+ });
+ format!("data: {value}\n\ndata: [DONE]\n\n")
+ } else {
+ format!("data: {value}\n\n")
+ };
+ Frame::data(Bytes::from(output))
+}
+
+fn ret_non_stream(id: &str, created: i64, content: &str) -> Bytes {
+ let res_body = json!({
+ "id": id,
+ "object": "chat.completion",
+ "created": created,
+ "model": "gpt-3.5-turbo",
+ "choices": [
+ {
+ "index": 0,
+ "message": {
+ "role": "assistant",
+ "content": content,
+ },
+ "finish_reason": "stop",
+ },
+ ],
+ "usage": {
+ "prompt_tokens": 0,
+ "completion_tokens": 0,
+ "total_tokens": 0,
+ },
+ });
+ Bytes::from(res_body.to_string())
+}
+
+fn ret_err<T: std::fmt::Display>(err: T) -> AppResponse {
+ let data = json!({
+ "error": {
+ "message": err.to_string(),
+ "type": "invalid_request_error",
+ },
+ });
+ Response::builder()
+ .status(StatusCode::OK)
+ .header("Content-Type", "application/json")
+ .body(Full::new(Bytes::from(data.to_string())).boxed())
+ .unwrap()
+}