summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/azure_openai.rs12
-rw-r--r--src/client/common.rs18
-rw-r--r--src/client/localai.rs10
-rw-r--r--src/client/mod.rs4
-rw-r--r--src/client/model.rs (renamed from src/client/model_info.rs)30
-rw-r--r--src/client/openai.rs10
6 files changed, 42 insertions, 42 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index f8a9dae..d1dc43b 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -1,5 +1,5 @@
use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
-use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, ModelInfo};
+use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, Model};
use crate::utils::PromptKind;
@@ -42,14 +42,14 @@ impl AzureOpenAIClient {
),
];
- pub fn list_models(local_config: &AzureOpenAIConfig, index: usize) -> Vec<ModelInfo> {
- let client = Self::name(local_config);
+ pub fn list_models(local_config: &AzureOpenAIConfig, client_index: usize) -> Vec<Model> {
+ let client_name = Self::name(local_config);
local_config
.models
.iter()
.map(|v| {
- ModelInfo::new(index, client, &v.name)
+ Model::new(client_index, client_name, &v.name)
.set_max_tokens(v.max_tokens)
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
@@ -70,11 +70,11 @@ impl AzureOpenAIClient {
let api_base = self.get_api_base()?;
- let body = openai_build_body(data, self.model_info.name.clone());
+ let body = openai_build_body(data, self.model.llm_name.clone());
let url = format!(
"{}/openai/deployments/{}/chat/completions?api-version=2023-05-15",
- &api_base, self.model_info.name
+ &api_base, self.model.llm_name
);
let builder = client.post(url).header("api-key", api_key).json(&body);
diff --git a/src/client/common.rs b/src/client/common.rs
index 0dc637b..464e46a 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -46,16 +46,16 @@ macro_rules! register_client {
pub struct $client {
global_config: $crate::config::GlobalConfig,
config: $config,
- model_info: $crate::client::ModelInfo,
+ model: $crate::client::Model,
}
impl $client {
pub const NAME: &str = $name;
pub fn init(global_config: $crate::config::GlobalConfig) -> Option<Box<dyn Client>> {
- let model_info = global_config.read().model_info.clone();
+ let model = global_config.read().model.clone();
let config = {
- if let ClientConfig::$config_key(c) = &global_config.read().clients[model_info.index] {
+ if let ClientConfig::$config_key(c) = &global_config.read().clients[model.client_index] {
c.clone()
} else {
return None;
@@ -64,7 +64,7 @@ macro_rules! register_client {
Some(Box::new(Self {
global_config,
config,
- model_info,
+ model,
}))
}
@@ -79,11 +79,11 @@ macro_rules! register_client {
None
$(.or_else(|| $client::init(config.clone())))+
.ok_or_else(|| {
- let model_info = config.read().model_info.clone();
+ let model = config.read().model.clone();
anyhow::anyhow!(
- "Unknown client {} at config.clients[{}]",
- &model_info.client,
- &model_info.index
+ "Unknown client '{}' at config.clients[{}]",
+ &model.client_name,
+ &model.client_index
)
})
}
@@ -101,7 +101,7 @@ macro_rules! register_client {
anyhow::bail!("Unknown client {}", client)
}
- pub fn list_models(config: &$crate::config::Config) -> Vec<$crate::client::ModelInfo> {
+ pub fn list_models(config: &$crate::config::Config) -> Vec<$crate::client::Model> {
config
.clients
.iter()
diff --git a/src/client/localai.rs b/src/client/localai.rs
index 5cc12cc..eb4de65 100644
--- a/src/client/localai.rs
+++ b/src/client/localai.rs
@@ -1,5 +1,5 @@
use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
-use super::{ExtraConfig, LocalAIClient, PromptType, SendData, ModelInfo};
+use super::{ExtraConfig, LocalAIClient, PromptType, SendData, Model};
use crate::utils::PromptKind;
@@ -41,14 +41,14 @@ impl LocalAIClient {
),
];
- pub fn list_models(local_config: &LocalAIConfig, index: usize) -> Vec<ModelInfo> {
- let client = Self::name(local_config);
+ pub fn list_models(local_config: &LocalAIConfig, client_index: usize) -> Vec<Model> {
+ let client_name = Self::name(local_config);
local_config
.models
.iter()
.map(|v| {
- ModelInfo::new(index, client, &v.name)
+ Model::new(client_index, client_name, &v.name)
.set_max_tokens(v.max_tokens)
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
@@ -58,7 +58,7 @@ impl LocalAIClient {
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_key = self.get_api_key().ok();
- let body = openai_build_body(data, self.model_info.name.clone());
+ let body = openai_build_body(data, self.model.llm_name.clone());
let chat_endpoint = self
.config
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 19a0875..7ac9aa0 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -1,11 +1,11 @@
#[macro_use]
mod common;
mod message;
-mod model_info;
+mod model;
pub use common::*;
pub use message::*;
-pub use model_info::*;
+pub use model::*;
register_client!(
(openai, "openai", OpenAI, OpenAIConfig, OpenAIClient),
diff --git a/src/client/model_info.rs b/src/client/model.rs
index 9e74951..d00ad46 100644
--- a/src/client/model_info.rs
+++ b/src/client/model.rs
@@ -7,31 +7,35 @@ 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 struct Model {
+ pub client_index: usize,
+ pub client_name: String,
+ pub llm_name: String,
pub max_tokens: Option<usize>,
pub tokens_count_factors: TokensCountFactors,
}
-impl Default for ModelInfo {
+impl Default for Model {
fn default() -> Self {
- ModelInfo::new(0, "", "")
+ Model::new(0, "", "")
}
}
-impl ModelInfo {
- pub fn new(index: usize, client: &str, name: &str) -> Self {
+impl Model {
+ pub fn new(client_index: usize, client_name: &str, name: &str) -> Self {
Self {
- index,
- client: client.into(),
- name: name.into(),
+ client_index,
+ client_name: client_name.into(),
+ llm_name: name.into(),
max_tokens: None,
tokens_count_factors: Default::default(),
}
}
+ pub fn id(&self) -> String {
+ format!("{}:{}", self.client_name, self.llm_name)
+ }
+
pub fn set_max_tokens(mut self, max_tokens: Option<usize>) -> Self {
match max_tokens {
None | Some(0) => self.max_tokens = None,
@@ -45,10 +49,6 @@ impl ModelInfo {
self
}
- pub fn id(&self) -> String {
- format!("{}:{}", self.client, self.name)
- }
-
pub fn messages_tokens(&self, messages: &[Message]) -> usize {
messages.iter().map(|v| count_tokens(&v.content)).sum()
}
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 5589d2d..f1243f0 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,6 +1,6 @@
use super::{
ExtraConfig, OpenAIClient, PromptType, SendData,
- ModelInfo, TokensCountFactors,
+ Model, TokensCountFactors,
};
use crate::{
@@ -44,12 +44,12 @@ impl OpenAIClient {
pub const PROMPTS: [PromptType<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
- pub fn list_models(local_config: &OpenAIConfig, index: usize) -> Vec<ModelInfo> {
- let client = Self::name(local_config);
+ pub fn list_models(local_config: &OpenAIConfig, client_index: usize) -> Vec<Model> {
+ let client_name = Self::name(local_config);
MODELS
.into_iter()
.map(|(name, max_tokens)| {
- ModelInfo::new(index, client, name)
+ Model::new(client_index, client_name, name)
.set_max_tokens(Some(max_tokens))
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
@@ -59,7 +59,7 @@ impl OpenAIClient {
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_key = self.get_api_key()?;
- let body = openai_build_body(data, self.model_info.name.clone());
+ let body = openai_build_body(data, self.model.llm_name.clone());
let env_prefix = Self::name(&self.config).to_uppercase();
let api_base = env::var(format!("{env_prefix}_API_BASE"))