summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-02 10:45:11 +0800
committerGitHub <noreply@github.com>2023-11-02 10:45:11 +0800
commit7c6841782d36faacc2aa3616dc3ed1b9403fa26e (patch)
tree12d51728e056dbd4e57fc5977b7c51b2b60c06ae /src/client
parent444f4ebe9de8aa68afdf8d39b734a8543b9b5c4b (diff)
downloadaichat-7c6841782d36faacc2aa3616dc3ed1b9403fa26e.tar.gz
refactor: improve code quanity (#197)
- move model_info.rs/message.rs to clients/ - rename SharedConfig to GlobalConfig
Diffstat (limited to 'src/client')
-rw-r--r--src/client/azure_openai.rs4
-rw-r--r--src/client/common.rs26
-rw-r--r--src/client/localai.rs4
-rw-r--r--src/client/message.rs52
-rw-r--r--src/client/mod.rs4
-rw-r--r--src/client/model_info.rs80
-rw-r--r--src/client/openai.rs2
7 files changed, 155 insertions, 17 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index fe3ec0f..f8a9dae 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -1,7 +1,7 @@
use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
-use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData};
+use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, ModelInfo};
-use crate::{config::ModelInfo, utils::PromptKind};
+use crate::utils::PromptKind;
use anyhow::{anyhow, Result};
use async_trait::async_trait;
diff --git a/src/client/common.rs b/src/client/common.rs
index 0d0c0e2..a7844f3 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,8 +1,12 @@
+use super::{openai::OpenAIConfig, ClientConfig, Message};
+
use crate::{
- config::{Message, SharedConfig},
+ config::GlobalConfig,
render::ReplyHandler,
- repl::AbortSignal,
- utils::{init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, PromptKind},
+ utils::{
+ init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, AbortSignal,
+ PromptKind,
+ },
};
use anyhow::{Context, Result};
@@ -13,8 +17,6 @@ use serde_json::{json, Value};
use std::{env, time::Duration};
use tokio::time::sleep;
-use super::{openai::OpenAIConfig, ClientConfig};
-
#[macro_export]
macro_rules! register_client {
(
@@ -42,15 +44,15 @@ macro_rules! register_client {
$(
#[derive(Debug)]
pub struct $client {
- global_config: $crate::config::SharedConfig,
+ global_config: $crate::config::GlobalConfig,
config: $config,
- model_info: $crate::config::ModelInfo,
+ model_info: $crate::client::ModelInfo,
}
impl $client {
pub const NAME: &str = $name;
- pub fn init(global_config: $crate::config::SharedConfig) -> Option<Box<dyn Client>> {
+ pub fn init(global_config: $crate::config::GlobalConfig) -> Option<Box<dyn Client>> {
let model_info = global_config.read().model_info.clone();
let config = {
if let ClientConfig::$config_key(c) = &global_config.read().clients[model_info.index] {
@@ -73,7 +75,7 @@ macro_rules! register_client {
)+
- pub fn init_client(config: $crate::config::SharedConfig) -> anyhow::Result<Box<dyn Client>> {
+ pub fn init_client(config: $crate::config::GlobalConfig) -> anyhow::Result<Box<dyn Client>> {
None
$(.or_else(|| $client::init(config.clone())))+
.ok_or_else(|| {
@@ -99,7 +101,7 @@ macro_rules! register_client {
anyhow::bail!("Unknown client {}", client)
}
- pub fn all_models(config: &$crate::config::Config) -> Vec<$crate::config::ModelInfo> {
+ pub fn all_models(config: &$crate::config::Config) -> Vec<$crate::client::ModelInfo> {
config
.clients
.iter()
@@ -122,7 +124,7 @@ macro_rules! openai_compatible_client {
fn config(
&self,
) -> (
- &$crate::config::SharedConfig,
+ &$crate::config::GlobalConfig,
&Option<$crate::client::ExtraConfig>,
) {
(&self.global_config, &self.config.extra)
@@ -169,7 +171,7 @@ macro_rules! config_get_fn {
#[async_trait]
pub trait Client {
- fn config(&self) -> (&SharedConfig, &Option<ExtraConfig>);
+ fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>);
fn build_client(&self) -> Result<ReqwestClient> {
let mut builder = ReqwestClient::builder();
diff --git a/src/client/localai.rs b/src/client/localai.rs
index 796b574..5cc12cc 100644
--- a/src/client/localai.rs
+++ b/src/client/localai.rs
@@ -1,7 +1,7 @@
use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
-use super::{ExtraConfig, LocalAIClient, PromptType, SendData};
+use super::{ExtraConfig, LocalAIClient, PromptType, SendData, ModelInfo};
-use crate::{config::ModelInfo, utils::PromptKind};
+use crate::utils::PromptKind;
use anyhow::Result;
use async_trait::async_trait;
diff --git a/src/client/message.rs b/src/client/message.rs
new file mode 100644
index 0000000..55b2663
--- /dev/null
+++ b/src/client/message.rs
@@ -0,0 +1,52 @@
+use serde::{Deserialize, Serialize};
+
+#[derive(Debug, Clone, Deserialize, Serialize)]
+pub struct Message {
+ pub role: MessageRole,
+ pub content: String,
+}
+
+impl Message {
+ pub fn new(content: &str) -> Self {
+ Self {
+ role: MessageRole::User,
+ content: content.to_string(),
+ }
+ }
+}
+
+#[derive(Debug, Clone, Copy, Deserialize, Serialize)]
+#[serde(rename_all = "snake_case")]
+pub enum MessageRole {
+ System,
+ Assistant,
+ User,
+}
+
+#[allow(dead_code)]
+impl MessageRole {
+ pub fn is_system(&self) -> bool {
+ matches!(self, MessageRole::System)
+ }
+
+ pub fn is_user(&self) -> bool {
+ matches!(self, MessageRole::User)
+ }
+
+ pub fn is_assistant(&self) -> bool {
+ matches!(self, MessageRole::Assistant)
+ }
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ #[test]
+ fn test_serde() {
+ assert_eq!(
+ serde_json::to_string(&Message::new("Hello World")).unwrap(),
+ "{\"role\":\"user\",\"content\":\"Hello World\"}"
+ );
+ }
+}
diff --git a/src/client/mod.rs b/src/client/mod.rs
index e55055d..19a0875 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -1,7 +1,11 @@
#[macro_use]
mod common;
+mod message;
+mod model_info;
pub use common::*;
+pub use message::*;
+pub use model_info::*;
register_client!(
(openai, "openai", OpenAI, OpenAIConfig, OpenAIClient),
diff --git a/src/client/model_info.rs b/src/client/model_info.rs
new file mode 100644
index 0000000..7a52e63
--- /dev/null
+++ b/src/client/model_info.rs
@@ -0,0 +1,80 @@
+use super::message::Message;
+
+use crate::utils::count_tokens;
+
+use anyhow::{bail, Result};
+
+pub type TokensCountFactors = (usize, usize); // (per-messages, bias)
+
+#[derive(Debug, Clone)]
+pub struct ModelInfo {
+ pub client: String,
+ pub name: String,
+ pub index: usize,
+ pub max_tokens: Option<usize>,
+ pub tokens_count_factors: TokensCountFactors,
+}
+
+impl Default for ModelInfo {
+ fn default() -> Self {
+ ModelInfo::new(0, "", "")
+ }
+}
+
+impl ModelInfo {
+ pub fn new(index: usize, client: &str, name: &str) -> Self {
+ Self {
+ index,
+ client: client.into(),
+ name: name.into(),
+ max_tokens: None,
+ tokens_count_factors: Default::default(),
+ }
+ }
+
+ pub fn set_max_tokens(mut self, max_tokens: Option<usize>) -> Self {
+ match max_tokens {
+ None | Some(0) => self.max_tokens = None,
+ _ => self.max_tokens = max_tokens,
+ }
+ self
+ }
+
+ pub fn set_tokens_count_factors(mut self, tokens_count_factors: TokensCountFactors) -> Self {
+ self.tokens_count_factors = tokens_count_factors;
+ self
+ }
+
+ pub fn full_name(&self) -> String {
+ format!("{}:{}", self.client, self.name)
+ }
+
+ pub fn messages_tokens(&self, messages: &[Message]) -> usize {
+ messages.iter().map(|v| count_tokens(&v.content)).sum()
+ }
+
+ pub fn total_tokens(&self, messages: &[Message]) -> usize {
+ if messages.is_empty() {
+ return 0;
+ }
+ let num_messages = messages.len();
+ let message_tokens = self.messages_tokens(messages);
+ let (per_messages, _) = self.tokens_count_factors;
+ if messages[num_messages - 1].role.is_user() {
+ num_messages * per_messages + message_tokens
+ } else {
+ (num_messages - 1) * per_messages + message_tokens
+ }
+ }
+
+ pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> {
+ let (_, bias) = self.tokens_count_factors;
+ let total_tokens = self.total_tokens(messages) + bias;
+ if let Some(max_tokens) = self.max_tokens {
+ if total_tokens >= max_tokens {
+ bail!("Exceed max tokens limit")
+ }
+ }
+ Ok(())
+ }
+}
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 98de5a8..5589d2d 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,9 +1,9 @@
use super::{
ExtraConfig, OpenAIClient, PromptType, SendData,
+ ModelInfo, TokensCountFactors,
};
use crate::{
- config::{ModelInfo, TokensCountFactors},
render::ReplyHandler,
utils::PromptKind,
};