summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-29 19:28:54 +0800
committerGitHub <noreply@github.com>2024-04-29 19:28:54 +0800
commit4d1c53384b751e39e8f5c9d3c512adedca552fdf (patch)
tree01334c5c00e796d9b63698e31cc58c5a38498185 /src/client
parent4ddccc361c04592d54d796b24f72508a5234ffb5 (diff)
downloadaichat-4d1c53384b751e39e8f5c9d3c512adedca552fdf.tar.gz
refactor: prompts for generating config file (#463)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/azure_openai.rs2
-rw-r--r--src/client/claude.rs2
-rw-r--r--src/client/cloudflare.rs4
-rw-r--r--src/client/cohere.rs2
-rw-r--r--src/client/common.rs2
-rw-r--r--src/client/ollama.rs8
-rw-r--r--src/client/openai_compatible.rs6
-rw-r--r--src/client/vertexai.rs2
8 files changed, 16 insertions, 12 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index 0ec0c14..dd83ae1 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -27,7 +27,7 @@ impl AzureOpenAIClient {
(
"models[].max_input_tokens",
"Max Input Tokens:",
- true,
+ false,
PromptKind::Integer,
),
];
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 72ed405..be68af1 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -26,7 +26,7 @@ impl ClaudeClient {
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
- [("api_key", "API Key:", false, PromptKind::String)];
+ [("api_key", "API Key:", true, PromptKind::String)];
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_key = self.get_api_key().ok();
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index 09f020a..bb638fb 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -27,8 +27,8 @@ impl CloudflareClient {
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 2] = [
- ("account_id", "Account ID:", false, PromptKind::String),
- ("api_key", "API Key:", false, PromptKind::String),
+ ("account_id", "Account ID:", true, PromptKind::String),
+ ("api_key", "API Key:", true, PromptKind::String),
];
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 8586b9d..e2a23b0 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -25,7 +25,7 @@ impl CohereClient {
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
- [("api_key", "API Key:", false, PromptKind::String)];
+ [("api_key", "API Key:", true, PromptKind::String)];
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_key = self.get_api_key()?;
diff --git a/src/client/common.rs b/src/client/common.rs
index 85255ee..7eb8be8 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -203,7 +203,7 @@ macro_rules! openai_compatible_client {
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
- [("api_key", "API Key:", false, PromptKind::String)];
+ [("api_key", "API Key:", true, PromptKind::String)];
fn request_builder(
&self,
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index ca9d05d..158bbbe 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -14,7 +14,7 @@ use serde_json::{json, Value};
#[derive(Debug, Clone, Deserialize, Default)]
pub struct OllamaConfig {
pub name: Option<String>,
- pub api_base: String,
+ pub api_base: Option<String>,
pub api_auth: Option<String>,
pub chat_endpoint: Option<String>,
pub models: Vec<ModelConfig>,
@@ -22,11 +22,12 @@ pub struct OllamaConfig {
}
impl OllamaClient {
+ config_get_fn!(api_base, get_api_base);
config_get_fn!(api_auth, get_api_auth);
pub const PROMPTS: [PromptType<'static>; 4] = [
("api_base", "API Base:", true, PromptKind::String),
- ("api_auth", "API Key:", false, PromptKind::String),
+ ("api_auth", "API Auth:", false, PromptKind::String),
("models[].name", "Model Name:", true, PromptKind::String),
(
"models[].max_input_tokens",
@@ -37,6 +38,7 @@ impl OllamaClient {
];
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
+ let api_base = self.get_api_base()?;
let api_auth = self.get_api_auth().ok();
let mut body = build_body(data, &self.model)?;
@@ -44,7 +46,7 @@ impl OllamaClient {
let chat_endpoint = self.config.chat_endpoint.as_deref().unwrap_or("/api/chat");
- let url = format!("{}{chat_endpoint}", self.config.api_base);
+ let url = format!("{api_base}{chat_endpoint}");
debug!("Ollama Request: {url} {body}");
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 304e748..ba3a751 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -10,7 +10,7 @@ use serde::Deserialize;
#[derive(Debug, Clone, Deserialize)]
pub struct OpenAICompatibleConfig {
pub name: Option<String>,
- pub api_base: String,
+ pub api_base: Option<String>,
pub api_key: Option<String>,
pub chat_endpoint: Option<String>,
pub models: Vec<ModelConfig>,
@@ -18,6 +18,7 @@ pub struct OpenAICompatibleConfig {
}
impl OpenAICompatibleClient {
+ config_get_fn!(api_base, get_api_base);
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 5] = [
@@ -34,6 +35,7 @@ impl OpenAICompatibleClient {
];
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
+ let api_base = self.get_api_base()?;
let api_key = self.get_api_key().ok();
let mut body = openai_build_body(data, &self.model);
@@ -45,7 +47,7 @@ impl OpenAICompatibleClient {
.as_deref()
.unwrap_or("/chat/completions");
- let url = format!("{}{chat_endpoint}", self.config.api_base);
+ let url = format!("{api_base}{chat_endpoint}");
debug!("OpenAICompatible Request: {url} {body}");
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 529c822..3ac917d 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -34,7 +34,7 @@ impl VertexAIClient {
pub const PROMPTS: [PromptType<'static>; 2] = [
("project_id", "Project ID", true, PromptKind::String),
- ("location", "Global Location", true, PromptKind::String),
+ ("location", "Location", true, PromptKind::String),
];
fn request_builder(