summaryrefslogtreecommitdiffstats
path: root/src/client/openai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-28 06:04:36 +0800
committerGitHub <noreply@github.com>2024-07-28 06:04:36 +0800
commit0e740d81e94505bd57036755abaaecb12c3b26e3 (patch)
tree49000370fb12e4e5e1f1bd4d145104f2f502aa16 /src/client/openai.rs
parentf5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff)
downloadaichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/openai.rs')
-rw-r--r--src/client/openai.rs80
1 files changed, 49 insertions, 31 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 2b83b7d..ec9cb9d 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -25,45 +25,66 @@ impl OpenAIClient {
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
+}
- fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> {
- let api_key = self.get_api_key()?;
- let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string());
+impl_client_trait!(
+ OpenAIClient,
+ (
+ prepare_chat_completions,
+ openai_chat_completions,
+ openai_chat_completions_streaming
+ ),
+ (prepare_embeddings, openai_embeddings),
+ (noop_prepare_rerank, noop_rerank),
+);
- let url = format!("{api_base}/chat/completions");
+fn prepare_chat_completions(
+ self_: &OpenAIClient,
+ data: ChatCompletionsData,
+) -> Result<RequestData> {
+ let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
- let body = openai_build_chat_completions_body(data, &self.model);
+ let url = format!("{api_base}/chat/completions");
- let mut request_data = RequestData::new(url, body);
+ let body = openai_build_chat_completions_body(data, &self_.model);
- request_data.bearer_auth(api_key);
- if let Some(organization_id) = &self.config.organization_id {
- request_data.header("OpenAI-Organization", organization_id);
- }
+ let mut request_data = RequestData::new(url, body);
- Ok(request_data)
+ request_data.bearer_auth(api_key);
+ if let Some(organization_id) = &self_.config.organization_id {
+ request_data.header("OpenAI-Organization", organization_id);
}
- fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> {
- let api_key = self.get_api_key()?;
- let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string());
+ Ok(request_data)
+}
- let url = format!("{api_base}/embeddings");
+fn prepare_embeddings(self_: &OpenAIClient, data: EmbeddingsData) -> Result<RequestData> {
+ let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
- let body = openai_build_embeddings_body(data, &self.model);
+ let url = format!("{api_base}/embeddings");
- let mut request_data = RequestData::new(url, body);
+ let body = openai_build_embeddings_body(data, &self_.model);
- request_data.bearer_auth(api_key);
- if let Some(organization_id) = &self.config.organization_id {
- request_data.header("OpenAI-Organization", organization_id);
- }
+ let mut request_data = RequestData::new(url, body);
- Ok(request_data)
+ request_data.bearer_auth(api_key);
+ if let Some(organization_id) = &self_.config.organization_id {
+ request_data.header("OpenAI-Organization", organization_id);
}
+
+ Ok(request_data)
}
-pub async fn openai_chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
+pub async fn openai_chat_completions(
+ builder: RequestBuilder,
+ _model: &Model,
+) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -78,6 +99,7 @@ pub async fn openai_chat_completions(builder: RequestBuilder) -> Result<ChatComp
pub async fn openai_chat_completions_streaming(
builder: RequestBuilder,
handler: &mut SseHandler,
+ _model: &Model,
) -> Result<()> {
let mut function_index = 0;
let mut function_name = String::new();
@@ -133,7 +155,10 @@ pub async fn openai_chat_completions_streaming(
sse_stream(builder, handle).await
}
-pub async fn openai_embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
+pub async fn openai_embeddings(
+ builder: RequestBuilder,
+ _model: &Model,
+) -> Result<EmbeddingsOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -277,10 +302,3 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOu
};
Ok(output)
}
-
-impl_client_trait!(
- OpenAIClient,
- openai_chat_completions,
- openai_chat_completions_streaming,
- openai_embeddings
-);