summaryrefslogtreecommitdiffstats
path: root/src/client/openai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-10-31 06:56:00 +0800
committerGitHub <noreply@github.com>2023-10-31 06:56:00 +0800
commit84004fd5767abc94a5360b742e7fd0c8a79c6c55 (patch)
tree539e0e642b6acfde44d445ba8178f781e1c98afb /src/client/openai.rs
parent36380475d5bd923f537e1aa29cad9f63bd138d40 (diff)
downloadaichat-84004fd5767abc94a5360b742e7fd0c8a79c6c55.tar.gz
feat: support OPENAI_API_BASE (#186)
Diffstat (limited to 'src/client/openai.rs')
-rw-r--r--src/client/openai.rs9
1 files changed, 7 insertions, 2 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 4ccb1b7..e0d321d 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -14,7 +14,7 @@ use serde_json::{json, Value};
use std::env;
use std::time::Duration;
-const API_URL: &str = "https://api.openai.com/v1/chat/completions";
+const API_BASE: &str = "https://api.openai.com/v1";
#[derive(Debug)]
pub struct OpenAIClient {
@@ -161,7 +161,12 @@ impl OpenAIClient {
.with_context(|| "Failed to build client")?
};
- let mut builder = client.post(API_URL).bearer_auth(api_key).json(&body);
+ let api_base = env::var("OPENAI_API_BASE")
+ .ok()
+ .unwrap_or_else(|| API_BASE.to_string());
+ let url = format!("{api_base}/chat/completions");
+
+ let mut builder = client.post(url).bearer_auth(api_key).json(&body);
if let Some(organization_id) = &self.local_config.organization_id {
builder = builder.header("OpenAI-Organization", organization_id);