diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-28 06:04:36 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-28 06:04:36 +0800 |
| commit | 0e740d81e94505bd57036755abaaecb12c3b26e3 (patch) | |
| tree | 49000370fb12e4e5e1f1bd4d145104f2f502aa16 /src/client/openai.rs | |
| parent | f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff) | |
| download | aichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz | |
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/openai.rs')
| -rw-r--r-- | src/client/openai.rs | 80 |
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 -); |
