From 912773c25a113f49c5df63cd3a8086d38c75103e Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 24 Sep 2024 07:42:24 +0800 Subject: refactor: embeddings/rerank fn accept ref data (#878) --- src/client/openai_compatible.rs | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) (limited to 'src/client/openai_compatible.rs') diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index b2797e7..f2b7ae3 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -66,7 +66,10 @@ fn prepare_chat_completions( Ok(request_data) } -fn prepare_embeddings(self_: &OpenAICompatibleClient, data: EmbeddingsData) -> Result { +fn prepare_embeddings( + self_: &OpenAICompatibleClient, + data: &EmbeddingsData, +) -> Result { let api_key = self_.get_api_key().ok(); let api_base = get_api_base_ext(self_)?; @@ -83,7 +86,7 @@ fn prepare_embeddings(self_: &OpenAICompatibleClient, data: EmbeddingsData) -> R Ok(request_data) } -fn prepare_rerank(self_: &OpenAICompatibleClient, data: RerankData) -> Result { +fn prepare_rerank(self_: &OpenAICompatibleClient, data: &RerankData) -> Result { let api_key = self_.get_api_key().ok(); let api_base = get_api_base_ext(self_)?; @@ -145,7 +148,7 @@ pub struct GenericRerankResBody { pub results: RerankOutput, } -pub fn generic_build_rerank_body(data: RerankData, model: &Model) -> Value { +pub fn generic_build_rerank_body(data: &RerankData, model: &Model) -> Value { let RerankData { query, documents, @@ -158,9 +161,9 @@ pub fn generic_build_rerank_body(data: RerankData, model: &Model) -> Value { "documents": documents, }); if model.client_name() == "voyageai" { - body["top_k"] = top_n.into() + body["top_k"] = (*top_n).into() } else { - body["top_n"] = top_n.into() + body["top_n"] = (*top_n).into() } body } -- cgit v1.2.3